Views
No views yet
onnx-community/gemma-4-E4B-it-ONNX. Multimodal heads (vision, audio) and non-Q4 quantization variants are stripped; the text decoder + embed table at INT4 (MatMulNBits) are kept as-is.config.json # text-only (text_config promoted)
generation_config.json
tokenizer.json + tokenizer_config.json + chat_template.jinja
onnx/
embed_tokens_q4.onnx (+ every .onnx_data* sidecar)
decoder_model_merged_q4.onnx (+ every .onnx_data* sidecar)embed_tokens_q4 — input_ids → inputs_embeds, per_layer_inputsdecoder_model_merged_q4 — inputs_embeds, per_layer_inputs, attention_mask, position_ids, past_key_values.* → logits, present.*ORTModelForCausalLM.from_pretrained(...).generate() does not orchestrate this 2-graph pattern as of optimum==1.x / optimum==2.x dev. Use the custom loop below.1import onnxruntime as ort
2import numpy as np
3from transformers import AutoTokenizer
4from huggingface_hub import snapshot_download
5
6local = snapshot_download("tss-deposium/gemma-4-E4B-text-only-onnx-int4")
7tok = AutoTokenizer.from_pretrained(local)
8
9providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] if ort.get_device() == "GPU" else ["CPUExecutionProvider"]
10embed_session = ort.InferenceSession(f"{local}/onnx/embed_tokens_q4.onnx", providers=providers)
11decoder_session = ort.InferenceSession(f"{local}/onnx/decoder_model_merged_q4.onnx", providers=providers)
12
13# --- Discover I/O signatures (Gemma 4 specifics: per_layer_inputs, KV cache layout) ---
14DEC_INPUTS = {x.name: x for x in decoder_session.get_inputs()}
15DEC_OUTPUTS = {x.name: x for x in decoder_session.get_outputs()}
16EMB_OUTPUTS = [x.name for x in embed_session.get_outputs()]
17PAST_KEYS = sorted([n for n in DEC_INPUTS if n.startswith("past_key_values")])
18PRESENT_KEYS = sorted([n for n in DEC_OUTPUTS if n.startswith(("present.", "present_key_values."))])
19HAS_PER_LAYER = "per_layer_inputs" in DEC_INPUTS
20HAS_NUM_LOGITS_TO_KEEP = "num_logits_to_keep" in DEC_INPUTS
21present_to_past = dict(zip(PRESENT_KEYS, PAST_KEYS))
22
23def _ort_dtype(t):
24 return {"tensor(float16)": np.float16, "tensor(float)": np.float32,
25 "tensor(int64)": np.int64, "tensor(int32)": np.int32, "tensor(bool)": np.bool_}[t]
26
27def _init_past_kv(batch=1):
28 config_dims = {}
29 for n in PAST_KEYS:
30 for i, d in enumerate(DEC_INPUTS[n].shape):
31 if isinstance(d, int) and d > 0: config_dims.setdefault(i, d)
32 pkv = {}
33 for n in PAST_KEYS:
34 meta = DEC_INPUTS[n]
35 shape = []
36 for i, d in enumerate(meta.shape):
37 if isinstance(d, int): shape.append(d)
38 elif i == 0: shape.append(batch)
39 elif i == 2: shape.append(0)
40 elif i in config_dims: shape.append(config_dims[i])
41 else: raise RuntimeError(f"unresolved symbolic dim {d} in {n}")
42 pkv[n] = np.zeros(shape, dtype=_ort_dtype(meta.type))
43 return pkv
44
45def _run_embed(input_ids):
46 outs = embed_session.run(None, {embed_session.get_inputs()[0].name: input_ids})
47 out_map = dict(zip(EMB_OUTPUTS, outs))
48 embeds = next(out_map[n] for n in EMB_OUTPUTS if "embed" in n.lower())
49 per_layer = next((out_map[n] for n in EMB_OUTPUTS if "per_layer" in n.lower()), None)
50 return embeds, per_layer
51
52def _run_decoder(embeds, per_layer, mask, pos, past):
53 feeds = {
54 "inputs_embeds": embeds.astype(_ort_dtype(DEC_INPUTS["inputs_embeds"].type)),
55 "attention_mask": mask.astype(_ort_dtype(DEC_INPUTS["attention_mask"].type)),
56 "position_ids": pos.astype(_ort_dtype(DEC_INPUTS["position_ids"].type)),
57 }
58 if HAS_PER_LAYER:
59 feeds["per_layer_inputs"] = per_layer.astype(_ort_dtype(DEC_INPUTS["per_layer_inputs"].type))
60 if HAS_NUM_LOGITS_TO_KEEP:
61 feeds["num_logits_to_keep"] = np.array(1, dtype=_ort_dtype(DEC_INPUTS["num_logits_to_keep"].type))
62 feeds.update(past)
63 outs = decoder_session.run(None, feeds)
64 out_names = [o.name for o in decoder_session.get_outputs()]
65 name_to_arr = dict(zip(out_names, outs))
66 return name_to_arr.get("logits", outs[0]), {present_to_past[n]: name_to_arr[n] for n in PRESENT_KEYS}
67
68def generate(prompt: str, max_new_tokens: int = 64) -> str:
69 msgs = [{"role": "user", "content": prompt}]
70 try:
71 text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True, enable_thinking=False)
72 except TypeError:
73 text = tok.apply_chat_template(msgs, tokenize=False, add_generation_prompt=True)
74 ids = tok(text, return_tensors="np").input_ids.astype(np.int64)
75
76 import json
77 gen_cfg_path = f"{local}/generation_config.json"
78 eos_raw = json.load(open(gen_cfg_path)).get("eos_token_id", tok.eos_token_id)
79 eos_ids = set(eos_raw) if isinstance(eos_raw, list) else {eos_raw}
80
81 past = _init_past_kv(batch=1)
82 seq = ids.shape[1]
83 mask = np.ones((1, seq), dtype=np.int64)
84 pos = np.arange(seq, dtype=np.int64)[None, :]
85
86 embeds, per_layer = _run_embed(ids)
87 logits, past = _run_decoder(embeds, per_layer, mask, pos, past)
88 next_id = int(np.argmax(logits[:, -1, :]))
89 out = [next_id]
90 cur = seq
91 while len(out) < max_new_tokens and next_id not in eos_ids:
92 cur += 1
93 tok_arr = np.array([[next_id]], dtype=np.int64)
94 mask = np.concatenate([mask, np.ones((1, 1), dtype=np.int64)], axis=1)
95 pos = np.array([[cur - 1]], dtype=np.int64)
96 embeds, per_layer = _run_embed(tok_arr)
97 logits, past = _run_decoder(embeds, per_layer, mask, pos, past)
98 next_id = int(np.argmax(logits[:, -1, :]))
99 out.append(next_id)
100 return tok.decode(out, skip_special_tokens=True)
101
102print(generate("Bonjour, explique en une phrase ce que tu es."))onnx-community/gemma-4-E4B-it-ONNXtheseedship/deposium-turbov3/docs/gemma4_e4b_text_only_onnx_int4_export.ipynb