Merge pull request #18 from xmarre/codex/fix-flux-prediction-leak

Fix FLUX forecast residency for batched multi-call steps
This commit is contained in:
xmarre
2026-03-31 22:55:23 +02:00
committed by GitHub
2 changed files with 97 additions and 75 deletions
+51 -11
View File
@@ -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)
+46 -64
View File
@@ -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: