From a73142188a1a9cd1748ef20388ade05ec4bf4a53 Mon Sep 17 00:00:00 2001 From: blepping Date: Wed, 3 Apr 2024 20:03:39 -0600 Subject: [PATCH] Only increment steps/delegate to callback when the sigma changes --- restart_sampling.py | 19 +++++++++++++++---- 1 file changed, 15 insertions(+), 4 deletions(-) diff --git a/restart_sampling.py b/restart_sampling.py index 542a1f1..192f13f 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -457,14 +457,25 @@ class KSamplerRestartWrapper: return result with trange(self.total_steps, disable=disable) as pbar: + last_cb_sigma = None - def callback_wrapper(x): - nonlocal step + def callback_wrapper(cb_state): + nonlocal step, last_cb_sigma + curr_sigma = cb_state.get("sigma") + curr_sigma = ( + curr_sigma.item() + if isinstance(curr_sigma, torch.Tensor) + else curr_sigma + ) + if last_cb_sigma is not None and curr_sigma == last_cb_sigma: + # No change since last time we were called, so we won't track it as a step. + return step += 1 pbar.update(1) - x["i"] = step + cb_state["i"] = step + last_cb_sigma = curr_sigma if callback is not None: - callback(x) + callback(cb_state) # Convenience function for code reuse. def sampler_function(x, sigs):