From 70f21e181f6ef58a0b5b7ebc6d88768532cd3be7 Mon Sep 17 00:00:00 2001 From: xmarre Date: Mon, 30 Mar 2026 15:26:25 +0200 Subject: [PATCH] Handle split FLUX calls within a solver step --- README.md | 3 + comfyui_spectrum/flux.py | 6 +- comfyui_spectrum/forecast.py | 4 + comfyui_spectrum/runtime.py | 144 +++++++++++++++++++++++++---------- tests/smoke_runtime.py | 94 +++++++++++++++++++---- 5 files changed, 195 insertions(+), 56 deletions(-) diff --git a/README.md b/README.md index fe05769..5bce349 100644 --- a/README.md +++ b/README.md @@ -40,6 +40,9 @@ Supported: - LoRAs on the normal model path - standard `transformer_options` patch chains - standard FLUX control residuals +- ComfyUI runs that sometimes split one logical solver step into multiple internal FLUX calls + +For split-step runs, Spectrum now aggregates actual hidden features across the sub-calls and falls back to the real path for any forecast step whose current call shape no longer matches the cached full-batch history. This avoids false run-wide disables while keeping the forecast path conservative. Not included: diff --git a/comfyui_spectrum/flux.py b/comfyui_spectrum/flux.py index a35bdf7..a547b1b 100644 --- a/comfyui_spectrum/flux.py +++ b/comfyui_spectrum/flux.py @@ -398,9 +398,10 @@ def _run_flux_forward_with_spectrum( extra_kwargs["modulation_dims_img"] = modulation_dims txt_vec = vec[:batch] + call_id: Optional[int] = None if step_ctx is not None: _, run_id, solver_step_id, actual_forward = step_ctx - runtime.register_model_hook_call( + call_id = runtime.register_model_hook_call( run_id, solver_step_id, expected_shape=expected_feature_shape, @@ -411,6 +412,7 @@ def _run_flux_forward_with_spectrum( run_id, solver_step_id, expected_shape=expected_feature_shape, + call_id=call_id, ) if pred_feature is not None: if runtime.cfg.debug: @@ -551,7 +553,7 @@ def _run_flux_forward_with_spectrum( prehead_feature = img[:, txt.shape[1] :, ...] if run_id is not None and solver_step_id is not None: - runtime.observe_actual_feature(run_id, solver_step_id, prehead_feature) + runtime.observe_actual_feature(run_id, solver_step_id, prehead_feature, call_id=call_id) final_kwargs = {} if modulation_dims is not None: diff --git a/comfyui_spectrum/forecast.py b/comfyui_spectrum/forecast.py index 46b4599..a77eac7 100644 --- a/comfyui_spectrum/forecast.py +++ b/comfyui_spectrum/forecast.py @@ -38,6 +38,10 @@ class ChebyshevSpectrumForecaster: self.ridge_lambda = float(ridge_lambda) self.max_history = int(max_history) + @property + def feature_shape(self) -> Optional[torch.Size]: + return self._feature_shape + def ready(self, min_points: Optional[int] = None) -> bool: needed = max(2, int(min_points) if min_points is not None else self.degree + 1) return len(self._history) >= needed diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index 3555161..5588041 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -1,7 +1,7 @@ from __future__ import annotations import math -from dataclasses import dataclass +from dataclasses import dataclass, field from typing import Any, Dict, Optional import torch @@ -38,12 +38,14 @@ class _ActiveStep: solver_step_id: int time_coord: float decision: Dict[str, Any] - expected_shape: Optional[tuple[int, ...]] = None - branch_signature: Optional[tuple[Any, ...]] = None + feature_tail_shape: Optional[tuple[int, ...]] = None hook_call_count: int = 0 - observed_actual: bool = False - used_forecast: bool = False - predicted_feature: Optional[torch.Tensor] = None + call_expected_shapes: list[tuple[int, ...]] = field(default_factory=list) + call_branch_signatures: list[Optional[tuple[Any, ...]]] = field(default_factory=list) + call_observed_actual: list[bool] = field(default_factory=list) + 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) class SpectrumRuntime: @@ -108,7 +110,8 @@ class SpectrumRuntime: self.curr_ws = float(self.cfg.window_size) self.forecaster.reset() for step in self._active_steps.values(): - step.predicted_feature = None + step.call_predicted_features = [None] * len(step.call_predicted_features) + step.call_used_forecast = [False] * len(step.call_used_forecast) self.stats.current_window = self.curr_ws self.stats.forecast_disabled = True self.stats.disable_reason = reason @@ -240,7 +243,8 @@ class SpectrumRuntime: return step.decision def step_used_forecast(self, run_id: int, solver_step_id: int) -> bool: - return self._require_active_step(run_id, solver_step_id).used_forecast + step = self._require_active_step(run_id, solver_step_id) + return any(step.call_used_forecast) def register_model_hook_call( self, @@ -249,31 +253,44 @@ class SpectrumRuntime: *, expected_shape: tuple[int, ...], branch_signature: Optional[tuple[Any, ...]] = None, - ) -> None: + ) -> int: step = self._require_active_step(run_id, solver_step_id) + shape = tuple(expected_shape) + tail_shape = shape[1:] step.hook_call_count += 1 - if step.hook_call_count > 1: - self._disable_forecasting("multiple model-hook calls observed within one solver step") - if step.expected_shape is None: - step.expected_shape = tuple(expected_shape) - elif tuple(expected_shape) != step.expected_shape: + if step.feature_tail_shape is None: + step.feature_tail_shape = tail_shape + elif tail_shape != step.feature_tail_shape: self._disable_forecasting("model-hook feature shape changed within one solver step") - if branch_signature is None: - return - if step.branch_signature is None: - step.branch_signature = branch_signature - elif branch_signature != step.branch_signature: - self._disable_forecasting("model-hook branch signature changed within one solver step") + if branch_signature is not None: + for prev_branch_signature in step.call_branch_signatures: + if prev_branch_signature is not None and prev_branch_signature != branch_signature: + self._disable_forecasting("model-hook branch signature changed within one solver step") + break - def observe_actual_feature(self, run_id: int, solver_step_id: int, feature: torch.Tensor) -> None: + step.call_expected_shapes.append(shape) + step.call_branch_signatures.append(branch_signature) + step.call_observed_actual.append(False) + step.call_used_forecast.append(False) + step.call_actual_features.append(None) + step.call_predicted_features.append(None) + return len(step.call_expected_shapes) - 1 + + def observe_actual_feature( + self, + run_id: int, + solver_step_id: int, + feature: torch.Tensor, + *, + call_id: Optional[int] = None, + ) -> None: step = self._require_active_step(run_id, solver_step_id) - step.observed_actual = True - step.used_forecast = False - step.predicted_feature = None - if self.forecast_disabled: - return - self.forecaster.update(step.time_coord, feature) + resolved_call_id = self._resolve_call_id(step, call_id) + step.call_observed_actual[resolved_call_id] = True + step.call_used_forecast[resolved_call_id] = False + step.call_predicted_features[resolved_call_id] = None + step.call_actual_features[resolved_call_id] = feature.detach() def predict_feature( self, @@ -281,41 +298,77 @@ class SpectrumRuntime: solver_step_id: int, *, expected_shape: Optional[tuple[int, ...]] = None, + call_id: Optional[int] = None, ) -> Optional[torch.Tensor]: step = self._require_active_step(run_id, solver_step_id) + resolved_call_id = self._resolve_call_id(step, call_id) if step.decision["actual_forward"]: return None if self.forecast_disabled or not self.forecaster.ready(self.min_fit_points): return None - if step.predicted_feature is None: - step.predicted_feature = self.forecaster.predict( + target_shape = tuple(expected_shape) if expected_shape is not None else step.call_expected_shapes[resolved_call_id] + history_shape = self.forecaster.feature_shape + 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 history_shape[0] != target_shape[0]: + return None + + if step.hook_call_count > 1: + 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 + step.call_predicted_features[resolved_call_id] = predicted_feature - if expected_shape is not None and tuple(step.predicted_feature.shape) != tuple(expected_shape): - self._disable_forecasting("predicted feature shape did not match the current solver-step input") - return None - - step.used_forecast = True - return step.predicted_feature + step.call_used_forecast[resolved_call_id] = True + return step.call_predicted_features[resolved_call_id] 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) - if bool(used_forecast): - step.used_forecast = True - if step.decision["actual_forward"] and not step.observed_actual: + requested_actual_forward = bool(step.decision["actual_forward"]) + if bool(used_forecast) and step.call_used_forecast: + step.call_used_forecast[-1] = True + + observed_actual = any(step.call_observed_actual) + used_forecast_any = any(step.call_used_forecast) + if step.decision["actual_forward"] and not observed_actual: self._disable_forecasting("solver step requested an actual forward but no actual feature was observed") - if not step.observed_actual and not step.used_forecast: + if not observed_actual and not used_forecast_any: self._disable_forecasting("solver step finished without an actual feature or a forecasted feature") - if step.used_forecast: + 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 else: - if not self.forecast_disabled and step.solver_step_id >= self.cfg.warmup_steps: + actual_parts = [part for part in step.call_actual_features if part is not None] + if actual_parts and not self.forecast_disabled: + combined_feature = actual_parts[0] if len(actual_parts) == 1 else torch.cat(actual_parts, dim=0) + try: + self.forecaster.update(step.time_coord, combined_feature) + except ValueError: + self._disable_forecasting("combined actual feature shape changed across solver steps") + if ( + requested_actual_forward + and not self.forecast_disabled + and step.solver_step_id >= self.cfg.warmup_steps + ): self.curr_ws = round(self.curr_ws + float(self.cfg.flex_window), 6) self.num_consecutive_cached_steps = 0 self.stats.actual_forward_count += 1 @@ -324,6 +377,15 @@ class SpectrumRuntime: self.stats.current_window = self.curr_ws self._active_steps.pop(int(solver_step_id), None) + @staticmethod + def _resolve_call_id(step: _ActiveStep, call_id: Optional[int]) -> int: + if not step.call_expected_shapes: + raise RuntimeError("Spectrum solver step has no active model-hook call.") + resolved = len(step.call_expected_shapes) - 1 if call_id is None else int(call_id) + if resolved < 0 or resolved >= len(step.call_expected_shapes): + raise RuntimeError(f"Spectrum solver-step call id {resolved} is not active.") + return resolved + def _require_active_step(self, run_id: int, solver_step_id: int) -> _ActiveStep: if self._active_run is None or self._active_run.run_id != int(run_id): raise RuntimeError("Spectrum runtime is not inside the requested sampling run.") diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index e78fdc8..01290eb 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -170,7 +170,46 @@ def test_unsupported_sampler_disables_forecast() -> None: runtime.end_run(run_id) -def test_inconsistent_hook_shape_disables_forecast() -> None: +def test_batch_split_falls_back_to_actual_without_disabling_run() -> 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): + decision = runtime.begin_solver_step( + run_id, + step_id, + runtime.time_coord_for_step(step_id), + total_steps, + ) + runtime.register_model_hook_call(run_id, step_id, expected_shape=(2, 8, 4)) + runtime.observe_actual_feature(run_id, step_id, torch.randn(2, 8, 4)) + 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 + call_id = runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 8, 4)) + predicted = runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=call_id) + assert predicted is None + assert runtime.stats.forecast_disabled is False + runtime.observe_actual_feature(run_id, 5, torch.randn(1, 8, 4), call_id=call_id) + runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 8, 4)) + runtime.observe_actual_feature(run_id, 5, torch.randn(1, 8, 4)) + runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False) + + assert runtime.stats.forecasted_count == 0 + assert runtime.stats.actual_forward_count == 6 + assert decision["actual_forward"] is True + runtime.end_run(run_id) + + +def test_nonbatch_shape_mismatch_disables_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) @@ -193,17 +232,17 @@ def test_inconsistent_hook_shape_disables_forecast() -> None: runtime.time_coord_for_step(5), total_steps, ) - runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 8, 4)) - predicted = runtime.predict_feature(run_id, 5, expected_shape=(2, 8, 4)) + call_id = runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 9, 4)) + predicted = runtime.predict_feature(run_id, 5, expected_shape=(1, 9, 4), call_id=call_id) assert predicted is None assert runtime.stats.forecast_disabled is True assert runtime.stats.disable_reason == "predicted feature shape did not match the current solver-step input" - runtime.observe_actual_feature(run_id, 5, torch.randn(1, 8, 4)) + runtime.observe_actual_feature(run_id, 5, torch.randn(1, 9, 4), call_id=call_id) runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False) runtime.end_run(run_id) -def test_multiple_hook_calls_disable_forecast() -> None: +def test_multiple_hook_calls_are_aggregated_on_actual_steps() -> 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) @@ -216,8 +255,10 @@ def test_multiple_hook_calls_disable_forecast() -> None: runtime.time_coord_for_step(step_id), total_steps, ) - runtime.register_model_hook_call(run_id, step_id, expected_shape=(1, 8, 4)) - runtime.observe_actual_feature(run_id, step_id, torch.randn(1, 8, 4)) + first_id = runtime.register_model_hook_call(run_id, step_id, expected_shape=(1, 8, 4)) + runtime.observe_actual_feature(run_id, step_id, torch.randn(1, 8, 4), call_id=first_id) + second_id = runtime.register_model_hook_call(run_id, step_id, expected_shape=(1, 8, 4)) + runtime.observe_actual_feature(run_id, step_id, torch.randn(1, 8, 4), call_id=second_id) runtime.finalize_solver_step(run_id, step_id, used_forecast=False) decision = runtime.begin_solver_step( @@ -226,11 +267,36 @@ def test_multiple_hook_calls_disable_forecast() -> None: runtime.time_coord_for_step(5), total_steps, ) - runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 8, 4)) - runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 8, 4)) + call_id = runtime.register_model_hook_call(run_id, 5, expected_shape=(2, 8, 4)) + predicted = runtime.predict_feature(run_id, 5, expected_shape=(2, 8, 4), call_id=call_id) + assert predicted is not None + assert predicted.shape == (2, 8, 4) + runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=True) + runtime.end_run(run_id) + + +def test_multicall_branch_signature_change_disables_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, + ) + first_id = runtime.register_model_hook_call( + run_id, 0, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0, 1)),) + ) + runtime.observe_actual_feature(run_id, 0, torch.randn(1, 8, 4), call_id=first_id) + second_id = runtime.register_model_hook_call( + run_id, 0, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1, 0)),) + ) assert runtime.stats.forecast_disabled is True - assert runtime.stats.disable_reason == "multiple model-hook calls observed within one solver step" - runtime.observe_actual_feature(run_id, 5, torch.randn(1, 8, 4)) + assert runtime.stats.disable_reason == "model-hook branch signature changed within one solver step" + runtime.observe_actual_feature(run_id, 0, torch.randn(1, 8, 4), call_id=second_id) runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False) runtime.end_run(run_id) @@ -339,8 +405,10 @@ def main() -> None: test_forecast_fallback_reconciles_bookkeeping() test_observe_actual_feature_clears_forecast_latch() test_unsupported_sampler_disables_forecast() - test_inconsistent_hook_shape_disables_forecast() - test_multiple_hook_calls_disable_forecast() + 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_multicall_branch_signature_change_disables_forecast() test_nonuniform_schedule_coords_are_used() test_forecaster_respects_nonuniform_coords() test_flux_sampler_contract_only_allows_euler()