MCPcopy Create free account
hub / github.com/openai/shap-e / forward

Method forward

shap_e/models/generation/perceiver.py:224–244  ·  view source on GitHub ↗

: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)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 2

timestep_embeddingFunction · 0.85
toMethod · 0.80

Tested by

no test coverage detected