Support per-token video and audio latent noise masks on MiniMax-H3 (#15375)

This commit is contained in:
Barish Ozbay
2026-08-17 23:18:18 -04:00
committed by GitHub
parent c1739380c6
commit ff6c8a8af1
3 changed files with 144 additions and 18 deletions

View File

@@ -74,6 +74,17 @@ def _axis_from_sqrt_area(dim, patch, sqrt_area):
return (torch.arange(n, dtype=torch.float64) * (ratio / n) + (1.0 - ratio) / 2.0) * 32.0
def mask_row_values(mask, latent_t, lat_h, lat_w):
# [T, H, W] denoise mask (1 = generate) -> per-2x2-patch-row float in [0, 1],
# None when every row fully generates
m = torch.nn.functional.pad(mask, (0, lat_w - mask.shape[-1], 0, lat_h - mask.shape[-2]), mode="replicate")
m = m.reshape(latent_t, lat_h // 2, 2, lat_w // 2, 2).amax(dim=(2, 4))
values = m.reshape(-1)
if bool((values >= 1.0 - 1e-3).all()):
return None
return values
def _frame_grid(h, w):
# area-normalized (h, w) coordinates of one latent frame's 2x2-patch rows
area = math.sqrt(h * w)
@@ -212,17 +223,22 @@ class AdalnProj(nn.Module):
return x.chunk(self.expand, dim=-1)
def _mod_row(vecs, row, dtype):
# row is a mod-row index, or a per-token LongTensor of mod-row indices
return vecs[row].to(dtype)
def _mod_scale_shift(h, shift, scale, segments):
# segments: [(start, stop, mod_row)] covering h contiguously.
for a, b, row in segments:
h[a:b].mul_(1.0 + scale[row].to(h.dtype)).add_(shift[row].to(h.dtype))
h[a:b].mul_(1.0 + _mod_row(scale, row, h.dtype)).add_(_mod_row(shift, row, h.dtype))
return h
def _mod_gate(x, gate, other, segments):
# other is the fresh attn/mlp output: accumulate the gated residual into the stream in place, one fused kernel per segment
for a, b, row in segments:
x[a:b].addcmul_(other[a:b], gate[row].to(x.dtype))
x[a:b].addcmul_(other[a:b], _mod_row(gate, row, x.dtype))
return x
@@ -288,13 +304,15 @@ class FinalLayer(nn.Module):
self.audio_out = operations.Linear(hidden, audio_dim, bias=True, dtype=torch.float32, device=device)
def forward(self, x, t_emb, video_seg, audio_seg):
# video_seg / audio_seg: (start, stop, timestep_row) of the target streams
# video_seg / audio_seg: (start, stop, row) of the target streams, where row
# is a mod-row index or a per-token blend (see _mod_row)
shift, scale = self.adaln_proj(t_emb)
va, vb, vrow = video_seg
aa, ab, arow = audio_seg
hv = (self.norm(x[va:vb]) * (1.0 + scale[vrow]) + shift[vrow]).to(torch.float32)
ha = (self.norm(x[aa:ab]) * (1.0 + scale[arow]) + shift[arow]).to(torch.float32)
return self.video_out(hv), self.audio_out(ha)
def mod(seg):
a, b, row = seg
return (self.norm(x[a:b]) * (1.0 + _mod_row(scale, row, scale.dtype)) + _mod_row(shift, row, shift.dtype)).to(torch.float32)
return self.video_out(mod(video_seg)), self.audio_out(mod(audio_seg))
class PackedLayout:
@@ -506,7 +524,7 @@ class MiniMaxH3Model(nn.Module):
rows.append(r.to(device))
return torch.cat(rows, dim=0) if rows else None
def forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, **kwargs):
def forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, denoise_mask=None, audio_denoise_mask=None, **kwargs):
# the sampler carries the audio as (sigma_v / sigma_a) * x_audio; undo it outside
# the wrappers so they and the network see the stream's own latent and velocity
scale = float((minimax_payload or {}).get("audio_scale", 1.0))
@@ -523,7 +541,8 @@ class MiniMaxH3Model(nn.Module):
self._forward,
self,
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options)
).execute(x, timestep, context, transformer_options, minimax_payload=minimax_payload, **kwargs)
).execute(x, timestep, context, transformer_options, minimax_payload=minimax_payload,
denoise_mask=denoise_mask, audio_denoise_mask=audio_denoise_mask, **kwargs)
if scale != 1.0:
# d/d(sigma_v) of the carried variable
@@ -531,7 +550,7 @@ class MiniMaxH3Model(nn.Module):
+ (1.0 + (scale - 1.0) * sigma_a).to(out[1].dtype) * out[1])
return out
def _forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, **kwargs):
def _forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, denoise_mask=None, audio_denoise_mask=None, **kwargs):
video_x, audio_x = x[0], x[1]
orig_t, orig_h, orig_w = video_x.shape[2], video_x.shape[3], video_x.shape[4]
video_x = comfy.ldm.common_dit.pad_to_patch_size(video_x, self.patch_size)
@@ -561,16 +580,47 @@ class MiniMaxH3Model(nn.Module):
# distinct timesteps are known analytically: text/pad follow video, cond rows pin near 1
vis_aug = float(payload.get("visual_cond_noise_aug", VISUAL_COND_TIMESTEP))
aud_aug = float(payload.get("audio_cond_noise_aug", AUDIO_COND_TIMESTEP))
has_vis_cond = any(k in ("cond", "ref_img") for _, _, k in layout.segments)
has_aud_cond = any(k in ("cond_audio", "ref_audio") for _, _, k in layout.segments)
seg_t = {"text": t_v, "video": t_v, "audio": t_a,
"cond": max(t_v, vis_aug), "ref_img": max(t_v, vis_aug),
"cond_audio": max(t_a, aud_aug), "ref_audio": max(t_a, aud_aug)}
unique_t = sorted({t_v, t_a} | ({seg_t["cond"]} if has_vis_cond else set())
| ({seg_t["ref_audio"]} if has_aud_cond else set()))
# masked rows run at their own strength: mask value m puts a row at sigma = m * sigma_stream,
# so its label is 1 - m * sigma, clamped at the cond timestep for fully preserved rows
t_pin_v = max(t_v, VISUAL_COND_TIMESTEP)
t_pin_a = max(t_a, AUDIO_COND_TIMESTEP)
video_rows_t = None
audio_rows_t = None
if denoise_mask is not None:
m = mask_row_values(denoise_mask[0, 0].to(torch.float32), latent_t, lat_h, lat_w)
if m is not None:
rows_t = (1.0 - m * sigma_v.to(m.device)).clamp(max=t_pin_v)
if rows_t.unique().numel() == 1:
seg_t["video"] = float(rows_t[0])
else:
video_rows_t = rows_t
if audio_denoise_mask is not None:
m = audio_denoise_mask[0, 0].to(torch.float32).reshape(-1)
if not bool((m >= 1.0 - 1e-3).all()):
sigma_a = 1.0 - t_a
rows_t = (1.0 - m * sigma_a).clamp(max=t_pin_a)
if rows_t.unique().numel() == 1:
seg_t["audio"] = float(rows_t[0])
else:
audio_rows_t = rows_t
unique_t = sorted({t_v, t_a} | {seg_t[k] for _, _, k in layout.segments}
| (set(video_rows_t.unique().tolist()) if video_rows_t is not None else set())
| (set(audio_rows_t.unique().tolist()) if audio_rows_t is not None else set()))
t_row = {t: i for i, t in enumerate(unique_t)}
seg_tag = {"text": 1, "video": 0, "audio": 2, "cond": 0, "ref_img": 0, "cond_audio": 2, "ref_audio": 2}
def rows_to_mod_index(rows_t, tag):
# per-row timestep values -> per-row mod-row indices into the t_emb table
levels = rows_t.unique()
base = torch.tensor([t_row[v] * 3 + tag for v in levels.tolist()],
dtype=torch.long, device=rows_t.device)
return base[torch.searchsorted(levels, rows_t)]
text_tags = payload.get("text_token_tags")
mod_segments = []
for a, b, kind in layout.segments:
@@ -583,6 +633,10 @@ class MiniMaxH3Model(nn.Module):
if i == b - a or tags[i] != tags[run_start]:
mod_segments.append((a + run_start, a + i, row_base + int(tags[run_start])))
run_start = i
elif kind == "video" and video_rows_t is not None:
mod_segments.append((a, b, rows_to_mod_index(video_rows_t, seg_tag[kind])))
elif kind == "audio" and audio_rows_t is not None:
mod_segments.append((a, b, rows_to_mod_index(audio_rows_t, seg_tag[kind])))
else:
mod_segments.append((a, b, row_base + seg_tag[kind]))
@@ -659,8 +713,16 @@ class MiniMaxH3Model(nn.Module):
comfy.model_prefetch.prefetch_queue_pop(prefetch_queue, device, None)
# target streams are single contiguous segments (audio then video, last two)
video_seg = next((a, b, t_row[seg_t["video"]]) for a, b, k in layout.segments if k == "video")
audio_seg = next((a, b, t_row[seg_t["audio"]]) for a, b, k in layout.segments if k == "audio")
va, vb, _ = next(s for s in layout.segments if s[2] == "video")
aa, ab, _ = next(s for s in layout.segments if s[2] == "audio")
if video_rows_t is not None:
video_seg = (va, vb, rows_to_mod_index(video_rows_t, 0) // 3)
else:
video_seg = (va, vb, t_row[seg_t["video"]])
if audio_rows_t is not None:
audio_seg = (aa, ab, rows_to_mod_index(audio_rows_t, 0) // 3)
else:
audio_seg = (aa, ab, t_row[seg_t["audio"]])
v, a = self.final_layer(h, t_emb, video_seg, audio_seg)
video_out = unpatchify_video(v, latent_t, lat_h // 2, lat_w // 2, self.latents_dim, self.patch_size)

View File

@@ -2179,6 +2179,11 @@ class MiniMaxH3(BaseModel):
payload["seed"] = kwargs.get("seed", 0)
# same value process_latent_in/out used, so the model never undoes a scale that was not applied
payload["audio_scale"] = self.audio_scale()
denoise_mask = kwargs.get("denoise_mask", None)
if denoise_mask is not None:
out.update(self._denoise_mask_conds(denoise_mask, latent_shapes))
if cross_attn is not None and latent_shapes is not None and len(latent_shapes) > 1:
# packed layout built once per sampling run, h/w rounded up to the DiT's 2x2 patch
vs = latent_shapes[0]
@@ -2189,6 +2194,65 @@ class MiniMaxH3(BaseModel):
out['minimax_payload'] = comfy.conds.CONDConstant(payload)
return out
def _pool_masks_to_token_grid(self, masks):
# pool the per-pixel masks to the label grid with amax: video per 2x2 DiT patch, audio per latent frame
video_mask = masks[0]
h, w = video_mask.shape[-2:]
ph, pw = self.diffusion_model.patch_size[1:]
lead = video_mask.shape[:-2]
video_mask = torch.nn.functional.pad(video_mask.reshape((-1,) + video_mask.shape[-3:]), (0, -w % pw, 0, -h % ph), mode="replicate")
video_mask = video_mask.reshape(lead + video_mask.shape[-2:])
video_mask = video_mask.reshape(video_mask.shape[:-2] + (video_mask.shape[-2] // ph, ph, video_mask.shape[-1] // pw, pw)).amax(dim=(-3, -1))
pooled = [video_mask.repeat_interleave(ph, dim=-2).repeat_interleave(pw, dim=-1)[..., :h, :w]]
if len(masks) > 1:
audio_mask = masks[1].amax(dim=1, keepdim=True)
pooled.append(audio_mask.expand_as(masks[1]).contiguous())
return pooled
def _token_grid_masks(self, denoise_mask, latent_shapes):
masks = utils.unpack_latents(denoise_mask, latent_shapes)
return [torch.ceil(mask * 256.0) / 256.0 for mask in self._pool_masks_to_token_grid(masks)]
def _denoise_mask_values(self, denoise_mask, latent_shapes):
if latent_shapes is None or len(latent_shapes) < 2:
return {}
masks = self._token_grid_masks(denoise_mask, latent_shapes)
out = {}
if torch.amin(masks[0]).item() < 1.0 - 1e-3:
out['denoise_mask'] = masks[0][:1, :1].clone()
if torch.amin(masks[1]).item() < 1.0 - 1e-3:
out['audio_denoise_mask'] = masks[1][:1].amax(dim=1, keepdim=True)
return out
def _denoise_mask_conds(self, denoise_mask, latent_shapes):
return {name: comfy.conds.CONDRegular(value) for name, value in self._denoise_mask_values(denoise_mask, latent_shapes).items()}
def scale_latent_inpaint(self, sigma, noise, latent_image, x=None, denoise_mask=None, **kwargs):
# preserved regions run at the cond timestep, inject them at cond strength
shapes = self.latent_shapes
if shapes is None or len(shapes) < 2:
return super().scale_latent_inpaint(sigma=sigma, noise=noise, latent_image=latent_image, **kwargs)
cleans = utils.unpack_latents(latent_image, shapes)
noises = utils.unpack_latents(noise, shapes)
aug = comfy.ldm.minimax.model.VISUAL_COND_TIMESTEP # H3's video timestep is 0.999 by default
cleans[0] = aug * cleans[0] + (1.0 - aug) * noises[0]
scale = self.audio_scale()
if scale != 1.0:
# the sampler carries audio as (sigma_v / sigma_a) * x_audio and latent_image
# holds audio_scale * x_audio, so rescale for the model to see it clean
model_sampling = self.model_sampling
sigma_v = sigma.clamp(min=1e-6)
sigma_a = comfy.ldm.minimax.model.time_shift_sigma(sigma_v, model_sampling.shift, model_sampling.audio_shift)
factor = (sigma_v / sigma_a) / scale
cleans[1] = cleans[1] * factor.view(factor.shape[:1] + (1,) * (cleans[1].ndim - 1)).to(cleans[1].dtype)
injected = utils.pack_latents(cleans)[0]
if x is None or denoise_mask is None:
return injected
token_grid_mask = utils.pack_latents(self._token_grid_masks(denoise_mask, shapes))[0]
x_blend_weight = (token_grid_mask - denoise_mask) / (1.0 - denoise_mask).clamp(min=1e-6)
x_blend_weight = torch.where(denoise_mask < 1.0, x_blend_weight.clamp(0.0, 1.0), torch.zeros_like(x_blend_weight))
return injected + x_blend_weight.to(injected.dtype) * (x - injected)
class TripoSplat(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.triposplat.model.LatentSeqMMFlowModel)

View File

@@ -636,7 +636,7 @@ class KSamplerX0Inpaint:
if "denoise_mask_function" in model_options:
denoise_mask = model_options["denoise_mask_function"](sigma, denoise_mask, extra_options={"model": self.inner_model, "sigmas": self.sigmas})
latent_mask = 1. - denoise_mask
x = x * denoise_mask + self.inner_model.inner_model.scale_latent_inpaint(x=x, sigma=sigma, noise=self.noise, latent_image=self.latent_image) * latent_mask
x = x * denoise_mask + self.inner_model.inner_model.scale_latent_inpaint(x=x, sigma=sigma, noise=self.noise, latent_image=self.latent_image, denoise_mask=denoise_mask) * latent_mask
out = self.inner_model(x, sigma, model_options=model_options, seed=seed)
if denoise_mask is not None:
out = out * denoise_mask + self.latent_image * latent_mask