VRAM Calculator — Java 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 Java implementation — the same logic the interactive tool runs, in a shareable, citable form.
// vram-calculator — Java port: estimate the VRAM an LLM needs (weights + KV cache).
import java.util.List;
/** Supported weight quantizations. bytesPerParam = bytes stored per weight
* (Q4_K_M = 4.85 bits/weight, the llama.cpp mix). */
enum Quant {
FP32(4), FP16(2), BF16(2), INT8(1), INT4(0.5), Q4_K_M(4.85 / 8);
final double bytesPerParam;
Quant(double bytesPerParam) { this.bytesPerParam = bytesPerParam; }
}
/** Architecture and serving knobs; every field must be >= 1. */
record VramOptions(int layers, int kvHeads, int headDim, int batch, int kvBytes) {
/** Default attention architecture: a modern GQA-style layout (fp16 KV). */
static VramOptions defaults() { return new VramOptions(32, 8, 128, 1, 2); }
}
/** One VRAM estimate, in decimal gigabytes (GB = 10^9, matching how "70B"
* and GPU sizes are quoted). */
record VramBreakdown(Quant quant, double bytesPerParam, double weightsGB, double kvCacheGB) {
double totalGB() { return weightsGB + kvCacheGB; }
}
/** A GPU memory tier, and one card scored against a total. */
record GpuCard(String name, double sizeGB) {}
record GpuFit(String name, double sizeGB, boolean fits, double headroomGB) {}
final class VramCalculator {
static final double GB = 1e9;
/** Common GPU memory tiers, from consumer boards to datacenter cards. */
static final List<GpuCard> GPU_CARDS = List.of(
new GpuCard("RTX 3060 Ti / RTX 4060 / RX 7600", 8),
new GpuCard("RTX 3060 12 GB / RTX 4070", 12),
new GpuCard("RTX 4060 Ti 16 GB / RTX 5080", 16),
new GpuCard("RTX 3090 / RTX 4090", 24),
new GpuCard("RTX A6000 / L40S", 48),
new GpuCard("A100 80 GB / H100 / H200", 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 IllegalArgumentException on paramsB ≤ 0, negative context, or
* any option below 1 (context 0 is allowed — no context, no cache).
*/
static VramBreakdown vram(double paramsB, Quant quant, int context, VramOptions o) {
if (!Double.isFinite(paramsB) || paramsB <= 0)
throw new IllegalArgumentException("paramsB must be finite > 0 (got " + paramsB + ")");
if (context < 0) throw new IllegalArgumentException("context must be >= 0 (got " + context + ")");
if (o.layers() < 1 || o.kvHeads() < 1 || o.headDim() < 1 || o.batch() < 1 || o.kvBytes() < 1)
throw new IllegalArgumentException("every option must be >= 1 (got " + o + ")");
double weightsGB = paramsB * 1e9 * quant.bytesPerParam / GB;
double kvCacheGB = 2.0 * o.layers() * context * o.kvHeads() * o.headDim() * o.kvBytes() * o.batch() / GB;
return new VramBreakdown(quant, quant.bytesPerParam, weightsGB, kvCacheGB);
}
/** Score every card against a total footprint. `fits` is inclusive: a
* total exactly equal to the card size fits (headroom 0). */
static List<GpuFit> gpuFits(double totalGB) {
if (!Double.isFinite(totalGB) || totalGB < 0)
throw new IllegalArgumentException("totalGB must be finite >= 0 (got " + totalGB + ")");
return GPU_CARDS.stream()
.map(c -> new GpuFit(c.name(), c.sizeGB(), c.sizeGB() >= totalGB, c.sizeGB() - totalGB))
.toList();
}
public static void main(String[] args) {
// 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.
VramOptions llama70 = new VramOptions(80, 8, 128, 1, 2);
VramBreakdown r70 = vram(70, Quant.Q4_K_M, 4096, llama70);
VramBreakdown r8 = vram(8, Quant.FP16, 8192, VramOptions.defaults());
System.out.printf("70B q4_K_M @4096 (80 layers): %.2f GB weights + %.2f GB KV = %.2f GB total%n",
r70.weightsGB(), r70.kvCacheGB(), r70.totalGB());
System.out.printf("8B fp16 @8192: %.2f GB weights + %.2f GB KV = %.2f GB total%n",
r8.weightsGB(), r8.kvCacheGB(), r8.totalGB());
System.out.printf("GPU fit for %.2f GB:%n", r8.totalGB());
for (GpuFit f : gpuFits(r8.totalGB()))
System.out.printf(" %-36s %s (%.2f GB %s)%n", f.name(), f.fits() ? "fits" : "too small",
Math.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 →