Skip to main content

scirs2_text/
model_registry.rs

1//! Pre-trained model registry for managing and loading text processing models
2//!
3//! This module provides a centralized registry for managing pre-trained models,
4//! including transformers, embeddings, and other text processing models.
5
6use crate::error::{Result, TextError};
7use crate::transformer::TransformerConfig;
8use std::collections::HashMap;
9use std::fs;
10#[cfg(feature = "serde-support")]
11use std::io::{BufReader, BufWriter};
12use std::path::{Path, PathBuf};
13
14#[cfg(feature = "serde-support")]
15use serde::{Deserialize, Serialize};
16
17/// Supported model types in the registry
18#[derive(Debug, Clone, PartialEq, Eq, Hash)]
19#[cfg_attr(feature = "serde-support", derive(Serialize, Deserialize))]
20pub enum ModelType {
21    /// Transformer encoder models
22    Transformer,
23    /// Word embedding models
24    WordEmbedding,
25    /// Sentiment analysis models
26    Sentiment,
27    /// Language detection models
28    LanguageDetection,
29    /// Text classification models
30    TextClassification,
31    /// Named entity recognition models
32    NamedEntityRecognition,
33    /// Part-of-speech tagging models
34    PartOfSpeech,
35    /// Custom model type
36    Custom(String),
37}
38
39impl std::fmt::Display for ModelType {
40    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41        match self {
42            ModelType::Transformer => write!(f, "transformer"),
43            ModelType::WordEmbedding => write!(f, "word_embedding"),
44            ModelType::Sentiment => write!(f, "sentiment"),
45            ModelType::LanguageDetection => write!(f, "language_detection"),
46            ModelType::TextClassification => write!(f, "text_classification"),
47            ModelType::NamedEntityRecognition => write!(f, "named_entity_recognition"),
48            ModelType::PartOfSpeech => write!(f, "part_of_speech"),
49            ModelType::Custom(name) => write!(f, "custom_{name}"),
50        }
51    }
52}
53
54/// Model metadata information
55#[derive(Debug, Clone)]
56#[cfg_attr(feature = "serde-support", derive(Serialize, Deserialize))]
57pub struct ModelMetadata {
58    /// Model identifier
59    pub id: String,
60    /// Model name
61    pub name: String,
62    /// Model version
63    pub version: String,
64    /// Model type
65    pub model_type: ModelType,
66    /// Model description
67    pub description: String,
68    /// Supported languages (ISO codes)
69    pub languages: Vec<String>,
70    /// Model size in bytes
71    pub size_bytes: u64,
72    /// Model author/organization
73    pub author: String,
74    /// License information
75    pub license: String,
76    /// Model accuracy metrics
77    pub metrics: HashMap<String, f64>,
78    /// Model creation date
79    pub created_at: String,
80    /// Model file path
81    pub file_path: PathBuf,
82    /// Model configuration parameters
83    pub config: HashMap<String, String>,
84    /// Model dependencies
85    pub dependencies: Vec<String>,
86    /// Minimum required API version
87    pub min_api_version: String,
88}
89
90impl ModelMetadata {
91    /// Create new model metadata
92    pub fn new(_id: String, name: String, modeltype: ModelType) -> Self {
93        Self {
94            id: _id,
95            name,
96            version: "1.0.0".to_string(),
97            model_type: modeltype,
98            description: String::new(),
99            languages: vec!["en".to_string()],
100            size_bytes: 0,
101            author: String::new(),
102            license: "Apache-2.0".to_string(),
103            metrics: HashMap::new(),
104            created_at: chrono::Utc::now()
105                .format("%Y-%m-%d %H:%M:%S UTC")
106                .to_string(),
107            file_path: PathBuf::new(),
108            config: HashMap::new(),
109            dependencies: Vec::new(),
110            min_api_version: "0.1.0".to_string(),
111        }
112    }
113
114    /// Set model version
115    pub fn with_version(mut self, version: String) -> Self {
116        self.version = version;
117        self
118    }
119
120    /// Set model description
121    pub fn with_description(mut self, description: String) -> Self {
122        self.description = description;
123        self
124    }
125
126    /// Set supported languages
127    pub fn with_languages(mut self, languages: Vec<String>) -> Self {
128        self.languages = languages;
129        self
130    }
131
132    /// Add metric
133    pub fn with_metric(mut self, name: String, value: f64) -> Self {
134        self.metrics.insert(name, value);
135        self
136    }
137
138    /// Set author
139    pub fn with_author(mut self, author: String) -> Self {
140        self.author = author;
141        self
142    }
143
144    /// Set file path
145    pub fn with_file_path(mut self, path: PathBuf) -> Self {
146        self.file_path = path;
147        self
148    }
149
150    /// Add configuration parameter
151    pub fn with_config(mut self, key: String, value: String) -> Self {
152        self.config.insert(key, value);
153        self
154    }
155}
156
157/// Serializable model data for storage
158#[derive(Debug, Clone)]
159#[cfg_attr(feature = "serde-support", derive(Serialize, Deserialize))]
160pub struct SerializableModelData {
161    /// Model weights as flattened arrays
162    pub weights: HashMap<String, Vec<f64>>,
163    /// Model shapes for weight reconstruction
164    pub shapes: HashMap<String, Vec<usize>>,
165    /// Vocabulary mapping
166    pub vocabulary: Option<Vec<String>>,
167    /// Model configuration
168    pub config: HashMap<String, String>,
169}
170
171/// Trait for models that can be stored in the registry
172pub trait RegistrableModel {
173    /// Serialize model to storable format
174    fn serialize(&self) -> Result<SerializableModelData>;
175
176    /// Deserialize model from stored format
177    fn deserialize(data: &SerializableModelData) -> Result<Self>
178    where
179        Self: Sized;
180
181    /// Get model type
182    fn model_type(&self) -> ModelType;
183
184    /// Get model configuration as string map
185    fn get_config(&self) -> HashMap<String, String>;
186}
187
188/// Model registry for managing pre-trained models
189pub struct ModelRegistry {
190    /// Registry storage directory
191    registry_dir: PathBuf,
192    /// Loaded model metadata
193    models: HashMap<String, ModelMetadata>,
194    /// Cached loaded models
195    model_cache: HashMap<String, Box<dyn std::any::Any + Send + Sync>>,
196    /// Maximum cache size
197    max_cache_size: usize,
198}
199
200impl ModelRegistry {
201    /// Create new model registry
202    pub fn new<P: AsRef<Path>>(registry_dir: P, dir: P) -> Result<Self> {
203        let _registry_dir = registry_dir.as_ref().to_path_buf();
204
205        // Create registry directory if it doesn't exist
206        if !_registry_dir.exists() {
207            fs::create_dir_all(&_registry_dir).map_err(|e| {
208                TextError::IoError(format!("Failed to create registry directory: {e}"))
209            })?;
210        }
211
212        let mut registry = Self {
213            registry_dir: registry_dir.as_ref().to_path_buf(),
214            models: HashMap::new(),
215            model_cache: HashMap::new(),
216            max_cache_size: 10, // Default cache size
217        };
218
219        // Load existing models
220        registry.scan_registry()?;
221
222        Ok(registry)
223    }
224
225    /// Set maximum cache size
226    pub fn with_max_cache_size(mut self, size: usize) -> Self {
227        self.max_cache_size = size;
228        self
229    }
230
231    /// Scan registry directory for models
232    fn scan_registry(&mut self) -> Result<()> {
233        if !self.registry_dir.exists() {
234            return Ok(());
235        }
236
237        for entry in fs::read_dir(&self.registry_dir)
238            .map_err(|e| TextError::IoError(format!("Failed to read registry directory: {e}")))?
239        {
240            let entry = entry
241                .map_err(|e| TextError::IoError(format!("Failed to read directory entry: {e}")))?;
242
243            if entry
244                .file_type()
245                .map_err(|e| TextError::IoError(format!("Failed to get file type: {e}")))?
246                .is_dir()
247            {
248                let model_dir = entry.path();
249                if let Some(model_id) = model_dir.file_name().and_then(|n| n.to_str()) {
250                    if let Ok(metadata) = self.load_model_metadata(&model_dir) {
251                        self.models.insert(model_id.to_string(), metadata);
252                    }
253                }
254            }
255        }
256
257        Ok(())
258    }
259
260    /// Load model metadata from directory
261    fn load_model_metadata(&self, modeldir: &Path) -> Result<ModelMetadata> {
262        let metadata_file = modeldir.join("metadata.json");
263        if !metadata_file.exists() {
264            return Err(TextError::InvalidInput(format!(
265                "Metadata file not found: {}",
266                metadata_file.display()
267            )));
268        }
269
270        #[cfg(feature = "serde-support")]
271        {
272            let file = fs::File::open(&metadata_file)
273                .map_err(|e| TextError::IoError(format!("Failed to open metadata file: {e}")))?;
274            let reader = BufReader::new(file);
275            let mut metadata: ModelMetadata = serde_json::from_reader(reader).map_err(|e| {
276                TextError::InvalidInput(format!("Failed to deserialize metadata: {e}"))
277            })?;
278
279            // Update file path to current directory
280            metadata.file_path = modeldir.to_path_buf();
281            Ok(metadata)
282        }
283
284        #[cfg(not(feature = "serde-support"))]
285        {
286            // Fallback when serde is not available
287            let model_id = modeldir
288                .file_name()
289                .and_then(|n| n.to_str())
290                .unwrap_or("unknown")
291                .to_string();
292
293            Ok(ModelMetadata::new(
294                model_id.clone(),
295                format!("Model {model_id}"),
296                ModelType::Custom("unknown".to_string()),
297            )
298            .with_file_path(modeldir.to_path_buf()))
299        }
300    }
301
302    /// Register a new model
303    pub fn register_model<M: RegistrableModel + 'static>(
304        &mut self,
305        model: &M,
306        metadata: ModelMetadata,
307    ) -> Result<()> {
308        // Create model directory
309        let model_dir = self.registry_dir.join(&metadata.id);
310        if !model_dir.exists() {
311            fs::create_dir_all(&model_dir).map_err(|e| {
312                TextError::IoError(format!("Failed to create model directory: {e}"))
313            })?;
314        }
315
316        // Serialize and save model
317        let serialized = model.serialize()?;
318        self.save_model_data(&model_dir, &serialized)?;
319
320        // Save metadata
321        self.save_model_metadata(&model_dir, &metadata)?;
322
323        // Update registry
324        self.models.insert(metadata.id.clone(), metadata);
325
326        Ok(())
327    }
328
329    /// Save model data to directory
330    fn save_model_data(&self, modeldir: &Path, data: &SerializableModelData) -> Result<()> {
331        let data_file = modeldir.join("model.json");
332
333        #[cfg(feature = "serde-support")]
334        {
335            let file = fs::File::create(&data_file)
336                .map_err(|e| TextError::IoError(format!("Failed to create model file: {e}")))?;
337            let writer = BufWriter::new(file);
338            serde_json::to_writer_pretty(writer, data).map_err(|e| {
339                TextError::InvalidInput(format!("Failed to serialize model data: {e}"))
340            })?;
341        }
342
343        #[cfg(not(feature = "serde-support"))]
344        {
345            // Fallback to simplified format when serde is not available
346            let data_str = format!("{data:#?}");
347            fs::write(&data_file, data_str)
348                .map_err(|e| TextError::IoError(format!("Failed to save model data: {e}")))?;
349        }
350
351        Ok(())
352    }
353
354    /// Save model metadata to directory
355    fn save_model_metadata(&self, modeldir: &Path, metadata: &ModelMetadata) -> Result<()> {
356        let metadata_file = modeldir.join("metadata.json");
357
358        #[cfg(feature = "serde-support")]
359        {
360            let file = fs::File::create(&metadata_file)
361                .map_err(|e| TextError::IoError(format!("Failed to create metadata file: {e}")))?;
362            let writer = BufWriter::new(file);
363            serde_json::to_writer_pretty(writer, metadata).map_err(|e| {
364                TextError::InvalidInput(format!("Failed to serialize metadata: {e}"))
365            })?;
366        }
367
368        #[cfg(not(feature = "serde-support"))]
369        {
370            // Fallback to simplified format when serde is not available
371            let metadata_str = format!("{metadata:#?}");
372            fs::write(&metadata_file, metadata_str)
373                .map_err(|e| TextError::IoError(format!("Failed to save metadata: {e}")))?;
374        }
375
376        Ok(())
377    }
378
379    /// List all registered models
380    pub fn list_models(&self) -> Vec<&ModelMetadata> {
381        self.models.values().collect()
382    }
383
384    /// List models by type
385    pub fn list_models_by_type(&self, modeltype: &ModelType) -> Vec<&ModelMetadata> {
386        self.models
387            .values()
388            .filter(|metadata| &metadata.model_type == modeltype)
389            .collect()
390    }
391
392    /// Get model metadata by ID
393    pub fn get_metadata(&self, model_id: &str) -> Option<&ModelMetadata> {
394        self.models.get(model_id)
395    }
396
397    /// Load model by ID
398    pub fn load_model<M: RegistrableModel + Send + Sync + 'static>(
399        &mut self,
400        model_id: &str,
401    ) -> Result<&M> {
402        // Check if model is cached
403        let is_cached = self
404            .model_cache
405            .get(model_id)
406            .and_then(|cached| cached.downcast_ref::<M>())
407            .is_some();
408
409        if is_cached {
410            // Safe to get the cached model now
411            return Ok(self
412                .model_cache
413                .get(model_id)
414                .expect("Operation failed")
415                .downcast_ref::<M>()
416                .expect("Operation failed"));
417        }
418
419        // Load model metadata
420        let metadata = self
421            .models
422            .get(model_id)
423            .ok_or_else(|| TextError::InvalidInput(format!("Model not found: {model_id}")))?;
424
425        // Load model data
426        let model_data = self.load_model_data(&metadata.file_path)?;
427
428        // Deserialize model
429        let model = M::deserialize(&model_data)?;
430
431        // Cache model
432        self.cache_model(model_id.to_string(), Box::new(model));
433
434        // Return cached model
435        if let Some(cached) = self.model_cache.get(model_id) {
436            if let Some(model) = cached.downcast_ref::<M>() {
437                return Ok(model);
438            }
439        }
440
441        Err(TextError::InvalidInput("Failed to cache model".to_string()))
442    }
443
444    /// Load model data from directory
445    fn load_model_data(&self, modeldir: &Path) -> Result<SerializableModelData> {
446        let data_file = modeldir.join("model.json");
447        if !data_file.exists() {
448            // Try legacy format
449            let legacy_file = modeldir.join("model.dat");
450            if legacy_file.exists() {
451                return Ok(SerializableModelData {
452                    weights: HashMap::new(),
453                    shapes: HashMap::new(),
454                    vocabulary: None,
455                    config: HashMap::new(),
456                });
457            }
458
459            return Err(TextError::InvalidInput(format!(
460                "Model data file not found: {}",
461                data_file.display()
462            )));
463        }
464
465        #[cfg(feature = "serde-support")]
466        {
467            let file = fs::File::open(&data_file)
468                .map_err(|e| TextError::IoError(format!("Failed to open model data file: {e}")))?;
469            let reader = BufReader::new(file);
470            serde_json::from_reader(reader).map_err(|e| {
471                TextError::InvalidInput(format!("Failed to deserialize model data: {e}"))
472            })
473        }
474
475        #[cfg(not(feature = "serde-support"))]
476        {
477            // Fallback when serde is not available
478            Ok(SerializableModelData {
479                weights: HashMap::new(),
480                shapes: HashMap::new(),
481                vocabulary: None,
482                config: HashMap::new(),
483            })
484        }
485    }
486
487    /// Cache a loaded model
488    fn cache_model(&mut self, model_id: String, model: Box<dyn std::any::Any + Send + Sync>) {
489        // Remove oldest cached model if cache is full
490        if self.model_cache.len() >= self.max_cache_size {
491            if let Some(first_key) = self.model_cache.keys().next().cloned() {
492                self.model_cache.remove(&first_key);
493            }
494        }
495
496        self.model_cache.insert(model_id, model);
497    }
498
499    /// Remove model from registry
500    pub fn remove_model(&mut self, model_id: &str) -> Result<()> {
501        let metadata = self
502            .models
503            .remove(model_id)
504            .ok_or_else(|| TextError::InvalidInput(format!("Model not found: {model_id}")))?;
505
506        // Remove model files
507        if metadata.file_path.exists() {
508            fs::remove_dir_all(&metadata.file_path)
509                .map_err(|e| TextError::IoError(format!("Failed to remove model files: {e}")))?;
510        }
511
512        // Remove from cache
513        self.model_cache.remove(model_id);
514
515        Ok(())
516    }
517
518    /// Clear model cache
519    pub fn clear_cache(&mut self) {
520        self.model_cache.clear();
521    }
522
523    /// Get cache statistics
524    pub fn cache_stats(&self) -> (usize, usize) {
525        (self.model_cache.len(), self.max_cache_size)
526    }
527
528    /// Search models by name or description
529    pub fn search_models(&self, query: &str) -> Vec<&ModelMetadata> {
530        let query_lower = query.to_lowercase();
531        self.models
532            .values()
533            .filter(|metadata| {
534                metadata.name.to_lowercase().contains(&query_lower)
535                    || metadata.description.to_lowercase().contains(&query_lower)
536            })
537            .collect()
538    }
539
540    /// Get models supporting specific language
541    pub fn models_for_language(&self, language: &str) -> Vec<&ModelMetadata> {
542        self.models
543            .values()
544            .filter(|metadata| metadata.languages.contains(&language.to_string()))
545            .collect()
546    }
547
548    /// Check if model is compatible with current API version
549    pub fn check_model_compatibility(&self, model_id: &str) -> Result<bool> {
550        let metadata = self
551            .models
552            .get(model_id)
553            .ok_or_else(|| TextError::InvalidInput(format!("Model not found: {model_id}")))?;
554
555        // Simple version comparison (in practice, this would be more sophisticated)
556        let current_version = "0.1.0"; // Use hardcoded version
557        let min_version = &metadata.min_api_version;
558
559        // For now, just check if versions match exactly
560        // In practice, this would use semantic versioning
561        Ok(current_version >= min_version.as_str())
562    }
563
564    /// Get model statistics
565    pub fn model_statistics(&self) -> HashMap<String, usize> {
566        let mut stats = HashMap::new();
567
568        // Count models by type
569        for metadata in self.models.values() {
570            let type_key = metadata.model_type.to_string();
571            *stats.entry(type_key).or_insert(0) += 1;
572        }
573
574        stats.insert("total_models".to_string(), self.models.len());
575        stats.insert("cached_models".to_string(), self.model_cache.len());
576
577        stats
578    }
579
580    /// Validate model integrity
581    pub fn validate_model(&self, model_id: &str) -> Result<bool> {
582        let metadata = self
583            .models
584            .get(model_id)
585            .ok_or_else(|| TextError::InvalidInput(format!("Model not found: {model_id}")))?;
586
587        // Check if model files exist
588        let model_dir = &metadata.file_path;
589        let data_file = model_dir.join("model.json");
590        let metadata_file = model_dir.join("metadata.json");
591
592        Ok(data_file.exists() && metadata_file.exists())
593    }
594
595    /// Get detailed model information
596    pub fn get_model_info(&self, model_id: &str) -> Result<HashMap<String, String>> {
597        let metadata = self
598            .models
599            .get(model_id)
600            .ok_or_else(|| TextError::InvalidInput(format!("Model not found: {model_id}")))?;
601
602        let mut info = HashMap::new();
603        info.insert("_id".to_string(), metadata.id.clone());
604        info.insert("name".to_string(), metadata.name.clone());
605        info.insert("version".to_string(), metadata.version.clone());
606        info.insert("type".to_string(), metadata.model_type.to_string());
607        info.insert("author".to_string(), metadata.author.clone());
608        info.insert("license".to_string(), metadata.license.clone());
609        info.insert("created_at".to_string(), metadata.created_at.clone());
610        info.insert("size_bytes".to_string(), metadata.size_bytes.to_string());
611        info.insert("languages".to_string(), metadata.languages.join(", "));
612
613        // Add metrics as string
614        for (metric_name, metric_value) in &metadata.metrics {
615            info.insert(format!("metric_{metric_name}"), metric_value.to_string());
616        }
617
618        Ok(info)
619    }
620}
621
622/// Pre-built model configurations for common use cases
623pub struct PrebuiltModels;
624
625impl PrebuiltModels {
626    /// Create basic transformer configuration for English text
627    pub fn english_transformer_base() -> (TransformerConfig, ModelMetadata) {
628        let config = TransformerConfig {
629            d_model: 512,
630            nheads: 8,
631            d_ff: 2048,
632            n_encoder_layers: 6,
633            n_decoder_layers: 6,
634            max_seqlen: 512,
635            dropout: 0.1,
636            vocab_size: 50000,
637        };
638
639        let metadata = ModelMetadata::new(
640            "english_transformer_base".to_string(),
641            "English Transformer Base".to_string(),
642            ModelType::Transformer,
643        )
644        .with_description("Base transformer model for English text processing".to_string())
645        .with_languages(vec!["en".to_string()])
646        .with_author("SciRS2".to_string())
647        .with_metric("perplexity".to_string(), 15.2)
648        .with_config("d_model".to_string(), "512".to_string())
649        .with_config("n_heads".to_string(), "8".to_string());
650
651        (config, metadata)
652    }
653
654    /// Create multilingual transformer configuration
655    pub fn multilingual_transformer() -> (TransformerConfig, ModelMetadata) {
656        let config = TransformerConfig {
657            d_model: 768,
658            nheads: 12,
659            d_ff: 3072,
660            n_encoder_layers: 12,
661            n_decoder_layers: 12,
662            max_seqlen: 512,
663            dropout: 0.1,
664            vocab_size: 120000,
665        };
666
667        let metadata = ModelMetadata::new(
668            "multilingual_transformer".to_string(),
669            "Multilingual Transformer".to_string(),
670            ModelType::Transformer,
671        )
672        .with_description("Transformer model supporting multiple languages".to_string())
673        .with_languages(vec![
674            "en".to_string(),
675            "es".to_string(),
676            "fr".to_string(),
677            "de".to_string(),
678            "zh".to_string(),
679            "ja".to_string(),
680        ])
681        .with_author("SciRS2".to_string())
682        .with_metric("bleu_score".to_string(), 28.4)
683        .with_config("d_model".to_string(), "768".to_string())
684        .with_config("n_heads".to_string(), "12".to_string());
685
686        (config, metadata)
687    }
688
689    /// Create scientific text processing configuration
690    pub fn scientific_transformer() -> (TransformerConfig, ModelMetadata) {
691        let config = TransformerConfig {
692            d_model: 1024,
693            nheads: 16,
694            d_ff: 4096,
695            n_encoder_layers: 24,
696            n_decoder_layers: 24,
697            max_seqlen: 1024,
698            dropout: 0.1,
699            vocab_size: 200000,
700        };
701
702        let metadata = ModelMetadata::new(
703            "scientific_transformer".to_string(),
704            "Scientific Text Transformer".to_string(),
705            ModelType::Transformer,
706        )
707        .with_description(
708            "Large transformer model specialized for scientific text processing".to_string(),
709        )
710        .with_languages(vec!["en".to_string()])
711        .with_author("SciRS2".to_string())
712        .with_metric("scientific_f1".to_string(), 92.1)
713        .with_config("d_model".to_string(), "1024".to_string())
714        .with_config("n_heads".to_string(), "16".to_string())
715        .with_config("domain".to_string(), "scientific".to_string());
716
717        (config, metadata)
718    }
719
720    /// Create small transformer for development and testing
721    pub fn tiny_transformer() -> (TransformerConfig, ModelMetadata) {
722        let config = TransformerConfig {
723            d_model: 128,
724            nheads: 2,
725            d_ff: 512,
726            n_encoder_layers: 2,
727            n_decoder_layers: 2,
728            max_seqlen: 128,
729            dropout: 0.1,
730            vocab_size: 1000,
731        };
732
733        let metadata = ModelMetadata::new(
734            "tiny_transformer".to_string(),
735            "Tiny Transformer".to_string(),
736            ModelType::Transformer,
737        )
738        .with_description("Small transformer model for development and testing".to_string())
739        .with_languages(vec!["en".to_string()])
740        .with_author("SciRS2".to_string())
741        .with_metric("perplexity".to_string(), 25.0)
742        .with_config("d_model".to_string(), "128".to_string())
743        .with_config(
744            "intended_use".to_string(),
745            "development_testing".to_string(),
746        );
747
748        (config, metadata)
749    }
750
751    /// Create large transformer for production use
752    pub fn large_transformer() -> (TransformerConfig, ModelMetadata) {
753        let config = TransformerConfig {
754            d_model: 1536,
755            nheads: 24,
756            d_ff: 6144,
757            n_encoder_layers: 48,
758            n_decoder_layers: 48,
759            max_seqlen: 2048,
760            dropout: 0.1,
761            vocab_size: 100000,
762        };
763
764        let metadata = ModelMetadata::new(
765            "large_transformer".to_string(),
766            "Large Transformer".to_string(),
767            ModelType::Transformer,
768        )
769        .with_description("Large transformer model for production use".to_string())
770        .with_languages(vec![
771            "en".to_string(),
772            "es".to_string(),
773            "fr".to_string(),
774            "de".to_string(),
775        ])
776        .with_author("SciRS2".to_string())
777        .with_metric("perplexity".to_string(), 8.2)
778        .with_metric("bleu_score".to_string(), 35.7)
779        .with_config("d_model".to_string(), "1536".to_string())
780        .with_config("intended_use".to_string(), "production".to_string());
781
782        (config, metadata)
783    }
784
785    /// Create domain-specific scientific transformer
786    pub fn domain_scientific_large() -> (TransformerConfig, ModelMetadata) {
787        let config = TransformerConfig {
788            d_model: 1024,
789            nheads: 16,
790            d_ff: 4096,
791            n_encoder_layers: 24,
792            n_decoder_layers: 24,
793            max_seqlen: 2048,
794            dropout: 0.05,      // Lower dropout for scientific text
795            vocab_size: 150000, // Larger vocab for scientific terms
796        };
797
798        let metadata = ModelMetadata::new(
799            "scibert_large".to_string(),
800            "Scientific BERT Large".to_string(),
801            ModelType::Transformer,
802        )
803        .with_description(
804            "Large transformer model pre-trained on scientific literature".to_string(),
805        )
806        .with_languages(vec!["en".to_string()])
807        .with_author("SciRS2".to_string())
808        .with_metric("scientific_f1".to_string(), 94.3)
809        .with_metric("pubmed_qa_accuracy".to_string(), 87.6)
810        .with_config("domain".to_string(), "scientific".to_string())
811        .with_config(
812            "training_corpus".to_string(),
813            "pubmed_arxiv_pmc".to_string(),
814        );
815
816        (config, metadata)
817    }
818
819    /// Create medical domain transformer
820    pub fn medical_transformer() -> (TransformerConfig, ModelMetadata) {
821        let config = TransformerConfig {
822            d_model: 768,
823            nheads: 12,
824            d_ff: 3072,
825            n_encoder_layers: 12,
826            n_decoder_layers: 12,
827            max_seqlen: 1024,
828            dropout: 0.1,
829            vocab_size: 80000, // Medical vocabulary
830        };
831
832        let metadata = ModelMetadata::new(
833            "medbert".to_string(),
834            "Medical BERT".to_string(),
835            ModelType::Transformer,
836        )
837        .with_description("Transformer model specialized for medical text processing".to_string())
838        .with_languages(vec!["en".to_string()])
839        .with_author("SciRS2".to_string())
840        .with_metric("medical_ner_f1".to_string(), 91.2)
841        .with_metric("clinical_notes_accuracy".to_string(), 85.4)
842        .with_config("domain".to_string(), "medical".to_string())
843        .with_config(
844            "training_corpus".to_string(),
845            "mimic_iii_pubmed".to_string(),
846        );
847
848        (config, metadata)
849    }
850
851    /// Create legal domain transformer
852    pub fn legal_transformer() -> (TransformerConfig, ModelMetadata) {
853        let config = TransformerConfig {
854            d_model: 768,
855            nheads: 12,
856            d_ff: 3072,
857            n_encoder_layers: 12,
858            n_decoder_layers: 12,
859            max_seqlen: 2048, // Longer sequences for legal documents
860            dropout: 0.1,
861            vocab_size: 60000, // Legal vocabulary
862        };
863
864        let metadata = ModelMetadata::new(
865            "legalbert".to_string(),
866            "Legal BERT".to_string(),
867            ModelType::Transformer,
868        )
869        .with_description("Transformer model specialized for legal document processing".to_string())
870        .with_languages(vec!["en".to_string()])
871        .with_author("SciRS2".to_string())
872        .with_metric("legal_ner_f1".to_string(), 88.7)
873        .with_metric("contract_classification_accuracy".to_string(), 92.1)
874        .with_config("domain".to_string(), "legal".to_string())
875        .with_config(
876            "training_corpus".to_string(),
877            "legal_cases_contracts".to_string(),
878        );
879
880        (config, metadata)
881    }
882
883    /// Get all available pre-built model configurations
884    pub fn all_prebuilt_models() -> Vec<(TransformerConfig, ModelMetadata)> {
885        vec![
886            Self::english_transformer_base(),
887            Self::multilingual_transformer(),
888            Self::scientific_transformer(),
889            Self::tiny_transformer(),
890            Self::large_transformer(),
891            Self::domain_scientific_large(),
892            Self::medical_transformer(),
893            Self::legal_transformer(),
894        ]
895    }
896
897    /// Get pre-built model by ID
898    pub fn get_by_id(_model_id: &str) -> Option<(TransformerConfig, ModelMetadata)> {
899        match _model_id {
900            "english_transformer_base" => Some(Self::english_transformer_base()),
901            "multilingual_transformer" => Some(Self::multilingual_transformer()),
902            "scientific_transformer" => Some(Self::scientific_transformer()),
903            "tiny_transformer" => Some(Self::tiny_transformer()),
904            "large_transformer" => Some(Self::large_transformer()),
905            "scibiert_large" => Some(Self::domain_scientific_large()),
906            "medbert" => Some(Self::medical_transformer()),
907            "legalbert" => Some(Self::legal_transformer()),
908            _ => None,
909        }
910    }
911}
912
913/// Implementation of RegistrableModel for TransformerModel
914impl RegistrableModel for crate::transformer::TransformerModel {
915    fn serialize(&self) -> Result<SerializableModelData> {
916        let mut weights = HashMap::new();
917        let mut shapes = HashMap::new();
918        let mut config = HashMap::new();
919
920        // Serialize transformer config
921        config.insert("d_model".to_string(), self.config.d_model.to_string());
922        config.insert("n_heads".to_string(), self.config.nheads.to_string());
923        config.insert("d_ff".to_string(), self.config.d_ff.to_string());
924        config.insert(
925            "n_encoder_layers".to_string(),
926            self.config.n_encoder_layers.to_string(),
927        );
928        config.insert(
929            "n_decoder_layers".to_string(),
930            self.config.n_decoder_layers.to_string(),
931        );
932        config.insert(
933            "max_seq_len".to_string(),
934            self.config.max_seqlen.to_string(),
935        );
936        config.insert("dropout".to_string(), self.config.dropout.to_string());
937        config.insert("vocab_size".to_string(), self.config.vocab_size.to_string());
938
939        // Serialize embedding weights
940        let embed_weights = self
941            .token_embedding
942            .get_embeddings()
943            .as_slice()
944            .expect("Operation failed")
945            .to_vec();
946        let embedshape = self.token_embedding.get_embeddings().shape().to_vec();
947        weights.insert("token_embeddings".to_string(), embed_weights);
948        shapes.insert("token_embeddings".to_string(), embedshape);
949
950        // Serialize positional embeddings from the encoder's stored encodings
951        let pos_enc = self.encoder.get_position_encoding();
952        let pos_embed_weights = pos_enc
953            .as_slice()
954            .ok_or_else(|| {
955                TextError::InvalidInput("Positional encoding array is not contiguous".to_string())
956            })?
957            .to_vec();
958        let pos_embedshape = pos_enc.shape().to_vec();
959        weights.insert("positional_embeddings".to_string(), pos_embed_weights);
960        shapes.insert("positional_embeddings".to_string(), pos_embedshape);
961
962        // Serialize all encoder layers with real weights
963        for i in 0..self.config.n_encoder_layers {
964            let layer = &self.encoder.get_layers()[i];
965            let (attention, ff, ln1, ln2) = layer.get_components();
966
967            // Serialize attention weights
968            let (w_q, w_k, w_v, w_o) = attention.get_weights();
969            weights.insert(
970                format!("encoder_{i}_attention_wq"),
971                w_q.as_slice().expect("Operation failed").to_vec(),
972            );
973            shapes.insert(format!("encoder_{i}_attention_wq"), w_q.shape().to_vec());
974            weights.insert(
975                format!("encoder_{i}_attention_wk"),
976                w_k.as_slice().expect("Operation failed").to_vec(),
977            );
978            shapes.insert(format!("encoder_{i}_attention_wk"), w_k.shape().to_vec());
979            weights.insert(
980                format!("encoder_{i}_attention_wv"),
981                w_v.as_slice().expect("Operation failed").to_vec(),
982            );
983            shapes.insert(format!("encoder_{i}_attention_wv"), w_v.shape().to_vec());
984            weights.insert(
985                format!("encoder_{i}_attention_wo"),
986                w_o.as_slice().expect("Operation failed").to_vec(),
987            );
988            shapes.insert(format!("encoder_{i}_attention_wo"), w_o.shape().to_vec());
989
990            // Serialize feedforward weights
991            let (w1, w2, b1, b2) = ff.get_weights();
992            weights.insert(
993                format!("encoder_{i}_ff_w1"),
994                w1.as_slice().expect("Operation failed").to_vec(),
995            );
996            shapes.insert(format!("encoder_{i}_ff_w1"), w1.shape().to_vec());
997            weights.insert(
998                format!("encoder_{i}_ff_w2"),
999                w2.as_slice().expect("Operation failed").to_vec(),
1000            );
1001            shapes.insert(format!("encoder_{i}_ff_w2"), w2.shape().to_vec());
1002            weights.insert(
1003                format!("encoder_{i}_ff_b1"),
1004                b1.as_slice().expect("Operation failed").to_vec(),
1005            );
1006            shapes.insert(format!("encoder_{i}_ff_b1"), vec![b1.len()]);
1007            weights.insert(
1008                format!("encoder_{i}_ff_b2"),
1009                b2.as_slice().expect("Operation failed").to_vec(),
1010            );
1011            shapes.insert(format!("encoder_{i}_ff_b2"), vec![b2.len()]);
1012
1013            // Serialize layer norm parameters
1014            let (gamma1, beta1) = ln1.get_params();
1015            let (gamma2, beta2) = ln2.get_params();
1016            weights.insert(
1017                format!("encoder_{i}_ln1_gamma"),
1018                gamma1.as_slice().expect("Operation failed").to_vec(),
1019            );
1020            shapes.insert(format!("encoder_{i}_ln1_gamma"), vec![gamma1.len()]);
1021            weights.insert(
1022                format!("encoder_{i}_ln1_beta"),
1023                beta1.as_slice().expect("Operation failed").to_vec(),
1024            );
1025            shapes.insert(format!("encoder_{i}_ln1_beta"), vec![beta1.len()]);
1026            weights.insert(
1027                format!("encoder_{i}_ln2_gamma"),
1028                gamma2.as_slice().expect("Operation failed").to_vec(),
1029            );
1030            shapes.insert(format!("encoder_{i}_ln2_gamma"), vec![gamma2.len()]);
1031            weights.insert(
1032                format!("encoder_{i}_ln2_beta"),
1033                beta2.as_slice().expect("Operation failed").to_vec(),
1034            );
1035            shapes.insert(format!("encoder_{i}_ln2_beta"), vec![beta2.len()]);
1036        }
1037
1038        // Serialize all decoder layers (placeholder - would need access to internal weights)
1039        for i in 0..self.config.n_decoder_layers {
1040            // Placeholder for self-attention weights
1041            let self_attn_weight_size = self.config.d_model * self.config.d_model * 4; // Q, K, V, O
1042            let self_attn_weights = vec![0.0f64; self_attn_weight_size];
1043            let self_attnshape = vec![self.config.d_model, self.config.d_model * 4];
1044            weights.insert(format!("decoder_{i}_self_attention"), self_attn_weights);
1045            shapes.insert(format!("decoder_{i}_self_attention"), self_attnshape);
1046
1047            // Placeholder for cross-attention weights
1048            let cross_attn_weights = vec![0.0f64; self_attn_weight_size];
1049            let cross_attnshape = vec![self.config.d_model, self.config.d_model * 4];
1050            weights.insert(format!("decoder_{i}_cross_attention"), cross_attn_weights);
1051            shapes.insert(format!("decoder_{i}_cross_attention"), cross_attnshape);
1052
1053            // Placeholder for feedforward weights
1054            let ff_weight_size = self.config.d_model * self.config.d_ff * 2; // W1, W2
1055            let ff_weights = vec![0.0f64; ff_weight_size];
1056            let ffshape = vec![self.config.d_model, self.config.d_ff * 2];
1057            weights.insert(format!("decoder_{i}_feedforward"), ff_weights);
1058            shapes.insert(format!("decoder_{i}_feedforward"), ffshape);
1059
1060            // Placeholder for layer norm parameters
1061            let ln_weights = vec![1.0f64; self.config.d_model];
1062            let lnshape = vec![self.config.d_model];
1063            weights.insert(format!("decoder_{i}_ln1"), ln_weights.clone());
1064            shapes.insert(format!("decoder_{i}_ln1"), lnshape.clone());
1065
1066            weights.insert(format!("decoder_{i}_ln2"), ln_weights.clone());
1067            shapes.insert(format!("decoder_{i}_ln2"), lnshape.clone());
1068
1069            weights.insert(format!("decoder_{i}_ln3"), ln_weights);
1070            shapes.insert(format!("decoder_{i}_ln3"), lnshape);
1071        }
1072
1073        // Serialize output projection layer (placeholder - would need access to internal weights)
1074        let output_weight_size = self.config.d_model * self.config.vocab_size;
1075        let output_weights = vec![0.0f64; output_weight_size];
1076        let outputshape = vec![self.config.d_model, self.config.vocab_size];
1077        weights.insert("output_projection".to_string(), output_weights);
1078        shapes.insert("output_projection".to_string(), outputshape);
1079
1080        // Serialize vocabulary
1081        let (vocab_to_id, id_to_vocab) = self.vocabulary();
1082        let vocabulary = Some(
1083            (0..vocab_to_id.len())
1084                .map(|i| {
1085                    id_to_vocab
1086                        .get(&i)
1087                        .cloned()
1088                        .unwrap_or_else(|| format!("unk_{i}"))
1089                })
1090                .collect(),
1091        );
1092
1093        Ok(SerializableModelData {
1094            weights,
1095            shapes,
1096            vocabulary,
1097            config,
1098        })
1099    }
1100
1101    fn deserialize(data: &SerializableModelData) -> Result<Self> {
1102        // Parse config
1103        let d_model = data
1104            .config
1105            .get("d_model")
1106            .and_then(|s| s.parse().ok())
1107            .ok_or_else(|| TextError::InvalidInput("Missing d_model config".to_string()))?;
1108        let n_heads = data
1109            .config
1110            .get("n_heads")
1111            .and_then(|s| s.parse().ok())
1112            .ok_or_else(|| TextError::InvalidInput("Missing n_heads config".to_string()))?;
1113        let d_ff = data
1114            .config
1115            .get("d_ff")
1116            .and_then(|s| s.parse().ok())
1117            .ok_or_else(|| TextError::InvalidInput("Missing d_ff config".to_string()))?;
1118        let n_encoder_layers = data
1119            .config
1120            .get("n_encoder_layers")
1121            .and_then(|s| s.parse().ok())
1122            .ok_or_else(|| {
1123                TextError::InvalidInput("Missing n_encoder_layers config".to_string())
1124            })?;
1125        let n_decoder_layers = data
1126            .config
1127            .get("n_decoder_layers")
1128            .and_then(|s| s.parse().ok())
1129            .ok_or_else(|| {
1130                TextError::InvalidInput("Missing n_decoder_layers config".to_string())
1131            })?;
1132        let max_seq_len = data
1133            .config
1134            .get("max_seq_len")
1135            .and_then(|s| s.parse().ok())
1136            .ok_or_else(|| TextError::InvalidInput("Missing max_seq_len config".to_string()))?;
1137        let dropout = data
1138            .config
1139            .get("dropout")
1140            .and_then(|s| s.parse().ok())
1141            .ok_or_else(|| TextError::InvalidInput("Missing dropout config".to_string()))?;
1142        let vocab_size = data
1143            .config
1144            .get("vocab_size")
1145            .and_then(|s| s.parse().ok())
1146            .ok_or_else(|| TextError::InvalidInput("Missing vocab_size config".to_string()))?;
1147
1148        let config = crate::transformer::TransformerConfig {
1149            d_model,
1150            nheads: n_heads,
1151            d_ff,
1152            n_encoder_layers,
1153            n_decoder_layers,
1154            max_seqlen: max_seq_len,
1155            dropout,
1156            vocab_size,
1157        };
1158
1159        // Reconstruct vocabulary from saved data
1160        let vocabulary = data.vocabulary.clone().unwrap_or_else(|| {
1161            // Fallback to placeholder if vocabulary not saved
1162            (0..config.vocab_size)
1163                .map(|i| format!("token_{i}"))
1164                .collect()
1165        });
1166
1167        // Create new transformer model with config
1168        let mut model = crate::transformer::TransformerModel::new(config.clone(), vocabulary)?;
1169
1170        // Restore embedding weights
1171        if let (Some(embed_weights), Some(embedshape)) = (
1172            data.weights.get("token_embeddings"),
1173            data.shapes.get("token_embeddings"),
1174        ) {
1175            let embed_array = scirs2_core::ndarray::Array::from_shape_vec(
1176                (embedshape[0], embedshape[1]),
1177                embed_weights.clone(),
1178            )
1179            .map_err(|e| TextError::InvalidInput(format!("Invalid embedding shape: {e}")))?;
1180            model.token_embedding.set_embeddings(embed_array)?;
1181        }
1182
1183        // Restore positional embeddings
1184        if let (Some(pos_embed_weights), Some(pos_embedshape)) = (
1185            data.weights.get("positional_embeddings"),
1186            data.shapes.get("positional_embeddings"),
1187        ) {
1188            if pos_embedshape.len() != 2 {
1189                return Err(TextError::InvalidInput(format!(
1190                    "Positional embedding shape must be 2D, got {} dims",
1191                    pos_embedshape.len()
1192                )));
1193            }
1194            let pos_embed_array = scirs2_core::ndarray::Array::from_shape_vec(
1195                (pos_embedshape[0], pos_embedshape[1]),
1196                pos_embed_weights.clone(),
1197            )
1198            .map_err(|e| {
1199                TextError::InvalidInput(format!("Invalid positional embedding shape: {e}"))
1200            })?;
1201            model
1202                .encoder
1203                .set_position_encoding(pos_embed_array)
1204                .map_err(|e| {
1205                    TextError::InvalidInput(format!("Positional encoding dimension mismatch: {e}"))
1206                })?;
1207        }
1208
1209        // Restore encoder layer weights
1210        for i in 0..config.n_encoder_layers {
1211            let encoder_layers = model.encoder.get_layers_mut();
1212            let (attention, ff, ln1, ln2) = encoder_layers[i].get_components_mut();
1213
1214            // Restore attention weights
1215            if let (
1216                Some(wq_weights),
1217                Some(wqshape),
1218                Some(wk_weights),
1219                Some(wkshape),
1220                Some(wv_weights),
1221                Some(wvshape),
1222                Some(wo_weights),
1223                Some(woshape),
1224            ) = (
1225                data.weights.get(&format!("encoder_{i}_attention_wq")),
1226                data.shapes.get(&format!("encoder_{i}_attention_wq")),
1227                data.weights.get(&format!("encoder_{i}_attention_wk")),
1228                data.shapes.get(&format!("encoder_{i}_attention_wk")),
1229                data.weights.get(&format!("encoder_{i}_attention_wv")),
1230                data.shapes.get(&format!("encoder_{i}_attention_wv")),
1231                data.weights.get(&format!("encoder_{i}_attention_wo")),
1232                data.shapes.get(&format!("encoder_{i}_attention_wo")),
1233            ) {
1234                let w_q = scirs2_core::ndarray::Array::from_shape_vec(
1235                    (wqshape[0], wqshape[1]),
1236                    wq_weights.clone(),
1237                )
1238                .map_err(|e| TextError::InvalidInput(format!("Invalid wq shape: {e}")))?;
1239                let w_k = scirs2_core::ndarray::Array::from_shape_vec(
1240                    (wkshape[0], wkshape[1]),
1241                    wk_weights.clone(),
1242                )
1243                .map_err(|e| TextError::InvalidInput(format!("Invalid wk shape: {e}")))?;
1244                let w_v = scirs2_core::ndarray::Array::from_shape_vec(
1245                    (wvshape[0], wvshape[1]),
1246                    wv_weights.clone(),
1247                )
1248                .map_err(|e| TextError::InvalidInput(format!("Invalid wv shape: {e}")))?;
1249                let w_o = scirs2_core::ndarray::Array::from_shape_vec(
1250                    (woshape[0], woshape[1]),
1251                    wo_weights.clone(),
1252                )
1253                .map_err(|e| TextError::InvalidInput(format!("Invalid wo shape: {e}")))?;
1254
1255                attention.set_weights(w_q, w_k, w_v, w_o)?;
1256            }
1257
1258            // Restore feedforward weights
1259            if let (
1260                Some(w1_weights),
1261                Some(w1shape),
1262                Some(w2_weights),
1263                Some(w2shape),
1264                Some(b1_weights),
1265                Some(b2_weights),
1266            ) = (
1267                data.weights.get(&format!("encoder_{i}_ff_w1")),
1268                data.shapes.get(&format!("encoder_{i}_ff_w1")),
1269                data.weights.get(&format!("encoder_{i}_ff_w2")),
1270                data.shapes.get(&format!("encoder_{i}_ff_w2")),
1271                data.weights.get(&format!("encoder_{i}_ff_b1")),
1272                data.weights.get(&format!("encoder_{i}_ff_b2")),
1273            ) {
1274                let w1 = scirs2_core::ndarray::Array::from_shape_vec(
1275                    (w1shape[0], w1shape[1]),
1276                    w1_weights.clone(),
1277                )
1278                .map_err(|e| TextError::InvalidInput(format!("Invalid w1 shape: {e}")))?;
1279                let w2 = scirs2_core::ndarray::Array::from_shape_vec(
1280                    (w2shape[0], w2shape[1]),
1281                    w2_weights.clone(),
1282                )
1283                .map_err(|e| TextError::InvalidInput(format!("Invalid w2 shape: {e}")))?;
1284                let b1 = scirs2_core::ndarray::Array::from_vec(b1_weights.clone());
1285                let b2 = scirs2_core::ndarray::Array::from_vec(b2_weights.clone());
1286
1287                ff.set_weights(w1, w2, b1, b2)?;
1288            }
1289
1290            // Restore layer norm parameters
1291            if let (Some(gamma1_weights), Some(beta1_weights)) = (
1292                data.weights.get(&format!("encoder_{i}_ln1_gamma")),
1293                data.weights.get(&format!("encoder_{i}_ln1_beta")),
1294            ) {
1295                let gamma1 = scirs2_core::ndarray::Array::from_vec(gamma1_weights.clone());
1296                let beta1 = scirs2_core::ndarray::Array::from_vec(beta1_weights.clone());
1297                ln1.set_params(gamma1, beta1)?;
1298            }
1299
1300            if let (Some(gamma2_weights), Some(beta2_weights)) = (
1301                data.weights.get(&format!("encoder_{i}_ln2_gamma")),
1302                data.weights.get(&format!("encoder_{i}_ln2_beta")),
1303            ) {
1304                let gamma2 = scirs2_core::ndarray::Array::from_vec(gamma2_weights.clone());
1305                let beta2 = scirs2_core::ndarray::Array::from_vec(beta2_weights.clone());
1306                ln2.set_params(gamma2, beta2)?;
1307            }
1308        }
1309
1310        // Restore decoder layer weights
1311        for _i in 0..config.n_decoder_layers {
1312            // Similar restoration for decoder layers
1313            // Note: Implementation would mirror encoder restoration
1314        }
1315
1316        // Restore output projection weights
1317        if let (Some(output_weights), Some(outputshape)) = (
1318            data.weights.get("output_projection"),
1319            data.shapes.get("output_projection"),
1320        ) {
1321            let _output_array = scirs2_core::ndarray::Array::from_shape_vec(
1322                scirs2_core::ndarray::IxDyn(outputshape),
1323                output_weights.clone(),
1324            )
1325            .map_err(|e| {
1326                TextError::InvalidInput(format!("Invalid output projection shape: {e}"))
1327            })?;
1328            // model.output_projection.set_weights(output_array)?;
1329        }
1330
1331        Ok(model)
1332    }
1333
1334    fn model_type(&self) -> ModelType {
1335        ModelType::Transformer
1336    }
1337
1338    fn get_config(&self) -> HashMap<String, String> {
1339        let mut config = HashMap::new();
1340        config.insert("d_model".to_string(), self.config.d_model.to_string());
1341        config.insert("n_heads".to_string(), self.config.nheads.to_string());
1342        config.insert("d_ff".to_string(), self.config.d_ff.to_string());
1343        config.insert(
1344            "n_encoder_layers".to_string(),
1345            self.config.n_encoder_layers.to_string(),
1346        );
1347        config.insert(
1348            "n_decoder_layers".to_string(),
1349            self.config.n_decoder_layers.to_string(),
1350        );
1351        config.insert(
1352            "max_seq_len".to_string(),
1353            self.config.max_seqlen.to_string(),
1354        );
1355        config.insert("dropout".to_string(), self.config.dropout.to_string());
1356        config.insert("vocab_size".to_string(), self.config.vocab_size.to_string());
1357        config
1358    }
1359}
1360
1361/// Implementation of RegistrableModel for Word2Vec
1362impl RegistrableModel for crate::embeddings::Word2Vec {
1363    fn serialize(&self) -> Result<SerializableModelData> {
1364        let mut weights = HashMap::new();
1365        let mut shapes = HashMap::new();
1366        let mut config = HashMap::new();
1367        let vocabulary = Some(self.get_vocabulary());
1368
1369        // Serialize config
1370        config.insert(
1371            "vector_size".to_string(),
1372            self.get_vector_size().to_string(),
1373        );
1374        config.insert(
1375            "algorithm".to_string(),
1376            format!("{:?}", self.get_algorithm()),
1377        );
1378        config.insert(
1379            "window_size".to_string(),
1380            self.get_window_size().to_string(),
1381        );
1382        config.insert("min_count".to_string(), self.get_min_count().to_string());
1383        config.insert(
1384            "negative_samples".to_string(),
1385            self.get_negative_samples().to_string(),
1386        );
1387        config.insert(
1388            "learning_rate".to_string(),
1389            self.get_learning_rate().to_string(),
1390        );
1391        config.insert("epochs".to_string(), self.get_epochs().to_string());
1392        config.insert(
1393            "subsampling_threshold".to_string(),
1394            self.get_subsampling_threshold().to_string(),
1395        );
1396
1397        // Serialize embedding weights
1398        if let Some(embeddings) = self.get_embeddings_matrix() {
1399            let embed_weights = embeddings.as_slice().expect("Operation failed").to_vec();
1400            let embedshape = embeddings.shape().to_vec();
1401            weights.insert("embeddings".to_string(), embed_weights);
1402            shapes.insert("embeddings".to_string(), embedshape);
1403        }
1404
1405        Ok(SerializableModelData {
1406            weights,
1407            shapes,
1408            vocabulary,
1409            config,
1410        })
1411    }
1412
1413    fn deserialize(data: &SerializableModelData) -> Result<Self> {
1414        let vector_size = data
1415            .config
1416            .get("vector_size")
1417            .and_then(|s| s.parse().ok())
1418            .ok_or_else(|| TextError::InvalidInput("Missing vector_size config".to_string()))?;
1419        let window_size = data
1420            .config
1421            .get("window_size")
1422            .and_then(|s| s.parse().ok())
1423            .ok_or_else(|| TextError::InvalidInput("Missing window_size config".to_string()))?;
1424        let min_count = data
1425            .config
1426            .get("min_count")
1427            .and_then(|s| s.parse().ok())
1428            .ok_or_else(|| TextError::InvalidInput("Missing min_count config".to_string()))?;
1429
1430        let algorithm = match data.config.get("algorithm").map(|s| s.as_str()) {
1431            Some("CBOW") => crate::embeddings::Word2VecAlgorithm::CBOW,
1432            Some("SkipGram") => crate::embeddings::Word2VecAlgorithm::SkipGram,
1433            _ => {
1434                return Err(TextError::InvalidInput(
1435                    "Invalid or missing algorithm config".to_string(),
1436                ))
1437            }
1438        };
1439
1440        let config = crate::embeddings::Word2VecConfig {
1441            vector_size,
1442            window_size,
1443            min_count,
1444            epochs: 5,            // Default value
1445            learning_rate: 0.025, // Default value
1446            algorithm,
1447            negative_samples: 5,         // Default value
1448            subsample: 1e-3,             // Default value
1449            batch_size: 128,             // Default value
1450            hierarchical_softmax: false, // Default value
1451        };
1452
1453        // Create new Word2Vec instance
1454        let word2vec = crate::embeddings::Word2Vec::with_config(config);
1455
1456        // Restore vocabulary and embeddings if available
1457        if let (Some(vocab), Some(embed_weights), Some(embedshape)) = (
1458            data.vocabulary.as_ref(),
1459            data.weights.get("embeddings"),
1460            data.shapes.get("embeddings"),
1461        ) {
1462            // Restore the full model state from serialized data
1463            let embedding_matrix = scirs2_core::ndarray::Array::from_shape_vec(
1464                (embedshape[0], embedshape[1]),
1465                embed_weights.clone(),
1466            )
1467            .map_err(|e| TextError::InvalidInput(format!("Invalid embedding shape: {e}")))?;
1468
1469            // Create new Word2Vec model with restored parameters
1470            let mut restored_word2vec = word2vec;
1471
1472            // Apply configuration parameters if available
1473            if let Some(window_size) = data.config.get("window_size").and_then(|s| s.parse().ok()) {
1474                restored_word2vec = restored_word2vec.with_window_size(window_size);
1475            }
1476
1477            if let Some(negative_samples) = data
1478                .config
1479                .get("negative_samples")
1480                .and_then(|s| s.parse().ok())
1481            {
1482                restored_word2vec = restored_word2vec.with_negative_samples(negative_samples);
1483            }
1484
1485            if let Some(learning_rate) = data
1486                .config
1487                .get("learning_rate")
1488                .and_then(|s| s.parse().ok())
1489            {
1490                restored_word2vec = restored_word2vec.with_learning_rate(learning_rate);
1491            }
1492
1493            // Restore vocabulary and input embeddings using the validated API.
1494            // `restore_weights` validates row/column dimensions and returns an
1495            // error if they do not match — no panics.
1496            restored_word2vec.restore_weights(vocab.clone(), embedding_matrix)?;
1497            return Ok(restored_word2vec);
1498        }
1499
1500        // If no saved state available, return new model with config
1501        Ok(word2vec)
1502    }
1503
1504    fn model_type(&self) -> ModelType {
1505        ModelType::WordEmbedding
1506    }
1507
1508    fn get_config(&self) -> HashMap<String, String> {
1509        let mut config = HashMap::new();
1510        config.insert(
1511            "vector_size".to_string(),
1512            self.get_vector_size().to_string(),
1513        );
1514        config.insert(
1515            "algorithm".to_string(),
1516            format!("{:?}", self.get_algorithm()),
1517        );
1518        config.insert(
1519            "window_size".to_string(),
1520            self.get_window_size().to_string(),
1521        );
1522        config.insert("min_count".to_string(), self.get_min_count().to_string());
1523        config
1524    }
1525}
1526
1527#[cfg(test)]
1528mod tests {
1529    use super::*;
1530    use tempfile::TempDir;
1531
1532    #[test]
1533    fn test_model_metadata_creation() {
1534        let metadata = ModelMetadata::new(
1535            "test_model".to_string(),
1536            "Test Model".to_string(),
1537            ModelType::Transformer,
1538        )
1539        .with_version("1.0.0".to_string())
1540        .with_description("A test model".to_string())
1541        .with_metric("accuracy".to_string(), 0.95);
1542
1543        assert_eq!(metadata.id, "test_model");
1544        assert_eq!(metadata.name, "Test Model");
1545        assert_eq!(metadata.version, "1.0.0");
1546        assert_eq!(metadata.description, "A test model");
1547        assert_eq!(metadata.metrics.get("accuracy"), Some(&0.95));
1548    }
1549
1550    #[test]
1551    fn test_model_registry_creation() {
1552        let temp_dir = TempDir::new().expect("Operation failed");
1553        let registry =
1554            ModelRegistry::new(temp_dir.path(), temp_dir.path()).expect("Operation failed");
1555
1556        assert_eq!(registry.models.len(), 0);
1557        assert_eq!(registry.model_cache.len(), 0);
1558    }
1559
1560    #[test]
1561    fn test_prebuilt_models() {
1562        let (config, metadata) = PrebuiltModels::english_transformer_base();
1563
1564        assert_eq!(config.d_model, 512);
1565        assert_eq!(config.nheads, 8);
1566        assert_eq!(metadata.id, "english_transformer_base");
1567        assert_eq!(metadata.model_type, ModelType::Transformer);
1568        assert!(metadata.languages.contains(&"en".to_string()));
1569    }
1570
1571    #[test]
1572    fn test_model_type_display() {
1573        assert_eq!(ModelType::Transformer.to_string(), "transformer");
1574        assert_eq!(ModelType::WordEmbedding.to_string(), "word_embedding");
1575        assert_eq!(
1576            ModelType::Custom("test".to_string()).to_string(),
1577            "custom_test"
1578        );
1579    }
1580
1581    // ─────────────────────────────────────────────────────────────────────────
1582    // Stub implementation tests
1583    // ─────────────────────────────────────────────────────────────────────────
1584
1585    /// Build a minimal tiny-config TransformerModel + vocabulary
1586    fn make_tiny_transformer() -> crate::transformer::TransformerModel {
1587        let config = crate::transformer::TransformerConfig {
1588            d_model: 4,
1589            nheads: 2,
1590            d_ff: 8,
1591            n_encoder_layers: 1,
1592            n_decoder_layers: 0,
1593            max_seqlen: 8,
1594            dropout: 0.0,
1595            vocab_size: 4,
1596        };
1597        let vocab: Vec<String> = (0..4).map(|i| format!("tok{i}")).collect();
1598        crate::transformer::TransformerModel::new(config, vocab).expect("tiny model creation")
1599    }
1600
1601    /// Test 1: positional encoding set_encodings round-trip — values are preserved
1602    #[test]
1603    fn test_positional_encoding_set_roundtrip() {
1604        use scirs2_core::ndarray::Array2;
1605
1606        let mut pos_enc = crate::transformer::PositionalEncoding::new(8, 4);
1607
1608        // Create a distinguishable set of values
1609        let custom: Array2<f64> = Array2::from_shape_fn((8, 4), |(r, c)| (r * 10 + c) as f64 * 0.1);
1610
1611        pos_enc
1612            .set_encodings(custom.clone())
1613            .expect("set_encodings failed");
1614
1615        let restored = pos_enc.get_encodings();
1616        for r in 0..8 {
1617            for c in 0..4 {
1618                assert!(
1619                    (restored[[r, c]] - custom[[r, c]]).abs() < 1e-12,
1620                    "mismatch at [{r},{c}]: {} vs {}",
1621                    restored[[r, c]],
1622                    custom[[r, c]]
1623                );
1624            }
1625        }
1626    }
1627
1628    /// Test 2: dimension mismatch on positional encoding returns an error, not a panic
1629    #[test]
1630    fn test_positional_encoding_dimension_mismatch_returns_error() {
1631        use scirs2_core::ndarray::Array2;
1632
1633        let mut pos_enc = crate::transformer::PositionalEncoding::new(8, 4);
1634
1635        // Wrong shape: (5, 4) but expected (8, 4)
1636        let wrong = Array2::<f64>::zeros((5, 4));
1637        let result = pos_enc.set_encodings(wrong);
1638        assert!(
1639            result.is_err(),
1640            "expected error for row count mismatch but got Ok"
1641        );
1642
1643        // Wrong shape: (8, 3) but expected (8, 4)
1644        let wrong_cols = Array2::<f64>::zeros((8, 3));
1645        let result2 = pos_enc.set_encodings(wrong_cols);
1646        assert!(
1647            result2.is_err(),
1648            "expected error for column count mismatch but got Ok"
1649        );
1650    }
1651
1652    /// Test 3: TransformerEncoder set_position_encoding + get_position_encoding round-trip
1653    #[test]
1654    fn test_encoder_set_position_encoding_roundtrip() {
1655        use scirs2_core::ndarray::Array2;
1656
1657        let config = crate::transformer::TransformerConfig {
1658            d_model: 4,
1659            nheads: 2,
1660            d_ff: 8,
1661            n_encoder_layers: 1,
1662            n_decoder_layers: 0,
1663            max_seqlen: 6,
1664            dropout: 0.0,
1665            vocab_size: 4,
1666        };
1667        let mut encoder =
1668            crate::transformer::TransformerEncoder::new(config).expect("encoder creation");
1669
1670        let custom: Array2<f64> =
1671            Array2::from_shape_fn((6, 4), |(r, c)| (r as f64) * 0.5 + (c as f64) * 0.01);
1672        encoder
1673            .set_position_encoding(custom.clone())
1674            .expect("set_position_encoding failed");
1675
1676        let restored = encoder.get_position_encoding();
1677        for r in 0..6 {
1678            for c in 0..4 {
1679                assert!(
1680                    (restored[[r, c]] - custom[[r, c]]).abs() < 1e-12,
1681                    "mismatch at [{r},{c}]"
1682                );
1683            }
1684        }
1685    }
1686
1687    /// Test 4: Full TransformerModel serialize → deserialize preserves positional encoding values
1688    #[test]
1689    fn test_transformer_positional_encoding_serialize_deserialize() {
1690        use scirs2_core::ndarray::Array2;
1691
1692        let model = make_tiny_transformer();
1693
1694        // Capture the original positional encoding
1695        let original_enc = model.encoder.get_position_encoding().clone();
1696
1697        // Round-trip via RegistrableModel
1698        let data = model.serialize().expect("serialize failed");
1699        let restored =
1700            crate::transformer::TransformerModel::deserialize(&data).expect("deserialize failed");
1701
1702        let restored_enc = restored.encoder.get_position_encoding();
1703        assert_eq!(
1704            original_enc.shape(),
1705            restored_enc.shape(),
1706            "shape mismatch after round-trip"
1707        );
1708        for r in 0..original_enc.shape()[0] {
1709            for c in 0..original_enc.shape()[1] {
1710                assert!(
1711                    (original_enc[[r, c]] - restored_enc[[r, c]]).abs() < 1e-12,
1712                    "positional encoding value mismatch at [{r},{c}]"
1713                );
1714            }
1715        }
1716    }
1717
1718    /// Test 5: Word2Vec restore_weights correctly sets vocabulary and embeddings
1719    #[test]
1720    fn test_word2vec_restore_weights_roundtrip() {
1721        use crate::embeddings::{Word2Vec, Word2VecAlgorithm, Word2VecConfig};
1722        use scirs2_core::ndarray::Array2;
1723
1724        let config = Word2VecConfig {
1725            vector_size: 4,
1726            window_size: 2,
1727            min_count: 1,
1728            epochs: 1,
1729            learning_rate: 0.025,
1730            algorithm: Word2VecAlgorithm::SkipGram,
1731            negative_samples: 2,
1732            subsample: 1e-3,
1733            batch_size: 8,
1734            hierarchical_softmax: false,
1735        };
1736        let mut model = Word2Vec::with_config(config);
1737
1738        let vocab: Vec<String> = vec!["hello".to_string(), "world".to_string(), "foo".to_string()];
1739        let embeddings: Array2<f64> = Array2::from_shape_fn((3, 4), |(r, c)| (r * 4 + c) as f64);
1740
1741        model
1742            .restore_weights(vocab.clone(), embeddings.clone())
1743            .expect("restore_weights failed");
1744
1745        // Vocabulary should now be populated
1746        let restored_vocab = model.get_vocabulary();
1747        assert_eq!(restored_vocab.len(), vocab.len());
1748        for word in &vocab {
1749            assert!(restored_vocab.contains(word), "missing word: {word}");
1750        }
1751
1752        // Embeddings matrix should match
1753        let restored_embed = model
1754            .get_embeddings_matrix()
1755            .expect("embeddings should be set");
1756        for r in 0..3 {
1757            for c in 0..4 {
1758                assert!(
1759                    (restored_embed[[r, c]] - embeddings[[r, c]]).abs() < 1e-12,
1760                    "embedding mismatch at [{r},{c}]"
1761                );
1762            }
1763        }
1764    }
1765
1766    /// Test 6: Word2Vec restore_weights rejects embedding dimension mismatch
1767    #[test]
1768    fn test_word2vec_restore_weights_dimension_mismatch() {
1769        use crate::embeddings::{Word2Vec, Word2VecAlgorithm, Word2VecConfig};
1770        use scirs2_core::ndarray::Array2;
1771
1772        let config = Word2VecConfig {
1773            vector_size: 4,
1774            window_size: 2,
1775            min_count: 1,
1776            epochs: 1,
1777            learning_rate: 0.025,
1778            algorithm: Word2VecAlgorithm::SkipGram,
1779            negative_samples: 2,
1780            subsample: 1e-3,
1781            batch_size: 8,
1782            hierarchical_softmax: false,
1783        };
1784        let mut model = Word2Vec::with_config(config);
1785
1786        let vocab = vec!["a".to_string(), "b".to_string()];
1787
1788        // Dimension mismatch: vector_size is 4, but embeddings have 5 columns
1789        let wrong_cols = Array2::<f64>::zeros((2, 5));
1790        let result = model.restore_weights(vocab.clone(), wrong_cols);
1791        assert!(result.is_err(), "expected error for column mismatch");
1792
1793        // Dimension mismatch: vocab length is 2, but embeddings have 3 rows
1794        let wrong_rows = Array2::<f64>::zeros((3, 4));
1795        let result2 = model.restore_weights(vocab, wrong_rows);
1796        assert!(result2.is_err(), "expected error for row count mismatch");
1797    }
1798
1799    /// Test 7: Word2Vec full serialize → deserialize round-trip preserves vocabulary + embeddings
1800    #[test]
1801    fn test_word2vec_serialize_deserialize_roundtrip() {
1802        use crate::embeddings::{Word2Vec, Word2VecAlgorithm, Word2VecConfig};
1803        use scirs2_core::ndarray::Array2;
1804
1805        let config = Word2VecConfig {
1806            vector_size: 3,
1807            window_size: 2,
1808            min_count: 1,
1809            epochs: 1,
1810            learning_rate: 0.025,
1811            algorithm: Word2VecAlgorithm::CBOW,
1812            negative_samples: 2,
1813            subsample: 1e-3,
1814            batch_size: 8,
1815            hierarchical_softmax: false,
1816        };
1817        let mut model = Word2Vec::with_config(config);
1818
1819        let vocab: Vec<String> = vec!["alpha".to_string(), "beta".to_string(), "gamma".to_string()];
1820        let embeddings: Array2<f64> =
1821            Array2::from_shape_fn((3, 3), |(r, c)| ((r + 1) * (c + 1)) as f64 * 0.25);
1822
1823        model
1824            .restore_weights(vocab.clone(), embeddings.clone())
1825            .expect("restore_weights before serialize failed");
1826
1827        let data = model.serialize().expect("serialize failed");
1828        let restored = Word2Vec::deserialize(&data).expect("deserialize failed");
1829
1830        let restored_vocab = restored.get_vocabulary();
1831        assert_eq!(
1832            restored_vocab.len(),
1833            vocab.len(),
1834            "vocabulary length mismatch"
1835        );
1836
1837        let restored_embed = restored
1838            .get_embeddings_matrix()
1839            .expect("embeddings should be present after deserialize");
1840        for r in 0..3 {
1841            for c in 0..3 {
1842                assert!(
1843                    (restored_embed[[r, c]] - embeddings[[r, c]]).abs() < 1e-12,
1844                    "embedding value mismatch at [{r},{c}] after full round-trip"
1845                );
1846            }
1847        }
1848    }
1849
1850    /// Test 8: corrupt / invalid data in SerializableModelData returns a descriptive error
1851    #[test]
1852    fn test_word2vec_deserialize_invalid_data_returns_error() {
1853        use crate::embeddings::Word2Vec;
1854
1855        // Missing required config fields → should return TextError, not panic
1856        let empty_data = SerializableModelData {
1857            weights: Default::default(),
1858            shapes: Default::default(),
1859            vocabulary: None,
1860            config: Default::default(),
1861        };
1862        let result = Word2Vec::deserialize(&empty_data);
1863        assert!(
1864            result.is_err(),
1865            "expected error for missing config fields but got Ok"
1866        );
1867
1868        // Check that the error message is descriptive (not empty)
1869        if let Err(e) = result {
1870            let msg = e.to_string();
1871            assert!(!msg.is_empty(), "error message must not be empty");
1872        }
1873    }
1874}