Henry Ndubuaku
Model card: Cactus Hybrid quickstart + collection/GitHub links
3504e19 verified
|
Raw History Blame
6.29 kB
metadata
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:

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

# 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

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.