| 13 | |
| 14 | |
| 15 | def 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 | |
| 31 | def sample_latents( |