为什么标准的推理基准总在100并发、1024令牌提示下报告吞吐量?因为那是公共推理端点的样子。但一个只有少量代理、每次调用简短且停驻在用户等待中的工作负载,形状完全不同。

这项判断并非共识。Zimbres 的系列测量提出相反证据:代理群“彼此交换消息,使每用户任务的令牌量增加一个数量级以上,且没有人在等任何一个 token,此时单 token 延迟失去大部分意义,总吞吐量才是硬约束”。两者都对,但对应不同的拓扑——带人在末端的串行代理回路受延迟束缚,并行机器到机器的代理群则受吞吐量限制。同时覆盖两者的指标是良好吞吐(goodput):在每 token 延迟目标约束下的吞吐量。文章的基准采用的正是这一框架,但并未声称从核级别取得生产级 goodput 结果。

打开网易新闻 查看精彩图片

这样的负载正好适合单颗芯片——前提是你得先把模型跑起来。麻烦出在这里。

量化感知训练(QAT)导出的 Gemma 4 E2B 模型在 TPU 上无法被 vLLM 加载。检查 safetensors 头部可以发现:BF16 导出为全部 35 层都提供了 self_attn.k_norm,而 QAT 导出只把它放在 15 个非 KV 共享层上。两种配置都声明 num_kv_shared_layers 为 20。第 15 到 34 层复用了低层计算的 K/V,它们没有任何 K 侧参数,因此给它们一个 k‑norm 毫无意义。QAT 导出的结构在架构上更诚实,而原来的 checkpoint 之所以能加载,仅仅因为它碰巧带着那些无用的张量。问题已提交为 tpu‑inference #3225,修复方案是跳过为共享层实例化 K/V 侧参数,而不是无条件要求它们。

既然没法直接用 QAT 模型,要么服务 BF16 并放弃量化优势,要么亲手写推理路径。我选择了后者。

ports/gemma4/jax_e_model.py 是一份纯 JAX 版的 Gemma 4 E2B 实现——路径上不经过 PyTorch,也不经过 torch_xla。权重从 safetensors 直接转为 JAX PyTree(通过 safetensors.flax 原生处理 bfloat16),测试中甚至断言 torch 在加载过程中从未进入 sys.modules。如果你习惯的是那种偷偷通过 AutoModelForCausalLM 加载权重再转换的“JAX”推理,这个版本完全不同。

KV 共享并没有在加载时打补丁,而是直接编码在层定义里。first_kv_shared_layer_idx 属性标记首个共享层的索引(35‑20=15),kv_share_map 方法为每一层返回其 KV 状态的源层。这样,共享关系成为模型结构本身的一部分,不再依赖外部脚本。