MCPcopy Create free account
hub / github.com/NVlabs/DiffusionNFT / ddim_step

Function ddim_step

flow_grpo/diffusers_patch/solver.py:159–185  ·  view source on GitHub ↗
(
    model_output: torch.Tensor,
    latents: torch.Tensor,
    eta: float,
    sigmas: torch.Tensor,
    index: int,
    prev_sample: torch.Tensor,
)

Source from the content-addressed store, hash-verified

157
158
159def ddim_step(
160 model_output: torch.Tensor,
161 latents: torch.Tensor,
162 eta: float,
163 sigmas: torch.Tensor,
164 index: int,
165 prev_sample: torch.Tensor,
166):
167 model_output = convert_model_output(model_output, latents, sigmas, step_index=index)
168 prev_sample, prev_sample_mean, std_dev_t, dt_sqrt = ddim_update(
169 model_output,
170 sigmas.to(torch.float64),
171 index,
172 latents,
173 eta=eta,
174 )
175
176 # Compute log_prob
177 log_prob = (
178 -((prev_sample.detach() - prev_sample_mean) ** 2) / (2 * ((std_dev_t * dt_sqrt) ** 2))
179 - torch.log(std_dev_t * dt_sqrt)
180 - torch.log(torch.sqrt(2 * torch.as_tensor(math.pi)))
181 )
182
183 # mean along all but batch dimension
184 log_prob = log_prob.mean(dim=tuple(range(1, log_prob.ndim)))
185 return prev_sample, model_output, log_prob
186
187
188@dataclass

Callers 1

run_samplingFunction · 0.85

Calls 3

convert_model_outputFunction · 0.85
ddim_updateFunction · 0.85
toMethod · 0.80

Tested by

no test coverage detected