Configuration Parsing Warning:In config.json: "num_experts" must be a number

Cactus Hybrid — Gemma 4 E2B

A small, on-device model is fast and private, but sometimes wrong. At Cactus we post-train models to know when they are wrong: we ship probes inside the checkpoint that score every answer with a confidence between 0 and 1, returned as structured data (never parsed out of the answer text). Answer on-device when confidence is high; re-route to a bigger model when it's low:

if confidence < 0.85:
    answer = ask_a_bigger_model(prompt)

This repo is google/gemma-4-E2B-it plus the handoff probe: a small head (weight prefix handoff_probe.*) that scores every generation with confidence = 1 - p_wrong. The base weights are byte-identical to the stock checkpoint (same keys); the repo adds eleven probe tensors, a remote-code model class and a custom_generate recipe. Stock engine commands work unchanged — you only add --trust-remote-code / trust_remote_code=True.

Benchmarks

Gemma 4 E2B Hybrid, the smallest Gemma model, matches Gemini 3.1 Flash-Lite on most benchmarks by routing only 15–35% of queries to Flash-Lite and running the rest itself:

Benchmark Handoff to match Flash-Lite (FP16) At 4-bit At 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% n/a

Quantisation quality is measured on Cactus Quants, which performs well at uniform quantization; developers are encouraged to benchmark Unsloth, GGUF, and MLX quantization independently.

Quickstart

# 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)

Load the model with an explicit .to(device), not device_map="auto": the probe scores generations outside the module forward() path, so weights that accelerate offloads (left on the meta device) crash the confidence read.

Serving with stock transformers

transformers serve --trust-remote-code
# then request model "Cactus-Compute/gemma-4-e2b-it-hybrid" via the OpenAI-compatible API

or interactively:

transformers chat Cactus-Compute/gemma-4-e2b-it-hybrid --trust-remote-code

How the confidence reaches you (in-band trailer)

transformers serve cannot add response fields, so the score travels in-band: the assistant content's final line is

\n[[hybrid:confidence=0.7812]]

always exactly 4 decimals, ASCII, confidence in [0, 1]. Strip the final [[hybrid:...]] line before display and parse the float for routing. The trailer is emitted in both streaming and non-streaming modes.

Limitations

  • --continuous-batching: not supported — the CB scheduler bypasses generate(), so no probe runs and no trailer is emitted. Serve without --continuous-batching to get confidence scores.
  • The trailer (and confidence) is only produced for single-sequence decoding: batch size 1, no beam search, no assisted/speculative decoding. Unsupported modes fall back to stock behavior (no trailer).
  • The probe scores at most the first 1024 generated tokens.
  • The trailer's token ids are appended to the returned sequences, so reported completion token counts include the trailer (a handful of tokens).

More Python APIs

# Structured API: clean sequences + raw float (no in-band trailer).
sequences, confidence = model.generate_with_confidence(inputs, max_new_tokens=512)
print(confidence)                     # e.g. 0.7812
print(model.last_confidence)          # same value

# Stock generate (custom_generate recipe): plain tensor + in-band trailer.
sequences = model.generate(**inputs, max_new_tokens=512)

# Suppress the trailer while keeping stock behavior:
sequences = model.generate(**inputs, max_new_tokens=512, emit_trailer=False)

Probe contract

  • Input: float32 [T, 1536] — output of decoder layer index 28 (config.probe_layer), captured at the position that predicts each generated token: row 0 = last prompt position at prefill, row t = position captured at generation step t. Only the first 1024 rows are scored.
  • Math (float32): x = LayerNorm(x, eps=1e-5) * norm.weight + norm.bias; p = relu(x @ proj.weight.T + proj.bias); s = p @ attn_query / sqrt(32); w = softmax_T(s - max); pooled = w @ p; h = relu(head.0 @ pooled + b); h = relu(head.2 @ h + b); logit = head.4 @ h + b; p_wrong = sigmoid(logit); confidence = 1 - p_wrong.
  • Capture uses a forward hook that keeps only one [1, 1536] row per decode step — full hidden-state stacks are never materialized.

Repo contents

File Purpose
configuration_gemma_4_e2b_it_hybrid.py Gemma4E2BItHybridConfig (stock Gemma-4 text config + probe hyperparams)
modeling_gemma_4_e2b_it_hybrid.py Gemma4E2BItHybridForCausalLM (stock Gemma4ForCausalLM + handoff_probe.*)
custom_generate/generate.py stock decode loop + confidence + in-band trailer
model*.safetensors base weights (identical keys) + handoff_probe.* tensors
gemma_4_e2b_it_hybrid.py single-file mlx-lm model, wired via config.json's model_file

Routing quality (AUROC)

AUROC measures how well the probe separates wrong answers from right ones (higher = better, 0.5 is random, 1.0 is perfect):

Hold-out Modality Cactus Hybrid Token Entropy
MMLU text MCQ 0.770 0.697
MMLU-Pro text MCQ 0.771 0.692
ARC-Easy text MCQ 0.888 0.655
ARC-Challenge text MCQ 0.834 0.646
GSM8K (3-shot) text gen 0.782 0.731
MMBench-EN-Dev vision MCQ 0.840 0.435
ChartQA vision QA 0.779 0.615
DocVQA vision QA 0.781 0.512
MMAU audio MCQ 0.789 0.517
GigaSpeech audio 0.876 0.343
Earnings-22 audio 0.839 0.323
LibriSpeech audio 0.822 0.427
Mean 0.814 0.549

The strongest result: the probe was trained on zero audio data, yet achieves 0.79–0.88 AUROC on four audio benchmarks (two transcription, one audio MCQ, one out-of-domain transcription). This rules out surface-level explanations: the probe is reading a modality-independent correctness signal from the hidden state, not memorizing patterns from training data.

All formats

All Cactus Hybrid builds live in the Cactus Hybrid collection: Transformers · GGUF / llama.cpp · MLX · Cactus engine. Copy-paste quickstarts for every engine: github.com/cactus-compute/cactus-hybrid.

License

Gemma is provided under and subject to the Gemma Terms of Use. This derivative includes the Cactus handoff probe head.

Downloads last month
146
Safetensors
Model size
5B params
Tensor type
F32
·
BF16
·
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for Cactus-Compute/gemma-4-e2b-it-hybrid

Finetuned
(293)
this model
Quantizations
2 models

Collection including Cactus-Compute/gemma-4-e2b-it-hybrid