Merge pull request #6 from xmarre/codex/fix-split-forecast-path

Fix split forecast handling for unlabeled subcalls
This commit is contained in:
xmarre
2026-03-30 23:50:31 +02:00
committed by GitHub
2 changed files with 120 additions and 18 deletions
+41 -18
View File
@@ -51,6 +51,8 @@ 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
used_forecast_any: bool = False
class SpectrumRuntime:
@@ -122,6 +124,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
@@ -254,7 +257,7 @@ class SpectrumRuntime:
def step_used_forecast(self, run_id: int, solver_step_id: int) -> bool:
step = self._require_active_step(run_id, solver_step_id)
return any(step.call_used_forecast)
return step.used_forecast_any or any(step.call_used_forecast)
@staticmethod
def _split_branch_signature(
@@ -367,6 +370,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
if history_shape is not None:
history_shape = tuple(history_shape)
if history_shape[1:] != target_shape[1:]:
@@ -374,15 +378,18 @@ 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)
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,29 @@ 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:
self._disable_forecasting("forecasted solver step batch layout changed within one solver step")
return None
else:
predicted_feature = self.forecaster.predict(
time_coord=step.time_coord,
@@ -417,6 +433,7 @@ class SpectrumRuntime:
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]
def abort_solver_step(self, run_id: int, solver_step_id: int) -> None:
@@ -425,6 +442,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:
@@ -432,7 +450,7 @@ class SpectrumRuntime:
requested_actual_forward = bool(step.decision["actual_forward"])
observed_actual = any(step.call_observed_actual)
used_forecast_any = any(step.call_used_forecast)
used_forecast_any = step.used_forecast_any or any(step.call_used_forecast)
if bool(used_forecast) and not observed_actual:
used_forecast_any = True
if step.decision["actual_forward"] and not observed_actual:
@@ -441,10 +459,15 @@ 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 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]:
self._disable_forecasting("forecasted solver step batch layout changed within one solver step")
if used_forecast_any:
self.num_consecutive_cached_steps += 1
+79
View File
@@ -327,6 +327,42 @@ def test_split_forecast_step_uses_spectrum_across_subcalls() -> None:
runtime.end_run(run_id)
def test_split_forecast_step_without_layout_labels_falls_back_to_actual() -> 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 None
runtime.observe_actual_feature(run_id, 5, torch.randn(1, 8, 4), call_id=first_id)
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 None
runtime.observe_actual_feature(run_id, 5, torch.randn(1, 8, 4), call_id=second_id)
runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False)
assert runtime.stats.forecast_disabled is False
assert runtime.stats.forecasted_count == 0
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)
@@ -416,6 +452,47 @@ def test_mixed_batch_layout_presence_disables_forecast() -> None:
runtime.end_run(run_id)
def test_split_forecast_failure_after_first_slice_still_counts_as_mixed_path() -> 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,
)
first_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)),)
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(1, 8, 4), call_id=first_id)
second_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)),)
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(1, 8, 4), call_id=second_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,
)
first_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)),)
)
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 None
runtime.observe_actual_feature(run_id, 5, torch.ones(1, 8, 4), call_id=second_id)
runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False)
assert runtime.stats.disable_reason == "solver step mixed forecasted and actual model-hook paths"
runtime.end_run(run_id)
def test_aborted_solver_step_is_discarded_without_disabling_forecast() -> None:
runtime = make_runtime()
sample_sigmas = torch.linspace(1.0, 0.0, 51)
@@ -541,9 +618,11 @@ 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_falls_back_to_actual()
test_duplicate_batch_labels_can_be_reordered_for_forecast()
test_topology_change_disables_forecast()
test_mixed_batch_layout_presence_disables_forecast()
test_split_forecast_failure_after_first_slice_still_counts_as_mixed_path()
test_aborted_solver_step_is_discarded_without_disabling_forecast()
test_nonuniform_schedule_coords_are_used()
test_forecaster_respects_nonuniform_coords()