Enable more single-eval samplers for Spectrum forecasting

This commit is contained in:
xmarre
2026-04-04 06:57:54 +02:00
parent bd99d473e0
commit 151e6719c8
2 changed files with 46 additions and 6 deletions
+11 -1
View File
@@ -9,7 +9,17 @@ from .config import SpectrumConfig
from .runtime import SpectrumRuntime
LOG = logging.getLogger(__name__)
_SUPPORTED_SINGLE_EVAL_SAMPLERS = frozenset({"sample_euler", "sample_euler_flow", "euler_flow"})
_SUPPORTED_SINGLE_EVAL_SAMPLERS = frozenset(
{
"sample_euler",
"sample_euler_ancestral",
"sample_euler_flow",
"sample_lcm",
"sample_dpmpp_2m_sde",
"sample_dpmpp_3m_sde",
"euler_flow",
}
)
def _clone_model(model: Any) -> Any:
+35 -5
View File
@@ -37,21 +37,41 @@ def make_runtime(**overrides) -> SpectrumRuntime:
return SpectrumRuntime(cfg)
def test_supported_single_eval_sampler_names_include_euler_flow() -> None:
def test_supported_single_eval_sampler_names_include_supported_variants() -> None:
assert "sample_euler" in _SUPPORTED_SINGLE_EVAL_SAMPLERS
assert "sample_euler_ancestral" in _SUPPORTED_SINGLE_EVAL_SAMPLERS
assert "sample_euler_flow" in _SUPPORTED_SINGLE_EVAL_SAMPLERS
assert "sample_lcm" in _SUPPORTED_SINGLE_EVAL_SAMPLERS
assert "sample_dpmpp_2m_sde" in _SUPPORTED_SINGLE_EVAL_SAMPLERS
assert "sample_dpmpp_3m_sde" in _SUPPORTED_SINGLE_EVAL_SAMPLERS
assert "euler_flow" in _SUPPORTED_SINGLE_EVAL_SAMPLERS
class _FunctionSampler:
def __init__(self, fn):
self.sampler_function = fn
def sample_euler_ancestral(*args, **kwargs):
raise NotImplementedError
def sample_euler_flow(*args, **kwargs):
raise NotImplementedError
def sample_lcm(*args, **kwargs):
raise NotImplementedError
def sample_dpmpp_2m_sde(*args, **kwargs):
raise NotImplementedError
def sample_dpmpp_3m_sde(*args, **kwargs):
raise NotImplementedError
EulerFlowSampler = type("euler_flow", (), {})
assert _supports_solver_step_tracking(_FunctionSampler(sample_euler_ancestral)) is True
assert _supports_solver_step_tracking(_FunctionSampler(sample_euler_flow)) is True
assert _supports_solver_step_tracking(_FunctionSampler(sample_lcm)) is True
assert _supports_solver_step_tracking(_FunctionSampler(sample_dpmpp_2m_sde)) is True
assert _supports_solver_step_tracking(_FunctionSampler(sample_dpmpp_3m_sde)) is True
assert _supports_solver_step_tracking(EulerFlowSampler()) is True
@@ -612,8 +632,18 @@ def test_forecaster_respects_nonuniform_coords() -> None:
assert torch.allclose(pred, torch.tensor([1.0]), atol=1e-5)
def test_flux_sampler_contract_allows_euler_and_euler_flow() -> None:
assert _SUPPORTED_SINGLE_EVAL_SAMPLERS == frozenset({"sample_euler", "sample_euler_flow", "euler_flow"})
def test_flux_sampler_contract_allows_supported_single_eval_variants() -> None:
assert _SUPPORTED_SINGLE_EVAL_SAMPLERS == frozenset(
{
"sample_euler",
"sample_euler_ancestral",
"sample_euler_flow",
"sample_lcm",
"sample_dpmpp_2m_sde",
"sample_dpmpp_3m_sde",
"euler_flow",
}
)
def test_tail_actual_steps_force_real_forwards() -> None:
@@ -689,7 +719,7 @@ def test_forecast_feature_sanitization_stats_only_report_real_violations() -> No
def main() -> None:
test_supported_single_eval_sampler_names_include_euler_flow()
test_supported_single_eval_sampler_names_include_supported_variants()
test_forecaster_recomputes_coeff_on_update_not_predict()
test_solver_step_scheduler()
test_forecast_fallback_reconciles_bookkeeping()
@@ -708,7 +738,7 @@ def main() -> None:
test_aborted_solver_step_is_discarded_without_disabling_forecast()
test_nonuniform_schedule_coords_are_used()
test_forecaster_respects_nonuniform_coords()
test_flux_sampler_contract_allows_euler_and_euler_flow()
test_flux_sampler_contract_allows_supported_single_eval_variants()
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()