Dynamatrix/DiffBIR-OpenXLab
0
1from typing import overload2import torch3from torch.nn import functional as F4 5 6class Guidance:7 8 def __init__(self, scale, type, t_start, t_stop, space, repeat, loss_type):9 self.scale = scale10 self.type = type11 self.t_start = t_start12 self.t_stop = t_stop13 self.target = None14 self.space = space15 self.repeat = repeat16 self.loss_type = loss_type17 18 def load_target(self, target):19 self.target = target20 21 def __call__(self, target_x0, pred_x0, t):22 if self.t_stop < t and t < self.t_start:23 # print("sampling with classifier guidance")24 # avoid propagating gradient out of this scope25 pred_x0 = pred_x0.detach().clone()26 target_x0 = target_x0.detach().clone()27 return self.scale * self._forward(target_x0, pred_x0)28 else:29 return None30 31 @overload32 def _forward(self, target_x0, pred_x0): ...33 34 35class MSEGuidance(Guidance):36 37 def __init__(self, scale, type, t_start, t_stop, space, repeat, loss_type) -> None:38 super().__init__(39 scale, type, t_start, t_stop, space, repeat, loss_type40 )41 42 @torch.enable_grad()43 def _forward(self, target_x0: torch.Tensor, pred_x0: torch.Tensor):44 # inputs: [-1, 1], nchw, rgb45 pred_x0.requires_grad_(True)46 47 if self.loss_type == "mse":48 loss = (pred_x0 - target_x0).pow(2).mean((1, 2, 3)).sum()49 elif self.loss_type == "downsample_mse":50 # FIXME: scale_factor should be 1/4, not 451 lr_pred_x0 = F.interpolate(pred_x0, scale_factor=4, mode="bicubic")52 lr_target_x0 = F.interpolate(target_x0, scale_factor=4, mode="bicubic")53 loss = (lr_pred_x0 - lr_target_x0).pow(2).mean((1, 2, 3)).sum()54 else:55 raise ValueError(self.loss_type)56 57 print(f"loss = {loss.item()}")58 return -torch.autograd.grad(loss, pred_x0)[0]59 