diff --git a/comfyui_spectrum/flux.py b/comfyui_spectrum/flux.py index a547b1b..e857624 100644 --- a/comfyui_spectrum/flux.py +++ b/comfyui_spectrum/flux.py @@ -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, diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index ef556a1..5ead3d0 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -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 diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index d22a268..37a1903 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -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()