multimodalart HF Staff commited on
Commit
dd73dc8
verified
1 Parent(s): 0a3117f

Upload h3_nvfp4.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. h3_nvfp4.py +964 -0
h3_nvfp4.py ADDED
@@ -0,0 +1,964 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Blackwell-native MiniMax-H3 transformer for the pruned ComfyUI NVFP4 checkpoint.
2
+
3
+ The public diffusers checkpoint spends 13.04B of its 33.12B parameters on per-block
4
+ AdaLN projections. ComfyUI's pruned checkpoint replaces those projections with an
5
+ interpolated 1025-point timestep curve, fuses Q/K/V, and stores the four large linear
6
+ layers in every block as NVFP4. This adapter keeps diffusers' packed-sequence contract
7
+ so the rest of the split Space (schedulers, VAEs and remote conditioner) stays unchanged.
8
+
9
+ The kernel/layout conventions follow ComfyUI's Apache-2.0 implementation:
10
+ https://github.com/Comfy-Org/ComfyUI/blob/master/comfy/ldm/minimax/model.py
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import json
16
+ import math
17
+ import os
18
+ from types import SimpleNamespace
19
+
20
+ import torch
21
+ import torch.nn as nn
22
+ import torch.nn.functional as F
23
+ import comfy_kitchen as kitchen
24
+ from comfy_kitchen.tensor import QuantizedTensor, TensorCoreNVFP4Layout
25
+ from diffusers.models.attention_dispatch import dispatch_attention_fn
26
+
27
+ try:
28
+ import triton
29
+ import triton.language as tl
30
+ except ImportError: # PyTorch CUDA wheels include Triton; retain a portable fallback for source inspection/tests.
31
+ triton = None
32
+ tl = None
33
+
34
+
35
+ NVFP4_REPO = os.environ.get("H3_NVFP4_REPO", "lilcheaty/MiniMax-H3-NVFP4")
36
+ NVFP4_FILE = os.environ.get("H3_NVFP4_FILE", "minimax_h3_fl2va_pruned_nvfp4.safetensors")
37
+
38
+ HIDDEN = 5376
39
+ HEADS = 56
40
+ HEAD_DIM = 128
41
+ FFN = 14336
42
+ TEXT_DIM = 5120
43
+ TIME_DIM = 8
44
+ VIDEO_DIM = 24 * 1 * 2 * 2
45
+ AUDIO_DIM = 32
46
+ LAYERS = 50
47
+ REFINER_LAYERS = 2
48
+ EPS = 1e-5
49
+
50
+ # EasyCache is the conservative profile. The Ultra Fast profile uses a bounded linear residual forecast:
51
+ # three exact warmup evaluations, at most three forecasts in a row, and two exact tail evaluations. Unlike blind
52
+ # output reuse, forecasting follows the local denoising trajectory while making the amount of saved work predictable.
53
+ EASYCACHE_THRESHOLD = max(0.0, float(os.environ.get("H3_EASYCACHE_THRESHOLD", "0.10")))
54
+ EASYCACHE_START = min(1.0, max(0.0, float(os.environ.get("H3_EASYCACHE_START", "0.15"))))
55
+ EASYCACHE_END = min(1.0, max(EASYCACHE_START, float(os.environ.get("H3_EASYCACHE_END", "0.95"))))
56
+ EASYCACHE_SUBSAMPLE = max(1, int(os.environ.get("H3_EASYCACHE_SUBSAMPLE", "8")))
57
+ FIRST_BLOCK_THRESHOLD = max(0.0, float(os.environ.get("H3_FIRST_BLOCK_THRESHOLD", "0.08")))
58
+ FIRST_BLOCK_DENSE_START = max(1, int(os.environ.get("H3_FIRST_BLOCK_DENSE_START", "3")))
59
+ FIRST_BLOCK_DENSE_END = max(1, int(os.environ.get("H3_FIRST_BLOCK_DENSE_END", "2")))
60
+ FORECAST_BLEND = min(1.0, max(0.0, float(os.environ.get("H3_FORECAST_BLEND", "0.65"))))
61
+ FUSED_ADALN = os.environ.get("H3_FUSED_ADALN", "0") == "1" and triton is not None
62
+ SOL_ATTN = os.environ.get("H3_SOL_ATTN", "1") == "1"
63
+ SOL_ATTN_TAU = float(os.environ.get("H3_SOL_ATTN_TAU", "1.0"))
64
+ SOL_ATTN_DENSE_STEPS = max(0, int(os.environ.get("H3_SOL_ATTN_DENSE_STEPS", "10")))
65
+ SOL_ATTN_DENSE_LAYERS = max(0, int(os.environ.get("H3_SOL_ATTN_DENSE_LAYERS", "2")))
66
+ SOL_ATTN_MIN_TOKENS = max(0, int(os.environ.get("H3_SOL_ATTN_MIN_TOKENS", "8192")))
67
+
68
+
69
+ if triton is not None:
70
+
71
+ @triton.jit
72
+ def _adaln_modulate_kernel(
73
+ x, shift, scale, row_ids, elements: tl.constexpr, hidden: tl.constexpr, modulation_stride: tl.constexpr
74
+ ):
75
+ offsets = tl.program_id(0) * 256 + tl.arange(0, 256)
76
+ mask = offsets < elements
77
+ columns = offsets % hidden
78
+ rows = offsets // hidden
79
+ modulation_rows = tl.load(row_ids + rows, mask=mask, other=0)
80
+ modulation_offsets = modulation_rows * modulation_stride + columns
81
+ values = tl.load(x + offsets, mask=mask)
82
+ shifts = tl.load(shift + modulation_offsets, mask=mask)
83
+ scales = tl.load(scale + modulation_offsets, mask=mask)
84
+ tl.store(x + offsets, values * (1.0 + scales) + shifts, mask=mask)
85
+
86
+ @triton.jit
87
+ def _adaln_gate_kernel(
88
+ x, update, gate, row_ids, elements: tl.constexpr, hidden: tl.constexpr, modulation_stride: tl.constexpr
89
+ ):
90
+ offsets = tl.program_id(0) * 256 + tl.arange(0, 256)
91
+ mask = offsets < elements
92
+ columns = offsets % hidden
93
+ rows = offsets // hidden
94
+ modulation_rows = tl.load(row_ids + rows, mask=mask, other=0)
95
+ gates = tl.load(gate + modulation_rows * modulation_stride + columns, mask=mask)
96
+ values = tl.load(x + offsets, mask=mask)
97
+ updates = tl.load(update + offsets, mask=mask)
98
+ tl.store(x + offsets, values + updates * gates, mask=mask)
99
+
100
+
101
+ class H3StepCache:
102
+ """ComfyUI EasyCache-style adaptive reuse of a complete H3 denoising result.
103
+
104
+ This caches the model residual, not the generated video. A request with a new prompt, seed, canvas or keyframe
105
+ starts from an empty cache. Decisions use a sparse sample of generated video latent rows, while the reused
106
+ residual contains every video and audio row so their joint denoising trajectory stays coupled.
107
+ """
108
+
109
+ def __init__(self):
110
+ self.total_steps = 0
111
+ self.step = 0
112
+ self.skipped = 0
113
+ self.profile = "balanced"
114
+ self.consecutive_skips = 0
115
+ self.last_actual_step = None
116
+ self.previous_input = None
117
+ self.previous_output = None
118
+ self.previous_output_norm = None
119
+ self.relative_rate = None
120
+ self.accumulated_change = None
121
+ self.video_residual = None
122
+ self.audio_residual = None
123
+ self.video_residual_slope = None
124
+ self.audio_residual_slope = None
125
+ self.pending_input = None
126
+ self.pending_input_change = None
127
+ self.pending_track = False
128
+ self.head_residual = None
129
+ self.tail_residual = None
130
+ self.first_block_output = None
131
+
132
+ def begin(self, total_steps: int | None, profile: str = "balanced") -> None:
133
+ self.__init__()
134
+ self.total_steps = max(0, int(total_steps or 0))
135
+ self.profile = str(profile or "balanced").lower()
136
+
137
+ @property
138
+ def enabled(self) -> bool:
139
+ return self.profile != "exact" and self.total_steps > 2
140
+
141
+ def _forecast(self, video_input, audio_input):
142
+ distance = max(1, self.step - int(self.last_actual_step or 0))
143
+ video_residual = self.video_residual
144
+ audio_residual = self.audio_residual
145
+ if self.video_residual_slope is not None:
146
+ video_residual = video_residual + self.video_residual_slope * (distance * FORECAST_BLEND)
147
+ audio_residual = audio_residual + self.audio_residual_slope * (distance * FORECAST_BLEND)
148
+ self.skipped += 1
149
+ self.consecutive_skips += 1
150
+ self.step += 1
151
+ return video_input + video_residual, audio_input + audio_residual
152
+
153
+ def try_reuse(self, video_input, audio_input, condition_rows: int):
154
+ self.pending_input = None
155
+ self.pending_input_change = None
156
+ self.pending_track = False
157
+ if not self.enabled:
158
+ return None
159
+
160
+ # Balanced uses NVIDIA's H3 FirstBlockCache below, after block 0 has produced a high-signal residual.
161
+ # Only the deliberately aggressive Ultra profile forecasts a whole transformer call before block 0.
162
+ if not self.profile.startswith("ultra"):
163
+ return None
164
+
165
+ # Ultra Fast is deliberately bounded: no more than three forecasts can separate exact transformer calls, and
166
+ # the high-noise warmup plus low-noise tail remain exact. At the default 16 steps this executes 7 full DiT
167
+ # evaluations instead of 16 while still sampling the original 16-step scheduler trajectory.
168
+ if self.profile.startswith("ultra"):
169
+ can_forecast = (
170
+ self.step >= 3
171
+ and self.step < self.total_steps - 2
172
+ and self.consecutive_skips < 3
173
+ and self.last_actual_step is not None
174
+ and self.video_residual is not None
175
+ and self.audio_residual is not None
176
+ and self.video_residual.shape == video_input.shape
177
+ and self.audio_residual.shape == audio_input.shape
178
+ )
179
+ if can_forecast:
180
+ return self._forecast(video_input, audio_input)
181
+ return None
182
+
183
+ if EASYCACHE_THRESHOLD <= 0.0:
184
+ return None
185
+
186
+ end_step = math.floor(self.total_steps * EASYCACHE_END)
187
+ if self.step >= end_step:
188
+ return None
189
+
190
+ # Condition latents are static. Excluding them makes the change estimate reflect the generated trajectory.
191
+ sampled_input = video_input[0, condition_rows::EASYCACHE_SUBSAMPLE].detach().float()
192
+ self.pending_input = sampled_input
193
+ self.pending_track = True
194
+ if self.previous_input is not None:
195
+ self.pending_input_change = (sampled_input - self.previous_input).abs().mean()
196
+
197
+ start_step = math.ceil(self.total_steps * EASYCACHE_START)
198
+ can_reuse = (
199
+ self.step >= start_step
200
+ and self.pending_input_change is not None
201
+ and self.relative_rate is not None
202
+ and self.previous_output_norm is not None
203
+ and self.video_residual is not None
204
+ and self.audio_residual is not None
205
+ and self.video_residual.shape == video_input.shape
206
+ and self.audio_residual.shape == audio_input.shape
207
+ )
208
+ if not can_reuse:
209
+ return None
210
+
211
+ estimated_change = self.relative_rate * self.pending_input_change
212
+ estimated_change = estimated_change / self.previous_output_norm.clamp_min(1e-6)
213
+ accumulated = estimated_change if self.accumulated_change is None else self.accumulated_change + estimated_change
214
+ if bool((accumulated < EASYCACHE_THRESHOLD).item()):
215
+ self.accumulated_change = accumulated
216
+ self.skipped += 1
217
+ self.step += 1
218
+ return video_input + self.video_residual, audio_input + self.audio_residual
219
+ return None
220
+
221
+ def first_block_decision(self, block_input: torch.Tensor, block_output: torch.Tensor) -> bool:
222
+ """Return True when blocks 1..49 can reuse their previous joint residual.
223
+
224
+ This is the single-GPU equivalent of NVIDIA Sol-Engine's H3 FirstBlockCache at threshold 0.08. The first
225
+ block is always evaluated. Its normalized residual change is a much stronger predictor than raw latent
226
+ motion, while the cached tail residual still covers the complete text/video/audio packed sequence.
227
+ """
228
+ if not self.enabled or self.profile.startswith("ultra"):
229
+ return False
230
+ if FIRST_BLOCK_THRESHOLD <= 0.0:
231
+ self.first_block_output = block_output.detach().clone()
232
+ return False
233
+
234
+ keep_dense = self.step < FIRST_BLOCK_DENSE_START or self.step >= self.total_steps - FIRST_BLOCK_DENSE_END
235
+ residual = block_output - block_input
236
+ reusable = (
237
+ not keep_dense
238
+ and self.head_residual is not None
239
+ and self.tail_residual is not None
240
+ and self.tail_residual.shape == block_output.shape
241
+ )
242
+ should_reuse = False
243
+ if reusable:
244
+ difference = (residual - self.head_residual).abs().mean()
245
+ reference = self.head_residual.abs().mean().clamp_min(1e-8)
246
+ should_reuse = bool(((difference / reference) <= FIRST_BLOCK_THRESHOLD).item())
247
+
248
+ if should_reuse:
249
+ self.skipped += 1
250
+ self.consecutive_skips += 1
251
+ self.step += 1
252
+ return True
253
+
254
+ # This engine's residual/gate operations update `packed` in place. Preserve the head output before later
255
+ # blocks mutate the same storage; diffusers' reference blocks are out-of-place and do not need this clone.
256
+ self.first_block_output = block_output.detach().clone()
257
+ self.head_residual = residual.detach()
258
+ return False
259
+
260
+ def update_first_block_tail(self, final_block_output: torch.Tensor) -> None:
261
+ if self.first_block_output is None:
262
+ return
263
+ self.tail_residual = (final_block_output - self.first_block_output).detach()
264
+ self.last_actual_step = self.step
265
+ self.consecutive_skips = 0
266
+ self.step += 1
267
+ self.first_block_output = None
268
+
269
+ def update(self, video_input, audio_input, video_output, audio_output, condition_rows: int) -> None:
270
+ # Balanced's clock and state are updated at the block-stack boundary by FirstBlockCache.
271
+ if self.enabled and not self.profile.startswith("ultra"):
272
+ return
273
+ if self.pending_track:
274
+ sampled_output = video_output[0, condition_rows::EASYCACHE_SUBSAMPLE].detach().float()
275
+ if self.previous_output is not None and self.pending_input_change is not None:
276
+ output_change = (sampled_output - self.previous_output).abs().mean()
277
+ self.relative_rate = output_change / self.pending_input_change.clamp_min(1e-6)
278
+ self.previous_input = self.pending_input.clone()
279
+ self.previous_output = sampled_output.clone()
280
+ self.previous_output_norm = sampled_output.abs().mean()
281
+ if not self.profile.startswith("ultra"):
282
+ self.video_residual = (video_output - video_input).detach()
283
+ self.audio_residual = (audio_output - audio_input).detach()
284
+ self.accumulated_change = None
285
+ if self.profile.startswith("ultra"):
286
+ new_video_residual = (video_output - video_input).detach()
287
+ new_audio_residual = (audio_output - audio_input).detach()
288
+ if self.video_residual is not None and self.last_actual_step is not None:
289
+ gap = max(1, self.step - self.last_actual_step)
290
+ self.video_residual_slope = (new_video_residual - self.video_residual) / gap
291
+ self.audio_residual_slope = (new_audio_residual - self.audio_residual) / gap
292
+ self.video_residual = new_video_residual
293
+ self.audio_residual = new_audio_residual
294
+ self.last_actual_step = self.step
295
+ self.consecutive_skips = 0
296
+ self.step += 1
297
+ self.pending_input = None
298
+ self.pending_input_change = None
299
+ self.pending_track = False
300
+
301
+ def finish(self) -> dict:
302
+ stats = {
303
+ "steps": self.step,
304
+ "computed": max(0, self.step - self.skipped),
305
+ "forecasted": self.skipped,
306
+ "profile": self.profile,
307
+ }
308
+ if self.enabled and self.step:
309
+ computed = max(1, self.step - self.skipped)
310
+ print(
311
+ f"[h3-nvfp4] adaptive step cache skipped {self.skipped}/{self.step} transformer evaluations "
312
+ f"({self.step / computed:.2f}x denoiser-work reduction)",
313
+ flush=True,
314
+ )
315
+ self.begin(None)
316
+ return stats
317
+
318
+
319
+ class H3SolAttention:
320
+ """NVIDIA Sol-Attn policy adapted to H3's single-GPU packed attention.
321
+
322
+ The packed prefix (text, conditioning video and generated audio) remains an exact KV sink and its query rows are
323
+ recomputed densely. Only target-video query/key interactions become sparse, after ten dense denoising steps and
324
+ outside the first two transformer blocks. Any unavailable/JIT-failing backend falls back to cuDNN for the request.
325
+ """
326
+
327
+ def __init__(self):
328
+ self.enabled = SOL_ATTN
329
+ self.step = 0
330
+ self.video_start = 0
331
+ self.sparse_calls = 0
332
+ self.dense_calls = 0
333
+ self.failure = None
334
+
335
+ def begin(self):
336
+ self.step = 0
337
+ self.video_start = 0
338
+ self.sparse_calls = 0
339
+ self.dense_calls = 0
340
+ self.failure = None
341
+
342
+ def observe(self, video_indices: torch.Tensor, sequence: int, step: int) -> None:
343
+ self.step = int(step)
344
+ if not self.video_start:
345
+ deltas = video_indices[1:] - video_indices[:-1]
346
+ breaks = (deltas != 1).nonzero().flatten()
347
+ start = int(breaks[-1]) + 1 if len(breaks) else 0
348
+ self.video_start = int(video_indices[start]) if video_indices.numel() else sequence
349
+
350
+ def __call__(self, query, key, value, layer: int):
351
+ tokens = int(query.shape[1])
352
+ if (
353
+ not self.enabled
354
+ or self.failure is not None
355
+ or self.step < SOL_ATTN_DENSE_STEPS
356
+ or layer < SOL_ATTN_DENSE_LAYERS
357
+ or tokens < SOL_ATTN_MIN_TOKENS
358
+ or not 0 < self.video_start < tokens
359
+ ):
360
+ self.dense_calls += 1
361
+ return None
362
+ try:
363
+ from sol_attn import sol_attn
364
+
365
+ q, k, v = (tensor.contiguous() for tensor in (query, key, value))
366
+ attended = sol_attn(
367
+ q,
368
+ k,
369
+ v,
370
+ tau=SOL_ATTN_TAU,
371
+ thresh_type="diag",
372
+ kv_splits=1,
373
+ sink_start=0,
374
+ sink_tokens=self.video_start,
375
+ )
376
+ # An exact KV sink does not make the prefix's own queries dense. H3 jointly generates audio in that
377
+ # prefix, so reproduce those rows with exact attention as NVIDIA's H3 integration does.
378
+ prefix = self.video_start
379
+ dense_prefix = F.scaled_dot_product_attention(
380
+ q[:, :prefix].transpose(1, 2),
381
+ k.transpose(1, 2),
382
+ v.transpose(1, 2),
383
+ dropout_p=0.0,
384
+ is_causal=False,
385
+ ).transpose(1, 2)
386
+ attended[:, :prefix] = dense_prefix
387
+ self.sparse_calls += 1
388
+ return attended
389
+ except Exception as error:
390
+ self.failure = f"{type(error).__name__}: {error}"
391
+ print(f"[h3-sol-attn] falling back to dense attention: {self.failure}", flush=True)
392
+ self.dense_calls += 1
393
+ return None
394
+
395
+
396
+ def _quant_config(handle, prefix: str) -> dict | None:
397
+ key = f"{prefix}.comfy_quant"
398
+ if key not in handle.keys():
399
+ return None
400
+ return json.loads(handle.get_tensor(key).numpy().tobytes())
401
+
402
+
403
+ class H3Linear(nn.Module):
404
+ """A plain or comfy-kitchen NVFP4 linear, selected by checkpoint metadata."""
405
+
406
+ def __init__(
407
+ self,
408
+ in_features: int,
409
+ out_features: int,
410
+ bias: bool = False,
411
+ compute_dtype: torch.dtype | None = None,
412
+ ):
413
+ super().__init__()
414
+ self.in_features = in_features
415
+ self.out_features = out_features
416
+ self.compute_dtype = compute_dtype
417
+ self.register_parameter("weight", None)
418
+ self.register_parameter("bias", None)
419
+ self.register_buffer("input_scale", None)
420
+ self.register_buffer("pre_quant_scale", None)
421
+ self.quantized = False
422
+ self.full_precision_mm = False
423
+
424
+ def load(self, handle, prefix: str) -> None:
425
+ config = _quant_config(handle, prefix)
426
+ weight = handle.get_tensor(f"{prefix}.weight")
427
+
428
+ if config is None:
429
+ self.weight = nn.Parameter(
430
+ weight if self.compute_dtype is None else weight.to(self.compute_dtype), requires_grad=False
431
+ )
432
+ elif config.get("format") == "nvfp4":
433
+ block_scale = handle.get_tensor(f"{prefix}.weight_scale")
434
+ if block_scale.dtype == torch.uint8:
435
+ block_scale = block_scale.view(torch.float8_e4m3fn)
436
+ tensor_scale = handle.get_tensor(f"{prefix}.weight_scale_2").float()
437
+ params = TensorCoreNVFP4Layout.Params(
438
+ scale=tensor_scale,
439
+ block_scale=block_scale,
440
+ orig_dtype=torch.bfloat16,
441
+ orig_shape=(self.out_features, self.in_features),
442
+ )
443
+ quantized = QuantizedTensor(weight.to(torch.uint8), "TensorCoreNVFP4Layout", params)
444
+ self.weight = nn.Parameter(quantized, requires_grad=False)
445
+ self.quantized = True
446
+ self.full_precision_mm = bool(config.get("full_precision_matrix_mult", False))
447
+ for name in ("input_scale", "pre_quant_scale"):
448
+ key = f"{prefix}.{name}"
449
+ if key in handle.keys():
450
+ setattr(self, name, handle.get_tensor(key))
451
+ else:
452
+ raise ValueError(f"Unsupported quantization on {prefix}: {config}")
453
+
454
+ bias_key = f"{prefix}.bias"
455
+ if bias_key in handle.keys():
456
+ bias = handle.get_tensor(bias_key)
457
+ self.bias = nn.Parameter(
458
+ bias if self.compute_dtype is None else bias.to(self.compute_dtype), requires_grad=False
459
+ )
460
+
461
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
462
+ if self.pre_quant_scale is not None:
463
+ hidden_states = hidden_states * self.pre_quant_scale.to(
464
+ device=hidden_states.device, dtype=hidden_states.dtype
465
+ )
466
+ if not self.quantized:
467
+ hidden_states = hidden_states.to(self.weight.dtype)
468
+ return F.linear(
469
+ hidden_states,
470
+ self.weight,
471
+ self.bias,
472
+ )
473
+
474
+ if self.full_precision_mm:
475
+ # Some AWQ checkpoints use NVFP4 as a compact weight format but deliberately retain BF16 activations and
476
+ # GEMMs. Dequantization is layer-local, so residency stays compact without adding activation error.
477
+ weight = self.weight.dequantize().to(hidden_states.dtype)
478
+ return F.linear(hidden_states, weight, None if self.bias is None else self.bias.to(hidden_states.dtype))
479
+
480
+ shape = hidden_states.shape
481
+ flat = hidden_states.reshape(-1, shape[-1])
482
+ scale = None if self.input_scale is None else self.input_scale.to(flat.device)
483
+ quantized_input = QuantizedTensor.from_float(flat, "TensorCoreNVFP4Layout", scale=scale)
484
+ output = F.linear(
485
+ quantized_input,
486
+ self.weight,
487
+ None if self.bias is None else self.bias.to(hidden_states.dtype),
488
+ )
489
+ return output.reshape(*shape[:-1], self.out_features)
490
+
491
+
492
+ class H3RMSNorm(nn.Module):
493
+ def __init__(self, width: int, eps: float = EPS):
494
+ super().__init__()
495
+ self.width = width
496
+ self.eps = eps
497
+ self.register_parameter("weight", None)
498
+
499
+ def load(self, handle, prefix: str) -> None:
500
+ self.weight = nn.Parameter(handle.get_tensor(f"{prefix}.weight"), requires_grad=False)
501
+
502
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
503
+ return F.rms_norm(
504
+ hidden_states,
505
+ (self.width,),
506
+ self.weight,
507
+ self.eps,
508
+ )
509
+
510
+
511
+ class H3Attention(nn.Module):
512
+ def __init__(self):
513
+ super().__init__()
514
+ self.qkv_proj = H3Linear(HIDDEN, 3 * HEADS * HEAD_DIM)
515
+ self.q_norm = H3RMSNorm(HEAD_DIM)
516
+ self.k_norm = H3RMSNorm(HEAD_DIM)
517
+ self.out_proj = H3Linear(HEADS * HEAD_DIM, HIDDEN)
518
+
519
+ def load(self, handle, prefix: str) -> None:
520
+ self.qkv_proj.load(handle, f"{prefix}.qkv_proj")
521
+ self.q_norm.load(handle, f"{prefix}.q_norm")
522
+ self.k_norm.load(handle, f"{prefix}.k_norm")
523
+ self.out_proj.load(handle, f"{prefix}.out_proj")
524
+
525
+ def forward(self, hidden_states, rope_table, backend: str, sparse=None, layer: int = -1):
526
+ sequence = hidden_states.shape[0]
527
+ qkv = self.qkv_proj(hidden_states)
528
+ query, key, value = qkv.split(HEADS * HEAD_DIM, dim=-1)
529
+ query = query.view(1, sequence, HEADS, HEAD_DIM)
530
+ key = key.view(1, sequence, HEADS, HEAD_DIM)
531
+ value = value.view(1, sequence, HEADS, HEAD_DIM)
532
+
533
+ # One in-place kernel replaces Q RMSNorm, K RMSNorm and both partial RoPE applications.
534
+ kitchen.rms_rope_split_half_(
535
+ query,
536
+ key,
537
+ rope_table,
538
+ self.q_norm.weight,
539
+ self.k_norm.weight,
540
+ epsilon=self.q_norm.eps,
541
+ rot_dim=rope_table.shape[-3] * 2,
542
+ )
543
+ attended = sparse(query, key, value, layer) if sparse is not None else None
544
+ if attended is None:
545
+ attended = dispatch_attention_fn(
546
+ query,
547
+ key,
548
+ value,
549
+ attn_mask=None,
550
+ dropout_p=0.0,
551
+ is_causal=False,
552
+ backend=backend,
553
+ )
554
+ return self.out_proj(attended.reshape(sequence, HEADS * HEAD_DIM))
555
+
556
+
557
+ class H3MLP(nn.Module):
558
+ def __init__(self):
559
+ super().__init__()
560
+ self.fc1 = H3Linear(HIDDEN, 2 * FFN)
561
+ self.fc2 = H3Linear(FFN, HIDDEN)
562
+
563
+ def load(self, handle, prefix: str) -> None:
564
+ self.fc1.load(handle, f"{prefix}.fc1")
565
+ self.fc2.load(handle, f"{prefix}.fc2")
566
+
567
+ def forward(self, hidden_states):
568
+ gate, up = self.fc1(hidden_states).chunk(2, dim=-1)
569
+ return self.fc2(F.silu(gate).mul_(up))
570
+
571
+
572
+ class H3RefinerBlock(nn.Module):
573
+ def __init__(self):
574
+ super().__init__()
575
+ self.norm1 = H3RMSNorm(HIDDEN)
576
+ self.attn = H3Attention()
577
+ self.norm2 = H3RMSNorm(HIDDEN)
578
+ self.mlp = H3MLP()
579
+
580
+ def load(self, handle, prefix: str) -> None:
581
+ self.norm1.load(handle, f"{prefix}.norm1")
582
+ self.attn.load(handle, f"{prefix}.attn")
583
+ self.norm2.load(handle, f"{prefix}.norm2")
584
+ self.mlp.load(handle, f"{prefix}.mlp")
585
+
586
+
587
+ class H3AdaLN(nn.Module):
588
+ def __init__(self, expand: int, modalities: int):
589
+ super().__init__()
590
+ self.expand = expand
591
+ self.modalities = modalities
592
+ # Curve checkpoints deliberately evaluate interpolation and modulation projection in FP32. Expanding the
593
+ # checkpoint's tiny FP16 [*, 8] matrices once at load avoids 51 request-step casts.
594
+ self.linear = H3Linear(
595
+ TIME_DIM, expand * HIDDEN * modalities, bias=True, compute_dtype=torch.float32
596
+ )
597
+
598
+ def load(self, handle, prefix: str) -> None:
599
+ self.linear.load(handle, f"{prefix}.linear")
600
+
601
+ def forward(self, time_embedding, output_dtype=None):
602
+ projected = self.linear(time_embedding)
603
+ if output_dtype is not None:
604
+ # One contiguous conversion is numerically identical to converting the six chunk views independently,
605
+ # and removes five CUDA launches from every one of the 50 blocks.
606
+ projected = projected.to(output_dtype)
607
+ projected = projected.view(-1, self.expand * HIDDEN)
608
+ return projected.chunk(self.expand, dim=-1)
609
+
610
+
611
+ class H3Block(nn.Module):
612
+ def __init__(self):
613
+ super().__init__()
614
+ self.norm1 = H3RMSNorm(HIDDEN)
615
+ self.attn = H3Attention()
616
+ self.norm2 = H3RMSNorm(HIDDEN)
617
+ self.mlp = H3MLP()
618
+ self.adaln_proj = H3AdaLN(6, 3)
619
+
620
+ def load(self, handle, prefix: str) -> None:
621
+ self.norm1.load(handle, f"{prefix}.norm1")
622
+ self.attn.load(handle, f"{prefix}.attn")
623
+ self.norm2.load(handle, f"{prefix}.norm2")
624
+ self.mlp.load(handle, f"{prefix}.mlp")
625
+ self.adaln_proj.load(handle, f"{prefix}.adaln_proj")
626
+
627
+
628
+ class H3FinalLayer(nn.Module):
629
+ def __init__(self):
630
+ super().__init__()
631
+ self.norm = H3RMSNorm(HIDDEN)
632
+ self.adaln_proj = H3AdaLN(2, 1)
633
+ self.video_out = H3Linear(HIDDEN, VIDEO_DIM, bias=True, compute_dtype=torch.float32)
634
+ self.audio_out = H3Linear(HIDDEN, AUDIO_DIM, bias=True, compute_dtype=torch.float32)
635
+
636
+ def load(self, handle, prefix: str) -> None:
637
+ self.norm.load(handle, f"{prefix}.norm")
638
+ self.adaln_proj.load(handle, f"{prefix}.adaln_proj")
639
+ self.video_out.load(handle, f"{prefix}.video_out")
640
+ self.audio_out.load(handle, f"{prefix}.audio_out")
641
+
642
+
643
+ class H3NVFP4Transformer(nn.Module):
644
+ """Diffusers-compatible H3 transformer backed by fused comfy-kitchen NVFP4 kernels."""
645
+
646
+ def __init__(self):
647
+ super().__init__()
648
+ # The modular pipeline reads these values through the diffusers component config rather than inspecting the
649
+ # module itself. Keep the public transformer contract even though this lean adapter is not a ConfigMixin.
650
+ self.config = SimpleNamespace(
651
+ patch_size=(1, 2, 2),
652
+ in_channels=24,
653
+ audio_in_channels=AUDIO_DIM,
654
+ text_dim=TEXT_DIM,
655
+ )
656
+ self.video_patch_proj = H3Linear(VIDEO_DIM, HIDDEN, bias=True, compute_dtype=torch.float32)
657
+ self.audio_patch_proj = H3Linear(AUDIO_DIM, HIDDEN, bias=True, compute_dtype=torch.float32)
658
+ self.condition_proj = H3Linear(TEXT_DIM, HIDDEN, bias=True)
659
+ self.token_refiner = nn.ModuleList([H3RefinerBlock() for _ in range(REFINER_LAYERS)])
660
+ self.token_refiner_norm = H3RMSNorm(HIDDEN)
661
+ self.blocks = nn.ModuleList([H3Block() for _ in range(LAYERS)])
662
+ self.final_layer = H3FinalLayer()
663
+ self.register_buffer("adaln_t_table", None)
664
+ self.register_buffer("rope_inv_freq", None)
665
+ self.attention_backend = "_native_cudnn"
666
+ self._text_cache = None
667
+ self._rope_cache = None
668
+ self._segment_cache = None
669
+ self._condition_video_rows = None
670
+ self._condition_video_embedding = None
671
+ self._output_indices = None
672
+ self._generated_rows = None
673
+ self._step_cache = H3StepCache()
674
+ self._sol_attention = H3SolAttention()
675
+
676
+ @property
677
+ def dtype(self) -> torch.dtype:
678
+ """Match ModelMixin's placement contract used by ModularPipeline.to()."""
679
+ return self.condition_proj.weight.dtype
680
+
681
+ @property
682
+ def device(self) -> torch.device:
683
+ return self.adaln_t_table.device
684
+
685
+ def load(self, path: str) -> None:
686
+ from safetensors import safe_open
687
+
688
+ with safe_open(path, framework="pt", device="cpu") as handle:
689
+ self.video_patch_proj.load(handle, "video_patch_proj")
690
+ self.audio_patch_proj.load(handle, "audio_patch_proj")
691
+ self.condition_proj.load(handle, "condition_proj")
692
+ for index, block in enumerate(self.token_refiner):
693
+ block.load(handle, f"token_refiner.blocks.{index}")
694
+ self.token_refiner_norm.load(handle, "token_refiner.final_norm")
695
+ for index, block in enumerate(self.blocks):
696
+ block.load(handle, f"blocks.{index}")
697
+ self.final_layer.load(handle, "final_layer")
698
+ self.adaln_t_table = handle.get_tensor("adaln_t_table")
699
+ self.rope_inv_freq = handle.get_tensor("rope.inv_freq")
700
+ # Every loaded tensor is already a frozen Parameter (or a buffer). Avoid mutating the quantized tensor
701
+ # subclass through a redundant requires_grad_ dispatch.
702
+ self.eval()
703
+
704
+ def set_attention_backend(self, backend: str) -> None:
705
+ self.attention_backend = backend
706
+
707
+ def begin_request(self, total_steps: int | None = None, profile: str = "balanced") -> None:
708
+ self._text_cache = None
709
+ self._rope_cache = None
710
+ self._segment_cache = None
711
+ self._condition_video_rows = None
712
+ self._condition_video_embedding = None
713
+ self._output_indices = None
714
+ self._generated_rows = None
715
+ self._step_cache.begin(total_steps, profile)
716
+ self._sol_attention.begin()
717
+
718
+ def end_request(self) -> dict:
719
+ stats = self._step_cache.finish()
720
+ stats["sol_sparse_calls"] = self._sol_attention.sparse_calls
721
+ stats["sol_dense_calls"] = self._sol_attention.dense_calls
722
+ stats["sol_failure"] = self._sol_attention.failure
723
+ self._text_cache = None
724
+ self._rope_cache = None
725
+ self._segment_cache = None
726
+ self._condition_video_rows = None
727
+ self._condition_video_embedding = None
728
+ self._output_indices = None
729
+ self._generated_rows = None
730
+ return stats
731
+
732
+ def _refine_text(self, text_states: torch.Tensor) -> torch.Tensor:
733
+ key = (text_states.data_ptr(), tuple(text_states.shape), text_states.device)
734
+ if self._text_cache is not None and self._text_cache[0] == key:
735
+ return self._text_cache[1]
736
+ hidden = self.condition_proj(text_states)
737
+ # Text is tiny compared with the video sequence; use the same fused QKV path with an identity RoPE omitted.
738
+ for block in self.token_refiner:
739
+ residual = hidden
740
+ normalized = block.norm1(hidden)
741
+ qkv = block.attn.qkv_proj(normalized)
742
+ query, key_states, value = qkv.split(HEADS * HEAD_DIM, dim=-1)
743
+ query = block.attn.q_norm(query.view(1, -1, HEADS, HEAD_DIM))
744
+ key_states = block.attn.k_norm(key_states.view(1, -1, HEADS, HEAD_DIM))
745
+ value = value.view(1, -1, HEADS, HEAD_DIM)
746
+ attended = dispatch_attention_fn(
747
+ query,
748
+ key_states,
749
+ value,
750
+ attn_mask=None,
751
+ dropout_p=0.0,
752
+ is_causal=False,
753
+ backend=self.attention_backend,
754
+ ).reshape(-1, HEADS * HEAD_DIM)
755
+ hidden = residual + block.attn.out_proj(attended)
756
+ hidden = hidden + block.mlp(block.norm2(hidden))
757
+ hidden = self.token_refiner_norm(hidden)
758
+ self._text_cache = (key, hidden)
759
+ return hidden
760
+
761
+ def _rope(self, position_ids: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
762
+ key = (position_ids.data_ptr(), tuple(position_ids.shape), position_ids.device, dtype)
763
+ if self._rope_cache is not None and self._rope_cache[0] == key:
764
+ return self._rope_cache[1]
765
+ positions = position_ids.to(torch.float32)
766
+ frequencies = positions.unsqueeze(-1) * self.rope_inv_freq.to(position_ids.device).view(1, 1, -1)
767
+ temporal, height, width = frequencies.unbind(dim=1)
768
+ angles = torch.cat((temporal, height, width), dim=-1)
769
+ cosine, sine = angles.cos(), angles.sin()
770
+ table = torch.stack((cosine, -sine, sine, cosine), dim=-1)
771
+ table = table.reshape(1, position_ids.shape[0], 1, angles.shape[-1], 2, 2).to(dtype)
772
+ self._rope_cache = (key, table)
773
+ return table
774
+
775
+ def _time_embedding(self, timestep: torch.Tensor) -> torch.Tensor:
776
+ table = self.adaln_t_table.to(timestep.device)
777
+ position = timestep.float().clamp(0.0, 1.0) * (table.shape[0] - 1)
778
+ lower = position.floor().long().clamp(max=table.shape[0] - 2)
779
+ return torch.lerp(table[lower], table[lower + 1], (position - lower).unsqueeze(1))
780
+
781
+ def _segments(self, indices: torch.Tensor):
782
+ if self._segment_cache is None:
783
+ host = indices.detach().cpu()
784
+ changes = (host[1:] != host[:-1]).nonzero().flatten().add(1).tolist()
785
+ bounds = [0, *changes, len(host)]
786
+ # Python row ids avoid indexing modulation tensors with CUDA scalar tensors in every block.
787
+ self._segment_cache = [
788
+ (start, stop, int(host[start])) for start, stop in zip(bounds[:-1], bounds[1:])
789
+ ]
790
+ return self._segment_cache
791
+
792
+ def _video_layout(self, video_indices: torch.Tensor) -> int:
793
+ """Number of leading, static keyframe-patch rows in the video latent tensor."""
794
+ if self._condition_video_rows is None:
795
+ host = video_indices.detach().cpu()
796
+ discontinuities = (host[1:] - host[:-1] != 1).nonzero().flatten()
797
+ self._condition_video_rows = int(discontinuities[0]) + 1 if len(discontinuities) else 0
798
+ return self._condition_video_rows
799
+
800
+ def _project_video(self, hidden_states: torch.Tensor, condition_rows: int, dtype: torch.dtype) -> torch.Tensor:
801
+ source = hidden_states[0]
802
+ if condition_rows == 0:
803
+ return self.video_patch_proj(source.float()).to(dtype)
804
+ if self._condition_video_embedding is None:
805
+ self._condition_video_embedding = self.video_patch_proj(source[:condition_rows].float()).to(dtype)
806
+ generated = self.video_patch_proj(source[condition_rows:].float()).to(dtype)
807
+ return torch.cat((self._condition_video_embedding, generated), dim=0)
808
+
809
+ @staticmethod
810
+ def _modulate(hidden, shift, scale, row_ids, segments):
811
+ if FUSED_ADALN and hidden.is_cuda and hidden.is_contiguous():
812
+ _adaln_modulate_kernel[(triton.cdiv(hidden.numel(), 256),)](
813
+ hidden, shift, scale, row_ids, hidden.numel(), HIDDEN, shift.stride(0), num_warps=4
814
+ )
815
+ return hidden
816
+ for start, stop, row in segments:
817
+ hidden[start:stop].mul_(1.0 + scale[row]).add_(shift[row])
818
+ return hidden
819
+
820
+ @staticmethod
821
+ def _gate(hidden, update, gate, row_ids, segments):
822
+ if FUSED_ADALN and hidden.is_cuda and hidden.is_contiguous() and update.is_contiguous():
823
+ _adaln_gate_kernel[(triton.cdiv(hidden.numel(), 256),)](
824
+ hidden, update, gate, row_ids, hidden.numel(), HIDDEN, gate.stride(0), num_warps=4
825
+ )
826
+ return hidden
827
+ for start, stop, row in segments:
828
+ hidden[start:stop].addcmul_(update[start:stop], gate[row])
829
+ return hidden
830
+
831
+ def forward(
832
+ self,
833
+ hidden_states,
834
+ audio_hidden_states,
835
+ encoder_hidden_states,
836
+ timestep,
837
+ timestep_indices,
838
+ token_tags,
839
+ position_ids,
840
+ video_indices,
841
+ audio_indices,
842
+ text_indices,
843
+ attention_kwargs=None,
844
+ return_dict=True,
845
+ ):
846
+ from diffusers.models.transformers.transformer_minimax_h3 import MiniMaxH3TransformerOutput
847
+
848
+ if hidden_states.shape[0] != 1:
849
+ raise ValueError("The NVFP4 MiniMax-H3 engine supports batch size 1.")
850
+
851
+ condition_rows = self._video_layout(video_indices)
852
+ reused = self._step_cache.try_reuse(hidden_states, audio_hidden_states, condition_rows)
853
+ if reused is not None:
854
+ video_output, audio_output = reused
855
+ if not return_dict:
856
+ return video_output, audio_output
857
+ return MiniMaxH3TransformerOutput(sample=video_output, audio_sample=audio_output)
858
+
859
+ text = self._refine_text(encoder_hidden_states[0].to(torch.bfloat16))
860
+ video = self._project_video(hidden_states, condition_rows, text.dtype)
861
+ audio = self.audio_patch_proj(audio_hidden_states[0].float()).to(text.dtype)
862
+ # Text, video and audio indices partition the packed sequence, so initialization would only add a full HBM
863
+ # write before the three index copies overwrite every row.
864
+ packed = text.new_empty((position_ids.shape[0], HIDDEN))
865
+ packed.index_copy_(0, text_indices, text)
866
+ packed.index_copy_(0, video_indices, video)
867
+ packed.index_copy_(0, audio_indices, audio)
868
+
869
+ time_embedding = self._time_embedding(timestep)
870
+ adaln_indices = timestep_indices * 3 + token_tags.clamp(min=0)
871
+ segments = self._segments(adaln_indices)
872
+ rope = self._rope(position_ids, packed.dtype)
873
+ use_sol_attention = self._step_cache.profile != "exact" and self._sol_attention.enabled
874
+ self._sol_attention.observe(video_indices, packed.shape[0], self._step_cache.step)
875
+
876
+ reused_tail = False
877
+ for layer, block in enumerate(self.blocks):
878
+ if layer == 0:
879
+ # Block 0 writes its residual updates in place, so retain the pre-block value for the official FBC
880
+ # signal `(head_output - head_input)`.
881
+ block_input = packed.detach().clone()
882
+ # One conversion per small modulation table, rather than one conversion per sequence segment.
883
+ modulations = block.adaln_proj(time_embedding, packed.dtype)
884
+ shift_attn, scale_attn, gate_attn, shift_mlp, scale_mlp, gate_mlp = modulations
885
+ normalized = self._modulate(block.norm1(packed), shift_attn, scale_attn, adaln_indices, segments)
886
+ packed = self._gate(
887
+ packed,
888
+ block.attn(
889
+ normalized,
890
+ rope,
891
+ self.attention_backend,
892
+ self._sol_attention if use_sol_attention else None,
893
+ layer,
894
+ ),
895
+ gate_attn,
896
+ adaln_indices,
897
+ segments,
898
+ )
899
+ normalized = self._modulate(block.norm2(packed), shift_mlp, scale_mlp, adaln_indices, segments)
900
+ packed = self._gate(packed, block.mlp(normalized), gate_mlp, adaln_indices, segments)
901
+
902
+ if layer == 0:
903
+ if self._step_cache.first_block_decision(block_input, packed):
904
+ packed = packed + self._step_cache.tail_residual
905
+ reused_tail = True
906
+ break
907
+ if layer == len(self.blocks) - 1 and not reused_tail:
908
+ self._step_cache.update_first_block_tail(packed)
909
+
910
+ shift, scale = self.final_layer.adaln_proj(time_embedding)
911
+
912
+ # Keyframe output rows are discarded by the scheduler. Avoid their FP32 output projection and put zeros in
913
+ # those unused slots to retain the pipeline's expected tensor shape.
914
+ generated_video_indices = video_indices[condition_rows:]
915
+ if self._output_indices is None:
916
+ self._generated_rows = generated_video_indices.shape[0]
917
+ self._output_indices = torch.cat((generated_video_indices, audio_indices))
918
+ generated_rows = self._generated_rows
919
+ normalized_output = self.final_layer.norm(packed.index_select(0, self._output_indices))
920
+ video_times = timestep_indices.index_select(0, generated_video_indices)
921
+ video_hidden = normalized_output[:generated_rows]
922
+ video_hidden = video_hidden * (1.0 + scale.index_select(0, video_times)) + shift.index_select(0, video_times)
923
+ generated_video_output = self.final_layer.video_out(video_hidden.float())
924
+ if condition_rows:
925
+ video_output = generated_video_output.new_zeros((1, hidden_states.shape[1], VIDEO_DIM))
926
+ video_output[0, condition_rows:] = generated_video_output
927
+ else:
928
+ video_output = generated_video_output.unsqueeze(0)
929
+
930
+ audio_times = timestep_indices.index_select(0, audio_indices)
931
+ audio_hidden = normalized_output[generated_rows:]
932
+ audio_hidden = audio_hidden * (1.0 + scale.index_select(0, audio_times)) + shift.index_select(0, audio_times)
933
+ audio_output = self.final_layer.audio_out(audio_hidden.float()).unsqueeze(0)
934
+
935
+ self._step_cache.update(
936
+ hidden_states,
937
+ audio_hidden_states,
938
+ video_output,
939
+ audio_output,
940
+ condition_rows,
941
+ )
942
+
943
+ if not return_dict:
944
+ return video_output, audio_output
945
+ return MiniMaxH3TransformerOutput(sample=video_output, audio_sample=audio_output)
946
+
947
+
948
+ def load_transformer() -> H3NVFP4Transformer:
949
+ if torch.version.cuda is None or int(torch.version.cuda.split(".")[0]) < 13:
950
+ raise RuntimeError("NVFP4 requires the CUDA 13 PyTorch build.")
951
+ from huggingface_hub import hf_hub_download
952
+
953
+ path = hf_hub_download(repo_id=NVFP4_REPO, filename=NVFP4_FILE)
954
+ transformer = H3NVFP4Transformer()
955
+ transformer.load(path)
956
+ print(f"[h3-nvfp4] loaded {NVFP4_REPO}/{NVFP4_FILE}", flush=True)
957
+ return transformer
958
+
959
+
960
+ def status() -> str:
961
+ return (
962
+ f"NVFP4 路 linear residual forecast {FORECAST_BLEND:g} / adaptive cache {EASYCACHE_THRESHOLD:g} 路 "
963
+ f"pruned AdaLN curve 路 fused QKV/QK-norm/RoPE 路 `{NVFP4_REPO}`"
964
+ )