--- 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.