diff --git a/README.md b/README.md index 6e97822..fe05769 100644 --- a/README.md +++ b/README.md @@ -125,6 +125,14 @@ Number of initial real forwards before forecasting is allowed. Recommended default: `5` +### `tail_actual_steps` + +Number of final solver steps forced to stay on the real path. + +Practical default for quality-sensitive runs: `3` + +This protects the refinement tail where late-step forecast bias tends to show up first as smoother microdetail, especially in edit-conditioned runs. + ### `max_history` Cap for cached real-forward feature points used for the fit. @@ -145,6 +153,7 @@ Enables lightweight logging during patch install and a per-run summary of actual - `window_size = 2.0` - `flex_window = 0.75` - `warmup_steps = 5` +- `tail_actual_steps = 3` ### More aggressive @@ -154,6 +163,7 @@ Enables lightweight logging during patch install and a per-run summary of actual - `window_size = 2.0` - `flex_window = 3.0` - `warmup_steps = 5` +- `tail_actual_steps = 3` ## Design notes @@ -187,6 +197,7 @@ This repo normalizes the Chebyshev basis against the detected schedule length fr - It depends on current ComfyUI FLUX internals staying broadly compatible with the present `forward_orig` signature. - It is designed to coexist with standard transformer patch chains, but it is **not guaranteed** to compose with other custom nodes that also replace FLUX `forward_orig` directly. - The scheduler is faithful to the official adaptive-window strategy, but one safety approximation is added: forecasting is held back until enough real points exist to fit the chosen Chebyshev degree. +- The last few refinement steps can also be reserved as actual-only with `tail_actual_steps` to reduce late-step texture loss. - No claims are made here about exact paper speedups inside arbitrary ComfyUI workflows. Sampler choice, guidance path, ControlNet usage, resolution, and other wrappers all affect real wall-clock results. ## Validation / smoke test diff --git a/comfyui_spectrum/config.py b/comfyui_spectrum/config.py index 48d5610..c64f9cb 100644 --- a/comfyui_spectrum/config.py +++ b/comfyui_spectrum/config.py @@ -13,6 +13,7 @@ class SpectrumConfig: window_size: float = 2.0 flex_window: float = 0.75 warmup_steps: int = 5 + tail_actual_steps: int = 3 max_history: int = 128 debug: bool = False @@ -31,6 +32,8 @@ class SpectrumConfig: raise ValueError("flex_window must be >= 0.") if int(self.warmup_steps) < 0: raise ValueError("warmup_steps must be >= 0.") + if int(self.tail_actual_steps) < 0: + raise ValueError("tail_actual_steps must be >= 0.") if int(self.max_history) < max(8, int(self.degree) + 1): raise ValueError("max_history must be at least max(8, degree + 1).") return self diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index cdf3b4c..3555161 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -141,6 +141,15 @@ class SpectrumRuntime: raise RuntimeError(f"Spectrum solver step {solver_step_id} is outside the active schedule.") return float(self._active_run.schedule_coords[idx]) + def _is_tail_actual_step(self, solver_step_id: int) -> bool: + if self._active_run is None: + return False + tail_actual_steps = int(self.cfg.tail_actual_steps) + if tail_actual_steps <= 0: + return False + tail_start = max(0, self._active_run.total_steps - tail_actual_steps) + return int(solver_step_id) >= tail_start + def start_run(self, sample_sigmas: torch.Tensor, sampler_name: str, *, supports_solver_steps: bool) -> int: self.run_id += 1 schedule_values, schedule_coords = self._build_schedule_coords(sample_sigmas) @@ -194,15 +203,17 @@ class SpectrumRuntime: self._disable_forecasting("solver-step ids are not sequential within the sampling run") actual_forward = True + tail_actual_only = self._is_tail_actual_step(int(solver_step_id)) if ( not self.forecast_disabled + and not tail_actual_only and int(solver_step_id) >= self.cfg.warmup_steps and self.forecaster.ready(self.min_fit_points) ): ws_floor = max(1, int(math.floor(self.curr_ws))) actual_forward = ((self.num_consecutive_cached_steps + 1) % ws_floor) == 0 - if self.forecast_disabled or not self.forecaster.ready(self.min_fit_points): + if self.forecast_disabled or tail_actual_only or not self.forecaster.ready(self.min_fit_points): actual_forward = True decision = { @@ -272,6 +283,8 @@ class SpectrumRuntime: expected_shape: Optional[tuple[int, ...]] = None, ) -> Optional[torch.Tensor]: step = self._require_active_step(run_id, solver_step_id) + if step.decision["actual_forward"]: + return None if self.forecast_disabled or not self.forecaster.ready(self.min_fit_points): return None diff --git a/nodes.py b/nodes.py index 0fcd7d9..f0cdaa6 100644 --- a/nodes.py +++ b/nodes.py @@ -16,6 +16,7 @@ class SpectrumApplyFlux: "window_size": ("FLOAT", {"default": 2.0, "min": 1.0, "max": 16.0, "step": 0.05}), "flex_window": ("FLOAT", {"default": 0.75, "min": 0.0, "max": 8.0, "step": 0.05}), "warmup_steps": ("INT", {"default": 5, "min": 0, "max": 32, "step": 1}), + "tail_actual_steps": ("INT", {"default": 3, "min": 0, "max": 32, "step": 1}), "max_history": ("INT", {"default": 128, "min": 8, "max": 512, "step": 1}), "debug": ("BOOLEAN", {"default": False}), } @@ -36,6 +37,7 @@ class SpectrumApplyFlux: window_size, flex_window, warmup_steps, + tail_actual_steps, max_history, debug, ): @@ -51,6 +53,7 @@ class SpectrumApplyFlux: window_size=float(window_size), flex_window=float(flex_window), warmup_steps=int(warmup_steps), + tail_actual_steps=int(tail_actual_steps), max_history=int(max_history), debug=bool(debug), ).validate() diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index 0abebfd..e78fdc8 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -20,16 +20,19 @@ from comfyui_spectrum.flux import ( from comfyui_spectrum.runtime import SpectrumRuntime -def make_runtime() -> SpectrumRuntime: - cfg = SpectrumConfig( - blend_weight=0.5, - degree=4, - ridge_lambda=0.1, - window_size=2.0, - flex_window=0.75, - warmup_steps=5, - max_history=128, - ).validate() +def make_runtime(**overrides) -> SpectrumRuntime: + cfg_kwargs = { + "blend_weight": 0.5, + "degree": 4, + "ridge_lambda": 0.1, + "window_size": 2.0, + "flex_window": 0.75, + "warmup_steps": 5, + "tail_actual_steps": 3, + "max_history": 128, + } + cfg_kwargs.update(overrides) + cfg = SpectrumConfig(**cfg_kwargs).validate() return SpectrumRuntime(cfg) @@ -259,6 +262,54 @@ def test_flux_sampler_contract_only_allows_euler() -> None: assert _SUPPORTED_SINGLE_EVAL_SAMPLERS == frozenset({"sample_euler"}) +def test_tail_actual_steps_force_real_forwards() -> None: + runtime = make_runtime( + degree=1, + ridge_lambda=0.1, + window_size=2.0, + flex_window=0.75, + warmup_steps=2, + tail_actual_steps=3, + max_history=16, + ) + sample_sigmas = torch.linspace(1.0, 0.0, 9) + run_id = runtime.start_run(sample_sigmas, "sample_euler", supports_solver_steps=True) + total_steps = len(sample_sigmas) - 1 + + for step_id in range(total_steps): + 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=(1, 8, 4)) + + if step_id >= total_steps - 3: + assert decision["actual_forward"] is True + assert runtime.predict_feature(run_id, step_id, expected_shape=(1, 8, 4)) is None + runtime.observe_actual_feature(run_id, step_id, torch.randn(1, 8, 4)) + runtime.finalize_solver_step(run_id, step_id, used_forecast=False) + continue + + if step_id < 2: + assert decision["actual_forward"] is True + runtime.observe_actual_feature(run_id, step_id, torch.randn(1, 8, 4)) + runtime.finalize_solver_step(run_id, step_id, used_forecast=False) + continue + + if decision["actual_forward"]: + runtime.observe_actual_feature(run_id, step_id, torch.randn(1, 8, 4)) + runtime.finalize_solver_step(run_id, step_id, used_forecast=False) + else: + assert runtime.predict_feature(run_id, step_id, expected_shape=(1, 8, 4)) is not None + runtime.finalize_solver_step(run_id, step_id, used_forecast=True) + + assert runtime.stats.forecasted_count > 0 + assert runtime.stats.actual_forward_count == 6 + runtime.end_run(run_id) + + def test_forecast_feature_is_sanitized_before_fp16_final_layer() -> None: feature = torch.tensor([float("nan"), float("inf"), float("-inf"), 70000.0, -70000.0, 123.5]) sanitized = _sanitize_forecast_feature_for_final_layer(feature, torch.float16) @@ -293,6 +344,7 @@ def main() -> None: test_nonuniform_schedule_coords_are_used() test_forecaster_respects_nonuniform_coords() test_flux_sampler_contract_only_allows_euler() + test_tail_actual_steps_force_real_forwards() test_forecast_feature_is_sanitized_before_fp16_final_layer() test_forecast_feature_sanitization_stats_only_report_real_violations() print("ok")