(self, in_channels: int, alpha: float = 0, merge_strategy: str = "learned")
| 107 | |
| 108 | class VideoBlock(AttnBlock): |
| 109 | def __init__(self, in_channels: int, alpha: float = 0, merge_strategy: str = "learned"): |
| 110 | super().__init__(in_channels) |
| 111 | # no context, single headed, as in base class |
| 112 | self.time_mix_block = VideoTransformerBlock( |
| 113 | dim=in_channels, |
| 114 | n_heads=1, |
| 115 | d_head=in_channels, |
| 116 | checkpoint=False, |
| 117 | ff_in=True, |
| 118 | attn_mode="softmax", |
| 119 | ) |
| 120 | |
| 121 | time_embed_dim = self.in_channels * 4 |
| 122 | self.video_time_embed = torch.nn.Sequential( |
| 123 | torch.nn.Linear(self.in_channels, time_embed_dim), |
| 124 | torch.nn.SiLU(), |
| 125 | torch.nn.Linear(time_embed_dim, self.in_channels), |
| 126 | ) |
| 127 | |
| 128 | self.merge_strategy = merge_strategy |
| 129 | if self.merge_strategy == "fixed": |
| 130 | self.register_buffer("mix_factor", torch.Tensor([alpha])) |
| 131 | elif self.merge_strategy == "learned": |
| 132 | self.register_parameter("mix_factor", torch.nn.Parameter(torch.Tensor([alpha]))) |
| 133 | else: |
| 134 | raise ValueError(f"unknown merge strategy {self.merge_strategy}") |
| 135 | |
| 136 | def forward(self, x, timesteps, skip_video=False): |
| 137 | if skip_video: |
nothing calls this directly
no test coverage detected