File size: 6,293 Bytes
94981e6 3504e19 94981e6 3504e19 94981e6 3504e19 94981e6 3504e19 94981e6 3504e19 94981e6 3504e19 94981e6 3504e19 94981e6 3504e19 94981e6 3504e19 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 | ---
license: gemma
base_model: google/gemma-4-E2B-it
pipeline_tag: text-generation
tags:
- gemma4
- hybrid
- custom_code
- custom_generate
---
# 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 —
`0.85` is a good threshold:
```python
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`.
## Quickstart
```python
# 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
```bash
transformers serve --trust-remote-code
# then request model "Cactus-Compute/gemma-4-e2b-it-hybrid" via the OpenAI-compatible API
```
or interactively:
```bash
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
```python
# 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` |
## All formats
All Cactus Hybrid builds live in the
[Cactus Hybrid collection](https://huggingface.co/collections/Cactus-Compute/cactus-hybrid-6a60da4551074db058e8bb64):
[Transformers](https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid) ·
[GGUF / llama.cpp](https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid-GGUF) ·
[MLX](https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid-mlx) ·
[Cactus engine](https://huggingface.co/Cactus-Compute/gemma-4-E2B-it).
Copy-paste quickstarts for every engine:
[github.com/cactus-compute/cactus-hybrid](https://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.
|