Merge pull request #2 from xmarre/codex/diagnose-approximation-artifact-bias
Add tail_actual_steps option to force final real forwards
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user