From b6b80d0b7964ae81efda3daba5154d5fcb58e86e Mon Sep 17 00:00:00 2001 From: xmarre Date: Tue, 31 Mar 2026 00:46:01 +0200 Subject: [PATCH] Fix split forecast chunk label handling --- comfyui_spectrum/flux.py | 5 ++- comfyui_spectrum/runtime.py | 72 ++++++++++++++++++++++--------------- tests/smoke_runtime.py | 69 +++++++++++++++++++++++++---------- 3 files changed, 99 insertions(+), 47 deletions(-) diff --git a/comfyui_spectrum/flux.py b/comfyui_spectrum/flux.py index e857624..dbf2952 100644 --- a/comfyui_spectrum/flux.py +++ b/comfyui_spectrum/flux.py @@ -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 diff --git a/comfyui_spectrum/runtime.py b/comfyui_spectrum/runtime.py index 5af8378..18bf27f 100644 --- a/comfyui_spectrum/runtime.py +++ b/comfyui_spectrum/runtime.py @@ -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])) diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index 6bb25cb..3d6d2ed 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -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()