VRAM Calculator — Swift 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 Swift implementation — the same logic the interactive tool runs, in a shareable, citable form.
// vram-calculator — Swift port: estimate the VRAM an LLM needs (weights + KV cache).
import Foundation
/// Supported weight quantizations. bytesPerParam = bytes stored per weight
/// (Q4_K_M = 4.85 bits/weight, the llama.cpp mix).
enum Quant: String, CaseIterable {
case fp32, fp16, bf16, int8, int4
case q4KM = "q4_K_M"
var bytesPerParam: Double {
switch self {
case .fp32: 4
case .fp16, .bf16: 2
case .int8: 1
case .int4: 0.5
case .q4KM: 4.85 / 8
}
}
}
/// Architecture and serving knobs; every field must be >= 1. Defaults form
/// the modern GQA-style layout (Llama-2-70B overrides layers to 80).
struct VramOptions {
var layers = 32 // transformer layers (blocks)
var kvHeads = 8 // key/value heads after GQA
var headDim = 128 // dimension of one attention head
var batch = 1 // sequences served concurrently; multiplies the KV cache
var kvBytes = 2 // bytes per KV element (fp16 K and V tensors)
}
/// One VRAM estimate, in decimal gigabytes (GB = 10^9, matching how "70B"
/// and GPU sizes are quoted).
struct VramBreakdown {
var quant: Quant
var bytesPerParam: Double
var weightsGB: Double
var kvCacheGB: Double
var totalGB: Double { weightsGB + kvCacheGB }
}
/// A GPU memory tier, and one card scored against a total.
struct GpuCard { let name: String; let sizeGB: Double }
struct GpuFit { let name: String; let sizeGB: Double; let fits: Bool; let headroomGB: Double }
enum VramError: Error { case badParams(String) }
let gb = 1e9
/// Common GPU memory tiers, from consumer boards to datacenter cards.
let gpuCards = [
GpuCard(name: "RTX 3060 Ti / RTX 4060 / RX 7600", sizeGB: 8),
GpuCard(name: "RTX 3060 12 GB / RTX 4070", sizeGB: 12),
GpuCard(name: "RTX 4060 Ti 16 GB / RTX 5080", sizeGB: 16),
GpuCard(name: "RTX 3090 / RTX 4090", sizeGB: 24),
GpuCard(name: "RTX A6000 / L40S", sizeGB: 48),
GpuCard(name: "A100 80 GB / H100 / H200", sizeGB: 80),
]
/// Estimate the VRAM footprint of a model: weights plus KV cache.
///
/// weightsGB = paramsB × bytesPerParam
/// kvCacheGB = 2 × layers × context × kvHeads × headDim × kvBytes × batch / 1e9
///
/// Throws on paramsB ≤ 0, negative context, or any option below 1
/// (context 0 is allowed — no context, no cache).
func vram(_ paramsB: Double, _ quant: Quant, _ context: Int, opts: VramOptions = VramOptions()) throws -> VramBreakdown {
guard paramsB.isFinite, paramsB > 0 else { throw VramError.badParams("paramsB must be finite > 0 (got \(paramsB))") }
guard context >= 0 else { throw VramError.badParams("context must be >= 0 (got \(context))") }
guard opts.layers >= 1, opts.kvHeads >= 1, opts.headDim >= 1, opts.batch >= 1, opts.kvBytes >= 1 else {
throw VramError.badParams("every option must be >= 1")
}
let weightsGB = paramsB * 1e9 * quant.bytesPerParam / gb
let kvCacheGB = 2 * Double(opts.layers * context * opts.kvHeads * opts.headDim * opts.kvBytes * opts.batch) / gb
return VramBreakdown(quant: quant, bytesPerParam: quant.bytesPerParam, weightsGB: weightsGB, kvCacheGB: kvCacheGB)
}
/// Score every card against a total footprint. `fits` is inclusive: a total
/// exactly equal to the card size fits (headroom 0).
func gpuFits(_ totalGB: Double, cards: [GpuCard] = gpuCards) throws -> [GpuFit] {
guard totalGB.isFinite, totalGB >= 0 else { throw VramError.badParams("totalGB must be finite >= 0 (got \(totalGB))") }
return cards.map { GpuFit(name: $0.name, sizeGB: $0.sizeGB, fits: $0.sizeGB >= totalGB, headroomGB: $0.sizeGB - totalGB) }
}
// Demo — the lib's canonical vectors: a 70B Llama-2-shape build (80 layers)
// and an 8B default-arch build, then the GPU fit table for the 8B total.
let r70 = try vram(70, .q4KM, 4096, opts: VramOptions(layers: 80))
let r8 = try vram(8, .fp16, 8192)
print(String(format: "70B q4_K_M @4096 (80 layers): %.2f GB weights + %.2f GB KV = %.2f GB total", r70.weightsGB, r70.kvCacheGB, r70.totalGB))
print(String(format: "8B fp16 @8192: %.2f GB weights + %.2f GB KV = %.2f GB total", r8.weightsGB, r8.kvCacheGB, r8.totalGB))
print(String(format: "GPU fit for %.2f GB:", r8.totalGB))
for f in try gpuFits(r8.totalGB) {
print(String(format: " %-36s %@ (%.2f GB %@)", f.name, f.fits ? "fits" : "too small", abs(f.headroomGB), f.fits ? "headroom" : "short"))
}
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 →