Fix split forecast guards for unlabeled subcalls

This commit is contained in:
xmarre
2026-03-30 22:28:42 +02:00
parent 4108a3210f
commit f92860d962
2 changed files with 82 additions and 17 deletions
+49 -17
View File
@@ -51,6 +51,7 @@ class _ActiveStep:
call_predicted_features: list[Optional[torch.Tensor]] = field(default_factory=list)
predicted_full_feature: Optional[torch.Tensor] = None
prediction_row_positions: Optional[dict[int, deque[int]]] = None
prediction_next_row: int = 0
class SpectrumRuntime:
@@ -122,6 +123,7 @@ class SpectrumRuntime:
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
self.stats.current_window = self.curr_ws
self.stats.forecast_disabled = True
self.stats.disable_reason = reason
@@ -367,22 +369,27 @@ 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
if history_shape is not None:
history_shape = tuple(history_shape)
if history_shape[1:] != target_shape[1:]:
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
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)
or (history_shape[0] != target_shape[0])
)
if self._history_batch_labels is None and step.hook_call_count > 1:
return None
if self._history_batch_labels is not None and target_batch_labels is None:
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
):
return None
if step.call_predicted_features[resolved_call_id] is None:
if self._history_batch_labels is not None:
if needs_full_prediction:
if step.predicted_full_feature is None:
predicted_full_feature = self.forecaster.predict(
time_coord=step.time_coord,
@@ -392,20 +399,38 @@ class SpectrumRuntime:
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_row_positions = self._build_label_positions(self._history_batch_labels)
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
if target_batch_labels is None or step.prediction_row_positions is None:
return None
order = []
for label in target_batch_labels:
positions = step.prediction_row_positions.get(int(label))
if not positions:
if target_batch_labels is not None and self._history_batch_labels is not None:
if step.prediction_row_positions is None:
if any(step.call_used_forecast):
self._disable_forecasting("forecasted solver step batch layout changed within one solver step")
return None
order.append(positions.popleft())
predicted_feature = step.predicted_full_feature[order, ...]
order = []
for label in target_batch_labels:
positions = step.prediction_row_positions.get(int(label))
if not positions:
if any(step.call_used_forecast):
self._disable_forecasting("forecasted solver step batch layout changed within one solver step")
return None
order.append(positions.popleft())
predicted_feature = step.predicted_full_feature[order, ...]
else:
if step.prediction_row_positions is not None and any(step.call_used_forecast):
self._disable_forecasting("forecasted solver step batch layout changed within one solver step")
return None
start = step.prediction_next_row
end = start + target_shape[0]
if step.predicted_full_feature is None or end > step.predicted_full_feature.shape[0]:
if any(step.call_used_forecast):
self._disable_forecasting("forecasted solver step batch layout changed within one solver step")
return None
predicted_feature = step.predicted_full_feature[start:end, ...]
step.prediction_next_row = end
else:
predicted_feature = self.forecaster.predict(
time_coord=step.time_coord,
@@ -425,6 +450,7 @@ class SpectrumRuntime:
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
self._active_steps.pop(int(solver_step_id), None)
def finalize_solver_step(self, run_id: int, solver_step_id: int, *, used_forecast: bool) -> None:
@@ -441,6 +467,12 @@ class SpectrumRuntime:
self._disable_forecasting("solver step finished without an actual feature or a forecasted feature")
if any(not (obs or used) for obs, used in zip(step.call_observed_actual, step.call_used_forecast)):
self._disable_forecasting("solver step finished with an incomplete model-hook call")
if 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]:
self._disable_forecasting("forecasted solver step batch layout changed within one solver step")
if observed_actual and used_forecast_any:
self._disable_forecasting("solver step mixed forecasted and actual model-hook paths")
+33
View File
@@ -327,6 +327,38 @@ def test_split_forecast_step_uses_spectrum_across_subcalls() -> None:
runtime.end_run(run_id)
def test_split_forecast_step_without_layout_labels_uses_sequential_full_prediction() -> None:
runtime = make_runtime()
sample_sigmas = torch.linspace(1.0, 0.0, 51)
run_id = runtime.start_run(sample_sigmas, "sample_euler", supports_solver_steps=True)
total_steps = len(sample_sigmas) - 1
for step_id in range(5):
runtime.begin_solver_step(
run_id,
step_id,
runtime.time_coord_for_step(step_id),
total_steps,
)
call_id = runtime.register_model_hook_call(run_id, step_id, expected_shape=(2, 8, 4), branch_signature=None)
runtime.observe_actual_feature(run_id, step_id, torch.randn(2, 8, 4), call_id=call_id)
runtime.finalize_solver_step(run_id, step_id, used_forecast=False)
decision = runtime.begin_solver_step(
run_id,
5,
runtime.time_coord_for_step(5),
total_steps,
)
assert decision["actual_forward"] is False
first_id = runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 8, 4), branch_signature=None)
assert runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=first_id) is not None
second_id = runtime.register_model_hook_call(run_id, 5, expected_shape=(1, 8, 4), branch_signature=None)
assert runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=second_id) is not None
runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=True)
runtime.end_run(run_id)
def test_duplicate_batch_labels_can_be_reordered_for_forecast() -> None:
runtime = make_runtime()
sample_sigmas = torch.linspace(1.0, 0.0, 51)
@@ -541,6 +573,7 @@ def main() -> None:
test_nonbatch_shape_mismatch_disables_forecast()
test_multiple_hook_calls_are_aggregated_on_actual_steps()
test_split_forecast_step_uses_spectrum_across_subcalls()
test_split_forecast_step_without_layout_labels_uses_sequential_full_prediction()
test_duplicate_batch_labels_can_be_reordered_for_forecast()
test_topology_change_disables_forecast()
test_mixed_batch_layout_presence_disables_forecast()