Handle split forecast steps and abort interrupted solver steps

This commit is contained in:
xmarre
2026-03-30 21:35:38 +02:00
parent d00b05b9a1
commit c56d43da02
3 changed files with 137 additions and 23 deletions
+14 -5
View File
@@ -229,16 +229,25 @@ def _install_sampler_level_wrappers(model: Any, runtime: SpectrumRuntime) -> Non
patched_model_options = _copy_model_options_with_step_context(
effective_model_options, wrapper_runtime, decision
)
step_aborted = False
try:
return executor(x, timestep, patched_model_options, seed)
finally:
wrapper_runtime.finalize_solver_step(
except BaseException:
step_aborted = True
wrapper_runtime.abort_solver_step(
decision["run_id"],
decision["solver_step_id"],
used_forecast=wrapper_runtime.step_used_forecast(
decision["run_id"], decision["solver_step_id"]
),
)
raise
finally:
if not step_aborted:
wrapper_runtime.finalize_solver_step(
decision["run_id"],
decision["solver_step_id"],
used_forecast=wrapper_runtime.step_used_forecast(
decision["run_id"], decision["solver_step_id"]
),
)
comfy.patcher_extension.add_wrapper(
comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
+53 -18
View File
@@ -49,6 +49,8 @@ class _ActiveStep:
call_used_forecast: list[bool] = field(default_factory=list)
call_actual_features: list[Optional[torch.Tensor]] = field(default_factory=list)
call_predicted_features: list[Optional[torch.Tensor]] = field(default_factory=list)
predicted_full_feature: Optional[torch.Tensor] = None
prediction_row_positions: Optional[dict[int, deque[int]]] = None
class SpectrumRuntime:
@@ -118,6 +120,8 @@ class SpectrumRuntime:
for step in self._active_steps.values():
step.call_predicted_features = [None] * len(step.call_predicted_features)
step.call_used_forecast = [False] * len(step.call_used_forecast)
step.predicted_full_feature = None
step.prediction_row_positions = None
self.stats.current_window = self.curr_ws
self.stats.forecast_disabled = True
self.stats.disable_reason = reason
@@ -338,6 +342,13 @@ class SpectrumRuntime:
order = [source_positions[int(label)].popleft() for label in target_labels]
return feature[order, ...]
@staticmethod
def _build_label_positions(labels: tuple[int, ...]) -> dict[int, deque[int]]:
positions: dict[int, deque[int]] = defaultdict(deque)
for idx, label in enumerate(labels):
positions[int(label)].append(idx)
return positions
def predict_feature(
self,
run_id: int,
@@ -361,37 +372,61 @@ class SpectrumRuntime:
if history_shape[1:] != target_shape[1:]:
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
return None
if history_shape[0] != target_shape[0]:
if self._history_batch_labels is None and history_shape[0] != target_shape[0]:
return None
if step.hook_call_count > 1:
if self._history_batch_labels is None and step.hook_call_count > 1:
return None
if self._history_batch_labels is not None and target_batch_labels is None:
return None
if step.call_predicted_features[resolved_call_id] is None:
predicted_feature = self.forecaster.predict(
time_coord=step.time_coord,
blend_weight=self.cfg.blend_weight,
)
if tuple(predicted_feature.shape) != target_shape:
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
return None
if self._history_batch_labels is not None and target_batch_labels is not None:
reordered = self._reorder_feature_to_labels(
predicted_feature,
self._history_batch_labels,
target_batch_labels,
if self._history_batch_labels is not None:
if step.predicted_full_feature is None:
predicted_full_feature = self.forecaster.predict(
time_coord=step.time_coord,
blend_weight=self.cfg.blend_weight,
)
if history_shape is not None and tuple(predicted_full_feature.shape) != history_shape:
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
return None
step.predicted_full_feature = predicted_full_feature
step.prediction_row_positions = self._build_label_positions(self._history_batch_labels)
if target_batch_labels is None or step.prediction_row_positions is None:
return None
order = []
for label in target_batch_labels:
positions = step.prediction_row_positions.get(int(label))
if not positions:
if any(step.call_used_forecast):
self._disable_forecasting("forecasted solver step batch layout changed within one solver step")
return None
order.append(positions.popleft())
predicted_feature = step.predicted_full_feature[order, ...]
else:
predicted_feature = self.forecaster.predict(
time_coord=step.time_coord,
blend_weight=self.cfg.blend_weight,
)
if reordered is None:
if tuple(predicted_feature.shape) != target_shape:
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
return None
predicted_feature = reordered
step.call_predicted_features[resolved_call_id] = predicted_feature
step.call_used_forecast[resolved_call_id] = True
return step.call_predicted_features[resolved_call_id]
def abort_solver_step(self, run_id: int, solver_step_id: int) -> None:
step = self._require_active_step(run_id, solver_step_id)
step.call_predicted_features = [None] * len(step.call_predicted_features)
step.call_used_forecast = [False] * len(step.call_used_forecast)
step.predicted_full_feature = None
step.prediction_row_positions = None
self._active_steps.pop(int(solver_step_id), None)
def finalize_solver_step(self, run_id: int, solver_step_id: int, *, used_forecast: bool) -> None:
step = self._require_active_step(run_id, solver_step_id)
requested_actual_forward = bool(step.decision["actual_forward"])
@@ -404,14 +439,14 @@ class SpectrumRuntime:
self._disable_forecasting("solver step requested an actual forward but no actual feature was observed")
if not observed_actual and not used_forecast_any:
self._disable_forecasting("solver step finished without an actual feature or a forecasted feature")
if any(not (obs or used) for obs, used in zip(step.call_observed_actual, step.call_used_forecast)):
self._disable_forecasting("solver step finished with an incomplete model-hook call")
if observed_actual and used_forecast_any:
self._disable_forecasting("solver step mixed forecasted and actual model-hook paths")
used_forecast_any = False
if used_forecast_any:
if step.hook_call_count > 1:
self._disable_forecasting("forecasted solver step re-entered the model hook")
self.num_consecutive_cached_steps += 1
self.stats.forecasted_count += 1
step.decision["actual_forward"] = False
+70
View File
@@ -276,6 +276,57 @@ def test_multiple_hook_calls_are_aggregated_on_actual_steps() -> None:
runtime.end_run(run_id)
def test_split_forecast_step_uses_spectrum_across_subcalls() -> None:
runtime = make_runtime()
sample_sigmas = torch.linspace(1.0, 0.0, 51)
run_id = runtime.start_run(sample_sigmas, "sample_euler", supports_solver_steps=True)
total_steps = len(sample_sigmas) - 1
for step_id in range(5):
runtime.begin_solver_step(
run_id,
step_id,
runtime.time_coord_for_step(step_id),
total_steps,
)
first_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)),)
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(1, 8, 4) * 10.0, call_id=first_id)
second_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)),)
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(1, 8, 4) * 20.0, call_id=second_id)
runtime.finalize_solver_step(run_id, step_id, used_forecast=False)
decision = runtime.begin_solver_step(
run_id,
5,
runtime.time_coord_for_step(5),
total_steps,
)
assert decision["actual_forward"] is False
first_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)),)
)
first_pred = runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=first_id)
assert first_pred is not None
assert first_pred.shape == (1, 8, 4)
second_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)),)
)
second_pred = runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=second_id)
assert second_pred is not None
assert second_pred.shape == (1, 8, 4)
runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=True)
assert runtime.stats.forecast_disabled is False
assert runtime.stats.forecasted_count == 1
runtime.end_run(run_id)
def test_duplicate_batch_labels_can_be_reordered_for_forecast() -> None:
runtime = make_runtime()
sample_sigmas = torch.linspace(1.0, 0.0, 51)
@@ -365,6 +416,23 @@ def test_mixed_batch_layout_presence_disables_forecast() -> None:
runtime.end_run(run_id)
def test_aborted_solver_step_is_discarded_without_disabling_forecast() -> None:
runtime = make_runtime()
sample_sigmas = torch.linspace(1.0, 0.0, 51)
run_id = runtime.start_run(sample_sigmas, "sample_euler", supports_solver_steps=True)
total_steps = len(sample_sigmas) - 1
decision = runtime.begin_solver_step(
run_id,
0,
runtime.time_coord_for_step(0),
total_steps,
)
runtime.abort_solver_step(decision["run_id"], decision["solver_step_id"])
assert runtime.stats.forecast_disabled is False
runtime.end_run(run_id)
def test_nonuniform_schedule_coords_are_used() -> None:
runtime = make_runtime()
sample_sigmas = torch.tensor([10.0, 9.0, 1.0, 0.0])
@@ -472,9 +540,11 @@ def main() -> None:
test_batch_split_falls_back_to_actual_without_disabling_run()
test_nonbatch_shape_mismatch_disables_forecast()
test_multiple_hook_calls_are_aggregated_on_actual_steps()
test_split_forecast_step_uses_spectrum_across_subcalls()
test_duplicate_batch_labels_can_be_reordered_for_forecast()
test_topology_change_disables_forecast()
test_mixed_batch_layout_presence_disables_forecast()
test_aborted_solver_step_is_discarded_without_disabling_forecast()
test_nonuniform_schedule_coords_are_used()
test_forecaster_respects_nonuniform_coords()
test_flux_sampler_contract_only_allows_euler()