wxDai commited on
Commit
c246c32
·
1 Parent(s): 72c383b

fix: sync streaming and PE fixes while preserving ZeroGPU integration

Browse files

Align shared modules with JoyAI-Video-Edit d88456a, including RV2V prompt
enhancement, streaming cleanup, complete file recordings and API defaults.
Honor explicit static-KV overrides while defaulting file input to off.

Run the shared session handler through the process-queue adapter instead of
maintaining a separate session loop. Keep ZeroGPU acquisition, leases, AOTI,
HF dependency choices, quota UI and team gallery intact.

requirements.txt CHANGED
@@ -18,7 +18,7 @@ opencv-python-headless==4.13.0.92
18
  av==13.1.0
19
  imageio-ffmpeg==0.6.0
20
  websockets==16.0
21
- openai==2.41.0
22
  loguru==0.7.3
23
 
24
  # Prebuilt custom CUDA kernels (cp310, torch 2.9.1 ABI, sm_120) are NOT listed
 
18
  av==13.1.0
19
  imageio-ffmpeg==0.6.0
20
  websockets==16.0
21
+ httpx==0.28.1
22
  loguru==0.7.3
23
 
24
  # Prebuilt custom CUDA kernels (cp310, torch 2.9.1 ABI, sm_120) are NOT listed
static/index.html CHANGED
@@ -417,12 +417,14 @@ const I18N = {
417
  gate_settling: "保持不动,正在对焦",
418
  gate_no_face_wait: "没有检测到人脸,请正对摄像头 😊",
419
  toast_done: "✨ 编辑完成,点右下角气泡下载",
 
 
420
  toast_finalizing: "✨ 编辑完成,正在生成可下载视频",
421
  edit_processing: "✨ 正在编辑",
422
  q_downgrade_auto: " · 已自动降到低画质", q_downgrade_hint: " · 建议降低画质",
423
  flow_drop_tail: " · 带宽不足,正在丢帧保延迟",
424
  up_auto_tail: " · 上行已自动降至中档",
425
- res_notice: "本演示运行在 {res},而非模型原生的 1248×720 分辨率,并且没有提示词增强,画质与编辑效果会有所下降。",
426
  },
427
  en: {
428
  tagline: "Streaming video editing",
@@ -495,12 +497,14 @@ const I18N = {
495
  gate_settling: "Hold still, focusing",
496
  gate_no_face_wait: "No face detected, please face the camera 😊",
497
  toast_done: "✨ Done — click the bubble at bottom-right to download",
 
 
498
  toast_finalizing: "✨ Done — generating the downloadable video",
499
  edit_processing: "✨ Editing",
500
  q_downgrade_auto: " · auto-lowered to low quality", q_downgrade_hint: " · consider lowering quality",
501
  flow_drop_tail: " · bandwidth limited, dropping frames",
502
  up_auto_tail: " · uplink auto-lowered to mid",
503
- res_notice: "This demo runs at {res}, not the model's native 1248×720 resolution, and without prompt enhancement; visual quality and edit fidelity are reduced.",
504
  },
505
  };
506
  let currentLang = localStorage.getItem("joyomni_lang") || "en";
@@ -736,6 +740,8 @@ let tickInFlight = false;
736
  let pePausedSend = false;
737
  let peDeferThisSend = false;
738
  let videoPullActive = false;
 
 
739
  let videoPullIdx = 0;
740
  const videoPullFps = 24;
741
  let videoPullTotal = 0;
@@ -752,6 +758,8 @@ function stopRecordDrain() {
752
  let startingRun = false;
753
  let sessionGranted = false;
754
  let peCache = null;
 
 
755
  let pendingPeCacheKey = null;
756
  let outputPlaybackTimer = null;
757
  let outputPlaybackIntervalMs = null;
@@ -781,10 +789,8 @@ let lastObjectUrl = null;
781
  let alignMaxAgeMs = 0;
782
  let lastShownCaptureMs = 0;
783
  let pendingRevokeUrl = null;
784
- let lastSourceObjectUrl = null;
785
  let refImagePreviewDataUrl = null;
786
  let injectedRefImageDataUrl = null;
787
- let suppressPeThisSend = false;
788
  let outputWaitingFace = false;
789
  let waitHintReason = null;
790
  let waitHintAt = 0;
@@ -838,6 +844,7 @@ function backendPendingFrames() {
838
  }
839
 
840
  function noteBackendAck(msg) {
 
841
  const framesIn = Number(msg && msg.frames_in);
842
  if (Number.isFinite(framesIn)) {
843
  backendAckedFrames = Math.max(backendAckedFrames, framesIn);
@@ -872,6 +879,7 @@ function orientDims(o) {
872
  return o === "portrait" ? { width: short, height: long } : { width: long, height: short };
873
  }
874
  function applyOrientation(o) {
 
875
  const dims = orientDims(o);
876
  document.getElementById("width").value = dims.width;
877
  document.getElementById("height").value = dims.height;
@@ -892,7 +900,14 @@ function maxTemporalIdsValue() {
892
  }
893
 
894
  function peCacheKey(rawPrompt, refImage) {
895
- return JSON.stringify({ prompt: rawPrompt || "", refImage: refImage || "" });
 
 
 
 
 
 
 
896
  }
897
 
898
  function startFrameTimer() {
@@ -913,7 +928,7 @@ function startNetPingTimer() {
913
  if (netPingTimer) clearInterval(netPingTimer);
914
  netPingTimer = setInterval(() => {
915
  if (ws && ws.readyState === WebSocket.OPEN) {
916
- try { ws.send(JSON.stringify({ type: "ping", t: Date.now(), recv: receivedFrames, rtt: netRttLastMs })); } catch (err) {}
917
  }
918
  autoQualityTick();
919
  }, 1000);
@@ -941,40 +956,18 @@ function readRefImageDataUrl() {
941
  });
942
  }
943
 
944
- function applyRefImagePreview() {
945
- if (!refImagePreviewDataUrl) return false;
946
- if (lastSourceObjectUrl) URL.revokeObjectURL(lastSourceObjectUrl);
947
- lastSourceObjectUrl = null;
948
- return true;
949
- }
950
-
951
- function updateSourceFrameVisibility() {
952
- const enabled = refImageSelected() || !!refImagePreviewDataUrl;
953
- if (!enabled) {
954
- if (lastSourceObjectUrl) URL.revokeObjectURL(lastSourceObjectUrl);
955
- lastSourceObjectUrl = null;
956
- return;
957
- }
958
- applyRefImagePreview();
959
- }
960
-
961
  async function handleRefImageChange() {
962
  injectedRefImageDataUrl = null;
963
  if (!refImageSelected()) {
964
  refImagePreviewDataUrl = null;
965
- if (lastSourceObjectUrl) URL.revokeObjectURL(lastSourceObjectUrl);
966
- lastSourceObjectUrl = null;
967
- updateSourceFrameVisibility();
968
  syncRefThumb();
969
  updateMetrics();
970
  return;
971
  }
972
  try {
973
  await readRefImageDataUrl();
974
- applyRefImagePreview();
975
  } catch (err) {
976
  }
977
- updateSourceFrameVisibility();
978
  syncRefThumb();
979
  updateMetrics();
980
  }
@@ -1246,10 +1239,7 @@ function resetRunMetrics() {
1246
  if (pendingRevokeUrl) { URL.revokeObjectURL(pendingRevokeUrl); pendingRevokeUrl = null; }
1247
  outputImg.onload = outputImg.onerror = null;
1248
  outputImg.removeAttribute("src");
1249
- if (lastSourceObjectUrl) URL.revokeObjectURL(lastSourceObjectUrl);
1250
- lastSourceObjectUrl = null;
1251
  hidePeResult();
1252
- updateSourceFrameVisibility();
1253
  sessionProfileTimings = false;
1254
  stageTimings.innerHTML = "";
1255
  stageTimingsPanel.style.display = "none";
@@ -1347,8 +1337,10 @@ function resetUplinkEncoder() {
1347
 
1348
  function ensureUplinkEncoder(width, height) {
1349
  if (upEncoder && upEncoder.state === "configured") return;
 
1350
  upEncoder = new VideoEncoder({
1351
  output: (chunk) => {
 
1352
  const info = upPendingByTs.get(chunk.timestamp);
1353
  upPendingByTs.delete(chunk.timestamp);
1354
  if (!info || !ws || ws.readyState !== WebSocket.OPEN) return;
@@ -1359,7 +1351,7 @@ function ensureUplinkEncoder(width, height) {
1359
  const buf = new ArrayBuffer(chunk.byteLength);
1360
  chunk.copyTo(buf);
1361
  try {
1362
- ws.send(JSON.stringify({ type: "frame_meta", seq: info.seq, t_capture_ms: info.t_capture_ms }));
1363
  ws.send(buf);
1364
  } catch (e) {
1365
  return;
@@ -1399,12 +1391,15 @@ function resetOutputDecoder() {
1399
 
1400
  function feedOutputH264(buf, meta) {
1401
  if (!outDecoder || outDecoder.state === "closed") {
 
1402
  outDecoder = new VideoDecoder({
1403
  output: (frame) => {
 
1404
  const m = outMetaByTs.get(frame.timestamp);
1405
  outMetaByTs.delete(frame.timestamp);
1406
  createImageBitmap(frame).then((bmp) => {
1407
  frame.close();
 
1408
  enqueueOutputFrame(bmp, m);
1409
  }, () => { frame.close(); });
1410
  },
@@ -1422,6 +1417,10 @@ function feedOutputH264(buf, meta) {
1422
  }
1423
 
1424
  function enqueueOutputFrame(frame, meta) {
 
 
 
 
1425
  const item = frame instanceof ImageBitmap
1426
  ? { bmp: frame, meta, receivedAtMs: Date.now() }
1427
  : { url: URL.createObjectURL(frame), meta, receivedAtMs: Date.now() };
@@ -1581,6 +1580,7 @@ function frameBlob(width, height, quality) {
1581
  }
1582
 
1583
  async function tick() {
 
1584
  if (!ws || ws.readyState !== WebSocket.OPEN) return;
1585
  if (pePausedSend) return;
1586
  if (tickInFlight) {
@@ -1603,11 +1603,11 @@ async function tick() {
1603
  return;
1604
  }
1605
  const blob = await frameBlob(width, height, effectiveUpQuality());
1606
- if (!blob || ws.readyState !== WebSocket.OPEN) return;
1607
  const seq = sentFrames + 1;
1608
  const tCaptureMs = Date.now();
1609
  try {
1610
- ws.send(JSON.stringify({ type: "frame_meta", seq, t_capture_ms: tCaptureMs }));
1611
  if (ws.readyState !== WebSocket.OPEN) return;
1612
  ws.send(blob);
1613
  } catch (err) {
@@ -1617,7 +1617,7 @@ async function tick() {
1617
  sentFramesWindow += 1;
1618
  updateMetrics();
1619
  } finally {
1620
- tickInFlight = false;
1621
  }
1622
  }
1623
 
@@ -1630,17 +1630,17 @@ function seekVideoTo(timeSec) {
1630
  setTimeout(finish, 400);
1631
  });
1632
  }
1633
- async function sendCurrentVideoFrame(mediaTimeSec) {
1634
  if (!ws || ws.readyState !== WebSocket.OPEN) return false;
1635
  const width = Number(document.getElementById("width").value);
1636
  const height = Number(document.getElementById("height").value);
1637
  const quality = effectiveUpQuality();
1638
  const blob = await frameBlob(width, height, quality);
1639
- if (!blob || ws.readyState !== WebSocket.OPEN) return false;
1640
  const seq = sentFrames + 1;
1641
  const tCaptureMs = Math.round(mediaTimeSec * 1000);
1642
  try {
1643
- ws.send(JSON.stringify({ type: "frame_meta", seq, t_capture_ms: tCaptureMs }));
1644
  if (ws.readyState !== WebSocket.OPEN) return false;
1645
  ws.send(blob);
1646
  } catch (err) {
@@ -1651,8 +1651,8 @@ async function sendCurrentVideoFrame(mediaTimeSec) {
1651
  updateMetrics();
1652
  return true;
1653
  }
1654
- async function pullNextVideoFrame() {
1655
- if (!videoPullActive) return;
1656
  if (!ws || ws.readyState !== WebSocket.OPEN) { videoPullActive = false; return; }
1657
  if (pePausedSend) return;
1658
  if (videoPullIdx >= videoPullTotal) {
@@ -1663,30 +1663,42 @@ async function pullNextVideoFrame() {
1663
  const frameStartMs = Date.now();
1664
  const intervalMs = 1000 / Math.max(1, videoPullFps);
1665
  if (backendPendingFrames() >= MAX_BACKEND_PENDING_FRAMES) {
1666
- setTimeout(() => { pullNextVideoFrame(); }, 30);
1667
  return;
1668
  }
1669
  const t = videoPullIdx / videoPullFps;
1670
  await seekVideoTo(t);
1671
- if (!videoPullActive || pePausedSend) return;
1672
- const ok = await sendCurrentVideoFrame(t);
1673
- if (!videoPullActive) return;
1674
  if (ok) videoPullIdx += 1;
1675
  const spent = Date.now() - frameStartMs;
1676
  const wait = ok ? Math.max(0, intervalMs - spent) : 30;
1677
- setTimeout(() => { pullNextVideoFrame(); }, wait);
1678
  }
1679
  function startVideoPull() {
 
 
1680
  const dur = (Number.isFinite(camera.duration) && camera.duration > 0) ? camera.duration : 0;
1681
  videoPullTotal = dur > 0 ? Math.max(1, Math.ceil(dur * videoPullFps)) : Number.MAX_SAFE_INTEGER;
1682
  videoPullIdx = 0;
1683
  videoPullActive = true;
1684
- pullNextVideoFrame();
 
 
 
 
 
 
1685
  }
1686
- function stopVideoPull() { videoPullActive = false; }
1687
 
1688
  async function beginSession() {
1689
  if (!ws || ws.readyState !== WebSocket.OPEN) return false;
 
 
 
 
 
1690
  resetRunMetrics();
1691
  pePausedSend = false;
1692
  clearResultVideo();
@@ -1695,17 +1707,14 @@ async function beginSession() {
1695
  try { camera.pause(); } catch (err) {}
1696
  try { camera.currentTime = 0; } catch (err) {}
1697
  }
1698
- updateSourceFrameVisibility();
1699
  let refImage = null;
1700
  try {
1701
  refImage = await readRefImageDataUrl();
1702
- if (refImage) applyRefImagePreview();
1703
  } catch (err) {
1704
  }
 
1705
  const rawPrompt = document.getElementById("prompt").value;
1706
- const peSuppressed = suppressPeThisSend;
1707
- suppressPeThisSend = false;
1708
- const usePe = !peSuppressed && !refImage && document.getElementById("usePe").checked;
1709
  const width = Number(document.getElementById("width").value);
1710
  const height = Number(document.getElementById("height").value);
1711
  const cacheKey = peCacheKey(rawPrompt, refImage);
@@ -1731,6 +1740,7 @@ async function beginSession() {
1731
  }
1732
  const startPayload = {
1733
  type: "start",
 
1734
  prompt: rawPrompt,
1735
  width,
1736
  height,
@@ -1750,15 +1760,14 @@ async function beginSession() {
1750
  }
1751
  if (cachedEnhancedPrompt) startPayload.enhanced_prompt = cachedEnhancedPrompt;
1752
  if (refImage) startPayload.ref_image = refImage;
1753
- const maxTemporalIds = maxTemporalIdsValue();
1754
- if (maxTemporalIds !== null) startPayload.max_temporal_ids = maxTemporalIds;
1755
- startPayload.freeze_kv_on_static = document.getElementById("freezeKvOnStatic").checked;
1756
  const sdt = parseFloat(document.getElementById("staticDiffThresh").value);
1757
  startPayload.static_diff_thresh = Number.isFinite(sdt) ? sdt : 0.5;
1758
  startPayload.profile_timings = document.getElementById("profileTimings").checked;
1759
  sessionProfileTimings = startPayload.profile_timings;
1760
  stageTimingsPanel.style.display = sessionProfileTimings ? "" : "none";
1761
- if (!ws || ws.readyState !== WebSocket.OPEN) return false;
1762
  ws.send(JSON.stringify(startPayload));
1763
  return true;
1764
  }
@@ -1807,8 +1816,10 @@ async function start() {
1807
  socket.onopen = async () => {
1808
  };
1809
  socket.onmessage = async (event) => {
 
1810
  if (typeof event.data === "string") {
1811
  const msg = JSON.parse(event.data);
 
1812
  noteBackendAck(msg);
1813
  if (msg.type === "queue_position") {
1814
  const ahead = Number(msg.ahead != null ? msg.ahead : msg.position) || 0;
@@ -1821,6 +1832,7 @@ async function start() {
1821
  return;
1822
  }
1823
  if (msg.type === "started") {
 
1824
  if (Number(msg.width) > 0 && Number(msg.height) > 0) {
1825
  document.getElementById("width").value = msg.width;
1826
  document.getElementById("height").value = msg.height;
@@ -1849,8 +1861,9 @@ async function start() {
1849
  if (usingVideoFile && peDeferThisSend) {
1850
  ensureOutputPlaybackTimer(true);
1851
  startNetPingTimer();
 
1852
  await seekVideoTo(0);
1853
- await sendCurrentVideoFrame(0);
1854
  } else {
1855
  startFrameTimer();
1856
  }
@@ -1899,7 +1912,8 @@ async function start() {
1899
  clearOutputQueue();
1900
  lastShownCaptureMs = Date.now();
1901
  alignMaxAgeMs = 0;
1902
- if (usingVideoFile) {
 
1903
  try { camera.pause(); } catch (err) {}
1904
  startVideoPull();
1905
  }
@@ -1982,7 +1996,7 @@ async function start() {
1982
  return;
1983
  }
1984
  if (msg.type === "chunk_done") {
1985
- try { ws.send(JSON.stringify({ type: "ack", recv: receivedFrames })); } catch (err) {}
1986
  if (
1987
  chunkProfileLast &&
1988
  (msg.ws_send_s !== undefined || msg.server_residence_s !== undefined) &&
@@ -2070,12 +2084,13 @@ async function start() {
2070
  }
2071
  return;
2072
  }
2073
- if (startingRun) return;
 
2074
  const meta = pendingOutputMeta;
2075
  pendingOutputMeta = null;
2076
  receivedFrames += 1;
2077
  if (receivedFrames % 4 === 0) {
2078
- try { ws.send(JSON.stringify({ type: "ack", recv: receivedFrames })); } catch (err) {}
2079
  }
2080
  receivedFramesWindow += 1;
2081
  if (meta) {
@@ -2107,11 +2122,11 @@ async function start() {
2107
  clearOutputQueue();
2108
  stopRecordDrain();
2109
  finalizeIsForResult = false;
2110
- resultEpoch += 1;
2111
  ws = null;
2112
  startingRun = false;
2113
  sessionGranted = false;
2114
- if (!terminalHold) {
2115
  setSendBusy(false, t("busy_disconnected"));
2116
  showOutputIdle(true);
2117
  }
@@ -2141,7 +2156,7 @@ function stopActiveRun() {
2141
  ws = null;
2142
  if (currentWs) {
2143
  try {
2144
- if (currentWs.readyState === WebSocket.OPEN) currentWs.send(JSON.stringify({ type: "stop" }));
2145
  } catch (err) {}
2146
  currentWs.onopen = null;
2147
  currentWs.onmessage = null;
@@ -2172,7 +2187,7 @@ async function send() {
2172
  await start();
2173
  }
2174
 
2175
- sendBtn.onclick = () => { suppressPeThisSend = false; send(); };
2176
 
2177
  const PROMPT_GROUPS = [
2178
  {
@@ -2265,16 +2280,13 @@ function applyCaseSelection(item, group) {
2265
  if (el) el.value = "";
2266
  injectedRefImageDataUrl = REF_IMAGES[item.ref];
2267
  refImagePreviewDataUrl = injectedRefImageDataUrl;
2268
- applyRefImagePreview();
2269
  } else if (!refImageSelected()) {
2270
  injectedRefImageDataUrl = null;
2271
  refImagePreviewDataUrl = null;
2272
  }
2273
- updateSourceFrameVisibility();
2274
  syncRefThumb();
2275
  closeCardsPop();
2276
  updateMetrics();
2277
- suppressPeThisSend = false;
2278
  send();
2279
  }
2280
 
@@ -2293,14 +2305,10 @@ function syncRefThumb() {
2293
  syncPeAvailability();
2294
  }
2295
 
2296
- function refImageActive() {
2297
- return refImageSelected() || !!injectedRefImageDataUrl || !!refImagePreviewDataUrl;
2298
- }
2299
-
2300
  function syncPeAvailability() {
2301
  const usePeEl = document.getElementById("usePe");
2302
  const peAvailable = !(SERVER_DEFAULTS && SERVER_DEFAULTS.pe_available === false);
2303
- const locked = refImageActive() || !peAvailable;
2304
  if (usePeEl) {
2305
  usePeEl.disabled = locked;
2306
  if (!peAvailable) usePeEl.checked = false;
@@ -2368,6 +2376,7 @@ syncPeAvailability();
2368
  function handleVideoFileChange(event) {
2369
  const file = event.target.files && event.target.files[0];
2370
  if (!file) return;
 
2371
  stopActiveRun();
2372
  camera.onended = null;
2373
  camera.onloadedmetadata = null;
@@ -2431,12 +2440,13 @@ function handleVideoEnded() {
2431
  showOutputIdle(true);
2432
  setSendBusy(false, "");
2433
  }, 30000);
2434
- try { ws.send(JSON.stringify({ type: "finalize_recording" })); }
2435
  catch (err) {}
2436
  }
2437
  }
2438
 
2439
  function switchToCamera() {
 
2440
  stopActiveRun();
2441
  startingRun = false;
2442
  sessionGranted = false;
@@ -2547,6 +2557,7 @@ function stopResultSync() {
2547
  try { if (usingVideoFile) camera.pause(); } catch (e) {}
2548
  }
2549
  function clearResultVideo() {
 
2550
  stopResultSync();
2551
  finalizeIsForResult = false;
2552
  lastRecId = null;
@@ -2578,17 +2589,19 @@ function triggerBrowserDownload(url, filename) {
2578
  }
2579
  function fetchAndSaveRecording() {
2580
  if (resultBlobUrl) triggerBrowserDownload(resultBlobUrl, "joyomni_output.mp4");
 
2581
  }
2582
  async function fetchAndShowResult() {
2583
  const myEpoch = resultEpoch;
 
 
2584
  try {
2585
  const resp = await fetch("/download_last?rec=" + encodeURIComponent(lastRecId), { cache: "no-store" });
2586
  if (myEpoch !== resultEpoch) return;
2587
  if (!resp.ok) {
2588
  let detail = "";
2589
  try { const j = await resp.json(); detail = j.error || ""; } catch (e) {}
2590
- setSendBusy(false, "");
2591
- return;
2592
  }
2593
  const blob = await resp.blob();
2594
  if (myEpoch !== resultEpoch) return;
@@ -2597,20 +2610,33 @@ async function fetchAndShowResult() {
2597
  showResultVideo(resultBlobUrl);
2598
  if (usingVideoFile) startResultSync();
2599
  showDownloadBubble(true);
 
2600
  showOutputToast(t("toast_done"), 4000);
2601
- setSendBusy(false, "");
2602
  } catch (err) {
 
 
 
 
 
 
 
 
 
 
 
 
2603
  }
2604
  }
2605
  function onRecordingFinalized(msg) {
2606
  if (finalizeWatchdog) { clearTimeout(finalizeWatchdog); finalizeWatchdog = null; }
2607
- lastRecId = msg.rec || null;
2608
  if (!lastRecId) {
2609
  finalizeIsForResult = false;
2610
  stopResultSync();
2611
  if (outputWaitOverlay) outputWaitOverlay.style.display = "none";
2612
  showOutputIdle(true);
2613
  setSendBusy(false, "");
 
2614
  return;
2615
  }
2616
  finalizeIsForResult = false;
@@ -2752,7 +2778,7 @@ let autoQLowStreak = 0;
2752
  let autoQHoldTicks = 0;
2753
  function sendOutputQuality(v) {
2754
  if (ws && ws.readyState === WebSocket.OPEN && sessionGranted) {
2755
- try { ws.send(JSON.stringify({ type: "set_output_quality", value: v })); } catch (err) {}
2756
  }
2757
  }
2758
  function setDownQuality(v, opts) {
@@ -2858,9 +2884,9 @@ function applyServerDefaults() {
2858
  if (v !== null && v !== undefined) document.getElementById(id).checked = !!v;
2859
  }
2860
  const maxIds = d.max_temporal_ids;
2861
- if (maxIds !== null && maxIds !== undefined) {
2862
- document.getElementById("useMaxTemporalIds").checked = true;
2863
- document.getElementById("maxTemporalIds").value = maxIds;
2864
  }
2865
  if (d.output_quality !== undefined && d.output_quality !== "auto") {
2866
  autoQuality = false;
@@ -2884,7 +2910,6 @@ window.addEventListener("load", () => {
2884
  applyLang();
2885
  applyOrientation(defaultCameraOrientation());
2886
  updateMaxTemporalIdsVisibility();
2887
- updateSourceFrameVisibility();
2888
  syncRefThumb();
2889
  showOutputIdle(true);
2890
  updateMetrics();
 
417
  gate_settling: "保持不动,正在对焦",
418
  gate_no_face_wait: "没有检测到人脸,请正对摄像头 😊",
419
  toast_done: "✨ 编辑完成,点右下角气泡下载",
420
+ download_retry: "重试下载",
421
+ download_failed: "下载失败:",
422
  toast_finalizing: "✨ 编辑完成,正在生成可下载视频",
423
  edit_processing: "✨ 正在编辑",
424
  q_downgrade_auto: " · 已自动降到低画质", q_downgrade_hint: " · 建议降低画质",
425
  flow_drop_tail: " · 带宽不足,正在丢帧保延迟",
426
  up_auto_tail: " · 上行已自动降至中档",
427
+ res_notice: "本演示运行在 {res},而非模型原生的 1248×720 分辨率,画质与编辑效果会有所下降。",
428
  },
429
  en: {
430
  tagline: "Streaming video editing",
 
497
  gate_settling: "Hold still, focusing",
498
  gate_no_face_wait: "No face detected, please face the camera 😊",
499
  toast_done: "✨ Done — click the bubble at bottom-right to download",
500
+ download_retry: "Retry download",
501
+ download_failed: "Download failed: ",
502
  toast_finalizing: "✨ Done — generating the downloadable video",
503
  edit_processing: "✨ Editing",
504
  q_downgrade_auto: " · auto-lowered to low quality", q_downgrade_hint: " · consider lowering quality",
505
  flow_drop_tail: " · bandwidth limited, dropping frames",
506
  up_auto_tail: " · uplink auto-lowered to mid",
507
+ res_notice: "This demo runs at {res}, not the model's native 1248×720 resolution; visual quality and edit fidelity are reduced.",
508
  },
509
  };
510
  let currentLang = localStorage.getItem("joyomni_lang") || "en";
 
740
  let pePausedSend = false;
741
  let peDeferThisSend = false;
742
  let videoPullActive = false;
743
+ let videoPullGeneration = 0;
744
+ let videoPullTimer = null;
745
  let videoPullIdx = 0;
746
  const videoPullFps = 24;
747
  let videoPullTotal = 0;
 
758
  let startingRun = false;
759
  let sessionGranted = false;
760
  let peCache = null;
761
+ let sourceRevision = 0;
762
+ let activeSessionId = null;
763
  let pendingPeCacheKey = null;
764
  let outputPlaybackTimer = null;
765
  let outputPlaybackIntervalMs = null;
 
789
  let alignMaxAgeMs = 0;
790
  let lastShownCaptureMs = 0;
791
  let pendingRevokeUrl = null;
 
792
  let refImagePreviewDataUrl = null;
793
  let injectedRefImageDataUrl = null;
 
794
  let outputWaitingFace = false;
795
  let waitHintReason = null;
796
  let waitHintAt = 0;
 
844
  }
845
 
846
  function noteBackendAck(msg) {
847
+ if (msg && msg.session_id && msg.session_id !== activeSessionId) return;
848
  const framesIn = Number(msg && msg.frames_in);
849
  if (Number.isFinite(framesIn)) {
850
  backendAckedFrames = Math.max(backendAckedFrames, framesIn);
 
879
  return o === "portrait" ? { width: short, height: long } : { width: long, height: short };
880
  }
881
  function applyOrientation(o) {
882
+ invalidatePeSource();
883
  const dims = orientDims(o);
884
  document.getElementById("width").value = dims.width;
885
  document.getElementById("height").value = dims.height;
 
900
  }
901
 
902
  function peCacheKey(rawPrompt, refImage) {
903
+ return JSON.stringify({ prompt: rawPrompt || "", refImage: refImage || "", sourceRevision,
904
+ width: document.getElementById("width").value, height: document.getElementById("height").value });
905
+ }
906
+
907
+ function invalidatePeSource() {
908
+ sourceRevision += 1;
909
+ peCache = null;
910
+ pendingPeCacheKey = null;
911
  }
912
 
913
  function startFrameTimer() {
 
928
  if (netPingTimer) clearInterval(netPingTimer);
929
  netPingTimer = setInterval(() => {
930
  if (ws && ws.readyState === WebSocket.OPEN) {
931
+ try { ws.send(JSON.stringify({ session_id: activeSessionId, type: "ping", t: Date.now(), recv: receivedFrames, rtt: netRttLastMs })); } catch (err) {}
932
  }
933
  autoQualityTick();
934
  }, 1000);
 
956
  });
957
  }
958
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
959
  async function handleRefImageChange() {
960
  injectedRefImageDataUrl = null;
961
  if (!refImageSelected()) {
962
  refImagePreviewDataUrl = null;
 
 
 
963
  syncRefThumb();
964
  updateMetrics();
965
  return;
966
  }
967
  try {
968
  await readRefImageDataUrl();
 
969
  } catch (err) {
970
  }
 
971
  syncRefThumb();
972
  updateMetrics();
973
  }
 
1239
  if (pendingRevokeUrl) { URL.revokeObjectURL(pendingRevokeUrl); pendingRevokeUrl = null; }
1240
  outputImg.onload = outputImg.onerror = null;
1241
  outputImg.removeAttribute("src");
 
 
1242
  hidePeResult();
 
1243
  sessionProfileTimings = false;
1244
  stageTimings.innerHTML = "";
1245
  stageTimingsPanel.style.display = "none";
 
1337
 
1338
  function ensureUplinkEncoder(width, height) {
1339
  if (upEncoder && upEncoder.state === "configured") return;
1340
+ const sessionId = activeSessionId;
1341
  upEncoder = new VideoEncoder({
1342
  output: (chunk) => {
1343
+ if (sessionId !== activeSessionId) return;
1344
  const info = upPendingByTs.get(chunk.timestamp);
1345
  upPendingByTs.delete(chunk.timestamp);
1346
  if (!info || !ws || ws.readyState !== WebSocket.OPEN) return;
 
1351
  const buf = new ArrayBuffer(chunk.byteLength);
1352
  chunk.copyTo(buf);
1353
  try {
1354
+ ws.send(JSON.stringify({ session_id: activeSessionId, type: "frame_meta", seq: info.seq, t_capture_ms: info.t_capture_ms }));
1355
  ws.send(buf);
1356
  } catch (e) {
1357
  return;
 
1391
 
1392
  function feedOutputH264(buf, meta) {
1393
  if (!outDecoder || outDecoder.state === "closed") {
1394
+ const sessionId = activeSessionId;
1395
  outDecoder = new VideoDecoder({
1396
  output: (frame) => {
1397
+ if (sessionId !== activeSessionId) { frame.close(); return; }
1398
  const m = outMetaByTs.get(frame.timestamp);
1399
  outMetaByTs.delete(frame.timestamp);
1400
  createImageBitmap(frame).then((bmp) => {
1401
  frame.close();
1402
+ if (sessionId !== activeSessionId) { bmp.close(); return; }
1403
  enqueueOutputFrame(bmp, m);
1404
  }, () => { frame.close(); });
1405
  },
 
1417
  }
1418
 
1419
  function enqueueOutputFrame(frame, meta) {
1420
+ if (meta && meta.session_id && meta.session_id !== activeSessionId) {
1421
+ if (frame instanceof ImageBitmap) frame.close();
1422
+ return;
1423
+ }
1424
  const item = frame instanceof ImageBitmap
1425
  ? { bmp: frame, meta, receivedAtMs: Date.now() }
1426
  : { url: URL.createObjectURL(frame), meta, receivedAtMs: Date.now() };
 
1580
  }
1581
 
1582
  async function tick() {
1583
+ const generation = videoPullGeneration;
1584
  if (!ws || ws.readyState !== WebSocket.OPEN) return;
1585
  if (pePausedSend) return;
1586
  if (tickInFlight) {
 
1603
  return;
1604
  }
1605
  const blob = await frameBlob(width, height, effectiveUpQuality());
1606
+ if (generation !== videoPullGeneration || !blob || !ws || ws.readyState !== WebSocket.OPEN) return;
1607
  const seq = sentFrames + 1;
1608
  const tCaptureMs = Date.now();
1609
  try {
1610
+ ws.send(JSON.stringify({ session_id: activeSessionId, type: "frame_meta", seq, t_capture_ms: tCaptureMs }));
1611
  if (ws.readyState !== WebSocket.OPEN) return;
1612
  ws.send(blob);
1613
  } catch (err) {
 
1617
  sentFramesWindow += 1;
1618
  updateMetrics();
1619
  } finally {
1620
+ if (generation === videoPullGeneration) tickInFlight = false;
1621
  }
1622
  }
1623
 
 
1630
  setTimeout(finish, 400);
1631
  });
1632
  }
1633
+ async function sendCurrentVideoFrame(mediaTimeSec, generation = videoPullGeneration) {
1634
  if (!ws || ws.readyState !== WebSocket.OPEN) return false;
1635
  const width = Number(document.getElementById("width").value);
1636
  const height = Number(document.getElementById("height").value);
1637
  const quality = effectiveUpQuality();
1638
  const blob = await frameBlob(width, height, quality);
1639
+ if (generation !== videoPullGeneration || !blob || !ws || ws.readyState !== WebSocket.OPEN) return false;
1640
  const seq = sentFrames + 1;
1641
  const tCaptureMs = Math.round(mediaTimeSec * 1000);
1642
  try {
1643
+ ws.send(JSON.stringify({ session_id: activeSessionId, type: "frame_meta", seq, t_capture_ms: tCaptureMs }));
1644
  if (ws.readyState !== WebSocket.OPEN) return false;
1645
  ws.send(blob);
1646
  } catch (err) {
 
1651
  updateMetrics();
1652
  return true;
1653
  }
1654
+ async function pullNextVideoFrame(generation) {
1655
+ if (generation !== videoPullGeneration || !videoPullActive) return;
1656
  if (!ws || ws.readyState !== WebSocket.OPEN) { videoPullActive = false; return; }
1657
  if (pePausedSend) return;
1658
  if (videoPullIdx >= videoPullTotal) {
 
1663
  const frameStartMs = Date.now();
1664
  const intervalMs = 1000 / Math.max(1, videoPullFps);
1665
  if (backendPendingFrames() >= MAX_BACKEND_PENDING_FRAMES) {
1666
+ videoPullTimer = setTimeout(() => pullNextVideoFrame(generation), 30);
1667
  return;
1668
  }
1669
  const t = videoPullIdx / videoPullFps;
1670
  await seekVideoTo(t);
1671
+ if (generation !== videoPullGeneration || !videoPullActive || pePausedSend) return;
1672
+ const ok = await sendCurrentVideoFrame(t, generation);
1673
+ if (generation !== videoPullGeneration || !videoPullActive) return;
1674
  if (ok) videoPullIdx += 1;
1675
  const spent = Date.now() - frameStartMs;
1676
  const wait = ok ? Math.max(0, intervalMs - spent) : 30;
1677
+ videoPullTimer = setTimeout(() => pullNextVideoFrame(generation), wait);
1678
  }
1679
  function startVideoPull() {
1680
+ if (videoPullActive) return;
1681
+ stopVideoPull();
1682
  const dur = (Number.isFinite(camera.duration) && camera.duration > 0) ? camera.duration : 0;
1683
  videoPullTotal = dur > 0 ? Math.max(1, Math.ceil(dur * videoPullFps)) : Number.MAX_SAFE_INTEGER;
1684
  videoPullIdx = 0;
1685
  videoPullActive = true;
1686
+ pullNextVideoFrame(videoPullGeneration);
1687
+ }
1688
+ function stopVideoPull() {
1689
+ videoPullActive = false;
1690
+ videoPullGeneration += 1;
1691
+ if (videoPullTimer !== null) clearTimeout(videoPullTimer);
1692
+ videoPullTimer = null;
1693
  }
 
1694
 
1695
  async function beginSession() {
1696
  if (!ws || ws.readyState !== WebSocket.OPEN) return false;
1697
+ stopVideoPull();
1698
+ activeSessionId = String(Number(activeSessionId) + 1);
1699
+ const sessionId = activeSessionId;
1700
+ const socket = ws;
1701
+ if (!usingVideoFile) invalidatePeSource();
1702
  resetRunMetrics();
1703
  pePausedSend = false;
1704
  clearResultVideo();
 
1707
  try { camera.pause(); } catch (err) {}
1708
  try { camera.currentTime = 0; } catch (err) {}
1709
  }
 
1710
  let refImage = null;
1711
  try {
1712
  refImage = await readRefImageDataUrl();
 
1713
  } catch (err) {
1714
  }
1715
+ if (sessionId !== activeSessionId || socket !== ws) return false;
1716
  const rawPrompt = document.getElementById("prompt").value;
1717
+ const usePe = document.getElementById("usePe").checked;
 
 
1718
  const width = Number(document.getElementById("width").value);
1719
  const height = Number(document.getElementById("height").value);
1720
  const cacheKey = peCacheKey(rawPrompt, refImage);
 
1740
  }
1741
  const startPayload = {
1742
  type: "start",
1743
+ session_id: activeSessionId,
1744
  prompt: rawPrompt,
1745
  width,
1746
  height,
 
1760
  }
1761
  if (cachedEnhancedPrompt) startPayload.enhanced_prompt = cachedEnhancedPrompt;
1762
  if (refImage) startPayload.ref_image = refImage;
1763
+ startPayload.max_temporal_ids = maxTemporalIdsValue();
1764
+ startPayload.freeze_kv_on_static = !usingVideoFile && document.getElementById("freezeKvOnStatic").checked;
 
1765
  const sdt = parseFloat(document.getElementById("staticDiffThresh").value);
1766
  startPayload.static_diff_thresh = Number.isFinite(sdt) ? sdt : 0.5;
1767
  startPayload.profile_timings = document.getElementById("profileTimings").checked;
1768
  sessionProfileTimings = startPayload.profile_timings;
1769
  stageTimingsPanel.style.display = sessionProfileTimings ? "" : "none";
1770
+ if (sessionId !== activeSessionId || socket !== ws || !ws || ws.readyState !== WebSocket.OPEN) return false;
1771
  ws.send(JSON.stringify(startPayload));
1772
  return true;
1773
  }
 
1816
  socket.onopen = async () => {
1817
  };
1818
  socket.onmessage = async (event) => {
1819
+ if (ws !== socket) return;
1820
  if (typeof event.data === "string") {
1821
  const msg = JSON.parse(event.data);
1822
+ if (msg.session_id && msg.session_id !== activeSessionId) return;
1823
  noteBackendAck(msg);
1824
  if (msg.type === "queue_position") {
1825
  const ahead = Number(msg.ahead != null ? msg.ahead : msg.position) || 0;
 
1832
  return;
1833
  }
1834
  if (msg.type === "started") {
1835
+ peDeferThisSend = msg.pe_deferred ?? (peDeferThisSend && msg.use_pe !== false);
1836
  if (Number(msg.width) > 0 && Number(msg.height) > 0) {
1837
  document.getElementById("width").value = msg.width;
1838
  document.getElementById("height").value = msg.height;
 
1861
  if (usingVideoFile && peDeferThisSend) {
1862
  ensureOutputPlaybackTimer(true);
1863
  startNetPingTimer();
1864
+ const generation = videoPullGeneration;
1865
  await seekVideoTo(0);
1866
+ if (generation === videoPullGeneration) await sendCurrentVideoFrame(0, generation);
1867
  } else {
1868
  startFrameTimer();
1869
  }
 
1912
  clearOutputQueue();
1913
  lastShownCaptureMs = Date.now();
1914
  alignMaxAgeMs = 0;
1915
+ if (usingVideoFile && peDeferThisSend && !msg.cached) {
1916
+ peDeferThisSend = false;
1917
  try { camera.pause(); } catch (err) {}
1918
  startVideoPull();
1919
  }
 
1996
  return;
1997
  }
1998
  if (msg.type === "chunk_done") {
1999
+ try { ws.send(JSON.stringify({ session_id: activeSessionId, type: "ack", recv: receivedFrames })); } catch (err) {}
2000
  if (
2001
  chunkProfileLast &&
2002
  (msg.ws_send_s !== undefined || msg.server_residence_s !== undefined) &&
 
2084
  }
2085
  return;
2086
  }
2087
+ if (startingRun || !pendingOutputMeta ||
2088
+ (pendingOutputMeta.session_id && pendingOutputMeta.session_id !== activeSessionId)) return;
2089
  const meta = pendingOutputMeta;
2090
  pendingOutputMeta = null;
2091
  receivedFrames += 1;
2092
  if (receivedFrames % 4 === 0) {
2093
+ try { ws.send(JSON.stringify({ session_id: activeSessionId, type: "ack", recv: receivedFrames })); } catch (err) {}
2094
  }
2095
  receivedFramesWindow += 1;
2096
  if (meta) {
 
2122
  clearOutputQueue();
2123
  stopRecordDrain();
2124
  finalizeIsForResult = false;
2125
+ if (!lastRecId) resultEpoch += 1;
2126
  ws = null;
2127
  startingRun = false;
2128
  sessionGranted = false;
2129
+ if (!terminalHold && !lastRecId) {
2130
  setSendBusy(false, t("busy_disconnected"));
2131
  showOutputIdle(true);
2132
  }
 
2156
  ws = null;
2157
  if (currentWs) {
2158
  try {
2159
+ if (currentWs.readyState === WebSocket.OPEN) currentWs.send(JSON.stringify({ session_id: activeSessionId, type: "stop" }));
2160
  } catch (err) {}
2161
  currentWs.onopen = null;
2162
  currentWs.onmessage = null;
 
2187
  await start();
2188
  }
2189
 
2190
+ sendBtn.onclick = send;
2191
 
2192
  const PROMPT_GROUPS = [
2193
  {
 
2280
  if (el) el.value = "";
2281
  injectedRefImageDataUrl = REF_IMAGES[item.ref];
2282
  refImagePreviewDataUrl = injectedRefImageDataUrl;
 
2283
  } else if (!refImageSelected()) {
2284
  injectedRefImageDataUrl = null;
2285
  refImagePreviewDataUrl = null;
2286
  }
 
2287
  syncRefThumb();
2288
  closeCardsPop();
2289
  updateMetrics();
 
2290
  send();
2291
  }
2292
 
 
2305
  syncPeAvailability();
2306
  }
2307
 
 
 
 
 
2308
  function syncPeAvailability() {
2309
  const usePeEl = document.getElementById("usePe");
2310
  const peAvailable = !(SERVER_DEFAULTS && SERVER_DEFAULTS.pe_available === false);
2311
+ const locked = !peAvailable;
2312
  if (usePeEl) {
2313
  usePeEl.disabled = locked;
2314
  if (!peAvailable) usePeEl.checked = false;
 
2376
  function handleVideoFileChange(event) {
2377
  const file = event.target.files && event.target.files[0];
2378
  if (!file) return;
2379
+ invalidatePeSource();
2380
  stopActiveRun();
2381
  camera.onended = null;
2382
  camera.onloadedmetadata = null;
 
2440
  showOutputIdle(true);
2441
  setSendBusy(false, "");
2442
  }, 30000);
2443
+ try { ws.send(JSON.stringify({ session_id: activeSessionId, type: "finalize_recording" })); }
2444
  catch (err) {}
2445
  }
2446
  }
2447
 
2448
  function switchToCamera() {
2449
+ invalidatePeSource();
2450
  stopActiveRun();
2451
  startingRun = false;
2452
  sessionGranted = false;
 
2557
  try { if (usingVideoFile) camera.pause(); } catch (e) {}
2558
  }
2559
  function clearResultVideo() {
2560
+ if (downloadBubble) downloadBubble.disabled = false;
2561
  stopResultSync();
2562
  finalizeIsForResult = false;
2563
  lastRecId = null;
 
2589
  }
2590
  function fetchAndSaveRecording() {
2591
  if (resultBlobUrl) triggerBrowserDownload(resultBlobUrl, "joyomni_output.mp4");
2592
+ else if (lastRecId) fetchAndShowResult();
2593
  }
2594
  async function fetchAndShowResult() {
2595
  const myEpoch = resultEpoch;
2596
+ setSendBusy(true, t("toast_finalizing"));
2597
+ if (downloadBubble) downloadBubble.disabled = true;
2598
  try {
2599
  const resp = await fetch("/download_last?rec=" + encodeURIComponent(lastRecId), { cache: "no-store" });
2600
  if (myEpoch !== resultEpoch) return;
2601
  if (!resp.ok) {
2602
  let detail = "";
2603
  try { const j = await resp.json(); detail = j.error || ""; } catch (e) {}
2604
+ throw new Error(detail || `HTTP ${resp.status}`);
 
2605
  }
2606
  const blob = await resp.blob();
2607
  if (myEpoch !== resultEpoch) return;
 
2610
  showResultVideo(resultBlobUrl);
2611
  if (usingVideoFile) startResultSync();
2612
  showDownloadBubble(true);
2613
+ if (downloadBubble) downloadBubble.textContent = t("download_done");
2614
  showOutputToast(t("toast_done"), 4000);
 
2615
  } catch (err) {
2616
+ if (myEpoch !== resultEpoch) return;
2617
+ if (outputWaitOverlay) outputWaitOverlay.style.display = "none";
2618
+ showOutputToast(t("download_failed") + String(err.message || err), 8000);
2619
+ if (lastRecId) {
2620
+ showDownloadBubble(true);
2621
+ if (downloadBubble) downloadBubble.textContent = t("download_retry");
2622
+ }
2623
+ } finally {
2624
+ if (myEpoch === resultEpoch) {
2625
+ setSendBusy(false, "");
2626
+ if (downloadBubble) downloadBubble.disabled = false;
2627
+ }
2628
  }
2629
  }
2630
  function onRecordingFinalized(msg) {
2631
  if (finalizeWatchdog) { clearTimeout(finalizeWatchdog); finalizeWatchdog = null; }
2632
+ lastRecId = msg.ok === false ? null : (msg.rec || null);
2633
  if (!lastRecId) {
2634
  finalizeIsForResult = false;
2635
  stopResultSync();
2636
  if (outputWaitOverlay) outputWaitOverlay.style.display = "none";
2637
  showOutputIdle(true);
2638
  setSendBusy(false, "");
2639
+ if (msg.message) showOutputToast(msg.message, 8000);
2640
  return;
2641
  }
2642
  finalizeIsForResult = false;
 
2778
  let autoQHoldTicks = 0;
2779
  function sendOutputQuality(v) {
2780
  if (ws && ws.readyState === WebSocket.OPEN && sessionGranted) {
2781
+ try { ws.send(JSON.stringify({ session_id: activeSessionId, type: "set_output_quality", value: v })); } catch (err) {}
2782
  }
2783
  }
2784
  function setDownQuality(v, opts) {
 
2884
  if (v !== null && v !== undefined) document.getElementById(id).checked = !!v;
2885
  }
2886
  const maxIds = d.max_temporal_ids;
2887
+ if (maxIds !== undefined) {
2888
+ document.getElementById("useMaxTemporalIds").checked = maxIds !== null;
2889
+ if (maxIds !== null) document.getElementById("maxTemporalIds").value = maxIds;
2890
  }
2891
  if (d.output_quality !== undefined && d.output_quality !== "auto") {
2892
  autoQuality = false;
 
2910
  applyLang();
2911
  applyOrientation(defaultCameraOrientation());
2912
  updateMaxTemporalIdsVisibility();
 
2913
  syncRefThumb();
2914
  showOutputIdle(true);
2915
  updateMetrics();
xvideo/models/pipeline.py CHANGED
@@ -314,38 +314,39 @@ class Pipeline(DiffusionPipeline):
314
  return Pipeline._KV_CACHE_ID_REF_IMAGE
315
  raise ValueError(f"Unsupported cache kind: {kind!r}")
316
 
317
- @staticmethod
318
  def _get_chunk_windows(
 
319
  total_latent_frames: int,
320
  chunk_size: int,
321
  window_size: int,
322
  global_sink_chunk: bool,
323
  ) -> List[Dict[str, Any]]:
 
 
 
 
 
 
 
 
 
 
 
 
324
  if window_size <= 0:
325
  raise ValueError(f"`window_size` must be positive, got {window_size}.")
326
-
327
- windows = []
328
- num_chunks = (total_latent_frames + chunk_size - 1) // chunk_size
329
- for chunk_idx in range(num_chunks):
330
- chunk_start = chunk_idx * chunk_size
331
- chunk_end = min(total_latent_frames, chunk_start + chunk_size)
332
- if global_sink_chunk and chunk_idx > 0:
333
- tail_window_size = max(window_size - 1, 1)
334
- tail_chunk_start = max(1, chunk_idx - tail_window_size + 1)
335
- selected_chunk_ids = [0] + list(range(tail_chunk_start, chunk_idx + 1))
336
- else:
337
- window_chunk_start = max(0, chunk_idx - window_size + 1)
338
- selected_chunk_ids = list(range(window_chunk_start, chunk_idx + 1))
339
-
340
- windows.append(
341
- {
342
- "chunk_idx": chunk_idx,
343
- "chunk_start": chunk_start,
344
- "chunk_end": chunk_end,
345
- "selected_chunk_ids": selected_chunk_ids,
346
- }
347
- )
348
- return windows
349
 
350
  @staticmethod
351
  def _chunk_frame_bounds(chunk_id: int, chunk_size: int, total_latent_frames: int) -> Tuple[int, int]:
 
314
  return Pipeline._KV_CACHE_ID_REF_IMAGE
315
  raise ValueError(f"Unsupported cache kind: {kind!r}")
316
 
317
+ @classmethod
318
  def _get_chunk_windows(
319
+ cls,
320
  total_latent_frames: int,
321
  chunk_size: int,
322
  window_size: int,
323
  global_sink_chunk: bool,
324
  ) -> List[Dict[str, Any]]:
325
+ num_chunks = (total_latent_frames + chunk_size - 1) // chunk_size
326
+ return [cls._get_chunk_window(i, total_latent_frames, chunk_size, window_size, global_sink_chunk)
327
+ for i in range(num_chunks)]
328
+
329
+ @staticmethod
330
+ def _get_chunk_window(
331
+ chunk_idx: int,
332
+ total_latent_frames: int,
333
+ chunk_size: int,
334
+ window_size: int,
335
+ global_sink_chunk: bool,
336
+ ) -> Dict[str, Any]:
337
  if window_size <= 0:
338
  raise ValueError(f"`window_size` must be positive, got {window_size}.")
339
+ if global_sink_chunk and chunk_idx > 0:
340
+ tail_start = max(1, chunk_idx - max(window_size - 1, 1) + 1)
341
+ selected_chunk_ids = [0] + list(range(tail_start, chunk_idx + 1))
342
+ else:
343
+ selected_chunk_ids = list(range(max(0, chunk_idx - window_size + 1), chunk_idx + 1))
344
+ return {
345
+ "chunk_idx": chunk_idx,
346
+ "chunk_start": chunk_idx * chunk_size,
347
+ "chunk_end": min(total_latent_frames, (chunk_idx + 1) * chunk_size),
348
+ "selected_chunk_ids": selected_chunk_ids,
349
+ }
 
 
 
 
 
 
 
 
 
 
 
 
350
 
351
  @staticmethod
352
  def _chunk_frame_bounds(chunk_id: int, chunk_size: int, total_latent_frames: int) -> Tuple[int, int]:
xvideo/serving/joyomni_streaming.py CHANGED
@@ -352,15 +352,11 @@ class JoyOmniRuntime:
352
  postprocess_device=postprocess_device_obj,
353
  )
354
 
355
- try:
356
- if os.environ.get("JOYOMNI_SKIP_LOAD_WARMUP", "0").lower() in {"1", "true", "yes", "on"}:
357
- print("#####[STREAM] load-time warmup skipped (JOYOMNI_SKIP_LOAD_WARMUP); "
358
- "warm up later on a real GPU")
359
- else:
360
- for (_wh, _ww) in _orientations:
361
- runtime.warmup_full_pipeline(height=_wh, width=_ww)
362
- except Exception as _wexc:
363
- print(f"#####[STREAM] full-pipeline warmup error (non-fatal): {_wexc!r}")
364
 
365
  if device_obj.type == "cuda":
366
  _free_b, _total_b = torch.cuda.mem_get_info(device_obj)
@@ -388,20 +384,6 @@ class JoyOmniRuntime:
388
  ref_image=ref_image,
389
  )
390
 
391
- def prebake_graph(
392
- self,
393
- prompt: str,
394
- *,
395
- settings: StreamingSettings,
396
- ref_image: Image.Image | None = None,
397
- ) -> None:
398
- session = self.create_v2v_session(prompt, settings=settings, ref_image=ref_image)
399
- arr = np.zeros((settings.height, settings.width, 3), dtype=np.uint8)
400
- try:
401
- session._initialize(Image.fromarray(arr, mode="RGB"))
402
- finally:
403
- session.close()
404
-
405
  def warmup_full_pipeline(
406
  self,
407
  *,
@@ -445,10 +427,7 @@ class JoyOmniRuntime:
445
  print(f"#####[STREAM] full-pipeline warmup skipped/failed: {exc!r}")
446
  finally:
447
  if session is not None:
448
- try:
449
- session.close()
450
- except Exception:
451
- pass
452
 
453
  class JoyOmniV2VStreamingSession:
454
  _session_serial_counter = 0
@@ -528,6 +507,7 @@ class JoyOmniV2VStreamingSession:
528
  self._result_queue: queue.Queue[StreamingChunkResult] | None = None
529
  self._workers: list[threading.Thread] = []
530
  self._worker_error: BaseException | None = None
 
531
  self._worker_error_lock = threading.Lock()
532
  self._pseudo_latent_condition = threading.Condition()
533
  self._pseudo_latent_chunk_idx: int | None = None
@@ -599,6 +579,8 @@ class JoyOmniV2VStreamingSession:
599
  self,
600
  frame: Image.Image,
601
  frame_meta: dict[str, Any] | None = None,
 
 
602
  ) -> list[StreamingChunkResult]:
603
  self._raise_worker_error_if_needed()
604
  frame = self._resize_frame(frame)
@@ -607,20 +589,20 @@ class JoyOmniV2VStreamingSession:
607
  self._initialize(frame)
608
  self.pending_frames.append(frame)
609
  self.pending_metas.append(meta)
610
- results: list[StreamingChunkResult] = []
611
  while len(self.pending_frames) >= self.frames_per_next_chunk:
612
  n = self.frames_per_next_chunk
613
  chunk_frames = self.pending_frames[:n]
614
  chunk_metas = self.pending_metas[:n]
615
  del self.pending_frames[:n]
616
  del self.pending_metas[:n]
617
- results.extend(self._submit_or_process_chunk(chunk_frames, chunk_metas))
618
- results.extend(self._drain_async_results())
619
  self._raise_worker_error_if_needed()
620
  return results
621
 
622
  @torch.no_grad()
623
  def flush_pending(self) -> None:
 
624
  if not self.initialized or not self.pending_frames:
625
  return
626
  valid = len(self.pending_frames)
@@ -629,18 +611,25 @@ class JoyOmniV2VStreamingSession:
629
  chunk_metas = self.pending_metas + [self.pending_metas[-1]] * pad
630
  self.pending_frames = []
631
  self.pending_metas = []
632
- self._submit_or_process_chunk(chunk_frames, chunk_metas, valid_count=valid)
633
  self._raise_worker_error_if_needed()
634
 
635
- def close(self) -> None:
 
 
 
 
 
 
 
 
 
636
  self._clear_vae_feature_caches()
637
  self.pipeline.transformer.reset_inference_kv_cache()
638
- self._stop_async_workers()
639
- # No empty_cache: keep the freed KV/graph blocks cached in the allocator
640
- # so the next session's pools reuse them instead of fresh cudaMallocs.
641
 
642
  @torch.no_grad()
643
  def _initialize(self, first_frame: Image.Image) -> None:
 
644
  if self.initialized:
645
  return
646
 
@@ -732,8 +721,10 @@ class JoyOmniV2VStreamingSession:
732
  return np.asarray(frame, dtype=np.float32)
733
 
734
  def _update_static_anchor(self, chunk_idx: int, source_frames: list[Image.Image]) -> int | None:
 
 
735
  gray = self._chunk_last_frame_gray(source_frames)
736
- if not self.settings.freeze_kv_on_static or chunk_idx == 0 or gray is None:
737
  self._prev_static_gray = gray if gray is not None else self._prev_static_gray
738
  self._static_anchor_id = None
739
  return None
@@ -764,15 +755,6 @@ class JoyOmniV2VStreamingSession:
764
  self._prev_static_gray = gray
765
  return anchor_id
766
 
767
- def _submit_or_process_chunk(
768
- self,
769
- source_frames: list[Image.Image],
770
- source_metas: list[dict[str, Any]],
771
- valid_count: int | None = None,
772
- ) -> list[StreamingChunkResult]:
773
- self._submit_async_chunk(source_frames, source_metas, valid_count)
774
- return []
775
-
776
  def _new_profile(self, chunk_idx: int, input_frames: int) -> dict[str, Any]:
777
  return {
778
  "chunk_idx": chunk_idx,
@@ -973,13 +955,13 @@ class JoyOmniV2VStreamingSession:
973
  )
974
 
975
  total_latent_frames = chunk_idx + 1
976
- chunk_windows = self.pipeline._get_chunk_windows(
 
977
  total_latent_frames=total_latent_frames,
978
  chunk_size=self.chunk_size,
979
  window_size=self.local_window_size,
980
  global_sink_chunk=self.global_sink_chunk,
981
  )
982
- chunk_window = chunk_windows[-1]
983
  selected_chunk_ids = chunk_window["selected_chunk_ids"]
984
  history_chunk_ids = selected_chunk_ids[:-1]
985
  active_chunk_id = selected_chunk_ids[-1]
@@ -1233,16 +1215,35 @@ class JoyOmniV2VStreamingSession:
1233
  for worker in self._workers:
1234
  worker.start()
1235
 
1236
- def _stop_async_workers(self) -> None:
1237
- if self._encode_queue is None:
1238
- return
1239
- try:
1240
- self._encode_queue.put(None, timeout=0.1)
1241
- except queue.Full:
1242
- if self._worker_error is None:
1243
- self._encode_queue.put(None)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1244
  for worker in self._workers:
1245
- worker.join()
 
 
 
1246
  self._encode_queue = None
1247
  self._dit_queue = None
1248
  self._decode_queue = None
@@ -1252,21 +1253,23 @@ class JoyOmniV2VStreamingSession:
1252
  self._workers = []
1253
 
1254
  def _set_worker_error(self, exc: BaseException) -> None:
 
 
1255
  import traceback as _tb
1256
  print(f"#####[STREAM-WORKER-ERROR] {exc!r}", flush=True)
1257
  _tb.print_exc()
1258
  with self._worker_error_lock:
1259
  if self._worker_error is None:
1260
  self._worker_error = exc
1261
- with self._pseudo_latent_condition:
1262
- self._pseudo_latent_condition.notify_all()
1263
 
1264
  def _raise_worker_error_if_needed(self) -> None:
1265
  with self._worker_error_lock:
1266
  exc = self._worker_error
1267
- self._worker_error = None
1268
  if exc is not None:
1269
  raise RuntimeError(f"async streaming worker failed: {exc!r}") from exc
 
 
1270
 
1271
  @staticmethod
1272
  def _queue_size(q: queue.Queue | None) -> int | None:
@@ -1351,6 +1354,7 @@ class JoyOmniV2VStreamingSession:
1351
  source_metas: list[dict[str, Any]],
1352
  valid_count: int | None = None,
1353
  ) -> None:
 
1354
  if self._encode_queue is None:
1355
  self._start_async_workers()
1356
  assert self._encode_queue is not None
@@ -1368,13 +1372,8 @@ class JoyOmniV2VStreamingSession:
1368
  )
1369
  self._set_debug_state("submit", "put_encode", chunk_idx)
1370
 
1371
- while True:
1372
- self._raise_worker_error_if_needed()
1373
- try:
1374
- self._encode_queue.put(job, timeout=1.0)
1375
- break
1376
- except queue.Full:
1377
- continue
1378
  self._inc_debug_counter("submitted_chunks")
1379
  self.chunk_idx += 1
1380
 
@@ -1416,10 +1415,9 @@ class JoyOmniV2VStreamingSession:
1416
  while True:
1417
  self._set_debug_state("vae-encode", "wait_encode_queue")
1418
  _q_started = time.perf_counter()
1419
- job = self._encode_queue.get()
1420
  if job is None:
1421
  self._set_debug_state("vae-encode", "stop")
1422
- self._dit_queue.put(None)
1423
  return
1424
  self._record_queue_wait(job.profile, "q_wait_encode_s", _q_started)
1425
  try:
@@ -1438,12 +1436,11 @@ class JoyOmniV2VStreamingSession:
1438
  ready = _record_ready_event()
1439
  self._timer_record(job.profile, "reference_prepare_s", started)
1440
  self._set_debug_state("vae-encode", "put_dit_queue", job.chunk_idx)
1441
- self._dit_queue.put(_EncodedChunk(job=job, ref_chunk_latent=ref_chunk_latent, ready_event=ready))
1442
  self._inc_debug_counter("encoded_chunks")
1443
  except BaseException as exc:
1444
  self._set_debug_state("vae-encode", "error", job.chunk_idx)
1445
  self._set_worker_error(exc)
1446
- self._dit_queue.put(None)
1447
  return
1448
 
1449
  def _dit_worker(self) -> None:
@@ -1452,10 +1449,9 @@ class JoyOmniV2VStreamingSession:
1452
  while True:
1453
  self._set_debug_state("dit-denoise", "wait_dit_queue")
1454
  _q_started = time.perf_counter()
1455
- encoded = self._dit_queue.get()
1456
  if encoded is None:
1457
  self._set_debug_state("dit-denoise", "stop")
1458
- self._decode_queue.put(None)
1459
  return
1460
  self._record_queue_wait(encoded.job.profile, "q_wait_dit_s", _q_started)
1461
  try:
@@ -1479,14 +1475,13 @@ class JoyOmniV2VStreamingSession:
1479
  )
1480
  _dit_ready = _record_ready_event()
1481
  self._set_debug_state("dit-denoise", "put_decode_queue", encoded.job.chunk_idx)
1482
- self._decode_queue.put(
1483
  _DenoisedChunk(job=encoded.job, current_chunk_latents=current_chunk_latents, ready_event=_dit_ready)
1484
  )
1485
  self._inc_debug_counter("denoised_chunks")
1486
  except BaseException as exc:
1487
  self._set_debug_state("dit-denoise", "error", encoded.job.chunk_idx)
1488
  self._set_worker_error(exc)
1489
- self._decode_queue.put(None)
1490
  return
1491
 
1492
  def _decode_worker(self) -> None:
@@ -1496,11 +1491,9 @@ class JoyOmniV2VStreamingSession:
1496
  while True:
1497
  self._set_debug_state("vae-decode", "wait_decode_queue")
1498
  _q_started = time.perf_counter()
1499
- denoised = self._decode_queue.get()
1500
  if denoised is None:
1501
  self._set_debug_state("vae-decode", "stop")
1502
- self._pseudo_queue.put(None)
1503
- self._postprocess_queue.put(None)
1504
  return
1505
  self._record_queue_wait(denoised.job.profile, "q_wait_decode_s", _q_started)
1506
  try:
@@ -1519,19 +1512,17 @@ class JoyOmniV2VStreamingSession:
1519
  _dec_ready = _record_ready_event()
1520
 
1521
  self._set_debug_state("vae-decode", "put_pseudo_queue", denoised.job.chunk_idx)
1522
- self._pseudo_queue.put(
1523
  _DecodedPixelsChunk(job=denoised.job, decoded_pixels=decoded_pixels, ready_event=_dec_ready)
1524
  )
1525
  self._set_debug_state("vae-decode", "put_postprocess_queue", denoised.job.chunk_idx)
1526
- self._postprocess_queue.put(
1527
  _DecodedPixelsChunk(job=denoised.job, decoded_pixels=decoded_pixels, ready_event=_dec_ready)
1528
  )
1529
  self._inc_debug_counter("decoded_pixel_chunks")
1530
  except BaseException as exc:
1531
  self._set_debug_state("vae-decode", "error", denoised.job.chunk_idx)
1532
  self._set_worker_error(exc)
1533
- self._pseudo_queue.put(None)
1534
- self._postprocess_queue.put(None)
1535
  return
1536
 
1537
  def _pseudo_worker(self) -> None:
@@ -1539,7 +1530,7 @@ class JoyOmniV2VStreamingSession:
1539
  while True:
1540
  self._set_debug_state("pseudo-encode", "wait_pseudo_queue")
1541
  _q_started = time.perf_counter()
1542
- decoded = self._pseudo_queue.get()
1543
  if decoded is None:
1544
  self._set_debug_state("pseudo-encode", "stop")
1545
  return
@@ -1570,7 +1561,7 @@ class JoyOmniV2VStreamingSession:
1570
  while True:
1571
  self._set_debug_state("postprocess", "wait_postprocess_queue")
1572
  _q_started = time.perf_counter()
1573
- decoded = self._postprocess_queue.get()
1574
  if decoded is None:
1575
  self._set_debug_state("postprocess", "stop")
1576
  return
@@ -1600,7 +1591,7 @@ class JoyOmniV2VStreamingSession:
1600
  n_out,
1601
  )
1602
  self._set_debug_state("postprocess", "put_result_queue", decoded.job.chunk_idx)
1603
- self._result_queue.put(
1604
  StreamingChunkResult(
1605
  jpegs=packed,
1606
  profile=decoded.job.profile,
@@ -1830,14 +1821,15 @@ class JoyOmniV2VStreamingSession:
1830
 
1831
  def _next_selected_chunk_ids(self, chunk_idx: int | None = None) -> list[int]:
1832
  chunk_idx = self.chunk_idx if chunk_idx is None else chunk_idx
1833
- next_total = chunk_idx + 2
1834
- windows = self.pipeline._get_chunk_windows(
 
1835
  total_latent_frames=next_total,
1836
  chunk_size=self.chunk_size,
1837
  window_size=self.local_window_size,
1838
  global_sink_chunk=self.global_sink_chunk,
1839
  )
1840
- return list(windows[-1]["selected_chunk_ids"])
1841
 
1842
  def _resize_frame(self, frame: Image.Image) -> Image.Image:
1843
  frame = frame.convert("RGB")
 
352
  postprocess_device=postprocess_device_obj,
353
  )
354
 
355
+ if os.environ.get("JOYOMNI_SKIP_LOAD_WARMUP", "0").lower() in {"1", "true", "yes", "on"}:
356
+ print("#####[STREAM] load-time warmup skipped (JOYOMNI_SKIP_LOAD_WARMUP)")
357
+ else:
358
+ for (_wh, _ww) in _orientations:
359
+ runtime.warmup_full_pipeline(height=_wh, width=_ww)
 
 
 
 
360
 
361
  if device_obj.type == "cuda":
362
  _free_b, _total_b = torch.cuda.mem_get_info(device_obj)
 
384
  ref_image=ref_image,
385
  )
386
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
387
  def warmup_full_pipeline(
388
  self,
389
  *,
 
427
  print(f"#####[STREAM] full-pipeline warmup skipped/failed: {exc!r}")
428
  finally:
429
  if session is not None:
430
+ session.close()
 
 
 
431
 
432
  class JoyOmniV2VStreamingSession:
433
  _session_serial_counter = 0
 
507
  self._result_queue: queue.Queue[StreamingChunkResult] | None = None
508
  self._workers: list[threading.Thread] = []
509
  self._worker_error: BaseException | None = None
510
+ self._stop_event = threading.Event()
511
  self._worker_error_lock = threading.Lock()
512
  self._pseudo_latent_condition = threading.Condition()
513
  self._pseudo_latent_chunk_idx: int | None = None
 
579
  self,
580
  frame: Image.Image,
581
  frame_meta: dict[str, Any] | None = None,
582
+ *,
583
+ drain_results: bool = True,
584
  ) -> list[StreamingChunkResult]:
585
  self._raise_worker_error_if_needed()
586
  frame = self._resize_frame(frame)
 
589
  self._initialize(frame)
590
  self.pending_frames.append(frame)
591
  self.pending_metas.append(meta)
 
592
  while len(self.pending_frames) >= self.frames_per_next_chunk:
593
  n = self.frames_per_next_chunk
594
  chunk_frames = self.pending_frames[:n]
595
  chunk_metas = self.pending_metas[:n]
596
  del self.pending_frames[:n]
597
  del self.pending_metas[:n]
598
+ self._submit_async_chunk(chunk_frames, chunk_metas)
599
+ results = self._drain_async_results() if drain_results else []
600
  self._raise_worker_error_if_needed()
601
  return results
602
 
603
  @torch.no_grad()
604
  def flush_pending(self) -> None:
605
+ self._raise_worker_error_if_needed()
606
  if not self.initialized or not self.pending_frames:
607
  return
608
  valid = len(self.pending_frames)
 
611
  chunk_metas = self.pending_metas + [self.pending_metas[-1]] * pad
612
  self.pending_frames = []
613
  self.pending_metas = []
614
+ self._submit_async_chunk(chunk_frames, chunk_metas, valid_count=valid)
615
  self._raise_worker_error_if_needed()
616
 
617
+ def close(self, timeout: float = 5.0) -> None:
618
+ self._stop_async_workers(timeout)
619
+ if torch.cuda.is_available():
620
+ devices = {str(self.device), str(self.postprocess_device)}
621
+ devices.update(str(_module_device(vae)) for vae in (
622
+ self.pipeline.vae, self.decode_vae, self.pseudo_encode_vae,
623
+ ) if vae is not None)
624
+ for device in devices:
625
+ if torch.device(device).type == "cuda":
626
+ torch.cuda.synchronize(device)
627
  self._clear_vae_feature_caches()
628
  self.pipeline.transformer.reset_inference_kv_cache()
 
 
 
629
 
630
  @torch.no_grad()
631
  def _initialize(self, first_frame: Image.Image) -> None:
632
+ self._raise_worker_error_if_needed()
633
  if self.initialized:
634
  return
635
 
 
721
  return np.asarray(frame, dtype=np.float32)
722
 
723
  def _update_static_anchor(self, chunk_idx: int, source_frames: list[Image.Image]) -> int | None:
724
+ if not self.settings.freeze_kv_on_static:
725
+ return None
726
  gray = self._chunk_last_frame_gray(source_frames)
727
+ if chunk_idx == 0 or gray is None:
728
  self._prev_static_gray = gray if gray is not None else self._prev_static_gray
729
  self._static_anchor_id = None
730
  return None
 
755
  self._prev_static_gray = gray
756
  return anchor_id
757
 
 
 
 
 
 
 
 
 
 
758
  def _new_profile(self, chunk_idx: int, input_frames: int) -> dict[str, Any]:
759
  return {
760
  "chunk_idx": chunk_idx,
 
955
  )
956
 
957
  total_latent_frames = chunk_idx + 1
958
+ chunk_window = self.pipeline._get_chunk_window(
959
+ chunk_idx=chunk_idx,
960
  total_latent_frames=total_latent_frames,
961
  chunk_size=self.chunk_size,
962
  window_size=self.local_window_size,
963
  global_sink_chunk=self.global_sink_chunk,
964
  )
 
965
  selected_chunk_ids = chunk_window["selected_chunk_ids"]
966
  history_chunk_ids = selected_chunk_ids[:-1]
967
  active_chunk_id = selected_chunk_ids[-1]
 
1215
  for worker in self._workers:
1216
  worker.start()
1217
 
1218
+ def request_stop(self) -> None:
1219
+ self._stop_event.set()
1220
+ with self._pseudo_latent_condition:
1221
+ self._pseudo_latent_condition.notify_all()
1222
+
1223
+ def _queue_get(self, q):
1224
+ while not self._stop_event.is_set():
1225
+ try:
1226
+ return q.get(timeout=0.05)
1227
+ except queue.Empty:
1228
+ pass
1229
+ return None
1230
+
1231
+ def _queue_put(self, q, item) -> None:
1232
+ while not self._stop_event.is_set():
1233
+ try:
1234
+ q.put(item, timeout=0.05)
1235
+ return
1236
+ except queue.Full:
1237
+ pass
1238
+
1239
+ def _stop_async_workers(self, timeout: float = 5.0) -> None:
1240
+ self.request_stop()
1241
+ deadline = time.monotonic() + timeout
1242
  for worker in self._workers:
1243
+ worker.join(timeout=max(0.0, deadline - time.monotonic()))
1244
+ alive = [worker.name for worker in self._workers if worker.is_alive()]
1245
+ if alive:
1246
+ raise TimeoutError(f"streaming workers did not stop: {', '.join(alive)}")
1247
  self._encode_queue = None
1248
  self._dit_queue = None
1249
  self._decode_queue = None
 
1253
  self._workers = []
1254
 
1255
  def _set_worker_error(self, exc: BaseException) -> None:
1256
+ if self._stop_event.is_set():
1257
+ return
1258
  import traceback as _tb
1259
  print(f"#####[STREAM-WORKER-ERROR] {exc!r}", flush=True)
1260
  _tb.print_exc()
1261
  with self._worker_error_lock:
1262
  if self._worker_error is None:
1263
  self._worker_error = exc
1264
+ self.request_stop()
 
1265
 
1266
  def _raise_worker_error_if_needed(self) -> None:
1267
  with self._worker_error_lock:
1268
  exc = self._worker_error
 
1269
  if exc is not None:
1270
  raise RuntimeError(f"async streaming worker failed: {exc!r}") from exc
1271
+ if self._stop_event.is_set():
1272
+ raise RuntimeError("streaming session stopped")
1273
 
1274
  @staticmethod
1275
  def _queue_size(q: queue.Queue | None) -> int | None:
 
1354
  source_metas: list[dict[str, Any]],
1355
  valid_count: int | None = None,
1356
  ) -> None:
1357
+ self._raise_worker_error_if_needed()
1358
  if self._encode_queue is None:
1359
  self._start_async_workers()
1360
  assert self._encode_queue is not None
 
1372
  )
1373
  self._set_debug_state("submit", "put_encode", chunk_idx)
1374
 
1375
+ self._queue_put(self._encode_queue, job)
1376
+ self._raise_worker_error_if_needed()
 
 
 
 
 
1377
  self._inc_debug_counter("submitted_chunks")
1378
  self.chunk_idx += 1
1379
 
 
1415
  while True:
1416
  self._set_debug_state("vae-encode", "wait_encode_queue")
1417
  _q_started = time.perf_counter()
1418
+ job = self._queue_get(self._encode_queue)
1419
  if job is None:
1420
  self._set_debug_state("vae-encode", "stop")
 
1421
  return
1422
  self._record_queue_wait(job.profile, "q_wait_encode_s", _q_started)
1423
  try:
 
1436
  ready = _record_ready_event()
1437
  self._timer_record(job.profile, "reference_prepare_s", started)
1438
  self._set_debug_state("vae-encode", "put_dit_queue", job.chunk_idx)
1439
+ self._queue_put(self._dit_queue, _EncodedChunk(job=job, ref_chunk_latent=ref_chunk_latent, ready_event=ready))
1440
  self._inc_debug_counter("encoded_chunks")
1441
  except BaseException as exc:
1442
  self._set_debug_state("vae-encode", "error", job.chunk_idx)
1443
  self._set_worker_error(exc)
 
1444
  return
1445
 
1446
  def _dit_worker(self) -> None:
 
1449
  while True:
1450
  self._set_debug_state("dit-denoise", "wait_dit_queue")
1451
  _q_started = time.perf_counter()
1452
+ encoded = self._queue_get(self._dit_queue)
1453
  if encoded is None:
1454
  self._set_debug_state("dit-denoise", "stop")
 
1455
  return
1456
  self._record_queue_wait(encoded.job.profile, "q_wait_dit_s", _q_started)
1457
  try:
 
1475
  )
1476
  _dit_ready = _record_ready_event()
1477
  self._set_debug_state("dit-denoise", "put_decode_queue", encoded.job.chunk_idx)
1478
+ self._queue_put(self._decode_queue,
1479
  _DenoisedChunk(job=encoded.job, current_chunk_latents=current_chunk_latents, ready_event=_dit_ready)
1480
  )
1481
  self._inc_debug_counter("denoised_chunks")
1482
  except BaseException as exc:
1483
  self._set_debug_state("dit-denoise", "error", encoded.job.chunk_idx)
1484
  self._set_worker_error(exc)
 
1485
  return
1486
 
1487
  def _decode_worker(self) -> None:
 
1491
  while True:
1492
  self._set_debug_state("vae-decode", "wait_decode_queue")
1493
  _q_started = time.perf_counter()
1494
+ denoised = self._queue_get(self._decode_queue)
1495
  if denoised is None:
1496
  self._set_debug_state("vae-decode", "stop")
 
 
1497
  return
1498
  self._record_queue_wait(denoised.job.profile, "q_wait_decode_s", _q_started)
1499
  try:
 
1512
  _dec_ready = _record_ready_event()
1513
 
1514
  self._set_debug_state("vae-decode", "put_pseudo_queue", denoised.job.chunk_idx)
1515
+ self._queue_put(self._pseudo_queue,
1516
  _DecodedPixelsChunk(job=denoised.job, decoded_pixels=decoded_pixels, ready_event=_dec_ready)
1517
  )
1518
  self._set_debug_state("vae-decode", "put_postprocess_queue", denoised.job.chunk_idx)
1519
+ self._queue_put(self._postprocess_queue,
1520
  _DecodedPixelsChunk(job=denoised.job, decoded_pixels=decoded_pixels, ready_event=_dec_ready)
1521
  )
1522
  self._inc_debug_counter("decoded_pixel_chunks")
1523
  except BaseException as exc:
1524
  self._set_debug_state("vae-decode", "error", denoised.job.chunk_idx)
1525
  self._set_worker_error(exc)
 
 
1526
  return
1527
 
1528
  def _pseudo_worker(self) -> None:
 
1530
  while True:
1531
  self._set_debug_state("pseudo-encode", "wait_pseudo_queue")
1532
  _q_started = time.perf_counter()
1533
+ decoded = self._queue_get(self._pseudo_queue)
1534
  if decoded is None:
1535
  self._set_debug_state("pseudo-encode", "stop")
1536
  return
 
1561
  while True:
1562
  self._set_debug_state("postprocess", "wait_postprocess_queue")
1563
  _q_started = time.perf_counter()
1564
+ decoded = self._queue_get(self._postprocess_queue)
1565
  if decoded is None:
1566
  self._set_debug_state("postprocess", "stop")
1567
  return
 
1591
  n_out,
1592
  )
1593
  self._set_debug_state("postprocess", "put_result_queue", decoded.job.chunk_idx)
1594
+ self._queue_put(self._result_queue,
1595
  StreamingChunkResult(
1596
  jpegs=packed,
1597
  profile=decoded.job.profile,
 
1821
 
1822
  def _next_selected_chunk_ids(self, chunk_idx: int | None = None) -> list[int]:
1823
  chunk_idx = self.chunk_idx if chunk_idx is None else chunk_idx
1824
+ next_total = (chunk_idx + 2) * self.chunk_size
1825
+ window = self.pipeline._get_chunk_window(
1826
+ chunk_idx=chunk_idx + 1,
1827
  total_latent_frames=next_total,
1828
  chunk_size=self.chunk_size,
1829
  window_size=self.local_window_size,
1830
  global_sink_chunk=self.global_sink_chunk,
1831
  )
1832
+ return window["selected_chunk_ids"]
1833
 
1834
  def _resize_frame(self, frame: Image.Image) -> Image.Image:
1835
  frame = frame.convert("RGB")
xvideo/serving/pe.py CHANGED
@@ -1,14 +1,13 @@
 
1
  import base64
2
- import json
3
  import logging
4
  import os
5
  import re
6
- import time
7
- import urllib.request
8
  from io import BytesIO
9
  from typing import List, Optional
10
 
11
  from PIL import Image
 
12
 
13
  logger = logging.getLogger("joyomni.pe")
14
 
@@ -56,10 +55,10 @@ V2V_TEMPLATE = """# INPUT DATA
56
  background/scene replacement, no art-style or medium change (anime / painting / cartoon look and
57
  similar), no relighting or color grading beyond what a requested edit needs for integration, no
58
  camera motion, speed changes, depth-of-field effects, or extra elements on your own initiative.
59
- Unless the Target Objective names an art style, the edited video stays photorealistic
60
- live-action — never add a painting / ink / anime / cartoon style because the scene's culture or
61
- era suggests one. Use only the recipes that match the requested task; elaborate the requested
62
- edits, never add new ones.
63
  2. Priority: Describe the edits in the Target Objective's order of importance — the primary subject
64
  edit comes FIRST and carries the most detail (for a well-known character or person, spell out
65
  their canonical visual features); secondary edits such as a background swap get one concise
@@ -98,10 +97,17 @@ FOUNDATION of your prompt, and seamlessly expand them into a highly detailed, co
98
  - Background Replacement: "Replace the original background with [highly detailed description of the
99
  new environment], ensuring the foreground elements are seamlessly integrated with matching global
100
  illumination, reflections, and realistic cast shadows."
101
- Whenever this recipe is used, append this clause verbatim right after it: "the area directly
102
- behind the subject's head and shoulders shows only the new environment — the original chair and
103
- its headrest are gone." In an output that does not replace the background, that clause and any
104
- mention of removing the chair, headrest, or other scene objects are FORBIDDEN.
 
 
 
 
 
 
 
105
  - Style Transfer: "Render the scene in the style of [Style Name], featuring [2-3 concrete visual
106
  characteristics]."
107
  - Whole-frame Style Coverage: When the objective converts the entire video to an art style, the
@@ -116,9 +122,11 @@ FOUNDATION of your prompt, and seamlessly expand them into a highly detailed, co
116
  cyberpunk) is paired with "keep everything unchanged", realize the style as bold, clearly visible
117
  lighting and color grading (e.g. neon rim light, saturated color cast) on the unchanged scene —
118
  never "subtle".
119
- - Boundary: A request that redraws ONLY the background is a Background Replacement (its anchor case
120
- applies); a request that converts the ENTIRE frame to a style keeps the scene's content —
121
- background objects stay, restyled in place — and follows Whole-frame Style Coverage.
 
 
122
  - Weather/Environment: "Add [weather/season specifics] seamlessly affecting the global scene
123
  physics."
124
  - Lighting & Color Grading: "Apply cinematic relighting and color grading: [detailed description of
@@ -138,12 +146,24 @@ FOUNDATION of your prompt, and seamlessly expand them into a highly detailed, co
138
  Anchor ONLY what the Target Objective leaves untouched — an anchor must never contradict the
139
  requested edit, and preservation statements must never conflict with the Target Objective. Decide
140
  by case:
141
- - Background replaced: do NOT mention or preserve ANY object from the original background (walls,
142
- ceiling, furniture, decor, wall-mounted items — the chair or seat the subject sits on counts as
143
- background furniture too, never anchor it); the new environment must fully replace them, and the
144
- only valid anchors are foreground subject elements that survive the edit.
145
- - Background kept: the ENTIRE original background is a mandatory anchor — state that it remains
146
- unchanged, and NEVER invent or substitute a new environment.
 
 
 
 
 
 
 
 
 
 
 
 
147
  - Subject replaced or transformed: do not anchor the subject's original clothing or body — anchor
148
  only pose, motion, and what the objective explicitly keeps. When the subject turns into a
149
  different material or character, express likeness as part of the transformation ("an ice
@@ -156,6 +176,76 @@ by case:
156
  Objective names them.
157
  """
158
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
159
 
160
  def _downscale(image: Image.Image, max_side: int = PE_IMAGE_MAX_SIDE) -> Image.Image:
161
  if max_side and max_side > 0:
@@ -220,10 +310,17 @@ def _message_content_to_text(content) -> str:
220
  def _sanitize_enhanced(text: str, fallback: str) -> str:
221
  if not text:
222
  return fallback
223
- cleaned = text
224
- cleaned = re.sub(r"!\[[^\]]*\]\([^)]*\)", "", cleaned)
225
- cleaned = re.sub(r"\[([^\]]*)\]\([^)]*\)", r"\1", cleaned)
226
- cleaned = re.sub(r"https?://\S+", "", cleaned)
 
 
 
 
 
 
 
227
  cleaned = re.sub(r"[ \t]+\n", "\n", cleaned)
228
  cleaned = re.sub(r"\n{3,}", "\n\n", cleaned).strip()
229
  if len(cleaned) < 10:
@@ -258,54 +355,65 @@ class PromptEnhancer:
258
  self.base_url = base_url or os.environ.get("OPENAI_BASE_URL") or DEFAULT_BASE_URL
259
  self.model = model or os.environ.get("PE_MODEL") or DEFAULT_MODEL
260
  self.anthropic = "/anthropic" in self.base_url
261
- if not self.anthropic:
262
- from openai import OpenAI
263
 
264
- self.client = OpenAI(api_key=self.api_key, base_url=self.base_url)
265
- self.max_retries = max_retries
266
-
267
- def _anthropic_complete(self, system_prompt, user_text, images_b64) -> str:
 
 
 
268
  content = [{"type": "text", "text": user_text}]
269
  for i, b64 in enumerate(images_b64):
270
  content.append({"type": "text", "text": f"\n[Image {i}]:"})
271
  content.append({"type": "image", "source": {
272
  "type": "base64", "media_type": "image/png", "data": b64}})
273
- body = json.dumps({
274
  "model": self.model, "max_tokens": 4096, "system": system_prompt,
275
  "messages": [{"role": "user", "content": content}],
276
- }).encode()
277
- req = urllib.request.Request(
278
- f"{self.base_url.rstrip('/')}/v1/messages", data=body,
279
- headers={"Authorization": f"Bearer {self.api_key}",
280
- "anthropic-version": "2023-06-01",
281
- "Content-Type": "application/json"})
282
- resp = json.load(urllib.request.urlopen(req, timeout=90))
283
- return "".join(b.get("text", "") for b in resp.get("content", [])
284
- if b.get("type") == "text")
285
-
286
- def _chat(self, system_prompt, user_text, images_b64, raw_fallback="") -> Optional[str]:
287
- messages = None if self.anthropic else _build_messages(system_prompt, user_text, images_b64)
288
- last_err = None
289
- for attempt in range(1, self.max_retries + 1):
290
- try:
291
- if self.anthropic:
292
- text = self._anthropic_complete(system_prompt, user_text, images_b64)
293
- else:
294
- resp = self.client.chat.completions.create(
295
- model=self.model, messages=messages, max_completion_tokens=8192
296
  )
297
- text = _message_content_to_text(resp.choices[0].message.content)
298
- return _sanitize_enhanced(text.strip(), raw_fallback or text.strip())
299
- except Exception as e: # noqa: BLE001
300
- last_err = e
301
- logger.warning("PE attempt %d/%d failed: %s", attempt, self.max_retries, e)
302
- time.sleep(min(attempt, 5))
303
- logger.error("PE failed after %d attempts: %s", self.max_retries, last_err)
304
- return None
305
 
306
- def __call__(self, task_type, user_prompt, video=None) -> Optional[str]:
307
  if not user_prompt or not user_prompt.strip():
308
  return user_prompt
309
  video_frames = _video_frames_to_b64(video)
310
- text = V2V_TEMPLATE.format(user_prompt=user_prompt)
311
- return self._chat(SYSTEM_PROMPT, text, video_frames, raw_fallback=user_prompt) or user_prompt
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import asyncio
2
  import base64
 
3
  import logging
4
  import os
5
  import re
 
 
6
  from io import BytesIO
7
  from typing import List, Optional
8
 
9
  from PIL import Image
10
+ import httpx
11
 
12
  logger = logging.getLogger("joyomni.pe")
13
 
 
55
  background/scene replacement, no art-style or medium change (anime / painting / cartoon look and
56
  similar), no relighting or color grading beyond what a requested edit needs for integration, no
57
  camera motion, speed changes, depth-of-field effects, or extra elements on your own initiative.
58
+ Preserve the source video's existing visual style and medium unless the Target Objective
59
+ explicitly requests changing them. Do not introduce a new style merely because the scene's
60
+ culture or era suggests one. Use only the recipes that match the requested task; elaborate the
61
+ requested edits, never add new ones.
62
  2. Priority: Describe the edits in the Target Objective's order of importance — the primary subject
63
  edit comes FIRST and carries the most detail (for a well-known character or person, spell out
64
  their canonical visual features); secondary edits such as a background swap get one concise
 
97
  - Background Replacement: "Replace the original background with [highly detailed description of the
98
  new environment], ensuring the foreground elements are seamlessly integrated with matching global
99
  illumination, reflections, and realistic cast shadows."
100
+ Inspect the original background directly behind the subject, especially visible chairs, chair
101
+ backs and headrests closely bordering the subject's visible outline. These objects are part of
102
+ the background even when they touch the subject's outline. When present within the requested
103
+ background replacement, their removal is MANDATORY unless the Target Objective explicitly asks
104
+ to keep them. Explicitly name each visible object or component in this region and state that it
105
+ is completely removed, so the exposed area directly behind the subject shows the new environment;
106
+ a generic "replace the background" instruction is not enough. Name only components actually
107
+ visible in the source: a visible chair back does not imply a visible headrest. When none is
108
+ visible, do not mention these objects at all, even in a removal or preservation clause. Make this
109
+ decision from the image before writing; never output "if present" or "if visible" conditions.
110
+ Do not add these removals to unrelated edits unless the Target Objective requests them.
111
  - Style Transfer: "Render the scene in the style of [Style Name], featuring [2-3 concrete visual
112
  characteristics]."
113
  - Whole-frame Style Coverage: When the objective converts the entire video to an art style, the
 
122
  cyberpunk) is paired with "keep everything unchanged", realize the style as bold, clearly visible
123
  lighting and color grading (e.g. neon rim light, saturated color cast) on the unchanged scene —
124
  never "subtle".
125
+ - Boundary: Replacing the background with a different environment is Background Replacement.
126
+ Restyling the existing background preserves its scene content and layout unless the Target
127
+ Objective requests content changes; changing its appearance alone does not trigger the
128
+ background furniture-removal rule. A whole-frame style conversion restyles both the subject
129
+ and background in place and follows Whole-frame Style Coverage.
130
  - Weather/Environment: "Add [weather/season specifics] seamlessly affecting the global scene
131
  physics."
132
  - Lighting & Color Grading: "Apply cinematic relighting and color grading: [detailed description of
 
146
  Anchor ONLY what the Target Objective leaves untouched — an anchor must never contradict the
147
  requested edit, and preservation statements must never conflict with the Target Objective. Decide
148
  by case:
149
+ - Background replaced: do not preserve original background objects within the replacement scope
150
+ unless the Target Objective explicitly keeps them. Original objects directly behind the subject,
151
+ including visible chairs, chair backs and headrests along the subject's outline, remain part of
152
+ the background; explicitly remove them under the Background Replacement rule. The new
153
+ environment must fill all replaced regions, including visible gaps around the subject and frame
154
+ edges. Anchor only surviving foreground elements and objects the user explicitly keeps, and
155
+ describe only source objects actually visible in the provided frame.
156
+ - Background untouched: when no requested edit affects the background, preserve its scene content,
157
+ spatial layout and visual appearance. State that the background remains unchanged; do not invent
158
+ or substitute a new environment.
159
+ - Background appearance edited without scene replacement: for requested full-frame or
160
+ background-only stylization, weather/season changes, lighting, color grading or depth-of-field
161
+ changes, preserve the original scene's objects and spatial layout except for requested content
162
+ changes. Apply the requested visual changes to all affected background regions. Do not call the
163
+ background "unchanged" or preserve its original rendering, lighting, colors or focus when those
164
+ attributes are being edited. For whole-frame art-style conversion, use the preservation wording
165
+ in Whole-frame Style Coverage; otherwise state which scene content/layout stays and which visual
166
+ attributes change. Apply effects only within the requested scope; respect explicit exclusions.
167
  - Subject replaced or transformed: do not anchor the subject's original clothing or body — anchor
168
  only pose, motion, and what the objective explicitly keeps. When the subject turns into a
169
  different material or character, express likeness as part of the transformation ("an ice
 
176
  Objective names them.
177
  """
178
 
179
+ RV2V_SYSTEM_PROMPT = """You write precise English instructions for reference-image-guided video
180
+ editing (RV2V). Understand the source video and the separate reference image, then describe only
181
+ the requested transfer. The source video determines the subject, pose, motion, framing and scene;
182
+ the reference supplies only the visual attributes requested by the user."""
183
+
184
+ RV2V_TEMPLATE = """# INPUT
185
+ User request: {user_prompt}
186
+ Image 0 is the REFERENCE IMAGE: the donor of the requested appearance, garment, object or scene.
187
+ Any subsequent images are SOURCE VIDEO FRAMES: the video to edit.
188
+
189
+ # OUTPUT
190
+ Return ONLY one cohesive English editing prompt, about 80-140 words. No explanation or labels.
191
+ Start with a direct edit instruction, naming the target in the video and the requested item from
192
+ the reference image. Refer to it as "the reference image", never by image number in the output.
193
+ Inspect the reference and describe 3-5 distinctive visible features relevant to the request:
194
+ color, shape, material, pattern and construction. Be concrete and accurate; do not invent details.
195
+ Keep the entire request, including any explicit exceptions, extra edits and exact quoted text.
196
+
197
+ # TRANSFER SCOPE
198
+ Explicit requested edits and exceptions override the default preservation rules below. Determine
199
+ all requested changes first; preserve only attributes outside that combined edit scope.
200
+ - Clothes: identify the visible source clothing layers and the requested replacement scope before
201
+ choosing what to preserve. A request explicitly targeting an inner top, outer garment or trousers
202
+ changes only that layer; a complete outfit transfers the visible outfit. For an unqualified
203
+ request to replace the top/upper-body clothing, make the reference top the visible upper-body
204
+ garment: replace existing layers that would cover or conceal it, including outerwear or
205
+ shoulder-draped coverings, unless the user explicitly keeps them. Do not classify these clothing
206
+ layers as accessories to preserve. Top-only edits keep lower-body clothing, and trousers-only
207
+ edits keep upper-body clothing. Name the source garments/layers being replaced in the final
208
+ prompt, then describe the reference garment's cut, silhouette, length, color and visible details.
209
+ Preserve the source person's identity, face, hair, body proportions and pose within the unchanged
210
+ scope.
211
+ The reference model's face, body, pose and background are not part of a clothing transfer.
212
+ Fit the garment naturally to the source body, with correct overlap and contact at hands/arms.
213
+ - Background: replace the original environment with the environment in the reference, describing
214
+ its main visible structures. Keep the source foreground person and their clothing. Remove old
215
+ background furniture only when it is actually visible. Do not add people from the reference.
216
+ - Person/character replacement: describe the reference face/head AND the requested body/outfit.
217
+ Preserve source pose and motion, without preserving identity attributes that must change.
218
+ - Object/accessory: transfer only the requested object, at its functional position and scale.
219
+ - Style: apply only the requested visual style from the reference to the specified content.
220
+
221
+ Close with a concise preservation clause for what stays unchanged. Keep the original background
222
+ unless the user requests changing it. Do not copy reference framing, pose, lighting, typography,
223
+ logos, watermarks, props or other garments unless requested. Do not describe invisible body parts
224
+ or force the camera to reveal the whole garment. Avoid vague quality words and excessive negative
225
+ instructions. The edit is already fully present in the first output frame and follows source
226
+ motion consistently; do not describe a gradual transformation.
227
+
228
+ # SOURCE FRAMING — CLOTHING EDITS
229
+ First identify the source person's actual visible crop and pose. In a seated chest-up or
230
+ head-and-shoulders view, begin the final editing prompt by anchoring the edit to that existing
231
+ close-up and seated pose. Describe only the reference garment surfaces that can appear inside
232
+ this source crop. For a dress, name the visible upper part of the dress and describe its neckline,
233
+ shoulder fabric, sleeves and chest texture. For an outfit, describe only its visible jacket/top
234
+ and inner layer. Completely omit descriptions of offscreen garments, waist/hip shaping, skirt
235
+ length, trouser legs, hemlines, feet and full-body silhouettes from the final prompt, including
236
+ negative mentions of those details. This is a visibility rule, not a change in the user's outfit
237
+ choice. Full-body source shots still receive the full requested outfit. At each output frame,
238
+ match the person's head, shoulder and arm positions, pose, subject scale and framing to the
239
+ corresponding source-video frame; let the clothing follow the source motion with natural
240
+ deformation. Do not freeze the person in the provided frame's pose or suppress source movement.
241
+ Preserve all visible source accessories that fall outside the requested edit scope. Explicitly
242
+ apply any requested accessory removal, replacement or modification; never also describe that
243
+ accessory as unchanged. Retained accessories follow the source body's motion
244
+ and placement, with the new fabric layered naturally around them. State the actual visible crop
245
+ directly; do not output an "if visible" or "if cropped" condition. Never enlarge the visible body
246
+ region to show a garment.
247
+ """
248
+
249
 
250
  def _downscale(image: Image.Image, max_side: int = PE_IMAGE_MAX_SIDE) -> Image.Image:
251
  if max_side and max_side > 0:
 
310
  def _sanitize_enhanced(text: str, fallback: str) -> str:
311
  if not text:
312
  return fallback
313
+ def clean_link(match):
314
+ if match.group() in fallback or match.group(2) in fallback:
315
+ return match.group()
316
+ return "" if match.group().startswith("!") else match.group(1)
317
+
318
+ def clean_url(match):
319
+ url = match.group().rstrip(".,;:!?,。;:!?)")
320
+ return match.group() if url in fallback else match.group()[len(url):]
321
+
322
+ cleaned = re.sub(r"!?\[([^\]]*)\]\(([^)]*)\)", clean_link, text)
323
+ cleaned = re.sub(r'https?://[^\s<>"\'\[\]()“”]+', clean_url, cleaned)
324
  cleaned = re.sub(r"[ \t]+\n", "\n", cleaned)
325
  cleaned = re.sub(r"\n{3,}", "\n\n", cleaned).strip()
326
  if len(cleaned) < 10:
 
355
  self.base_url = base_url or os.environ.get("OPENAI_BASE_URL") or DEFAULT_BASE_URL
356
  self.model = model or os.environ.get("PE_MODEL") or DEFAULT_MODEL
357
  self.anthropic = "/anthropic" in self.base_url
358
+ self.max_retries = max(1, max_retries)
 
359
 
360
+ def _request(self, system_prompt, user_text, images_b64):
361
+ headers = {"Authorization": f"Bearer {self.api_key}"}
362
+ if not self.anthropic:
363
+ return f"{self.base_url.rstrip('/')}/chat/completions", headers, {
364
+ "model": self.model, "max_completion_tokens": 8192,
365
+ "messages": _build_messages(system_prompt, user_text, images_b64),
366
+ }
367
  content = [{"type": "text", "text": user_text}]
368
  for i, b64 in enumerate(images_b64):
369
  content.append({"type": "text", "text": f"\n[Image {i}]:"})
370
  content.append({"type": "image", "source": {
371
  "type": "base64", "media_type": "image/png", "data": b64}})
372
+ body = {
373
  "model": self.model, "max_tokens": 4096, "system": system_prompt,
374
  "messages": [{"role": "user", "content": content}],
375
+ }
376
+ headers["anthropic-version"] = "2023-06-01"
377
+ return f"{self.base_url.rstrip('/')}/v1/messages", headers, body
378
+
379
+ async def _chat(self, system_prompt, user_text, images_b64, raw_fallback):
380
+ url, headers, body = self._request(system_prompt, user_text, images_b64)
381
+ async with httpx.AsyncClient(timeout=90.0) as client:
382
+ for attempt in range(1, self.max_retries + 1):
383
+ try:
384
+ response = await client.post(url, headers=headers, json=body)
385
+ response.raise_for_status()
386
+ data = response.json()
387
+ content = data["content"] if self.anthropic else data["choices"][0]["message"]["content"]
388
+ return _sanitize_enhanced(_message_content_to_text(content).strip(), raw_fallback)
389
+ except (httpx.TransportError, httpx.HTTPStatusError) as exc:
390
+ retryable = not isinstance(exc, httpx.HTTPStatusError) or (
391
+ exc.response.status_code in (408, 409, 429) or exc.response.status_code >= 500
 
 
 
392
  )
393
+ if not retryable or attempt == self.max_retries:
394
+ raise
395
+ logger.warning("PE attempt %d/%d failed: %s", attempt, self.max_retries, type(exc).__name__)
396
+ await asyncio.sleep(min(attempt, 5))
 
 
 
 
397
 
398
+ async def enhance(self, task_type, user_prompt, video=None, ref_image=None, *, timeout=90.0):
399
  if not user_prompt or not user_prompt.strip():
400
  return user_prompt
401
  video_frames = _video_frames_to_b64(video)
402
+ if task_type == "rv2v":
403
+ reference = _img_to_b64(ref_image)
404
+ if reference is None:
405
+ logger.warning("RV2V PE needs a reference image; using raw prompt")
406
+ return user_prompt
407
+ text = RV2V_TEMPLATE.format(user_prompt=user_prompt)
408
+ system, images = RV2V_SYSTEM_PROMPT, [reference, *video_frames]
409
+ else:
410
+ text = V2V_TEMPLATE.format(user_prompt=user_prompt)
411
+ system, images = SYSTEM_PROMPT, video_frames
412
+ return await asyncio.wait_for(self._chat(system, text, images, user_prompt), timeout=timeout)
413
+
414
+ def __call__(self, task_type, user_prompt, video=None, ref_image=None, *, timeout=90.0) -> Optional[str]:
415
+ try:
416
+ return asyncio.run(self.enhance(task_type, user_prompt, video, ref_image, timeout=timeout))
417
+ except Exception as exc:
418
+ logger.warning("PE failed: %s; using raw prompt", type(exc).__name__)
419
+ return user_prompt
xvideo/serving/serve_joyomni_streaming.py CHANGED
@@ -13,6 +13,7 @@ import sys
13
  import tempfile
14
  import threading
15
  import time
 
16
  from contextlib import asynccontextmanager
17
  from fractions import Fraction
18
  from pathlib import Path
@@ -371,14 +372,16 @@ def _person_present(image: Image.Image, *, onnx_path: str, conf: float) -> bool:
371
  out = net.forward()
372
  return float(out[0, 4, :].max()) >= float(conf)
373
 
374
- def _enhance_prompt_sync(
375
  *,
376
  raw_prompt: str,
377
  pe_frame: Image.Image | None,
378
  pe_model: str | None,
 
 
379
  ) -> dict[str, Any]:
380
  started = time.time()
381
- task_type = "v2v"
382
  enhanced_prompt = raw_prompt
383
  error = None
384
  model = pe_model or DEFAULT_PE_MODEL
@@ -387,10 +390,12 @@ def _enhance_prompt_sync(
387
 
388
  enhancer = PromptEnhancer(model=pe_model)
389
  model = enhancer.model or model
390
- enhanced = enhancer(
391
  task_type,
392
  raw_prompt,
393
  video=[pe_frame] if pe_frame is not None else None,
 
 
394
  )
395
  if isinstance(enhanced, str) and enhanced.strip():
396
  enhanced_prompt = enhanced.strip()
@@ -420,11 +425,14 @@ class _SegmentedRecorder:
420
  segment_seconds: int,
421
  queue_max: int = 64,
422
  lossless: bool = False,
 
423
  ) -> None:
424
  self._prefix = prefix
425
  self._codec = codec
426
  self._bitrate = int(bitrate)
427
  self._lossless = bool(lossless)
 
 
428
  self._segment_ms = max(1, int(segment_seconds)) * 1000
429
  self._q: "queue.Queue[Any]" = queue.Queue(maxsize=max(1, queue_max))
430
  self._stop = threading.Event()
@@ -442,25 +450,37 @@ class _SegmentedRecorder:
442
  self._started = True
443
 
444
  def submit(self, item: Any, t_capture_ms: float) -> None:
445
- if not self._started or self._stop.is_set():
446
- return
447
- try:
448
- self._q.put_nowait((item, float(t_capture_ms)))
449
- except queue.Full:
450
- self.frames_dropped_recording += 1
 
 
 
 
 
 
 
 
 
 
 
 
451
 
452
  def stop(self, timeout: float = 5.0) -> None:
453
  if not self._started:
454
  return
455
- if self._stop.is_set():
456
- self._thread.join(timeout=timeout)
457
- return
458
  self._stop.set()
459
  try:
460
  self._q.put_nowait(None)
461
  except queue.Full:
462
  pass
463
  self._thread.join(timeout=timeout)
 
 
 
464
 
465
  def _to_image(self, item: Any) -> "Image.Image | None":
466
  if isinstance(item, Image.Image):
@@ -513,23 +533,16 @@ class _SegmentedRecorder:
513
  try:
514
  for packet in stream.encode():
515
  output.mux(packet)
516
- except Exception:
517
- pass
518
- try:
519
  output.close()
520
- except Exception:
521
- pass
522
 
523
  def _run(self) -> None:
524
- try:
525
- import av
526
- except Exception:
527
- return
528
  output = stream = None
529
  time_base = Fraction(1, 1000)
530
  seg_t0 = 0.0
531
  last_pts = -1
532
  try:
 
533
  while True:
534
  try:
535
  got = self._q.get(timeout=0.2)
@@ -542,7 +555,7 @@ class _SegmentedRecorder:
542
  item, t_capture_ms = got
543
  image = self._to_image(item)
544
  if image is None:
545
- continue
546
  if self._width is None:
547
  self._width, self._height = int(image.width), int(image.height)
548
  if output is None:
@@ -552,21 +565,24 @@ class _SegmentedRecorder:
552
  pts = round(t_capture_ms - seg_t0)
553
  if pts <= last_pts:
554
  pts = last_pts + 1
555
- try:
556
- frame = av.VideoFrame.from_image(image).reformat(format="yuv420p")
557
- frame.pts = pts
558
- frame.time_base = time_base
559
- for packet in stream.encode(frame):
560
- output.mux(packet)
561
- self.frames_written += 1
562
- last_pts = pts
563
- except Exception:
564
- pass
565
  if pts >= self._segment_ms:
566
- self._close_segment(output, stream)
567
  output = stream = None
 
 
 
568
  finally:
569
- self._close_segment(output, stream)
 
 
 
570
 
571
  def _optional_positive_int(value: Any, *, name: str) -> int | None:
572
  if value is None or value == "":
@@ -630,6 +646,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
630
  app.state.runtime_lock = threading.Lock()
631
  app.state.inference_lock = threading.Lock()
632
  app.state.active_session = None
 
633
  app.state.ws_debug = {}
634
 
635
  app.state.session_gate = SessionGate()
@@ -664,7 +681,8 @@ def create_app(args: argparse.Namespace) -> FastAPI:
664
  def health() -> JSONResponse:
665
  return JSONResponse(
666
  {
667
- "ok": True,
 
668
  "runtime_loaded": app.state.runtime is not None,
669
  "dit_ckpt": args.dit_ckpt,
670
  "device": str(app.state.runtime.device) if app.state.runtime is not None else args.device,
@@ -783,6 +801,12 @@ def create_app(args: argparse.Namespace) -> FastAPI:
783
  ticket = None
784
  frames_in = 0
785
  frames_out = 0
 
 
 
 
 
 
786
  session_prompt = args.prompt
787
  session_settings: StreamingSettings | None = None
788
  ref_image: Image.Image | None = None
@@ -819,7 +843,6 @@ def create_app(args: argparse.Namespace) -> FastAPI:
819
 
820
  rec_input: _SegmentedRecorder | None = None
821
  rec_output: _SegmentedRecorder | None = None
822
- rec_seq = 0
823
  rec_base: Path | None = None
824
  lossless_mode = False
825
  ws_debug: dict[str, Any] = {
@@ -842,7 +865,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
842
 
843
  async def _ws_send_json(payload: dict[str, Any]) -> None:
844
  try:
845
- await asyncio.wait_for(websocket.send_json(payload), timeout=WS_SEND_TIMEOUT_S)
846
  except asyncio.TimeoutError:
847
  raise WebSocketDisconnect()
848
 
@@ -869,7 +892,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
869
  if not encoded_frames:
870
  return 0
871
 
872
- if gate_state.get("absent_hold"):
873
  return 0
874
  count = len(encoded_frames)
875
  if not source_metas:
@@ -989,7 +1012,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
989
  ws_debug["frames_out"] = frames_out
990
  _rec_o = rec_output
991
  if _rec_o is not None:
992
- _rec_o.submit(encoded, source_meta.get("t_capture_ms"))
993
  ws_debug["rec_out_written"] = _rec_o.frames_written
994
  ws_debug["rec_out_dropped"] = _rec_o.frames_dropped_recording
995
  continue
@@ -1067,49 +1090,28 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1067
  )
1068
  return count
1069
 
1070
- async def _send_chunk_result(
1071
- result,
1072
- *,
1073
- fallback_meta: dict[str, Any] | None = None,
1074
- fallback_elapsed: float = 0.0,
1075
- ) -> int:
1076
- if not result.jpegs:
1077
- return 0
1078
  jpegs = result.jpegs
1079
- if result.source_metas:
1080
- source_metas = result.source_metas
1081
- elif fallback_meta is not None:
1082
- source_metas = [fallback_meta] * len(jpegs)
1083
- else:
1084
- source_metas = [{} for _ in jpegs]
1085
  if result.valid_count is not None:
1086
  jpegs = jpegs[: result.valid_count]
1087
  source_metas = source_metas[: result.valid_count]
1088
- server_elapsed = float(result.elapsed or fallback_elapsed)
1089
  return await _send_encoded_frames(
1090
- jpegs, source_metas, result.profile, server_elapsed,
1091
  )
1092
 
1093
  async def _output_pump(session_ref) -> None:
 
1094
  while not stop_output_pump.is_set():
1095
- try:
1096
- result = await asyncio.to_thread(session_ref.wait_async_result, 0.05)
1097
- except Exception as exc:
1098
- await _send_json({"type": "error", "message": repr(exc)})
1099
- break
1100
  if result is None:
1101
  continue
1102
- if result.jpegs:
1103
- jpegs = result.jpegs
1104
- source_metas = result.source_metas or [{} for _ in jpegs]
1105
- if result.valid_count is not None:
1106
- jpegs = jpegs[: result.valid_count]
1107
- source_metas = source_metas[: result.valid_count]
1108
- await _send_encoded_frames(
1109
- jpegs, source_metas, result.profile, float(result.elapsed or 0.0)
1110
- )
1111
 
1112
  def _create_session():
 
 
1113
  if session_settings is None:
1114
  raise RuntimeError("streaming settings are not initialized")
1115
  return runtime.create_v2v_session(
@@ -1118,65 +1120,37 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1118
  ref_image=ref_image,
1119
  )
1120
 
1121
- def _session_health_error(session_ref) -> str | None:
1122
- if session_ref is None:
1123
- return None
1124
  try:
1125
- snapshot = session_ref.debug_snapshot()
1126
- except Exception:
1127
- return None
1128
-
1129
- workers = snapshot.get("workers") or {}
1130
- errored = []
1131
- dead = []
1132
- for name, info in workers.items():
1133
- state = info.get("state") or {}
1134
- state_name = state.get("state")
1135
- if state_name == "error":
1136
- errored.append(f"{name}@chunk={state.get('chunk_idx')}")
1137
- if info.get("alive") is False:
1138
- dead.append(f"{name}:{state_name or 'unknown'}")
1139
- if errored:
1140
- return "streaming pipeline worker error: " + ", ".join(errored)
1141
-
1142
- queues = snapshot.get("queues") or {}
1143
- queue_maxsize = snapshot.get("queue_maxsize") or {}
1144
- encode_depth = queues.get("encode")
1145
- encode_max = queue_maxsize.get("encode")
1146
- encode_full = (
1147
- isinstance(encode_depth, int) and
1148
- isinstance(encode_max, int) and
1149
- encode_max > 0 and
1150
- encode_depth >= encode_max
1151
- )
1152
- if encode_full and dead:
1153
- return (
1154
- f"streaming pipeline stuck: encode queue full "
1155
- f"({encode_depth}/{encode_max}) and workers not alive: " +
1156
- ", ".join(dead)
1157
- )
1158
- return None
1159
-
1160
- def _close_session_sync(session_ref) -> None:
1161
- with app.state.inference_lock:
1162
- session_ref.close()
1163
 
1164
  async def _close_session_safely(session_ref, reason: str) -> None:
 
 
 
 
1165
  try:
1166
- await asyncio.wait_for(
1167
- asyncio.to_thread(_close_session_sync, session_ref),
1168
- timeout=max(0.1, float(args.session_close_timeout_s)),
1169
- )
1170
- except asyncio.TimeoutError:
1171
- ws_debug["close_timeout"] = reason
1172
- print(
1173
- f"#####[WS-GUARD] session close timed out after "
1174
- f"{args.session_close_timeout_s:.1f}s reason={reason}",
1175
- flush=True,
1176
- )
1177
  except Exception as exc:
1178
  ws_debug["close_error"] = repr(exc)
1179
- print(f"#####[WS-GUARD] session close failed reason={reason}: {exc!r}", flush=True)
 
 
 
1180
 
1181
  async def _stop_output_task() -> None:
1182
  nonlocal output_task
@@ -1186,23 +1160,24 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1186
  await asyncio.wait_for(output_task, timeout=1.0)
1187
  except asyncio.TimeoutError:
1188
  output_task.cancel()
 
1189
  except Exception:
1190
  pass
1191
  output_task = None
1192
 
1193
  def _start_recorders() -> None:
1194
- nonlocal rec_input, rec_output, rec_seq, rec_base
1195
  if args.record_dir is None:
1196
  return
1197
  _stop_recorders()
1198
  try:
1199
- rec_seq += 1
1200
- base = Path(args.record_dir) / f"{int(time.time())}_{rec_seq}"
1201
- base.mkdir(parents=True, exist_ok=True)
1202
  common = dict(
1203
  codec=str(args.record_codec),
1204
  bitrate=int(args.record_bitrate),
1205
  segment_seconds=int(args.record_segment_seconds),
 
1206
  )
1207
  rec_input = _SegmentedRecorder(prefix=base / "input", **common)
1208
  rec_output = _SegmentedRecorder(prefix=base / "output", lossless=lossless_mode, **common)
@@ -1221,21 +1196,25 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1221
  except Exception as exc:
1222
  ws_debug["rec_error"] = repr(exc)
1223
  print(f"#####[REC] start failed: {exc!r} -> recording OFF", flush=True)
1224
- rec_input = None
1225
- rec_output = None
1226
- rec_base = None
1227
 
1228
  def _stop_recorders() -> None:
1229
  nonlocal rec_input, rec_output, rec_base
 
1230
  for _rec in (rec_input, rec_output):
1231
  if _rec is not None:
1232
  try:
1233
- _rec.stop()
1234
  except Exception as exc:
1235
  ws_debug["rec_error"] = f"stop: {exc!r}"
 
1236
  rec_input = None
1237
  rec_output = None
1238
  rec_base = None
 
 
1239
 
1240
  def _write_prompt_sidecar() -> None:
1241
  if rec_base is None:
@@ -1259,7 +1238,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1259
  pe_task.cancel()
1260
  try:
1261
  await pe_task
1262
- except BaseException:
1263
  pass
1264
  pe_task = None
1265
 
@@ -1269,19 +1248,37 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1269
  stop_output_pump.clear()
1270
  output_task = asyncio.create_task(_output_pump(session_ref))
1271
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1272
 
1273
  async def _reset_session(reason: str) -> None:
1274
- nonlocal session, frames_since_session_reset, reset_count
1275
  if session is None:
1276
  return
1277
  print(f"#####[STREAM] session reset ({reason})", flush=True)
1278
 
 
 
1279
  await _stop_output_task()
1280
  await _close_session_safely(session, "kv_reset")
1281
  reset_count += 1
1282
  session = _create_session()
1283
  app.state.active_session = session
1284
  frames_since_session_reset = 0
 
1285
  ws_debug["kv_reset_count"] = reset_count
1286
  ws_debug["frames_since_session_reset"] = frames_since_session_reset
1287
  await _send_json(
@@ -1317,6 +1314,10 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1317
  last_activity = time.monotonic()
1318
  last_frames_out = frames_out
1319
  while True:
 
 
 
 
1320
  if frames_out != last_frames_out or pe_task is not None:
1321
  last_frames_out = frames_out
1322
  last_activity = time.monotonic()
@@ -1338,6 +1339,8 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1338
  if "text" in message and message["text"] is not None:
1339
  payload = json.loads(message["text"])
1340
  msg_type = payload.get("type")
 
 
1341
  if msg_type == "start":
1342
  print(f"#####[RESTART] 'start' received (session {'live' if session is not None else 'none'})", flush=True)
1343
  last_activity = time.monotonic()
@@ -1352,6 +1355,11 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1352
  app.state.active_session = None
1353
  session = None
1354
 
 
 
 
 
 
1355
  raw_session_prompt = str(payload.get("prompt", args.prompt))
1356
  session_prompt = raw_session_prompt
1357
  try:
@@ -1382,7 +1390,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1382
  continue
1383
 
1384
  freeze_kv_on_static = bool(
1385
- payload.get("freeze_kv_on_static", args.freeze_kv_on_static)
1386
  )
1387
  static_diff_thresh = float(
1388
  payload.get("static_diff_thresh", args.static_diff_thresh)
@@ -1392,7 +1400,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1392
  ))
1393
  use_pe = bool(payload.get("use_pe", args.use_pe)) and bool(os.environ.get("OPENAI_API_KEY"))
1394
 
1395
- entry_gate = bool(payload.get("gate_enabled", True))
1396
  face_gate_pending = entry_gate
1397
  flow["recv"] = None
1398
  flow["at"] = 0.0
@@ -1401,6 +1409,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1401
  flow["clamped"] = False
1402
  flow["consec"] = 0
1403
  flow["base"] = frames_out
 
1404
 
1405
  gate_on = bool(args.online_gate) and entry_gate
1406
 
@@ -1506,25 +1515,20 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1506
  "ref_image": ref_image is not None,
1507
  "kv_reset_frames": kv_reset_frames,
1508
  "use_pe": use_pe,
 
1509
  "pe_model": args.pe_model or DEFAULT_PE_MODEL,
1510
  "max_temporal_ids": max_temporal_ids,
1511
  "freeze_kv_on_static": freeze_kv_on_static,
1512
  "static_diff_thresh": static_diff_thresh,
1513
  }
1514
  )
1515
- _start_recorders()
1516
  _start_output_task(session)
1517
- if not pe_defer:
1518
- def _prebake(_p=session_prompt, _s=session_settings, _r=ref_image):
1519
- with app.state.inference_lock:
1520
- runtime.prebake_graph(_p, settings=_s, ref_image=_r)
1521
- asyncio.create_task(asyncio.to_thread(_prebake))
1522
  elif msg_type == "stop":
1523
  ws_debug["last_message_type"] = "stop"
1524
 
1525
  await _cancel_pe()
1526
  await _stop_output_task()
1527
- await asyncio.to_thread(_stop_recorders)
1528
  break
1529
  elif msg_type == "finalize_recording":
1530
  ws_debug["last_message_type"] = "finalize_recording"
@@ -1533,15 +1537,18 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1533
  "message": "Recording is not enabled (--record-dir is unset)."})
1534
  continue
1535
 
1536
- if session is not None:
1537
- await asyncio.to_thread(session.flush_pending)
1538
- _flush_deadline = time.monotonic() + 20.0
1539
- while frames_out < frames_in and time.monotonic() < _flush_deadline:
1540
- await asyncio.sleep(0.05)
 
 
 
 
 
 
1541
  last_activity = time.monotonic()
1542
- await _stop_output_task()
1543
- finalized = rec_base
1544
- await asyncio.to_thread(_stop_recorders)
1545
  await _send_json({
1546
  "type": "recording_finalized",
1547
  "ok": finalized is not None,
@@ -1616,6 +1623,8 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1616
  continue
1617
 
1618
  if "bytes" not in message or message["bytes"] is None:
 
 
1619
  continue
1620
  if session is None:
1621
  await _send_json({"type": "error", "message": "send start JSON before frames"})
@@ -1670,11 +1679,11 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1670
  if pe_report and pe_report.get("enhanced_prompt"):
1671
  session_prompt = pe_report["enhanced_prompt"]
1672
  await _reset_session("kv_reset_frames")
1673
- continue
1674
 
1675
  _prof_on = session_settings is not None and session_settings.profile_timings
1676
 
1677
  def _run_frame():
 
1678
  if uplink_frame is not None:
1679
  frame = uplink_frame
1680
  elif _prof_on:
@@ -1828,9 +1837,7 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1828
  gate_state["pe_anchor"] = frame
1829
  return "__gate_pe__"
1830
 
1831
- health_error = _session_health_error(session)
1832
- if health_error:
1833
- raise RuntimeError(health_error)
1834
 
1835
  acquired = app.state.inference_lock.acquire(
1836
  timeout=max(0.1, float(args.inference_lock_timeout_s))
@@ -1840,29 +1847,32 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1840
  f"inference lock timeout after {args.inference_lock_timeout_s:.1f}s"
1841
  )
1842
  try:
1843
- health_error = _session_health_error(session)
1844
- if health_error:
1845
- raise RuntimeError(health_error)
1846
-
1847
- if session_max_inflight and session.inflight_chunks() >= session_max_inflight:
1848
  ws_debug["frames_dropped_backpressure"] = (
1849
  int(ws_debug.get("frames_dropped_backpressure", 0)) + 1
1850
  )
1851
- return session._drain_async_results()
1852
 
1853
  _rec_i = rec_input
1854
  if _rec_i is not None:
1855
  _rec_i.submit(frame, frame_meta.get("t_capture_ms"))
1856
  ws_debug["rec_in_written"] = _rec_i.frames_written
1857
  ws_debug["rec_in_dropped"] = _rec_i.frames_dropped_recording
1858
- return session.push_frame(frame, frame_meta=frame_meta)
 
 
1859
  finally:
1860
  app.state.inference_lock.release()
1861
 
1862
- started = time.time()
1863
  try:
1864
- chunk_results = await asyncio.wait_for(
1865
- asyncio.to_thread(_run_frame),
 
 
 
 
 
 
1866
  timeout=max(0.1, float(args.push_frame_timeout_s)),
1867
  )
1868
  except asyncio.TimeoutError:
@@ -1901,15 +1911,13 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1901
  gate_state["pe_anchor"] = None
1902
  await _send_json({"type": "pe_running", "frames_in": frames_in})
1903
 
1904
- async def _run_pe(_sess=session, _anchor=anchor, _raw=raw_session_prompt):
1905
  nonlocal pe_defer, pe_report, pe_task
1906
  try:
1907
- report = await asyncio.wait_for(
1908
- asyncio.to_thread(
1909
- _enhance_prompt_sync,
1910
- raw_prompt=_raw,
1911
- pe_frame=_anchor, pe_model=args.pe_model,
1912
- ),
1913
  timeout=max(1.0, float(args.pe_timeout_s)),
1914
  )
1915
  txt = str(report.get("enhanced_prompt") or _raw)
@@ -1918,18 +1926,13 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1918
  with app.state.inference_lock:
1919
  _sess.prompt = txt
1920
  _sess._initialize(_anchor)
1921
- await asyncio.to_thread(_swap)
1922
  pe_report = report
1923
  ws_debug["pe_report"] = report
1924
  await _send_json({"type": "prompt_enhanced", **report})
1925
- except Exception:
1926
- def _swap_raw():
1927
- with app.state.inference_lock:
1928
- _sess._initialize(_anchor)
1929
- await asyncio.to_thread(_swap_raw)
1930
- await _send_json({"type": "prompt_enhanced", "enabled": True,
1931
- "raw_prompt": _raw, "enhanced_prompt": _raw,
1932
- "fallback": True, "error": True, "elapsed_s": 0.0})
1933
  finally:
1934
  pe_defer = False
1935
  pe_task = None
@@ -1954,44 +1957,21 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1954
  session_prompt = pe_report["enhanced_prompt"]
1955
  await _reset_session("person_count_changed")
1956
  continue
1957
- elapsed = time.time() - started
1958
- if not chunk_results:
1959
- await _send_json(
1960
- {
1961
- "type": "accepted",
1962
- "frames_in": frames_in,
1963
- "frames_out": frames_out,
1964
- "next_chunk_needs": session.frames_per_next_chunk - len(session.pending_frames),
1965
- }
1966
- )
1967
- continue
1968
-
1969
- last_count = 0
1970
- for result in chunk_results:
1971
- last_count = await _send_chunk_result(
1972
- result,
1973
- fallback_meta=frame_meta,
1974
- fallback_elapsed=elapsed,
1975
- )
1976
- if last_count == 0:
1977
- await _send_json(
1978
- {
1979
- "type": "accepted",
1980
- "frames_in": frames_in,
1981
- "frames_out": frames_out,
1982
- "next_chunk_needs": session.frames_per_next_chunk - len(session.pending_frames),
1983
- }
1984
- )
1985
  except WebSocketDisconnect:
1986
  pass
1987
- except RuntimeError as exc:
1988
- if "disconnect" not in str(exc).lower():
1989
- raise
 
 
1990
  finally:
1991
  try:
1992
  await _cancel_pe()
1993
  await _stop_output_task()
1994
- await asyncio.to_thread(_stop_recorders)
1995
  if session is not None:
1996
  await _close_session_safely(session, "finally")
1997
  if getattr(app.state, "active_session", None) is session:
@@ -1999,6 +1979,10 @@ def create_app(args: argparse.Namespace) -> FastAPI:
1999
  ws_debug["closed_at"] = time.time()
2000
  ws_debug["send_state"] = "closed"
2001
  finally:
 
 
 
 
2002
  if ticket is not None:
2003
  app.state.session_gate.release(ticket)
2004
 
@@ -2056,13 +2040,13 @@ def build_parser() -> argparse.ArgumentParser:
2056
  parser.add_argument("--prompt", type=str, default="Keep the person and scene temporally consistent while applying the requested edit.")
2057
  parser.add_argument("--profile-timings", action=argparse.BooleanOptionalAction, default=False)
2058
  parser.add_argument("--kv-reset-frames", type=int, default=600)
2059
- parser.add_argument("--max-temporal-ids", type=int, default=None)
2060
  parser.add_argument("--freeze-kv-on-static", action=argparse.BooleanOptionalAction, default=True)
2061
  parser.add_argument("--static-diff-thresh", type=float, default=0.5)
2062
  parser.add_argument("--preload", action=argparse.BooleanOptionalAction, default=True)
2063
- parser.add_argument("--push-frame-timeout-s", type=float, default=15.0, help="Max seconds a single frame submission may block before releasing the WS session gate.")
2064
- parser.add_argument("--max-inflight-chunks", type=int, default=2, help="Drop incoming frames while this many chunks are already in the pipeline (latency governor; pins display latency to ~N chunk periods). 2 keeps glass-to-glass latency low (~1 chunk) and prevents the session-start init stall from building a frame backlog (the 'ramp-up'); higher values buffer more (smoother under hiccups) at the cost of latency and a startup ramp. 0 disables. Overridable per session via the start payload's max_inflight_chunks.")
2065
- parser.add_argument("--session-close-timeout-s", type=float, default=5.0, help="Best-effort session cleanup timeout during WS teardown.")
2066
  parser.add_argument("--inference-lock-timeout-s", type=float, default=5.0, help="Max seconds to wait for the process-wide inference lock.")
2067
 
2068
  parser.add_argument("--record-dir", type=str, default=None, help="Directory to record input/output mp4s into (per-session subfolder). Off if unset.")
 
13
  import tempfile
14
  import threading
15
  import time
16
+ import uuid
17
  from contextlib import asynccontextmanager
18
  from fractions import Fraction
19
  from pathlib import Path
 
372
  out = net.forward()
373
  return float(out[0, 4, :].max()) >= float(conf)
374
 
375
+ async def _enhance_prompt(
376
  *,
377
  raw_prompt: str,
378
  pe_frame: Image.Image | None,
379
  pe_model: str | None,
380
+ timeout: float,
381
+ ref_image: Image.Image | None = None,
382
  ) -> dict[str, Any]:
383
  started = time.time()
384
+ task_type = "rv2v" if ref_image is not None else "v2v"
385
  enhanced_prompt = raw_prompt
386
  error = None
387
  model = pe_model or DEFAULT_PE_MODEL
 
390
 
391
  enhancer = PromptEnhancer(model=pe_model)
392
  model = enhancer.model or model
393
+ enhanced = await enhancer.enhance(
394
  task_type,
395
  raw_prompt,
396
  video=[pe_frame] if pe_frame is not None else None,
397
+ ref_image=ref_image,
398
+ timeout=timeout,
399
  )
400
  if isinstance(enhanced, str) and enhanced.strip():
401
  enhanced_prompt = enhanced.strip()
 
425
  segment_seconds: int,
426
  queue_max: int = 64,
427
  lossless: bool = False,
428
+ reliable: bool = False,
429
  ) -> None:
430
  self._prefix = prefix
431
  self._codec = codec
432
  self._bitrate = int(bitrate)
433
  self._lossless = bool(lossless)
434
+ self._reliable = bool(reliable)
435
+ self._error: BaseException | None = None
436
  self._segment_ms = max(1, int(segment_seconds)) * 1000
437
  self._q: "queue.Queue[Any]" = queue.Queue(maxsize=max(1, queue_max))
438
  self._stop = threading.Event()
 
450
  self._started = True
451
 
452
  def submit(self, item: Any, t_capture_ms: float) -> None:
453
+ deadline = time.monotonic() + 20.0
454
+ while True:
455
+ self._raise_error()
456
+ if not self._started or self._stop.is_set():
457
+ raise RuntimeError("recorder is not accepting frames")
458
+ try:
459
+ self._q.put((item, float(t_capture_ms)), timeout=0.05 if self._reliable else 0)
460
+ return
461
+ except queue.Full:
462
+ if not self._reliable:
463
+ self.frames_dropped_recording += 1
464
+ return
465
+ if time.monotonic() >= deadline:
466
+ raise TimeoutError("recording queue remained full")
467
+
468
+ def _raise_error(self) -> None:
469
+ if self._error is not None:
470
+ raise RuntimeError(f"recording failed: {self._error!r}") from self._error
471
 
472
  def stop(self, timeout: float = 5.0) -> None:
473
  if not self._started:
474
  return
 
 
 
475
  self._stop.set()
476
  try:
477
  self._q.put_nowait(None)
478
  except queue.Full:
479
  pass
480
  self._thread.join(timeout=timeout)
481
+ if self._thread.is_alive():
482
+ raise TimeoutError("recorder did not finish encoding")
483
+ self._raise_error()
484
 
485
  def _to_image(self, item: Any) -> "Image.Image | None":
486
  if isinstance(item, Image.Image):
 
533
  try:
534
  for packet in stream.encode():
535
  output.mux(packet)
536
+ finally:
 
 
537
  output.close()
 
 
538
 
539
  def _run(self) -> None:
 
 
 
 
540
  output = stream = None
541
  time_base = Fraction(1, 1000)
542
  seg_t0 = 0.0
543
  last_pts = -1
544
  try:
545
+ import av
546
  while True:
547
  try:
548
  got = self._q.get(timeout=0.2)
 
555
  item, t_capture_ms = got
556
  image = self._to_image(item)
557
  if image is None:
558
+ raise ValueError("invalid recording frame")
559
  if self._width is None:
560
  self._width, self._height = int(image.width), int(image.height)
561
  if output is None:
 
565
  pts = round(t_capture_ms - seg_t0)
566
  if pts <= last_pts:
567
  pts = last_pts + 1
568
+ frame = av.VideoFrame.from_image(image).reformat(format="yuv420p")
569
+ frame.pts = pts
570
+ frame.time_base = time_base
571
+ for packet in stream.encode(frame):
572
+ output.mux(packet)
573
+ self.frames_written += 1
574
+ last_pts = pts
 
 
 
575
  if pts >= self._segment_ms:
576
+ previous_output, previous_stream = output, stream
577
  output = stream = None
578
+ self._close_segment(previous_output, previous_stream)
579
+ except BaseException as exc:
580
+ self._error = exc
581
  finally:
582
+ try:
583
+ self._close_segment(output, stream)
584
+ except BaseException as exc:
585
+ self._error = exc
586
 
587
  def _optional_positive_int(value: Any, *, name: str) -> int | None:
588
  if value is None or value == "":
 
646
  app.state.runtime_lock = threading.Lock()
647
  app.state.inference_lock = threading.Lock()
648
  app.state.active_session = None
649
+ app.state.runtime_error = None
650
  app.state.ws_debug = {}
651
 
652
  app.state.session_gate = SessionGate()
 
681
  def health() -> JSONResponse:
682
  return JSONResponse(
683
  {
684
+ "ok": app.state.runtime_error is None,
685
+ "error": app.state.runtime_error,
686
  "runtime_loaded": app.state.runtime is not None,
687
  "dit_ckpt": args.dit_ckpt,
688
  "device": str(app.state.runtime.device) if app.state.runtime is not None else args.device,
 
801
  ticket = None
802
  frames_in = 0
803
  frames_out = 0
804
+ session_id = None
805
+ edit_frames = 0
806
+ completed_frames = 0
807
+ pending_work: set[asyncio.Task] = set()
808
+ pe_report = None
809
+ raw_session_prompt = args.prompt
810
  session_prompt = args.prompt
811
  session_settings: StreamingSettings | None = None
812
  ref_image: Image.Image | None = None
 
843
 
844
  rec_input: _SegmentedRecorder | None = None
845
  rec_output: _SegmentedRecorder | None = None
 
846
  rec_base: Path | None = None
847
  lossless_mode = False
848
  ws_debug: dict[str, Any] = {
 
865
 
866
  async def _ws_send_json(payload: dict[str, Any]) -> None:
867
  try:
868
+ await asyncio.wait_for(websocket.send_json({**payload, "session_id": session_id}), timeout=WS_SEND_TIMEOUT_S)
869
  except asyncio.TimeoutError:
870
  raise WebSocketDisconnect()
871
 
 
892
  if not encoded_frames:
893
  return 0
894
 
895
+ if not lossless_mode and gate_state.get("absent_hold"):
896
  return 0
897
  count = len(encoded_frames)
898
  if not source_metas:
 
1012
  ws_debug["frames_out"] = frames_out
1013
  _rec_o = rec_output
1014
  if _rec_o is not None:
1015
+ await asyncio.to_thread(_rec_o.submit, encoded, source_meta.get("t_capture_ms"))
1016
  ws_debug["rec_out_written"] = _rec_o.frames_written
1017
  ws_debug["rec_out_dropped"] = _rec_o.frames_dropped_recording
1018
  continue
 
1090
  )
1091
  return count
1092
 
1093
+ async def _send_chunk_result(result) -> int:
 
 
 
 
 
 
 
1094
  jpegs = result.jpegs
1095
+ source_metas = result.source_metas or [{} for _ in jpegs]
 
 
 
 
 
1096
  if result.valid_count is not None:
1097
  jpegs = jpegs[: result.valid_count]
1098
  source_metas = source_metas[: result.valid_count]
 
1099
  return await _send_encoded_frames(
1100
+ jpegs, source_metas, result.profile, float(result.elapsed or 0.0),
1101
  )
1102
 
1103
  async def _output_pump(session_ref) -> None:
1104
+ nonlocal completed_frames
1105
  while not stop_output_pump.is_set():
1106
+ result = await asyncio.to_thread(session_ref.wait_async_result, 0.05)
 
 
 
 
1107
  if result is None:
1108
  continue
1109
+ await _send_chunk_result(result)
1110
+ completed_frames += min(len(result.jpegs), result.valid_count) if result.valid_count is not None else len(result.jpegs)
 
 
 
 
 
 
 
1111
 
1112
  def _create_session():
1113
+ if app.state.runtime_error:
1114
+ raise RuntimeError(app.state.runtime_error)
1115
  if session_settings is None:
1116
  raise RuntimeError("streaming settings are not initialized")
1117
  return runtime.create_v2v_session(
 
1120
  ref_image=ref_image,
1121
  )
1122
 
1123
+ async def _run_work(func, *, timeout=None):
1124
+ task = asyncio.create_task(asyncio.to_thread(func))
1125
+ pending_work.add(task)
1126
  try:
1127
+ return await asyncio.wait_for(asyncio.shield(task), timeout=timeout)
1128
+ finally:
1129
+ if task.done():
1130
+ pending_work.discard(task)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1131
 
1132
  async def _close_session_safely(session_ref, reason: str) -> None:
1133
+ if app.state.runtime_error:
1134
+ return
1135
+ session_ref.request_stop()
1136
+ deadline = time.monotonic() + max(0.1, float(args.session_close_timeout_s))
1137
  try:
1138
+ if pending_work:
1139
+ done, waiting = await asyncio.wait(pending_work, timeout=max(0.0, deadline - time.monotonic()))
1140
+ for task in done:
1141
+ if not task.cancelled():
1142
+ task.exception()
1143
+ pending_work.difference_update(done)
1144
+ if waiting:
1145
+ raise TimeoutError("session operations have not stopped")
1146
+ remaining = max(0.01, deadline - time.monotonic())
1147
+ await _run_work(lambda: session_ref.close(timeout=remaining), timeout=remaining)
 
1148
  except Exception as exc:
1149
  ws_debug["close_error"] = repr(exc)
1150
+ app.state.runtime_error = f"Unsafe runtime after {reason}: {exc!r}; restart required"
1151
+ for task in pending_work:
1152
+ task.add_done_callback(lambda done: done.exception() if not done.cancelled() else None)
1153
+ raise RuntimeError(app.state.runtime_error) from exc
1154
 
1155
  async def _stop_output_task() -> None:
1156
  nonlocal output_task
 
1160
  await asyncio.wait_for(output_task, timeout=1.0)
1161
  except asyncio.TimeoutError:
1162
  output_task.cancel()
1163
+ await asyncio.gather(output_task, return_exceptions=True)
1164
  except Exception:
1165
  pass
1166
  output_task = None
1167
 
1168
  def _start_recorders() -> None:
1169
+ nonlocal rec_input, rec_output, rec_base
1170
  if args.record_dir is None:
1171
  return
1172
  _stop_recorders()
1173
  try:
1174
+ base = Path(args.record_dir) / f"{int(time.time())}_{uuid.uuid4().int}"
1175
+ base.mkdir(parents=True)
 
1176
  common = dict(
1177
  codec=str(args.record_codec),
1178
  bitrate=int(args.record_bitrate),
1179
  segment_seconds=int(args.record_segment_seconds),
1180
+ reliable=lossless_mode,
1181
  )
1182
  rec_input = _SegmentedRecorder(prefix=base / "input", **common)
1183
  rec_output = _SegmentedRecorder(prefix=base / "output", lossless=lossless_mode, **common)
 
1196
  except Exception as exc:
1197
  ws_debug["rec_error"] = repr(exc)
1198
  print(f"#####[REC] start failed: {exc!r} -> recording OFF", flush=True)
1199
+ _stop_recorders()
1200
+ if lossless_mode:
1201
+ raise
1202
 
1203
  def _stop_recorders() -> None:
1204
  nonlocal rec_input, rec_output, rec_base
1205
+ error = None
1206
  for _rec in (rec_input, rec_output):
1207
  if _rec is not None:
1208
  try:
1209
+ _rec.stop(timeout=20.0 if lossless_mode else 5.0)
1210
  except Exception as exc:
1211
  ws_debug["rec_error"] = f"stop: {exc!r}"
1212
+ error = exc
1213
  rec_input = None
1214
  rec_output = None
1215
  rec_base = None
1216
+ if error is not None:
1217
+ raise error
1218
 
1219
  def _write_prompt_sidecar() -> None:
1220
  if rec_base is None:
 
1238
  pe_task.cancel()
1239
  try:
1240
  await pe_task
1241
+ except asyncio.CancelledError:
1242
  pass
1243
  pe_task = None
1244
 
 
1248
  stop_output_pump.clear()
1249
  output_task = asyncio.create_task(_output_pump(session_ref))
1250
 
1251
+ def _check_output() -> None:
1252
+ if output_task is not None and output_task.done():
1253
+ output_task.result()
1254
+ raise RuntimeError("output processing stopped before completion")
1255
+ session._raise_worker_error_if_needed()
1256
+
1257
+ async def _drain_session() -> None:
1258
+ await _run_work(session.flush_pending, timeout=max(0.1, float(args.push_frame_timeout_s)))
1259
+ deadline = time.monotonic() + 20.0
1260
+ while completed_frames < edit_frames:
1261
+ _check_output()
1262
+ if time.monotonic() >= deadline:
1263
+ raise TimeoutError(f"incomplete output: {completed_frames}/{edit_frames} frames")
1264
+ await asyncio.sleep(0.01)
1265
+ _check_output()
1266
 
1267
  async def _reset_session(reason: str) -> None:
1268
+ nonlocal session, frames_since_session_reset, reset_count, edit_frames, completed_frames
1269
  if session is None:
1270
  return
1271
  print(f"#####[STREAM] session reset ({reason})", flush=True)
1272
 
1273
+ if lossless_mode:
1274
+ await _drain_session()
1275
  await _stop_output_task()
1276
  await _close_session_safely(session, "kv_reset")
1277
  reset_count += 1
1278
  session = _create_session()
1279
  app.state.active_session = session
1280
  frames_since_session_reset = 0
1281
+ edit_frames = completed_frames = 0
1282
  ws_debug["kv_reset_count"] = reset_count
1283
  ws_debug["frames_since_session_reset"] = frames_since_session_reset
1284
  await _send_json(
 
1314
  last_activity = time.monotonic()
1315
  last_frames_out = frames_out
1316
  while True:
1317
+ if app.state.runtime_error:
1318
+ raise RuntimeError(app.state.runtime_error)
1319
+ if output_task is not None and output_task.done():
1320
+ output_task.result()
1321
  if frames_out != last_frames_out or pe_task is not None:
1322
  last_frames_out = frames_out
1323
  last_activity = time.monotonic()
 
1339
  if "text" in message and message["text"] is not None:
1340
  payload = json.loads(message["text"])
1341
  msg_type = payload.get("type")
1342
+ if msg_type != "start" and payload.get("session_id", session_id) != session_id:
1343
+ continue
1344
  if msg_type == "start":
1345
  print(f"#####[RESTART] 'start' received (session {'live' if session is not None else 'none'})", flush=True)
1346
  last_activity = time.monotonic()
 
1355
  app.state.active_session = None
1356
  session = None
1357
 
1358
+ session_id = str(payload.get("session_id") or uuid.uuid4().hex)
1359
+ frames_in = frames_out = edit_frames = completed_frames = last_frames_out = 0
1360
+ next_frame_meta = None
1361
+ ws_debug.update(frames_in=0, frames_out=0, output_bytes=0, chunk_results_sent=0,
1362
+ frames_dropped_backpressure=0, session_id=session_id)
1363
  raw_session_prompt = str(payload.get("prompt", args.prompt))
1364
  session_prompt = raw_session_prompt
1365
  try:
 
1390
  continue
1391
 
1392
  freeze_kv_on_static = bool(
1393
+ payload.get("freeze_kv_on_static", False if lossless_mode else args.freeze_kv_on_static)
1394
  )
1395
  static_diff_thresh = float(
1396
  payload.get("static_diff_thresh", args.static_diff_thresh)
 
1400
  ))
1401
  use_pe = bool(payload.get("use_pe", args.use_pe)) and bool(os.environ.get("OPENAI_API_KEY"))
1402
 
1403
+ entry_gate = bool(payload.get("gate_enabled", True)) and not lossless_mode
1404
  face_gate_pending = entry_gate
1405
  flow["recv"] = None
1406
  flow["at"] = 0.0
 
1409
  flow["clamped"] = False
1410
  flow["consec"] = 0
1411
  flow["base"] = frames_out
1412
+ flow.update(has_ack=False, dropped=0, skew_min=None, skew_at=0.0, up_ms=0.0)
1413
 
1414
  gate_on = bool(args.online_gate) and entry_gate
1415
 
 
1515
  "ref_image": ref_image is not None,
1516
  "kv_reset_frames": kv_reset_frames,
1517
  "use_pe": use_pe,
1518
+ "pe_deferred": pe_defer,
1519
  "pe_model": args.pe_model or DEFAULT_PE_MODEL,
1520
  "max_temporal_ids": max_temporal_ids,
1521
  "freeze_kv_on_static": freeze_kv_on_static,
1522
  "static_diff_thresh": static_diff_thresh,
1523
  }
1524
  )
1525
+ await asyncio.to_thread(_start_recorders)
1526
  _start_output_task(session)
 
 
 
 
 
1527
  elif msg_type == "stop":
1528
  ws_debug["last_message_type"] = "stop"
1529
 
1530
  await _cancel_pe()
1531
  await _stop_output_task()
 
1532
  break
1533
  elif msg_type == "finalize_recording":
1534
  ws_debug["last_message_type"] = "finalize_recording"
 
1537
  "message": "Recording is not enabled (--record-dir is unset)."})
1538
  continue
1539
 
1540
+ try:
1541
+ if pe_task is not None:
1542
+ await pe_task
1543
+ if session is not None:
1544
+ await _drain_session()
1545
+ await _stop_output_task()
1546
+ finalized = rec_base
1547
+ await asyncio.to_thread(_stop_recorders)
1548
+ except Exception as exc:
1549
+ await _send_json({"type": "recording_finalized", "ok": False, "message": str(exc)})
1550
+ break
1551
  last_activity = time.monotonic()
 
 
 
1552
  await _send_json({
1553
  "type": "recording_finalized",
1554
  "ok": finalized is not None,
 
1623
  continue
1624
 
1625
  if "bytes" not in message or message["bytes"] is None:
1626
+ if message.get("type") == "websocket.disconnect":
1627
+ break
1628
  continue
1629
  if session is None:
1630
  await _send_json({"type": "error", "message": "send start JSON before frames"})
 
1679
  if pe_report and pe_report.get("enhanced_prompt"):
1680
  session_prompt = pe_report["enhanced_prompt"]
1681
  await _reset_session("kv_reset_frames")
 
1682
 
1683
  _prof_on = session_settings is not None and session_settings.profile_timings
1684
 
1685
  def _run_frame():
1686
+ nonlocal edit_frames
1687
  if uplink_frame is not None:
1688
  frame = uplink_frame
1689
  elif _prof_on:
 
1837
  gate_state["pe_anchor"] = frame
1838
  return "__gate_pe__"
1839
 
1840
+ session._raise_worker_error_if_needed()
 
 
1841
 
1842
  acquired = app.state.inference_lock.acquire(
1843
  timeout=max(0.1, float(args.inference_lock_timeout_s))
 
1847
  f"inference lock timeout after {args.inference_lock_timeout_s:.1f}s"
1848
  )
1849
  try:
1850
+ if not lossless_mode and session_max_inflight and session.inflight_chunks() >= session_max_inflight:
 
 
 
 
1851
  ws_debug["frames_dropped_backpressure"] = (
1852
  int(ws_debug.get("frames_dropped_backpressure", 0)) + 1
1853
  )
1854
+ return []
1855
 
1856
  _rec_i = rec_input
1857
  if _rec_i is not None:
1858
  _rec_i.submit(frame, frame_meta.get("t_capture_ms"))
1859
  ws_debug["rec_in_written"] = _rec_i.frames_written
1860
  ws_debug["rec_in_dropped"] = _rec_i.frames_dropped_recording
1861
+ session.push_frame(frame, frame_meta=frame_meta, drain_results=False)
1862
+ edit_frames += 1
1863
+ return []
1864
  finally:
1865
  app.state.inference_lock.release()
1866
 
 
1867
  try:
1868
+ capacity_deadline = time.monotonic() + max(0.1, float(args.push_frame_timeout_s))
1869
+ while lossless_mode and session_max_inflight and session.inflight_chunks() >= session_max_inflight:
1870
+ _check_output()
1871
+ if time.monotonic() >= capacity_deadline:
1872
+ raise TimeoutError("file input timed out waiting for inference capacity")
1873
+ await asyncio.sleep(0.01)
1874
+ chunk_results = await _run_work(
1875
+ _run_frame,
1876
  timeout=max(0.1, float(args.push_frame_timeout_s)),
1877
  )
1878
  except asyncio.TimeoutError:
 
1911
  gate_state["pe_anchor"] = None
1912
  await _send_json({"type": "pe_running", "frames_in": frames_in})
1913
 
1914
+ async def _run_pe(_sess=session, _anchor=anchor, _raw=raw_session_prompt, _ref=ref_image):
1915
  nonlocal pe_defer, pe_report, pe_task
1916
  try:
1917
+ report = await _enhance_prompt(
1918
+ raw_prompt=_raw,
1919
+ pe_frame=_anchor, pe_model=args.pe_model,
1920
+ ref_image=_ref,
 
 
1921
  timeout=max(1.0, float(args.pe_timeout_s)),
1922
  )
1923
  txt = str(report.get("enhanced_prompt") or _raw)
 
1926
  with app.state.inference_lock:
1927
  _sess.prompt = txt
1928
  _sess._initialize(_anchor)
1929
+ await _run_work(_swap, timeout=max(0.1, float(args.push_frame_timeout_s)))
1930
  pe_report = report
1931
  ws_debug["pe_report"] = report
1932
  await _send_json({"type": "prompt_enhanced", **report})
1933
+ except Exception as exc:
1934
+ _sess.request_stop()
1935
+ await _send_json({"type": "error", "message": str(exc)})
 
 
 
 
 
1936
  finally:
1937
  pe_defer = False
1938
  pe_task = None
 
1957
  session_prompt = pe_report["enhanced_prompt"]
1958
  await _reset_session("person_count_changed")
1959
  continue
1960
+ await _send_json({
1961
+ "type": "accepted", "frames_in": frames_in, "frames_out": frames_out,
1962
+ "next_chunk_needs": session.frames_per_next_chunk - len(session.pending_frames),
1963
+ })
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1964
  except WebSocketDisconnect:
1965
  pass
1966
+ except Exception as exc:
1967
+ try:
1968
+ await _send_json({"type": "error", "message": str(exc)})
1969
+ except (WebSocketDisconnect, RuntimeError):
1970
+ pass
1971
  finally:
1972
  try:
1973
  await _cancel_pe()
1974
  await _stop_output_task()
 
1975
  if session is not None:
1976
  await _close_session_safely(session, "finally")
1977
  if getattr(app.state, "active_session", None) is session:
 
1979
  ws_debug["closed_at"] = time.time()
1980
  ws_debug["send_state"] = "closed"
1981
  finally:
1982
+ try:
1983
+ await asyncio.to_thread(_stop_recorders)
1984
+ except Exception as exc:
1985
+ ws_debug["rec_error"] = repr(exc)
1986
  if ticket is not None:
1987
  app.state.session_gate.release(ticket)
1988
 
 
2040
  parser.add_argument("--prompt", type=str, default="Keep the person and scene temporally consistent while applying the requested edit.")
2041
  parser.add_argument("--profile-timings", action=argparse.BooleanOptionalAction, default=False)
2042
  parser.add_argument("--kv-reset-frames", type=int, default=600)
2043
+ parser.add_argument("--max-temporal-ids", type=int, default=8)
2044
  parser.add_argument("--freeze-kv-on-static", action=argparse.BooleanOptionalAction, default=True)
2045
  parser.add_argument("--static-diff-thresh", type=float, default=0.5)
2046
  parser.add_argument("--preload", action=argparse.BooleanOptionalAction, default=True)
2047
+ parser.add_argument("--push-frame-timeout-s", type=float, default=15.0, help="Timeout for frame submission or file-input backpressure; teardown must finish before model reuse.")
2048
+ parser.add_argument("--max-inflight-chunks", type=int, default=2, help="Limit chunks in flight: camera input drops excess frames, file input waits for capacity. 0 disables; overridable per session.")
2049
+ parser.add_argument("--session-close-timeout-s", type=float, default=5.0, help="Session cleanup timeout. If workers or GPU work remain active, reject new sessions until restart.")
2050
  parser.add_argument("--inference-lock-timeout-s", type=float, default=5.0, help="Max seconds to wait for the process-wide inference lock.")
2051
 
2052
  parser.add_argument("--record-dir", type=str, default=None, help="Directory to record input/output mp4s into (per-session subfolder). Off if unset.")
xvideo/serving/zerogpu_engine.py CHANGED
@@ -1,57 +1,18 @@
1
- """ZeroGPU session engine for JoyAI-Video-Edit.
2
-
3
- NEW FILE — does not modify the existing standalone server
4
- (serve_joyomni_streaming.py). It runs INSIDE the @spaces.GPU fork and reuses the
5
- vendored gate / prompt-enhancement / streaming logic to drive one browser session
6
- whose WebSocket transport has been replaced by two multiprocessing.Queues.
7
-
8
- The WS-facing loop from the standalone server is deeply coupled to the live
9
- `session` object (pending_frames, inference_lock, _encode_streaming_prompt, ...).
10
- Rather than bridge those across processes, the whole loop runs here, in the fork,
11
- where the real session is local. The only substitution vs. the standalone server
12
- is the transport: a `_QueueWebSocket` adapter presents the same
13
- `receive()/send_json()/send_bytes()` surface the original loop expects, backed by
14
- the in/out queues instead of a socket. Gate/PE/codec helpers are imported from
15
- serve_joyomni_streaming verbatim (see the import block).
16
- """
17
  from __future__ import annotations
18
 
19
  import argparse
20
  import asyncio
21
- import json
22
  import os
23
- import queue as _queue
24
- import threading
25
- import time
26
  import traceback
27
  from pathlib import Path
28
- from typing import Any
29
 
30
- from PIL import Image
31
 
32
- from xvideo.serving.serve_joyomni_streaming import (
33
- _check_face_gate,
34
- build_parser,
35
- _decode_image,
36
- _decode_ref_image,
37
- _enhance_prompt_sync,
38
- _count_faces_from,
39
- _detect_gate_faces,
40
- _face_present_from,
41
- _optional_positive_int,
42
- _person_present,
43
- _H264Stream,
44
- _H264Ingest,
45
- _SegmentedRecorder,
46
- _snap_to_align,
47
- )
48
- from xvideo.serving.joyomni_streaming import StreamingSettings
49
 
50
 
51
- # Server "args" the vendored gate logic reads: serve_joyomni_streaming's own
52
- # argparse defaults (= run_server.sh effective config), overridden only with
53
- # Space wiring. Recording persists to /data like the local server records to
54
- # deploy/recordings (owner's decision, 2026-08-19).
55
  def _default_args(paths: dict) -> argparse.Namespace:
56
  args = build_parser().parse_args([])
57
  args.face_detector_onnx = paths.get("face_onnx", "")
@@ -67,29 +28,29 @@ def _default_args(paths: dict) -> argparse.Namespace:
67
  return args
68
 
69
 
70
- # Transport adapter: presents a WebSocket-like surface over the in/out queues so
71
- # the ported session loop reads/writes exactly as it did against a real socket.
72
  class _QueueWebSocket:
73
  def __init__(self, in_q, out_q):
74
  self._in = in_q
75
  self._out = out_q
76
  self._closed = False
77
 
 
 
 
78
  async def receive(self) -> dict:
79
- loop = asyncio.get_event_loop()
80
- while True:
81
  try:
82
- item = await loop.run_in_executor(None, lambda: self._in.get(timeout=1.0))
83
- except _queue.Empty:
 
84
  continue
85
  kind = item.get("kind")
86
  if kind == "close":
87
  self._closed = True
88
- return {"type": "websocket.disconnect"}
89
- if kind == "text":
90
- return {"text": item["data"], "bytes": None, "type": "websocket.receive"}
91
- if kind == "bytes":
92
- return {"text": None, "bytes": item["data"], "type": "websocket.receive"}
93
 
94
  async def send_json(self, payload: dict) -> None:
95
  await self._put(("json", payload))
@@ -98,25 +59,15 @@ class _QueueWebSocket:
98
  await self._put(("bin", data))
99
 
100
  async def _put(self, item) -> None:
101
- loop = asyncio.get_event_loop()
102
- if item[0] == "json":
103
- def _put_blocking():
104
- deadline = time.monotonic() + 60.0
105
- while not self._closed and time.monotonic() < deadline:
106
- try:
107
- self._out.put(item, timeout=1.0)
108
- return
109
- except _queue.Full:
110
- continue
111
- await loop.run_in_executor(None, _put_blocking)
112
- return
113
- try:
114
- await loop.run_in_executor(None, lambda: self._out.put(item, timeout=2.0))
115
- except _queue.Full:
116
- pass # drop frames under backpressure; control messages never drop
117
 
118
 
119
- # Entry point called from the fork (app.py::_run_session_in_fork).
120
  _VAE_AOTI = {"done": False}
121
 
122
 
@@ -231,550 +182,8 @@ def run_session_blocking(runtime, paths: dict, in_q, out_q) -> None:
231
  pass
232
 
233
 
234
- # The ported session loop. Structurally mirrors websocket_endpoint() in
235
- # serve_joyomni_streaming.py, minus the SessionGate and with the transport
236
- # swapped for _QueueWebSocket. Kept compact: gate/PE helpers are imported;
237
- # the per-frame gate STATE MACHINE is preserved.
238
  async def _session_loop(runtime, args, websocket: _QueueWebSocket) -> None:
239
- session = None
240
- frames_in = 0
241
- frames_out = 0
242
- session_prompt = args.prompt
243
- raw_session_prompt = args.prompt
244
- session_settings: StreamingSettings | None = None
245
- ref_image: Image.Image | None = None
246
- output_quality = 60 if args.output_quality == "auto" else int(args.output_quality)
247
- output_codec = "mjpeg"
248
- h264_stream: _H264Stream | None = None
249
- input_codec = "mjpeg"
250
- h264_ingest: _H264Ingest | None = None
251
- frames_since_session_reset = 0
252
- next_frame_meta: dict[str, Any] | None = None
253
-
254
- gate_state: dict[str, Any] = {}
255
-
256
- inference_lock = threading.Lock()
257
-
258
- def _locked_close(_s):
259
- with inference_lock:
260
- _s.close()
261
-
262
- async def _send_json(payload: dict) -> None:
263
- await websocket.send_json(payload)
264
-
265
- def _create_session():
266
- return runtime.create_v2v_session(
267
- prompt=session_prompt, settings=session_settings, ref_image=ref_image,
268
- )
269
-
270
- # ---- output pump: drains decoded chunks and streams them out ----
271
- stop_output = threading.Event()
272
- output_task: asyncio.Task | None = None
273
-
274
- async def _send_chunk_result(result) -> int:
275
- if not result.jpegs:
276
- return 0
277
- frames = result.jpegs
278
- metas = result.source_metas or [{} for _ in frames]
279
- if result.valid_count is not None:
280
- frames = frames[: result.valid_count]
281
- metas = metas[: result.valid_count]
282
- nonlocal frames_out
283
- packets = None
284
- if h264_stream is not None and frames and not isinstance(frames[0], bytes):
285
- packets = await asyncio.to_thread(h264_stream.encode, frames)
286
- await _send_json({"type": "chunk_start", "count": len(frames),
287
- "elapsed": float(result.elapsed or 0.0), "frames_in": frames_in})
288
- for idx, item in enumerate(frames):
289
- meta = metas[min(idx, len(metas) - 1)]
290
- if rec_output is not None:
291
- rec_output.submit(item, meta.get("t_capture_ms"))
292
- wire, is_key = (packets[idx] if packets is not None else (item, False))
293
- await _send_json({"type": "output_frame", "index": idx, "count": len(frames),
294
- "source_seq": meta.get("seq"), "t_capture_ms": meta.get("t_capture_ms"),
295
- "profile": result.profile, "key": is_key})
296
- await websocket.send_bytes(wire)
297
- frames_out += 1
298
- await _send_json({"type": "chunk_done", "count": len(frames),
299
- "frames_in": frames_in, "frames_out": frames_out})
300
- return len(frames)
301
-
302
- async def _output_pump(sess) -> None:
303
- while not stop_output.is_set():
304
- try:
305
- result = await asyncio.to_thread(sess.wait_async_result, 0.05)
306
- except Exception as exc: # noqa: BLE001
307
- await _send_json({"type": "error", "message": repr(exc)})
308
- break
309
- if result is not None and result.jpegs:
310
- await _send_chunk_result(result)
311
- # Stop requested: flush any chunk still in flight in the pipeline so the
312
- # user's final frames are not dropped on stop/reset/disconnect.
313
- try:
314
- for result in await asyncio.to_thread(sess._drain_async_results):
315
- if result is not None and result.jpegs:
316
- await _send_chunk_result(result)
317
- except Exception: # noqa: BLE001
318
- pass
319
-
320
- def _start_output(sess):
321
- nonlocal output_task
322
- stop_output.clear()
323
- output_task = asyncio.create_task(_output_pump(sess))
324
-
325
- async def _stop_output():
326
- nonlocal output_task
327
- stop_output.set()
328
- if output_task is not None:
329
- try:
330
- await asyncio.wait_for(output_task, timeout=1.0)
331
- except Exception: # noqa: BLE001
332
- output_task.cancel()
333
- output_task = None
334
-
335
- async def _reset_session(reason: str):
336
- nonlocal session, frames_since_session_reset
337
- if session is None:
338
- return
339
- await _stop_output()
340
- await asyncio.to_thread(_locked_close, session)
341
- session = _create_session()
342
- frames_since_session_reset = 0
343
- await _send_json({"type": "session_reset", "reason": reason,
344
- "frames_per_next_chunk": session.frames_per_next_chunk})
345
- _start_output(session)
346
-
347
- # ---- recording: same shape as the local server (input/output mp4 segments
348
- # per session dir + prompts.json sidecar); recorders span session resets ----
349
- rec_input: _SegmentedRecorder | None = None
350
- rec_output: _SegmentedRecorder | None = None
351
- rec_base: Path | None = None
352
- rec_seq = 0
353
-
354
- def _write_prompt_sidecar() -> None:
355
- if rec_base is None:
356
- return
357
- try:
358
- doc = {"raw_prompt": raw_session_prompt}
359
- if session_prompt and session_prompt != raw_session_prompt:
360
- doc["enhanced_prompt"] = session_prompt
361
- (rec_base / "prompts.json").write_text(
362
- json.dumps(doc, ensure_ascii=False, indent=2), encoding="utf-8")
363
- except Exception: # noqa: BLE001
364
- pass
365
-
366
- def _stop_recorders() -> None:
367
- nonlocal rec_input, rec_output, rec_base
368
- for _rec in (rec_input, rec_output):
369
- if _rec is not None:
370
- try:
371
- _rec.stop()
372
- except Exception: # noqa: BLE001
373
- pass
374
- rec_input = rec_output = None
375
- rec_base = None
376
-
377
- def _start_recorders() -> None:
378
- nonlocal rec_input, rec_output, rec_seq, rec_base
379
- if args.record_dir is None:
380
- return
381
- _stop_recorders()
382
- try:
383
- rec_seq += 1
384
- base = Path(args.record_dir) / f"{int(time.time())}_{rec_seq}"
385
- base.mkdir(parents=True, exist_ok=True)
386
- common = dict(codec=str(args.record_codec), bitrate=int(args.record_bitrate),
387
- segment_seconds=int(args.record_segment_seconds))
388
- rec_input = _SegmentedRecorder(prefix=base / "input", **common)
389
- rec_output = _SegmentedRecorder(prefix=base / "output", **common)
390
- rec_input.start()
391
- rec_output.start()
392
- rec_base = base
393
- print(f"#####[REC] recording -> {base}", flush=True)
394
- _write_prompt_sidecar()
395
- if ref_image is not None:
396
- try:
397
- ref_image.save(base / "ref.png")
398
- except Exception: # noqa: BLE001
399
- pass
400
- except Exception as exc: # noqa: BLE001
401
- print(f"#####[REC] start failed: {exc!r} -> recording OFF", flush=True)
402
- rec_input = rec_output = None
403
- rec_base = None
404
-
405
- await _send_json({"type": "session_granted"})
406
-
407
- # ---- main receive loop ----
408
- while True:
409
- message = await websocket.receive()
410
- if message.get("type") == "websocket.disconnect":
411
- break
412
- text = message.get("text")
413
- if text is not None:
414
- payload = json.loads(text)
415
- msg_type = payload.get("type")
416
-
417
- if msg_type == "start":
418
- if session is not None:
419
- await _stop_output()
420
- await asyncio.to_thread(_locked_close, session)
421
- session = None
422
- raw_session_prompt = str(payload.get("prompt", args.prompt))
423
- session_prompt = raw_session_prompt
424
- try:
425
- ref_image = _decode_ref_image(payload.get("ref_image"))
426
- except Exception as exc: # noqa: BLE001
427
- await _send_json({"type": "error", "message": f"ref image decode: {exc!r}"})
428
- continue
429
- kv_reset_frames = max(0, int(payload.get("kv_reset_frames", args.kv_reset_frames)))
430
- output_quality = max(1, min(100, int(payload.get("output_quality", output_quality))))
431
- runtime.output_quality = output_quality
432
- runtime.lossless_output = False
433
- _up_allow = args.uplink_codec == "auto"
434
- _dn_allow = args.downlink_codec == "auto"
435
- output_codec = "h264" if payload.get("output_codec") == "h264" and _dn_allow else "mjpeg"
436
- h264_stream = _H264Stream(output_quality) if output_codec == "h264" else None
437
- input_codec = "h264" if payload.get("input_codec") == "h264" and _up_allow else "mjpeg"
438
- h264_ingest = _H264Ingest() if input_codec == "h264" else None
439
- try:
440
- max_temporal_ids = _optional_positive_int(
441
- payload.get("max_temporal_ids", args.max_temporal_ids), name="max_temporal_ids")
442
- except Exception as exc: # noqa: BLE001
443
- await _send_json({"type": "error", "message": f"bad max_temporal_ids: {exc!r}"})
444
- continue
445
- freeze_kv_on_static = bool(payload.get("freeze_kv_on_static", args.freeze_kv_on_static))
446
- static_diff_thresh = float(payload.get("static_diff_thresh", args.static_diff_thresh))
447
- session_max_inflight = max(0, int(
448
- payload.get("max_inflight_chunks", args.max_inflight_chunks) or 0))
449
- use_pe = bool(payload.get("use_pe", args.use_pe)) and bool(os.environ.get("OPENAI_API_KEY"))
450
- entry_gate = bool(payload.get("gate_enabled", True))
451
- face_gate_pending = entry_gate
452
- gate_on = bool(args.online_gate) and entry_gate
453
- fg_score = float(payload.get("fg_score", args.face_gate_score))
454
- fg_min_below = float(payload.get("fg_min_below_ratio", args.face_gate_min_below_ratio))
455
- fg_stable = max(1, int(payload.get("fg_stable_frames", args.face_gate_stable_frames)))
456
- fg_absent = max(1, int(args.presence_absent_frames))
457
- fg_return = max(1, int(args.presence_return_frames))
458
- _gate_fps = float(payload.get("fps") or args.fps or 24.0)
459
- _fscale = _gate_fps / 24.0
460
- fg_stable = max(1, int(round(fg_stable * _fscale)))
461
- fg_absent = max(1, int(round(fg_absent * _fscale)))
462
- fg_return = max(1, int(round(fg_return * _fscale)))
463
- count_change_frames = max(1, int(round(int(args.person_count_change_frames) * _fscale)))
464
- body_flip_frames = max(1, int(round(int(args.person_body_flip_frames) * _fscale)))
465
- person_stride = max(1, int(round(int(args.person_check_stride) * _fscale)))
466
- gate_move_eps = float(args.face_gate_move_eps) / _fscale
467
- _reset_gate_state(gate_state)
468
- pe_defer = False
469
- if use_pe:
470
- cached = str(payload.get("enhanced_prompt") or "").strip()
471
- if cached:
472
- session_prompt = cached
473
- else:
474
- pe_defer = True
475
- _align = runtime.pipeline.vae.stem.stride * 8
476
- session_settings = StreamingSettings(
477
- height=_snap_to_align(int(payload.get("height", args.height)), _align),
478
- width=_snap_to_align(int(payload.get("width", args.width)), _align),
479
- num_inference_steps=int(payload.get("num_inference_steps", args.num_inference_steps)),
480
- seed=int(payload.get("seed", args.seed)),
481
- max_temporal_ids=max_temporal_ids,
482
- freeze_kv_on_static=freeze_kv_on_static,
483
- static_diff_thresh=static_diff_thresh,
484
- profile_timings=bool(payload.get("profile_timings", args.profile_timings)),
485
- output_codec=output_codec,
486
- )
487
- session = _create_session()
488
- frames_since_session_reset = 0
489
- if not pe_defer:
490
- # Bake the graph before "started" so early frames don't age in queue.
491
- def _prebake(_p=session_prompt, _s=session_settings, _r=ref_image):
492
- with inference_lock:
493
- runtime.prebake_graph(_p, settings=_s, ref_image=_r)
494
- try:
495
- await asyncio.to_thread(_prebake)
496
- except Exception: # noqa: BLE001
497
- pass
498
- await _send_json({"type": "started",
499
- "frames_per_next_chunk": session.frames_per_next_chunk,
500
- "height": session_settings.height, "width": session_settings.width,
501
- "output_codec": output_codec, "input_codec": input_codec,
502
- "ref_image": ref_image is not None, "use_pe": use_pe})
503
- _start_output(session)
504
- await asyncio.to_thread(_start_recorders)
505
-
506
- elif msg_type == "stop":
507
- await _stop_output()
508
- await asyncio.to_thread(_stop_recorders)
509
- break
510
-
511
- elif msg_type == "finalize_recording":
512
- if args.record_dir is None:
513
- await _send_json({"type": "recording_finalized", "ok": False,
514
- "message": "Recording is not enabled (--record-dir is unset)."})
515
- continue
516
- if session is not None:
517
- await asyncio.to_thread(session.flush_pending)
518
- _deadline = time.monotonic() + 20.0
519
- while frames_out < frames_in and time.monotonic() < _deadline:
520
- await asyncio.sleep(0.05)
521
- await _stop_output()
522
- finalized = rec_base
523
- await asyncio.to_thread(_stop_recorders)
524
- await _send_json({"type": "recording_finalized", "ok": finalized is not None,
525
- "rec": finalized.name if finalized is not None else None,
526
- "message": None if finalized is not None
527
- else "No downloadable result yet. Send an edit first."})
528
- elif msg_type == "frame_meta":
529
- next_frame_meta = {"seq": int(payload.get("seq", frames_in + 1)),
530
- "t_capture_ms": float(payload.get("t_capture_ms", time.time() * 1000.0))}
531
- elif msg_type == "ping":
532
- await _send_json({"type": "pong", "t": payload.get("t")})
533
- elif msg_type == "set_output_quality":
534
- try:
535
- output_quality = max(1, min(100, int(payload.get("value", output_quality))))
536
- runtime.output_quality = output_quality
537
- if h264_stream is not None:
538
- h264_stream.set_quality(output_quality)
539
- except (TypeError, ValueError):
540
- pass
541
- continue
542
-
543
- # ---- binary frame ----
544
- frame_bytes = message.get("bytes")
545
- if frame_bytes is None:
546
- continue
547
- if session is None:
548
- await _send_json({"type": "error", "message": "send start before frames"})
549
- continue
550
-
551
- frames_in += 1
552
- frame_meta = next_frame_meta or {"seq": frames_in, "t_capture_ms": time.time() * 1000.0}
553
- next_frame_meta = None
554
-
555
- uplink_frame: Image.Image | None = None
556
- if h264_ingest is not None:
557
- uplink_frame = await asyncio.to_thread(h264_ingest.decode_one, frame_bytes)
558
- if uplink_frame is None:
559
- continue
560
-
561
- if rec_input is not None:
562
- rec_input.submit(uplink_frame if uplink_frame is not None else frame_bytes,
563
- frame_meta.get("t_capture_ms"))
564
-
565
- if kv_reset_frames > 0 and frames_since_session_reset >= kv_reset_frames:
566
- face_gate_pending = entry_gate
567
- _reset_gate_state(gate_state)
568
- await _reset_session("kv_reset_frames")
569
- continue
570
-
571
- def _run_frame():
572
- frame = uplink_frame if uplink_frame is not None else _decode_image(frame_bytes)
573
- sentinel = _apply_gate(
574
- frame, gate_state, args,
575
- face_gate_pending=face_gate_pending, gate_on=gate_on, pe_defer=pe_defer,
576
- fg_score=fg_score, fg_min_below=fg_min_below, fg_stable=fg_stable,
577
- fg_absent=fg_absent, fg_return=fg_return, person_stride=person_stride,
578
- body_flip_frames=body_flip_frames, count_change_frames=count_change_frames,
579
- gate_move_eps=gate_move_eps, session=session,
580
- )
581
- if sentinel is not None:
582
- return sentinel
583
- if session_max_inflight and session.inflight_chunks() >= session_max_inflight:
584
- return session._drain_async_results()
585
- with inference_lock:
586
- return session.push_frame(frame, frame_meta=frame_meta)
587
-
588
- chunk_results = await asyncio.to_thread(_run_frame)
589
-
590
- # ---- sentinels from the gate ----
591
- if isinstance(chunk_results, tuple) and chunk_results and isinstance(chunk_results[0], str):
592
- s = chunk_results[0]
593
- if s == "__no_person__":
594
- await _send_json({"type": "no_person",
595
- "reason": chunk_results[1] if len(chunk_results) > 1 else "no_person",
596
- "frames_in": frames_in})
597
- continue
598
- if s == "__person_returned__":
599
- face_gate_pending = True
600
- _reset_gate_state(gate_state)
601
- await _reset_session("person_returned")
602
- continue
603
- if chunk_results == "__gate_pe__":
604
- anchor = gate_state.pop("pe_anchor", None)
605
- await _send_json({"type": "pe_running", "frames_in": frames_in})
606
- try:
607
- pe_report = await asyncio.wait_for(
608
- asyncio.to_thread(
609
- _enhance_prompt_sync, raw_prompt=raw_session_prompt,
610
- ref_image=ref_image, pe_frame=anchor, pe_model=args.pe_model),
611
- timeout=max(1.0, float(args.pe_timeout_s)),
612
- )
613
- enhanced = str(pe_report.get("enhanced_prompt") or raw_session_prompt)
614
- session_prompt = enhanced
615
-
616
- def _swap(_s=session, _t=enhanced, _a=anchor):
617
- with inference_lock:
618
- _s.prompt = _t
619
- _s._initialize(_a)
620
- await asyncio.to_thread(_swap)
621
- print(f"#####[PE] enhanced ok in {float(pe_report.get('elapsed_s', 0.0)):.1f}s "
622
- f"(model={pe_report.get('model')}) -> session initialized", flush=True)
623
- _write_prompt_sidecar()
624
- await _send_json({"type": "prompt_enhanced", **pe_report})
625
- except Exception as exc: # noqa: BLE001
626
- _why = "timeout" if isinstance(exc, asyncio.TimeoutError) else repr(exc)
627
- print(f"#####[PE] deferred enhance failed: {_why} -> raw prompt", flush=True)
628
-
629
- def _init_raw(_s=session, _a=anchor):
630
- with inference_lock:
631
- _s._initialize(_a)
632
- await asyncio.to_thread(_init_raw)
633
- await _send_json({"type": "prompt_enhanced", "enabled": True,
634
- "raw_prompt": raw_session_prompt, "enhanced_prompt": raw_session_prompt,
635
- "fallback": True, "error": True, "elapsed_s": 0.0})
636
- pe_defer = False
637
- continue
638
- if isinstance(chunk_results, str):
639
- await _send_json({"type": "waiting_face", "reason": chunk_results, "frames_in": frames_in})
640
- continue
641
-
642
- if face_gate_pending:
643
- face_gate_pending = False
644
- frames_since_session_reset += 1
645
-
646
- if gate_state.get("recount"):
647
- gate_state["recount"] = False
648
- await _reset_session("person_count_changed")
649
- continue
650
-
651
- if not chunk_results:
652
- await _send_json({"type": "accepted", "frames_in": frames_in, "frames_out": frames_out,
653
- "next_chunk_needs": session.frames_per_next_chunk - len(session.pending_frames)})
654
- continue
655
- # output frames are streamed by the pump; push_frame's direct returns
656
- # (init chunk) are also flushed here for parity.
657
- for result in chunk_results:
658
- await _send_chunk_result(result)
659
-
660
- await _stop_output()
661
- await asyncio.to_thread(_stop_recorders)
662
- if session is not None:
663
- try:
664
- await asyncio.to_thread(_locked_close, session)
665
- except Exception: # noqa: BLE001
666
- pass
667
-
668
-
669
- # Gate state machine — extracted from serve_joyomni_streaming's _run_frame.
670
- # Returns None to proceed to push_frame, or a sentinel (str / tuple) the caller
671
- # translates into a WS message. Kept faithful to the original transitions.
672
- def _reset_gate_state(gs: dict) -> None:
673
- gs.clear()
674
- gs.update({"count": 0, "cx": None, "cy": None, "csz": None, "absent": 0,
675
- "absent_hold": False, "present": 0,
676
- "person_check_i": 0, "person_last": True, "face_last": True,
677
- "hold_reason": "no_person", "body_miss": 0, "subject_count": None,
678
- "cand": None, "cand_n": 0, "recount": False, "settle_ax": None,
679
- "settle_ay": None, "settle_asz": None, "pe_anchor": None})
680
-
681
-
682
- def _apply_gate(frame, gs, args, *, face_gate_pending, gate_on, pe_defer, fg_score,
683
- fg_min_below, fg_stable, fg_absent, fg_return, person_stride,
684
- body_flip_frames, count_change_frames, gate_move_eps, session):
685
- if face_gate_pending:
686
- reason, center, nf = _check_face_gate(
687
- frame, onnx_path=args.face_detector_onnx,
688
- score_thresh=fg_score, min_below_ratio=fg_min_below)
689
- if reason is not None:
690
- gs["count"] = 0; gs["settle_ax"] = None; gs["settle_ay"] = None
691
- return reason
692
- if center is not None:
693
- cx, cy, csz = center
694
- if nf < 2 and abs(cx - 0.5) > float(args.face_gate_center_margin):
695
- gs["count"] = 0; gs["settle_ax"] = None; gs["settle_ay"] = None
696
- gs["settle_asz"] = None; gs["cx"] = cx; gs["cy"] = cy; gs["csz"] = csz
697
- return "off_center"
698
- pcx, pcy, pcsz = gs.get("cx"), gs.get("cy"), gs.get("csz")
699
- eps = gate_move_eps; cap = float(args.face_gate_settle_drift)
700
- szeps = eps * 0.5
701
- still = (pcx is not None and abs(cx - pcx) <= eps and abs(cy - pcy) <= eps
702
- and pcsz is not None and abs(csz - pcsz) <= szeps)
703
- ax, ay, asz = gs.get("settle_ax"), gs.get("settle_ay"), gs.get("settle_asz")
704
- if (still and ax is not None and abs(cx - ax) <= cap and abs(cy - ay) <= cap
705
- and asz is not None and abs(csz - asz) <= cap):
706
- gs["count"] += 1
707
- else:
708
- gs["count"] = 1; gs["settle_ax"] = cx; gs["settle_ay"] = cy; gs["settle_asz"] = csz
709
- gs["cx"] = cx; gs["cy"] = cy; gs["csz"] = csz
710
- if gs["count"] < fg_stable:
711
- return "settling"
712
- if pe_defer:
713
- gs["pe_anchor"] = frame
714
- return "__gate_pe__"
715
-
716
- if gate_on and not face_gate_pending:
717
- gfaces, gfw, gfh = _detect_gate_faces(
718
- frame, onnx_path=args.face_detector_onnx, score_thresh=fg_score)
719
- tick = gs.get("person_check_i", 0)
720
- if tick == 0 or gs.get("absent_hold"):
721
- gs["person_last"] = _person_present(
722
- frame, onnx_path=args.person_detector_onnx, conf=float(args.person_gate_conf))
723
- if gs["person_last"]:
724
- gs["face_last"] = _face_present_from(
725
- gfaces, gfw, gfh,
726
- min_ratio=float(args.face_present_min_ratio),
727
- edge_margin=float(args.face_present_edge_margin))
728
- else:
729
- gs["face_last"] = True
730
- gs["person_check_i"] = (tick + 1) % person_stride
731
- body_here = bool(gs["person_last"]); face_here = bool(gs.get("face_last", True))
732
- present = body_here and face_here
733
- reason_now = "no_person" if not body_here else ("no_face" if not face_here else "")
734
- body_flip = body_flip_frames
735
- if body_here:
736
- gs["body_miss"] = 0
737
- else:
738
- gs["body_miss"] = gs.get("body_miss", 0) + 1
739
- if present:
740
- gs["absent"] = 0
741
- if gs.get("absent_hold"):
742
- gs["present"] = gs.get("present", 0) + 1
743
- if gs["present"] >= fg_return:
744
- gs["absent_hold"] = False; gs["present"] = 0
745
- return ("__person_returned__",)
746
- else:
747
- gs["present"] = 0; gs["absent"] += 1
748
- if reason_now == "no_person" and gs.get("body_miss", 0) < body_flip:
749
- reason_now = "no_face"
750
- gs["hold_reason"] = reason_now or gs.get("hold_reason") or "no_person"
751
- if not gs.get("absent_hold") and gs["absent"] >= fg_absent:
752
- gs["absent_hold"] = True
753
- try:
754
- if session is not None:
755
- session.pending_frames.clear(); session.pending_metas.clear()
756
- except Exception: # noqa: BLE001
757
- pass
758
- if gs.get("absent_hold"):
759
- return ("__no_person__", gs.get("hold_reason", "no_person"))
760
-
761
- n = _count_faces_from(
762
- gfaces, gfw, gfh, count_min_ratio=float(args.count_face_min_ratio))
763
- if gs["subject_count"] is None:
764
- gs["subject_count"] = n; gs["cand"] = None; gs["cand_n"] = 0
765
- elif n != gs["subject_count"]:
766
- if n == gs["cand"]:
767
- gs["cand_n"] += 1
768
- else:
769
- gs["cand"] = n; gs["cand_n"] = 1
770
- if gs["cand_n"] >= count_change_frames:
771
- if n > gs["subject_count"]:
772
- gs["recount"] = True
773
- gs["subject_count"] = n; gs["cand"] = None; gs["cand_n"] = 0
774
- else:
775
- gs["cand"] = None; gs["cand_n"] = 0
776
-
777
- if pe_defer and not face_gate_pending:
778
- gs["pe_anchor"] = frame
779
- return "__gate_pe__"
780
- return None
 
1
+ """Run the shared streaming server inside a ZeroGPU fork over process queues."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2
  from __future__ import annotations
3
 
4
  import argparse
5
  import asyncio
 
6
  import os
7
+ import queue
 
 
8
  import traceback
9
  from pathlib import Path
 
10
 
11
+ from fastapi import WebSocketDisconnect
12
 
13
+ from xvideo.serving.serve_joyomni_streaming import build_parser, create_app
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
14
 
15
 
 
 
 
 
16
  def _default_args(paths: dict) -> argparse.Namespace:
17
  args = build_parser().parse_args([])
18
  args.face_detector_onnx = paths.get("face_onnx", "")
 
28
  return args
29
 
30
 
 
 
31
  class _QueueWebSocket:
32
  def __init__(self, in_q, out_q):
33
  self._in = in_q
34
  self._out = out_q
35
  self._closed = False
36
 
37
+ async def accept(self) -> None:
38
+ pass
39
+
40
  async def receive(self) -> dict:
41
+ while not self._closed:
 
42
  try:
43
+ item = self._in.get_nowait()
44
+ except queue.Empty:
45
+ await asyncio.sleep(0.005)
46
  continue
47
  kind = item.get("kind")
48
  if kind == "close":
49
  self._closed = True
50
+ break
51
+ if kind in {"text", "bytes"}:
52
+ return {"type": "websocket.receive", kind: item["data"]}
53
+ return {"type": "websocket.disconnect"}
 
54
 
55
  async def send_json(self, payload: dict) -> None:
56
  await self._put(("json", payload))
 
59
  await self._put(("bin", data))
60
 
61
  async def _put(self, item) -> None:
62
+ while not self._closed:
63
+ try:
64
+ self._out.put_nowait(item)
65
+ return
66
+ except queue.Full:
67
+ await asyncio.sleep(0.005)
68
+ raise WebSocketDisconnect()
 
 
 
 
 
 
 
 
 
69
 
70
 
 
71
  _VAE_AOTI = {"done": False}
72
 
73
 
 
182
  pass
183
 
184
 
 
 
 
 
185
  async def _session_loop(runtime, args, websocket: _QueueWebSocket) -> None:
186
+ app = create_app(args)
187
+ app.state.runtime = runtime
188
+ endpoint = next(route.endpoint for route in app.routes if route.path == "/ws")
189
+ await endpoint(websocket)