| 48 | |
| 49 | class IntermediateFeatureNormalizer(IntermediateFeatureNormalizerBase): |
| 50 | def __init__(self, num_intermediates: int, embed_dim: int, rot_per_layer: bool = False, dtype: torch.dtype = torch.float32): |
| 51 | super().__init__() |
| 52 | self.register_buffer('alphas', torch.ones(num_intermediates, dtype=dtype)) |
| 53 | |
| 54 | rot = torch.eye(embed_dim, dtype=dtype) |
| 55 | if rot_per_layer: |
| 56 | rot = rot.unsqueeze(0).repeat(num_intermediates, 1, 1) |
| 57 | |
| 58 | self.register_buffer('rotation', rot.contiguous()) |
| 59 | self.register_buffer('means', torch.zeros(num_intermediates, embed_dim, dtype=dtype)) |
| 60 | |
| 61 | def forward(self, x: torch.Tensor, index: int, rot_index: int = None, skip: Optional[int] = None) -> InterFeatState: |
| 62 | if rot_index is None: |