| 73 | |
| 74 | class VideoDiffusionLoss(StandardDiffusionLoss): |
| 75 | def __init__(self, block_scale=None, block_size=None, min_snr_value=None, fixed_frames=0, |
| 76 | cond_inds=None, |
| 77 | cond_inds_prob=None, |
| 78 | apply_cond_aug=None, |
| 79 | max_aug_t=700, # * For V2 |
| 80 | contained_aug_t=False, |
| 81 | exclude_cond_from_loss=True, |
| 82 | loss_weight_short_dcl=False, # * Not supported |
| 83 | loss_weight_multi_dcl=False, # * already include short_dcl |
| 84 | k_multi_dcl=1, |
| 85 | discount_multi_dcl=1, |
| 86 | normalize_dcl=False, |
| 87 | detach_dcl_prev=False, |
| 88 | detach_dcl_post=False, |
| 89 | **kwargs): |
| 90 | self.fixed_frames = fixed_frames |
| 91 | self.block_scale = block_scale |
| 92 | self.block_size = block_size |
| 93 | self.min_snr_value = min_snr_value |
| 94 | super().__init__(**kwargs) |
| 95 | |
| 96 | self.cond_inds = cond_inds # Nested list. e.g., [0, 1, 2] should be [ [0, 1, 2] ] |
| 97 | |
| 98 | self.cond_inds_prob = cond_inds_prob # probability of different conditioning schemes, None for uniform distribution |
| 99 | assert cond_inds_prob is None or len(cond_inds_prob) == len(cond_inds), f"Invalid cond_inds_prob: {cond_inds_prob}" |
| 100 | |
| 101 | self.apply_cond_aug = apply_cond_aug # apply aug on conditioning frames |
| 102 | assert self.apply_cond_aug in [None, 'V1', 'V2'], f"Invalid apply_cond_aug: {self.apply_cond_aug}" |
| 103 | |
| 104 | self.max_aug_t = max_aug_t # * For V2 |
| 105 | self.contained_aug_t = contained_aug_t # contain the aug_t to be no greater than t (on prediction frames) |
| 106 | self.exclude_cond_from_loss = exclude_cond_from_loss |
| 107 | |
| 108 | if loss_weight_short_dcl is not False: |
| 109 | raise NotImplementedError("loss_weight_short_dcl is not supported.") |
| 110 | |
| 111 | self.loss_weight_multi_dcl = loss_weight_multi_dcl |
| 112 | self.k_multi_dcl = k_multi_dcl |
| 113 | self.discount_multi_dcl = discount_multi_dcl |
| 114 | self.normalize_dcl = normalize_dcl |
| 115 | |
| 116 | self.detach_dcl_prev = detach_dcl_prev |
| 117 | self.detach_dcl_post = detach_dcl_post |
| 118 | assert not (self.detach_dcl_prev and self.detach_dcl_post), "Only one detach option can be enabled." |
| 119 | |
| 120 | def __call__(self, network, denoiser, conditioner, input, batch): |
| 121 | cond = conditioner(batch) |