:param x: an [N x C x T] tensor. :param t: an [N] tensor. :return: an [N x C' x T] tensor.
(self, x: torch.Tensor, t: torch.Tensor)
| 222 | self.output_proj.bias.zero_() |
| 223 | |
| 224 | def forward(self, x: torch.Tensor, t: torch.Tensor): |
| 225 | """ |
| 226 | :param x: an [N x C x T] tensor. |
| 227 | :param t: an [N] tensor. |
| 228 | :return: an [N x C' x T] tensor. |
| 229 | """ |
| 230 | assert x.shape[-1] == self.decoder.n_ctx |
| 231 | t_embed = self.time_embed(timestep_embedding(t, self.encoder.width)) |
| 232 | data = self.input_proj(x.permute(0, 2, 1)) + t_embed[:, None] |
| 233 | data = self.ln_pre(data) |
| 234 | |
| 235 | l = torch.arange(self.n_latent).to(x.device) |
| 236 | h = self.latent_embed(timestep_embedding(l, self.decoder.width)) |
| 237 | h = h.unsqueeze(0).repeat(x.shape[0], 1, 1) |
| 238 | |
| 239 | h = self.encoder(h, data) |
| 240 | h = self.processor(h) |
| 241 | h = self.decoder(data, h) |
| 242 | h = self.ln_post(h) |
| 243 | h = self.output_proj(h) |
| 244 | return h.permute(0, 2, 1) |
nothing calls this directly
no test coverage detected