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:
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user