multimodalart HF Staff commited on
Commit
9f6fb7b
ยท
verified ยท
1 Parent(s): 89cea9b

Match the reference's directed mask and sentence register

Browse files

Checked against github.com/Danzer1xxxxChan/H3-World (code/abot/action_script.py and code/diffsynth_h3_action.patch).

- caption_for: port annotate_from_keys9. The camera clause is never omitted โ€” `camera follows him` while moving, `camera holds steady` while still โ€” and the pan adverb attaches to the pan clause rather than after a join. Previously 11 of 17 presets produced a sentence shape absent from training; the default example gave 30 of 37 latent frames such a shape.
- directed_attention: implement mask_mod's leak_out and leak_in. A sentence row as a key is now readable only by itself and its own frame's video rows; as a query it reads its own frame's video rows only. Previously every sentence read every other sentence in all 50 blocks, breaking the cut-vertex property the reference documents as required.
- token refiner: attend block-diagonally, one segment per sentence, matching the reference's split refiner_cu. The LoRA adapts refiner_blocks.{0,1}, which were trained under that segmentation.
- _masked_lse: zero all-masked rows. Once the mask bites, most rows see no sentence key, and softmax over an all -inf row is NaN.
- version the processor patch guard so a reload cannot keep a stale closure.

Remaining divergence, documented in the README: sentences are still encoded jointly through the shared conditioner rather than independently.

Files changed (2) hide show
  1. README.md +41 -10
  2. app.py +182 -50
README.md CHANGED
@@ -53,16 +53,20 @@ So a 124-frame request (37 latent frames) is conditioned on a prompt that looks
53
 
54
  ```
55
  A third-person view of a man walking through a city intersection...
56
- the man walks forward
57
- the man walks forward
58
  ...
59
- the man walks forward and strafes right
60
  ...
61
  the man stands still, camera pans right sharply
62
  ```
63
 
64
- This Space builds that text from the key script for you โ€” `caption_for()` in `app.py` โ€” and shows you the resolved
65
- per-frame captions under **Conditioning** after every run.
 
 
 
 
66
 
67
  ## The directed attention mask
68
 
@@ -83,19 +87,46 @@ off its flash kernel for the whole 21k-row sequence). The keys are split into th
83
  | **B** | everything after the text block (~99% of keys) | flash, unmasked |
84
 
85
  Each returns its output and its log-sum-exp; the three are recombined with the online-softmax identity, which is
86
- numerically identical to one masked softmax over the full row. Only **video** queries are restricted (frame *i*'s rows
87
- see only sentence *i*); the captions themselves still see everything โ€” that is the "directed" part.
 
 
 
 
 
 
 
 
 
 
88
 
89
  The mask is toggleable in **Advanced** so you can see the difference. With it off, the same script produces a video
90
  that drifts through a blur of every action at once.
91
 
92
- Two guards, both in `app.py`:
93
 
94
  * The token spans are located by re-tokenizing prefixes of the prompt. If a BPE merge straddles a sentence boundary
95
  the spans would be wrong, so `build_conditioning_text()` verifies `cuts[-1] == total` and **refuses to mask** rather
96
  than mask the wrong rows.
97
- * `MiniMaxH3TokenRefinerBlock` runs the same attention module over the *text stream alone*. The processor detects that
98
- (the sequence is too short to contain the video block) and falls straight through to the stock path.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
99
 
100
  ## Architecture โ€” why two Spaces
101
 
 
53
 
54
  ```
55
  A third-person view of a man walking through a city intersection...
56
+ the man walks forward, camera follows him
57
+ the man walks forward, camera follows him
58
  ...
59
+ the man walks forward and strafes right, camera follows him
60
  ...
61
  the man stands still, camera pans right sharply
62
  ```
63
 
64
+ Every sentence ends in a camera clause โ€” there is no "no camera" shape in the training data. With no camera key
65
+ held the reference still states what the camera is doing: `camera follows him` while the character is moving,
66
+ `camera holds steady` while he is not. `caption_for()` in `app.py` is a port of the reference's
67
+ `annotate_from_keys9` (`code/abot/action_script.py`), sentence-for-sentence across all 130 reachable key
68
+ combinations, so the strings the model receives here are the strings it was trained on. The resolved per-frame
69
+ captions are shown under **Conditioning** after every run.
70
 
71
  ## The directed attention mask
72
 
 
87
  | **B** | everything after the text block (~99% of keys) | flash, unmasked |
88
 
89
  Each returns its output and its log-sum-exp; the three are recombined with the online-softmax identity, which is
90
+ numerically identical to one masked softmax over the full row.
91
+
92
+ **What "directed" means.** The reference's `mask_mod` cuts the annotations' *outgoing* edges, and only those:
93
+
94
+ * a sentence row, **as a key**, is readable by itself and by the video rows of its own latent frame โ€” the static
95
+ prompt, the keyframe anchors, the audio and **every other sentence** are blocked from reading it;
96
+ * a sentence row, **as a query**, reads the video rows of its own latent frame only. It still reads the static
97
+ prompt, the anchors and the audio freely โ€” that is information flowing *into* the sentence, giving it scene
98
+ grounding, not a bypass around frame *i*.
99
+
100
+ Blocking `A_j โ†’ A_k` is the load-bearing half. With it, `V_k` is a cut vertex between `A_k` and the rest of the
101
+ sequence; without it `A_j`'s content reaches `A_k` first, and `V_k` reading `A_k` is reading `A_j` too.
102
 
103
  The mask is toggleable in **Advanced** so you can see the difference. With it off, the same script produces a video
104
  that drifts through a blur of every action at once.
105
 
106
+ Three guards, all in `app.py`:
107
 
108
  * The token spans are located by re-tokenizing prefixes of the prompt. If a BPE merge straddles a sentence boundary
109
  the spans would be wrong, so `build_conditioning_text()` verifies `cuts[-1] == total` and **refuses to mask** rather
110
  than mask the wrong rows.
111
+ * Once the mask actually bites, most rows see *no* sentence key at all, so the masked region's log-sum-exp is `-inf`
112
+ for them. The merge treats that as weight zero, but `softmax` over an all-`-inf` row is `NaN` โ€” so `_masked_lse`
113
+ zeroes those rows explicitly.
114
+ * `MiniMaxH3TokenRefinerBlock` runs the same attention module over the *text stream alone*, which the processor
115
+ detects by its length (`cap_end` is the whole text length by construction). That stream is attended
116
+ **block-diagonally** โ€” one segment for the scene prompt, one per sentence โ€” matching the reference's split
117
+ `refiner_cu`. Running it as a single document lets sentences mix *before* the DiT backbone, and no mask on the DiT
118
+ side can close an upstream leak. The LoRA adapts `token_refiner.refiner_blocks.{0,1}`, so those deltas were trained
119
+ under this segmentation.
120
+
121
+ ### Known divergence from the reference
122
+
123
+ One difference remains, and it is upstream of everything above. The reference encodes **each sentence
124
+ independently** through the text encoder and concatenates the row blocks after the head, caching by string so a
125
+ repeated sentence is bit-identical. This Space builds one prompt string and sends it to the shared
126
+ `multimodalart/qwen3vl-conditioner`, so sentence *k* is contextualised against the image and every preceding
127
+ sentence before the DiT ever sees it. Closing it needs a raw text-encode endpoint on the conditioner, which it does
128
+ not currently expose. Related: the reference also gives annotation rows a mirrored positional offset
129
+ (`origin = text_len - s_k[-1] - 1`), which is not reachable from this side of the split either.
130
 
131
  ## Architecture โ€” why two Spaces
132
 
app.py CHANGED
@@ -14,10 +14,16 @@ Two things beyond a plain ``load_lora_weights`` are needed, and this Space imple
14
 
15
  2. **A directed attention mask.** MiniMax-H3 denoises one packed sequence
16
  ``[text | keyframe anchors | audio | video]`` under *full* self-attention, so by default every video
17
- row sees every sentence and the per-frame binding is lost. The patch below restricts each video
18
- row of latent frame ``f`` to sentence ``f`` โ€” and only among the sentences; the scene prompt, the
19
- keyframe anchors, the audio rows and the whole video block stay fully visible. It is *directed*:
20
- the sentences themselves are not restricted.
 
 
 
 
 
 
21
 
22
  Implementing that as a dense ``[S, S]`` mask would force SDPA off its flash kernel. Instead the
23
  attention is split by key region and recombined with an online-softmax (log-sum-exp) merge, which
@@ -134,41 +140,76 @@ PRESETS: dict[str, str] = {
134
  }
135
 
136
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
137
  def caption_for(keys: str, subject: str = DEFAULT_SUBJECT) -> str:
138
  """Render one latent frame's key state as the sentence H3-World was trained to read.
139
 
140
- The register is the model card's: ``"the man walks forward, camera pans left sharply"`` โ€” a
141
- locomotion clause from ``W/A/S/D``, an optional camera clause from ``I/J/K/L``, and ``F`` turning
142
- the camera adverb from ``slowly`` to ``sharply``.
 
 
 
143
  """
144
- held = set(keys.upper())
145
-
146
- motion = []
147
- if "W" in held:
148
- motion.append("walks forward")
149
- elif "S" in held:
150
- motion.append("walks backward")
151
- if "A" in held:
152
- motion.append("strafes left")
153
- elif "D" in held:
154
- motion.append("strafes right")
155
- body = " and ".join(motion) if motion else "stands still"
156
-
157
- camera = []
158
- if "J" in held:
159
- camera.append("camera pans left")
160
- elif "L" in held:
161
- camera.append("camera pans right")
162
- if "K" in held:
163
- camera.append("camera tilts up")
164
- elif "I" in held:
165
- camera.append("camera tilts down")
166
-
167
- sentence = f"{subject.strip() or DEFAULT_SUBJECT} {body}"
168
- if camera:
169
- adverb = "sharply" if "F" in held else "slowly"
170
- sentence += ", " + ", ".join(camera) + f" {adverb}"
171
- return sentence
172
 
173
 
174
  def parse_script(script: str, num_slots: int) -> list[str]:
@@ -315,6 +356,8 @@ def directed_plan(num_text_tokens: int, num_prompt_tokens: int, cuts, num_latent
315
  # โ”€โ”€ The directed-mask attention patch โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
316
 
317
  DIRECTED: dict = {"plan": None}
 
 
318
 
319
 
320
  def _lse_flash(q, k, v, scale):
@@ -370,7 +413,11 @@ def _masked_lse(query, key, value, mask, scale, chunk: int = 2048):
370
  if mask is not None:
371
  scores.masked_fill_(~mask[start:stop].unsqueeze(0).unsqueeze(0), float("-inf"))
372
  lse[:, :, start:stop] = torch.logsumexp(scores, dim=-1)
373
- probs = torch.softmax(scores, dim=-1).to(v.dtype)
 
 
 
 
374
  out[:, start:stop] = torch.matmul(probs, v).transpose(1, 2).to(out.dtype)
375
  del scores, probs, q
376
  return out, lse
@@ -396,35 +443,115 @@ def _merge_lse(out_a, lse_a, out_b, lse_b, chunk: int = 4096):
396
  return out_a, lse_a
397
 
398
 
399
- def directed_attention(query, key, value, plan):
400
- """Full self-attention with each video row's view of the *sentences* narrowed to its own.
401
 
402
- Returns ``None`` when ``plan`` does not describe this call's sequence โ€” the token refiner runs the
403
- same attention module over the text stream alone, and it must keep the unmasked path.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
404
  """
405
  length = query.shape[1]
406
  cap_start, cap_end = plan["cap_start"], plan["cap_end"]
407
  rows_per_frame = plan["rows_per_frame"]
408
  num_latent_frames = plan["num_latent_frames"]
 
 
 
 
 
 
 
409
  # The video block is the tail of the packed sequence, so its start follows from the geometry
410
  # alone โ€” no need to re-derive the keyframe or audio row counts here.
411
  video_start = length - num_latent_frames * rows_per_frame
412
  if video_start <= cap_end:
413
  return None
414
 
415
- scale = query.shape[-1] ** -0.5
 
 
 
 
416
  out, lse = _flash_lse(query, key[:, cap_end:], value[:, cap_end:], scale) # region B
417
  if cap_start > 0: # region A
418
  out, lse = _merge_lse(out, lse, *_flash_lse(query, key[:, :cap_start], value[:, :cap_start], scale))
419
 
420
- # region C โ€” the sentences, the only place the mask bites.
421
- visible = torch.ones(length, cap_end - cap_start, dtype=torch.bool, device=query.device)
422
- for index, (start, stop) in enumerate(plan["spans"]):
423
- rows = slice(video_start + index * rows_per_frame, video_start + (index + 1) * rows_per_frame)
424
- visible[rows] = False
425
- visible[rows, start - cap_start : stop - cap_start] = True
 
426
  out_c, lse_c = _masked_lse(query, key[:, cap_start:cap_end], value[:, cap_start:cap_end], visible, scale)
427
  out, _ = _merge_lse(out, lse, out_c, lse_c)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
428
  return out
429
 
430
 
@@ -437,9 +564,13 @@ def install_directed_processor() -> None:
437
  """
438
  from diffusers.models.transformers import transformer_minimax_h3 as h3
439
 
440
- if getattr(h3.MiniMaxH3AttnProcessor, "_h3world_directed", False):
 
 
 
441
  return
442
- base_call = h3.MiniMaxH3AttnProcessor.__call__
 
443
 
444
  def __call__(self, attn, hidden_states, rotary_emb=None, attention_mask=None):
445
  plan = DIRECTED.get("plan")
@@ -463,8 +594,9 @@ def install_directed_processor() -> None:
463
  attended = attended.flatten(2, 3).type_as(query)
464
  return attn.to_out[1](attn.to_out[0](attended))
465
 
 
466
  h3.MiniMaxH3AttnProcessor.__call__ = __call__
467
- h3.MiniMaxH3AttnProcessor._h3world_directed = True
468
 
469
 
470
  # โ”€โ”€ LoRA loading (weight folding) โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
 
14
 
15
  2. **A directed attention mask.** MiniMax-H3 denoises one packed sequence
16
  ``[text | keyframe anchors | audio | video]`` under *full* self-attention, so by default every video
17
+ row sees every sentence and the per-frame binding is lost. The patch below cuts the sentences'
18
+ *outgoing* edges, exactly as the reference's ``mask_mod`` does: sentence ``f`` is readable only by
19
+ itself and by the video rows of latent frame ``f`` โ€” the scene prompt, the keyframe anchors, the
20
+ audio and **every other sentence** are blocked from reading it โ€” and sentence ``f`` as a query reads
21
+ frame ``f``'s video rows only, while still reading the prompt, anchors and audio freely. That last
22
+ asymmetry is what "directed" means: information may flow *into* a sentence for scene grounding, but
23
+ nothing may route *around* frame ``f`` to reach it. Blocking ``A_j -> A_k`` is the load-bearing
24
+ half; without it ``A_j``'s content reaches ``A_k`` first and ``V_k`` reading ``A_k`` reads ``A_j``
25
+ too. The token refiner runs the same rule as one segment per sentence, since it sees the text
26
+ stream before the backbone does and no mask here could close a leak that happened there.
27
 
28
  Implementing that as a dense ``[S, S]`` mask would force SDPA off its flash kernel. Instead the
29
  attention is split by key region and recombined with an online-softmax (log-sum-exp) merge, which
 
140
  }
141
 
142
 
143
+ # The annotation vocabulary, ported verbatim from the author's own
144
+ # ``code/abot/action_script.py`` (github.com/Danzer1xxxxChan/H3-World). Training and
145
+ # inference in the reference run through the same `annotate_from_keys9`, so any drift here
146
+ # is a drift away from the only sentence shapes the LoRA ever saw.
147
+ MOTION = {"W": "walks forward", "S": "walks backward",
148
+ "A": "strafes left", "D": "strafes right"}
149
+ MOTION_ORDER = ("W", "S", "A", "D")
150
+ MOTION_IDLE = "stands still"
151
+ PAN_KEY = {"J": "left", "L": "right"}
152
+ TILT_KEY = {"I": "tilts down", "K": "tilts up"}
153
+ CAMERA_IDLE = "holds steady"
154
+ CAMERA_FOLLOW = "follows him"
155
+ FAST_KEY = "F"
156
+
157
+
158
+ def _purify(held: dict[str, bool], pairs) -> None:
159
+ """Cancel a pair of mutually exclusive keys that ended up both set.
160
+
161
+ The training pooling takes amax over a 4-frame window, so W-then-S inside one bin sets
162
+ both bits; the reference cancels them rather than emitting "walks forward and walks
163
+ backward". Reachable here through raw key entry (``WS``).
164
+ """
165
+ for a, b in pairs:
166
+ if held[a] and held[b]:
167
+ held[a] = held[b] = False
168
+
169
+
170
+ def _motion_clause(held: dict[str, bool]) -> str:
171
+ _purify(held, (("W", "S"), ("A", "D")))
172
+ words = [MOTION[name] for name in MOTION_ORDER if held[name]]
173
+ return " and ".join(words) if words else MOTION_IDLE
174
+
175
+
176
+ def _camera_clause(held: dict[str, bool], moving: bool) -> str:
177
+ """The camera clause is **never** omitted.
178
+
179
+ When no camera key is held the reference still says what the camera is doing โ€”
180
+ ``follows him`` while the character moves, ``holds steady`` while he does not โ€” so
181
+ every sentence in training ends in a camera clause. Dropping it (which this Space used
182
+ to do) puts the majority of latent frames in a sentence shape the model never saw.
183
+ """
184
+ _purify(held, (("J", "L"), ("I", "K")))
185
+ parts = []
186
+ for key, side in PAN_KEY.items():
187
+ if held[key]:
188
+ parts.append(f"pans {side} {'sharply' if held[FAST_KEY] else 'slowly'}")
189
+ for key, word in TILT_KEY.items():
190
+ if held[key]:
191
+ parts.append(word)
192
+ if parts:
193
+ return " and ".join(parts)
194
+ return CAMERA_FOLLOW if moving else CAMERA_IDLE
195
+
196
+
197
  def caption_for(keys: str, subject: str = DEFAULT_SUBJECT) -> str:
198
  """Render one latent frame's key state as the sentence H3-World was trained to read.
199
 
200
+ ``"the man walks forward and strafes right, camera pans left sharply"`` โ€” a locomotion
201
+ clause from ``W/A/S/D``, a camera clause from ``I/J/K/L`` with ``F`` turning its adverb
202
+ from ``slowly`` to ``sharply``, and the two joined by ``", camera "``.
203
+
204
+ ``subject`` stays configurable, but the camera clause's ``follows him`` is left exactly
205
+ as trained rather than agreed with it โ€” fidelity to the register beats grammar here.
206
  """
207
+ upper = keys.upper()
208
+ held = {name: (name in upper) for name in KEYS}
209
+ # each clause purifies its own copy: cancelling J/L must not disturb W/S
210
+ motion = _motion_clause(dict(held))
211
+ camera = _camera_clause(dict(held), motion != MOTION_IDLE)
212
+ return f"{subject.strip() or DEFAULT_SUBJECT} {motion}, camera {camera}"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
213
 
214
 
215
  def parse_script(script: str, num_slots: int) -> list[str]:
 
356
  # โ”€โ”€ The directed-mask attention patch โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
357
 
358
  DIRECTED: dict = {"plan": None}
359
+ # Bumped whenever the patched processor's behaviour changes; see `install_directed_processor`.
360
+ DIRECTED_PATCH_VERSION = "leak_out+leak_in+refiner_segments"
361
 
362
 
363
  def _lse_flash(q, k, v, scale):
 
413
  if mask is not None:
414
  scores.masked_fill_(~mask[start:stop].unsqueeze(0).unsqueeze(0), float("-inf"))
415
  lse[:, :, start:stop] = torch.logsumexp(scores, dim=-1)
416
+ # A row with NO visible key in this region is normal once the mask actually bites
417
+ # (static text, the keyframe cond and the audio rows see no sentence at all). Its
418
+ # log-sum-exp is -inf, which the merge handles as weight zero โ€” but softmax over
419
+ # an all -inf row is NaN, and NaN * 0 stays NaN, so it has to be zeroed here.
420
+ probs = torch.nan_to_num(torch.softmax(scores, dim=-1), nan=0.0).to(v.dtype)
421
  out[:, start:stop] = torch.matmul(probs, v).transpose(1, 2).to(out.dtype)
422
  del scores, probs, q
423
  return out, lse
 
443
  return out_a, lse_a
444
 
445
 
446
+ def _refiner_segmented(query, key, value, plan, scale):
447
+ """Block-diagonal attention over the text stream: the head, then one sentence each.
448
 
449
+ MiniMax-H3's token refiner runs before the DiT backbone, over the text rows alone. The
450
+ reference splits its ``refiner_cu`` into one segment per sentence, and is explicit
451
+ about why: run as a single document, *"by the time A_k reaches the DiT backbone, its
452
+ representation already contains information from other action sentences, and no mask
453
+ on the DiT side can close that upstream leak."* The LoRA adapts
454
+ ``token_refiner.refiner_blocks.{0,1}``, so those deltas were trained under this
455
+ segmentation.
456
+ """
457
+ cap_start = plan["cap_start"]
458
+ segments = ([(0, cap_start)] if cap_start > 0 else []) + list(plan["spans"])
459
+ out = torch.zeros_like(query)
460
+ for start, stop in segments:
461
+ if stop <= start:
462
+ continue
463
+ q = query[:, start:stop]
464
+ piece, _ = _flash_lse(q, key[:, start:stop], value[:, start:stop], scale)
465
+ out[:, start:stop] = piece.to(out.dtype)
466
+ return out
467
+
468
+
469
+ def directed_attention(query, key, value, plan):
470
+ """Self-attention under H3-World's directed mask.
471
+
472
+ Two shapes reach this function. The **token refiner**'s stream is the text rows alone
473
+ (``length == cap_end``) and gets segmented block-diagonally. The **packed sequence**
474
+ ``[text | keyframe anchors | audio | video]`` gets the directed mask proper, which is
475
+ two rules, both from the reference's ``mask_mod``:
476
+
477
+ * ``leak_out`` โ€” an annotation row as a *key* is readable only by itself and by the
478
+ video rows of its own latent frame. Static text, the keyframe cond, the audio and
479
+ **every other annotation** are blocked from reading it. This is the load-bearing one:
480
+ the reference's own comment says *"'A_k never reads another A_j' is required for this
481
+ to hold โ€” otherwise A_j's content would flow into A_k first, and V_k reading A_k
482
+ would be reading A_j too, and the cut-vertex property would break immediately."*
483
+ * ``leak_in`` โ€” an annotation row as a *query* reads the video rows of its own latent
484
+ frame only. It still reads static text, the cond and the audio freely: that is
485
+ information flowing *into* the sentence, giving it scene grounding, not a bypass.
486
+
487
+ "Directed" therefore means only the annotations' outgoing edges are cut โ€” not, as an
488
+ earlier version of this file had it, that the annotations are unrestricted as keys.
489
+
490
+ Returns ``None`` when ``plan`` does not describe this call's sequence at all.
491
  """
492
  length = query.shape[1]
493
  cap_start, cap_end = plan["cap_start"], plan["cap_end"]
494
  rows_per_frame = plan["rows_per_frame"]
495
  num_latent_frames = plan["num_latent_frames"]
496
+ scale = query.shape[-1] ** -0.5
497
+
498
+ # `cap_end` is the whole text length by construction (see `directed_plan`), so the
499
+ # refiner's own stream identifies itself by its length.
500
+ if length == cap_end:
501
+ return _refiner_segmented(query, key, value, plan, scale)
502
+
503
  # The video block is the tail of the packed sequence, so its start follows from the geometry
504
  # alone โ€” no need to re-derive the keyframe or audio row counts here.
505
  video_start = length - num_latent_frames * rows_per_frame
506
  if video_start <= cap_end:
507
  return None
508
 
509
+ spans = plan["spans"]
510
+
511
+ def video_rows(index):
512
+ return slice(video_start + index * rows_per_frame, video_start + (index + 1) * rows_per_frame)
513
+
514
  out, lse = _flash_lse(query, key[:, cap_end:], value[:, cap_end:], scale) # region B
515
  if cap_start > 0: # region A
516
  out, lse = _merge_lse(out, lse, *_flash_lse(query, key[:, :cap_start], value[:, :cap_start], scale))
517
 
518
+ # region C โ€” the sentence keys. Default DENIED: only the sentence's own rows and the
519
+ # video rows of its own latent frame may read it.
520
+ visible = torch.zeros(length, cap_end - cap_start, dtype=torch.bool, device=query.device)
521
+ for index, (start, stop) in enumerate(spans):
522
+ column = slice(start - cap_start, stop - cap_start)
523
+ visible[start:stop, column] = True # same_ann: A_k reads A_k
524
+ visible[video_rows(index), column] = True # v_reads_own: V_k reads A_k
525
  out_c, lse_c = _masked_lse(query, key[:, cap_start:cap_end], value[:, cap_start:cap_end], visible, scale)
526
  out, _ = _merge_lse(out, lse, out_c, lse_c)
527
+
528
+ # `leak_in`: the annotation rows' own output is rebuilt from scratch, because region B
529
+ # above gave them every video row and they may only have their own frame's. They are a
530
+ # contiguous slice and there are only a few hundred of them, so recomputing is cheaper
531
+ # than trying to subtract the wrong contribution back out of a merged softmax.
532
+ q_ann = query[:, cap_start:cap_end]
533
+ out_ann, lse_ann = _flash_lse( # cond + audio, unrestricted
534
+ q_ann, key[:, cap_end:video_start], value[:, cap_end:video_start], scale)
535
+ if cap_start > 0: # static text, unrestricted
536
+ out_ann, lse_ann = _merge_lse(
537
+ out_ann, lse_ann, *_flash_lse(q_ann, key[:, :cap_start], value[:, :cap_start], scale))
538
+ own = torch.zeros(cap_end - cap_start, cap_end - cap_start, dtype=torch.bool, device=query.device)
539
+ for start, stop in spans: # its own sentence only
540
+ local = slice(start - cap_start, stop - cap_start)
541
+ own[local, local] = True
542
+ out_ann, lse_ann = _merge_lse(
543
+ out_ann, lse_ann,
544
+ *_masked_lse(q_ann, key[:, cap_start:cap_end], value[:, cap_start:cap_end], own, scale))
545
+ # its own latent frame's video rows, one small attention per frame
546
+ for index, (start, stop) in enumerate(spans):
547
+ local = slice(start - cap_start, stop - cap_start)
548
+ rows = video_rows(index)
549
+ piece, piece_lse = _flash_lse(q_ann[:, local], key[:, rows], value[:, rows], scale)
550
+ merged, merged_lse = _merge_lse(
551
+ out_ann[:, local].clone(), lse_ann[:, :, local].clone(), piece, piece_lse)
552
+ out_ann[:, local] = merged
553
+ lse_ann[:, :, local] = merged_lse
554
+ out[:, cap_start:cap_end] = out_ann.to(out.dtype)
555
  return out
556
 
557
 
 
564
  """
565
  from diffusers.models.transformers import transformer_minimax_h3 as h3
566
 
567
+ # Version the guard, not just its presence: a reload re-executes this module but the
568
+ # patched class object survives it, so a bare boolean would keep the PREVIOUS closure
569
+ # โ€” still pointing at the previous `directed_attention` โ€” installed forever.
570
+ if getattr(h3.MiniMaxH3AttnProcessor, "_h3world_directed", None) == DIRECTED_PATCH_VERSION:
571
  return
572
+ base_call = getattr(h3.MiniMaxH3AttnProcessor, "_h3world_base_call",
573
+ h3.MiniMaxH3AttnProcessor.__call__)
574
 
575
  def __call__(self, attn, hidden_states, rotary_emb=None, attention_mask=None):
576
  plan = DIRECTED.get("plan")
 
594
  attended = attended.flatten(2, 3).type_as(query)
595
  return attn.to_out[1](attn.to_out[0](attended))
596
 
597
+ h3.MiniMaxH3AttnProcessor._h3world_base_call = base_call
598
  h3.MiniMaxH3AttnProcessor.__call__ = __call__
599
+ h3.MiniMaxH3AttnProcessor._h3world_directed = DIRECTED_PATCH_VERSION
600
 
601
 
602
  # โ”€โ”€ LoRA loading (weight folding) โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€