| 17 | |
| 18 | class VideoResBlock(ResnetBlock): |
| 19 | def __init__( |
| 20 | self, |
| 21 | out_channels, |
| 22 | *args, |
| 23 | dropout=0.0, |
| 24 | video_kernel_size=3, |
| 25 | alpha=0.0, |
| 26 | merge_strategy="learned", |
| 27 | **kwargs, |
| 28 | ): |
| 29 | super().__init__(out_channels=out_channels, dropout=dropout, *args, **kwargs) |
| 30 | if video_kernel_size is None: |
| 31 | video_kernel_size = [3, 1, 1] |
| 32 | self.time_stack = ResBlock( |
| 33 | channels=out_channels, |
| 34 | emb_channels=0, |
| 35 | dropout=dropout, |
| 36 | dims=3, |
| 37 | use_scale_shift_norm=False, |
| 38 | use_conv=False, |
| 39 | up=False, |
| 40 | down=False, |
| 41 | kernel_size=video_kernel_size, |
| 42 | use_checkpoint=False, |
| 43 | skip_t_emb=True, |
| 44 | ) |
| 45 | |
| 46 | self.merge_strategy = merge_strategy |
| 47 | if self.merge_strategy == "fixed": |
| 48 | self.register_buffer("mix_factor", torch.Tensor([alpha])) |
| 49 | elif self.merge_strategy == "learned": |
| 50 | self.register_parameter("mix_factor", torch.nn.Parameter(torch.Tensor([alpha]))) |
| 51 | else: |
| 52 | raise ValueError(f"unknown merge strategy {self.merge_strategy}") |
| 53 | |
| 54 | def get_alpha(self, bs): |
| 55 | if self.merge_strategy == "fixed": |