Skip to content

VRAM Calculator — Go 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 Go 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: Go (1.21+, standard library only)
// Source: CosmoDev polyglot showcase port of the VRAM Calculator tool,
// ported from src/lib/vramCalculator.ts (the canonical TypeScript
// implementation). The live Go CLI twin — cli/vram-calculator/ in the
// CosmoDev repo, powering `cosmodev vram-calculator` — is the
// authoritative Go source; this snippet restates its core math.
// 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:
//   - Weights: params × bytes-per-param (the quantization's footprint).
//   - KV cache: 2 (K and V) × layers × context × kvHeads × headDim × batch
//     × bytes per KV element.
package vramcalculator

import (
	"fmt"
	"math"
	"strings"
)

// Quant is a supported weight quantization.
type Quant string

const (
	QuantFp32 Quant = "fp32"
	QuantFp16 Quant = "fp16"
	QuantBf16 Quant = "bf16"
	QuantInt8 Quant = "int8"
	QuantInt4 Quant = "int4"
	QuantQ4KM Quant = "q4_K_M"
)

// Quants lists every valid Quant, in the order the UI lists them.
var Quants = []Quant{QuantFp32, QuantFp16, QuantBf16, QuantInt8, QuantInt4, QuantQ4KM}

// QuantBytes is the quantization table: bytes stored per weight per format
// (Q4_K_M = 4.85 bits/weight, the llama.cpp mix).
var QuantBytes = map[Quant]float64{
	QuantFp32: 4,
	QuantFp16: 2,
	QuantBf16: 2,
	QuantInt8: 1,
	QuantInt4: 0.5,
	QuantQ4KM: 4.85 / 8,
}

// DefaultArch is the default attention architecture: a modern GQA-style
// layout (32 layers, 8 KV heads, 128-dim heads). Override per model.
var DefaultArch = struct {
	Layers  int
	KvHeads int
	HeadDim int
}{Layers: 32, KvHeads: 8, HeadDim: 128}

// DefaultKvBytes is the bytes per KV-cache element (fp16 K and V tensors)
// unless overridden.
const DefaultKvBytes = 2

// GB is a decimal gigabyte.
const GB = 1e9

// Options configures Vram. The zero value means "use every default" — a
// pointer field distinguishes an explicit override from the default.
type Options struct {
	Layers  *int // Transformer layers (blocks). Default 32.
	KvHeads *int // Key/value heads after GQA. Default 8.
	HeadDim *int // Dimension of one attention head. Default 128.
	Batch   *int // Concurrent sequences; multiplies the KV cache. Default 1.
	KvBytes *int // Bytes per KV element. Default 2 (fp16).
}

// Breakdown is the estimated VRAM footprint of a model.
type Breakdown struct {
	Quant         Quant
	BytesPerParam float64
	WeightsGB     float64
	KvCacheGB     float64
	TotalGB       float64
}

func requireFiniteMin(name string, value float64, min float64) (float64, error) {
	if math.IsNaN(value) || math.IsInf(value, 0) || value < min {
		return 0, fmt.Errorf("%s must be a finite number ≥ %g (got %v)", name, min, value)
	}
	return value, nil
}

func optInt(p *int, def int, name string, min int) (float64, error) {
	v := def
	if p != nil {
		v = *p
	}
	return requireFiniteMin(name, float64(v), float64(min))
}

// Vram estimates the VRAM footprint of a model: weights plus KV cache.
//
//	weightsGB = paramsB × bytesPerParam
//	kvCacheGB = 2 × layers × context × kvHeads × headDim × kvBytes × batch / 1e9
//
// Returns an error on paramsB ≤ 0, an unknown quantization, negative
// context, or any option below 1 (context 0 is allowed — no context, no
// cache).
func Vram(paramsB float64, quant Quant, context float64, opts Options) (Breakdown, error) {
	if math.IsNaN(paramsB) || math.IsInf(paramsB, 0) || paramsB <= 0 {
		return Breakdown{}, fmt.Errorf("paramsB must be a finite number > 0 (got %v)", paramsB)
	}
	bytesPerParam, ok := QuantBytes[quant]
	if !ok {
		return Breakdown{}, fmt.Errorf("unknown quantization %q — expected one of %s", string(quant), joinQuants())
	}
	if _, err := requireFiniteMin("context", context, 0); err != nil {
		return Breakdown{}, err
	}
	layers, err := optInt(opts.Layers, DefaultArch.Layers, "layers", 1)
	if err != nil {
		return Breakdown{}, err
	}
	kvHeads, err := optInt(opts.KvHeads, DefaultArch.KvHeads, "kvHeads", 1)
	if err != nil {
		return Breakdown{}, err
	}
	headDim, err := optInt(opts.HeadDim, DefaultArch.HeadDim, "headDim", 1)
	if err != nil {
		return Breakdown{}, err
	}
	batch, err := optInt(opts.Batch, 1, "batch", 1)
	if err != nil {
		return Breakdown{}, err
	}
	kvBytes, err := optInt(opts.KvBytes, DefaultKvBytes, "kvBytes", 1)
	if err != nil {
		return Breakdown{}, err
	}

	weightsGB := paramsB * 1e9 * bytesPerParam / GB
	kvCacheGB := 2 * layers * context * kvHeads * headDim * kvBytes * batch / GB
	return Breakdown{
		Quant:         quant,
		BytesPerParam: bytesPerParam,
		WeightsGB:     weightsGB,
		KvCacheGB:     kvCacheGB,
		TotalGB:       weightsGB + kvCacheGB,
	}, nil
}

func joinQuants() string {
	parts := make([]string, len(Quants))
	for i, q := range Quants {
		parts[i] = string(q)
	}
	return strings.Join(parts, ", ")
}

// GpuCard is one GPU memory tier.
type GpuCard struct {
	Name   string
	SizeGB float64
}

// GpuCards lists common tiers, from consumer boards to datacenter cards.
var GpuCards = []GpuCard{
	{Name: "RTX 3060 Ti / RTX 4060 / RX 7600", SizeGB: 8},
	{Name: "RTX 3060 12 GB / RTX 4070", SizeGB: 12},
	{Name: "RTX 4060 Ti 16 GB / RTX 5080", SizeGB: 16},
	{Name: "RTX 3090 / RTX 4090", SizeGB: 24},
	{Name: "RTX A6000 / L40S", SizeGB: 48},
	{Name: "A100 80 GB / H100 / H200", SizeGB: 80},
}

// GpuFit scores one card against a total footprint.
type GpuFit struct {
	Name       string
	SizeGB     float64
	Fits       bool
	HeadroomGB float64
}

// GpuFits scores every card against a total footprint. Fits is inclusive: a
// total exactly equal to the card size fits (headroom 0). Pass a custom
// cards list to score other tiers (nil = GpuCards).
func GpuFits(totalGB float64, cards []GpuCard) ([]GpuFit, error) {
	if _, err := requireFiniteMin("totalGB", totalGB, 0); err != nil {
		return nil, err
	}
	if cards == nil {
		cards = GpuCards
	}
	rows := make([]GpuFit, len(cards))
	for i, c := range cards {
		headroom := c.SizeGB - totalGB
		rows[i] = GpuFit{Name: c.Name, SizeGB: c.SizeGB, Fits: headroom >= 0, HeadroomGB: headroom}
	}
	return rows, nil
}

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 →