//! # RTX Hub - Model Registry and Versioning System //! //! RTX Hub provides a comprehensive model registry and versioning system for //! RustyTorch production deployments. It handles model packaging, versioning, //! metadata management, and dependency tracking. //! //! ## Core Features //! //! - Content-addressable model storage //! - Semantic versioning with dependency resolution //! - Model metadata and configuration management //! - Multi-backend storage support (local, S3, GCS, Azure) //! - Model validation and integrity checking //! - Registry synchronization and replication //! - HuggingFace-style pretrained model loading API //! - Comprehensive model search and discovery //! - SafeTensors file format support //! - HuggingFace Hub compatibility //! //! ## Quick Start //! //! Load a model with a single line: //! //! ```rust,ignore //! use rtx_hub::load; //! //! // Load a model from RustyTorch Hub //! let model = load("rustytorch/llama-3-8b-q4").await?; //! //! // Or from HuggingFace Hub //! let model = load("meta-llama/Llama-3-8B").await?; //! ``` //! //! ## Loading Options //! //! ```rust,ignore //! use rtx_hub::{load_with_config, LoadConfig}; //! //! let config = LoadConfig::default() //! .with_revision("v1.0.0") //! .with_dtype(DType::BF16) //! .with_device("cuda:0"); //! //! let model = load_with_config("rustytorch/llama-3-8b", config).await?; //! ``` #![deny(clippy::unwrap_used)] #![cfg_attr(test, allow(clippy::unwrap_used))] pub mod discovery; pub mod error; pub mod huggingface; pub mod model; pub mod model_card; pub mod packaging; pub mod pretrained; pub mod registry; pub mod safetensors; pub mod storage; pub mod versioning; pub mod weights; pub use discovery::{ ModelDiscovery, ModelMetrics, ModelSummary, SearchQuery, SearchResult, SortBy, }; pub use error::{HubError, HubResult}; pub use huggingface::{ HFGenerationConfig, HFHubClient, HFModelConfig, HFTokenizerConfig, RTXActivation, RTXArchitecture, RTXDType, RTXModelConfig, convert_hf_config, }; pub use model::{ModelId, ModelInfo, ModelMetadata, ModelStatus}; pub use model_card::{ CO2Emissions, MetricResult, ModelCardBuilder, ModelCardDoc, ModelIndex, Widget, }; pub use packaging::{ModelPackage, ModelPackager, PackagingOptions}; pub use pretrained::{FromPretrained, LocalCache, PretrainedConfig, PretrainedLoader}; pub use registry::{ModelRegistry, QueryOptions, Registry, RegistryConfig, SortField}; pub use safetensors::{ SafeTensors, SafeTensorsBuilder, SafeTensorsDType, ShardedSafeTensors, TensorInfo, }; pub use storage::{StorageBackend, StorageConfig}; pub use versioning::{ DependencySpec, ModelVersion, ResolvedDependency, VersionConstraint, VersionResolver, }; pub use weights::{ShardedWeights, WeightManager, WeightShard}; use std::path::PathBuf; use tracing::info; // ============================================================================ // One-Command Model Loading API // ============================================================================ /// Configuration for model loading. #[derive(Debug, Clone, Default)] pub struct LoadConfig { /// Git-like revision (e.g., "main", "v1.0.0") pub revision: Option, /// Target data type for weights pub dtype: Option, /// Target device (e.g., "cuda:0", "metal", "cpu") pub device: Option, /// Custom cache directory pub cache_dir: Option, /// Force re-download even if cached pub force_download: bool, /// Offline mode (local files only) pub offline: bool, /// Authentication token for private models pub auth_token: Option, /// Use HuggingFace Hub instead of RustyTorch Hub pub use_hf_hub: bool, /// Trust remote code (for custom model classes) pub trust_remote_code: bool, } impl LoadConfig { /// Create a new load configuration with defaults. pub fn new() -> Self { Self::default() } /// Set the revision/tag. pub fn with_revision(mut self, revision: impl Into) -> Self { self.revision = Some(revision.into()); self } /// Set the target data type. pub fn with_dtype(mut self, dtype: RTXDType) -> Self { self.dtype = Some(dtype); self } /// Set the target device. pub fn with_device(mut self, device: impl Into) -> Self { self.device = Some(device.into()); self } /// Set the cache directory. pub fn with_cache_dir(mut self, cache_dir: impl Into) -> Self { self.cache_dir = Some(cache_dir.into()); self } /// Enable force download. pub fn with_force_download(mut self, force: bool) -> Self { self.force_download = force; self } /// Enable offline mode. pub fn with_offline(mut self, offline: bool) -> Self { self.offline = offline; self } /// Set authentication token. pub fn with_auth_token(mut self, token: impl Into) -> Self { self.auth_token = Some(token.into()); self } /// Use HuggingFace Hub. pub fn with_hf_hub(mut self) -> Self { self.use_hf_hub = true; self } /// Trust remote code. pub fn with_trust_remote_code(mut self, trust: bool) -> Self { self.trust_remote_code = trust; self } /// Get the cache directory, using default if not set. pub fn get_cache_dir(&self) -> PathBuf { self.cache_dir.clone().unwrap_or_else(|| { dirs::home_dir() .unwrap_or_else(|| PathBuf::from(".")) .join(".cache") .join("rtx-hub") }) } /// Convert to PretrainedConfig. pub fn to_pretrained_config(&self) -> PretrainedConfig { PretrainedConfig { revision: self.revision.clone(), cache_dir: self.cache_dir.clone(), force_download: self.force_download, resume_download: true, local_files_only: self.offline, use_auth_token: self.auth_token.clone(), mirror: None, } } } /// Loaded model with configuration and weights path. #[derive(Debug)] pub struct LoadedModel { /// Model ID pub model_id: String, /// Path to cached model files pub path: PathBuf, /// Model configuration (converted from HF format) pub config: RTXModelConfig, /// Original HuggingFace configuration pub hf_config: HFModelConfig, /// Tokenizer configuration (if available) pub tokenizer_config: Option, /// Generation configuration (if available) pub generation_config: Option, /// Load configuration used pub load_config: LoadConfig, } impl LoadedModel { /// Get the path to the weights file. pub fn weights_path(&self) -> PathBuf { // Check for sharded model first let index_path = self.path.join("model.safetensors.index.json"); if index_path.exists() { return self.path.clone(); } // Single file let safetensors_path = self.path.join("model.safetensors"); if safetensors_path.exists() { return safetensors_path; } // Fallback to PyTorch format self.path.join("pytorch_model.bin") } /// Check if model uses sharded weights. pub fn is_sharded(&self) -> bool { self.path.join("model.safetensors.index.json").exists() } /// Load SafeTensors weights. pub async fn load_safetensors(&self) -> HubResult { let weights_path = self.weights_path(); if weights_path.is_dir() { // Sharded model - return error, use load_sharded_safetensors instead return Err(HubError::InvalidPackage { reason: "Model uses sharded weights. Use load_sharded_safetensors() instead." .to_string(), }); } SafeTensors::load(&weights_path).await } /// Load sharded SafeTensors weights. pub async fn load_sharded_safetensors(&self) -> HubResult { ShardedSafeTensors::load_from_dir(&self.path).await } /// Get model architecture. pub fn architecture(&self) -> RTXArchitecture { self.config.architecture } /// Get vocabulary size. pub fn vocab_size(&self) -> usize { self.config.vocab_size } /// Get hidden size. pub fn hidden_size(&self) -> usize { self.config.hidden_size } /// Get number of layers. pub fn num_layers(&self) -> usize { self.config.num_layers } /// Get number of attention heads. pub fn num_heads(&self) -> usize { self.config.num_heads } } /// Load a model with default configuration. /// /// This is the simplest way to load a model: /// /// ```rust,ignore /// let model = rtx_hub::load("rustytorch/llama-3-8b-q4").await?; /// ``` /// /// Model IDs can be: /// - RustyTorch Hub: `rustytorch/model-name` /// - HuggingFace Hub: `meta-llama/Llama-3-8B` /// - Local path: `/path/to/model` pub async fn load(model_id: &str) -> HubResult { load_with_config(model_id, LoadConfig::default()).await } /// Load a model with custom configuration. /// /// ```rust,ignore /// let config = LoadConfig::default() /// .with_revision("v1.0.0") /// .with_dtype(RTXDType::BF16); /// /// let model = rtx_hub::load_with_config("rustytorch/llama-3-8b", config).await?; /// ``` pub async fn load_with_config(model_id: &str, config: LoadConfig) -> HubResult { info!("Loading model: {}", model_id); // Check if it's a local path let local_path = PathBuf::from(model_id); if local_path.exists() && local_path.is_dir() { return load_from_local(model_id, local_path, config).await; } // Determine which hub to use let use_hf = config.use_hf_hub || is_hf_model_id(model_id); if use_hf { load_from_hf_hub(model_id, config).await } else { load_from_rtx_hub(model_id, config).await } } /// Check if model ID looks like a HuggingFace model. fn is_hf_model_id(model_id: &str) -> bool { // HuggingFace models typically have format "org/model" or "user/model" // RustyTorch hub uses "rustytorch/model" prefix if model_id.starts_with("rustytorch/") { return false; } // Known HuggingFace organizations let hf_orgs = [ "meta-llama/", "mistralai/", "openai/", "google/", "facebook/", "microsoft/", "EleutherAI/", "huggingface/", "bigscience/", "stabilityai/", "Qwen/", ]; hf_orgs.iter().any(|org| model_id.starts_with(org)) } /// Load model from local directory. async fn load_from_local( model_id: &str, path: PathBuf, config: LoadConfig, ) -> HubResult { info!("Loading local model from: {}", path.display()); // Load HuggingFace config let config_path = path.join("config.json"); if !config_path.exists() { return Err(HubError::InvalidPackage { reason: format!("config.json not found in {}", path.display()), }); } let config_content = tokio::fs::read_to_string(&config_path).await?; let hf_config: HFModelConfig = serde_json::from_str(&config_content)?; // Convert to RTX config let rtx_config = convert_hf_config(&hf_config)?; // Load optional configs let tokenizer_config = load_optional_json::(&path.join("tokenizer_config.json")).await; let generation_config = load_optional_json::(&path.join("generation_config.json")).await; Ok(LoadedModel { model_id: model_id.to_string(), path, config: rtx_config, hf_config, tokenizer_config, generation_config, load_config: config, }) } /// Load model from HuggingFace Hub. async fn load_from_hf_hub(model_id: &str, config: LoadConfig) -> HubResult { info!("Loading from HuggingFace Hub: {}", model_id); let cache_dir = config.get_cache_dir(); let mut client = HFHubClient::new(cache_dir); if let Some(ref token) = config.auth_token { client = client.with_token(token); } let pretrained_config = config.to_pretrained_config(); let model_path = client.download_model(model_id, &pretrained_config).await?; // Load configs let hf_config = client.load_config(&model_path).await?; let rtx_config = convert_hf_config(&hf_config)?; let tokenizer_config = client.load_tokenizer_config(&model_path).await.ok(); let generation_config = client.load_generation_config(&model_path).await.ok(); Ok(LoadedModel { model_id: model_id.to_string(), path: model_path, config: rtx_config, hf_config, tokenizer_config, generation_config, load_config: config, }) } /// Load model from RustyTorch Hub. async fn load_from_rtx_hub(model_id: &str, config: LoadConfig) -> HubResult { info!("Loading from RustyTorch Hub: {}", model_id); // For now, fall back to HuggingFace Hub loading // In the future, this will connect to the RustyTorch model registry load_from_hf_hub(model_id, config.clone()) .await .map_err(|_| HubError::ModelNotFound { model_id: format!( "{} (RustyTorch Hub not available, try with .with_hf_hub())", model_id ), }) } /// Load optional JSON file. async fn load_optional_json(path: &PathBuf) -> Option { if !path.exists() { return None; } match tokio::fs::read_to_string(path).await { Ok(content) => serde_json::from_str(&content).ok(), Err(_) => None, } } // ============================================================================ // Model Loading Utilities // ============================================================================ /// List available models in cache. pub async fn list_cached_models(cache_dir: Option) -> HubResult> { let cache_dir = cache_dir.unwrap_or_else(|| { dirs::home_dir() .unwrap_or_else(|| PathBuf::from(".")) .join(".cache") .join("rtx-hub") .join("models") }); if !cache_dir.exists() { return Ok(Vec::new()); } let mut models = Vec::new(); let mut entries = tokio::fs::read_dir(&cache_dir).await?; while let Some(entry) = entries.next_entry().await? { let name = entry.file_name().to_string_lossy().to_string(); // Convert back from cache format (namespace--model) to model ID let model_id = name.replace("--", "/"); models.push(model_id); } Ok(models) } /// Clear cached models. pub async fn clear_cache(cache_dir: Option) -> HubResult<()> { let cache_dir = cache_dir.unwrap_or_else(|| { dirs::home_dir() .unwrap_or_else(|| PathBuf::from(".")) .join(".cache") .join("rtx-hub") }); if cache_dir.exists() { tokio::fs::remove_dir_all(&cache_dir).await?; } info!("Cache cleared: {}", cache_dir.display()); Ok(()) } /// Get cache size in bytes. pub async fn cache_size(cache_dir: Option) -> HubResult { let cache_dir = cache_dir.unwrap_or_else(|| { dirs::home_dir() .unwrap_or_else(|| PathBuf::from(".")) .join(".cache") .join("rtx-hub") }); if !cache_dir.exists() { return Ok(0); } calculate_dir_size(&cache_dir).await } /// Calculate directory size recursively. async fn calculate_dir_size(dir: &PathBuf) -> HubResult { let mut total_size = 0u64; let mut entries = tokio::fs::read_dir(dir).await?; while let Some(entry) = entries.next_entry().await? { let path = entry.path(); let metadata = tokio::fs::metadata(&path).await?; if metadata.is_file() { total_size += metadata.len(); } else if metadata.is_dir() { total_size += Box::pin(calculate_dir_size(&path)).await?; } } Ok(total_size) } #[cfg(test)] mod tests { use super::*; use tempfile::TempDir; /// Integration test setup helper pub struct TestRegistry { pub registry: ModelRegistry, pub temp_dir: TempDir, } impl TestRegistry { pub async fn new() -> HubResult { let temp_dir = TempDir::new()?; let config = RegistryConfig { storage: StorageConfig::Local { base_path: temp_dir.path().to_path_buf(), }, database_url: format!("sqlite://{}/registry.db", temp_dir.path().display()), enable_validation: true, enable_compression: true, ..Default::default() }; let registry = ModelRegistry::new(config).await?; Ok(Self { registry, temp_dir }) } } }