echo-infinity / model /dmd.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
3e936b2 verified
Raw
History Blame Contribute Delete
10.1 kB
import torch.nn.functional as F
from typing import Optional, Tuple
import torch
import time
from model.base import SelfForcingModel
import torch.distributed as dist
def _pixels_to_videos_0_255(pixels: torch.Tensor) -> torch.Tensor:
return ((pixels.float() + 1.0) * 127.5).clamp(0, 255).to(torch.uint8)
class DMD(SelfForcingModel):
def __init__(self, args, device):
super().__init__(args, device)
self.num_frame_per_block = getattr(args, 'num_frame_per_block', 1)
self.same_step_across_blocks = getattr(args, 'same_step_across_blocks', True)
self.min_num_training_frames = getattr(args, 'min_num_training_frames', 21)
self.num_training_frames = getattr(args, 'num_training_frames', 21)
if self.num_frame_per_block > 1:
self.generator.model.num_frame_per_block = self.num_frame_per_block
self.independent_first_frame = getattr(args, 'independent_first_frame', False)
if self.independent_first_frame:
self.generator.model.independent_first_frame = True
if args.gradient_checkpointing:
self.generator.enable_gradient_checkpointing()
self.fake_score.enable_gradient_checkpointing()
self.inference_pipeline: SelfForcingTrainingPipeline = None
self.num_train_timestep = args.num_train_timestep
self.min_step = int(0.02 * self.num_train_timestep)
self.max_step = int(0.98 * self.num_train_timestep)
if hasattr(args, 'real_guidance_scale'):
self.real_guidance_scale = args.real_guidance_scale
self.fake_guidance_scale = args.fake_guidance_scale
else:
self.real_guidance_scale = args.guidance_scale
self.fake_guidance_scale = 0.0
self.timestep_shift = getattr(args, 'timestep_shift', 1.0)
self.ts_schedule = getattr(args, 'ts_schedule', True)
self.ts_schedule_max = getattr(args, 'ts_schedule_max', False)
self.min_score_timestep = getattr(args, 'min_score_timestep', 0)
if getattr(self.scheduler, 'alphas_cumprod', None) is not None:
self.scheduler.alphas_cumprod = self.scheduler.alphas_cumprod.to(device)
else:
self.scheduler.alphas_cumprod = None
def _compute_kl_grad(self, noisy_image_or_video: torch.Tensor, estimated_clean_image_or_video: torch.Tensor, timestep: torch.Tensor, conditional_dict: dict, unconditional_dict: dict, normalization: bool=True) -> Tuple[torch.Tensor, dict]:
_, pred_fake_image_cond = self.fake_score(noisy_image_or_video=noisy_image_or_video, conditional_dict=conditional_dict, timestep=timestep)
if self.fake_guidance_scale != 0.0:
_, pred_fake_image_uncond = self.fake_score(noisy_image_or_video=noisy_image_or_video, conditional_dict=unconditional_dict, timestep=timestep)
pred_fake_image = pred_fake_image_cond + (pred_fake_image_cond - pred_fake_image_uncond) * self.fake_guidance_scale
else:
pred_fake_image = pred_fake_image_cond
_, pred_real_image_cond = self.real_score(noisy_image_or_video=noisy_image_or_video, conditional_dict=conditional_dict, timestep=timestep)
_, pred_real_image_uncond = self.real_score(noisy_image_or_video=noisy_image_or_video, conditional_dict=unconditional_dict, timestep=timestep)
pred_real_image = pred_real_image_cond + (pred_real_image_cond - pred_real_image_uncond) * self.real_guidance_scale
grad = pred_fake_image - pred_real_image
if normalization:
p_real = estimated_clean_image_or_video - pred_real_image
normalizer = torch.abs(p_real).mean(dim=[1, 2, 3, 4], keepdim=True)
grad = grad / normalizer
grad = torch.nan_to_num(grad)
return (grad, {'dmdtrain_gradient_norm': torch.mean(torch.abs(grad)).detach(), 'timestep': timestep.detach()})
def compute_distribution_matching_loss(self, image_or_video: torch.Tensor, conditional_dict: dict, unconditional_dict: dict, gradient_mask: Optional[torch.Tensor]=None, denoised_timestep_from: int=0, denoised_timestep_to: int=0, text_prompts: Optional[list]=None) -> Tuple[torch.Tensor, dict]:
original_latent = image_or_video
batch_size, num_frame = image_or_video.shape[:2]
with torch.no_grad():
min_timestep = denoised_timestep_to if self.ts_schedule and denoised_timestep_to is not None else self.min_score_timestep
max_timestep = denoised_timestep_from if self.ts_schedule_max and denoised_timestep_from is not None else self.num_train_timestep
timestep = self._get_timestep(min_timestep, max_timestep, batch_size, num_frame, self.num_frame_per_block, uniform_timestep=True)
if self.timestep_shift > 1:
timestep = self.timestep_shift * (timestep / 1000) / (1 + (self.timestep_shift - 1) * (timestep / 1000)) * 1000
timestep = timestep.clamp(self.min_step, self.max_step)
noise = torch.randn_like(image_or_video)
noisy_latent = self.scheduler.add_noise(image_or_video.flatten(0, 1), noise.flatten(0, 1), timestep.flatten(0, 1)).detach().unflatten(0, (batch_size, num_frame))
grad, dmd_log_dict = self._compute_kl_grad(noisy_image_or_video=noisy_latent, estimated_clean_image_or_video=original_latent, timestep=timestep, conditional_dict=conditional_dict, unconditional_dict=unconditional_dict)
if gradient_mask is not None:
base_dmd_loss = 0.5 * F.mse_loss(original_latent.double()[gradient_mask], (original_latent.double() - grad.double()).detach()[gradient_mask], reduction='mean')
else:
base_dmd_loss = 0.5 * F.mse_loss(original_latent.double(), (original_latent.double() - grad.double()).detach(), reduction='mean')
dmd_loss = base_dmd_loss
return (dmd_loss, dmd_log_dict)
def generator_loss(self, image_or_video_shape, conditional_dict: dict, unconditional_dict: dict, clean_latent: torch.Tensor, initial_latent: torch.Tensor=None, text_prompts: Optional[list]=None) -> Tuple[torch.Tensor, dict]:
slice_last_frames = getattr(self.args, 'slice_last_frames', 21)
_t_gen_start = time.time()
pred_image, gradient_mask, denoised_timestep_from, denoised_timestep_to = self._run_generator(image_or_video_shape=image_or_video_shape, conditional_dict=conditional_dict, initial_latent=initial_latent, slice_last_frames=slice_last_frames)
gen_time = time.time() - _t_gen_start
_t_loss_start = time.time()
dmd_loss, dmd_log_dict = self.compute_distribution_matching_loss(image_or_video=pred_image, conditional_dict=conditional_dict, unconditional_dict=unconditional_dict, gradient_mask=gradient_mask, denoised_timestep_from=denoised_timestep_from, denoised_timestep_to=denoised_timestep_to, text_prompts=text_prompts)
try:
loss_val = dmd_loss.item()
except Exception:
loss_val = float('nan')
loss_time = time.time() - _t_loss_start
dmd_log_dict.update({'gen_time': gen_time, 'loss_time': loss_time})
return (dmd_loss, dmd_log_dict)
def critic_loss(self, image_or_video_shape, conditional_dict: dict, unconditional_dict: dict, clean_latent: torch.Tensor, initial_latent: torch.Tensor=None) -> Tuple[torch.Tensor, dict]:
slice_last_frames = getattr(self.args, 'slice_last_frames', 21)
_t_gen_start = time.time()
with torch.no_grad():
generated_image, _, denoised_timestep_from, denoised_timestep_to = self._run_generator(image_or_video_shape=image_or_video_shape, conditional_dict=conditional_dict, initial_latent=initial_latent, slice_last_frames=slice_last_frames)
gen_time = time.time() - _t_gen_start
batch_size, num_frame = generated_image.shape[:2]
_t_loss_start = time.time()
min_timestep = denoised_timestep_to if self.ts_schedule and denoised_timestep_to is not None else self.min_score_timestep
max_timestep = denoised_timestep_from if self.ts_schedule_max and denoised_timestep_from is not None else self.num_train_timestep
critic_timestep = self._get_timestep(min_timestep, max_timestep, batch_size, num_frame, self.num_frame_per_block, uniform_timestep=True)
if self.timestep_shift > 1:
critic_timestep = self.timestep_shift * (critic_timestep / 1000) / (1 + (self.timestep_shift - 1) * (critic_timestep / 1000)) * 1000
critic_timestep = critic_timestep.clamp(self.min_step, self.max_step)
critic_noise = torch.randn_like(generated_image)
noisy_generated_image = self.scheduler.add_noise(generated_image.flatten(0, 1), critic_noise.flatten(0, 1), critic_timestep.flatten(0, 1)).unflatten(0, (batch_size, num_frame))
_, pred_fake_image = self.fake_score(noisy_image_or_video=noisy_generated_image, conditional_dict=conditional_dict, timestep=critic_timestep)
if self.args.denoising_loss_type == 'flow':
from utils.wan_wrapper import WanDiffusionWrapper
flow_pred = WanDiffusionWrapper._convert_x0_to_flow_pred(scheduler=self.scheduler, x0_pred=pred_fake_image.flatten(0, 1), xt=noisy_generated_image.flatten(0, 1), timestep=critic_timestep.flatten(0, 1))
pred_fake_noise = None
else:
flow_pred = None
pred_fake_noise = self.scheduler.convert_x0_to_noise(x0=pred_fake_image.flatten(0, 1), xt=noisy_generated_image.flatten(0, 1), timestep=critic_timestep.flatten(0, 1)).unflatten(0, (batch_size, num_frame))
denoising_loss = self.denoising_loss_func(x=generated_image.flatten(0, 1), x_pred=pred_fake_image.flatten(0, 1), noise=critic_noise.flatten(0, 1), noise_pred=pred_fake_noise, alphas_cumprod=self.scheduler.alphas_cumprod, timestep=critic_timestep.flatten(0, 1), flow_pred=flow_pred)
try:
loss_val = denoising_loss.item()
except Exception:
loss_val = float('nan')
loss_time = time.time() - _t_loss_start
critic_log_dict = {'critic_timestep': critic_timestep.detach(), 'gen_time': gen_time, 'loss_time': loss_time}
return (denoising_loss, critic_log_dict)