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 (