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.