Fix split forecast chunk label handling

This commit is contained in:
xmarre
2026-03-31 00:46:01 +02:00
parent 1c8f9244ad
commit b6b80d0b79
3 changed files with 99 additions and 47 deletions
+4 -1
View File
@@ -103,7 +103,10 @@ def _build_branch_signature(transformer_options: Dict[str, Any]) -> Optional[tup
uuids = transformer_options.get("uuids")
if uuids is not None:
signature.append(("uuids_len", len(uuids)))
try:
signature.append(("uuids", tuple(uuids)))
except Exception:
signature.append(("uuids", tuple(str(u) for u in uuids)))
if not signature:
return None
+44 -28
View File
@@ -44,13 +44,13 @@ class _ActiveStep:
hook_call_count: int = 0
call_expected_shapes: list[tuple[int, ...]] = field(default_factory=list)
call_branch_signatures: list[Optional[tuple[Any, ...]]] = field(default_factory=list)
call_batch_labels: list[Optional[tuple[int, ...]]] = field(default_factory=list)
call_batch_labels: list[Optional[tuple[Any, ...]]] = field(default_factory=list)
call_observed_actual: list[bool] = field(default_factory=list)
call_used_forecast: list[bool] = field(default_factory=list)
call_actual_features: list[Optional[torch.Tensor]] = field(default_factory=list)
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_row_positions: Optional[dict[Any, deque[int]]] = None
prediction_next_row: int = 0
used_forecast_any: bool = False
@@ -67,7 +67,7 @@ class SpectrumRuntime:
self.stats = RuntimeStats(current_window=float(self.cfg.window_size))
self._active_run: Optional[_ActiveRun] = None
self._active_steps: Dict[int, _ActiveStep] = {}
self._history_batch_labels: Optional[tuple[int, ...]] = None
self._history_batch_labels: Optional[tuple[Any, ...]] = None
self._reset_scheduler_state()
@property
@@ -262,20 +262,34 @@ class SpectrumRuntime:
@staticmethod
def _split_branch_signature(
branch_signature: Optional[tuple[Any, ...]],
) -> tuple[tuple[Any, ...], Optional[tuple[int, ...]]]:
) -> tuple[tuple[Any, ...], Optional[tuple[Any, ...]], bool]:
if branch_signature is None:
return (), None
return (), None, False
topology_entries: list[Any] = []
batch_labels: Optional[tuple[int, ...]] = None
cond_labels: Optional[tuple[Any, ...]] = None
uuids: Optional[tuple[Any, ...]] = None
for entry in branch_signature:
if isinstance(entry, tuple) and len(entry) == 2 and entry[0] == "cond_or_uncond":
try:
batch_labels = tuple(int(v) for v in entry[1])
cond_labels = tuple(int(v) for v in entry[1])
except Exception:
batch_labels = tuple(entry[1])
cond_labels = tuple(entry[1])
elif isinstance(entry, tuple) and len(entry) == 2 and entry[0] == "uuids":
try:
uuids = tuple(entry[1])
except Exception:
uuids = tuple(str(v) for v in entry[1])
else:
topology_entries.append(entry)
return tuple(topology_entries), batch_labels
if cond_labels is not None and uuids is not None and len(cond_labels) == len(uuids):
batch_labels = tuple((cond_labels[i], uuids[i]) for i in range(len(cond_labels)))
return tuple(topology_entries), batch_labels, True
if cond_labels is not None:
return tuple(topology_entries), cond_labels, False
if uuids is not None:
return tuple(topology_entries), tuple(("uuid", u) for u in uuids), True
return tuple(topology_entries), None, False
def register_model_hook_call(
self,
@@ -288,7 +302,7 @@ class SpectrumRuntime:
step = self._require_active_step(run_id, solver_step_id)
shape = tuple(expected_shape)
tail_shape = shape[1:]
topology_signature, batch_labels = self._split_branch_signature(branch_signature)
topology_signature, batch_labels, can_expand_batch_labels = self._split_branch_signature(branch_signature)
step.hook_call_count += 1
if step.feature_tail_shape is None:
step.feature_tail_shape = tail_shape
@@ -300,8 +314,14 @@ class SpectrumRuntime:
elif topology_signature != step.topology_signature:
self._disable_forecasting("model-hook branch signature changed within one solver step")
if batch_labels is not None and len(batch_labels) != shape[0]:
batch_labels = None
if batch_labels is not None:
if len(batch_labels) == shape[0]:
pass
elif can_expand_batch_labels and len(batch_labels) > 0 and shape[0] % len(batch_labels) == 0:
rows_per_label = shape[0] // len(batch_labels)
batch_labels = tuple(label for label in batch_labels for _ in range(rows_per_label))
else:
batch_labels = None
step.call_expected_shapes.append(shape)
step.call_branch_signatures.append(branch_signature)
@@ -324,14 +344,15 @@ class SpectrumRuntime:
resolved_call_id = self._resolve_call_id(step, call_id)
step.call_observed_actual[resolved_call_id] = True
step.call_used_forecast[resolved_call_id] = False
step.used_forecast_any = any(step.call_used_forecast)
step.call_predicted_features[resolved_call_id] = None
step.call_actual_features[resolved_call_id] = feature.detach()
@staticmethod
def _reorder_feature_to_labels(
feature: torch.Tensor,
source_labels: tuple[int, ...],
target_labels: tuple[int, ...],
source_labels: tuple[Any, ...],
target_labels: tuple[Any, ...],
) -> Optional[torch.Tensor]:
if feature.shape[0] != len(source_labels) or len(source_labels) != len(target_labels):
return None
@@ -339,17 +360,17 @@ class SpectrumRuntime:
return feature
if sorted(source_labels) != sorted(target_labels):
return None
source_positions: dict[int, deque[int]] = defaultdict(deque)
source_positions: dict[Any, deque[int]] = defaultdict(deque)
for idx, label in enumerate(source_labels):
source_positions[int(label)].append(idx)
order = [source_positions[int(label)].popleft() for label in target_labels]
source_positions[label].append(idx)
order = [source_positions[label].popleft() for label in target_labels]
return feature[order, ...]
@staticmethod
def _build_label_positions(labels: tuple[int, ...]) -> dict[int, deque[int]]:
positions: dict[int, deque[int]] = defaultdict(deque)
def _build_label_positions(labels: tuple[Any, ...]) -> dict[Any, deque[int]]:
positions: dict[Any, deque[int]] = defaultdict(deque)
for idx, label in enumerate(labels):
positions[int(label)].append(idx)
positions[label].append(idx)
return positions
def predict_feature(
@@ -407,20 +428,15 @@ class SpectrumRuntime:
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 = []
for label in target_batch_labels:
positions = step.prediction_row_positions.get(int(label))
positions = step.prediction_row_positions.get(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(
@@ -483,7 +499,7 @@ class SpectrumRuntime:
if any(labeled_parts) and not all(labeled_parts):
self._disable_forecasting("model-hook batch layout changed within one solver step")
elif all(labeled_parts):
rows: list[tuple[int, int, torch.Tensor]] = []
rows: list[tuple[Any, int, torch.Tensor]] = []
arrival = 0
for labels, part in zip(actual_labels, actual_parts):
assert labels is not None
@@ -491,7 +507,7 @@ class SpectrumRuntime:
self._disable_forecasting("model-hook batch layout changed within one solver step")
break
for row_idx, label in enumerate(labels):
rows.append((int(label), arrival, part[row_idx : row_idx + 1]))
rows.append((label, arrival, part[row_idx : row_idx + 1]))
arrival += 1
if rows:
rows.sort(key=lambda item: (item[0], item[1]))
+51 -18
View File
@@ -54,7 +54,7 @@ def test_solver_step_scheduler() -> None:
run_id,
step_id,
expected_shape=(1, 8, 4),
branch_signature=(("cond_or_uncond", (0, 1)),),
branch_signature=(("cond_or_uncond", (0, 1)), ("uuids", ("u0", "u1"))),
)
runtime.observe_actual_feature(decision["run_id"], decision["solver_step_id"], torch.randn(1, 8, 4))
runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False)
@@ -69,7 +69,7 @@ def test_solver_step_scheduler() -> None:
run_id,
5,
expected_shape=(1, 8, 4),
branch_signature=(("cond_or_uncond", (0, 1)),),
branch_signature=(("cond_or_uncond", (0, 1)), ("uuids", ("u0", "u1"))),
)
predicted = runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4))
assert predicted is not None
@@ -184,11 +184,11 @@ def test_batch_split_falls_back_to_actual_without_disabling_run() -> None:
total_steps,
)
first_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)),)
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)), ("uuids", ("u1",)))
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(1, 8, 4) * 10.0, 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,)),)
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)), ("uuids", ("u0",)))
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(1, 8, 4) * 20.0, call_id=second_id)
runtime.finalize_solver_step(run_id, step_id, used_forecast=False)
@@ -201,7 +201,7 @@ def test_batch_split_falls_back_to_actual_without_disabling_run() -> None:
)
assert decision["actual_forward"] is False
call_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(2, 8, 4), branch_signature=(("cond_or_uncond", (0, 1)),)
run_id, 5, expected_shape=(2, 8, 4), branch_signature=(("cond_or_uncond", (0, 1)), ("uuids", ("u0", "u1")))
)
predicted = runtime.predict_feature(run_id, 5, expected_shape=(2, 8, 4), call_id=call_id)
assert predicted is not None
@@ -290,11 +290,11 @@ def test_split_forecast_step_uses_spectrum_across_subcalls() -> None:
total_steps,
)
first_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)),)
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)), ("uuids", ("u1",)))
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(1, 8, 4) * 10.0, 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,)),)
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)), ("uuids", ("u0",)))
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(1, 8, 4) * 20.0, call_id=second_id)
runtime.finalize_solver_step(run_id, step_id, used_forecast=False)
@@ -308,14 +308,14 @@ def test_split_forecast_step_uses_spectrum_across_subcalls() -> None:
assert decision["actual_forward"] is False
first_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)),)
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)), ("uuids", ("u0",)))
)
first_pred = runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=first_id)
assert first_pred is not None
assert first_pred.shape == (1, 8, 4)
second_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)),)
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)), ("uuids", ("u1",)))
)
second_pred = runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=second_id)
assert second_pred is not None
@@ -363,6 +363,38 @@ def test_split_forecast_step_without_layout_labels_falls_back_to_actual() -> Non
runtime.end_run(run_id)
def test_chunk_level_cond_or_uncond_and_uuid_labels_expand_to_rows_for_split_forecast() -> 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=(4, 8, 4),
branch_signature=(("cond_or_uncond", (1, 0)), ("uuids", ("u1", "u0"))),
)
runtime.observe_actual_feature(run_id, step_id, torch.randn(4, 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=(2, 8, 4), branch_signature=(("cond_or_uncond", (0,)), ("uuids", ("u0",))),
)
assert runtime.predict_feature(run_id, 5, expected_shape=(2, 8, 4), call_id=first_id) is not None
second_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(2, 8, 4), branch_signature=(("cond_or_uncond", (1,)), ("uuids", ("u1",))),
)
assert runtime.predict_feature(run_id, 5, expected_shape=(2, 8, 4), call_id=second_id) is not None
runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=True)
assert runtime.stats.forecasted_count == 1
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)
@@ -377,11 +409,11 @@ def test_duplicate_batch_labels_can_be_reordered_for_forecast() -> None:
total_steps,
)
first_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(2, 8, 4), branch_signature=(("cond_or_uncond", (1, 1)),)
run_id, step_id, expected_shape=(2, 8, 4), branch_signature=(("cond_or_uncond", (1, 1)), ("uuids", ("u1a", "u1b")))
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(2, 8, 4) * 10.0, call_id=first_id)
second_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(2, 8, 4), branch_signature=(("cond_or_uncond", (0, 0)),)
run_id, step_id, expected_shape=(2, 8, 4), branch_signature=(("cond_or_uncond", (0, 0)), ("uuids", ("u0a", "u0b")))
)
runtime.observe_actual_feature(run_id, step_id, torch.ones(2, 8, 4) * 20.0, call_id=second_id)
runtime.finalize_solver_step(run_id, step_id, used_forecast=False)
@@ -393,7 +425,7 @@ def test_duplicate_batch_labels_can_be_reordered_for_forecast() -> None:
total_steps,
)
call_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(4, 8, 4), branch_signature=(("cond_or_uncond", (0, 0, 1, 1)),)
run_id, 5, expected_shape=(4, 8, 4), branch_signature=(("cond_or_uncond", (0, 0, 1, 1)), ("uuids", ("u0a", "u0b", "u1a", "u1b")))
)
predicted = runtime.predict_feature(run_id, 5, expected_shape=(4, 8, 4), call_id=call_id)
assert predicted is not None
@@ -415,11 +447,11 @@ def test_topology_change_disables_forecast() -> None:
total_steps,
)
first_id = runtime.register_model_hook_call(
run_id, 0, expected_shape=(1, 8, 4), branch_signature=(("uuids_len", 1), ("cond_or_uncond", (0,)))
run_id, 0, expected_shape=(1, 8, 4), branch_signature=(("hooks_id", 1), ("cond_or_uncond", (0,)), ("uuids", ("u0",)))
)
runtime.observe_actual_feature(run_id, 0, torch.randn(1, 8, 4), call_id=first_id)
second_id = runtime.register_model_hook_call(
run_id, 0, expected_shape=(1, 8, 4), branch_signature=(("uuids_len", 2), ("cond_or_uncond", (1,)))
run_id, 0, expected_shape=(1, 8, 4), branch_signature=(("hooks_id", 2), ("cond_or_uncond", (1,)), ("uuids", ("u1",)))
)
assert runtime.stats.forecast_disabled is True
assert runtime.stats.disable_reason == "model-hook branch signature changed within one solver step"
@@ -443,7 +475,7 @@ def test_mixed_batch_layout_presence_disables_forecast() -> None:
first_id = runtime.register_model_hook_call(run_id, 0, expected_shape=(1, 8, 4), branch_signature=None)
runtime.observe_actual_feature(run_id, 0, torch.randn(1, 8, 4), call_id=first_id)
second_id = runtime.register_model_hook_call(
run_id, 0, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)),)
run_id, 0, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)), ("uuids", ("u1",)))
)
runtime.observe_actual_feature(run_id, 0, torch.randn(1, 8, 4), call_id=second_id)
runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=False)
@@ -466,11 +498,11 @@ def test_split_forecast_failure_after_first_slice_still_counts_as_mixed_path() -
total_steps,
)
first_id = runtime.register_model_hook_call(
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)),)
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (1,)), ("uuids", ("u1",)))
)
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,)),)
run_id, step_id, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)), ("uuids", ("u0",)))
)
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)
@@ -482,7 +514,7 @@ def test_split_forecast_failure_after_first_slice_still_counts_as_mixed_path() -
total_steps,
)
first_id = runtime.register_model_hook_call(
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)),)
run_id, 5, expected_shape=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)), ("uuids", ("u0",)))
)
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)
@@ -619,6 +651,7 @@ def main() -> None:
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_chunk_level_cond_or_uncond_and_uuid_labels_expand_to_rows_for_split_forecast()
test_duplicate_batch_labels_can_be_reordered_for_forecast()
test_topology_change_disables_forecast()
test_mixed_batch_layout_presence_disables_forecast()