Diagnose Spectrum approximation bias

This commit is contained in:
xmarre
2026-03-18 11:03:34 +01:00
parent 68e3d81d0e
commit 6269091889
5 changed files with 93 additions and 11 deletions
+11
View File
@@ -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
+3
View File
@@ -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
+14 -1
View File
@@ -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
+3
View File
@@ -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()
+62 -10
View File
@@ -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")