diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index 5ead3d0..5af8378 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -51,6 +51,8 @@ class _ActiveStep: 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 + prediction_next_row: int = 0 + used_forecast_any: bool = False class SpectrumRuntime: @@ -122,6 +124,7 @@ class SpectrumRuntime: step.call_used_forecast = [False] * len(step.call_used_forecast) step.predicted_full_feature = None step.prediction_row_positions = None + step.prediction_next_row = 0 self.stats.current_window = self.curr_ws self.stats.forecast_disabled = True self.stats.disable_reason = reason @@ -254,7 +257,7 @@ class SpectrumRuntime: def step_used_forecast(self, run_id: int, solver_step_id: int) -> bool: step = self._require_active_step(run_id, solver_step_id) - return any(step.call_used_forecast) + return step.used_forecast_any or any(step.call_used_forecast) @staticmethod def _split_branch_signature( @@ -367,6 +370,7 @@ class SpectrumRuntime: target_shape = tuple(expected_shape) if expected_shape is not None else step.call_expected_shapes[resolved_call_id] target_batch_labels = step.call_batch_labels[resolved_call_id] history_shape = self.forecaster.feature_shape + needs_full_prediction = False if history_shape is not None: history_shape = tuple(history_shape) if history_shape[1:] != target_shape[1:]: @@ -374,15 +378,18 @@ class SpectrumRuntime: return None if self._history_batch_labels is None and history_shape[0] != target_shape[0]: return None + needs_full_prediction = (self._history_batch_labels is not None) - 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: + if ( + self._history_batch_labels is None + and step.hook_call_count > 1 + and not needs_full_prediction + and step.predicted_full_feature is None + ): return None if step.call_predicted_features[resolved_call_id] is None: - if self._history_batch_labels is not None: + if needs_full_prediction: if step.predicted_full_feature is None: predicted_full_feature = self.forecaster.predict( time_coord=step.time_coord, @@ -392,20 +399,29 @@ class SpectrumRuntime: 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) + step.prediction_next_row = 0 + if self._history_batch_labels is not None: + step.prediction_row_positions = self._build_label_positions(self._history_batch_labels) + else: + step.prediction_row_positions = None - 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 target_batch_labels is not None and self._history_batch_labels is not None: + if step.prediction_row_positions is None: 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, ...] + 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: + self._disable_forecasting("forecasted solver step batch layout changed within one solver step") + return None else: predicted_feature = self.forecaster.predict( time_coord=step.time_coord, @@ -417,6 +433,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: @@ -425,6 +442,7 @@ class SpectrumRuntime: step.call_used_forecast = [False] * len(step.call_used_forecast) step.predicted_full_feature = None step.prediction_row_positions = None + step.prediction_next_row = 0 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: @@ -432,7 +450,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: @@ -441,10 +459,15 @@ class SpectrumRuntime: 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 + elif used_forecast_any and step.predicted_full_feature is not None: + if step.prediction_row_positions is not None: + if any(positions for positions in step.prediction_row_positions.values()): + self._disable_forecasting("forecasted solver step batch layout changed within one solver step") + elif step.prediction_next_row != step.predicted_full_feature.shape[0]: + self._disable_forecasting("forecasted solver step batch layout changed within one solver step") if used_forecast_any: self.num_consecutive_cached_steps += 1 diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index 37a1903..6bb25cb 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -327,6 +327,42 @@ def test_split_forecast_step_uses_spectrum_across_subcalls() -> None: runtime.end_run(run_id) +def test_split_forecast_step_without_layout_labels_falls_back_to_actual() -> 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, + ) + call_id = runtime.register_model_hook_call(run_id, step_id, expected_shape=(2, 8, 4), branch_signature=None) + runtime.observe_actual_feature(run_id, step_id, torch.randn(2, 8, 4), call_id=call_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=None) + assert runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=first_id) is None + runtime.observe_actual_feature(run_id, 5, torch.randn(1, 8, 4), call_id=first_id) + 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.randn(1, 8, 4), call_id=second_id) + runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False) + assert runtime.stats.forecast_disabled is False + assert runtime.stats.forecasted_count == 0 + 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) @@ -416,6 +452,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) @@ -541,9 +618,11 @@ def main() -> None: test_nonbatch_shape_mismatch_disables_forecast() test_multiple_hook_calls_are_aggregated_on_actual_steps() test_split_forecast_step_uses_spectrum_across_subcalls() + test_split_forecast_step_without_layout_labels_falls_back_to_actual() 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()