Added pure lpips dfe

This commit is contained in:
Jaret Burkett
2026-05-31 11:52:03 -06:00
parent 212cfe998a
commit e5439509b5

View File

@@ -4,6 +4,7 @@ import os
from torch import nn
from safetensors.torch import load_file
import torch.nn.functional as F
import torch.utils.checkpoint as ckpt
from diffusers import AutoencoderTiny
from transformers import AutoImageProcessor, AutoModel, SiglipImageProcessor, SiglipVisionModel
import lpips
@@ -1179,6 +1180,163 @@ class DiffusionFeatureExtractor9(nn.Module):
return loss_perceptual
class DiffusionFeatureExtractor10(nn.Module):
def __init__(
self,
device=torch.device("cuda"),
dtype=torch.bfloat16,
vae=None,
sd=None,
partial_step: bool = False
):
super().__init__()
self.version = 10
self.sd_ref = weakref.ref(sd) if sd is not None else None
self.lpips_model = lpips.LPIPS(net='vgg')
self.lpips_model = self.lpips_model.to(device, dtype=torch.float32)
self.losses = {}
self.log_every = 100
self.step = 0
self.do_partial_step = partial_step
def _vgg_slices(self, x):
# run the lpips vgg backbone slice-by-slice so we can gradient
# checkpoint each slice. checkpointing activates whenever grads are
# enabled, so it does not require the module to be in train mode.
net = self.lpips_model.net
slices = [net.slice1, net.slice2, net.slice3, net.slice4, net.slice5]
outs = []
h = x
for s in slices:
if torch.is_grad_enabled():
h = ckpt.checkpoint(s, h, use_reentrant=False)
else:
h = s(h)
outs.append(h)
return outs
def get_lpips_features(self, tensors_0_1):
device = self.lpips_model.scaling_layer.shift.device
tensors_n1p1 = (tensors_0_1 * 2) - 1
def get_lpips_features(img): # -1 to 1
in0_input = self.lpips_model.scaling_layer(img)
outs0 = self._vgg_slices(in0_input)
feats_list = []
for kk in range(self.lpips_model.L):
feats_list.append(lpips.normalize_tensor(outs0[kk]))
return feats_list
lpips_feat_list = [x for x in get_lpips_features(
tensors_n1p1.to(device, dtype=torch.float32))]
return lpips_feat_list
def forward(
self,
noise,
noise_pred,
noisy_latents,
timesteps,
batch: DataLoaderBatchDTO,
scheduler: CustomFlowMatchEulerDiscreteScheduler,
model=None
):
dtype = torch.bfloat16
device = self.sd_ref().vae.device
tensors = batch.tensor.to(device, dtype=dtype)
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, :, :]
is_video = True
if len(tensors.shape) == 5:
# batch is different
# (B, T, C, H, W)
# only take first time
tensors = tensors[:, 0, :, :, :]
with torch.no_grad():
tv = timesteps.to(noise_pred.device).to(noise_pred.dtype) / 1000.0
# expand shape to match noise_pred
while len(tv.shape) < len(noise_pred.shape):
tv = tv.unsqueeze(-1)
with torch.no_grad():
target_0_1 = (tensors + 1) / 2 # 0 to 1
if not self.do_partial_step:
# step latent
x0 = noisy_latents - tv * noise_pred
stepped_latents = x0
# min 0.001
tv = torch.clamp(tv, min=0.001)
else:
# step is random 0.1 to 0.25
step = torch.rand_like(tv) * 0.15 + 0.1
next_step = tv - step
next_step = torch.clamp(next_step, min=0.0)
stepped_latents = noisy_latents + (next_step - tv) * noise_pred
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)
# add noise
target_latents = (1.0 - next_step) * target_latents + next_step * noise
target_n1p1 = self.sd_ref().decode_latents(target_latents)
target_0_1 = (target_n1p1 + 1) / 2 # 0 to 1
latents = stepped_latents.to(self.sd_ref().vae.device, dtype=self.sd_ref().vae.dtype)
tensors_n1p1 = self.sd_ref().decode_latents(latents)
pred_images = (tensors_n1p1 + 1) / 2 # 0 to 1
with torch.no_grad():
target_feats = self.get_lpips_features(target_0_1.float())
pred_feats = self.get_lpips_features(pred_images.float())
velocity_equiv_weight = (1.0 / torch.clamp(tv, min=0.1) ** 2)
loss_perceptual = 0
for idx, pred_feat in enumerate(pred_feats):
perceptual_loss = torch.nn.functional.mse_loss(
pred_feat.float(), target_feats[idx].float(), reduction="none"
)
# mean over channels/spatial per sample, keep batch dim to weight by timestep
perceptual_loss = perceptual_loss.mean(dim=[1, 2, 3], keepdim=True)
loss_perceptual = loss_perceptual + (perceptual_loss * velocity_equiv_weight).mean()
if self.do_partial_step:
loss_perceptual = loss_perceptual * 10.0
if 'loss' not in self.losses:
self.losses['loss'] = loss_perceptual.item()
else:
self.losses['loss'] += loss_perceptual.item()
with torch.no_grad():
if self.step % self.log_every == 0 and self.step > 0:
print(f"DFE losses:")
for key in self.losses:
self.losses[key] /= self.log_every
# print in 2.000e-01 format
print(f" - {key}: {self.losses[key]:.3e}")
self.losses[key] = 0.0
# total_loss += mse_loss
self.step += 1
return loss_perceptual
def load_dfe(model_path, vae=None, sd: 'BaseModel' = None) -> DiffusionFeatureExtractor:
if model_path == "v3":
dfe = DiffusionFeatureExtractor3(vae=vae)
@@ -1208,6 +1366,10 @@ def load_dfe(model_path, vae=None, sd: 'BaseModel' = None) -> DiffusionFeatureEx
dfe = DiffusionFeatureExtractor9(vae=vae, sd=sd)
dfe.eval()
return dfe
if model_path == "v10":
dfe = DiffusionFeatureExtractor10(vae=vae, sd=sd)
dfe.eval()
return dfe
if not os.path.exists(model_path):
raise FileNotFoundError(f"Model file not found: {model_path}")
# if it ende with safetensors