pablovela5620 Claude Fable 5 commited on
Commit
326f401
·
1 Parent(s): 3ec40ad

Make Stop actually halt the pipeline via a run-scoped sentinel

Browse files

Gradio's cancels only closes the request generator; the pipeline keeps
running on a worker thread locally and in a forked child on ZeroGPU.
Stop now writes a sentinel file under the run's data dir, and the
motion/denoise hooks — the only code the pipeline re-enters — raise
RunCancelled when they see it.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>

Files changed (2) hide show
  1. fdanyone_app.py +42 -1
  2. tests/test_app_helpers.py +38 -0
fdanyone_app.py CHANGED
@@ -522,6 +522,28 @@ class Session:
522
  """Whether the last ``_pump`` call's pipeline thread returned."""
523
 
524
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
525
  T = TypeVar("T")
526
 
527
 
@@ -570,6 +592,7 @@ def _motion_hook(link: Link, session: Session) -> Callable[[str, dict[str, objec
570
  recording: rr.RecordingStream = link.recording
571
 
572
  def on_stage(stage: str, payload: dict[str, object]) -> None:
 
573
  if stage == "bboxes":
574
  boxes: Float[torch.Tensor, "frames 4"] = payload["bbx_xyxy"] # pyrefly: ignore
575
  for index, box in enumerate(boxes.numpy()[: INFERENCE.num_frames]):
@@ -604,6 +627,7 @@ def _denoise_hook(link: Link, session: Session) -> Callable[[int, tuple[int, ...
604
  view_indices: tuple[int, ...],
605
  x0_hat: Float[torch.Tensor, "views 48 latent_t latent_h latent_w"],
606
  ) -> None:
 
607
  if PREVIEW_DECODER is None:
608
  return
609
  plan: tuple[tuple[int, int], ...] = preview_slice_plan(x0_hat.shape[2])
@@ -844,6 +868,7 @@ def smoke_motion_phase(link: Link, session: Session) -> Iterator[str]:
844
 
845
  rng: np.random.Generator = np.random.default_rng(0)
846
  for index in range(0, 121, 10):
 
847
  _set_frame(link.recording, index, session.spec.fps)
848
  x0: float = 180.0 + 2.0 * index
849
  link.recording.log(
@@ -864,6 +889,7 @@ def smoke_generate_phase(link: Link, session: Session) -> Iterator[str]:
864
  link.recording.send_blueprint(diffusion_blueprint(), make_active=True)
865
  yield "[smoke] diffusion"
866
  for step in range(4):
 
867
  link.recording.set_time(DIFFUSION_TIMELINE, sequence=step)
868
  for view in range(VIEWS):
869
  noise: UInt8[np.ndarray, "160 88 3"] = rng.integers(
@@ -947,6 +973,16 @@ def run_gpu(spec: RunSpec) -> Iterator[tuple[dict | None, bytes | None, str]]:
947
  yield session.summary, link.read(), "Publishing the result."
948
 
949
 
 
 
 
 
 
 
 
 
 
 
950
  def publish_cpu(spec: RunSpec, summary: dict | None) -> Iterator[tuple[bytes | None, str]]:
951
  """Attach the finished rig and its six videos, off the GPU allocation."""
952
 
@@ -1052,7 +1088,12 @@ def build_demo() -> gr.Blocks:
1052
  prepared = started.success(prepare_cpu, spec, [viewer, status])
1053
  generated = prepared.success(run_gpu, spec, [summary, viewer, status])
1054
  published = generated.success(publish_cpu, [spec, summary], [viewer, status])
1055
- stop_button.click(None, cancels=[started, prepared, generated, published])
 
 
 
 
 
1056
  return demo
1057
 
1058
 
 
522
  """Whether the last ``_pump`` call's pipeline thread returned."""
523
 
524
 
525
+ class RunCancelled(FourDAnyoneError):
526
+ """Raised inside a pipeline hook once the run's stop sentinel appears."""
527
+
528
+
529
+ def _stop_path(spec: RunSpec) -> Path:
530
+ return spec.data_dir / "stop-requested"
531
+
532
+
533
+ def _check_stop(spec: RunSpec) -> None:
534
+ """Abort the pipeline from inside one of its hooks after Stop is pressed.
535
+
536
+ Gradio's ``cancels`` only closes the request's generator; the pipeline keeps
537
+ running — on a worker thread here, in a forked ZeroGPU child on the Space.
538
+ A file under the run's own directory is the one channel that reaches both,
539
+ and the hooks are the only code of ours the pipeline re-enters, so they are
540
+ where the run can be unwound.
541
+ """
542
+
543
+ if _stop_path(spec).exists():
544
+ raise RunCancelled("Stopped.")
545
+
546
+
547
  T = TypeVar("T")
548
 
549
 
 
592
  recording: rr.RecordingStream = link.recording
593
 
594
  def on_stage(stage: str, payload: dict[str, object]) -> None:
595
+ _check_stop(session.spec)
596
  if stage == "bboxes":
597
  boxes: Float[torch.Tensor, "frames 4"] = payload["bbx_xyxy"] # pyrefly: ignore
598
  for index, box in enumerate(boxes.numpy()[: INFERENCE.num_frames]):
 
627
  view_indices: tuple[int, ...],
628
  x0_hat: Float[torch.Tensor, "views 48 latent_t latent_h latent_w"],
629
  ) -> None:
630
+ _check_stop(session.spec)
631
  if PREVIEW_DECODER is None:
632
  return
633
  plan: tuple[tuple[int, int], ...] = preview_slice_plan(x0_hat.shape[2])
 
868
 
869
  rng: np.random.Generator = np.random.default_rng(0)
870
  for index in range(0, 121, 10):
871
+ _check_stop(session.spec)
872
  _set_frame(link.recording, index, session.spec.fps)
873
  x0: float = 180.0 + 2.0 * index
874
  link.recording.log(
 
889
  link.recording.send_blueprint(diffusion_blueprint(), make_active=True)
890
  yield "[smoke] diffusion"
891
  for step in range(4):
892
+ _check_stop(session.spec)
893
  link.recording.set_time(DIFFUSION_TIMELINE, sequence=step)
894
  for view in range(VIEWS):
895
  noise: UInt8[np.ndarray, "160 88 3"] = rng.integers(
 
973
  yield session.summary, link.read(), "Publishing the result."
974
 
975
 
976
+ def request_stop(spec: RunSpec | None) -> str:
977
+ """Write the run's stop sentinel; the pipeline unwinds at its next hook."""
978
+
979
+ if spec is None:
980
+ return "Nothing is running."
981
+ spec.data_dir.mkdir(parents=True, exist_ok=True)
982
+ _stop_path(spec).touch()
983
+ return "Stopping — the pipeline halts at its next step."
984
+
985
+
986
  def publish_cpu(spec: RunSpec, summary: dict | None) -> Iterator[tuple[bytes | None, str]]:
987
  """Attach the finished rig and its six videos, off the GPU allocation."""
988
 
 
1088
  prepared = started.success(prepare_cpu, spec, [viewer, status])
1089
  generated = prepared.success(run_gpu, spec, [summary, viewer, status])
1090
  published = generated.success(publish_cpu, [spec, summary], [viewer, status])
1091
+ # `cancels` detaches the UI at once, but the pipeline itself only halts
1092
+ # when it reads the sentinel `request_stop` writes: the worker thread
1093
+ # (and on ZeroGPU, the forked child) never sees a Gradio cancellation.
1094
+ stop_button.click(
1095
+ request_stop, spec, status, cancels=[started, prepared, generated, published]
1096
+ )
1097
  return demo
1098
 
1099
 
tests/test_app_helpers.py CHANGED
@@ -110,6 +110,44 @@ def test_pump_yields_hook_labels_and_returns_the_work_result() -> None:
110
  assert session.worker_finished is True
111
 
112
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
113
  def test_pump_reraises_a_worker_failure_on_the_caller_thread() -> None:
114
  """A pipeline error must surface where Gradio can turn it into a message."""
115
 
 
110
  assert session.worker_finished is True
111
 
112
 
113
+ def test_stop_sentinel_unwinds_the_pump_worker(tmp_path: Path) -> None:
114
+ """Stop must reach the pipeline thread itself, not just Gradio's generator.
115
+
116
+ ``cancels`` closes the request generator, but the worker thread (or the
117
+ forked ZeroGPU child) never hears it; the sentinel checked by the hooks is
118
+ the only thing that actually frees the GPU.
119
+ """
120
+
121
+ session: fdanyone_app.Session = _pump_session()
122
+ spec: fdanyone_app.RunSpec = fdanyone_app.RunSpec(
123
+ token=session.spec.token,
124
+ video_path=session.spec.video_path,
125
+ start_time=0.0,
126
+ seed=0,
127
+ fps=Fraction(25, 1),
128
+ data_dir=tmp_path,
129
+ )
130
+ session = fdanyone_app.Session(spec)
131
+ assert fdanyone_app.request_stop(None) == "Nothing is running."
132
+ fdanyone_app.request_stop(spec)
133
+ assert (tmp_path / "stop-requested").exists()
134
+
135
+ ticks: list[int] = []
136
+
137
+ def work() -> str:
138
+ # What every pipeline hook does on entry, once per stage or step.
139
+ for index in range(1000):
140
+ fdanyone_app._check_stop(spec)
141
+ ticks.append(index)
142
+ return "finished"
143
+
144
+ with pytest.raises(fdanyone_app.RunCancelled):
145
+ list(fdanyone_app._pump(session, work))
146
+ assert ticks == []
147
+ # The worker joined before re-raising, so the thread is truly gone.
148
+ assert session.worker_finished is True
149
+
150
+
151
  def test_pump_reraises_a_worker_failure_on_the_caller_thread() -> None:
152
  """A pipeline error must surface where Gradio can turn it into a message."""
153