Henry Ndubuaku commited on
Commit
3504e19
·
verified ·
1 Parent(s): 723cf0b

Model card: Cactus Hybrid quickstart + collection/GitHub links

Browse files
Files changed (1) hide show
  1. README.md +60 -37
README.md CHANGED
@@ -9,20 +9,53 @@ tags:
9
  - custom_generate
10
  ---
11
 
12
- # gemma-4-e2b-it-hybrid
13
 
14
- `Cactus-Compute/gemma-4-e2b-it-hybrid` is **google/gemma-4-E2B-it plus a handoff probe**: a
15
- small head (weight prefix `handoff_probe.*`) that scores every generation with
 
 
 
 
16
 
 
 
 
17
  ```
18
- confidence = 1 - p_wrong
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
19
  ```
20
 
21
- so a client can decide when to keep the on-device answer and when to hand off to
22
- a cloud model. The base weights are byte-identical to the stock checkpoint (same
23
- keys); the repo adds eleven probe tensors, a remote-code model class and a
24
- `custom_generate` recipe. **Stock engine commands work unchanged** — you only add
25
- `--trust-remote-code` / `trust_remote_code=True`.
26
 
27
  ## Serving with stock transformers
28
 
@@ -62,34 +95,19 @@ trailer is emitted in both streaming and non-streaming modes.
62
  - The trailer's token ids are appended to the returned sequences, so reported
63
  completion token counts include the trailer (a handful of tokens).
64
 
65
- ## Python usage
66
 
67
  ```python
68
- from transformers import AutoModelForCausalLM, AutoTokenizer
69
-
70
- model = AutoModelForCausalLM.from_pretrained(
71
- "Cactus-Compute/gemma-4-e2b-it-hybrid", trust_remote_code=True, dtype="auto"
72
- )
73
- tokenizer = AutoTokenizer.from_pretrained("Cactus-Compute/gemma-4-e2b-it-hybrid")
74
- inputs = tokenizer.apply_chat_template(
75
- [{"role": "user", "content": "Explain why this contract clause is risky."}],
76
- add_generation_prompt=True, return_tensors="pt",
77
- ).to(model.device)
78
-
79
  # Structured API: clean sequences + raw float (no in-band trailer).
80
  sequences, confidence = model.generate_with_confidence(inputs, max_new_tokens=512)
81
  print(confidence) # e.g. 0.7812
82
  print(model.last_confidence) # same value
83
 
84
  # Stock generate (custom_generate recipe): plain tensor + in-band trailer.
85
- sequences = model.generate(inputs, max_new_tokens=512)
86
-
87
- # Or opt into a rich return with the raw float:
88
- out = model.generate(inputs, max_new_tokens=512, return_confidence=True)
89
- out.sequences, out.confidence, out.trailer_text
90
 
91
  # Suppress the trailer while keeping stock behavior:
92
- sequences = model.generate(inputs, max_new_tokens=512, emit_trailer=False)
93
  ```
94
 
95
  ## Probe contract
@@ -107,14 +125,6 @@ sequences = model.generate(inputs, max_new_tokens=512, emit_trailer=False)
107
  - Capture uses a forward hook that keeps only one `[1, 1536]` row per decode
108
  step — full hidden-state stacks are never materialized.
109
 
110
- ## MLX (Apple Silicon)
111
-
112
- The repo ships a single-file `mlx-lm` model (`gemma_4_e2b_it_hybrid.py`),
113
- referenced by `"model_file"` in `config.json` (mlx-lm >= 0.30.1). After
114
- generation, read `model.last_confidence` (or `model.confidence(num_tokens=N)`).
115
- An mlx-lm `Model` cannot inject tokens into the stream, so there is **no
116
- in-band trailer** on MLX today — confidence is Python-API only there.
117
-
118
  ## Repo contents
119
 
120
  | File | Purpose |
@@ -125,5 +135,18 @@ in-band trailer** on MLX today — confidence is Python-API only there.
125
  | `model*.safetensors` | base weights (identical keys) + `handoff_probe.*` tensors |
126
  | `gemma_4_e2b_it_hybrid.py` | single-file `mlx-lm` model, wired via config.json's `model_file` |
127
 
128
- Requires `transformers>=5,<6`. Gemma model use remains subject to the Gemma
129
- terms.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9
  - custom_generate
10
  ---
11
 
12
+ # Cactus Hybrid — Gemma 4 E2B
13
 
14
+ A small, on-device model is fast and private, but sometimes wrong. At Cactus we
15
+ post-train models to *know when they are wrong*: we ship probes inside the
16
+ checkpoint that score every answer with a **confidence** between 0 and 1,
17
+ returned as structured data (never parsed out of the answer text). Answer
18
+ on-device when confidence is high; re-route to a bigger model when it's low —
19
+ `0.85` is a good threshold:
20
 
21
+ ```python
22
+ if confidence < 0.85:
23
+ answer = ask_a_bigger_model(prompt)
24
  ```
25
+
26
+ This repo is **google/gemma-4-E2B-it plus the handoff probe**: a small head
27
+ (weight prefix `handoff_probe.*`) that scores every generation with
28
+ `confidence = 1 - p_wrong`. The base weights are byte-identical to the stock
29
+ checkpoint (same keys); the repo adds eleven probe tensors, a remote-code model
30
+ class and a `custom_generate` recipe. **Stock engine commands work unchanged** —
31
+ you only add `--trust-remote-code` / `trust_remote_code=True`.
32
+
33
+ ## Quickstart
34
+
35
+ ```python
36
+ # pip install "transformers>=5.5.4,<5.6" torch (5.14+ segfaults on this checkpoint)
37
+ import torch
38
+ from transformers import AutoModelForCausalLM, AutoTokenizer
39
+
40
+ model_id = "Cactus-Compute/gemma-4-e2b-it-hybrid"
41
+ device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
42
+
43
+ tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
44
+ model = AutoModelForCausalLM.from_pretrained(model_id, trust_remote_code=True, dtype="auto").to(device)
45
+
46
+ messages = [{"role": "user", "content": "What is the capital of France?"}]
47
+ inputs = tokenizer.apply_chat_template(
48
+ messages, add_generation_prompt=True, return_tensors="pt", return_dict=True
49
+ ).to(device)
50
+ out = model.generate(**inputs, return_confidence=True, max_new_tokens=512)
51
+
52
+ print(tokenizer.decode(out.sequences[0][inputs["input_ids"].shape[-1]:], skip_special_tokens=True))
53
+ print("confidence:", out.confidence)
54
  ```
55
 
56
+ Load the model with an explicit `.to(device)`, not `device_map="auto"`: the
57
+ probe scores generations outside the module `forward()` path, so weights that
58
+ accelerate offloads (left on the `meta` device) crash the confidence read.
 
 
59
 
60
  ## Serving with stock transformers
61
 
 
95
  - The trailer's token ids are appended to the returned sequences, so reported
96
  completion token counts include the trailer (a handful of tokens).
97
 
98
+ ## More Python APIs
99
 
100
  ```python
 
 
 
 
 
 
 
 
 
 
 
101
  # Structured API: clean sequences + raw float (no in-band trailer).
102
  sequences, confidence = model.generate_with_confidence(inputs, max_new_tokens=512)
103
  print(confidence) # e.g. 0.7812
104
  print(model.last_confidence) # same value
105
 
106
  # Stock generate (custom_generate recipe): plain tensor + in-band trailer.
107
+ sequences = model.generate(**inputs, max_new_tokens=512)
 
 
 
 
108
 
109
  # Suppress the trailer while keeping stock behavior:
110
+ sequences = model.generate(**inputs, max_new_tokens=512, emit_trailer=False)
111
  ```
112
 
113
  ## Probe contract
 
125
  - Capture uses a forward hook that keeps only one `[1, 1536]` row per decode
126
  step — full hidden-state stacks are never materialized.
127
 
 
 
 
 
 
 
 
 
128
  ## Repo contents
129
 
130
  | File | Purpose |
 
135
  | `model*.safetensors` | base weights (identical keys) + `handoff_probe.*` tensors |
136
  | `gemma_4_e2b_it_hybrid.py` | single-file `mlx-lm` model, wired via config.json's `model_file` |
137
 
138
+ ## All formats
139
+
140
+ All Cactus Hybrid builds live in the
141
+ [Cactus Hybrid collection](https://huggingface.co/collections/Cactus-Compute/cactus-hybrid-6a60da4551074db058e8bb64):
142
+ [Transformers](https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid) ·
143
+ [GGUF / llama.cpp](https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid-GGUF) ·
144
+ [MLX](https://huggingface.co/Cactus-Compute/gemma-4-e2b-it-hybrid-mlx) ·
145
+ [Cactus engine](https://huggingface.co/Cactus-Compute/gemma-4-E2B-it).
146
+ Copy-paste quickstarts for every engine:
147
+ [github.com/cactus-compute/cactus-hybrid](https://github.com/cactus-compute/cactus-hybrid).
148
+
149
+ ## License
150
+
151
+ Gemma is provided under and subject to the Gemma Terms of Use. This derivative
152
+ includes the Cactus handoff probe head.