From f926b81d0ed1e5be503bb65633d514ecda358fa9 Mon Sep 17 00:00:00 2001 From: xmarre Date: Tue, 31 Mar 2026 19:17:35 +0200 Subject: [PATCH] Preserve forecast output device across CPU history updates --- comfyui_spectrum/forecast.py | 24 ++++++++++++++++++++---- comfyui_spectrum/runtime.py | 8 +++++++- 2 files changed, 27 insertions(+), 5 deletions(-) diff --git a/comfyui_spectrum/forecast.py b/comfyui_spectrum/forecast.py index 1f21b0a..17bc315 100644 --- a/comfyui_spectrum/forecast.py +++ b/comfyui_spectrum/forecast.py @@ -71,7 +71,15 @@ class ChebyshevSpectrumForecaster: needed = max(2, int(min_points) if min_points is not None else self.degree + 1) return len(self._history) >= needed - def update(self, time_coord: float, feature: torch.Tensor) -> None: + def update( + self, + time_coord: float, + feature: torch.Tensor, + *, + predict_device: Optional[torch.device] = None, + output_device: Optional[torch.device] = None, + output_dtype: Optional[torch.dtype] = None, + ) -> None: feat = feature.detach() if self._feature_shape is None: self._feature_shape = feat.shape @@ -80,9 +88,17 @@ class ChebyshevSpectrumForecaster: raise ValueError( f"Spectrum feature shape changed from {tuple(self._feature_shape)} to {tuple(feat.shape)}." ) - self._feature_dtype = feat.dtype - self._predict_device = feat.device - self._output_device = feat.device + self._feature_dtype = output_dtype if output_dtype is not None else feat.dtype + if predict_device is not None: + self._predict_device = predict_device + elif self._predict_device is None: + self._predict_device = feat.device + if output_device is not None: + self._output_device = output_device + elif self._output_device is None: + self._output_device = feat.device + if self._predict_device is None: + self._predict_device = self._output_device feature_flat = feat.reshape(-1).to(device="cpu", dtype=torch.float32, copy=True) basis_row = self._build_design( diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index cd23974..7d13e15 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -561,7 +561,13 @@ class SpectrumRuntime: device=target_device, dtype=target_dtype if target_dtype is not None else combined_feature.dtype, ) - self.forecaster.update(step.time_coord, combined_feature) + self.forecaster.update( + step.time_coord, + combined_feature, + predict_device=step.actual_feature_device, + output_device=step.actual_feature_device, + output_dtype=step.actual_feature_dtype, + ) except ValueError: self._disable_forecasting("combined actual feature shape changed across solver steps") if (