VRAM-Caclulator / app.py
3v324v23's picture
Recognize ROCmFP4/ROCmFPX GGUF quants
44bdb86
Raw
History Blame
15.7 kB
"""llama.cpp VRAM Calculator — Hugging Face Space.
A Gradio app that estimates VRAM usage for a Hugging Face GGUF model given
quantization type, context length, KV cache options, YaRN context-extension
parameters, MTP heads, and a multi-GPU budget. Architecture is auto-fetched
from the GGUF header (range-read, no full download) with a manual-override
tab and presets for offline use.
"""
from __future__ import annotations
import gradio as gr
from huggingface_hub import HfApi
import spaces # noqa: F401 — present so ZeroGPU detects a GPU-aware Space
from vramcalc import (
QUANT_BPW,
ModelArch,
Inputs,
estimate,
command_preview,
format_bytes,
quant_from_filename,
parse_hf_range,
)
from vramcalc.presets import PRESETS, PRESET_NAMES
QUANT_CHOICES = list(QUANT_BPW.keys())
CACHE_DTYPES = ["f16", "bf16", "f32", "q8_0", "q8_1", "q4_0", "q4_1", "q5_0", "q5_1"]
COMPUTE_DTYPES = ["f16", "bf16", "f32"]
def list_gguf_files(repo_id: str, hf_token: str):
"""List *.gguf files in a repo and auto-detect their quant names.
Returns gr.update objects (not bare lists) so Gradio 6 reliably re-renders
the dropdown and pre-selects the first file. A plain list return no longer
repopulates a Dropdown under SSR in Gradio 6.
"""
if not repo_id or not repo_id.strip():
return gr.update(choices=[], value=None), "Enter a Hugging Face repo id."
try:
api = HfApi(token=hf_token or None)
files = api.list_repo_files(repo_id=repo_id.strip())
except Exception as e: # noqa: BLE001
return gr.update(choices=[], value=None), f"Error listing repo: {e}"
ggufs = sorted(f for f in files if f.lower().endswith(".gguf"))
if not ggufs:
return gr.update(choices=[], value=None), f"No .gguf files found in {repo_id!r}."
choices = []
for f in ggufs:
q = quant_from_filename(f)
label = f"{f}" + (f" [{q}]" if q else "")
choices.append((label, f))
return (
gr.update(choices=choices, value=ggufs[0]),
f"Found {len(ggufs)} GGUF file(s).",
)
def fetch_arch(repo_id: str, filename: str, hf_token: str):
"""Range-read a GGUF header from HF and return editable arch fields.
Also returns a gr.update for the Quantization dropdown that syncs it to the
quant detected from the filename (e.g. ROCmFP4/ROCmFPX) — a no-op when the
filename has no recognized quant token, so the user's manual pick is kept.
The GGUF header's tensor types for ROCm formats use custom IDs stock
llama.cpp does not map, so the filename is the reliable source here.
"""
detected = quant_from_filename(filename) if filename else None
if detected and detected not in QUANT_CHOICES:
detected = None
quant_update = gr.update(value=detected) if detected else gr.update()
if not repo_id or not filename:
return _empty_arch_fields(), "Pick a GGUF file first.", quant_update
try:
meta = parse_hf_range(
repo_id.strip(), filename, token=hf_token or None
)
except Exception as e: # noqa: BLE001
return _empty_arch_fields(), f"Error reading GGUF header: {e}", quant_update
quant_note = f", quant {detected}" if detected else ""
if not meta.n_layer:
return (
_arch_to_fields(meta),
"Parsed header but architecture fields look empty; "
"edit them manually below.",
quant_update,
)
return (
_arch_to_fields(meta),
f"Fetched {meta.architecture or 'model'}: "
f"{meta.n_layer} layers, {meta.n_embd} embd, "
f"{meta.n_head}/{meta.n_head_kv} heads, ctx {meta.training_ctx}, "
f"params {meta.params or 'n/a'}{quant_note}.",
quant_update,
)
def _empty_arch_fields():
return _arch_to_fields(ModelArch())
def _arch_to_fields(m: ModelArch):
return [
m.name, m.architecture, m.n_layer, m.n_embd, m.n_head, m.n_head_kv,
m.training_ctx, m.params, m.rope_freq_base, m.n_expert,
m.n_expert_used, m.n_mtp,
]
ARCH_FIELD_NAMES = [
"name", "architecture", "n_layer", "n_embd", "n_head", "n_head_kv",
"training_ctx", "params", "rope_freq_base", "n_expert",
"n_expert_used", "n_mtp",
]
def _fields_to_arch(fields) -> ModelArch:
values = {k: v for k, v in zip(ARCH_FIELD_NAMES, fields)}
return ModelArch(
name=str(values["name"] or ""),
architecture=str(values["architecture"] or ""),
n_layer=int(values["n_layer"] or 0),
n_embd=int(values["n_embd"] or 0),
n_head=int(values["n_head"] or 0),
n_head_kv=int(values["n_head_kv"] or 0),
training_ctx=int(values["training_ctx"] or 0),
params=int(values["params"] or 0),
rope_freq_base=float(values["rope_freq_base"] or 10000.0),
n_expert=int(values["n_expert"] or 0),
n_expert_used=int(values["n_expert_used"] or 0),
n_mtp=int(values["n_mtp"] or 0),
)
def load_preset(name: str):
if name and name in PRESETS:
return _arch_to_fields(PRESETS[name]), f"Loaded preset: {name}"
return _empty_arch_fields(), ""
def _parse_gpu_list(text: str) -> list[float]:
out = []
for tok in (text or "").replace(";", ",").split(","):
tok = tok.strip()
if tok:
try:
out.append(float(tok))
except ValueError:
pass
return out or [24.0]
@spaces.GPU
def compute(
arch_fields,
quant, n_ctx, cache_dtype, flash_attn, compute_dtype,
n_batch, n_prompt,
rope_freq_scale, yarn_ext_factor, yarn_attn_factor,
yarn_beta_fast, yarn_beta_slow,
gpu_vram_text, kv_on_largest, safety_margin,
):
arch = _fields_to_arch(arch_fields)
gpus = _parse_gpu_list(gpu_vram_text)
inp = Inputs(
quant=quant,
n_ctx=int(n_ctx),
cache_dtype=cache_dtype,
flash_attn=bool(flash_attn),
compute_dtype=compute_dtype,
n_batch=int(n_batch),
n_prompt=int(n_prompt),
rope_freq_scale=float(rope_freq_scale),
yarn_ext_factor=float(yarn_ext_factor),
yarn_attn_factor=float(yarn_attn_factor),
yarn_beta_fast=float(yarn_beta_fast),
yarn_beta_slow=float(yarn_beta_slow),
gpu_vram_gb=gpus,
kv_on_largest=bool(kv_on_largest),
safety_margin_pct=float(safety_margin),
)
if arch.n_layer <= 0 or arch.n_embd <= 0 or arch.params <= 0:
return (
"⚠️ Model architecture is incomplete. Fill n_layer, n_embd, and "
"params (or fetch from a GGUF / load a preset).",
"", "",
)
bd = estimate(arch, inp)
# breakdown table
rows = [
["Weights (GGUF, " + quant + ")", format_bytes(bd.weights_bytes)],
["KV cache (" + cache_dtype + ")", format_bytes(bd.kv_cache_bytes)],
["Compute / scratch", format_bytes(bd.compute_scratch_bytes)],
["MTP overhead", format_bytes(bd.mtp_overhead_bytes)],
["GGUF header / overhead", format_bytes(bd.gguf_overhead_bytes)],
["Safety margin (" + str(inp.safety_margin_pct) + "%)",
format_bytes(bd.safety_margin_bytes)],
["**Total**", f"**{format_bytes(bd.total_bytes)}**"],
]
breakdown_md = "| Component | Size |\n|---|---|\n" + "\n".join(
f"| {a} | {b} |" for a, b in rows
)
# warnings
warn_md = ""
if bd.warnings:
warn_md = "\n\n**⚠️ YaRN / context notes:**\n" + "\n".join(
f"- {w}" for w in bd.warnings
)
eff_md = (
f"\n\nEffective context (training_ctx / rope_freq_scale): "
f"**{bd.effective_context}**"
)
# per-GPU table
gpu_md = ""
if bd.gpu and bd.gpu.assignments:
a = bd.gpu.assignments
header = "| GPU | VRAM | Weights | KV+Compute | Used | Free | Fits? |"
sep = "|---|---|---|---|---|---|---|"
body = []
for g in a:
kv = format_bytes(g.kv_compute_bytes) if g.kv_compute_bytes else "—"
badge = "✅" if g.fits else "❌"
body.append(
f"| {g.index} | {format_bytes(g.vram_bytes)} | "
f"{format_bytes(g.weight_bytes)} | {kv} | "
f"{format_bytes(g.used_bytes)} | {format_bytes(g.free_bytes)} "
f"| {badge} |"
)
total_badge = "✅ all fit" if bd.gpu.all_fit else "❌ over budget"
gpu_md = (
"**Per-GPU split (estimate, proportional weight split):**\n\n"
+ header + "\n" + sep + "\n" + "\n".join(body)
+ f"\n\nTotal VRAM: {format_bytes(bd.gpu.total_vram_bytes)} · "
f"Total used: {format_bytes(bd.gpu.total_used_bytes)} · "
f"{total_badge} · KV on GPU {bd.gpu.kv_gpu_index}"
)
cmd = command_preview(arch, inp)
summary = (
f"**{arch.name or arch.architecture or 'Model'}** @ {quant}, "
f"ctx {inp.n_ctx} ({cache_dtype} KV"
+ (", FA" if inp.flash_attn else ", no FA")
+ f"), {len(gpus)} GPU(s) → "
f"**{format_bytes(bd.total_bytes)}** total"
)
return summary + "\n\n" + breakdown_md + eff_md + warn_md + "\n\n" + gpu_md, cmd, cmd
def build_ui():
with gr.Blocks(title="llama.cpp VRAM Calculator") as demo:
gr.Markdown(
"# 🦀 llama.cpp VRAM Calculator\n"
"Estimate VRAM for a Hugging Face GGUF model: quant size, "
"context, KV cache options, YaRN context extension, MTP heads, "
"and multi-GPU split. Architecture is auto-fetched from the GGUF "
"header (range-read — no full model download)."
)
arch_state = gr.State(_empty_arch_fields())
with gr.Row():
with gr.Column(scale=1):
gr.Markdown("### 1. Model source")
with gr.Tab("Auto-fetch from HF"):
repo_id = gr.Textbox(
label="HF repo id",
placeholder="e.g. bartowski/Llama-3-8B-Instruct-GGUF",
)
hf_token = gr.Textbox(
label="HF token (optional, for gated/private repos)",
type="password",
)
list_btn = gr.Button("List GGUF files")
file_picker = gr.Dropdown(
label="GGUF file", choices=[], interactive=True
)
fetch_btn = gr.Button("Fetch architecture from GGUF header")
fetch_status = gr.Markdown("")
with gr.Tab("Presets / manual"):
preset_dd = gr.Dropdown(
label="Quick preset", choices=PRESET_NAMES, interactive=True
)
load_preset_btn = gr.Button("Load preset")
gr.Markdown("### Architecture (editable)")
arch_inputs = [
gr.Textbox(label="name", value=""),
gr.Textbox(label="architecture", value=""),
gr.Number(label="n_layer", value=0, precision=0),
gr.Number(label="n_embd", value=0, precision=0),
gr.Number(label="n_head", value=0, precision=0),
gr.Number(label="n_head_kv", value=0, precision=0),
gr.Number(label="training_ctx", value=0, precision=0),
gr.Number(label="params", value=0, precision=0),
gr.Number(label="rope_freq_base", value=10000.0),
gr.Number(label="n_expert (MoE)", value=0, precision=0),
gr.Number(label="n_expert_used", value=0, precision=0),
gr.Number(label="n_mtp", value=0, precision=0),
]
with gr.Column(scale=1):
gr.Markdown("### 2. Inference options")
quant = gr.Dropdown(
label="Quantization", choices=QUANT_CHOICES, value="Q4_K_M"
)
n_ctx = gr.Number(label="Target context (n_ctx)", value=8192, precision=0)
with gr.Row():
cache_dtype = gr.Dropdown(
label="KV cache dtype", choices=CACHE_DTYPES, value="f16"
)
compute_dtype = gr.Dropdown(
label="Compute dtype", choices=COMPUTE_DTYPES, value="f16"
)
flash_attn = gr.Checkbox(label="Flash attention", value=True)
with gr.Row():
n_batch = gr.Number(label="n_batch", value=512, precision=0)
n_prompt = gr.Number(label="n_prompt (active)", value=0, precision=0)
gr.Markdown("### YaRN / RoPE context extension")
with gr.Row():
rope_freq_scale = gr.Number(label="rope_freq_scale", value=1.0)
yarn_ext_factor = gr.Number(label="yarn_ext_factor", value=-1.0)
yarn_attn_factor = gr.Number(label="yarn_attn_factor", value=1.0)
with gr.Row():
yarn_beta_fast = gr.Number(label="yarn_beta_fast", value=32.0)
yarn_beta_slow = gr.Number(label="yarn_beta_slow", value=1.0)
gr.Markdown("### 3. Multi-GPU budget")
gpu_vram_text = gr.Textbox(
label="Per-GPU VRAM (GB, comma-separated)",
value="24",
placeholder="e.g. 24,24,16",
)
with gr.Row():
kv_on_largest = gr.Checkbox(
label="Place KV cache on largest GPU", value=False
)
safety_margin = gr.Number(label="Safety margin %", value=5.0)
compute_btn = gr.Button("Compute VRAM", variant="primary")
gr.Markdown("### Results")
result_md = gr.Markdown("")
with gr.Accordion("llama.cpp launch command preview", open=False):
cmd_md = gr.Markdown("")
cmd_text = gr.Textbox(
label="Command (copyable)", lines=8, interactive=False
)
# --- wiring ---
def _store_arch(*fields):
return list(fields)
arch_inputs_and_state = [*arch_inputs, arch_state]
# keep arch_state synced whenever arch fields change
for comp in arch_inputs:
comp.change(
fn=_store_arch, inputs=arch_inputs, outputs=arch_state
)
list_btn.click(
fn=list_gguf_files, inputs=[repo_id, hf_token],
outputs=[file_picker, fetch_status],
)
fetch_btn.click(
fn=fetch_arch, inputs=[repo_id, file_picker, hf_token],
outputs=[*arch_inputs, fetch_status, quant],
).then(
fn=_store_arch, inputs=arch_inputs, outputs=arch_state
)
load_preset_btn.click(
fn=load_preset, inputs=[preset_dd],
outputs=[*arch_inputs, fetch_status],
).then(
fn=_store_arch, inputs=arch_inputs, outputs=arch_state
)
compute_btn.click(
fn=compute,
inputs=[
arch_state, quant, n_ctx, cache_dtype, flash_attn, compute_dtype,
n_batch, n_prompt,
rope_freq_scale, yarn_ext_factor, yarn_attn_factor,
yarn_beta_fast, yarn_beta_slow,
gpu_vram_text, kv_on_largest, safety_margin,
],
outputs=[result_md, cmd_md, cmd_text],
)
return demo
demo = build_ui()
if __name__ == "__main__":
demo.launch()