MCPcopy Create free account
hub / github.com/QWTforGithub/T2LDM / sample

Method sample

models/diffusion/discrete_time.py:190–215  ·  view source on GitHub ↗
(
        self,
        batch_size: int,
        num_steps: int,
        progress: bool = True,
        rng: list[torch.Generator] | torch.Generator | None = None,
        return_noise: bool = False,
        mode: Literal["ddpm", "ddim"] = "ddpm",
        text_features: Tensor | None = None,
        text_null_features: Tensor | None = None,
    )

Source from the content-addressed store, hash-verified

188
189 @torch.inference_mode()
190 def sample(
191 self,
192 batch_size: int,
193 num_steps: int,
194 progress: bool = True,
195 rng: list[torch.Generator] | torch.Generator | None = None,
196 return_noise: bool = False,
197 mode: Literal["ddpm", "ddim"] = "ddpm",
198 text_features: Tensor | None = None,
199 text_null_features: Tensor | None = None,
200 ):
201 noise = self.randn(batch_size, *self.sampling_shape, rng=rng, device=self.device)
202 x = noise
203
204 tqdm_kwargs = dict(desc="sampling", leave=False, disable=not progress)
205 for timestep in tqdm(list(reversed(range(num_steps))), **tqdm_kwargs):
206 timesteps = torch.full((batch_size,), timestep, device=self.device).long()
207 x = self.p_sample(
208 x,
209 timesteps,
210 text_features=text_features,
211 text_null_features=text_null_features,
212 mode=mode
213 )
214
215 return noise, x if return_noise else x

Callers

nothing calls this directly

Calls 3

p_sampleMethod · 0.95
tqdmFunction · 0.85
randnMethod · 0.80

Tested by

no test coverage detected