| 279 | |
| 280 | |
| 281 | class AlphaBlender(nn.Module): |
| 282 | strategies = ["learned", "fixed", "learned_with_images"] |
| 283 | |
| 284 | def __init__( |
| 285 | self, |
| 286 | alpha: float, |
| 287 | merge_strategy: str = "learned_with_images", |
| 288 | rearrange_pattern: str = "b t -> (b t) 1 1", |
| 289 | ): |
| 290 | super().__init__() |
| 291 | self.merge_strategy = merge_strategy |
| 292 | self.rearrange_pattern = rearrange_pattern |
| 293 | |
| 294 | assert merge_strategy in self.strategies, f"merge_strategy needs to be in {self.strategies}" |
| 295 | |
| 296 | if self.merge_strategy == "fixed": |
| 297 | self.register_buffer("mix_factor", torch.Tensor([alpha])) |
| 298 | elif self.merge_strategy == "learned" or self.merge_strategy == "learned_with_images": |
| 299 | self.register_parameter("mix_factor", torch.nn.Parameter(torch.Tensor([alpha]))) |
| 300 | else: |
| 301 | raise ValueError(f"unknown merge strategy {self.merge_strategy}") |
| 302 | |
| 303 | def get_alpha(self, image_only_indicator: torch.Tensor) -> torch.Tensor: |
| 304 | if self.merge_strategy == "fixed": |
| 305 | alpha = self.mix_factor |
| 306 | elif self.merge_strategy == "learned": |
| 307 | alpha = torch.sigmoid(self.mix_factor) |
| 308 | elif self.merge_strategy == "learned_with_images": |
| 309 | assert image_only_indicator is not None, "need image_only_indicator ..." |
| 310 | alpha = torch.where( |
| 311 | image_only_indicator.bool(), |
| 312 | torch.ones(1, 1, device=image_only_indicator.device), |
| 313 | rearrange(torch.sigmoid(self.mix_factor), "... -> ... 1"), |
| 314 | ) |
| 315 | alpha = rearrange(alpha, self.rearrange_pattern) |
| 316 | else: |
| 317 | raise NotImplementedError |
| 318 | return alpha |
| 319 | |
| 320 | def forward( |
| 321 | self, |
| 322 | x_spatial: torch.Tensor, |
| 323 | x_temporal: torch.Tensor, |
| 324 | image_only_indicator: Optional[torch.Tensor] = None, |
| 325 | ) -> torch.Tensor: |
| 326 | alpha = self.get_alpha(image_only_indicator) |
| 327 | x = alpha.to(x_spatial.dtype) * x_spatial + (1.0 - alpha).to(x_spatial.dtype) * x_temporal |
| 328 | return x |