Fix sampler issues for audio with minimax, support more samplers. (#15243)

This commit is contained in:
Jukka Seppänen
2026-08-06 23:36:34 +03:00
committed by GitHub
parent 2eb609766a
commit bdcb886a47
7 changed files with 129 additions and 30 deletions

View File

@@ -9,8 +9,9 @@ The packed sequence is:
Timestep domain: the model receives the *video* sigma from the sampler and
derives per-token timesteps t = 1 - sigma internally; the audio stream runs on
its own shifted schedule (sigma_shift video 12.0 / audio 3.0), mapped from the
video sigma in closed form. The audio velocity is returned scaled by the
schedule map's derivative d(sigma_a)/d(sigma_v).
video sigma in closed form. The sampler carries the audio latent scaled onto the
video schedule (ModelSamplingAV); forward() undoes that scale and converts the
velocity back, so _forward only ever sees the stream's own latent.
"""
import math
@@ -38,17 +39,6 @@ def time_shift_sigma(sigma, from_shift, to_shift):
return to_shift * base / (1.0 + (to_shift - 1.0) * base)
def time_shift_slope(sigma, from_shift, to_shift):
"""d(sigma_to)/d(sigma_from) at the same base-grid point.
Scaling a stream's returned velocity by this slope makes the flat ODE that
any sampler integrates on the from-schedule equal to that stream's true ODE
on its own schedule.
"""
base = sigma / (from_shift + sigma * (1.0 - from_shift))
return (to_shift * (1.0 + (from_shift - 1.0) * base) ** 2) / (from_shift * (1.0 + (to_shift - 1.0) * base) ** 2)
def patchify_video(latent, patch_size=(1, 2, 2)):
# [B, C, T, H, W] -> [B*t*h*w, C*pt*ph*pw]
b, c, t_full, h_full, w_full = latent.shape
@@ -496,12 +486,30 @@ class MiniMaxH3Model(nn.Module):
return torch.cat(rows, dim=0) if rows else None
def forward(self, x, timestep, context, transformer_options={}, minimax_payload=None, **kwargs):
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
# 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))
audio_x = x[1]
if scale != 1.0:
shift_v = float(transformer_options.get("minimax_h3_sigma_shift_video", self.sigma_shift_video))
shift_a = float(transformer_options.get("minimax_h3_sigma_shift_audio", self.sigma_shift_audio))
sigma_v = (timestep.flatten()[0] / 1000.0).float().clamp(min=1e-6)
sigma_a = time_shift_sigma(sigma_v, shift_v, shift_a)
audio_x = audio_x * (sigma_a / sigma_v).to(audio_x.dtype)
x = [x[0], audio_x]
out = comfy.patcher_extension.WrapperExecutor.new_class_executor(
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)
if scale != 1.0:
# d/d(sigma_v) of the carried variable
out[1] = ((1.0 - scale) * audio_x
+ (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):
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]
@@ -639,8 +647,4 @@ class MiniMaxH3Model(nn.Module):
video_out = video_out[:, :, :orig_t, :orig_h, :orig_w]
audio_out = unpack_audio(a)
# The sampler integrates the flat ODE dX/dsigma_v = (X - denoised)/sigma_v.
# Scaling the audio velocity by d(sigma_a)/d(sigma_v) makes that ODE equal
# to the audio stream's true ODE on its own shifted schedule.
slope_a = time_shift_slope(sigma_v, shift_v, shift_a).to(audio_out.dtype)
return [-video_out.to(video_x.dtype), (-slope_a) * audio_out.to(audio_x.dtype)]
return [-video_out.to(video_x.dtype), -audio_out.to(audio_x.dtype)]

View File

@@ -22,6 +22,7 @@ import torch
import logging
import comfy.ldm.lightricks.av_model
import comfy.ldm.minimax.model
import comfy.nested_tensor
import comfy.ldm.lightricks.symmetric_patchifier
import comfy.context_windows
from comfy.ldm.modules.diffusionmodules.openaimodel import UNetModel, Timestep
@@ -100,6 +101,7 @@ class ModelType(Enum):
FLOW_COSMOS = 10
IMG_TO_IMG_FLOW = 11
V_PREDICTION_DDPM = 12
FLOW_AV = 13
def model_sampling(model_config, model_type):
@@ -136,6 +138,9 @@ def model_sampling(model_config, model_type):
c = comfy.model_sampling.IMG_TO_IMG_FLOW
elif model_type == ModelType.V_PREDICTION_DDPM:
c = comfy.model_sampling.V_PREDICTION_DDPM
elif model_type == ModelType.FLOW_AV:
c = comfy.model_sampling.CONST
s = comfy.model_sampling.ModelSamplingAV
class ModelSampling(s, c):
pass
@@ -180,6 +185,7 @@ class BaseModel(torch.nn.Module):
self.model_type = model_type
self.model_sampling = model_sampling(model_config, model_type)
self.latent_shapes = None # set by the sampler for models that pack several streams into one latent
self.adm_channels = unet_config.get("adm_in_channels", None)
if self.adm_channels is None:
@@ -2065,9 +2071,33 @@ class Hunyuan3Dv2_1(BaseModel):
return out
class MiniMaxH3(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLOW, device=None):
def __init__(self, model_config, model_type=ModelType.FLOW_AV, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.minimax.model.MiniMaxH3Model)
def audio_scale(self):
"""Scale the sampler carries the audio stream at, 1.0 when not sampling the packed latent."""
if self.latent_shapes is None or len(self.latent_shapes) < 2:
return 1.0
return self.model_sampling.audio_scale
def _scale_audio_slice(self, latent, scale):
# the sampler carries the audio stream scaled onto the video schedule
if scale == 1.0:
return latent
if latent.is_nested: # the x0 output hands back the unpacked view
streams = latent.unbind()
return comfy.nested_tensor.NestedTensor([streams[0], streams[1] * scale] + list(streams[2:]))
n = math.prod(self.latent_shapes[0][1:])
latent = latent.clone()
latent[..., n:] *= scale
return latent
def process_latent_in(self, latent):
return self._scale_audio_slice(super().process_latent_in(latent), self.audio_scale())
def process_latent_out(self, latent):
return super().process_latent_out(self._scale_audio_slice(latent, 1.0 / self.audio_scale()))
def extra_conds(self, **kwargs):
out = super().extra_conds(**kwargs)
cross_attn = kwargs.get("cross_attn", None)
@@ -2102,6 +2132,8 @@ class MiniMaxH3(BaseModel):
if kwargs.get("minimax_audio_cond_noise_aug", None) is not None:
payload["audio_cond_noise_aug"] = kwargs["minimax_audio_cond_noise_aug"]
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()
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]

View File

@@ -325,6 +325,27 @@ class ModelSamplingDiscreteFlow(torch.nn.Module):
return 0.0
return time_snr_shift(self.shift, 1.0 - percent)
class ModelSamplingAV(ModelSamplingDiscreteFlow):
"""Flow sampling for packed audio-video latents whose audio stream has its own flow shift.
Carrying the audio latent scaled onto the video schedule makes the pack an ordinary
single-schedule flow latent whose audio target is scaled by audio_scale.
"""
def __init__(self, model_config=None):
super().__init__(model_config)
sampling_settings = model_config.sampling_settings if model_config is not None else {}
self.audio_shift = sampling_settings.get("audio_shift", None)
def set_parameters(self, shift=1.0, audio_shift=None, timesteps=1000, multiplier=1000):
self.audio_shift = audio_shift
super().set_parameters(shift=shift, timesteps=timesteps, multiplier=multiplier)
@property
def audio_scale(self):
if self.audio_shift is None:
return 1.0
return self.shift / self.audio_shift
class StableCascadeSampling(ModelSamplingDiscrete):
def __init__(self, model_config=None):
super().__init__()

View File

@@ -1218,6 +1218,8 @@ class CFGGuider:
return sampling_function(self.inner_model, x, timestep, self.conds.get("negative", None), self.conds.get("positive", None), self.cfg, model_options=model_options, seed=seed)
def inner_sample(self, noise, latent_image, device, sampler, sigmas, denoise_mask, callback, disable_pbar, seed, latent_shapes=None):
self.inner_model.latent_shapes = latent_shapes
if latent_image is not None and torch.count_nonzero(latent_image) > 0: #Don't shift the empty latent image.
latent_image = self.inner_model.process_latent_in(latent_image)

View File

@@ -963,6 +963,7 @@ class MiniMaxH3(supported_models_base.BASE):
sampling_settings = {
"shift": 12.0,
"audio_shift": 3.0,
}
unet_extra_config = {}