Skip to content

VRAM Calculator — Rust source

Estimate the VRAM an LLM needs — weights by quantization plus the KV cache for your context and batch — and see which consumer and datacenter GPUs hold it.

This is the Rust implementation — the same logic the interactive tool runs, in a shareable, citable form.

//! vram_calculator — estimate the VRAM an LLM needs (weights + KV cache).
//!
//! Language: Rust (edition 2021, standard library only)
//! Source:   CosmoDev polyglot showcase port of the VRAM Calculator tool,
//!           ported from src/lib/vramCalculator.ts (the canonical TypeScript
//!           implementation).
//! Tool page: https://dev.cosmolabs.org/tools/vram-calculator
//! License:  display source — part of CosmoDev's polyglot tool pages.
//!
//! Sizes use decimal gigabytes (GB = 10^9 bytes). Two components are
//! estimated: weights (params × bytes-per-param) and the KV cache
//! (2 × layers × context × kv_heads × head_dim × batch × bytes per KV
//! element). Activations and CUDA-context overhead are not modeled — treat
//! the total as a floor and leave headroom on the card.

use std::fmt;

/// A supported weight quantization.
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Quant {
    Fp32,
    Fp16,
    Bf16,
    Int8,
    Int4,
    /// llama.cpp Q4_K_M mix (~4.85 bits per weight).
    Q4Km,
}

impl Quant {
    /// Every variant, in the order the UI lists them.
    pub const ALL: [Quant; 6] = [
        Quant::Fp32,
        Quant::Fp16,
        Quant::Bf16,
        Quant::Int8,
        Quant::Int4,
        Quant::Q4Km,
    ];

    /// Parse a Quant from its UI id ("fp32", "q4_K_M", …).
    pub fn from_id(id: &str) -> Option<Quant> {
        match id {
            "fp32" => Some(Quant::Fp32),
            "fp16" => Some(Quant::Fp16),
            "bf16" => Some(Quant::Bf16),
            "int8" => Some(Quant::Int8),
            "int4" => Some(Quant::Int4),
            "q4_K_M" => Some(Quant::Q4Km),
            _ => None,
        }
    }

    /// Bytes stored per weight (Q4_K_M = 4.85 bits/weight).
    pub fn bytes_per_param(self) -> f64 {
        match self {
            Quant::Fp32 => 4.0,
            Quant::Fp16 => 2.0,
            Quant::Bf16 => 2.0,
            Quant::Int8 => 1.0,
            Quant::Int4 => 0.5,
            Quant::Q4Km => 4.85 / 8.0,
        }
    }
}

impl fmt::Display for Quant {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        let id = match self {
            Quant::Fp32 => "fp32",
            Quant::Fp16 => "fp16",
            Quant::Bf16 => "bf16",
            Quant::Int8 => "int8",
            Quant::Int4 => "int4",
            Quant::Q4Km => "q4_K_M",
        };
        f.write_str(id)
    }
}

/// Default attention architecture: a modern GQA-style layout (32 layers,
/// 8 KV heads, 128-dim heads). Override per model.
pub const DEFAULT_LAYERS: f64 = 32.0;
pub const DEFAULT_KV_HEADS: f64 = 8.0;
pub const DEFAULT_HEAD_DIM: f64 = 128.0;

/// Bytes per KV-cache element (fp16 K and V tensors) unless overridden.
pub const DEFAULT_KV_BYTES: f64 = 2.0;

/// A decimal gigabyte.
pub const GB: f64 = 1e9;

/// Options for [`vram`]; `None` everywhere means "use every default".
#[derive(Debug, Clone, Copy, Default)]
pub struct Options {
    /// Transformer layers (blocks). Default 32.
    pub layers: Option<f64>,
    /// Key/value heads after GQA. Default 8.
    pub kv_heads: Option<f64>,
    /// Dimension of one attention head. Default 128.
    pub head_dim: Option<f64>,
    /// Concurrent sequences; multiplies the KV cache. Default 1.
    pub batch: Option<f64>,
    /// Bytes per KV-cache element. Default 2 (fp16).
    pub kv_bytes: Option<f64>,
}

/// The estimated VRAM footprint of a model.
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Breakdown {
    pub quant: Quant,
    /// Bytes per weight for the chosen quantization.
    pub bytes_per_param: f64,
    /// Weights footprint in GB.
    pub weights_gb: f64,
    /// KV-cache footprint in GB.
    pub kv_cache_gb: f64,
    /// `weights_gb + kv_cache_gb`.
    pub total_gb: f64,
}

/// One GPU memory tier.
#[derive(Debug, Clone, Copy)]
pub struct GpuCard {
    pub name: &'static str,
    pub size_gb: f64,
}

/// Common GPU memory tiers, from consumer boards to datacenter cards.
pub const GPU_CARDS: [GpuCard; 6] = [
    GpuCard { name: "RTX 3060 Ti / RTX 4060 / RX 7600", size_gb: 8.0 },
    GpuCard { name: "RTX 3060 12 GB / RTX 4070", size_gb: 12.0 },
    GpuCard { name: "RTX 4060 Ti 16 GB / RTX 5080", size_gb: 16.0 },
    GpuCard { name: "RTX 3090 / RTX 4090", size_gb: 24.0 },
    GpuCard { name: "RTX A6000 / L40S", size_gb: 48.0 },
    GpuCard { name: "A100 80 GB / H100 / H200", size_gb: 80.0 },
];

/// One card scored against a total footprint.
#[derive(Debug, Clone, Copy)]
pub struct GpuFit {
    pub name: &'static str,
    pub size_gb: f64,
    /// True when the total fits with zero or more GB to spare.
    pub fits: bool,
    /// `size_gb − total_gb`; negative when the card is too small.
    pub headroom_gb: f64,
}

/// Errors from invalid input — never a panic.
#[derive(Debug, Clone, PartialEq)]
pub enum VramError {
    /// `params_b` was NaN/∞/≤ 0.
    InvalidParams(f64),
    /// The quantization id is not in `Quant::ALL`.
    UnknownQuant(String),
    /// A named value was non-finite or below its minimum.
    InvalidValue(String),
}

impl fmt::Display for VramError {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        match self {
            VramError::InvalidParams(v) => {
                write!(f, "paramsB must be a finite number > 0 (got {v})")
            }
            VramError::UnknownQuant(q) => {
                let ids: Vec<String> = Quant::ALL.iter().map(|q| q.to_string()).collect();
                write!(f, "unknown quantization {q:?} — expected one of {}", ids.join(", "))
            }
            VramError::InvalidValue(msg) => f.write_str(msg),
        }
    }
}

impl std::error::Error for VramError {}

fn require_finite_min(name: &str, value: f64, min: f64) -> Result<f64, VramError> {
    if value.is_nan() || value.is_infinite() || value < min {
        return Err(VramError::InvalidValue(format!(
            "{name} must be a finite number ≥ {min} (got {value})"
        )));
    }
    Ok(value)
}

/// Estimate the VRAM footprint of a model: weights plus KV cache.
///
/// ```text
/// weights_gb  = params_b × bytes_per_param
/// kv_cache_gb = 2 × layers × context × kv_heads × head_dim × kv_bytes × batch / 1e9
/// ```
///
/// Errors on `params_b ≤ 0`, an unknown quantization, negative context, or
/// any option below 1 (context 0 is allowed — no context, no cache).
pub fn vram(params_b: f64, quant: Quant, context: f64, opts: &Options) -> Result<Breakdown, VramError> {
    if params_b.is_nan() || params_b.is_infinite() || params_b <= 0.0 {
        return Err(VramError::InvalidParams(params_b));
    }
    let bytes_per_param = quant.bytes_per_param();
    require_finite_min("context", context, 0.0)?;
    let layers = require_finite_min("layers", opts.layers.unwrap_or(DEFAULT_LAYERS), 1.0)?;
    let kv_heads = require_finite_min("kvHeads", opts.kv_heads.unwrap_or(DEFAULT_KV_HEADS), 1.0)?;
    let head_dim = require_finite_min("headDim", opts.head_dim.unwrap_or(DEFAULT_HEAD_DIM), 1.0)?;
    let batch = require_finite_min("batch", opts.batch.unwrap_or(1.0), 1.0)?;
    let kv_bytes = require_finite_min("kvBytes", opts.kv_bytes.unwrap_or(DEFAULT_KV_BYTES), 1.0)?;

    let weights_gb = params_b * 1e9 * bytes_per_param / GB;
    let kv_cache_gb = 2.0 * layers * context * kv_heads * head_dim * kv_bytes * batch / GB;
    Ok(Breakdown {
        quant,
        bytes_per_param,
        weights_gb,
        kv_cache_gb,
        total_gb: weights_gb + kv_cache_gb,
    })
}

/// Parse a quantization id and estimate — the convenience entry point for
/// string-driven callers (mirrors the TS lib's string-keyed table).
pub fn vram_by_id(
    params_b: f64,
    quant_id: &str,
    context: f64,
    opts: &Options,
) -> Result<Breakdown, VramError> {
    let quant = Quant::from_id(quant_id).ok_or_else(|| VramError::UnknownQuant(quant_id.to_string()))?;
    vram(params_b, quant, context, opts)
}

/// Score every card against a total footprint. `fits` is inclusive: a total
/// exactly equal to the card size fits (headroom 0). Pass a custom `cards`
/// slice to score other tiers.
pub fn gpu_fits(total_gb: f64, cards: &[GpuCard]) -> Result<Vec<GpuFit>, VramError> {
    require_finite_min("totalGB", total_gb, 0.0)?;
    let cards = if cards.is_empty() { &GPU_CARDS[..] } else { cards };
    Ok(cards
        .iter()
        .map(|c| {
            let headroom = c.size_gb - total_gb;
            GpuFit {
                name: c.name,
                size_gb: c.size_gb,
                fits: headroom >= 0.0,
                headroom_gb: headroom,
            }
        })
        .collect())
}

Also available in 13 other languages

Every CosmoDev tool ships its pure logic in TypeScript (web) and Go (CLI), with authored implementations in a dozen-plus languages — the same contract, ported. Compare all languages side by side →