{"slug": "kv-cache-size-calculator", "title": "KV Cache Size Calculator", "summary": "A developer released a Python tool that calculates and plots KV cache size versus context length from HuggingFace config.json files. The tool supports standard MHA/GQA, MLA, hybrid architectures, and sliding window with per-layer compression, and can load configs from local files, URLs, or HuggingFace/ModelScope repo IDs.", "body_md": "| #!/usr/bin/env python3 | |\n| \"\"\"根据 HuggingFace config.json 计算并绘制 KV cache 大小随上下文长度的变化。 | |\n| 支持的注意力结构: | |\n| - 标准 MHA / GQA: KV = 2 * L * n_kv_heads * head_dim * dtype_bytes | |\n| - MLA (kv_lora_rank, 如 DeepSeek-V3 / Kimi / GLM): | |\n| 每层只缓存压缩的 c_kv (kv_lora_rank) 和 rope key (qk_rope_head_dim), | |\n| V 由 c_kv 重构不单独占用缓存: | |\n| KV = L * (kv_lora_rank + qk_rope_head_dim) * kv_bytes | |\n| - 混合架构 (linear_attention + 周期性 full_attention): | |\n| 只有 full attention 层的缓存随上下文线性增长; | |\n| linear attention (KDA / gated delta net 等) 是常数大小的递归状态 | |\n| (delta 规则状态 + 短卷积状态, 默认按 fp32 估算并叠加, --no-linear-state 可关闭)。 | |\n| - 滑窗 + 按层压缩 KV (DeepSeek V4 的 CSA/HCA, compress_ratios): | |\n| sliding_kv = active_layers * sliding_window * head_dim * kv_bytes (常数) | |\n| compressed_kv = Σ_{ratio>0} floor((tokens-1)/ratio) * head_dim * kv_bytes | |\n| indexer_kv = ratio4_layers * floor(tokens/4) * index_head_dim * fp4_bytes | |\n| 生产环境默认 KV 为 FP8 (1 B), indexer 为 FP4 (0.5 B)。 | |\n| 用法: | |\n| python3 kv_cache_calc.py configs/Kimi-K3.json configs/DeepSeek-V4-Flash.json | |\n| python3 kv_cache_calc.py moonshotai/Kimi-K3 # HuggingFace repo id | |\n| python3 kv_cache_calc.py https://.../config.json # 直接 URL | |\n| python3 kv_cache_calc.py cfg.json --check \"Kimi-K3:1048576:27.41851\" | |\n| \"\"\" | |\n| import argparse | |\n| import json | |\n| import math | |\n| import os | |\n| import sys | |\n| import urllib.request | |\n| from collections import Counter | |\n| DTYPE_BYTES = { | |\n| \"float32\": 4, \"float\": 4, \"fp32\": 4, | |\n| \"float16\": 2, \"fp16\": 2, \"half\": 2, | |\n| \"bfloat16\": 2, \"bf16\": 2, | |\n| \"float8_e4m3fn\": 1, \"float8_e5m2\": 1, \"fp8\": 1, \"int8\": 1, | |\n| \"float4\": 0.5, \"fp4\": 0.5, \"nvfp4\": 0.5, \"mxfp4\": 0.5, | |\n| } | |\n| # 层类型名字中包含这些关键字的层, 其 KV cache 随上下文线性增长 | |\n| FULL_ATTN_TYPES = (\"full_attention\", \"deepseek_sparse_attention\") | |\n| def dtype_bytes(name: str, default: float = 2) -> float: | |\n| return DTYPE_BYTES.get(str(name).lower().replace(\"torch.\", \"\"), default) | |\n| MS_ORG_ALIASES = { # ModelScope 上与 HuggingFace 组织名不同的映射 | |\n| \"zai-org\": \"ZhipuAI\", # 智谱 | |\n| } | |\n| def load_config(src: str, hub: str = \"auto\") -> dict: | |\n| \"\"\"从本地文件、URL 或模型仓库 id 加载 config.json. | |\n| hub: \"huggingface\" | \"modelscope\" | \"auto\"(先 HF 后 MS) | |\n| \"\"\" | |\n| if os.path.exists(src): | |\n| with open(src) as f: | |\n| return json.load(f) | |\n| if src.startswith(\"http\"): | |\n| with urllib.request.urlopen(src) as r: | |\n| return json.load(r) | |\n| urls = [] | |\n| if hub in (\"huggingface\", \"auto\"): | |\n| urls.append(f\"https://huggingface.co/{src}/resolve/main/config.json\") | |\n| if hub in (\"modelscope\", \"auto\"): | |\n| org, _, name = src.partition(\"/\") | |\n| alias = MS_ORG_ALIASES.get(org, org) | |\n| if alias != org: # 别名优先, 原名留作回退 | |\n| urls.append(f\"https://modelscope.cn/models/{alias}/{name}/resolve/master/config.json\") | |\n| urls.append(f\"https://modelscope.cn/models/{src}/resolve/master/config.json\") | |\n| last_err = None | |\n| for url in urls: | |\n| try: | |\n| print(f\"正在下载 {url} ...\", file=sys.stderr) | |\n| with urllib.request.urlopen(url) as r: | |\n| return json.load(r) | |\n| except Exception as e: # 换下一个 hub/别名重试 | |\n| last_err = e | |\n| raise RuntimeError(f\"无法获取 {src} 的 config.json: {last_err}\") | |\n| def model_name(src: str, cfg: dict) -> str: | |\n| if os.path.exists(src): | |\n| base = os.path.splitext(os.path.basename(src))[0] | |\n| else: | |\n| base = src.rstrip(\"/\").split(\"/\")[-1] | |\n| if base == \"config.json\": | |\n| base = src.split(\"/\")[-3] | |\n| arch = (cfg.get(\"architectures\") or [\"unknown\"])[0] | |\n| return f\"{base} ({arch})\" | |\n| def full_attn_layers(text: dict, cfg: dict) -> list: | |\n| \"\"\"确定哪些层是 full attention (KV 随上下文增长), 返回层索引列表\"\"\" | |\n| # 1) 显式的 layer_types 列表 (如 Qwen flash-next) | |\n| if \"layer_types\" in text: | |\n| return [i for i, t in enumerate(text[\"layer_types\"]) | |\n| if any(k in t for k in FULL_ATTN_TYPES)] | |\n| # 2) linear_attn_config.full_attn_layers (如 Kimi K3, GLM) | |\n| lac = text.get(\"linear_attn_config\") or cfg.get(\"linear_attn_config\") or {} | |\n| if \"full_attn_layers\" in lac: | |\n| return list(lac[\"full_attn_layers\"]) | |\n| # 3) 普通稠密模型: 所有层 | |\n| return list(range(text.get(\"num_hidden_layers\", 0))) | |\n| def linear_state_bytes(text: dict, n_linear: int, state_bytes: float) -> float: | |\n| \"\"\"估算每个 linear attention 层的常数递归状态 (可选叠加项). | |\n| 包含两部分, 均为常数大小 (不随上下文增长): | |\n| - delta 规则状态 (num_heads, d_k, d_v), 生产实现通常为 fp32 | |\n| - 短卷积状态 (kernel_size x 各投影维度) | |\n| \"\"\" | |\n| if n_linear <= 0: | |\n| return 0 | |\n| if \"linear_num_value_heads\" in text: # Qwen 风格 gated deltanet: K/V head 分列 | |\n| dstate = (text[\"linear_num_value_heads\"] * text[\"linear_key_head_dim\"] | |\n| * text[\"linear_value_head_dim\"]) | |\n| kernel = int(text.get(\"linear_conv_kernel_dim\", 4)) | |\n| conv = kernel * (text[\"linear_num_key_heads\"] * text[\"linear_key_head_dim\"] | |\n| + text[\"linear_num_value_heads\"] * text[\"linear_value_head_dim\"]) | |\n| return (dstate + conv) * state_bytes * n_linear | |\n| lac = text.get(\"linear_attn_config\") or {} | |\n| if lac: # Kimi / GLM 风格 KDA | |\n| h, d = lac.get(\"num_heads\", 0), lac.get(\"head_dim\", 0) | |\n| kernel = int(lac.get(\"short_conv_kernel_size\", 4)) | |\n| return (h * d * d + kernel * h * d) * state_bytes * n_linear | |\n| return 0 # 信息不足, 不估算 | |\n| def make_info(src, cfg, *, dtype, db, n_layers, n_full, n_linear, n_mtp, | |\n| per_token, const, max_pos, kind, layer_desc=None, | |\n| parts=None, extra_desc=None): | |\n| \"\"\"构造统一的模型信息; parts 为 [(名称, fn(tokens)->bytes)], 用于分项展示\"\"\" | |\n| def fn(tokens, _pt=per_token, _c=const): | |\n| return _pt * tokens + _c | |\n| if parts is None: | |\n| parts = [(\"KV cache\", fn)] | |\n| total_fn = (lambda t, _p=parts: sum(f(t) for _, f in _p)) if len(parts) > 1 else parts[0][1] | |\n| return { | |\n| \"name\": model_name(src, cfg), \"kind\": kind, \"dtype\": dtype, \"db\": db, | |\n| \"n_layers\": n_layers, \"n_full\": n_full, \"n_linear\": n_linear, | |\n| \"n_mtp\": n_mtp, \"per_token_bytes\": per_token, \"const_bytes\": const, | |\n| \"per_layer_bytes\": per_token // max(n_full, 1), | |\n| \"max_pos\": max_pos, \"layer_desc\": layer_desc, \"parts\": parts, | |\n| \"fn\": total_fn, \"extra_desc\": extra_desc, | |\n| } | |\n| def analyze_compress(text, src, cfg, dtype, kv_db, indexer_db): | |\n| \"\"\"DeepSeek V4 风格: 每层 = 滑窗 + 按 compress_ratio 压缩的 latent KV; | |\n| 压缩比最小的层带 indexer (稀疏打分器), 另有一份压缩缓存.\"\"\" | |\n| ratios = [int(r) for r in text[\"compress_ratios\"]] | |\n| n_layers = text.get(\"num_hidden_layers\", len(ratios)) | |\n| main = ratios[:n_layers] | |\n| hd = text[\"head_dim\"] | |\n| win = int(text.get(\"sliding_window\") or 0) | |\n| idx_hd = int(text.get(\"index_head_dim\") or 0) | |\n| nonzero = [r for r in main if r > 0] | |\n| # 官方实现中 indexer 挂在 compress_ratio==4 的层上; 泛化为压缩比最小的层 | |\n| idx_ratio = min(nonzero) if (idx_hd and nonzero) else None | |\n| n_idx = sum(1 for r in main if idx_ratio and r == idx_ratio) | |\n| active = len(main) # 滑窗缓存存在于每个主层 (含 ratio=0 的纯滑窗层) | |\n| const_kv = active * win * hd * kv_db | |\n| def kv_fn(t): | |\n| # 只统计已完成的压缩组 (进行中的组在滑窗/压缩状态里) | |\n| return sum(max(t - 1, 0) // r for r in nonzero) * hd * kv_db + const_kv | |\n| def idx_fn(t): | |\n| return n_idx * (t // idx_ratio) * idx_hd * indexer_db if idx_ratio else 0 | |\n| parts = [(\"KV cache (滑窗+压缩)\", kv_fn)] | |\n| if idx_hd: | |\n| parts.append((f\"Indexer cache (FP4, {n_idx}层x1/{idx_ratio})\", idx_fn)) | |\n| dist = \", \".join(f\"{c}层x1/{r}\" for r, c in sorted(Counter(nonzero).items())) | |\n| kind = (f\"压缩KV稀疏注意力 (latent head_dim={hd}, window={win}\" | |\n| + (f\", indexer={idx_hd}维x1/{idx_ratio}\" if idx_ratio else \"\") | |\n| + f\", KV精度={kv_db}B, indexer精度={indexer_db}B)\") | |\n| layer_desc = (f\"共 {n_layers} 层 = {dist}, 滑窗 {win}\" | |\n| + (\"; MTP 层仅滑窗, 不计入\" if len(ratios) > n_layers else \"\")) | |\n| return make_info(src, cfg, dtype=dtype, db=kv_db, n_layers=n_layers, | |\n| n_full=len(nonzero), n_linear=active - len(nonzero), | |\n| n_mtp=len(ratios) - n_layers if len(ratios) > n_layers else 0, | |\n| per_token=0, const=const_kv, # per_token 由 fn 精确给出 | |\n| max_pos=text.get(\"max_position_embeddings\"), | |\n| kind=kind, layer_desc=layer_desc, parts=parts) | |\n| def default_kv_dtype(text: dict, cfg: dict, model_dtype: str) -> str: | |\n| \"\"\"按官方配置推断 KV cache 精度: | |\n| - fp8 量化且未排除 attention 输入投影 (k/v) --> fp8 (如 MiMo/DSv4) | |\n| - attention 被显式排除量化 (如 GLM 的 attn_mha、Kimi 的 self_attn) --> 模型精度 bf16 | |\n| - 无量化配置 --> 模型精度 | |\n| \"\"\" | |\n| q = text.get(\"quantization_config\") or cfg.get(\"quantization_config\") or {} | |\n| if q.get(\"quant_method\") == \"fp8\": | |\n| ign = [str(x) for x in (q.get(\"ignored_layers\") or q.get(\"modules_to_not_convert\") | |\n| or q.get(\"ignore\") or [])] | |\n| attn_excluded = any(\"attn\" in s and \"o_proj\" not in s for s in ign) | |\n| if not attn_excluded: | |\n| return \"fp8\" | |\n| return str(model_dtype).lower().replace(\"torch.\", \"\") | |\n| def analyze(src: str, cfg: dict, dtype_override: str | None, | |\n| kv_dtype_override: str | None, indexer_dtype: str, | |\n| include_linear_state: bool, state_dtype: str) -> dict: | |\n| text = cfg.get(\"text_config\", cfg) | |\n| dtype = dtype_override or text.get(\"dtype\") or text.get(\"torch_dtype\") or \"bfloat16\" | |\n| db = dtype_bytes(dtype) | |\n| if text.get(\"compress_ratios\"): # DeepSeek V4 风格 CSA/HCA | |\n| kv_d = kv_dtype_override or \"fp8\" # 生产默认 FP8 attention cache | |\n| return analyze_compress(text, src, cfg, kv_d, dtype_bytes(kv_d, 1), | |\n| dtype_bytes(indexer_dtype, 0.5)) | |\n| kv_d = kv_dtype_override or default_kv_dtype(text, cfg, dtype) | |\n| db = dtype_bytes(kv_d) | |\n| # MiMo V2 风格: hybrid_layer_pattern 逐层标注, 0=full/global attention, | |\n| # 1=滑窗(SWA)层; 滑窗层只保留常数大小的窗口缓存 (如 128 个 token). | |\n| # 官方口径: MiMo-V2.5-Pro 为 10 GA + 60 SWA (6:1), 长上下文缓存省 ~7 倍. | |\n| if text.get(\"hybrid_layer_pattern\"): | |\n| pattern = [int(x) for x in text[\"hybrid_layer_pattern\"]] | |\n| n_full = pattern.count(0) | |\n| n_swa = len(pattern) - n_full | |\n| kv_heads = text.get(\"num_key_value_heads\", text.get(\"num_attention_heads\", 0)) | |\n| head_dim = text.get(\"head_dim\") or text.get(\"hidden_size\", 0) // max(text.get(\"num_attention_heads\", 1), 1) | |\n| v_head_dim = text.get(\"v_head_dim\", head_dim) # 部分模型 QK/V head_dim 不同 | |\n| win = int(text.get(\"sliding_window\") or text.get(\"sliding_window_size\") or 0) | |\n| per_entry = (head_dim + v_head_dim) * kv_heads * db | |\n| const = n_swa * win * per_entry | |\n| per_token = n_full * per_entry | |\n| kind = (f\"GQA+滑窗混合 (n_kv_heads={kv_heads}, qk_head_dim={head_dim}, \" | |\n| f\"v_head_dim={v_head_dim}, full层={n_full}, 滑窗层={n_swa}x窗{win})\") | |\n| layer_desc = (f\"共 {len(pattern)} 层 = {n_full} 层 full attention + \" | |\n| f\"{n_swa} 层滑窗 (常数 {human(const)})\") | |\n| return make_info(src, cfg, dtype=kv_d, db=db, n_layers=len(pattern), | |\n| n_full=n_full, n_linear=n_swa, n_mtp=0, | |\n| per_token=per_token, const=const, | |\n| max_pos=text.get(\"max_position_embeddings\"), kind=kind, | |\n| layer_desc=layer_desc) | |\n| full_idx = full_attn_layers(text, cfg) | |\n| n_layers = text.get(\"num_hidden_layers\", len(full_idx)) | |\n| n_full = len(full_idx) | |\n| # MTP / nextn 层: 若其注意力是 full attention, 同样占用随上下文增长的缓存 | |\n| mtp = text.get(\"num_nextn_predict_layers\") or text.get(\"mtp_num_hidden_layers\") or 0 | |\n| mtp_cfg = text.get(\"mtp\") or {} | |\n| if mtp and mtp_cfg.get(\"layer_types\") and not any( | |\n| k in t for k in FULL_ATTN_TYPES for t in mtp_cfg[\"layer_types\"]): | |\n| mtp = 0 | |\n| if text.get(\"kv_lora_rank\"): # MLA: 缓存 c_kv + k_pe, V 不单独存 | |\n| mla_rank = text[\"kv_lora_rank\"] | |\n| rope_dim = text.get(\"qk_rope_head_dim\", 0) | |\n| per_layer = (mla_rank + rope_dim) * db | |\n| kind = f\"MLA (kv_lora_rank={mla_rank}, qk_rope_head_dim={rope_dim})\" | |\n| else: # 标准 MHA / GQA | |\n| kv_heads = text.get(\"num_key_value_heads\", text.get(\"num_attention_heads\", 0)) | |\n| head_dim = text.get(\"head_dim\") | |\n| if not head_dim: | |\n| head_dim = (text.get(\"qk_nope_head_dim\", 0) or 0) + \\ | |\n| (text.get(\"qk_rope_head_dim\", 0) or 0) | |\n| if not head_dim: | |\n| head_dim = text.get(\"hidden_size\", 0) // max(text.get(\"num_attention_heads\", 1), 1) | |\n| per_layer = 2 * kv_heads * head_dim * db # K 和 V 各一份 | |\n| kind = f\"GQA (n_kv_heads={kv_heads}, head_dim={head_dim})\" | |\n| per_token = per_layer * (n_full + mtp) | |\n| n_linear = n_layers - n_full | |\n| const = linear_state_bytes(text, n_linear, dtype_bytes(state_dtype)) \\ | |\n| if include_linear_state else 0 | |\n| # DSA 稀疏层 (deepseek_sparse_attention) 的 indexer 也有一份随上下文增长的 | |\n| # 压缩缓存 (每 index_kpool 个 token 存 index_head_dim 维), 如 GLM / DSv3.2 系 | |\n| parts = None | |\n| idx_hd, kpool = text.get(\"index_head_dim\"), text.get(\"index_kpool\") | |\n| n_sparse = sum(1 for t in text.get(\"layer_types\", []) if \"sparse\" in t) if idx_hd else 0 | |\n| if n_sparse and idx_hd and kpool: | |\n| def idx_fn(t, _n=n_sparse, _r=int(kpool), _d=int(idx_hd), _b=db): | |\n| return _n * (t // _r) * _d * _b | |\n| parts = [(\"KV cache\", lambda t, _pt=per_token, _c=const: _pt * t + _c), | |\n| (f\"Indexer cache ({n_sparse}层x1/{kpool})\", idx_fn)] | |\n| kind += f\" + DSA稀疏 (indexer {idx_hd}维x1/{kpool})\" | |\n| return make_info(src, cfg, dtype=kv_d, db=db, n_layers=n_layers, n_full=n_full, | |\n| n_linear=n_linear, n_mtp=mtp, per_token=per_token, const=const, | |\n| max_pos=text.get(\"max_position_embeddings\"), kind=kind, | |\n| parts=parts) | |\n| def human(n: float) -> str: | |\n| for unit in (\"B\", \"KiB\", \"MiB\", \"GiB\", \"TiB\"): | |\n| if abs(n) < 1024 or unit == \"TiB\": | |\n| return f\"{n:.0f} B\" if unit == \"B\" else f\"{n:.4f} {unit}\" | |\n| n /= 1024 | |\n| def kv_bytes(info: dict, ctx: int, batch: int = 1) -> float: | |\n| return info[\"fn\"](ctx) * batch | |\n| def print_report(info: dict, batch: int, checks: list) -> None: | |\n| print(\"=\" * 78) | |\n| print(f\"模型: {info['name']}\") | |\n| print(f\"注意力类型: {info['kind']}, dtype={info['dtype']}\") | |\n| if info.get(\"layer_desc\"): | |\n| print(f\"层数: {info['layer_desc']}\") | |\n| else: | |\n| print(f\"层数: 共 {info['n_layers']} 层 = {info['n_full']} 层 full attention\" | |\n| + (f\" + {info['n_mtp']} 层 MTP\" if info[\"n_mtp\"] else \"\") | |\n| + (f\" + {info['n_linear']} 层 linear attention (常数状态, 未计入)\" | |\n| if info[\"n_linear\"] else \"\")) | |\n| max_ctx = info[\"max_pos\"] or 131072 | |\n| per_tok = info[\"fn\"](max_ctx) / max_ctx | |\n| if info[\"per_token_bytes\"]: | |\n| line = f\"每层缓存: {human(info['per_layer_bytes'])} / token\" | |\n| if info[\"const_bytes\"]: | |\n| line += f\"\\n每 token 缓存: {human(info['per_token_bytes'])} + 常数 {human(info['const_bytes'])}\" | |\n| else: | |\n| line += f\"\\n每 token 缓存: {human(info['per_token_bytes'])}\" | |\n| print(line) | |\n| else: | |\n| print(f\"每 token 缓存: {human(per_tok)} (按最大上下文折算)\") | |\n| if len(info[\"parts\"]) > 1: | |\n| for label, f in info[\"parts\"]: | |\n| print(f\" - {label:<28} @ {max_ctx:,}: {human(f(max_ctx))}\") | |\n| print(f\" {'合计':<30} @ {max_ctx:,}: {human(info['fn'](max_ctx))}\") | |\n| print(f\"最大上下文: {max_ctx:,} tokens\") | |\n| print(\"-\" * 78) | |\n| print(f\"{'上下文长度':>14} | {'KV cache / 序列':>18}\" + | |\n| (f\" | {'x batch='+str(batch):>18}\" if batch > 1 else \"\")) | |\n| points = {1024, 4096, 8192, 16384, 32768, 65536, 131072, | |\n| 262144, 524288, 1048576, max_ctx} | |\n| points = sorted(p for p in points if p <= max_ctx) | |\n| for ctx in points: | |\n| b = kv_bytes(info, ctx, batch) | |\n| line = f\"{ctx:>14,} | {human(b):>18}\" | |\n| if batch > 1: | |\n| line += f\" | {human(b * batch):>18}\" | |\n| print(line) | |\n| print(\"-\" * 78) | |\n| for c in checks: | |\n| name, ctx, expected = c | |\n| if name.lower() not in info[\"name\"].lower(): | |\n| continue | |\n| got = kv_bytes(info, ctx, batch) / 2**30 | |\n| diff = got - expected | |\n| print(f\"核对 [{name} @ {ctx:,}]: 计算 {got:.5f} GiB vs 参考 {expected:.5f} GiB\" | |\n| f\" (差 {diff:+.5f} GiB, {diff/expected*100:+.3f}%)\") | |\n| print() | |\n| def short_name(info: dict) -> str: | |\n| return info[\"name\"].split(\" (\")[0] | |\n| def fmt_per_tok(b: float) -> str: | |\n| return f\"{b/1024:.2f}\".rstrip(\"0\").rstrip(\".\") + \" KiB/tok\" if b >= 1024 \\ | |\n| else f\"{b:.0f} B/tok\" | |\n| def plot(infos: list, batch: int, out: str, max_ctx_cap: int | None): | |\n| import matplotlib | |\n| matplotlib.use(\"Agg\") | |\n| import matplotlib.pyplot as plt | |\n| fig, ax = plt.subplots(figsize=(10, 7)) | |\n| markers = [\"o\", \"s\", \"^\", \"D\", \"v\", \"p\", \"h\", \"X\"] | |\n| endpoints = [] # (axvline x, 最终 y, 标签) 用于错开标注 | |\n| for i, info in enumerate(infos): | |\n| max_ctx = min(info[\"max_pos\"] or 131072, max_ctx_cap or math.inf) | |\n| lo, hi = 10, int(math.log2(max_ctx)) | |\n| xs = [2**k for k in range(lo, hi + 1)] | |\n| ys = [kv_bytes(info, x, batch) / 2**30 for x in xs] | |\n| per_tok = info[\"fn\"](max_ctx) / max_ctx | |\n| ax.plot(xs, ys, marker=markers[i % len(markers)], ms=4, lw=1.5, | |\n| label=f\"{short_name(info)} {fmt_per_tok(per_tok)} · KV {info['dtype']}\") | |\n| endpoints.append((xs[-1], ys[-1], f\"{ys[-1]:.2f}\")) | |\n| if info[\"max_pos\"] and info[\"max_pos\"] <= (max_ctx_cap or math.inf): | |\n| ax.axvline(info[\"max_pos\"], ls=\":\", lw=0.8, alpha=0.5) | |\n| # 端点标注按 y 排序后做简单的错位, 避免重叠 | |\n| endpoints.sort(key=lambda e: e[1]) | |\n| placed = [] | |\n| for x, y, text in endpoints: | |\n| for py in placed: # 对数轴上保持约 25% 的间距, placed 已按升序 | |\n| if abs(math.log10(y) - math.log10(py)) < 0.22: | |\n| y = py * 1.25 | |\n| placed.append(y) | |\n| ax.annotate(text, (x, y), textcoords=\"offset points\", xytext=(-6, 0), | |\n| ha=\"right\", va=\"center\", fontsize=8, alpha=0.8, | |\n| bbox=dict(fc=\"white\", ec=\"none\", alpha=0.5, pad=0.6)) | |\n| ax.set_xscale(\"log\", base=2) | |\n| ax.set_yscale(\"log\") | |\n| ax.set_xlabel(\"Context length (tokens)\") | |\n| ax.set_ylabel(\"KV cache size (GiB)\") | |\n| ax.grid(True, which=\"both\", ls=\"--\", alpha=0.3) | |\n| ax.legend(fontsize=8.5, ncol=3, loc=\"lower center\", | |\n| bbox_to_anchor=(0.5, 1.01), frameon=False) | |\n| fig.suptitle(\"KV cache vs context length\" + (f\" (batch={batch})\" if batch > 1 else \"\"), | |\n| y=0.975) | |\n| fig.tight_layout(rect=[0, 0, 1, 0.92]) | |\n| fig.savefig(out, dpi=300) | |\n| svg = os.path.splitext(out)[0] + \".svg\" | |\n| fig.savefig(svg) | |\n| print(f\"图已保存到 {out} (300 dpi) 和 {svg}\") | |\n| def main(): | |\n| ap = argparse.ArgumentParser(description=\"根据 config.json 计算/绘制 KV cache 随上下文长度的变化\") | |\n| ap.add_argument(\"configs\", nargs=\"+\", help=\"config.json 路径 / 模型仓库 id / URL\") | |\n| ap.add_argument(\"--hub\", choices=[\"huggingface\", \"modelscope\", \"auto\"], | |\n| default=\"auto\", help=\"模型仓库来源 (默认 auto: HF 失败回退 ModelScope)\") | |\n| ap.add_argument(\"--dtype\", help=\"覆盖模型 dtype (如 bf16/fp16)\", default=None) | |\n| ap.add_argument(\"--kv-dtype\", help=\"覆盖 KV cache 精度 (如 bf16/fp8/fp4/int8)\", default=None) | |\n| ap.add_argument(\"--indexer-dtype\", help=\"indexer 缓存精度 (默认 fp4)\", default=\"fp4\") | |\n| ap.add_argument(\"--batch\", type=int, default=1, help=\"batch 大小 (默认 1)\") | |\n| ap.add_argument(\"--max-context\", type=int, default=None, help=\"绘制曲线的最大上下文长度\") | |\n| ap.add_argument(\"--no-linear-state\", action=\"store_false\", dest=\"linear_state\", | |\n| help=\"不计入 linear attention 层的常数递归状态估算\") | |\n| ap.set_defaults(linear_state=True) | |\n| ap.add_argument(\"--state-dtype\", default=\"float32\", help=\"递归状态 dtype (默认 float32)\") | |\n| ap.add_argument(\"--check\", action=\"append\", default=[], | |\n| help='核对数据, 格式 \"名字:上下文长度:GiB数值\", 可多次指定') | |\n| ap.add_argument(\"-o\", \"--out\", default=\"kv_cache.png\", help=\"输出图片路径\") | |\n| args = ap.parse_args() | |\n| checks = [] | |\n| for c in args.check: | |\n| name, ctx, val = c.rsplit(\":\", 2) | |\n| checks.append((name, int(ctx), float(val))) | |\n| infos = [] | |\n| for src in args.configs: | |\n| info = analyze(src, load_config(src, args.hub), args.dtype, args.kv_dtype, | |\n| args.indexer_dtype, args.linear_state, args.state_dtype) | |\n| infos.append(info) | |\n| print_report(info, args.batch, checks) | |\n| if args.out: | |\n| plot(infos, args.batch, args.out, args.max_context) | |\n| if __name__ == \"__main__\": | |\n| main() |", "url": "https://wpnews.pro/news/kv-cache-size-calculator", "canonical_source": "https://gist.github.com/wszqkzqk/2702ba95fd20e9176b366ce1575a556c", "published_at": "2026-08-30 07:11:28+00:00", "updated_at": "2026-08-30 08:22:15.357875+00:00", "lang": "en", "topics": ["developer-tools", "large-language-models", "ai-infrastructure"], "entities": ["HuggingFace", "ModelScope", "DeepSeek", "Kimi", "GLM"], "alternates": {"html": "https://wpnews.pro/news/kv-cache-size-calculator", "markdown": "https://wpnews.pro/news/kv-cache-size-calculator.md", "text": "https://wpnews.pro/news/kv-cache-size-calculator.txt", "jsonld": "https://wpnews.pro/news/kv-cache-size-calculator.jsonld"}}