diff --git a/python/sglang/multimodal_gen/configs/models/dits/stablediffusion3.py b/python/sglang/multimodal_gen/configs/models/dits/stablediffusion3.py index cbe3141b7..feb32ebec 100644 --- a/python/sglang/multimodal_gen/configs/models/dits/stablediffusion3.py +++ b/python/sglang/multimodal_gen/configs/models/dits/stablediffusion3.py @@ -35,3 +35,4 @@ class StableDiffusion3TransformerConfig(DiTConfig): arch_config: StableDiffusion3TransformerArchConfig = field( default_factory=StableDiffusion3TransformerArchConfig ) + prefix: str = "SD3" diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index c5636c507..7390f8f62 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -207,6 +207,10 @@ class SamplingParams: None # TeaCacheParams or WanTeaCacheParams, set by model-specific subclass ) + # Spectrum parameters + enable_spectrum: bool = False + spectrum_params: Any = None # SpectrumParams + # Profiling profile: bool = field(default=False, metadata={"batch_sig_exclude": True}) num_profiled_timesteps: int = field(default=5, metadata={"batch_sig_exclude": True}) @@ -330,6 +334,16 @@ class SamplingParams: if env_steps is not None and self.num_inference_steps is not None: self.num_inference_steps = int(env_steps) + if self.enable_spectrum and isinstance(self.spectrum_params, dict): + from sglang.multimodal_gen.configs.sample.spectrum import SpectrumParams + + self.spectrum_params = SpectrumParams(**self.spectrum_params) + + if self.enable_spectrum and self.spectrum_params is None: + from sglang.multimodal_gen.configs.sample.spectrum import SpectrumParams + + self.spectrum_params = SpectrumParams() + def build_request_extra(self) -> dict[str, Any]: """Return optional request-scoped extras for downstream pipeline stages.""" extra = {} @@ -547,6 +561,11 @@ class SamplingParams: f"boundary_ratio must be within [0, 1], got {self.boundary_ratio!r}" ) + if self.enable_teacache and self.enable_spectrum: + raise ValueError( + "enable_teacache and enable_spectrum are mutually exclusive; enable only one." + ) + RLRolloutArgs.validate_sampling_params(self) def check_sampling_param(self): @@ -869,6 +888,69 @@ class SamplingParams: "--enable-teacache", action="store_true", ) + add_argument( + "--enable-spectrum", + action="store_true", + ) + add_argument("--w", type=float) + add_argument( + "--taylor-order", + "--taylor_order", + dest="taylor_order", + type=int, + ) + add_argument( + "--history-size", + "--history_size", + dest="history_size", + type=int, + ) + add_argument( + "--spectrum-window-size", + "--spectrum_window_size", + "--window-size", + "--window_size", + dest="spectrum_window_size", + type=float, + help="Spectrum initial skip window size.", + ) + add_argument( + "--spectrum-flex-window", + "--spectrum_flex_window", + "--flex-window", + "--flex_window", + dest="spectrum_flex_window", + type=float, + help="Spectrum adaptive window growth slope.", + ) + add_argument( + "--spectrum-warmup-steps", + "--spectrum_warmup_steps", + dest="spectrum_warmup_steps", + type=int, + help="Spectrum warmup denoising steps before caching.", + ) + add_argument( + "--spectrum-m", + "--spectrum_m", + dest="spectrum_m", + type=int, + help="Spectrum Chebyshev polynomial degree (M).", + ) + add_argument( + "--spectrum-lam", + "--spectrum_lam", + dest="spectrum_lam", + type=float, + help="Spectrum ridge regularization strength.", + ) + add_argument( + "--spectrum-tau-num-steps", + "--spectrum_tau_num_steps", + dest="spectrum_tau_num_steps", + type=int, + help="Spectrum tau normalization horizon.", + ) # profiling add_argument( @@ -1299,6 +1381,34 @@ class SamplingParams: } if isinstance(cli_args.get("seed"), list) and len(cli_args["seed"]) == 1: cli_args["seed"] = cli_args["seed"][0] + + spectrum_overrides = {} + spectrum_flag_map = { + "w": "w", + "taylor_order": "taylor_order", + "window_size": "spectrum_window_size", + "flex_window": "spectrum_flex_window", + "history_size": "history_size", + "warmup_steps": "spectrum_warmup_steps", + "m": "spectrum_m", + "lam": "spectrum_lam", + "tau_num_steps": "spectrum_tau_num_steps", + } + for field_name, arg_name in spectrum_flag_map.items(): + if hasattr(args, arg_name) and getattr(args, arg_name) is not None: + spectrum_overrides[field_name] = getattr(args, arg_name) + if spectrum_overrides: + if not cli_args.get("enable_spectrum", False): + logger.info( + "Spectrum override flags were provided without --enable-spectrum; " + "auto-enabling Spectrum caching." + ) + cli_args["enable_spectrum"] = True + existing_spectrum_params = cli_args.get("spectrum_params") + if isinstance(existing_spectrum_params, dict): + spectrum_overrides = {**existing_spectrum_params, **spectrum_overrides} + cli_args["spectrum_params"] = spectrum_overrides + return cli_args def output_file_path(self): diff --git a/python/sglang/multimodal_gen/configs/sample/spectrum.py b/python/sglang/multimodal_gen/configs/sample/spectrum.py new file mode 100644 index 000000000..b3cf0d7ed --- /dev/null +++ b/python/sglang/multimodal_gen/configs/sample/spectrum.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +from __future__ import annotations + +from dataclasses import dataclass + +from sglang.multimodal_gen.configs.sample.sampling_params import CacheParams + + +@dataclass +class SpectrumParams(CacheParams): + """ + Parameters for [Spectrum](https://arxiv.org/abs/2603.01623) spectral feature + forecasting. + + Attributes: + cache_type: (`str`, defaults to `spectrum`): + A string labeling these parameters as belonging to spectrum. + window_size (`float`, defaults to `2.0`): + Initial skip window size (paper: N). After warmup, a real forward is + run every `floor(window_size)` consecutive cached steps; when a real + forward runs, the window grows by `flex_window`. Recommended: `2.0`. + flex_window (`float`, defaults to `0.75`): + Adaptive window slope (paper: alpha). Controls how quickly skip intervals + widen as sampling progresses. Recommended: `0.75` for ~3.5x speedup, + `3.0` for ~5x speedup (fewer network evaluations, lower quality). + warmup_steps (`int`, defaults to `5`): + Number of initial forward passes that always run the full DiT before + Spectrum caching begins. Early steps define global structure and are + kept exact in the paper setup. + w (`float`, defaults to `1.0`): + Blend weight between a local Taylor predictor and the Chebyshev + spectral predictor (``(1 - w) * h_taylor + w * h_cheb``). This + mix is not in the paper; ``w=1.0`` is pure Chebyshev and matches + the [reference repo](https://github.com/hanjq17/Spectrum#-text-to-image-t2i) + default. They recommend ``0.5``–``1.0`` for extra robustness. + lam (`float`, defaults to `0.1`): + Ridge regularization strength (lambda) when fitting Chebyshev + coefficients over recent cached features. + m (`int`, defaults to `4`): + Number of Chebyshev bases (polynomial degree M). The design matrix has + `M + 1` columns. + history_size (`int`, defaults to `100`): + Maximum number of recent (step, feature) pairs kept for coefficient + fitting (paper reference code uses a sliding window of up to 100). + tau_num_steps (`int`, defaults to `50`): + Denominator for mapping denoising step indices to Chebyshev time + ``tau in [-1, 1]``. The reference repo hardcodes 50 regardless of + ``num_inference_steps``; keep this at 50 unless you know you need + another horizon. + taylor_order (`int`, defaults to `1`): + Order of the local discrete Taylor fallback used in the blend when + `w < 1.0`. Supported values: `1`, `2`, or `3`. + separate_cfg_branches (not a field — determined by the DiT model): + Wan, Hunyuan, and SD3 maintain independent Spectrum state per CFG + branch. Other ``CachableDiT`` models use one counter for all + forwards; see ``get_total_forward_steps()``. + """ + + cache_type: str = "spectrum" + window_size: float = 2.0 + flex_window: float = 0.75 + warmup_steps: int = 5 + w: float = 1.0 + lam: float = 0.1 + m: int = 4 + history_size: int = 100 + tau_num_steps: int = 50 + taylor_order: int = 1 + + def get_total_forward_steps( + self, num_inference_steps: int, do_cfg: bool, separate_cfg_branches: bool + ) -> int: + """How many DiT forward calls one Spectrum counter sees per generation. + + Used by (1) ``ChebyshevForecaster`` to map step indices onto [-1, 1] and + by (2) ``begin_spectrum_step`` to wrap branch counters at end-of-run. + Not used for logging. + + - **Separate counters** (Wan, Hunyuan, SD3): each CFG branch has its own + ``spectrum_cnt`` / forecaster. Every denoising step triggers one cond + forward and one uncond forward, but each counter only advances on its + branch → ``num_inference_steps`` calls per counter. + + - **Single counter** (FLUX, …): one counter interleaves cond and + uncond forwards when CFG is enabled → ``2 * num_inference_steps`` calls. + FLUX.1-dev uses embedded guidance (no true CFG), so this path is + normally single-branch in practice. + + When CFG is off, every model performs one forward per denoising step. + """ + if do_cfg and not separate_cfg_branches: + return num_inference_steps * 2 + return num_inference_steps diff --git a/python/sglang/multimodal_gen/runtime/cache/__init__.py b/python/sglang/multimodal_gen/runtime/cache/__init__.py index 62f0f8457..eb95a42fb 100644 --- a/python/sglang/multimodal_gen/runtime/cache/__init__.py +++ b/python/sglang/multimodal_gen/runtime/cache/__init__.py @@ -6,6 +6,7 @@ This module provides various caching strategies to accelerate diffusion transformer (DiT) inference: - TeaCache: Temporal similarity-based caching for diffusion models +- Spectrum: Chebyshev spectral feature forecasting for step skipping - cache-dit integration: Block-level caching with DBCache and TaylorSeer """ @@ -16,12 +17,15 @@ from sglang.multimodal_gen.runtime.cache.cache_dit_integration import ( enable_cache_on_transformer, get_scm_mask, ) +from sglang.multimodal_gen.runtime.cache.spectrum import SpectrumMixin from sglang.multimodal_gen.runtime.cache.teacache import TeaCacheContext, TeaCacheMixin __all__ = [ # TeaCache (always available) "TeaCacheContext", "TeaCacheMixin", + # Spectrum (always available) + "SpectrumMixin", # cache-dit integration (lazy-loaded, requires cache-dit package) "CacheDitConfig", "enable_cache_on_transformer", diff --git a/python/sglang/multimodal_gen/runtime/cache/spectrum.py b/python/sglang/multimodal_gen/runtime/cache/spectrum.py new file mode 100644 index 000000000..abc03f84f --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/cache/spectrum.py @@ -0,0 +1,645 @@ +# SPDX-License-Identifier: Apache-2.0 +""" +Spectrum: Adaptive Spectral Feature Forecasting for diffusion sampling acceleration. + +Training-free step skipping with Chebyshev polynomial ridge regression over +denoiser block outputs. See https://arxiv.org/abs/2603.01623 +""" + +from __future__ import annotations + +import logging +import math +from dataclasses import dataclass +from typing import TYPE_CHECKING, Optional + +import torch +import torch.nn as nn + +if TYPE_CHECKING: + from sglang.multimodal_gen.configs.sample.spectrum import SpectrumParams + + +logger = logging.getLogger(__name__) + + +def _flatten(x: torch.Tensor) -> tuple[torch.Tensor, torch.Size]: + """Reshape tensor to (1, -1) for ridge regression, preserving original shape.""" + return x.reshape(1, -1), x.shape + + +def _unflatten(x_flat: torch.Tensor, shape: torch.Size) -> torch.Tensor: + """Restore flattened tensor to original shape.""" + return x_flat.reshape(shape) + + +def _is_missing_linalg_backend_error(err: RuntimeError) -> bool: + """Detect the "no LAPACK / no MAGMA" RuntimeError PyTorch raises when the + build lacks a linear algebra backend for the tensor's device. + + Some PyTorch builds (notably several ROCm wheels) ship without LAPACK for + CPU tensors *and* without MAGMA for CUDA/HIP tensors, making + ``torch.linalg.cholesky``/``torch.cholesky_solve`` unusable on any device. + """ + msg = str(err) + return "LAPACK" in msg or "MAGMA" in msg + + +def _cholesky_lower(a: torch.Tensor) -> torch.Tensor: + """Cholesky factorization ``A = L @ L.T`` without ``torch.linalg``. + + Used as a fallback when the backend has no LAPACK/MAGMA support. ``a`` is + expected to be tiny (Chebyshev design matrix size), so the Python-level + loop is cheap. + """ + n = a.shape[0] + L = torch.zeros_like(a) + for i in range(n): + for j in range(i + 1): + s = L[i, :j] @ L[j, :j] + if i == j: + L[i, j] = torch.sqrt(a[i, i] - s) + else: + L[i, j] = (a[i, j] - s) / L[j, j] + return L + + +def _cholesky_solve_lower(l: torch.Tensor, b: torch.Tensor) -> torch.Tensor: + """Solve ``(L @ L.T) @ x = b`` given lower-triangular Cholesky factor ``l``.""" + n = l.shape[0] + y = torch.zeros_like(b) + for i in range(n): + y[i] = (b[i] - l[i, :i] @ y[:i]) / l[i, i] + x = torch.zeros_like(b) + for i in reversed(range(n)): + x[i] = (y[i] - l[i + 1 :, i] @ x[i + 1 :]) / l[i, i] + return x + + +def _ridge_cholesky_solve(xtx: torch.Tensor, xth: torch.Tensor) -> torch.Tensor: + """Solve the ridge normal equations ``xtx @ coef = xth`` via Cholesky. + + Prefers ``torch.linalg.cholesky``/``torch.cholesky_solve``, retrying once + with jitter if ``xtx`` is only marginally non-positive-definite. Falls + back to a manual, LAPACK/MAGMA-free Cholesky solve when the backend has + no linear algebra support at all for this device (see + `_is_missing_linalg_backend_error`). + """ + p = xtx.shape[0] + eye = torch.eye(p, device=xtx.device, dtype=xtx.dtype) + try: + chol = torch.linalg.cholesky(xtx) + return torch.cholesky_solve(xth, chol) + except RuntimeError as e: + if _is_missing_linalg_backend_error(e): + chol = _cholesky_lower(xtx) + if not torch.isfinite(chol).all(): + jitter = 1e-6 * xtx.diag().mean() + chol = _cholesky_lower(xtx + jitter * eye) + return _cholesky_solve_lower(chol, xth) + jitter = 1e-6 * xtx.diag().mean() + chol = torch.linalg.cholesky(xtx + jitter * eye) + return torch.cholesky_solve(xth, chol) + + +class ChebyshevForecaster(nn.Module): + """Chebyshev-basis ridge regression forecaster over diffusion step index. + + Keeps a sliding window of (step, flattened_feature) pairs from recent *real* + DiT forwards. On predict(), fits ridge regression coefficients for Chebyshev + bases T_0..T_M on normalized time, then evaluates at the requested step. + """ + + def __init__( + self, + M: int = 4, + K: int = 100, + lam: float = 0.1, + num_steps: int = 50, + device: Optional[torch.device] = None, + feature_shape: Optional[torch.Size] = None, + ) -> None: + super().__init__() + self.M = M + self.K = K + self.lam = lam + self.num_steps = num_steps + # Preallocated (K,) / (K, N) ring buffer, written in place by update(). + self.register_buffer("t_buf", torch.empty(0)) + self._H_buf: Optional[torch.Tensor] = None + self._write_idx = 0 + self._count = 0 + self._shape: Optional[torch.Size] = feature_shape + self._feature_dtype: Optional[torch.dtype] = None + self._coef: Optional[torch.Tensor] = None + self.device_ref = device + self._tau_scale = 2.0 / float(self.num_steps) if self.num_steps != 0 else 0.0 + + @property + def P(self) -> int: + return self.M + 1 + + def _taus(self, t: torch.Tensor) -> torch.Tensor: + """Normalize timesteps to [-1, 1] range for Chebyshev basis.""" + if self.num_steps == 0: + return torch.zeros_like(t) + return t * self._tau_scale - 1.0 + + def _build_design(self, taus: torch.Tensor) -> torch.Tensor: + """Build Chebyshev basis design matrix [T_0(tau), T_1(tau), ..., T_M(tau)].""" + taus = taus.reshape(-1, 1) + k = taus.shape[0] + t0 = torch.ones((k, 1), device=taus.device, dtype=taus.dtype) + if self.M == 0: + return t0 + t1 = taus + cols = [t0, t1] + for _ in range(2, self.M + 1): + cols.append(2 * taus * cols[-1] - cols[-2]) + return torch.cat(cols[: self.M + 1], dim=1) + + def update(self, t: float | torch.Tensor, h: torch.Tensor) -> None: + device = self.device_ref or h.device + t_tensor = torch.as_tensor(t, dtype=torch.float32, device=device) + h_flat, shape = _flatten(h) + if self._feature_dtype is None: + self._feature_dtype = h_flat.dtype + h_flat = h_flat.to(device).to(torch.float32) + if self._shape is None: + self._shape = shape + else: + assert shape == self._shape, "Spectrum feature shape must remain constant" + + if self._H_buf is None: + # Preallocate the ring buffer once the flattened feature size is known + self.t_buf = torch.zeros(self.K, dtype=torch.float32, device=device) + self._H_buf = torch.zeros( + self.K, h_flat.shape[1], dtype=torch.float32, device=device + ) + + self.t_buf[self._write_idx] = t_tensor + self._H_buf[self._write_idx] = h_flat[0] + self._write_idx = (self._write_idx + 1) % self.K + self._count = min(self._count + 1, self.K) + + # Invalidate cached ridge coefficients; will be refit on next predict() + self._coef = None + + def _recent(self, n: int) -> tuple[torch.Tensor, torch.Tensor]: + """Return the last `min(n, count)` valid (t, h) entries in chronological + order (oldest first, most recent last). + + Only used by the discrete Taylor blend, which needs true temporal + order for its finite differences (unlike the ridge fit, see + `_fit_if_needed`). `n` is small (<= taylor_order + 1), so even the + rare wraparound `cat` below is cheap. + """ + count = min(self._count, n) + start = (self._write_idx - count) % self.K + if start + count <= self.K: + # Contiguous window -- cheap view, no copy. + return self.t_buf[start : start + count], self._H_buf[start : start + count] + # Wraps past the end of the physical buffer -- must reorder. + tail = self.K - start + t = torch.cat([self.t_buf[start:], self.t_buf[: count - tail]]) + h = torch.cat([self._H_buf[start:], self._H_buf[: count - tail]]) + return t, h + + def ready(self) -> bool: + return self._count >= 1 + + def _fit_if_needed(self) -> None: + """Fit ridge regression coefficients on cached (t, h) pairs if not cached.""" + if self._coef is not None: + return + assert self.ready() + assert self._H_buf is not None + assert self._feature_dtype is not None + feature_dtype = self._feature_dtype + t, h = self.t_buf[: self._count], self._H_buf[: self._count] + taus = self._taus(t) + # Ridge solve in fp32; autocast would keep matmuls in bf16 and break + # torch.cholesky_solve dtype requirements. + with torch.autocast(device_type=self._H_buf.device.type, enabled=False): + x = self._build_design(taus).to(torch.float32) + p = x.shape[1] + lam_i = self.lam * torch.eye(p, device=x.device, dtype=x.dtype) + xt = x.transpose(0, 1) + xtx = xt @ x + lam_i + xth = xt @ h + self._coef = _ridge_cholesky_solve(xtx, xth).to(feature_dtype) + + @torch.no_grad() + def predict(self, t_star: float | torch.Tensor) -> torch.Tensor: + assert self._shape is not None + device = self.t_buf.device + t_star = torch.as_tensor(t_star, dtype=torch.float32, device=device) + self._fit_if_needed() + assert self._coef is not None + tau_star = self._taus(t_star) + x_star = self._build_design(tau_star[None]).to(self._coef.dtype) + h_flat = x_star @ self._coef + return _unflatten(h_flat, self._shape) + + +class SpectrumForecaster(nn.Module): + """Chebyshev + discrete Taylor blend forecaster. + + The paper uses pure Chebyshev; the reference repo blends with a local + discrete Taylor predictor. ``w=1`` is Chebyshev-only; lower ``w`` adds Taylor. + """ + + def __init__( + self, + cheb: ChebyshevForecaster, + *, + taylor_order: int = 1, + w: float = 1.0, + ) -> None: + super().__init__() + self.cheb = cheb + self.taylor_order = taylor_order + self.w = w + + @torch.no_grad() + def _local_taylor_discrete(self, t_star: torch.Tensor) -> torch.Tensor: + """Predict hidden state at t_star using discrete Taylor expansion from recent real steps.""" + assert self.cheb._H_buf is not None + assert self.cheb._shape is not None + assert self.cheb._feature_dtype is not None + feature_dtype = self.cheb._feature_dtype + t, h = self.cheb._recent(self.taylor_order + 1) + h_i = h[-1] + if t.numel() < 2: + return _unflatten(h_i.reshape(1, -1), self.cheb._shape).to(feature_dtype) + h_im1 = h[-2] + t_i = t[-1] + t_im1 = t[-2] + dh1 = h_i - h_im1 + dt_last = (t_i - t_im1).clamp_min(1e-8) + k = ((t_star - t_i) / dt_last).to(h_i.dtype) + out = h_i + k * dh1 + if self.taylor_order >= 2 and t.numel() >= 3: + h_im2 = h[-3] + d2 = h_i - 2 * h_im1 + h_im2 + out = out + 0.5 * k * (k - 1.0) * d2 + if self.taylor_order >= 3 and t.numel() >= 4: + h_im3 = h[-4] + d3 = h_i - 3 * h_im1 + 3 * h_im2 - h_im3 + out = out + (k * (k - 1.0) * (k - 2.0) / 6.0) * d3 + return _unflatten(out.reshape(1, -1), self.cheb._shape).to(feature_dtype) + + @torch.no_grad() + def predict(self, t_star: float | torch.Tensor) -> torch.Tensor: + """Blend Chebyshev regression and local Taylor predictions.""" + device = self.cheb.t_buf.device + t_star = torch.as_tensor(t_star, dtype=torch.float32, device=device) + if self.w >= 1.0: + return self.cheb.predict(t_star) + elif self.w <= 0.0: + return self._local_taylor_discrete(t_star) + return torch.lerp( + self._local_taylor_discrete(t_star), self.cheb.predict(t_star), self.w + ) + + def update(self, t, h) -> None: + self.cheb.update(t, h) + + def ready(self) -> bool: + return self.cheb.ready() + + +@dataclass +class SpectrumContext: + current_step: int + num_inference_steps: int + total_forward_steps: int + do_cfg: bool + is_cfg_negative: bool + spectrum_params: SpectrumParams + debug: bool + + +class SpectrumMixin: + """Mixin providing Spectrum step-skipping and feature forecasting. + + Wired into ``CachableDiT`` (see ``runtime/models/dits/base.py``). Concrete + models call three hooks from ``forward()`` around their transformer blocks: + + - ``begin_spectrum_step()`` — advance skip schedule; returns whether to run + blocks or forecast instead. + - ``spectrum_record_features()`` — after a real forward, store block outputs. + - ``spectrum_predict_features()`` — on skipped steps, return forecasted outputs. + + Models with separate CFG branches (Wan, Hunyuan, SD3) list their config prefix + in ``_CFG_SUPPORTED_PREFIXES`` so cond/uncond maintain independent counters and + forecasters. All other ``CachableDiT`` subclasses share one counter. + """ + + # DiT config prefixes that run separate cond/uncond forwards (see TeaCache). + _CFG_SUPPORTED_PREFIXES: set[str] = {"wan", "hunyuan", "sd3"} + + def _init_spectrum_state(self) -> None: + """Initialize Spectrum state variables. Dual-branch models (Wan, Hunyuan) track separate + cond/uncond forecasters and counters; others share one. + """ + # Positive branch (cond) or single-branch state + self.spectrum_cnt = 0 + self.spectrum_num_consecutive_cached_steps = 0 + self.spectrum_curr_ws: Optional[float] = None + self.spectrum_forecaster: Optional[SpectrumForecaster] = None + self.spectrum_real_steps = 0 + self.spectrum_skipped_steps = 0 + self.spectrum_shadow_rel_l2_sum = 0.0 + self.spectrum_shadow_rel_l2_count = 0 + + # Negative branch (uncond) state for dual-branch CFG models + self.spectrum_cnt_negative = 0 + self.spectrum_num_consecutive_cached_steps_negative = 0 + self.spectrum_curr_ws_negative: Optional[float] = None + self.spectrum_forecaster_negative: Optional[SpectrumForecaster] = None + self.spectrum_real_steps_negative = 0 + self.spectrum_skipped_steps_negative = 0 + self.spectrum_shadow_rel_l2_sum_negative = 0.0 + self.spectrum_shadow_rel_l2_count_negative = 0 + + # Runtime branch tracking + self.spectrum_is_cfg_negative = False + self._spectrum_ctx: Optional[SpectrumContext] = None + prefix = getattr(self.config, "prefix", "").lower() + self._spectrum_supports_cfg_cache = prefix in self._CFG_SUPPORTED_PREFIXES + + def reset_spectrum_state(self, spectrum_params: SpectrumParams) -> None: + self.spectrum_cnt = 0 + self.spectrum_num_consecutive_cached_steps = 0 + self.spectrum_curr_ws = spectrum_params.window_size + self.spectrum_forecaster = None + self.spectrum_is_cfg_negative = False + self.spectrum_real_steps = 0 + self.spectrum_skipped_steps = 0 + self.spectrum_shadow_rel_l2_sum = 0.0 + self.spectrum_shadow_rel_l2_count = 0 + if self._spectrum_supports_cfg_cache: + self.spectrum_cnt_negative = 0 + self.spectrum_num_consecutive_cached_steps_negative = 0 + self.spectrum_curr_ws_negative = spectrum_params.window_size + self.spectrum_forecaster_negative = None + self.spectrum_real_steps_negative = 0 + self.spectrum_skipped_steps_negative = 0 + self.spectrum_shadow_rel_l2_sum_negative = 0.0 + self.spectrum_shadow_rel_l2_count_negative = 0 + + def _get_spectrum_branch_state(self) -> tuple[int, int, float]: + """Get schedule state for current branch (cond or uncond).""" + if self.spectrum_is_cfg_negative and self._spectrum_supports_cfg_cache: + return ( + self.spectrum_cnt_negative, + self.spectrum_num_consecutive_cached_steps_negative, + self.spectrum_curr_ws_negative or 0.0, + ) + return ( + self.spectrum_cnt, + self.spectrum_num_consecutive_cached_steps, + self.spectrum_curr_ws or 0.0, + ) + + def _set_spectrum_branch_state( + self, cnt: int, consecutive: int, curr_ws: float + ) -> None: + """Set schedule state for current branch (cond or uncond).""" + if self.spectrum_is_cfg_negative and self._spectrum_supports_cfg_cache: + self.spectrum_cnt_negative = cnt + self.spectrum_num_consecutive_cached_steps_negative = consecutive + self.spectrum_curr_ws_negative = curr_ws + else: + self.spectrum_cnt = cnt + self.spectrum_num_consecutive_cached_steps = consecutive + self.spectrum_curr_ws = curr_ws + + def _get_spectrum_forecaster(self) -> Optional[SpectrumForecaster]: + """Get forecaster for current branch (cond or uncond).""" + if self.spectrum_is_cfg_negative and self._spectrum_supports_cfg_cache: + return self.spectrum_forecaster_negative + return self.spectrum_forecaster + + def _set_spectrum_forecaster(self, forecaster: SpectrumForecaster) -> None: + """Set forecaster for current branch (cond or uncond).""" + if self.spectrum_is_cfg_negative and self._spectrum_supports_cfg_cache: + self.spectrum_forecaster_negative = forecaster + else: + self.spectrum_forecaster = forecaster + + def _get_spectrum_context(self) -> Optional[SpectrumContext]: + """Retrieve current Spectrum context from forward batch. Returns None if disabled.""" + from sglang.multimodal_gen.runtime.managers.forward_context import ( + get_forward_context, + ) + + try: + forward_context = get_forward_context() + except AssertionError: + return None + forward_batch = forward_context.forward_batch + if ( + forward_batch is None + or not forward_batch.enable_spectrum + or forward_batch.spectrum_params is None + ): + return None + + spectrum_params = forward_batch.spectrum_params + do_cfg = forward_batch.do_classifier_free_guidance + is_cfg_negative = forward_batch.is_cfg_negative + num_inference_steps = forward_batch.num_inference_steps + total_forward_steps = spectrum_params.get_total_forward_steps( + num_inference_steps, + do_cfg, + self._spectrum_supports_cfg_cache, + ) + + return SpectrumContext( + current_step=forward_context.current_timestep, + num_inference_steps=num_inference_steps, + total_forward_steps=total_forward_steps, + do_cfg=do_cfg, + is_cfg_negative=is_cfg_negative, + spectrum_params=spectrum_params, + debug=bool(getattr(forward_batch, "debug", False)), + ) + + def _record_spectrum_step_stat(self, actual_forward: bool) -> None: + if self.spectrum_is_cfg_negative and self._spectrum_supports_cfg_cache: + if actual_forward: + self.spectrum_real_steps_negative += 1 + else: + self.spectrum_skipped_steps_negative += 1 + return + + if actual_forward: + self.spectrum_real_steps += 1 + else: + self.spectrum_skipped_steps += 1 + + def _record_shadow_error_stat(self, rel_l2: float) -> None: + if self.spectrum_is_cfg_negative and self._spectrum_supports_cfg_cache: + self.spectrum_shadow_rel_l2_sum_negative += rel_l2 + self.spectrum_shadow_rel_l2_count_negative += 1 + return + + self.spectrum_shadow_rel_l2_sum += rel_l2 + self.spectrum_shadow_rel_l2_count += 1 + + def _emit_spectrum_summary(self, ctx: SpectrumContext) -> None: + if not ctx.debug: + return + + if self.spectrum_is_cfg_negative and self._spectrum_supports_cfg_cache: + real = self.spectrum_real_steps_negative + skipped = self.spectrum_skipped_steps_negative + err_sum = self.spectrum_shadow_rel_l2_sum_negative + err_count = self.spectrum_shadow_rel_l2_count_negative + branch = "negative" + else: + real = self.spectrum_real_steps + skipped = self.spectrum_skipped_steps + err_sum = self.spectrum_shadow_rel_l2_sum + err_count = self.spectrum_shadow_rel_l2_count + branch = "positive" + + total = real + skipped + skip_ratio = (skipped / total) if total > 0 else 0.0 + avg_rel_l2 = (err_sum / err_count) if err_count > 0 else float("nan") + logger.info( + "[Spectrum/%s] total=%d real=%d skipped=%d skip_ratio=%.3f shadow_rel_l2_avg=%s n=%d window_start=%.3f flex=%.3f", + branch, + total, + real, + skipped, + skip_ratio, + f"{avg_rel_l2:.4f}" if err_count > 0 else "NA", + err_count, + float(ctx.spectrum_params.window_size), + float(ctx.spectrum_params.flex_window), + ) + + def begin_spectrum_step(self) -> bool: + """Advance Spectrum schedule. Returns True when transformer blocks should run. + + Schedule (after ``warmup_steps`` real forwards): + - ``window = floor(curr_ws)`` — run a real forward when + ``(consecutive_cached + 1) % window == 0``, otherwise skip. + - Each real forward increases ``curr_ws`` by ``flex_window``, widening + gaps over time (paper alpha). + """ + ctx = self._get_spectrum_context() + self._spectrum_ctx = ctx + if ctx is None: + # Spectrum disabled — always run blocks (normal DiT path). + return True + + # Reset at the very first denoising step of each generation. + # Only the positive (or sole) branch triggers the reset so that: + # - single-branch models (FLUX, Hunyuan embedded guidance) reset once. + # - dual-branch models (Wan with true CFG) reset both counters from the + # positive-branch call and leave the negative-branch call unaffected, + # keeping both branches synchronised (both start at cnt=0 → cnt=1). + # Doing the reset here (not inside _get_spectrum_context) guarantees it + # fires exactly once per step, preventing the double-reset that would + # desync the two branches. + if ctx.current_step == 0 and not ctx.is_cfg_negative: + if ctx.debug: + logger.info( + "[Spectrum] Debug mode enables shadow-error validation; runtime perf is not representative of non-debug runs." + ) + self.reset_spectrum_state(ctx.spectrum_params) + + self.spectrum_is_cfg_negative = ctx.is_cfg_negative + params = ctx.spectrum_params + cnt, consecutive, curr_ws = self._get_spectrum_branch_state() + + # Warmup: first ``warmup_steps`` calls on this branch always run the DiT. + actual_forward = True + if cnt >= params.warmup_steps: + # After warmup, skip most steps and only run a real forward every + # ``window`` cached steps. ``consecutive`` counts skips since last real. + window = max(1, math.floor(curr_ws)) + actual_forward = (consecutive + 1) % window == 0 + if actual_forward: + # Widen the gap for the next stretch (paper alpha / flex_window). + curr_ws += params.flex_window + curr_ws = round(curr_ws, 3) + + # One denoising forward completed on this branch. + cnt += 1 + if actual_forward: + consecutive = 0 + else: + consecutive += 1 + self._record_spectrum_step_stat(actual_forward) + + # End-of-run wrap: after ``total_steps`` forwards on this branch, reset + # counters so state does not leak if the same module is reused. (A fresh + # run also resets via ``reset_spectrum_state`` at denoising timestep 0.) + total_steps = ctx.total_forward_steps + if cnt >= total_steps: + self._emit_spectrum_summary(ctx) + cnt = 0 + consecutive = 0 + curr_ws = params.window_size + + self._set_spectrum_branch_state(cnt, consecutive, curr_ws) + return actual_forward + + def spectrum_record_features(self, features: torch.Tensor) -> None: + """Append block outputs from a real forward to the branch forecaster.""" + ctx = self._spectrum_ctx or self._get_spectrum_context() + if ctx is None: + return + + params = ctx.spectrum_params + forecaster = self._get_spectrum_forecaster() + step_idx = float(ctx.current_step) + + # Initialize forecaster on first real forward + if forecaster is None: + cheb = ChebyshevForecaster( + M=params.m, + K=params.history_size, + lam=params.lam, + num_steps=params.tau_num_steps, + device=features.device, + feature_shape=features.shape, + ) + forecaster = SpectrumForecaster( + cheb, taylor_order=params.taylor_order, w=params.w + ) + self._set_spectrum_forecaster(forecaster) + + # In debug mode, compute shadow prediction error for validation + if ctx.debug and forecaster.ready() and step_idx >= float(params.warmup_steps): + predicted = forecaster.predict(step_idx).to( + dtype=features.dtype, device=features.device + ) + pred_f = predicted.float() + feat_f = features.detach().float() + denom = torch.norm(feat_f).item() + if denom > 0: + rel_l2 = (torch.norm(pred_f - feat_f).item()) / denom + self._record_shadow_error_stat(rel_l2) + + # Update forecaster with actual features from this real step + forecaster.update(step_idx, features.detach()) + self._spectrum_ctx = None + + def spectrum_predict_features(self, template: torch.Tensor) -> torch.Tensor: + """Return forecasted block outputs for a skipped step (same shape as template).""" + forecaster = self._get_spectrum_forecaster() + if forecaster is None or not forecaster.ready(): + return template + ctx = self._spectrum_ctx or self._get_spectrum_context() + if ctx is None: + return template + step_idx = float(ctx.current_step) + predicted = forecaster.predict(step_idx) + self._spectrum_ctx = None + return predicted.to(dtype=template.dtype, device=template.device) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/base.py b/python/sglang/multimodal_gen/runtime/models/dits/base.py index 9816f5fb0..75094ac98 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/base.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/base.py @@ -9,6 +9,9 @@ from torch import nn from sglang.multimodal_gen.configs.models import DiTConfig +# NOTE: SpectrumMixin lives in runtime.cache.spectrum +from sglang.multimodal_gen.runtime.cache.spectrum import SpectrumMixin + # NOTE: TeaCacheContext and TeaCacheMixin have been moved to # sglang.multimodal_gen.runtime.cache.teacache # For backwards compatibility, re-export from the new location @@ -87,11 +90,13 @@ class BaseDiT(nn.Module, ABC): return next(self.parameters()).device -class CachableDiT(TeaCacheMixin, BaseDiT): +class CachableDiT(SpectrumMixin, TeaCacheMixin, BaseDiT): """ - An intermediate base class that adds TeaCache optimization functionality to DiT models. + Base class for DiT models that support inference-time cache accelerators. + + Inherits ``SpectrumMixin`` (Chebyshev step skipping) and ``TeaCacheMixin`` + (temporal L1 similarity caching) plus ``BaseDiT`` core functionality. - Inherits TeaCacheMixin for cache logic and BaseDiT for core DiT functionality. """ # These are required class attributes that should be overridden by concrete implementations @@ -110,6 +115,7 @@ class CachableDiT(TeaCacheMixin, BaseDiT): def __init__(self, config: DiTConfig, **kwargs) -> None: super().__init__(config, **kwargs) + self._init_spectrum_state() self._init_teacache_state() @classmethod diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index 793cf0338..a5423edd5 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -66,6 +66,7 @@ from sglang.multimodal_gen.runtime.layers.visual_embedding import ( CombinedTimestepGuidanceTextProjEmbeddings, CombinedTimestepTextProjEmbeddings, ) +from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) @@ -1149,24 +1150,35 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): ip_hidden_states = self.encoder_hid_proj(ip_adapter_image_embeds) joint_attention_kwargs.update({"ip_hidden_states": ip_hidden_states}) - for block in self.transformer_blocks: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - freqs_cis=freqs_cis, - joint_attention_kwargs=joint_attention_kwargs, - num_replicated_prefix=num_replicated_prefix, - ) - for block in self.single_transformer_blocks: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - freqs_cis=singles_freqs_cis, - joint_attention_kwargs=joint_attention_kwargs, - num_replicated_prefix=num_replicated_prefix, - ) + forward_batch = get_forward_context().forward_batch + spectrum_enabled = forward_batch is not None and forward_batch.enable_spectrum + run_transformer_blocks = ( + self.begin_spectrum_step() if spectrum_enabled else True + ) + if run_transformer_blocks: + for block in self.transformer_blocks: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + freqs_cis=freqs_cis, + joint_attention_kwargs=joint_attention_kwargs, + num_replicated_prefix=num_replicated_prefix, + ) + for block in self.single_transformer_blocks: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + freqs_cis=singles_freqs_cis, + joint_attention_kwargs=joint_attention_kwargs, + num_replicated_prefix=num_replicated_prefix, + ) + if spectrum_enabled: + self.spectrum_record_features(hidden_states) + else: + if spectrum_enabled: + hidden_states = self.spectrum_predict_features(hidden_states) hidden_states = self.norm_out(hidden_states, temb) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py index 2f329c564..16fed4a6c 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/hunyuanvideo.py @@ -662,6 +662,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixi forward_context = get_forward_context() forward_batch = forward_context.forward_batch enable_teacache = forward_batch is not None and forward_batch.enable_teacache + enable_spectrum = forward_batch is not None and forward_batch.enable_spectrum if guidance is None: guidance = torch.tensor( @@ -744,11 +745,10 @@ class HunyuanVideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixi freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None - should_skip_forward = self.should_skip_forward_for_cached_states( - img=img, vec=vec - ) - - if should_skip_forward: + run_transformer_blocks = self.begin_spectrum_step() + if enable_spectrum and not run_transformer_blocks: + img = self.spectrum_predict_features(img) + elif self.should_skip_forward_for_cached_states(img=img, vec=vec): img = self.retrieve_cached_states(img) else: if enable_teacache: @@ -786,6 +786,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixi if enable_teacache: self.maybe_cache_states(img, original_img) + if enable_spectrum: + self.spectrum_record_features(img) # Final layer processing img = self.final_layer(img, vec) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/stablediffusion3.py b/python/sglang/multimodal_gen/runtime/models/dits/stablediffusion3.py index 419d41af8..82236c919 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/stablediffusion3.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/stablediffusion3.py @@ -17,6 +17,7 @@ from diffusers.models.normalization import AdaLayerNormContinuous from sglang.multimodal_gen.configs.models.dits.stablediffusion3 import ( StableDiffusion3TransformerConfig, ) +from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload import ( LayerwiseOffloadableModuleMixin, ) @@ -128,26 +129,37 @@ class SD3Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): else: interval_control = 0 - for index_block, block in enumerate(self.transformer_blocks): - if index_block not in skip_layer_set: - encoder_embeddings, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_embeddings, - temb=temb, - joint_attention_kwargs=joint_attention_kwargs, - ) + forward_batch = get_forward_context().forward_batch + spectrum_enabled = forward_batch is not None and forward_batch.enable_spectrum + run_transformer_blocks = ( + self.begin_spectrum_step() if spectrum_enabled else True + ) + if run_transformer_blocks: + for index_block, block in enumerate(self.transformer_blocks): + if index_block not in skip_layer_set: + encoder_embeddings, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_embeddings, + temb=temb, + joint_attention_kwargs=joint_attention_kwargs, + ) - # controlnet residual - if ( - block_controlnet_hidden_states is not None - and block.context_pre_only is False - ): - hidden_states = ( - hidden_states - + block_controlnet_hidden_states[ - int(index_block / interval_control) - ] - ) + # controlnet residual + if ( + block_controlnet_hidden_states is not None + and block.context_pre_only is False + ): + hidden_states = ( + hidden_states + + block_controlnet_hidden_states[ + int(index_block / interval_control) + ] + ) + if spectrum_enabled: + self.spectrum_record_features(hidden_states) + else: + if spectrum_enabled: + hidden_states = self.spectrum_predict_features(hidden_states) hidden_states = self.norm_out(hidden_states, temb) hidden_states = self.proj_out(hidden_states) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py index 907871973..a00431e06 100755 --- a/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/wanvideo.py @@ -1010,6 +1010,7 @@ class WanTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin): self.enable_teacache = ( forward_batch is not None and forward_batch.enable_teacache ) + enable_spectrum = forward_batch is not None and forward_batch.enable_spectrum orig_dtype = hidden_states.dtype if not isinstance(encoder_hidden_states, torch.Tensor): @@ -1143,12 +1144,14 @@ class WanTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin): assert encoder_hidden_states.dtype == orig_dtype # 4. Transformer blocks - # if caching is enabled, we might be able to skip the forward pass + run_transformer_blocks = self.begin_spectrum_step() should_skip_forward = self.should_skip_forward_for_cached_states( timestep_proj=timestep_proj, temb=temb ) - if should_skip_forward: + if enable_spectrum and not run_transformer_blocks: + hidden_states = self.spectrum_predict_features(hidden_states) + elif should_skip_forward: hidden_states = self.retrieve_cached_states(hidden_states) else: # if teacache is enabled, we need to cache the original hidden states @@ -1162,6 +1165,8 @@ class WanTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin): # if teacache is enabled, we need to cache the original hidden states if self.enable_teacache: self.maybe_cache_states(hidden_states, original_hidden_states) + if enable_spectrum: + self.spectrum_record_features(hidden_states) self.cnt += 1 if sequence_shard_enabled: diff --git a/python/sglang/multimodal_gen/test/unit/test_sampling_params.py b/python/sglang/multimodal_gen/test/unit/test_sampling_params.py index cfdf58f2a..fe1dad199 100644 --- a/python/sglang/multimodal_gen/test/unit/test_sampling_params.py +++ b/python/sglang/multimodal_gen/test/unit/test_sampling_params.py @@ -92,6 +92,12 @@ class TestSamplingParamsValidate(unittest.TestCase): with self.assertRaisesRegex(ValueError, r"boundary_ratio"): SamplingParams(boundary_ratio=math.nan) + def test_teacache_and_spectrum_are_mutually_exclusive(self): + with self.assertRaisesRegex( + ValueError, r"enable_teacache and enable_spectrum are mutually exclusive" + ): + SamplingParams(enable_teacache=True, enable_spectrum=True) + class TestSamplingParamsSubclass(unittest.TestCase): def test_flux_defaults_resolution_when_not_provided(self): @@ -243,6 +249,38 @@ class TestSamplingParamsCliArgs(unittest.TestCase): self._parse_cli_kwargs(["--quality", "high"])["quality"], "high" ) + def test_get_cli_args_maps_spectrum_prefixed_flags(self): + kwargs = self._parse_cli_kwargs( + [ + "--enable-spectrum", + "--spectrum-window-size", + "2.5", + "--spectrum-flex-window", + "0.9", + "--spectrum-warmup-steps", + "6", + "--spectrum-m", + "3", + "--spectrum-lam", + "0.2", + "--spectrum-tau-num-steps", + "42", + ] + ) + + self.assertTrue(kwargs["enable_spectrum"]) + self.assertEqual( + kwargs["spectrum_params"], + { + "window_size": 2.5, + "flex_window": 0.9, + "warmup_steps": 6, + "m": 3, + "lam": 0.2, + "tau_num_steps": 42, + }, + ) + def test_qwen_image_cli_path_preserves_model_defaults(self): params = self._make_qwen_image_params([]) diff --git a/python/sglang/multimodal_gen/test/unit/test_spectrum.py b/python/sglang/multimodal_gen/test/unit/test_spectrum.py new file mode 100644 index 000000000..e7b1591dc --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_spectrum.py @@ -0,0 +1,92 @@ +# SPDX-License-Identifier: Apache-2.0 +import unittest + +import torch + +from sglang.multimodal_gen.runtime.cache.spectrum import ( + ChebyshevForecaster, + SpectrumForecaster, +) + + +class TestSpectrumForecaster(unittest.TestCase): + def test_chebyshev_fit_and_predict(self) -> None: + """Predict returns the expected feature shape after fitting on prior steps.""" + forecaster = ChebyshevForecaster( + M=2, K=10, lam=0.1, num_steps=10, feature_shape=(4, 8) + ) + for step in range(6): + forecaster.update(float(step), torch.randn(4, 8)) + predicted = forecaster.predict(6.0) + self.assertEqual(predicted.shape, (4, 8)) + + def test_chebyshev_fit_and_predict_bfloat16(self) -> None: + """bfloat16 inputs produce bfloat16 predictions.""" + forecaster = ChebyshevForecaster( + M=2, K=10, lam=0.1, num_steps=10, feature_shape=(4, 8) + ) + for step in range(6): + forecaster.update(float(step), torch.randn(4, 8, dtype=torch.bfloat16)) + predicted = forecaster.predict(6.0) + self.assertEqual(predicted.shape, (4, 8)) + self.assertEqual(predicted.dtype, torch.bfloat16) + + @unittest.skipUnless(torch.cuda.is_available(), "requires CUDA autocast") + def test_chebyshev_fit_and_predict_bfloat16_under_autocast(self) -> None: + """CUDA bf16 autocast keeps prediction working and preserves bf16 output.""" + forecaster = ChebyshevForecaster( + M=2, K=10, lam=0.1, num_steps=10, feature_shape=(4, 8) + ).cuda() + with torch.autocast(device_type="cuda", dtype=torch.bfloat16): + for step in range(6): + forecaster.update( + float(step), torch.randn(4, 8, device="cuda", dtype=torch.bfloat16) + ) + predicted = forecaster.predict(6.0) + self.assertEqual(predicted.shape, (4, 8)) + self.assertEqual(predicted.dtype, torch.bfloat16) + + def test_spectrum_blend_predict(self) -> None: + """The blended Spectrum forecaster returns the expected output shape.""" + cheb = ChebyshevForecaster( + M=2, K=10, lam=0.1, num_steps=10, feature_shape=(2, 3) + ) + blend = SpectrumForecaster(cheb, taylor_order=1, w=0.5) + for step in range(4): + blend.update(float(step), torch.ones(2, 3) * step) + out = blend.predict(4.0) + self.assertEqual(out.shape, (2, 3)) + + def test_chebyshev_prediction_error_is_bounded_on_smooth_signal(self) -> None: + """Prediction error stays very low on a deterministic smooth signal.""" + + def smooth_feature(step: float) -> torch.Tensor: + # Linear trend should be modeled accurately by M=1 Chebyshev basis. + base = 0.5 + 0.125 * step + return torch.tensor( + [[base, base + 0.1], [0.75 * base, -0.5 * base]], + dtype=torch.float32, + ) + + forecaster = ChebyshevForecaster( + M=1, K=16, lam=1e-6, num_steps=50, feature_shape=(2, 2) + ) + for step in range(8): + forecaster.update(float(step), smooth_feature(float(step))) + + target = smooth_feature(8.0) + predicted = forecaster.predict(8.0) + rel_l2 = torch.norm(predicted - target) / torch.norm(target) + self.assertLess(rel_l2.item(), 1e-3) + + def test_chebyshev_tau_horizon_matches_reference(self) -> None: + """Tau normalization matches the fixed 50-step horizon used by the reference.""" + f50 = ChebyshevForecaster(M=0, num_steps=50, feature_shape=(1,)) + f20 = ChebyshevForecaster(M=0, num_steps=20, feature_shape=(1,)) + t = torch.tensor([10.0]) + self.assertAlmostEqual(f50._taus(t).item(), -0.6, places=5) + self.assertAlmostEqual(f20._taus(t).item(), 0.0, places=5) + + +if __name__ == "__main__": + unittest.main()