MCPcopy Create free account
hub / github.com/CompVis/diff2flow / decode

Method decode

diff2flow/ddim.py:284–324  ·  view source on GitHub ↗
(
        self,
        model,
        x_latent,
        t_start=None,
        ddim_steps=100,
        use_original_steps=True,
        model_kwargs=None,
        progress=True,
        clip_denoised=False
    )

Source from the content-addressed store, hash-verified

282
283 @torch.no_grad()
284 def decode(
285 self,
286 model,
287 x_latent,
288 t_start=None,
289 ddim_steps=100,
290 use_original_steps=True,
291 model_kwargs=None,
292 progress=True,
293 clip_denoised=False
294 ):
295 bs, dev = x_latent.shape[0], x_latent.device
296 model_kwargs = model_kwargs or {}
297
298 self.make_schedule(ddim_num_steps=ddim_steps, device=dev, ddim_eta=0.0, verbose=False)
299
300 timesteps = np.arange(self.ddpm_num_timesteps) if use_original_steps else self.ddim_timesteps
301 if t_start is not None:
302 timesteps = timesteps[:t_start]
303
304 time_range = np.flip(timesteps)
305 total_steps = timesteps.shape[0]
306
307 iterator = tqdm(time_range, desc='Decoding image', total=total_steps, disable=not progress)
308
309 x_dec = x_latent
310 for i, step in enumerate(iterator):
311 index = total_steps - i - 1
312 ts = torch.full((bs,), step, device=dev, dtype=torch.long)
313
314 x_dec, _ = self.p_sample_ddim(
315 model=model,
316 x=x_dec,
317 t=ts,
318 index=index,
319 model_kwargs=model_kwargs,
320 use_original_steps=use_original_steps,
321 clip_denoised=clip_denoised
322 )
323
324 return x_dec

Callers

nothing calls this directly

Calls 2

make_scheduleMethod · 0.95
p_sample_ddimMethod · 0.95

Tested by

no test coverage detected