MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / __init__

Method __init__

sat/sgm/modules/diffusionmodules/loss.py:75–118  ·  view source on GitHub ↗
(self, block_scale=None, block_size=None, min_snr_value=None, fixed_frames=0, 
                 cond_inds=None,
                 cond_inds_prob=None,
                 apply_cond_aug=None,
                 max_aug_t=700,  # * For V2
                 contained_aug_t=False,
                 exclude_cond_from_loss=True,
                 loss_weight_short_dcl=False,  # * Not supported
                 loss_weight_multi_dcl=False,  # * already include short_dcl
                 k_multi_dcl=1,
                 discount_multi_dcl=1,
                 normalize_dcl=False,
                 detach_dcl_prev=False,
                 detach_dcl_post=False,
                **kwargs)

Source from the content-addressed store, hash-verified

73
74class 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)

Callers 1

__init__Method · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected