Add files using upload-large-folder tool
Browse files- .gitattributes +1 -0
- README.md +129 -0
- chat_template.jinja +386 -0
- config.json +96 -0
- configuration_gemma_4_e2b_it_hybrid.py +46 -0
- custom_generate/generate.py +178 -0
- custom_generate/requirements.txt +1 -0
- gemma_4_e2b_it_hybrid.py +219 -0
- generation_config.json +14 -0
- model-00001-of-00003.safetensors +3 -0
- model-00002-of-00003.safetensors +3 -0
- model-00003-of-00003.safetensors +3 -0
- model.safetensors.index.json +618 -0
- modeling_gemma_4_e2b_it_hybrid.py +165 -0
- tokenizer.json +3 -0
- tokenizer_config.json +120 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
tokenizer.json filter=lfs diff=lfs merge=lfs -text
|
README.md
ADDED
|
@@ -0,0 +1,129 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: gemma
|
| 3 |
+
base_model: google/gemma-4-E2B-it
|
| 4 |
+
pipeline_tag: text-generation
|
| 5 |
+
tags:
|
| 6 |
+
- gemma4
|
| 7 |
+
- hybrid
|
| 8 |
+
- custom_code
|
| 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 |
+
|
| 29 |
+
```bash
|
| 30 |
+
transformers serve --trust-remote-code
|
| 31 |
+
# then request model "Cactus-Compute/gemma-4-e2b-it-hybrid" via the OpenAI-compatible API
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
or interactively:
|
| 35 |
+
|
| 36 |
+
```bash
|
| 37 |
+
transformers chat Cactus-Compute/gemma-4-e2b-it-hybrid --trust-remote-code
|
| 38 |
+
```
|
| 39 |
+
|
| 40 |
+
### How the confidence reaches you (in-band trailer)
|
| 41 |
+
|
| 42 |
+
`transformers serve` cannot add response fields, so the score travels **in-band**:
|
| 43 |
+
the assistant content's final line is
|
| 44 |
+
|
| 45 |
+
```
|
| 46 |
+
\n[[hybrid:confidence=0.7812]]
|
| 47 |
+
```
|
| 48 |
+
|
| 49 |
+
always exactly 4 decimals, ASCII, `confidence` in `[0, 1]`. Strip the final
|
| 50 |
+
`[[hybrid:...]]` line before display and parse the float for routing. The
|
| 51 |
+
trailer is emitted in both streaming and non-streaming modes.
|
| 52 |
+
|
| 53 |
+
### Limitations
|
| 54 |
+
|
| 55 |
+
- `--continuous-batching`: not supported — the CB scheduler bypasses
|
| 56 |
+
`generate()`, so no probe runs and **no trailer is emitted**. Serve without
|
| 57 |
+
`--continuous-batching` to get confidence scores.
|
| 58 |
+
- The trailer (and confidence) is only produced for single-sequence decoding:
|
| 59 |
+
batch size 1, no beam search, no assisted/speculative decoding. Unsupported
|
| 60 |
+
modes fall back to stock behavior (no trailer).
|
| 61 |
+
- The probe scores at most the first 1024 generated tokens.
|
| 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
|
| 96 |
+
|
| 97 |
+
- Input: float32 `[T, 1536]` — output of decoder layer index 28
|
| 98 |
+
(`config.probe_layer`), captured at the position that predicts each generated
|
| 99 |
+
token: row 0 = last prompt position at prefill, row t = position captured at
|
| 100 |
+
generation step t. Only the first 1024 rows are scored.
|
| 101 |
+
- Math (float32): `x = LayerNorm(x, eps=1e-5) * norm.weight + norm.bias`;
|
| 102 |
+
`p = relu(x @ proj.weight.T + proj.bias)`;
|
| 103 |
+
`s = p @ attn_query / sqrt(32)`; `w = softmax_T(s - max)`; `pooled = w @ p`;
|
| 104 |
+
`h = relu(head.0 @ pooled + b)`; `h = relu(head.2 @ h + b)`;
|
| 105 |
+
`logit = head.4 @ h + b`; `p_wrong = sigmoid(logit)`;
|
| 106 |
+
`confidence = 1 - p_wrong`.
|
| 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 |
|
| 121 |
+
|---|---|
|
| 122 |
+
| `configuration_gemma_4_e2b_it_hybrid.py` | `Gemma4E2BItHybridConfig` (stock Gemma-4 text config + probe hyperparams) |
|
| 123 |
+
| `modeling_gemma_4_e2b_it_hybrid.py` | `Gemma4E2BItHybridForCausalLM` (stock `Gemma4ForCausalLM` + `handoff_probe.*`) |
|
| 124 |
+
| `custom_generate/generate.py` | stock decode loop + confidence + in-band trailer |
|
| 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.
|
chat_template.jinja
ADDED
|
@@ -0,0 +1,386 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{#
|
| 2 |
+
Template: Google Gemma 4 Canonical Chat Template
|
| 3 |
+
Author: Google Gemma Engineering Team
|
| 4 |
+
Published: 2026-07-09
|
| 5 |
+
Context: Fixed tool-calling loops, turn closures, and thinking content-ordering.
|
| 6 |
+
#}
|
| 7 |
+
{%- macro format_parameters(properties, required, filter_keys=false) -%}
|
| 8 |
+
{%- set standard_keys = ['description', 'type', 'properties', 'required', 'nullable'] -%}
|
| 9 |
+
{%- set ns = namespace(found_first=false) -%}
|
| 10 |
+
{%- for key, value in properties | dictsort -%}
|
| 11 |
+
{%- set add_comma = false -%}
|
| 12 |
+
{%- if not filter_keys or key not in standard_keys -%}
|
| 13 |
+
{%- if ns.found_first %},{% endif -%}
|
| 14 |
+
{%- set ns.found_first = true -%}
|
| 15 |
+
{{ key }}:{
|
| 16 |
+
{%- if value['description'] -%}
|
| 17 |
+
description:<|"|>{{ value['description'] }}<|"|>
|
| 18 |
+
{%- set add_comma = true -%}
|
| 19 |
+
{%- endif -%}
|
| 20 |
+
{%- if value['type'] | upper == 'STRING' -%}
|
| 21 |
+
{%- if value['enum'] -%}
|
| 22 |
+
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
|
| 23 |
+
enum:{{ format_argument(value['enum']) }}
|
| 24 |
+
{%- endif -%}
|
| 25 |
+
{%- elif value['type'] | upper == 'ARRAY' -%}
|
| 26 |
+
{%- if value['items'] is mapping and value['items'] -%}
|
| 27 |
+
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
|
| 28 |
+
items:{
|
| 29 |
+
{%- set ns_items = namespace(found_first=false) -%}
|
| 30 |
+
{%- for item_key, item_value in value['items'] | dictsort -%}
|
| 31 |
+
{%- if item_value is not none -%}
|
| 32 |
+
{%- if ns_items.found_first %},{% endif -%}
|
| 33 |
+
{%- set ns_items.found_first = true -%}
|
| 34 |
+
{%- if item_key == 'properties' -%}
|
| 35 |
+
properties:{
|
| 36 |
+
{%- if item_value is mapping -%}
|
| 37 |
+
{{- format_parameters(item_value, value['items']['required'] | default([])) -}}
|
| 38 |
+
{%- endif -%}
|
| 39 |
+
}
|
| 40 |
+
{%- elif item_key == 'required' -%}
|
| 41 |
+
required:[
|
| 42 |
+
{%- for req_item in item_value -%}
|
| 43 |
+
<|"|>{{- req_item -}}<|"|>
|
| 44 |
+
{%- if not loop.last %},{% endif -%}
|
| 45 |
+
{%- endfor -%}
|
| 46 |
+
]
|
| 47 |
+
{%- elif item_key == 'type' -%}
|
| 48 |
+
{%- if item_value is string -%}
|
| 49 |
+
type:{{ format_argument(item_value | upper) }}
|
| 50 |
+
{%- else -%}
|
| 51 |
+
type:{{ format_argument(item_value | map('upper') | list) }}
|
| 52 |
+
{%- endif -%}
|
| 53 |
+
{%- else -%}
|
| 54 |
+
{{ item_key }}:{{ format_argument(item_value) }}
|
| 55 |
+
{%- endif -%}
|
| 56 |
+
{%- endif -%}
|
| 57 |
+
{%- endfor -%}
|
| 58 |
+
}
|
| 59 |
+
{%- endif -%}
|
| 60 |
+
{%- endif -%}
|
| 61 |
+
{%- if value['nullable'] %}
|
| 62 |
+
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
|
| 63 |
+
nullable:true
|
| 64 |
+
{%- endif -%}
|
| 65 |
+
{%- if value['type'] | upper == 'OBJECT' -%}
|
| 66 |
+
{%- if value['properties'] is defined and value['properties'] is mapping -%}
|
| 67 |
+
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
|
| 68 |
+
properties:{
|
| 69 |
+
{{- format_parameters(value['properties'], value['required'] | default([])) -}}
|
| 70 |
+
}
|
| 71 |
+
{%- elif value is mapping -%}
|
| 72 |
+
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
|
| 73 |
+
properties:{
|
| 74 |
+
{{- format_parameters(value, value['required'] | default([]), filter_keys=true) -}}
|
| 75 |
+
}
|
| 76 |
+
{%- endif -%}
|
| 77 |
+
{%- if value['required'] -%}
|
| 78 |
+
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
|
| 79 |
+
required:[
|
| 80 |
+
{%- for item in value['required'] | default([]) -%}
|
| 81 |
+
<|"|>{{- item -}}<|"|>
|
| 82 |
+
{%- if not loop.last %},{% endif -%}
|
| 83 |
+
{%- endfor -%}
|
| 84 |
+
]
|
| 85 |
+
{%- endif -%}
|
| 86 |
+
{%- endif -%}
|
| 87 |
+
{%- if add_comma %},{%- else -%} {%- set add_comma = true -%} {% endif -%}
|
| 88 |
+
type:<|"|>{{ value['type'] | upper }}<|"|>}
|
| 89 |
+
{%- endif -%}
|
| 90 |
+
{%- endfor -%}
|
| 91 |
+
{%- endmacro -%}
|
| 92 |
+
{%- macro format_function_declaration(tool_data) -%}
|
| 93 |
+
declaration:{{- tool_data['function']['name'] -}}{description:<|"|>{{- tool_data['function']['description'] -}}<|"|>
|
| 94 |
+
{%- set params = tool_data['function']['parameters'] -%}
|
| 95 |
+
{%- if params -%}
|
| 96 |
+
,parameters:{
|
| 97 |
+
{%- if params['properties'] -%}
|
| 98 |
+
properties:{ {{- format_parameters(params['properties'], params['required']) -}} },
|
| 99 |
+
{%- endif -%}
|
| 100 |
+
{%- if params['required'] -%}
|
| 101 |
+
required:[
|
| 102 |
+
{%- for item in params['required'] -%}
|
| 103 |
+
<|"|>{{- item -}}<|"|>
|
| 104 |
+
{{- ',' if not loop.last -}}
|
| 105 |
+
{%- endfor -%}
|
| 106 |
+
],
|
| 107 |
+
{%- endif -%}
|
| 108 |
+
{%- if params['type'] -%}
|
| 109 |
+
type:<|"|>{{- params['type'] | upper -}}<|"|>}
|
| 110 |
+
{%- endif -%}
|
| 111 |
+
{%- endif -%}
|
| 112 |
+
{%- if 'response' in tool_data['function'] -%}
|
| 113 |
+
{%- set response_declaration = tool_data['function']['response'] -%}
|
| 114 |
+
,response:{
|
| 115 |
+
{%- if response_declaration['description'] -%}
|
| 116 |
+
description:<|"|>{{- response_declaration['description'] -}}<|"|>,
|
| 117 |
+
{%- endif -%}
|
| 118 |
+
{%- if response_declaration['type'] | upper == 'OBJECT' -%}
|
| 119 |
+
type:<|"|>{{- response_declaration['type'] | upper -}}<|"|>}
|
| 120 |
+
{%- endif -%}
|
| 121 |
+
{%- endif -%}
|
| 122 |
+
}
|
| 123 |
+
{%- endmacro -%}
|
| 124 |
+
{%- macro format_argument(argument, escape_keys=True) -%}
|
| 125 |
+
{%- if argument is none -%}
|
| 126 |
+
{{- 'null' -}}
|
| 127 |
+
{%- elif argument is string -%}
|
| 128 |
+
{{- '<|"|>' + argument + '<|"|>' -}}
|
| 129 |
+
{%- elif argument is boolean -%}
|
| 130 |
+
{{- 'true' if argument else 'false' -}}
|
| 131 |
+
{%- elif argument is mapping -%}
|
| 132 |
+
{{- '{' -}}
|
| 133 |
+
{%- set ns = namespace(found_first=false) -%}
|
| 134 |
+
{%- for key, value in argument | dictsort -%}
|
| 135 |
+
{%- if ns.found_first %},{% endif -%}
|
| 136 |
+
{%- set ns.found_first = true -%}
|
| 137 |
+
{%- if escape_keys -%}
|
| 138 |
+
{{- '<|"|>' + key + '<|"|>' -}}
|
| 139 |
+
{%- else -%}
|
| 140 |
+
{{- key -}}
|
| 141 |
+
{%- endif -%}
|
| 142 |
+
:{{- format_argument(value, escape_keys=escape_keys) -}}
|
| 143 |
+
{%- endfor -%}
|
| 144 |
+
{{- '}' -}}
|
| 145 |
+
{%- elif argument is sequence -%}
|
| 146 |
+
{{- '[' -}}
|
| 147 |
+
{%- for item in argument -%}
|
| 148 |
+
{{- format_argument(item, escape_keys=escape_keys) -}}
|
| 149 |
+
{%- if not loop.last %},{% endif -%}
|
| 150 |
+
{%- endfor -%}
|
| 151 |
+
{{- ']' -}}
|
| 152 |
+
{%- else -%}
|
| 153 |
+
{{- argument -}}
|
| 154 |
+
{%- endif -%}
|
| 155 |
+
{%- endmacro -%}
|
| 156 |
+
{%- macro strip_thinking(text) -%}
|
| 157 |
+
{%- set ns = namespace(result='') -%}
|
| 158 |
+
{%- for part in text.split('<channel|>') -%}
|
| 159 |
+
{%- if '<|channel>' in part -%}
|
| 160 |
+
{%- set ns.result = ns.result + part.split('<|channel>')[0] -%}
|
| 161 |
+
{%- else -%}
|
| 162 |
+
{%- set ns.result = ns.result + part -%}
|
| 163 |
+
{%- endif -%}
|
| 164 |
+
{%- endfor -%}
|
| 165 |
+
{{- ns.result | trim -}}
|
| 166 |
+
{%- endmacro -%}
|
| 167 |
+
|
| 168 |
+
{%- macro format_tool_response_block(tool_name, response) -%}
|
| 169 |
+
{{- '<|tool_response>' -}}
|
| 170 |
+
{%- if response is mapping -%}
|
| 171 |
+
{{- 'response:' + tool_name + '{' -}}
|
| 172 |
+
{%- for key, value in response | dictsort -%}
|
| 173 |
+
{{- key -}}:{{- format_argument(value, escape_keys=False) -}}
|
| 174 |
+
{%- if not loop.last %},{% endif -%}
|
| 175 |
+
{%- endfor -%}
|
| 176 |
+
{{- '}' -}}
|
| 177 |
+
{%- else -%}
|
| 178 |
+
{{- 'response:' + tool_name + '{value:' + format_argument(response, escape_keys=False) + '}' -}}
|
| 179 |
+
{%- endif -%}
|
| 180 |
+
{{- '<tool_response|>' -}}
|
| 181 |
+
{%- endmacro -%}
|
| 182 |
+
|
| 183 |
+
{#- ===== SETUP ===== -#}
|
| 184 |
+
{%- set ns = namespace(prev_message_type=None, prev_non_tool_role=None) -%}
|
| 185 |
+
{%- set loop_messages = messages -%}
|
| 186 |
+
{%- set enable_thinking = enable_thinking | default(false) -%}
|
| 187 |
+
{%- set preserve_thinking = preserve_thinking | default(false) -%}
|
| 188 |
+
{{- bos_token -}}
|
| 189 |
+
{#- Handle System/Tool Definitions Block -#}
|
| 190 |
+
{%- if enable_thinking or tools or (messages and messages[0]['role'] in ['system', 'developer']) -%}
|
| 191 |
+
{{- '<|turn>system\n' -}}
|
| 192 |
+
{#- Inject Thinking token at the very top of the FIRST system turn -#}
|
| 193 |
+
{%- if enable_thinking -%}
|
| 194 |
+
{{- '<|think|>\n' -}}
|
| 195 |
+
{%- set ns.prev_message_type = 'think' -%}
|
| 196 |
+
{%- endif -%}
|
| 197 |
+
{%- if messages and messages[0]['role'] in ['system', 'developer'] -%}
|
| 198 |
+
{%- if messages[0]['content'] is string -%}
|
| 199 |
+
{{- messages[0]['content'] | trim -}}
|
| 200 |
+
{%- elif messages[0]['content'] is sequence -%}
|
| 201 |
+
{%- for item in messages[0]['content'] -%}
|
| 202 |
+
{{- item['text'] | trim + ' '-}}
|
| 203 |
+
{%- endfor -%}
|
| 204 |
+
{%- endif -%}
|
| 205 |
+
{%- set loop_messages = messages[1:] -%}
|
| 206 |
+
{%- endif -%}
|
| 207 |
+
{%- if tools -%}
|
| 208 |
+
{%- for tool in tools %}
|
| 209 |
+
{{- '<|tool>' -}}
|
| 210 |
+
{{- format_function_declaration(tool) | trim -}}
|
| 211 |
+
{{- '<tool|>' -}}
|
| 212 |
+
{%- endfor %}
|
| 213 |
+
{%- set ns.prev_message_type = 'tool' -%}
|
| 214 |
+
{%- endif -%}
|
| 215 |
+
{{- '<turn|>\n' -}}
|
| 216 |
+
{%- endif %}
|
| 217 |
+
|
| 218 |
+
{#- Pre-scan: find last user message index for reasoning guard -#}
|
| 219 |
+
{%- set ns_turn = namespace(last_user_idx=-1) -%}
|
| 220 |
+
{%- for i in range(loop_messages | length) -%}
|
| 221 |
+
{%- if loop_messages[i]['role'] == 'user' -%}
|
| 222 |
+
{%- set ns_turn.last_user_idx = i -%}
|
| 223 |
+
{%- endif -%}
|
| 224 |
+
{%- endfor -%}
|
| 225 |
+
|
| 226 |
+
{#- Loop through messages -#}
|
| 227 |
+
{%- for message in loop_messages -%}
|
| 228 |
+
{%- if message['role'] != 'tool' -%}
|
| 229 |
+
{%- set ns.prev_message_type = None -%}
|
| 230 |
+
{%- set role = 'model' if message['role'] == 'assistant' else message['role'] -%}
|
| 231 |
+
{#- Detect continuation using tracked state — O(1) instead of O(n) backward scan -#}
|
| 232 |
+
{%- set continue_same_model_turn = (role == 'model' and ns.prev_non_tool_role == 'assistant') -%}
|
| 233 |
+
{%- if not continue_same_model_turn -%}
|
| 234 |
+
{{- '<|turn>' + role + '\n' }}
|
| 235 |
+
{%- endif -%}
|
| 236 |
+
|
| 237 |
+
{#- Render reasoning/reasoning_content as thinking channel -#}
|
| 238 |
+
{%- set thinking_text = message.get('reasoning') or message.get('reasoning_content') -%}
|
| 239 |
+
{%- set thinking_gate = (loop.index0 > ns_turn.last_user_idx) or (preserve_thinking and message.get('tool_calls')) -%}
|
| 240 |
+
{%- if thinking_text and thinking_gate -%}
|
| 241 |
+
{{- '<|channel>thought\n' + thinking_text + '\n<channel|>' -}}
|
| 242 |
+
{%- endif -%}
|
| 243 |
+
|
| 244 |
+
{%- if message.get('tool_calls') -%}
|
| 245 |
+
{%- for tool_call in message.get('tool_calls') -%}
|
| 246 |
+
{%- set function = tool_call['function'] -%}
|
| 247 |
+
{{- '<|tool_call>call:' + function['name'] + '{' -}}
|
| 248 |
+
{%- if function['arguments'] is mapping -%}
|
| 249 |
+
{%- set ns_args = namespace(found_first=false) -%}
|
| 250 |
+
{%- for key, value in function['arguments'] | dictsort -%}
|
| 251 |
+
{%- if ns_args.found_first %},{% endif -%}
|
| 252 |
+
{%- set ns_args.found_first = true -%}
|
| 253 |
+
{{- key -}}:{{- format_argument(value, escape_keys=False) -}}
|
| 254 |
+
{%- endfor -%}
|
| 255 |
+
{%- elif function['arguments'] is none -%}
|
| 256 |
+
{%- else -%}
|
| 257 |
+
{{- raise_exception(
|
| 258 |
+
"chat_template: tool_calls[].function.arguments must be a "
|
| 259 |
+
"JSON object (mapping), not a string. Deserialize arguments "
|
| 260 |
+
"before passing to the template."
|
| 261 |
+
) -}}
|
| 262 |
+
{%- endif -%}
|
| 263 |
+
{{- '}<tool_call|>' -}}
|
| 264 |
+
{%- endfor -%}
|
| 265 |
+
{%- set ns.prev_message_type = 'tool_call' -%}
|
| 266 |
+
{%- endif -%}
|
| 267 |
+
|
| 268 |
+
{%- set ns_tr_out = namespace(flag=false) -%}
|
| 269 |
+
{%- if message.get('tool_responses') -%}
|
| 270 |
+
{#- Legacy: tool_responses embedded on the assistant message (Google/Gemma native) -#}
|
| 271 |
+
{%- for tool_response in message.get('tool_responses') -%}
|
| 272 |
+
{{- format_tool_response_block(tool_response['name'] | default('unknown', true), tool_response['response']) -}}
|
| 273 |
+
{%- set ns_tr_out.flag = true -%}
|
| 274 |
+
{%- set ns.prev_message_type = 'tool_response' -%}
|
| 275 |
+
{%- endfor -%}
|
| 276 |
+
{%- elif message.get('tool_calls') -%}
|
| 277 |
+
{#- OpenAI Chat Completions: forward-scan consecutive role:tool messages -#}
|
| 278 |
+
{%- set ns_tool_scan = namespace(stopped=false) -%}
|
| 279 |
+
{%- for k in range(loop.index0 + 1, loop_messages | length) -%}
|
| 280 |
+
{%- if ns_tool_scan.stopped -%}
|
| 281 |
+
{%- elif loop_messages[k]['role'] != 'tool' -%}
|
| 282 |
+
{%- set ns_tool_scan.stopped = true -%}
|
| 283 |
+
{%- else -%}
|
| 284 |
+
{%- set follow = loop_messages[k] -%}
|
| 285 |
+
{#- Resolve tool_call_id to function name -#}
|
| 286 |
+
{%- set ns_tname = namespace(name=follow.get('name') or 'unknown') -%}
|
| 287 |
+
{%- for tc in message.get('tool_calls') -%}
|
| 288 |
+
{%- if tc.get('id') == follow.get('tool_call_id') -%}
|
| 289 |
+
{%- set ns_tname.name = tc['function']['name'] -%}
|
| 290 |
+
{%- endif -%}
|
| 291 |
+
{%- endfor -%}
|
| 292 |
+
{#- Handle content as string or content-parts array -#}
|
| 293 |
+
{%- set tool_body = follow.get('content') -%}
|
| 294 |
+
{%- if tool_body is string -%}
|
| 295 |
+
{{- format_tool_response_block(ns_tname.name, tool_body) -}}
|
| 296 |
+
{%- elif tool_body is sequence and tool_body is not string -%}
|
| 297 |
+
{%- set ns_txt = namespace(s='') -%}
|
| 298 |
+
{%- for part in tool_body -%}
|
| 299 |
+
{%- if part.get('type') == 'text' -%}
|
| 300 |
+
{%- set ns_txt.s = ns_txt.s + (part.get('text') | default('')) -%}
|
| 301 |
+
{%- endif -%}
|
| 302 |
+
{%- endfor -%}
|
| 303 |
+
{{- format_tool_response_block(ns_tname.name, ns_txt.s) -}}
|
| 304 |
+
{%- for part in tool_body -%}
|
| 305 |
+
{%- if part.get('type') in ['image', 'image_url'] -%}
|
| 306 |
+
{{- '<|image|>' -}}
|
| 307 |
+
{%- elif part.get('type') in ['audio', 'input_audio'] -%}
|
| 308 |
+
{{- '<|audio|>' -}}
|
| 309 |
+
{%- elif part.get('type') == 'video' -%}
|
| 310 |
+
{{- '<|video|>' -}}
|
| 311 |
+
{%- endif -%}
|
| 312 |
+
{%- endfor -%}
|
| 313 |
+
{%- else -%}
|
| 314 |
+
{{- format_tool_response_block(ns_tname.name, tool_body) -}}
|
| 315 |
+
{%- endif -%}
|
| 316 |
+
{%- set ns_tr_out.flag = true -%}
|
| 317 |
+
{%- set ns.prev_message_type = 'tool_response' -%}
|
| 318 |
+
{%- endif -%}
|
| 319 |
+
{%- endfor -%}
|
| 320 |
+
{%- endif -%}
|
| 321 |
+
|
| 322 |
+
{%- set captured_content -%}
|
| 323 |
+
{%- if message.get('content') is string -%}
|
| 324 |
+
{%- if role == 'model' -%}
|
| 325 |
+
{{- strip_thinking(message['content']) -}}
|
| 326 |
+
{%- else -%}
|
| 327 |
+
{{- message['content'] | trim -}}
|
| 328 |
+
{%- endif -%}
|
| 329 |
+
{%- elif message.get('content') is sequence -%}
|
| 330 |
+
{%- for item in message['content'] -%}
|
| 331 |
+
{%- if item.get('type') == 'text' -%}
|
| 332 |
+
{%- if role == 'model' -%}
|
| 333 |
+
{{- strip_thinking(item['text']) -}}
|
| 334 |
+
{%- else -%}
|
| 335 |
+
{{- item['text'] | trim -}}
|
| 336 |
+
{%- endif -%}
|
| 337 |
+
{%- elif item.get('type') in ['image', 'image_url'] -%}
|
| 338 |
+
{{- '<|image|>' -}}
|
| 339 |
+
{%- elif item.get('type') in ['audio', 'input_audio'] -%}
|
| 340 |
+
{{- '<|audio|>' -}}
|
| 341 |
+
{%- elif item.get('type') == 'video' -%}
|
| 342 |
+
{{- '<|video|>' -}}
|
| 343 |
+
{%- endif -%}
|
| 344 |
+
{%- endfor -%}
|
| 345 |
+
{%- endif -%}
|
| 346 |
+
{%- endset -%}
|
| 347 |
+
|
| 348 |
+
{{- captured_content -}}
|
| 349 |
+
{%- set has_content = captured_content | trim | length > 0 -%}
|
| 350 |
+
|
| 351 |
+
{#- Forward-scan: find next non-tool message role for continuation detection -#}
|
| 352 |
+
{%- set next_nt = namespace(role=None, found=false) -%}
|
| 353 |
+
{%- for j in range(loop.index0 + 1, loop_messages | length) -%}
|
| 354 |
+
{%- if not next_nt.found -%}
|
| 355 |
+
{%- if loop_messages[j]['role'] != 'tool' -%}
|
| 356 |
+
{%- set next_nt.role = loop_messages[j]['role'] -%}
|
| 357 |
+
{%- set next_nt.found = true -%}
|
| 358 |
+
{%- endif -%}
|
| 359 |
+
{%- endif -%}
|
| 360 |
+
{%- endfor -%}
|
| 361 |
+
|
| 362 |
+
{%- set continues_into_next = (
|
| 363 |
+
role == 'model'
|
| 364 |
+
and next_nt.role == 'assistant'
|
| 365 |
+
and (not message.get('tool_calls') or ns_tr_out.flag)
|
| 366 |
+
) -%}
|
| 367 |
+
|
| 368 |
+
{%- if ns.prev_message_type == 'tool_call' and not ns_tr_out.flag -%}
|
| 369 |
+
{{- '<|tool_response>' -}}
|
| 370 |
+
{%- elif continues_into_next -%}
|
| 371 |
+
{%- elif not (ns_tr_out.flag and not has_content and not next_nt.found) -%}
|
| 372 |
+
{{- '<turn|>\n' -}}
|
| 373 |
+
{%- endif -%}
|
| 374 |
+
|
| 375 |
+
{#- Track previous non-tool role for next iteration (avoids O(n) backward scan) -#}
|
| 376 |
+
{%- set ns.prev_non_tool_role = message['role'] -%}
|
| 377 |
+
{%- endif -%}
|
| 378 |
+
{%- endfor -%}
|
| 379 |
+
|
| 380 |
+
{%- if add_generation_prompt -%}
|
| 381 |
+
{%- if ns.prev_message_type != 'tool_response' and ns.prev_message_type != 'tool_call' -%}
|
| 382 |
+
{{- '<|turn>model\n' -}}
|
| 383 |
+
{%- elif ns.prev_message_type == 'tool_response' and enable_thinking -%}
|
| 384 |
+
{{- '<|channel>thought\n' -}}
|
| 385 |
+
{%- endif -%}
|
| 386 |
+
{%- endif -%}
|
config.json
ADDED
|
@@ -0,0 +1,96 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"Gemma4E2BItHybridForCausalLM"
|
| 4 |
+
],
|
| 5 |
+
"attention_bias": false,
|
| 6 |
+
"attention_dropout": 0.0,
|
| 7 |
+
"attention_k_eq_v": false,
|
| 8 |
+
"auto_map": {
|
| 9 |
+
"AutoConfig": "configuration_gemma_4_e2b_it_hybrid.Gemma4E2BItHybridConfig",
|
| 10 |
+
"AutoModelForCausalLM": "modeling_gemma_4_e2b_it_hybrid.Gemma4E2BItHybridForCausalLM"
|
| 11 |
+
},
|
| 12 |
+
"bos_token_id": 2,
|
| 13 |
+
"dtype": "bfloat16",
|
| 14 |
+
"enable_moe_block": false,
|
| 15 |
+
"eos_token_id": 1,
|
| 16 |
+
"expert_intermediate_size": null,
|
| 17 |
+
"final_logit_softcapping": 30.0,
|
| 18 |
+
"global_head_dim": 512,
|
| 19 |
+
"head_dim": 256,
|
| 20 |
+
"hidden_activation": "gelu_pytorch_tanh",
|
| 21 |
+
"hidden_size": 1536,
|
| 22 |
+
"hidden_size_per_layer_input": 256,
|
| 23 |
+
"initializer_range": 0.02,
|
| 24 |
+
"intermediate_size": 6144,
|
| 25 |
+
"layer_types": [
|
| 26 |
+
"sliding_attention",
|
| 27 |
+
"sliding_attention",
|
| 28 |
+
"sliding_attention",
|
| 29 |
+
"sliding_attention",
|
| 30 |
+
"full_attention",
|
| 31 |
+
"sliding_attention",
|
| 32 |
+
"sliding_attention",
|
| 33 |
+
"sliding_attention",
|
| 34 |
+
"sliding_attention",
|
| 35 |
+
"full_attention",
|
| 36 |
+
"sliding_attention",
|
| 37 |
+
"sliding_attention",
|
| 38 |
+
"sliding_attention",
|
| 39 |
+
"sliding_attention",
|
| 40 |
+
"full_attention",
|
| 41 |
+
"sliding_attention",
|
| 42 |
+
"sliding_attention",
|
| 43 |
+
"sliding_attention",
|
| 44 |
+
"sliding_attention",
|
| 45 |
+
"full_attention",
|
| 46 |
+
"sliding_attention",
|
| 47 |
+
"sliding_attention",
|
| 48 |
+
"sliding_attention",
|
| 49 |
+
"sliding_attention",
|
| 50 |
+
"full_attention",
|
| 51 |
+
"sliding_attention",
|
| 52 |
+
"sliding_attention",
|
| 53 |
+
"sliding_attention",
|
| 54 |
+
"sliding_attention",
|
| 55 |
+
"full_attention",
|
| 56 |
+
"sliding_attention",
|
| 57 |
+
"sliding_attention",
|
| 58 |
+
"sliding_attention",
|
| 59 |
+
"sliding_attention",
|
| 60 |
+
"full_attention"
|
| 61 |
+
],
|
| 62 |
+
"max_position_embeddings": 131072,
|
| 63 |
+
"model_file": "gemma_4_e2b_it_hybrid.py",
|
| 64 |
+
"model_type": "gemma-4-e2b-it-hybrid",
|
| 65 |
+
"num_attention_heads": 8,
|
| 66 |
+
"num_experts": null,
|
| 67 |
+
"num_global_key_value_heads": null,
|
| 68 |
+
"num_hidden_layers": 35,
|
| 69 |
+
"num_key_value_heads": 1,
|
| 70 |
+
"num_kv_shared_layers": 20,
|
| 71 |
+
"pad_token_id": 0,
|
| 72 |
+
"probe_feature_size": 1536,
|
| 73 |
+
"probe_layer": 28,
|
| 74 |
+
"probe_max_tokens": 1024,
|
| 75 |
+
"probe_proj_dim": 32,
|
| 76 |
+
"rms_norm_eps": 1e-06,
|
| 77 |
+
"rope_parameters": {
|
| 78 |
+
"full_attention": {
|
| 79 |
+
"partial_rotary_factor": 0.25,
|
| 80 |
+
"rope_theta": 1000000.0,
|
| 81 |
+
"rope_type": "proportional"
|
| 82 |
+
},
|
| 83 |
+
"sliding_attention": {
|
| 84 |
+
"rope_theta": 10000.0,
|
| 85 |
+
"rope_type": "default"
|
| 86 |
+
}
|
| 87 |
+
},
|
| 88 |
+
"sliding_window": 512,
|
| 89 |
+
"tie_word_embeddings": true,
|
| 90 |
+
"top_k_experts": null,
|
| 91 |
+
"use_bidirectional_attention": null,
|
| 92 |
+
"use_cache": true,
|
| 93 |
+
"use_double_wide_mlp": true,
|
| 94 |
+
"vocab_size": 262144,
|
| 95 |
+
"vocab_size_per_layer_input": 262144
|
| 96 |
+
}
|
configuration_gemma_4_e2b_it_hybrid.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Configuration for Cactus-Compute/gemma-4-e2b-it-hybrid.
|
| 2 |
+
|
| 3 |
+
`Gemma4E2BItHybridConfig` is the stock Gemma-4 text config plus the hyperparameters of
|
| 4 |
+
the handoff probe that scores every generation with ``confidence = 1 - p_wrong``.
|
| 5 |
+
Everything the base model needs is inherited unchanged, so the base weights load
|
| 6 |
+
with identical keys.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
from transformers.models.gemma4.configuration_gemma4 import Gemma4TextConfig
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class Gemma4E2BItHybridConfig(Gemma4TextConfig):
|
| 13 |
+
"""Gemma-4 text config extended with handoff-probe hyperparameters.
|
| 14 |
+
|
| 15 |
+
Extra fields (all serialized to ``config.json``):
|
| 16 |
+
|
| 17 |
+
- ``probe_layer`` (`int`, defaults to 28): zero-indexed decoder layer whose
|
| 18 |
+
output feeds the probe. Must be ``< num_hidden_layers``.
|
| 19 |
+
- ``probe_feature_size`` (`int`, defaults to 1536): width of the captured
|
| 20 |
+
hidden states; must equal ``hidden_size``.
|
| 21 |
+
- ``probe_max_tokens`` (`int`, defaults to 1024): the probe scores at most
|
| 22 |
+
the first ``probe_max_tokens`` generated-token rows.
|
| 23 |
+
- ``probe_proj_dim`` (`int`, defaults to 32): width of the probe's attention
|
| 24 |
+
projection (fixed by the released checkpoint).
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
model_type = "gemma-4-e2b-it-hybrid"
|
| 28 |
+
|
| 29 |
+
probe_layer: int = 28
|
| 30 |
+
probe_feature_size: int = 1536
|
| 31 |
+
probe_max_tokens: int = 1024
|
| 32 |
+
probe_proj_dim: int = 32
|
| 33 |
+
|
| 34 |
+
def __post_init__(self, **kwargs):
|
| 35 |
+
super().__post_init__(**kwargs)
|
| 36 |
+
if not 0 <= self.probe_layer < self.num_hidden_layers:
|
| 37 |
+
raise ValueError(
|
| 38 |
+
f"probe_layer={self.probe_layer} must be in [0, num_hidden_layers="
|
| 39 |
+
f"{self.num_hidden_layers})"
|
| 40 |
+
)
|
| 41 |
+
# Note: probe_feature_size == hidden_size is enforced by the model, not
|
| 42 |
+
# here — `to_diff_dict()` default-constructs the config class, so this
|
| 43 |
+
# class must stay constructible with pure defaults.
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
__all__ = ["Gemma4E2BItHybridConfig"]
|
custom_generate/generate.py
ADDED
|
@@ -0,0 +1,178 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Hub ``custom_generate`` override for Cactus-Compute/gemma-4-e2b-it-hybrid.
|
| 2 |
+
|
| 3 |
+
Runs the stock decode loop, scores the generation with the handoff probe, and
|
| 4 |
+
transports ``confidence = 1 - p_wrong`` in-band: the assistant text gains a
|
| 5 |
+
final line ``\\n[[hybrid:confidence=0.7812]]`` (exactly 4 decimals, ASCII).
|
| 6 |
+
|
| 7 |
+
Return contract:
|
| 8 |
+
|
| 9 |
+
- default: a plain ``sequences`` tensor — required by ``transformers serve``,
|
| 10 |
+
which slices ``sequences[0, input_len:]`` on the raw return. The trailer's
|
| 11 |
+
token ids are appended to the returned sequences, and, when a ``streamer``
|
| 12 |
+
is passed, the trailer is also pushed through it, so serve emits the trailer
|
| 13 |
+
in both streaming and non-streaming modes.
|
| 14 |
+
- ``return_confidence=True``: a ``GenerateWithConfidenceOutput`` with clean
|
| 15 |
+
sequences (no trailer) and the raw ``confidence`` float.
|
| 16 |
+
- ``emit_trailer=False``: suppress the in-band trailer (Python API opt-out).
|
| 17 |
+
|
| 18 |
+
Generations the probe contract does not cover (batch > 1, beam search,
|
| 19 |
+
assisted decoding) fall back to the stock behavior: no trailer, confidence
|
| 20 |
+
``None``.
|
| 21 |
+
"""
|
| 22 |
+
|
| 23 |
+
from dataclasses import dataclass
|
| 24 |
+
from typing import Any, Optional
|
| 25 |
+
|
| 26 |
+
import torch
|
| 27 |
+
|
| 28 |
+
TRAILER_TEMPLATE = "\n[[hybrid:confidence={confidence:.4f}]]"
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
@dataclass
|
| 32 |
+
class GenerateWithConfidenceOutput:
|
| 33 |
+
"""Rich result for ``generate(..., return_confidence=True)``."""
|
| 34 |
+
|
| 35 |
+
sequences: torch.LongTensor
|
| 36 |
+
confidence: Optional[float]
|
| 37 |
+
trailer_text: Optional[str]
|
| 38 |
+
generate_output: Any # the raw return of the stock `generate`
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
class _DeferredEndStreamer:
|
| 42 |
+
"""Forwards ``put`` but swallows ``end`` so the trailer can be pushed
|
| 43 |
+
through the real streamer after the stock decode loop finishes."""
|
| 44 |
+
|
| 45 |
+
def __init__(self, inner):
|
| 46 |
+
self._inner = inner
|
| 47 |
+
|
| 48 |
+
def put(self, value):
|
| 49 |
+
self._inner.put(value)
|
| 50 |
+
|
| 51 |
+
def end(self):
|
| 52 |
+
pass
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _resolve_tokenizer(model, kwargs):
|
| 56 |
+
"""Best-effort tokenizer for encoding the trailer. Never raises."""
|
| 57 |
+
tok = kwargs.get("tokenizer")
|
| 58 |
+
if tok is None:
|
| 59 |
+
tok = getattr(model, "_hybrid_trailer_tokenizer", None)
|
| 60 |
+
if tok is None:
|
| 61 |
+
try:
|
| 62 |
+
from transformers import AutoTokenizer
|
| 63 |
+
|
| 64 |
+
name_or_path = getattr(model.config, "_name_or_path", "") or ""
|
| 65 |
+
if name_or_path:
|
| 66 |
+
tok = AutoTokenizer.from_pretrained(name_or_path)
|
| 67 |
+
model._hybrid_trailer_tokenizer = tok
|
| 68 |
+
except Exception:
|
| 69 |
+
return None
|
| 70 |
+
# Unwrap processors down to the underlying tokenizer.
|
| 71 |
+
inner = getattr(tok, "tokenizer", None)
|
| 72 |
+
if inner is not None and hasattr(inner, "encode"):
|
| 73 |
+
tok = inner
|
| 74 |
+
return tok
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def _encode_trailer(tokenizer, text):
|
| 78 |
+
"""Encode the trailer to token ids. Returns [] when encoding fails."""
|
| 79 |
+
if tokenizer is None:
|
| 80 |
+
return []
|
| 81 |
+
try:
|
| 82 |
+
ids = tokenizer(text, add_special_tokens=False)["input_ids"]
|
| 83 |
+
except Exception:
|
| 84 |
+
try:
|
| 85 |
+
ids = tokenizer.encode(text, add_special_tokens=False)
|
| 86 |
+
except Exception:
|
| 87 |
+
return []
|
| 88 |
+
if ids and isinstance(ids[0], list): # batched encoding shape
|
| 89 |
+
ids = ids[0]
|
| 90 |
+
return [int(i) for i in ids]
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def generate(*args, model=None, **kwargs):
|
| 94 |
+
"""Stock generation + handoff-probe confidence with an in-band trailer."""
|
| 95 |
+
if model is None:
|
| 96 |
+
raise TypeError("custom generate requires the `model` keyword argument")
|
| 97 |
+
if args:
|
| 98 |
+
if len(args) > 1:
|
| 99 |
+
raise TypeError(
|
| 100 |
+
"gemma-4-e2b-it-hybrid generate accepts at most one positional argument (`inputs`); "
|
| 101 |
+
"pass everything else as keyword arguments"
|
| 102 |
+
)
|
| 103 |
+
if kwargs.get("inputs") is not None:
|
| 104 |
+
raise TypeError("got `inputs` both positionally and as a keyword argument")
|
| 105 |
+
kwargs["inputs"] = args[0]
|
| 106 |
+
|
| 107 |
+
return_confidence = bool(kwargs.pop("return_confidence", False))
|
| 108 |
+
emit_trailer = bool(kwargs.pop("emit_trailer", True))
|
| 109 |
+
streamer = kwargs.pop("streamer", None)
|
| 110 |
+
|
| 111 |
+
# The class attribute is the stock GenerationMixin.generate; the instance
|
| 112 |
+
# attribute is this function (installed by `from_pretrained`), so calling
|
| 113 |
+
# through the class avoids recursion.
|
| 114 |
+
stock_generate = type(model).generate
|
| 115 |
+
|
| 116 |
+
scorable = kwargs.get("assistant_model") is None and hasattr(model, "probe_capture")
|
| 117 |
+
wrapped_streamer = _DeferredEndStreamer(streamer) if streamer is not None else None
|
| 118 |
+
|
| 119 |
+
try:
|
| 120 |
+
if scorable:
|
| 121 |
+
with model.probe_capture() as rows:
|
| 122 |
+
output = stock_generate(model, streamer=wrapped_streamer, **kwargs)
|
| 123 |
+
confidence = model.confidence_from_rows(rows)
|
| 124 |
+
else:
|
| 125 |
+
output = stock_generate(model, streamer=wrapped_streamer, **kwargs)
|
| 126 |
+
confidence = None
|
| 127 |
+
except BaseException:
|
| 128 |
+
if streamer is not None:
|
| 129 |
+
try:
|
| 130 |
+
streamer.end()
|
| 131 |
+
except Exception:
|
| 132 |
+
pass
|
| 133 |
+
raise
|
| 134 |
+
|
| 135 |
+
sequences = output if isinstance(output, torch.Tensor) else getattr(output, "sequences", None)
|
| 136 |
+
|
| 137 |
+
if return_confidence:
|
| 138 |
+
if streamer is not None:
|
| 139 |
+
streamer.end()
|
| 140 |
+
trailer_text = (
|
| 141 |
+
TRAILER_TEMPLATE.format(confidence=confidence) if confidence is not None else None
|
| 142 |
+
)
|
| 143 |
+
return GenerateWithConfidenceOutput(
|
| 144 |
+
sequences=sequences,
|
| 145 |
+
confidence=confidence,
|
| 146 |
+
trailer_text=trailer_text,
|
| 147 |
+
generate_output=output,
|
| 148 |
+
)
|
| 149 |
+
|
| 150 |
+
trailer_ids: list = []
|
| 151 |
+
if (
|
| 152 |
+
emit_trailer
|
| 153 |
+
and confidence is not None
|
| 154 |
+
and sequences is not None
|
| 155 |
+
and sequences.ndim == 2
|
| 156 |
+
and sequences.shape[0] == 1
|
| 157 |
+
):
|
| 158 |
+
trailer_text = TRAILER_TEMPLATE.format(confidence=confidence)
|
| 159 |
+
trailer_ids = _encode_trailer(_resolve_tokenizer(model, kwargs), trailer_text)
|
| 160 |
+
|
| 161 |
+
if trailer_ids:
|
| 162 |
+
trailer_tensor = torch.tensor(
|
| 163 |
+
[trailer_ids], dtype=sequences.dtype, device=sequences.device
|
| 164 |
+
)
|
| 165 |
+
new_sequences = torch.cat([sequences, trailer_tensor], dim=1)
|
| 166 |
+
if isinstance(output, torch.Tensor):
|
| 167 |
+
output = new_sequences
|
| 168 |
+
else:
|
| 169 |
+
try:
|
| 170 |
+
output.sequences = new_sequences
|
| 171 |
+
except Exception:
|
| 172 |
+
pass
|
| 173 |
+
if streamer is not None:
|
| 174 |
+
streamer.put(trailer_tensor[0].cpu())
|
| 175 |
+
|
| 176 |
+
if streamer is not None:
|
| 177 |
+
streamer.end()
|
| 178 |
+
return output
|
custom_generate/requirements.txt
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
transformers>=5,<6
|
gemma_4_e2b_it_hybrid.py
ADDED
|
@@ -0,0 +1,219 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Cactus-Compute/gemma-4-e2b-it-hybrid — single-file mlx-lm model (``model_file`` mechanism).
|
| 2 |
+
|
| 3 |
+
Drop this file into the converted MLX repo and add ``"model_file":
|
| 4 |
+
"gemma_4_e2b_it_hybrid.py"`` to its ``config.json`` (mlx-lm >= 0.30.1, the mechanism
|
| 5 |
+
from mlx-lm PR #830). ``mlx_lm.load`` then builds ``Model``/``ModelArgs`` from
|
| 6 |
+
this file instead of the built-in ``gemma4_text`` classes.
|
| 7 |
+
|
| 8 |
+
The model is the stock ``mlx_lm.models.gemma4_text`` Gemma-4 text model plus
|
| 9 |
+
the handoff probe (weight prefix ``handoff_probe.*``). During generation the
|
| 10 |
+
output of decoder layer ``probe_layer`` is captured at the position that
|
| 11 |
+
predicts each generated token; after generation
|
| 12 |
+
|
| 13 |
+
``model.last_confidence`` -> float in [0, 1] or None
|
| 14 |
+
|
| 15 |
+
exposes ``confidence = 1 - p_wrong``. All probe math runs in float32.
|
| 16 |
+
|
| 17 |
+
Capture semantics (tuned to ``mlx_lm.generate_step``):
|
| 18 |
+
|
| 19 |
+
- multi-token (prefill-chunk) forwards reset the capture buffer and are never
|
| 20 |
+
kept; ``generate_step`` always feeds the final prompt token through a
|
| 21 |
+
1-token step, and that forward's last position is row 0 of the contract;
|
| 22 |
+
- ``generate_step`` pipelines one forward ahead, so the buffer ends with one
|
| 23 |
+
lookahead row for a token that is never emitted; ``last_confidence`` drops
|
| 24 |
+
that final row. Use ``model.confidence(num_tokens=N)`` to score exactly the
|
| 25 |
+
first N generated tokens instead.
|
| 26 |
+
|
| 27 |
+
Limitations:
|
| 28 |
+
|
| 29 |
+
- Server-side surfacing is Python-API only for MLX today: an mlx-lm ``Model``
|
| 30 |
+
cannot inject tokens into the stream, so there is no in-band
|
| 31 |
+
``[[hybrid:confidence=...]]`` trailer here (unlike the transformers repo).
|
| 32 |
+
- Speculative decoding is not supported (draft verification feeds multiple
|
| 33 |
+
tokens per forward, which resets the buffer).
|
| 34 |
+
- Prompts of exactly 2 tokens leave one extra leading row in the buffer (the
|
| 35 |
+
1-token prefill chunk is indistinguishable from a decode step). Chat-template
|
| 36 |
+
prompts are always far longer.
|
| 37 |
+
"""
|
| 38 |
+
|
| 39 |
+
import math
|
| 40 |
+
from dataclasses import dataclass
|
| 41 |
+
from typing import Any, List, Optional
|
| 42 |
+
|
| 43 |
+
import mlx.core as mx
|
| 44 |
+
import mlx.nn as nn
|
| 45 |
+
|
| 46 |
+
from mlx_lm.models import gemma4_text
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
@dataclass
|
| 50 |
+
class ModelArgs(gemma4_text.ModelArgs):
|
| 51 |
+
model_type: str = "gemma-4-e2b-it-hybrid"
|
| 52 |
+
probe_layer: int = 28
|
| 53 |
+
probe_feature_size: int = 1536
|
| 54 |
+
probe_max_tokens: int = 1024
|
| 55 |
+
probe_proj_dim: int = 32
|
| 56 |
+
|
| 57 |
+
def __post_init__(self):
|
| 58 |
+
super().__post_init__()
|
| 59 |
+
if not 0 <= self.probe_layer < self.num_hidden_layers:
|
| 60 |
+
raise ValueError(
|
| 61 |
+
f"probe_layer={self.probe_layer} must be in "
|
| 62 |
+
f"[0, num_hidden_layers={self.num_hidden_layers})"
|
| 63 |
+
)
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
# Hidden widths of the probe MLP head, fixed by the released checkpoint.
|
| 67 |
+
PROBE_HEAD_DIMS = (128, 64)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
class HandoffProbe(nn.Module):
|
| 71 |
+
"""The released handoff probe; parameter tree matches the checkpoint keys
|
| 72 |
+
``norm.{weight,bias}``, ``proj.{weight,bias}``, ``attn_query``,
|
| 73 |
+
``head.{0,2,4}.{weight,bias}``."""
|
| 74 |
+
|
| 75 |
+
def __init__(self, feature_size: int, proj_dim: int = 32):
|
| 76 |
+
super().__init__()
|
| 77 |
+
h1, h2 = PROBE_HEAD_DIMS
|
| 78 |
+
self.norm = nn.LayerNorm(feature_size, eps=1e-5)
|
| 79 |
+
self.proj = nn.Linear(feature_size, proj_dim)
|
| 80 |
+
self.attn_query = mx.zeros((proj_dim,))
|
| 81 |
+
self.head = [
|
| 82 |
+
nn.Linear(proj_dim, h1),
|
| 83 |
+
nn.ReLU(),
|
| 84 |
+
nn.Linear(h1, h2),
|
| 85 |
+
nn.ReLU(),
|
| 86 |
+
nn.Linear(h2, 1),
|
| 87 |
+
]
|
| 88 |
+
|
| 89 |
+
def p_wrong(self, hidden_states: mx.array, max_tokens: int = 1024) -> mx.array:
|
| 90 |
+
"""Score ``[T, feature_size]`` rows; float32 math per the contract."""
|
| 91 |
+
x = hidden_states[:max_tokens].astype(mx.float32)
|
| 92 |
+
x = mx.fast.layer_norm(
|
| 93 |
+
x,
|
| 94 |
+
self.norm.weight.astype(mx.float32),
|
| 95 |
+
self.norm.bias.astype(mx.float32),
|
| 96 |
+
1e-5,
|
| 97 |
+
)
|
| 98 |
+
projected = mx.maximum(
|
| 99 |
+
x @ self.proj.weight.astype(mx.float32).T + self.proj.bias.astype(mx.float32),
|
| 100 |
+
0.0,
|
| 101 |
+
)
|
| 102 |
+
scores = projected @ self.attn_query.astype(mx.float32)
|
| 103 |
+
scores = scores / math.sqrt(projected.shape[-1])
|
| 104 |
+
weights = mx.softmax(scores - scores.max(), axis=-1)
|
| 105 |
+
pooled = weights @ projected
|
| 106 |
+
|
| 107 |
+
h0, h2_, h4 = self.head[0], self.head[2], self.head[4]
|
| 108 |
+
h = mx.maximum(
|
| 109 |
+
pooled @ h0.weight.astype(mx.float32).T + h0.bias.astype(mx.float32), 0.0
|
| 110 |
+
)
|
| 111 |
+
h = mx.maximum(
|
| 112 |
+
h @ h2_.weight.astype(mx.float32).T + h2_.bias.astype(mx.float32), 0.0
|
| 113 |
+
)
|
| 114 |
+
logit = h @ h4.weight.astype(mx.float32).T + h4.bias.astype(mx.float32)
|
| 115 |
+
return mx.sigmoid(logit)[0]
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
class _ProbeState:
|
| 119 |
+
"""Plain-object row buffer, invisible to the MLX module tree."""
|
| 120 |
+
|
| 121 |
+
def __init__(self):
|
| 122 |
+
self.rows: List[mx.array] = []
|
| 123 |
+
|
| 124 |
+
def reset(self):
|
| 125 |
+
self.rows.clear()
|
| 126 |
+
|
| 127 |
+
|
| 128 |
+
class _CaptureDecoderLayer(gemma4_text.DecoderLayer):
|
| 129 |
+
"""Stock decoder layer that reports its output to a capture sink.
|
| 130 |
+
|
| 131 |
+
Same submodule tree as ``DecoderLayer``, so weight keys are unchanged.
|
| 132 |
+
"""
|
| 133 |
+
|
| 134 |
+
def __call__(self, x, *args, **kwargs):
|
| 135 |
+
out = super().__call__(x, *args, **kwargs)
|
| 136 |
+
sink = getattr(self, "capture_sink", None)
|
| 137 |
+
if sink is not None:
|
| 138 |
+
sink(out[0])
|
| 139 |
+
return out
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
class Model(gemma4_text.Model):
|
| 143 |
+
def __init__(self, args: ModelArgs):
|
| 144 |
+
super().__init__(args)
|
| 145 |
+
self.args = args
|
| 146 |
+
self.handoff_probe = HandoffProbe(args.probe_feature_size, args.probe_proj_dim)
|
| 147 |
+
self._probe_state = _ProbeState()
|
| 148 |
+
|
| 149 |
+
capture = _CaptureDecoderLayer(args, layer_idx=args.probe_layer)
|
| 150 |
+
capture.capture_sink = self._capture_row
|
| 151 |
+
self.model.layers[args.probe_layer] = capture
|
| 152 |
+
|
| 153 |
+
# --- capture -----------------------------------------------------------
|
| 154 |
+
|
| 155 |
+
def _capture_row(self, hidden: mx.array) -> None:
|
| 156 |
+
"""Keep the last position of each 1-token probe-layer forward.
|
| 157 |
+
|
| 158 |
+
``generate_step``'s prefill loop always leaves the final prompt token
|
| 159 |
+
to a 1-token ``_step`` call, so multi-token (prefill-chunk) forwards
|
| 160 |
+
never produce a contract row — they just reset the buffer. Row 0 is
|
| 161 |
+
the step that consumes the last prompt token (= last prompt position).
|
| 162 |
+
"""
|
| 163 |
+
rows = self._probe_state.rows
|
| 164 |
+
if hidden.shape[1] > 1:
|
| 165 |
+
rows.clear()
|
| 166 |
+
return
|
| 167 |
+
# Keep one row beyond the window so dropping the lookahead row cannot
|
| 168 |
+
# lose a real one.
|
| 169 |
+
if len(rows) <= self.args.probe_max_tokens:
|
| 170 |
+
rows.append(hidden[:, -1, :])
|
| 171 |
+
|
| 172 |
+
def __call__(
|
| 173 |
+
self,
|
| 174 |
+
inputs: mx.array,
|
| 175 |
+
cache=None,
|
| 176 |
+
input_embeddings: Optional[mx.array] = None,
|
| 177 |
+
per_layer_inputs: Optional[mx.array] = None,
|
| 178 |
+
):
|
| 179 |
+
# A fresh cache (or none) means a new generation: reset the buffer.
|
| 180 |
+
if cache is None or (len(cache) > 0 and getattr(cache[0], "offset", 0) == 0):
|
| 181 |
+
self._probe_state.reset()
|
| 182 |
+
return super().__call__(
|
| 183 |
+
inputs,
|
| 184 |
+
cache=cache,
|
| 185 |
+
input_embeddings=input_embeddings,
|
| 186 |
+
per_layer_inputs=per_layer_inputs,
|
| 187 |
+
)
|
| 188 |
+
|
| 189 |
+
# --- scoring -----------------------------------------------------------
|
| 190 |
+
|
| 191 |
+
def reset_probe(self) -> None:
|
| 192 |
+
"""Clear the capture buffer (e.g. between manual forward calls)."""
|
| 193 |
+
self._probe_state.reset()
|
| 194 |
+
|
| 195 |
+
def confidence(self, num_tokens: Optional[int] = None) -> Optional[float]:
|
| 196 |
+
"""``1 - p_wrong`` for the captured generation.
|
| 197 |
+
|
| 198 |
+
``num_tokens`` scores exactly the first N captured rows. Without it,
|
| 199 |
+
the final row is dropped to discard ``generate_step``'s one-step
|
| 200 |
+
lookahead forward. Returns ``None`` when nothing (scoreable) was
|
| 201 |
+
captured.
|
| 202 |
+
"""
|
| 203 |
+
rows = list(self._probe_state.rows)
|
| 204 |
+
if not rows or rows[0].shape[0] != 1:
|
| 205 |
+
return None
|
| 206 |
+
if num_tokens is not None:
|
| 207 |
+
rows = rows[:num_tokens]
|
| 208 |
+
elif len(rows) >= 2:
|
| 209 |
+
rows = rows[:-1]
|
| 210 |
+
if not rows:
|
| 211 |
+
return None
|
| 212 |
+
stacked = mx.concatenate(rows, axis=0) # [T, features]
|
| 213 |
+
p_wrong = self.handoff_probe.p_wrong(stacked, self.args.probe_max_tokens)
|
| 214 |
+
return float(1.0 - p_wrong.item())
|
| 215 |
+
|
| 216 |
+
@property
|
| 217 |
+
def last_confidence(self) -> Optional[float]:
|
| 218 |
+
"""Confidence of the most recent generation (None before any)."""
|
| 219 |
+
return self.confidence()
|
generation_config.json
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"bos_token_id": 2,
|
| 3 |
+
"do_sample": true,
|
| 4 |
+
"eos_token_id": [
|
| 5 |
+
1,
|
| 6 |
+
106,
|
| 7 |
+
50
|
| 8 |
+
],
|
| 9 |
+
"pad_token_id": 0,
|
| 10 |
+
"temperature": 1.0,
|
| 11 |
+
"top_k": 64,
|
| 12 |
+
"top_p": 0.95,
|
| 13 |
+
"transformers_version": "5.5.0.dev0"
|
| 14 |
+
}
|
model-00001-of-00003.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:3da4a9ebb8bada980bdfbc213ee28464d6ce39647f3d6dfcb2833b76e592b8d8
|
| 3 |
+
size 805306504
|
model-00002-of-00003.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:794f283c42b320dee5824838f9d33ebcdd07cc2ce2cf23fca279dd0de362a824
|
| 3 |
+
size 4697620632
|
model-00003-of-00003.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:6b979b9af49171e398743e21a205e14ae562aca42a767c1daeba12d6599a77fa
|
| 3 |
+
size 3792302306
|
model.safetensors.index.json
ADDED
|
@@ -0,0 +1,618 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"metadata": {
|
| 3 |
+
"total_size": 9295159114
|
| 4 |
+
},
|
| 5 |
+
"weight_map": {
|
| 6 |
+
"handoff_probe.attn_query": "model-00003-of-00003.safetensors",
|
| 7 |
+
"handoff_probe.head.0.bias": "model-00003-of-00003.safetensors",
|
| 8 |
+
"handoff_probe.head.0.weight": "model-00003-of-00003.safetensors",
|
| 9 |
+
"handoff_probe.head.2.bias": "model-00003-of-00003.safetensors",
|
| 10 |
+
"handoff_probe.head.2.weight": "model-00003-of-00003.safetensors",
|
| 11 |
+
"handoff_probe.head.4.bias": "model-00003-of-00003.safetensors",
|
| 12 |
+
"handoff_probe.head.4.weight": "model-00003-of-00003.safetensors",
|
| 13 |
+
"handoff_probe.norm.bias": "model-00003-of-00003.safetensors",
|
| 14 |
+
"handoff_probe.norm.weight": "model-00003-of-00003.safetensors",
|
| 15 |
+
"handoff_probe.proj.bias": "model-00003-of-00003.safetensors",
|
| 16 |
+
"handoff_probe.proj.weight": "model-00003-of-00003.safetensors",
|
| 17 |
+
"model.embed_tokens.weight": "model-00001-of-00003.safetensors",
|
| 18 |
+
"model.embed_tokens_per_layer.weight": "model-00002-of-00003.safetensors",
|
| 19 |
+
"model.layers.0.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 20 |
+
"model.layers.0.layer_scalar": "model-00003-of-00003.safetensors",
|
| 21 |
+
"model.layers.0.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 22 |
+
"model.layers.0.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 23 |
+
"model.layers.0.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 24 |
+
"model.layers.0.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 25 |
+
"model.layers.0.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 26 |
+
"model.layers.0.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 27 |
+
"model.layers.0.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 28 |
+
"model.layers.0.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 29 |
+
"model.layers.0.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 30 |
+
"model.layers.0.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 31 |
+
"model.layers.0.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 32 |
+
"model.layers.0.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 33 |
+
"model.layers.0.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 34 |
+
"model.layers.0.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 35 |
+
"model.layers.0.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 36 |
+
"model.layers.1.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 37 |
+
"model.layers.1.layer_scalar": "model-00003-of-00003.safetensors",
|
| 38 |
+
"model.layers.1.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 39 |
+
"model.layers.1.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 40 |
+
"model.layers.1.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 41 |
+
"model.layers.1.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 42 |
+
"model.layers.1.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 43 |
+
"model.layers.1.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 44 |
+
"model.layers.1.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 45 |
+
"model.layers.1.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 46 |
+
"model.layers.1.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 47 |
+
"model.layers.1.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 48 |
+
"model.layers.1.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 49 |
+
"model.layers.1.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 50 |
+
"model.layers.1.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 51 |
+
"model.layers.1.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 52 |
+
"model.layers.1.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 53 |
+
"model.layers.10.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 54 |
+
"model.layers.10.layer_scalar": "model-00003-of-00003.safetensors",
|
| 55 |
+
"model.layers.10.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 56 |
+
"model.layers.10.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 57 |
+
"model.layers.10.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 58 |
+
"model.layers.10.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 59 |
+
"model.layers.10.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 60 |
+
"model.layers.10.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 61 |
+
"model.layers.10.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 62 |
+
"model.layers.10.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 63 |
+
"model.layers.10.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 64 |
+
"model.layers.10.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 65 |
+
"model.layers.10.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 66 |
+
"model.layers.10.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 67 |
+
"model.layers.10.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 68 |
+
"model.layers.10.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 69 |
+
"model.layers.10.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 70 |
+
"model.layers.11.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 71 |
+
"model.layers.11.layer_scalar": "model-00003-of-00003.safetensors",
|
| 72 |
+
"model.layers.11.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 73 |
+
"model.layers.11.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 74 |
+
"model.layers.11.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 75 |
+
"model.layers.11.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 76 |
+
"model.layers.11.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 77 |
+
"model.layers.11.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 78 |
+
"model.layers.11.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 79 |
+
"model.layers.11.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 80 |
+
"model.layers.11.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 81 |
+
"model.layers.11.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 82 |
+
"model.layers.11.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 83 |
+
"model.layers.11.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 84 |
+
"model.layers.11.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 85 |
+
"model.layers.11.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 86 |
+
"model.layers.11.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 87 |
+
"model.layers.12.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 88 |
+
"model.layers.12.layer_scalar": "model-00003-of-00003.safetensors",
|
| 89 |
+
"model.layers.12.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 90 |
+
"model.layers.12.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 91 |
+
"model.layers.12.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 92 |
+
"model.layers.12.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 93 |
+
"model.layers.12.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 94 |
+
"model.layers.12.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 95 |
+
"model.layers.12.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 96 |
+
"model.layers.12.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 97 |
+
"model.layers.12.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 98 |
+
"model.layers.12.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 99 |
+
"model.layers.12.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 100 |
+
"model.layers.12.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 101 |
+
"model.layers.12.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 102 |
+
"model.layers.12.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 103 |
+
"model.layers.12.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 104 |
+
"model.layers.13.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 105 |
+
"model.layers.13.layer_scalar": "model-00003-of-00003.safetensors",
|
| 106 |
+
"model.layers.13.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 107 |
+
"model.layers.13.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 108 |
+
"model.layers.13.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 109 |
+
"model.layers.13.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 110 |
+
"model.layers.13.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 111 |
+
"model.layers.13.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 112 |
+
"model.layers.13.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 113 |
+
"model.layers.13.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 114 |
+
"model.layers.13.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 115 |
+
"model.layers.13.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 116 |
+
"model.layers.13.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 117 |
+
"model.layers.13.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 118 |
+
"model.layers.13.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 119 |
+
"model.layers.13.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 120 |
+
"model.layers.13.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 121 |
+
"model.layers.14.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 122 |
+
"model.layers.14.layer_scalar": "model-00003-of-00003.safetensors",
|
| 123 |
+
"model.layers.14.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 124 |
+
"model.layers.14.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 125 |
+
"model.layers.14.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 126 |
+
"model.layers.14.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 127 |
+
"model.layers.14.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 128 |
+
"model.layers.14.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 129 |
+
"model.layers.14.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 130 |
+
"model.layers.14.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 131 |
+
"model.layers.14.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 132 |
+
"model.layers.14.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 133 |
+
"model.layers.14.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 134 |
+
"model.layers.14.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 135 |
+
"model.layers.14.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 136 |
+
"model.layers.14.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 137 |
+
"model.layers.14.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 138 |
+
"model.layers.15.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 139 |
+
"model.layers.15.layer_scalar": "model-00003-of-00003.safetensors",
|
| 140 |
+
"model.layers.15.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 141 |
+
"model.layers.15.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 142 |
+
"model.layers.15.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 143 |
+
"model.layers.15.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 144 |
+
"model.layers.15.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 145 |
+
"model.layers.15.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 146 |
+
"model.layers.15.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 147 |
+
"model.layers.15.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 148 |
+
"model.layers.15.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 149 |
+
"model.layers.15.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 150 |
+
"model.layers.15.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 151 |
+
"model.layers.15.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 152 |
+
"model.layers.15.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 153 |
+
"model.layers.15.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 154 |
+
"model.layers.15.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 155 |
+
"model.layers.16.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 156 |
+
"model.layers.16.layer_scalar": "model-00003-of-00003.safetensors",
|
| 157 |
+
"model.layers.16.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 158 |
+
"model.layers.16.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 159 |
+
"model.layers.16.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 160 |
+
"model.layers.16.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 161 |
+
"model.layers.16.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 162 |
+
"model.layers.16.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 163 |
+
"model.layers.16.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 164 |
+
"model.layers.16.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 165 |
+
"model.layers.16.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 166 |
+
"model.layers.16.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 167 |
+
"model.layers.16.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 168 |
+
"model.layers.16.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 169 |
+
"model.layers.16.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 170 |
+
"model.layers.16.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 171 |
+
"model.layers.16.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 172 |
+
"model.layers.17.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 173 |
+
"model.layers.17.layer_scalar": "model-00003-of-00003.safetensors",
|
| 174 |
+
"model.layers.17.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 175 |
+
"model.layers.17.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 176 |
+
"model.layers.17.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 177 |
+
"model.layers.17.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 178 |
+
"model.layers.17.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 179 |
+
"model.layers.17.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 180 |
+
"model.layers.17.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 181 |
+
"model.layers.17.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 182 |
+
"model.layers.17.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 183 |
+
"model.layers.17.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 184 |
+
"model.layers.17.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 185 |
+
"model.layers.17.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 186 |
+
"model.layers.17.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 187 |
+
"model.layers.17.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 188 |
+
"model.layers.17.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 189 |
+
"model.layers.18.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 190 |
+
"model.layers.18.layer_scalar": "model-00003-of-00003.safetensors",
|
| 191 |
+
"model.layers.18.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 192 |
+
"model.layers.18.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 193 |
+
"model.layers.18.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 194 |
+
"model.layers.18.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 195 |
+
"model.layers.18.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 196 |
+
"model.layers.18.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 197 |
+
"model.layers.18.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 198 |
+
"model.layers.18.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 199 |
+
"model.layers.18.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 200 |
+
"model.layers.18.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 201 |
+
"model.layers.18.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 202 |
+
"model.layers.18.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 203 |
+
"model.layers.18.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 204 |
+
"model.layers.18.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 205 |
+
"model.layers.18.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 206 |
+
"model.layers.19.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 207 |
+
"model.layers.19.layer_scalar": "model-00003-of-00003.safetensors",
|
| 208 |
+
"model.layers.19.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 209 |
+
"model.layers.19.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 210 |
+
"model.layers.19.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 211 |
+
"model.layers.19.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 212 |
+
"model.layers.19.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 213 |
+
"model.layers.19.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 214 |
+
"model.layers.19.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 215 |
+
"model.layers.19.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 216 |
+
"model.layers.19.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 217 |
+
"model.layers.19.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 218 |
+
"model.layers.19.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 219 |
+
"model.layers.19.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 220 |
+
"model.layers.19.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 221 |
+
"model.layers.19.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 222 |
+
"model.layers.19.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 223 |
+
"model.layers.2.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 224 |
+
"model.layers.2.layer_scalar": "model-00003-of-00003.safetensors",
|
| 225 |
+
"model.layers.2.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 226 |
+
"model.layers.2.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 227 |
+
"model.layers.2.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 228 |
+
"model.layers.2.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 229 |
+
"model.layers.2.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 230 |
+
"model.layers.2.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 231 |
+
"model.layers.2.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 232 |
+
"model.layers.2.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 233 |
+
"model.layers.2.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 234 |
+
"model.layers.2.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 235 |
+
"model.layers.2.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 236 |
+
"model.layers.2.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 237 |
+
"model.layers.2.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 238 |
+
"model.layers.2.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 239 |
+
"model.layers.2.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 240 |
+
"model.layers.20.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 241 |
+
"model.layers.20.layer_scalar": "model-00003-of-00003.safetensors",
|
| 242 |
+
"model.layers.20.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 243 |
+
"model.layers.20.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 244 |
+
"model.layers.20.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 245 |
+
"model.layers.20.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 246 |
+
"model.layers.20.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 247 |
+
"model.layers.20.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 248 |
+
"model.layers.20.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 249 |
+
"model.layers.20.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 250 |
+
"model.layers.20.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 251 |
+
"model.layers.20.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 252 |
+
"model.layers.20.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 253 |
+
"model.layers.20.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 254 |
+
"model.layers.20.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 255 |
+
"model.layers.20.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 256 |
+
"model.layers.20.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 257 |
+
"model.layers.21.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 258 |
+
"model.layers.21.layer_scalar": "model-00003-of-00003.safetensors",
|
| 259 |
+
"model.layers.21.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 260 |
+
"model.layers.21.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 261 |
+
"model.layers.21.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 262 |
+
"model.layers.21.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 263 |
+
"model.layers.21.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 264 |
+
"model.layers.21.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 265 |
+
"model.layers.21.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 266 |
+
"model.layers.21.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 267 |
+
"model.layers.21.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 268 |
+
"model.layers.21.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 269 |
+
"model.layers.21.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 270 |
+
"model.layers.21.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 271 |
+
"model.layers.21.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 272 |
+
"model.layers.21.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 273 |
+
"model.layers.21.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 274 |
+
"model.layers.22.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 275 |
+
"model.layers.22.layer_scalar": "model-00003-of-00003.safetensors",
|
| 276 |
+
"model.layers.22.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 277 |
+
"model.layers.22.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 278 |
+
"model.layers.22.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 279 |
+
"model.layers.22.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 280 |
+
"model.layers.22.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 281 |
+
"model.layers.22.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 282 |
+
"model.layers.22.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 283 |
+
"model.layers.22.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 284 |
+
"model.layers.22.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 285 |
+
"model.layers.22.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 286 |
+
"model.layers.22.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 287 |
+
"model.layers.22.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 288 |
+
"model.layers.22.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 289 |
+
"model.layers.22.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 290 |
+
"model.layers.22.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 291 |
+
"model.layers.23.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 292 |
+
"model.layers.23.layer_scalar": "model-00003-of-00003.safetensors",
|
| 293 |
+
"model.layers.23.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 294 |
+
"model.layers.23.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 295 |
+
"model.layers.23.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 296 |
+
"model.layers.23.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 297 |
+
"model.layers.23.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 298 |
+
"model.layers.23.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 299 |
+
"model.layers.23.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 300 |
+
"model.layers.23.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 301 |
+
"model.layers.23.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 302 |
+
"model.layers.23.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 303 |
+
"model.layers.23.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 304 |
+
"model.layers.23.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 305 |
+
"model.layers.23.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 306 |
+
"model.layers.23.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 307 |
+
"model.layers.23.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 308 |
+
"model.layers.24.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 309 |
+
"model.layers.24.layer_scalar": "model-00003-of-00003.safetensors",
|
| 310 |
+
"model.layers.24.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 311 |
+
"model.layers.24.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 312 |
+
"model.layers.24.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 313 |
+
"model.layers.24.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 314 |
+
"model.layers.24.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 315 |
+
"model.layers.24.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 316 |
+
"model.layers.24.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 317 |
+
"model.layers.24.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 318 |
+
"model.layers.24.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 319 |
+
"model.layers.24.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 320 |
+
"model.layers.24.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 321 |
+
"model.layers.24.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 322 |
+
"model.layers.24.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 323 |
+
"model.layers.24.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 324 |
+
"model.layers.24.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 325 |
+
"model.layers.25.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 326 |
+
"model.layers.25.layer_scalar": "model-00003-of-00003.safetensors",
|
| 327 |
+
"model.layers.25.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 328 |
+
"model.layers.25.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 329 |
+
"model.layers.25.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 330 |
+
"model.layers.25.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 331 |
+
"model.layers.25.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 332 |
+
"model.layers.25.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 333 |
+
"model.layers.25.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 334 |
+
"model.layers.25.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 335 |
+
"model.layers.25.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 336 |
+
"model.layers.25.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 337 |
+
"model.layers.25.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 338 |
+
"model.layers.25.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 339 |
+
"model.layers.25.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 340 |
+
"model.layers.25.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 341 |
+
"model.layers.25.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 342 |
+
"model.layers.26.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 343 |
+
"model.layers.26.layer_scalar": "model-00003-of-00003.safetensors",
|
| 344 |
+
"model.layers.26.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 345 |
+
"model.layers.26.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 346 |
+
"model.layers.26.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 347 |
+
"model.layers.26.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 348 |
+
"model.layers.26.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 349 |
+
"model.layers.26.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 350 |
+
"model.layers.26.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 351 |
+
"model.layers.26.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 352 |
+
"model.layers.26.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 353 |
+
"model.layers.26.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 354 |
+
"model.layers.26.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 355 |
+
"model.layers.26.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 356 |
+
"model.layers.26.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 357 |
+
"model.layers.26.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 358 |
+
"model.layers.26.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 359 |
+
"model.layers.27.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 360 |
+
"model.layers.27.layer_scalar": "model-00003-of-00003.safetensors",
|
| 361 |
+
"model.layers.27.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 362 |
+
"model.layers.27.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 363 |
+
"model.layers.27.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 364 |
+
"model.layers.27.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 365 |
+
"model.layers.27.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 366 |
+
"model.layers.27.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 367 |
+
"model.layers.27.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 368 |
+
"model.layers.27.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 369 |
+
"model.layers.27.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 370 |
+
"model.layers.27.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 371 |
+
"model.layers.27.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 372 |
+
"model.layers.27.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 373 |
+
"model.layers.27.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 374 |
+
"model.layers.27.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 375 |
+
"model.layers.27.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 376 |
+
"model.layers.28.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 377 |
+
"model.layers.28.layer_scalar": "model-00003-of-00003.safetensors",
|
| 378 |
+
"model.layers.28.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 379 |
+
"model.layers.28.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 380 |
+
"model.layers.28.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 381 |
+
"model.layers.28.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 382 |
+
"model.layers.28.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 383 |
+
"model.layers.28.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 384 |
+
"model.layers.28.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 385 |
+
"model.layers.28.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 386 |
+
"model.layers.28.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 387 |
+
"model.layers.28.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 388 |
+
"model.layers.28.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 389 |
+
"model.layers.28.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 390 |
+
"model.layers.28.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 391 |
+
"model.layers.28.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 392 |
+
"model.layers.28.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 393 |
+
"model.layers.29.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 394 |
+
"model.layers.29.layer_scalar": "model-00003-of-00003.safetensors",
|
| 395 |
+
"model.layers.29.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 396 |
+
"model.layers.29.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 397 |
+
"model.layers.29.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 398 |
+
"model.layers.29.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 399 |
+
"model.layers.29.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 400 |
+
"model.layers.29.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 401 |
+
"model.layers.29.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 402 |
+
"model.layers.29.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 403 |
+
"model.layers.29.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 404 |
+
"model.layers.29.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 405 |
+
"model.layers.29.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 406 |
+
"model.layers.29.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 407 |
+
"model.layers.29.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 408 |
+
"model.layers.29.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 409 |
+
"model.layers.29.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 410 |
+
"model.layers.3.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 411 |
+
"model.layers.3.layer_scalar": "model-00003-of-00003.safetensors",
|
| 412 |
+
"model.layers.3.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 413 |
+
"model.layers.3.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 414 |
+
"model.layers.3.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 415 |
+
"model.layers.3.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 416 |
+
"model.layers.3.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 417 |
+
"model.layers.3.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 418 |
+
"model.layers.3.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 419 |
+
"model.layers.3.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 420 |
+
"model.layers.3.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 421 |
+
"model.layers.3.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 422 |
+
"model.layers.3.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 423 |
+
"model.layers.3.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 424 |
+
"model.layers.3.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 425 |
+
"model.layers.3.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 426 |
+
"model.layers.3.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 427 |
+
"model.layers.30.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 428 |
+
"model.layers.30.layer_scalar": "model-00003-of-00003.safetensors",
|
| 429 |
+
"model.layers.30.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 430 |
+
"model.layers.30.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 431 |
+
"model.layers.30.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 432 |
+
"model.layers.30.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 433 |
+
"model.layers.30.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 434 |
+
"model.layers.30.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 435 |
+
"model.layers.30.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 436 |
+
"model.layers.30.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 437 |
+
"model.layers.30.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 438 |
+
"model.layers.30.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 439 |
+
"model.layers.30.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 440 |
+
"model.layers.30.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 441 |
+
"model.layers.30.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 442 |
+
"model.layers.30.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 443 |
+
"model.layers.30.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 444 |
+
"model.layers.31.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 445 |
+
"model.layers.31.layer_scalar": "model-00003-of-00003.safetensors",
|
| 446 |
+
"model.layers.31.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 447 |
+
"model.layers.31.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 448 |
+
"model.layers.31.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 449 |
+
"model.layers.31.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 450 |
+
"model.layers.31.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 451 |
+
"model.layers.31.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 452 |
+
"model.layers.31.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 453 |
+
"model.layers.31.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 454 |
+
"model.layers.31.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 455 |
+
"model.layers.31.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 456 |
+
"model.layers.31.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 457 |
+
"model.layers.31.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 458 |
+
"model.layers.31.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 459 |
+
"model.layers.31.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 460 |
+
"model.layers.31.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 461 |
+
"model.layers.32.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 462 |
+
"model.layers.32.layer_scalar": "model-00003-of-00003.safetensors",
|
| 463 |
+
"model.layers.32.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 464 |
+
"model.layers.32.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 465 |
+
"model.layers.32.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 466 |
+
"model.layers.32.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 467 |
+
"model.layers.32.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 468 |
+
"model.layers.32.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 469 |
+
"model.layers.32.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 470 |
+
"model.layers.32.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 471 |
+
"model.layers.32.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 472 |
+
"model.layers.32.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 473 |
+
"model.layers.32.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 474 |
+
"model.layers.32.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 475 |
+
"model.layers.32.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 476 |
+
"model.layers.32.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 477 |
+
"model.layers.32.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 478 |
+
"model.layers.33.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 479 |
+
"model.layers.33.layer_scalar": "model-00003-of-00003.safetensors",
|
| 480 |
+
"model.layers.33.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 481 |
+
"model.layers.33.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 482 |
+
"model.layers.33.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 483 |
+
"model.layers.33.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 484 |
+
"model.layers.33.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 485 |
+
"model.layers.33.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 486 |
+
"model.layers.33.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 487 |
+
"model.layers.33.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 488 |
+
"model.layers.33.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 489 |
+
"model.layers.33.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 490 |
+
"model.layers.33.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 491 |
+
"model.layers.33.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 492 |
+
"model.layers.33.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 493 |
+
"model.layers.33.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 494 |
+
"model.layers.33.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 495 |
+
"model.layers.34.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 496 |
+
"model.layers.34.layer_scalar": "model-00003-of-00003.safetensors",
|
| 497 |
+
"model.layers.34.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 498 |
+
"model.layers.34.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 499 |
+
"model.layers.34.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 500 |
+
"model.layers.34.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 501 |
+
"model.layers.34.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 502 |
+
"model.layers.34.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 503 |
+
"model.layers.34.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 504 |
+
"model.layers.34.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 505 |
+
"model.layers.34.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 506 |
+
"model.layers.34.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 507 |
+
"model.layers.34.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 508 |
+
"model.layers.34.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 509 |
+
"model.layers.34.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 510 |
+
"model.layers.34.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 511 |
+
"model.layers.34.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 512 |
+
"model.layers.4.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 513 |
+
"model.layers.4.layer_scalar": "model-00003-of-00003.safetensors",
|
| 514 |
+
"model.layers.4.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 515 |
+
"model.layers.4.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 516 |
+
"model.layers.4.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 517 |
+
"model.layers.4.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 518 |
+
"model.layers.4.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 519 |
+
"model.layers.4.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 520 |
+
"model.layers.4.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 521 |
+
"model.layers.4.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 522 |
+
"model.layers.4.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 523 |
+
"model.layers.4.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 524 |
+
"model.layers.4.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 525 |
+
"model.layers.4.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 526 |
+
"model.layers.4.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 527 |
+
"model.layers.4.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 528 |
+
"model.layers.4.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 529 |
+
"model.layers.5.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 530 |
+
"model.layers.5.layer_scalar": "model-00003-of-00003.safetensors",
|
| 531 |
+
"model.layers.5.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 532 |
+
"model.layers.5.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 533 |
+
"model.layers.5.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 534 |
+
"model.layers.5.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 535 |
+
"model.layers.5.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 536 |
+
"model.layers.5.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 537 |
+
"model.layers.5.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 538 |
+
"model.layers.5.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 539 |
+
"model.layers.5.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 540 |
+
"model.layers.5.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 541 |
+
"model.layers.5.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 542 |
+
"model.layers.5.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 543 |
+
"model.layers.5.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 544 |
+
"model.layers.5.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 545 |
+
"model.layers.5.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 546 |
+
"model.layers.6.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 547 |
+
"model.layers.6.layer_scalar": "model-00003-of-00003.safetensors",
|
| 548 |
+
"model.layers.6.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 549 |
+
"model.layers.6.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 550 |
+
"model.layers.6.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 551 |
+
"model.layers.6.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 552 |
+
"model.layers.6.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 553 |
+
"model.layers.6.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 554 |
+
"model.layers.6.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 555 |
+
"model.layers.6.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 556 |
+
"model.layers.6.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 557 |
+
"model.layers.6.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 558 |
+
"model.layers.6.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 559 |
+
"model.layers.6.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 560 |
+
"model.layers.6.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 561 |
+
"model.layers.6.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 562 |
+
"model.layers.6.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 563 |
+
"model.layers.7.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 564 |
+
"model.layers.7.layer_scalar": "model-00003-of-00003.safetensors",
|
| 565 |
+
"model.layers.7.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 566 |
+
"model.layers.7.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 567 |
+
"model.layers.7.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 568 |
+
"model.layers.7.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 569 |
+
"model.layers.7.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 570 |
+
"model.layers.7.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 571 |
+
"model.layers.7.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 572 |
+
"model.layers.7.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 573 |
+
"model.layers.7.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 574 |
+
"model.layers.7.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 575 |
+
"model.layers.7.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 576 |
+
"model.layers.7.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 577 |
+
"model.layers.7.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 578 |
+
"model.layers.7.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 579 |
+
"model.layers.7.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 580 |
+
"model.layers.8.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 581 |
+
"model.layers.8.layer_scalar": "model-00003-of-00003.safetensors",
|
| 582 |
+
"model.layers.8.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 583 |
+
"model.layers.8.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 584 |
+
"model.layers.8.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 585 |
+
"model.layers.8.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 586 |
+
"model.layers.8.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 587 |
+
"model.layers.8.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 588 |
+
"model.layers.8.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 589 |
+
"model.layers.8.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 590 |
+
"model.layers.8.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 591 |
+
"model.layers.8.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 592 |
+
"model.layers.8.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 593 |
+
"model.layers.8.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 594 |
+
"model.layers.8.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 595 |
+
"model.layers.8.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 596 |
+
"model.layers.8.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 597 |
+
"model.layers.9.input_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 598 |
+
"model.layers.9.layer_scalar": "model-00003-of-00003.safetensors",
|
| 599 |
+
"model.layers.9.mlp.down_proj.weight": "model-00003-of-00003.safetensors",
|
| 600 |
+
"model.layers.9.mlp.gate_proj.weight": "model-00003-of-00003.safetensors",
|
| 601 |
+
"model.layers.9.mlp.up_proj.weight": "model-00003-of-00003.safetensors",
|
| 602 |
+
"model.layers.9.per_layer_input_gate.weight": "model-00003-of-00003.safetensors",
|
| 603 |
+
"model.layers.9.per_layer_projection.weight": "model-00003-of-00003.safetensors",
|
| 604 |
+
"model.layers.9.post_attention_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 605 |
+
"model.layers.9.post_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 606 |
+
"model.layers.9.post_per_layer_input_norm.weight": "model-00003-of-00003.safetensors",
|
| 607 |
+
"model.layers.9.pre_feedforward_layernorm.weight": "model-00003-of-00003.safetensors",
|
| 608 |
+
"model.layers.9.self_attn.k_norm.weight": "model-00003-of-00003.safetensors",
|
| 609 |
+
"model.layers.9.self_attn.k_proj.weight": "model-00003-of-00003.safetensors",
|
| 610 |
+
"model.layers.9.self_attn.o_proj.weight": "model-00003-of-00003.safetensors",
|
| 611 |
+
"model.layers.9.self_attn.q_norm.weight": "model-00003-of-00003.safetensors",
|
| 612 |
+
"model.layers.9.self_attn.q_proj.weight": "model-00003-of-00003.safetensors",
|
| 613 |
+
"model.layers.9.self_attn.v_proj.weight": "model-00003-of-00003.safetensors",
|
| 614 |
+
"model.norm.weight": "model-00003-of-00003.safetensors",
|
| 615 |
+
"model.per_layer_model_projection.weight": "model-00003-of-00003.safetensors",
|
| 616 |
+
"model.per_layer_projection_norm.weight": "model-00003-of-00003.safetensors"
|
| 617 |
+
}
|
| 618 |
+
}
|
modeling_gemma_4_e2b_it_hybrid.py
ADDED
|
@@ -0,0 +1,165 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Cactus-Compute/gemma-4-e2b-it-hybrid — Gemma-4 causal LM with a handoff probe.
|
| 2 |
+
|
| 3 |
+
``Gemma4E2BItHybridForCausalLM`` is the stock ``Gemma4ForCausalLM`` plus a small
|
| 4 |
+
"handoff probe" head (weight prefix ``handoff_probe.*``) that scores each
|
| 5 |
+
generation with ``confidence = 1 - p_wrong``. Base weights keep identical keys,
|
| 6 |
+
so the checkpoint is the stock checkpoint with eleven extra probe tensors.
|
| 7 |
+
|
| 8 |
+
Probe contract (checkpoint layer 28, float32 math):
|
| 9 |
+
|
| 10 |
+
- input: ``[T, 1536]`` — output of decoder layer index ``config.probe_layer``
|
| 11 |
+
at the position that predicts each generated token (row 0 = last prompt
|
| 12 |
+
position at prefill, row t = position captured at generation step t). Only
|
| 13 |
+
the first ``config.probe_max_tokens`` rows are used.
|
| 14 |
+
- ``x = LayerNorm(x, eps=1e-5) * norm.weight + norm.bias``
|
| 15 |
+
- ``p = relu(x @ proj.weight.T + proj.bias)``
|
| 16 |
+
- ``s = p @ attn_query / sqrt(probe_proj_dim)``; ``w = softmax_T(s)``
|
| 17 |
+
- ``pooled = w @ p``
|
| 18 |
+
- ``h = relu(head.0 @ pooled); h = relu(head.2 @ h); logit = head.4 @ h``
|
| 19 |
+
- ``p_wrong = sigmoid(logit)``; ``confidence = 1 - p_wrong``
|
| 20 |
+
|
| 21 |
+
Layer capture uses a forward hook that keeps only the last position of each
|
| 22 |
+
decode step (a ``[batch, hidden]`` row), never the full hidden-state stack.
|
| 23 |
+
"""
|
| 24 |
+
|
| 25 |
+
import math
|
| 26 |
+
from contextlib import contextmanager
|
| 27 |
+
|
| 28 |
+
import torch
|
| 29 |
+
import torch.nn.functional as F
|
| 30 |
+
from torch import nn
|
| 31 |
+
|
| 32 |
+
from transformers.generation.utils import GenerationMixin
|
| 33 |
+
from transformers.models.gemma4.modeling_gemma4 import Gemma4ForCausalLM
|
| 34 |
+
|
| 35 |
+
try:
|
| 36 |
+
from .configuration_gemma_4_e2b_it_hybrid import Gemma4E2BItHybridConfig
|
| 37 |
+
except ImportError: # direct (non-package) execution
|
| 38 |
+
from configuration_gemma_4_e2b_it_hybrid import Gemma4E2BItHybridConfig
|
| 39 |
+
|
| 40 |
+
# Hidden widths of the probe MLP head, fixed by the released checkpoint.
|
| 41 |
+
PROBE_HEAD_DIMS = (128, 64)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
class HandoffProbe(nn.Module):
|
| 45 |
+
"""The released handoff probe. Weight keys match the probe checkpoint:
|
| 46 |
+
|
| 47 |
+
``norm.{weight,bias}``, ``proj.{weight,bias}``, ``attn_query``,
|
| 48 |
+
``head.{0,2,4}.{weight,bias}``.
|
| 49 |
+
"""
|
| 50 |
+
|
| 51 |
+
def __init__(self, feature_size: int, proj_dim: int = 32) -> None:
|
| 52 |
+
super().__init__()
|
| 53 |
+
h1, h2 = PROBE_HEAD_DIMS
|
| 54 |
+
self.norm = nn.LayerNorm(feature_size, eps=1e-5)
|
| 55 |
+
self.proj = nn.Linear(feature_size, proj_dim)
|
| 56 |
+
self.attn_query = nn.Parameter(torch.zeros(proj_dim))
|
| 57 |
+
self.head = nn.Sequential(
|
| 58 |
+
nn.Linear(proj_dim, h1),
|
| 59 |
+
nn.ReLU(),
|
| 60 |
+
nn.Linear(h1, h2),
|
| 61 |
+
nn.ReLU(),
|
| 62 |
+
nn.Linear(h2, 1),
|
| 63 |
+
)
|
| 64 |
+
|
| 65 |
+
@torch.no_grad()
|
| 66 |
+
def p_wrong(self, hidden_states: torch.Tensor, max_tokens: int = 1024) -> float:
|
| 67 |
+
"""Score ``[T, feature_size]`` generated-token hidden states.
|
| 68 |
+
|
| 69 |
+
All math runs in float32 on the probe's own device, regardless of the
|
| 70 |
+
dtype the module weights were loaded in; the input may arrive on any
|
| 71 |
+
device (capture hooks store rows on CPU).
|
| 72 |
+
"""
|
| 73 |
+
if hidden_states.ndim != 2:
|
| 74 |
+
raise ValueError(f"expected [tokens, features], got {tuple(hidden_states.shape)}")
|
| 75 |
+
if hidden_states.shape[0] == 0:
|
| 76 |
+
raise ValueError("cannot score an empty generation")
|
| 77 |
+
x = hidden_states[:max_tokens].to(self.norm.weight.device, torch.float32)
|
| 78 |
+
|
| 79 |
+
x = F.layer_norm(
|
| 80 |
+
x, (x.shape[-1],), self.norm.weight.float(), self.norm.bias.float(), self.norm.eps
|
| 81 |
+
)
|
| 82 |
+
projected = F.relu(F.linear(x, self.proj.weight.float(), self.proj.bias.float()))
|
| 83 |
+
scores = projected @ self.attn_query.float() / math.sqrt(projected.shape[-1])
|
| 84 |
+
weights = torch.softmax(scores, dim=0)
|
| 85 |
+
pooled = weights @ projected
|
| 86 |
+
|
| 87 |
+
h = F.relu(F.linear(pooled, self.head[0].weight.float(), self.head[0].bias.float()))
|
| 88 |
+
h = F.relu(F.linear(h, self.head[2].weight.float(), self.head[2].bias.float()))
|
| 89 |
+
logit = F.linear(h, self.head[4].weight.float(), self.head[4].bias.float())
|
| 90 |
+
return float(torch.sigmoid(logit)[0].item())
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
class Gemma4E2BItHybridForCausalLM(Gemma4ForCausalLM):
|
| 94 |
+
"""Stock Gemma-4 causal LM plus the ``handoff_probe.*`` scoring head."""
|
| 95 |
+
|
| 96 |
+
config_class = Gemma4E2BItHybridConfig
|
| 97 |
+
config: Gemma4E2BItHybridConfig
|
| 98 |
+
|
| 99 |
+
def __init__(self, config: Gemma4E2BItHybridConfig) -> None:
|
| 100 |
+
if config.probe_feature_size != config.hidden_size:
|
| 101 |
+
raise ValueError(
|
| 102 |
+
f"probe_feature_size={config.probe_feature_size} must equal "
|
| 103 |
+
f"hidden_size={config.hidden_size}"
|
| 104 |
+
)
|
| 105 |
+
super().__init__(config)
|
| 106 |
+
self.handoff_probe = HandoffProbe(
|
| 107 |
+
config.probe_feature_size, getattr(config, "probe_proj_dim", 32)
|
| 108 |
+
)
|
| 109 |
+
#: Confidence of the most recent scored generation (``None`` before any).
|
| 110 |
+
self.last_confidence: float | None = None
|
| 111 |
+
|
| 112 |
+
@contextmanager
|
| 113 |
+
def probe_capture(self):
|
| 114 |
+
"""Capture probe-layer rows during generation.
|
| 115 |
+
|
| 116 |
+
Yields a list that fills with one ``[batch, hidden]`` float32 CPU tensor
|
| 117 |
+
per decode step (row 0 comes from the prefill forward at the last prompt
|
| 118 |
+
position). At most ``config.probe_max_tokens`` rows are kept.
|
| 119 |
+
"""
|
| 120 |
+
rows: list[torch.Tensor] = []
|
| 121 |
+
layer = self.model.layers[self.config.probe_layer]
|
| 122 |
+
max_rows = self.config.probe_max_tokens
|
| 123 |
+
|
| 124 |
+
def hook(module, args, output):
|
| 125 |
+
if len(rows) < max_rows:
|
| 126 |
+
hidden = output[0] if isinstance(output, tuple) else output
|
| 127 |
+
rows.append(hidden[:, -1, :].detach().to(torch.float32).cpu())
|
| 128 |
+
|
| 129 |
+
handle = layer.register_forward_hook(hook)
|
| 130 |
+
try:
|
| 131 |
+
yield rows
|
| 132 |
+
finally:
|
| 133 |
+
handle.remove()
|
| 134 |
+
|
| 135 |
+
def confidence_from_rows(self, rows: list[torch.Tensor]) -> float | None:
|
| 136 |
+
"""Turn captured probe rows into ``confidence = 1 - p_wrong``.
|
| 137 |
+
|
| 138 |
+
Returns ``None`` when the capture is empty or not a single-sequence
|
| 139 |
+
generation (batch size > 1, beam search, ...), for which the probe
|
| 140 |
+
contract is undefined.
|
| 141 |
+
"""
|
| 142 |
+
if not rows or any(row.shape[0] != 1 for row in rows):
|
| 143 |
+
return None
|
| 144 |
+
states = torch.cat(rows, dim=0) # [T, hidden]
|
| 145 |
+
p_wrong = self.handoff_probe.p_wrong(states, self.config.probe_max_tokens)
|
| 146 |
+
self.last_confidence = 1.0 - p_wrong
|
| 147 |
+
return self.last_confidence
|
| 148 |
+
|
| 149 |
+
def generate_with_confidence(self, *args, **kwargs):
|
| 150 |
+
"""Run the stock ``generate`` and score it with the handoff probe.
|
| 151 |
+
|
| 152 |
+
Returns ``(sequences, confidence)`` where ``sequences`` is exactly what
|
| 153 |
+
the stock ``generate`` returns for the given arguments (no in-band
|
| 154 |
+
trailer is appended) and ``confidence`` is a float in ``[0, 1]``, or
|
| 155 |
+
``None`` when the generation cannot be scored (batch > 1, beams, or
|
| 156 |
+
assisted decoding).
|
| 157 |
+
"""
|
| 158 |
+
kwargs.pop("return_confidence", None)
|
| 159 |
+
with self.probe_capture() as rows:
|
| 160 |
+
sequences = GenerationMixin.generate(self, *args, **kwargs)
|
| 161 |
+
confidence = self.confidence_from_rows(rows)
|
| 162 |
+
return sequences, confidence
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
__all__ = ["Gemma4E2BItHybridForCausalLM", "HandoffProbe"]
|
tokenizer.json
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:cc8d3a0ce36466ccc1278bf987df5f71db1719b9ca6b4118264f45cb627bfe0f
|
| 3 |
+
size 32169626
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,120 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"audio_token": "<|audio|>",
|
| 3 |
+
"backend": "tokenizers",
|
| 4 |
+
"boa_token": "<|audio>",
|
| 5 |
+
"boi_token": "<|image>",
|
| 6 |
+
"bos_token": "<bos>",
|
| 7 |
+
"eoa_token": "<audio|>",
|
| 8 |
+
"eoc_token": "<channel|>",
|
| 9 |
+
"eoi_token": "<image|>",
|
| 10 |
+
"eos_token": "<eos>",
|
| 11 |
+
"eot_token": "<turn|>",
|
| 12 |
+
"escape_token": "<|\"|>",
|
| 13 |
+
"etc_token": "<tool_call|>",
|
| 14 |
+
"etd_token": "<tool|>",
|
| 15 |
+
"etr_token": "<tool_response|>",
|
| 16 |
+
"extra_special_tokens": [
|
| 17 |
+
"<|video|>"
|
| 18 |
+
],
|
| 19 |
+
"image_token": "<|image|>",
|
| 20 |
+
"mask_token": "<mask>",
|
| 21 |
+
"model_max_length": 1000000000000000019884624838656,
|
| 22 |
+
"pad_token": "<pad>",
|
| 23 |
+
"padding_side": "left",
|
| 24 |
+
"processor_class": "Gemma4Processor",
|
| 25 |
+
"response_schema": {
|
| 26 |
+
"type": "object",
|
| 27 |
+
"properties": {
|
| 28 |
+
"role": {
|
| 29 |
+
"const": "assistant"
|
| 30 |
+
},
|
| 31 |
+
"thinking": {
|
| 32 |
+
"type": "string"
|
| 33 |
+
},
|
| 34 |
+
"content": {
|
| 35 |
+
"type": "string"
|
| 36 |
+
},
|
| 37 |
+
"tool_calls": {
|
| 38 |
+
"x-regex-iterator": "<\\|tool_call>(.*?)<tool_call\\|>",
|
| 39 |
+
"type": "array",
|
| 40 |
+
"items": {
|
| 41 |
+
"type": "object",
|
| 42 |
+
"properties": {
|
| 43 |
+
"type": {
|
| 44 |
+
"const": "function"
|
| 45 |
+
},
|
| 46 |
+
"function": {
|
| 47 |
+
"type": "object",
|
| 48 |
+
"x-regex": "call\\:(?P<name>\\w+)(?P<arguments>\\{.*\\})",
|
| 49 |
+
"properties": {
|
| 50 |
+
"name": {
|
| 51 |
+
"type": "string"
|
| 52 |
+
},
|
| 53 |
+
"arguments": {
|
| 54 |
+
"type": "object",
|
| 55 |
+
"x-parser": "gemma4-tool-call",
|
| 56 |
+
"additionalProperties": {}
|
| 57 |
+
}
|
| 58 |
+
}
|
| 59 |
+
}
|
| 60 |
+
}
|
| 61 |
+
}
|
| 62 |
+
}
|
| 63 |
+
},
|
| 64 |
+
"x-regex": "(\\<\\|channel\\>thought\\n(?P<thinking>.*?)\\<channel\\|\\>)?(?P<tool_calls>\\<\\|tool_call\\>.*\\<tool_call\\|\\>)?(?P<content>(?:(?!\\<turn\\|\\>)(?!\\<\\|tool_response\\>).)+)?(?:\\<turn\\|\\>|\\<\\|tool_response\\>)?"
|
| 65 |
+
},
|
| 66 |
+
"response_template": {
|
| 67 |
+
"defaults": {
|
| 68 |
+
"role": "assistant"
|
| 69 |
+
},
|
| 70 |
+
"fields": {
|
| 71 |
+
"content": {
|
| 72 |
+
"close": [
|
| 73 |
+
"<turn|>",
|
| 74 |
+
"<|tool_response>",
|
| 75 |
+
"<eos>"
|
| 76 |
+
],
|
| 77 |
+
"content": "text"
|
| 78 |
+
},
|
| 79 |
+
"thinking": {
|
| 80 |
+
"close": "<channel|>",
|
| 81 |
+
"content": "text",
|
| 82 |
+
"open": "<|channel>thought\n"
|
| 83 |
+
},
|
| 84 |
+
"tool_calls": {
|
| 85 |
+
"close": "<tool_call|>",
|
| 86 |
+
"content": "json",
|
| 87 |
+
"content_args": {
|
| 88 |
+
"string_delims": [
|
| 89 |
+
[
|
| 90 |
+
"<|\"|>",
|
| 91 |
+
"<|\"|>"
|
| 92 |
+
]
|
| 93 |
+
],
|
| 94 |
+
"unquoted_keys": true
|
| 95 |
+
},
|
| 96 |
+
"open_pattern": "<\\|tool_call>call:(?P<name>\\w+)",
|
| 97 |
+
"repeats": true,
|
| 98 |
+
"transform": {
|
| 99 |
+
"function": {
|
| 100 |
+
"arguments": "{content}",
|
| 101 |
+
"name": "{name}"
|
| 102 |
+
},
|
| 103 |
+
"type": "function"
|
| 104 |
+
}
|
| 105 |
+
}
|
| 106 |
+
},
|
| 107 |
+
"start_anchor": [
|
| 108 |
+
"<|turn>model\n",
|
| 109 |
+
"<tool_response|>"
|
| 110 |
+
]
|
| 111 |
+
},
|
| 112 |
+
"soc_token": "<|channel>",
|
| 113 |
+
"sot_token": "<|turn>",
|
| 114 |
+
"stc_token": "<|tool_call>",
|
| 115 |
+
"std_token": "<|tool>",
|
| 116 |
+
"str_token": "<|tool_response>",
|
| 117 |
+
"think_token": "<|think|>",
|
| 118 |
+
"tokenizer_class": "GemmaTokenizer",
|
| 119 |
+
"unk_token": "<unk>"
|
| 120 |
+
}
|