(self, z, temperature=1.0, cfg=1.0)
| 33 | return loss.mean() |
| 34 | |
| 35 | def sample(self, z, temperature=1.0, cfg=1.0): |
| 36 | # diffusion loss sampling |
| 37 | if not cfg == 1.0: |
| 38 | noise = torch.randn(z.shape[0] // 2, self.in_channels).cuda() |
| 39 | noise = torch.cat([noise, noise], dim=0) |
| 40 | model_kwargs = dict(c=z, cfg_scale=cfg) |
| 41 | sample_fn = self.net.forward_with_cfg |
| 42 | else: |
| 43 | noise = torch.randn(z.shape[0], self.in_channels).cuda() |
| 44 | model_kwargs = dict(c=z) |
| 45 | sample_fn = self.net.forward |
| 46 | |
| 47 | sampled_token_latent = self.gen_diffusion.p_sample_loop( |
| 48 | sample_fn, noise.shape, noise, clip_denoised=False, model_kwargs=model_kwargs, progress=False, |
| 49 | temperature=temperature |
| 50 | ) |
| 51 | |
| 52 | return sampled_token_latent |
| 53 | |
| 54 | |
| 55 | def modulate(x, shift, scale): |
no test coverage detected