Fixed for DFEs with pixelspace and video models

This commit is contained in:
Jaret Burkett
2026-07-27 11:12:02 -06:00
parent 0e17841767
commit fb204b7677
4 changed files with 92 additions and 53 deletions

View File

@@ -796,6 +796,9 @@ class SDTrainer(BaseSDTrainProcess):
tv = torch.clamp(tv, min=0.001)
# step latent, use here or with do_fft_loss
if self.sd.x0_pred:
t0 = noise_pred
else:
t0 = noisy_latents - tv * noise_pred
if self.train_config.t0_loss_target:

View File

@@ -197,6 +197,9 @@ class BaseModel:
# if a mask is passed, do the loss with the mask. May be set false for models that use a mask for other reasons.
self.do_masked_loss = True
# if the model outputs an x0 prediction (clean latent)
self.x0_pred = False
# properties for old arch for backwards compatibility
@property
def unet(self):

View File

@@ -16,6 +16,13 @@ from toolkit.models.sapiens2 import Sapiens2
import huggingface_hub
def _fold_frames_to_batch(x: torch.Tensor) -> torch.Tensor:
"""(B, C, T, H, W) -> (B*T, C, H, W), each sample's frames contiguous -- so the 2D
feature losses run on EVERY frame of a video instead of only the first."""
b, c, t, h, w = x.shape
return x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
class ResBlock(nn.Module):
def __init__(self, in_channels, out_channels):
super().__init__()
@@ -521,18 +528,19 @@ class DiffusionFeatureExtractor4(nn.Module):
is_video = False
# stack time for video models on the batch dimension
if len(noise_pred.shape) == 5:
# B, C, T, H, W = images.shape
# only take first time
noise = noise[:, :, 0, :, :]
noise_pred = noise_pred[:, :, 0, :, :]
noisy_latents = noisy_latents[:, :, 0, :, :]
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
# frames, and repeat the per-sample timestep for each of its frames
num_frames = noise_pred.shape[2]
noise = _fold_frames_to_batch(noise)
noise_pred = _fold_frames_to_batch(noise_pred)
noisy_latents = _fold_frames_to_batch(noisy_latents)
timesteps = timesteps.repeat_interleave(num_frames)
is_video = True
if len(tensors.shape) == 5:
# batch is different
# (B, T, C, H, W)
# only take first time
tensors = tensors[:, 0, :, :, :]
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
# the frame order of the folded predictions above
tensors = tensors.reshape(-1, *tensors.shape[2:])
if model is not None and hasattr(model, 'get_stepped_pred'):
stepped_latents = model.get_stepped_pred(noise_pred, noise)
@@ -745,18 +753,19 @@ class DiffusionFeatureExtractor6(nn.Module):
is_video = False
# stack time for video models on the batch dimension
if len(noise_pred.shape) == 5:
# B, C, T, H, W = images.shape
# only take first time
noise = noise[:, :, 0, :, :]
noise_pred = noise_pred[:, :, 0, :, :]
noisy_latents = noisy_latents[:, :, 0, :, :]
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
# frames, and repeat the per-sample timestep for each of its frames
num_frames = noise_pred.shape[2]
noise = _fold_frames_to_batch(noise)
noise_pred = _fold_frames_to_batch(noise_pred)
noisy_latents = _fold_frames_to_batch(noisy_latents)
timesteps = timesteps.repeat_interleave(num_frames)
is_video = True
if len(tensors.shape) == 5:
# batch is different
# (B, T, C, H, W)
# only take first time
tensors = tensors[:, 0, :, :, :]
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
# the frame order of the folded predictions above
tensors = tensors.reshape(-1, *tensors.shape[2:])
with torch.no_grad():
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
@@ -914,18 +923,19 @@ class DiffusionFeatureExtractor7(nn.Module):
is_video = False
# stack time for video models on the batch dimension
if len(noise_pred.shape) == 5:
# B, C, T, H, W = images.shape
# only take first time
noise = noise[:, :, 0, :, :]
noise_pred = noise_pred[:, :, 0, :, :]
noisy_latents = noisy_latents[:, :, 0, :, :]
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
# frames, and repeat the per-sample timestep for each of its frames
num_frames = noise_pred.shape[2]
noise = _fold_frames_to_batch(noise)
noise_pred = _fold_frames_to_batch(noise_pred)
noisy_latents = _fold_frames_to_batch(noisy_latents)
timesteps = timesteps.repeat_interleave(num_frames)
is_video = True
if len(tensors.shape) == 5:
# batch is different
# (B, T, C, H, W)
# only take first time
tensors = tensors[:, 0, :, :, :]
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
# the frame order of the folded predictions above
tensors = tensors.reshape(-1, *tensors.shape[2:])
with torch.no_grad():
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
@@ -937,6 +947,9 @@ class DiffusionFeatureExtractor7(nn.Module):
target_0_1 = (tensors + 1) / 2 # 0 to 1
if not self.do_partial_step:
if getattr(self.sd_ref(), "x0_pred", False):
x0 = noise_pred
else:
# step latent
x0 = noisy_latents - tv * noise_pred
stepped_latents = x0
@@ -952,6 +965,9 @@ class DiffusionFeatureExtractor7(nn.Module):
with torch.no_grad():
# make a noisy target at next timestep
target_latents = batch.latents.to(self.sd_ref().vae.device, dtype=self.sd_ref().vae.dtype)
if target_latents.dim() == 5:
# fold frames to match the folded noise/predictions
target_latents = _fold_frames_to_batch(target_latents)
# add noise
target_latents = (1.0 - next_step) * target_latents + next_step * noise
target_n1p1 = self.sd_ref().decode_latents(target_latents)
@@ -1110,18 +1126,19 @@ class DiffusionFeatureExtractor9(nn.Module):
is_video = False
# stack time for video models on the batch dimension
if len(noise_pred.shape) == 5:
# B, C, T, H, W = images.shape
# only take first time
noise = noise[:, :, 0, :, :]
noise_pred = noise_pred[:, :, 0, :, :]
noisy_latents = noisy_latents[:, :, 0, :, :]
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
# frames, and repeat the per-sample timestep for each of its frames
num_frames = noise_pred.shape[2]
noise = _fold_frames_to_batch(noise)
noise_pred = _fold_frames_to_batch(noise_pred)
noisy_latents = _fold_frames_to_batch(noisy_latents)
timesteps = timesteps.repeat_interleave(num_frames)
is_video = True
if len(tensors.shape) == 5:
# batch is different
# (B, T, C, H, W)
# only take first time
tensors = tensors[:, 0, :, :, :]
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
# the frame order of the folded predictions above
tensors = tensors.reshape(-1, *tensors.shape[2:])
with torch.no_grad():
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
@@ -1133,6 +1150,9 @@ class DiffusionFeatureExtractor9(nn.Module):
target_0_1 = (tensors + 1) / 2 # 0 to 1
if not self.do_partial_step:
if getattr(self.sd_ref(), "x0_pred", False):
x0 = noise_pred
else:
# step latent
x0 = noisy_latents - tv * noise_pred
stepped_latents = x0
@@ -1148,6 +1168,9 @@ class DiffusionFeatureExtractor9(nn.Module):
with torch.no_grad():
# make a noisy target at next timestep
target_latents = batch.latents.to(self.sd_ref().vae.device, dtype=self.sd_ref().vae.dtype)
if target_latents.dim() == 5:
# fold frames to match the folded noise/predictions
target_latents = _fold_frames_to_batch(target_latents)
# add noise
target_latents = (1.0 - next_step) * target_latents + next_step * noise
target_n1p1 = self.sd_ref().decode_latents(target_latents)
@@ -1266,18 +1289,19 @@ class DiffusionFeatureExtractor10(nn.Module):
is_video = False
# stack time for video models on the batch dimension
if len(noise_pred.shape) == 5:
# B, C, T, H, W = images.shape
# only take first time
noise = noise[:, :, 0, :, :]
noise_pred = noise_pred[:, :, 0, :, :]
noisy_latents = noisy_latents[:, :, 0, :, :]
# (B, C, T, H, W): fold every frame into the batch dim so the loss covers all
# frames, and repeat the per-sample timestep for each of its frames
num_frames = noise_pred.shape[2]
noise = _fold_frames_to_batch(noise)
noise_pred = _fold_frames_to_batch(noise_pred)
noisy_latents = _fold_frames_to_batch(noisy_latents)
timesteps = timesteps.repeat_interleave(num_frames)
is_video = True
if len(tensors.shape) == 5:
# batch is different
# (B, T, C, H, W)
# only take first time
tensors = tensors[:, 0, :, :, :]
# batch tensor is frames-first (B, T, C, H, W): fold to (B*T, C, H, W), matching
# the frame order of the folded predictions above
tensors = tensors.reshape(-1, *tensors.shape[2:])
with torch.no_grad():
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
@@ -1289,6 +1313,9 @@ class DiffusionFeatureExtractor10(nn.Module):
target_0_1 = (tensors + 1) / 2 # 0 to 1
if not self.do_partial_step:
if getattr(self.sd_ref(), "x0_pred", False):
x0 = noise_pred
else:
# step latent
x0 = noisy_latents - tv * noise_pred
stepped_latents = x0
@@ -1304,6 +1331,9 @@ class DiffusionFeatureExtractor10(nn.Module):
with torch.no_grad():
# make a noisy target at next timestep
target_latents = batch.latents.to(self.sd_ref().vae.device, dtype=self.sd_ref().vae.dtype)
if target_latents.dim() == 5:
# fold frames to match the folded noise/predictions
target_latents = _fold_frames_to_batch(target_latents)
# add noise
target_latents = (1.0 - next_step) * target_latents + next_step * noise
target_n1p1 = self.sd_ref().decode_latents(target_latents)

View File

@@ -235,6 +235,9 @@ class StableDiffusion:
# if a mask is passed, do the loss with the mask. May be set false for models that use a mask for other reasons.
self.do_masked_loss = True
# if the model outputs an x0 prediction (clean latent)
self.x0_pred = False
# properties for old arch for backwards compatibility
@property
def is_xl(self):