Spaces:
Running on Zero
Match the reference's directed mask and sentence register
Browse filesChecked 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.
|
@@ -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 |
-
|
| 65 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.
|
| 87 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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 |
-
*
|
| 98 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
|
|
@@ -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
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 141 |
-
|
| 142 |
-
|
|
|
|
|
|
|
|
|
|
| 143 |
"""
|
| 144 |
-
|
| 145 |
-
|
| 146 |
-
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 400 |
-
"""
|
| 401 |
|
| 402 |
-
|
| 403 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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
|
| 421 |
-
|
| 422 |
-
|
| 423 |
-
|
| 424 |
-
|
| 425 |
-
visible[
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
| 441 |
return
|
| 442 |
-
base_call = h3.MiniMaxH3AttnProcessor
|
|
|
|
| 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 =
|
| 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) โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
|