(self, *, device: torch.device, wrapped: nn.Module, n_ctx: int, d_latent: int)
| 6 | |
| 7 | class SplitVectorDiffusion(nn.Module): |
| 8 | def __init__(self, *, device: torch.device, wrapped: nn.Module, n_ctx: int, d_latent: int): |
| 9 | super().__init__() |
| 10 | self.device = device |
| 11 | self.n_ctx = n_ctx |
| 12 | self.d_latent = d_latent |
| 13 | self.wrapped = wrapped |
| 14 | |
| 15 | if hasattr(self.wrapped, "cached_model_kwargs"): |
| 16 | self.cached_model_kwargs = self.wrapped.cached_model_kwargs |
| 17 | |
| 18 | def forward(self, x: torch.Tensor, t: torch.Tensor, **kwargs): |
| 19 | h = x.reshape(x.shape[0], self.n_ctx, -1).permute(0, 2, 1) |
nothing calls this directly
no outgoing calls
no test coverage detected