Apply sigma to contrastive guidance to balance loss better. Prevent noise grads on images/non audio datasets. Prep for training adapters on MiniMax H3

This commit is contained in:
Jaret Burkett
2026-08-05 13:27:12 -06:00
parent 9065951da3
commit 1e1418b22c
8 changed files with 163 additions and 10 deletions

View File

@@ -736,7 +736,15 @@ class SDTrainer(BaseSDTrainProcess):
if isinstance(guidance_scale, list):
guidance_scale = torch.tensor(guidance_scale).to(target.device, dtype=target.dtype)
guidance_scale = guidance_scale.view(-1, 1, 1, 1) if not is_video else guidance_scale.view(-1, 1, 1, 1, 1)
if self.train_config.guidance_loss_schedule == 'sigma':
# the (target - uncond) sample direction carries s * fresh_noise
# that nothing can predict at low sigma, so decay the
# extrapolation toward a plain flow target as sigma falls
sigma = (timesteps.to(target.device) / 1000.0).to(target.dtype)
sigma = sigma.view(-1, 1, 1, 1) if not is_video else sigma.view(-1, 1, 1, 1, 1)
guidance_scale = 1.0 + (guidance_scale - 1.0) * sigma
unconditional_target = unconditional_target * alpha
target = unconditional_target + guidance_scale * (target - unconditional_target)
@@ -762,6 +770,16 @@ class SDTrainer(BaseSDTrainProcess):
audio_target.device, dtype=audio_target.dtype
).view(-1, *audio_dims)
if self.train_config.guidance_loss_schedule == 'sigma':
# audio streams can run on their own remapped sigma
audio_sigma = getattr(batch, 'audio_sigma', None)
if audio_sigma is None:
audio_sigma = timesteps / 1000.0
audio_sigma = audio_sigma.to(
audio_target.device, dtype=audio_target.dtype
).view(-1, *audio_dims)
audio_guidance_scale = 1.0 + (audio_guidance_scale - 1.0) * audio_sigma
batch.audio_target = (
audio_uncond + audio_guidance_scale * (audio_target - audio_uncond)
).to(batch.audio_target.dtype).detach()