一颗 TPU 芯片跑八个智能体:用原生 JAX 服务小型智能体负载
来源:dev.to — 2026-07-29
📋 概述
主流服务基准都以并发 100、1024-token 提示词衡量吞吐,但小型智能体负载截然不同:低并发高单流价值、延迟串行累积、上下文单调膨胀、输出必须结构化。作者针对这一被忽视的工作负载,用纯 JAX 写了一个 Gemma 4 E2B 推理引擎(无 PyTorch、无 torch_xla),在单颗 v6e 芯片上服务智能体循环。文章还揭示了一个 QAT 检查点无法在 vLLM TPU 加载的 bug(#3225),并展示静态 KV 缓存与逐位校验的正确性验证方法。
🔑 核心要点
- 被忽视的负载:小型智能体是低并发、高单流价值、延迟串行累积,与传统吞吐基准完全不同
- goodput 框架:结合吞吐与每 token 延迟目标,同时覆盖串行延迟敏感与并行吞吐受限两种拓扑
- 纯 JAX 引擎:safetensors 到 JAX PyTree 直接加载,测试断言 torch 从不进入 sys.modules
- 修复 QAT 加载 bug:tpu-inference #3225 中 QAT 检查点因 KV 共享层被错误要求 k_norm 而无法加载,修法是跳过共享层实例化
- 静态 KV 缓存:用 dynamic_update_slice 写静态缓存,通过缓存解码与全模型重跑逐位对比验证正确性
💡 金句
A malformed tool call isn’t a quality regression, it’s a crash.
👍 0
👎 0
← 返回 Dev.to 首页