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

Method __init__

shap_e/models/generation/latent_diffusion.py:8–16  ·  view source on GitHub ↗
(self, *, device: torch.device, wrapped: nn.Module, n_ctx: int, d_latent: int)

Source from the content-addressed store, hash-verified

6
7class 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)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected