"""Assemble a full VRAM breakdown and a llama.cpp launch-command preview.""" from __future__ import annotations from dataclasses import dataclass, field from .quant import weight_bytes from .kv import kv_cache_bytes, compute_scratch_bytes from .yarn import yarn_effective_context, yarn_warnings from .gpu import gpu_split, GpuSplitResult @dataclass class ModelArch: """Architecture parameters (fetched from GGUF or entered manually).""" name: str = "" architecture: str = "" n_layer: int = 0 n_embd: int = 0 n_head: int = 0 n_head_kv: int = 0 training_ctx: int = 0 params: int = 0 # total params including MoE experts rope_freq_base: float = 10000.0 n_expert: int = 0 # MoE n_expert_used: int = 0 n_mtp: int = 0 # MTP heads (e.g. DeepSeek-V3 = 1) @dataclass class Inputs: quant: str = "Q4_K_M" n_ctx: int = 8192 cache_dtype: str = "f16" flash_attn: bool = True compute_dtype: str = "f16" n_batch: int = 512 n_prompt: int = 0 # YaRN rope_freq_scale: float = 1.0 yarn_ext_factor: float = -1.0 yarn_attn_factor: float = 1.0 yarn_beta_fast: float = 32.0 yarn_beta_slow: float = 1.0 # GPU gpu_vram_gb: list[float] = field(default_factory=lambda: [24.0]) kv_on_largest: bool = False # margin safety_margin_pct: float = 5.0 @dataclass class Breakdown: weights_bytes: float = 0.0 kv_cache_bytes: float = 0.0 compute_scratch_bytes: float = 0.0 mtp_overhead_bytes: float = 0.0 gguf_overhead_bytes: float = 0.0 safety_margin_bytes: float = 0.0 total_bytes: float = 0.0 effective_context: int = 0 warnings: list[str] = field(default_factory=list) gpu: GpuSplitResult | None = None def estimate(arch: ModelArch, inp: Inputs) -> Breakdown: weights = weight_bytes(arch.params, inp.quant) # KV cache (with MTP layers folded in) kv = kv_cache_bytes( n_layer=arch.n_layer, n_embd=arch.n_embd, n_head=arch.n_head, n_head_kv=arch.n_head_kv, n_ctx=inp.n_ctx, cache_dtype=inp.cache_dtype, n_mtp=arch.n_mtp, flash_attn=inp.flash_attn, ) # compute / activation scratch scratch = compute_scratch_bytes( n_layer=arch.n_layer, n_embd=arch.n_embd, n_head=arch.n_head, n_head_kv=arch.n_head_kv, n_batch=inp.n_batch, compute_dtype=inp.compute_dtype, cache_dtype=inp.cache_dtype, flash_attn=inp.flash_attn, n_mtp=arch.n_mtp, ) # isolate MTP overhead for display: the extra layer's weights + its KV share mtp_overhead = 0.0 if arch.n_mtp > 0: # approximate extra weight as one layer's worth: params/n_layer * bpw/8 if arch.n_layer > 0: per_layer_params = arch.params / arch.n_layer mtp_weights = per_layer_params * ( __import__("vramcalc.quant", fromlist=["QUANT_BPW"]).QUANT_BPW[inp.quant] / 8.0 ) else: mtp_weights = 0.0 # KV for the extra layers head_dim = arch.n_embd // arch.n_head if arch.n_head > 0 else 0 mtp_kv = ( arch.n_mtp * inp.n_ctx * 2 * arch.n_head_kv * head_dim * __import__("vramcalc.kv", fromlist=["cache_dtype_bytes"]).cache_dtype_bytes(inp.cache_dtype) ) mtp_overhead = mtp_weights + mtp_kv # GGUF header / alignment overhead: small constant per file, rough estimate gguf_overhead = max(arch.n_layer * 4096, 1 << 20) # >= 1 MiB subtotal = weights + kv + scratch + gguf_overhead margin = subtotal * (inp.safety_margin_pct / 100.0) total = subtotal + margin eff_ctx = yarn_effective_context(arch.training_ctx, inp.rope_freq_scale) warns = yarn_warnings( training_ctx=arch.training_ctx, target_ctx=inp.n_ctx, rope_freq_scale=inp.rope_freq_scale, yarn_ext_factor=inp.yarn_ext_factor, yarn_attn_factor=inp.yarn_attn_factor, ) gpu_vram_bytes = [int(g * (1 << 30)) for g in inp.gpu_vram_gb] gpu = gpu_split( gpu_vram_bytes=gpu_vram_bytes, weights_bytes=weights, kv_compute_bytes=kv + scratch, kv_on_largest=inp.kv_on_largest, ) return Breakdown( weights_bytes=weights, kv_cache_bytes=kv, compute_scratch_bytes=scratch, mtp_overhead_bytes=mtp_overhead, gguf_overhead_bytes=gguf_overhead, safety_margin_bytes=margin, total_bytes=total, effective_context=eff_ctx, warnings=warns, gpu=gpu, ) def format_bytes(n: float) -> str: """Human-readable byte size.""" n = float(n) if n < 0: return "-" + format_bytes(-n) units = [("GiB", 1 << 30), ("MiB", 1 << 20), ("KiB", 1 << 10)] for label, size in units: if n >= size: return f"{n / size:.2f} {label}" return f"{n:.0f} B" # AMD ROCm floating-point block quants that stock llama.cpp cannot run. They # require the pinned ciru-ai/ROCmFPX runner (see # https://github.com/ciru-ai/ROCmFPX). Listed so command_preview can flag it. ROCMFP_QUANTS = {"ROCmFP4", "ROCmFPX"} def command_preview(arch: ModelArch, inp: Inputs) -> str: """Generate a llama.cpp launch command from the current inputs.""" is_rocmfp = inp.quant in ROCMFP_QUANTS runner = "rocmfpx-llama-server" if is_rocmfp else "llama-server" parts = [runner] parts.append(f"-m model-{inp.quant}.gguf") parts.append(f"-c {inp.n_ctx}") parts.append("-ngl 999") # full offload assumption parts.append(f"-b {inp.n_batch}") if inp.cache_dtype != "f16": parts.append(f"--cache-type-k {inp.cache_dtype}") parts.append(f"--cache-type-v {inp.cache_dtype}") if inp.flash_attn: parts.append("--flash-attn") if len(inp.gpu_vram_gb) > 1: # tensor-split proportional to vram split = ",".join(f"{g}" for g in inp.gpu_vram_gb) parts.append(f"--tensor-split {split}") # YaRN / rope yarn_args = [] if inp.rope_freq_scale != 1.0: yarn_args.append(f"--rope-freq-scale {inp.rope_freq_scale}") if arch.rope_freq_base != 10000.0: yarn_args.append(f"--rope-freq-base {arch.rope_freq_base}") if inp.yarn_ext_factor >= 0.0: yarn_args.append(f"--yarn-ext-factor {inp.yarn_ext_factor}") if inp.yarn_attn_factor != 1.0: yarn_args.append(f"--yarn-attn-factor {inp.yarn_attn_factor}") if inp.yarn_beta_fast != 32.0: yarn_args.append(f"--yarn-beta-fast {inp.yarn_beta_fast}") if inp.yarn_beta_slow != 1.0: yarn_args.append(f"--yarn-beta-slow {inp.yarn_beta_slow}") if inp.n_ctx > arch.training_ctx and arch.training_ctx > 0: yarn_args.append("--rope-scaling yarn") if yarn_args: parts.extend(yarn_args) if arch.n_mtp > 0: parts.append(f"--mtp {arch.n_mtp}") if is_rocmfp: parts.append( f"# NOTE: {inp.quant} needs the ciru-ai/ROCmFPX runner " f"(stock llama.cpp cannot read these tensor types)." ) return " \\\n ".join(parts)