Merge pull request #17 from xmarre/codex/fix-flux-forecast-vram-spike

Fix FLUX forecast VRAM spikes in labeled multi-call steps
This commit is contained in:
xmarre
2026-03-31 21:42:24 +02:00
committed by GitHub
2 changed files with 124 additions and 26 deletions
+41
View File
@@ -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,
+83 -26
View File
@@ -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