Cactus 发布 Gemma 4 E2B Hybrid:可在设备端为每个回答输出置信度分数,低分时自动路由至更大模型
Show HN: 《仙人掌杂交种》:我们教4岁的杰玛如何分辨对错
Cactus 推出基于 Gemma 4 的混合模型“Cactus Hybrid”,在模型检查点内嵌入置信度探针,为每个生成答案输出 0-1 之间的结构化置信度分数。高置信度时在设备端直接回答,低分时可自动路由至更大模型。该探针在零音频训练数据下,于四个音频基准上达到 0.79-0.88 AUROC,远超 token 熵基线(均值 0.549),且 MIT 协议开源。
把置信度探针塞进 Gemma 4 E2B,让模型自己判断该不该求助大模型,仅路由 15-35% 的查询就能持平 Flash-Lite,对离线/隐私优先的端侧产品是个省成本的新思路。
Cactus Hybrid
一个小型的端侧模型速度快且保护隐私,但有时会出错。在 Cactus,我们对模型进行后训练,让它们知道自己何时出错:我们在 checkpoint 中内置了探针,为每个答案打出一个介于 0 到 1 之间的置信度分数,以结构化数据的形式返回(绝不从答案文本中解析提取)。置信度高时在端侧作答;置信度低时,你可以将请求重新路由到更大的模型:
if confidence < 0.85:
answer = ask_a_bigger_model(prompt)
我们以 Gemma 4 E2B Hybrid 启动这次发布,所有构建版本都托管在 Hugging Face 上的Cactus Hybrid 集合中。
Gemma 4 E2B hybrid 是最小的 Gemma 模型,仅将 15–35% 的查询路由到 Gemini 3.1 Flash-Lite,其余由自身处理,在大多数基准测试上就能与 Gemini 3.1 Flash-Lite 持平。
| 基准测试 | 为匹配 Flash-Lite 而交接的比例(FP16) | 4-bit 量化下 | 3-bit 量化下 |
|---|---|---|---|
| ChartQA | 15–20% | 25–30% | 40–50% |
| MMBench | 30–35% | 40–45% | 50–55% |
| LibriSpeech | 25–30% | 35–40% | 55–65% |
| GigaSpeech | 30–35% | 40–45% | 50–55% |
| MMAU | 30–35% | 35–40% | 50–55% |
| MMLU-Pro | 45–55% | ~90% | 不适用 |
- 注:量化质量是在 Cactus Quants 上测得的,该方法在均匀量化下表现良好。
- 鼓励开发者分别针对 Unsloth、GGUF 和 MLX 量化进行基准测试。
Cactus
# pip install cactus-compute
import json
from cactus.bindings.cactus import cactus_complete, cactus_init
from cactus.cli.download import download_bundle
lm = cactus_init(str(download_bundle("Cactus-Compute/gemma-4-E2B-it")))
result = cactus_complete(
lm,
[{"role": "user", "content": "What is the capital of France?"}],
json.dumps({"max_tokens": 512, "auto_handoff": False}),
None,
lambda *_: None,
)
print(result["response"].strip())
print("confidence:", result["confidence"])
MLX
# pip install mlx-lm
import re
from mlx_lm import load, generate
model, tokenizer = load(
"Cactus-Compute/gemma-4-e2b-it-hybrid-mlx",
tokenizer_config={"trust_remote_code": True},
)
messages = [{"role": "user", "content": "What is the capital of France?"}]
answer = generate(
model,
tokenizer,
prompt=tokenizer.apply_chat_template(messages, add_generation_prompt=True),
max_tokens=512,
)
# the checkpoint reasons before answering; keep only the final answer
answer = re.split(r"<\|?channel\|?>", answer)[-1]
answer = re.sub(r"^(thought|final)\b\s*", "", answer).strip()
print(answer)
print("confidence:", model.last_confidence)
Transformers
# pip install "transformers>=5.5.4,<5.6" torch (5.14+ segfaults on this checkpoint)
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
model_id = "Cactus-Compute/gemma-4-e2b-it-hybrid"
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(model_id, trust_remote_code=True, dtype="auto").to(device)
messages = [{"role": "user", "content": "What is the capital of France?"}]
inputs = tokenizer.apply_chat_template(
messages, add_generation_prompt=True, return_tensors="pt", return_dict=True
).to(device)
out = model.generate(**inputs, return_confidence=True, max_new_tokens=512)
print(tokenizer.decode(out.sequences[0][inputs["input_ids"].shape[-1]:], skip_special_tokens=True))
print("confidence:", out.confidence)
加载模型时使用显式的 .to(device),而不是 device_map="auto":该探针在模块 forward() 路径之外对生成结果进行评分,因此那些加速卸载的权重(留在 meta 设备上)会导致置信度读取崩溃。
llama.cpp
llama.cpp 是 C++ 编写的,因此该探针是一个需要编译进引擎的补丁(参见 patches/llama.cpp/)。先构建打过补丁的服务器一次:
git clone https://github.com/cactus-compute/cactus-hybrid && cd cactus-hybrid
./patches/llama.cpp/install.sh && rehash
然后像使用任何 llama-server 一样对它进行服务和查询——响应会携带一个顶层的 confidence 字段:
llama-server -hf Cactus-Compute/gemma-4-e2b-it-hybrid-GGUF:Q4_K_M --jinja
curl -s http://localhost:8080/v1/chat/completions \
-d '{"messages":[{"role":"user","content":"What is the capital of France?"}],"max_tokens":512}' \
| jq '{answer: .choices[0].message.content, confidence}'
路由质量(AUROC)
Gemma 4 E2B Hybrid AUROC 衡量该机制将错误答案与正确答案区分开来的能力(越高越好,0.5 为随机,1.0 为完美):
| 留出集 | 模态 | Cactus Hybrid | Token 熵 |
|---|---|---|---|
| MMLU | 文本 MCQ | 0.770 | 0.697 |
| MMLU-Pro | 文本 MCQ | 0.771 | 0.692 |
| ARC-Easy | 文本 MCQ | 0.888 | 0.655 |
| ARC-Challenge | 文本 MCQ | 0.834 | 0.646 |
| GSM8K(3-shot) | 文本生成 | 0.782 | 0.731 |
| MMBench-EN-Dev | 视觉多选题 | 0.840 | 0.435 |
| ChartQA | 视觉问答 | 0.779 | 0.615 |
| DocVQA | 视觉问答 | 0.781 | 0.512 |
| MMAU | 音频多选题 | 0.789 | 0.517 |
| GigaSpeech | 音频 | 0.876 | 0.343 |
| Earnings-22 | 音频 | 0.839 | 0.323 |
| LibriSpeech | 音频 | 0.822 | 0.427 |
| 均值 | 0.814 | 0.549 |
最强的结果:该探针是在零音频数据上训练的,却在四个音频基准(两个转录、一个音频 MCQ、一个域外转录)上达到了 0.79–0.88 AUROC。
这排除了表层解释,该探针是从隐藏状态中读取一种与模态无关的正确性信号,而不是记忆训练数据中的模式。
采用 MIT 许可证。Gemma 模型的使用须遵守 Gemma 条款。
来源:Hacker News 热门(buzzing.cc 中文翻译) · github.com