6.7 KiB
6.7 KiB
Model Hub
RustyTorch++ provides a unified model hub for loading models from multiple sources with a single line of code.
Quick Start
use rtx_hub::load;
// Load from HuggingFace Hub
let model = load("meta-llama/Llama-3.2-1B").await?;
// Load from RustyTorch Hub
let model = load("rustytorch/llama-3-8b-q4").await?;
// Load from local directory
let model = load("/path/to/model").await?;
One-Command Loading
The load() function automatically:
- Detects the source (HuggingFace, RustyTorch, or local)
- Downloads and caches models
- Parses configuration files
- Converts HuggingFace configs to RustyTorch format
let model = rtx_hub::load("meta-llama/Llama-3.2-3B").await?;
// Access model information
println!("Architecture: {:?}", model.architecture());
println!("Hidden size: {}", model.hidden_size());
println!("Layers: {}", model.num_layers());
println!("Vocab size: {}", model.vocab_size());
Configuration Options
use rtx_hub::{load_with_config, LoadConfig, RTXDType};
let config = LoadConfig::new()
.with_revision("v1.0.0") // Specific version
.with_dtype(RTXDType::BF16) // Target dtype
.with_device("cuda:0") // Target device
.with_cache_dir("/data/models") // Custom cache
.with_auth_token("hf_xxx") // Private models
.with_offline(true); // Use cached only
let model = load_with_config("mistralai/Mistral-7B-v0.1", config).await?;
LoadConfig Options
| Option | Description |
|---|---|
with_revision(rev) |
Git revision (branch, tag, commit) |
with_dtype(dtype) |
Target data type (F32, F16, BF16, FP8) |
with_device(dev) |
Target device ("cuda:0", "metal", "cpu") |
with_cache_dir(path) |
Custom cache directory |
with_force_download(bool) |
Re-download even if cached |
with_offline(bool) |
Use only cached models |
with_auth_token(token) |
HuggingFace auth token |
with_hf_hub() |
Force HuggingFace Hub |
with_trust_remote_code(bool) |
Trust custom model code |
Working with LoadedModel
let model = rtx_hub::load("meta-llama/Llama-3.2-1B").await?;
// Model configuration (converted to RTX format)
let config = &model.config;
println!("Architecture: {:?}", config.architecture);
println!("Hidden size: {}", config.hidden_size);
println!("Num heads: {}", config.num_heads);
println!("Num KV heads: {}", config.num_kv_heads); // GQA support
// Original HuggingFace config
let hf_config = &model.hf_config;
println!("Model type: {:?}", hf_config.model_type);
println!("RoPE theta: {:?}", hf_config.rope_theta);
// Tokenizer config (if available)
if let Some(tok) = &model.tokenizer_config {
println!("Max length: {:?}", tok.model_max_length);
println!("Chat template: {:?}", tok.chat_template.is_some());
}
// Generation config (if available)
if let Some(gen) = &model.generation_config {
println!("Temperature: {:?}", gen.temperature);
println!("Top-p: {:?}", gen.top_p);
}
Loading Weights
SafeTensors (Recommended)
// Single file model
if !model.is_sharded() {
let weights = model.load_safetensors().await?;
for name in weights.tensor_names() {
let info = weights.tensor_info(name).unwrap();
let data = weights.tensor_data(name).unwrap();
println!("{}: {:?} ({} bytes)", name, info.shape, data.len());
}
}
// Sharded model (multiple files)
if model.is_sharded() {
let weights = model.load_sharded_safetensors().await?;
println!("Loaded {} shards", weights.num_shards());
println!("Total tensors: {}", weights.num_tensors());
// Get specific tensor (searches all shards)
if let Some((info, data)) = weights.tensor_data("model.layers.0.self_attn.q_proj.weight") {
println!("Q proj shape: {:?}", info.shape);
}
}
Direct SafeTensors Loading
use rtx_hub::{SafeTensors, ShardedSafeTensors};
// Load single file
let weights = SafeTensors::load("/path/to/model.safetensors").await?;
// Load from bytes
let bytes = std::fs::read("/path/to/model.safetensors")?;
let weights = SafeTensors::from_bytes(&bytes)?;
// Validate integrity
weights.validate()?;
// Load sharded model from directory
let weights = ShardedSafeTensors::load_from_dir("/path/to/model/").await?;
Cache Management
use rtx_hub::{list_cached_models, cache_size, clear_cache};
// List cached models
let models = list_cached_models(None).await?;
for model_id in models {
println!("Cached: {}", model_id);
}
// Get cache size
let size = cache_size(None).await?;
println!("Cache: {:.2} GB", size as f64 / 1_000_000_000.0);
// Clear cache
clear_cache(None).await?;
// Custom cache directory
let custom = PathBuf::from("/data/models");
let models = list_cached_models(Some(custom)).await?;
Supported Architectures
The hub automatically detects and supports:
| Architecture | Model Type | Example |
|---|---|---|
| LLaMA | llama |
meta-llama/Llama-3.2-* |
| Mistral | mistral |
mistralai/Mistral-7B-* |
| Mixtral | mixtral |
mistralai/Mixtral-8x7B-* |
| GPT-2 | gpt2 |
openai-community/gpt2 |
| GPT-NeoX | gpt_neox |
EleutherAI/pythia-* |
| Falcon | falcon |
tiiuae/falcon-* |
| Phi | phi, phi3 |
microsoft/Phi-3-* |
| Qwen | qwen, qwen2 |
Qwen/Qwen2-* |
| BERT | bert |
google-bert/bert-* |
| RoBERTa | roberta |
FacebookAI/roberta-* |
| T5 | t5 |
google-t5/t5-* |
| Mamba | mamba |
state-spaces/mamba-* |
| Gemma | gemma |
google/gemma-* |
Environment Variables
| Variable | Description |
|---|---|
HF_TOKEN |
HuggingFace authentication token |
HUGGING_FACE_HUB_TOKEN |
Alternative token variable |
RTX_HUB_CACHE |
Custom cache directory |
Example: Complete Pipeline
use rtx_hub::{load_with_config, LoadConfig, RTXDType};
#[tokio::main]
async fn main() -> anyhow::Result<()> {
// Load model
let config = LoadConfig::new()
.with_dtype(RTXDType::BF16);
let model = load_with_config("meta-llama/Llama-3.2-3B", config).await?;
println!("Loaded: {}", model.model_id);
println!(" Architecture: {:?}", model.architecture());
println!(" Hidden: {}", model.hidden_size());
println!(" Layers: {}", model.num_layers());
println!(" Heads: {}", model.num_heads());
// Load weights
let weights = if model.is_sharded() {
model.load_sharded_safetensors().await?
} else {
// ... handle single file
};
// Now use weights with rtx-inference
// let engine = InferenceEngine::from_weights(model.config, weights)?;
Ok(())
}
Next Steps
- LLM Inference - Use loaded models for inference
- Speculative Decoding - 2-4x faster generation
- Quantization - Reduce memory usage