Spaces:
Running on Zero
Running on Zero
File size: 16,371 Bytes
4d565a4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 | """First-block cache for MiniMax-H3: skip the trunk on steps whose block-0 residual barely moved.
Block 0 and the final AdaLN head run at the true timestep on **every** step; only blocks 1..49 are skipped, and only
while the residual they would have been handed looks like the one from the last step that actually ran. The decision
signal is the relative L1 between this step's block-0 residual (`block0_out - block0_in`) and the residual of the last
*computed* step, over the whole packed sequence. On a skip the trunk's contribution is replayed as a cached residual,
`(final_trunk_out - block0_out)` of that computed step.
Ported from `duckyshell/ComfyUI-MiniMaxH3-FirstBlockCache` (`nodes.py` @ 725973c) — same signal, same protected window
(10%-95% of the schedule, converted through the video shift of 12.0) and the same cap of two consecutive skips, which
is their "H3 Safe" preset at threshold 0.08.
**The audio exemption is not theirs.** duckyshell has no audio term anywhere: the decision is ~98% video by row count
and every audio row rides the stale trunk residual, which is what costs the soundtrack its energy. Audio runs its own
schedule (shift 3 against video's 12, diverging by up to 11x in rate across 20 steps), so on a skip the audio rows of
the trunk output are instead a 2-point **linear** extrapolation of the last two actually-computed audio features, in
the audio sigma coordinate — `xmarre/ComfyUI-Spectrum-MiniMax-H3`'s `audio_blend_weight=0.0` path, applied where
Spectrum applies it (the post-trunk hidden feature, ahead of the head that still runs at the true timestep) and fixed
to the right coordinate. Following ComfyUI PR #15390, no carried audio tensor is ever mutated: every write is a fresh
tensor out of `index_copy`.
It earns its place. Running this Space's own request at threshold 0.08 with `H3_FBC_AUDIO_EXEMPT=0` — the duckyshell
mechanism verbatim, same 13 skipped forwards, same video to within 0.2 dB — costs the soundtrack **16.7% of its RMS**
(0.0483 against an uncached 0.0580), while the exemption holds it at 1.04x. That is the same direction the offline
study measured on a near-silent clip, three times the size on one with real audio energy.
**A threshold does not travel across step counts.** The signature shrinks as the schedule is subdivided, so the same
number gates far more loosely at more steps. Measured on this Space — 960x544, 124 frames, one image reference,
seed 42, AoTI blocks, at its **default 28 steps** (27 forwards) — against the same request with `H3_FBC=0`:
| threshold | skipped | denoise loop | end to end | audio RMS vs uncached |
|---|---|---|---|---|
| 0.03 | 0 / 27 | 1.02x | 1.06x | 1.000 (bitwise-identical audio) |
| 0.05 | 9 / 27 | 1.46x | 1.36x | 0.956 |
| 0.08 | 13 / 27 | 2.19x | 1.82x | 1.040 |
The offline study calibrated 0.08 over a 20-step schedule, where it skipped 7 of 19 forwards. At 28 steps that same
0.08 skips 13 of 27 — nearly half — and the sampled video visibly re-rolls its background detail. **0.05 is the
default** because it reproduces the skip fraction that study validated (33% here against 37% there); 0.08 is the
aggressive setting. Below the signature floor — 0.03 skipped nothing at all here — a threshold buys nothing and still
pays for the signature, so lower is not safer, it is just slower.
A cached request is not the uncached one: the trajectory moves, so the video is a different sample of the same prompt
(same shot, same subject, same quality — different signage and background detail). `H3_FBC=0` restores today's output
exactly, and is worth reaching for when a request has to reproduce a specific earlier result.
This composes with `h3_aoti`: that module patches each of the 50 blocks' own `forward` and they stay a real
`ModuleList`, so skipping the trunk simply does not call blocks 1..49 that step. `LazyAOTIModel` rebinds its constants
whenever the weights dict it is handed changes identity, and the first forward of a request never skips, so every
block has bound its own weights before any step is cached.
"""
from __future__ import annotations
import contextlib
import os
import types
ENABLED = os.environ.get("H3_FBC", "1") == "1"
# Relative-L1 gate on the block-0 residual, calibrated at this Space's default 28 steps — see the table above.
THRESHOLD = float(os.environ.get("H3_FBC_THRESHOLD", "0.05"))
# duckyshell's cap. Without it the gate compares against an ever-older computed step and drifts away unbounded.
MAX_CONSECUTIVE_HITS = int(os.environ.get("H3_FBC_MAX_CONSECUTIVE", "2"))
# The protected head and tail of the schedule, as fractions, converted to sigma through the video shift below.
START_PERCENT = float(os.environ.get("H3_FBC_START_PERCENT", "0.10"))
END_PERCENT = float(os.environ.get("H3_FBC_END_PERCENT", "0.95"))
# `MiniMaxH3SetTimestepsStep` builds the video schedule at shift 12.0 and the audio one at 3.0.
VIDEO_SHIFT = float(os.environ.get("H3_FBC_VIDEO_SHIFT", "12.0"))
AUDIO_EXEMPT = os.environ.get("H3_FBC_AUDIO_EXEMPT", "1") == "1"
# Every keyword `MiniMaxH3LoopDenoiser` passes. It filters the packed-sequence layout through
# `inspect.signature(transformer.forward).parameters`, so a replacement forward that drops a name silently stops
# receiving it; `install` refuses rather than let that happen quietly.
FORWARD_PARAMETERS = (
"hidden_states",
"audio_hidden_states",
"encoder_hidden_states",
"timestep",
"timestep_indices",
"token_tags",
"position_ids",
"video_indices",
"audio_indices",
"text_indices",
"attention_kwargs",
"return_dict",
)
def status() -> str:
return (
f"first-block cache **on** · threshold `{THRESHOLD}` · audio exemption "
f"{'on' if AUDIO_EXEMPT else 'off'}"
if ENABLED
else "first-block cache **off** (`H3_FBC=1` to skip the trunk on steady steps)"
)
def _shifted_sigma(u: float, shift: float) -> float:
return shift * u / (1.0 + (shift - 1.0) * u)
def _rel_l1(current, previous) -> float:
numerator = (current.float() - previous.float()).abs().mean()
denominator = previous.float().abs().mean().clamp(min=1e-8)
return float((numerator / denominator).item())
class _State:
def __init__(self, threshold: float, steps: int, audio_exempt: bool):
self.threshold = threshold
self.steps = steps
self.audio_exempt = audio_exempt
# duckyshell reads the window as sigma bounds: a flow model's sigma at `u = 1 - percent`, shifted.
self.start_sigma = _shifted_sigma(1.0 - START_PERCENT, VIDEO_SHIFT)
self.end_sigma = _shifted_sigma(1.0 - END_PERCENT, VIDEO_SHIFT)
self.original = None
self.failed = False
self.consecutive_hits = 0
self.prev_first_residual = None
self.tail_residual = None
self.audio_history = [] # [(sigma_audio, audio rows of the trunk output)], newest last, at most two
self.computed = 0
self.skipped = 0
def _cached_forward(
self,
state,
hidden_states,
audio_hidden_states,
encoder_hidden_states,
timestep,
timestep_indices,
token_tags,
position_ids,
video_indices,
audio_indices,
text_indices,
return_dict,
):
"""`MiniMaxH3Transformer3DModel.forward` with the block loop split at block 0.
Everything outside the loop is that method verbatim, at the `diffusers` commit `requirements.txt` pins; keep the
two in step when the pin moves.
"""
import torch
from diffusers.models.transformers.transformer_minimax_h3 import (
MINIMAX_H3_MODALITY_NUM,
MiniMaxH3TransformerOutput,
)
sequence_length = position_ids.shape[0]
rotary_emb = self.rope(position_ids)
video_embeds = self.proj_in(hidden_states.to(self.proj_in.weight.dtype))
audio_embeds = self.audio_proj_in(audio_hidden_states.to(self.audio_proj_in.weight.dtype))
text_embeds = self.context_embedder(encoder_hidden_states.to(self.context_embedder.weight.dtype))
text_embeds = self.token_refiner(text_embeds)
packed = text_embeds.new_zeros((text_embeds.shape[0], sequence_length, text_embeds.shape[-1]))
packed = packed.index_copy(1, text_indices, text_embeds)
packed = packed.index_copy(1, video_indices, video_embeds.to(text_embeds.dtype))
packed = packed.index_copy(1, audio_indices, audio_embeds.to(text_embeds.dtype))
temb = self.time_proj(timestep)
temb = self.time_embedder(temb.to(self.time_embedder.linear_1.weight.dtype))
adaln_indices = timestep_indices * MINIMAX_H3_MODALITY_NUM + token_tags.clamp(min=0)
attention_mask = None
is_pad = token_tags < 0
if bool(is_pad.any()):
attention_mask = is_pad[None, :] == is_pad[:, None]
blocks = self.transformer_blocks
block0_out = blocks[0](packed, temb, adaln_indices, rotary_emb, attention_mask)
first_residual = block0_out - packed
# The generated rows trail their modality's index list, so the last row of each carries that stream's live noise
# level. The scheduler exposes `timesteps = 1 - sigmas[:-1]`.
sigma_video = 1.0 - float(timestep[timestep_indices[video_indices[-1]]].item())
sigma_audio = 1.0 - float(timestep[timestep_indices[audio_indices[-1]]].item())
use_cache = False
if (
state.prev_first_residual is not None
and state.tail_residual is not None
and state.prev_first_residual.shape == first_residual.shape
and state.consecutive_hits < MAX_CONSECUTIVE_HITS
and state.end_sigma <= sigma_video <= state.start_sigma
):
use_cache = _rel_l1(first_residual, state.prev_first_residual) <= state.threshold
if use_cache:
state.consecutive_hits += 1
state.skipped += 1
trunk_out = block0_out + state.tail_residual
if state.audio_exempt and state.audio_history:
sigma_prev, feature_prev = state.audio_history[-1]
if len(state.audio_history) == 2 and abs(sigma_prev - state.audio_history[-2][0]) > 1e-8:
sigma_prev2, feature_prev2 = state.audio_history[-2]
ratio = (sigma_audio - sigma_prev) / (sigma_prev - sigma_prev2)
audio_feature = feature_prev + (feature_prev - feature_prev2) * ratio
else:
audio_feature = feature_prev
trunk_out = trunk_out.index_copy(1, audio_indices, audio_feature.to(trunk_out.dtype))
else:
state.consecutive_hits = 0
state.computed += 1
trunk_out = block0_out
for block in blocks[1:]:
trunk_out = block(trunk_out, temb, adaln_indices, rotary_emb, attention_mask)
state.tail_residual = (trunk_out - block0_out).detach()
state.prev_first_residual = first_residual.detach()
if state.audio_exempt:
state.audio_history.append((sigma_audio, trunk_out.index_select(1, audio_indices).detach().float()))
state.audio_history = state.audio_history[-2:]
out = self.norm_out(trunk_out, temb, timestep_indices).to(self.proj_out.weight.dtype)
video_output = self.proj_out(out).index_select(1, video_indices)
audio_output = self.audio_proj_out(out).index_select(1, audio_indices)
if not return_dict:
return (video_output, audio_output)
return MiniMaxH3TransformerOutput(sample=video_output, audio_sample=audio_output)
def _forward(
self,
hidden_states,
audio_hidden_states,
encoder_hidden_states,
timestep,
timestep_indices,
token_tags,
position_ids,
video_indices,
audio_indices,
text_indices,
attention_kwargs=None,
return_dict: bool = True,
):
"""The installed forward. Anything it cannot serve — a LoRA scale, an unexpected layout, a bug — is handed to the
original forward instead, for this call and every later one, so a cached request can degrade to an uncached one but
never to a failed one."""
state = self._h3_fbc
original = dict(
hidden_states=hidden_states,
audio_hidden_states=audio_hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep,
timestep_indices=timestep_indices,
token_tags=token_tags,
position_ids=position_ids,
video_indices=video_indices,
audio_indices=audio_indices,
text_indices=text_indices,
attention_kwargs=attention_kwargs,
return_dict=return_dict,
)
# `apply_lora_scale` decorates the real forward and this one is not it, so a request that actually scales a LoRA
# goes down the original path rather than silently losing its scale.
if state.failed or (attention_kwargs or {}).get("scale") is not None:
return state.original(**original)
try:
return _cached_forward(
self,
state,
hidden_states,
audio_hidden_states,
encoder_hidden_states,
timestep,
timestep_indices,
token_tags,
position_ids,
video_indices,
audio_indices,
text_indices,
return_dict,
)
except Exception as error:
state.failed = True
print(f"[h3-fbc] disabled for this request ({type(error).__name__}: {error}); running uncached", flush=True)
return state.original(**original)
def install(transformer, steps: int = 0, threshold: float = THRESHOLD, audio_exempt: bool = AUDIO_EXEMPT) -> bool:
"""Bind the caching forward onto `transformer`. Returns whether it went on.
`accelerate`'s `add_hook_to_module` — what `ComponentsManager.enable_auto_cpu_offload` installs — moves the real
forward to `_old_forward` and puts its own onload wrapper in `forward`. Replacing `forward` there would step over
the wrapper and run the block stack against weights still on the host, so the replacement goes into `_old_forward`
whenever the hook is present.
"""
import inspect
if getattr(transformer, "_h3_fbc", None) is not None:
return True
hooked = hasattr(transformer, "_hf_hook") and hasattr(transformer, "_old_forward")
current = transformer._old_forward if hooked else transformer.forward
missing = [name for name in FORWARD_PARAMETERS if name not in inspect.signature(current).parameters]
if missing:
print(f"[h3-fbc] this transformer's forward has no {missing}; running uncached", flush=True)
return False
if not hasattr(transformer, "transformer_blocks") or len(transformer.transformer_blocks) < 2:
print("[h3-fbc] no block stack to skip; running uncached", flush=True)
return False
state = _State(threshold, steps, audio_exempt)
state.original = current
transformer._h3_fbc = state
bound = types.MethodType(_forward, transformer)
if hooked:
transformer._old_forward = bound
else:
transformer.forward = bound
return True
def uninstall(transformer) -> None:
state = getattr(transformer, "_h3_fbc", None)
if state is None:
return
if hasattr(transformer, "_hf_hook") and hasattr(transformer, "_old_forward"):
transformer._old_forward = state.original
else:
transformer.__dict__.pop("forward", None)
del transformer._h3_fbc
total = state.computed + state.skipped
if total:
print(
f"[h3-fbc] {state.skipped}/{total} forwards served from cache "
f"(threshold {state.threshold}, audio exemption {'on' if state.audio_exempt else 'off'})",
flush=True,
)
@contextlib.contextmanager
def enabled(transformer, steps: int = 0):
"""Cache the trunk for the duration of one request. The state is per-request by construction — a residual only ever
means something within the schedule it was measured on — and nothing in here can raise into the request."""
installed = False
if ENABLED:
try:
installed = install(transformer, steps=steps)
except Exception as error:
print(f"[h3-fbc] install failed ({type(error).__name__}: {error}); running uncached", flush=True)
try:
yield installed
finally:
if installed:
try:
uninstall(transformer)
except Exception as error:
print(f"[h3-fbc] uninstall failed ({type(error).__name__}: {error})", flush=True)
|