Fix split forecast chunk label handling
This commit is contained in:
@@ -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
@@ -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
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user