| 731 | strategies = ["learned", "fixed", "learned_with_images"] |
| 732 | |
| 733 | def __init__( |
| 734 | self, |
| 735 | alpha: float, |
| 736 | merge_strategy: str = "learned_with_images", |
| 737 | switch_spatial_to_temporal_mix: bool = False, |
| 738 | ): |
| 739 | super().__init__() |
| 740 | self.merge_strategy = merge_strategy |
| 741 | self.switch_spatial_to_temporal_mix = switch_spatial_to_temporal_mix # For TemporalVAE |
| 742 | |
| 743 | if merge_strategy not in self.strategies: |
| 744 | raise ValueError(f"merge_strategy needs to be in {self.strategies}") |
| 745 | |
| 746 | if self.merge_strategy == "fixed": |
| 747 | self.register_buffer("mix_factor", torch.Tensor([alpha])) |
| 748 | elif self.merge_strategy == "learned" or self.merge_strategy == "learned_with_images": |
| 749 | self.register_parameter("mix_factor", torch.nn.Parameter(torch.Tensor([alpha]))) |
| 750 | else: |
| 751 | raise ValueError(f"Unknown merge strategy {self.merge_strategy}") |
| 752 | |
| 753 | def get_alpha(self, image_only_indicator: torch.Tensor, ndims: int) -> torch.Tensor: |
| 754 | if self.merge_strategy == "fixed": |