| 3 | |
| 4 | @torch.no_grad() |
| 5 | def sample_images(nr_images, model_ema, model_config, nr_iters=100, extra_args={}, callback=None): |
| 6 | model_ema.eval() |
| 7 | sigma_min = model_config['sigma_min'] |
| 8 | sigma_max = model_config['sigma_max'] |
| 9 | size = model_config['input_size'] |
| 10 | n_per_proc = nr_images |
| 11 | x = torch.randn([1, n_per_proc, model_config['input_channels'], size[0], size[1]]).cuda() |
| 12 | x = x[0] * sigma_max |
| 13 | model_fn = model_ema |
| 14 | sigmas = K.sampling.get_sigmas_karras(nr_iters, sigma_min, sigma_max, rho=7., device="cuda") |
| 15 | x_0 = K.sampling.sample_dpmpp_2m_sde(model_fn, x, sigmas, extra_args=extra_args, eta=0.0, solver_type='heun', disable=False, callback=callback) |
| 16 | return x_0 |
| 17 | |
| 18 | @torch.no_grad() |
| 19 | #samples using classifier free guidance and only enables the cfg_val when the sigma is within the interval. |