Skip to content

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 →