Files
rustytorch/docs/book/src/getting-started/model-hub.md
T
2026-03-04 00:08:42 +00:00

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

// 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