Merge pull request #5 from xmarre/codex/fix-split-forecast-steps
Fix split forecast solver steps and abort lifecycle handling
This commit is contained in:
@@ -229,16 +229,25 @@ def _install_sampler_level_wrappers(model: Any, runtime: SpectrumRuntime) -> Non
|
||||
patched_model_options = _copy_model_options_with_step_context(
|
||||
effective_model_options, wrapper_runtime, decision
|
||||
)
|
||||
step_aborted = False
|
||||
try:
|
||||
return executor(x, timestep, patched_model_options, seed)
|
||||
finally:
|
||||
wrapper_runtime.finalize_solver_step(
|
||||
except BaseException:
|
||||
step_aborted = True
|
||||
wrapper_runtime.abort_solver_step(
|
||||
decision["run_id"],
|
||||
decision["solver_step_id"],
|
||||
used_forecast=wrapper_runtime.step_used_forecast(
|
||||
decision["run_id"], decision["solver_step_id"]
|
||||
),
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
if not step_aborted:
|
||||
wrapper_runtime.finalize_solver_step(
|
||||
decision["run_id"],
|
||||
decision["solver_step_id"],
|
||||
used_forecast=wrapper_runtime.step_used_forecast(
|
||||
decision["run_id"], decision["solver_step_id"]
|
||||
),
|
||||
)
|
||||
|
||||
comfy.patcher_extension.add_wrapper(
|
||||
comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
|
||||
|
||||
+53
-18
@@ -49,6 +49,8 @@ class _ActiveStep:
|
||||
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
|
||||
|
||||
|
||||
class SpectrumRuntime:
|
||||
@@ -118,6 +120,8 @@ class SpectrumRuntime:
|
||||
for step in self._active_steps.values():
|
||||
step.call_predicted_features = [None] * len(step.call_predicted_features)
|
||||
step.call_used_forecast = [False] * len(step.call_used_forecast)
|
||||
step.predicted_full_feature = None
|
||||
step.prediction_row_positions = None
|
||||
self.stats.current_window = self.curr_ws
|
||||
self.stats.forecast_disabled = True
|
||||
self.stats.disable_reason = reason
|
||||
@@ -338,6 +342,13 @@ class SpectrumRuntime:
|
||||
order = [source_positions[int(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)
|
||||
for idx, label in enumerate(labels):
|
||||
positions[int(label)].append(idx)
|
||||
return positions
|
||||
|
||||
def predict_feature(
|
||||
self,
|
||||
run_id: int,
|
||||
@@ -361,37 +372,61 @@ class SpectrumRuntime:
|
||||
if history_shape[1:] != target_shape[1:]:
|
||||
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
|
||||
return None
|
||||
if history_shape[0] != target_shape[0]:
|
||||
if self._history_batch_labels is None and history_shape[0] != target_shape[0]:
|
||||
return None
|
||||
|
||||
if step.hook_call_count > 1:
|
||||
if self._history_batch_labels is None and step.hook_call_count > 1:
|
||||
return None
|
||||
|
||||
if self._history_batch_labels is not None and target_batch_labels is None:
|
||||
return None
|
||||
|
||||
if step.call_predicted_features[resolved_call_id] is None:
|
||||
predicted_feature = self.forecaster.predict(
|
||||
time_coord=step.time_coord,
|
||||
blend_weight=self.cfg.blend_weight,
|
||||
)
|
||||
if tuple(predicted_feature.shape) != target_shape:
|
||||
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
|
||||
return None
|
||||
if self._history_batch_labels is not None and target_batch_labels is not None:
|
||||
reordered = self._reorder_feature_to_labels(
|
||||
predicted_feature,
|
||||
self._history_batch_labels,
|
||||
target_batch_labels,
|
||||
if self._history_batch_labels is not None:
|
||||
if step.predicted_full_feature is None:
|
||||
predicted_full_feature = self.forecaster.predict(
|
||||
time_coord=step.time_coord,
|
||||
blend_weight=self.cfg.blend_weight,
|
||||
)
|
||||
if history_shape is not None and tuple(predicted_full_feature.shape) != history_shape:
|
||||
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
|
||||
return None
|
||||
step.predicted_full_feature = predicted_full_feature
|
||||
step.prediction_row_positions = self._build_label_positions(self._history_batch_labels)
|
||||
|
||||
if target_batch_labels is None or step.prediction_row_positions is None:
|
||||
return None
|
||||
|
||||
order = []
|
||||
for label in target_batch_labels:
|
||||
positions = step.prediction_row_positions.get(int(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:
|
||||
predicted_feature = self.forecaster.predict(
|
||||
time_coord=step.time_coord,
|
||||
blend_weight=self.cfg.blend_weight,
|
||||
)
|
||||
if reordered is None:
|
||||
if tuple(predicted_feature.shape) != target_shape:
|
||||
self._disable_forecasting("predicted feature shape did not match the current solver-step input")
|
||||
return None
|
||||
predicted_feature = reordered
|
||||
step.call_predicted_features[resolved_call_id] = predicted_feature
|
||||
|
||||
step.call_used_forecast[resolved_call_id] = True
|
||||
return step.call_predicted_features[resolved_call_id]
|
||||
|
||||
def abort_solver_step(self, run_id: int, solver_step_id: int) -> None:
|
||||
step = self._require_active_step(run_id, solver_step_id)
|
||||
step.call_predicted_features = [None] * len(step.call_predicted_features)
|
||||
step.call_used_forecast = [False] * len(step.call_used_forecast)
|
||||
step.predicted_full_feature = None
|
||||
step.prediction_row_positions = None
|
||||
self._active_steps.pop(int(solver_step_id), None)
|
||||
|
||||
def finalize_solver_step(self, run_id: int, solver_step_id: int, *, used_forecast: bool) -> None:
|
||||
step = self._require_active_step(run_id, solver_step_id)
|
||||
requested_actual_forward = bool(step.decision["actual_forward"])
|
||||
@@ -404,14 +439,14 @@ class SpectrumRuntime:
|
||||
self._disable_forecasting("solver step requested an actual forward but no actual feature was observed")
|
||||
if not observed_actual and not used_forecast_any:
|
||||
self._disable_forecasting("solver step finished without an actual feature or a forecasted feature")
|
||||
if any(not (obs or used) for obs, used in zip(step.call_observed_actual, step.call_used_forecast)):
|
||||
self._disable_forecasting("solver step finished with an incomplete model-hook call")
|
||||
|
||||
if observed_actual and used_forecast_any:
|
||||
self._disable_forecasting("solver step mixed forecasted and actual model-hook paths")
|
||||
used_forecast_any = False
|
||||
|
||||
if used_forecast_any:
|
||||
if step.hook_call_count > 1:
|
||||
self._disable_forecasting("forecasted solver step re-entered the model hook")
|
||||
self.num_consecutive_cached_steps += 1
|
||||
self.stats.forecasted_count += 1
|
||||
step.decision["actual_forward"] = False
|
||||
|
||||
@@ -276,6 +276,57 @@ def test_multiple_hook_calls_are_aggregated_on_actual_steps() -> None:
|
||||
runtime.end_run(run_id)
|
||||
|
||||
|
||||
def test_split_forecast_step_uses_spectrum_across_subcalls() -> 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) * 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,)),)
|
||||
)
|
||||
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)
|
||||
|
||||
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=(1, 8, 4), branch_signature=(("cond_or_uncond", (0,)),)
|
||||
)
|
||||
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,)),)
|
||||
)
|
||||
second_pred = runtime.predict_feature(run_id, 5, expected_shape=(1, 8, 4), call_id=second_id)
|
||||
assert second_pred is not None
|
||||
assert second_pred.shape == (1, 8, 4)
|
||||
|
||||
runtime.finalize_solver_step(decision["run_id"], decision["solver_step_id"], used_forecast=True)
|
||||
assert runtime.stats.forecast_disabled is False
|
||||
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)
|
||||
@@ -365,6 +416,23 @@ def test_mixed_batch_layout_presence_disables_forecast() -> None:
|
||||
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)
|
||||
run_id = runtime.start_run(sample_sigmas, "sample_euler", supports_solver_steps=True)
|
||||
total_steps = len(sample_sigmas) - 1
|
||||
|
||||
decision = runtime.begin_solver_step(
|
||||
run_id,
|
||||
0,
|
||||
runtime.time_coord_for_step(0),
|
||||
total_steps,
|
||||
)
|
||||
runtime.abort_solver_step(decision["run_id"], decision["solver_step_id"])
|
||||
assert runtime.stats.forecast_disabled is False
|
||||
runtime.end_run(run_id)
|
||||
|
||||
|
||||
def test_nonuniform_schedule_coords_are_used() -> None:
|
||||
runtime = make_runtime()
|
||||
sample_sigmas = torch.tensor([10.0, 9.0, 1.0, 0.0])
|
||||
@@ -472,9 +540,11 @@ def main() -> None:
|
||||
test_batch_split_falls_back_to_actual_without_disabling_run()
|
||||
test_nonbatch_shape_mismatch_disables_forecast()
|
||||
test_multiple_hook_calls_are_aggregated_on_actual_steps()
|
||||
test_split_forecast_step_uses_spectrum_across_subcalls()
|
||||
test_duplicate_batch_labels_can_be_reordered_for_forecast()
|
||||
test_topology_change_disables_forecast()
|
||||
test_mixed_batch_layout_presence_disables_forecast()
|
||||
test_aborted_solver_step_is_discarded_without_disabling_forecast()
|
||||
test_nonuniform_schedule_coords_are_used()
|
||||
test_forecaster_respects_nonuniform_coords()
|
||||
test_flux_sampler_contract_only_allows_euler()
|
||||
|
||||
Reference in New Issue
Block a user