Added pure lpips dfe
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user