(
self,
samples: Float[Tensor, "*batch dim"],
)
| 26 | self.register_buffer("phases", phases, persistent=False) |
| 27 | |
| 28 | def forward( |
| 29 | self, |
| 30 | samples: Float[Tensor, "*batch dim"], |
| 31 | ) -> Float[Tensor, "*batch embedded_dim"]: |
| 32 | samples = einsum(samples, self.frequencies, "... d, f p -> ... d f p") |
| 33 | return rearrange(torch.sin(samples + self.phases), "... d f p -> ... (d f p)") |
| 34 | |
| 35 | def d_out(self, dimensionality: int): |
| 36 | return self.frequencies.numel() * dimensionality |
nothing calls this directly
no outgoing calls
no test coverage detected