Preserve mixed-path forecast state after split-step failures

This commit is contained in:
xmarre
2026-03-30 23:01:28 +02:00
parent 083b3d44f0
commit e1ca258a5f
2 changed files with 45 additions and 1 deletions
+3 -1
View File
@@ -52,6 +52,7 @@ class _ActiveStep:
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:
@@ -444,6 +445,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:
@@ -460,7 +462,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:
+42
View File
@@ -448,6 +448,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)
@@ -577,6 +618,7 @@ def main() -> None:
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()