From c9d713a52de17e38a21db4ce9f363b2cb0e4afd5 Mon Sep 17 00:00:00 2001 From: xmarre Date: Tue, 31 Mar 2026 21:32:09 +0200 Subject: [PATCH] Reduce forecast VRAM and add FLUX debug logging --- comfyui_spectrum/flux.py | 41 ++++++++++++++ comfyui_spectrum/runtime.py | 109 +++++++++++++++++++++++++++--------- 2 files changed, 124 insertions(+), 26 deletions(-) diff --git a/comfyui_spectrum/flux.py b/comfyui_spectrum/flux.py index 36ef29f..bd6b596 100644 --- a/comfyui_spectrum/flux.py +++ b/comfyui_spectrum/flux.py @@ -92,6 +92,29 @@ def _forecast_feature_sanitization_stats(feature: torch.Tensor, dtype: torch.dty } +def _debug_log_flux_forecast_context( + runtime: SpectrumRuntime, + *, + stage: str, + run_id: int, + solver_step_id: int, + expected_feature_shape: Tuple[int, ...], + post_input_patches_len: int, + timestep_zero_index: Optional[Sequence[Tuple[int, int]]], +) -> None: + if not runtime.cfg.debug: + return + LOG.warning( + "Spectrum flux forecast context run_id=%s step=%s stage=%s expected_shape=%s post_input_patches=%s timestep_zero_index=%s", + run_id, + solver_step_id, + stage, + expected_feature_shape, + post_input_patches_len, + timestep_zero_index is not None, + ) + + def _build_branch_signature(transformer_options: Dict[str, Any]) -> Optional[tuple[Any, ...]]: signature = [] cond_or_uncond = transformer_options.get("cond_or_uncond") @@ -383,6 +406,15 @@ def _run_flux_forward_with_spectrum( if step_ctx is not None and hidden_dim is not None and not post_input_patches: _, run_id, solver_step_id, actual_forward = step_ctx expected_feature_shape = (raw_img.shape[0], raw_img.shape[1], hidden_dim) + _debug_log_flux_forecast_context( + runtime, + stage="pre_img_in", + run_id=run_id, + solver_step_id=solver_step_id, + expected_feature_shape=expected_feature_shape, + post_input_patches_len=len(post_input_patches), + timestep_zero_index=timestep_zero_index, + ) vec_orig = vec txt_vec = vec modulation_dims = None @@ -471,6 +503,15 @@ def _run_flux_forward_with_spectrum( if step_ctx is not None and call_id is None: _, run_id, solver_step_id, actual_forward = step_ctx + _debug_log_flux_forecast_context( + runtime, + stage="post_img_in", + run_id=run_id, + solver_step_id=solver_step_id, + expected_feature_shape=expected_feature_shape, + post_input_patches_len=len(post_input_patches), + timestep_zero_index=timestep_zero_index, + ) call_id = runtime.register_model_hook_call( run_id, solver_step_id, diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index 31e1473..5875614 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -1,5 +1,6 @@ from __future__ import annotations +import logging import math from collections import defaultdict, deque from dataclasses import dataclass, field @@ -10,6 +11,8 @@ import torch from .config import SpectrumConfig from .forecast import ChebyshevSpectrumForecaster +LOG = logging.getLogger(__name__) + @dataclass(slots=True) class RuntimeStats: @@ -49,6 +52,7 @@ class _ActiveStep: call_used_forecast: list[bool] = field(default_factory=list) 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 @@ -123,6 +127,7 @@ class SpectrumRuntime: self._history_batch_labels = None for step in self._active_steps.values(): 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 @@ -334,6 +339,7 @@ class SpectrumRuntime: step.call_used_forecast.append(False) step.call_actual_features.append(None) step.call_predicted_features.append(None) + step.call_prediction_rows.append(None) return len(step.call_expected_shapes) - 1 def observe_actual_feature( @@ -350,6 +356,7 @@ class SpectrumRuntime: step.call_used_forecast[resolved_call_id] = False step.used_forecast_any = any(step.call_used_forecast) step.call_predicted_features[resolved_call_id] = None + step.call_prediction_rows[resolved_call_id] = None if step.actual_feature_device is None: step.actual_feature_device = feature.device if step.actual_feature_dtype is None: @@ -381,6 +388,22 @@ class SpectrumRuntime: positions[label].append(idx) return positions + @staticmethod + def _tensor_shape(feature: Optional[torch.Tensor]) -> Optional[tuple[int, ...]]: + if feature is None: + 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 predict_feature( self, run_id: int, @@ -409,6 +432,22 @@ class SpectrumRuntime: return None needs_full_prediction = (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", + run_id, + solver_step_id, + resolved_call_id, + step.hook_call_count, + 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), + self._tensor_shape(step.call_predicted_features[resolved_call_id]), + step.call_prediction_rows[resolved_call_id], + ) + if ( self._history_batch_labels is None and step.hook_call_count > 1 @@ -417,23 +456,27 @@ class SpectrumRuntime: ): return None - if step.call_predicted_features[resolved_call_id] is None: - if 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 + cached_prediction = 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 @@ -443,26 +486,40 @@ class SpectrumRuntime: if not positions: return None order.append(positions.popleft()) - predicted_feature = step.predicted_full_feature[order, ...] + prediction_rows = tuple(order) + step.call_prediction_rows[resolved_call_id] = prediction_rows else: return None - else: - predicted_feature = self.forecaster.predict( - time_coord=step.time_coord, - 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 + predicted_feature = self._select_prediction_rows(step.predicted_full_feature, prediction_rows) + else: + predicted_feature = self.forecaster.predict( + time_coord=step.time_coord, + 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 step.call_used_forecast[resolved_call_id] = True step.used_forecast_any = True - return step.call_predicted_features[resolved_call_id] + 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", + 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], + ) + return predicted_feature def abort_solver_step(self, run_id: int, solver_step_id: int) -> None: step = self._require_active_step(run_id, solver_step_id) 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