diff --git a/comfyui_spectrum/forecast.py b/comfyui_spectrum/forecast.py index a53a1ee..78d068c 100644 --- a/comfyui_spectrum/forecast.py +++ b/comfyui_spectrum/forecast.py @@ -1,7 +1,7 @@ from __future__ import annotations from dataclasses import dataclass -from typing import List, Optional +from typing import List, Optional, Sequence import torch @@ -315,22 +315,55 @@ class ChebyshevSpectrumForecaster: self._coeff_device = self._coeff.to(device=self._predict_device, dtype=self._predict_dtype) return self._coeff_device - def _linear_prediction(self, time_coord: float) -> torch.Tensor: - if self._latest_feature_flat_device is None or self._latest_time_coord is None: + @staticmethod + def _select_rows(tensor: torch.Tensor, rows: tuple[int, ...], *, dim: int) -> torch.Tensor: + if len(rows) == 0: + raise RuntimeError("Spectrum forecaster received an empty row selection.") + if len(rows) == tensor.shape[dim] and all(row == idx for idx, row in enumerate(rows)): + return tensor + start = rows[0] + if all(row == start + offset for offset, row in enumerate(rows)): + return tensor.narrow(dim, start, len(rows)) + index = torch.tensor(rows, device=tensor.device, dtype=torch.long) + return tensor.index_select(dim, index) + + def _normalize_prediction_rows(self, rows: Optional[Sequence[int]]) -> tuple[int, ...]: + if self._feature_shape is None: raise RuntimeError("Spectrum forecaster has no cached feature history.") + batch = int(self._feature_shape[0]) + if rows is None: + return tuple(range(batch)) + resolved = tuple(int(row) for row in rows) + if not resolved: + raise RuntimeError("Spectrum forecaster received an empty row selection.") + for row in resolved: + if row < 0 or row >= batch: + raise RuntimeError( + f"Spectrum forecaster row selection {resolved} is outside the cached batch size {batch}." + ) + return resolved + + def _linear_prediction_rows(self, time_coord: float, rows: tuple[int, ...]) -> torch.Tensor: + if self._feature_shape is None or self._latest_feature_flat_device is None or self._latest_time_coord is None: + raise RuntimeError("Spectrum forecaster has no cached feature history.") + last = self._select_rows(self._latest_feature_flat_device.reshape(self._feature_shape), rows, dim=0) if self._previous_feature_flat_device is None or self._previous_time_coord is None: - return self._latest_feature_flat_device + return last delta_coord = self._latest_time_coord - self._previous_time_coord if abs(delta_coord) <= 1e-12: - return self._latest_feature_flat_device + return last + prev = self._select_rows(self._previous_feature_flat_device.reshape(self._feature_shape), rows, dim=0) k = (float(time_coord) - float(self._latest_time_coord)) / float(delta_coord) - last_f = self._latest_feature_flat_device - prev_f = self._previous_feature_flat_device - return last_f + k * (last_f - prev_f) + return last + k * (last - prev) - def predict(self, time_coord: float, blend_weight: float) -> torch.Tensor: + def predict_rows( + self, + time_coord: float, + rows: Optional[Sequence[int]], + blend_weight: float, + ) -> torch.Tensor: if ( self._feature_shape is None or self._feature_dtype is None @@ -344,16 +377,23 @@ class ChebyshevSpectrumForecaster: raise RuntimeError("Spectrum forecaster is not ready yet.") degree, _ = self._ensure_coeff() + resolved_rows = self._normalize_prediction_rows(rows) + subset_shape = (len(resolved_rows), *tuple(self._feature_shape[1:])) coeff_device = self._ensure_coeff_device() coord_star = torch.tensor([float(time_coord)], device=self._predict_device, dtype=torch.float32) design_star = self._build_design(coord_star, degree).to(dtype=coeff_device.dtype) - spectral = (design_star @ coeff_device).reshape(self._feature_shape) + coeff_view = coeff_device.reshape(coeff_device.shape[0], *tuple(self._feature_shape)) + coeff_rows = self._select_rows(coeff_view, resolved_rows, dim=1) + spectral = (design_star @ coeff_rows.reshape(coeff_rows.shape[0], -1)).reshape(subset_shape) blend = float(blend_weight) if blend >= (1.0 - 1e-12): out = spectral else: - linear = self._linear_prediction(time_coord).reshape(self._feature_shape) + linear = self._linear_prediction_rows(time_coord, resolved_rows) out = blend * spectral + (1.0 - blend) * linear return out.to(device=self._output_device, dtype=self._feature_dtype) + + def predict(self, time_coord: float, blend_weight: float) -> torch.Tensor: + return self.predict_rows(time_coord=time_coord, rows=None, blend_weight=blend_weight) diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index 5875614..bd3cf5d 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -53,9 +53,7 @@ class _ActiveStep: call_actual_features: list[Optional[torch.Tensor]] = field(default_factory=list) call_predicted_features: list[Optional[torch.Tensor]] = field(default_factory=list) call_prediction_rows: list[Optional[tuple[int, ...]]] = field(default_factory=list) - predicted_full_feature: Optional[torch.Tensor] = None prediction_row_positions: Optional[dict[Any, deque[int]]] = None - prediction_next_row: int = 0 used_forecast_any: bool = False actual_feature_device: Optional[torch.device] = None actual_feature_dtype: Optional[torch.dtype] = None @@ -129,9 +127,7 @@ class SpectrumRuntime: step.call_predicted_features = [None] * len(step.call_predicted_features) step.call_prediction_rows = [None] * len(step.call_prediction_rows) step.call_used_forecast = [False] * len(step.call_used_forecast) - step.predicted_full_feature = None step.prediction_row_positions = None - step.prediction_next_row = 0 step.actual_feature_device = None step.actual_feature_dtype = None self.stats.current_window = self.curr_ws @@ -394,15 +390,34 @@ class SpectrumRuntime: return None return tuple(feature.shape) - @staticmethod - def _select_prediction_rows(feature: torch.Tensor, rows: tuple[int, ...]) -> torch.Tensor: - if len(rows) == feature.shape[0] and all(row == idx for idx, row in enumerate(rows)): - return feature - start = rows[0] - if all(row == start + offset for offset, row in enumerate(rows)): - return feature[start : start + len(rows), ...] - index = torch.tensor(rows, device=feature.device, dtype=torch.long) - return feature.index_select(0, index) + def _prediction_rows_for_call( + self, + step: _ActiveStep, + resolved_call_id: int, + ) -> Optional[tuple[int, ...]]: + cached_rows = step.call_prediction_rows[resolved_call_id] + if cached_rows is not None: + return cached_rows + target_batch_labels = step.call_batch_labels[resolved_call_id] + if target_batch_labels is None or self._history_batch_labels is None: + return None + if step.prediction_row_positions is None: + step.prediction_row_positions = self._build_label_positions(self._history_batch_labels) + trial_positions = { + label: deque(position_list) + for label, position_list in step.prediction_row_positions.items() + } + + order = [] + for label in target_batch_labels: + positions = trial_positions.get(label) + if not positions: + return None + order.append(positions.popleft()) + step.prediction_row_positions = trial_positions + prediction_rows = tuple(order) + step.call_prediction_rows[resolved_call_id] = prediction_rows + return prediction_rows def predict_feature( self, @@ -422,7 +437,7 @@ class SpectrumRuntime: target_shape = tuple(expected_shape) if expected_shape is not None else step.call_expected_shapes[resolved_call_id] target_batch_labels = step.call_batch_labels[resolved_call_id] history_shape = self.forecaster.feature_shape - needs_full_prediction = False + needs_row_selection = False if history_shape is not None: history_shape = tuple(history_shape) if history_shape[1:] != target_shape[1:]: @@ -430,11 +445,11 @@ class SpectrumRuntime: return None if self._history_batch_labels is None and history_shape[0] != target_shape[0]: return None - needs_full_prediction = (self._history_batch_labels is not None) + needs_row_selection = (self._history_batch_labels is not None) if self.cfg.debug: LOG.warning( - "Spectrum forecast request run_id=%s step=%s call=%s hook_calls=%s target_shape=%s target_has_labels=%s history_has_labels=%s needs_full_prediction=%s predicted_full_shape=%s cached_call_shape=%s cached_rows=%s", + "Spectrum forecast request run_id=%s step=%s call=%s hook_calls=%s target_shape=%s target_has_labels=%s history_has_labels=%s needs_row_selection=%s cached_call_shape=%s cached_rows=%s", run_id, solver_step_id, resolved_call_id, @@ -442,8 +457,7 @@ class SpectrumRuntime: target_shape, target_batch_labels is not None, self._history_batch_labels is not None, - needs_full_prediction, - self._tensor_shape(step.predicted_full_feature), + needs_row_selection, self._tensor_shape(step.call_predicted_features[resolved_call_id]), step.call_prediction_rows[resolved_call_id], ) @@ -451,66 +465,39 @@ class SpectrumRuntime: if ( self._history_batch_labels is None and step.hook_call_count > 1 - and not needs_full_prediction - and step.predicted_full_feature is None + and not needs_row_selection ): return None - cached_prediction = step.call_predicted_features[resolved_call_id] + cached_prediction = None if needs_row_selection else step.call_predicted_features[resolved_call_id] if cached_prediction is not None: predicted_feature = cached_prediction - elif needs_full_prediction: - if step.predicted_full_feature is None: - predicted_full_feature = self.forecaster.predict( - time_coord=step.time_coord, - blend_weight=self.cfg.blend_weight, - ) - if history_shape is not None and tuple(predicted_full_feature.shape) != history_shape: - self._disable_forecasting("predicted feature shape did not match the current solver-step input") - return None - step.predicted_full_feature = predicted_full_feature - step.prediction_next_row = 0 - if self._history_batch_labels is not None: - step.prediction_row_positions = self._build_label_positions(self._history_batch_labels) - else: - step.prediction_row_positions = None - - prediction_rows = step.call_prediction_rows[resolved_call_id] - if prediction_rows is None: - if target_batch_labels is not None and self._history_batch_labels is not None: - if step.prediction_row_positions is None: - return None - order = [] - for label in target_batch_labels: - positions = step.prediction_row_positions.get(label) - if not positions: - return None - order.append(positions.popleft()) - prediction_rows = tuple(order) - step.call_prediction_rows[resolved_call_id] = prediction_rows - else: - return None - predicted_feature = self._select_prediction_rows(step.predicted_full_feature, prediction_rows) else: - predicted_feature = self.forecaster.predict( + prediction_rows: Optional[tuple[int, ...]] = None + if needs_row_selection: + prediction_rows = self._prediction_rows_for_call(step, resolved_call_id) + if prediction_rows is None: + return None + predicted_feature = self.forecaster.predict_rows( time_coord=step.time_coord, + rows=prediction_rows, blend_weight=self.cfg.blend_weight, ) if tuple(predicted_feature.shape) != target_shape: self._disable_forecasting("predicted feature shape did not match the current solver-step input") return None - step.call_predicted_features[resolved_call_id] = predicted_feature + if not needs_row_selection: + step.call_predicted_features[resolved_call_id] = predicted_feature step.call_used_forecast[resolved_call_id] = True step.used_forecast_any = True if self.cfg.debug: LOG.warning( - "Spectrum forecast result run_id=%s step=%s call=%s predicted_shape=%s predicted_full_shape=%s cached_call_shape=%s selected_rows=%s", + "Spectrum forecast result run_id=%s step=%s call=%s predicted_shape=%s cached_call_shape=%s selected_rows=%s", run_id, solver_step_id, resolved_call_id, self._tensor_shape(predicted_feature), - self._tensor_shape(step.predicted_full_feature), self._tensor_shape(step.call_predicted_features[resolved_call_id]), step.call_prediction_rows[resolved_call_id], ) @@ -521,9 +508,7 @@ class SpectrumRuntime: step.call_predicted_features = [None] * len(step.call_predicted_features) step.call_prediction_rows = [None] * len(step.call_prediction_rows) step.call_used_forecast = [False] * len(step.call_used_forecast) - step.predicted_full_feature = None step.prediction_row_positions = None - step.prediction_next_row = 0 step.actual_feature_device = None step.actual_feature_dtype = None self._active_steps.pop(int(solver_step_id), None) @@ -545,11 +530,8 @@ class SpectrumRuntime: if observed_actual and used_forecast_any: self._disable_forecasting("solver step mixed forecasted and actual model-hook paths") used_forecast_any = False - elif used_forecast_any and step.predicted_full_feature is not None: - if step.prediction_row_positions is not None: - if any(positions for positions in step.prediction_row_positions.values()): - self._disable_forecasting("forecasted solver step batch layout changed within one solver step") - elif step.prediction_next_row != step.predicted_full_feature.shape[0]: + elif used_forecast_any and step.prediction_row_positions is not None: + if any(positions for positions in step.prediction_row_positions.values()): self._disable_forecasting("forecasted solver step batch layout changed within one solver step") if used_forecast_any: