Preserve mixed-path forecast state after split-step failures
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user