Handle split FLUX calls within a solver step

This commit is contained in:
xmarre
2026-03-30 15:26:25 +02:00
parent 2a203487c1
commit 70f21e181f
5 changed files with 195 additions and 56 deletions
+3
View File
@@ -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:
+4 -2
View File
@@ -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:
+4
View File
@@ -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
+103 -41
View File
@@ -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.")
+81 -13
View File
@@ -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()