Support per-token video and audio latent noise masks on MiniMax-H3
Video masks snap to the 2x2 latent patch grid, audio masks to whole latent frames, both rounded to binary.
This commit is contained in:
parent
88fec4b605
commit
3c4749976b
|
|
@ -74,6 +74,16 @@ 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_targets(mask, latent_t, lat_h, lat_w):
|
||||
# [T, H, W] denoise mask (1 = generate) -> per-2x2-patch-row bool, None when every row 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))
|
||||
target = m.reshape(-1) >= 0.5
|
||||
if bool(target.all()):
|
||||
return None
|
||||
return target
|
||||
|
||||
|
||||
def _frame_grid(h, w):
|
||||
# area-normalized (h, w) coordinates of one latent frame's 2x2-patch rows
|
||||
area = math.sqrt(h * w)
|
||||
|
|
@ -199,17 +209,25 @@ class AdalnProj(nn.Module):
|
|||
return x.chunk(self.expand, dim=-1)
|
||||
|
||||
|
||||
def _mod_row(vecs, row, dtype):
|
||||
# row is a mod-row index, or (target_row, pin_row, weight[n,1]) blending two rows per token
|
||||
if isinstance(row, tuple):
|
||||
rt, rp, w = row
|
||||
return torch.lerp(vecs[rp], vecs[rt], w.to(vecs.dtype)).to(dtype)
|
||||
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
|
||||
|
||||
|
||||
|
|
@ -275,13 +293,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:
|
||||
|
|
@ -485,7 +505,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))
|
||||
|
|
@ -502,7 +522,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
|
||||
|
|
@ -510,7 +531,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)
|
||||
|
|
@ -541,13 +562,33 @@ 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 == "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),
|
||||
"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()))
|
||||
|
||||
# rows that are preserved by the noise mask run at the cond timestep
|
||||
t_pin_v = max(t_v, VISUAL_COND_TIMESTEP)
|
||||
t_pin_a = max(t_a, AUDIO_COND_TIMESTEP)
|
||||
video_w = None
|
||||
audio_w = None
|
||||
if denoise_mask is not None:
|
||||
targets = mask_row_targets(denoise_mask[0, 0].to(torch.float32), latent_t, lat_h, lat_w)
|
||||
if targets is not None:
|
||||
if bool(targets.any()):
|
||||
video_w = targets.to(torch.float32).unsqueeze(1) # [n, 1], 1 = generate
|
||||
else:
|
||||
seg_t["video"] = t_pin_v
|
||||
if audio_denoise_mask is not None:
|
||||
targets = audio_denoise_mask[0, 0].to(torch.float32).reshape(-1) >= 0.5
|
||||
if not bool(targets.all()):
|
||||
if bool(targets.any()):
|
||||
audio_w = targets.to(torch.float32).unsqueeze(1)
|
||||
else:
|
||||
seg_t["audio"] = t_pin_a
|
||||
|
||||
unique_t = sorted({t_v, t_a} | {seg_t[k] for _, _, k in layout.segments}
|
||||
| ({t_pin_v} if video_w is not None else set())
|
||||
| ({t_pin_a} if audio_w 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, "ref_audio": 2}
|
||||
|
||||
|
|
@ -563,6 +604,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_w is not None:
|
||||
mod_segments.append((a, b, (row_base + seg_tag[kind], t_row[t_pin_v] * 3 + seg_tag[kind], video_w)))
|
||||
elif kind == "audio" and audio_w is not None:
|
||||
mod_segments.append((a, b, (row_base + seg_tag[kind], t_row[t_pin_a] * 3 + seg_tag[kind], audio_w)))
|
||||
else:
|
||||
mod_segments.append((a, b, row_base + seg_tag[kind]))
|
||||
|
||||
|
|
@ -639,8 +684,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_w is not None:
|
||||
video_seg = (va, vb, (t_row[seg_t["video"]], t_row[t_pin_v], video_w))
|
||||
else:
|
||||
video_seg = (va, vb, t_row[seg_t["video"]])
|
||||
if audio_w is not None:
|
||||
audio_seg = (aa, ab, (t_row[seg_t["audio"]], t_row[t_pin_a], audio_w))
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -252,6 +252,9 @@ class BaseModel(torch.nn.Module):
|
|||
def process_timestep(self, timestep, **kwargs):
|
||||
return timestep
|
||||
|
||||
def process_denoise_mask(self, denoise_masks):
|
||||
return denoise_masks
|
||||
|
||||
def get_dtype(self):
|
||||
return self.diffusion_model.dtype
|
||||
|
||||
|
|
@ -2134,6 +2137,15 @@ 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 and latent_shapes is not None and len(latent_shapes) > 1:
|
||||
masks = utils.unpack_latents(denoise_mask, latent_shapes)
|
||||
if torch.amin(masks[0]).item() < 0.5:
|
||||
out['denoise_mask'] = comfy.conds.CONDRegular(masks[0][:1, :1].clone())
|
||||
if torch.amin(masks[1]).item() < 0.5:
|
||||
out['audio_denoise_mask'] = comfy.conds.CONDRegular(masks[1][:1, :1].clone())
|
||||
|
||||
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]
|
||||
|
|
@ -2144,6 +2156,40 @@ class MiniMaxH3(BaseModel):
|
|||
out['minimax_payload'] = comfy.conds.CONDConstant(payload)
|
||||
return out
|
||||
|
||||
def process_denoise_mask(self, denoise_masks):
|
||||
# snap the video mask to the DiT patch grid and the audio mask to whole latent
|
||||
# frames so a row's timestep label matches its content
|
||||
vm = denoise_masks[0]
|
||||
h, w = vm.shape[-2:]
|
||||
ph, pw = self.diffusion_model.patch_size[1:]
|
||||
vm = torch.nn.functional.pad(vm, (0, -w % pw, 0, -h % ph))
|
||||
vm = (vm.reshape(vm.shape[:-2] + (vm.shape[-2] // ph, ph, vm.shape[-1] // pw, pw)).amax(dim=(-3, -1)) >= 0.5).to(vm.dtype)
|
||||
denoise_masks[0] = vm.repeat_interleave(ph, dim=-2).repeat_interleave(pw, dim=-1)[..., :h, :w]
|
||||
if len(denoise_masks) > 1:
|
||||
am = denoise_masks[1]
|
||||
denoise_masks[1] = (am.amax(dim=1, keepdim=True) >= 0.5).to(am.dtype).expand_as(am).contiguous()
|
||||
return denoise_masks
|
||||
|
||||
def scale_latent_inpaint(self, sigma, noise, latent_image, **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
|
||||
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
|
||||
ms = self.model_sampling
|
||||
sigma_v = sigma.clamp(min=1e-6)
|
||||
sigma_a = comfy.ldm.minimax.model.time_shift_sigma(sigma_v, ms.shift, ms.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)
|
||||
return utils.pack_latents(cleans)[0]
|
||||
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -1307,6 +1307,8 @@ class CFGGuider:
|
|||
for i in range(len(denoise_masks)):
|
||||
denoise_masks[i] = comfy.sampler_helpers.prepare_mask(denoise_masks[i], latent_shapes[i], self.model_patcher.load_device)
|
||||
|
||||
denoise_masks = self.model_patcher.model.process_denoise_mask(denoise_masks)
|
||||
|
||||
if len(denoise_masks) > 1:
|
||||
denoise_mask, _ = comfy.utils.pack_latents(denoise_masks)
|
||||
else:
|
||||
|
|
|
|||
Loading…
Reference in New Issue