diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index 8d14149..23f6325 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -52,6 +52,7 @@ class _ActiveStep: predicted_full_feature: Optional[torch.Tensor] = None prediction_row_positions: Optional[dict[int, deque[int]]] = None prediction_next_row: int = 0 + used_forecast_any: bool = False class SpectrumRuntime: @@ -444,6 +445,7 @@ class SpectrumRuntime: step.call_predicted_features[resolved_call_id] = predicted_feature step.call_used_forecast[resolved_call_id] = True + step.used_forecast_any = True return step.call_predicted_features[resolved_call_id] def abort_solver_step(self, run_id: int, solver_step_id: int) -> None: @@ -460,7 +462,7 @@ class SpectrumRuntime: requested_actual_forward = bool(step.decision["actual_forward"]) observed_actual = any(step.call_observed_actual) - used_forecast_any = any(step.call_used_forecast) + used_forecast_any = step.used_forecast_any or any(step.call_used_forecast) if bool(used_forecast) and not observed_actual: used_forecast_any = True if step.decision["actual_forward"] and not observed_actual: diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index dd86878..0194e64 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -448,6 +448,47 @@ def test_mixed_batch_layout_presence_disables_forecast() -> None: runtime.end_run(run_id) +def test_split_forecast_failure_after_first_slice_still_counts_as_mixed_path() -> 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), 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), 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, + ) + first_id = runtime.register_model_hook_call( + run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)),) + ) + assert runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=first_id) is not None + second_id = runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 8, 4), branch_signature=None) + assert runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=second_id) is None + runtime.observe_actual_feature(run_id, 5, torch.ones(1, 8, 4), call_id=second_id) + runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False) + assert runtime.stats.disable_reason == "solver step mixed forecasted and actual model-hook paths" + 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) @@ -577,6 +618,7 @@ def main() -> None: test_duplicate_batch_labels_can_be_reordered_for_forecast() test_topology_change_disables_forecast() test_mixed_batch_layout_presence_disables_forecast() + test_split_forecast_failure_after_first_slice_still_counts_as_mixed_path() test_aborted_solver_step_is_discarded_without_disabling_forecast() test_nonuniform_schedule_coords_are_used() test_forecaster_respects_nonuniform_coords()