diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index 5ead3d0..0ad46b7 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -51,6 +51,7 @@ 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 class SpectrumRuntime: @@ -122,6 +123,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 @@ -367,22 +369,27 @@ 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:]: self._disable_forecasting("predicted feature shape did not match the current solver-step input") 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) + or (history_shape[0] != target_shape[0]) + ) - 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,38 @@ 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: + if step.prediction_row_positions is not None and any(step.call_used_forecast): + self._disable_forecasting("forecasted solver step batch layout changed within one solver step") + return None + start = step.prediction_next_row + end = start + target_shape[0] + if step.predicted_full_feature is None or end > step.predicted_full_feature.shape[0]: + if any(step.call_used_forecast): + self._disable_forecasting("forecasted solver step batch layout changed within one solver step") + return None + predicted_feature = step.predicted_full_feature[start:end, ...] + step.prediction_next_row = end else: predicted_feature = self.forecaster.predict( time_coord=step.time_coord, @@ -425,6 +450,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: @@ -441,6 +467,12 @@ 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 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 observed_actual and used_forecast_any: self._disable_forecasting("solver step mixed forecasted and actual model-hook paths") diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index 37a1903..dd86878 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -327,6 +327,38 @@ def test_split_forecast_step_uses_spectrum_across_subcalls() -> None: runtime.end_run(run_id) +def test_split_forecast_step_without_layout_labels_uses_sequential_full_prediction() -> 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 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 not None + runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=True) + 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) @@ -541,6 +573,7 @@ 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_uses_sequential_full_prediction() test_duplicate_batch_labels_can_be_reordered_for_forecast() test_topology_change_disables_forecast() test_mixed_batch_layout_presence_disables_forecast()