Files
rustytorch/crates/integration/rtx-candle/src/hub.rs
T
2026-03-04 00:08:42 +00:00

278 lines
7.3 KiB
Rust

//! `HuggingFace` Hub integration for model downloading
use crate::error::{CandleError, Result};
use serde::{Deserialize, Serialize};
use std::path::{Path, PathBuf};
use tracing::{debug, info};
/// Configuration for `HuggingFace` Hub access
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct HubConfig {
/// `HuggingFace` API token (optional, for private models)
pub token: Option<String>,
/// Cache directory for downloaded models
pub cache_dir: PathBuf,
/// Revision/branch to download (default: "main")
pub revision: String,
/// Prefer safetensors format
pub prefer_safetensors: bool,
}
impl Default for HubConfig {
fn default() -> Self {
let cache_dir = dirs::cache_dir()
.unwrap_or_else(|| PathBuf::from("."))
.join("rtx-candle")
.join("models");
Self {
token: std::env::var("HF_TOKEN").ok(),
cache_dir,
revision: "main".to_string(),
prefer_safetensors: true,
}
}
}
impl HubConfig {
/// Set the API token
pub fn with_token(mut self, token: impl Into<String>) -> Self {
self.token = Some(token.into());
self
}
/// Set the cache directory
pub fn with_cache_dir(mut self, dir: impl Into<PathBuf>) -> Self {
self.cache_dir = dir.into();
self
}
/// Set the revision
pub fn with_revision(mut self, revision: impl Into<String>) -> Self {
self.revision = revision.into();
self
}
}
/// Model file information
#[derive(Debug, Clone)]
pub struct ModelFile {
/// File name
pub name: String,
/// File size in bytes
pub size: u64,
/// File type
pub file_type: ModelFileType,
}
/// Type of model file
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ModelFileType {
/// `SafeTensors` weights
SafeTensors,
/// `PyTorch` weights
PyTorch,
/// GGUF quantized weights
Gguf,
/// Tokenizer config
TokenizerConfig,
/// Model config
ModelConfig,
/// Other file
Other,
}
impl ModelFileType {
/// Detect file type from name
pub fn from_filename(name: &str) -> Self {
if name.ends_with(".safetensors") {
Self::SafeTensors
} else if name.ends_with(".bin") || name.ends_with(".pt") || name.ends_with(".pth") {
Self::PyTorch
} else if name.ends_with(".gguf") {
Self::Gguf
} else if name.contains("tokenizer") && name.ends_with(".json") {
Self::TokenizerConfig
} else if name == "config.json" {
Self::ModelConfig
} else {
Self::Other
}
}
}
/// Download a model from `HuggingFace` Hub
///
/// Returns the path to the downloaded model directory.
pub fn download_model(model_id: &str, config: &HubConfig) -> Result<PathBuf> {
info!("Downloading model from HuggingFace Hub: {}", model_id);
// Sanitize model ID for directory name
let model_dir_name = model_id.replace('/', "--");
let model_dir = config
.cache_dir
.join(&model_dir_name)
.join(&config.revision);
// Check if already cached
if model_dir.exists() {
debug!("Model already cached at: {}", model_dir.display());
return Ok(model_dir);
}
// Create cache directory
std::fs::create_dir_all(&model_dir)?;
// In a real implementation, this would:
// 1. Query HuggingFace API for model files
// 2. Download weight files (preferring safetensors)
// 3. Download config.json and tokenizer files
// 4. Cache everything locally
// For now, return placeholder path
// Real implementation would use hf-hub crate or direct API calls
info!("Model would be downloaded to: {}", model_dir.display());
// Placeholder: create marker file
let marker = model_dir.join(".downloading");
std::fs::write(&marker, model_id)?;
Err(CandleError::hub(format!(
"Hub download not implemented. Would download {} to {}",
model_id,
model_dir.display()
)))
}
/// Get the path to a cached model
pub fn get_cached_model(model_id: &str, config: &HubConfig) -> Option<PathBuf> {
let model_dir_name = model_id.replace('/', "--");
let model_dir = config
.cache_dir
.join(&model_dir_name)
.join(&config.revision);
if model_dir.exists() {
Some(model_dir)
} else {
None
}
}
/// List all cached models
pub fn list_cached_models(config: &HubConfig) -> Result<Vec<String>> {
let mut models = Vec::new();
if !config.cache_dir.exists() {
return Ok(models);
}
for entry in std::fs::read_dir(&config.cache_dir)? {
let entry = entry?;
if entry.file_type()?.is_dir() {
let name = entry.file_name().to_string_lossy().replace("--", "/");
models.push(name);
}
}
Ok(models)
}
/// Clear the model cache
pub fn clear_cache(config: &HubConfig) -> Result<()> {
if config.cache_dir.exists() {
std::fs::remove_dir_all(&config.cache_dir)?;
info!("Cleared model cache: {}", config.cache_dir.display());
}
Ok(())
}
/// Get cache size in bytes
pub fn cache_size(config: &HubConfig) -> Result<u64> {
if !config.cache_dir.exists() {
return Ok(0);
}
let mut total = 0;
for entry in walkdir(config.cache_dir.as_path())? {
let metadata = entry.metadata()?;
if metadata.is_file() {
total += metadata.len();
}
}
Ok(total)
}
/// Simple directory walker
fn walkdir(path: &Path) -> Result<Vec<std::fs::DirEntry>> {
let mut entries = Vec::new();
fn walk_recursive(path: &Path, entries: &mut Vec<std::fs::DirEntry>) -> std::io::Result<()> {
for entry in std::fs::read_dir(path)? {
let entry = entry?;
let entry_path = entry.path();
entries.push(entry);
if entry_path.is_dir() {
walk_recursive(&entry_path, entries)?;
}
}
Ok(())
}
walk_recursive(path, &mut entries)?;
Ok(entries)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_hub_config_default() {
let config = HubConfig::default();
assert_eq!(config.revision, "main");
assert!(config.prefer_safetensors);
}
#[test]
fn test_hub_config_builder() {
let config = HubConfig::default()
.with_token("test_token")
.with_revision("v1.0");
assert_eq!(config.token, Some("test_token".to_string()));
assert_eq!(config.revision, "v1.0");
}
#[test]
fn test_model_file_type() {
assert_eq!(
ModelFileType::from_filename("model.safetensors"),
ModelFileType::SafeTensors
);
assert_eq!(
ModelFileType::from_filename("pytorch_model.bin"),
ModelFileType::PyTorch
);
assert_eq!(
ModelFileType::from_filename("model.gguf"),
ModelFileType::Gguf
);
assert_eq!(
ModelFileType::from_filename("config.json"),
ModelFileType::ModelConfig
);
}
#[test]
fn test_get_cached_model_not_exists() {
let config = HubConfig {
cache_dir: PathBuf::from("/nonexistent"),
..Default::default()
};
assert!(get_cached_model("test/model", &config).is_none());
}
}