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

Function uncond_guide_model

shap_e/diffusion/sample.py:15–28  ·  view source on GitHub ↗
(
    model: Callable[..., torch.Tensor], scale: float
)

Source from the content-addressed store, hash-verified

13
14
15def uncond_guide_model(
16 model: Callable[..., torch.Tensor], scale: float
17) -> Callable[..., torch.Tensor]:
18 def model_fn(x_t, ts, **kwargs):
19 half = x_t[: len(x_t) // 2]
20 combined = torch.cat([half, half], dim=0)
21 model_out = model(combined, ts, **kwargs)
22 eps, rest = model_out[:, :3], model_out[:, 3:]
23 cond_eps, uncond_eps = torch.chunk(eps, 2, dim=0)
24 half_eps = uncond_eps + scale * (cond_eps - uncond_eps)
25 eps = torch.cat([half_eps, half_eps], dim=0)
26 return torch.cat([eps, rest], dim=1)
27
28 return model_fn
29
30
31def sample_latents(

Callers 1

sample_latentsFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected