summerMC commited on
Commit
efc38b5
·
verified ·
1 Parent(s): 701dcba

Upload Qwen3.5 UNI MAX 9B checkpoint with benchmark results

Browse files
.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 CHANGED
@@ -1,3 +1,253 @@
1
  ---
2
- license: apache-2.0
 
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
+ library_name: transformers
3
+ pipeline_tag: text-generation
4
+ tags:
5
+ - qwen
6
+ - qwen3.5
7
+ - text-generation
8
+ - recurrent
9
+ - linear-attention
10
+ - cuda
11
+ - custom-code
12
+ base_model: Qwen/Qwen3.5-9B
13
  ---
14
+
15
+ # Qwen3.5 UNI MAX 9B
16
+
17
+ This repository contains a UNI MAX conversion of
18
+ `Qwen/Qwen3.5-9B`.
19
+
20
+ The source hybrid attention topology is converted to a recurrent
21
+ linear-attention runtime. Full-attention layers are replaced through
22
+ layer-wise distillation from neighboring native GDN/linear-attention
23
+ layers.
24
+
25
+ ## Architecture
26
+
27
+ - Base model: `Qwen/Qwen3.5-9B`
28
+ - Hidden layers: `32`
29
+ - Runtime topology: linear attention
30
+ - Converted source full-attention layers: `3, 7, 11, 15, 19, 23, 27, 31`
31
+ - Custom runtime: `qwen35_unimax_engine`
32
+ - Kernel acceleration: enabled for the benchmark below
33
+
34
+ The checkpoint should be used with the matching UNI MAX runtime code.
35
+ It is not claimed to be numerically identical to the original
36
+ full-attention model.
37
+
38
+ ## Distillation
39
+
40
+ The converted full-attention layers were initialized from neighboring
41
+ native GDN layers and then optimized with sequential layer-wise
42
+ distillation.
43
+
44
+ The current build used 20 optimization steps per converted layer.
45
+
46
+ ## Benchmark environment
47
+
48
+ | Item | Value |
49
+ |---|---|
50
+ | GPU | NVIDIA RTX PRO 6000 Blackwell Server Edition |
51
+ | GPU VRAM | 94.97 GiB |
52
+ | Peak allocated VRAM during benchmark | 70.14 GiB |
53
+ | PyTorch | 2.11.0+cu130 |
54
+ | CUDA | 13.0 |
55
+ | CUDA capability | 12.0 |
56
+ | NVIDIA driver | 580.82.07 |
57
+ | Decode steps | 64 |
58
+ | Warmup runs | 2 |
59
+ | Measured runs | 5 |
60
+ | Context lengths | 256, 1024, 4096, 16384 |
61
+ | Benchmark wall time | 90.62 s |
62
+
63
+ ## Benchmark results
64
+
65
+ | Metric | Value |
66
+ |---|---:|
67
+ | verified | None |
68
+ | selected_mode | exact |
69
+ | selected_block | 8 |
70
+
71
+ The benchmark above was generated directly from
72
+ `benchmark_uni_max()` on the hardware shown above.
73
+
74
+ Because throughput, latency, kernel selection, quantization support,
75
+ and memory consumption depend strongly on the GPU, CUDA/PyTorch
76
+ versions, context length, and runtime configuration, these numbers
77
+ should not be treated as hardware-independent performance claims.
78
+
79
+ ## Raw benchmark output
80
+
81
+ ```json
82
+ {
83
+ "official": [
84
+ [
85
+ 256,
86
+ 9771.298185735035,
87
+ 61.8274717963079
88
+ ],
89
+ [
90
+ 1024,
91
+ 17823.304562533922,
92
+ 61.66807965901939
93
+ ],
94
+ [
95
+ 4096,
96
+ 20324.797811420387,
97
+ 60.59438884144146
98
+ ],
99
+ [
100
+ 16384,
101
+ 18413.893708085165,
102
+ 58.409657369508345
103
+ ]
104
+ ],
105
+ "safe": [
106
+ [
107
+ 256,
108
+ 10715.681982051035,
109
+ 79.45237658001217
110
+ ],
111
+ [
112
+ 1024,
113
+ 19627.025186990726,
114
+ 79.44713049745303
115
+ ],
116
+ [
117
+ 4096,
118
+ 22395.258154021976,
119
+ 79.44376760906928
120
+ ],
121
+ [
122
+ 16384,
123
+ 22677.037604190627,
124
+ 79.4362706749314
125
+ ]
126
+ ],
127
+ "speed": [
128
+ [
129
+ 256,
130
+ 10573.784778868314,
131
+ 79.44571045597908,
132
+ 0.28697000016109087,
133
+ 10452.495043114082
134
+ ],
135
+ [
136
+ 1024,
137
+ 19569.254713139573,
138
+ 79.44294428720825,
139
+ 0.2903700001297693,
140
+ 19462.145367010704
141
+ ],
142
+ [
143
+ 4096,
144
+ 22337.05592154273,
145
+ 79.44607524961788,
146
+ 0.29715000027863425,
147
+ 22300.547693510274
148
+ ],
149
+ [
150
+ 16384,
151
+ 22578.659939700647,
152
+ 79.43958559875574,
153
+ 0.32848999990164884,
154
+ 22568.420127905207
155
+ ]
156
+ ],
157
+ "verified": null,
158
+ "selected_mode": "exact",
159
+ "selected_block": 8,
160
+ "candidates": [
161
+ {
162
+ "mode": "fp8",
163
+ "accepted": false,
164
+ "agreement": 0.1875,
165
+ "prefix": 10,
166
+ "replay_tok_s": 84.39886475159058,
167
+ "block": 8,
168
+ "reason": "agreement 0.188 < 1.000"
169
+ },
170
+ {
171
+ "mode": "int8",
172
+ "accepted": false,
173
+ "agreement": 1.0,
174
+ "prefix": 64,
175
+ "replay_tok_s": 64.56140914184326,
176
+ "block": 2,
177
+ "reason": "not faster than current 79.45 tok/s"
178
+ }
179
+ ]
180
+ }
181
+ ```
182
+
183
+ The machine-readable benchmark record is also included as
184
+ `benchmark_results.json`.
185
+
186
+ ## Runtime configuration
187
+
188
+ The benchmark evaluated:
189
+
190
+ - CUDA kernels enabled
191
+ - graph candidates: `(2, 4, 8)`
192
+ - quantization candidates: `('fp8', 'int8')`
193
+ - quality verification steps: `64`
194
+ - quality context: `64`
195
+ - minimum fast-path agreement: `1.0`
196
+
197
+ The runtime may select a different execution path depending on
198
+ hardware support and quality verification.
199
+
200
+ ## Usage
201
+
202
+ Install the matching `qwen35_unimax_engine` runtime before loading the
203
+ checkpoint.
204
+
205
+ ```python
206
+ from qwen35_unimax_engine import UniMaxPipeline
207
+
208
+ pipe = UniMaxPipeline.from_pretrained(
209
+ "YOUR_USERNAME/qwen35-unimax-engine-9b-v2_1",
210
+ mode="auto",
211
+ graph_candidates=(2, 4, 8),
212
+ quant_modes=("fp8", "int8"),
213
+ quality_steps=64,
214
+ min_agreement=1.0,
215
+ use_kernels=True,
216
+ )
217
+
218
+ print(pipe.runtime_summary())
219
+
220
+ output = pipe(
221
+ "Explain recurrent linear attention:",
222
+ max_new_tokens=128,
223
+ )
224
+
225
+ print(output)
226
+ ```
227
+
228
+ ## Validation
229
+
230
+ Before publication, the local checkpoint was checked with the UNI MAX
231
+ checkpoint validator and benchmark path.
232
+
233
+ Users should independently evaluate task quality before deploying the
234
+ converted model. Layer-wise agreement tests are narrower than a full
235
+ language-model evaluation suite.
236
+
237
+ ## Limitations
238
+
239
+ This is a converted and distilled model, not the unmodified
240
+ `Qwen/Qwen3.5-9B` checkpoint.
241
+
242
+ Replacing full attention with recurrent linear attention changes the
243
+ model architecture and can affect accuracy, long-context behavior,
244
+ reasoning, generation quality, and numerical output.
245
+
246
+ Benchmark results are specific to the hardware and software
247
+ environment documented above.
248
+
249
+ ## Base model
250
+
251
+ Derived from `Qwen/Qwen3.5-9B`. Refer to the upstream model repository
252
+ for the original model documentation, license, intended use, and
253
+ limitations.
benchmark_results.json ADDED
@@ -0,0 +1,140 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model": "/content/qwen35_uni_max_work/qwen35-unimax-engine-9b-v2_1",
3
+ "source_model": "Qwen/Qwen3.5-9B",
4
+ "environment": {
5
+ "gpu": "NVIDIA RTX PRO 6000 Blackwell Server Edition",
6
+ "vram_gib": 94.971,
7
+ "peak_allocated_vram_gib": 70.141,
8
+ "torch": "2.11.0+cu130",
9
+ "cuda": "13.0",
10
+ "cuda_capability": [
11
+ 12,
12
+ 0
13
+ ],
14
+ "nvidia_driver": "580.82.07",
15
+ "python": "3.13.15"
16
+ },
17
+ "benchmark": {
18
+ "contexts": [
19
+ 256,
20
+ 1024,
21
+ 4096,
22
+ 16384
23
+ ],
24
+ "decode_steps": 64,
25
+ "warmup": 2,
26
+ "runs": 5,
27
+ "graph_candidates": [
28
+ 2,
29
+ 4,
30
+ 8
31
+ ],
32
+ "quant_modes": [
33
+ "fp8",
34
+ "int8"
35
+ ],
36
+ "quality_steps": 64,
37
+ "quality_context": 64,
38
+ "min_agreement": 1.0,
39
+ "wall_time_seconds": 90.61699619499996
40
+ },
41
+ "results": {
42
+ "official": [
43
+ [
44
+ 256,
45
+ 9771.298185735035,
46
+ 61.8274717963079
47
+ ],
48
+ [
49
+ 1024,
50
+ 17823.304562533922,
51
+ 61.66807965901939
52
+ ],
53
+ [
54
+ 4096,
55
+ 20324.797811420387,
56
+ 60.59438884144146
57
+ ],
58
+ [
59
+ 16384,
60
+ 18413.893708085165,
61
+ 58.409657369508345
62
+ ]
63
+ ],
64
+ "safe": [
65
+ [
66
+ 256,
67
+ 10715.681982051035,
68
+ 79.45237658001217
69
+ ],
70
+ [
71
+ 1024,
72
+ 19627.025186990726,
73
+ 79.44713049745303
74
+ ],
75
+ [
76
+ 4096,
77
+ 22395.258154021976,
78
+ 79.44376760906928
79
+ ],
80
+ [
81
+ 16384,
82
+ 22677.037604190627,
83
+ 79.4362706749314
84
+ ]
85
+ ],
86
+ "speed": [
87
+ [
88
+ 256,
89
+ 10573.784778868314,
90
+ 79.44571045597908,
91
+ 0.28697000016109087,
92
+ 10452.495043114082
93
+ ],
94
+ [
95
+ 1024,
96
+ 19569.254713139573,
97
+ 79.44294428720825,
98
+ 0.2903700001297693,
99
+ 19462.145367010704
100
+ ],
101
+ [
102
+ 4096,
103
+ 22337.05592154273,
104
+ 79.44607524961788,
105
+ 0.29715000027863425,
106
+ 22300.547693510274
107
+ ],
108
+ [
109
+ 16384,
110
+ 22578.659939700647,
111
+ 79.43958559875574,
112
+ 0.32848999990164884,
113
+ 22568.420127905207
114
+ ]
115
+ ],
116
+ "verified": null,
117
+ "selected_mode": "exact",
118
+ "selected_block": 8,
119
+ "candidates": [
120
+ {
121
+ "mode": "fp8",
122
+ "accepted": false,
123
+ "agreement": 0.1875,
124
+ "prefix": 10,
125
+ "replay_tok_s": 84.39886475159058,
126
+ "block": 8,
127
+ "reason": "agreement 0.188 < 1.000"
128
+ },
129
+ {
130
+ "mode": "int8",
131
+ "accepted": false,
132
+ "agreement": 1.0,
133
+ "prefix": 64,
134
+ "replay_tok_s": 64.56140914184326,
135
+ "block": 2,
136
+ "reason": "not faster than current 79.45 tok/s"
137
+ }
138
+ ]
139
+ }
140
+ }
chat_template.jinja ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {%- set image_count = namespace(value=0) %}
2
+ {%- set video_count = namespace(value=0) %}
3
+ {%- macro render_content(content, do_vision_count, is_system_content=false) %}
4
+ {%- if content is string %}
5
+ {{- content }}
6
+ {%- elif content is iterable and content is not mapping %}
7
+ {%- for item in content %}
8
+ {%- if 'image' in item or 'image_url' in item or item.type == 'image' %}
9
+ {%- if is_system_content %}
10
+ {{- raise_exception('System message cannot contain images.') }}
11
+ {%- endif %}
12
+ {%- if do_vision_count %}
13
+ {%- set image_count.value = image_count.value + 1 %}
14
+ {%- endif %}
15
+ {%- if add_vision_id %}
16
+ {{- 'Picture ' ~ image_count.value ~ ': ' }}
17
+ {%- endif %}
18
+ {{- '<|vision_start|><|image_pad|><|vision_end|>' }}
19
+ {%- elif 'video' in item or item.type == 'video' %}
20
+ {%- if is_system_content %}
21
+ {{- raise_exception('System message cannot contain videos.') }}
22
+ {%- endif %}
23
+ {%- if do_vision_count %}
24
+ {%- set video_count.value = video_count.value + 1 %}
25
+ {%- endif %}
26
+ {%- if add_vision_id %}
27
+ {{- 'Video ' ~ video_count.value ~ ': ' }}
28
+ {%- endif %}
29
+ {{- '<|vision_start|><|video_pad|><|vision_end|>' }}
30
+ {%- elif 'text' in item %}
31
+ {{- item.text }}
32
+ {%- else %}
33
+ {{- raise_exception('Unexpected item type in content.') }}
34
+ {%- endif %}
35
+ {%- endfor %}
36
+ {%- elif content is none or content is undefined %}
37
+ {{- '' }}
38
+ {%- else %}
39
+ {{- raise_exception('Unexpected content type.') }}
40
+ {%- endif %}
41
+ {%- endmacro %}
42
+ {%- if not messages %}
43
+ {{- raise_exception('No messages provided.') }}
44
+ {%- endif %}
45
+ {%- if tools and tools is iterable and tools is not mapping %}
46
+ {{- '<|im_start|>system\n' }}
47
+ {{- "# Tools\n\nYou have access to the following functions:\n\n<tools>" }}
48
+ {%- for tool in tools %}
49
+ {{- "\n" }}
50
+ {{- tool | tojson }}
51
+ {%- endfor %}
52
+ {{- "\n</tools>" }}
53
+ {{- '\n\nIf you choose to call a function ONLY reply in the following format with NO suffix:\n\n<tool_call>\n<function=example_function_name>\n<parameter=example_parameter_1>\nvalue_1\n</parameter>\n<parameter=example_parameter_2>\nThis is the value for the second parameter\nthat can span\nmultiple lines\n</parameter>\n</function>\n</tool_call>\n\n<IMPORTANT>\nReminder:\n- Function calls MUST follow the specified format: an inner <function=...></function> block must be nested within <tool_call></tool_call> XML tags\n- Required parameters MUST be specified\n- You may provide optional reasoning for your function call in natural language BEFORE the function call, but NOT after\n- If there is no function call available, answer the question like normal with your current knowledge and do not tell the user about function calls\n</IMPORTANT>' }}
54
+ {%- if messages[0].role == 'system' %}
55
+ {%- set content = render_content(messages[0].content, false, true)|trim %}
56
+ {%- if content %}
57
+ {{- '\n\n' + content }}
58
+ {%- endif %}
59
+ {%- endif %}
60
+ {{- '<|im_end|>\n' }}
61
+ {%- else %}
62
+ {%- if messages[0].role == 'system' %}
63
+ {%- set content = render_content(messages[0].content, false, true)|trim %}
64
+ {{- '<|im_start|>system\n' + content + '<|im_end|>\n' }}
65
+ {%- endif %}
66
+ {%- endif %}
67
+ {%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %}
68
+ {%- for message in messages[::-1] %}
69
+ {%- set index = (messages|length - 1) - loop.index0 %}
70
+ {%- if ns.multi_step_tool and message.role == "user" %}
71
+ {%- set content = render_content(message.content, false)|trim %}
72
+ {%- if not(content.startswith('<tool_response>') and content.endswith('</tool_response>')) %}
73
+ {%- set ns.multi_step_tool = false %}
74
+ {%- set ns.last_query_index = index %}
75
+ {%- endif %}
76
+ {%- endif %}
77
+ {%- endfor %}
78
+ {%- if ns.multi_step_tool %}
79
+ {{- raise_exception('No user query found in messages.') }}
80
+ {%- endif %}
81
+ {%- for message in messages %}
82
+ {%- set content = render_content(message.content, true)|trim %}
83
+ {%- if message.role == "system" %}
84
+ {%- if not loop.first %}
85
+ {{- raise_exception('System message must be at the beginning.') }}
86
+ {%- endif %}
87
+ {%- elif message.role == "user" %}
88
+ {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }}
89
+ {%- elif message.role == "assistant" %}
90
+ {%- set reasoning_content = '' %}
91
+ {%- if message.reasoning_content is string %}
92
+ {%- set reasoning_content = message.reasoning_content %}
93
+ {%- else %}
94
+ {%- if '</think>' in content %}
95
+ {%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %}
96
+ {%- set content = content.split('</think>')[-1].lstrip('\n') %}
97
+ {%- endif %}
98
+ {%- endif %}
99
+ {%- set reasoning_content = reasoning_content|trim %}
100
+ {%- if loop.index0 > ns.last_query_index %}
101
+ {{- '<|im_start|>' + message.role + '\n<think>\n' + reasoning_content + '\n</think>\n\n' + content }}
102
+ {%- else %}
103
+ {{- '<|im_start|>' + message.role + '\n' + content }}
104
+ {%- endif %}
105
+ {%- if message.tool_calls and message.tool_calls is iterable and message.tool_calls is not mapping %}
106
+ {%- for tool_call in message.tool_calls %}
107
+ {%- if tool_call.function is defined %}
108
+ {%- set tool_call = tool_call.function %}
109
+ {%- endif %}
110
+ {%- if loop.first %}
111
+ {%- if content|trim %}
112
+ {{- '\n\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
113
+ {%- else %}
114
+ {{- '<tool_call>\n<function=' + tool_call.name + '>\n' }}
115
+ {%- endif %}
116
+ {%- else %}
117
+ {{- '\n<tool_call>\n<function=' + tool_call.name + '>\n' }}
118
+ {%- endif %}
119
+ {%- if tool_call.arguments is defined %}
120
+ {%- for args_name, args_value in tool_call.arguments|items %}
121
+ {{- '<parameter=' + args_name + '>\n' }}
122
+ {%- set args_value = args_value | tojson | safe if args_value is mapping or (args_value is sequence and args_value is not string) else args_value | string %}
123
+ {{- args_value }}
124
+ {{- '\n</parameter>\n' }}
125
+ {%- endfor %}
126
+ {%- endif %}
127
+ {{- '</function>\n</tool_call>' }}
128
+ {%- endfor %}
129
+ {%- endif %}
130
+ {{- '<|im_end|>\n' }}
131
+ {%- elif message.role == "tool" %}
132
+ {%- if loop.previtem and loop.previtem.role != "tool" %}
133
+ {{- '<|im_start|>user' }}
134
+ {%- endif %}
135
+ {{- '\n<tool_response>\n' }}
136
+ {{- content }}
137
+ {{- '\n</tool_response>' }}
138
+ {%- if not loop.last and loop.nextitem.role != "tool" %}
139
+ {{- '<|im_end|>\n' }}
140
+ {%- elif loop.last %}
141
+ {{- '<|im_end|>\n' }}
142
+ {%- endif %}
143
+ {%- else %}
144
+ {{- raise_exception('Unexpected message role.') }}
145
+ {%- endif %}
146
+ {%- endfor %}
147
+ {%- if add_generation_prompt %}
148
+ {{- '<|im_start|>assistant\n' }}
149
+ {%- if enable_thinking is defined and enable_thinking is false %}
150
+ {{- '<think>\n\n</think>\n\n' }}
151
+ {%- else %}
152
+ {{- '<think>\n' }}
153
+ {%- endif %}
154
+ {%- endif %}
config.json ADDED
@@ -0,0 +1,133 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "Qwen35GDN24ForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "attn_output_gate": true,
8
+ "auto_map": {
9
+ "AutoConfig": "configuration_qwen35_gdn24.Qwen35GDN24Config",
10
+ "AutoModelForCausalLM": "modeling_qwen35_gdn24.Qwen35GDN24ForCausalLM"
11
+ },
12
+ "bos_token_id": null,
13
+ "converted_layers": [
14
+ 3,
15
+ 7,
16
+ 11,
17
+ 15,
18
+ 19,
19
+ 23,
20
+ 27,
21
+ 31
22
+ ],
23
+ "dtype": "bfloat16",
24
+ "eos_token_id": 248044,
25
+ "full_attention_interval": 4,
26
+ "gdn24_format_version": 1,
27
+ "head_dim": 256,
28
+ "hidden_act": "silu",
29
+ "hidden_size": 4096,
30
+ "initializer_range": 0.02,
31
+ "intermediate_size": 12288,
32
+ "layer_types": [
33
+ "linear_attention",
34
+ "linear_attention",
35
+ "linear_attention",
36
+ "linear_attention",
37
+ "linear_attention",
38
+ "linear_attention",
39
+ "linear_attention",
40
+ "linear_attention",
41
+ "linear_attention",
42
+ "linear_attention",
43
+ "linear_attention",
44
+ "linear_attention",
45
+ "linear_attention",
46
+ "linear_attention",
47
+ "linear_attention",
48
+ "linear_attention",
49
+ "linear_attention",
50
+ "linear_attention",
51
+ "linear_attention",
52
+ "linear_attention",
53
+ "linear_attention",
54
+ "linear_attention",
55
+ "linear_attention",
56
+ "linear_attention",
57
+ "linear_attention",
58
+ "linear_attention",
59
+ "linear_attention",
60
+ "linear_attention",
61
+ "linear_attention",
62
+ "linear_attention",
63
+ "linear_attention",
64
+ "linear_attention"
65
+ ],
66
+ "linear_conv_kernel_dim": 4,
67
+ "linear_key_head_dim": 128,
68
+ "linear_num_key_heads": 16,
69
+ "linear_num_value_heads": 32,
70
+ "linear_value_head_dim": 128,
71
+ "mamba_ssm_dtype": "float32",
72
+ "max_position_embeddings": 262144,
73
+ "mlp_only_layers": [],
74
+ "model_type": "qwen3_5_gdn24",
75
+ "mtp_num_hidden_layers": 1,
76
+ "mtp_use_dedicated_embeddings": false,
77
+ "num_attention_heads": 16,
78
+ "num_hidden_layers": 32,
79
+ "num_key_value_heads": 4,
80
+ "pad_token_id": null,
81
+ "partial_rotary_factor": 0.25,
82
+ "rms_norm_eps": 1e-06,
83
+ "rope_parameters": {
84
+ "mrope_interleaved": true,
85
+ "mrope_section": [
86
+ 11,
87
+ 11,
88
+ 10
89
+ ],
90
+ "partial_rotary_factor": 0.25,
91
+ "rope_theta": 10000000,
92
+ "rope_type": "default"
93
+ },
94
+ "source_layer_types": [
95
+ "linear_attention",
96
+ "linear_attention",
97
+ "linear_attention",
98
+ "full_attention",
99
+ "linear_attention",
100
+ "linear_attention",
101
+ "linear_attention",
102
+ "full_attention",
103
+ "linear_attention",
104
+ "linear_attention",
105
+ "linear_attention",
106
+ "full_attention",
107
+ "linear_attention",
108
+ "linear_attention",
109
+ "linear_attention",
110
+ "full_attention",
111
+ "linear_attention",
112
+ "linear_attention",
113
+ "linear_attention",
114
+ "full_attention",
115
+ "linear_attention",
116
+ "linear_attention",
117
+ "linear_attention",
118
+ "full_attention",
119
+ "linear_attention",
120
+ "linear_attention",
121
+ "linear_attention",
122
+ "full_attention",
123
+ "linear_attention",
124
+ "linear_attention",
125
+ "linear_attention",
126
+ "full_attention"
127
+ ],
128
+ "source_model": "Qwen/Qwen3.5-9B",
129
+ "tie_word_embeddings": false,
130
+ "transformers_version": "5.16.0.dev0",
131
+ "use_cache": true,
132
+ "vocab_size": 248320
133
+ }
configuration_qwen35_gdn24.py ADDED
@@ -0,0 +1,93 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from transformers.models.qwen3_5.configuration_qwen3_5 import Qwen3_5TextConfig
4
+
5
+
6
+ class Qwen35GDN24Config(Qwen3_5TextConfig):
7
+ """Qwen3.5 text config whose runtime token mixer is Gated DeltaNet in all layers.
8
+
9
+ ``source_layer_types`` preserves the official/source Qwen3.5 topology while
10
+ ``layer_types`` is always the homogeneous runtime topology. The runtime
11
+ topology is normalized *before* calling the Transformers base config so
12
+ base-class post-init/validation can never restore or observe the source
13
+ ``full_attention`` entries as runtime layers.
14
+ """
15
+
16
+ model_type = "qwen3_5_gdn24"
17
+
18
+ def __init__(
19
+ self,
20
+ *args,
21
+ source_layer_types: list[str] | None = None,
22
+ converted_layers: list[int] | None = None,
23
+ source_model: str = "Qwen/Qwen3.5-9B",
24
+ gdn24_format_version: int = 1,
25
+ **kwargs,
26
+ ):
27
+ # Keep the serialized/incoming topology separate from the runtime one.
28
+ # On a fresh conversion this is the official hybrid topology. On
29
+ # reload, source_layer_types is serialized explicitly and therefore
30
+ # remains authoritative even though layer_types is already all-linear.
31
+ incoming_layer_types = kwargs.pop("layer_types", None)
32
+ if source_layer_types is None:
33
+ source_layer_types = incoming_layer_types
34
+
35
+ if source_layer_types is None:
36
+ # Preserve the base model's default topology only as conversion
37
+ # metadata. Qwen3.5 defaults to full attention every fourth layer.
38
+ n = int(kwargs.get("num_hidden_layers", len(source_layer_types) if source_layer_types is not None else 32))
39
+ interval = int(kwargs.get("full_attention_interval", 4))
40
+ if interval > 0:
41
+ source_layer_types = [
42
+ "linear_attention" if (i + 1) % interval else "full_attention"
43
+ for i in range(n)
44
+ ]
45
+ else:
46
+ source_layer_types = ["linear_attention"] * n
47
+ else:
48
+ source_layer_types = list(source_layer_types)
49
+ n = int(kwargs.get("num_hidden_layers", len(source_layer_types)))
50
+
51
+ if len(source_layer_types) != n:
52
+ raise ValueError(
53
+ f"source_layer_types length ({len(source_layer_types)}) must equal "
54
+ f"num_hidden_layers ({n})"
55
+ )
56
+ if any(t not in {"linear_attention", "full_attention"} for t in source_layer_types):
57
+ raise ValueError(f"unsupported source layer type: {source_layer_types}")
58
+
59
+ expected = [i for i, t in enumerate(source_layer_types) if t == "full_attention"]
60
+ if converted_layers is None:
61
+ converted_layers = expected
62
+ normalized_converted = sorted(int(i) for i in converted_layers)
63
+ if normalized_converted != expected:
64
+ raise ValueError(
65
+ f"converted_layers must equal all source full-attention layers: {expected}"
66
+ )
67
+
68
+ # Critical fix: normalize BEFORE base config construction. Newer
69
+ # Transformers Qwen3.5 configs perform layer-type handling in
70
+ # __post_init__, so fixing the value only after super().__init__ is too
71
+ # late and is version-sensitive.
72
+ kwargs["layer_types"] = ["linear_attention"] * n
73
+ super().__init__(*args, **kwargs)
74
+
75
+ # Assert the invariant explicitly instead of silently shipping a hybrid
76
+ # runtime config if upstream Transformers changes behavior again.
77
+ self.layer_types = ["linear_attention"] * int(self.num_hidden_layers)
78
+ if len(self.layer_types) != n or set(self.layer_types) != {"linear_attention"}:
79
+ raise RuntimeError("failed to normalize GDN24 runtime layer_types")
80
+
81
+ self.source_layer_types = source_layer_types
82
+ self.converted_layers = normalized_converted
83
+ self.source_model = str(source_model)
84
+ self.gdn24_format_version = int(gdn24_format_version)
85
+
86
+ self.architectures = ["Qwen35GDN24ForCausalLM"]
87
+ self.auto_map = {
88
+ "AutoConfig": "configuration_qwen35_gdn24.Qwen35GDN24Config",
89
+ "AutoModelForCausalLM": "modeling_qwen35_gdn24.Qwen35GDN24ForCausalLM",
90
+ }
91
+
92
+
93
+ Qwen35GDN24Config.register_for_auto_class("AutoConfig")
engine.py ADDED
@@ -0,0 +1,430 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ """Persistent CUDA-Graph greedy decoder for the all-recurrent GDN24 runtime.
4
+
5
+ The core idea is specific to a fully recurrent model: the cache has fixed tensor
6
+ addresses and token position is implicit in state evolution. We can therefore
7
+ capture multiple autoregressive steps into one CUDA Graph and let the graph feed
8
+ its own argmax token into the next step.
9
+ """
10
+
11
+ from dataclasses import dataclass
12
+ import time
13
+ from typing import Any
14
+
15
+ import torch
16
+
17
+
18
+ @dataclass
19
+ class _LayerSnapshot:
20
+ conv: dict[int, torch.Tensor]
21
+ recurrent: dict[int, torch.Tensor]
22
+ has_previous_state: dict[int, bool]
23
+ conv_initialized: dict[int, bool]
24
+ recurrent_initialized: dict[int, bool]
25
+
26
+
27
+ @dataclass
28
+ class CacheSnapshot:
29
+ seen_tokens: int
30
+ layers: list[_LayerSnapshot]
31
+
32
+
33
+ def snapshot_cache(cache) -> CacheSnapshot:
34
+ layers: list[_LayerSnapshot] = []
35
+ for layer in cache.layers:
36
+ conv = {int(i): t.detach().clone() for i, t in layer.conv_states.items() if t is not None}
37
+ recurrent = {
38
+ int(i): t.detach().clone() for i, t in layer.recurrent_states.items() if t is not None
39
+ }
40
+ layers.append(
41
+ _LayerSnapshot(
42
+ conv=conv,
43
+ recurrent=recurrent,
44
+ has_previous_state=dict(layer.has_previous_state),
45
+ conv_initialized=dict(layer.is_conv_states_initialized),
46
+ recurrent_initialized=dict(layer.is_recurrent_states_initialized),
47
+ )
48
+ )
49
+ return CacheSnapshot(seen_tokens=int(cache.seen_tokens), layers=layers)
50
+
51
+
52
+ def restore_cache_(cache, snap: CacheSnapshot) -> None:
53
+ if len(cache.layers) != len(snap.layers):
54
+ raise ValueError("cache topology changed while restoring snapshot")
55
+ for layer, state in zip(cache.layers, snap.layers):
56
+ for i, src in state.conv.items():
57
+ dst = layer.conv_states[i]
58
+ if dst is None or dst.shape != src.shape:
59
+ raise ValueError(f"conv state storage changed at state {i}")
60
+ dst.copy_(src)
61
+ for i, src in state.recurrent.items():
62
+ dst = layer.recurrent_states[i]
63
+ if dst is None or dst.shape != src.shape:
64
+ raise ValueError(f"recurrent state storage changed at state {i}")
65
+ dst.copy_(src)
66
+ layer.has_previous_state.update(state.has_previous_state)
67
+ layer.is_conv_states_initialized.update(state.conv_initialized)
68
+ layer.is_recurrent_states_initialized.update(state.recurrent_initialized)
69
+ cache.seen_tokens = int(snap.seen_tokens)
70
+
71
+
72
+ class SuperTurboGraphDecoder:
73
+ """Persistent, self-feeding greedy CUDA Graph decoder.
74
+
75
+ A single replay advances `block_size` recurrent decode steps. The graph owns
76
+ one persistent recurrent cache, so it can be reused across prompts: `reset()`
77
+ preserves all CUDA addresses, prefill writes the new prompt state, then graph
78
+ replay continues autoregressively from that state.
79
+ """
80
+
81
+ def __init__(
82
+ self,
83
+ model,
84
+ *,
85
+ block_size: int = 8,
86
+ warmup_steps: int = 2,
87
+ graph_pool=None,
88
+ ):
89
+ if not torch.cuda.is_available():
90
+ raise RuntimeError("MAX TURBO CUDA Graph decoding requires CUDA")
91
+ if block_size < 1:
92
+ raise ValueError("block_size must be >= 1")
93
+ self.model = model.eval()
94
+ self.block_size = int(block_size)
95
+ self.warmup_steps = int(warmup_steps)
96
+ self.graph_pool = graph_pool
97
+ self.device = model.get_input_embeddings().weight.device
98
+ if self.device.type != "cuda":
99
+ raise RuntimeError(f"model must be on CUDA, got {self.device}")
100
+
101
+ self.cache = self.model.make_recurrent_cache()
102
+ self.static_token = torch.zeros((1, 1), dtype=torch.long, device=self.device)
103
+ self.output_tokens = torch.empty((1, self.block_size), dtype=torch.long, device=self.device)
104
+ self.graph: torch.cuda.CUDAGraph | None = None
105
+ self.capture_seconds: float | None = None
106
+ self.capture_error: str | None = None
107
+ self.capture_allocated_delta_mib: float = 0.0
108
+ self.capture_reserved_delta_mib: float = 0.0
109
+ self._captured = False
110
+ self._capture_attempted = False
111
+
112
+ @torch.inference_mode()
113
+ def _greedy_step(self) -> torch.Tensor:
114
+ # Private clock flag is consumed by Qwen35GDN24Model and never reaches GDN kernels.
115
+ return self.model.greedy_step(
116
+ self.static_token,
117
+ self.cache,
118
+ advance_cache_clock=False,
119
+ )
120
+
121
+ @torch.inference_mode()
122
+ def _unrolled_block(self) -> None:
123
+ for i in range(self.block_size):
124
+ next_token = self._greedy_step()
125
+ self.output_tokens[:, i : i + 1].copy_(next_token)
126
+ self.static_token.copy_(next_token)
127
+
128
+ @torch.inference_mode()
129
+ def capture(self) -> bool:
130
+ if self._captured:
131
+ return True
132
+ if self._capture_attempted:
133
+ return False
134
+ self._capture_attempted = True
135
+ try:
136
+ # Initialize every native LinearAttentionLayer cache with stable addresses.
137
+ self.cache.reset()
138
+ dummy = torch.zeros((1, 1), dtype=torch.long, device=self.device)
139
+ first = self.model(
140
+ input_ids=dummy,
141
+ past_key_values=self.cache,
142
+ use_cache=True,
143
+ logits_to_keep=1,
144
+ ).logits[:, -1, :].argmax(dim=-1, keepdim=True)
145
+ self.static_token.copy_(first)
146
+ torch.cuda.synchronize()
147
+
148
+ baseline = snapshot_cache(self.cache)
149
+ baseline_token = self.static_token.detach().clone()
150
+
151
+ # Warm single-token recurrent kernels and Triton fusions before capture.
152
+ for _ in range(max(self.warmup_steps, 0)):
153
+ self._unrolled_block()
154
+ torch.cuda.synchronize()
155
+ restore_cache_(self.cache, baseline)
156
+ self.static_token.copy_(baseline_token)
157
+
158
+ graph = torch.cuda.CUDAGraph()
159
+ torch.cuda.synchronize()
160
+ alloc_before = torch.cuda.memory_allocated(self.device)
161
+ reserve_before = torch.cuda.memory_reserved(self.device)
162
+ t0 = time.perf_counter()
163
+ graph_kwargs = {} if self.graph_pool is None else {"pool": self.graph_pool}
164
+ with torch.cuda.graph(graph, **graph_kwargs):
165
+ self._unrolled_block()
166
+ torch.cuda.synchronize()
167
+ self.capture_seconds = time.perf_counter() - t0
168
+ self.capture_allocated_delta_mib = max(0, torch.cuda.memory_allocated(self.device) - alloc_before) / 2**20
169
+ self.capture_reserved_delta_mib = max(0, torch.cuda.memory_reserved(self.device) - reserve_before) / 2**20
170
+
171
+ # Capture executes once. Restore the exact pre-capture numerical state
172
+ # without replacing any tensor objects/addresses recorded by the graph.
173
+ restore_cache_(self.cache, baseline)
174
+ self.static_token.copy_(baseline_token)
175
+ if hasattr(graph, "instantiate"):
176
+ try:
177
+ graph.instantiate()
178
+ except Exception:
179
+ pass
180
+ self.graph = graph
181
+ self._captured = True
182
+ self.cache.reset()
183
+ return True
184
+ except Exception as exc:
185
+ self.capture_error = f"{type(exc).__name__}: {exc}"
186
+ self.graph = None
187
+ self._captured = False
188
+ try:
189
+ self.cache.reset()
190
+ except Exception:
191
+ pass
192
+ return False
193
+
194
+ @torch.inference_mode()
195
+ def reset(self) -> None:
196
+ self.cache.reset()
197
+
198
+ @torch.inference_mode()
199
+ def prefill(self, input_ids: torch.Tensor) -> tuple[torch.Tensor, Any]:
200
+ """Prefill into the persistent graph cache and return the first next token."""
201
+ self.cache.reset()
202
+ out = self.model(
203
+ input_ids=input_ids,
204
+ past_key_values=self.cache,
205
+ use_cache=True,
206
+ logits_to_keep=1,
207
+ )
208
+ token = torch.argmax(out.logits[:, -1, :], dim=-1, keepdim=True)
209
+ return token, out
210
+
211
+ @torch.inference_mode()
212
+ def decode_forwards(self, first_token: torch.Tensor, steps: int) -> torch.Tensor:
213
+ """Run `steps` recurrent forward passes, matching the benchmark's decode metric.
214
+
215
+ The returned tensor contains the last generated token. For the fast path,
216
+ `steps` should be divisible by `block_size`; a short eager tail handles any
217
+ remainder.
218
+ """
219
+ steps = int(steps)
220
+ if steps < 0:
221
+ raise ValueError("steps must be >= 0")
222
+ if steps == 0:
223
+ return first_token
224
+ if not self._captured:
225
+ return self._decode_eager(first_token, steps)
226
+
227
+ self.static_token.copy_(first_token)
228
+ full_blocks, remainder = divmod(steps, self.block_size)
229
+ for _ in range(full_blocks):
230
+ self.graph.replay()
231
+ # The graph deliberately does not update the Python token clock.
232
+ self.cache.advance(full_blocks * self.block_size)
233
+
234
+ if remainder:
235
+ token = self.static_token
236
+ for _ in range(remainder):
237
+ out = self.model(
238
+ input_ids=token,
239
+ past_key_values=self.cache,
240
+ use_cache=True,
241
+ logits_to_keep=1,
242
+ )
243
+ token = torch.argmax(out.logits[:, -1, :], dim=-1, keepdim=True)
244
+ self.static_token.copy_(token)
245
+ return self.static_token
246
+
247
+ @torch.inference_mode()
248
+ def _decode_eager(self, first_token: torch.Tensor, steps: int) -> torch.Tensor:
249
+ token = first_token
250
+ for _ in range(steps):
251
+ out = self.model(
252
+ input_ids=token,
253
+ past_key_values=self.cache,
254
+ use_cache=True,
255
+ logits_to_keep=1,
256
+ )
257
+ token = torch.argmax(out.logits[:, -1, :], dim=-1, keepdim=True)
258
+ return token
259
+
260
+ @torch.inference_mode()
261
+ def decode_tokens(self, first_token: torch.Tensor, steps: int) -> torch.Tensor:
262
+ """Return every greedy token produced by ``steps`` recurrent forwards.
263
+
264
+ Unlike :meth:`decode_forwards`, this materializes the token sequence and is
265
+ intended for user generation, quality checks, and UNI MAX verification.
266
+ The persistent recurrent cache remains graph-address-stable.
267
+ """
268
+ steps = int(steps)
269
+ if steps < 0:
270
+ raise ValueError("steps must be >= 0")
271
+ if steps == 0:
272
+ return torch.empty((first_token.shape[0], 0), dtype=torch.long, device=first_token.device)
273
+
274
+ self.static_token.copy_(first_token)
275
+ pieces: list[torch.Tensor] = []
276
+ if self._captured:
277
+ full_blocks, remainder = divmod(steps, self.block_size)
278
+ for _ in range(full_blocks):
279
+ self.graph.replay()
280
+ pieces.append(self.output_tokens.detach().clone())
281
+ if full_blocks:
282
+ self.cache.advance(full_blocks * self.block_size)
283
+ token = self.static_token
284
+ else:
285
+ remainder = steps
286
+ token = first_token
287
+
288
+ for _ in range(remainder):
289
+ token = self.model.greedy_step(token, self.cache)
290
+ pieces.append(token.detach().clone())
291
+ if pieces:
292
+ self.static_token.copy_(pieces[-1][:, -1:])
293
+ return torch.cat(pieces, dim=1)
294
+
295
+ @torch.inference_mode()
296
+ def replay_block_tokens(self, first_token: torch.Tensor) -> torch.Tensor:
297
+ """Replay exactly one captured block and return its generated tokens."""
298
+ if not self._captured:
299
+ if not self.capture():
300
+ raise RuntimeError(f"CUDA Graph capture failed: {self.capture_error}")
301
+ self.static_token.copy_(first_token)
302
+ self.graph.replay()
303
+ self.cache.advance(self.block_size)
304
+ return self.output_tokens.detach().clone()
305
+
306
+ @torch.inference_mode()
307
+ def validate_against_eager(self, first_token: torch.Tensor) -> bool:
308
+ """Bit-exact greedy-token check for one captured block."""
309
+ if not self._captured and not self.capture():
310
+ raise RuntimeError(f"CUDA Graph capture failed: {self.capture_error}")
311
+ snap = snapshot_cache(self.cache)
312
+ start = first_token.detach().clone()
313
+ try:
314
+ token = start
315
+ eager_tokens = []
316
+ for _ in range(self.block_size):
317
+ token = self.model.greedy_step(
318
+ token, self.cache, advance_cache_clock=False
319
+ )
320
+ eager_tokens.append(token.detach().clone())
321
+ eager = torch.cat(eager_tokens, dim=1)
322
+
323
+ restore_cache_(self.cache, snap)
324
+ self.static_token.copy_(start)
325
+ self.graph.replay()
326
+ torch.cuda.synchronize()
327
+ graphed = self.output_tokens.detach().clone()
328
+ if not torch.equal(eager, graphed):
329
+ raise RuntimeError(
330
+ f"CUDA Graph greedy mismatch: eager={eager.tolist()} graph={graphed.tolist()}"
331
+ )
332
+ return True
333
+ finally:
334
+ restore_cache_(self.cache, snap)
335
+ self.static_token.copy_(start)
336
+
337
+ @property
338
+ def ready(self) -> bool:
339
+ return self._captured
340
+
341
+
342
+ @dataclass(frozen=True)
343
+ class GraphTuneResult:
344
+ block_size: int
345
+ decode_tok_s: float
346
+ capture_seconds: float
347
+ graph_allocated_mib: float = 0.0
348
+ graph_reserved_mib: float = 0.0
349
+
350
+
351
+ @torch.inference_mode()
352
+ def autotune_graph_decoder(
353
+ model,
354
+ *,
355
+ candidates: tuple[int, ...] = (4, 8, 16),
356
+ warmup_replays: int = 2,
357
+ timed_replays: int = 6,
358
+ ) -> tuple[SuperTurboGraphDecoder, list[GraphTuneResult]]:
359
+ """Capture and benchmark several self-feeding graph block sizes.
360
+
361
+ The search measures steady-state graph replay only; graph capture time is
362
+ reported separately and excluded. Every candidate is numerically checked
363
+ against eager greedy decoding before it can win.
364
+ """
365
+ if not candidates:
366
+ raise ValueError("at least one graph block candidate is required")
367
+ results: list[GraphTuneResult] = []
368
+ best: SuperTurboGraphDecoder | None = None
369
+ best_speed = -1.0
370
+ failures: list[str] = []
371
+ shared_pool = torch.cuda.graph_pool_handle() if hasattr(torch.cuda, "graph_pool_handle") else None
372
+
373
+ for raw_block in candidates:
374
+ block = int(raw_block)
375
+ if block < 1:
376
+ failures.append(f"x{block}: invalid block size")
377
+ continue
378
+ engine = SuperTurboGraphDecoder(model, block_size=block, warmup_steps=2, graph_pool=shared_pool)
379
+ if not engine.capture():
380
+ failures.append(f"x{block}: {engine.capture_error}")
381
+ continue
382
+ # Start from a valid persistent recurrent state, then verify exact greedy
383
+ # token agreement before any timing result is accepted.
384
+ seed = torch.zeros((1, 8), dtype=torch.long, device=engine.device)
385
+ first, _ = engine.prefill(seed)
386
+ engine.validate_against_eager(first)
387
+ engine.reset()
388
+ first, _ = engine.prefill(seed)
389
+ engine.static_token.copy_(first)
390
+
391
+ for _ in range(max(int(warmup_replays), 0)):
392
+ engine.graph.replay()
393
+ torch.cuda.synchronize()
394
+
395
+ t0 = time.perf_counter()
396
+ for _ in range(max(int(timed_replays), 1)):
397
+ engine.graph.replay()
398
+ torch.cuda.synchronize()
399
+ dt = time.perf_counter() - t0
400
+ forwards = block * max(int(timed_replays), 1)
401
+ speed = forwards / dt
402
+ result = GraphTuneResult(
403
+ block_size=block,
404
+ decode_tok_s=float(speed),
405
+ capture_seconds=float(engine.capture_seconds or 0.0),
406
+ graph_allocated_mib=float(engine.capture_allocated_delta_mib),
407
+ graph_reserved_mib=float(engine.capture_reserved_delta_mib),
408
+ )
409
+ results.append(result)
410
+ if speed > best_speed:
411
+ best_speed = speed
412
+ previous_best = best
413
+ best = engine
414
+ if previous_best is not None:
415
+ del previous_best
416
+ torch.cuda.empty_cache()
417
+ else:
418
+ # Drop non-winning graph pools as early as possible.
419
+ del engine
420
+ torch.cuda.empty_cache()
421
+
422
+ if best is None:
423
+ reason = "; ".join(failures) if failures else "no valid candidates"
424
+ raise RuntimeError(f"MAX TURBO graph autotune failed: {reason}")
425
+ best.reset()
426
+ return best, results
427
+
428
+
429
+ # MAX TURBO public name; keep the v3 class name as a compatibility alias.
430
+ MaxTurboGraphDecoder = SuperTurboGraphDecoder
fused_ops.py ADDED
@@ -0,0 +1,220 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ """Small fused CUDA/Triton operators for GDN24 MAX TURBO.
4
+
5
+ The runtime deliberately keeps a pure-Torch fallback so checkpoints remain
6
+ portable. Triton is used only when it is already available through the CUDA
7
+ PyTorch stack; it is not a checkpoint dependency.
8
+ """
9
+
10
+ import torch
11
+ from torch import nn
12
+ import torch.nn.functional as F
13
+
14
+ try: # Triton ships with CUDA PyTorch builds used by Colab.
15
+ import triton
16
+ import triton.language as tl
17
+ _HAS_TRITON = True
18
+ except Exception: # pragma: no cover - CPU/source validation path
19
+ triton = None
20
+ tl = None
21
+ _HAS_TRITON = False
22
+
23
+
24
+ if _HAS_TRITON:
25
+ @triton.jit
26
+ def _rmsnorm_kernel(x_ptr, w_ptr, y_ptr, n_cols: tl.constexpr, eps: tl.constexpr, BLOCK: tl.constexpr):
27
+ row = tl.program_id(0)
28
+ offs = tl.arange(0, BLOCK)
29
+ mask = offs < n_cols
30
+ x = tl.load(x_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32)
31
+ w = tl.load(w_ptr + offs, mask=mask, other=0.0).to(tl.float32)
32
+ var = tl.sum(x * x, axis=0) / n_cols
33
+ rstd = tl.rsqrt(var + eps)
34
+ y = x * rstd * (1.0 + w)
35
+ tl.store(y_ptr + row * n_cols + offs, y, mask=mask)
36
+
37
+
38
+ @triton.jit
39
+ def _add_rmsnorm_kernel(
40
+ x_ptr,
41
+ update_ptr,
42
+ w_ptr,
43
+ sum_ptr,
44
+ norm_ptr,
45
+ n_cols: tl.constexpr,
46
+ eps: tl.constexpr,
47
+ BLOCK: tl.constexpr,
48
+ ):
49
+ row = tl.program_id(0)
50
+ offs = tl.arange(0, BLOCK)
51
+ mask = offs < n_cols
52
+ x = tl.load(x_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32)
53
+ u = tl.load(update_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32)
54
+ w = tl.load(w_ptr + offs, mask=mask, other=0.0).to(tl.float32)
55
+ s = x + u
56
+ var = tl.sum(s * s, axis=0) / n_cols
57
+ rstd = tl.rsqrt(var + eps)
58
+ n = s * rstd * (1.0 + w)
59
+ tl.store(sum_ptr + row * n_cols + offs, s, mask=mask)
60
+ tl.store(norm_ptr + row * n_cols + offs, n, mask=mask)
61
+
62
+
63
+ @triton.jit
64
+ def _add_final_rmsnorm_kernel(
65
+ x_ptr,
66
+ update_ptr,
67
+ w_ptr,
68
+ norm_ptr,
69
+ n_cols: tl.constexpr,
70
+ eps: tl.constexpr,
71
+ BLOCK: tl.constexpr,
72
+ ):
73
+ row = tl.program_id(0)
74
+ offs = tl.arange(0, BLOCK)
75
+ mask = offs < n_cols
76
+ x = tl.load(x_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32)
77
+ u = tl.load(update_ptr + row * n_cols + offs, mask=mask, other=0.0).to(tl.float32)
78
+ w = tl.load(w_ptr + offs, mask=mask, other=0.0).to(tl.float32)
79
+ s = x + u
80
+ var = tl.sum(s * s, axis=0) / n_cols
81
+ rstd = tl.rsqrt(var + eps)
82
+ n = s * rstd * (1.0 + w)
83
+ tl.store(norm_ptr + row * n_cols + offs, n, mask=mask)
84
+
85
+
86
+ @triton.jit
87
+ def _silu_mul_kernel(a_ptr, b_ptr, out_ptr, n_elements, BLOCK: tl.constexpr):
88
+ offs = tl.program_id(0) * BLOCK + tl.arange(0, BLOCK)
89
+ mask = offs < n_elements
90
+ a = tl.load(a_ptr + offs, mask=mask, other=0.0).to(tl.float32)
91
+ b = tl.load(b_ptr + offs, mask=mask, other=0.0).to(tl.float32)
92
+ # SiLU(a) = a * sigmoid(a)
93
+ sig = 1.0 / (1.0 + tl.exp(-a))
94
+ out = a * sig * b
95
+ tl.store(out_ptr + offs, out, mask=mask)
96
+
97
+
98
+ def _can_triton(x: torch.Tensor) -> bool:
99
+ return bool(_HAS_TRITON and x.is_cuda and x.is_contiguous())
100
+
101
+
102
+ def qwen_rmsnorm(x: torch.Tensor, weight: torch.Tensor, eps: float) -> torch.Tensor:
103
+ """Exact Qwen3.5 offset RMSNorm: norm(x) * (1 + weight)."""
104
+ d = x.shape[-1]
105
+ if _can_triton(x) and weight.is_cuda and weight.is_contiguous():
106
+ y = torch.empty_like(x)
107
+ x2 = x.view(-1, d)
108
+ y2 = y.view(-1, d)
109
+ block = triton.next_power_of_2(d)
110
+ _rmsnorm_kernel[(x2.shape[0],)](x2, weight, y2, n_cols=d, eps=float(eps), BLOCK=block)
111
+ return y
112
+ xf = x.float()
113
+ y = xf * torch.rsqrt(xf.square().mean(dim=-1, keepdim=True) + eps)
114
+ y = y * (1.0 + weight.float())
115
+ return y.to(dtype=x.dtype)
116
+
117
+
118
+ def add_rmsnorm(
119
+ x: torch.Tensor,
120
+ update: torch.Tensor,
121
+ weight: torch.Tensor,
122
+ eps: float,
123
+ ) -> tuple[torch.Tensor, torch.Tensor]:
124
+ """Fuse residual addition with Qwen3.5 offset RMSNorm.
125
+
126
+ Returns `(x + update, rmsnorm(x + update))`.
127
+ """
128
+ d = x.shape[-1]
129
+ if _can_triton(x) and update.is_contiguous() and weight.is_cuda and weight.is_contiguous():
130
+ summed = torch.empty_like(x)
131
+ normed = torch.empty_like(x)
132
+ x2 = x.view(-1, d)
133
+ u2 = update.view(-1, d)
134
+ s2 = summed.view(-1, d)
135
+ n2 = normed.view(-1, d)
136
+ block = triton.next_power_of_2(d)
137
+ _add_rmsnorm_kernel[(x2.shape[0],)](
138
+ x2, u2, weight, s2, n2, n_cols=d, eps=float(eps), BLOCK=block
139
+ )
140
+ return summed, normed
141
+ summed = x + update
142
+ xf = summed.float()
143
+ normed = xf * torch.rsqrt(xf.square().mean(dim=-1, keepdim=True) + eps)
144
+ normed = normed * (1.0 + weight.float())
145
+ return summed, normed.to(dtype=x.dtype)
146
+
147
+
148
+ def add_final_rmsnorm(
149
+ x: torch.Tensor,
150
+ update: torch.Tensor,
151
+ weight: torch.Tensor,
152
+ eps: float,
153
+ ) -> torch.Tensor:
154
+ """Fuse final residual addition and RMSNorm when the unnormalized sum is not needed."""
155
+ d = x.shape[-1]
156
+ if _can_triton(x) and update.is_contiguous() and weight.is_cuda and weight.is_contiguous():
157
+ normed = torch.empty_like(x)
158
+ x2 = x.view(-1, d)
159
+ u2 = update.view(-1, d)
160
+ n2 = normed.view(-1, d)
161
+ block = triton.next_power_of_2(d)
162
+ _add_final_rmsnorm_kernel[(x2.shape[0],)](
163
+ x2, u2, weight, n2, n_cols=d, eps=float(eps), BLOCK=block
164
+ )
165
+ return normed
166
+ summed = x + update
167
+ xf = summed.float()
168
+ normed = xf * torch.rsqrt(xf.square().mean(dim=-1, keepdim=True) + eps)
169
+ normed = normed * (1.0 + weight.float())
170
+ return normed.to(dtype=x.dtype)
171
+
172
+
173
+ def silu_mul(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
174
+ """Fused SwiGLU pointwise core: SiLU(a) * b."""
175
+ if _can_triton(a) and b.is_contiguous() and a.shape == b.shape:
176
+ out = torch.empty_like(a)
177
+ n = a.numel()
178
+ block = 256
179
+ _silu_mul_kernel[(triton.cdiv(n, block),)](a, b, out, n_elements=n, BLOCK=block)
180
+ return out
181
+ return F.silu(a) * b
182
+
183
+
184
+ class SuperRMSNorm(nn.Module):
185
+ """Checkpoint-compatible replacement for Qwen3_5RMSNorm."""
186
+
187
+ def __init__(self, dim: int, eps: float = 1e-6):
188
+ super().__init__()
189
+ self.eps = float(eps)
190
+ # Qwen3.5 uses zero-centered weights and multiplies by (1 + weight).
191
+ self.weight = nn.Parameter(torch.zeros(dim))
192
+
193
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
194
+ return qwen_rmsnorm(x, self.weight, self.eps)
195
+
196
+ def extra_repr(self) -> str:
197
+ return f"{tuple(self.weight.shape)}, eps={self.eps}"
198
+
199
+
200
+ class SuperSwiGLUMLP(nn.Module):
201
+ """Checkpoint-compatible Qwen3.5 dense MLP with a fused SwiGLU pointwise kernel."""
202
+
203
+ def __init__(self, config):
204
+ super().__init__()
205
+ self.hidden_size = config.hidden_size
206
+ self.intermediate_size = config.intermediate_size
207
+ self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
208
+ self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False)
209
+ self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False)
210
+ self.hidden_act = str(config.hidden_act)
211
+
212
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
213
+ gate = self.gate_proj(x)
214
+ up = self.up_proj(x)
215
+ if self.hidden_act == "silu":
216
+ hidden = silu_mul(gate, up)
217
+ else:
218
+ from transformers.activations import ACT2FN
219
+ hidden = ACT2FN[self.hidden_act](gate) * up
220
+ return self.down_proj(hidden)
gdn24_metadata.json ADDED
@@ -0,0 +1,26 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "source_model": "Qwen/Qwen3.5-9B",
3
+ "architecture": "32x native Qwen3.5 GatedDeltaNet",
4
+ "converted_layers": [
5
+ 3,
6
+ 7,
7
+ 11,
8
+ 15,
9
+ 19,
10
+ 23,
11
+ 27,
12
+ 31
13
+ ],
14
+ "steps_per_converted_layer": 20,
15
+ "seq_len": 128,
16
+ "lr": 0.0001,
17
+ "seed": 1234,
18
+ "runtime_adapters": 0,
19
+ "uni_max_runtime": 1,
20
+ "phase_policy": "exact_bf16_prefill+autotuned_decode",
21
+ "fusions": [
22
+ "residual+rmsnorm",
23
+ "swiglu"
24
+ ],
25
+ "cuda_graph_decode_compatible": true
26
+ }
generation_config.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "eos_token_id": 248044,
4
+ "output_attentions": false,
5
+ "output_hidden_states": false,
6
+ "transformers_version": "5.16.0.dev0",
7
+ "use_cache": true
8
+ }
model-00001-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:51ca0c41567cc617239c12d88281ca02552f6fffe0846ecd1e81f83f08d9b72b
3
+ size 4991268352
model-00002-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:84d4ac920ebc03c163d7347a18c9c43e26f583b0c078e44cc90516fa76a12251
3
+ size 4939275072
model-00003-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c7b24fdcdf6f78f754e07f6335f3fa7dbdc651bb33002aa943158ff0c1f634dd
3
+ size 4906162280
model-00004-of-00004.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c7bdd2633dc3b0f8ba6434550d3aab1f0d44dffb3ea88dffdf2517809d1a9cf8
3
+ size 3209884384
model.safetensors.index.json ADDED
@@ -0,0 +1,459 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "metadata": {
3
+ "total_parameters": 9023268864,
4
+ "total_size": 18046537728
5
+ },
6
+ "weight_map": {
7
+ "lm_head.weight": "model-00004-of-00004.safetensors",
8
+ "model.embed_tokens.weight": "model-00001-of-00004.safetensors",
9
+ "model.layers.0.input_layernorm.weight": "model-00001-of-00004.safetensors",
10
+ "model.layers.0.linear_attn.A_log": "model-00001-of-00004.safetensors",
11
+ "model.layers.0.linear_attn.conv1d.weight": "model-00001-of-00004.safetensors",
12
+ "model.layers.0.linear_attn.dt_bias": "model-00001-of-00004.safetensors",
13
+ "model.layers.0.linear_attn.in_proj_a.weight": "model-00001-of-00004.safetensors",
14
+ "model.layers.0.linear_attn.in_proj_b.weight": "model-00001-of-00004.safetensors",
15
+ "model.layers.0.linear_attn.in_proj_qkv.weight": "model-00001-of-00004.safetensors",
16
+ "model.layers.0.linear_attn.in_proj_z.weight": "model-00001-of-00004.safetensors",
17
+ "model.layers.0.linear_attn.norm.weight": "model-00001-of-00004.safetensors",
18
+ "model.layers.0.linear_attn.out_proj.weight": "model-00001-of-00004.safetensors",
19
+ "model.layers.0.mlp.down_proj.weight": "model-00001-of-00004.safetensors",
20
+ "model.layers.0.mlp.gate_proj.weight": "model-00001-of-00004.safetensors",
21
+ "model.layers.0.mlp.up_proj.weight": "model-00001-of-00004.safetensors",
22
+ "model.layers.0.post_attention_layernorm.weight": "model-00001-of-00004.safetensors",
23
+ "model.layers.1.input_layernorm.weight": "model-00001-of-00004.safetensors",
24
+ "model.layers.1.linear_attn.A_log": "model-00001-of-00004.safetensors",
25
+ "model.layers.1.linear_attn.conv1d.weight": "model-00001-of-00004.safetensors",
26
+ "model.layers.1.linear_attn.dt_bias": "model-00001-of-00004.safetensors",
27
+ "model.layers.1.linear_attn.in_proj_a.weight": "model-00001-of-00004.safetensors",
28
+ "model.layers.1.linear_attn.in_proj_b.weight": "model-00001-of-00004.safetensors",
29
+ "model.layers.1.linear_attn.in_proj_qkv.weight": "model-00001-of-00004.safetensors",
30
+ "model.layers.1.linear_attn.in_proj_z.weight": "model-00001-of-00004.safetensors",
31
+ "model.layers.1.linear_attn.norm.weight": "model-00001-of-00004.safetensors",
32
+ "model.layers.1.linear_attn.out_proj.weight": "model-00001-of-00004.safetensors",
33
+ "model.layers.1.mlp.down_proj.weight": "model-00001-of-00004.safetensors",
34
+ "model.layers.1.mlp.gate_proj.weight": "model-00001-of-00004.safetensors",
35
+ "model.layers.1.mlp.up_proj.weight": "model-00001-of-00004.safetensors",
36
+ "model.layers.1.post_attention_layernorm.weight": "model-00001-of-00004.safetensors",
37
+ "model.layers.10.input_layernorm.weight": "model-00002-of-00004.safetensors",
38
+ "model.layers.10.linear_attn.A_log": "model-00002-of-00004.safetensors",
39
+ "model.layers.10.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
40
+ "model.layers.10.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
41
+ "model.layers.10.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
42
+ "model.layers.10.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
43
+ "model.layers.10.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
44
+ "model.layers.10.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
45
+ "model.layers.10.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
46
+ "model.layers.10.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
47
+ "model.layers.10.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
48
+ "model.layers.10.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
49
+ "model.layers.10.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
50
+ "model.layers.10.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
51
+ "model.layers.11.input_layernorm.weight": "model-00002-of-00004.safetensors",
52
+ "model.layers.11.linear_attn.A_log": "model-00002-of-00004.safetensors",
53
+ "model.layers.11.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
54
+ "model.layers.11.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
55
+ "model.layers.11.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
56
+ "model.layers.11.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
57
+ "model.layers.11.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
58
+ "model.layers.11.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
59
+ "model.layers.11.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
60
+ "model.layers.11.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
61
+ "model.layers.11.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
62
+ "model.layers.11.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
63
+ "model.layers.11.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
64
+ "model.layers.11.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
65
+ "model.layers.12.input_layernorm.weight": "model-00002-of-00004.safetensors",
66
+ "model.layers.12.linear_attn.A_log": "model-00002-of-00004.safetensors",
67
+ "model.layers.12.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
68
+ "model.layers.12.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
69
+ "model.layers.12.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
70
+ "model.layers.12.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
71
+ "model.layers.12.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
72
+ "model.layers.12.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
73
+ "model.layers.12.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
74
+ "model.layers.12.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
75
+ "model.layers.12.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
76
+ "model.layers.12.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
77
+ "model.layers.12.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
78
+ "model.layers.12.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
79
+ "model.layers.13.input_layernorm.weight": "model-00002-of-00004.safetensors",
80
+ "model.layers.13.linear_attn.A_log": "model-00002-of-00004.safetensors",
81
+ "model.layers.13.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
82
+ "model.layers.13.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
83
+ "model.layers.13.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
84
+ "model.layers.13.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
85
+ "model.layers.13.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
86
+ "model.layers.13.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
87
+ "model.layers.13.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
88
+ "model.layers.13.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
89
+ "model.layers.13.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
90
+ "model.layers.13.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
91
+ "model.layers.13.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
92
+ "model.layers.13.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
93
+ "model.layers.14.input_layernorm.weight": "model-00002-of-00004.safetensors",
94
+ "model.layers.14.linear_attn.A_log": "model-00002-of-00004.safetensors",
95
+ "model.layers.14.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
96
+ "model.layers.14.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
97
+ "model.layers.14.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
98
+ "model.layers.14.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
99
+ "model.layers.14.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
100
+ "model.layers.14.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
101
+ "model.layers.14.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
102
+ "model.layers.14.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
103
+ "model.layers.14.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
104
+ "model.layers.14.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
105
+ "model.layers.14.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
106
+ "model.layers.14.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
107
+ "model.layers.15.input_layernorm.weight": "model-00002-of-00004.safetensors",
108
+ "model.layers.15.linear_attn.A_log": "model-00002-of-00004.safetensors",
109
+ "model.layers.15.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
110
+ "model.layers.15.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
111
+ "model.layers.15.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
112
+ "model.layers.15.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
113
+ "model.layers.15.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
114
+ "model.layers.15.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
115
+ "model.layers.15.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
116
+ "model.layers.15.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
117
+ "model.layers.15.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
118
+ "model.layers.15.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
119
+ "model.layers.15.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
120
+ "model.layers.15.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
121
+ "model.layers.16.input_layernorm.weight": "model-00002-of-00004.safetensors",
122
+ "model.layers.16.linear_attn.A_log": "model-00002-of-00004.safetensors",
123
+ "model.layers.16.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
124
+ "model.layers.16.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
125
+ "model.layers.16.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
126
+ "model.layers.16.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
127
+ "model.layers.16.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
128
+ "model.layers.16.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
129
+ "model.layers.16.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
130
+ "model.layers.16.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
131
+ "model.layers.16.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
132
+ "model.layers.16.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
133
+ "model.layers.16.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
134
+ "model.layers.16.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
135
+ "model.layers.17.input_layernorm.weight": "model-00002-of-00004.safetensors",
136
+ "model.layers.17.linear_attn.A_log": "model-00002-of-00004.safetensors",
137
+ "model.layers.17.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
138
+ "model.layers.17.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
139
+ "model.layers.17.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
140
+ "model.layers.17.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
141
+ "model.layers.17.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
142
+ "model.layers.17.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
143
+ "model.layers.17.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
144
+ "model.layers.17.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
145
+ "model.layers.17.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
146
+ "model.layers.17.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
147
+ "model.layers.17.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
148
+ "model.layers.17.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
149
+ "model.layers.18.input_layernorm.weight": "model-00003-of-00004.safetensors",
150
+ "model.layers.18.linear_attn.A_log": "model-00002-of-00004.safetensors",
151
+ "model.layers.18.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
152
+ "model.layers.18.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
153
+ "model.layers.18.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
154
+ "model.layers.18.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
155
+ "model.layers.18.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
156
+ "model.layers.18.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
157
+ "model.layers.18.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
158
+ "model.layers.18.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
159
+ "model.layers.18.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
160
+ "model.layers.18.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
161
+ "model.layers.18.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
162
+ "model.layers.18.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
163
+ "model.layers.19.input_layernorm.weight": "model-00003-of-00004.safetensors",
164
+ "model.layers.19.linear_attn.A_log": "model-00003-of-00004.safetensors",
165
+ "model.layers.19.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
166
+ "model.layers.19.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
167
+ "model.layers.19.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
168
+ "model.layers.19.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
169
+ "model.layers.19.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
170
+ "model.layers.19.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
171
+ "model.layers.19.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
172
+ "model.layers.19.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
173
+ "model.layers.19.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
174
+ "model.layers.19.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
175
+ "model.layers.19.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
176
+ "model.layers.19.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
177
+ "model.layers.2.input_layernorm.weight": "model-00001-of-00004.safetensors",
178
+ "model.layers.2.linear_attn.A_log": "model-00001-of-00004.safetensors",
179
+ "model.layers.2.linear_attn.conv1d.weight": "model-00001-of-00004.safetensors",
180
+ "model.layers.2.linear_attn.dt_bias": "model-00001-of-00004.safetensors",
181
+ "model.layers.2.linear_attn.in_proj_a.weight": "model-00001-of-00004.safetensors",
182
+ "model.layers.2.linear_attn.in_proj_b.weight": "model-00001-of-00004.safetensors",
183
+ "model.layers.2.linear_attn.in_proj_qkv.weight": "model-00001-of-00004.safetensors",
184
+ "model.layers.2.linear_attn.in_proj_z.weight": "model-00001-of-00004.safetensors",
185
+ "model.layers.2.linear_attn.norm.weight": "model-00001-of-00004.safetensors",
186
+ "model.layers.2.linear_attn.out_proj.weight": "model-00001-of-00004.safetensors",
187
+ "model.layers.2.mlp.down_proj.weight": "model-00001-of-00004.safetensors",
188
+ "model.layers.2.mlp.gate_proj.weight": "model-00001-of-00004.safetensors",
189
+ "model.layers.2.mlp.up_proj.weight": "model-00001-of-00004.safetensors",
190
+ "model.layers.2.post_attention_layernorm.weight": "model-00001-of-00004.safetensors",
191
+ "model.layers.20.input_layernorm.weight": "model-00003-of-00004.safetensors",
192
+ "model.layers.20.linear_attn.A_log": "model-00003-of-00004.safetensors",
193
+ "model.layers.20.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
194
+ "model.layers.20.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
195
+ "model.layers.20.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
196
+ "model.layers.20.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
197
+ "model.layers.20.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
198
+ "model.layers.20.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
199
+ "model.layers.20.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
200
+ "model.layers.20.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
201
+ "model.layers.20.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
202
+ "model.layers.20.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
203
+ "model.layers.20.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
204
+ "model.layers.20.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
205
+ "model.layers.21.input_layernorm.weight": "model-00003-of-00004.safetensors",
206
+ "model.layers.21.linear_attn.A_log": "model-00003-of-00004.safetensors",
207
+ "model.layers.21.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
208
+ "model.layers.21.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
209
+ "model.layers.21.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
210
+ "model.layers.21.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
211
+ "model.layers.21.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
212
+ "model.layers.21.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
213
+ "model.layers.21.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
214
+ "model.layers.21.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
215
+ "model.layers.21.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
216
+ "model.layers.21.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
217
+ "model.layers.21.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
218
+ "model.layers.21.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
219
+ "model.layers.22.input_layernorm.weight": "model-00003-of-00004.safetensors",
220
+ "model.layers.22.linear_attn.A_log": "model-00003-of-00004.safetensors",
221
+ "model.layers.22.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
222
+ "model.layers.22.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
223
+ "model.layers.22.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
224
+ "model.layers.22.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
225
+ "model.layers.22.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
226
+ "model.layers.22.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
227
+ "model.layers.22.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
228
+ "model.layers.22.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
229
+ "model.layers.22.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
230
+ "model.layers.22.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
231
+ "model.layers.22.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
232
+ "model.layers.22.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
233
+ "model.layers.23.input_layernorm.weight": "model-00003-of-00004.safetensors",
234
+ "model.layers.23.linear_attn.A_log": "model-00003-of-00004.safetensors",
235
+ "model.layers.23.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
236
+ "model.layers.23.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
237
+ "model.layers.23.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
238
+ "model.layers.23.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
239
+ "model.layers.23.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
240
+ "model.layers.23.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
241
+ "model.layers.23.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
242
+ "model.layers.23.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
243
+ "model.layers.23.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
244
+ "model.layers.23.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
245
+ "model.layers.23.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
246
+ "model.layers.23.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
247
+ "model.layers.24.input_layernorm.weight": "model-00003-of-00004.safetensors",
248
+ "model.layers.24.linear_attn.A_log": "model-00003-of-00004.safetensors",
249
+ "model.layers.24.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
250
+ "model.layers.24.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
251
+ "model.layers.24.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
252
+ "model.layers.24.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
253
+ "model.layers.24.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
254
+ "model.layers.24.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
255
+ "model.layers.24.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
256
+ "model.layers.24.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
257
+ "model.layers.24.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
258
+ "model.layers.24.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
259
+ "model.layers.24.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
260
+ "model.layers.24.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
261
+ "model.layers.25.input_layernorm.weight": "model-00003-of-00004.safetensors",
262
+ "model.layers.25.linear_attn.A_log": "model-00003-of-00004.safetensors",
263
+ "model.layers.25.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
264
+ "model.layers.25.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
265
+ "model.layers.25.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
266
+ "model.layers.25.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
267
+ "model.layers.25.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
268
+ "model.layers.25.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
269
+ "model.layers.25.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
270
+ "model.layers.25.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
271
+ "model.layers.25.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
272
+ "model.layers.25.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
273
+ "model.layers.25.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
274
+ "model.layers.25.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
275
+ "model.layers.26.input_layernorm.weight": "model-00003-of-00004.safetensors",
276
+ "model.layers.26.linear_attn.A_log": "model-00003-of-00004.safetensors",
277
+ "model.layers.26.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
278
+ "model.layers.26.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
279
+ "model.layers.26.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
280
+ "model.layers.26.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
281
+ "model.layers.26.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
282
+ "model.layers.26.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
283
+ "model.layers.26.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
284
+ "model.layers.26.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
285
+ "model.layers.26.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
286
+ "model.layers.26.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
287
+ "model.layers.26.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
288
+ "model.layers.26.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
289
+ "model.layers.27.input_layernorm.weight": "model-00003-of-00004.safetensors",
290
+ "model.layers.27.linear_attn.A_log": "model-00003-of-00004.safetensors",
291
+ "model.layers.27.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
292
+ "model.layers.27.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
293
+ "model.layers.27.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
294
+ "model.layers.27.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
295
+ "model.layers.27.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
296
+ "model.layers.27.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
297
+ "model.layers.27.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
298
+ "model.layers.27.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
299
+ "model.layers.27.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
300
+ "model.layers.27.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
301
+ "model.layers.27.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
302
+ "model.layers.27.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
303
+ "model.layers.28.input_layernorm.weight": "model-00003-of-00004.safetensors",
304
+ "model.layers.28.linear_attn.A_log": "model-00003-of-00004.safetensors",
305
+ "model.layers.28.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
306
+ "model.layers.28.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
307
+ "model.layers.28.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
308
+ "model.layers.28.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
309
+ "model.layers.28.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
310
+ "model.layers.28.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
311
+ "model.layers.28.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
312
+ "model.layers.28.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
313
+ "model.layers.28.mlp.down_proj.weight": "model-00003-of-00004.safetensors",
314
+ "model.layers.28.mlp.gate_proj.weight": "model-00003-of-00004.safetensors",
315
+ "model.layers.28.mlp.up_proj.weight": "model-00003-of-00004.safetensors",
316
+ "model.layers.28.post_attention_layernorm.weight": "model-00003-of-00004.safetensors",
317
+ "model.layers.29.input_layernorm.weight": "model-00004-of-00004.safetensors",
318
+ "model.layers.29.linear_attn.A_log": "model-00003-of-00004.safetensors",
319
+ "model.layers.29.linear_attn.conv1d.weight": "model-00003-of-00004.safetensors",
320
+ "model.layers.29.linear_attn.dt_bias": "model-00003-of-00004.safetensors",
321
+ "model.layers.29.linear_attn.in_proj_a.weight": "model-00003-of-00004.safetensors",
322
+ "model.layers.29.linear_attn.in_proj_b.weight": "model-00003-of-00004.safetensors",
323
+ "model.layers.29.linear_attn.in_proj_qkv.weight": "model-00003-of-00004.safetensors",
324
+ "model.layers.29.linear_attn.in_proj_z.weight": "model-00003-of-00004.safetensors",
325
+ "model.layers.29.linear_attn.norm.weight": "model-00003-of-00004.safetensors",
326
+ "model.layers.29.linear_attn.out_proj.weight": "model-00003-of-00004.safetensors",
327
+ "model.layers.29.mlp.down_proj.weight": "model-00004-of-00004.safetensors",
328
+ "model.layers.29.mlp.gate_proj.weight": "model-00004-of-00004.safetensors",
329
+ "model.layers.29.mlp.up_proj.weight": "model-00004-of-00004.safetensors",
330
+ "model.layers.29.post_attention_layernorm.weight": "model-00004-of-00004.safetensors",
331
+ "model.layers.3.input_layernorm.weight": "model-00001-of-00004.safetensors",
332
+ "model.layers.3.linear_attn.A_log": "model-00001-of-00004.safetensors",
333
+ "model.layers.3.linear_attn.conv1d.weight": "model-00001-of-00004.safetensors",
334
+ "model.layers.3.linear_attn.dt_bias": "model-00001-of-00004.safetensors",
335
+ "model.layers.3.linear_attn.in_proj_a.weight": "model-00001-of-00004.safetensors",
336
+ "model.layers.3.linear_attn.in_proj_b.weight": "model-00001-of-00004.safetensors",
337
+ "model.layers.3.linear_attn.in_proj_qkv.weight": "model-00001-of-00004.safetensors",
338
+ "model.layers.3.linear_attn.in_proj_z.weight": "model-00001-of-00004.safetensors",
339
+ "model.layers.3.linear_attn.norm.weight": "model-00001-of-00004.safetensors",
340
+ "model.layers.3.linear_attn.out_proj.weight": "model-00001-of-00004.safetensors",
341
+ "model.layers.3.mlp.down_proj.weight": "model-00001-of-00004.safetensors",
342
+ "model.layers.3.mlp.gate_proj.weight": "model-00001-of-00004.safetensors",
343
+ "model.layers.3.mlp.up_proj.weight": "model-00001-of-00004.safetensors",
344
+ "model.layers.3.post_attention_layernorm.weight": "model-00001-of-00004.safetensors",
345
+ "model.layers.30.input_layernorm.weight": "model-00004-of-00004.safetensors",
346
+ "model.layers.30.linear_attn.A_log": "model-00004-of-00004.safetensors",
347
+ "model.layers.30.linear_attn.conv1d.weight": "model-00004-of-00004.safetensors",
348
+ "model.layers.30.linear_attn.dt_bias": "model-00004-of-00004.safetensors",
349
+ "model.layers.30.linear_attn.in_proj_a.weight": "model-00004-of-00004.safetensors",
350
+ "model.layers.30.linear_attn.in_proj_b.weight": "model-00004-of-00004.safetensors",
351
+ "model.layers.30.linear_attn.in_proj_qkv.weight": "model-00004-of-00004.safetensors",
352
+ "model.layers.30.linear_attn.in_proj_z.weight": "model-00004-of-00004.safetensors",
353
+ "model.layers.30.linear_attn.norm.weight": "model-00004-of-00004.safetensors",
354
+ "model.layers.30.linear_attn.out_proj.weight": "model-00004-of-00004.safetensors",
355
+ "model.layers.30.mlp.down_proj.weight": "model-00004-of-00004.safetensors",
356
+ "model.layers.30.mlp.gate_proj.weight": "model-00004-of-00004.safetensors",
357
+ "model.layers.30.mlp.up_proj.weight": "model-00004-of-00004.safetensors",
358
+ "model.layers.30.post_attention_layernorm.weight": "model-00004-of-00004.safetensors",
359
+ "model.layers.31.input_layernorm.weight": "model-00004-of-00004.safetensors",
360
+ "model.layers.31.linear_attn.A_log": "model-00004-of-00004.safetensors",
361
+ "model.layers.31.linear_attn.conv1d.weight": "model-00004-of-00004.safetensors",
362
+ "model.layers.31.linear_attn.dt_bias": "model-00004-of-00004.safetensors",
363
+ "model.layers.31.linear_attn.in_proj_a.weight": "model-00004-of-00004.safetensors",
364
+ "model.layers.31.linear_attn.in_proj_b.weight": "model-00004-of-00004.safetensors",
365
+ "model.layers.31.linear_attn.in_proj_qkv.weight": "model-00004-of-00004.safetensors",
366
+ "model.layers.31.linear_attn.in_proj_z.weight": "model-00004-of-00004.safetensors",
367
+ "model.layers.31.linear_attn.norm.weight": "model-00004-of-00004.safetensors",
368
+ "model.layers.31.linear_attn.out_proj.weight": "model-00004-of-00004.safetensors",
369
+ "model.layers.31.mlp.down_proj.weight": "model-00004-of-00004.safetensors",
370
+ "model.layers.31.mlp.gate_proj.weight": "model-00004-of-00004.safetensors",
371
+ "model.layers.31.mlp.up_proj.weight": "model-00004-of-00004.safetensors",
372
+ "model.layers.31.post_attention_layernorm.weight": "model-00004-of-00004.safetensors",
373
+ "model.layers.4.input_layernorm.weight": "model-00001-of-00004.safetensors",
374
+ "model.layers.4.linear_attn.A_log": "model-00001-of-00004.safetensors",
375
+ "model.layers.4.linear_attn.conv1d.weight": "model-00001-of-00004.safetensors",
376
+ "model.layers.4.linear_attn.dt_bias": "model-00001-of-00004.safetensors",
377
+ "model.layers.4.linear_attn.in_proj_a.weight": "model-00001-of-00004.safetensors",
378
+ "model.layers.4.linear_attn.in_proj_b.weight": "model-00001-of-00004.safetensors",
379
+ "model.layers.4.linear_attn.in_proj_qkv.weight": "model-00001-of-00004.safetensors",
380
+ "model.layers.4.linear_attn.in_proj_z.weight": "model-00001-of-00004.safetensors",
381
+ "model.layers.4.linear_attn.norm.weight": "model-00001-of-00004.safetensors",
382
+ "model.layers.4.linear_attn.out_proj.weight": "model-00001-of-00004.safetensors",
383
+ "model.layers.4.mlp.down_proj.weight": "model-00001-of-00004.safetensors",
384
+ "model.layers.4.mlp.gate_proj.weight": "model-00001-of-00004.safetensors",
385
+ "model.layers.4.mlp.up_proj.weight": "model-00001-of-00004.safetensors",
386
+ "model.layers.4.post_attention_layernorm.weight": "model-00001-of-00004.safetensors",
387
+ "model.layers.5.input_layernorm.weight": "model-00001-of-00004.safetensors",
388
+ "model.layers.5.linear_attn.A_log": "model-00001-of-00004.safetensors",
389
+ "model.layers.5.linear_attn.conv1d.weight": "model-00001-of-00004.safetensors",
390
+ "model.layers.5.linear_attn.dt_bias": "model-00001-of-00004.safetensors",
391
+ "model.layers.5.linear_attn.in_proj_a.weight": "model-00001-of-00004.safetensors",
392
+ "model.layers.5.linear_attn.in_proj_b.weight": "model-00001-of-00004.safetensors",
393
+ "model.layers.5.linear_attn.in_proj_qkv.weight": "model-00001-of-00004.safetensors",
394
+ "model.layers.5.linear_attn.in_proj_z.weight": "model-00001-of-00004.safetensors",
395
+ "model.layers.5.linear_attn.norm.weight": "model-00001-of-00004.safetensors",
396
+ "model.layers.5.linear_attn.out_proj.weight": "model-00001-of-00004.safetensors",
397
+ "model.layers.5.mlp.down_proj.weight": "model-00001-of-00004.safetensors",
398
+ "model.layers.5.mlp.gate_proj.weight": "model-00001-of-00004.safetensors",
399
+ "model.layers.5.mlp.up_proj.weight": "model-00001-of-00004.safetensors",
400
+ "model.layers.5.post_attention_layernorm.weight": "model-00001-of-00004.safetensors",
401
+ "model.layers.6.input_layernorm.weight": "model-00002-of-00004.safetensors",
402
+ "model.layers.6.linear_attn.A_log": "model-00001-of-00004.safetensors",
403
+ "model.layers.6.linear_attn.conv1d.weight": "model-00001-of-00004.safetensors",
404
+ "model.layers.6.linear_attn.dt_bias": "model-00001-of-00004.safetensors",
405
+ "model.layers.6.linear_attn.in_proj_a.weight": "model-00001-of-00004.safetensors",
406
+ "model.layers.6.linear_attn.in_proj_b.weight": "model-00001-of-00004.safetensors",
407
+ "model.layers.6.linear_attn.in_proj_qkv.weight": "model-00001-of-00004.safetensors",
408
+ "model.layers.6.linear_attn.in_proj_z.weight": "model-00001-of-00004.safetensors",
409
+ "model.layers.6.linear_attn.norm.weight": "model-00001-of-00004.safetensors",
410
+ "model.layers.6.linear_attn.out_proj.weight": "model-00001-of-00004.safetensors",
411
+ "model.layers.6.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
412
+ "model.layers.6.mlp.gate_proj.weight": "model-00001-of-00004.safetensors",
413
+ "model.layers.6.mlp.up_proj.weight": "model-00001-of-00004.safetensors",
414
+ "model.layers.6.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
415
+ "model.layers.7.input_layernorm.weight": "model-00002-of-00004.safetensors",
416
+ "model.layers.7.linear_attn.A_log": "model-00002-of-00004.safetensors",
417
+ "model.layers.7.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
418
+ "model.layers.7.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
419
+ "model.layers.7.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
420
+ "model.layers.7.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
421
+ "model.layers.7.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
422
+ "model.layers.7.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
423
+ "model.layers.7.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
424
+ "model.layers.7.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
425
+ "model.layers.7.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
426
+ "model.layers.7.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
427
+ "model.layers.7.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
428
+ "model.layers.7.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
429
+ "model.layers.8.input_layernorm.weight": "model-00002-of-00004.safetensors",
430
+ "model.layers.8.linear_attn.A_log": "model-00002-of-00004.safetensors",
431
+ "model.layers.8.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
432
+ "model.layers.8.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
433
+ "model.layers.8.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
434
+ "model.layers.8.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
435
+ "model.layers.8.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
436
+ "model.layers.8.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
437
+ "model.layers.8.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
438
+ "model.layers.8.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
439
+ "model.layers.8.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
440
+ "model.layers.8.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
441
+ "model.layers.8.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
442
+ "model.layers.8.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
443
+ "model.layers.9.input_layernorm.weight": "model-00002-of-00004.safetensors",
444
+ "model.layers.9.linear_attn.A_log": "model-00002-of-00004.safetensors",
445
+ "model.layers.9.linear_attn.conv1d.weight": "model-00002-of-00004.safetensors",
446
+ "model.layers.9.linear_attn.dt_bias": "model-00002-of-00004.safetensors",
447
+ "model.layers.9.linear_attn.in_proj_a.weight": "model-00002-of-00004.safetensors",
448
+ "model.layers.9.linear_attn.in_proj_b.weight": "model-00002-of-00004.safetensors",
449
+ "model.layers.9.linear_attn.in_proj_qkv.weight": "model-00002-of-00004.safetensors",
450
+ "model.layers.9.linear_attn.in_proj_z.weight": "model-00002-of-00004.safetensors",
451
+ "model.layers.9.linear_attn.norm.weight": "model-00002-of-00004.safetensors",
452
+ "model.layers.9.linear_attn.out_proj.weight": "model-00002-of-00004.safetensors",
453
+ "model.layers.9.mlp.down_proj.weight": "model-00002-of-00004.safetensors",
454
+ "model.layers.9.mlp.gate_proj.weight": "model-00002-of-00004.safetensors",
455
+ "model.layers.9.mlp.up_proj.weight": "model-00002-of-00004.safetensors",
456
+ "model.layers.9.post_attention_layernorm.weight": "model-00002-of-00004.safetensors",
457
+ "model.norm.weight": "model-00004-of-00004.safetensors"
458
+ }
459
+ }
modeling_qwen35_gdn24.py ADDED
@@ -0,0 +1,338 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from typing import Any
4
+
5
+ import torch
6
+ from torch import nn
7
+
8
+ from transformers.cache_utils import Cache, LinearAttentionLayer
9
+ from transformers.generation import GenerationMixin
10
+ from transformers.masking_utils import create_recurrent_attention_mask
11
+ from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
12
+ from transformers.models.qwen3_5.modeling_qwen3_5 import (
13
+ Qwen3_5GatedDeltaNet,
14
+ Qwen3_5PreTrainedModel,
15
+ )
16
+
17
+ from .configuration_qwen35_gdn24 import Qwen35GDN24Config
18
+ from .fused_ops import SuperRMSNorm, SuperSwiGLUMLP, add_final_rmsnorm, add_rmsnorm
19
+
20
+
21
+ class GDN24Cache(Cache):
22
+ """24 native LinearAttentionLayer states plus an explicit logical token clock.
23
+
24
+ Hugging Face's generic all-linear DynamicCache has no attention layer from
25
+ which to infer sequence length. GDN itself does not need K/V length, but
26
+ GenerationMixin does need a logical prefix length. We therefore keep the
27
+ recurrent states native and add only a scalar clock.
28
+ """
29
+
30
+ _qwen35_gdn24_cache_protocol = 1
31
+ is_compileable = False
32
+
33
+ def __init__(self, config: Qwen35GDN24Config):
34
+ number_of_states = getattr(config, "number_of_conv_states", 1)
35
+ layers = [
36
+ LinearAttentionLayer(number_of_states=number_of_states)
37
+ for _ in range(config.num_hidden_layers)
38
+ ]
39
+ super().__init__(layers=layers)
40
+ self.seen_tokens = 0
41
+
42
+ def advance(self, token_count: int) -> None:
43
+ self.seen_tokens += int(token_count)
44
+
45
+ def get_seq_length(self, layer_idx: int = 0) -> int:
46
+ return int(self.seen_tokens)
47
+
48
+ def get_query_offset(self, layer_idx: int = 0) -> int:
49
+ return int(self.seen_tokens)
50
+
51
+ def get_mask_sizes(self, query_length: int, layer_idx: int) -> tuple[int, int]:
52
+ # Every layer is recurrent; no K/V sequence dimension is materialized.
53
+ return int(query_length), 0
54
+
55
+ def get_max_length(self, layer_idx: int | None = None) -> int:
56
+ return -1
57
+
58
+ def reset(self) -> None:
59
+ super().reset()
60
+ self.seen_tokens = 0
61
+
62
+ def crop(self, tokens_to_remove: int) -> None:
63
+ if tokens_to_remove == 0:
64
+ return
65
+ raise NotImplementedError(
66
+ "GDN24 recurrent state is not invertible; speculative rollback requires state snapshots."
67
+ )
68
+
69
+
70
+ class Qwen35GDN24DecoderLayer(nn.Module):
71
+ """Qwen3.5 decoder block with Gated DeltaNet as the token mixer in every layer."""
72
+
73
+ def __init__(self, config: Qwen35GDN24Config, layer_idx: int):
74
+ super().__init__()
75
+ self.layer_idx = int(layer_idx)
76
+ self.linear_attn = Qwen3_5GatedDeltaNet(config, layer_idx)
77
+ self.mlp = SuperSwiGLUMLP(config)
78
+ self.input_layernorm = SuperRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
79
+ self.post_attention_layernorm = SuperRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
80
+
81
+ def forward(
82
+ self,
83
+ hidden_states: torch.Tensor,
84
+ attention_mask: torch.Tensor | None = None,
85
+ past_key_values: Cache | None = None,
86
+ **kwargs: Any,
87
+ ) -> torch.Tensor:
88
+ residual = hidden_states
89
+ hidden_states = self.input_layernorm(hidden_states)
90
+ hidden_states = self.linear_attn(
91
+ hidden_states=hidden_states,
92
+ cache_params=past_key_values,
93
+ attention_mask=attention_mask,
94
+ **kwargs,
95
+ )
96
+ hidden_states = residual + hidden_states
97
+ residual = hidden_states
98
+ hidden_states = self.post_attention_layernorm(hidden_states)
99
+ hidden_states = self.mlp(hidden_states)
100
+ return residual + hidden_states
101
+
102
+
103
+ class Qwen35GDN24PreTrainedModel(Qwen3_5PreTrainedModel):
104
+ config_class = Qwen35GDN24Config
105
+ base_model_prefix = "model"
106
+ supports_gradient_checkpointing = False
107
+ _skip_keys_device_placement = ["past_key_values"]
108
+ _supports_static_cache = False
109
+ _is_stateful = True
110
+ _no_split_modules = ["Qwen35GDN24DecoderLayer"]
111
+
112
+
113
+ class Qwen35GDN24Model(Qwen35GDN24PreTrainedModel):
114
+ def __init__(self, config: Qwen35GDN24Config):
115
+ super().__init__(config)
116
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, config.pad_token_id)
117
+ self.layers = nn.ModuleList(
118
+ [Qwen35GDN24DecoderLayer(config, i) for i in range(config.num_hidden_layers)]
119
+ )
120
+ self.norm = SuperRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
121
+ # No rotary module: all 24 mixers are recurrent and position is implicit in state evolution.
122
+ self.post_init()
123
+
124
+ def get_input_embeddings(self):
125
+ return self.embed_tokens
126
+
127
+ def set_input_embeddings(self, value):
128
+ self.embed_tokens = value
129
+
130
+ def make_recurrent_cache(self) -> GDN24Cache:
131
+ return GDN24Cache(self.config)
132
+
133
+ @staticmethod
134
+ def _compatible_cache(cache: object) -> bool:
135
+ return (
136
+ isinstance(cache, Cache)
137
+ and getattr(cache, "_qwen35_gdn24_cache_protocol", None) == 1
138
+ and hasattr(cache, "advance")
139
+ and hasattr(cache, "layers")
140
+ )
141
+
142
+ def forward(
143
+ self,
144
+ input_ids: torch.LongTensor | None = None,
145
+ attention_mask: torch.Tensor | None = None,
146
+ position_ids: torch.LongTensor | None = None,
147
+ past_key_values: Cache | None = None,
148
+ inputs_embeds: torch.FloatTensor | None = None,
149
+ use_cache: bool | None = None,
150
+ output_hidden_states: bool | None = None,
151
+ return_dict: bool | None = None,
152
+ _advance_cache_clock: bool = True,
153
+ **kwargs: Any,
154
+ ) -> BaseModelOutputWithPast | tuple:
155
+ if (input_ids is None) == (inputs_embeds is None):
156
+ raise ValueError("Specify exactly one of input_ids or inputs_embeds")
157
+
158
+ use_cache = self.config.use_cache if use_cache is None else use_cache
159
+ output_hidden_states = (
160
+ self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
161
+ )
162
+ return_dict = self.config.return_dict if return_dict is None else return_dict
163
+
164
+ if inputs_embeds is None:
165
+ inputs_embeds = self.embed_tokens(input_ids)
166
+ seq_len = inputs_embeds.shape[1]
167
+
168
+ if use_cache and past_key_values is None:
169
+ past_key_values = self.make_recurrent_cache()
170
+ elif past_key_values is not None and not self._compatible_cache(past_key_values):
171
+ raise TypeError(
172
+ "Qwen35GDN24 requires its recurrent-cache protocol object; "
173
+ f"got {type(past_key_values).__module__}.{type(past_key_values).__name__}. "
174
+ "Use model.make_recurrent_cache()."
175
+ )
176
+
177
+ # Common generation/prefill fast path: no padding means no recurrent mask
178
+ # construction at all. This removes Python work from every single-token
179
+ # decode step and is exactly equivalent to create_recurrent_attention_mask(None).
180
+ recurrent_mask = None
181
+ if attention_mask is not None:
182
+ recurrent_mask = create_recurrent_attention_mask(
183
+ config=self.config,
184
+ inputs_embeds=inputs_embeds,
185
+ attention_mask=attention_mask,
186
+ past_key_values=past_key_values,
187
+ )
188
+
189
+ # MAX TURBO residual pipeline. It is algebraically identical to the
190
+ # standard pre-norm decoder, but fuses residual addition with the next
191
+ # RMSNorm and fuses SwiGLU's pointwise core.
192
+ hidden_states = inputs_embeds
193
+ all_hidden_states = () if output_hidden_states else None
194
+ if len(self.layers) > 0:
195
+ if output_hidden_states:
196
+ all_hidden_states += (hidden_states,)
197
+ normed = self.layers[0].input_layernorm(hidden_states)
198
+ for i, layer in enumerate(self.layers):
199
+ mixed = layer.linear_attn(
200
+ hidden_states=normed,
201
+ cache_params=past_key_values,
202
+ attention_mask=recurrent_mask,
203
+ use_cache=use_cache,
204
+ **kwargs,
205
+ )
206
+ attn_residual, mlp_input = add_rmsnorm(
207
+ hidden_states,
208
+ mixed,
209
+ layer.post_attention_layernorm.weight,
210
+ layer.post_attention_layernorm.eps,
211
+ )
212
+ mlp_update = layer.mlp(mlp_input)
213
+ if i + 1 < len(self.layers):
214
+ next_layer = self.layers[i + 1]
215
+ hidden_states, normed = add_rmsnorm(
216
+ attn_residual,
217
+ mlp_update,
218
+ next_layer.input_layernorm.weight,
219
+ next_layer.input_layernorm.eps,
220
+ )
221
+ if output_hidden_states:
222
+ all_hidden_states += (hidden_states,)
223
+ else:
224
+ hidden_states = add_final_rmsnorm(
225
+ attn_residual, mlp_update, self.norm.weight, self.norm.eps
226
+ )
227
+ else:
228
+ hidden_states = self.norm(hidden_states)
229
+
230
+ if output_hidden_states:
231
+ all_hidden_states += (hidden_states,)
232
+
233
+ # CUDA Graph replay cannot rerun Python scalar mutations. MAX TURBO
234
+ # passes _advance_cache_clock=False inside a captured graph and advances
235
+ # the logical clock once per replay from the host.
236
+ if use_cache and past_key_values is not None and _advance_cache_clock:
237
+ past_key_values.advance(seq_len)
238
+
239
+ if not return_dict:
240
+ values = (hidden_states, past_key_values)
241
+ if output_hidden_states:
242
+ values += (all_hidden_states,)
243
+ return values
244
+
245
+ return BaseModelOutputWithPast(
246
+ last_hidden_state=hidden_states,
247
+ past_key_values=past_key_values,
248
+ hidden_states=all_hidden_states,
249
+ attentions=None,
250
+ )
251
+
252
+
253
+ class Qwen35GDN24ForCausalLM(Qwen35GDN24PreTrainedModel, GenerationMixin):
254
+ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
255
+ _keys_to_ignore_on_load_unexpected = [r"^model\.visual.*", r"^mtp.*"]
256
+
257
+ @classmethod
258
+ def _supports_default_dynamic_cache(cls) -> bool:
259
+ return False
260
+
261
+ def __init__(self, config: Qwen35GDN24Config):
262
+ super().__init__(config)
263
+ self.model = Qwen35GDN24Model(config)
264
+ self.vocab_size = config.vocab_size
265
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
266
+ self.post_init()
267
+
268
+ def get_input_embeddings(self):
269
+ return self.model.embed_tokens
270
+
271
+ def set_input_embeddings(self, value):
272
+ self.model.embed_tokens = value
273
+
274
+ def get_output_embeddings(self):
275
+ return self.lm_head
276
+
277
+ def set_output_embeddings(self, value):
278
+ self.lm_head = value
279
+
280
+ def make_recurrent_cache(self) -> GDN24Cache:
281
+ return self.model.make_recurrent_cache()
282
+
283
+ @torch.inference_mode()
284
+ def greedy_step(
285
+ self,
286
+ input_ids: torch.LongTensor,
287
+ past_key_values: Cache,
288
+ *,
289
+ advance_cache_clock: bool = True,
290
+ ) -> torch.LongTensor:
291
+ outputs = self.model(
292
+ input_ids=input_ids,
293
+ past_key_values=past_key_values,
294
+ use_cache=True,
295
+ _advance_cache_clock=advance_cache_clock,
296
+ )
297
+ logits = self.lm_head(outputs.last_hidden_state[:, -1, :])
298
+ return torch.argmax(logits, dim=-1, keepdim=True)
299
+
300
+ def forward(
301
+ self,
302
+ input_ids: torch.LongTensor | None = None,
303
+ attention_mask: torch.Tensor | None = None,
304
+ position_ids: torch.LongTensor | None = None,
305
+ past_key_values: Cache | None = None,
306
+ inputs_embeds: torch.FloatTensor | None = None,
307
+ labels: torch.LongTensor | None = None,
308
+ use_cache: bool | None = None,
309
+ logits_to_keep: int | torch.Tensor = 0,
310
+ **kwargs: Any,
311
+ ) -> CausalLMOutputWithPast:
312
+ outputs = self.model(
313
+ input_ids=input_ids,
314
+ attention_mask=attention_mask,
315
+ position_ids=position_ids,
316
+ past_key_values=past_key_values,
317
+ inputs_embeds=inputs_embeds,
318
+ use_cache=use_cache,
319
+ **kwargs,
320
+ )
321
+ hidden_states = outputs.last_hidden_state
322
+ slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
323
+ logits = self.lm_head(hidden_states[:, slice_indices, :])
324
+
325
+ loss = None
326
+ if labels is not None:
327
+ loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.vocab_size, **kwargs)
328
+
329
+ return CausalLMOutputWithPast(
330
+ loss=loss,
331
+ logits=logits,
332
+ past_key_values=outputs.past_key_values,
333
+ hidden_states=outputs.hidden_states,
334
+ attentions=None,
335
+ )
336
+
337
+
338
+ Qwen35GDN24ForCausalLM.register_for_auto_class("AutoModelForCausalLM")
quantization.py ADDED
@@ -0,0 +1,230 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ """Runtime-only quantization policies for UNI MAX.
4
+
5
+ UNI MAX v1.1 distinguishes two very different notions:
6
+
7
+ * ``full`` quantization (MAX TURBO compatible): every dense MLP projection plus
8
+ the LM head. This is fast, but it changes the hidden trajectory that writes
9
+ future GDN recurrent states. A cache copied from a BF16-prefill model is
10
+ therefore *not* generally a valid state for this quantized dynamical system.
11
+
12
+ * ``state_safe`` quantization (UNI MAX): only the LM head and the MLP of the
13
+ final decoder layer. These operators execute strictly after the final
14
+ persistent GDN state update for each token. Hence they cannot change any
15
+ conv/recurrent cache tensor. If their greedy token agrees with the exact
16
+ model, the next-token recurrent state remains exactly synchronized.
17
+
18
+ Quantization is inference-only and is never written back into the checkpoint.
19
+ """
20
+
21
+ from dataclasses import dataclass
22
+ from typing import Callable
23
+
24
+ import torch
25
+ from torch import nn
26
+
27
+
28
+ @dataclass(frozen=True)
29
+ class QuantizationReport:
30
+ mode: str
31
+ applied: bool
32
+ reason: str | None
33
+ targeted_linears: int
34
+ cuda_capability: tuple[int, int] | None
35
+ policy: str = "full"
36
+
37
+
38
+ def _cuda_capability() -> tuple[int, int] | None:
39
+ if not torch.cuda.is_available():
40
+ return None
41
+ return tuple(int(x) for x in torch.cuda.get_device_capability())
42
+
43
+
44
+ def _is_dense_mlp_linear(module: nn.Module, fqn: str) -> bool:
45
+ if not isinstance(module, nn.Linear):
46
+ return False
47
+ return ".mlp." in fqn and fqn.rsplit(".", 1)[-1] in {"gate_proj", "up_proj", "down_proj"}
48
+
49
+
50
+ def _target_dense_decode_linears(module: nn.Module, fqn: str) -> bool:
51
+ """MAX-compatible full dense target set; native GDN projections excluded."""
52
+ if not isinstance(module, nn.Linear):
53
+ return False
54
+ return fqn == "lm_head" or _is_dense_mlp_linear(module, fqn)
55
+
56
+
57
+ def _final_layer_index(model: nn.Module) -> int:
58
+ cfg = getattr(model, "config", None)
59
+ n = getattr(cfg, "num_hidden_layers", None)
60
+ if n is not None:
61
+ return int(n) - 1
62
+ layers = getattr(getattr(model, "model", None), "layers", None)
63
+ if layers is None:
64
+ raise ValueError("cannot determine final decoder layer index")
65
+ return len(layers) - 1
66
+
67
+
68
+ def make_state_safe_filter(model: nn.Module) -> Callable[[nn.Module, str], bool]:
69
+ """Return the causal-state-compatible UNI quantization filter.
70
+
71
+ The eligible set is exactly:
72
+ * ``lm_head``
73
+ * ``model.layers.<last>.mlp.{gate_proj,up_proj,down_proj}``
74
+
75
+ No module whose output can influence a persistent recurrent-state write is
76
+ eligible.
77
+ """
78
+ last = _final_layer_index(model)
79
+ prefix = f"model.layers.{last}.mlp."
80
+
81
+ def _filter(module: nn.Module, fqn: str) -> bool:
82
+ if not isinstance(module, nn.Linear):
83
+ return False
84
+ if fqn == "lm_head":
85
+ return True
86
+ return fqn.startswith(prefix) and fqn.rsplit(".", 1)[-1] in {
87
+ "gate_proj",
88
+ "up_proj",
89
+ "down_proj",
90
+ }
91
+
92
+ return _filter
93
+
94
+
95
+ def count_target_linears(model: nn.Module) -> int:
96
+ return sum(1 for fqn, mod in model.named_modules() if _target_dense_decode_linears(mod, fqn))
97
+
98
+
99
+ def count_state_safe_linears(model: nn.Module) -> int:
100
+ filt = make_state_safe_filter(model)
101
+ return sum(1 for fqn, mod in model.named_modules() if filt(mod, fqn))
102
+
103
+
104
+ def state_safe_target_names(model: nn.Module) -> tuple[str, ...]:
105
+ filt = make_state_safe_filter(model)
106
+ return tuple(fqn for fqn, mod in model.named_modules() if filt(mod, fqn))
107
+
108
+
109
+ def available_quant_modes() -> tuple[str, ...]:
110
+ modes = ["exact"]
111
+ try:
112
+ import torchao # noqa: F401
113
+ modes.extend(["fp8", "int8"])
114
+ except Exception:
115
+ pass
116
+ return tuple(modes)
117
+
118
+
119
+ def _apply_quantization(model: nn.Module, mode: str, *, policy: str) -> QuantizationReport:
120
+ mode = str(mode).lower().strip()
121
+ policy = str(policy).lower().strip()
122
+ capability = _cuda_capability()
123
+
124
+ if policy == "state_safe":
125
+ filter_fn = make_state_safe_filter(model)
126
+ targeted = count_state_safe_linears(model)
127
+ elif policy == "full":
128
+ filter_fn = _target_dense_decode_linears
129
+ targeted = count_target_linears(model)
130
+ else:
131
+ return QuantizationReport(mode, False, f"unknown quantization policy: {policy}", 0, capability, policy)
132
+
133
+ if mode in {"", "none", "exact", "bf16"}:
134
+ return QuantizationReport("exact", False, None, targeted, capability, policy)
135
+
136
+ try:
137
+ from torchao.quantization import quantize_
138
+ except Exception as exc: # pragma: no cover - runtime dependent
139
+ return QuantizationReport(
140
+ mode,
141
+ False,
142
+ f"torchao unavailable: {type(exc).__name__}: {exc}",
143
+ targeted,
144
+ capability,
145
+ policy,
146
+ )
147
+
148
+ if targeted == 0:
149
+ return QuantizationReport(mode, False, "no eligible dense linear layers found", 0, capability, policy)
150
+
151
+ try:
152
+ if mode == "fp8":
153
+ if capability is None or capability < (8, 9):
154
+ return QuantizationReport(
155
+ mode,
156
+ False,
157
+ f"FP8 requires CUDA SM 8.9+, got {capability}",
158
+ targeted,
159
+ capability,
160
+ policy,
161
+ )
162
+ from torchao.quantization import Float8DynamicActivationFloat8WeightConfig, PerTensor
163
+
164
+ config = Float8DynamicActivationFloat8WeightConfig(granularity=PerTensor())
165
+ elif mode == "int8":
166
+ from torchao.quantization import Int8WeightOnlyConfig
167
+
168
+ config = Int8WeightOnlyConfig()
169
+ else:
170
+ return QuantizationReport(
171
+ mode,
172
+ False,
173
+ f"unknown quantization mode: {mode}",
174
+ targeted,
175
+ capability,
176
+ policy,
177
+ )
178
+
179
+ quantize_(model, config, filter_fn=filter_fn)
180
+ return QuantizationReport(mode, True, None, targeted, capability, policy)
181
+ except Exception as exc: # pragma: no cover - hardware/runtime specific
182
+ return QuantizationReport(
183
+ mode,
184
+ False,
185
+ f"{type(exc).__name__}: {exc}",
186
+ targeted,
187
+ capability,
188
+ policy,
189
+ )
190
+
191
+
192
+ def apply_max_quantization(model: nn.Module, mode: str) -> QuantizationReport:
193
+ """Legacy MAX TURBO policy: all MLP projections + LM head."""
194
+ return _apply_quantization(model, mode, policy="full")
195
+
196
+
197
+ def apply_uni_quantization(model: nn.Module, mode: str) -> QuantizationReport:
198
+ """UNI MAX state-compatible policy: final MLP + LM head only."""
199
+ return _apply_quantization(model, mode, policy="state_safe")
200
+
201
+
202
+ @torch.inference_mode()
203
+ def greedy_token_trace(model, input_ids: torch.Tensor, steps: int = 32) -> torch.Tensor:
204
+ steps = int(steps)
205
+ if steps < 1:
206
+ return torch.empty((input_ids.shape[0], 0), dtype=torch.long, device=input_ids.device)
207
+ cache = model.make_recurrent_cache()
208
+ out = model(input_ids=input_ids, past_key_values=cache, use_cache=True, logits_to_keep=1)
209
+ token = out.logits[:, -1, :].argmax(dim=-1, keepdim=True)
210
+ tokens = [token]
211
+ for _ in range(steps - 1):
212
+ token = model.greedy_step(token, cache)
213
+ tokens.append(token)
214
+ return torch.cat(tokens, dim=1)
215
+
216
+
217
+ def token_agreement(reference: torch.Tensor, candidate: torch.Tensor) -> tuple[float, int]:
218
+ if reference.shape != candidate.shape:
219
+ raise ValueError(f"token trace shape mismatch: {reference.shape} vs {candidate.shape}")
220
+ if reference.numel() == 0:
221
+ return 1.0, 0
222
+ eq = reference.eq(candidate)
223
+ agreement = float(eq.float().mean().item())
224
+ flat = eq.reshape(-1).tolist()
225
+ prefix = 0
226
+ for ok in flat:
227
+ if not ok:
228
+ break
229
+ prefix += 1
230
+ return agreement, prefix
tokenizer.json ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8818dc7a3be5f461790e3a81703816f482925f6c2ff9fef5a9fc4b821e5051f2
3
+ size 19989423
tokenizer_config.json ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "audio_bos_token": "<|audio_start|>",
4
+ "audio_eos_token": "<|audio_end|>",
5
+ "audio_token": "<|audio_pad|>",
6
+ "backend": "tokenizers",
7
+ "bos_token": null,
8
+ "clean_up_tokenization_spaces": false,
9
+ "eos_token": "<|im_end|>",
10
+ "errors": "replace",
11
+ "image_token": "<|image_pad|>",
12
+ "is_local": false,
13
+ "local_files_only": false,
14
+ "model_max_length": 262144,
15
+ "model_specific_special_tokens": {
16
+ "audio_bos_token": "<|audio_start|>",
17
+ "audio_eos_token": "<|audio_end|>",
18
+ "audio_token": "<|audio_pad|>",
19
+ "image_token": "<|image_pad|>",
20
+ "video_token": "<|video_pad|>",
21
+ "vision_bos_token": "<|vision_start|>",
22
+ "vision_eos_token": "<|vision_end|>"
23
+ },
24
+ "pad_token": "<|endoftext|>",
25
+ "pretokenize_regex": "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+",
26
+ "split_special_tokens": false,
27
+ "tokenizer_class": "Qwen2Tokenizer",
28
+ "unk_token": null,
29
+ "video_token": "<|video_pad|>",
30
+ "vision_bos_token": "<|vision_start|>",
31
+ "vision_eos_token": "<|vision_end|>"
32
+ }
uni_engine.py ADDED
@@ -0,0 +1,290 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ """UNI MAX: unified exact-prefill / quantized-decode recurrent inference.
4
+
5
+ The GDN24 cache topology is identical between the exact BF16 model and a model
6
+ whose **state-safe tail** (final MLP + LM head) is quantized. UNI MAX v1.1 exploits that causal invariant:
7
+
8
+ 1. prefill the prompt with the exact BF16 model;
9
+ 2. copy only the fixed recurrent cache state into a state-compatible persistent decode engine;
10
+ 3. decode with the fastest CUDA-Graph candidate (typically FP8 on NVIDIA L4);
11
+ 4. optionally verify FP8 draft blocks with exact BF16 chunk forwards to recover
12
+ exact greedy-token semantics.
13
+
14
+ No KV sequence is copied: the bridge is O(1) in context length. Full-MLP quantization is deliberately excluded from this bridge because it changes the hidden trajectory that writes future recurrent states.
15
+ """
16
+
17
+ from dataclasses import dataclass
18
+ import time
19
+ from typing import Any
20
+
21
+ import torch
22
+
23
+ from .engine import MaxTurboGraphDecoder, snapshot_cache, restore_cache_
24
+
25
+
26
+ @dataclass(frozen=True)
27
+ class BridgeReport:
28
+ seconds: float
29
+ mib: float
30
+ seen_tokens: int
31
+
32
+
33
+ @dataclass(frozen=True)
34
+ class VerifyReport:
35
+ generated_tokens: int
36
+ drafted_tokens: int
37
+ accepted_draft_tokens: int
38
+ rejected_blocks: int
39
+ verifier_blocks: int
40
+
41
+ @property
42
+ def acceptance_rate(self) -> float:
43
+ if self.drafted_tokens <= 0:
44
+ return 1.0
45
+ return self.accepted_draft_tokens / self.drafted_tokens
46
+
47
+
48
+ def cache_payload_bytes(cache) -> int:
49
+ total = 0
50
+ for layer in cache.layers:
51
+ for table_name in ("conv_states", "recurrent_states"):
52
+ table = getattr(layer, table_name, {})
53
+ for tensor in table.values():
54
+ if tensor is not None:
55
+ total += tensor.numel() * tensor.element_size()
56
+ return int(total)
57
+
58
+
59
+ @torch.inference_mode()
60
+ def copy_cache_(dst, src) -> None:
61
+ """Copy recurrent state without reallocating destination tensors.
62
+
63
+ Destination storage must already be materialized (CUDA-Graph capture does
64
+ this). Tensor addresses are preserved, so captured graphs remain valid.
65
+ """
66
+ if len(dst.layers) != len(src.layers):
67
+ raise ValueError("cache topology mismatch")
68
+ for d_layer, s_layer in zip(dst.layers, src.layers):
69
+ for table_name in ("conv_states", "recurrent_states"):
70
+ d_table = getattr(d_layer, table_name)
71
+ s_table = getattr(s_layer, table_name)
72
+ for idx, s_tensor in s_table.items():
73
+ if s_tensor is None:
74
+ continue
75
+ d_tensor = d_table.get(idx)
76
+ if d_tensor is None:
77
+ raise RuntimeError(
78
+ f"destination cache storage is not materialized: {table_name}[{idx}]"
79
+ )
80
+ if d_tensor.shape != s_tensor.shape or d_tensor.dtype != s_tensor.dtype:
81
+ raise RuntimeError(
82
+ f"cache state mismatch for {table_name}[{idx}]: "
83
+ f"dst={tuple(d_tensor.shape)}/{d_tensor.dtype}, "
84
+ f"src={tuple(s_tensor.shape)}/{s_tensor.dtype}"
85
+ )
86
+ d_tensor.copy_(s_tensor)
87
+ d_layer.has_previous_state.clear()
88
+ d_layer.has_previous_state.update(dict(s_layer.has_previous_state))
89
+ d_layer.is_conv_states_initialized.clear()
90
+ d_layer.is_conv_states_initialized.update(dict(s_layer.is_conv_states_initialized))
91
+ d_layer.is_recurrent_states_initialized.clear()
92
+ d_layer.is_recurrent_states_initialized.update(dict(s_layer.is_recurrent_states_initialized))
93
+ dst.seen_tokens = int(src.seen_tokens)
94
+
95
+
96
+ class UniMaxEngine:
97
+ """Phase-specialized recurrent inference engine.
98
+
99
+ ``exact_model`` always handles prefill. ``decode_model`` can be the same model (UNI-SAFE) or a state-safe quantized clone (UNI-SPEED / UNI-EXACT).
100
+ """
101
+
102
+ def __init__(
103
+ self,
104
+ exact_model,
105
+ decode_model,
106
+ decode_graph: MaxTurboGraphDecoder,
107
+ ):
108
+ self.exact_model = exact_model.eval()
109
+ self.decode_model = decode_model.eval()
110
+ self.decode_graph = decode_graph
111
+ self.device = self.exact_model.get_input_embeddings().weight.device
112
+ if self.device != self.decode_graph.device:
113
+ raise ValueError("exact and decode engines must live on the same CUDA device")
114
+ self.exact_cache = self.exact_model.make_recurrent_cache()
115
+ self.last_bridge: BridgeReport | None = None
116
+
117
+ @torch.inference_mode()
118
+ def reset(self) -> None:
119
+ self.exact_cache.reset()
120
+ self.decode_graph.reset()
121
+ self.last_bridge = None
122
+
123
+ @torch.inference_mode()
124
+ def _bridge(self) -> BridgeReport:
125
+ if self.device.type == "cuda":
126
+ torch.cuda.synchronize(self.device)
127
+ t0 = time.perf_counter()
128
+ copy_cache_(self.decode_graph.cache, self.exact_cache)
129
+ if self.device.type == "cuda":
130
+ torch.cuda.synchronize(self.device)
131
+ dt = time.perf_counter() - t0
132
+ report = BridgeReport(
133
+ seconds=float(dt),
134
+ mib=cache_payload_bytes(self.exact_cache) / 2**20,
135
+ seen_tokens=int(self.exact_cache.seen_tokens),
136
+ )
137
+ self.last_bridge = report
138
+ return report
139
+
140
+ @torch.inference_mode()
141
+ def prefill_exact(self, input_ids: torch.Tensor) -> tuple[torch.Tensor, Any, BridgeReport]:
142
+ """Exact BF16 prefill, then O(1)-context state-compatible recurrent bridge."""
143
+ self.exact_cache.reset()
144
+ out = self.exact_model(
145
+ input_ids=input_ids,
146
+ past_key_values=self.exact_cache,
147
+ use_cache=True,
148
+ logits_to_keep=1,
149
+ )
150
+ first = torch.argmax(out.logits[:, -1, :], dim=-1, keepdim=True)
151
+ bridge = self._bridge()
152
+ self.decode_graph.static_token.copy_(first)
153
+ return first, out, bridge
154
+
155
+ @torch.inference_mode()
156
+ def decode_fast(self, first_token: torch.Tensor, recurrent_forwards: int) -> torch.Tensor:
157
+ return self.decode_graph.decode_forwards(first_token, int(recurrent_forwards))
158
+
159
+ @torch.inference_mode()
160
+ def generate_fast(self, input_ids: torch.Tensor, max_new_tokens: int) -> torch.Tensor:
161
+ """UNI-SPEED generation: exact first token, quantized graph thereafter."""
162
+ n = int(max_new_tokens)
163
+ if n <= 0:
164
+ return torch.empty((input_ids.shape[0], 0), dtype=torch.long, device=input_ids.device)
165
+ first, _, _ = self.prefill_exact(input_ids)
166
+ if n == 1:
167
+ return first
168
+ tail = self.decode_graph.decode_tokens(first, n - 1)
169
+ return torch.cat([first, tail], dim=1)
170
+
171
+ @torch.inference_mode()
172
+ def decode_verified(
173
+ self,
174
+ first_token: torch.Tensor,
175
+ recurrent_forwards: int,
176
+ *,
177
+ draft_block: int | None = None,
178
+ ) -> tuple[torch.Tensor, VerifyReport]:
179
+ """Verify ``recurrent_forwards`` tokens after ``first_token``.
180
+
181
+ ``prefill_exact`` must have been called immediately before this method so
182
+ both exact and decode caches represent the same prompt state.
183
+ """
184
+ remaining = int(recurrent_forwards)
185
+ if remaining <= 0:
186
+ empty = torch.empty((first_token.shape[0], 0), dtype=torch.long, device=first_token.device)
187
+ return empty, VerifyReport(0, 0, 0, 0, 0)
188
+ if first_token.shape[0] != 1:
189
+ raise ValueError("UNI-EXACT currently supports batch size 1")
190
+
191
+ current = first_token
192
+ generated: list[torch.Tensor] = []
193
+ k_default = max(1, int(draft_block or self.decode_graph.block_size))
194
+ drafted = accepted_drafts = rejected_blocks = verifier_blocks = 0
195
+
196
+ while remaining > 0:
197
+ k = min(k_default, remaining)
198
+ exact_before = snapshot_cache(self.exact_cache)
199
+ draft = self.decode_graph.decode_tokens(current, k)
200
+ drafted += k
201
+
202
+ verify_input = current if k == 1 else torch.cat([current, draft[:, :-1]], dim=1)
203
+ verify_out = self.exact_model(
204
+ input_ids=verify_input,
205
+ past_key_values=self.exact_cache,
206
+ use_cache=True,
207
+ logits_to_keep=k,
208
+ )
209
+ exact_pred = torch.argmax(verify_out.logits[:, -k:, :], dim=-1)
210
+ verifier_blocks += 1
211
+ mismatch_positions = (~exact_pred.eq(draft)[0]).nonzero(as_tuple=False)
212
+
213
+ if mismatch_positions.numel() == 0:
214
+ generated.append(draft.detach().clone())
215
+ accepted_drafts += k
216
+ current = draft[:, -1:]
217
+ remaining -= k
218
+ continue
219
+
220
+ m = int(mismatch_positions[0, 0].item())
221
+ accepted_drafts += m
222
+ rejected_blocks += 1
223
+ restore_cache_(self.exact_cache, exact_before)
224
+ replay_input = current if m == 0 else torch.cat([current, draft[:, :m]], dim=1)
225
+ replay = self.exact_model(
226
+ input_ids=replay_input,
227
+ past_key_values=self.exact_cache,
228
+ use_cache=True,
229
+ logits_to_keep=1,
230
+ )
231
+ corrected = torch.argmax(replay.logits[:, -1, :], dim=-1, keepdim=True)
232
+ if m:
233
+ generated.append(draft[:, :m].detach().clone())
234
+ generated.append(corrected.detach().clone())
235
+ current = corrected
236
+ emitted = m + 1
237
+ remaining -= emitted
238
+ copy_cache_(self.decode_graph.cache, self.exact_cache)
239
+ self.decode_graph.static_token.copy_(current)
240
+
241
+ out = torch.cat(generated, dim=1)
242
+ return out, VerifyReport(
243
+ generated_tokens=int(out.shape[1]),
244
+ drafted_tokens=int(drafted),
245
+ accepted_draft_tokens=int(accepted_drafts),
246
+ rejected_blocks=int(rejected_blocks),
247
+ verifier_blocks=int(verifier_blocks),
248
+ )
249
+
250
+ @torch.inference_mode()
251
+ def generate_verified(
252
+ self,
253
+ input_ids: torch.Tensor,
254
+ max_new_tokens: int,
255
+ *,
256
+ draft_block: int | None = None,
257
+ ) -> tuple[torch.Tensor, VerifyReport]:
258
+ """UNI-EXACT generation with exact greedy-token semantics."""
259
+ n = int(max_new_tokens)
260
+ if n <= 0:
261
+ empty = torch.empty((input_ids.shape[0], 0), dtype=torch.long, device=input_ids.device)
262
+ return empty, VerifyReport(0, 0, 0, 0, 0)
263
+ first, _, _ = self.prefill_exact(input_ids)
264
+ if n == 1:
265
+ return first, VerifyReport(1, 0, 0, 0, 0)
266
+ tail, report = self.decode_verified(first, n - 1, draft_block=draft_block)
267
+ full = torch.cat([first, tail], dim=1)
268
+ return full, VerifyReport(
269
+ generated_tokens=int(full.shape[1]),
270
+ drafted_tokens=report.drafted_tokens,
271
+ accepted_draft_tokens=report.accepted_draft_tokens,
272
+ rejected_blocks=report.rejected_blocks,
273
+ verifier_blocks=report.verifier_blocks,
274
+ )
275
+
276
+
277
+
278
+ @torch.inference_mode()
279
+ def exact_greedy_tokens(model, input_ids: torch.Tensor, max_new_tokens: int) -> torch.Tensor:
280
+ n = int(max_new_tokens)
281
+ if n <= 0:
282
+ return torch.empty((input_ids.shape[0], 0), dtype=torch.long, device=input_ids.device)
283
+ cache = model.make_recurrent_cache()
284
+ out = model(input_ids=input_ids, past_key_values=cache, use_cache=True, logits_to_keep=1)
285
+ token = torch.argmax(out.logits[:, -1, :], dim=-1, keepdim=True)
286
+ pieces = [token]
287
+ for _ in range(n - 1):
288
+ token = model.greedy_step(token, cache)
289
+ pieces.append(token)
290
+ return torch.cat(pieces, dim=1)