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