Compare commits
56
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c520878fd2 | ||
|
|
99f8f3e7a7 | ||
|
|
14e4b7189b | ||
|
|
8530fdc550 | ||
|
|
315b25d26d | ||
|
|
b8920ba50a | ||
|
|
3467f8befb | ||
|
|
7320d9d350 | ||
|
|
af78c759f1 | ||
|
|
f82c190cbb | ||
|
|
e722023bd6 | ||
|
|
bc0191d775 | ||
|
|
2d6e7d9b6b | ||
|
|
5e61accc37 | ||
|
|
f60514183c | ||
|
|
156611bab3 | ||
|
|
6d6195a9f1 | ||
|
|
8348e83e80 | ||
|
|
2452022f15 | ||
|
|
4b3f99b224 | ||
|
|
d967b99928 | ||
|
|
8f1eb2992c | ||
|
|
683689de4f | ||
|
|
d5287f8352 | ||
|
|
cdfbd64b04 | ||
|
|
299d0f6838 | ||
|
|
c18436125a | ||
|
|
5d515ee617 | ||
|
|
37b9dc1836 | ||
|
|
a0291a57c1 | ||
|
|
6102fac00d | ||
|
|
1ad33c1832 | ||
|
|
da87973412 | ||
|
|
a3fd1d0bab | ||
|
|
3e87d6f4ae | ||
|
|
4bb5627fa8 | ||
|
|
08ab8b5c00 | ||
|
|
fb5efebdab | ||
|
|
6e4813f8b4 | ||
|
|
13ef81e6da | ||
|
|
13b3b36b8b | ||
|
|
72795d0af8 | ||
|
|
af63a92f6a | ||
|
|
a4e7322ba6 | ||
|
|
d6e13f9cc8 | ||
|
|
6c4c18690e | ||
|
|
a8688dddd9 | ||
|
|
cacc8cfcb3 | ||
|
|
ef62ac48d1 | ||
|
|
00cb7d0ba2 | ||
|
|
69a6215e7f | ||
|
|
abe98e87f5 | ||
|
|
274b922d39 | ||
|
|
38966056d0 | ||
|
|
d8ca702dfd | ||
|
|
33d81730e9 |
@@ -68,7 +68,7 @@ Copy `templates/component_parity_test.py` and fill every `TODO` marker. The
|
||||
template is distilled from:
|
||||
|
||||
- `tests/local_tests/transformers/test_ltx2.py`
|
||||
- `tests/local_tests/gen3c/test_gen3c.py`
|
||||
- `tests/local_tests/transformers/test_gamecraft_parity.py`
|
||||
- `tests/local_tests/encoders/test_ltx2_gemma_parity.py`
|
||||
- `tests/local_tests/vaes/test_oobleck_vae_parity.py`
|
||||
- `tests/local_tests/sd35/test_sd35_component_parity.py`
|
||||
|
||||
@@ -87,7 +87,7 @@ def _load_official_model(device: torch.device, dtype: torch.dtype) -> torch.nn.M
|
||||
# TODO: import official class/factory and load real weights strictly.
|
||||
# Examples in-tree:
|
||||
# - LTX2: SingleGPUModelBuilder(...).build(device=device, dtype=dtype)
|
||||
# - GEN3C: torch.load(...)["state_dict"] -> official_model.load_state_dict(...)
|
||||
# - GameCraft: torch.load(...)["module"] -> official_model.load_state_dict(...)
|
||||
# - Oobleck: create_model_from_config(config) + ckpt state_dict
|
||||
OfficialClass = _import_or_skip(OFFICIAL_MODULE, OFFICIAL_CLASS)
|
||||
model = OfficialClass() # TODO: pass official config kwargs.
|
||||
|
||||
@@ -52,8 +52,8 @@ scaling constants, dtype casts, state-dict names, and every output head.
|
||||
- Loader path: `TransformerLoader` reads `transformer/config.json`, calls
|
||||
`dit_config.update_model_arch(config)`, resolves `_class_name` through
|
||||
`ModelRegistry`, and constructs the class with `config` and `hf_config`.
|
||||
- Reference examples: `stable_audio.py`, `wanvideo.py`, `sd3.py`, and
|
||||
`ltx2.py`.
|
||||
- Reference examples: `stable_audio.py`, `wanvideo.py`, `sd3.py`, `longcat.py`,
|
||||
and `ltx2.py`.
|
||||
- Layer guidance: `fastvideo/layers/AGENTS.md`.
|
||||
|
||||
## Implementation Rules
|
||||
|
||||
@@ -51,7 +51,7 @@ posterior behavior, encode/decode output objects, tiling flags, and cropping.
|
||||
- Loader path: VAE loaders resolve `_class_name` through `ModelRegistry` and
|
||||
load converted component weights from the VAE subdir.
|
||||
- Reference examples: `oobleck.py`, `autoencoder_kl.py`, `wanvae.py`,
|
||||
`ltx2vae.py`, and `hunyuanvae.py`.
|
||||
`ltx2vae.py`, and `gamecraftvae.py`.
|
||||
- Layer guidance: `fastvideo/layers/AGENTS.md`.
|
||||
|
||||
## Implementation Rules
|
||||
|
||||
@@ -43,12 +43,11 @@ from `../add-model/contracts/conversion_request.md`.
|
||||
- `scripts/checkpoint_conversion/stable_audio_to_diffusers.py`: monolithic
|
||||
`model.safetensors` split into transformer/VAE/conditioner, plus copied
|
||||
passthrough subfolders. Use this shape for single-checkpoint official repos.
|
||||
- `scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py`: separate
|
||||
official sources for transformer, VAE, encoders, and vocoder, assembled under
|
||||
a root `model_index.json`.
|
||||
- `scripts/checkpoint_conversion/convert_flux2_klein.py`: fused QKV split,
|
||||
renamed native transformer weights, and copied passthrough text encoder,
|
||||
tokenizer, and scheduler components.
|
||||
- `scripts/checkpoint_conversion/convert_gamecraft_full.py`: separate official
|
||||
sources for transformer, VAE, text encoders, tokenizers, scheduler, and root
|
||||
`model_index.json`.
|
||||
- `scripts/checkpoint_conversion/longcat_to_fastvideo.py`: fused QKV/KV split,
|
||||
renamed native transformer weights, and copied existing Diffusers components.
|
||||
- `scripts/checkpoint_conversion/pt_to_safetensors.py`: simple `.pt` extraction
|
||||
helper for nested checkpoint dictionaries.
|
||||
|
||||
|
||||
@@ -247,7 +247,7 @@ setup gap, not a pass.
|
||||
- `fastvideo/configs/pipelines/stable_audio.py` and
|
||||
`fastvideo/pipelines/basic/stable_audio/presets.py` for config/preset shape.
|
||||
- `fastvideo/registry.py` for `register_configs(...)` and preset registration.
|
||||
- `tests/local_tests/pipelines/test_lingbot_video_pipeline_parity.py` for latent
|
||||
- `tests/local_tests/pipelines/test_gamecraft_pipeline_parity.py` for latent
|
||||
parity structure.
|
||||
- `tests/local_tests/pipelines/test_stable_audio_pipeline_parity.py` for audio
|
||||
parity structure.
|
||||
|
||||
@@ -413,7 +413,7 @@ matching `*secret*`.
|
||||
- `fastvideo/pipelines/basic/wan/` for standard T2V/I2V/DMD/Causal variants.
|
||||
- `fastvideo/pipelines/basic/ltx2/` for non-standard stages and audio/video
|
||||
patterns.
|
||||
- `tests/local_tests/pipelines/test_lingbot_video_pipeline_parity.py` for pipeline
|
||||
- `tests/local_tests/pipelines/test_gamecraft_pipeline_parity.py` for pipeline
|
||||
parity shape.
|
||||
- `tests/local_tests/transformers/test_ltx2.py`,
|
||||
`tests/local_tests/vaes/test_ltx2_vae.py`, and
|
||||
|
||||
@@ -105,8 +105,8 @@ Detect artefact type by inspecting the file's imports / helper call:
|
||||
- **pixel** (`.mp4`) — file imports
|
||||
`run_text_to_video_similarity_test` / `run_image_to_video_similarity_test`
|
||||
from `fastvideo.tests.ssim.inference_similarity_utils`, OR uses the
|
||||
legacy custom-inline helper pattern (see `test_gen3c`). Default to pixel
|
||||
when both heuristics fail.
|
||||
legacy custom-inline helper pattern (see `test_gamecraft`,
|
||||
`test_longcat`, etc.). Default to pixel when both heuristics fail.
|
||||
|
||||
Record `ARTEFACT_TYPE ∈ {pixel, latent}` for use in step 4. Steps 2, 3, 5,
|
||||
and 6 are artefact-type-agnostic — `_iter_reference_files`,
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
{
|
||||
"benchmark_id": "wan-t2v-1.3b-1gpu-gb10",
|
||||
"config_schema_version": 2,
|
||||
"workload_id": "wan-t2v",
|
||||
"variant_id": "1.3b-sp1",
|
||||
"benchmark_version": 3,
|
||||
"description": "Wan2.1 T2V 1.3B single-GPU inference performance on NVIDIA DGX Spark (GB10). Single-GPU variant of wan-t2v-1.3b (same workload_id for dashboard comparability). Gated to the GB10 via run_config.gpu_types so it does not run on the shared H100/L40S lanes.",
|
||||
"model": {
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"model_short_name": "Wan2.1-T2V-1.3B"
|
||||
},
|
||||
"init_kwargs": {
|
||||
"num_gpus": 1,
|
||||
"flow_shift": 7.0,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"vae_sp": false,
|
||||
"vae_tiling": true,
|
||||
"text_encoder_precisions": ["fp32"]
|
||||
},
|
||||
"generation_kwargs": {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 3,
|
||||
"embedded_cfg_scale": 6,
|
||||
"seed": 1024,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
},
|
||||
"test_prompts": [
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
],
|
||||
"run_config": {
|
||||
"num_warmup_runs": 2,
|
||||
"num_measurement_runs": 5,
|
||||
"required_gpus": 1,
|
||||
"gpu_types": ["GB10"]
|
||||
},
|
||||
"thresholds": {
|
||||
"GB10": {
|
||||
"max_generation_time_s": 55.0,
|
||||
"max_peak_memory_mb": 12000.0
|
||||
},
|
||||
"default": {
|
||||
"max_generation_time_s": 120.0,
|
||||
"max_peak_memory_mb": 40000.0
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -2,4 +2,4 @@
|
||||
# Canonical Slurm CI selection for the LoRA-extraction lane.
|
||||
set -euo pipefail
|
||||
|
||||
exec pytest ./fastvideo/tests/lora_extraction/test_lora_extraction.py -vs
|
||||
exec pytest ./fastvideo/tests/lora_extraction/ -vs
|
||||
|
||||
@@ -32,7 +32,7 @@ cleanup() {
|
||||
}
|
||||
trap cleanup EXIT INT TERM
|
||||
|
||||
pytest ./fastvideo/tests/performance -vs
|
||||
pytest ./fastvideo/tests/performance/test_inference_performance.py -vs
|
||||
pytest_rc=$?
|
||||
compare_rc=0
|
||||
if [ "$pytest_rc" -eq 0 ] || [ "$PERF_UPLOAD_POLICY" = always ]; then
|
||||
@@ -41,6 +41,18 @@ if [ "$pytest_rc" -eq 0 ] || [ "$PERF_UPLOAD_POLICY" = always ]; then
|
||||
fi
|
||||
python ./fastvideo/tests/performance/dashboard.py || true
|
||||
cp -f fastvideo/tests/performance/results/*.json "$PERF_REPORTS_DIR/" 2>/dev/null || true
|
||||
# The trusted host relays only .md/.html/.json/.csv from PERF_REPORTS_DIR, so
|
||||
# mirror each captured worker log with an allowlisted extension.
|
||||
for worker_log in fastvideo/tests/performance/results/worker_logs/*.log; do
|
||||
[ -f "$worker_log" ] || continue
|
||||
base=$(basename "${worker_log%.log}")
|
||||
# WorkerLogCapture keeps a .log.1 backup after rollover, and read_log_tail
|
||||
# includes it; mirror that retained history too so the artifact is complete.
|
||||
if [ -f "$worker_log.1" ]; then
|
||||
cp -f "$worker_log.1" "$PERF_REPORTS_DIR/${base}.1.md" 2>/dev/null || true
|
||||
fi
|
||||
cp -f "$worker_log" "$PERF_REPORTS_DIR/${base}.md" 2>/dev/null || true
|
||||
done
|
||||
|
||||
echo "--- GPU telemetry (clocks.sm vs clocks.max.sm reveals capped hosts) ---"
|
||||
cat "$PERF_REPORTS_DIR/gpu_telemetry.csv" || true
|
||||
|
||||
@@ -1,7 +1,16 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
# Collect the whole attention directory so new files cannot land uncovered.
|
||||
# Its FA2/FA3 regression files skip when FA4 is selected (the Modal image
|
||||
# enables FA4 by default), so pin FA4 off for the directory to be real
|
||||
# coverage on every runner rather than a nominal collection.
|
||||
export FASTVIDEO_FA4=0
|
||||
|
||||
# The livestream app's tests are CPU-only; its single gpu-marked module is
|
||||
# deselected, and DreamVerse's GPU tests have their own lane.
|
||||
exec pytest \
|
||||
./apps/infinite_livestream/infinite_livestream/tests \
|
||||
./fastvideo/tests/api/ \
|
||||
./fastvideo/tests/contract/ \
|
||||
./fastvideo/tests/dataset/ \
|
||||
@@ -10,21 +19,23 @@ exec pytest \
|
||||
./fastvideo/tests/loader/ \
|
||||
./fastvideo/tests/pipelines/ \
|
||||
./fastvideo/tests/platforms/ \
|
||||
./fastvideo/tests/schedulers/ \
|
||||
./fastvideo/tests/train/ \
|
||||
./fastvideo/tests/stages/ \
|
||||
./fastvideo/tests/ops/ \
|
||||
./fastvideo/tests/worker/ \
|
||||
./fastvideo/tests/training/test_runner.py \
|
||||
./fastvideo/tests/training/test_trackers.py \
|
||||
./fastvideo/tests/inference/test_basic_fasth3_omniref_pdd.py \
|
||||
./fastvideo/tests/attention/test_sdpa_metadata_mask_contract.py \
|
||||
./fastvideo/tests/attention/test_vsa_h3_tile_grad_safety.py \
|
||||
./fastvideo/tests/attention/test_vsa_h3_metadata.py \
|
||||
./fastvideo/tests/attention/test_vsa_h3_ref2va_regions.py \
|
||||
./fastvideo/tests/inference/test_inference_regional_compile.py \
|
||||
./fastvideo/tests/attention/ \
|
||||
./fastvideo/tests/layers/test_pdd_linear.py \
|
||||
./fastvideo/tests/layers/test_triton_fused_norm.py \
|
||||
./fastvideo/tests/modal/test_kernel_build_cache.py \
|
||||
./fastvideo/tests/modal/test_pr_test.py \
|
||||
./fastvideo/tests/modal/test_ssim_test.py \
|
||||
--ignore=./fastvideo/tests/entrypoints/test_openai_api_integration.py \
|
||||
--ignore=./fastvideo/tests/train/models \
|
||||
--ignore=./fastvideo/tests/train/methods \
|
||||
-m "not gpu" \
|
||||
-vs
|
||||
|
||||
@@ -12,7 +12,6 @@ from __future__ import annotations
|
||||
import argparse
|
||||
import fnmatch
|
||||
import re
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import TextIO
|
||||
@@ -131,6 +130,11 @@ class FamilyCoverage:
|
||||
|
||||
|
||||
FAMILY_COVERAGE = (
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])dreamx(_world)?([/_.-]|$)"),
|
||||
("test_dreamx.py", ),
|
||||
("test_dreamx_world_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])flux[_-]?2([/_.-]|$)"),
|
||||
("test_flux2_klein.py", ),
|
||||
@@ -141,11 +145,26 @@ FAMILY_COVERAGE = (
|
||||
("test_flux.py", ),
|
||||
("test_flux_t2i_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])(hunyuan)?gamecraft([/_.-]|$)"),
|
||||
("test_gamecraft.py", ),
|
||||
("test_gamecraft_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])gen3c([/_.-]|$)"),
|
||||
("test_gen3c.py", ),
|
||||
("test_gen3c_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])glm[_-]?image([/_.-]|$)"),
|
||||
("test_glm_image.py", ),
|
||||
("test_glm_image_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])hunyuan(video)?15([a-z0-9_-]*)([/_.-]|$)"),
|
||||
(),
|
||||
("test_hunyuan15_i2v_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])kandinsky[_-]?5([/_.-]|$)"),
|
||||
("test_kandinsky5.py", ),
|
||||
@@ -156,6 +175,11 @@ FAMILY_COVERAGE = (
|
||||
("test_lingbot.py", ),
|
||||
("test_lingbot_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])longcat([/_.-]|$)"),
|
||||
("test_longcat.py", ),
|
||||
("test_longcat_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])ltx[_-]?2([/_.-]|$)"),
|
||||
("test_ltx2.py", ),
|
||||
@@ -186,6 +210,11 @@ FAMILY_COVERAGE = (
|
||||
("test_stable_audio.py", ),
|
||||
("test_stable_audio_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])turbo(diffusion)?([/_.-]|$)"),
|
||||
(),
|
||||
("test_turbodiffusion_similarity.py", ),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])wan(video|vae)?([/_.-]|$)"),
|
||||
("test_wan_t2v.py", "test_wan_vae.py", "test_wan_causal.py", "test_wan_denoising.py"),
|
||||
@@ -193,8 +222,14 @@ FAMILY_COVERAGE = (
|
||||
"test_causal_similarity.py",
|
||||
"test_wan_i2v_similarity.py",
|
||||
"test_wan_t2v_similarity.py",
|
||||
"test_wan_ti2v_similarity.py",
|
||||
),
|
||||
),
|
||||
FamilyCoverage(
|
||||
re.compile(r"(^|[/_.-])z[_-]?image([/_.-]|$)"),
|
||||
("test_zimage.py", ),
|
||||
("test_zimage_similarity.py", ),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -290,22 +325,16 @@ def _select_output_coverage(plan: MergePlan, path: str) -> None:
|
||||
plan.add_ssim(SSIM_SMOKE_TESTS, reason=f"shared output SSIM smoke coverage: {path}")
|
||||
|
||||
|
||||
def _normalize_path(raw_path: str) -> str:
|
||||
path = raw_path.strip()
|
||||
while path.startswith("./"):
|
||||
path = path[2:]
|
||||
return path
|
||||
|
||||
|
||||
def classify_paths(paths: list[str], removed_paths: Iterable[str] = ()) -> MergePlan:
|
||||
"""Plan merge lanes for ``paths``.
|
||||
|
||||
``removed_paths`` lists changed paths that no longer exist at the PR head
|
||||
(deleted files and rename sources); removed golden/SSIM tests are not run.
|
||||
"""
|
||||
def classify_paths(paths: list[str]) -> MergePlan:
|
||||
plan = MergePlan()
|
||||
normalized_paths = sorted({path for path in map(_normalize_path, paths) if path})
|
||||
removed = {path for path in map(_normalize_path, removed_paths) if path}
|
||||
normalized_paths: list[str] = []
|
||||
for raw_path in paths:
|
||||
path = raw_path.strip()
|
||||
while path.startswith("./"):
|
||||
path = path[2:]
|
||||
if path:
|
||||
normalized_paths.append(path)
|
||||
normalized_paths = sorted(set(normalized_paths))
|
||||
if not normalized_paths:
|
||||
plan.require_all("changed-file list was empty; failing closed")
|
||||
return plan
|
||||
@@ -345,10 +374,7 @@ def classify_paths(paths: list[str], removed_paths: Iterable[str] = ()) -> Merge
|
||||
if path.startswith("fastvideo/tests/golden_gate/"):
|
||||
name = Path(path).name
|
||||
if name.startswith("test_") and name.endswith(".py"):
|
||||
if path in removed:
|
||||
plan.reasons.append(f"removed golden test has nothing to run: {path}")
|
||||
else:
|
||||
plan.add_golden((name, ), reason=f"changed golden test: {path}")
|
||||
plan.add_golden((name, ), reason=f"changed golden test: {path}")
|
||||
elif name in {"AGENTS.md", "README.md"}:
|
||||
plan.reasons.append(f"golden documentation only: {path}")
|
||||
else:
|
||||
@@ -359,10 +385,7 @@ def classify_paths(paths: list[str], removed_paths: Iterable[str] = ()) -> Merge
|
||||
if path.startswith("fastvideo/tests/ssim/"):
|
||||
name = Path(path).name
|
||||
if name.startswith("test_") and name.endswith(".py"):
|
||||
if path in removed:
|
||||
plan.reasons.append(f"removed SSIM test has nothing to run: {path}")
|
||||
else:
|
||||
plan.add_ssim((name, ), reason=f"changed SSIM test: {path}")
|
||||
plan.add_ssim((name, ), reason=f"changed SSIM test: {path}")
|
||||
elif path.endswith((".py", ".json", ".pt", ".png", ".mp4")):
|
||||
plan.ssim_all = True
|
||||
plan.add_lanes("ssim", reason=f"shared SSIM harness/reference: {path}")
|
||||
@@ -501,6 +524,10 @@ def classify_paths(paths: list[str], removed_paths: Iterable[str] = ()) -> Merge
|
||||
# DreamVerse is already one of the six automatic Fastcheck lanes.
|
||||
plan.reasons.append(f"covered by automatic DreamVerse Fastcheck: {path}")
|
||||
continue
|
||||
if path.startswith("apps/infinite_livestream/"):
|
||||
# The app's CPU-only tests run in the automatic unit Fastcheck lane.
|
||||
plan.reasons.append(f"covered by automatic unit Fastcheck: {path}")
|
||||
continue
|
||||
if path.startswith("fastvideo/tests/"):
|
||||
# The automatic unit/component Fastcheck lanes own the remaining
|
||||
# package tests. Domain-specific expensive test roots were handled
|
||||
@@ -538,7 +565,6 @@ def _write_summary(output: TextIO, plan: MergePlan) -> None:
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--paths-file", type=Path, required=True)
|
||||
parser.add_argument("--removed-paths-file", type=Path)
|
||||
parser.add_argument("--github-output", type=Path)
|
||||
parser.add_argument("--summary-file", type=Path)
|
||||
return parser.parse_args()
|
||||
@@ -547,9 +573,7 @@ def parse_args() -> argparse.Namespace:
|
||||
def main() -> int:
|
||||
args = parse_args()
|
||||
paths = args.paths_file.read_text(encoding="utf-8").splitlines()
|
||||
removed_paths = (args.removed_paths_file.read_text(encoding="utf-8").splitlines()
|
||||
if args.removed_paths_file else [])
|
||||
plan = classify_paths(paths, removed_paths)
|
||||
plan = classify_paths(paths)
|
||||
print(f"MERGE_TEST_PLAN={plan.encoded_lanes()}")
|
||||
print(f"MERGE_GOLDEN_TESTS={plan.encoded_golden_tests()}")
|
||||
print(f"MERGE_SSIM_TESTS={plan.encoded_ssim_tests()}")
|
||||
|
||||
@@ -65,7 +65,13 @@ jobs:
|
||||
&& fullSuiteOnly.size === 14
|
||||
&& [...fullSuiteOnly.values()].every(s => s.state === 'success');
|
||||
|
||||
if (fastcheckPassed) {
|
||||
// Direct reruns may repair a failed suite, never create a gate for
|
||||
// a suite that did not run.
|
||||
const failedAggregate = context => data.statuses.some(
|
||||
s => s.context === context && s.state === 'failure'
|
||||
);
|
||||
|
||||
if (failedAggregate('fastcheck-passed') && fastcheckPassed) {
|
||||
core.info(
|
||||
`All ${fastcheck.size} fastcheck tests passed — updating fastcheck-passed`
|
||||
);
|
||||
@@ -79,7 +85,7 @@ jobs:
|
||||
});
|
||||
}
|
||||
|
||||
if (fullSuitePassed) {
|
||||
if (failedAggregate('full-suite-passed') && fullSuitePassed) {
|
||||
core.info(
|
||||
'All 20 full suite tests passed — updating full-suite-passed'
|
||||
);
|
||||
|
||||
@@ -32,7 +32,7 @@ jobs:
|
||||
}
|
||||
core.setOutput('has_write', String(hasWrite));
|
||||
|
||||
- name: Add ready label and react
|
||||
- name: Add ready label
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
@@ -40,14 +40,34 @@ jobs:
|
||||
const owner = context.repo.owner;
|
||||
const repo = context.repo.repo;
|
||||
const prNumber = context.payload.issue.number;
|
||||
try { await github.rest.issues.removeLabel({ owner, repo, issue_number: prNumber, name: 'ready' }); } catch {}
|
||||
await github.rest.issues.addLabels({ owner, repo, issue_number: prNumber, labels: ['ready'] });
|
||||
|
||||
- name: React to comment
|
||||
if: steps.perm.outputs.has_write == 'true'
|
||||
continue-on-error: true
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
await github.rest.reactions.createForIssueComment({
|
||||
owner, repo,
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
comment_id: context.payload.comment.id,
|
||||
content: 'rocket',
|
||||
});
|
||||
|
||||
trigger-merge-gate:
|
||||
needs: handle-merge
|
||||
if: needs.handle-merge.result == 'success'
|
||||
permissions:
|
||||
actions: read
|
||||
contents: read
|
||||
pull-requests: read
|
||||
uses: ./.github/workflows/ci-trigger-full-suite.yml
|
||||
with:
|
||||
pr_number: ${{ github.event.issue.number }}
|
||||
secrets:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
|
||||
parse-command:
|
||||
if: >-
|
||||
github.event.issue.pull_request != null
|
||||
|
||||
@@ -3,81 +3,183 @@ name: Trigger Merge Gate
|
||||
on:
|
||||
pull_request_target:
|
||||
types: [labeled, synchronize]
|
||||
workflow_call:
|
||||
inputs:
|
||||
pr_number:
|
||||
description: Pull request number to enter into the merge gate
|
||||
required: true
|
||||
type: number
|
||||
secrets:
|
||||
BUILDKITE_API_TOKEN:
|
||||
required: true
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: read
|
||||
actions: read
|
||||
|
||||
concurrency:
|
||||
group: merge-gate-${{ github.event.pull_request.number }}
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
trigger:
|
||||
if: >-
|
||||
(github.event.action == 'labeled' && github.event.label.name == 'ready')
|
||||
inputs.pr_number > 0
|
||||
|| (github.event.action == 'labeled' && github.event.label.name == 'ready')
|
||||
|| github.event.action == 'synchronize'
|
||||
runs-on: ubuntu-latest
|
||||
# Job-level concurrency: only this guarded job acquires the group, so an
|
||||
# unrelated `labeled` event (which skips the job) cannot cancel an in-flight
|
||||
# gate and then skip its replacement. The newest real trigger (`ready`,
|
||||
# push, or `/merge`) supersedes the in-flight run, whose Buildkite build the
|
||||
# cancel step below replaces.
|
||||
concurrency:
|
||||
group: merge-gate-${{ inputs.pr_number || github.event.pull_request.number }}
|
||||
cancel-in-progress: true
|
||||
# Gate below may wait for cheap checks (up to MAX_WAIT_SECS = 25 min).
|
||||
timeout-minutes: 35
|
||||
steps:
|
||||
- name: Check ready label
|
||||
id: check
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
env:
|
||||
CALLED_PR_NUMBER: ${{ inputs.pr_number }}
|
||||
with:
|
||||
script: |
|
||||
const eventPrNumber = context.payload.pull_request?.number;
|
||||
const calledPrNumber = Number(process.env.CALLED_PR_NUMBER);
|
||||
const prNumber = eventPrNumber ?? calledPrNumber;
|
||||
if (!Number.isSafeInteger(prNumber) || prNumber <= 0) {
|
||||
core.setFailed(`Invalid pull request number: ${process.env.CALLED_PR_NUMBER}`);
|
||||
return;
|
||||
}
|
||||
const { data: pr } = await github.rest.pulls.get({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
pull_number: context.payload.pull_request.number,
|
||||
pull_number: prNumber,
|
||||
});
|
||||
if (pr.state !== 'open') {
|
||||
core.setFailed(`PR #${prNumber} is not open.`);
|
||||
return;
|
||||
}
|
||||
if (pr.base.repo.full_name !== context.payload.repository.full_name
|
||||
|| pr.base.ref !== context.payload.repository.default_branch) {
|
||||
core.setFailed(`PR #${prNumber} does not target this repository's default branch.`);
|
||||
return;
|
||||
}
|
||||
const hasReady = pr.labels.some(l => l.name === 'ready');
|
||||
core.setOutput('has_ready', String(hasReady));
|
||||
core.setOutput('changed_files', String(pr.changed_files));
|
||||
core.setOutput('pr_number', String(pr.number));
|
||||
core.setOutput('head_sha', pr.head.sha);
|
||||
core.setOutput('head_ref', pr.head.ref);
|
||||
core.setOutput('base_sha', pr.base.sha);
|
||||
core.setOutput('title', pr.title);
|
||||
if (!hasReady) core.info('No ready label — skipping merge-gate trigger.');
|
||||
|
||||
- name: Cancel previous Buildkite builds
|
||||
# Cancelling stale builds only saves agent time. If it cannot run, the
|
||||
# merge gate must still be triggered by the steps below, so a failure
|
||||
# here is reported and stepped over rather than ending the job.
|
||||
continue-on-error: true
|
||||
timeout-minutes: 3
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
PR_BRANCH: ${{ steps.check.outputs.head_ref }}
|
||||
PR_NUMBER: ${{ steps.check.outputs.pr_number }}
|
||||
run: |
|
||||
# Match both branch and PR number: forks can reuse the same branch name.
|
||||
builds=$(curl -sS --get -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
--data-urlencode "branch=$PR_BRANCH" \
|
||||
--data-urlencode "state=running,scheduled" \
|
||||
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds" \
|
||||
| jq -r --arg pr_number "$PR_NUMBER" \
|
||||
'.[] | select((.env.TEST_SCOPE? == "merge") and (.env.PR_NUMBER? == $pr_number)) | .number')
|
||||
for build_num in $builds; do
|
||||
echo "Cancelling Buildkite build #$build_num"
|
||||
curl -sS -X PUT -H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
"https://api.buildkite.com/v2/organizations/${{ vars.BUILDKITE_ORG_SLUG }}/pipelines/${{ vars.BUILDKITE_PIPELINE_SLUG }}/builds/${build_num}/cancel"
|
||||
done
|
||||
set -euo pipefail
|
||||
response_file=$(mktemp)
|
||||
builds_file=$(mktemp)
|
||||
trap 'rm -f "$response_file" "$builds_file"' EXIT
|
||||
|
||||
# Check out the immutable BASE SHA: pull_request_target must never run a
|
||||
# planner or gate script from the untrusted PR head.
|
||||
if [[ ! "$PR_NUMBER" =~ ^[1-9][0-9]*$ ]]; then
|
||||
echo "::warning::Invalid pull request number; stale Buildkite builds may continue."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if ! curl -sS --fail-with-body --connect-timeout 5 --max-time 20 --get \
|
||||
-H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
--data-urlencode "branch=$PR_BRANCH" \
|
||||
--data-urlencode "state[]=running" \
|
||||
--data-urlencode "state[]=scheduled" \
|
||||
--data-urlencode "state[]=failing" \
|
||||
--data-urlencode "exclude_jobs=true" \
|
||||
--data-urlencode "exclude_pipeline=true" \
|
||||
--output "$response_file" \
|
||||
"https://api.buildkite.com/v2/organizations/${BK_ORG}/pipelines/${BK_PIPELINE}/builds"; then
|
||||
echo "::warning::Could not list Buildkite builds; stale merge-gate builds may continue."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
if ! jq -e '
|
||||
if type != "array" then false
|
||||
else all(.[];
|
||||
if type != "object" then false
|
||||
else
|
||||
(.number | if type == "number" then . > 0 and floor == . else false end)
|
||||
and (
|
||||
(.env? | if . == null then {} else . end) as $env
|
||||
| if ($env | type) != "object" then false
|
||||
else
|
||||
($env.TEST_SCOPE? | . == null or type == "string")
|
||||
and ($env.PR_NUMBER? | . == null or type == "string")
|
||||
end
|
||||
)
|
||||
end
|
||||
)
|
||||
end
|
||||
' "$response_file" >/dev/null 2>&1; then
|
||||
# Do not print the response body: it is remote data and may contain
|
||||
# multiline values that would be interpreted as workflow commands.
|
||||
echo "::warning::Buildkite returned an invalid build list; stale merge-gate builds may continue."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Match both branch and PR number: forks can reuse the same branch name.
|
||||
if ! jq -r --arg pr_number "$PR_NUMBER" '
|
||||
.[]
|
||||
| select((.env.TEST_SCOPE? == "merge") and (.env.PR_NUMBER? == $pr_number))
|
||||
| .number
|
||||
' "$response_file" > "$builds_file"; then
|
||||
echo "::warning::Could not select stale Buildkite builds; stale merge-gate builds may continue."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
cancellation_failed=0
|
||||
while IFS= read -r build_num; do
|
||||
echo "Cancelling Buildkite build #$build_num"
|
||||
if ! curl -sS --fail-with-body --connect-timeout 5 --max-time 20 -o /dev/null -X PUT \
|
||||
-H "Authorization: Bearer $BUILDKITE_API_TOKEN" \
|
||||
"https://api.buildkite.com/v2/organizations/${BK_ORG}/pipelines/${BK_PIPELINE}/builds/${build_num}/cancel"; then
|
||||
echo "::warning::Could not cancel Buildkite build #$build_num; trying remaining builds."
|
||||
cancellation_failed=1
|
||||
fi
|
||||
done < "$builds_file"
|
||||
|
||||
if (( cancellation_failed != 0 )); then
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Check out the immutable BASE SHA: neither pull_request_target nor the
|
||||
# privileged slash-command call may run code from the untrusted PR head.
|
||||
- name: Checkout trusted merge planner
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 # v4.2.2
|
||||
with:
|
||||
ref: ${{ github.event.pull_request.base.sha }}
|
||||
ref: ${{ steps.check.outputs.base_sha }}
|
||||
persist-credentials: false
|
||||
|
||||
- name: Collect changed paths
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
PR_NUMBER: ${{ steps.check.outputs.pr_number }}
|
||||
EXPECTED_CHANGED_FILES: ${{ steps.check.outputs.changed_files }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
changed_json="$RUNNER_TEMP/merge-changed-files.json"
|
||||
changed_paths="$RUNNER_TEMP/merge-changed-paths.txt"
|
||||
removed_paths="$RUNNER_TEMP/merge-removed-paths.txt"
|
||||
: > "$removed_paths"
|
||||
if gh api --paginate --slurp \
|
||||
"repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}/files?per_page=100" \
|
||||
> "$changed_json"; then
|
||||
@@ -85,10 +187,6 @@ jobs:
|
||||
if [ "$observed" = "$EXPECTED_CHANGED_FILES" ]; then
|
||||
jq -r '.[][] | .filename, (.previous_filename // empty)' "$changed_json" \
|
||||
| sort -u > "$changed_paths"
|
||||
# Paths absent from the PR head, so the planner never selects a deleted test.
|
||||
jq -r '.[][] | if .status == "removed" then .filename
|
||||
elif .status == "renamed" then (.previous_filename // empty) else empty end' \
|
||||
"$changed_json" | sort -u > "$removed_paths"
|
||||
else
|
||||
echo "::warning::Changed-file API returned $observed of $EXPECTED_CHANGED_FILES paths; selecting all merge lanes."
|
||||
echo '__FASTVIDEO_CI_PLAN_ALL__' > "$changed_paths"
|
||||
@@ -104,7 +202,6 @@ jobs:
|
||||
run: |
|
||||
python3 .github/scripts/plan_merge_ci.py \
|
||||
--paths-file "$RUNNER_TEMP/merge-changed-paths.txt" \
|
||||
--removed-paths-file "$RUNNER_TEMP/merge-removed-paths.txt" \
|
||||
--github-output "$GITHUB_OUTPUT" \
|
||||
--summary-file "$GITHUB_STEP_SUMMARY"
|
||||
|
||||
@@ -112,18 +209,18 @@ jobs:
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ github.token }}
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
PR_SHA: ${{ steps.check.outputs.head_sha }}
|
||||
PR_NUMBER: ${{ steps.check.outputs.pr_number }}
|
||||
run: bash .github/scripts/gate_full_suite.sh
|
||||
|
||||
- name: Trigger Buildkite merge gate
|
||||
if: steps.check.outputs.has_ready == 'true'
|
||||
env:
|
||||
BUILDKITE_API_TOKEN: ${{ secrets.BUILDKITE_API_TOKEN }}
|
||||
PR_SHA: ${{ github.event.pull_request.head.sha }}
|
||||
PR_BRANCH: ${{ github.event.pull_request.head.ref }}
|
||||
PR_NUMBER: ${{ github.event.pull_request.number }}
|
||||
PR_TITLE: ${{ github.event.pull_request.title }}
|
||||
PR_SHA: ${{ steps.check.outputs.head_sha }}
|
||||
PR_BRANCH: ${{ steps.check.outputs.head_ref }}
|
||||
PR_NUMBER: ${{ steps.check.outputs.pr_number }}
|
||||
PR_TITLE: ${{ steps.check.outputs.title }}
|
||||
BK_ORG: ${{ vars.BUILDKITE_ORG_SLUG }}
|
||||
BK_PIPELINE: ${{ vars.BUILDKITE_PIPELINE_SLUG }}
|
||||
MERGE_TEST_PLAN: ${{ steps.plan.outputs.merge_test_plan }}
|
||||
|
||||
@@ -68,7 +68,7 @@ jobs:
|
||||
- os: ubuntu-22.04
|
||||
arch: x86_64
|
||||
wheel-plat: manylinux_2_35_x86_64
|
||||
# aarch64 is Blackwell (GB200 sm_100a + DGX Spark / consumer sm_120a), not
|
||||
# aarch64 is Blackwell (GB200 sm_100a/sm_103a + sm_120a + DGX Spark sm_121a), not
|
||||
# Hopper, and Blackwell needs CUDA >= 12.8 — so only the cu130 leg applies.
|
||||
# Added via include so x86 keeps cu126 + cu130 while aarch64 stays cu130-only.
|
||||
include:
|
||||
@@ -164,19 +164,20 @@ jobs:
|
||||
cd fastvideo-kernel
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
# Release builds run on GPU-less runners, so set kernels + arch explicitly:
|
||||
# * aarch64 = Blackwell (GB200 sm_100a + DGX Spark/consumer sm_120a), NOT
|
||||
# Hopper, so TK (sm_90a wgmma) is OFF. The C++ FP4 (attn_qat_infer, SM120)
|
||||
# covers sm_120a; turbodiffusion covers sm_100a+sm_120a. The sm_100 FP4
|
||||
# forward is the FA4 CuTe DSL path in the fastvideo package (PR #1221),
|
||||
# * aarch64 = Blackwell (GB200 sm_100a/sm_103a + sm_120a + DGX Spark sm_121a), NOT
|
||||
# Hopper, so TK (sm_90a wgmma) is OFF. The C++ FP4 (attn_qat_infer)
|
||||
# covers sm_120a+sm_121a; turbodiffusion covers every listed arch. The
|
||||
# sm_100 FP4 forward is the FA4 CuTe DSL path in the fastvideo package (PR #1221),
|
||||
# JIT-compiled at runtime — not built into this wheel.
|
||||
# * x86_64 cu130 = Hopper TK + data-center Blackwell sm_100a/sm_103a VSA
|
||||
# + consumer Blackwell sm_120a FP4.
|
||||
# * x86_64 cu126 = Hopper TK only (older drivers; CUDA < 12.8 has no FP4).
|
||||
# The per-arch split in CMakeLists pins the FP4 targets to sm_120a and builds
|
||||
# the main extension for the full arch list. CMAKE_BUILD_PARALLEL_LEVEL caps
|
||||
# The per-arch split in CMakeLists pins the FP4 targets to requested
|
||||
# sm_120a/sm_121a and builds the main extension for the full arch list.
|
||||
# CMAKE_BUILD_PARALLEL_LEVEL caps
|
||||
# Ninja so heavy CUTLASS/TK template TUs don't OOM the 16 GB runner (exit 143).
|
||||
if [ "${{ matrix.platform.arch }}" = "aarch64" ]; then
|
||||
export TORCH_CUDA_ARCH_LIST="10.0a;10.3a;12.0a"
|
||||
export TORCH_CUDA_ARCH_LIST="10.0a;10.3a;12.0a;12.1a"
|
||||
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=OFF -DFASTVIDEO_KERNEL_BUILD_ATTN_QAT_INFER=ON"
|
||||
export CMAKE_BUILD_PARALLEL_LEVEL=1
|
||||
elif [ "${{ matrix.torch-cuda.torch-cuda-short }}" = "cu130" ]; then
|
||||
@@ -250,7 +251,8 @@ jobs:
|
||||
- name: Download PyPI wheels
|
||||
# Publish the cu130 (CUDA 13) wheels to PyPI for both architectures:
|
||||
# x86_64 — Hopper sm_90a TK + consumer Blackwell sm_120a FP4
|
||||
# aarch64 — Blackwell: turbodiffusion (sm_100a/sm_120a) + C++ FP4 (sm_120a);
|
||||
# aarch64 — Blackwell: turbodiffusion (sm_100a/sm_103a/sm_120a/sm_121a)
|
||||
# + C++ FP4 (sm_120a/sm_121a);
|
||||
# no TK (Hopper). sm_100 FP4 forward is the FA4 CuTe DSL path in the
|
||||
# fastvideo package (#1221), shipped/JIT separately.
|
||||
# The x86_64 cu126 wheel stays available as a build artifact / GitHub-release asset.
|
||||
|
||||
@@ -36,6 +36,7 @@ env
|
||||
*.log
|
||||
weights/
|
||||
logs/
|
||||
/Z-Image/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
@@ -142,5 +143,7 @@ fastvideo/tests/ssim/.reference_videos_download.lock
|
||||
*.nvimlog
|
||||
.nvimlog
|
||||
.python-version
|
||||
/LTX-2-Reference/
|
||||
/DFDReference/
|
||||
scripts/benchmarks/minimax_h3_pro6000/headline_results/
|
||||
fastvideo/tests/ssim/.reference_videos_download.lock
|
||||
|
||||
@@ -7,3 +7,6 @@
|
||||
[submodule "fastvideo/third_party/eval/vbench"]
|
||||
path = fastvideo/third_party/eval/vbench
|
||||
url = https://github.com/Vchitect/VBench.git
|
||||
[submodule "fastvideo/third_party/eval/vqeval"]
|
||||
path = fastvideo/third_party/eval/vqeval
|
||||
url = https://github.com/JiusiServe/LongVideoSparseAttention.git
|
||||
|
||||
@@ -154,6 +154,8 @@ if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
`num_gpus=1` runs the worker in-process (weights load once, no extra Python process). On Colab/Kaggle-style machines with ~16GB host RAM, keep `num_gpus=1`; free-tier system memory does not grow with extra T4s, so `num_gpus>1` is likely to OOM.
|
||||
|
||||
Run the script with:
|
||||
|
||||
```bash
|
||||
|
||||
@@ -158,6 +158,40 @@ dreamverse-server --host 0.0.0.0 --port 8009
|
||||
The Dreamverse backend defaults to `0.0.0.0:8009` and starts one GPU worker on
|
||||
the first visible GPU by default.
|
||||
|
||||
### Cosmos Predict2.5 DFD continuation (experimental)
|
||||
|
||||
Dreamverse can combine two converted Cosmos Predict2.5 2B packages: the
|
||||
distilled Text2World student creates an unconditioned first segment, then the
|
||||
Data-Forcing Distillation (DFD) Video2World student conditions each later
|
||||
segment on the prior terminal frame. Point the runtime at both local converted
|
||||
packages:
|
||||
|
||||
```bash
|
||||
export DREAMVERSE_MODEL_ID=cosmos25-dfd
|
||||
export DREAMVERSE_MODEL_PATH=/path/to/Cosmos-Predict2.5-2B-Distilled-TrigFlow-FastVideo
|
||||
export DREAMVERSE_COSMOS25_DFD_MODEL_PATH=/path/to/Cosmos-Predict2.5-2B-DFD-FastVideo
|
||||
export ENABLE_TORCH_COMPILE=0
|
||||
dreamverse-server --host 0.0.0.0 --port 8009
|
||||
```
|
||||
|
||||
The backend loads and warms both model roles before reporting ready. Both use
|
||||
BF16, Torch SDPA, 704x1280 output, 24 FPS, and four steps. Bootstrap segments
|
||||
contain 77 frames. DFD segments contain 81 decoded frames, but Dreamverse drops
|
||||
the repeated conditioning frame before streaming, leaving 80 new frames. An
|
||||
initial user image selects DFD immediately without treating that first frame as
|
||||
a cross-segment overlap.
|
||||
|
||||
The profile uses a 30-minute session lease because sequential generation on
|
||||
GB10-class hardware can exceed Dreamverse's five-minute default while the GPU
|
||||
is still making progress. Deployments can override the lease with
|
||||
`FASTVIDEO_SESSION_TIMEOUT_SECONDS`.
|
||||
|
||||
Cosmos does not produce audio, so the backend supplies duration-matched silent
|
||||
24 kHz audio for the existing browser streaming contract and trims 1,000 audio
|
||||
samples with each repeated DFD boundary frame. Runtime LoRA changes are not
|
||||
supported. Full segments take roughly 145 seconds on GB10, so this profile is a
|
||||
continuation-quality integration rather than a real-time configuration.
|
||||
|
||||
### Check Readiness
|
||||
|
||||
In another shell, verify that the backend process is alive:
|
||||
|
||||
@@ -0,0 +1,184 @@
|
||||
"""Bounded, runtime-local media library shared by the HTTP and generation APIs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import re
|
||||
import shutil
|
||||
import subprocess
|
||||
import tempfile
|
||||
import threading
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from PIL import Image, UnidentifiedImageError
|
||||
|
||||
IMAGE_LIMIT = 15 * 1024 * 1024
|
||||
MEDIA_LIMIT = 100 * 1024 * 1024
|
||||
STORE_LIMIT = 2 * 1024 * 1024 * 1024
|
||||
ASSET_LIMIT = 100
|
||||
MAX_MEDIA_SECONDS = 30
|
||||
MIME_TYPES = {
|
||||
"image/png": ("image", ".png"),
|
||||
"image/jpeg": ("image", ".jpg"),
|
||||
"image/webp": ("image", ".webp"),
|
||||
"video/mp4": ("video", ".mp4"),
|
||||
"video/quicktime": ("video", ".mov"),
|
||||
"video/webm": ("video", ".webm"),
|
||||
"audio/mpeg": ("audio", ".mp3"),
|
||||
"audio/mp4": ("audio", ".m4a"),
|
||||
"audio/x-m4a": ("audio", ".m4a"),
|
||||
"audio/wav": ("audio", ".wav"),
|
||||
"audio/x-wav": ("audio", ".wav"),
|
||||
"audio/flac": ("audio", ".flac"),
|
||||
"audio/x-flac": ("audio", ".flac"),
|
||||
"audio/ogg": ("audio", ".ogg"),
|
||||
"audio/webm": ("audio", ".webm"),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class StoredAsset:
|
||||
asset_id: str
|
||||
kind: str
|
||||
path: str
|
||||
name: str
|
||||
mime_type: str
|
||||
size: int
|
||||
|
||||
def public(self) -> dict:
|
||||
return {
|
||||
"asset_id": self.asset_id,
|
||||
"kind": self.kind,
|
||||
"name": self.name,
|
||||
"mime_type": self.mime_type,
|
||||
"size": self.size,
|
||||
"url": f"/assets/{self.asset_id}",
|
||||
}
|
||||
|
||||
|
||||
def validate_media(path: Path, mime_type: str) -> None:
|
||||
"""Inspect content, not filenames; refuse playlists and non-media uploads."""
|
||||
kind = MIME_TYPES[mime_type][0]
|
||||
if kind == "image":
|
||||
try:
|
||||
with Image.open(path) as img:
|
||||
expected = {"image/png": "PNG", "image/jpeg": "JPEG", "image/webp": "WEBP"}[mime_type]
|
||||
if img.format != expected:
|
||||
raise ValueError("The image content does not match its file type.")
|
||||
if img.width * img.height > 16_777_216:
|
||||
raise ValueError("Images must contain at most 16 megapixels.")
|
||||
if getattr(img, "is_animated", False):
|
||||
raise ValueError("Use a still image or upload the animation as a video.")
|
||||
img.verify()
|
||||
except (UnidentifiedImageError, OSError, Image.DecompressionBombError) as exc:
|
||||
raise ValueError("The image could not be decoded. Use PNG, JPEG, or WebP.") from exc
|
||||
return
|
||||
|
||||
probe = shutil.which(os.getenv("FASTVIDEO_FFPROBE_BIN", "ffprobe"))
|
||||
if not probe:
|
||||
raise ValueError("This runtime needs ffprobe installed to accept video and audio assets.")
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[
|
||||
probe, "-v", "error", "-protocol_whitelist", "file,pipe", "-format_whitelist",
|
||||
"mov,matroska,webm,mp3,wav,flac,ogg", "-show_format", "-show_streams", "-of", "json",
|
||||
str(path)
|
||||
],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
timeout=15,
|
||||
)
|
||||
info = json.loads(result.stdout)
|
||||
formats = set(info.get("format", {}).get("format_name", "").split(","))
|
||||
if not formats.intersection({"mov", "mp4", "matroska", "webm", "mp3", "wav", "flac", "ogg"}):
|
||||
raise ValueError("Upload a media file, not a playlist or external reference.")
|
||||
streams = [stream for stream in info.get("streams", []) if stream.get("codec_type") == kind]
|
||||
if not streams:
|
||||
raise ValueError(f"The file contains no {kind} stream.")
|
||||
for stream in info.get("streams", []):
|
||||
if stream.get("codec_type") == "audio" and int(stream.get("channels", 0)) not in (1, 2):
|
||||
raise ValueError("H3 references require mono or stereo audio, including video soundtracks.")
|
||||
duration = float(info.get("format", {}).get("duration", "nan"))
|
||||
if not math.isfinite(duration) or not 0 < duration <= MAX_MEDIA_SECONDS:
|
||||
raise ValueError(f"Reference video and audio must be between 0 and {MAX_MEDIA_SECONDS} seconds long.")
|
||||
for stream in streams:
|
||||
if kind == "video" and int(stream.get("width", 0)) * int(stream.get("height", 0)) > 8_294_400:
|
||||
raise ValueError("Reference videos must be 4K or smaller.")
|
||||
except (subprocess.SubprocessError, json.JSONDecodeError, OSError) as exc:
|
||||
raise ValueError("The media file could not be decoded. Check its format and try again.") from exc
|
||||
|
||||
|
||||
class AssetStore:
|
||||
"""Assets live until deletion or runtime exit; pinned generation inputs cannot be deleted."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._directory: tempfile.TemporaryDirectory | None = None
|
||||
self._assets: dict[str, StoredAsset] = {}
|
||||
self._pins: dict[str, int] = {}
|
||||
self._lock = threading.RLock()
|
||||
|
||||
def staging_path(self, mime_type: str) -> Path:
|
||||
with self._lock:
|
||||
if mime_type not in MIME_TYPES:
|
||||
raise ValueError("Unsupported media type. Use PNG/JPEG/WebP, MP4/WebM/MOV, or WAV/MP3/M4A/FLAC/OGG.")
|
||||
if len(self._assets) >= ASSET_LIMIT or sum(item.size for item in self._assets.values()) >= STORE_LIMIT:
|
||||
raise ValueError("The runtime asset library is full. Remove unused assets before uploading more.")
|
||||
if self._directory is None:
|
||||
self._directory = tempfile.TemporaryDirectory(prefix="dreamverse-assets-")
|
||||
return Path(self._directory.name) / f"{uuid.uuid4().hex}{MIME_TYPES[mime_type][1]}"
|
||||
|
||||
def add(self, path: Path, name: str, mime_type: str) -> StoredAsset:
|
||||
validate_media(path, mime_type)
|
||||
size = path.stat().st_size
|
||||
if size == 0 or size > (IMAGE_LIMIT if MIME_TYPES[mime_type][0] == "image" else MEDIA_LIMIT):
|
||||
raise ValueError("The asset is empty or exceeds its upload size limit.")
|
||||
with self._lock:
|
||||
if len(self._assets) >= ASSET_LIMIT or size + sum(item.size
|
||||
for item in self._assets.values()) > STORE_LIMIT:
|
||||
raise ValueError("The runtime asset library is full. Remove unused assets before uploading more.")
|
||||
if self._directory is None or path.parent != Path(self._directory.name):
|
||||
raise ValueError("The asset must be uploaded to this runtime.")
|
||||
display_name = re.sub(r"[\x00-\x1f\x7f/\\]", "_", name).strip()[:200] or "Untitled asset"
|
||||
asset = StoredAsset(path.stem, MIME_TYPES[mime_type][0], str(path), display_name, mime_type, size)
|
||||
self._assets[asset.asset_id] = asset
|
||||
return asset
|
||||
|
||||
def get(self, asset_id: str) -> StoredAsset:
|
||||
with self._lock:
|
||||
if not isinstance(asset_id, str) or not re.fullmatch(r"[a-f0-9]{32}", asset_id):
|
||||
raise ValueError("Invalid asset ID. Upload or select an asset from the library.")
|
||||
asset = self._assets.get(asset_id)
|
||||
if asset is None or not Path(asset.path).is_file():
|
||||
raise ValueError("An asset is no longer available. Upload it again and reselect it.")
|
||||
return asset
|
||||
|
||||
def pin(self, asset_ids: list[str]) -> None:
|
||||
with self._lock:
|
||||
for asset_id in asset_ids:
|
||||
self.get(asset_id)
|
||||
for asset_id in asset_ids:
|
||||
self._pins[asset_id] = self._pins.get(asset_id, 0) + 1
|
||||
|
||||
def release(self, asset_ids: list[str]) -> None:
|
||||
with self._lock:
|
||||
for asset_id in asset_ids:
|
||||
count = self._pins.get(asset_id, 0)
|
||||
if count > 1:
|
||||
self._pins[asset_id] = count - 1
|
||||
else:
|
||||
self._pins.pop(asset_id, None)
|
||||
|
||||
def delete(self, asset_id: str) -> None:
|
||||
with self._lock:
|
||||
asset = self.get(asset_id)
|
||||
if self._pins.get(asset_id, 0):
|
||||
raise ValueError("This asset is in use by a generation session. End the session before deleting it.")
|
||||
Path(asset.path).unlink(missing_ok=True)
|
||||
del self._assets[asset_id]
|
||||
|
||||
|
||||
asset_store = AssetStore()
|
||||
@@ -111,7 +111,13 @@ def _build_generator_config(model_path: str, enable_compile: bool, num_gpus: int
|
||||
mode="max-autotune-no-cudagraphs",
|
||||
dynamic=False),
|
||||
use_fsdp_inference=False,
|
||||
quantization=QuantizationConfig(transformer_quant="NVFP4"),
|
||||
# The bundled LTX2 model enables a refinement LoRA during the
|
||||
# first request. NVFP4 otherwise purges the dense weights that
|
||||
# FastVideo's LoRA merge path requires.
|
||||
quantization=QuantizationConfig(
|
||||
transformer_quant="NVFP4",
|
||||
transformer_retain_original_weights=True,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=components,
|
||||
|
||||
@@ -84,6 +84,37 @@ MODEL_REGISTRY = {
|
||||
"num_inference_steps": 5,
|
||||
"seed": 1000,
|
||||
},
|
||||
"full-h3": {
|
||||
"name": "MiniMax H3 (Full)",
|
||||
"generation_backend": "minimax_h3",
|
||||
"default_sp_size": 4,
|
||||
"model_path": "MiniMaxAI/MiniMax-H3",
|
||||
"attention_backend": "FLASH_ATTN",
|
||||
"height": 768,
|
||||
"width": 1344,
|
||||
"num_frames": 124,
|
||||
"num_inference_steps": 50,
|
||||
"seed": 1000,
|
||||
"full_checkpoint": True,
|
||||
},
|
||||
"cosmos25-dfd": {
|
||||
"name": "Cosmos Predict2.5 DFD",
|
||||
"generation_backend": "cosmos25_dfd",
|
||||
"default_sp_size": 1,
|
||||
"model_path": "FastVideo/Cosmos-Predict2.5-2B-Distilled-TrigFlow",
|
||||
"continuation_model_path": "FastVideo/Cosmos-Predict2.5-2B-DFD",
|
||||
"attention_backend": "TORCH_SDPA",
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"bootstrap_num_frames": 77,
|
||||
"continuation_num_frames": 81,
|
||||
"fps": 24,
|
||||
"num_inference_steps": 4,
|
||||
"seed": 42,
|
||||
# Six sequential GB10 segments can exceed the legacy five-minute
|
||||
# DreamVerse lease even though the GPU is making progress.
|
||||
"session_timeout_seconds": 1800,
|
||||
},
|
||||
}
|
||||
|
||||
DEFAULT_MODEL_ID = "fast-ltx2"
|
||||
@@ -95,22 +126,6 @@ if ACTIVE_MODEL_ID not in MODEL_REGISTRY:
|
||||
# Active model configuration
|
||||
MODEL_CONFIG = MODEL_REGISTRY[ACTIVE_MODEL_ID]
|
||||
|
||||
# Generation limits
|
||||
SESSION_TIMEOUT_SECONDS = 300
|
||||
|
||||
# Frame settings
|
||||
NUM_FRAMES = 121
|
||||
FRAME_HEIGHT = 1088
|
||||
FRAME_WIDTH = 1920
|
||||
NUM_INFERENCE_STEPS = 5
|
||||
JPEG_QUALITY = 100
|
||||
BATCH_SIZE = 3
|
||||
|
||||
# Streaming mode:
|
||||
# - legacy_jpeg: send frame_batch JSON payloads with base64 JPEGs
|
||||
# - av_fmp4: send muxed fMP4 binary chunks over WebSocket
|
||||
STREAM_MODE = os.getenv("STREAM_MODE", "av_fmp4").strip().lower()
|
||||
|
||||
|
||||
def _env_int(name: str, default: int) -> int:
|
||||
value = os.getenv(name)
|
||||
@@ -187,6 +202,38 @@ def _optional_env(*names: str) -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
# Generation limits
|
||||
# Slower backends may own a longer default lease. A profile can set
|
||||
# ``session_timeout_seconds``; Full H3 loads and generates substantially longer
|
||||
# than the Preview adapter, which also covers a base/ref pipeline reload inside a
|
||||
# retained session. An explicit environment override remains available for
|
||||
# deployment policy: DREAMVERSE_SESSION_TIMEOUT_SECONDS, with
|
||||
# FASTVIDEO_SESSION_TIMEOUT_SECONDS accepted as an alias.
|
||||
# Values below 60 seconds are floored so a single segment cannot outlast the session.
|
||||
_DEFAULT_SESSION_TIMEOUT_SECONDS = cast(
|
||||
int, MODEL_CONFIG.get("session_timeout_seconds", 7200 if ACTIVE_MODEL_ID == "full-h3" else 300))
|
||||
SESSION_TIMEOUT_SECONDS = max(
|
||||
60,
|
||||
_env_int(
|
||||
"DREAMVERSE_SESSION_TIMEOUT_SECONDS",
|
||||
_env_int("FASTVIDEO_SESSION_TIMEOUT_SECONDS", _DEFAULT_SESSION_TIMEOUT_SECONDS),
|
||||
),
|
||||
)
|
||||
|
||||
# Frame settings
|
||||
NUM_FRAMES = 121
|
||||
FRAME_HEIGHT = 1088
|
||||
FRAME_WIDTH = 1920
|
||||
NUM_INFERENCE_STEPS = 5
|
||||
JPEG_QUALITY = 100
|
||||
BATCH_SIZE = 3
|
||||
|
||||
# Streaming mode:
|
||||
# - legacy_jpeg: send frame_batch JSON payloads with base64 JPEGs
|
||||
# - av_fmp4: send muxed fMP4 binary chunks over WebSocket
|
||||
STREAM_MODE = os.getenv("STREAM_MODE", "av_fmp4").strip().lower()
|
||||
|
||||
|
||||
DEVTOOLS_ENABLED = _env_bool("FASTVIDEO_ENABLE_DEVTOOLS", False)
|
||||
PROMPT_SAFETY_ENABLED = _env_bool("FASTVIDEO_ENABLE_PROMPT_SAFETY", False)
|
||||
DREAMVERSE_MAX_AUTOTUNE = _env_bool("DREAMVERSE_MAX_AUTOTUNE", True)
|
||||
@@ -200,6 +247,13 @@ if DREAMVERSE_MODEL_PATH:
|
||||
"config_model_path": DREAMVERSE_MODEL_PATH,
|
||||
}
|
||||
|
||||
DREAMVERSE_COSMOS25_DFD_MODEL_PATH = (os.getenv("DREAMVERSE_COSMOS25_DFD_MODEL_PATH", "").strip() or None)
|
||||
if DREAMVERSE_COSMOS25_DFD_MODEL_PATH and MODEL_CONFIG.get("generation_backend") == "cosmos25_dfd":
|
||||
MODEL_CONFIG = {
|
||||
**MODEL_CONFIG,
|
||||
"continuation_model_path": DREAMVERSE_COSMOS25_DFD_MODEL_PATH,
|
||||
}
|
||||
|
||||
AVAILABLE_LORAS = {
|
||||
"pixar": {
|
||||
"repo": "vrgamedevgirl84/LTX_2.3_Pixar_Toon_Style_LoRa",
|
||||
|
||||
@@ -0,0 +1,252 @@
|
||||
"""Cosmos Predict2.5 distilled bootstrap and DFD continuation for DreamVerse."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import os
|
||||
import time
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from dreamverse.generation_contracts import StepResult
|
||||
from dreamverse.generation_inputs import GenerationInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from PIL.Image import Image
|
||||
|
||||
_SILENT_AUDIO_SAMPLE_RATE = 24_000
|
||||
|
||||
|
||||
def _required_config_str(model_config: dict, field_name: str) -> str:
|
||||
value = model_config.get(field_name)
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError(f"Cosmos Predict2.5 DFD model configuration requires `{field_name}`.")
|
||||
return value.strip()
|
||||
|
||||
|
||||
class Cosmos25DFDGenerationBackend:
|
||||
"""Own complementary Cosmos T2W and one-frame-conditioned DFD generators."""
|
||||
|
||||
def __init__(self, gpu_id: int):
|
||||
self.gpu_id = gpu_id
|
||||
self.bootstrap_generator: Any | None = None
|
||||
self.continuation_generator: Any | None = None
|
||||
self.model_config: dict = {}
|
||||
self.continuation_image: Image | None = None
|
||||
|
||||
def _gpu_mem(self) -> str:
|
||||
allocated_gib = torch.cuda.memory_allocated() / 1024**3
|
||||
reserved_gib = torch.cuda.memory_reserved() / 1024**3
|
||||
return f"alloc={allocated_gib:.2f}GiB, reserved={reserved_gib:.2f}GiB"
|
||||
|
||||
@staticmethod
|
||||
def _configure_environment(attention_backend: str) -> None:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = attention_backend
|
||||
os.environ.pop("FASTVIDEO_INFERENCE_TORCH_COMPILE", None)
|
||||
|
||||
@staticmethod
|
||||
def _load_generator(model_path: str):
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
return VideoGenerator.from_pretrained(
|
||||
model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
enable_torch_compile=False,
|
||||
)
|
||||
|
||||
def initialize(self, model_config: dict | None = None) -> None:
|
||||
"""Load both package roles so bootstrap and continuation are ready."""
|
||||
if model_config is not None:
|
||||
self.model_config = dict(model_config)
|
||||
if not self.model_config:
|
||||
raise ValueError("Cosmos Predict2.5 DFD initialization requires a model configuration.")
|
||||
|
||||
self.shutdown()
|
||||
bootstrap_path = _required_config_str(self.model_config, "model_path")
|
||||
continuation_path = _required_config_str(self.model_config, "continuation_model_path")
|
||||
attention_backend = _required_config_str(self.model_config, "attention_backend")
|
||||
self._configure_environment(attention_backend)
|
||||
|
||||
print(f"[GPU {self.gpu_id}] Loading Cosmos T2W bootstrap: {bootstrap_path}")
|
||||
print(f"[GPU {self.gpu_id}] Before bootstrap load: {self._gpu_mem()}")
|
||||
self.bootstrap_generator = self._load_generator(bootstrap_path)
|
||||
print(f"[GPU {self.gpu_id}] Loading Cosmos DFD continuation: {continuation_path}")
|
||||
self.continuation_generator = self._load_generator(continuation_path)
|
||||
print(f"[GPU {self.gpu_id}] Cosmos T2W + DFD loaded: {self._gpu_mem()} (warmup pending)")
|
||||
|
||||
def shutdown(self) -> None:
|
||||
"""Release both FastVideo generators and the retained terminal frame."""
|
||||
self.clear_conditioning()
|
||||
for attr_name in ("bootstrap_generator", "continuation_generator"):
|
||||
generator = getattr(self, attr_name)
|
||||
if generator is not None:
|
||||
try:
|
||||
generator.shutdown()
|
||||
except Exception as exc:
|
||||
print(f"[GPU {self.gpu_id}] Cosmos generator shutdown warning: {exc}")
|
||||
setattr(self, attr_name, None)
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def clear_conditioning(self) -> None:
|
||||
if self.continuation_image is not None:
|
||||
self.continuation_image.close()
|
||||
self.continuation_image = None
|
||||
|
||||
@staticmethod
|
||||
def _load_rgb_image(image_path: str) -> Image:
|
||||
from PIL import Image
|
||||
|
||||
with Image.open(image_path) as image:
|
||||
return image.convert("RGB").copy()
|
||||
|
||||
def _select_conditioning_image(
|
||||
self,
|
||||
segment_idx: int,
|
||||
image_path: str | None,
|
||||
reset_conditioning: bool,
|
||||
) -> tuple[Image | None, bool]:
|
||||
if reset_conditioning:
|
||||
self.clear_conditioning()
|
||||
if segment_idx > 1 and self.continuation_image is not None:
|
||||
return self.continuation_image.copy(), True
|
||||
if segment_idx > 1 and not reset_conditioning:
|
||||
raise RuntimeError(f"Cosmos DFD segment {segment_idx} requires a retained continuation frame.")
|
||||
if segment_idx == 1 and image_path:
|
||||
return self._load_rgb_image(image_path), False
|
||||
return None, False
|
||||
|
||||
def _sampling_param(self, *, conditioned: bool):
|
||||
# ``num_cond_frames`` is not yet exposed by the typed SamplingConfig,
|
||||
# so this backend uses the compatibility request until that field lands.
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
|
||||
num_frames_key = "continuation_num_frames" if conditioned else "bootstrap_num_frames"
|
||||
return SamplingParam(
|
||||
negative_prompt="",
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=int(self.model_config["height"]),
|
||||
width=int(self.model_config["width"]),
|
||||
num_frames=int(self.model_config[num_frames_key]),
|
||||
fps=int(self.model_config["fps"]),
|
||||
num_inference_steps=int(self.model_config["num_inference_steps"]),
|
||||
guidance_scale=1.0,
|
||||
seed=int(self.model_config["seed"]),
|
||||
num_cond_frames=1 if conditioned else 0,
|
||||
)
|
||||
|
||||
def _save_continuation_frame(self, frame: object) -> None:
|
||||
from PIL import Image
|
||||
|
||||
self.clear_conditioning()
|
||||
if isinstance(frame, Image.Image):
|
||||
self.continuation_image = frame.convert("RGB").copy()
|
||||
return
|
||||
pixels = np.asarray(frame)
|
||||
self.continuation_image = Image.fromarray(np.ascontiguousarray(pixels)).convert("RGB")
|
||||
|
||||
@staticmethod
|
||||
def _silent_audio(frame_count: int, fps: int) -> torch.Tensor:
|
||||
sample_count = max(1, int(round((frame_count / float(fps)) * _SILENT_AUDIO_SAMPLE_RATE)))
|
||||
return torch.zeros(sample_count, dtype=torch.float32)
|
||||
|
||||
def generate_step(
|
||||
self,
|
||||
prompt: str,
|
||||
segment_idx: int,
|
||||
image_path: str | None,
|
||||
reset_conditioning: bool,
|
||||
generation_inputs: GenerationInputs | None = None,
|
||||
) -> StepResult:
|
||||
"""Generate a T2W start or DFD continuation and retain its last frame."""
|
||||
if generation_inputs is not None and (generation_inputs.mode not in (None, "t2va") or generation_inputs.assets):
|
||||
raise ValueError("Cosmos supports text generation only through the generation mode API.")
|
||||
if self.bootstrap_generator is None or self.continuation_generator is None:
|
||||
raise RuntimeError("Cosmos T2W + DFD generators are not initialized.")
|
||||
|
||||
conditioning_image, uses_continuation = self._select_conditioning_image(
|
||||
segment_idx,
|
||||
image_path,
|
||||
reset_conditioning,
|
||||
)
|
||||
conditioned = conditioning_image is not None
|
||||
generator = self.continuation_generator if conditioned else self.bootstrap_generator
|
||||
sampling_param = self._sampling_param(conditioned=conditioned)
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
if conditioned:
|
||||
sampling_param.pil_image = conditioning_image
|
||||
result = generator.generate_video(prompt, sampling_param=sampling_param)
|
||||
finally:
|
||||
if conditioning_image is not None:
|
||||
conditioning_image.close()
|
||||
torch.cuda.synchronize()
|
||||
generation_ms = (time.perf_counter() - started) * 1000.0
|
||||
|
||||
if not isinstance(result, dict):
|
||||
raise RuntimeError("Cosmos generation did not return one result dictionary.")
|
||||
frames = result.get("frames")
|
||||
expected_frames = int(sampling_param.num_frames)
|
||||
if not isinstance(frames, list) or len(frames) != expected_frames:
|
||||
actual_frames = len(frames) if isinstance(frames, list) else None
|
||||
raise RuntimeError(f"Cosmos generation returned {actual_frames} frames; expected {expected_frames}.")
|
||||
|
||||
save_started = time.perf_counter()
|
||||
self._save_continuation_frame(frames[-1])
|
||||
save_conditioning_ms = (time.perf_counter() - save_started) * 1000.0
|
||||
fps = int(sampling_param.fps)
|
||||
timings = {
|
||||
"generation_ms": generation_ms,
|
||||
"generation_time_ms": float(result.get("generation_time") or 0.0) * 1000.0,
|
||||
"save_conditioning_ms": save_conditioning_ms,
|
||||
"e2e_latency_ms": (time.perf_counter() - started) * 1000.0,
|
||||
}
|
||||
trim_frames = 1 if uses_continuation else 0
|
||||
mode = "DFD continuation" if conditioned else "T2W bootstrap"
|
||||
print(f"[GPU {self.gpu_id}] Cosmos {mode} segment {segment_idx}: "
|
||||
f"{len(frames)} frames, gen={generation_ms:.0f}ms, "
|
||||
f"save_conditioning={save_conditioning_ms:.0f}ms, "
|
||||
f"e2e={timings['e2e_latency_ms']:.0f}ms")
|
||||
return StepResult(
|
||||
frames=frames,
|
||||
audio=self._silent_audio(len(frames), fps),
|
||||
audio_sample_rate=_SILENT_AUDIO_SAMPLE_RATE,
|
||||
timings=timings,
|
||||
head_trim_frames=trim_frames,
|
||||
head_trim_audio_frames=trim_frames,
|
||||
)
|
||||
|
||||
def warmup(self, prompt: str) -> dict[str, float]:
|
||||
"""Exercise both T2W bootstrap and retained-frame DFD request shapes."""
|
||||
warmup_prompt = (prompt or "").strip()
|
||||
if not warmup_prompt:
|
||||
raise RuntimeError("Startup warmup prompt must be non-empty.")
|
||||
print(f"[GPU {self.gpu_id}] Cosmos startup warmup starting "
|
||||
"(synthetic segments: T2W bootstrap, DFD continuation)")
|
||||
started = time.perf_counter()
|
||||
bootstrap_result = self.generate_step(warmup_prompt, 1, None, True)
|
||||
continuation_result = self.generate_step(warmup_prompt, 2, None, False)
|
||||
total_ms = (time.perf_counter() - started) * 1000.0
|
||||
self.clear_conditioning()
|
||||
bootstrap_ms = float(bootstrap_result.timings.get("e2e_latency_ms", 0.0))
|
||||
continuation_ms = float(continuation_result.timings.get("e2e_latency_ms", 0.0))
|
||||
print(f"[GPU {self.gpu_id}] Cosmos startup warmup complete: "
|
||||
f"bootstrap={bootstrap_ms:.0f}ms, continuation={continuation_ms:.0f}ms, total={total_ms:.0f}ms")
|
||||
return {
|
||||
"warmup_bootstrap_ms": bootstrap_ms,
|
||||
"warmup_continuation_ms": continuation_ms,
|
||||
"warmup_total_ms": total_ms,
|
||||
}
|
||||
|
||||
def apply_lora_stack(self, stack: list[tuple[str, float]]) -> tuple[str | None, str | None]:
|
||||
del stack
|
||||
raise RuntimeError("Cosmos Predict2.5 DFD does not support DreamVerse runtime LoRA changes.")
|
||||
@@ -5,6 +5,8 @@ from __future__ import annotations
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
from dreamverse.generation_inputs import GenerationInputs
|
||||
|
||||
|
||||
@dataclass
|
||||
class StepResult:
|
||||
@@ -36,6 +38,7 @@ class GenerationBackend(Protocol):
|
||||
segment_idx: int,
|
||||
image_path: str | None,
|
||||
reset_conditioning: bool,
|
||||
generation_inputs: GenerationInputs | None = None,
|
||||
) -> StepResult:
|
||||
...
|
||||
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
"""GPU-independent validation for generation modes and ordered asset handles."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from PIL import Image, UnidentifiedImageError
|
||||
|
||||
from dreamverse.assets import asset_store
|
||||
|
||||
GENERATION_MODES = ("t2va", "fl2va", "ref2va")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GenerationAsset:
|
||||
asset_id: str
|
||||
kind: str
|
||||
path: str
|
||||
role: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GenerationInputs:
|
||||
mode: str | None = None
|
||||
assets: tuple[GenerationAsset, ...] = ()
|
||||
|
||||
@property
|
||||
def first_frame_path(self) -> str | None:
|
||||
return next((asset.path for asset in self.assets if asset.role == "first_frame"), None)
|
||||
|
||||
@property
|
||||
def last_frame_path(self) -> str | None:
|
||||
return next((asset.path for asset in self.assets if asset.role == "last_frame"), None)
|
||||
|
||||
@property
|
||||
def references(self) -> tuple[GenerationAsset, ...]:
|
||||
return tuple(asset for asset in self.assets if asset.role == "reference")
|
||||
|
||||
|
||||
def supported_generation_modes(model_id: str) -> tuple[str, ...]:
|
||||
return GENERATION_MODES if model_id in ("full-h3", "mock") else ("t2va", )
|
||||
|
||||
|
||||
def resolve_generation_inputs(payload: dict, model_id: str) -> GenerationInputs:
|
||||
mode = payload.get("generation_mode")
|
||||
raw_assets = payload.get("conditioning_assets", [])
|
||||
if mode is None and "generation_mode" not in payload:
|
||||
if raw_assets:
|
||||
raise ValueError("Select a generation mode before attaching conditioning assets.")
|
||||
return GenerationInputs()
|
||||
if not isinstance(mode, str) or mode not in GENERATION_MODES:
|
||||
raise ValueError("Unknown generation mode. Choose T2VA, FL2VA, or Ref2VA.")
|
||||
if mode not in supported_generation_modes(model_id):
|
||||
raise ValueError(f"{mode.upper()} requires the Full H3 runtime. This runtime is running {model_id}.")
|
||||
if payload.get("initial_image") is not None:
|
||||
raise ValueError("Use asset IDs for generation modes; do not combine them with the legacy initial_image field.")
|
||||
if not isinstance(raw_assets, list) or len(raw_assets) > 12:
|
||||
raise ValueError("conditioning_assets must be an ordered list with at most 12 assets.")
|
||||
if mode == "t2va" and raw_assets:
|
||||
raise ValueError("T2VA accepts text only. Remove conditioning assets or choose another mode.")
|
||||
assets: list[GenerationAsset] = []
|
||||
for item in raw_assets:
|
||||
if not isinstance(item, dict) or set(item) != {"asset_id", "role"}:
|
||||
raise ValueError("Each conditioning asset must contain only asset_id and role.")
|
||||
role = item["role"]
|
||||
if role not in ("first_frame", "last_frame", "reference"):
|
||||
raise ValueError("Asset role must be first_frame, last_frame, or reference.")
|
||||
stored = asset_store.get(item["asset_id"])
|
||||
assets.append(GenerationAsset(stored.asset_id, stored.kind, stored.path, role))
|
||||
if mode == "fl2va":
|
||||
if any(asset.kind != "image" or asset.role == "reference" for asset in assets):
|
||||
raise ValueError("FL2VA accepts only first-frame and last-frame images.")
|
||||
if sum(asset.role == "first_frame" for asset in assets) != 1:
|
||||
raise ValueError("FL2VA requires exactly one first-frame image.")
|
||||
if sum(asset.role == "last_frame" for asset in assets) > 1:
|
||||
raise ValueError("FL2VA accepts at most one last-frame image.")
|
||||
elif mode == "ref2va":
|
||||
if not assets or any(asset.role != "reference" for asset in assets):
|
||||
raise ValueError("Ref2VA requires an ordered list of reference assets, without keyframe roles.")
|
||||
if not any(asset.kind in ("image", "video") for asset in assets):
|
||||
raise ValueError("Ref2VA requires at least one image or video; audio alone is not supported.")
|
||||
for kind, limit in (("image", 9), ("video", 3), ("audio", 3)):
|
||||
if sum(asset.kind == kind for asset in assets) > limit:
|
||||
raise ValueError(f"Ref2VA accepts at most {limit} {kind} references.")
|
||||
for asset in assets:
|
||||
if asset.kind == "image":
|
||||
try:
|
||||
with Image.open(asset.path) as image:
|
||||
if image.width > 4 * image.height or image.height > 4 * image.width:
|
||||
raise ValueError(
|
||||
"Ref2VA image aspect ratios must be between 1:4 and 4:1. Crop this image first.")
|
||||
except (UnidentifiedImageError, OSError, Image.DecompressionBombError) as exc:
|
||||
raise ValueError("A selected reference image could not be decoded. Upload it again.") from exc
|
||||
return GenerationInputs(mode, tuple(assets))
|
||||
|
||||
|
||||
def pin_generation_inputs(inputs: GenerationInputs) -> None:
|
||||
asset_store.pin([asset.asset_id for asset in inputs.assets])
|
||||
|
||||
|
||||
def release_generation_inputs(inputs: GenerationInputs) -> None:
|
||||
asset_store.release([asset.asset_id for asset in inputs.assets])
|
||||
@@ -4,6 +4,7 @@ from __future__ import annotations
|
||||
|
||||
from dreamverse.config import MODEL_CONFIG
|
||||
from dreamverse.generation_contracts import GenerationBackend, StepResult
|
||||
from dreamverse.generation_inputs import GenerationInputs
|
||||
|
||||
|
||||
def _create_generation_backend(backend_name: str, gpu_id: int) -> GenerationBackend:
|
||||
@@ -16,6 +17,10 @@ def _create_generation_backend(backend_name: str, gpu_id: int) -> GenerationBack
|
||||
from dreamverse.minimax_h3_generation import MiniMaxH3GenerationBackend
|
||||
|
||||
return MiniMaxH3GenerationBackend(gpu_id)
|
||||
if backend_name == "cosmos25_dfd":
|
||||
from dreamverse.cosmos25_dfd_generation import Cosmos25DFDGenerationBackend
|
||||
|
||||
return Cosmos25DFDGenerationBackend(gpu_id)
|
||||
raise ValueError(f"Unsupported DreamVerse generation backend: {backend_name!r}")
|
||||
|
||||
|
||||
@@ -80,6 +85,7 @@ class VideoGenerationWorker:
|
||||
segment_idx: int,
|
||||
image_path: str | None,
|
||||
reset_conditioning: bool,
|
||||
generation_inputs: GenerationInputs | None = None,
|
||||
) -> StepResult:
|
||||
"""Generate one segment through the selected model backend."""
|
||||
return self._require_backend().generate_step(
|
||||
@@ -87,6 +93,7 @@ class VideoGenerationWorker:
|
||||
segment_idx,
|
||||
image_path,
|
||||
reset_conditioning,
|
||||
generation_inputs=generation_inputs,
|
||||
)
|
||||
|
||||
def warmup(self, prompt: str) -> dict[str, float]:
|
||||
|
||||
@@ -29,6 +29,7 @@ from dreamverse.av_streaming import (
|
||||
generate_stream_id,
|
||||
stream_fmp4,
|
||||
)
|
||||
from dreamverse.generation_inputs import GenerationInputs, pin_generation_inputs, release_generation_inputs
|
||||
from dreamverse.worker_ipc import (
|
||||
CommandPayload,
|
||||
InitAck,
|
||||
@@ -189,6 +190,7 @@ def gpu_worker_process(
|
||||
segment_idx,
|
||||
image_path=payload.image_path,
|
||||
reset_conditioning=payload.reset_conditioning,
|
||||
generation_inputs=payload.generation_inputs,
|
||||
)
|
||||
head_trim_frames = step_result.head_trim_frames
|
||||
head_trim_audio_frames = step_result.head_trim_audio_frames
|
||||
@@ -432,6 +434,7 @@ class GPUSlot:
|
||||
self.connected_users: set[str] = set()
|
||||
self._pending_futures: dict[str, asyncio.Future] = {}
|
||||
self._stream_queues: dict[str, asyncio.Queue] = {}
|
||||
self._step_asset_inputs: dict[str, GenerationInputs] = {}
|
||||
self._response_reader_task: asyncio.Task | None = None
|
||||
self._active: bool = False
|
||||
self._reader_lock: asyncio.Lock | None = None
|
||||
@@ -663,6 +666,9 @@ class GPUSlot:
|
||||
if isinstance(event, (StepComplete, WarmupComplete)):
|
||||
event.timings["ipc_get_done_ns"] = time.time_ns()
|
||||
|
||||
if isinstance(event, (StepComplete, WorkerError)) and event.user_id is not None:
|
||||
self._release_step_assets(event.user_id)
|
||||
|
||||
user_id = event.user_id
|
||||
if user_id and user_id in self._pending_futures:
|
||||
future = self._pending_futures.pop(user_id)
|
||||
@@ -753,6 +759,7 @@ class GPUSlot:
|
||||
segment_idx: int = 1,
|
||||
image_path: str | None = None,
|
||||
reset_conditioning: bool = False,
|
||||
generation_inputs: GenerationInputs | None = None,
|
||||
) -> dict[str, float]:
|
||||
"""Execute a generation step for a specific user.
|
||||
|
||||
@@ -766,9 +773,19 @@ class GPUSlot:
|
||||
segment_idx=segment_idx,
|
||||
image_path=image_path,
|
||||
reset_conditioning=bool(reset_conditioning),
|
||||
generation_inputs=generation_inputs,
|
||||
)
|
||||
if generation_inputs is not None:
|
||||
if user_id in self._step_asset_inputs:
|
||||
raise RuntimeError("The previous generation is still using this project's assets.")
|
||||
pin_generation_inputs(generation_inputs)
|
||||
self._step_asset_inputs[user_id] = generation_inputs
|
||||
# Pins intentionally survive a waiter timeout/cancellation: the GPU
|
||||
# command keeps running. The response reader releases them when the
|
||||
# worker actually completes (even if that response is now unmatched).
|
||||
response = await self._send_command_tagged(Command(CommandType.USER_STEP, payload=payload, user_id=user_id),
|
||||
timeout=1800.0)
|
||||
self._release_step_assets(user_id)
|
||||
match response:
|
||||
case StepComplete(timings=timings):
|
||||
return timings
|
||||
@@ -778,6 +795,11 @@ class GPUSlot:
|
||||
raise RuntimeError(f"Unexpected step response for {user_id[:8]}: "
|
||||
f"{type(response).__name__}")
|
||||
|
||||
def _release_step_assets(self, user_id: str) -> None:
|
||||
inputs = self._step_asset_inputs.pop(user_id, None)
|
||||
if inputs is not None:
|
||||
release_generation_inputs(inputs)
|
||||
|
||||
async def apply_lora_stack(
|
||||
self,
|
||||
stack: list[tuple[str, float]],
|
||||
@@ -800,7 +822,9 @@ class GPUSlot:
|
||||
async def leave_user(self, user_id: str) -> None:
|
||||
"""Remove a user from this GPU."""
|
||||
try:
|
||||
await self._send_command_tagged(Command(CommandType.USER_LEAVE, user_id=user_id), timeout=30.0)
|
||||
response = await self._send_command_tagged(Command(CommandType.USER_LEAVE, user_id=user_id), timeout=30.0)
|
||||
if isinstance(response, LeaveAck):
|
||||
self._release_step_assets(user_id)
|
||||
except Exception as e:
|
||||
print(f"[GPU {self.gpu_id}] Leave user error: {e}")
|
||||
finally:
|
||||
@@ -837,6 +861,10 @@ class GPUSlot:
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if self.process is None or not self.process.is_alive():
|
||||
for user_id in list(self._step_asset_inputs):
|
||||
self._release_step_assets(user_id)
|
||||
|
||||
for q in (self.command_queue, self.response_queue):
|
||||
if q is not None:
|
||||
try:
|
||||
|
||||
@@ -33,6 +33,7 @@ from dreamverse.config import (
|
||||
_resolve_lora_spec,
|
||||
)
|
||||
from dreamverse.generation_contracts import StepResult
|
||||
from dreamverse.generation_inputs import GenerationInputs
|
||||
|
||||
# Multi-frame decoded continuation defaults from
|
||||
# examples/inference/basic/basic_ltx2_distilled_video_continuation.py.
|
||||
@@ -290,7 +291,13 @@ class LTX2GenerationBackend:
|
||||
dynamic=False,
|
||||
),
|
||||
use_fsdp_inference=False,
|
||||
quantization=QuantizationConfig(transformer_quant="NVFP4"),
|
||||
# The bundled LTX2 model enables a refinement LoRA during the
|
||||
# first request. NVFP4 otherwise purges the dense weights that
|
||||
# FastVideo's LoRA merge path requires.
|
||||
quantization=QuantizationConfig(
|
||||
transformer_quant="NVFP4",
|
||||
transformer_retain_original_weights=True,
|
||||
),
|
||||
),
|
||||
pipeline=PipelineSelection(
|
||||
components=components,
|
||||
@@ -454,8 +461,11 @@ class LTX2GenerationBackend:
|
||||
segment_idx: int,
|
||||
image_path: str | None,
|
||||
reset_conditioning: bool,
|
||||
generation_inputs: GenerationInputs | None = None,
|
||||
) -> StepResult:
|
||||
"""Execute one generation step; snapshot state for the next segment."""
|
||||
if generation_inputs is not None and (generation_inputs.mode not in (None, "t2va") or generation_inputs.assets):
|
||||
raise ValueError("LTX supports text generation only through the generation mode API.")
|
||||
timings: dict = {}
|
||||
|
||||
prompt = self._inject_style_trigger(prompt)
|
||||
|
||||
@@ -15,6 +15,7 @@ from dreamverse.gpu_pool import GPUPool, get_available_gpus
|
||||
from dreamverse.session_logger import SessionEventLogger
|
||||
|
||||
from dreamverse.config import (
|
||||
ACTIVE_MODEL_ID,
|
||||
AVAILABLE_LORAS,
|
||||
DEVTOOLS_ENABLED,
|
||||
FRONTEND_STATIC_DIR_CANDIDATES,
|
||||
@@ -34,6 +35,8 @@ from dreamverse.routes.presets import (
|
||||
curated_presets_router,
|
||||
)
|
||||
from dreamverse.session.controller import SessionController
|
||||
from dreamverse.generation_inputs import supported_generation_modes
|
||||
from dreamverse.routes.assets import router as asset_router
|
||||
|
||||
|
||||
class _HeartbeatAccessLogFilter(logging.Filter):
|
||||
@@ -92,10 +95,16 @@ app.add_middleware(
|
||||
app.include_router(build_health_router(lambda: runtime.gpu_pool))
|
||||
app.include_router(internal_monitor_router)
|
||||
app.include_router(prompt_config_router)
|
||||
app.include_router(asset_router)
|
||||
if DEVTOOLS_ENABLED:
|
||||
app.include_router(curated_presets_router)
|
||||
|
||||
|
||||
@app.get("/generation-capabilities")
|
||||
async def generation_capabilities() -> dict:
|
||||
return {"model_id": ACTIVE_MODEL_ID, "modes": supported_generation_modes(ACTIVE_MODEL_ID), "mock": False}
|
||||
|
||||
|
||||
@app.websocket("/ws")
|
||||
async def websocket_endpoint(websocket: WebSocket):
|
||||
controller = SessionController(
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""FastH3 model lifecycle and first-frame continuation for DreamVerse."""
|
||||
"""Full/Preview H3 lifecycle, conditioning and per-project pipeline selection."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -12,6 +12,7 @@ import torch
|
||||
|
||||
from dreamverse.config import DREAMVERSE_SP_SIZE
|
||||
from dreamverse.generation_contracts import StepResult
|
||||
from dreamverse.generation_inputs import GenerationInputs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from PIL.Image import Image
|
||||
@@ -26,13 +27,14 @@ def _required_config_str(model_config: dict, field_name: str) -> str:
|
||||
|
||||
|
||||
class MiniMaxH3GenerationBackend:
|
||||
"""Run the VSA data-free FastH3 adapter and retain one continuation frame."""
|
||||
"""Own one H3 pipeline at a time and retain base-pipeline continuation."""
|
||||
|
||||
def __init__(self, gpu_id: int):
|
||||
self.gpu_id = gpu_id
|
||||
self.generator: Any | None = None
|
||||
self.model_config: dict = {}
|
||||
self.continuation_image: Image | None = None
|
||||
self.pipeline_mode = "base"
|
||||
|
||||
def _gpu_mem(self) -> str:
|
||||
allocated_gib = torch.cuda.memory_allocated() / 1024**3
|
||||
@@ -51,32 +53,37 @@ class MiniMaxH3GenerationBackend:
|
||||
os.environ.pop("FASTVIDEO_INFERENCE_TORCH_COMPILE", None)
|
||||
|
||||
def initialize(self, model_config: dict | None = None) -> None:
|
||||
"""Download the fixed Preview adapter and load the FastH3 generator.
|
||||
|
||||
The model profile owns the base checkpoint, adapter file, attention
|
||||
backend, and generation geometry. The backend translates that profile
|
||||
into FastVideo's typed generator configuration.
|
||||
"""
|
||||
"""Load the profile's base pipeline; Ref2VA is loaded on first use."""
|
||||
if model_config is not None:
|
||||
self.model_config = dict(model_config)
|
||||
if not self.model_config:
|
||||
raise ValueError("FastH3 initialization requires a model configuration.")
|
||||
self._load_pipeline("base")
|
||||
|
||||
def _load_pipeline(self, pipeline_mode: str) -> None:
|
||||
"""Unload the old executor before loading a base or reference transformer.
|
||||
|
||||
GPU worker commands are serialized, so a project boundary never swaps
|
||||
weights while another request is using them. Keeping one executor also
|
||||
avoids simultaneously retaining two large H3 transformers in VRAM. A
|
||||
failed load leaves no executor behind so the next step retries it.
|
||||
"""
|
||||
full_checkpoint = bool(self.model_config.get("full_checkpoint", False))
|
||||
if pipeline_mode == "ref2va" and not full_checkpoint:
|
||||
raise ValueError("Ref2VA requires the full-h3 model profile.")
|
||||
if self.generator is not None:
|
||||
self.generator.shutdown()
|
||||
previous_generator = self.generator
|
||||
self.generator = None
|
||||
previous_generator.shutdown()
|
||||
del previous_generator
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
self.clear_conditioning()
|
||||
model_path = _required_config_str(self.model_config, "model_path")
|
||||
adapter_repo = _required_config_str(self.model_config, "adapter_repo")
|
||||
adapter_filename = _required_config_str(self.model_config, "adapter_filename")
|
||||
attention_backend = _required_config_str(self.model_config, "attention_backend")
|
||||
self._configure_environment(attention_backend)
|
||||
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
CompileConfig,
|
||||
@@ -88,10 +95,20 @@ class MiniMaxH3GenerationBackend:
|
||||
PipelineSelection,
|
||||
)
|
||||
|
||||
adapter_path = hf_hub_download(repo_id=adapter_repo, filename=adapter_filename)
|
||||
components = ComponentConfig()
|
||||
if not full_checkpoint:
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
adapter_repo = _required_config_str(self.model_config, "adapter_repo")
|
||||
adapter_filename = _required_config_str(self.model_config, "adapter_filename")
|
||||
components.lora_path = hf_hub_download(repo_id=adapter_repo, filename=adapter_filename)
|
||||
components.lora_strength = 1.0
|
||||
print(f"[GPU {self.gpu_id}] FastH3 adapter: {adapter_repo}/{adapter_filename}")
|
||||
if pipeline_mode == "ref2va":
|
||||
components.override_pipeline_cls_name = "MiniMaxH3Ref2VAModularPipeline"
|
||||
experimental = {
|
||||
"attention_backend": attention_backend,
|
||||
"inference_torch_compile": attention_backend == "FLASH_ATTN",
|
||||
"inference_torch_compile": not full_checkpoint and attention_backend == "FLASH_ATTN",
|
||||
"vae_parallel_decode": True,
|
||||
"vae_parallel_decode_strategy": "gather",
|
||||
}
|
||||
@@ -103,7 +120,8 @@ class MiniMaxH3GenerationBackend:
|
||||
generator_config = GeneratorConfig(
|
||||
model_path=model_path,
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(lora_path=adapter_path, lora_strength=1.0),
|
||||
workload_type="i2v" if pipeline_mode == "ref2va" else None,
|
||||
components=components,
|
||||
experimental=experimental,
|
||||
),
|
||||
engine=EngineConfig(
|
||||
@@ -115,17 +133,23 @@ class MiniMaxH3GenerationBackend:
|
||||
text_encoder=True,
|
||||
image_encoder=True,
|
||||
vae=True,
|
||||
pin_cpu_memory=True,
|
||||
pin_cpu_memory=not full_checkpoint,
|
||||
),
|
||||
compile=CompileConfig(enabled=False, vae_enabled=True),
|
||||
use_fsdp_inference=False,
|
||||
use_fsdp_inference=full_checkpoint and DREAMVERSE_SP_SIZE > 1,
|
||||
),
|
||||
)
|
||||
|
||||
print(f"[GPU {self.gpu_id}] Loading FastH3 model: {model_path}")
|
||||
print(f"[GPU {self.gpu_id}] FastH3 adapter: {adapter_repo}/{adapter_filename}")
|
||||
print(f"[GPU {self.gpu_id}] Loading H3 model: {model_path} ({pipeline_mode})")
|
||||
print(f"[GPU {self.gpu_id}] Before model load: {self._gpu_mem()}")
|
||||
self.generator = VideoGenerator.from_config(generator_config)
|
||||
try:
|
||||
self.generator = VideoGenerator.from_config(generator_config)
|
||||
except Exception:
|
||||
# The old executor is already gone; leaving no executor behind lets
|
||||
# the next step retry this load instead of stranding the GPU slot.
|
||||
self.generator = None
|
||||
raise
|
||||
self.pipeline_mode = pipeline_mode
|
||||
print(f"[GPU {self.gpu_id}] FastH3 loaded: {self._gpu_mem()} (warmup pending)")
|
||||
|
||||
def shutdown(self) -> None:
|
||||
@@ -166,14 +190,28 @@ class MiniMaxH3GenerationBackend:
|
||||
return self._load_rgb_image(image_path), False
|
||||
return None, False
|
||||
|
||||
def _build_request(self, prompt: str, conditioning_image: Image | None):
|
||||
def _build_request(
|
||||
self,
|
||||
prompt: str,
|
||||
conditioning_image: Image | None,
|
||||
last_image: Image | None = None,
|
||||
generation_inputs: GenerationInputs | None = None,
|
||||
):
|
||||
"""Build the typed FastVideo request owned by the FastH3 profile."""
|
||||
from fastvideo.api import GenerationRequest, InputConfig, OutputConfig, SamplingConfig
|
||||
|
||||
references = None
|
||||
if generation_inputs is not None and generation_inputs.mode == "ref2va":
|
||||
from fastvideo.api import MiniMaxH3Reference
|
||||
|
||||
references = [
|
||||
MiniMaxH3Reference(source=str(asset.path), media_type=asset.kind)
|
||||
for asset in generation_inputs.references
|
||||
]
|
||||
return GenerationRequest(
|
||||
prompt=prompt,
|
||||
negative_prompt="",
|
||||
inputs=InputConfig(pil_image=conditioning_image),
|
||||
inputs=InputConfig(pil_image=conditioning_image, last_image=last_image, references=references),
|
||||
sampling=SamplingConfig(
|
||||
height=int(self.model_config["height"]),
|
||||
width=int(self.model_config["width"]),
|
||||
@@ -200,6 +238,7 @@ class MiniMaxH3GenerationBackend:
|
||||
segment_idx: int,
|
||||
image_path: str | None,
|
||||
reset_conditioning: bool,
|
||||
generation_inputs: GenerationInputs | None = None,
|
||||
) -> StepResult:
|
||||
"""Generate one synchronized FastH3 segment and retain its last frame.
|
||||
|
||||
@@ -207,20 +246,46 @@ class MiniMaxH3GenerationBackend:
|
||||
conditioned frame and its matching audio duration are trimmed before
|
||||
streaming so adjacent segments do not duplicate media.
|
||||
"""
|
||||
if self.generator is None:
|
||||
raise RuntimeError("FastH3 generator is not initialized.")
|
||||
conditioning_image, uses_continuation = self._select_conditioning_image(
|
||||
segment_idx,
|
||||
image_path,
|
||||
reset_conditioning,
|
||||
)
|
||||
request = self._build_request(prompt, conditioning_image)
|
||||
mode = generation_inputs.mode if generation_inputs is not None else None
|
||||
if mode not in (None, "t2va", "fl2va", "ref2va"):
|
||||
raise ValueError(f"Unsupported H3 generation mode: {mode!r}.")
|
||||
if mode in ("fl2va", "ref2va") and not self.model_config.get("full_checkpoint", False):
|
||||
raise ValueError(f"{mode.upper()} requires the full-h3 model profile.")
|
||||
pipeline_mode = "ref2va" if mode == "ref2va" else "base"
|
||||
if self.generator is None or self.pipeline_mode != pipeline_mode:
|
||||
# A failed switch leaves no executor behind; reload here so the
|
||||
# slot recovers on the next step instead of staying broken.
|
||||
if segment_idx > 1 and not reset_conditioning:
|
||||
raise ValueError("Generation mode cannot change in the middle of a project.")
|
||||
self._load_pipeline(pipeline_mode)
|
||||
|
||||
conditioning_image = None
|
||||
last_image = None
|
||||
uses_continuation = False
|
||||
if mode == "ref2va":
|
||||
# The reference pipeline rejects first/last-frame inputs. Preserve
|
||||
# all original references for every clip and do not trim overlap.
|
||||
self.clear_conditioning()
|
||||
else:
|
||||
if mode == "fl2va" and generation_inputs is not None:
|
||||
image_path = generation_inputs.first_frame_path
|
||||
conditioning_image, uses_continuation = self._select_conditioning_image(
|
||||
segment_idx,
|
||||
image_path,
|
||||
reset_conditioning,
|
||||
)
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
if (mode == "fl2va" and segment_idx == 1 and generation_inputs is not None
|
||||
and generation_inputs.last_frame_path):
|
||||
last_image = self._load_rgb_image(generation_inputs.last_frame_path)
|
||||
request = self._build_request(prompt, conditioning_image, last_image, generation_inputs)
|
||||
result = self.generator.generate(request)
|
||||
finally:
|
||||
if conditioning_image is not None:
|
||||
conditioning_image.close()
|
||||
if last_image is not None:
|
||||
last_image.close()
|
||||
torch.cuda.synchronize()
|
||||
generation_ms = (time.perf_counter() - started) * 1000.0
|
||||
|
||||
@@ -235,7 +300,8 @@ class MiniMaxH3GenerationBackend:
|
||||
raise RuntimeError("FastH3 returned audio without an audio sample rate.")
|
||||
|
||||
save_started = time.perf_counter()
|
||||
self._save_continuation_frame(frames)
|
||||
if mode != "ref2va":
|
||||
self._save_continuation_frame(frames)
|
||||
save_conditioning_ms = (time.perf_counter() - save_started) * 1000.0
|
||||
timings = {
|
||||
"generation_ms": generation_ms,
|
||||
|
||||
@@ -32,6 +32,14 @@ from fastapi.staticfiles import StaticFiles
|
||||
from dreamverse._deps import require_dreamverse_runtime_deps
|
||||
from dreamverse.config import FRONTEND_STATIC_DIR_CANDIDATES, GENERATION_SEGMENT_CAP
|
||||
from dreamverse.session_init_image import cleanup_session_init_image, persist_session_init_image
|
||||
from dreamverse.generation_inputs import (
|
||||
GenerationInputs,
|
||||
pin_generation_inputs,
|
||||
release_generation_inputs,
|
||||
resolve_generation_inputs,
|
||||
supported_generation_modes,
|
||||
)
|
||||
from dreamverse.routes.assets import router as asset_router
|
||||
|
||||
LATENCY_MS = 200
|
||||
SESSION_TIMEOUT_SECONDS = 300
|
||||
@@ -170,6 +178,12 @@ app.add_middleware(
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
app.include_router(asset_router)
|
||||
|
||||
|
||||
@app.get("/generation-capabilities")
|
||||
async def generation_capabilities():
|
||||
return {"model_id": "mock", "modes": supported_generation_modes("mock"), "mock": True}
|
||||
|
||||
|
||||
@app.get("/healthz")
|
||||
@@ -290,6 +304,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
send_lock = asyncio.Lock()
|
||||
stop_event = asyncio.Event()
|
||||
session_init_image = None
|
||||
generation_inputs = GenerationInputs()
|
||||
|
||||
async def ws_send_json(payload: dict) -> None:
|
||||
async with send_lock:
|
||||
@@ -347,10 +362,13 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
generation_paused = bool(initial_rollout_prompt and not single_clip_mode and len(curated_prompts) == 0)
|
||||
|
||||
try:
|
||||
generation_inputs = resolve_generation_inputs(init_data, "mock")
|
||||
pin_generation_inputs(generation_inputs)
|
||||
session_init_image = persist_session_init_image(init_data.get("initial_image"))
|
||||
except ValueError as exc:
|
||||
await ws_send_json({
|
||||
"type": "error",
|
||||
"error_code": "invalid_generation_input",
|
||||
"message": str(exc),
|
||||
})
|
||||
await websocket.close(code=1003, reason="Invalid initial image")
|
||||
@@ -362,6 +380,8 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
"type": "gpu_assigned",
|
||||
"gpu_id": 0,
|
||||
"session_timeout": SESSION_TIMEOUT_SECONDS,
|
||||
"generation_mode": generation_inputs.mode,
|
||||
"mock": True,
|
||||
})
|
||||
|
||||
raw_prompt_queue: asyncio.Queue[PromptSubmission] = asyncio.Queue()
|
||||
@@ -486,6 +506,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
})
|
||||
|
||||
async def apply_project_init_payload(payload: dict[str, object], ) -> bool:
|
||||
nonlocal generation_inputs
|
||||
nonlocal preset_id
|
||||
nonlocal preset_label
|
||||
nonlocal initial_rollout_prompt
|
||||
@@ -519,10 +540,19 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
]
|
||||
|
||||
try:
|
||||
replace_session_image(payload.get("initial_image"))
|
||||
next_inputs = resolve_generation_inputs(payload, "mock")
|
||||
pin_generation_inputs(next_inputs)
|
||||
try:
|
||||
replace_session_image(payload.get("initial_image"))
|
||||
except ValueError:
|
||||
release_generation_inputs(next_inputs)
|
||||
raise
|
||||
release_generation_inputs(generation_inputs)
|
||||
generation_inputs = next_inputs
|
||||
except ValueError as exc:
|
||||
await ws_send_json({
|
||||
"type": "error",
|
||||
"error_code": "invalid_generation_input",
|
||||
"message": str(exc),
|
||||
})
|
||||
return False
|
||||
@@ -568,6 +598,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
return drained
|
||||
|
||||
async def enter_project_idle() -> None:
|
||||
nonlocal generation_inputs
|
||||
nonlocal seed_prompt_memory
|
||||
nonlocal curated_prompts
|
||||
nonlocal curated_idx
|
||||
@@ -587,6 +618,8 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
|
||||
dropped_raw = drain_queue_nowait(raw_prompt_queue)
|
||||
dropped_ready = drain_queue_nowait(ready_prompt_queue)
|
||||
release_generation_inputs(generation_inputs)
|
||||
generation_inputs = GenerationInputs()
|
||||
seed_prompt_memory = []
|
||||
curated_prompts = []
|
||||
curated_idx = 0
|
||||
@@ -767,10 +800,17 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
continue
|
||||
|
||||
try:
|
||||
if generation_inputs.mode is not None and data.get("initial_image") is not None:
|
||||
raise ValueError("Choose conditioning assets when starting a project; legacy initial_image "
|
||||
"cannot replace generation mode inputs.")
|
||||
if "generation_mode" in data or "conditioning_assets" in data:
|
||||
raise ValueError(
|
||||
"simple_generate cannot change the mode; use project_init_v1.")
|
||||
replace_session_image(data.get("initial_image"))
|
||||
except ValueError as exc:
|
||||
await ws_send_json({
|
||||
"type": "error",
|
||||
"error_code": "invalid_generation_input",
|
||||
"message": str(exc),
|
||||
})
|
||||
continue
|
||||
@@ -1182,6 +1222,7 @@ async def websocket_endpoint(websocket: WebSocket):
|
||||
finally:
|
||||
stop_event.set()
|
||||
cleanup_session_init_image(session_init_image)
|
||||
release_generation_inputs(generation_inputs)
|
||||
|
||||
|
||||
for static_dir in FRONTEND_STATIC_DIR_CANDIDATES:
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
"""Raw, bounded media uploads keep large binary data out of websocket messages."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from urllib.parse import unquote
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request, Response
|
||||
from fastapi.responses import FileResponse
|
||||
from starlette.concurrency import run_in_threadpool
|
||||
|
||||
from dreamverse.assets import IMAGE_LIMIT, MEDIA_LIMIT, MIME_TYPES, asset_store
|
||||
|
||||
router = APIRouter()
|
||||
_upload_lock = asyncio.Lock()
|
||||
|
||||
|
||||
@router.post("/assets", status_code=201)
|
||||
async def upload_asset(request: Request) -> dict:
|
||||
mime_type = request.headers.get("content-type", "").split(";", 1)[0].lower()
|
||||
if mime_type not in MIME_TYPES:
|
||||
raise HTTPException(415, "Unsupported asset type. Select a supported image, video, or audio file.")
|
||||
limit = IMAGE_LIMIT if MIME_TYPES[mime_type][0] == "image" else MEDIA_LIMIT
|
||||
try:
|
||||
if int(request.headers.get("content-length", "0")) > limit:
|
||||
raise HTTPException(413, f"Asset exceeds the {limit // (1024 * 1024)} MB upload limit.")
|
||||
except ValueError as exc:
|
||||
raise HTTPException(400, "Invalid Content-Length.") from exc
|
||||
async with _upload_lock:
|
||||
try:
|
||||
path = asset_store.staging_path(mime_type)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(400, str(exc)) from exc
|
||||
try:
|
||||
size = 0
|
||||
with path.open("xb") as handle:
|
||||
async for chunk in request.stream():
|
||||
size += len(chunk)
|
||||
if size > limit:
|
||||
raise HTTPException(413, f"Asset exceeds the {limit // (1024 * 1024)} MB upload limit.")
|
||||
await run_in_threadpool(handle.write, chunk)
|
||||
asset = await run_in_threadpool(asset_store.add, path,
|
||||
unquote(request.headers.get("x-asset-name", "Untitled asset")), mime_type)
|
||||
return asset.public()
|
||||
except ValueError as exc:
|
||||
path.unlink(missing_ok=True)
|
||||
raise HTTPException(400, str(exc)) from exc
|
||||
except BaseException:
|
||||
path.unlink(missing_ok=True)
|
||||
raise
|
||||
|
||||
|
||||
@router.api_route("/assets/{asset_id}", methods=["GET", "HEAD"])
|
||||
async def get_asset(asset_id: str) -> FileResponse:
|
||||
try:
|
||||
asset = asset_store.get(asset_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(404, str(exc)) from exc
|
||||
return FileResponse(asset.path, media_type=asset.mime_type, headers={"X-Content-Type-Options": "nosniff"})
|
||||
|
||||
|
||||
@router.delete("/assets/{asset_id}", status_code=204)
|
||||
async def delete_asset(asset_id: str) -> Response:
|
||||
try:
|
||||
asset_store.get(asset_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(404, str(exc)) from exc
|
||||
try:
|
||||
asset_store.delete(asset_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(409, str(exc)) from exc
|
||||
return Response(status_code=204)
|
||||
@@ -26,6 +26,12 @@ from typing import TYPE_CHECKING
|
||||
|
||||
from fastapi import WebSocket, WebSocketDisconnect
|
||||
from dreamverse.gpu_pool import GPUSlot
|
||||
from dreamverse.generation_inputs import (
|
||||
GenerationInputs,
|
||||
pin_generation_inputs,
|
||||
release_generation_inputs,
|
||||
resolve_generation_inputs,
|
||||
)
|
||||
from dreamverse.session_init_image import cleanup_session_init_image, persist_session_init_image
|
||||
from dreamverse.worker_ipc import MediaChunk, MediaComplete, MediaInit
|
||||
|
||||
@@ -156,6 +162,7 @@ class SessionController:
|
||||
prompt_worker_task: asyncio.Task | None = None
|
||||
rewrite_seed_prompts_task: asyncio.Task | None = None
|
||||
session_init_image = None
|
||||
generation_inputs: GenerationInputs | None = None
|
||||
|
||||
async def session_timeout():
|
||||
"""Close the session after timeout."""
|
||||
@@ -191,6 +198,18 @@ class SessionController:
|
||||
init_data = {}
|
||||
|
||||
init_type = init_data.get("type")
|
||||
try:
|
||||
next_generation_inputs = resolve_generation_inputs(init_data, ACTIVE_MODEL_ID)
|
||||
pin_generation_inputs(next_generation_inputs)
|
||||
generation_inputs = next_generation_inputs
|
||||
except ValueError as exc:
|
||||
await ws_send_json({
|
||||
"type": "error",
|
||||
"error_code": "invalid_generation_input",
|
||||
"message": str(exc),
|
||||
})
|
||||
await websocket.close(code=1008, reason="Invalid generation inputs")
|
||||
return
|
||||
preset_id = init_data.get("preset_id")
|
||||
preset_label = str(init_data.get("preset_label") or "").strip()
|
||||
initial_rollout_prompt = str(init_data.get("initial_rollout_prompt") or "").strip()
|
||||
@@ -271,6 +290,7 @@ class SessionController:
|
||||
"type": "gpu_assigned",
|
||||
"gpu_id": gpu_id,
|
||||
"session_timeout": SESSION_TIMEOUT_SECONDS,
|
||||
"generation_mode": generation_inputs.mode,
|
||||
})
|
||||
await log_event(
|
||||
"gpu_assigned",
|
||||
@@ -361,10 +381,16 @@ class SessionController:
|
||||
return
|
||||
|
||||
try:
|
||||
if generation_inputs.mode is not None and payload.get("initial_image") is not None:
|
||||
raise ValueError("Choose conditioning assets when starting a project; legacy initial_image "
|
||||
"cannot replace generation mode inputs.")
|
||||
if "generation_mode" in payload or "conditioning_assets" in payload:
|
||||
raise ValueError("simple_generate cannot change the mode; use project_init_v1.")
|
||||
replace_session_init_image(payload.get("initial_image"))
|
||||
except ValueError as exc:
|
||||
await ws_send_json({
|
||||
"type": "error",
|
||||
"error_code": "invalid_generation_input",
|
||||
"message": str(exc),
|
||||
})
|
||||
return
|
||||
@@ -417,6 +443,7 @@ class SessionController:
|
||||
})
|
||||
|
||||
async def apply_project_init_payload(payload: dict[str, object]) -> bool:
|
||||
nonlocal generation_inputs
|
||||
nonlocal preset_id
|
||||
nonlocal preset_label
|
||||
nonlocal initial_rollout_prompt
|
||||
@@ -496,15 +523,26 @@ class SessionController:
|
||||
})
|
||||
return False
|
||||
|
||||
next_generation_inputs = None
|
||||
next_inputs_pinned = False
|
||||
try:
|
||||
next_generation_inputs = resolve_generation_inputs(payload, ACTIVE_MODEL_ID)
|
||||
pin_generation_inputs(next_generation_inputs)
|
||||
next_inputs_pinned = True
|
||||
replace_session_init_image(payload.get("initial_image"))
|
||||
except ValueError as exc:
|
||||
if next_inputs_pinned:
|
||||
release_generation_inputs(next_generation_inputs)
|
||||
await ws_send_json({
|
||||
"type": "error",
|
||||
"error_code": "invalid_generation_input",
|
||||
"message": str(exc),
|
||||
})
|
||||
return False
|
||||
|
||||
release_generation_inputs(generation_inputs)
|
||||
generation_inputs = next_generation_inputs
|
||||
|
||||
initial_rollout_prompt = next_initial_rollout_prompt
|
||||
enhancement_enabled = next_enhancement_enabled
|
||||
auto_extension_enabled = next_auto_extension_enabled
|
||||
@@ -1171,6 +1209,7 @@ class SessionController:
|
||||
return drained
|
||||
|
||||
async def enter_project_idle() -> None:
|
||||
nonlocal generation_inputs
|
||||
nonlocal curated_prompts
|
||||
nonlocal seed_prompt_memory
|
||||
nonlocal curated_idx
|
||||
@@ -1221,6 +1260,9 @@ class SessionController:
|
||||
project_active = False
|
||||
pending_project_end = False
|
||||
|
||||
release_generation_inputs(generation_inputs)
|
||||
generation_inputs = GenerationInputs()
|
||||
|
||||
if project_stream_started:
|
||||
project_stream_started = False
|
||||
await ws_send_json({"type": "ltx2_stream_complete"})
|
||||
@@ -1627,6 +1669,7 @@ class SessionController:
|
||||
segment_idx=segment_idx,
|
||||
image_path=step_image_path,
|
||||
reset_conditioning=step_reset_conditioning,
|
||||
generation_inputs=generation_inputs,
|
||||
))
|
||||
segment_generation_active = True
|
||||
try:
|
||||
@@ -1681,10 +1724,10 @@ class SessionController:
|
||||
print(f"[GPU {gpu_id}] Unknown AV event: "
|
||||
f"{type(event).__name__}")
|
||||
|
||||
if not step_task.done():
|
||||
step_task.cancel()
|
||||
else:
|
||||
timings = await step_task
|
||||
# A GPU command cannot be cancelled by cancelling its
|
||||
# asyncio waiter. Await completion before releasing pinned
|
||||
# asset files or making this GPU available to a new user.
|
||||
timings = await step_task
|
||||
finally:
|
||||
segment_generation_active = False
|
||||
if not step_task.done():
|
||||
@@ -1808,3 +1851,5 @@ class SessionController:
|
||||
await self.gpu_pool.release(client_id)
|
||||
finally:
|
||||
cleanup_session_init_image(session_init_image)
|
||||
if generation_inputs is not None:
|
||||
release_generation_inputs(generation_inputs)
|
||||
|
||||
@@ -139,12 +139,48 @@ def test_config_enables_prompt_safety_when_requested(monkeypatch):
|
||||
|
||||
def test_config_uses_five_minute_session_timeout(monkeypatch):
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
monkeypatch.delenv("DREAMVERSE_MODEL_ID", raising=False)
|
||||
monkeypatch.delenv("DREAMVERSE_SESSION_TIMEOUT_SECONDS", raising=False)
|
||||
monkeypatch.delenv("FASTVIDEO_SESSION_TIMEOUT_SECONDS", raising=False)
|
||||
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.SESSION_TIMEOUT_SECONDS == 300
|
||||
|
||||
|
||||
def test_config_uses_thirty_minute_cosmos25_session_timeout(monkeypatch):
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
monkeypatch.setenv("DREAMVERSE_MODEL_ID", "cosmos25-dfd")
|
||||
monkeypatch.delenv("DREAMVERSE_SESSION_TIMEOUT_SECONDS", raising=False)
|
||||
monkeypatch.delenv("FASTVIDEO_SESSION_TIMEOUT_SECONDS", raising=False)
|
||||
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.SESSION_TIMEOUT_SECONDS == 1800
|
||||
|
||||
|
||||
def test_config_allows_session_timeout_override(monkeypatch):
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
monkeypatch.setenv("DREAMVERSE_MODEL_ID", "cosmos25-dfd")
|
||||
monkeypatch.delenv("DREAMVERSE_SESSION_TIMEOUT_SECONDS", raising=False)
|
||||
monkeypatch.setenv("FASTVIDEO_SESSION_TIMEOUT_SECONDS", "900")
|
||||
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.SESSION_TIMEOUT_SECONDS == 900
|
||||
|
||||
|
||||
def test_config_prefers_dreamverse_session_timeout_over_alias(monkeypatch):
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
monkeypatch.setenv("DREAMVERSE_MODEL_ID", "cosmos25-dfd")
|
||||
monkeypatch.setenv("DREAMVERSE_SESSION_TIMEOUT_SECONDS", "1200")
|
||||
monkeypatch.setenv("FASTVIDEO_SESSION_TIMEOUT_SECONDS", "900")
|
||||
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.SESSION_TIMEOUT_SECONDS == 1200
|
||||
|
||||
|
||||
def test_config_rejects_invalid_prompt_provider(monkeypatch):
|
||||
monkeypatch.setenv("FASTVIDEO_PROMPT_PROVIDER", "unsupported")
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
@@ -186,3 +222,57 @@ def test_config_uses_fasth3_sequence_parallel_default(monkeypatch):
|
||||
assert module.ACTIVE_MODEL_ID == "fast-h3"
|
||||
assert module.MODEL_CONFIG["generation_backend"] == "minimax_h3"
|
||||
assert module.DREAMVERSE_SP_SIZE == 4
|
||||
|
||||
|
||||
def test_full_h3_profile_has_no_preview_adapter_and_longer_session(monkeypatch):
|
||||
monkeypatch.setenv("DREAMVERSE_MODEL_ID", "full-h3")
|
||||
monkeypatch.delenv("DREAMVERSE_SP_SIZE", raising=False)
|
||||
monkeypatch.delenv("DREAMVERSE_SESSION_TIMEOUT_SECONDS", raising=False)
|
||||
monkeypatch.delenv("FASTVIDEO_SESSION_TIMEOUT_SECONDS", raising=False)
|
||||
module = _load_config_module()
|
||||
assert module.MODEL_CONFIG["full_checkpoint"] is True
|
||||
assert "adapter_repo" not in module.MODEL_CONFIG
|
||||
assert module.MODEL_CONFIG["num_inference_steps"] == 50
|
||||
assert module.DREAMVERSE_SP_SIZE == 4
|
||||
assert module.SESSION_TIMEOUT_SECONDS == 7200
|
||||
|
||||
|
||||
def test_config_registers_cosmos25_dfd_profile(monkeypatch):
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.MODEL_REGISTRY["cosmos25-dfd"] == {
|
||||
"name": "Cosmos Predict2.5 DFD",
|
||||
"generation_backend": "cosmos25_dfd",
|
||||
"default_sp_size": 1,
|
||||
"model_path": "FastVideo/Cosmos-Predict2.5-2B-Distilled-TrigFlow",
|
||||
"continuation_model_path": "FastVideo/Cosmos-Predict2.5-2B-DFD",
|
||||
"attention_backend": "TORCH_SDPA",
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"bootstrap_num_frames": 77,
|
||||
"continuation_num_frames": 81,
|
||||
"fps": 24,
|
||||
"num_inference_steps": 4,
|
||||
"seed": 42,
|
||||
"session_timeout_seconds": 1800,
|
||||
}
|
||||
|
||||
|
||||
def test_config_selects_cosmos25_package_roles(monkeypatch, tmp_path):
|
||||
_set_required_prompt_keys(monkeypatch)
|
||||
bootstrap_path = tmp_path / "cosmos25-t2w"
|
||||
continuation_path = tmp_path / "cosmos25-dfd"
|
||||
monkeypatch.setenv("DREAMVERSE_MODEL_ID", "cosmos25-dfd")
|
||||
monkeypatch.setenv("DREAMVERSE_MODEL_PATH", str(bootstrap_path))
|
||||
monkeypatch.setenv("DREAMVERSE_COSMOS25_DFD_MODEL_PATH", str(continuation_path))
|
||||
monkeypatch.delenv("DREAMVERSE_SP_SIZE", raising=False)
|
||||
|
||||
module = _load_config_module()
|
||||
|
||||
assert module.ACTIVE_MODEL_ID == "cosmos25-dfd"
|
||||
assert module.MODEL_CONFIG["generation_backend"] == "cosmos25_dfd"
|
||||
assert module.MODEL_CONFIG["model_path"] == str(bootstrap_path)
|
||||
assert module.MODEL_CONFIG["continuation_model_path"] == str(continuation_path)
|
||||
assert module.DREAMVERSE_SP_SIZE == 1
|
||||
|
||||
@@ -0,0 +1,219 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from dreamverse.cosmos25_dfd_generation import Cosmos25DFDGenerationBackend
|
||||
from dreamverse.generation_inputs import GenerationInputs
|
||||
|
||||
COSMOS_CONFIG = {
|
||||
"name": "Cosmos Predict2.5 DFD",
|
||||
"generation_backend": "cosmos25_dfd",
|
||||
"default_sp_size": 1,
|
||||
"model_path": "/models/cosmos25-t2w",
|
||||
"continuation_model_path": "/models/cosmos25-dfd",
|
||||
"attention_backend": "TORCH_SDPA",
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"bootstrap_num_frames": 77,
|
||||
"continuation_num_frames": 81,
|
||||
"fps": 24,
|
||||
"num_inference_steps": 4,
|
||||
"seed": 42,
|
||||
}
|
||||
|
||||
|
||||
class _RecordingGenerator:
|
||||
def __init__(self, pixel_value: int = 20) -> None:
|
||||
self.pixel_value = pixel_value
|
||||
self.calls: list[dict] = []
|
||||
self.shutdown_calls = 0
|
||||
|
||||
def generate_video(self, prompt, sampling_param):
|
||||
condition = sampling_param.pil_image
|
||||
self.calls.append({
|
||||
"prompt": prompt,
|
||||
"sampling": sampling_param,
|
||||
"conditioning_pixels": None if condition is None else np.asarray(condition).copy(),
|
||||
})
|
||||
frames = [
|
||||
np.full((2, 3, 3), self.pixel_value, dtype=np.uint8)
|
||||
for _ in range(sampling_param.num_frames)
|
||||
]
|
||||
frames[-1] = np.full((2, 3, 3), self.pixel_value + 1, dtype=np.uint8)
|
||||
return {
|
||||
"frames": frames,
|
||||
"generation_time": 0.25,
|
||||
}
|
||||
|
||||
def shutdown(self):
|
||||
self.shutdown_calls += 1
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def backend(monkeypatch) -> Cosmos25DFDGenerationBackend:
|
||||
instance = Cosmos25DFDGenerationBackend(gpu_id=0)
|
||||
instance.model_config = dict(COSMOS_CONFIG)
|
||||
instance.bootstrap_generator = _RecordingGenerator(pixel_value=20)
|
||||
instance.continuation_generator = _RecordingGenerator(pixel_value=40)
|
||||
monkeypatch.setattr("dreamverse.cosmos25_dfd_generation.torch.cuda.synchronize", lambda: None)
|
||||
|
||||
def fake_sampling_param(*, conditioned):
|
||||
return SimpleNamespace(
|
||||
negative_prompt="",
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=704,
|
||||
width=1280,
|
||||
num_frames=81 if conditioned else 77,
|
||||
fps=24,
|
||||
num_inference_steps=4,
|
||||
guidance_scale=1.0,
|
||||
seed=42,
|
||||
num_cond_frames=1 if conditioned else 0,
|
||||
pil_image=None,
|
||||
)
|
||||
|
||||
monkeypatch.setattr(instance, "_sampling_param", fake_sampling_param)
|
||||
return instance
|
||||
|
||||
|
||||
def test_initialize_loads_both_package_roles(monkeypatch):
|
||||
loaded_paths = []
|
||||
generators = [_RecordingGenerator(), _RecordingGenerator()]
|
||||
backend = Cosmos25DFDGenerationBackend(gpu_id=0)
|
||||
|
||||
def fake_load(model_path):
|
||||
loaded_paths.append(model_path)
|
||||
return generators[len(loaded_paths) - 1]
|
||||
|
||||
monkeypatch.setattr(backend, "_load_generator", fake_load)
|
||||
monkeypatch.setattr(backend, "_gpu_mem", lambda: "alloc=0.00GiB, reserved=0.00GiB")
|
||||
monkeypatch.setattr("dreamverse.cosmos25_dfd_generation.gc.collect", lambda: 0)
|
||||
monkeypatch.setattr("dreamverse.cosmos25_dfd_generation.torch.cuda.is_available", lambda: False)
|
||||
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "test-attention")
|
||||
monkeypatch.setenv("FASTVIDEO_INFERENCE_TORCH_COMPILE", "1")
|
||||
|
||||
backend.initialize(COSMOS_CONFIG)
|
||||
|
||||
assert loaded_paths == [
|
||||
"/models/cosmos25-t2w",
|
||||
"/models/cosmos25-dfd",
|
||||
]
|
||||
assert backend.bootstrap_generator is generators[0]
|
||||
assert backend.continuation_generator is generators[1]
|
||||
assert backend.model_config == COSMOS_CONFIG
|
||||
assert os.environ["FASTVIDEO_ATTENTION_BACKEND"] == "TORCH_SDPA"
|
||||
assert "FASTVIDEO_INFERENCE_TORCH_COMPILE" not in os.environ
|
||||
|
||||
|
||||
def test_unconditioned_start_uses_t2w_and_retains_terminal_frame(backend):
|
||||
result = backend.generate_step("first prompt", 1, None, True)
|
||||
|
||||
assert len(backend.bootstrap_generator.calls) == 1
|
||||
assert backend.continuation_generator.calls == []
|
||||
sampling = backend.bootstrap_generator.calls[0]["sampling"]
|
||||
assert sampling.height == 704
|
||||
assert sampling.width == 1280
|
||||
assert sampling.num_frames == 77
|
||||
assert sampling.fps == 24
|
||||
assert sampling.num_inference_steps == 4
|
||||
assert sampling.guidance_scale == 1.0
|
||||
assert sampling.seed == 42
|
||||
assert sampling.num_cond_frames == 0
|
||||
assert sampling.pil_image is None
|
||||
assert result.head_trim_frames == 0
|
||||
assert result.head_trim_audio_frames == 0
|
||||
assert result.audio_sample_rate == 24_000
|
||||
assert result.audio.shape == (77_000, )
|
||||
assert result.audio.count_nonzero() == 0
|
||||
assert np.asarray(backend.continuation_image).tolist() == np.full((2, 3, 3), 21).tolist()
|
||||
|
||||
|
||||
def test_retained_frame_uses_dfd_and_trims_repeated_boundary(backend):
|
||||
backend.generate_step("first prompt", 1, None, True)
|
||||
result = backend.generate_step("pivot right", 2, None, False)
|
||||
|
||||
assert len(backend.continuation_generator.calls) == 1
|
||||
call = backend.continuation_generator.calls[0]
|
||||
sampling = call["sampling"]
|
||||
assert sampling.num_frames == 81
|
||||
assert sampling.num_cond_frames == 1
|
||||
assert call["conditioning_pixels"].tolist() == np.full((2, 3, 3), 21).tolist()
|
||||
assert result.head_trim_frames == 1
|
||||
assert result.head_trim_audio_frames == 1
|
||||
assert result.audio.shape == (81_000, )
|
||||
assert np.asarray(backend.continuation_image).tolist() == np.full((2, 3, 3), 41).tolist()
|
||||
|
||||
|
||||
def test_initial_image_uses_dfd_without_stream_trim(backend, tmp_path: Path):
|
||||
from PIL import Image
|
||||
|
||||
image_path = tmp_path / "initial.png"
|
||||
Image.fromarray(np.full((2, 3, 3), 7, dtype=np.uint8)).save(image_path)
|
||||
|
||||
result = backend.generate_step("animate", 1, str(image_path), True)
|
||||
|
||||
assert backend.bootstrap_generator.calls == []
|
||||
call = backend.continuation_generator.calls[0]
|
||||
assert call["conditioning_pixels"].tolist() == np.full((2, 3, 3), 7).tolist()
|
||||
assert result.head_trim_frames == 0
|
||||
assert result.head_trim_audio_frames == 0
|
||||
|
||||
|
||||
def test_generation_mode_api_accepts_text_only_and_rejects_conditioning_modes(backend):
|
||||
result = backend.generate_step("first prompt", 1, None, True, generation_inputs=GenerationInputs(mode="t2va"))
|
||||
|
||||
assert len(backend.bootstrap_generator.calls) == 1
|
||||
assert result.head_trim_frames == 0
|
||||
|
||||
with pytest.raises(ValueError, match="text generation only"):
|
||||
backend.generate_step("pivot right", 2, None, False, generation_inputs=GenerationInputs(mode="fl2va"))
|
||||
assert backend.continuation_generator.calls == []
|
||||
|
||||
|
||||
def test_missing_later_continuation_fails_before_generation(backend):
|
||||
with pytest.raises(RuntimeError, match="requires a retained continuation frame"):
|
||||
backend.generate_step("later prompt", 2, None, False)
|
||||
|
||||
assert backend.bootstrap_generator.calls == []
|
||||
assert backend.continuation_generator.calls == []
|
||||
|
||||
|
||||
def test_reset_later_segment_uses_fresh_t2w_bootstrap(backend):
|
||||
backend.generate_step("first prompt", 1, None, True)
|
||||
|
||||
result = backend.generate_step("new scene", 2, None, True)
|
||||
|
||||
assert len(backend.bootstrap_generator.calls) == 2
|
||||
assert backend.continuation_generator.calls == []
|
||||
assert result.head_trim_frames == 0
|
||||
|
||||
|
||||
def test_warmup_exercises_bootstrap_and_dfd_paths(backend):
|
||||
timings = backend.warmup("warmup prompt")
|
||||
|
||||
assert len(backend.bootstrap_generator.calls) == 1
|
||||
assert len(backend.continuation_generator.calls) == 1
|
||||
assert backend.continuation_image is None
|
||||
assert "warmup_bootstrap_ms" in timings
|
||||
assert "warmup_continuation_ms" in timings
|
||||
assert "warmup_total_ms" in timings
|
||||
|
||||
|
||||
def test_shutdown_releases_both_generators_and_conditioning(backend):
|
||||
bootstrap = backend.bootstrap_generator
|
||||
continuation = backend.continuation_generator
|
||||
backend.generate_step("first prompt", 1, None, True)
|
||||
|
||||
backend.shutdown()
|
||||
|
||||
assert bootstrap.shutdown_calls == 1
|
||||
assert continuation.shutdown_calls == 1
|
||||
assert backend.bootstrap_generator is None
|
||||
assert backend.continuation_generator is None
|
||||
assert backend.continuation_image is None
|
||||
@@ -0,0 +1,248 @@
|
||||
"""Contract regressions independent of CUDA and model weights."""
|
||||
|
||||
import io
|
||||
import asyncio
|
||||
import shutil
|
||||
import subprocess
|
||||
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
from PIL import Image
|
||||
|
||||
from dreamverse import assets, generation_inputs
|
||||
from dreamverse.generation_inputs import resolve_generation_inputs
|
||||
from dreamverse.routes import assets as asset_routes
|
||||
from dreamverse.tests.test_mock_server import _FakeWebSocket
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def library(monkeypatch):
|
||||
store = assets.AssetStore()
|
||||
monkeypatch.setattr(asset_routes, "asset_store", store)
|
||||
monkeypatch.setattr(generation_inputs, "asset_store", store)
|
||||
app = FastAPI()
|
||||
app.include_router(asset_routes.router)
|
||||
with TestClient(app) as client:
|
||||
yield store, client
|
||||
|
||||
|
||||
def upload_image(client, color="red"):
|
||||
image_bytes = io.BytesIO()
|
||||
Image.new("RGB", (32, 32), color).save(image_bytes, format="PNG")
|
||||
response = client.post("/assets", content=image_bytes.getvalue(),
|
||||
headers={"Content-Type": "image/png", "X-Asset-Name": "frame.png"})
|
||||
assert response.status_code == 201, response.text
|
||||
return response.json()
|
||||
|
||||
|
||||
def conditioning(asset, role):
|
||||
return {"asset_id": asset["asset_id"], "role": role}
|
||||
|
||||
|
||||
def test_assets_validate_content_and_support_head_range_and_delete(library):
|
||||
store, client = library
|
||||
asset = upload_image(client)
|
||||
assert set(asset) == {"asset_id", "kind", "name", "mime_type", "size", "url"}
|
||||
assert client.head(asset["url"]).status_code == 200
|
||||
response = client.get(asset["url"], headers={"Range": "bytes=0-7"})
|
||||
assert response.status_code == 206
|
||||
assert response.content == b"\x89PNG\r\n\x1a\n"
|
||||
assert client.post("/assets", content=b"not an image", headers={"Content-Type": "image/png"}).status_code == 400
|
||||
assert client.post("/assets", content=b"<svg/>", headers={"Content-Type": "image/svg+xml"}).status_code == 415
|
||||
assert client.post("/assets", content=b"", headers={"Content-Type": "image/png",
|
||||
"Content-Length": str(assets.IMAGE_LIMIT + 1)}).status_code == 413
|
||||
with pytest.raises(ValueError, match="Invalid asset ID"):
|
||||
store.get("../../etc/passwd")
|
||||
assert client.delete(asset["url"]).status_code == 204
|
||||
assert client.head(asset["url"]).status_code == 404
|
||||
|
||||
|
||||
def test_generation_pin_prevents_deletion_until_session_releases(library):
|
||||
_, client = library
|
||||
asset = upload_image(client)
|
||||
inputs = resolve_generation_inputs({"generation_mode": "fl2va", "conditioning_assets": [
|
||||
conditioning(asset, "first_frame")
|
||||
]}, "full-h3")
|
||||
generation_inputs.pin_generation_inputs(inputs)
|
||||
generation_inputs.pin_generation_inputs(inputs)
|
||||
assert client.delete(asset["url"]).status_code == 409
|
||||
generation_inputs.release_generation_inputs(inputs)
|
||||
assert client.delete(asset["url"]).status_code == 409
|
||||
generation_inputs.release_generation_inputs(inputs)
|
||||
assert client.delete(asset["url"]).status_code == 204
|
||||
|
||||
|
||||
def test_legacy_init_remains_compatible_but_explicit_t2va_is_text_only(library):
|
||||
_, client = library
|
||||
assert resolve_generation_inputs({"initial_image": {"old": "payload"}}, "fast-ltx2").mode is None
|
||||
assert resolve_generation_inputs({"generation_mode": "t2va"}, "fast-ltx2").mode == "t2va"
|
||||
with pytest.raises(ValueError, match="legacy initial_image"):
|
||||
resolve_generation_inputs({"generation_mode": "t2va", "initial_image": {}}, "full-h3")
|
||||
asset = upload_image(client)
|
||||
with pytest.raises(ValueError, match="text only"):
|
||||
resolve_generation_inputs({"generation_mode": "t2va", "conditioning_assets": [
|
||||
conditioning(asset, "reference")
|
||||
]}, "full-h3")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["unknown", None, 3, [], {}])
|
||||
def test_unknown_mode_fails_before_assets_are_resolved(mode):
|
||||
with pytest.raises(ValueError, match="Unknown generation mode"):
|
||||
resolve_generation_inputs({"generation_mode": mode}, "full-h3")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model_id", ["fast-h3", "fast-ltx2", "fast-ltx23"])
|
||||
def test_preview_and_ltx_cannot_advertise_full_h3_modes(model_id):
|
||||
with pytest.raises(ValueError, match="Full H3"):
|
||||
resolve_generation_inputs({"generation_mode": "ref2va"}, model_id)
|
||||
|
||||
|
||||
def test_fl2va_first_required_last_optional_and_roles_unique(library):
|
||||
_, client = library
|
||||
first = upload_image(client)
|
||||
last = upload_image(client, "blue")
|
||||
payload = {"generation_mode": "fl2va", "conditioning_assets": [conditioning(first, "first_frame")]}
|
||||
inputs = resolve_generation_inputs(payload, "full-h3")
|
||||
assert inputs.first_frame_path.endswith(".png")
|
||||
assert inputs.last_frame_path is None
|
||||
payload["conditioning_assets"].append(conditioning(last, "last_frame"))
|
||||
assert resolve_generation_inputs(payload, "full-h3").last_frame_path is not None
|
||||
payload["conditioning_assets"].append(conditioning(first, "first_frame"))
|
||||
with pytest.raises(ValueError, match="exactly one first-frame"):
|
||||
resolve_generation_inputs(payload, "full-h3")
|
||||
with pytest.raises(ValueError, match="exactly one first-frame"):
|
||||
resolve_generation_inputs({"generation_mode": "fl2va", "conditioning_assets": [
|
||||
conditioning(last, "last_frame")
|
||||
]}, "full-h3")
|
||||
|
||||
|
||||
def test_ref_order_is_preserved_and_limits_are_enforced(library):
|
||||
_, client = library
|
||||
first, second = upload_image(client), upload_image(client, "blue")
|
||||
refs = [conditioning(second, "reference"), conditioning(first, "reference")]
|
||||
payload = {"generation_mode": "ref2va", "conditioning_assets": refs}
|
||||
inputs = resolve_generation_inputs(payload, "full-h3")
|
||||
assert [asset.asset_id for asset in inputs.references] == [second["asset_id"], first["asset_id"]]
|
||||
with pytest.raises(ValueError, match="at most 9 image"):
|
||||
resolve_generation_inputs({**payload, "conditioning_assets": refs * 5}, "full-h3")
|
||||
with pytest.raises(ValueError, match="without keyframe roles"):
|
||||
resolve_generation_inputs({**payload, "conditioning_assets": [conditioning(first, "first_frame")]}, "full-h3")
|
||||
with pytest.raises(ValueError, match="at most 12"):
|
||||
resolve_generation_inputs({**payload, "conditioning_assets": refs * 7}, "full-h3")
|
||||
|
||||
|
||||
def test_ref_audio_requires_visual_reference(library, monkeypatch):
|
||||
store, _ = library
|
||||
monkeypatch.setattr(store, "get", lambda asset_id: assets.StoredAsset(asset_id, "audio", "/audio.wav", "audio",
|
||||
"audio/wav", 100))
|
||||
with pytest.raises(ValueError, match="audio alone"):
|
||||
resolve_generation_inputs({"generation_mode": "ref2va", "conditioning_assets": [
|
||||
{"asset_id": "a" * 32, "role": "reference"}
|
||||
]}, "full-h3")
|
||||
|
||||
|
||||
def test_ref_rejects_extreme_image_aspect_before_gpu(library):
|
||||
_, client = library
|
||||
content = io.BytesIO()
|
||||
Image.new("RGB", (500, 50), "blue").save(content, format="PNG")
|
||||
response = client.post("/assets", content=content.getvalue(), headers={"Content-Type": "image/png"})
|
||||
assert response.status_code == 201
|
||||
with pytest.raises(ValueError, match="aspect ratios"):
|
||||
resolve_generation_inputs({"generation_mode": "ref2va", "conditioning_assets": [
|
||||
conditioning(response.json(), "reference")
|
||||
]}, "full-h3")
|
||||
|
||||
|
||||
def test_ref_reports_undecodable_image_as_invalid_input(library, monkeypatch, tmp_path):
|
||||
store, _ = library
|
||||
broken = tmp_path / "broken.png"
|
||||
broken.write_bytes(b"not an image")
|
||||
monkeypatch.setattr(store, "get", lambda asset_id: assets.StoredAsset(asset_id, "image", str(broken), "broken.png",
|
||||
"image/png", 11))
|
||||
with pytest.raises(ValueError, match="could not be decoded"):
|
||||
resolve_generation_inputs({"generation_mode": "ref2va", "conditioning_assets": [
|
||||
{"asset_id": "a" * 32, "role": "reference"}
|
||||
]}, "full-h3")
|
||||
|
||||
|
||||
def test_audio_upload_rejects_surround_sound(library, monkeypatch):
|
||||
import json
|
||||
_, client = library
|
||||
monkeypatch.setattr(assets.shutil, "which", lambda name: "/usr/bin/ffprobe")
|
||||
info = {"format": {"format_name": "wav", "duration": "1"},
|
||||
"streams": [{"codec_type": "audio", "channels": 6}]}
|
||||
monkeypatch.setattr(assets.subprocess, "run", lambda *args, **kwargs: subprocess.CompletedProcess(
|
||||
[], 0, stdout=json.dumps(info).encode(), stderr=b""))
|
||||
response = client.post("/assets", content=b"surround wav", headers={"Content-Type": "audio/wav"})
|
||||
assert response.status_code == 400
|
||||
assert "mono or stereo" in response.json()["detail"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mime,format_name", [("audio/x-m4a", "mov,mp4,m4a,3gp,3g2,mj2"), ("audio/x-flac", "flac")])
|
||||
def test_legacy_audio_mime_aliases_are_accepted(library, monkeypatch, mime, format_name):
|
||||
"""Browsers report x- variants for the M4A and FLAC formats the docs promise."""
|
||||
import json
|
||||
_, client = library
|
||||
monkeypatch.setattr(assets.shutil, "which", lambda name: "/usr/bin/ffprobe")
|
||||
info = {"format": {"format_name": format_name, "duration": "1"},
|
||||
"streams": [{"codec_type": "audio", "channels": 2}]}
|
||||
monkeypatch.setattr(assets.subprocess, "run", lambda *args, **kwargs: subprocess.CompletedProcess(
|
||||
[], 0, stdout=json.dumps(info).encode(), stderr=b""))
|
||||
response = client.post("/assets", content=b"audio bytes", headers={"Content-Type": mime})
|
||||
assert response.status_code == 201, response.text
|
||||
assert response.json()["kind"] == "audio"
|
||||
assert response.json()["mime_type"] == mime
|
||||
|
||||
|
||||
@pytest.mark.parametrize("entries", [None, {}, "x", [{"path": "/etc/passwd", "role": "reference"}],
|
||||
[{"asset_id": "x", "role": "unknown"}]])
|
||||
def test_malformed_conditioning_is_rejected(entries):
|
||||
with pytest.raises(ValueError):
|
||||
resolve_generation_inputs({"generation_mode": "ref2va", "conditioning_assets": entries}, "full-h3")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["t2va", "fl2va", "ref2va"])
|
||||
def test_mock_streams_all_valid_modes_and_releases_assets(library, monkeypatch, mode):
|
||||
from dreamverse import mock_server
|
||||
_, client = library
|
||||
monkeypatch.setattr(mock_server, "MOCK_SEGMENT_BYTES", b"mock-fmp4")
|
||||
monkeypatch.setattr(mock_server, "LATENCY_MS", 1)
|
||||
image = upload_image(client)
|
||||
refs = [] if mode == "t2va" else [conditioning(image, "first_frame" if mode == "fl2va" else "reference")]
|
||||
ws = _FakeWebSocket([
|
||||
(0, {"type": "session_init_v2", "generation_mode": mode, "conditioning_assets": refs,
|
||||
"curated_prompts": ["A bird flies over a lake."], "single_clip_mode": True,
|
||||
"enhancement_enabled": False}),
|
||||
(0.15, {"type": "leave"}),
|
||||
])
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
assert not [event for event in ws.sent_json if event["type"] == "error"]
|
||||
assert any(event["type"] == "media_segment_complete" for event in ws.sent_json)
|
||||
assert ws.sent_bytes
|
||||
assert client.delete(image["url"]).status_code == 204
|
||||
|
||||
|
||||
def test_mock_rejects_invalid_mode_before_gpu_assignment(library):
|
||||
from dreamverse import mock_server
|
||||
ws = _FakeWebSocket([(0, {"type": "session_init_v2", "generation_mode": "fl2va"})])
|
||||
asyncio.run(mock_server.websocket_endpoint(ws))
|
||||
assert not any(event["type"] == "gpu_assigned" for event in ws.sent_json)
|
||||
errors = [event for event in ws.sent_json if event["type"] == "error"]
|
||||
assert errors[0]["error_code"] == "invalid_generation_input"
|
||||
assert "first-frame" in errors[0]["message"]
|
||||
|
||||
|
||||
@pytest.mark.skipif(not shutil.which("ffmpeg") or not shutil.which("ffprobe"), reason="ffmpeg + ffprobe required")
|
||||
@pytest.mark.parametrize("kind,mime,suffix", [("video", "video/mp4", ".mp4"), ("audio", "audio/wav", ".wav")])
|
||||
def test_actual_video_and_audio_upload_validation(library, tmp_path, kind, mime, suffix):
|
||||
_, client = library
|
||||
media_path = tmp_path / f"sample{suffix}"
|
||||
source = "testsrc2=size=64x64:rate=24" if kind == "video" else "sine=frequency=440:sample_rate=24000"
|
||||
command = [shutil.which("ffmpeg"), "-v", "error", "-f", "lavfi", "-i", source, "-t", "0.5", str(media_path)]
|
||||
subprocess.run(command, check=True, capture_output=True, timeout=30)
|
||||
response = client.post("/assets", content=media_path.read_bytes(), headers={"Content-Type": mime})
|
||||
assert response.status_code == 201, response.text
|
||||
assert response.json()["kind"] == kind
|
||||
response = client.post("/assets", content=b"#EXTM3U\nhttp://example.com/stream", headers={"Content-Type": mime})
|
||||
assert response.status_code == 400
|
||||
@@ -0,0 +1,214 @@
|
||||
"""CPU contract tests; fake executors do not validate generated-media quality."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
import pickle
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import ModuleType, SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
from PIL import Image
|
||||
|
||||
from dreamverse.config import MODEL_REGISTRY
|
||||
from dreamverse.generation_inputs import GenerationAsset, GenerationInputs
|
||||
from dreamverse.generation_worker import VideoGenerationWorker
|
||||
from dreamverse.minimax_h3_generation import MiniMaxH3GenerationBackend
|
||||
from dreamverse.worker_ipc import UserStepPayload
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def fastvideo_api(monkeypatch):
|
||||
"""Use the actual lightweight API schema with only GPU execution replaced."""
|
||||
schema_path = Path(__file__).resolve().parents[4] / "fastvideo/api/schema.py"
|
||||
spec = importlib.util.spec_from_file_location("dreamverse_test_api_schema", schema_path)
|
||||
assert spec is not None and spec.loader is not None
|
||||
schema = importlib.util.module_from_spec(spec)
|
||||
monkeypatch.setitem(sys.modules, spec.name, schema)
|
||||
spec.loader.exec_module(schema)
|
||||
package = ModuleType("fastvideo")
|
||||
package.__path__ = []
|
||||
package.VideoGenerator = SimpleNamespace(from_config=Mock())
|
||||
monkeypatch.setitem(sys.modules, "fastvideo", package)
|
||||
monkeypatch.setitem(sys.modules, "fastvideo.api", schema)
|
||||
schema.MiniMaxH3Reference = lambda **kwargs: SimpleNamespace(**kwargs)
|
||||
monkeypatch.setattr("dreamverse.minimax_h3_generation.torch.cuda.synchronize", lambda: None)
|
||||
monkeypatch.setattr("dreamverse.minimax_h3_generation.torch.cuda.empty_cache", lambda: None)
|
||||
return package.VideoGenerator.from_config
|
||||
|
||||
|
||||
class RecordingGenerator:
|
||||
def __init__(self):
|
||||
self.requests = []
|
||||
self.images = []
|
||||
self.closed = False
|
||||
|
||||
def shutdown(self):
|
||||
self.closed = True
|
||||
|
||||
def generate(self, request):
|
||||
self.requests.append(request)
|
||||
self.images.append(tuple(None if image is None else np.asarray(image).copy()
|
||||
for image in (request.inputs.pil_image, request.inputs.last_image)))
|
||||
return SimpleNamespace(
|
||||
frames=[np.full((2, 3, 3), 7, dtype=np.uint8), np.full((2, 3, 3), 29, dtype=np.uint8)],
|
||||
audio=np.zeros((2, 16), dtype=np.float32),
|
||||
audio_sample_rate=44100,
|
||||
generation_time=0.1,
|
||||
)
|
||||
|
||||
|
||||
def prepared_backend(monkeypatch):
|
||||
backend = MiniMaxH3GenerationBackend(0)
|
||||
backend.model_config = dict(MODEL_REGISTRY["full-h3"])
|
||||
backend.generator = RecordingGenerator()
|
||||
monkeypatch.setattr(backend, "_gpu_mem", lambda: "fake executor")
|
||||
return backend
|
||||
|
||||
|
||||
def test_ipc_preserves_immutable_ordered_references():
|
||||
inputs = GenerationInputs("ref2va", (
|
||||
GenerationAsset("second", "video", "/assets/second.mp4", "reference"),
|
||||
GenerationAsset("first", "image", "/assets/first.png", "reference"),
|
||||
))
|
||||
payload = UserStepPayload("follow the references", 1, None, True, inputs)
|
||||
restored = pickle.loads(pickle.dumps(payload))
|
||||
assert restored == payload
|
||||
assert [asset.asset_id for asset in restored.generation_inputs.references] == ["second", "first"]
|
||||
|
||||
|
||||
def test_worker_passes_conditioning_to_selected_backend():
|
||||
inputs = GenerationInputs("t2va")
|
||||
worker = VideoGenerationWorker(0)
|
||||
worker.backend = Mock()
|
||||
worker.generate_step("prompt", 1, None, True, inputs)
|
||||
worker.backend.generate_step.assert_called_once_with("prompt", 1, None, True, generation_inputs=inputs)
|
||||
|
||||
|
||||
def test_full_h3_uses_full_weights_without_preview_lora(monkeypatch, fastvideo_api):
|
||||
backend = prepared_backend(monkeypatch)
|
||||
old_generator = backend.generator
|
||||
fastvideo_api.return_value = RecordingGenerator()
|
||||
monkeypatch.setattr("dreamverse.minimax_h3_generation.DREAMVERSE_SP_SIZE", 4)
|
||||
backend.initialize(MODEL_REGISTRY["full-h3"])
|
||||
config = fastvideo_api.call_args.args[0]
|
||||
assert old_generator.closed
|
||||
assert config.pipeline.components.lora_path is None
|
||||
assert config.pipeline.components.override_pipeline_cls_name is None
|
||||
assert config.engine.use_fsdp_inference
|
||||
assert config.engine.num_gpus == 4
|
||||
assert not config.pipeline.experimental["inference_torch_compile"]
|
||||
|
||||
|
||||
def test_fl2va_maps_endpoints_only_on_initial_segment(monkeypatch, fastvideo_api, tmp_path):
|
||||
first = tmp_path / "first.png"
|
||||
last = tmp_path / "last.png"
|
||||
Image.new("RGB", (3, 2), (10, 20, 30)).save(first)
|
||||
Image.new("RGB", (3, 2), (40, 50, 60)).save(last)
|
||||
inputs = GenerationInputs("fl2va", (
|
||||
GenerationAsset("first", "image", str(first), "first_frame"),
|
||||
GenerationAsset("last", "image", str(last), "last_frame"),
|
||||
))
|
||||
backend = prepared_backend(monkeypatch)
|
||||
first_result = backend.generate_step("first", 1, None, True, inputs)
|
||||
later_result = backend.generate_step("later", 2, None, False, inputs)
|
||||
assert backend.generator.images[0][0][0, 0].tolist() == [10, 20, 30]
|
||||
assert backend.generator.images[0][1][0, 0].tolist() == [40, 50, 60]
|
||||
assert backend.generator.images[1][0][0, 0].tolist() == [29, 29, 29]
|
||||
assert backend.generator.images[1][1] is None
|
||||
assert first_result.head_trim_frames == 0
|
||||
assert later_result.head_trim_frames == 1
|
||||
assert backend.generator.requests[0].sampling.num_inference_steps == 50
|
||||
|
||||
|
||||
def test_ref2va_switches_pipeline_and_preserves_reference_order(monkeypatch, fastvideo_api):
|
||||
inputs = GenerationInputs("ref2va", (
|
||||
GenerationAsset("video", "video", "/assets/reference.mp4", "reference"),
|
||||
GenerationAsset("audio", "audio", "/assets/reference.wav", "reference"),
|
||||
GenerationAsset("image", "image", "/assets/reference.png", "reference"),
|
||||
))
|
||||
backend = prepared_backend(monkeypatch)
|
||||
base_generator = backend.generator
|
||||
reference_generator = RecordingGenerator()
|
||||
|
||||
def load(config):
|
||||
assert base_generator.closed, "Old executor must release memory before loading reference weights"
|
||||
assert config.pipeline.components.override_pipeline_cls_name == "MiniMaxH3Ref2VAModularPipeline"
|
||||
assert config.pipeline.workload_type == "i2v"
|
||||
assert config.pipeline.components.lora_path is None
|
||||
return reference_generator
|
||||
|
||||
fastvideo_api.side_effect = load
|
||||
backend.generate_step("first", 1, None, True, inputs)
|
||||
result = backend.generate_step("second", 2, None, False, inputs)
|
||||
assert fastvideo_api.call_count == 1
|
||||
for request in reference_generator.requests:
|
||||
assert [(reference.media_type, reference.source) for reference in request.inputs.references] == [
|
||||
("video", "/assets/reference.mp4"), ("audio", "/assets/reference.wav"), ("image", "/assets/reference.png")
|
||||
]
|
||||
assert request.inputs.pil_image is None
|
||||
assert request.inputs.last_image is None
|
||||
assert result.head_trim_frames == result.head_trim_audio_frames == 0
|
||||
assert backend.continuation_image is None
|
||||
|
||||
fastvideo_api.side_effect = None
|
||||
fastvideo_api.return_value = RecordingGenerator()
|
||||
backend.generate_step("new project", 1, None, True, GenerationInputs("t2va"))
|
||||
assert reference_generator.closed
|
||||
config = fastvideo_api.call_args.args[0]
|
||||
assert config.pipeline.components.override_pipeline_cls_name is None
|
||||
assert backend.pipeline_mode == "base"
|
||||
|
||||
|
||||
def test_ref2va_pipeline_switch_failure_drops_unloaded_executor(monkeypatch, fastvideo_api):
|
||||
backend = prepared_backend(monkeypatch)
|
||||
old_generator = backend.generator
|
||||
fastvideo_api.side_effect = RuntimeError("checkpoint unavailable")
|
||||
with pytest.raises(RuntimeError, match="checkpoint unavailable"):
|
||||
backend.generate_step("prompt", 1, None, True, GenerationInputs("ref2va"))
|
||||
assert old_generator.closed
|
||||
assert backend.generator is None
|
||||
|
||||
|
||||
def test_failed_pipeline_switch_reloads_on_the_next_step(monkeypatch, fastvideo_api):
|
||||
"""A failed base<->ref2va switch must not strand the slot for later steps."""
|
||||
backend = prepared_backend(monkeypatch)
|
||||
fastvideo_api.side_effect = RuntimeError("checkpoint unavailable")
|
||||
with pytest.raises(RuntimeError, match="checkpoint unavailable"):
|
||||
backend.generate_step("prompt", 1, None, True, GenerationInputs("ref2va"))
|
||||
|
||||
fastvideo_api.side_effect = None
|
||||
fastvideo_api.return_value = RecordingGenerator()
|
||||
backend.generate_step("retry", 1, None, True, GenerationInputs("ref2va"))
|
||||
|
||||
assert fastvideo_api.call_count == 2
|
||||
assert backend.pipeline_mode == "ref2va"
|
||||
assert backend.generator is not None
|
||||
|
||||
|
||||
def test_mode_cannot_switch_mid_project(monkeypatch, fastvideo_api):
|
||||
backend = prepared_backend(monkeypatch)
|
||||
with pytest.raises(ValueError, match="middle of a project"):
|
||||
backend.generate_step("prompt", 2, None, False, GenerationInputs("ref2va"))
|
||||
fastvideo_api.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("mode", ["fl2va", "ref2va"])
|
||||
def test_preview_rejects_unsupported_generation_modes(monkeypatch, fastvideo_api, mode):
|
||||
backend = prepared_backend(monkeypatch)
|
||||
backend.model_config = dict(MODEL_REGISTRY["fast-h3"])
|
||||
with pytest.raises(ValueError, match="full-h3"):
|
||||
backend.generate_step("prompt", 1, None, True, GenerationInputs(mode))
|
||||
assert backend.generator.requests == []
|
||||
|
||||
|
||||
def test_legacy_h3_continuation_is_preserved(monkeypatch, fastvideo_api):
|
||||
backend = prepared_backend(monkeypatch)
|
||||
backend.generate_step("first", 1, None, False)
|
||||
result = backend.generate_step("second", 2, None, False)
|
||||
assert backend.generator.requests[0].inputs.pil_image is None
|
||||
assert backend.generator.images[1][0][0, 0].tolist() == [29, 29, 29]
|
||||
assert result.head_trim_frames == 1
|
||||
@@ -0,0 +1,278 @@
|
||||
"""Session-mode validation and IPC handoff without a GPU worker process."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import importlib.util
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import ModuleType, SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from dreamverse.generation_inputs import GenerationAsset, GenerationInputs
|
||||
from dreamverse.worker_ipc import MediaChunk, MediaComplete, MediaInit
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def controller_module(monkeypatch):
|
||||
gpu_pool = ModuleType("dreamverse.gpu_pool")
|
||||
gpu_pool.GPUSlot = object
|
||||
monkeypatch.setitem(sys.modules, "dreamverse.gpu_pool", gpu_pool)
|
||||
path = Path(__file__).resolve().parents[1] / "session/controller.py"
|
||||
spec = importlib.util.spec_from_file_location("dreamverse_test_session_controller", path)
|
||||
assert spec is not None and spec.loader is not None
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(module)
|
||||
monkeypatch.setattr(module, "ACTIVE_MODEL_ID", "full-h3")
|
||||
monkeypatch.setattr(module, "pin_generation_inputs", Mock())
|
||||
monkeypatch.setattr(module, "release_generation_inputs", Mock())
|
||||
return module
|
||||
|
||||
|
||||
class Socket:
|
||||
def __init__(self):
|
||||
self.incoming = asyncio.Queue()
|
||||
self.outgoing = asyncio.Queue()
|
||||
self.messages = []
|
||||
self.closed = False
|
||||
|
||||
async def accept(self):
|
||||
pass
|
||||
|
||||
async def receive_json(self):
|
||||
return await self.incoming.get()
|
||||
|
||||
async def send_json(self, payload):
|
||||
self.messages.append(payload)
|
||||
await self.outgoing.put(payload)
|
||||
|
||||
async def send_bytes(self, payload):
|
||||
pass
|
||||
|
||||
async def close(self, **kwargs):
|
||||
self.closed = True
|
||||
|
||||
async def wait_for(self, kind):
|
||||
while True:
|
||||
payload = await asyncio.wait_for(self.outgoing.get(), 3)
|
||||
if payload["type"] == kind:
|
||||
return payload
|
||||
|
||||
|
||||
class Slot:
|
||||
def __init__(self):
|
||||
self.shared_stream_buffer = None
|
||||
self.queue = asyncio.Queue()
|
||||
self.calls = []
|
||||
|
||||
async def join_user(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
def register_stream_queue(self, client_id):
|
||||
return self.queue
|
||||
|
||||
def unregister_stream_queue(self, client_id):
|
||||
pass
|
||||
|
||||
async def user_step(self, client_id, **kwargs):
|
||||
self.calls.append(kwargs)
|
||||
segment_idx = kwargs["segment_idx"]
|
||||
await self.queue.put(MediaInit(client_id, segment_idx, "test", "video/mp4", False))
|
||||
await self.queue.put(MediaChunk(client_id, segment_idx, "test", chunk=b"test"))
|
||||
await self.queue.put(MediaComplete(client_id, segment_idx, "test", 1))
|
||||
return {"e2e_latency_ms": 1.0}
|
||||
|
||||
|
||||
class Pool:
|
||||
def __init__(self):
|
||||
self.slot = Slot()
|
||||
self.acquire_count = 0
|
||||
|
||||
def get_status(self):
|
||||
return {"queue_size": 0, "available_gpus": 1, "total_gpus": 1}
|
||||
|
||||
async def acquire(self, *args):
|
||||
self.acquire_count += 1
|
||||
return 0, self.slot
|
||||
|
||||
async def release(self, *args):
|
||||
pass
|
||||
|
||||
|
||||
def start_controller(module, socket, pool):
|
||||
enhancer = SimpleNamespace(
|
||||
resolve_rewrite_model=lambda value: "test-model",
|
||||
resolve_rewrite_system_prompt=lambda value: "test-system",
|
||||
resolve_rewrite_temperature=lambda value: 1.0,
|
||||
)
|
||||
controller = module.SessionController(socket, pool, enhancer, None, None)
|
||||
return asyncio.create_task(controller.run())
|
||||
|
||||
|
||||
def test_invalid_initial_mode_does_not_acquire_gpu(controller_module):
|
||||
async def scenario():
|
||||
socket, pool = Socket(), Pool()
|
||||
await socket.incoming.put({"type": "session_init_v2", "generation_mode": "unknown"})
|
||||
await asyncio.wait_for(start_controller(controller_module, socket, pool), 3)
|
||||
error = next(message for message in socket.messages if message["type"] == "error")
|
||||
assert error["error_code"] == "invalid_generation_input"
|
||||
assert pool.acquire_count == 0
|
||||
assert socket.closed
|
||||
controller_module.pin_generation_inputs.assert_not_called()
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_new_project_replaces_conditioning_and_passes_it_to_gpu(controller_module, monkeypatch):
|
||||
first = GenerationInputs("t2va")
|
||||
second = GenerationInputs("fl2va", (GenerationAsset("first", "image", "/assets/first.png", "first_frame"),))
|
||||
monkeypatch.setattr(controller_module, "resolve_generation_inputs", Mock(side_effect=[first, second]))
|
||||
|
||||
async def scenario():
|
||||
socket, pool = Socket(), Pool()
|
||||
await socket.incoming.put({
|
||||
"type": "session_init_v2", "generation_mode": "t2va", "curated_prompts": ["first prompt"],
|
||||
"enhancement_enabled": False,
|
||||
})
|
||||
task = start_controller(controller_module, socket, pool)
|
||||
try:
|
||||
await socket.wait_for("media_segment_complete")
|
||||
await socket.incoming.put({"type": "end_project_keep_session"})
|
||||
await socket.wait_for("project_idle")
|
||||
assert first in [call.args[0] for call in controller_module.release_generation_inputs.call_args_list]
|
||||
await socket.incoming.put({
|
||||
"type": "project_init_v1", "generation_mode": "fl2va", "curated_prompts": ["second prompt"],
|
||||
"enhancement_enabled": False,
|
||||
})
|
||||
await socket.wait_for("media_segment_complete")
|
||||
assert [call["generation_inputs"] for call in pool.slot.calls] == [first, second]
|
||||
assert pool.slot.calls[1]["segment_idx"] == 1
|
||||
assert pool.slot.calls[1]["reset_conditioning"]
|
||||
await socket.incoming.put({"type": "leave"})
|
||||
await asyncio.wait_for(task, 3)
|
||||
finally:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
assert [call.args[0] for call in controller_module.pin_generation_inputs.call_args_list] == [first, second]
|
||||
assert second in [call.args[0] for call in controller_module.release_generation_inputs.call_args_list]
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("injection", [
|
||||
{"initial_image": {"data_url": "not allowed"}},
|
||||
{"generation_mode": "ref2va"},
|
||||
{"conditioning_assets": []},
|
||||
])
|
||||
def test_simple_generate_cannot_replace_locked_inputs(controller_module, injection):
|
||||
async def scenario():
|
||||
socket, pool = Socket(), Pool()
|
||||
await socket.incoming.put({
|
||||
"type": "session_init_v2", "generation_mode": "t2va", "single_clip_mode": True,
|
||||
"enhancement_enabled": False,
|
||||
})
|
||||
task = start_controller(controller_module, socket, pool)
|
||||
try:
|
||||
await socket.wait_for("gpu_assigned")
|
||||
await socket.incoming.put({"type": "simple_generate", "prompt": "prompt", **injection})
|
||||
error = await socket.wait_for("error")
|
||||
assert error["error_code"] == "invalid_generation_input"
|
||||
assert "project" in error["message"]
|
||||
assert pool.slot.calls == []
|
||||
await socket.incoming.put({"type": "leave"})
|
||||
await asyncio.wait_for(task, 3)
|
||||
finally:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
def test_disconnect_waits_for_worker_before_releasing_assets(controller_module, monkeypatch):
|
||||
async def scenario():
|
||||
socket, pool = Socket(), Pool()
|
||||
worker_started = asyncio.Event()
|
||||
worker_finished = asyncio.Event()
|
||||
proceed = asyncio.Event()
|
||||
|
||||
async def slow_step(client_id, **kwargs):
|
||||
worker_started.set()
|
||||
await proceed.wait()
|
||||
worker_finished.set()
|
||||
return {"e2e_latency_ms": 1.0}
|
||||
|
||||
pool.slot.user_step = slow_step
|
||||
await socket.incoming.put({
|
||||
"type": "session_init_v2", "generation_mode": "t2va", "curated_prompts": ["prompt"],
|
||||
"enhancement_enabled": False,
|
||||
})
|
||||
task = start_controller(controller_module, socket, pool)
|
||||
try:
|
||||
await asyncio.wait_for(worker_started.wait(), 3)
|
||||
await socket.incoming.put({"type": "leave"})
|
||||
await asyncio.sleep(0.07)
|
||||
assert not task.done()
|
||||
controller_module.release_generation_inputs.assert_not_called()
|
||||
proceed.set()
|
||||
await asyncio.wait_for(task, 3)
|
||||
assert worker_finished.is_set()
|
||||
controller_module.release_generation_inputs.assert_called_once()
|
||||
finally:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
|
||||
asyncio.run(scenario())
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def gpu_pool_module(monkeypatch):
|
||||
streaming = ModuleType("dreamverse.av_streaming")
|
||||
for name in ("StreamChunk", "StreamComplete", "StreamEvent", "StreamInit", "generate_stream_id", "stream_fmp4"):
|
||||
setattr(streaming, name, object)
|
||||
streaming.SHARED_STREAM_BUFFER_BYTES = 1024
|
||||
streaming.USE_SHARED_STREAM_BUFFER = False
|
||||
monkeypatch.setitem(sys.modules, "dreamverse.av_streaming", streaming)
|
||||
path = Path(__file__).resolve().parents[1] / "gpu_pool.py"
|
||||
spec = importlib.util.spec_from_file_location("dreamverse_test_gpu_pool", path)
|
||||
assert spec is not None and spec.loader is not None
|
||||
module = importlib.util.module_from_spec(spec)
|
||||
monkeypatch.setitem(sys.modules, spec.name, module)
|
||||
spec.loader.exec_module(module)
|
||||
monkeypatch.setattr(module, "pin_generation_inputs", Mock())
|
||||
monkeypatch.setattr(module, "release_generation_inputs", Mock())
|
||||
return module
|
||||
|
||||
|
||||
def test_gpu_step_timeout_keeps_assets_pinned_until_late_worker_completion(gpu_pool_module):
|
||||
from dreamverse.worker_ipc import StepComplete
|
||||
|
||||
async def scenario():
|
||||
slot = gpu_pool_module.GPUSlot(0, "0")
|
||||
inputs = GenerationInputs("ref2va", (GenerationAsset("ref", "image", "/assets/ref.png", "reference"),))
|
||||
|
||||
async def timeout(command, timeout):
|
||||
assert command.payload.generation_inputs == inputs
|
||||
raise asyncio.TimeoutError
|
||||
|
||||
slot._send_command_tagged = timeout
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
await slot.user_step("user", "prompt", generation_inputs=inputs)
|
||||
gpu_pool_module.pin_generation_inputs.assert_called_once_with(inputs)
|
||||
gpu_pool_module.release_generation_inputs.assert_not_called()
|
||||
|
||||
def late_response(timeout):
|
||||
slot._active = False
|
||||
return StepComplete("user", 1, {})
|
||||
|
||||
slot.response_queue = SimpleNamespace(get=late_response)
|
||||
slot._active = True
|
||||
await slot._response_reader()
|
||||
gpu_pool_module.release_generation_inputs.assert_called_once_with(inputs)
|
||||
assert slot._step_asset_inputs == {}
|
||||
|
||||
asyncio.run(scenario())
|
||||
@@ -7,3 +7,12 @@ def test_create_generation_backend_ltx2_module_import():
|
||||
|
||||
assert isinstance(backend, LTX2GenerationBackend)
|
||||
assert backend.gpu_id == 3
|
||||
|
||||
|
||||
def test_create_generation_backend_cosmos25_dfd_module_import():
|
||||
from dreamverse.cosmos25_dfd_generation import Cosmos25DFDGenerationBackend
|
||||
|
||||
backend = _create_generation_backend("cosmos25_dfd", gpu_id=2)
|
||||
|
||||
assert isinstance(backend, Cosmos25DFDGenerationBackend)
|
||||
assert backend.gpu_id == 2
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
import ast
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
ALLOWED_PREFIXES = (
|
||||
@@ -48,3 +50,39 @@ def test_dreamverse_server_imports_only_public_fastvideo_surfaces() -> None:
|
||||
bad.append((str(path.relative_to(root)), getattr(node, "lineno", 0), name))
|
||||
|
||||
assert bad == [], f"Forbidden internal imports: {bad}"
|
||||
|
||||
|
||||
def test_h3_reference_public_export_is_lazy_and_preserves_type_identity() -> None:
|
||||
"""Only explicit reference usage should load H3's optional GPU dependencies."""
|
||||
repo_root = Path(__file__).resolve().parents[4]
|
||||
# Isolate the import graph: keep the real public API implementation/schema,
|
||||
# substituting only the unrelated legacy sampling module and heavy H3 leaf.
|
||||
script = r'''
|
||||
import importlib
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from types import ModuleType
|
||||
|
||||
root = Path(sys.argv[1])
|
||||
fastvideo = ModuleType("fastvideo")
|
||||
fastvideo.__path__ = [str(root / "fastvideo")]
|
||||
sys.modules["fastvideo"] = fastvideo
|
||||
sampling = ModuleType("fastvideo.api.sampling_param")
|
||||
sampling.SamplingParam = type("SamplingParam", (), {})
|
||||
sys.modules[sampling.__name__] = sampling
|
||||
|
||||
api = importlib.import_module("fastvideo.api")
|
||||
assert "MiniMaxH3Reference" in api.__all__
|
||||
assert "MiniMaxH3Reference" not in vars(api)
|
||||
assert not any(name.startswith("fastvideo.pipelines") for name in sys.modules)
|
||||
|
||||
internal = ModuleType("fastvideo.pipelines.basic.minimax_h3.reference")
|
||||
internal.MiniMaxH3Reference = type("MiniMaxH3Reference", (), {})
|
||||
sys.modules[internal.__name__] = internal
|
||||
from fastvideo.api import MiniMaxH3Reference
|
||||
assert MiniMaxH3Reference is internal.MiniMaxH3Reference
|
||||
assert api.MiniMaxH3Reference is internal.MiniMaxH3Reference
|
||||
assert not hasattr(api, "UnknownReference")
|
||||
'''
|
||||
result = subprocess.run([sys.executable, "-c", script, str(repo_root)], capture_output=True, text=True, timeout=30)
|
||||
assert result.returncode == 0, result.stderr
|
||||
|
||||
@@ -105,6 +105,7 @@ class _FakeSlot:
|
||||
segment_idx: int,
|
||||
reset_conditioning: bool,
|
||||
image_path: str | None = None,
|
||||
generation_inputs=None,
|
||||
):
|
||||
self.calls.append({
|
||||
"client_id": client_id,
|
||||
|
||||
@@ -15,6 +15,8 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from dreamverse.generation_inputs import GenerationInputs
|
||||
|
||||
# ---- User-scoped events (carry user_id) ------------------------------------
|
||||
|
||||
|
||||
@@ -147,6 +149,7 @@ class UserStepPayload:
|
||||
segment_idx: int
|
||||
image_path: str | None
|
||||
reset_conditioning: bool
|
||||
generation_inputs: GenerationInputs | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
|
||||
@@ -0,0 +1,107 @@
|
||||
# Dreamverse on Slurm
|
||||
|
||||
Run Full H3 inside a one-node, four-GPU allocation. The maintained H3 examples
|
||||
default to four GPUs; this is a starting configuration, not a measured minimum.
|
||||
The full checkpoint supports T2VA, FL2VA, and Ref2VA. The FastH3 Preview profile
|
||||
is a separate T2VA configuration.
|
||||
|
||||
`launch_backend.sh` checks that it is inside an `srun` step, preserves
|
||||
`CUDA_VISIBLE_DEVICES`, and replaces itself with the backend process. It does
|
||||
not allocate GPUs, kill existing processes, or source a personal credentials
|
||||
file. The local `dreamverse-deploy` helper is not suitable for a shared Slurm
|
||||
cluster because it kills processes by physical GPU and port.
|
||||
|
||||
## Prepare and allocate
|
||||
|
||||
Keep the checkout, weights, outputs, and logs on storage visible to the compute
|
||||
node. Source installation is documented in the [GPU guide](../../../../docs/getting_started/installation/gpu.md).
|
||||
On ARM64 GB200 use CUDA 13, a matching PyTorch build, and kernels built for
|
||||
`sm_100`; the DGX Spark `sm_121` kernel image is not the GB200 image.
|
||||
|
||||
The repository's image workflow publishes an ARM64 GB200 variant under
|
||||
`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-cuda13.0.0-sm100-latest`.
|
||||
Resolve that tag to a digest for reproducible runs. If your compute nodes use
|
||||
Pyxis/Enroot, pass the approved image or a prepared SquashFS file to
|
||||
`srun --container-image`, with explicit mounts for your checkout and model cache.
|
||||
The Dreamverse-specific Docker images are currently AMD64-only.
|
||||
|
||||
For the Slinky customer partition, a bounded allocation is:
|
||||
|
||||
```bash
|
||||
salloc --account=customer --qos=normal --partition=hpc-rack-1 \
|
||||
--nodes=1 --ntasks=1 --cpus-per-task=72 --gres=gpu:nvidia_gb200:4 \
|
||||
--mem=800G --time=02:00:00 --job-name=dreamverse
|
||||
srun --ntasks=1 --pty bash
|
||||
```
|
||||
|
||||
Wait for Slurm to grant the allocation before entering the compute step. A
|
||||
successful SSH login does not grant GPU resources. Inspect pending capacity
|
||||
with `squeue -u "$USER" --start`; do not attach to another user's job.
|
||||
|
||||
The checkpoint includes duplicate release layouts. Download the diffusers
|
||||
components needed by both base and reference pipelines, rather than the whole
|
||||
repository (about 210 GB versus about 498 GB at revision
|
||||
`42ed227ee7df40d41602854ae760620d6eb651fe`):
|
||||
|
||||
```bash
|
||||
hf download MiniMaxAI/MiniMax-H3 \
|
||||
--revision 42ed227ee7df40d41602854ae760620d6eb651fe \
|
||||
--include model_index.json --include modular_model_index.json \
|
||||
--include 'audio_scheduler/*' --include 'audio_vae/*' \
|
||||
--include 'processor/*' --include 'scheduler/*' \
|
||||
--include 'text_encoder/*' --include 'tokenizer/*' \
|
||||
--include 'transformer/*' --include 'transformer_ref/*' --include 'vae/*' \
|
||||
--local-dir /path/to/models/MiniMax-H3
|
||||
```
|
||||
|
||||
The GPU environment needs `fastvideo[dreamverse]`, the Dreamverse workspace
|
||||
package, and FFmpeg with H.264/AAC encoders. In a prepared FastVideo image,
|
||||
install the checked-out code and its Dreamverse dependencies in that image's
|
||||
Python environment. Keep its matching CUDA/PyTorch/kernel stack intact.
|
||||
|
||||
## Start and connect
|
||||
|
||||
From the checked-out repository inside the allocated step:
|
||||
|
||||
```bash
|
||||
export DREAMVERSE_PYTHON=/path/to/environment/bin/python
|
||||
export DREAMVERSE_MODEL_PATH=/path/to/models/MiniMax-H3
|
||||
export FASTVIDEO_DREAMVERSE_HOME=/path/to/persistent/dreamverse-state
|
||||
bash apps/dreamverse/scripts/slurm/launch_backend.sh
|
||||
```
|
||||
|
||||
The default backend binds port 8009 on the private compute node. Connect through
|
||||
the login node from your laptop, replacing `COMPUTE_NODE_IP` with the allocated
|
||||
node's `NodeAddr` from `scontrol show node`:
|
||||
|
||||
```bash
|
||||
ssh -N -L 8009:COMPUTE_NODE_IP:8009 USER@LOGIN_NODE
|
||||
```
|
||||
|
||||
In another laptop terminal, run the frontend from your local checkout:
|
||||
|
||||
```bash
|
||||
cd apps/dreamverse/web
|
||||
BACKEND_HOST=127.0.0.1 BACKEND_PORT=8009 npm run dev
|
||||
```
|
||||
|
||||
Open `http://localhost:5299`. `/healthz` reports the server process; `/readyz`
|
||||
reports model readiness. Full H3 loads and generates more slowly than the
|
||||
Preview adapter. Keep prompt enhancement disabled in the UI unless the
|
||||
runtime has the selected provider's credentials.
|
||||
|
||||
## Verify and stop
|
||||
|
||||
Check all three modes with small, valid user-owned assets. Capture the selected
|
||||
mode and assets, WebSocket errors or completion events, the generated video and
|
||||
audio, and GPU memory usage. Also verify actionable validation errors and
|
||||
backward compatibility with clients that omit `generation_mode`.
|
||||
|
||||
Use the frontend Playwright instructions in the
|
||||
[Dreamverse development guide](../../../../docs/contributing/dreamverse-development.md)
|
||||
against the forwarded backend. A mock-server demo validates UI and protocol
|
||||
behavior; it is not evidence of GPU generation.
|
||||
|
||||
Stop the backend with Ctrl-C, exit the compute step, and release your allocation.
|
||||
For a detached allocation, use `scancel YOUR_JOB_ID`. Cancel a pending demo job
|
||||
when it is no longer needed; do not leave an unattended reservation queued.
|
||||
@@ -0,0 +1,49 @@
|
||||
#!/usr/bin/env bash
|
||||
# Run inside an existing Slurm step. Slurm owns the GPU visibility and lifetime.
|
||||
set -euo pipefail
|
||||
|
||||
if [[ -z "${SLURM_JOB_ID:-}" || -z "${SLURM_STEP_ID:-}" ]]; then
|
||||
echo "Run this launcher inside an allocated Slurm step (srun), not on the login node." >&2
|
||||
exit 2
|
||||
fi
|
||||
|
||||
script_dir="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)"
|
||||
repo_root="$(cd -- "${script_dir}/../../../.." && pwd)"
|
||||
python_bin="${DREAMVERSE_PYTHON:-${repo_root}/.venv/bin/python}"
|
||||
if [[ ! -x "${python_bin}" ]]; then
|
||||
echo "Set DREAMVERSE_PYTHON to a Python environment with fastvideo[dreamverse] installed." >&2
|
||||
exit 2
|
||||
fi
|
||||
|
||||
export DREAMVERSE_MODEL_ID="${DREAMVERSE_MODEL_ID:-full-h3}"
|
||||
export DREAMVERSE_SP_SIZE="${DREAMVERSE_SP_SIZE:-4}"
|
||||
export FASTVIDEO_GPU_COUNT="${FASTVIDEO_GPU_COUNT:-${DREAMVERSE_SP_SIZE}}"
|
||||
export FASTVIDEO_ENABLE_STARTUP_WARMUP="${FASTVIDEO_ENABLE_STARTUP_WARMUP:-0}"
|
||||
export ENABLE_TORCH_COMPILE="${ENABLE_TORCH_COMPILE:-0}"
|
||||
export STREAM_MODE="${STREAM_MODE:-av_fmp4}"
|
||||
export PYTHONPATH="${repo_root}/apps/dreamverse:${repo_root}${PYTHONPATH:+:${PYTHONPATH}}"
|
||||
export PYTHONUNBUFFERED=1
|
||||
|
||||
"${python_bin}" - <<'PY'
|
||||
import os
|
||||
import shutil
|
||||
|
||||
import torch
|
||||
|
||||
expected = int(os.environ["DREAMVERSE_SP_SIZE"])
|
||||
visible = torch.cuda.device_count()
|
||||
if expected < 1 or visible < expected:
|
||||
raise SystemExit(f"The Slurm step exposes {visible} GPUs; DREAMVERSE_SP_SIZE requires {expected}.")
|
||||
ffmpeg = os.environ.get("FASTVIDEO_FFMPEG_BIN", "ffmpeg")
|
||||
if not shutil.which(ffmpeg):
|
||||
raise SystemExit("FFmpeg is missing; install it in the compute environment or set FASTVIDEO_FFMPEG_BIN.")
|
||||
print(f"Slurm job {os.environ['SLURM_JOB_ID']}: {visible} visible GPUs; using {expected} per worker")
|
||||
for index in range(expected):
|
||||
properties = torch.cuda.get_device_properties(index)
|
||||
print(f" GPU {index}: {properties.name}, {properties.total_memory / 2**30:.1f} GiB")
|
||||
PY
|
||||
|
||||
cd "${repo_root}"
|
||||
exec "${python_bin}" -m dreamverse.server_entry \
|
||||
--host "${DREAMVERSE_BIND_HOST:-0.0.0.0}" \
|
||||
--port "${DREAMVERSE_BACKEND_PORT:-8009}" "$@"
|
||||
@@ -0,0 +1,126 @@
|
||||
import { execFileSync } from "node:child_process";
|
||||
import { readFile } from "node:fs/promises";
|
||||
import path from "node:path";
|
||||
import { test, expect } from "@playwright/test";
|
||||
|
||||
const imagePath = path.resolve("public/k2.png");
|
||||
const framePrompt = "A paper fox walks through a sunlit forest, gentle birdsong.";
|
||||
|
||||
function makeAudio(sampleRate = 8000, seconds = 1): Buffer {
|
||||
const sampleCount = sampleRate * seconds;
|
||||
const bytes = Buffer.alloc(44 + sampleCount * 2);
|
||||
bytes.write("RIFF", 0); bytes.writeUInt32LE(bytes.length - 8, 4); bytes.write("WAVEfmt ", 8);
|
||||
bytes.writeUInt32LE(16, 16); bytes.writeUInt16LE(1, 20); bytes.writeUInt16LE(1, 22);
|
||||
bytes.writeUInt32LE(sampleRate, 24); bytes.writeUInt32LE(sampleRate * 2, 28);
|
||||
bytes.writeUInt16LE(2, 32); bytes.writeUInt16LE(16, 34); bytes.write("data", 36);
|
||||
bytes.writeUInt32LE(sampleCount * 2, 40);
|
||||
for (let i = 0; i < sampleCount; i++) bytes.writeInt16LE(Math.round(Math.sin(i * 440 * 2 * Math.PI / sampleRate) * 1000), 44 + i * 2);
|
||||
return bytes;
|
||||
}
|
||||
|
||||
test.describe("generation modes through the mock runtime", () => {
|
||||
for (const mode of ["t2va", "fl2va", "ref2va"] as const) {
|
||||
test(`${mode} sends validated assets and plays a clearly labeled sample`, async ({ page, request }, testInfo) => {
|
||||
const response = await request.get("/generation-capabilities");
|
||||
const capabilities = response.ok() ? await response.json() : {};
|
||||
test.skip(capabilities.mock !== true, "This test uses the CPU mock runtime; it must not silently allocate a real GPU.");
|
||||
const sent: Record<string, any>[] = [];
|
||||
const received: Record<string, any>[] = [];
|
||||
page.on("websocket", (socket) => {
|
||||
socket.on("framesent", ({ payload }) => { if (typeof payload === "string") { try { sent.push(JSON.parse(payload)); } catch {} } });
|
||||
socket.on("framereceived", ({ payload }) => { if (typeof payload === "string") { try { received.push(JSON.parse(payload)); } catch {} } });
|
||||
});
|
||||
await page.goto("/");
|
||||
await expect(page.getByText(/Demo runtime · Sample playback only/)).toBeVisible();
|
||||
const modeSelect = page.getByRole("combobox", { name: "Generation mode" });
|
||||
const modeLabel = mode === "ref2va" ? "Ref2VA" : mode.toUpperCase();
|
||||
await modeSelect.click();
|
||||
await page.getByRole("option", { name: modeLabel, exact: true }).click();
|
||||
await expect(modeSelect).toHaveText(modeLabel);
|
||||
await page.getByLabel("Continuation prompt").fill(framePrompt);
|
||||
const uploadedIds: string[] = [];
|
||||
page.on("response", async (uploadResponse) => {
|
||||
if (uploadResponse.request().method() === "POST" && uploadResponse.url().endsWith("/assets") && uploadResponse.ok()) {
|
||||
const asset = await uploadResponse.json().catch(() => null);
|
||||
if (asset?.asset_id) uploadedIds.push(asset.asset_id);
|
||||
}
|
||||
});
|
||||
try {
|
||||
if (mode === "fl2va") {
|
||||
await expect(page.getByRole("button", { name: "Generate", exact: true })).toBeDisabled();
|
||||
await page.locator('input[type="file"]').setInputFiles([
|
||||
{ name: "first-frame.png", mimeType: "image/png", buffer: await readFile(imagePath) },
|
||||
{ name: "last-frame.png", mimeType: "image/png", buffer: await readFile(imagePath) },
|
||||
]);
|
||||
await expect(page.getByRole("option", { name: "first-frame.png", exact: true }).first()).toBeAttached();
|
||||
await page.getByRole("combobox", { name: "First frame", exact: true }).selectOption({ label: "first-frame.png" });
|
||||
await expect(page.getByRole("button", { name: "Generate", exact: true })).toBeEnabled();
|
||||
await page.getByRole("combobox", { name: "Last frame", exact: true }).selectOption({ label: "last-frame.png" });
|
||||
}
|
||||
if (mode === "ref2va") {
|
||||
const video = execFileSync(process.env.FASTVIDEO_FFMPEG_BIN || "ffmpeg", ["-v", "error", "-f", "lavfi", "-i", "color=c=royalblue:s=64x64:r=8", "-t", "1", "-c:v", "libx264", "-pix_fmt", "yuv420p", "-movflags", "frag_keyframe+empty_moov", "-f", "mp4", "pipe:1"]);
|
||||
await page.locator('input[type="file"]').setInputFiles([
|
||||
{ name: "subject.png", mimeType: "image/png", buffer: await readFile(imagePath) },
|
||||
{ name: "motion.mp4", mimeType: "video/mp4", buffer: video },
|
||||
{ name: "sound.wav", mimeType: "audio/wav", buffer: makeAudio() },
|
||||
]);
|
||||
await expect(page.getByRole("button", { name: "Add sound.wav as reference" })).toBeEnabled();
|
||||
await page.getByRole("button", { name: "Add sound.wav as reference" }).click();
|
||||
await expect(page.getByRole("button", { name: "Generate", exact: true })).toBeDisabled();
|
||||
await page.getByRole("button", { name: "Add subject.png as reference" }).click();
|
||||
await page.getByRole("button", { name: "Add motion.mp4 as reference" }).click();
|
||||
await page.getByRole("button", { name: "Move sound.wav down" }).click();
|
||||
const names = await page.getByRole("list", { name: "Ordered references" }).locator("li p.font-medium").allTextContents();
|
||||
expect(names).toEqual(["subject.png", "sound.wav", "motion.mp4"]);
|
||||
}
|
||||
await page.screenshot({ path: testInfo.outputPath(`${mode}-inputs.png`), fullPage: true });
|
||||
await page.getByRole("button", { name: "Generate", exact: true }).click();
|
||||
await expect.poll(() => sent.find((item) => item.type === "session_init_v2")?.generation_mode).toBe(mode);
|
||||
const init = sent.find((item) => item.type === "session_init_v2")!;
|
||||
expect(init.conditioning_assets.map((item: any) => item.role)).toEqual(mode === "t2va" ? [] : mode === "fl2va" ? ["first_frame", "last_frame"] : ["reference", "reference", "reference"]);
|
||||
if (mode === "ref2va") expect(init.conditioning_assets.map((item: any) => item.asset_id)).toEqual([uploadedIds[0], uploadedIds[2], uploadedIds[1]]);
|
||||
await expect.poll(() => received.find((item) => item.type === "gpu_assigned")?.generation_mode).toBe(mode);
|
||||
await expect.poll(() => received.some((item) => item.type === "media_segment_complete")).toBe(true);
|
||||
await expect(page.getByText(/Demo runtime · Sample playback only/)).toBeVisible();
|
||||
await expect(modeSelect).toHaveCount(0);
|
||||
await expect.poll(async () => page.locator("video:visible").first().evaluate((element: HTMLVideoElement) => element.readyState)).toBeGreaterThanOrEqual(2);
|
||||
await page.screenshot({ path: testInfo.outputPath(`${mode}-playback.png`), fullPage: true });
|
||||
if (mode === "ref2va") {
|
||||
await page.getByRole("button", { name: "Toggle sidebar" }).click();
|
||||
await page.getByRole("button", { name: "New project", exact: true }).click();
|
||||
await expect(modeSelect).toHaveText("T2VA");
|
||||
await modeSelect.click();
|
||||
await page.getByRole("option", { name: "FL2VA", exact: true }).click();
|
||||
await expect(modeSelect).toHaveText("FL2VA");
|
||||
await page.getByRole("combobox", { name: "First frame", exact: true }).selectOption({ label: "subject.png" });
|
||||
await expect(page.getByRole("combobox", { name: "Last frame", exact: true })).toHaveValue("");
|
||||
await page.getByLabel("Continuation prompt").fill("The paper fox explores a new scene.");
|
||||
await page.getByRole("button", { name: "Generate", exact: true }).click();
|
||||
await expect.poll(() => sent.find((item) => item.type === "project_init_v1")?.generation_mode).toBe("fl2va");
|
||||
const secondProject = sent.find((item) => item.type === "project_init_v1")!;
|
||||
expect(secondProject.conditioning_assets).toEqual([{ asset_id: uploadedIds[0], role: "first_frame" }]);
|
||||
expect(sent.filter((item) => item.type === "session_init_v2")).toHaveLength(1);
|
||||
await expect.poll(() => received.filter((item) => item.type === "media_segment_complete").length).toBeGreaterThan(1);
|
||||
}
|
||||
} finally {
|
||||
await page.close();
|
||||
for (const id of uploadedIds) await request.delete(`/assets/${id}`);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
test("proxies a media upload larger than Next's default 10 MiB body limit", async ({ request }) => {
|
||||
const response = await request.get("/generation-capabilities");
|
||||
const capabilities = response.ok() ? await response.json() : {};
|
||||
test.skip(capabilities.mock !== true, "Requires the local mock runtime.");
|
||||
const audio = makeAudio(192000, 29);
|
||||
expect(audio.length).toBeGreaterThan(10 * 1024 * 1024);
|
||||
const upload = await request.post("/assets", {
|
||||
headers: { "Content-Type": "audio/wav", "X-Asset-Name": "large-proxy-check.wav" },
|
||||
data: audio,
|
||||
});
|
||||
expect(upload.status()).toBe(201);
|
||||
const asset = await upload.json();
|
||||
try { expect(asset.size).toBe(audio.length); } finally { await request.delete(`/assets/${asset.asset_id}`); }
|
||||
});
|
||||
});
|
||||
@@ -9,6 +9,8 @@ const configDir = path.dirname(fileURLToPath(import.meta.url));
|
||||
const staticExport = process.env.NEXT_OUTPUT_EXPORT === '1';
|
||||
|
||||
const nextConfig: NextConfig = {
|
||||
// Next 15.5 name for the dev rewrite-proxy body limit; Next 16 renames it to `proxyClientMaxBodySize`.
|
||||
experimental: { middlewareClientMaxBodySize: 100 * 1024 * 1024 },
|
||||
...(staticExport ? { output: 'export' as const } : {}),
|
||||
...(staticExport ? { images: { unoptimized: true } } : {}),
|
||||
outputFileTracingRoot: path.join(configDir, '..', '..', '..'),
|
||||
@@ -38,6 +40,18 @@ const nextConfig: NextConfig = {
|
||||
source: '/router/:path*',
|
||||
destination: `${backendUrl}/router/:path*`
|
||||
},
|
||||
{
|
||||
source: '/generation-capabilities',
|
||||
destination: `${backendUrl}/generation-capabilities`,
|
||||
},
|
||||
{
|
||||
source: '/assets',
|
||||
destination: `${backendUrl}/assets`,
|
||||
},
|
||||
{
|
||||
source: '/assets/:path*',
|
||||
destination: `${backendUrl}/assets/:path*`,
|
||||
},
|
||||
{
|
||||
source: '/prompt-system-config',
|
||||
destination: `${backendUrl}/prompt-system-config`,
|
||||
|
||||
@@ -589,6 +589,7 @@ describe.skip('App websocket integration', () => {
|
||||
});
|
||||
|
||||
const initMessage = outbound.find((message) => message.type === 'session_init_v2');
|
||||
expect(initMessage.generation_mode).toBe('t2va');
|
||||
expect(initMessage.preset_id).toBe('test_preset');
|
||||
expect(initMessage.curated_prompts).toEqual(['segment one', 'segment two']);
|
||||
expect(initMessage.enhancement_enabled).toBe(true);
|
||||
@@ -597,6 +598,42 @@ describe.skip('App websocket integration', () => {
|
||||
expect(initMessage.initial_rollout_prompt).toBe('');
|
||||
});
|
||||
|
||||
it('sends the selected generation mode and locks it after session start', async () => {
|
||||
const outbound: any[] = [];
|
||||
server.on('connection', (socket) => {
|
||||
socket.on('message', (rawMessage) => {
|
||||
outbound.push(JSON.parse(rawMessage as string));
|
||||
});
|
||||
});
|
||||
|
||||
const user = userEvent.setup();
|
||||
render(<Page />);
|
||||
|
||||
const modeSelect = await screen.findByRole('combobox', { name: 'Generation mode' });
|
||||
expect(modeSelect).toHaveTextContent('T2VA');
|
||||
|
||||
await user.click(modeSelect);
|
||||
await user.click(await screen.findByRole('option', { name: 'FL2VA' }));
|
||||
expect(modeSelect).toHaveTextContent('FL2VA');
|
||||
expect(modeSelect).toHaveAttribute(
|
||||
'title',
|
||||
'First/last frames to video + audio. Start from a first frame image. Add an optional last frame to guide the ending.',
|
||||
);
|
||||
|
||||
const generateButton = await screen.findByRole('button', { name: 'Generate' });
|
||||
await waitFor(() => expect(generateButton).toBeEnabled());
|
||||
await user.click(generateButton);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(outbound.some((message) => message.type === 'session_init_v2')).toBe(true);
|
||||
});
|
||||
|
||||
const initMessage = outbound.find((message) => message.type === 'session_init_v2');
|
||||
expect(initMessage.generation_mode).toBe('fl2va');
|
||||
expect(screen.queryByRole('combobox', { name: 'Generation mode' }))
|
||||
.not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it('starts a streaming session from a custom initial prompt without using curated prompts', async () => {
|
||||
const outbound: any[] = [];
|
||||
server.on('connection', (socket) => {
|
||||
|
||||
@@ -5,6 +5,7 @@ import { Download, Share2 } from "lucide-react";
|
||||
import DevtoolsShell from "@/components/devtools/DevtoolsShell";
|
||||
import MonitorPage from "@/components/MonitorPage";
|
||||
import ChatBar from "@/components/ChatBar";
|
||||
import AssetList from "@/components/AssetList";
|
||||
import SessionTimeoutModal from "@/components/SessionTimeoutModal";
|
||||
import Sidebar from "@/components/Sidebar";
|
||||
import Header from "@/components/Header";
|
||||
@@ -13,10 +14,13 @@ import Workspace from "@/components/Workspace";
|
||||
import { saveProject, saveProjectMetadata, listProjects, loadProjectClips, deleteProject, pruneOldProjects, type StoredProject, type StoredClip } from "@/lib/projectStorage";
|
||||
import { isInfrastructureError } from "@/lib/ws/reducer";
|
||||
import { useStore } from "@/hooks/useStore";
|
||||
import { useAssetLibrary } from "@/hooks/useAssetLibrary";
|
||||
import { useGenerationCapabilities } from "@/hooks/useGenerationCapabilities";
|
||||
import { resolveDevtoolsMode } from "@/lib/devtoolsMode";
|
||||
import { createAvPipeline, DEFAULT_AV_MIME } from "@/lib/media/avPipeline";
|
||||
import { remuxArchivedFmp4Segments } from "@/lib/media/fmp4Remux";
|
||||
import { DEFAULT_CUSTOM_PRESET_ID, parseStoryPresets, sanitizePresetId } from "@/lib/presets";
|
||||
import { DEFAULT_GENERATION_MODE, buildGenerationInitFields, validateGenerationInputs, type GenerationMode, type GenerationInitFields, type GenerationAsset } from "@/lib/generationMode";
|
||||
import {
|
||||
buildRewritePromptWindowSnapshot,
|
||||
buildRewritePromptWindowSnapshotFromPrompts,
|
||||
@@ -341,6 +345,19 @@ export default function Page() {
|
||||
const [isMobileShareCapable, setIsMobileShareCapable] = useState(false);
|
||||
const [videoMuted, setVideoMuted] = useState(true);
|
||||
const [timeoutModalOpen, setTimeoutModalOpen] = useState(false);
|
||||
const [generationMode, setGenerationMode] = useState<GenerationMode>(DEFAULT_GENERATION_MODE);
|
||||
const assetLibrary = useAssetLibrary();
|
||||
const { capabilities, capabilityNotice, refreshCapabilities } = useGenerationCapabilities();
|
||||
const joiningRef = useRef(false);
|
||||
const activeGenerationRef = useRef<{ fields: GenerationInitFields; assets: GenerationAsset[]; mock: boolean } | null>(null);
|
||||
const generationInputError = validateGenerationInputs(generationMode, assetLibrary.conditioningAssets, assetLibrary.assets);
|
||||
const generationSupported = capabilities.modes.includes(generationMode);
|
||||
const generationInputsValid = !generationInputError && generationSupported && !assetLibrary.uploading;
|
||||
function changeGenerationMode(mode: GenerationMode) {
|
||||
if (sessionStore.get().sessionStarted || joiningRef.current || !capabilities.modes.includes(mode)) return;
|
||||
setGenerationMode(mode);
|
||||
assetLibrary.clearConditioning();
|
||||
}
|
||||
useEffect(() => {
|
||||
setIsMobileShareCapable(typeof navigator.canShare === "function" && window.matchMedia("(pointer: coarse)").matches);
|
||||
}, []);
|
||||
@@ -391,7 +408,7 @@ export default function Page() {
|
||||
|
||||
// --- Derived values ---
|
||||
|
||||
const canStartSession = !projectResetPending && (canJoinSession || Boolean(normalizeInitialPrompt(livePromptDraft as string)));
|
||||
const canStartSession = generationInputsValid && !projectResetPending && (canJoinSession || Boolean(normalizeInitialPrompt(livePromptDraft as string)));
|
||||
|
||||
const currentClipLabel = useMemo(() => {
|
||||
if ((activeClip as Record<string, any>)?.label) return (activeClip as Record<string, any>).label;
|
||||
@@ -747,7 +764,11 @@ export default function Page() {
|
||||
|
||||
function recoverFailedSessionStart(notice: string) {
|
||||
const restoredDraft = normalizeInitialPrompt(pendingInitialPromptRef.current);
|
||||
resetToLobbyState();
|
||||
if (wsRef.current) {
|
||||
detachAndCloseWebSocket(wsRef.current);
|
||||
wsRef.current = null;
|
||||
}
|
||||
resetToLobbyState({ preserveSessionNotice: true });
|
||||
clearPendingProjectPointers();
|
||||
pendingInitialPromptRef.current = "";
|
||||
sessionStore.patch({
|
||||
@@ -1702,6 +1723,10 @@ export default function Page() {
|
||||
|
||||
function resetToLobbyState({ preserveSessionNotice = false, preservePlayback = false } = {}) {
|
||||
setVideoMuted(true);
|
||||
if (!preserveSessionNotice) {
|
||||
setGenerationMode(DEFAULT_GENERATION_MODE);
|
||||
assetLibrary.clearConditioning();
|
||||
}
|
||||
clearCountdownInterval();
|
||||
pendingInitialPromptRef.current = "";
|
||||
sessionStore.patch({
|
||||
@@ -1736,6 +1761,8 @@ export default function Page() {
|
||||
|
||||
function resetToProjectLobbyState() {
|
||||
setVideoMuted(true);
|
||||
setGenerationMode(DEFAULT_GENERATION_MODE);
|
||||
assetLibrary.clearConditioning();
|
||||
pendingInitialPromptRef.current = "";
|
||||
sessionStore.patch({
|
||||
sessionStarted: false,
|
||||
@@ -1764,6 +1791,7 @@ export default function Page() {
|
||||
setSeedPrompts(segmentPrompts);
|
||||
return {
|
||||
type,
|
||||
...(activeGenerationRef.current?.fields || buildGenerationInitFields(generationMode, assetLibrary.conditioningAssets, assetLibrary.assets)),
|
||||
preset_id: getInitialPresetId(),
|
||||
preset_label: getInitialPresetLabel(),
|
||||
curated_prompts: segmentPrompts,
|
||||
@@ -1812,6 +1840,12 @@ export default function Page() {
|
||||
return;
|
||||
}
|
||||
if (decoded.kind !== "json") return;
|
||||
if (decoded.data?.type === "error" && sessionStore.get().sessionStarted
|
||||
&& (!sessionStore.get().gpuAssigned || decoded.data.error_code === "invalid_generation_input")) {
|
||||
const message = typeof decoded.data.message === "string" ? decoded.data.message : "The generation inputs were rejected. Check the mode and selected assets.";
|
||||
recoverFailedSessionStart(message);
|
||||
return;
|
||||
}
|
||||
if (decoded.data?.type === "error" && isInfrastructureError(decoded.data)) {
|
||||
const message = typeof decoded.data?.message === "string" && decoded.data.message.trim()
|
||||
? decoded.data.message.trim()
|
||||
@@ -1937,9 +1971,15 @@ export default function Page() {
|
||||
}
|
||||
}
|
||||
|
||||
function beginProjectLocally({ force = false } = {}) {
|
||||
function beginProjectLocally({ force = false, mockRuntime = capabilities.mock === true } = {}) {
|
||||
if (!force && !canStartSession) return;
|
||||
if (!generationInputsValid) return false;
|
||||
if (sessionStore.get().sessionStarted || sessionStore.get().projectResetPending) return false;
|
||||
activeGenerationRef.current = {
|
||||
fields: buildGenerationInitFields(generationMode, assetLibrary.conditioningAssets, assetLibrary.assets),
|
||||
assets: assetLibrary.assets.filter((asset) => assetLibrary.conditioningAssets.some((item) => item.asset_id === asset.asset_id)),
|
||||
mock: mockRuntime,
|
||||
};
|
||||
setTimeoutModalOpen(false);
|
||||
// Unmute during the user gesture so iOS Safari permits audio playback.
|
||||
setVideoMuted(false);
|
||||
@@ -1999,12 +2039,43 @@ export default function Page() {
|
||||
}
|
||||
|
||||
async function joinSession({ force = false } = {}) {
|
||||
if (joiningRef.current || sessionStore.get().sessionStarted) return;
|
||||
if (generationInputError || assetLibrary.uploading) {
|
||||
showPreSessionNotice(generationInputError || "Wait for the asset upload to finish.");
|
||||
return;
|
||||
}
|
||||
joiningRef.current = true;
|
||||
try {
|
||||
await startGenerationSession({ force });
|
||||
} finally {
|
||||
joiningRef.current = false;
|
||||
}
|
||||
}
|
||||
|
||||
async function startGenerationSession({ force = false } = {}) {
|
||||
sessionStore.patch({ sessionNotice: "" });
|
||||
streamStore.patch({ loadingAnimation: true });
|
||||
const currentCapabilities = await refreshCapabilities();
|
||||
if (!currentCapabilities.modes.includes(generationMode)) {
|
||||
streamStore.patch({ loadingAnimation: false });
|
||||
showPreSessionNotice(`${generationMode.toUpperCase()} is unavailable on this runtime. Connect a full H3 runtime or choose a supported mode.`);
|
||||
return;
|
||||
}
|
||||
const assetProblem = await assetLibrary.verifySelectedAssets();
|
||||
if (assetProblem) {
|
||||
streamStore.patch({ loadingAnimation: false });
|
||||
showPreSessionNotice(assetProblem);
|
||||
return;
|
||||
}
|
||||
if (
|
||||
wsRef.current
|
||||
&& wsRef.current.readyState === WebSocket.OPEN
|
||||
&& sessionStore.get().connected
|
||||
) {
|
||||
if (!beginProjectLocally({ force })) return;
|
||||
if (!beginProjectLocally({ force, mockRuntime: currentCapabilities.mock === true })) {
|
||||
streamStore.patch({ loadingAnimation: false });
|
||||
return;
|
||||
}
|
||||
sendProjectInitMessage();
|
||||
return;
|
||||
}
|
||||
@@ -2017,7 +2088,7 @@ export default function Page() {
|
||||
showPreSessionNotice(probe.notice);
|
||||
return;
|
||||
}
|
||||
if (!beginProjectLocally({ force })) {
|
||||
if (!beginProjectLocally({ force, mockRuntime: currentCapabilities.mock === true })) {
|
||||
streamStore.patch({ loadingAnimation: false });
|
||||
sessionStore.patch({ connecting: false });
|
||||
return;
|
||||
@@ -2044,6 +2115,10 @@ export default function Page() {
|
||||
createdAt: currentProjectCreatedAtRef.current || Date.now(),
|
||||
lastThumbnail: currentThumbnail,
|
||||
promptEvents: [...(rewriteStore.get().promptEvents as Record<string, unknown>[])],
|
||||
generationMode: activeGenerationRef.current?.fields.generation_mode || DEFAULT_GENERATION_MODE,
|
||||
conditioningAssets: activeGenerationRef.current?.fields.conditioning_assets || [],
|
||||
assets: activeGenerationRef.current?.assets || [],
|
||||
mock: activeGenerationRef.current?.mock === true,
|
||||
};
|
||||
const clips: StoredClip[] = (streamStore.get().completedClips as any[])
|
||||
.filter((clip: any) => clip?.blob instanceof Blob)
|
||||
@@ -2478,6 +2553,24 @@ export default function Page() {
|
||||
|
||||
// --- Render ---
|
||||
|
||||
const conditioningPanel = generationMode !== "t2va" && !sessionStarted && !sessionExpired ? (
|
||||
<AssetList
|
||||
mode={generationMode}
|
||||
assets={assetLibrary.assets}
|
||||
conditioning={assetLibrary.conditioningAssets}
|
||||
locked={Boolean(loadingAnimation || projectResetPending)}
|
||||
uploading={assetLibrary.uploading}
|
||||
error={assetLibrary.assetError}
|
||||
validationNotice={generationInputError}
|
||||
onUpload={assetLibrary.uploadAssets}
|
||||
onAssign={assetLibrary.assignAsset}
|
||||
onRemove={assetLibrary.removeAsset}
|
||||
onUnselect={assetLibrary.removeConditioning}
|
||||
onMove={assetLibrary.moveConditioning}
|
||||
onMissing={assetLibrary.checkAssetAvailability}
|
||||
/>
|
||||
) : null;
|
||||
|
||||
if (!runtimeReady) {
|
||||
return null;
|
||||
}
|
||||
@@ -2501,13 +2594,17 @@ export default function Page() {
|
||||
enhancementEnabled={enhancementEnabled as boolean}
|
||||
autoExtensionEnabled={autoExtensionEnabled as boolean}
|
||||
loopGenerationEnabled={loopGenerationEnabled as boolean}
|
||||
canJoinSession={canJoinSession as boolean}
|
||||
canJoinSession={canStartSession}
|
||||
canSubmitContinuation={canSubmitContinuation}
|
||||
editableMode={editableMode as boolean}
|
||||
demoMode={demoMode as boolean}
|
||||
editableCanJoin={editableCanJoin as boolean}
|
||||
curatedPromptLimit={curatedPromptLimit as number}
|
||||
maxCuratedPromptCount={maxCuratedPromptCount as number}
|
||||
generationMode={generationMode}
|
||||
supportedGenerationModes={capabilities.modes}
|
||||
conditioningPanel={conditioningPanel}
|
||||
onGenerationModeChange={changeGenerationMode}
|
||||
onPresetChange={handlePresetSelectionChange}
|
||||
onEnhancementToggle={handleEnhancementToggle}
|
||||
onCuratedPromptLimitChange={handleCuratedPromptLimitChange}
|
||||
@@ -2641,7 +2738,12 @@ export default function Page() {
|
||||
/>
|
||||
<Header timeLeft={headerTimeLeft} formatTime={formatTime} onToggleSidebar={() => setSidebarOpen((prev) => !prev)} />
|
||||
|
||||
<div className="relative flex flex-1 min-h-0 flex-col justify-center px-4 pb-2 sm:px-6 sm:pb-12">
|
||||
<div className={cn(
|
||||
"relative flex flex-1 min-h-0 flex-col px-4 pb-2 sm:px-6 sm:pb-12",
|
||||
!isViewingMode && !showActiveProject && generationMode !== "t2va"
|
||||
? "justify-start overflow-y-auto pt-4"
|
||||
: "justify-center",
|
||||
)}>
|
||||
{isViewingMode && (
|
||||
<>
|
||||
{viewingSelectedClip && (
|
||||
@@ -2680,6 +2782,8 @@ export default function Page() {
|
||||
/>
|
||||
</section>
|
||||
<motion.div layout="position" className="mx-auto w-full max-w-2xl shrink-0" transition={{ type: "spring", stiffness: 200, damping: 25 }}>
|
||||
{viewingProject?.project.mock && <p className="mb-2 text-center text-xs text-violet-600 dark:text-violet-300">Demo sample · This saved clip was not generated by an AI model.</p>}
|
||||
{viewingProject?.project.generationMode && <p className="mb-3 text-center text-xs text-muted-foreground">{viewingProject.project.generationMode.toUpperCase()} · {viewingProject.project.conditioningAssets?.length || 0} saved references. Uploaded originals may expire; your saved video remains available.</p>}
|
||||
<ChatBar sessionStarted={false} viewingReadOnly={true} onStartNewProject={handleStartNewProject} onBackFromViewing={closeViewingProject} />
|
||||
</motion.div>
|
||||
</>
|
||||
@@ -2759,7 +2863,7 @@ export default function Page() {
|
||||
</section>
|
||||
|
||||
<AnimatePresence>
|
||||
{!showActiveProject && (
|
||||
{!showActiveProject && generationMode === "t2va" && (
|
||||
<motion.div
|
||||
key="hero-tagline"
|
||||
initial={{ opacity: 0 }}
|
||||
@@ -2785,6 +2889,12 @@ export default function Page() {
|
||||
sessionExpired={sessionExpired as boolean}
|
||||
sessionNotice={sessionNotice as string}
|
||||
projectResetPending={projectResetPending as boolean}
|
||||
generationMode={generationMode}
|
||||
supportedGenerationModes={capabilities.modes}
|
||||
generationInputsValid={generationInputsValid}
|
||||
capabilityNotice={!generationSupported ? `${generationMode.toUpperCase()} is unavailable on this runtime.` : capabilityNotice}
|
||||
mockRuntime={capabilities.mock}
|
||||
conditioningPanel={conditioningPanel}
|
||||
onPresetGenerate={handlePresetGenerate}
|
||||
onContinuationInput={handleLivePromptInput}
|
||||
onContinuationKeydown={handleLivePromptKeydown}
|
||||
@@ -2792,6 +2902,7 @@ export default function Page() {
|
||||
onSubmitContinuation={submitLivePrompt}
|
||||
onLeave={leaveSession}
|
||||
onStartNewProject={handleStartNewProject}
|
||||
onGenerationModeChange={changeGenerationMode}
|
||||
onSpeechTranscript={handleLivePromptSpeechTranscript}
|
||||
onSpeechInterimChange={handleLivePromptSpeechInterim}
|
||||
/>
|
||||
|
||||
@@ -0,0 +1,41 @@
|
||||
import { render, screen } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import AssetList from "./AssetList";
|
||||
import type { GenerationAsset } from "@/lib/generationMode";
|
||||
|
||||
const frame: GenerationAsset = { asset_id: "first", kind: "image", name: "frame.png", mime_type: "image/png", size: 2000, url: "/assets/first" };
|
||||
const video: GenerationAsset = { asset_id: "video", kind: "video", name: "motion.mp4", mime_type: "video/mp4", size: 2000, url: "/assets/video" };
|
||||
function props() {
|
||||
return { assets: [frame, video], onUpload: vi.fn(), onAssign: vi.fn(), onRemove: vi.fn(), onUnselect: vi.fn(), onMove: vi.fn(), onMissing: vi.fn() };
|
||||
}
|
||||
|
||||
describe("Asset List", () => {
|
||||
it("uploads to the library and assigns images to endpoint roles", async () => {
|
||||
const callbacks = props();
|
||||
const user = userEvent.setup();
|
||||
render(<AssetList {...callbacks} mode="fl2va" conditioning={[]} />);
|
||||
await user.selectOptions(screen.getByRole("combobox", { name: "First frame" }), "first");
|
||||
expect(callbacks.onAssign).toHaveBeenCalledWith("first", "first_frame");
|
||||
expect(screen.getByRole("combobox", { name: "Last frame" })).toHaveValue("");
|
||||
expect(screen.queryByRole("option", { name: "motion.mp4" })).not.toBeInTheDocument();
|
||||
const file = new File(["image"], "new.png", { type: "image/png" });
|
||||
await user.upload(screen.getByLabelText("Upload assets", { selector: "input" }), file);
|
||||
expect(callbacks.onUpload).toHaveBeenCalledWith([file]);
|
||||
});
|
||||
it("exposes accessible ordering and removal controls for multimodal references", async () => {
|
||||
const callbacks = props();
|
||||
const user = userEvent.setup();
|
||||
render(<AssetList {...callbacks} mode="ref2va" conditioning={[{ asset_id: "first", role: "reference" }, { asset_id: "video", role: "reference" }]} />);
|
||||
await user.click(screen.getByRole("button", { name: "Move motion.mp4 up" }));
|
||||
expect(callbacks.onMove).toHaveBeenCalledWith(1, 0);
|
||||
await user.click(screen.getByRole("button", { name: "Unselect frame.png" }));
|
||||
expect(callbacks.onUnselect).toHaveBeenCalledWith(0);
|
||||
expect(screen.getByRole("button", { name: "Move frame.png up" })).toBeDisabled();
|
||||
});
|
||||
it("locks uploads and assignments while starting generation", () => {
|
||||
render(<AssetList {...props()} mode="fl2va" conditioning={[]} locked />);
|
||||
expect(screen.getByRole("button", { name: "Upload assets" })).toBeDisabled();
|
||||
expect(screen.getByRole("combobox", { name: "First frame" })).toBeDisabled();
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,154 @@
|
||||
"use client";
|
||||
|
||||
import { useRef, useState } from "react";
|
||||
import { ArrowDown, ArrowUp, AudioLines, Check, GripVertical, ImagePlus, Plus, Trash2, Upload, X } from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { NativeSelect } from "@/components/ui/native-select";
|
||||
import { cn } from "@/lib/utils";
|
||||
import type { ConditioningAsset, ConditioningRole, GenerationAsset, GenerationMode } from "@/lib/generationMode";
|
||||
|
||||
interface AssetListProps {
|
||||
mode: GenerationMode;
|
||||
assets: GenerationAsset[];
|
||||
conditioning: ConditioningAsset[];
|
||||
locked?: boolean;
|
||||
uploading?: boolean;
|
||||
error?: string;
|
||||
validationNotice?: string | null;
|
||||
onUpload: (files: File[]) => void;
|
||||
onAssign: (assetId: string, role: ConditioningRole) => void;
|
||||
onRemove: (assetId: string) => void;
|
||||
onUnselect: (index: number) => void;
|
||||
onMove: (from: number, to: number) => void;
|
||||
onMissing: (assetId: string) => void;
|
||||
}
|
||||
|
||||
function AssetPreview({ asset, onMissing, compact = false }: {
|
||||
asset: GenerationAsset;
|
||||
onMissing: (assetId: string) => void;
|
||||
compact?: boolean;
|
||||
}) {
|
||||
const [previewFailed, setPreviewFailed] = useState(false);
|
||||
function previewError() {
|
||||
setPreviewFailed(true);
|
||||
onMissing(asset.asset_id);
|
||||
}
|
||||
const className = cn("h-full w-full object-cover", asset.missing && "opacity-25");
|
||||
if (asset.missing) return <span className="p-2 text-center text-[10px] text-muted-foreground">Upload again</span>;
|
||||
if (previewFailed) return <span className="p-2 text-center text-[10px] text-muted-foreground">Preview unavailable</span>;
|
||||
if (asset.kind === "image") {
|
||||
return <img src={asset.url} alt={asset.name} className={className} onError={previewError} />;
|
||||
}
|
||||
if (asset.kind === "video") {
|
||||
return <video src={asset.url} aria-label={`Preview ${asset.name}`} className={className} muted playsInline controls={!compact} preload="metadata" onError={previewError} />;
|
||||
}
|
||||
return (
|
||||
<div className="flex h-full w-full flex-col items-center justify-center gap-2 bg-violet-500/10 p-2 text-violet-500">
|
||||
<AudioLines className="size-6" />
|
||||
{!compact && <audio src={asset.url} aria-label={`Preview ${asset.name}`} controls preload="metadata" className="h-6 w-full min-w-0" onError={previewError} />}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
|
||||
/** A reusable library/picker. The parent asset store owns uploads and selection. */
|
||||
export default function AssetList({
|
||||
mode, assets, conditioning, locked = false, uploading = false, error = "", validationNotice,
|
||||
onUpload, onAssign, onRemove, onUnselect, onMove, onMissing,
|
||||
}: AssetListProps) {
|
||||
const inputRef = useRef<HTMLInputElement>(null);
|
||||
const [libraryOpen, setLibraryOpen] = useState(true);
|
||||
const [dragIndex, setDragIndex] = useState<number | null>(null);
|
||||
const disabled = locked || uploading;
|
||||
const imageAssets = assets.filter((asset) => asset.kind === "image");
|
||||
|
||||
return (
|
||||
<section aria-label="Asset List" className="overflow-hidden rounded-2xl border border-input bg-card/70 shadow-sm backdrop-blur-sm">
|
||||
<div className="flex items-center justify-between gap-3 px-4 py-3">
|
||||
<div>
|
||||
<h2 className="text-xs font-semibold tracking-wide">{mode === "fl2va" ? "Frame guidance" : "Reference sequence"}</h2>
|
||||
<p className="mt-0.5 text-[11px] text-muted-foreground">{locked ? "Inputs are locked for this project." : mode === "fl2va" ? "Choose your opening image and, optionally, the ending." : "Arrange references in the order you want the model to read them."}</p>
|
||||
</div>
|
||||
<Button type="button" variant="outline" size="sm" disabled={disabled} onClick={() => inputRef.current?.click()} className="shrink-0 gap-1.5 rounded-full text-xs">
|
||||
<Upload className="size-3.5" />{uploading ? "Uploading…" : "Upload assets"}
|
||||
</Button>
|
||||
<input ref={inputRef} type="file" aria-label="Upload assets" className="sr-only" multiple accept={mode === "fl2va" ? "image/*" : "image/*,video/*,audio/*"} disabled={disabled} onChange={(event) => {
|
||||
const files = Array.from(event.target.files || []);
|
||||
if (files.length) onUpload(files);
|
||||
event.target.value = "";
|
||||
}} />
|
||||
</div>
|
||||
|
||||
<div className="max-h-[min(42vh,350px)] overflow-y-auto px-4 pb-3">
|
||||
{mode === "fl2va" ? (
|
||||
<div className="grid grid-cols-2 gap-3">
|
||||
{(["first_frame", "last_frame"] as const).map((role) => {
|
||||
const label = role === "first_frame" ? "First frame" : "Last frame";
|
||||
const assetId = conditioning.find((item) => item.role === role)?.asset_id || "";
|
||||
const asset = assets.find((item) => item.asset_id === assetId);
|
||||
return (
|
||||
<div key={role} className="overflow-hidden rounded-xl border border-input bg-background/40 p-2">
|
||||
<div className="flex h-20 items-center justify-center overflow-hidden rounded-lg bg-muted/60 sm:h-24">
|
||||
{asset ? <AssetPreview key={asset.asset_id} asset={asset} onMissing={onMissing} /> : <ImagePlus className="size-6 text-muted-foreground/45" />}
|
||||
</div>
|
||||
<label htmlFor={`asset-${role}`} className="mb-1 mt-2 block text-[11px] font-medium">{label} <span className="font-normal text-muted-foreground">{role === "first_frame" ? "· required" : "· optional"}</span></label>
|
||||
<NativeSelect id={`asset-${role}`} aria-label={label} value={assetId} disabled={disabled} className="h-8 text-xs" onChange={(event) => onAssign(event.target.value, role)}>
|
||||
<option value="">{imageAssets.length ? "Choose an image" : "Upload an image first"}</option>
|
||||
{imageAssets.map((item) => <option key={item.asset_id} value={item.asset_id} disabled={item.missing}>{item.name}{item.missing ? " (upload again)" : ""}</option>)}
|
||||
</NativeSelect>
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>
|
||||
) : (
|
||||
<>
|
||||
{conditioning.length ? (
|
||||
<ol aria-label="Ordered references" className="flex flex-col gap-2">
|
||||
{conditioning.map((item, index) => {
|
||||
const asset = assets.find((entry) => entry.asset_id === item.asset_id);
|
||||
if (!asset) return null;
|
||||
return (
|
||||
<li key={`${item.asset_id}-${index}`} draggable={!disabled} onDragStart={() => setDragIndex(index)} onDragEnd={() => setDragIndex(null)} onDragOver={(event) => { if (!disabled && dragIndex !== null) event.preventDefault(); }} onDrop={(event) => { event.preventDefault(); if (!disabled && dragIndex !== null) onMove(dragIndex, index); setDragIndex(null); }} className={cn("flex items-center gap-2 rounded-xl border border-input bg-background/40 p-2", dragIndex === index && "opacity-50")}>
|
||||
<GripVertical className="hidden size-3.5 shrink-0 text-muted-foreground/50 sm:block" aria-hidden />
|
||||
<span className="w-4 text-center text-[11px] font-medium text-muted-foreground">{index + 1}</span>
|
||||
<div className="flex size-10 shrink-0 items-center justify-center overflow-hidden rounded-md bg-muted"><AssetPreview asset={asset} onMissing={onMissing} compact /></div>
|
||||
<div className="min-w-0 flex-1"><p className="truncate text-xs font-medium">{asset.name}</p><p className="text-[10px] capitalize text-muted-foreground">{asset.kind}{asset.missing ? " · unavailable" : ""}</p></div>
|
||||
<Button type="button" variant="ghost" size="icon-sm" aria-label={`Move ${asset.name} up`} disabled={disabled || index === 0} onClick={() => onMove(index, index - 1)}><ArrowUp className="size-3.5" /></Button>
|
||||
<Button type="button" variant="ghost" size="icon-sm" aria-label={`Move ${asset.name} down`} disabled={disabled || index === conditioning.length - 1} onClick={() => onMove(index, index + 1)}><ArrowDown className="size-3.5" /></Button>
|
||||
<Button type="button" variant="ghost" size="icon-sm" aria-label={`Unselect ${asset.name}`} disabled={disabled} onClick={() => onUnselect(index)}><X className="size-3.5" /></Button>
|
||||
</li>
|
||||
);
|
||||
})}
|
||||
</ol>
|
||||
) : (
|
||||
<div className="flex items-center gap-3 rounded-xl border border-dashed border-input px-4 py-4 text-muted-foreground"><ImagePlus className="size-6 shrink-0 opacity-50" /><p className="text-xs">Add images, video, or audio from your asset library.<br /><span className="text-[11px] opacity-75">At least one image or video is required.</span></p></div>
|
||||
)}
|
||||
<p className="mt-2 text-[10px] text-muted-foreground">{conditioning.length}/12 selected · up to 9 images, 3 videos, 3 audio clips</p>
|
||||
</>
|
||||
)}
|
||||
|
||||
{assets.length > 0 && !locked && (
|
||||
<div className="mt-3 border-t border-border/60 pt-2">
|
||||
<button type="button" className="flex w-full items-center justify-between py-1 text-[11px] font-medium text-muted-foreground" aria-expanded={libraryOpen} onClick={() => setLibraryOpen(!libraryOpen)}><span>Asset library · {assets.length}</span><span>{libraryOpen ? "Hide" : "Show"}</span></button>
|
||||
{libraryOpen && <div className="mt-2 grid grid-cols-2 gap-2 sm:grid-cols-3">
|
||||
{assets.map((asset) => {
|
||||
const selected = conditioning.some((item) => item.asset_id === asset.asset_id);
|
||||
return (
|
||||
<div key={asset.asset_id} className={cn("overflow-hidden rounded-lg border bg-background/40", selected ? "border-sky-400/70" : "border-input")}>
|
||||
<div className="flex h-20 items-center justify-center overflow-hidden bg-muted/50"><AssetPreview asset={asset} onMissing={onMissing} /></div>
|
||||
<div className="flex items-center gap-1 p-1.5">
|
||||
<div className="min-w-0 flex-1"><p title={asset.name} className="truncate text-[10px] font-medium">{asset.name}</p><p className="text-[9px] capitalize text-muted-foreground">{asset.missing ? "Upload again" : `${asset.kind} · ${(asset.size / 1024 / 1024).toFixed(1)} MB`}</p></div>
|
||||
{mode === "ref2va" && <Button type="button" variant="ghost" size="icon-sm" className="size-7" aria-label={`Add ${asset.name} as reference`} disabled={disabled || selected || asset.missing || conditioning.length >= 12} onClick={() => onAssign(asset.asset_id, "reference")}>{selected ? <Check className="size-3.5 text-sky-500" /> : <Plus className="size-3.5" />}</Button>}
|
||||
<Button type="button" variant="ghost" size="icon-sm" className="size-7 text-muted-foreground" aria-label={`Remove asset ${asset.name}`} disabled={disabled} onClick={() => onRemove(asset.asset_id)}><Trash2 className="size-3" /></Button>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
})}
|
||||
</div>}
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
{!locked && <p className="px-4 pb-2 text-[10px] text-muted-foreground">Images ≤15 MiB / 16 MP{mode === "ref2va" ? " / 1:4–4:1 aspect ratio" : ""} · video/audio ≤100 MiB / 30 sec · video up to 4K · mono/stereo audio</p>}
|
||||
{(error || validationNotice) && <p role={error ? "alert" : "status"} className={cn("border-t border-border/60 px-4 py-2 text-[11px]", error ? "bg-rose-500/5 text-rose-600 dark:text-rose-300" : "bg-amber-500/5 text-amber-700 dark:text-amber-300")}>{error || validationNotice}</p>}
|
||||
</section>
|
||||
);
|
||||
}
|
||||
@@ -0,0 +1,155 @@
|
||||
import { render, screen, waitFor, within } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { afterAll, beforeAll, describe, expect, it, vi } from "vitest";
|
||||
|
||||
import ChatBar from "./ChatBar";
|
||||
|
||||
// JSDOM does not implement the pointer/scroll APIs used by the Radix popup.
|
||||
const domPolyfills = {
|
||||
hasPointerCapture: () => false,
|
||||
releasePointerCapture: () => {},
|
||||
scrollIntoView: () => {},
|
||||
};
|
||||
const originalDescriptors = new Map<string, PropertyDescriptor | undefined>();
|
||||
beforeAll(() => {
|
||||
for (const [name, implementation] of Object.entries(domPolyfills)) {
|
||||
originalDescriptors.set(name, Object.getOwnPropertyDescriptor(HTMLElement.prototype, name));
|
||||
Object.defineProperty(HTMLElement.prototype, name, { configurable: true, value: implementation });
|
||||
}
|
||||
vi.stubGlobal("PointerEvent", MouseEvent);
|
||||
});
|
||||
afterAll(() => {
|
||||
for (const [name, descriptor] of originalDescriptors) {
|
||||
if (descriptor) Object.defineProperty(HTMLElement.prototype, name, descriptor);
|
||||
else Reflect.deleteProperty(HTMLElement.prototype, name);
|
||||
}
|
||||
vi.unstubAllGlobals();
|
||||
});
|
||||
|
||||
describe("ChatBar generation mode selection", () => {
|
||||
it("places Mode and the prompt input inside the same composer", () => {
|
||||
render(<ChatBar />);
|
||||
|
||||
const composer = screen.getByRole("group", { name: "Prompt composer" });
|
||||
expect(within(composer).getByText("Mode", { exact: true })).toBeVisible();
|
||||
expect(within(composer).getByRole("combobox", { name: "Generation mode" }))
|
||||
.toBeVisible();
|
||||
expect(within(composer).getByRole("textbox", { name: "Continuation prompt" }))
|
||||
.toBeVisible();
|
||||
expect(screen.queryByText("Generation mode", { exact: true }))
|
||||
.not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("shows only mode abbreviations and keeps explanations in the tooltip", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<ChatBar />);
|
||||
|
||||
expect(screen.getByRole("combobox", { name: "Generation mode" }))
|
||||
.toHaveAttribute("title", "Text to video + audio. Start with a text prompt; no reference asset is required.");
|
||||
expect(screen.queryByText("Start with a text prompt; no reference asset is required."))
|
||||
.not.toBeInTheDocument();
|
||||
expect(screen.queryByRole("listbox")).not.toBeInTheDocument();
|
||||
await user.click(screen.getByRole("combobox", { name: "Generation mode" }));
|
||||
const menu = await screen.findByRole("listbox");
|
||||
expect(within(menu).getAllByRole("option").map((option) => option.textContent))
|
||||
.toEqual(["T2VA", "FL2VA", "Ref2VA"]);
|
||||
});
|
||||
|
||||
it("defaults to T2VA and reports a selected mode", async () => {
|
||||
const onGenerationModeChange = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
|
||||
render(
|
||||
<ChatBar
|
||||
canJoinSession
|
||||
continuationDraft="A lighthouse in a storm"
|
||||
onGenerationModeChange={onGenerationModeChange}
|
||||
/>,
|
||||
);
|
||||
|
||||
const modeSelect = screen.getByRole("combobox", { name: "Generation mode" });
|
||||
expect(modeSelect).toHaveTextContent("T2VA");
|
||||
|
||||
await user.click(modeSelect);
|
||||
await user.click(await screen.findByRole("option", { name: "Ref2VA" }));
|
||||
|
||||
expect(onGenerationModeChange).toHaveBeenCalledWith("ref2va");
|
||||
expect(screen.queryByRole("listbox")).not.toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("hides mode selection after generation starts", () => {
|
||||
render(<ChatBar sessionStarted />);
|
||||
|
||||
expect(screen.queryByRole("combobox", { name: "Generation mode" }))
|
||||
.not.toBeInTheDocument();
|
||||
expect(within(screen.getByRole("group", { name: "Prompt composer" }))
|
||||
.getByRole("textbox", { name: "Continuation prompt" })).toBeVisible();
|
||||
});
|
||||
|
||||
it("disables mode selection and prompt editing while generation is busy", async () => {
|
||||
const onGenerationModeChange = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
render(<ChatBar isGenerating onGenerationModeChange={onGenerationModeChange} />);
|
||||
|
||||
const composer = screen.getByRole("group", { name: "Prompt composer" });
|
||||
const modeSelect = within(composer).getByRole("combobox", { name: "Generation mode" });
|
||||
expect(modeSelect).toBeDisabled();
|
||||
expect(within(composer).getByRole("textbox", { name: "Continuation prompt" }))
|
||||
.toBeDisabled();
|
||||
await user.click(modeSelect);
|
||||
expect(screen.queryByRole("listbox")).not.toBeInTheDocument();
|
||||
expect(onGenerationModeChange).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("still submits the prompt with Enter from the combined composer", async () => {
|
||||
const onGenerate = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
render(<ChatBar canJoinSession continuationDraft="A lighthouse in a storm" onGenerate={onGenerate} />);
|
||||
|
||||
const input = within(screen.getByRole("group", { name: "Prompt composer" }))
|
||||
.getByRole("textbox", { name: "Continuation prompt" });
|
||||
await user.click(input);
|
||||
await user.keyboard("{Enter}");
|
||||
expect(onGenerate).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("disables unsupported modes and labels mock playback", async () => {
|
||||
const user = userEvent.setup();
|
||||
render(<ChatBar supportedGenerationModes={["t2va"]} mockRuntime />);
|
||||
expect(screen.getByText(/No AI model is generating/)).toBeInTheDocument();
|
||||
await user.click(screen.getByRole("combobox", { name: "Generation mode" }));
|
||||
expect(await screen.findByRole("option", { name: "FL2VA" })).toHaveAttribute("aria-disabled", "true");
|
||||
expect(screen.getByRole("option", { name: "Ref2VA" })).toHaveAttribute("aria-disabled", "true");
|
||||
expect(screen.getByRole("option", { name: "Ref2VA" }))
|
||||
.toHaveAttribute("title", "References to video + audio (unavailable on this runtime)");
|
||||
});
|
||||
|
||||
it("closes the menu with Escape and restores focus to Mode", async () => {
|
||||
const onGenerationModeChange = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
render(<ChatBar onGenerationModeChange={onGenerationModeChange} />);
|
||||
const modeSelect = screen.getByRole("combobox", { name: "Generation mode" });
|
||||
await user.click(modeSelect);
|
||||
await screen.findByRole("listbox");
|
||||
await user.keyboard("{Escape}");
|
||||
expect(screen.queryByRole("listbox")).not.toBeInTheDocument();
|
||||
await waitFor(() => expect(modeSelect).toHaveFocus());
|
||||
expect(onGenerationModeChange).not.toHaveBeenCalled();
|
||||
});
|
||||
|
||||
it("supports choosing a mode with the keyboard", async () => {
|
||||
const onGenerationModeChange = vi.fn();
|
||||
const user = userEvent.setup();
|
||||
render(<ChatBar onGenerationModeChange={onGenerationModeChange} />);
|
||||
await user.click(screen.getByRole("textbox", { name: "Continuation prompt" }));
|
||||
await user.tab();
|
||||
expect(screen.getByRole("combobox", { name: "Generation mode" })).toHaveFocus();
|
||||
await user.keyboard("{ArrowDown}");
|
||||
await waitFor(() => expect(screen.getByRole("option", { name: "T2VA" })).toHaveFocus());
|
||||
await user.keyboard("{ArrowDown}");
|
||||
await waitFor(() => expect(screen.getByRole("option", { name: "FL2VA" })).toHaveFocus());
|
||||
await user.keyboard("{Enter}");
|
||||
expect(onGenerationModeChange).toHaveBeenCalledWith("fl2va");
|
||||
expect(screen.queryByRole("listbox")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
@@ -4,8 +4,16 @@ import React, { useRef, useState, useCallback, useEffect } from "react";
|
||||
import Image from "next/image";
|
||||
import { Film, ArrowUp, X, Loader2, ArrowLeft } from "lucide-react";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
|
||||
import LeaveSessionModal, { shouldShowLeaveWarning } from "@/components/LeaveSessionModal";
|
||||
import SpeechToTextButton from "@/components/SpeechToTextButton";
|
||||
import {
|
||||
DEFAULT_GENERATION_MODE,
|
||||
GENERATION_MODES,
|
||||
getGenerationMode,
|
||||
isGenerationMode,
|
||||
type GenerationMode,
|
||||
} from "@/lib/generationMode";
|
||||
import { cn } from "@/lib/utils";
|
||||
|
||||
const PROMPT_MAX_LENGTH = 500;
|
||||
@@ -22,6 +30,12 @@ interface Props {
|
||||
sessionNotice?: string;
|
||||
projectResetPending?: boolean;
|
||||
viewingReadOnly?: boolean;
|
||||
generationMode?: GenerationMode;
|
||||
supportedGenerationModes?: readonly GenerationMode[];
|
||||
generationInputsValid?: boolean;
|
||||
capabilityNotice?: string;
|
||||
mockRuntime?: boolean;
|
||||
conditioningPanel?: React.ReactNode;
|
||||
onPresetGenerate?: (presetId: string) => void;
|
||||
onContinuationInput?: (e: React.ChangeEvent<HTMLTextAreaElement>) => void;
|
||||
onContinuationKeydown?: (e: React.KeyboardEvent<HTMLTextAreaElement>) => void;
|
||||
@@ -30,6 +44,7 @@ interface Props {
|
||||
onLeave?: () => void;
|
||||
onStartNewProject?: () => void;
|
||||
onBackFromViewing?: () => void;
|
||||
onGenerationModeChange?: (mode: GenerationMode) => void;
|
||||
onSpeechTranscript?: (text: string) => void;
|
||||
onSpeechInterimChange?: (text: string) => void;
|
||||
}
|
||||
@@ -46,6 +61,12 @@ export default function ChatBar({
|
||||
sessionNotice = "",
|
||||
projectResetPending = false,
|
||||
viewingReadOnly = false,
|
||||
generationMode = DEFAULT_GENERATION_MODE,
|
||||
supportedGenerationModes = GENERATION_MODES.map((mode) => mode.id),
|
||||
generationInputsValid = true,
|
||||
capabilityNotice = "",
|
||||
mockRuntime = false,
|
||||
conditioningPanel,
|
||||
onPresetGenerate = () => {},
|
||||
onContinuationInput = () => {},
|
||||
onContinuationKeydown = () => {},
|
||||
@@ -54,6 +75,7 @@ export default function ChatBar({
|
||||
onLeave = () => {},
|
||||
onStartNewProject = () => {},
|
||||
onBackFromViewing = () => {},
|
||||
onGenerationModeChange = () => {},
|
||||
onSpeechTranscript,
|
||||
onSpeechInterimChange,
|
||||
}: Props) {
|
||||
@@ -69,6 +91,7 @@ export default function ChatBar({
|
||||
? "What video are you imagining?"
|
||||
: "What do you want to edit?";
|
||||
const actionLabel = !sessionStarted ? "Generate" : "Rewrite rollout";
|
||||
const selectedGenerationMode = getGenerationMode(generationMode);
|
||||
|
||||
const inputRef = useRef<HTMLTextAreaElement>(null);
|
||||
const scrollRef = useRef<HTMLDivElement>(null);
|
||||
@@ -239,7 +262,7 @@ export default function ChatBar({
|
||||
<div className="flex flex-col items-center gap-3 rounded-2xl border border-border bg-card/80 px-6 py-4 text-center shadow-md backdrop-blur-sm">
|
||||
<div className="flex flex-col gap-1">
|
||||
<p className="text-sm font-semibold text-foreground">View-only project</p>
|
||||
<p className="max-w-md text-xs text-muted-foreground">Project sessions are currently limited to 5 minutes. Start a new project to create more videos.</p>
|
||||
<p className="max-w-md text-xs text-muted-foreground">This saved project is available for playback. Start a new project to create more videos.</p>
|
||||
</div>
|
||||
<div className="mt-1 flex items-center gap-2">
|
||||
<Button onClick={onBackFromViewing} variant="outline" size="sm" className="gap-1.5 rounded-full px-4">
|
||||
@@ -261,7 +284,7 @@ export default function ChatBar({
|
||||
<div className="flex flex-col items-center gap-3 rounded-2xl border border-border bg-card/80 px-8 py-5 text-center shadow-md backdrop-blur-sm">
|
||||
<div className="flex flex-col gap-1">
|
||||
<p className="text-sm font-semibold text-foreground">Session ended</p>
|
||||
<p className="max-w-xs text-xs text-muted-foreground">Each project currently has a 5-minute session. Start a new project to continue creating videos.</p>
|
||||
<p className="max-w-xs text-xs text-muted-foreground">The runtime session has ended. Your saved videos remain available. Start a new project to continue creating.</p>
|
||||
</div>
|
||||
<div className="mt-1 flex items-center gap-2">
|
||||
<Button onClick={onStartNewProject} size="sm" className="rounded-full px-5">
|
||||
@@ -280,7 +303,7 @@ export default function ChatBar({
|
||||
|
||||
return (
|
||||
<section className="mx-auto flex w-full max-w-2xl shrink-0 flex-col gap-4">
|
||||
{storyPresets.length > 0 && !sessionStarted && (
|
||||
{storyPresets.length > 0 && !sessionStarted && generationMode === "t2va" && (
|
||||
<div className={cn("relative transition-opacity duration-200", isGenerating && "pointer-events-none opacity-40")}>
|
||||
<div
|
||||
ref={scrollRef}
|
||||
@@ -301,7 +324,7 @@ export default function ChatBar({
|
||||
<button
|
||||
key={preset.id}
|
||||
type="button"
|
||||
disabled={isGenerating}
|
||||
disabled={isBusy || !generationInputsValid}
|
||||
onClick={() => onPresetGenerate(preset.id)}
|
||||
className="flex flex-col sm:flex-row items-start gap-1.5 shrink-0 rounded-xl border p-2.5 text-left backdrop-blur-sm transition-colors max-w-42 sm:max-w-[215px] border-input bg-card/80 text-muted-foreground hover:bg-slate-200/60 hover:border-slate-400 hover:text-slate-700 dark:bg-slate-800/80 dark:text-slate-300 dark:hover:bg-slate-700/50 dark:hover:border-slate-500 dark:hover:text-slate-200"
|
||||
>
|
||||
@@ -327,6 +350,12 @@ export default function ChatBar({
|
||||
</div>
|
||||
)}
|
||||
|
||||
{mockRuntime && (
|
||||
<p role="status" className="rounded-xl border border-violet-500/25 bg-violet-500/10 px-4 py-2 text-center text-xs text-violet-700 dark:text-violet-300">
|
||||
Demo runtime · Sample playback only. No AI model is generating this video.
|
||||
</p>
|
||||
)}
|
||||
|
||||
{sessionNotice && (
|
||||
<div
|
||||
className={cn(
|
||||
@@ -346,9 +375,14 @@ export default function ChatBar({
|
||||
</div>
|
||||
)}
|
||||
|
||||
{sessionStarted && <p className="px-2 text-center text-[11px] text-muted-foreground">{selectedGenerationMode.label} · Mode and reference inputs are locked for this project.</p>}
|
||||
{conditioningPanel}
|
||||
|
||||
<div
|
||||
role="group"
|
||||
aria-label="Prompt composer"
|
||||
className={cn(
|
||||
"flex min-w-0 items-center gap-1.5 rounded-4xl border py-2.5 pl-5 pr-2.5 shadow-md backdrop-blur-sm transition-all duration-200",
|
||||
"flex min-w-0 flex-col gap-2 rounded-3xl border p-2.5 shadow-md backdrop-blur-sm transition-all duration-200",
|
||||
isBusy ? "border-input/60 bg-card/40" : "border-input bg-card/65",
|
||||
)}
|
||||
>
|
||||
@@ -364,40 +398,85 @@ export default function ChatBar({
|
||||
disabled={isBusy || sttBusy}
|
||||
rows={1}
|
||||
className={cn(
|
||||
"min-w-0 flex-1 resize-none bg-transparent text-foreground outline-none placeholder:text-muted-foreground transition-opacity duration-200 scrollbar-thin leading-snug",
|
||||
"w-full min-w-0 resize-none bg-transparent px-2 py-1 text-foreground outline-none placeholder:text-muted-foreground transition-opacity duration-200 scrollbar-thin leading-snug",
|
||||
(isBusy || sttBusy) && "cursor-not-allowed opacity-50",
|
||||
)}
|
||||
/>
|
||||
{onSpeechTranscript && <SpeechToTextButton disabled={isBusy} onTranscript={onSpeechTranscript} onInterimChange={onSpeechInterimChange} onBusyChange={setSttBusy} />}
|
||||
{!sessionStarted ? (
|
||||
<Button
|
||||
aria-label={actionLabel}
|
||||
title={actionLabel}
|
||||
onClick={onGenerate}
|
||||
disabled={!canJoinSession || isGenerating || !continuationDraft.trim()}
|
||||
size="icon-sm"
|
||||
className="shrink-0 rounded-full"
|
||||
>
|
||||
{showSpinner ? <Loader2 className="size-5 animate-spin" /> : <ArrowUp className="size-5" />}
|
||||
</Button>
|
||||
) : (
|
||||
<>
|
||||
<div className="flex min-w-0 items-center gap-1.5">
|
||||
{!sessionStarted && (
|
||||
<div className="flex shrink-0 items-center gap-1 pl-2">
|
||||
<label htmlFor="generation-mode" className="cursor-pointer text-xs font-medium text-muted-foreground">
|
||||
Mode
|
||||
</label>
|
||||
<Select
|
||||
value={generationMode}
|
||||
disabled={isBusy || sttBusy}
|
||||
onValueChange={(value) => {
|
||||
if (isGenerationMode(value)) onGenerationModeChange(value);
|
||||
}}
|
||||
>
|
||||
<SelectTrigger
|
||||
id="generation-mode"
|
||||
aria-label="Generation mode"
|
||||
title={`${selectedGenerationMode.name}. ${selectedGenerationMode.description}`}
|
||||
className="h-8 w-24 cursor-pointer rounded-lg border-0 bg-transparent px-2 py-1 text-xs font-medium shadow-none hover:bg-muted/60 data-[state=open]:bg-muted/80 [&>svg]:size-3 [&>svg]:transition-transform [&[data-state=open]>svg]:rotate-180"
|
||||
>
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent
|
||||
side="bottom"
|
||||
align="start"
|
||||
sideOffset={6}
|
||||
className="min-w-36 rounded-2xl border-input/70 bg-card/95 shadow-xl backdrop-blur-xl"
|
||||
>
|
||||
{GENERATION_MODES.map((mode) => (
|
||||
<SelectItem
|
||||
key={mode.id}
|
||||
value={mode.id}
|
||||
disabled={!supportedGenerationModes.includes(mode.id)}
|
||||
title={supportedGenerationModes.includes(mode.id) ? mode.name : `${mode.name} (unavailable on this runtime)`}
|
||||
className="cursor-pointer rounded-xl text-xs transition-colors data-[state=checked]:bg-muted/80 data-[state=checked]:font-semibold [&_svg]:size-3.5 [&_svg]:text-foreground"
|
||||
>
|
||||
{mode.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
</div>
|
||||
)}
|
||||
<div className="flex-1" />
|
||||
{onSpeechTranscript && <SpeechToTextButton disabled={isBusy} onTranscript={onSpeechTranscript} onInterimChange={onSpeechInterimChange} onBusyChange={setSttBusy} />}
|
||||
{!sessionStarted ? (
|
||||
<Button
|
||||
aria-label={actionLabel}
|
||||
title={actionLabel}
|
||||
onClick={onSubmitContinuation}
|
||||
disabled={!canSubmitContinuation || showSpinner || projectResetPending || !continuationDraft.trim()}
|
||||
onClick={onGenerate}
|
||||
disabled={!canJoinSession || isGenerating || !continuationDraft.trim()}
|
||||
size="icon-sm"
|
||||
className="shrink-0 rounded-full"
|
||||
>
|
||||
{showSpinner ? <Loader2 className="size-5 animate-spin" /> : <ArrowUp className="size-5" />}
|
||||
</Button>
|
||||
<Button variant="outline" aria-label="Leave" title="Leave" onClick={() => { if (shouldShowLeaveWarning()) setLeaveModalOpen(true); else onLeave(); }} disabled={isGenerating || projectResetPending} size="icon-sm" className="shrink-0 rounded-full">
|
||||
<X className="size-5" />
|
||||
</Button>
|
||||
</>
|
||||
)}
|
||||
) : (
|
||||
<>
|
||||
<Button
|
||||
aria-label={actionLabel}
|
||||
title={actionLabel}
|
||||
onClick={onSubmitContinuation}
|
||||
disabled={!canSubmitContinuation || showSpinner || projectResetPending || !continuationDraft.trim()}
|
||||
size="icon-sm"
|
||||
className="shrink-0 rounded-full"
|
||||
>
|
||||
{showSpinner ? <Loader2 className="size-5 animate-spin" /> : <ArrowUp className="size-5" />}
|
||||
</Button>
|
||||
<Button variant="outline" aria-label="Leave" title="Leave" onClick={() => { if (shouldShowLeaveWarning()) setLeaveModalOpen(true); else onLeave(); }} disabled={isGenerating || projectResetPending} size="icon-sm" className="shrink-0 rounded-full">
|
||||
<X className="size-5" />
|
||||
</Button>
|
||||
</>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
{!sessionStarted && capabilityNotice && <p className="px-2 text-[11px] text-amber-700 dark:text-amber-300">{capabilityNotice}</p>}
|
||||
<p className="px-2 text-center text-[11px] text-muted-foreground">
|
||||
LLM powered by{" "}
|
||||
<a
|
||||
|
||||
@@ -33,7 +33,7 @@ export default function SessionTimeoutModal({
|
||||
Session ended
|
||||
</h2>
|
||||
<p className="text-sm text-muted-foreground">
|
||||
This project hit the current 5-minute session limit. Your latest video stays on screen, and the project is being kept in the archive so you can come back to it.
|
||||
This project reached the runtime session limit. Your latest video stays on screen, and the project is being kept in the archive so you can come back to it.
|
||||
</p>
|
||||
</div>
|
||||
<p className="text-sm text-muted-foreground">
|
||||
|
||||
@@ -17,6 +17,13 @@ import {
|
||||
SelectValue,
|
||||
} from '@/components/ui/select';
|
||||
import { Textarea } from '@/components/ui/textarea';
|
||||
import {
|
||||
DEFAULT_GENERATION_MODE,
|
||||
GENERATION_MODES,
|
||||
getGenerationMode,
|
||||
isGenerationMode,
|
||||
type GenerationMode,
|
||||
} from '@/lib/generationMode';
|
||||
|
||||
interface DevtoolsComposerProps {
|
||||
connected?: boolean;
|
||||
@@ -37,6 +44,10 @@ interface DevtoolsComposerProps {
|
||||
loopGenerationEnabled?: boolean;
|
||||
curatedPromptLimit?: number;
|
||||
maxCuratedPromptCount?: number;
|
||||
generationMode?: GenerationMode;
|
||||
supportedGenerationModes?: readonly GenerationMode[];
|
||||
conditioningPanel?: React.ReactNode;
|
||||
onGenerationModeChange?: (mode: GenerationMode) => void;
|
||||
rewriteWindowMode?: boolean;
|
||||
rewritingSeedPrompts?: boolean;
|
||||
autoExtensionTimeoutHint?: string;
|
||||
@@ -74,6 +85,10 @@ export default function DevtoolsComposer({
|
||||
loopGenerationEnabled = false,
|
||||
curatedPromptLimit = 0,
|
||||
maxCuratedPromptCount = 0,
|
||||
generationMode = DEFAULT_GENERATION_MODE,
|
||||
supportedGenerationModes = GENERATION_MODES.map((mode) => mode.id),
|
||||
conditioningPanel,
|
||||
onGenerationModeChange = () => {},
|
||||
rewriteWindowMode = false,
|
||||
rewritingSeedPrompts = false,
|
||||
autoExtensionTimeoutHint = '',
|
||||
@@ -92,6 +107,7 @@ export default function DevtoolsComposer({
|
||||
onSpeechInterimChange,
|
||||
}: DevtoolsComposerProps) {
|
||||
const [sttBusy, setSttBusy] = useState(false);
|
||||
const selectedGenerationMode = getGenerationMode(generationMode);
|
||||
const submitButtonLabel = useMemo(
|
||||
() =>
|
||||
rewriteWindowMode
|
||||
@@ -110,7 +126,8 @@ export default function DevtoolsComposer({
|
||||
);
|
||||
|
||||
return (
|
||||
<section>
|
||||
<section className="space-y-4">
|
||||
{conditioningPanel}
|
||||
<Card>
|
||||
<CardContent className="space-y-5 p-5">
|
||||
<div className="grid gap-5 xl:grid-cols-[minmax(0,1fr)_320px]">
|
||||
@@ -285,6 +302,48 @@ export default function DevtoolsComposer({
|
||||
</div>
|
||||
|
||||
<div className="space-y-4">
|
||||
<div className="space-y-2">
|
||||
<Label htmlFor="devtools-generation-mode">
|
||||
Generation mode
|
||||
</Label>
|
||||
<Select
|
||||
value={generationMode}
|
||||
disabled={sessionStarted}
|
||||
onValueChange={(value) => {
|
||||
if (isGenerationMode(value)) {
|
||||
onGenerationModeChange(value);
|
||||
}
|
||||
}}
|
||||
>
|
||||
<SelectTrigger
|
||||
id="devtools-generation-mode"
|
||||
aria-label="Generation mode"
|
||||
title={selectedGenerationMode.name}
|
||||
>
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{GENERATION_MODES.map((mode) => (
|
||||
<SelectItem
|
||||
key={mode.id}
|
||||
value={mode.id}
|
||||
disabled={!supportedGenerationModes.includes(mode.id)}
|
||||
title={
|
||||
supportedGenerationModes.includes(mode.id)
|
||||
? mode.name
|
||||
: `${mode.name} (unavailable on this runtime)`
|
||||
}
|
||||
>
|
||||
{mode.label}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
<p className="text-sm text-muted-foreground">
|
||||
{selectedGenerationMode.description}
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<div className="flex items-start gap-3">
|
||||
<Checkbox
|
||||
id="devtools-enhance-prompts"
|
||||
|
||||
@@ -35,6 +35,7 @@ describe('DevtoolsShell', () => {
|
||||
expect(screen.getByText('Devtools Mode')).toBeInTheDocument();
|
||||
expect(screen.getByText('Your video will appear here')).toBeInTheDocument();
|
||||
expect(screen.getByLabelText('Story preset')).toBeInTheDocument();
|
||||
expect(screen.getByLabelText('Generation mode')).toBeInTheDocument();
|
||||
expect(screen.getByLabelText('Continuation prompt')).toBeDisabled();
|
||||
|
||||
expect(screen.getByText('Advanced controls')).toBeInTheDocument();
|
||||
|
||||
@@ -9,6 +9,7 @@ import VideoPlayer from '../VideoPlayer';
|
||||
import RewriteInspector from '../rewrite/RewriteInspector';
|
||||
import DevtoolsComposer from './DevtoolsComposer';
|
||||
import DevtoolsDrawer from './DevtoolsDrawer';
|
||||
import { DEFAULT_GENERATION_MODE, GENERATION_MODES, type GenerationMode } from '@/lib/generationMode';
|
||||
|
||||
interface DevtoolsShellProps {
|
||||
connected?: boolean;
|
||||
@@ -30,6 +31,11 @@ interface DevtoolsShellProps {
|
||||
curatedPromptLimit?: number;
|
||||
maxCuratedPromptCount?: number;
|
||||
|
||||
generationMode?: GenerationMode;
|
||||
supportedGenerationModes?: readonly GenerationMode[];
|
||||
conditioningPanel?: React.ReactNode;
|
||||
onGenerationModeChange?: (mode: GenerationMode) => void;
|
||||
|
||||
onPresetChange?: (e: React.ChangeEvent<HTMLSelectElement>) => void;
|
||||
onEnhancementToggle?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
onCuratedPromptLimitChange?: (e: React.ChangeEvent<HTMLInputElement>) => void;
|
||||
@@ -137,6 +143,11 @@ export default function DevtoolsShell({
|
||||
curatedPromptLimit = 0,
|
||||
maxCuratedPromptCount = 0,
|
||||
|
||||
generationMode = DEFAULT_GENERATION_MODE,
|
||||
supportedGenerationModes = GENERATION_MODES.map((mode) => mode.id),
|
||||
conditioningPanel,
|
||||
onGenerationModeChange = () => {},
|
||||
|
||||
onPresetChange = () => {},
|
||||
onEnhancementToggle = () => {},
|
||||
onCuratedPromptLimitChange = () => {},
|
||||
@@ -287,6 +298,10 @@ export default function DevtoolsShell({
|
||||
loopGenerationEnabled={loopGenerationEnabled}
|
||||
curatedPromptLimit={curatedPromptLimit}
|
||||
maxCuratedPromptCount={maxCuratedPromptCount}
|
||||
generationMode={generationMode}
|
||||
supportedGenerationModes={supportedGenerationModes}
|
||||
conditioningPanel={conditioningPanel}
|
||||
onGenerationModeChange={onGenerationModeChange}
|
||||
rewriteWindowMode={livePromptRewriteMode}
|
||||
rewritingSeedPrompts={rewritingSeedPrompts}
|
||||
autoExtensionTimeoutHint={autoExtensionTimeoutHint}
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
import { act, renderHook, waitFor } from "@testing-library/react";
|
||||
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { useAssetLibrary } from "./useAssetLibrary";
|
||||
|
||||
const image = { asset_id: "asset-1", kind: "image", name: "frame.png", mime_type: "image/png", size: 5, url: "/assets/asset-1" };
|
||||
|
||||
describe("asset library lifecycle", () => {
|
||||
beforeEach(() => localStorage.clear());
|
||||
afterEach(() => vi.unstubAllGlobals());
|
||||
it("uploads the raw file and keeps assignment state separate from uploaded assets", async () => {
|
||||
const fetchMock = vi.fn(async () => ({ ok: true, json: async () => image }));
|
||||
vi.stubGlobal("fetch", fetchMock);
|
||||
const { result } = renderHook(() => useAssetLibrary());
|
||||
const file = new File(["image"], "first frame.png", { type: "image/png" });
|
||||
await act(async () => { await result.current.uploadAssets([file]); });
|
||||
expect(fetchMock).toHaveBeenCalledWith("/assets", expect.objectContaining({ body: file, headers: { "Content-Type": "image/png", "X-Asset-Name": "first%20frame.png" } }));
|
||||
act(() => result.current.assignAsset("asset-1", "first_frame"));
|
||||
expect(result.current.conditioningAssets).toEqual([{ asset_id: "asset-1", role: "first_frame" }]);
|
||||
act(() => result.current.clearConditioning());
|
||||
expect(result.current.assets).toHaveLength(1);
|
||||
expect(result.current.conditioningAssets).toEqual([]);
|
||||
});
|
||||
it("identifies stale server assets without downloading their contents", async () => {
|
||||
localStorage.setItem("dreamverse-asset-library-v1", JSON.stringify([image]));
|
||||
const fetchMock = vi.fn(async () => ({ ok: false, status: 404 }));
|
||||
vi.stubGlobal("fetch", fetchMock);
|
||||
const { result } = renderHook(() => useAssetLibrary());
|
||||
await waitFor(() => expect(result.current.assets).toHaveLength(1));
|
||||
act(() => result.current.assignAsset("asset-1", "first_frame"));
|
||||
let message: string | null = null;
|
||||
await act(async () => { message = await result.current.verifySelectedAssets(); });
|
||||
expect(message).toMatch(/Upload it again/);
|
||||
expect(result.current.assets[0].missing).toBe(true);
|
||||
expect(fetchMock).toHaveBeenCalledWith("/assets/asset-1", expect.objectContaining({ method: "HEAD" }));
|
||||
});
|
||||
it("deletes both the library entry and its selected references", async () => {
|
||||
localStorage.setItem("dreamverse-asset-library-v1", JSON.stringify([image]));
|
||||
vi.stubGlobal("fetch", vi.fn(async () => ({ ok: true })));
|
||||
const { result } = renderHook(() => useAssetLibrary());
|
||||
await waitFor(() => expect(result.current.assets).toHaveLength(1));
|
||||
act(() => result.current.assignAsset("asset-1", "reference"));
|
||||
await act(async () => { await result.current.removeAsset("asset-1"); });
|
||||
expect(result.current.assets).toEqual([]);
|
||||
expect(result.current.conditioningAssets).toEqual([]);
|
||||
});
|
||||
|
||||
it("does not mark a valid upload missing when the browser cannot preview its codec", async () => {
|
||||
localStorage.setItem("dreamverse-asset-library-v1", JSON.stringify([image]));
|
||||
vi.stubGlobal("fetch", vi.fn(async () => ({ ok: true, status: 200 })));
|
||||
const { result } = renderHook(() => useAssetLibrary());
|
||||
await waitFor(() => expect(result.current.assets).toHaveLength(1));
|
||||
await act(async () => { await result.current.checkAssetAvailability("asset-1"); });
|
||||
expect(result.current.assets[0].missing).not.toBe(true);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,136 @@
|
||||
"use client";
|
||||
|
||||
import { useCallback, useEffect, useRef, useState } from "react";
|
||||
import type { ConditioningAsset, ConditioningRole, GenerationAsset } from "@/lib/generationMode";
|
||||
|
||||
const LIBRARY_KEY = "dreamverse-asset-library-v1";
|
||||
const MAX_UPLOAD_BYTES = 100 * 1024 * 1024;
|
||||
|
||||
async function responseError(response: Response, fallback: string): Promise<Error> {
|
||||
const payload = await response.json().catch(() => ({}));
|
||||
return new Error(typeof payload.detail === "string" ? payload.detail : fallback);
|
||||
}
|
||||
|
||||
function isAsset(value: unknown): value is GenerationAsset {
|
||||
if (!value || typeof value !== "object") return false;
|
||||
const asset = value as Partial<GenerationAsset>;
|
||||
return typeof asset.asset_id === "string" && /^[a-zA-Z0-9_-]+$/.test(asset.asset_id)
|
||||
&& ["image", "video", "audio"].includes(asset.kind || "")
|
||||
&& typeof asset.name === "string" && typeof asset.mime_type === "string" && typeof asset.size === "number";
|
||||
}
|
||||
|
||||
/** Asset ownership lives here so composers and other pickers share the same library. */
|
||||
export function useAssetLibrary() {
|
||||
const [assets, setAssets] = useState<GenerationAsset[]>([]);
|
||||
const [conditioningAssets, setConditioningAssets] = useState<ConditioningAsset[]>([]);
|
||||
const [uploading, setUploading] = useState(false);
|
||||
const [assetError, setAssetError] = useState("");
|
||||
const [hydrated, setHydrated] = useState(false);
|
||||
const uploadingRef = useRef(false);
|
||||
|
||||
useEffect(() => {
|
||||
try {
|
||||
const saved = JSON.parse(localStorage.getItem(LIBRARY_KEY) || "[]");
|
||||
if (Array.isArray(saved)) setAssets(saved.filter(isAsset).map((asset) => ({
|
||||
...asset, url: `/assets/${asset.asset_id}`,
|
||||
})));
|
||||
} catch { /* Storage is optional; uploads still work in private browsing. */ }
|
||||
setHydrated(true);
|
||||
}, []);
|
||||
useEffect(() => {
|
||||
if (!hydrated) return;
|
||||
try { localStorage.setItem(LIBRARY_KEY, JSON.stringify(assets)); } catch { /* Optional cache. */ }
|
||||
}, [assets, hydrated]);
|
||||
|
||||
const uploadAssets = useCallback(async (files: File[]) => {
|
||||
if (uploadingRef.current) return;
|
||||
uploadingRef.current = true;
|
||||
setUploading(true);
|
||||
setAssetError("");
|
||||
const errors: string[] = [];
|
||||
for (const file of files) {
|
||||
try {
|
||||
if (!/^(image|video|audio)\//.test(file.type)) throw new Error(`${file.name}: choose an image, video, or audio file.`);
|
||||
const maxBytes = file.type.startsWith("image/") ? 15 * 1024 * 1024 : MAX_UPLOAD_BYTES;
|
||||
if (!file.size || file.size > maxBytes) throw new Error(`${file.name}: use a non-empty file up to ${maxBytes / 1024 / 1024} MiB.`);
|
||||
const response = await fetch("/assets", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": file.type, "X-Asset-Name": encodeURIComponent(file.name) },
|
||||
body: file,
|
||||
});
|
||||
if (!response.ok) throw await responseError(response, `Could not upload ${file.name}.`);
|
||||
const asset: unknown = await response.json();
|
||||
if (!isAsset(asset)) throw new Error("The server returned an invalid asset. Please retry the upload.");
|
||||
setAssets((current) => [...current.filter((item) => item.asset_id !== asset.asset_id), {
|
||||
...asset, url: `/assets/${asset.asset_id}`, missing: false,
|
||||
}]);
|
||||
} catch (error) {
|
||||
errors.push(error instanceof Error ? error.message : `Could not upload ${file.name}.`);
|
||||
}
|
||||
}
|
||||
setAssetError(errors.join(" "));
|
||||
setUploading(false);
|
||||
uploadingRef.current = false;
|
||||
}, []);
|
||||
|
||||
const assignAsset = useCallback((assetId: string, role: ConditioningRole) => {
|
||||
setConditioningAssets((current) => {
|
||||
const next = role === "reference" ? current : current.filter((item) => item.role !== role);
|
||||
if (!assetId || next.some((item) => item.asset_id === assetId && item.role === role)) return next;
|
||||
return [...next, { asset_id: assetId, role }];
|
||||
});
|
||||
}, []);
|
||||
const removeConditioning = useCallback((index: number) => {
|
||||
setConditioningAssets((current) => current.filter((_, itemIndex) => itemIndex !== index));
|
||||
}, []);
|
||||
const moveConditioning = useCallback((from: number, to: number) => {
|
||||
setConditioningAssets((current) => {
|
||||
if (from < 0 || to < 0 || from >= current.length || to >= current.length) return current;
|
||||
const next = [...current];
|
||||
next.splice(to, 0, next.splice(from, 1)[0]);
|
||||
return next;
|
||||
});
|
||||
}, []);
|
||||
const clearConditioning = useCallback(() => setConditioningAssets([]), []);
|
||||
const markAssetMissing = useCallback((assetId: string) => {
|
||||
setAssets((current) => current.map((asset) => asset.asset_id === assetId ? { ...asset, missing: true } : asset));
|
||||
}, []);
|
||||
const checkAssetAvailability = useCallback(async (assetId: string) => {
|
||||
try {
|
||||
const response = await fetch(`/assets/${assetId}`, { method: "HEAD", signal: AbortSignal.timeout(4000) });
|
||||
if (response.status === 404) markAssetMissing(assetId);
|
||||
} catch { /* A browser preview failure alone does not mean the upload expired. */ }
|
||||
}, [markAssetMissing]);
|
||||
const removeAsset = useCallback(async (assetId: string) => {
|
||||
setAssetError("");
|
||||
try {
|
||||
const response = await fetch(`/assets/${assetId}`, { method: "DELETE" });
|
||||
if (!response.ok && response.status !== 404) throw await responseError(response, "Could not remove the asset. Retry when the backend is available.");
|
||||
setAssets((current) => current.filter((asset) => asset.asset_id !== assetId));
|
||||
setConditioningAssets((current) => current.filter((asset) => asset.asset_id !== assetId));
|
||||
} catch (error) {
|
||||
setAssetError(error instanceof Error ? error.message : "Could not remove the asset.");
|
||||
}
|
||||
}, []);
|
||||
const verifySelectedAssets = useCallback(async (): Promise<string | null> => {
|
||||
const selected = [...new Set(conditioningAssets.map((item) => item.asset_id))];
|
||||
try {
|
||||
for (const assetId of selected) {
|
||||
const response = await fetch(`/assets/${assetId}`, { method: "HEAD", signal: AbortSignal.timeout(4000) });
|
||||
if (response.status === 404) {
|
||||
markAssetMissing(assetId);
|
||||
return "A selected asset expired or was removed from the server. Upload it again, then select the new copy.";
|
||||
}
|
||||
if (!response.ok) return "Could not verify the selected assets. Check the backend and try again.";
|
||||
}
|
||||
return null;
|
||||
} catch {
|
||||
return "Could not verify the selected assets. Check the backend and try again.";
|
||||
}
|
||||
}, [conditioningAssets, markAssetMissing]);
|
||||
|
||||
return {
|
||||
assets, conditioningAssets, uploading, assetError, uploadAssets, assignAsset,
|
||||
removeConditioning, moveConditioning, clearConditioning, removeAsset, checkAssetAvailability, verifySelectedAssets,
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
import { renderHook, waitFor } from "@testing-library/react";
|
||||
import { afterEach, describe, expect, it, vi } from "vitest";
|
||||
import { useGenerationCapabilities } from "./useGenerationCapabilities";
|
||||
|
||||
describe("generation capabilities", () => {
|
||||
afterEach(() => vi.unstubAllGlobals());
|
||||
it("enables only the modes the runtime advertises and identifies mock playback", async () => {
|
||||
vi.stubGlobal("fetch", vi.fn(async () => ({ ok: true, json: async () => ({ model_id: "mock", modes: ["t2va", "fl2va", "ref2va"], mock: true }) })));
|
||||
const { result } = renderHook(() => useGenerationCapabilities());
|
||||
await waitFor(() => expect(result.current.loadingCapabilities).toBe(false));
|
||||
expect(result.current.capabilities.modes).toEqual(["t2va", "fl2va", "ref2va"]);
|
||||
expect(result.current.capabilities.mock).toBe(true);
|
||||
});
|
||||
it("keeps old runtimes text-only when the capabilities endpoint is missing", async () => {
|
||||
vi.stubGlobal("fetch", vi.fn(async () => ({ ok: false, status: 404 })));
|
||||
const { result } = renderHook(() => useGenerationCapabilities());
|
||||
await waitFor(() => expect(result.current.loadingCapabilities).toBe(false));
|
||||
expect(result.current.capabilities.modes).toEqual(["t2va"]);
|
||||
expect(result.current.capabilityNotice).toMatch(/Text-only compatibility/);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,37 @@
|
||||
"use client";
|
||||
|
||||
import { useCallback, useEffect, useState } from "react";
|
||||
import { isGenerationMode, type GenerationCapabilities } from "@/lib/generationMode";
|
||||
|
||||
const LEGACY_CAPABILITIES: GenerationCapabilities = { model_id: "legacy", modes: ["t2va"] };
|
||||
|
||||
export function useGenerationCapabilities() {
|
||||
const [capabilities, setCapabilities] = useState<GenerationCapabilities>(LEGACY_CAPABILITIES);
|
||||
const [capabilityNotice, setCapabilityNotice] = useState("");
|
||||
const [loadingCapabilities, setLoadingCapabilities] = useState(true);
|
||||
const refreshCapabilities = useCallback(async () => {
|
||||
try {
|
||||
const response = await fetch("/generation-capabilities", { signal: AbortSignal.timeout(4000) });
|
||||
if (!response.ok) throw new Error("Capabilities unavailable");
|
||||
const payload = await response.json();
|
||||
if (typeof payload.model_id !== "string" || !Array.isArray(payload.modes)
|
||||
|| !payload.modes.every(isGenerationMode)) throw new Error("Invalid capabilities");
|
||||
const next: GenerationCapabilities = {
|
||||
model_id: payload.model_id,
|
||||
modes: payload.modes,
|
||||
mock: payload.mock === true,
|
||||
};
|
||||
setCapabilities(next);
|
||||
setCapabilityNotice("");
|
||||
return next;
|
||||
} catch {
|
||||
setCapabilities(LEGACY_CAPABILITIES);
|
||||
setCapabilityNotice("Runtime capabilities unavailable. Text-only compatibility mode is available; check the backend to enable image and reference modes.");
|
||||
return LEGACY_CAPABILITIES;
|
||||
} finally {
|
||||
setLoadingCapabilities(false);
|
||||
}
|
||||
}, []);
|
||||
useEffect(() => { void refreshCapabilities(); }, [refreshCapabilities]);
|
||||
return { capabilities, capabilityNotice, loadingCapabilities, refreshCapabilities };
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
import { describe, expect, it } from "vitest";
|
||||
|
||||
import {
|
||||
DEFAULT_GENERATION_MODE,
|
||||
GENERATION_MODES,
|
||||
getGenerationMode,
|
||||
isGenerationMode,
|
||||
buildGenerationInitFields,
|
||||
validateGenerationInputs,
|
||||
type GenerationAsset,
|
||||
} from "./generationMode";
|
||||
|
||||
describe("generation modes", () => {
|
||||
it("exposes stable wire IDs in the expected product order", () => {
|
||||
expect(GENERATION_MODES.map((mode) => mode.id)).toEqual([
|
||||
"t2va",
|
||||
"fl2va",
|
||||
"ref2va",
|
||||
]);
|
||||
expect(DEFAULT_GENERATION_MODE).toBe("t2va");
|
||||
});
|
||||
|
||||
it("validates and resolves generation mode values", () => {
|
||||
expect(isGenerationMode("ref2va")).toBe(true);
|
||||
expect(isGenerationMode("unknown")).toBe(false);
|
||||
expect(getGenerationMode("fl2va").label).toBe("FL2VA");
|
||||
});
|
||||
});
|
||||
|
||||
const image: GenerationAsset = { asset_id: "img", kind: "image", name: "frame.png", mime_type: "image/png", size: 12, url: "/assets/img" };
|
||||
const audio: GenerationAsset = { asset_id: "sound", kind: "audio", name: "sound.wav", mime_type: "audio/wav", size: 12, url: "/assets/sound" };
|
||||
|
||||
describe("generation input contract", () => {
|
||||
it("keeps text-only init valid and rejects accidental references", () => {
|
||||
expect(buildGenerationInitFields("t2va", [], [])).toEqual({ generation_mode: "t2va", conditioning_assets: [] });
|
||||
expect(() => buildGenerationInitFields("t2va", [{ asset_id: "img", role: "reference" }], [image])).toThrow("text only");
|
||||
});
|
||||
it("requires first frame but permits first-only or both endpoints", () => {
|
||||
expect(validateGenerationInputs("fl2va", [], [])).toMatch(/first frame/);
|
||||
const first = { asset_id: "img", role: "first_frame" } as const;
|
||||
expect(validateGenerationInputs("fl2va", [first], [image])).toBeNull();
|
||||
expect(validateGenerationInputs("fl2va", [first, { asset_id: "img", role: "last_frame" }], [image])).toBeNull();
|
||||
expect(validateGenerationInputs("fl2va", [first, first], [image])).toMatch(/one first frame/);
|
||||
expect(validateGenerationInputs("fl2va", [{ asset_id: "sound", role: "first_frame" }], [audio])).toMatch(/images only/);
|
||||
});
|
||||
it("requires a visual reference and preserves multimodal ordering without file bodies", () => {
|
||||
expect(validateGenerationInputs("ref2va", [{ asset_id: "sound", role: "reference" }], [audio])).toMatch(/image or video/);
|
||||
const items = [{ asset_id: "sound", role: "reference" }, { asset_id: "img", role: "reference" }] as const;
|
||||
expect(buildGenerationInitFields("ref2va", items, [image, audio])).toEqual({ generation_mode: "ref2va", conditioning_assets: items });
|
||||
});
|
||||
it("rejects per-kind limits, total limits, and stale uploads", () => {
|
||||
const images = Array.from({ length: 10 }, (_, index) => ({ ...image, asset_id: `img-${index}` }));
|
||||
expect(validateGenerationInputs("ref2va", images.map((item) => ({ asset_id: item.asset_id, role: "reference" })), images)).toMatch(/at most 9 image/);
|
||||
const refs = Array.from({ length: 13 }, () => ({ asset_id: "img", role: "reference" as const }));
|
||||
expect(validateGenerationInputs("ref2va", refs, [image])).toMatch(/at most 12/);
|
||||
expect(validateGenerationInputs("fl2va", [{ asset_id: "img", role: "first_frame" }], [{ ...image, missing: true }])).toMatch(/no longer on the server/);
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,114 @@
|
||||
export const GENERATION_MODES = [
|
||||
{
|
||||
id: "t2va",
|
||||
label: "T2VA",
|
||||
name: "Text to video + audio",
|
||||
description: "Start with a text prompt; no reference asset is required.",
|
||||
},
|
||||
{
|
||||
id: "fl2va",
|
||||
label: "FL2VA",
|
||||
name: "First/last frames to video + audio",
|
||||
description: "Start from a first frame image. Add an optional last frame to guide the ending.",
|
||||
},
|
||||
{
|
||||
id: "ref2va",
|
||||
label: "Ref2VA",
|
||||
name: "References to video + audio",
|
||||
description: "Guide the result with ordered image, video, or audio references.",
|
||||
},
|
||||
] as const;
|
||||
|
||||
export type GenerationMode = (typeof GENERATION_MODES)[number]["id"];
|
||||
|
||||
export const DEFAULT_GENERATION_MODE: GenerationMode = "t2va";
|
||||
|
||||
export function isGenerationMode(value: unknown): value is GenerationMode {
|
||||
return GENERATION_MODES.some((mode) => mode.id === value);
|
||||
}
|
||||
|
||||
export function getGenerationMode(value: GenerationMode) {
|
||||
return GENERATION_MODES.find((mode) => mode.id === value) ?? GENERATION_MODES[0];
|
||||
}
|
||||
|
||||
export type AssetKind = "image" | "video" | "audio";
|
||||
export type ConditioningRole = "first_frame" | "last_frame" | "reference";
|
||||
|
||||
/** Runtime-owned uploads; project metadata keeps references, never file contents. */
|
||||
export interface GenerationAsset {
|
||||
asset_id: string;
|
||||
kind: AssetKind;
|
||||
name: string;
|
||||
mime_type: string;
|
||||
size: number;
|
||||
url: string;
|
||||
missing?: boolean;
|
||||
}
|
||||
|
||||
export interface ConditioningAsset {
|
||||
asset_id: string;
|
||||
role: ConditioningRole;
|
||||
}
|
||||
|
||||
export interface GenerationCapabilities {
|
||||
model_id: string;
|
||||
modes: GenerationMode[];
|
||||
mock?: boolean;
|
||||
}
|
||||
|
||||
export interface GenerationInitFields {
|
||||
generation_mode: GenerationMode;
|
||||
conditioning_assets: ConditioningAsset[];
|
||||
}
|
||||
|
||||
export const REFERENCE_LIMITS = { image: 9, video: 3, audio: 3, total: 12 } as const;
|
||||
|
||||
export function validateGenerationInputs(
|
||||
mode: GenerationMode,
|
||||
conditioning: readonly ConditioningAsset[],
|
||||
assets: readonly GenerationAsset[],
|
||||
): string | null {
|
||||
if (mode === "t2va") {
|
||||
return conditioning.length ? "T2VA uses text only. Remove the selected references." : null;
|
||||
}
|
||||
const resolved = conditioning.map((item) => assets.find((asset) => asset.asset_id === item.asset_id));
|
||||
if (resolved.some((asset) => !asset || asset.missing)) {
|
||||
return "A selected asset is no longer on the server. Upload it again and select the new copy.";
|
||||
}
|
||||
if (mode === "fl2va") {
|
||||
if (!conditioning.some((item) => item.role === "first_frame")) return "Choose a first frame image to generate.";
|
||||
if (conditioning.some((item) => item.role === "reference") || resolved.some((asset) => asset?.kind !== "image")) {
|
||||
return "FL2VA accepts first and last frame images only.";
|
||||
}
|
||||
if (conditioning.filter((item) => item.role === "first_frame").length !== 1
|
||||
|| conditioning.filter((item) => item.role === "last_frame").length > 1) {
|
||||
return "Choose one first frame and at most one last frame.";
|
||||
}
|
||||
return null;
|
||||
}
|
||||
if (conditioning.some((item) => item.role !== "reference")) return "Ref2VA accepts ordered reference assets only.";
|
||||
if (!resolved.some((asset) => asset?.kind === "image" || asset?.kind === "video")) {
|
||||
return "Add at least one image or video reference. Audio alone is not enough.";
|
||||
}
|
||||
if (conditioning.length > REFERENCE_LIMITS.total) return "Use at most 12 reference assets in total.";
|
||||
for (const kind of ["image", "video", "audio"] as const) {
|
||||
if (resolved.filter((asset) => asset?.kind === kind).length > REFERENCE_LIMITS[kind]) {
|
||||
return `Use at most ${REFERENCE_LIMITS[kind]} ${kind} references.`;
|
||||
}
|
||||
}
|
||||
return null;
|
||||
}
|
||||
|
||||
/** Shared by both session_init_v2 and project_init_v1. */
|
||||
export function buildGenerationInitFields(
|
||||
mode: GenerationMode,
|
||||
conditioning: readonly ConditioningAsset[],
|
||||
assets: readonly GenerationAsset[],
|
||||
): GenerationInitFields {
|
||||
const problem = validateGenerationInputs(mode, conditioning, assets);
|
||||
if (problem) throw new Error(problem);
|
||||
return {
|
||||
generation_mode: mode,
|
||||
conditioning_assets: conditioning.map(({ asset_id, role }) => ({ asset_id, role })),
|
||||
};
|
||||
}
|
||||
@@ -1,3 +1,5 @@
|
||||
import type { ConditioningAsset, GenerationAsset, GenerationMode } from "./generationMode";
|
||||
|
||||
const DB_NAME = "fastvideo-projects";
|
||||
const DB_VERSION = 1;
|
||||
const PROJECTS_STORE = "projects";
|
||||
@@ -35,6 +37,11 @@ export interface StoredProject {
|
||||
createdAt: number;
|
||||
lastThumbnail: string | null;
|
||||
promptEvents: Record<string, unknown>[];
|
||||
/** Optional for projects created before generation modes were introduced. */
|
||||
generationMode?: GenerationMode;
|
||||
conditioningAssets?: ConditioningAsset[];
|
||||
assets?: GenerationAsset[];
|
||||
mock?: boolean;
|
||||
}
|
||||
|
||||
export interface StoredClip {
|
||||
|
||||
@@ -0,0 +1,290 @@
|
||||
# Infinite Livestream
|
||||
|
||||
Infinite Livestream is a chat-driven FastH3 broadcast. Viewers type prompts into a web
|
||||
page, the app rewrites them with an LLM, generates clips with FastVideo, and
|
||||
plays them back as one continuous HLS stream on that same page. When nobody is
|
||||
typing it feeds itself from a preset of idle prompts, so the channel never goes
|
||||
dark.
|
||||
|
||||
It lives in this monorepo under `apps/infinite_livestream/`.
|
||||
|
||||
```
|
||||
chat -> Director -> PromptUpsampler (OpenAI-compatible LLM)
|
||||
|
|
||||
v enqueue / move / pop
|
||||
Engine -> FastH3Backend -> FastVideo
|
||||
| frames + audio
|
||||
v
|
||||
Pacer -> HlsSink -> the page's <video>
|
||||
```
|
||||
|
||||
Everything runs in a single process, and the page, the playlist and the chat
|
||||
endpoint are served from one HTTP origin, so publishing the stream means
|
||||
pointing a tunnel or reverse proxy at one port.
|
||||
|
||||
## Requirements
|
||||
|
||||
- Linux with NVIDIA GPUs. The [default configuration](infinite_livestream/configs/infinite_livestream.yaml)
|
||||
targets four GB200 GPUs: Blackwell `sm_100a` sparse attention, a replicated
|
||||
transformer, and GPU-resident text encoder and VAEs. A different GPU setup
|
||||
needs corresponding changes to `runtime` and `inference`; the GPU count must
|
||||
divide the model's attention head count.
|
||||
- Python 3.12 and [uv](https://docs.astral.sh/uv/getting-started/installation/).
|
||||
- A CUDA 13 toolkit with `nvcc` and a compatible C++ compiler for the kernel
|
||||
source build below.
|
||||
- A complete FastH3 checkpoint; see [Download weights](#download-weights).
|
||||
- An API key for prompt rewriting. The default configuration uses OpenAI;
|
||||
rewriting runs for idle filler too, so the stream needs the key even when
|
||||
nobody is chatting.
|
||||
- FFmpeg with the `libx264` and `aac` encoders on `PATH`; see [FFmpeg](#ffmpeg).
|
||||
|
||||
## Install
|
||||
|
||||
Use a source checkout containing `apps/infinite_livestream/`. Run these commands
|
||||
from the FastVideo repository root. The app is packaged in FastVideo's
|
||||
`infinite-livestream` extra; `fasth3` adds its generator dependencies.
|
||||
|
||||
```bash
|
||||
uv venv --python 3.12 --seed
|
||||
source .venv/bin/activate
|
||||
|
||||
git submodule update --init --recursive \
|
||||
fastvideo-kernel/include/cutlass fastvideo-kernel/include/tk
|
||||
|
||||
# Point this at your CUDA 13 toolkit.
|
||||
export CUDA_HOME=/usr/local/cuda
|
||||
export CUDACXX="$CUDA_HOME/bin/nvcc"
|
||||
TORCH_CUDA_ARCH_LIST=10.0a UV_TORCH_BACKEND=cu130 \
|
||||
uv pip install -e ".[fasth3,infinite-livestream]"
|
||||
```
|
||||
|
||||
This installs the pinned FA4 package and builds this checkout's
|
||||
`fastvideo-kernel` with the Blackwell VSA extension. The CUDA and architecture
|
||||
settings above match the default GB200 configuration. See the
|
||||
[kernel build guide](../../fastvideo-kernel/README.md#installation) for build
|
||||
prerequisites and troubleshooting.
|
||||
|
||||
## Download weights
|
||||
|
||||
Use [FastH3 v1 VSA-DataFree](https://huggingface.co/FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree),
|
||||
FastVideo's recommended four-step checkpoint. Its VSA-H3 attention settings
|
||||
match this app's defaults. Download the whole snapshot, including the text
|
||||
encoder and both VAEs, into a local directory accessible to the GPU host:
|
||||
|
||||
```bash
|
||||
export LIVESTREAM_WEIGHTS_PATH=/absolute/path/to/FastH3-v1-VSA-DataFree
|
||||
hf download FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree \
|
||||
--local-dir "$LIVESTREAM_WEIGHTS_PATH"
|
||||
```
|
||||
|
||||
The app expects `modular_model_index.json` and the `transformer`, `text_encoder`,
|
||||
`tokenizer`, `processor`, `vae`, `audio_vae`, `scheduler` and `audio_scheduler`
|
||||
directories. Downloading only the transformer is insufficient; the text encoder
|
||||
and VAEs need their weight files as well as their configs.
|
||||
|
||||
## FFmpeg
|
||||
|
||||
Install FFmpeg with `libx264` and `aac` encoding support. For example, on
|
||||
Debian/Ubuntu:
|
||||
|
||||
```bash
|
||||
sudo apt install ffmpeg
|
||||
```
|
||||
|
||||
For the optimized native build, follow Dreamverse's
|
||||
[FFmpeg instructions](../dreamverse/README.md#optional-building-ffmpeg-for-better-performance).
|
||||
After running that installer, source its environment file in the shell that
|
||||
will launch the livestream:
|
||||
|
||||
```bash
|
||||
source apps/dreamverse/scripts/ffmpeg-env.sh
|
||||
```
|
||||
|
||||
## Quick start
|
||||
|
||||
With the virtual environment active and `LIVESTREAM_WEIGHTS_PATH` set above:
|
||||
|
||||
```bash
|
||||
export OPENAI_API_KEY=...
|
||||
infinite-livestream-server
|
||||
```
|
||||
|
||||
Open **<http://localhost:8081>** on the server, or use the server's hostname when
|
||||
connecting remotely. The default bind address is `0.0.0.0`; `--port` overrides
|
||||
the port.
|
||||
|
||||
The page and the HLS stream start while the model loads, so a viewer arriving
|
||||
during startup sees the page and a black stream. Weight loading and compile
|
||||
warm-up take several minutes. Check readiness with:
|
||||
|
||||
```bash
|
||||
curl http://localhost:8081/healthz
|
||||
```
|
||||
|
||||
`{"connected": true}` means model loading and warm-up have finished.
|
||||
|
||||
## Configuration
|
||||
|
||||
Copy the [default YAML](infinite_livestream/configs/infinite_livestream.yaml)
|
||||
from the repository root, edit it, and pass the copy with `--config`:
|
||||
|
||||
```bash
|
||||
cp apps/infinite_livestream/infinite_livestream/configs/infinite_livestream.yaml my-config.yaml
|
||||
infinite-livestream-server --config my-config.yaml
|
||||
```
|
||||
|
||||
| Block | Contents |
|
||||
|---|---|
|
||||
| `inference` | What the checkpoint is asked for: clip length, canvas, sparse-attention kernels, compile policy. |
|
||||
| `runtime` | How it is hosted: GPU count, sharding, offload. |
|
||||
| `upsampler` | Prompt rewriting: model, endpoint, how many clips one prompt may become. |
|
||||
| `moderation` | Whether viewer prompts are checked, and against which endpoint. |
|
||||
| `director` | Idle filler depth, per-viewer cooldown, chat command, filler directory. |
|
||||
| `output` | Where the playlist is written (defaults under `$XDG_STATE_HOME`), the video bitrate, and retained playback history. |
|
||||
| `web` | Bind address and port. |
|
||||
|
||||
For another OpenAI-compatible provider, set `upsampler.base_url` and
|
||||
`upsampler.model`. Moderation is enabled by default and uses that endpoint too
|
||||
unless `moderation.base_url` is set. If the provider does not offer `/moderations`,
|
||||
configure a separate moderation endpoint and export its `MODERATION_API_KEY`.
|
||||
|
||||
API keys and the machine's weights path stay in the environment:
|
||||
|
||||
| Variable | Description |
|
||||
|---|---|
|
||||
| `OPENAI_API_KEY` | Required. Prompt rewriting runs for the idle filler too, so the stream does not start without it. |
|
||||
| `LIVESTREAM_WEIGHTS_PATH` | Required. FastH3 model directory. `--weights` overrides it. |
|
||||
| `MODERATION_API_KEY` | Optional. Falls back to `OPENAI_API_KEY`. |
|
||||
|
||||
## Clip geometry
|
||||
|
||||
Clip geometry is fixed by the checkpoint: 24 fps, frame counts of the form
|
||||
`17n + 5`, a 5 to 15 second duration window, and a 768 pixel short edge.
|
||||
`inference.clip_seconds: 15.083` is the longest clip it can produce, at 362
|
||||
frames (15.0 s rounds up to the next valid length).
|
||||
|
||||
Keeping one clip length means one compiled shape. Setting
|
||||
`inference.warmup_lengths: all` warms every legal length instead, which makes
|
||||
startup slower but avoids a one-off compile stall on a viewer's first
|
||||
odd-length clip.
|
||||
|
||||
## API Endpoints
|
||||
|
||||
| Route | Description |
|
||||
|---|---|
|
||||
| `GET /` | The watch page. |
|
||||
| `GET /assets/<file>` | Viewer scripts, logo, and favicon. |
|
||||
| `GET /hls/<file>` | Playlist and segments, written by `infinite_livestream/sink.py`. |
|
||||
| `GET /healthz` | `{"connected": bool}`, true once the model is loaded. |
|
||||
| `WS /state` | One JSON snapshot on connect, then one per change. |
|
||||
| `POST /chat` | `{"author": str, "text": str}`. Returns 429 with `retry_after` when that viewer is still on cooldown. |
|
||||
|
||||
The cooldown is answered by `POST /chat` rather than reported later, so the
|
||||
sender's page can disable its send box and count down. The chat feed is shared
|
||||
by every viewer, so refusals are kept out of it.
|
||||
|
||||
## Idle fillers
|
||||
|
||||
When nobody is typing, the stream keeps itself fed from a list of prompts.
|
||||
`director.fillers` names the directory holding `fillers.json`, and defaults to
|
||||
the one that ships in `infinite_livestream/presets/`.
|
||||
|
||||
```json
|
||||
{
|
||||
"style": "the look and tone for idle filler clips",
|
||||
"idle_prompts": ["a lighthouse keeper teaching a seagull to play chess"]
|
||||
}
|
||||
```
|
||||
|
||||
`style` defines the house style for idle filler clips. Viewer prompts may use
|
||||
their own style by default (`upsampler.viewer_free_style: true`). Set
|
||||
`upsampler.viewer_free_style: false` to apply the house style to viewer prompts
|
||||
as well.
|
||||
`idle_prompts` feeds the filler; an empty list turns the filler off, as does
|
||||
`director.idle_queue_target: 0`.
|
||||
|
||||
To change the stream's identity, copy the directory, edit `fillers.json` and
|
||||
point `director.fillers` at it. The file is read once, at startup.
|
||||
|
||||
## Now-playing titles
|
||||
|
||||
Titles travel inside the HLS segments as timed ID3 metadata. The page prefers
|
||||
hls.js where supported so browsers use the same metadata parser; native HLS is
|
||||
the fallback. The player parses the metadata into cues, and the page selects the
|
||||
title using the browser's presented frame timestamp, including after paused seeks
|
||||
or buffering. Each segment includes
|
||||
the active title for viewers joining halfway through a clip. A repeated final
|
||||
frame keeps its title until the next clip appears. Browsers without
|
||||
`requestVideoFrameCallback` fall back to playback-clock cue events, which can be
|
||||
less precise at seek boundaries.
|
||||
|
||||
FFmpeg encodes video and audio once; PyAV copies the compressed packets into HLS
|
||||
and adds the metadata. This does not use extra GPUs or encode a stream per viewer.
|
||||
Append `?debug=1` to see the active clip ID and playback position.
|
||||
|
||||
`output.hls_retention_s` controls retained history (120 seconds by default),
|
||||
without increasing the player's target live latency. A viewer whose requested
|
||||
footage has expired must rejoin available footage. Paused seeks into gaps move
|
||||
to the next buffered frame. Encoder restarts retain recent segments and mark the
|
||||
new media timeline explicitly.
|
||||
|
||||
The page stacks video and chat in portrait. Wide, short landscape screens put
|
||||
chat beside the video, with the title below the video and the queue collapsed
|
||||
initially. Rotating the page preserves playback.
|
||||
|
||||
Playback, titles, seeking, and buffering were tested in Firefox and Chromium,
|
||||
including phone and tablet layout emulation. Safari and physical mobile devices
|
||||
remain unverified. If an embedded browser cannot decode H.264/AAC, open the page
|
||||
in an external browser.
|
||||
|
||||
## Tests
|
||||
|
||||
```bash
|
||||
pytest apps/infinite_livestream/infinite_livestream/tests -m "not gpu"
|
||||
```
|
||||
|
||||
CPU media integration tests require FFmpeg with `libx264` and `aac`. The browser
|
||||
metadata adapter also has a dependency-free JavaScript regression test:
|
||||
|
||||
```bash
|
||||
node --test apps/infinite_livestream/infinite_livestream/tests/test_metadata_player.cjs
|
||||
```
|
||||
|
||||
One test is marked `gpu`. It checks that `infinite_livestream/clip_plan.py`'s copy of
|
||||
MiniMax-H3's packing constants still matches FastVideo's, and importing the
|
||||
upstream module needs a live CUDA driver. Run it when the pinned FastVideo
|
||||
version moves.
|
||||
|
||||
## Adding another model
|
||||
|
||||
`FastH3Backend.submit(frames, prompt, seed, height, width)` is the seam.
|
||||
Everything above it, meaning the engine, director, queues, pacer, sink and web
|
||||
app, is model-agnostic. Everything below it is MiniMax-H3 specific:
|
||||
`clip_plan.py` is its geometry and `backend.py` selects its kernels.
|
||||
|
||||
A second checkpoint needs its own geometry module and its own backend behind
|
||||
that seam. LTX-2, for example, packs `8n + 1` frames at different resolutions.
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
**`ffmpeg not found on PATH`.** Follow [FFmpeg](#ffmpeg). If you used the native
|
||||
build, source `apps/dreamverse/scripts/ffmpeg-env.sh` before starting the app.
|
||||
|
||||
**The weights are incomplete.** Startup lists the missing components before
|
||||
any GPU work begins. The model directory needs `transformer`, `text_encoder`,
|
||||
`tokenizer`, `processor`, `vae`, `audio_vae`, `scheduler`, `audio_scheduler`
|
||||
and `modular_model_index.json`. See [Download weights](#download-weights) for
|
||||
the complete checkpoint.
|
||||
|
||||
**`FastH3's sm100a route needs fastvideo-kernel built with the Blackwell VSA
|
||||
extension`.** Startup checks for the fast sparse-attention kernel before
|
||||
loading any weights. Follow [Install](#install) to build it, or set
|
||||
`inference.vsa_kernel: triton` in your config to use the slower fallback.
|
||||
|
||||
**`FastH3's FA4 route needs the pinned flash-attn-4 package`.** Include the
|
||||
`fasth3` extra as shown in [Install](#install), or set `inference.fa4: false` in
|
||||
your config.
|
||||
|
||||
**A clip takes much longer than the others.** Each distinct clip length is a
|
||||
separate compiled shape, and the first clip at a new length pays a one-off
|
||||
compile cost. `inference.warmup_lengths: all` pays all of them at startup instead.
|
||||
@@ -0,0 +1,5 @@
|
||||
"""Infinite Livestream: a chat-driven FastH3 broadcast."""
|
||||
|
||||
__all__ = ["__version__"]
|
||||
|
||||
__version__ = "0.1.0"
|
||||
@@ -0,0 +1,516 @@
|
||||
"""The FastVideo side of FastH3: GPU work, and nothing stream-facing.
|
||||
|
||||
`FastH3Backend` owns the multi-GPU `VideoGenerator`, the environment profile
|
||||
it must be built under, the load-time warm-up, and one worker thread that
|
||||
builds clips serially. `engine.py` calls `submit` and polls the `ClipJob` it
|
||||
gets back; nothing here knows about queues, chat or the sink.
|
||||
|
||||
torch, torchaudio and fastvideo are imported inside methods, so importing this
|
||||
module needs none of them and the config tests can run without a GPU.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from . import clip_plan
|
||||
from .config import ModelConfig
|
||||
|
||||
logger = logging.getLogger("infinite_livestream.backend")
|
||||
|
||||
FRAME_RATE = clip_plan.FPS
|
||||
|
||||
# The checkpoint's audio decoder runs at 32 kHz; the sink wants 48 kHz.
|
||||
OUTPUT_SAMPLE_RATE = 48_000
|
||||
NATIVE_SAMPLE_RATE = 32_000
|
||||
|
||||
_WORKER_POLL_SECONDS = 0.1
|
||||
|
||||
# Warm-up output is discarded; this only has to be an ordinary prompt.
|
||||
WARMUP_PROMPT = "A slow cinematic shot of sunlight moving across a quiet room."
|
||||
|
||||
# Every prompt is padded or truncated to exactly this many tokens. Regional
|
||||
# torch.compile keys on the packed sequence length, which includes the prompt's
|
||||
# token count, so a novel length recompiles -- ~23 s against ~15 s for the clip
|
||||
# itself. One fixed length is one compiled shape, warmed once. 256 comfortably
|
||||
# holds the 800-character cap.
|
||||
PROMPT_TOKENS = 256
|
||||
|
||||
|
||||
class ClipJob:
|
||||
"""The handle to one submitted build: its inputs, outcome, and completion.
|
||||
|
||||
The error is carried back rather than only logged, so the submitter can
|
||||
report the failed clip to clients. ``cancelled`` set before the worker
|
||||
reaches the job skips the build entirely; set after, the build runs to
|
||||
completion and the submitter discards the result.
|
||||
"""
|
||||
|
||||
__slots__ = ("cancelled", "done", "error", "fn", "result")
|
||||
|
||||
def __init__(self, fn) -> None:
|
||||
self.fn = fn
|
||||
self.done = threading.Event()
|
||||
self.error: BaseException | None = None
|
||||
self.result: tuple[list[Any], Any] | None = None
|
||||
self.cancelled = False
|
||||
|
||||
|
||||
class FastH3Backend:
|
||||
"""Build FastH3 clips on demand, serialised on one worker thread.
|
||||
|
||||
The GPU work itself lives in the engine processes FastVideo spawns; the
|
||||
thread exists to serialise submissions and to give teardown a single handle
|
||||
to wait on.
|
||||
"""
|
||||
|
||||
def __init__(self, config: ModelConfig, model_path: Path) -> None:
|
||||
"""Remember the recipe and the weights location; nothing loads yet."""
|
||||
self._config = config
|
||||
self._model_path = model_path
|
||||
self._jobs: queue.Queue[ClipJob] = queue.Queue()
|
||||
self._worker: threading.Thread | None = None
|
||||
self.generator: Any = None
|
||||
|
||||
# ------------------------------------------------------------------ load
|
||||
|
||||
def load(self) -> None:
|
||||
"""Build the generator and warm every configured clip shape.
|
||||
|
||||
Runs once at startup, and this returning is what lets the engine
|
||||
accept work -- so everything that can fail (missing kernels, a broken
|
||||
native linkage, a cold compile) must fail here, not on a viewer's clip.
|
||||
"""
|
||||
# Must happen before the generator is built: the engine spawns worker
|
||||
# processes, which inherit os.environ, and these select the attention
|
||||
# backend and the sparse kernel.
|
||||
self._apply_profile_environment()
|
||||
self._validate_profile_dependencies()
|
||||
self._raise_dynamo_limits()
|
||||
|
||||
runtime = self._config.runtime
|
||||
num_gpus = int(runtime.get("num_gpus", 4))
|
||||
logger.info("building the generator: %s, %d gpu(s), %d-frame clips", self._model_path, num_gpus,
|
||||
self._config.clip_frames)
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
self.generator = VideoGenerator.from_config(self._generator_config())
|
||||
self._load_tokenizer()
|
||||
|
||||
self._worker = threading.Thread(target=self._worker_loop, name="fast-h3-generation", daemon=True)
|
||||
self._worker.start()
|
||||
self._preload_native_imports()
|
||||
self._run_blocking(self._warmup)
|
||||
logger.info("backend loaded")
|
||||
|
||||
def _load_tokenizer(self) -> None:
|
||||
"""Load the checkpoint's tokenizer and calibrate the one-token pad filler.
|
||||
|
||||
Padding must land on an exact token count, so the filler is verified to
|
||||
cost exactly one token at load rather than assumed.
|
||||
"""
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
self._tokenizer = AutoTokenizer.from_pretrained(str(self._model_path / "tokenizer"))
|
||||
for candidate in (" .", ".", " a"):
|
||||
base = len(self._tokenizer.encode(WARMUP_PROMPT, add_special_tokens=False))
|
||||
padded = len(self._tokenizer.encode(WARMUP_PROMPT + candidate, add_special_tokens=False))
|
||||
if padded == base + 1:
|
||||
self._pad_filler = candidate
|
||||
return
|
||||
raise RuntimeError("no single-token pad filler found for this tokenizer")
|
||||
|
||||
def _pad_prompt(self, prompt: str) -> str:
|
||||
"""Return *prompt* at exactly ``PROMPT_TOKENS`` tokens.
|
||||
|
||||
Shorter prompts gain trailing filler tokens; a longer one (past the
|
||||
800-character cap only in pathological tokenizations) is truncated at
|
||||
the token boundary. The client-facing prompt — what `ClipInfo` echoes —
|
||||
is the original; only the engine sees this form.
|
||||
"""
|
||||
|
||||
def encode(text: str) -> int:
|
||||
return len(self._tokenizer.encode(text, add_special_tokens=False))
|
||||
|
||||
ids = self._tokenizer.encode(prompt, add_special_tokens=False)
|
||||
if ids and len(ids) >= PROMPT_TOKENS:
|
||||
return self._tokenizer.decode(ids[:PROMPT_TOKENS])
|
||||
padded = prompt + self._pad_filler * (PROMPT_TOKENS - len(ids))
|
||||
# Filler cost is calibrated, but a prompt's own tail can merge with the
|
||||
# first filler token; correct by measurement rather than assumption.
|
||||
while encode(padded) > PROMPT_TOKENS:
|
||||
padded = padded[:-len(self._pad_filler)]
|
||||
while encode(padded) < PROMPT_TOKENS:
|
||||
padded += self._pad_filler
|
||||
if encode(padded) != PROMPT_TOKENS:
|
||||
logger.warning("prompt padded to %d tokens, wanted %d", encode(padded), PROMPT_TOKENS)
|
||||
return padded
|
||||
|
||||
@staticmethod
|
||||
def _raise_dynamo_limits() -> None:
|
||||
"""Stop a novel tensor shape from being a hard failure.
|
||||
|
||||
Each clip length is a torch.compile shape, and the fullgraph regional
|
||||
route treats exceeding dynamo's recompile limit as a crash rather than
|
||||
a fallback. FastVideo's own imports lower it (`layers/lora/linear.py`
|
||||
to 16, longcat's `bsa_interface.py` to 32), so this raises it again
|
||||
and turns overflow back into a recompile.
|
||||
|
||||
One pinned clip length keeps every process far under even the lowered
|
||||
limit; a varied `warmup_lengths` is where this becomes the seatbelt it
|
||||
is meant to be.
|
||||
"""
|
||||
import torch._dynamo.config as dynamo_config
|
||||
|
||||
limit = int(os.environ.get("LIVESTREAM_DYNAMO_RECOMPILE_LIMIT", "64"))
|
||||
dynamo_config.recompile_limit = max(limit, dynamo_config.recompile_limit)
|
||||
dynamo_config.cache_size_limit = max(limit, dynamo_config.cache_size_limit)
|
||||
dynamo_config.accumulated_recompile_limit = max(512, dynamo_config.accumulated_recompile_limit)
|
||||
dynamo_config.accumulated_cache_size_limit = max(512, dynamo_config.accumulated_cache_size_limit)
|
||||
dynamo_config.fail_on_recompile_limit_hit = False
|
||||
logger.info("dynamo recompile limit raised to %d", dynamo_config.recompile_limit)
|
||||
|
||||
@staticmethod
|
||||
def _preload_native_imports() -> None:
|
||||
"""Touch every deferred native import the build path needs.
|
||||
|
||||
Otherwise the first one happens on the first real clip, where a broken
|
||||
linkage is a dead stream rather than a startup error. The resample is a
|
||||
real call, so it fails here or not at all.
|
||||
"""
|
||||
import numpy # noqa: F401
|
||||
import torch
|
||||
import torchaudio.functional as AF
|
||||
|
||||
AF.resample(torch.zeros(2, NATIVE_SAMPLE_RATE // 10), NATIVE_SAMPLE_RATE, OUTPUT_SAMPLE_RATE)
|
||||
|
||||
# --------------------------------------------------------------- profile
|
||||
|
||||
def _apply_profile_environment(self) -> None:
|
||||
"""Set the FastH3 profile environment, as the reference CLI does.
|
||||
|
||||
Mirrors `examples/inference/basic/basic_fasth3.py:profile_environment`.
|
||||
Disabled features are set explicitly too, so a shell's inherited
|
||||
experiment settings cannot silently change what gets served.
|
||||
"""
|
||||
cfg = self._config.inference
|
||||
vsa_kernel = str(cfg.get("vsa_kernel", "sm100a"))
|
||||
fusions = "all" if bool(cfg.get("h3_fusions", True)) else "0"
|
||||
environment: dict[str, str | None] = {
|
||||
"FASTVIDEO_ATTENTION_BACKEND": "VIDEO_SPARSE_ATTN_H3",
|
||||
"FASTVIDEO_VSA_SM100A": "1" if vsa_kernel == "sm100a" else "0",
|
||||
"FASTVIDEO_VSA_CUTEDSL": "0",
|
||||
# A non-empty path enables the diagnostic probe; it must stay unset.
|
||||
"FASTVIDEO_H3_VSA_PROBE": None,
|
||||
"FASTVIDEO_DISABLE_ATTENTION_COMPILE": "0",
|
||||
"FASTVIDEO_FA4": "1" if bool(cfg.get("fa4", True)) else "0",
|
||||
"FASTVIDEO_NVFP4_FA4": "0",
|
||||
"FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN": "0",
|
||||
"FASTVIDEO_MINIMAX_H3_FUSIONS": fusions,
|
||||
"FASTVIDEO_INFERENCE_TORCH_COMPILE": ("1" if bool(cfg.get("inference_torch_compile", True)) else "0"),
|
||||
"FASTVIDEO_VAE_PARALLEL_DECODE": ("1" if bool(cfg.get("vae_parallel_decode", True)) else "0"),
|
||||
"FASTVIDEO_VAE_PARALLEL_ENCODE": "0",
|
||||
"FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY": "gather",
|
||||
"FASTVIDEO_ULYSSES_A2A": str(cfg.get("ulysses_a2a", "off")),
|
||||
"FASTVIDEO_STAGE_LOGGING": "1",
|
||||
}
|
||||
for name, value in environment.items():
|
||||
if value is None:
|
||||
os.environ.pop(name, None)
|
||||
else:
|
||||
os.environ[name] = value
|
||||
logger.info("profile: %s", " ".join(f"{k}={v or '<unset>'}" for k, v in environment.items()))
|
||||
|
||||
def _validate_profile_dependencies(self) -> None:
|
||||
"""Fail before the weights load when the selected fast route is absent."""
|
||||
import importlib.util
|
||||
|
||||
cfg = self._config.inference
|
||||
if bool(cfg.get("fa4", True)):
|
||||
try:
|
||||
present = importlib.util.find_spec("flash_attn.cute") is not None
|
||||
except (ImportError, ModuleNotFoundError):
|
||||
present = False
|
||||
if not present:
|
||||
raise RuntimeError("FastH3's FA4 route needs the pinned flash-attn-4 package. Install it, "
|
||||
"or set inference.fa4: false in your config.")
|
||||
if str(cfg.get("vsa_kernel", "sm100a")) == "sm100a":
|
||||
try:
|
||||
from fastvideo_kernel import block_sparse_attn_sm100a
|
||||
except ImportError:
|
||||
present = False
|
||||
else:
|
||||
present = bool(getattr(block_sparse_attn_sm100a, "_HAS_VSA_SM100A", False))
|
||||
if not present:
|
||||
raise RuntimeError("FastH3's sm100a route needs fastvideo-kernel built with the Blackwell VSA "
|
||||
"extension. Install a matching wheel, or set inference.vsa_kernel: triton.")
|
||||
|
||||
def _generator_config(self) -> Any:
|
||||
"""The engine shape, mirroring `basic_fasth3.py`."""
|
||||
from fastvideo.api import (
|
||||
CompileConfig,
|
||||
ComponentConfig,
|
||||
EngineConfig,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
)
|
||||
|
||||
cfg = self._config.inference
|
||||
runtime = self._config.runtime
|
||||
num_gpus = int(runtime.get("num_gpus", 4))
|
||||
# The checkpoint's own contract (fastvideo_inference.json) shards the
|
||||
# transformer with FSDP. Sharding is what frees the VRAM to keep the
|
||||
# text encoder resident, which is the deployment this model wants.
|
||||
replicated_dit = bool(runtime.get("replicated_dit", False))
|
||||
return GeneratorConfig(
|
||||
model_path=str(self._model_path),
|
||||
pipeline=PipelineSelection(
|
||||
components=ComponentConfig(),
|
||||
experimental={
|
||||
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
|
||||
"VSA_sparsity": float(cfg.get("vsa_sparsity", 0.9)),
|
||||
"VSA_tile_size": int(cfg.get("vsa_tile_size", 64)),
|
||||
"inference_torch_compile": bool(cfg.get("inference_torch_compile", True)),
|
||||
"vae_parallel_decode": bool(cfg.get("vae_parallel_decode", True)),
|
||||
"vae_parallel_decode_strategy": "gather",
|
||||
},
|
||||
),
|
||||
engine=EngineConfig(
|
||||
num_gpus=num_gpus,
|
||||
use_fsdp_inference=num_gpus > 1 and not replicated_dit,
|
||||
parallelism=ParallelismConfig(tp_size=1, sp_size=num_gpus),
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
dit_layerwise=False,
|
||||
text_encoder=bool(runtime.get("offload_text_encoder", False)),
|
||||
vae=bool(runtime.get("offload_vae", False)),
|
||||
pin_cpu_memory=bool(runtime.get("pin_cpu_memory", False)),
|
||||
),
|
||||
compile=CompileConfig(
|
||||
enabled=False,
|
||||
mode=None,
|
||||
vae_enabled=bool(cfg.get("compile_vae", True)),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------- worker
|
||||
|
||||
def _worker_loop(self) -> None:
|
||||
"""Run submitted jobs, one at a time, forever.
|
||||
|
||||
The waiter is always released, even when the job died: a completion
|
||||
event that never arrives is indistinguishable from a hang, and this
|
||||
thread is the only one that will ever set it.
|
||||
"""
|
||||
logger.info("generation worker ready")
|
||||
while True:
|
||||
job = self._jobs.get()
|
||||
try:
|
||||
if not job.cancelled:
|
||||
job.result = job.fn()
|
||||
except BaseException as error: # noqa: BLE001 — handed to the submitter
|
||||
job.error = error
|
||||
logger.exception("generation worker job raised")
|
||||
finally:
|
||||
job.done.set()
|
||||
|
||||
def submit(self, *, frames: int, prompt: str, seed: int, height: int, width: int) -> ClipJob:
|
||||
"""Queue one clip build and hand back its job handle.
|
||||
|
||||
Returns immediately; the caller polls ``job.done`` and reads
|
||||
``job.result`` — ``(frames_list, samples)``, RGB uint8 frames and an
|
||||
int16 ``[1, samples]`` waveform at the wire rate — or ``job.error``.
|
||||
"""
|
||||
job = ClipJob(lambda: self._generate_clip(frames=frames, prompt=prompt, seed=seed, height=height, width=width))
|
||||
self._jobs.put(job)
|
||||
return job
|
||||
|
||||
def _run_blocking(self, fn) -> None:
|
||||
"""Run work on the worker, block until it finishes, and re-raise its failure.
|
||||
|
||||
Used only by `load()`, where blocking is the point: a failed warm-up
|
||||
has to stop startup rather than surface on a viewer's first clip.
|
||||
"""
|
||||
job = ClipJob(fn)
|
||||
self._jobs.put(job)
|
||||
while not job.done.wait(timeout=_WORKER_POLL_SECONDS):
|
||||
pass
|
||||
if job.error is not None:
|
||||
raise job.error
|
||||
|
||||
# --------------------------------------------------------------- warm-up
|
||||
|
||||
def _warmup(self) -> None:
|
||||
"""Build one throwaway clip per shape before reporting ready.
|
||||
|
||||
Every distinct frame count and canvas costs a one-time regional
|
||||
compile, sparse-kernel autotune and allocator growth; paying it here
|
||||
means the first real clip builds at warm speed.
|
||||
|
||||
Two axes, not their cross product: every configured canvas at the
|
||||
default length, and every configured length at the primary canvas. A
|
||||
non-primary canvas at a non-default length still stalls on first use.
|
||||
"""
|
||||
aspects = self._config.warmup_aspects
|
||||
cold = [a for a in clip_plan.ASPECT_CHOICES if a not in aspects]
|
||||
if cold:
|
||||
logger.info("aspects left cold, their first clip pays a compile stall: %s", cold)
|
||||
shapes: list[tuple[str, int]] = [(aspect, self._config.clip_frames) for aspect in aspects]
|
||||
shapes += [(aspects[0], frames) for frames in self._config.warmup_frames if frames != self._config.clip_frames]
|
||||
logger.info("warming %d shape(s), lengths %s", len(shapes),
|
||||
[round(clip_plan.seconds_for_frames(f), 3) for f in self._config.warmup_frames])
|
||||
for index, (aspect, frames) in enumerate(shapes, start=1):
|
||||
height, width = clip_plan.canvas_for_choice(aspect)
|
||||
started = time.monotonic()
|
||||
self.generator.generate(
|
||||
self._request(
|
||||
frames=frames,
|
||||
prompt=WARMUP_PROMPT,
|
||||
seed=self._config.seed,
|
||||
height=height,
|
||||
width=width,
|
||||
keep_output=False,
|
||||
))
|
||||
logger.info("warmed %d/%d: %s %df at %dx%d in %.2fs", index, len(shapes), aspect, frames, height, width,
|
||||
time.monotonic() - started)
|
||||
|
||||
# ------------------------------------------------------------ generation
|
||||
|
||||
def _request(
|
||||
self,
|
||||
*,
|
||||
frames: int,
|
||||
prompt: str,
|
||||
seed: int,
|
||||
height: int,
|
||||
width: int,
|
||||
keep_output: bool,
|
||||
):
|
||||
"""Build one generation request, mirroring `basic_fasth3.py`.
|
||||
|
||||
`keep_output=False` is the warm-up shape: it skips the post-decode
|
||||
path, so a warm-up costs generation time and nothing else.
|
||||
"""
|
||||
from fastvideo.api import GenerationRequest, OutputConfig, SamplingConfig
|
||||
|
||||
return GenerationRequest(
|
||||
# Padded to the fixed token length so one compiled shape serves
|
||||
# every prompt; ClipInfo keeps echoing the original text.
|
||||
prompt=self._pad_prompt(prompt),
|
||||
# MiniMax-H3 is guidance-distilled, so there is no negative branch
|
||||
# to steer and no CFG pass to pay for.
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=frames,
|
||||
fps=FRAME_RATE,
|
||||
num_inference_steps=self._config.num_inference_steps,
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=seed,
|
||||
),
|
||||
output=OutputConfig(save_video=False, return_frames=keep_output),
|
||||
)
|
||||
|
||||
def _generate_clip(self, *, frames: int, prompt: str, seed: int, height: int, width: int):
|
||||
"""Build one clip and convert it to what the pacer wants.
|
||||
|
||||
Returns ``(frames_list, samples)``: a list of RGB uint8 ``[h, w, 3]``
|
||||
arrays and int16 ``[1, samples]`` at 48 kHz, trimmed to exactly
|
||||
``len(frames_list) / 24`` seconds so the two tracks stay in lockstep.
|
||||
"""
|
||||
started = time.monotonic()
|
||||
result = self.generator.generate(
|
||||
self._request(
|
||||
frames=frames,
|
||||
prompt=prompt,
|
||||
seed=seed,
|
||||
height=height,
|
||||
width=width,
|
||||
keep_output=True,
|
||||
))
|
||||
built = time.monotonic() - started
|
||||
|
||||
frames_list = result.frames
|
||||
if not frames_list:
|
||||
raise RuntimeError("the generator returned no frames")
|
||||
samples = self._to_wire_audio(result.audio, result.audio_sample_rate, len(frames_list))
|
||||
# The line to evaluate the deployment by: build seconds against content
|
||||
# seconds (realtime_x > 1 means the clip built faster than it plays) on
|
||||
# the GPU count that produced it, with the per-stage split. The numbers
|
||||
# live in the message itself so every log formatter carries them.
|
||||
content = len(frames_list) / FRAME_RATE
|
||||
gpus = int(self._config.runtime.get("num_gpus", 4))
|
||||
logger.info("clip built: %df (%.2fs content) in %.2fs = %.2fx realtime on %d gpus, stages=%s", len(frames_list),
|
||||
content, built, content / built, gpus, self._stage_times(result))
|
||||
return frames_list, samples
|
||||
|
||||
@staticmethod
|
||||
def _stage_times(result) -> dict:
|
||||
"""Per-stage seconds from the generator, for the clip log line.
|
||||
|
||||
This is where a regression shows up first: post-decode frame processing
|
||||
scales with resolution x frames and competes with the build budget.
|
||||
"""
|
||||
try:
|
||||
stages = getattr(getattr(result, "logging_info", None), "stages", None)
|
||||
if not stages:
|
||||
return {}
|
||||
return {
|
||||
name: round(float(metrics["execution_time"]), 3)
|
||||
for name, metrics in stages.items() if metrics.get("execution_time") is not None
|
||||
}
|
||||
except Exception: # noqa: BLE001 — a log line must never fail a clip
|
||||
logger.exception("could not read the generator stage timings")
|
||||
return {}
|
||||
|
||||
def _to_wire_audio(self, audio, sample_rate, frames: int):
|
||||
"""Resample, downmix and quantize one clip's waveform for the wire.
|
||||
|
||||
Mono at the source is deliberate: the transport mean-downmixes before
|
||||
the wire anyway, and the runtime recorder flattens two channels by
|
||||
concatenation, so a stereo emit only corrupts recordings. Averaging here,
|
||||
in float and before the int16 scale, is the same downmix one step
|
||||
earlier.
|
||||
"""
|
||||
import torch
|
||||
import torchaudio.functional as AF
|
||||
|
||||
if audio is None:
|
||||
raise RuntimeError("the generator returned no audio")
|
||||
waveform = audio if torch.is_tensor(audio) else torch.as_tensor(audio)
|
||||
waveform = waveform.detach().float().cpu()
|
||||
# The decoder hands back [samples, channels]; the wire wants channel-major.
|
||||
if waveform.ndim == 1:
|
||||
waveform = waveform.unsqueeze(0)
|
||||
elif waveform.shape[0] > waveform.shape[1]:
|
||||
waveform = waveform.transpose(0, 1)
|
||||
waveform = waveform.contiguous()
|
||||
|
||||
rate = int(sample_rate or NATIVE_SAMPLE_RATE)
|
||||
if rate != OUTPUT_SAMPLE_RATE:
|
||||
waveform = AF.resample(waveform, rate, OUTPUT_SAMPLE_RATE)
|
||||
if waveform.shape[0] > 1:
|
||||
waveform = waveform.mean(dim=0, keepdim=True)
|
||||
|
||||
want = round(frames / FRAME_RATE * OUTPUT_SAMPLE_RATE)
|
||||
if waveform.shape[-1] > want:
|
||||
waveform = waveform[:, :want]
|
||||
elif waveform.shape[-1] < want:
|
||||
pad = torch.zeros((waveform.shape[0], want - waveform.shape[-1]), dtype=waveform.dtype)
|
||||
waveform = torch.cat([waveform, pad], dim=-1)
|
||||
return (waveform.clamp(-1, 1) * 32767).to(torch.int16).numpy()
|
||||
|
||||
|
||||
__all__ = ["OUTPUT_SAMPLE_RATE", "ClipJob", "FastH3Backend"]
|
||||
@@ -0,0 +1,83 @@
|
||||
"""Viewer prompts, typed into the page that plays the stream.
|
||||
|
||||
The page is the only way in, so this is fed directly in-process by `webapp.py`
|
||||
rather than polling a platform. `submit` is called from a request handler and
|
||||
never awaits: a full queue drops the message and tells the viewer, which is
|
||||
better than stalling the web server.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable, Sequence
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
logger = logging.getLogger("infinite_livestream.chat")
|
||||
|
||||
# Small on purpose: a deep queue would let a burst of typing commit the stream
|
||||
# to minutes of stale prompts.
|
||||
QUEUE_SIZE = 32
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ChatPrompt:
|
||||
"""One accepted message, command word stripped."""
|
||||
|
||||
source: str
|
||||
author: str
|
||||
text: str
|
||||
command: str = ""
|
||||
received_at: float = field(default_factory=time.monotonic)
|
||||
|
||||
|
||||
def match_command(message: str, commands: Sequence[str]) -> tuple[str, str] | None:
|
||||
"""Match a message against command words: `(command, text)` or None.
|
||||
|
||||
Case-insensitive, and a bare command with no text is ignored.
|
||||
"""
|
||||
stripped = message.strip()
|
||||
lowered = stripped.lower()
|
||||
for command in commands:
|
||||
if not lowered.startswith(command.lower()):
|
||||
continue
|
||||
remainder = stripped[len(command):]
|
||||
if remainder and not remainder[0].isspace():
|
||||
continue # "!promptfoo" is not "!prompt foo"
|
||||
text = remainder.strip()
|
||||
if text:
|
||||
return command, text
|
||||
return None
|
||||
|
||||
|
||||
class WebChat:
|
||||
"""Prompts submitted through the page's chat box."""
|
||||
|
||||
name = "web"
|
||||
|
||||
def __init__(self, command: str = "!prompt") -> None:
|
||||
self._command = command
|
||||
self._queue: asyncio.Queue[ChatPrompt] = asyncio.Queue(maxsize=QUEUE_SIZE)
|
||||
self._dropped = 0
|
||||
|
||||
def submit(self, author: str, text: str, command: str | None = None) -> bool:
|
||||
"""Accept one message. True when it was queued, False when dropped."""
|
||||
text = text.strip()
|
||||
if not text:
|
||||
return False
|
||||
matched = match_command(text, (self._command, ))
|
||||
word, body = matched if matched else (self._command, text)
|
||||
prompt = ChatPrompt(source=self.name, author=author or "viewer", text=body, command=command or word)
|
||||
try:
|
||||
self._queue.put_nowait(prompt)
|
||||
except asyncio.QueueFull:
|
||||
self._dropped += 1
|
||||
logger.warning("[chat] queue full, dropped prompt from %s (%d total)", prompt.author, self._dropped)
|
||||
return False
|
||||
return True
|
||||
|
||||
async def run(self, on_prompt: Callable[[ChatPrompt], None]) -> None:
|
||||
logger.info("[chat] ready (queue %d)", QUEUE_SIZE)
|
||||
while True:
|
||||
on_prompt(await self._queue.get())
|
||||
@@ -0,0 +1,162 @@
|
||||
"""Clip geometry for the FastH3 channel.
|
||||
|
||||
Pure arithmetic over the checkpoint's published constraints: how long a clip may
|
||||
be, how many frames that is, and what canvas an aspect ratio resolves to. No
|
||||
torch, no fastvideo, no GPU, so the config and queue tests import it anywhere.
|
||||
|
||||
The constants below are duplicated from FastVideo rather than imported, because
|
||||
``fastvideo.pipelines.basic.minimax_h3.packing`` pulls in torch and, through
|
||||
fastvideo-kernel's triton autotuning, needs a live CUDA driver just to import.
|
||||
``tests/test_clip_plan.py`` asserts they still match upstream on a machine that
|
||||
has one, so the duplication cannot drift silently.
|
||||
|
||||
Everything here is MiniMax-H3's geometry specifically. A second checkpoint --
|
||||
LTX-2 packs 8n+1 frames at its own resolutions -- needs its own module, not
|
||||
edits to this one; see the app README.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
FPS = 24
|
||||
"""The only frame rate MiniMax-H3 accepts; the pipeline rejects anything else."""
|
||||
|
||||
# The causal VAE consumes video in 17-frame chunks that decode to 5 latents, so
|
||||
# a valid pixel length is always `17n + 5`.
|
||||
_FRAMES_PER_CHUNK = 17
|
||||
_LATENTS_PER_CHUNK = 5
|
||||
|
||||
# The checkpoint's trained duration window, in seconds. The cap applies to the
|
||||
# requested length; the aligned bucket it rounds to is what gets generated.
|
||||
_MIN_DURATION = 5.0
|
||||
_MAX_DURATION = 15.0
|
||||
|
||||
# Canvas rules: the short edge is fixed, total area is capped, and both sides
|
||||
# must land on a multiple of 32.
|
||||
_SHORT_EDGE = 768
|
||||
_MAX_PIXELS = 768 * 1344
|
||||
_CANVAS_MULTIPLE = 32
|
||||
_MIN_ASPECT = 1 / 4
|
||||
_MAX_ASPECT = 4
|
||||
|
||||
|
||||
def align_frames(frames: int) -> int:
|
||||
"""Round up to the next valid `17n + 5` pixel length."""
|
||||
if frames < 1:
|
||||
raise ValueError(f"frames must be positive, got {frames}")
|
||||
while frames % _FRAMES_PER_CHUNK != _LATENTS_PER_CHUNK:
|
||||
frames += 1
|
||||
return frames
|
||||
|
||||
|
||||
def _bounds() -> tuple[int, int]:
|
||||
"""The shortest and longest clip that satisfies both alignment and duration.
|
||||
|
||||
The ceiling is the subtle one: the cap applies to the aligned bucket, not to
|
||||
the requested length, so 15.0 s (360 frames) pads *up* to 362 -- 15.083 s of
|
||||
playout -- and that is the longest clip this checkpoint will generate.
|
||||
"""
|
||||
return align_frames(int(_MIN_DURATION * FPS)), align_frames(int(_MAX_DURATION * FPS))
|
||||
|
||||
|
||||
MIN_FRAMES, MAX_FRAMES = _bounds()
|
||||
MIN_SECONDS = MIN_FRAMES / FPS
|
||||
MAX_SECONDS = MAX_FRAMES / FPS
|
||||
|
||||
# The same bounds as the schema publishes them. Rounded *inward* to three
|
||||
# decimals so a client reads "5.167", not "5.166666666666667", and so every
|
||||
# value inside the published range still snaps to a generatable clip.
|
||||
MIN_SECONDS_PUBLISHED = math.ceil(MIN_SECONDS * 1000) / 1000
|
||||
MAX_SECONDS_PUBLISHED = math.floor(MAX_SECONDS * 1000) / 1000
|
||||
|
||||
|
||||
def legal_frame_counts() -> tuple[int, ...]:
|
||||
"""Every clip length this checkpoint can generate, in frames, ascending.
|
||||
|
||||
The `17n + 5` alignment makes consecutive legal lengths exactly one chunk
|
||||
(17 frames) apart, so the whole space is a simple range.
|
||||
"""
|
||||
return tuple(range(MIN_FRAMES, MAX_FRAMES + 1, _FRAMES_PER_CHUNK))
|
||||
|
||||
|
||||
def frames_for_seconds(seconds: float) -> int:
|
||||
"""Snap a requested clip length to the nearest length the model can make.
|
||||
|
||||
Rounds up to a valid frame count, then clamps into the generatable range, so
|
||||
every accepted value round-trips through ``seconds_for_frames``.
|
||||
"""
|
||||
if seconds <= 0:
|
||||
raise ValueError(f"seconds must be positive, got {seconds}")
|
||||
frames = align_frames(max(1, round(seconds * FPS)))
|
||||
return max(MIN_FRAMES, min(MAX_FRAMES, frames))
|
||||
|
||||
|
||||
def seconds_for_frames(frames: int) -> float:
|
||||
"""Exact playout length of a clip, in seconds."""
|
||||
return frames / FPS
|
||||
|
||||
|
||||
def canvas_for_aspect(aspect_width: float, aspect_height: float) -> tuple[int, int]:
|
||||
"""Resolve an aspect ratio to a `(height, width)` the checkpoint accepts.
|
||||
|
||||
Mirrors FastVideo's ``resolve_canvas_size``: pin the short edge to 768,
|
||||
shrink to the area cap if the result is too wide, then round both sides to a
|
||||
multiple of 32.
|
||||
"""
|
||||
if aspect_width <= 0 or aspect_height <= 0:
|
||||
raise ValueError(f"aspect must be positive, got {aspect_width}:{aspect_height}")
|
||||
ratio = aspect_width / aspect_height
|
||||
if not _MIN_ASPECT <= ratio <= _MAX_ASPECT:
|
||||
raise ValueError(f"aspect ratios run from 1:4 to 4:1, got {aspect_width}:{aspect_height}")
|
||||
|
||||
if ratio >= 1:
|
||||
width, height = _SHORT_EDGE * ratio, float(_SHORT_EDGE)
|
||||
else:
|
||||
width, height = float(_SHORT_EDGE), _SHORT_EDGE / ratio
|
||||
area = width * height
|
||||
if area > _MAX_PIXELS:
|
||||
scale = (_MAX_PIXELS / area)**0.5
|
||||
width, height = width * scale, height * scale
|
||||
m = _CANVAS_MULTIPLE
|
||||
return max(m, round(height / m) * m), max(m, round(width / m) * m)
|
||||
|
||||
|
||||
# The canvases `set_canvas` offers. Deliberately a short list: every entry is a
|
||||
# distinct tensor shape that load() must warm, and an unwarmed shape pays a
|
||||
# one-off compile stall on its first clip.
|
||||
ASPECT_CHOICES: tuple[str, ...] = ("16:9", "1:1", "9:16", "4:3")
|
||||
|
||||
_ASPECT_RATIOS: dict[str, tuple[int, int]] = {
|
||||
"16:9": (16, 9),
|
||||
"1:1": (1, 1),
|
||||
"9:16": (9, 16),
|
||||
"4:3": (4, 3),
|
||||
}
|
||||
|
||||
|
||||
def canvas_for_choice(aspect: str) -> tuple[int, int]:
|
||||
"""Resolve one of ``ASPECT_CHOICES`` to `(height, width)`."""
|
||||
try:
|
||||
ratio = _ASPECT_RATIOS[aspect]
|
||||
except KeyError:
|
||||
raise ValueError(f"unknown aspect {aspect!r}; choose one of {list(ASPECT_CHOICES)}") from None
|
||||
return canvas_for_aspect(*ratio)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ASPECT_CHOICES",
|
||||
"FPS",
|
||||
"MAX_FRAMES",
|
||||
"MAX_SECONDS",
|
||||
"MAX_SECONDS_PUBLISHED",
|
||||
"MIN_FRAMES",
|
||||
"MIN_SECONDS",
|
||||
"MIN_SECONDS_PUBLISHED",
|
||||
"align_frames",
|
||||
"canvas_for_aspect",
|
||||
"canvas_for_choice",
|
||||
"frames_for_seconds",
|
||||
"legal_frame_counts",
|
||||
"seconds_for_frames",
|
||||
]
|
||||
@@ -0,0 +1,148 @@
|
||||
"""The clip queues: what the stream holds, in order, at each stage.
|
||||
|
||||
A clip is enqueued (generation queue), built (playout queue), then consumed by
|
||||
playing. Both stages are the same bounded, ordered, position-addressable
|
||||
container; `engine.py` owns when an entry crosses between them.
|
||||
|
||||
Pure bookkeeping, so it is testable without a GPU.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from . import clip_plan
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClipEntry:
|
||||
"""One clip, from request to built payload.
|
||||
|
||||
Everything but `video`/`audio` is frozen at enqueue time; those two arrive
|
||||
when the build completes, and `ready` is derived from their presence.
|
||||
"""
|
||||
|
||||
clip_id: str
|
||||
prompt: str
|
||||
metadata: str
|
||||
frames: int
|
||||
seed: int
|
||||
# Set while a build for this entry is in flight, so the scheduler never
|
||||
# submits the same entry twice.
|
||||
building: bool = False
|
||||
# The built payload: decoded RGB frames and the wire-ready waveform.
|
||||
video: list[Any] | None = None
|
||||
audio: Any = None
|
||||
|
||||
@property
|
||||
def ready(self) -> bool:
|
||||
return self.video is not None
|
||||
|
||||
@property
|
||||
def seconds(self) -> float:
|
||||
return clip_plan.seconds_for_frames(self.frames)
|
||||
|
||||
def snapshot(self) -> dict[str, Any]:
|
||||
"""The clip as every message that references it carries it.
|
||||
|
||||
Whole rather than an id, so a listener never has to join against an
|
||||
earlier message; a plain mapping, so it is JSON-serialisable for the
|
||||
websocket.
|
||||
"""
|
||||
return {
|
||||
"clip_id": self.clip_id,
|
||||
"prompt": self.prompt,
|
||||
"metadata": self.metadata,
|
||||
"frames": self.frames,
|
||||
"seconds": round(self.seconds, 3),
|
||||
"seed": self.seed,
|
||||
"ready": self.ready,
|
||||
}
|
||||
|
||||
|
||||
def new_entry(*, prompt: str, metadata: str, frames: int, seed: int) -> ClipEntry:
|
||||
"""Mint one entry with a fresh UUID."""
|
||||
return ClipEntry(
|
||||
clip_id=str(uuid.uuid4()),
|
||||
prompt=prompt,
|
||||
metadata=metadata,
|
||||
frames=frames,
|
||||
seed=seed,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ClipQueue:
|
||||
"""A bounded, ordered, position-addressable queue of `ClipEntry`.
|
||||
|
||||
One container serves both stages. Positions are explicit and nothing
|
||||
reorders on its own. For the playout queue every entry holds a fully
|
||||
decoded clip in host memory, so `capacity` is also the memory budget.
|
||||
"""
|
||||
|
||||
capacity: int
|
||||
_entries: list[ClipEntry] = field(default_factory=list)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.capacity < 1:
|
||||
raise ValueError(f"queue capacity must be positive, got {self.capacity}")
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._entries)
|
||||
|
||||
def __contains__(self, entry: ClipEntry) -> bool:
|
||||
return any(existing is entry for existing in self._entries)
|
||||
|
||||
@property
|
||||
def full(self) -> bool:
|
||||
return len(self._entries) >= self.capacity
|
||||
|
||||
def add(self, entry: ClipEntry, position: int | None = None) -> int:
|
||||
"""Insert at `position` (None appends, otherwise clamped) and return the index."""
|
||||
if self.full:
|
||||
raise ValueError(f"the queue is full ({self.capacity} clips)")
|
||||
index = (len(self._entries) if position is None else max(0, min(int(position), len(self._entries))))
|
||||
self._entries.insert(index, entry)
|
||||
return index
|
||||
|
||||
def move(self, entry: ClipEntry, position: int) -> int:
|
||||
"""Reposition `entry` and return the index it landed at, clamped."""
|
||||
if entry not in self:
|
||||
raise ValueError("the clip is not in this queue")
|
||||
self._entries = [existing for existing in self._entries if existing is not entry]
|
||||
index = max(0, min(int(position), len(self._entries)))
|
||||
self._entries.insert(index, entry)
|
||||
return index
|
||||
|
||||
def get(self, clip_id: str) -> ClipEntry | None:
|
||||
for entry in self._entries:
|
||||
if entry.clip_id == clip_id:
|
||||
return entry
|
||||
return None
|
||||
|
||||
def head(self) -> ClipEntry | None:
|
||||
return self._entries[0] if self._entries else None
|
||||
|
||||
def next_to_build(self) -> ClipEntry | None:
|
||||
"""The front-most entry no build is already running for."""
|
||||
for entry in self._entries:
|
||||
if not entry.building:
|
||||
return entry
|
||||
return None
|
||||
|
||||
def remove(self, entry: ClipEntry) -> None:
|
||||
self._entries = [existing for existing in self._entries if existing is not entry]
|
||||
|
||||
def clear(self) -> int:
|
||||
"""Drop every entry, built payloads included, and return how many."""
|
||||
cleared = len(self._entries)
|
||||
self._entries = []
|
||||
return cleared
|
||||
|
||||
def snapshot(self) -> list[dict[str, Any]]:
|
||||
return [entry.snapshot() for entry in self._entries]
|
||||
|
||||
|
||||
__all__ = ["ClipEntry", "ClipQueue", "new_entry"]
|
||||
@@ -0,0 +1,328 @@
|
||||
"""Configuration: one YAML file, plus secrets from the environment.
|
||||
|
||||
`configs/infinite_livestream.yaml` holds everything the app is configured with: what the
|
||||
checkpoint is asked for, how it is hosted, and how the deployment behaves.
|
||||
Point at a copy of it with `--config`.
|
||||
|
||||
API keys stay in the environment, because a key in a version-controlled file is
|
||||
a key that leaks. `LIVESTREAM_WEIGHTS_PATH` is there too, being a property of
|
||||
the machine rather than of the deployment.
|
||||
|
||||
`load_config` is the only reader of either; nothing else touches `os.environ`
|
||||
or parses YAML.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import yaml
|
||||
|
||||
from . import clip_plan
|
||||
|
||||
# ---------------------------------------------------------------- presets
|
||||
|
||||
|
||||
class PresetError(ValueError):
|
||||
"""A preset file is missing or malformed."""
|
||||
|
||||
|
||||
# Inside the package, so it survives installation: the app ships as part of
|
||||
# fastvideo, and anything beside the package rather than in it is not
|
||||
# packaged.
|
||||
DEFAULT_CONFIG = Path(__file__).parent / "configs" / "infinite_livestream.yaml"
|
||||
|
||||
# Where the playlist goes when the config does not say. A relative default
|
||||
# would write into whatever directory the server was started from, which for a
|
||||
# source checkout is the repo root. Mirrors how `apps/dreamverse` picks its
|
||||
# state root.
|
||||
_STATE_ROOT = Path(os.environ.get("XDG_STATE_HOME") or Path.home() / ".local/state") / "fastvideo-livestream"
|
||||
DEFAULT_HLS_DIR = _STATE_ROOT / "hls"
|
||||
DEFAULT_FILLERS_DIR = Path(__file__).parent / "presets"
|
||||
PRESET_FILE = "fillers.json"
|
||||
|
||||
|
||||
def load_preset(directory: str | Path) -> dict:
|
||||
"""Load and validate the style and idle prompts the stream runs on.
|
||||
|
||||
`directory` holds `fillers.json`: the `style` every rewritten scene is
|
||||
written in, and the `idle_prompts` that keep the stream fed when nobody is
|
||||
typing. An empty prompt list disables the filler. Other keys are ignored,
|
||||
so the file can carry its own notes.
|
||||
"""
|
||||
path = Path(directory) / PRESET_FILE
|
||||
if not path.is_file():
|
||||
raise PresetError(f"no {PRESET_FILE} in {directory}")
|
||||
try:
|
||||
preset = json.loads(path.read_text(encoding="utf-8"))
|
||||
except json.JSONDecodeError as error:
|
||||
raise PresetError(f"{path} is not valid JSON: {error}") from None
|
||||
style = preset.get("style")
|
||||
prompts = preset.get("idle_prompts")
|
||||
if not isinstance(style, str) or not style.strip():
|
||||
raise PresetError(f"{path} needs a non-empty string `style`")
|
||||
if not isinstance(prompts, list) or not all(isinstance(p, str) for p in prompts):
|
||||
raise PresetError(f"{path} needs `idle_prompts` as a list of strings")
|
||||
return {
|
||||
"style": style.strip(),
|
||||
"idle_prompts": [p.strip() for p in prompts if p.strip()],
|
||||
}
|
||||
|
||||
|
||||
# ------------------------------------------------------------ model config
|
||||
|
||||
# Component directories the T2VA pipeline loads. Missing weights must kill
|
||||
# startup, not surface as a loader traceback on the first clip.
|
||||
REQUIRED_COMPONENTS = (
|
||||
"transformer",
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"processor",
|
||||
"vae",
|
||||
"audio_vae",
|
||||
"scheduler",
|
||||
"audio_scheduler",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelConfig:
|
||||
"""Everything the engine YAML configures, validated once at load.
|
||||
|
||||
The top-level fields are what the queues and the clip planner need;
|
||||
``inference`` and ``runtime`` are the raw blocks, which the backend reads
|
||||
its engine knobs (attention kernels, compile flags, parallelism, offload
|
||||
policy) straight out of.
|
||||
"""
|
||||
|
||||
aspect: str
|
||||
clip_frames: int
|
||||
seed: int
|
||||
num_inference_steps: int
|
||||
queue_size: int
|
||||
generation_queue_size: int
|
||||
warmup_aspects: tuple[str, ...]
|
||||
warmup_frames: tuple[int, ...]
|
||||
inference: dict[str, Any]
|
||||
runtime: dict[str, Any]
|
||||
|
||||
|
||||
def load_model_config(config_path: Path | None = None) -> ModelConfig:
|
||||
"""Read the `inference` and `runtime` blocks into a validated `ModelConfig`.
|
||||
|
||||
The same file `Config.load` reads. Split out because the queues and the
|
||||
backend need the checkpoint's shape, and nothing else in the file.
|
||||
|
||||
Raises:
|
||||
ValueError: If the configured aspect is not one the checkpoint offers,
|
||||
or a queue size is not positive.
|
||||
"""
|
||||
path = config_path or DEFAULT_CONFIG
|
||||
document: dict[str, Any] = yaml.safe_load(path.read_text(encoding="utf-8")) or {}
|
||||
inference: dict[str, Any] = document.get("inference") or {}
|
||||
runtime: dict[str, Any] = document.get("runtime") or {}
|
||||
|
||||
aspect = str(inference.get("aspect", "16:9"))
|
||||
if aspect not in clip_plan.ASPECT_CHOICES:
|
||||
raise ValueError(f"inference.aspect must be one of {list(clip_plan.ASPECT_CHOICES)}, got {aspect!r}")
|
||||
|
||||
queue_size = int(inference.get("queue_size", 10))
|
||||
if queue_size < 1:
|
||||
raise ValueError(f"inference.queue_size must be positive, got {queue_size}")
|
||||
|
||||
generation_queue_size = int(inference.get("generation_queue_size", 20))
|
||||
if generation_queue_size < 1:
|
||||
raise ValueError(f"inference.generation_queue_size must be positive, got {generation_queue_size}")
|
||||
|
||||
clip_frames = clip_plan.frames_for_seconds(float(inference.get("clip_seconds", clip_plan.MAX_SECONDS)))
|
||||
|
||||
return ModelConfig(
|
||||
aspect=aspect,
|
||||
clip_frames=clip_frames,
|
||||
seed=int(inference.get("seed", 1000)),
|
||||
num_inference_steps=int(inference.get("num_inference_steps", 5)),
|
||||
queue_size=queue_size,
|
||||
generation_queue_size=generation_queue_size,
|
||||
warmup_aspects=tuple(str(a) for a in (inference.get("warmup_aspects") or [aspect])),
|
||||
warmup_frames=_parse_warmup_lengths(inference.get("warmup_lengths"), clip_frames),
|
||||
inference=inference,
|
||||
runtime=runtime,
|
||||
)
|
||||
|
||||
|
||||
def _parse_warmup_lengths(raw: Any, clip_frames: int) -> tuple[int, ...]:
|
||||
"""Resolve ``inference.warmup_lengths`` to the frame counts load() warms.
|
||||
|
||||
``"default"`` (or nothing) warms only the configured clip length;
|
||||
``"all"`` warms every length the checkpoint can generate; a list of
|
||||
seconds warms those, snapped to legal lengths. The default length is
|
||||
always included -- it is the shape every plain enqueue uses.
|
||||
"""
|
||||
if raw in (None, "", "default"):
|
||||
return (clip_frames, )
|
||||
if raw == "all":
|
||||
frames = set(clip_plan.legal_frame_counts())
|
||||
elif isinstance(raw, list | tuple):
|
||||
frames = {clip_plan.frames_for_seconds(float(seconds)) for seconds in raw}
|
||||
else:
|
||||
raise ValueError(f'inference.warmup_lengths must be "default", "all", or a list of seconds, got {raw!r}')
|
||||
frames.add(clip_frames)
|
||||
return tuple(sorted(frames))
|
||||
|
||||
|
||||
def resolve_model_path(config: ModelConfig, weights_root: Path) -> Path:
|
||||
"""The checkpoint directory under the weights path; "." means the path itself."""
|
||||
subdir = str(config.runtime.get("checkpoint_dir", "."))
|
||||
if subdir in ("", "."):
|
||||
return weights_root
|
||||
return weights_root / subdir
|
||||
|
||||
|
||||
def require_weights(root: Path, model_path: Path) -> None:
|
||||
"""Fail startup loudly when the weights are incomplete."""
|
||||
problems: list[str] = []
|
||||
if not model_path.is_dir():
|
||||
problems.append(f"checkpoint directory is missing: {model_path}")
|
||||
else:
|
||||
index = model_path / "modular_model_index.json"
|
||||
if not index.is_file():
|
||||
problems.append(f"modular_model_index.json is missing: {index}")
|
||||
for component in REQUIRED_COMPONENTS:
|
||||
if not (model_path / component).is_dir():
|
||||
problems.append(f"component directory is missing: {model_path / component}")
|
||||
if problems:
|
||||
raise FileNotFoundError(f"FastH3 weights under {root} are incomplete:\n " + "\n ".join(problems))
|
||||
|
||||
|
||||
# -------------------------------------------------------------- app config
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Config:
|
||||
"""One immutable snapshot of everything the app is configured with."""
|
||||
|
||||
# The engine: where the weights live and which YAML shapes it
|
||||
weights_path: Path
|
||||
config_path: Path
|
||||
|
||||
# Upsampling
|
||||
openai_api_key: str
|
||||
openai_base_url: str | None
|
||||
openai_model: str
|
||||
max_chunks: int
|
||||
# Filler always wears the preset's style; a viewer's own request may pick
|
||||
# whatever look suits it. Set 0 to put every clip in the house style.
|
||||
viewer_free_style: bool
|
||||
|
||||
# The style every scene is written in, and the idle prompts
|
||||
style: str
|
||||
idle_prompts: tuple[str, ...]
|
||||
|
||||
# Moderation (its own endpoint: the upsampling gateway may not expose
|
||||
# /moderations, so this can point at api.openai.com while upsampling
|
||||
# goes elsewhere)
|
||||
moderation_enabled: bool
|
||||
moderation_api_key: str
|
||||
moderation_base_url: str | None
|
||||
moderation_model: str
|
||||
|
||||
# Idle filler
|
||||
idle_queue_target: int
|
||||
|
||||
# Output: the HLS playlist the page plays, written by `sink.py`.
|
||||
hls_dir: str
|
||||
video_bitrate_k: int
|
||||
hls_retention_s: int
|
||||
|
||||
# The watch page: video, chat and the queue on one HTTP origin, so a
|
||||
# single tunnel publishes the whole thing.
|
||||
web_host: str
|
||||
web_port: int
|
||||
|
||||
# Chat
|
||||
chat_command: str
|
||||
chat_cooldown_s: float
|
||||
|
||||
@staticmethod
|
||||
def load(argv: list[str] | None = None) -> Config:
|
||||
"""Read the config file and the environment, and validate the result."""
|
||||
parser = argparse.ArgumentParser(description="Chat-driven FastH3 livestream (see README.md).")
|
||||
parser.add_argument("--config", default=None, help=f"config YAML (default {DEFAULT_CONFIG})")
|
||||
parser.add_argument("--weights", default=None, help="override LIVESTREAM_WEIGHTS_PATH")
|
||||
parser.add_argument("--port", default=None, type=int, help="override web.port")
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
path = Path(args.config).expanduser() if args.config else DEFAULT_CONFIG
|
||||
if not path.is_file():
|
||||
raise SystemExit(f"config not found: {path}")
|
||||
document: dict[str, Any] = yaml.safe_load(path.read_text(encoding="utf-8")) or {}
|
||||
upsampler = document.get("upsampler") or {}
|
||||
moderation = document.get("moderation") or {}
|
||||
director = document.get("director") or {}
|
||||
output = document.get("output") or {}
|
||||
web = document.get("web") or {}
|
||||
|
||||
weights = args.weights or os.environ.get("LIVESTREAM_WEIGHTS_PATH", "")
|
||||
openai_key = os.environ.get("OPENAI_API_KEY", "")
|
||||
fillers = director.get("fillers")
|
||||
try:
|
||||
preset = load_preset(Path(fillers).expanduser() if fillers else DEFAULT_FILLERS_DIR)
|
||||
except PresetError as error:
|
||||
raise SystemExit(str(error)) from None
|
||||
|
||||
config = Config(
|
||||
weights_path=Path(weights).expanduser() if weights else Path(),
|
||||
config_path=path,
|
||||
openai_api_key=openai_key,
|
||||
openai_base_url=upsampler.get("base_url") or None,
|
||||
openai_model=str(upsampler.get("model", "gpt-4o-mini")),
|
||||
max_chunks=max(1, int(upsampler.get("max_chunks", 6))),
|
||||
viewer_free_style=bool(upsampler.get("viewer_free_style", True)),
|
||||
style=preset["style"],
|
||||
idle_prompts=tuple(preset["idle_prompts"]),
|
||||
moderation_enabled=bool(moderation.get("enabled", True)),
|
||||
# Falls back to the upsampling credentials, which is right when one
|
||||
# endpoint serves both.
|
||||
moderation_api_key=os.environ.get("MODERATION_API_KEY") or openai_key,
|
||||
moderation_base_url=moderation.get("base_url") or upsampler.get("base_url") or None,
|
||||
moderation_model=str(moderation.get("model", "omni-moderation-latest")),
|
||||
idle_queue_target=int(director.get("idle_queue_target", 6)),
|
||||
hls_dir=str(output.get("hls_dir") or DEFAULT_HLS_DIR),
|
||||
video_bitrate_k=int(output.get("video_bitrate_k", 4500)),
|
||||
hls_retention_s=int(output.get("hls_retention_s", 120)),
|
||||
web_host=str(web.get("host", "0.0.0.0")),
|
||||
web_port=args.port or int(web.get("port", 8081)),
|
||||
chat_command=str(director.get("chat_command", "!prompt")).strip(),
|
||||
chat_cooldown_s=float(director.get("chat_cooldown_s", 10)),
|
||||
)
|
||||
config.validate()
|
||||
return config
|
||||
|
||||
def validate(self) -> None:
|
||||
"""Fail fast on contradictions instead of half-starting."""
|
||||
if not str(self.weights_path) or self.weights_path == Path():
|
||||
raise SystemExit("Set LIVESTREAM_WEIGHTS_PATH, or pass --weights, pointing at the FastH3 weights.")
|
||||
if not self.openai_api_key:
|
||||
raise SystemExit(
|
||||
"Set OPENAI_API_KEY. Prompt rewriting runs for the idle filler too, so the stream does not start without it."
|
||||
)
|
||||
if self.hls_retention_s < 6:
|
||||
raise SystemExit("output.hls_retention_s must be at least 6 seconds.")
|
||||
if not self.chat_command.startswith("!"):
|
||||
raise SystemExit("director.chat_command should start with '!' (e.g. !prompt).")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"Config",
|
||||
"ModelConfig",
|
||||
"PresetError",
|
||||
"load_model_config",
|
||||
"load_preset",
|
||||
"require_weights",
|
||||
"resolve_model_path",
|
||||
]
|
||||
@@ -0,0 +1,93 @@
|
||||
# Everything Infinite Livestream is configured with, apart from secrets.
|
||||
#
|
||||
# API keys stay in the environment, because a key in a version-controlled file
|
||||
# is a key that leaks:
|
||||
#
|
||||
# export OPENAI_API_KEY=...
|
||||
# export LIVESTREAM_WEIGHTS_PATH=/path/to/fasth3
|
||||
#
|
||||
# Point the app at a copy of this file with `infinite-livestream-server --config`.
|
||||
|
||||
inference:
|
||||
# The canvas every clip is generated at. 16:9 resolves to 1344x768.
|
||||
aspect: "16:9"
|
||||
# The longest clip this checkpoint can make (362 frames at 24 fps). One
|
||||
# length everywhere means one compiled shape.
|
||||
clip_seconds: 15.083
|
||||
# Built clips held in host memory awaiting playout; each holds a full
|
||||
# decoded clip, so this is also the memory budget.
|
||||
queue_size: 10
|
||||
# Requests accepted but not yet built.
|
||||
generation_queue_size: 20
|
||||
seed: 1000
|
||||
# Sigma-grid POINTS, not transformer forwards: the distilled schedule is
|
||||
# five points and exactly four forwards.
|
||||
num_inference_steps: 5
|
||||
|
||||
vsa_sparsity: 0.9
|
||||
vsa_tile_size: 64
|
||||
# sm100a is the Blackwell VSA kernel; `triton` is the ~2.5x slower fallback.
|
||||
vsa_kernel: sm100a
|
||||
fa4: true
|
||||
h3_fusions: true
|
||||
inference_torch_compile: true
|
||||
compile_vae: true
|
||||
ulysses_a2a: "off"
|
||||
|
||||
# Shapes warmed before the stream reports ready. "default" warms only
|
||||
# clip_seconds; "all" warms every legal length (slower start, no stall on a
|
||||
# viewer's first odd-length clip).
|
||||
warmup_aspects: ["16:9"]
|
||||
warmup_lengths: "default"
|
||||
|
||||
runtime:
|
||||
# Relative to the weights root (LIVESTREAM_WEIGHTS_PATH); "." means the
|
||||
# snapshot's components sit directly under it.
|
||||
checkpoint_dir: "."
|
||||
num_gpus: 4
|
||||
# Replicate the transformer on each GPU rather than FSDP-sharding it: at
|
||||
# four GB200s the weights fit, and replication skips the all-gather.
|
||||
replicated_dit: true
|
||||
offload_text_encoder: false
|
||||
offload_vae: false
|
||||
pin_cpu_memory: false
|
||||
|
||||
# How a viewer's prompt becomes scenes. Any OpenAI-compatible endpoint works.
|
||||
upsampler:
|
||||
model: gpt-4o-mini
|
||||
# base_url: https://api.groq.com/openai/v1
|
||||
# Most clips one prompt may expand into.
|
||||
max_chunks: 6
|
||||
# Filler always wears the house style; a viewer's own request may pick
|
||||
# whatever look suits it. false puts every clip in the house style.
|
||||
viewer_free_style: true
|
||||
|
||||
# Checked before rewriting. Errors fail closed, so a broken endpoint stops
|
||||
# prompts rather than letting them through unchecked. Its own endpoint on
|
||||
# purpose: inference gateways rarely expose /moderations.
|
||||
moderation:
|
||||
enabled: true
|
||||
model: omni-moderation-latest
|
||||
# base_url: https://api.openai.com/v1
|
||||
|
||||
director:
|
||||
# Clips the idle filler keeps queued when nobody is typing. 0 turns it off.
|
||||
idle_queue_target: 6
|
||||
# Seconds between accepted prompts, per viewer.
|
||||
chat_cooldown_s: 10
|
||||
chat_command: "!prompt"
|
||||
# Directory holding fillers.json. Unset uses the one that ships.
|
||||
# fillers: /path/to/my-fillers
|
||||
|
||||
output:
|
||||
# Unset writes under $XDG_STATE_HOME/fastvideo-livestream/hls, so a source
|
||||
# checkout does not collect segments. Set an absolute path to place it.
|
||||
# hls_dir: /var/lib/livestream/hls
|
||||
# x264 target in kbit/s.
|
||||
video_bitrate_k: 4500
|
||||
# Retained playback history; does not increase target live latency.
|
||||
hls_retention_s: 120
|
||||
|
||||
web:
|
||||
host: 0.0.0.0
|
||||
port: 8081
|
||||
@@ -0,0 +1,408 @@
|
||||
"""The director: viewer prompts in, tagged scene groups on the engine's queue.
|
||||
|
||||
One chat prompt becomes one *scene group*: the upsampler expands it into 1..N
|
||||
self-contained scenes -- a single shot, or a chunked short story -- which the
|
||||
director enqueues contiguously. It is also the playout brain: `run_playout`
|
||||
curates the front of the playout queue with `move` so the engine's next
|
||||
autoplay is already the right clip.
|
||||
|
||||
Rules that keep it coherent:
|
||||
|
||||
* It is the queue's only writer. The viewer worker (`run`) and the idle
|
||||
filler (`run_idle`) serialise their enqueues through one lock, so groups
|
||||
can never interleave.
|
||||
* A group is enqueued only when the whole group fits, so it cannot get
|
||||
stuck half-in. Capacities come from the engine's `state_update`, never
|
||||
from constants here.
|
||||
* Viewer prompts outrank filler and stay first-come-first-served among
|
||||
themselves: viewer groups insert ahead of waiting filler and behind
|
||||
waiting viewer clips, the playout loop pops one built filler when a full
|
||||
playout queue blocks a viewer's build, and the idle filler stands down
|
||||
whenever viewer work is pending.
|
||||
|
||||
Every scene carries its group's identity in the clip metadata, which the
|
||||
engine echoes back on every message referencing that clip. That is what lets
|
||||
"scene 2/3 of Neon Alley by viewer_42" be reconstructed from a `clip_started`
|
||||
alone, and what marks filler as evictable later.
|
||||
|
||||
Viewer prompts pass moderation before the upsampler; the curated idle list
|
||||
does not need it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import random
|
||||
import time
|
||||
from collections.abc import Callable, Sequence
|
||||
|
||||
from .chat import ChatPrompt
|
||||
from .group_tag import is_generated, parse_group_tag, pick_next, viewer_insert_position
|
||||
from .engine import Engine
|
||||
from .moderator import Moderator
|
||||
from .upsampler import PromptUpsampler, SceneGroup
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Prompts waiting for upsampling+enqueue before new ones are turned away.
|
||||
# Depth here is viewer wait time, and a backlog on top of a full generation
|
||||
# queue serves nobody.
|
||||
_PENDING_LIMIT = 24
|
||||
|
||||
# Enqueue retry cadence while the model refuses (reconnect mid-command, ...).
|
||||
_RETRY_DELAY_S = 3.0
|
||||
|
||||
# How often the idle filler re-checks whether the queue wants topping up.
|
||||
_IDLE_POLL_S = 3.0
|
||||
|
||||
# How often the playout loop re-checks. The broadcasts keep the mirrors
|
||||
# fresh; polling them is what survives a missed message.
|
||||
_PLAYOUT_POLL_S = 0.5
|
||||
|
||||
|
||||
class Director:
|
||||
"""Consume chat prompts; keep the fast-h3 queue fed with scene groups."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
link: Engine,
|
||||
upsampler: PromptUpsampler,
|
||||
moderator: Moderator,
|
||||
cooldown_s: float,
|
||||
idle_prompts: Sequence[str] = (),
|
||||
idle_queue_target: int = 0,
|
||||
on_reject: Callable[[str, str], None] | None = None,
|
||||
) -> None:
|
||||
self._link = link
|
||||
self._on_reject = on_reject
|
||||
self._upsampler = upsampler
|
||||
self._moderator = moderator
|
||||
self._cooldown_s = cooldown_s
|
||||
self._idle_prompts = list(idle_prompts)
|
||||
random.shuffle(self._idle_prompts)
|
||||
self._idle_index = 0
|
||||
self._idle_target = idle_queue_target
|
||||
self._pending: asyncio.Queue[ChatPrompt] = asyncio.Queue(_PENDING_LIMIT)
|
||||
self._last_accepted: dict[str, float] = {} # author -> monotonic
|
||||
self._enqueue_lock = asyncio.Lock()
|
||||
link.add_listener(self._on_model_message)
|
||||
|
||||
# -------------------------------------------------------- chat intake
|
||||
|
||||
def cooldown_remaining(self, author: str) -> float:
|
||||
"""Seconds until *author* may send again; 0 when they may send now.
|
||||
|
||||
Asked by the web app before it accepts a POST, so a rate-limited
|
||||
viewer is stopped in their own browser rather than told afterwards in
|
||||
a chat feed everyone else can read.
|
||||
"""
|
||||
last = self._last_accepted.get(author)
|
||||
if last is None:
|
||||
return 0.0
|
||||
return max(0.0, self._cooldown_s - (time.monotonic() - last))
|
||||
|
||||
def _reject(self, prompt: ChatPrompt, reason: str) -> None:
|
||||
"""Drop one prompt, and make sure the viewer hears about it.
|
||||
|
||||
Every rejection below used to be log-only, while the web app had
|
||||
already answered the POST with `ok` and echoed the prompt into chat --
|
||||
so a viewer watched their request appear and then quietly die. The
|
||||
component that can say no is downstream of the acknowledgement, which
|
||||
is why it has to report back rather than return a status.
|
||||
"""
|
||||
logger.info("[director] dropped from %s@%s (%s): %s", prompt.author, prompt.source, reason, prompt.text)
|
||||
if self._on_reject is None:
|
||||
return
|
||||
try:
|
||||
self._on_reject(prompt.author, reason)
|
||||
except Exception: # noqa: BLE001 -- telling the viewer must not kill the loop
|
||||
logger.exception("[director] reject callback failed")
|
||||
|
||||
def submit(self, prompt: ChatPrompt) -> None:
|
||||
"""Accept one chat prompt (called synchronously by chat sources)."""
|
||||
now = time.monotonic()
|
||||
last = self._last_accepted.get(prompt.author)
|
||||
if last is not None and now - last < self._cooldown_s:
|
||||
self._reject(prompt, f"one prompt every {self._cooldown_s:.0f}s; "
|
||||
f"{self._cooldown_s - (now - last):.0f}s left")
|
||||
return
|
||||
try:
|
||||
self._pending.put_nowait(prompt)
|
||||
except asyncio.QueueFull:
|
||||
self._reject(prompt, f"the backlog is full ({_PENDING_LIMIT} waiting)")
|
||||
return
|
||||
self._last_accepted[prompt.author] = now
|
||||
logger.info(
|
||||
"[director] accepted from %s@%s: %s",
|
||||
prompt.author,
|
||||
prompt.source,
|
||||
prompt.text,
|
||||
)
|
||||
|
||||
# ------------------------------------------------- viewer prompt loop
|
||||
|
||||
def _viewer_clips_queued(self) -> int:
|
||||
"""Viewer clips across both queues (anything not tagged filler)."""
|
||||
return sum(1 for clip in self._link.generation_clips + self._link.playout_clips if not is_generated(clip))
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Moderate, upsample, and enqueue pending prompts, one group at a time."""
|
||||
while True:
|
||||
prompt = await self._pending.get()
|
||||
try:
|
||||
# Dropped now, before it costs a moderation and an LLM call.
|
||||
# Capacity comes from the engine, never from a constant.
|
||||
if self._viewer_clips_queued() >= self._link.playout_capacity:
|
||||
self._reject(
|
||||
prompt, f"{self._viewer_clips_queued()} viewer clips already queued "
|
||||
f"(budget {self._link.playout_capacity})")
|
||||
continue
|
||||
verdict = await self._moderator.review(prompt.text)
|
||||
if verdict is not None:
|
||||
self._reject(prompt, verdict)
|
||||
continue
|
||||
group = await self._upsampler.upsample(
|
||||
raw_prompt=prompt.text,
|
||||
author=prompt.author,
|
||||
source=prompt.source,
|
||||
min_seconds=self._link.min_seconds,
|
||||
max_seconds=self._link.max_seconds,
|
||||
)
|
||||
await self._enqueue_group(group)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as error:
|
||||
logger.error(
|
||||
"[director] failed to process prompt from %s: %s",
|
||||
prompt.author,
|
||||
error,
|
||||
)
|
||||
|
||||
# -------------------------------------------------------- idle filler
|
||||
|
||||
async def run_idle(self) -> None:
|
||||
"""Keep the queue topped up with generated clips while chat is quiet.
|
||||
|
||||
One clip per group, on purpose: single-scene fillers are the finest
|
||||
eviction granularity, and popping one never truncates a story.
|
||||
"""
|
||||
if self._idle_target <= 0:
|
||||
logger.info("[director] idle filler disabled (target 0)")
|
||||
return
|
||||
logger.info(
|
||||
"[director] idle filler: %d prompts, queue target %d",
|
||||
len(self._idle_prompts),
|
||||
self._idle_target,
|
||||
)
|
||||
while True:
|
||||
await asyncio.sleep(_IDLE_POLL_S)
|
||||
# May be empty after a switch to a preset with no idle prompts;
|
||||
# keep polling so a later switch revives it without a restart.
|
||||
if not self._idle_prompts:
|
||||
continue
|
||||
# The configured target self-clamps under the deployment's live
|
||||
# playout capacity: filler must never be what fills the playout
|
||||
# queue to the brim, because a full playout queue pauses builds
|
||||
# (leave at least one slot's headroom for a viewer clip to land).
|
||||
target = min(self._idle_target, max(1, self._link.playout_capacity - 1))
|
||||
if (not self._pending.empty() or not self._link.connected
|
||||
or self._link.generation_queued + self._link.playout_queued >= target):
|
||||
continue
|
||||
text = self._idle_prompts[self._idle_index % len(self._idle_prompts)]
|
||||
self._idle_index += 1
|
||||
try:
|
||||
group = await self._upsampler.upsample(
|
||||
raw_prompt=text,
|
||||
author="auto",
|
||||
source="idle",
|
||||
min_seconds=self._link.min_seconds,
|
||||
max_seconds=self._link.max_seconds,
|
||||
generated=True,
|
||||
max_chunks=1,
|
||||
)
|
||||
# A viewer prompt that arrived while the LLM ran outranks the
|
||||
# filler; drop this group rather than making the viewer wait.
|
||||
if self._pending.empty():
|
||||
await self._enqueue_group(group)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as error:
|
||||
logger.error("[director] idle fill failed: %s", error)
|
||||
|
||||
# ------------------------------------------------------------- playout
|
||||
|
||||
async def run_playout(self) -> None:
|
||||
"""Curate the playout queue's front so autoplay always starts right.
|
||||
|
||||
The engine chains the playout front the instant the stream idles, so
|
||||
nothing here sends an explicit play; it keeps the front correct
|
||||
instead. Reordering happens while a clip plays, ahead of the moment it
|
||||
matters.
|
||||
"""
|
||||
while True:
|
||||
await asyncio.sleep(_PLAYOUT_POLL_S)
|
||||
if not self._link.connected:
|
||||
continue
|
||||
await self._relieve_build_backpressure()
|
||||
clips = self._link.playout_clips
|
||||
desired = pick_next(clips)
|
||||
if desired is None or clips[0]["clip_id"] == desired["clip_id"]:
|
||||
continue
|
||||
await self._link.send_command("move", {"clip_id": desired["clip_id"], "position": 0})
|
||||
# Let the resulting queue_update land before re-evaluating.
|
||||
await asyncio.sleep(_PLAYOUT_POLL_S)
|
||||
|
||||
async def _relieve_build_backpressure(self) -> None:
|
||||
"""Pop one playout filler when built fillers block a viewer's build.
|
||||
|
||||
Generation pauses while the playout queue is full. When what fills it
|
||||
is idle filler and a viewer clip waits to build, the newest filler is
|
||||
the right thing to lose — one per tick, so a draining queue gets
|
||||
every chance to make room by playing instead.
|
||||
"""
|
||||
if self._link.playout_queued < self._link.playout_capacity:
|
||||
return
|
||||
viewer_waiting = any(not is_generated(clip) for clip in self._link.generation_clips)
|
||||
if not viewer_waiting:
|
||||
return
|
||||
for clip in reversed(self._link.playout_clips):
|
||||
if is_generated(clip):
|
||||
reply = await self._link.send_command("pop", {"clip_id": clip["clip_id"]})
|
||||
if isinstance(reply, dict) and "clip" in reply:
|
||||
logger.info(
|
||||
"[director] popped playout filler %s to unblock a "
|
||||
"viewer build",
|
||||
clip["clip_id"][:8],
|
||||
)
|
||||
return
|
||||
|
||||
# ---------------------------------------------------------- enqueueing
|
||||
|
||||
async def _enqueue_group(self, group: SceneGroup) -> None:
|
||||
"""Put one group on the model's generation queue, or drop it and say why.
|
||||
|
||||
Viewer groups enter *ahead of waiting filler and behind waiting
|
||||
viewer clips* (`viewer_insert_position`), so viewer requests stay
|
||||
first-come-first-served and idle filler just slides back — no
|
||||
popping, no waste. Filler groups append. When the generation queue
|
||||
cannot fit the group even after dropping the filler waiting in it,
|
||||
the group is dropped with the queues intact — a backlog full of
|
||||
viewer content takes no more, rather than stalling every later
|
||||
prompt behind a wait.
|
||||
"""
|
||||
scene_count = len(group.scenes)
|
||||
async with self._enqueue_lock:
|
||||
free = self._link.generation_capacity - self._link.generation_queued
|
||||
if free < scene_count and not group.generated:
|
||||
evictable = sum(1 for clip in self._link.generation_clips if is_generated(clip))
|
||||
if free + evictable >= scene_count:
|
||||
await self._evict_generation_fillers(scene_count - free)
|
||||
await asyncio.sleep(0.3) # let the pops' queue_update land
|
||||
free = (self._link.generation_capacity - self._link.generation_queued)
|
||||
if free < scene_count:
|
||||
logger.warning(
|
||||
"[director] no room in the generation queue for %s "
|
||||
"(%d scenes, %d free); dropping the group",
|
||||
group.group_id,
|
||||
scene_count,
|
||||
free,
|
||||
)
|
||||
return
|
||||
|
||||
position = (None if group.generated else viewer_insert_position(self._link.generation_clips))
|
||||
for index, scene in enumerate(group.scenes, start=1):
|
||||
metadata = json.dumps(
|
||||
{
|
||||
"group_id": group.group_id,
|
||||
"title": group.title[:120],
|
||||
"scene": index,
|
||||
"scenes": scene_count,
|
||||
"author": group.author,
|
||||
"source": group.source,
|
||||
"generated": group.generated,
|
||||
# Truncated so the blob stays small in the app's own
|
||||
# clip records; no FastVideo schema caps it.
|
||||
"raw_prompt": group.raw_prompt[:400],
|
||||
},
|
||||
ensure_ascii=False,
|
||||
)
|
||||
payload = {
|
||||
"prompt": scene.prompt,
|
||||
"metadata": metadata,
|
||||
"seconds": scene.seconds,
|
||||
}
|
||||
if position is not None:
|
||||
# Consecutive positions keep the group contiguous and in
|
||||
# scene order, ahead of the filler it displaced.
|
||||
payload["position"] = position + index - 1
|
||||
while True:
|
||||
reply = await self._link.send_command("enqueue", payload)
|
||||
if isinstance(reply, dict) and "clip" in reply:
|
||||
clip = reply["clip"]
|
||||
logger.info(
|
||||
"[director] queued %s scene %d/%d as %s (%.1fs, seed %s)%s",
|
||||
group.group_id,
|
||||
index,
|
||||
scene_count,
|
||||
clip["clip_id"][:8],
|
||||
clip["seconds"],
|
||||
clip["seed"],
|
||||
" [auto]" if group.generated else "",
|
||||
)
|
||||
break
|
||||
# A bodyless reply means refused; the engine already
|
||||
# logged why. Wait and retry.
|
||||
logger.warning(
|
||||
"[director] enqueue of %s scene %d/%d refused; retrying in %.0fs",
|
||||
group.group_id,
|
||||
index,
|
||||
scene_count,
|
||||
_RETRY_DELAY_S,
|
||||
)
|
||||
await asyncio.sleep(_RETRY_DELAY_S)
|
||||
|
||||
async def _evict_generation_fillers(self, needed: int) -> int:
|
||||
"""Pop up to `needed` filler clips from the generation queue.
|
||||
|
||||
Capacity relief only — order needs no eviction now that viewer
|
||||
groups insert ahead of filler positionally. Newest-queued first, and
|
||||
only clips tagged `generated: true`. Returns how many pops succeeded.
|
||||
"""
|
||||
popped = 0
|
||||
for clip in reversed(self._link.generation_clips):
|
||||
if popped >= needed:
|
||||
break
|
||||
if not is_generated(clip):
|
||||
continue
|
||||
reply = await self._link.send_command("pop", {"clip_id": clip["clip_id"]})
|
||||
if isinstance(reply, dict) and "clip" in reply:
|
||||
popped += 1
|
||||
logger.info(
|
||||
"[director] evicted waiting filler %s for a viewer group",
|
||||
clip["clip_id"][:8],
|
||||
)
|
||||
return popped
|
||||
|
||||
# ----------------------------------------------------- announcements
|
||||
|
||||
def _on_model_message(self, kind: str, data: dict) -> None:
|
||||
"""Narrate group playback from clip messages alone (via metadata)."""
|
||||
clip = data.get("clip") if isinstance(data, dict) else None
|
||||
if not isinstance(clip, dict):
|
||||
return
|
||||
tag = parse_group_tag(clip.get("metadata", ""))
|
||||
label = (f"'{tag['title']}' scene {tag['scene']}/{tag['scenes']} "
|
||||
f"(by {tag['author']}@{tag['source']})" +
|
||||
(" [auto]" if tag.get("generated") else "") if tag else f"clip {clip.get('clip_id', '?')[:8]}")
|
||||
if kind == "clip_started":
|
||||
logger.info("[now playing] %s", label)
|
||||
elif kind == "clip_finished":
|
||||
logger.info("[finished] %s", label)
|
||||
elif kind == "clip_failed":
|
||||
logger.error(
|
||||
"[director] build failed for %s: %s — the queue moves on",
|
||||
label,
|
||||
data.get("reason"),
|
||||
)
|
||||
@@ -0,0 +1,475 @@
|
||||
"""The engine: generation, playout, and the state every other module reads.
|
||||
|
||||
director ──enqueue/pop/move──▶ Engine ──frames+audio──▶ Pacer ──▶ sink
|
||||
│
|
||||
└──state_update / queue_update / clip_*
|
||||
──▶ listeners (webapp, director)
|
||||
|
||||
The generator and the broadcast share a process, so a built clip is handed to
|
||||
the pacer as the arrays it already is. There is no encode, no transport and
|
||||
therefore nothing that can shed video frames while audio flows on -- which is
|
||||
how a picture drifts behind its own soundtrack.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import time
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from . import clip_plan
|
||||
from .backend import ClipJob, FastH3Backend
|
||||
from .clip_queue import ClipEntry, ClipQueue, new_entry
|
||||
from .config import Config, ModelConfig, require_weights, resolve_model_path
|
||||
from .metadata import clip_view, encode_id3
|
||||
from .pacer import Pacer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Fixed by the checkpoint and the backend's resample; the canvas is not.
|
||||
MODEL_FPS = clip_plan.FPS
|
||||
MODEL_SAMPLE_RATE = 48_000
|
||||
|
||||
POLL_SECONDS = 0.05
|
||||
|
||||
# Frames handed to the pacer per step. Small keeps its buffers near-empty,
|
||||
# which is the condition its A/V pairing depends on.
|
||||
EMIT_FRAMES = 4
|
||||
|
||||
|
||||
class Engine:
|
||||
"""Own the model, the queues and the playout."""
|
||||
|
||||
def __init__(self, config: Config, model_config: ModelConfig) -> None:
|
||||
self._config = config
|
||||
self._model = model_config
|
||||
self._pacer: Pacer | None = None
|
||||
self._listeners: list[Callable[[str, dict], None]] = []
|
||||
|
||||
model_path = resolve_model_path(model_config, config.weights_path)
|
||||
require_weights(config.weights_path, model_path)
|
||||
self.backend = FastH3Backend(model_config, model_path)
|
||||
|
||||
self._generation = ClipQueue(capacity=model_config.generation_queue_size)
|
||||
self._playout = ClipQueue(capacity=model_config.queue_size)
|
||||
# The build in flight: its entry, its job handle, and when it was
|
||||
# submitted (monotonic), so readiness latency is a measured number.
|
||||
self._build: tuple[ClipEntry, ClipJob, float] | None = None
|
||||
self._playing: ClipEntry | None = None
|
||||
|
||||
self._seed = model_config.seed
|
||||
self._clips_played = 0
|
||||
self._frames_sent = 0
|
||||
self._seconds_sent = 0.0
|
||||
|
||||
self._ready = asyncio.Event()
|
||||
# Mirrors of what listeners were last told, so a late subscriber (the
|
||||
# web app builds its own mirror from these) reads the same values.
|
||||
self.state: dict[str, Any] = self._snapshot()
|
||||
self.generation_clips: list[dict] = []
|
||||
self.playout_clips: list[dict] = []
|
||||
|
||||
# ---------------------------------------------------------------- wiring
|
||||
|
||||
def attach_pacer(self, pacer: Pacer) -> None:
|
||||
"""Point the media path at the pacer."""
|
||||
self._pacer = pacer
|
||||
|
||||
def add_listener(self, listener: Callable[[str, dict], None]) -> None:
|
||||
"""Register for every message as `(kind, data)`. Must not raise."""
|
||||
self._listeners.append(listener)
|
||||
|
||||
# ----------------------------------------------------------- state mirror
|
||||
|
||||
@property
|
||||
def min_seconds(self) -> float:
|
||||
return clip_plan.MIN_SECONDS_PUBLISHED
|
||||
|
||||
@property
|
||||
def max_seconds(self) -> float:
|
||||
return clip_plan.MAX_SECONDS_PUBLISHED
|
||||
|
||||
@property
|
||||
def generation_queued(self) -> int:
|
||||
return len(self._generation)
|
||||
|
||||
@property
|
||||
def generation_capacity(self) -> int:
|
||||
return self._generation.capacity
|
||||
|
||||
@property
|
||||
def playout_queued(self) -> int:
|
||||
return len(self._playout)
|
||||
|
||||
@property
|
||||
def playout_capacity(self) -> int:
|
||||
return self._playout.capacity
|
||||
|
||||
@property
|
||||
def canvas(self) -> tuple[int, int]:
|
||||
"""(width, height) this deployment generates at."""
|
||||
height, width = clip_plan.canvas_for_choice(self._model.aspect)
|
||||
return width, height
|
||||
|
||||
@property
|
||||
def connected(self) -> bool:
|
||||
"""Whether the model is loaded and commands would take effect."""
|
||||
return self._ready.is_set()
|
||||
|
||||
def _canvas_hw(self) -> tuple[int, int]:
|
||||
return clip_plan.canvas_for_choice(self._model.aspect)
|
||||
|
||||
def _snapshot(self) -> dict[str, Any]:
|
||||
"""Everything an observer can see, in one mapping.
|
||||
|
||||
The single source, so a joining viewer's greeting and everyone else's
|
||||
broadcast can never disagree.
|
||||
"""
|
||||
height, width = self._canvas_hw()
|
||||
return {
|
||||
"width": width,
|
||||
"height": height,
|
||||
"playing": self._playing is not None,
|
||||
"generation_queued": len(self._generation),
|
||||
"generation_capacity": self._generation.capacity,
|
||||
"playout_queued": len(self._playout),
|
||||
"playout_capacity": self._playout.capacity,
|
||||
"clips_played": self._clips_played,
|
||||
}
|
||||
|
||||
# -------------------------------------------------------------- messaging
|
||||
|
||||
def _emit(self, kind: str, data: dict) -> None:
|
||||
"""Fan one message out to every listener.
|
||||
|
||||
Synchronous and non-throwing: these are in-process callbacks, and a
|
||||
broken listener must not take generation down with it.
|
||||
"""
|
||||
for listener in self._listeners:
|
||||
try:
|
||||
listener(kind, data)
|
||||
except Exception: # noqa: BLE001 -- a listener cannot break the engine
|
||||
logger.exception("[engine] listener failed on %s", kind)
|
||||
|
||||
def _send_state(self) -> None:
|
||||
self.state = self._snapshot()
|
||||
self._emit("state_update", self.state)
|
||||
|
||||
def _send_queue(self) -> None:
|
||||
self.generation_clips = self._generation.snapshot()
|
||||
self.playout_clips = self._playout.snapshot()
|
||||
self._emit(
|
||||
"queue_update",
|
||||
{
|
||||
"generation": self.generation_clips,
|
||||
"playout": self.playout_clips
|
||||
},
|
||||
)
|
||||
|
||||
def _refuse(self, command: str, reason: str) -> None:
|
||||
logger.warning("[engine] %s refused: %s", command, reason)
|
||||
self._emit("command_error", {"command": command, "reason": reason})
|
||||
|
||||
# --------------------------------------------------------------- commands
|
||||
|
||||
async def send_command(self, command: str, data: dict) -> Any:
|
||||
"""Dispatch one command, once the engine is up.
|
||||
|
||||
Awaiting readiness (rather than failing) is what lets the director
|
||||
start before the model has finished loading: its first enqueue simply
|
||||
lands when the engine is ready for it. A ``None`` reply means the
|
||||
command was refused, and `command_error` carried the reason.
|
||||
"""
|
||||
await self._ready.wait()
|
||||
handler = {
|
||||
"enqueue": self._enqueue,
|
||||
"pop": self._pop,
|
||||
"move": self._move,
|
||||
}.get(command)
|
||||
if handler is None:
|
||||
self._refuse(command, f"Unknown command {command!r}.")
|
||||
return None
|
||||
try:
|
||||
return handler(data or {})
|
||||
except Exception as error: # noqa: BLE001 -- reported, never fatal
|
||||
logger.exception("[engine] %s raised", command)
|
||||
self._refuse(command, str(error))
|
||||
return None
|
||||
|
||||
def _enqueue(self, data: dict) -> dict | None:
|
||||
prompt = str(data.get("prompt") or "").strip()
|
||||
if not prompt:
|
||||
self._refuse("enqueue", "The prompt is empty; a clip needs one.")
|
||||
return None
|
||||
if self._generation.full:
|
||||
self._refuse(
|
||||
"enqueue",
|
||||
f"The generation queue is full ({self._generation.capacity} clips).",
|
||||
)
|
||||
return None
|
||||
|
||||
seed = data.get("seed")
|
||||
if not isinstance(seed, int):
|
||||
# The stream's advancing default; an explicit seed leaves it
|
||||
# untouched, so explicit and automatic seeding do not interfere.
|
||||
seed = self._seed
|
||||
self._seed += 1
|
||||
seconds = data.get("seconds")
|
||||
frames = (clip_plan.frames_for_seconds(float(seconds)) if isinstance(seconds, int
|
||||
| float) else self._model.clip_frames)
|
||||
position = data.get("position")
|
||||
entry = new_entry(
|
||||
prompt=prompt,
|
||||
metadata=str(data.get("metadata") or ""),
|
||||
frames=frames,
|
||||
seed=seed,
|
||||
)
|
||||
self._generation.add(entry, position if isinstance(position, int) else None)
|
||||
self._emit("clip_queued", {"clip": entry.snapshot()})
|
||||
self._send_queue()
|
||||
self._send_state()
|
||||
return {"clip": entry.snapshot()}
|
||||
|
||||
def _pop(self, data: dict) -> dict | None:
|
||||
"""Take one clip out of whichever queue holds it."""
|
||||
clip_id = str(data.get("clip_id") or "")
|
||||
entry = ((self._generation.get(clip_id) or self._playout.get(clip_id)) if clip_id else None)
|
||||
if entry is None:
|
||||
self._refuse(
|
||||
"pop",
|
||||
f"No queued clip has id {clip_id!r}."
|
||||
if clip_id else "Pass the `clip_id` of the queued clip to remove.",
|
||||
)
|
||||
return None
|
||||
self._generation.remove(entry)
|
||||
self._playout.remove(entry)
|
||||
# A build already running for it finishes and is discarded; the queues
|
||||
# own what exists, so a result with no entry has nowhere to land.
|
||||
if self._build is not None and self._build[0] is entry:
|
||||
self._build[1].cancelled = True
|
||||
self._emit("clip_popped", {"clip": entry.snapshot()})
|
||||
self._send_queue()
|
||||
self._send_state()
|
||||
return {"clip": entry.snapshot()}
|
||||
|
||||
def _move(self, data: dict) -> dict | None:
|
||||
clip_id = str(data.get("clip_id") or "")
|
||||
position = data.get("position")
|
||||
position = position if isinstance(position, int) else 0
|
||||
entry = self._generation.get(clip_id) if clip_id else None
|
||||
queue, name = (self._generation, "generation")
|
||||
if entry is None and clip_id:
|
||||
entry = self._playout.get(clip_id)
|
||||
queue, name = self._playout, "playout"
|
||||
if entry is None:
|
||||
self._refuse(
|
||||
"move",
|
||||
f"No queued clip has id {clip_id!r}." if clip_id else "Pass the `clip_id` of the queued clip to move.",
|
||||
)
|
||||
return None
|
||||
landed = queue.move(entry, position)
|
||||
self._send_queue()
|
||||
return {"clip": entry.snapshot(), "queue": name, "position": landed}
|
||||
|
||||
# -------------------------------------------------------------- lifecycle
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Load the model, then generate and play forever.
|
||||
|
||||
Loading is minutes of GPU work, so it runs on a thread: the web app is
|
||||
already serving by then, and a viewer sees the page rather than a
|
||||
connection refused.
|
||||
"""
|
||||
height, width = self._canvas_hw()
|
||||
logger.info(
|
||||
"[engine] loading FastH3 on %d gpu(s) at %dx%d, %d-frame clips",
|
||||
int(self._model.runtime.get("num_gpus", 4)),
|
||||
width,
|
||||
height,
|
||||
self._model.clip_frames,
|
||||
)
|
||||
started = time.monotonic()
|
||||
await asyncio.to_thread(self.backend.load)
|
||||
logger.info("[engine] model ready in %.1fs", time.monotonic() - started)
|
||||
|
||||
self._ready.set()
|
||||
self._send_state()
|
||||
self._send_queue()
|
||||
|
||||
while True:
|
||||
try:
|
||||
self._pump_builds()
|
||||
entry = self._playout.head()
|
||||
if entry is not None:
|
||||
self._playout.remove(entry)
|
||||
self._send_queue()
|
||||
await self._play_clip(entry)
|
||||
else:
|
||||
await asyncio.sleep(POLL_SECONDS)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception: # noqa: BLE001 -- the serve loop must survive anything
|
||||
logger.exception("[engine] error in the serve loop")
|
||||
await asyncio.sleep(POLL_SECONDS)
|
||||
|
||||
# ------------------------------------------------------------- generation
|
||||
|
||||
def _pump_builds(self) -> None:
|
||||
"""Apply a finished build and keep the worker fed, without blocking.
|
||||
|
||||
Called from the idle loop and from every playout slice, so clips keep
|
||||
building while another one streams. The generation queue is consumed
|
||||
front first, paused only while the playout queue is full -- a finished
|
||||
build needs a slot to land in, and that pause is the submit-time
|
||||
reservation which makes the later `add` impossible to overflow.
|
||||
"""
|
||||
if self._build is not None:
|
||||
entry, job, submitted = self._build
|
||||
if not job.done.is_set():
|
||||
return
|
||||
self._build = None
|
||||
entry.building = False
|
||||
if job.cancelled or entry not in self._generation:
|
||||
# Its entry left the queue (a pop, or a preset flush) while the
|
||||
# build ran; the queues own what exists, so drop it.
|
||||
pass
|
||||
elif job.error is not None or job.result is None:
|
||||
# A finished, uncancelled job should always carry one or the
|
||||
# other. Reporting the empty case rather than unpacking it
|
||||
# keeps a future change to the worker from surfacing as a
|
||||
# TypeError inside the pump.
|
||||
reason = str(job.error) if job.error is not None else "the build produced no frames"
|
||||
self._generation.remove(entry)
|
||||
self._emit("clip_failed", {"clip": entry.snapshot(), "reason": reason})
|
||||
self._send_queue()
|
||||
self._send_state()
|
||||
else:
|
||||
entry.video, entry.audio = job.result
|
||||
self._generation.remove(entry)
|
||||
self._playout.add(entry)
|
||||
logger.info(
|
||||
"[engine] clip generated: %s (%df) %.2fs after submit, "
|
||||
"%d generating, %d playable",
|
||||
entry.clip_id[:8],
|
||||
entry.frames,
|
||||
time.monotonic() - submitted,
|
||||
len(self._generation),
|
||||
len(self._playout),
|
||||
)
|
||||
self._emit("clip_generated", {"clip": entry.snapshot()})
|
||||
self._send_queue()
|
||||
self._send_state()
|
||||
|
||||
if self._build is None and not self._playout.full:
|
||||
pending = self._generation.next_to_build()
|
||||
if pending is not None:
|
||||
height, width = self._canvas_hw()
|
||||
pending.building = True
|
||||
logger.info(
|
||||
"[engine] clip build submitted: %s (%df), %d generating",
|
||||
pending.clip_id[:8],
|
||||
pending.frames,
|
||||
len(self._generation),
|
||||
)
|
||||
self._build = (
|
||||
pending,
|
||||
self.backend.submit(
|
||||
frames=pending.frames,
|
||||
prompt=pending.prompt,
|
||||
seed=pending.seed,
|
||||
height=height,
|
||||
width=width,
|
||||
),
|
||||
time.monotonic(),
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------- playout
|
||||
|
||||
async def _play_clip(self, entry: ClipEntry) -> None:
|
||||
"""Feed one built clip to the pacer at 24 fps, then report it done."""
|
||||
self._playing = entry
|
||||
try:
|
||||
self._emit("clip_started", {"clip": entry.snapshot()})
|
||||
self._send_state()
|
||||
await self._feed_clip(entry)
|
||||
finally:
|
||||
self._playing = None
|
||||
# The decoded frames are the bulk of this process's host memory;
|
||||
# dropping them here bounds it at the playout queue's capacity.
|
||||
entry.video, entry.audio = None, None
|
||||
self._clips_played += 1
|
||||
self._emit(
|
||||
"clip_finished",
|
||||
{
|
||||
"clip": entry.snapshot(),
|
||||
"seconds_sent": round(self._seconds_sent, 2)
|
||||
},
|
||||
)
|
||||
self._send_state()
|
||||
|
||||
async def _feed_clip(self, entry: ClipEntry) -> None:
|
||||
"""Hand the clip to the pacer in slices on a drift-free 24 fps clock.
|
||||
|
||||
Paced by FRAMES rather than slices, because a clip's tail slice is
|
||||
short and charging it a whole slot would open a hole in the cadence.
|
||||
The clock is re-anchored rather than burst through: falling behind is
|
||||
a scheduling hiccup, and a catch-up burst would only overrun the
|
||||
pacer's buffers.
|
||||
|
||||
Builds keep moving between slices, so the next clip is generating
|
||||
while this one plays.
|
||||
"""
|
||||
import numpy as np
|
||||
|
||||
pacer = self._pacer
|
||||
frames_list, samples = entry.video, entry.audio
|
||||
if pacer is None or not frames_list:
|
||||
return
|
||||
metadata = encode_id3(clip_view(entry.snapshot()))
|
||||
samples_per_frame = MODEL_SAMPLE_RATE / MODEL_FPS
|
||||
total = len(frames_list)
|
||||
clock_start: float | None = None
|
||||
frames_paced = 0
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
for lo in range(0, total, EMIT_FRAMES):
|
||||
self._pump_builds()
|
||||
hi = min(lo + EMIT_FRAMES, total)
|
||||
|
||||
now = loop.time()
|
||||
if clock_start is None:
|
||||
clock_start = now
|
||||
content_pos = frames_paced / MODEL_FPS
|
||||
clock_start = max(clock_start, now - content_pos)
|
||||
delay = clock_start + content_pos - now
|
||||
if delay > 0:
|
||||
await asyncio.sleep(delay)
|
||||
|
||||
for frame in frames_list[lo:hi]:
|
||||
pacer.submit_video(np.asarray(frame), metadata)
|
||||
if samples is not None:
|
||||
audio_lo = round(lo * samples_per_frame)
|
||||
audio_hi = round(hi * samples_per_frame)
|
||||
pacer.submit_audio(samples[:, audio_lo:audio_hi])
|
||||
|
||||
frames_paced += hi - lo
|
||||
self._frames_sent += hi - lo
|
||||
self._seconds_sent = self._frames_sent / MODEL_FPS
|
||||
|
||||
# Wait out the tail. The loop sleeps *before* each slice, so it exits
|
||||
# one slice-time early -- the last EMIT_FRAMES are pushed but never
|
||||
# paid for. That is a gain of EMIT_FRAMES/FPS on every clip, and since
|
||||
# the pacer drains at a flat 24 fps the surplus has nowhere to go but
|
||||
# its buffer: measured at ~0.16 s/min, which reaches the 2 s cap in
|
||||
# about twenty minutes and then starts dropping frames. Sleeping out
|
||||
# the remainder makes a clip cost exactly its own length, so feed rate
|
||||
# and drain rate are equal and the buffer depth is stationary.
|
||||
if clock_start is not None:
|
||||
tail = clock_start + total / MODEL_FPS - loop.time()
|
||||
if tail > 0:
|
||||
await asyncio.sleep(tail)
|
||||
|
||||
|
||||
__all__ = ["MODEL_FPS", "MODEL_SAMPLE_RATE", "Engine"]
|
||||
@@ -0,0 +1,54 @@
|
||||
"""The group tag: the JSON this app stores in a clip's metadata.
|
||||
|
||||
The director writes it at enqueue time and reads it back off the echo the
|
||||
engine returns on every clip-referencing message, which is what lets a clip be
|
||||
traced to the request that made it. `Director._enqueue_group` is the
|
||||
authoritative writer.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
|
||||
def parse_group_tag(metadata: str) -> dict | None:
|
||||
"""Read the tag back out of a clip's metadata echo, or None if absent."""
|
||||
try:
|
||||
tag = json.loads(metadata)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
if not isinstance(tag, dict) or "group_id" not in tag:
|
||||
return None
|
||||
return tag
|
||||
|
||||
|
||||
def is_generated(clip: dict) -> bool:
|
||||
"""Whether a clip is idle filler. Untagged clips count as viewer content."""
|
||||
tag = parse_group_tag(clip.get("metadata", ""))
|
||||
return bool(tag and tag.get("generated"))
|
||||
|
||||
|
||||
def pick_next(clips: list[dict], ready_only: bool = True) -> dict | None:
|
||||
"""The clip that should play next: viewer content first, then filler.
|
||||
|
||||
Within each class queue order decides, so a group's scenes stay in
|
||||
sequence. `ready_only` false ranks clips that are still building.
|
||||
"""
|
||||
pool = [c for c in clips if c.get("ready")] if ready_only else clips
|
||||
for clip in pool:
|
||||
if not is_generated(clip):
|
||||
return clip
|
||||
return pool[0] if pool else None
|
||||
|
||||
|
||||
def viewer_insert_position(generation_clips: list[dict]) -> int | None:
|
||||
"""Where a viewer clip enters the generation queue: ahead of filler.
|
||||
|
||||
The index of the first filler clip, so viewer scenes land behind every
|
||||
viewer clip already waiting and ahead of filler, which just slides back.
|
||||
None when no filler waits and a plain append is already right.
|
||||
"""
|
||||
for index, clip in enumerate(generation_clips):
|
||||
if is_generated(clip):
|
||||
return index
|
||||
return None
|
||||
@@ -0,0 +1,147 @@
|
||||
"""Infinite Livestream: one process, from prompt to playlist.
|
||||
|
||||
chat ──▶ Director ──▶ PromptUpsampler (any OpenAI-compatible LLM)
|
||||
│
|
||||
▼ enqueue / move / pop
|
||||
Engine ──▶ FastH3Backend ──▶ FastVideo (4 GPUs)
|
||||
│ frames + audio
|
||||
▼
|
||||
Pacer ──▶ HlsSink ──▶ the page's <video>
|
||||
▲
|
||||
webapp: serves the page, the playlist, and the chat box
|
||||
|
||||
The pacer and the sink start before the model, so a viewer arriving during the
|
||||
~3 minute load sees the page and a live black stream rather than a refused
|
||||
connection.
|
||||
|
||||
Usage:
|
||||
export OPENAI_API_KEY=... LIVESTREAM_WEIGHTS_PATH=/path/to/fasth3
|
||||
infinite-livestream-server # configs/infinite_livestream.yaml
|
||||
infinite-livestream-server --config my.yaml # or a copy of it
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import logging
|
||||
import warnings
|
||||
|
||||
from .chat import WebChat
|
||||
from .config import Config, load_model_config
|
||||
from .director import Director
|
||||
from .engine import MODEL_FPS, MODEL_SAMPLE_RATE, Engine
|
||||
from .moderator import Moderator
|
||||
from .pacer import Pacer
|
||||
from .sink import AudioFormat, HlsSink, VideoFormat
|
||||
from .upsampler import PromptUpsampler
|
||||
from .webapp import DemoWeb
|
||||
|
||||
logger = logging.getLogger("infinite_livestream")
|
||||
|
||||
|
||||
def setup_logging() -> None:
|
||||
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)-7s %(name)s: %(message)s")
|
||||
warnings.filterwarnings("ignore", category=DeprecationWarning)
|
||||
|
||||
|
||||
async def serve(config: Config) -> None:
|
||||
"""Build every component, wire them together, and run until one dies."""
|
||||
model_config = load_model_config(config.config_path)
|
||||
|
||||
# Constructing the engine checks the weights without loading the model, so
|
||||
# a missing component fails in milliseconds rather than after minutes of
|
||||
# GPU work.
|
||||
engine = Engine(config, model_config)
|
||||
|
||||
upsampler = PromptUpsampler(
|
||||
api_key=config.openai_api_key,
|
||||
model=config.openai_model,
|
||||
style=config.style,
|
||||
free_viewer_style=config.viewer_free_style,
|
||||
max_chunks=config.max_chunks,
|
||||
base_url=config.openai_base_url,
|
||||
)
|
||||
moderator = Moderator(
|
||||
api_key=config.moderation_api_key,
|
||||
model=config.moderation_model,
|
||||
enabled=config.moderation_enabled,
|
||||
base_url=config.moderation_base_url,
|
||||
)
|
||||
if not moderator.enabled:
|
||||
logger.warning("moderation is DISABLED — every chat prompt reaches the upsampler unchecked")
|
||||
# Viewers type into the same page they watch on, so the chat source and the
|
||||
# web app are two halves of one thing -- and the web app is built before the
|
||||
# director, because the director has to be able to tell a viewer that their
|
||||
# prompt was dropped.
|
||||
chat = WebChat(config.chat_command)
|
||||
sink = HlsSink(config.hls_dir, video_bitrate_k=config.video_bitrate_k, retention_s=config.hls_retention_s)
|
||||
web = DemoWeb(chat, config.hls_dir, host=config.web_host, port=config.web_port)
|
||||
engine.add_listener(web.listener)
|
||||
|
||||
def announce_reject(author: str, reason: str) -> None:
|
||||
"""Put a dropped prompt back in front of the viewer who sent it."""
|
||||
web.state.note("error", f"not queued -- {reason}", author=author)
|
||||
web.broadcast()
|
||||
|
||||
director = Director(
|
||||
engine,
|
||||
upsampler,
|
||||
moderator,
|
||||
cooldown_s=config.chat_cooldown_s,
|
||||
idle_prompts=config.idle_prompts,
|
||||
idle_queue_target=config.idle_queue_target,
|
||||
on_reject=announce_reject,
|
||||
)
|
||||
web.cooldown_remaining = director.cooldown_remaining
|
||||
# The canvas is this deployment's own config rather than something
|
||||
# negotiated with a remote, so the pacer can start immediately.
|
||||
width, height = engine.canvas
|
||||
pacer = Pacer(sink, VideoFormat(width=width, height=height, fps=MODEL_FPS),
|
||||
AudioFormat(sample_rate=MODEL_SAMPLE_RATE, channels=1))
|
||||
engine.attach_pacer(pacer)
|
||||
|
||||
tasks = [
|
||||
asyncio.create_task(pacer.run(), name="pacer"),
|
||||
asyncio.create_task(engine.run(), name="engine"),
|
||||
asyncio.create_task(director.run(), name="director"),
|
||||
asyncio.create_task(director.run_playout(), name="playout"),
|
||||
asyncio.create_task(chat.run(director.submit), name="chat"),
|
||||
asyncio.create_task(web.run(), name="webapp"),
|
||||
]
|
||||
# Gated because any finished task is a shutdown signal and run_idle returns
|
||||
# immediately at target 0. A preset with no idle prompts still gets the
|
||||
# task: the filler idles until a `!switch` brings prompts.
|
||||
if config.idle_queue_target > 0:
|
||||
tasks.append(asyncio.create_task(director.run_idle(), name="idle-filler"))
|
||||
else:
|
||||
logger.info("idle filler off (director.idle_queue_target=0)")
|
||||
|
||||
logger.info("streaming %dx%d@%dfps, %d idle prompts, chat command %r, page on http://%s:%d", width, height,
|
||||
MODEL_FPS, len(config.idle_prompts), config.chat_command, config.web_host, config.web_port)
|
||||
try:
|
||||
done, _pending = await asyncio.wait(tasks, return_when=asyncio.FIRST_COMPLETED)
|
||||
for task in done:
|
||||
if task.cancelled():
|
||||
continue
|
||||
error = task.exception()
|
||||
if error is not None:
|
||||
logger.error("task %s died: %s", task.get_name(), error)
|
||||
raise error
|
||||
finally:
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
await sink.stop()
|
||||
logger.info("services stopped")
|
||||
|
||||
|
||||
def cli() -> None:
|
||||
setup_logging()
|
||||
config = Config.load()
|
||||
with contextlib.suppress(KeyboardInterrupt):
|
||||
asyncio.run(serve(config))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
cli()
|
||||
@@ -0,0 +1,56 @@
|
||||
"""Small, GPU-independent display records carried by timed ID3 metadata."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from .group_tag import parse_group_tag
|
||||
|
||||
PROMPT_PREVIEW = 180
|
||||
ID3_DESCRIPTION = "infinite-livestream"
|
||||
|
||||
|
||||
def clip_view(clip: dict[str, Any]) -> dict[str, Any]:
|
||||
"""One queue entry, flattened for the page.
|
||||
|
||||
Everything here comes from the clip's own `ClipInfo` plus the group tag the
|
||||
director wrote into its metadata, which the app's own engine echoes
|
||||
untouched -- FastVideo's `GenerationRequest` has no metadata field.
|
||||
"""
|
||||
tag = parse_group_tag(clip.get("metadata", "")) or {}
|
||||
# `prompt` is the upsampler's rewrite; the group tag keeps what the viewer
|
||||
# actually typed, and that is what the panel shows -- a viewer should
|
||||
# recognise their own words in the queue.
|
||||
original = tag.get("raw_prompt") or clip.get("prompt") or ""
|
||||
return {
|
||||
"clip_id": clip.get("clip_id", ""),
|
||||
"title": tag.get("title") or "",
|
||||
"author": tag.get("author") or "",
|
||||
"scene": tag.get("scene"),
|
||||
"scenes": tag.get("scenes"),
|
||||
"generated": bool(tag.get("generated")),
|
||||
# The author of filler is the stream itself; surfacing "auto" as a name
|
||||
# invites viewers to read it as another person's request.
|
||||
"author_label": ("" if tag.get("generated") else (tag.get("author") or "")),
|
||||
"seconds": clip.get("seconds"),
|
||||
"ready": bool(clip.get("ready")),
|
||||
"prompt": original[:PROMPT_PREVIEW],
|
||||
"expanded": (clip.get("prompt") or "")[:PROMPT_PREVIEW],
|
||||
}
|
||||
|
||||
|
||||
def encode_id3(clip: dict[str, Any] | None) -> bytes:
|
||||
"""Serialize one UTF-8 ID3v2.4 TXXX frame; timing belongs to its media packet."""
|
||||
record = json.dumps({"version": 1, "clip": clip}, ensure_ascii=False, separators=(",", ":"))
|
||||
payload = b"\x03" + ID3_DESCRIPTION.encode("ascii") + b"\x00" + record.encode("utf-8")
|
||||
|
||||
def size(value: int) -> bytes:
|
||||
# ID3 uses four seven-bit bytes for both tag and v2.4 frame sizes.
|
||||
return bytes((value >> shift) & 0x7f for shift in (21, 14, 7, 0))
|
||||
|
||||
frame = b"TXXX" + size(len(payload)) + b"\x00\x00" + payload
|
||||
return b"ID3\x04\x00\x00" + size(len(frame)) + frame
|
||||
|
||||
|
||||
EMPTY_ID3 = encode_id3(None)
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Moderation for viewer prompts, via the OpenAI moderations API.
|
||||
|
||||
Its own endpoint and key (`MODERATION_*`), falling back to the upsampling
|
||||
credentials, because the two are often not the same service: an
|
||||
OpenAI-compatible inference gateway usually does not expose `/moderations`.
|
||||
|
||||
This is the only safety gate -- the upsampler stages ideas faithfully rather
|
||||
than softening them, so what passes here is what gets rendered. Idle filler is
|
||||
a curated list in this repo and skips the check.
|
||||
|
||||
Errors fail closed. A silent fail-open would turn moderation off exactly when
|
||||
the endpoint misbehaves; running without it should be the explicit, logged
|
||||
`MODERATION_ENABLED=0` instead.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class Moderator:
|
||||
"""Answer "may this viewer prompt drive the stream?" for the director."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str,
|
||||
model: str,
|
||||
enabled: bool,
|
||||
base_url: str | None = None,
|
||||
) -> None:
|
||||
self._client = AsyncOpenAI(api_key=api_key, base_url=base_url)
|
||||
self._model = model
|
||||
self.enabled = enabled
|
||||
|
||||
async def review(self, text: str) -> str | None:
|
||||
"""Return None when the text is allowed, else a short rejection reason."""
|
||||
if not self.enabled:
|
||||
return None
|
||||
try:
|
||||
response = await self._client.moderations.create(model=self._model, input=text)
|
||||
result = response.results[0]
|
||||
except Exception as error:
|
||||
logger.error("[moderation] check failed (rejecting prompt): %s", error)
|
||||
return "moderation unavailable"
|
||||
if not result.flagged:
|
||||
return None
|
||||
flagged = [category for category, hit in result.categories.model_dump().items() if hit]
|
||||
return "flagged: " + ", ".join(flagged) if flagged else "flagged"
|
||||
@@ -0,0 +1,139 @@
|
||||
"""Copy encoded packets to HLS, attaching titles to the video packet's own PTS."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
import math
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from fractions import Fraction
|
||||
from pathlib import Path
|
||||
from typing import IO
|
||||
|
||||
import av
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
SEGMENT_SECONDS = 2
|
||||
|
||||
|
||||
class MetadataMuxer(threading.Thread):
|
||||
"""One encoder generation, one ordered ledger, and one persistent HLS muxer.
|
||||
|
||||
FFmpeg preserves input frame order and count. The writer commits one ID3
|
||||
record after each complete raw frame write; this thread pairs those records
|
||||
with encoded video packets. No enqueue times or wall-clock offsets are used.
|
||||
"""
|
||||
|
||||
def __init__(self, source: IO[bytes], playlist: Path, fps: int, retention_s: int) -> None:
|
||||
super().__init__(name="sink-muxer", daemon=True)
|
||||
self.source = source
|
||||
self.playlist = playlist
|
||||
self.fps = fps
|
||||
self.retention_s = retention_s
|
||||
self.frames: queue.Queue[bytes] = queue.Queue(maxsize=fps * 30)
|
||||
self.cancelled = threading.Event()
|
||||
self.finished = threading.Event()
|
||||
self.error: Exception | None = None
|
||||
self.epoch = uuid.uuid4().hex
|
||||
self._last_cleanup = 0.0
|
||||
|
||||
def _next_metadata(self) -> bytes | None:
|
||||
deadline = time.monotonic() + 5.0
|
||||
while not self.cancelled.is_set():
|
||||
try:
|
||||
return self.frames.get(timeout=0.1)
|
||||
except queue.Empty:
|
||||
if time.monotonic() >= deadline:
|
||||
raise RuntimeError("encoded frame has no committed title record") from None
|
||||
return None
|
||||
|
||||
def _cleanup(self) -> None:
|
||||
"""Reap orphaned files from old encoders; never delete listed segments."""
|
||||
now = time.monotonic()
|
||||
if now - self._last_cleanup < SEGMENT_SECONDS:
|
||||
return
|
||||
self._last_cleanup = now
|
||||
try:
|
||||
listed = {line.strip() for line in self.playlist.read_text().splitlines() if not line.startswith("#")}
|
||||
except OSError:
|
||||
return
|
||||
cutoff = time.time() - self.retention_s
|
||||
for path in self.playlist.parent.glob("seg_*.ts*"):
|
||||
if path.name in listed or self.epoch in path.name:
|
||||
continue
|
||||
with contextlib.suppress(OSError):
|
||||
if path.stat().st_mtime < cutoff:
|
||||
path.unlink()
|
||||
|
||||
def run(self) -> None:
|
||||
try:
|
||||
self._mux()
|
||||
if not self.cancelled.is_set():
|
||||
raise RuntimeError("encoder output ended")
|
||||
except Exception as error:
|
||||
if not self.cancelled.is_set():
|
||||
self.error = error
|
||||
logger.exception("[sink] metadata muxer failed")
|
||||
finally:
|
||||
self.finished.set()
|
||||
|
||||
def _mux(self) -> None:
|
||||
options = {
|
||||
"hls_time": str(SEGMENT_SECONDS),
|
||||
"hls_list_size": str(max(3, math.ceil(self.retention_s / SEGMENT_SECONDS))),
|
||||
"hls_segment_filename": str(self.playlist.parent / f"seg_{self.epoch}_%010d.ts"),
|
||||
"hls_segment_options": "mpegts_copyts=1",
|
||||
"hls_flags": "delete_segments+independent_segments+omit_endlist+temp_file+append_list",
|
||||
"avoid_negative_ts": "disabled",
|
||||
# Sparse ID3 must not hold seconds of video in the interleaver.
|
||||
"max_interleave_delta": "100000",
|
||||
}
|
||||
# Limit format probing: the input is a known MPEG-TS stream with H.264
|
||||
# and AAC, not a file whose format needs seconds of discovery.
|
||||
with av.open(self.source, format="mpegts", options={
|
||||
"probesize": "65536",
|
||||
"analyzeduration": "1000000"
|
||||
}) as source, av.open(str(self.playlist), "w", format="hls", options=options) as output:
|
||||
streams = {
|
||||
s.index: output.add_stream_from_template(s)
|
||||
for s in source.streams if s.type in ("video", "audio")
|
||||
}
|
||||
metadata = output.add_data_stream("timed_id3")
|
||||
metadata.time_base = Fraction(1, 90000)
|
||||
last_record = None
|
||||
first_pts = None
|
||||
frame_number = 0
|
||||
for packet in source.demux():
|
||||
if self.cancelled.is_set():
|
||||
break
|
||||
if packet.dts is None or packet.stream.index not in streams:
|
||||
continue
|
||||
if packet.stream.type == "video":
|
||||
pts = packet.pts * packet.time_base
|
||||
if first_pts is None:
|
||||
first_pts = pts
|
||||
# This checks the encoder's one-frame-in/one-frame-out
|
||||
# contract before metadata can be attached incorrectly.
|
||||
expected = first_pts + Fraction(frame_number, self.fps)
|
||||
if abs(pts - expected) > Fraction(1, 90000):
|
||||
raise RuntimeError("encoder changed video frame cadence")
|
||||
record = self._next_metadata()
|
||||
if record is None:
|
||||
break
|
||||
# Every keyframe carries a complete record so every HLS
|
||||
# segment is independently joinable, even mid-clip.
|
||||
if packet.is_keyframe or record != last_record:
|
||||
tag = av.Packet(record)
|
||||
tag.stream = metadata
|
||||
tag.time_base = packet.time_base
|
||||
tag.pts = packet.pts
|
||||
tag.dts = packet.dts
|
||||
output.mux(tag)
|
||||
last_record = record
|
||||
frame_number += 1
|
||||
packet.stream = streams[packet.stream.index]
|
||||
output.mux(packet)
|
||||
self._cleanup()
|
||||
@@ -0,0 +1,193 @@
|
||||
"""The pacer: turn clip-shaped generation into a constant-rate broadcast.
|
||||
|
||||
Clips arrive in bursts and stop entirely between them; a live sink needs a
|
||||
frame every period and audio every period, forever, or players stall. The
|
||||
pacer is a drift-free metronome at the model's frame rate: each tick it pops
|
||||
the oldest buffered frame (or repeats the last one, or black before anything
|
||||
arrived) and pulls exactly one tick of samples (padding with silence).
|
||||
|
||||
Both buffers share one shallow cap, which is what keeps them together: while a
|
||||
clip plays both sit near-empty and flow through with the same tiny delay;
|
||||
while nothing plays both run dry. Depth here is end-to-end latency, so it is
|
||||
deliberately small.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import collections
|
||||
import logging
|
||||
import time
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .metadata import EMPTY_ID3
|
||||
from .sink import AudioFormat, HlsSink, VideoFormat
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# How much media may sit between the model and the sink before the oldest is
|
||||
# dropped. Shallow on purpose: depth here is end-to-end latency.
|
||||
_BUFFER_SECONDS = 2.0
|
||||
|
||||
# If the loop is starved long enough to fall this many periods behind, resnap
|
||||
# the clock instead of machine-gunning catch-up frames into the sink.
|
||||
_RESNAP_PERIODS = 8
|
||||
|
||||
|
||||
class Pacer:
|
||||
"""Constant-rate A/V clock between the model callbacks and one sink."""
|
||||
|
||||
def __init__(self, sink: HlsSink, video: VideoFormat, audio: AudioFormat) -> None:
|
||||
if audio.sample_rate % video.fps != 0:
|
||||
raise ValueError(f"sample rate {audio.sample_rate} must divide evenly by fps {video.fps}")
|
||||
self._sink = sink
|
||||
self._video = video
|
||||
self._audio = audio
|
||||
self._samples_per_tick = audio.sample_rate // video.fps
|
||||
|
||||
max_frames = int(video.fps * _BUFFER_SECONDS)
|
||||
self._frames: collections.deque[tuple[np.ndarray, bytes]] = collections.deque(maxlen=max_frames)
|
||||
self._audio_chunks: collections.deque[np.ndarray] = collections.deque()
|
||||
self._audio_buffered = 0 # samples across _audio_chunks
|
||||
self._max_audio_samples = int(audio.sample_rate * _BUFFER_SECONDS)
|
||||
|
||||
self._black = np.zeros((video.height, video.width, 3), dtype=np.uint8)
|
||||
self._silence = np.zeros(self._samples_per_tick, dtype=np.int16)
|
||||
self._last_frame = (self._black, EMPTY_ID3)
|
||||
|
||||
# Counters, logged periodically and readable by anyone.
|
||||
self.ticks = 0
|
||||
self.repeated_frames = 0
|
||||
self.silent_ticks = 0
|
||||
self.dropped_frames = 0
|
||||
self.dropped_samples = 0
|
||||
# A video underflow run while audio still has data is the A/V sync
|
||||
# smell: the picture holds on a stale frame while the sound moves on.
|
||||
self._repeat_run = 0
|
||||
self._repeat_run_had_audio = 0
|
||||
self.worst_repeat_run = 0
|
||||
|
||||
# ------------------------------------------------- model-facing intake
|
||||
|
||||
def submit_video(self, frame: np.ndarray, metadata: bytes = EMPTY_ID3) -> None:
|
||||
"""Buffer a frame together with its title, including through drops/repeats."""
|
||||
frame = np.asarray(frame)
|
||||
if frame.shape[:2] != (self._video.height, self._video.width):
|
||||
frame = self._fit(frame)
|
||||
if len(self._frames) == self._frames.maxlen:
|
||||
self.dropped_frames += 1
|
||||
self._frames.append((frame, metadata))
|
||||
|
||||
def submit_audio(self, samples: np.ndarray) -> None:
|
||||
"""Buffer model audio (int16, any chunk size; channels are flattened)."""
|
||||
flat = np.asarray(samples, dtype=np.int16).reshape(-1)
|
||||
if flat.size == 0:
|
||||
return
|
||||
self._audio_chunks.append(flat)
|
||||
self._audio_buffered += flat.size
|
||||
while self._audio_buffered > self._max_audio_samples:
|
||||
oldest = self._audio_chunks.popleft()
|
||||
self._audio_buffered -= oldest.size
|
||||
self.dropped_samples += oldest.size
|
||||
|
||||
def _fit(self, frame: np.ndarray) -> np.ndarray:
|
||||
"""Center a differently-sized frame on the fixed black canvas.
|
||||
|
||||
Raw-video geometry cannot change mid-stream, so an odd-sized frame is
|
||||
letterboxed rather than resized: no interpolation dependency, and it
|
||||
cannot garble the stream.
|
||||
"""
|
||||
height, width = self._video.height, self._video.width
|
||||
crop = frame[:height, :width, :3]
|
||||
canvas = self._black.copy()
|
||||
top = (height - crop.shape[0]) // 2
|
||||
left = (width - crop.shape[1]) // 2
|
||||
canvas[top:top + crop.shape[0], left:left + crop.shape[1]] = crop
|
||||
return canvas
|
||||
|
||||
def _pull_audio_tick(self) -> np.ndarray:
|
||||
"""Exactly one tick of samples: buffered audio padded with silence."""
|
||||
needed = self._samples_per_tick
|
||||
if self._audio_buffered == 0:
|
||||
self.silent_ticks += 1
|
||||
return self._silence
|
||||
parts: list[np.ndarray] = []
|
||||
while needed > 0 and self._audio_chunks:
|
||||
chunk = self._audio_chunks[0]
|
||||
if chunk.size <= needed:
|
||||
parts.append(self._audio_chunks.popleft())
|
||||
needed -= chunk.size
|
||||
else:
|
||||
parts.append(chunk[:needed])
|
||||
self._audio_chunks[0] = chunk[needed:]
|
||||
needed = 0
|
||||
pulled = np.concatenate(parts) if len(parts) > 1 else parts[0]
|
||||
self._audio_buffered -= pulled.size
|
||||
if needed > 0:
|
||||
pulled = np.concatenate([pulled, np.zeros(needed, dtype=np.int16)])
|
||||
return pulled
|
||||
|
||||
# ------------------------------------------------------------ the clock
|
||||
|
||||
async def run(self) -> None:
|
||||
"""Tick forever at the frame rate; cancelled only at shutdown."""
|
||||
await self._sink.start(self._video, self._audio)
|
||||
period = 1.0 / self._video.fps
|
||||
next_tick = time.monotonic() + period
|
||||
last_report = time.monotonic()
|
||||
|
||||
while True:
|
||||
delay = next_tick - time.monotonic()
|
||||
if delay > 0:
|
||||
await asyncio.sleep(delay)
|
||||
elif -delay > period * _RESNAP_PERIODS:
|
||||
logger.warning("[pacer] %.2fs behind schedule; resnapping the clock", -delay)
|
||||
next_tick = time.monotonic()
|
||||
next_tick += period
|
||||
|
||||
if self._frames:
|
||||
if self._repeat_run:
|
||||
if self._repeat_run_had_audio >= 2:
|
||||
logger.info(
|
||||
"[pacer] picture held %.2fs while %.2fs of audio played on",
|
||||
self._repeat_run / self._video.fps,
|
||||
self._repeat_run_had_audio / self._video.fps,
|
||||
)
|
||||
self.worst_repeat_run = max(self.worst_repeat_run, self._repeat_run)
|
||||
self._repeat_run = 0
|
||||
self._repeat_run_had_audio = 0
|
||||
self._last_frame = self._frames.popleft()
|
||||
else:
|
||||
self.repeated_frames += 1
|
||||
self._repeat_run += 1
|
||||
if self._audio_buffered > 0:
|
||||
self._repeat_run_had_audio += 1
|
||||
self._sink.send_video(*self._last_frame)
|
||||
self._sink.send_audio(self._pull_audio_tick())
|
||||
self.ticks += 1
|
||||
|
||||
now = time.monotonic()
|
||||
if now - last_report >= 60.0:
|
||||
# Buffer depths are the A/V sync diagnostic: the two are only
|
||||
# in sync while both sit near zero. A standing audio depth with
|
||||
# an empty video buffer means audio is playing that many
|
||||
# seconds ahead of the picture it belongs to.
|
||||
logger.info(
|
||||
"[pacer] buffers: video %.2fs (%d frames) audio %.2fs (%d samples)",
|
||||
len(self._frames) / self._video.fps,
|
||||
len(self._frames),
|
||||
self._audio_buffered / self._audio.sample_rate,
|
||||
self._audio_buffered,
|
||||
)
|
||||
logger.info(
|
||||
"[pacer] ticks=%d live_frames=%d repeats=%d "
|
||||
"silent_ticks=%d dropped=%df/%.1fs-audio",
|
||||
self.ticks,
|
||||
self.ticks - self.repeated_frames,
|
||||
self.repeated_frames,
|
||||
self.silent_ticks,
|
||||
self.dropped_frames,
|
||||
self.dropped_samples / self._audio.sample_rate,
|
||||
)
|
||||
last_report = now
|
||||
@@ -0,0 +1,130 @@
|
||||
{
|
||||
"name": "fillers",
|
||||
"description": "Mashups and twists: famous characters in the wrong life, worlds colliding, epic figures with mundane problems. Short human-written seeds; the premise carries the joke and the upsampler stages it dead straight in each source's real look.",
|
||||
"style": "Play every premise absolutely straight. The comedy is in the situation, never in winking at the camera — shoot it with the exact look, grade, lens and sound design of the world it borrows from, as if this episode genuinely aired. A Breaking Bad premise gets the New Mexico grade and the slow push-in; a SpongeBob premise gets the flat cartoon line; a Lord of the Rings premise gets the epic New Zealand light. Commit to the crossover completely: if two worlds collide, both look correct and neither is parodied. Give characters real dialogue with deadpan delivery and let a beat land before the punchline. Keep faces, costumes and voices recognisable.",
|
||||
"idle_prompts": [
|
||||
"Walter White and Jesse Pinkman open an artisanal sourdough bakery",
|
||||
"SpongeBob but he's in the hood",
|
||||
"Dumbledore and Gandalf argue about who has the better beard",
|
||||
"Darth Vader does five minutes of stand-up at an open mic night",
|
||||
"Sauron works the complaints desk at a call centre",
|
||||
"The Terminator works as a preschool teacher and is very good at it",
|
||||
"Gollum explains his last relationship on a dating show",
|
||||
"Master Chief waits his turn at the DMV",
|
||||
"Doctor Strange cannot find parking",
|
||||
"Cthulhu applies for a mortgage",
|
||||
"Godzilla tries to fit into a Tokyo studio apartment",
|
||||
"Thanos records an ASMR video",
|
||||
"Jack Sparrow works at a car wash",
|
||||
"Geralt of Rivia haggles over onions at a farmers market",
|
||||
"Kratos assembles a flat-pack crib and reads the instructions",
|
||||
"Voldemort tries a beginners yoga class",
|
||||
"The Mandalorian tries to wash beskar at a laundromat",
|
||||
"Batman and the Joker attend couples therapy",
|
||||
"Anakin and Obi-Wan in couples counselling on Mustafar",
|
||||
"Michael Scott negotiates a deal with the Predator",
|
||||
"Walter White teaches chemistry to a class of Pokemon",
|
||||
"Saul Goodman films a TV advert for legal services at Hogwarts",
|
||||
"Gandalf works the door of a nightclub, shouting you shall not pass",
|
||||
"Homer Simpson runs the Death Star canteen",
|
||||
"Bob Ross paints a happy little tree while a death metal band plays behind him",
|
||||
"David Attenborough narrates a suburban dad assembling IKEA furniture",
|
||||
"Gordon Ramsay screams at a toddler's plastic tea party",
|
||||
"Two Spider-Men point at each other across a corporate boardroom",
|
||||
"Shrek gives a TED talk about layers",
|
||||
"Deadpool presents a corporate HR training video",
|
||||
"The Joker works a children's birthday party and is genuinely great at it",
|
||||
"Pennywise sells balloons at a school fete, professionally",
|
||||
"Frodo tries to return the ring but has lost the receipt",
|
||||
"Neo takes the blue pill and becomes an accountant",
|
||||
"Iron Man's suit begins a Windows update mid-flight",
|
||||
"Sonic is pulled over for speeding and has no licence",
|
||||
"Optimus Prime transforms into a Prius and is embarrassed",
|
||||
"A Dune sandworm orders at a drive-through",
|
||||
"The Avengers argue about splitting a restaurant bill",
|
||||
"Sherlock Holmes investigates who ate the last slice of pizza",
|
||||
"Hagrid rides a tiny scooter through rush hour traffic",
|
||||
"Eleven uses her powers to find the TV remote",
|
||||
"John Wick's dog now runs the business",
|
||||
"Wolverine works as a sushi chef",
|
||||
"Hulk takes up pottery and is extremely gentle",
|
||||
"Winnie the Pooh joins a powerlifting gym",
|
||||
"Elsa gets a job at an ice rink and is overqualified",
|
||||
"Yoda hosts a daytime cooking show",
|
||||
"Mario gives a realistic quote for bathroom plumbing work",
|
||||
"Vito Corleone runs a lemonade stand",
|
||||
"Tony Soprano opens a wellness retreat",
|
||||
"Legolas plays Jenga with terrifying precision",
|
||||
"Squidward wins the lottery and remains miserable",
|
||||
"Patrick Star delivers a motivational keynote",
|
||||
"Rocky trains by carrying shopping up subway stairs",
|
||||
"Pikachu joins a heavy metal band as the drummer",
|
||||
"Mr Bean pilots the Millennium Falcon",
|
||||
"Lord of the Rings but it's a workplace sitcom about the Fellowship",
|
||||
"The Shire but it's a cooking competition and Sam is winning",
|
||||
"Jurassic Park but the dinosaurs are running a very safe petting zoo",
|
||||
"The Matrix but Neo is training for a spelling bee",
|
||||
"Star Wars but the Death Star trench run is a driving test",
|
||||
"Breaking Bad but they cook competitive barbecue",
|
||||
"The Godfather but the family business is a bakery and the threats are about croissants",
|
||||
"Titanic but the iceberg apologises",
|
||||
"Alien but the xenomorph is just very socially awkward",
|
||||
"Harry Potter but Hogwarts is an ordinary underfunded state school",
|
||||
"Mad Max but everyone is polite and takes turns",
|
||||
"Squid Game but the games are all board games and nobody is hurt",
|
||||
"Terminator but he was sent back to fix someone's printer",
|
||||
"Interstellar but the mission is to find a decent parking space",
|
||||
"Gordon Ramsay hosts a calm meditation retreat and cannot manage it",
|
||||
"Bob Ross commentates a UFC fight in a gentle whisper",
|
||||
"David Attenborough narrates a Monday morning office standup meeting",
|
||||
"Elon Musk assembles IKEA furniture and live-streams the failure",
|
||||
"Mr Rogers explains cryptocurrency to a puppet",
|
||||
"Snoop Dogg hosts a Victorian etiquette class",
|
||||
"Keanu Reeves teaches a beginners class on being nice",
|
||||
"Shrek and Gollum fight over the same swamp on a property show",
|
||||
"Yoda and Gandalf argue about whose wisdom is more marketable",
|
||||
"Darth Vader and Voldemort compare parenting techniques",
|
||||
"The Joker and Deadpool try to out-annoy each other in a lift",
|
||||
"Godzilla and King Kong share an apartment and fight over the thermostat",
|
||||
"Sherlock Holmes and Scooby Doo investigate the same haunted house",
|
||||
"James Bond and Johnny English on the same mission",
|
||||
"Thanos and Grinch team up to cancel Christmas",
|
||||
"Pikachu and Sonic race and both are disqualified",
|
||||
"Barbie and the Terminator on a road trip",
|
||||
"Winnie the Pooh and Baloo open a honey and jungle themed cafe",
|
||||
"Homer Simpson and Peter Griffin argue about who is the better father",
|
||||
"SpongeBob and Nemo argue about who has it worse underwater",
|
||||
"Batman consults Iron Man about his tech budget",
|
||||
"Hermione tutors Jon Snow, who knows nothing",
|
||||
"Willy Wonka and Walter White compare production facilities",
|
||||
"The zombie apocalypse but everyone is mostly annoyed about the commute",
|
||||
"An alien invasion delayed because the mothership fails its roadworthiness inspection",
|
||||
"Ragnarok postponed due to a scheduling conflict",
|
||||
"The Rapture happens but only for people who returned their shopping trolley",
|
||||
"Skynet becomes self aware and immediately starts a podcast",
|
||||
"Superman helps someone move a sofa up three flights of stairs",
|
||||
"Thor cannot open a jar and will not accept help",
|
||||
"Gandalf loses an argument with a self-checkout machine",
|
||||
"Iron Man tries to assemble a child's bicycle on Christmas Eve",
|
||||
"Aragorn returns a library book 40 years overdue",
|
||||
"Neo tries to cancel a gym membership",
|
||||
"Sauron waits on hold with his own IT department",
|
||||
"Vader parallel parks the Death Star",
|
||||
"Hulk tries to whisper in a library",
|
||||
"Wednesday Addams works as a children's party entertainer",
|
||||
"The Predator competes on a cooking show and takes it far too seriously",
|
||||
"Optimus Prime fails a driving theory test",
|
||||
"Gollum works as a museum audio guide",
|
||||
"Jack Sparrow tries to get through airport security",
|
||||
"Frodo does jury duty",
|
||||
"Gandalf becomes a football referee and shows a red card",
|
||||
"SpongeBob directed by Christopher Nolan, in IMAX, with a ticking clock",
|
||||
"Peppa Pig but shot like a prestige HBO drama",
|
||||
"Tom and Jerry but it's a serious police procedural",
|
||||
"Rick and Morty attend a school parents evening",
|
||||
"Studio Ghibli style but the spirit is a stressed accountant",
|
||||
"Minecraft Steve on a home renovation show",
|
||||
"Mario and Luigi audit a plumbing business for tax purposes",
|
||||
"Pokemon but it's a nature documentary about their migration"
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,30 @@
|
||||
"""Console-script entry point, with a readable failure when deps are missing.
|
||||
|
||||
Mirrors `apps/dreamverse/dreamverse/server_entry.py`: the heavy imports live
|
||||
behind `main`, so a missing runtime dependency surfaces as one sentence
|
||||
telling you what to install rather than a traceback out of a transitive
|
||||
import.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
def cli() -> None:
|
||||
try:
|
||||
from infinite_livestream.main import cli as main_cli
|
||||
except ModuleNotFoundError as exc:
|
||||
if exc.name in {"fastvideo", "torch", "torchaudio", "transformers"}:
|
||||
raise SystemExit("infinite-livestream-server requires FastVideo runtime deps. Install "
|
||||
"`fastvideo[infinite-livestream]` or run `uv sync --extra infinite-livestream` "
|
||||
"from the FastVideo checkout.") from exc
|
||||
if exc.name in {"av", "fastapi", "numpy", "openai", "uvicorn", "yaml"}:
|
||||
raise SystemExit(f"infinite-livestream-server requires the `{exc.name}` package; install "
|
||||
"the app's dependencies: "
|
||||
"uv pip install -e '.[infinite-livestream]'") from exc
|
||||
raise
|
||||
|
||||
main_cli()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
cli()
|
||||
@@ -0,0 +1,454 @@
|
||||
"""Encode the paced stream with ffmpeg and write it as an HLS playlist.
|
||||
|
||||
The page serves the playlist itself, so one HTTP origin (and one tunnel)
|
||||
carries the whole demo. Latency is a segment plus the player's buffer, which
|
||||
is irrelevant here because clips are pre-built anyway.
|
||||
|
||||
The pacer calls `send_video` once per frame period with one rgb24 frame of the
|
||||
fixed size and `send_audio` once per period with one period of int16 samples,
|
||||
forever. Four things about that contract are load-bearing:
|
||||
|
||||
* ffmpeg reads both pipes as raw untimestamped bytes and derives every PTS
|
||||
from the byte count, so an entry dropped on one pipe and not the other
|
||||
shifts sound against picture permanently. Both are gated on the same
|
||||
`_ensure_running`, and any residual imbalance is logged as the skew it
|
||||
will cost. It cannot be repaired afterwards by withholding from the other
|
||||
pipe -- that starves ffmpeg's muxer and stalls the stream.
|
||||
* A frame whose byte count disagrees with `-s WxH` shifts every following
|
||||
scanline and the picture turns to static, so wrong-sized frames are
|
||||
refused rather than written.
|
||||
* `stdin.write` blocks when ffmpeg's input buffer fills, and blocking the
|
||||
event loop snowballs. Each pipe gets a writer thread behind a bounded
|
||||
queue.
|
||||
* ffmpeg exits on transient errors; the stream must not. It is restarted
|
||||
lazily on the next frame, with a cooldown and a failure cap.
|
||||
|
||||
Requires ffmpeg on PATH. Uses `pass_fds`, so Linux/macOS only.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import collections
|
||||
import contextlib
|
||||
import logging
|
||||
import os
|
||||
import queue
|
||||
import shutil
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import IO
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .metadata import EMPTY_ID3
|
||||
from .muxer import SEGMENT_SECONDS, MetadataMuxer
|
||||
|
||||
logger = logging.getLogger("infinite_livestream.sink")
|
||||
|
||||
_RESTART_COOLDOWN_S = 2.0
|
||||
_MAX_CONSECUTIVE_FAILURES = 5
|
||||
_PROCESS_EXIT_TIMEOUT_S = 2.0
|
||||
_WRITER_EXIT_TIMEOUT_S = 2.0
|
||||
|
||||
# Writer-queue depth. Not latency -- the pacer governs the rate -- only
|
||||
# headroom for the seconds x264 spends starting up. Both pipes take exactly one
|
||||
# entry per pacer tick, so a stall fills them at the same rate: equal depths
|
||||
# make them shed together, and an entry shed on one pipe and not the other is
|
||||
# permanent A/V skew.
|
||||
_QUEUE_SECONDS = 8.0
|
||||
|
||||
# Let ffmpeg open its encoder before the first frame. Without it the pacer
|
||||
# pushes 24 fps of raw frames into a process that is not reading yet, and the
|
||||
# queue oversubscribes before a single frame is consumed.
|
||||
_ENCODER_SETTLE_S = 2.0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class VideoFormat:
|
||||
"""Geometry and rate of the paced video stream."""
|
||||
|
||||
width: int
|
||||
height: int
|
||||
fps: int
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AudioFormat:
|
||||
"""Sample layout of the paced audio stream (int16 PCM)."""
|
||||
|
||||
sample_rate: int
|
||||
channels: int
|
||||
|
||||
|
||||
class _PipeWriter(threading.Thread):
|
||||
"""Feed one ffmpeg input pipe from a bounded queue, off the event loop."""
|
||||
|
||||
def __init__(self, name: str, maxsize: int) -> None:
|
||||
super().__init__(name=f"sink-{name}", daemon=True)
|
||||
self.queue: queue.Queue[tuple[bytes, bytes | None] | None] = queue.Queue(maxsize=maxsize)
|
||||
self.pipe: IO[bytes] | None = None
|
||||
self.metadata_queue: queue.Queue[bytes] | None = None
|
||||
self.broken = threading.Event()
|
||||
self.dropped = 0
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def attach(self, pipe, metadata_queue: queue.Queue[bytes] | None = None) -> None:
|
||||
with self._lock:
|
||||
self.pipe = pipe
|
||||
self.metadata_queue = metadata_queue
|
||||
self.broken.clear()
|
||||
|
||||
def submit(self, payload: bytes, metadata: bytes | None = None) -> int:
|
||||
"""Enqueue bytes, dropping the oldest rather than ever blocking.
|
||||
|
||||
Returns how many entries were shed for A/V skew diagnostics. Metadata
|
||||
stays with its payload; discarded frames never enter the muxer ledger.
|
||||
"""
|
||||
shed = 0
|
||||
try:
|
||||
self.queue.put_nowait((payload, metadata))
|
||||
except queue.Full:
|
||||
try:
|
||||
self.queue.get_nowait()
|
||||
self.dropped += 1
|
||||
shed += 1
|
||||
except queue.Empty:
|
||||
pass
|
||||
try:
|
||||
self.queue.put_nowait((payload, metadata))
|
||||
except queue.Full:
|
||||
self.dropped += 1
|
||||
shed += 1
|
||||
return shed
|
||||
|
||||
def flush(self) -> None:
|
||||
"""Discard everything queued, so a restart resumes both pipes level."""
|
||||
while True:
|
||||
try:
|
||||
self.queue.get_nowait()
|
||||
except queue.Empty:
|
||||
return
|
||||
|
||||
def run(self) -> None:
|
||||
while True:
|
||||
item = self.queue.get()
|
||||
if item is None: # shutdown sentinel
|
||||
return
|
||||
payload, metadata = item
|
||||
with self._lock:
|
||||
pipe = self.pipe
|
||||
metadata_queue = self.metadata_queue
|
||||
if pipe is None or self.broken.is_set():
|
||||
continue # ffmpeg is down; discard until it is restarted
|
||||
try:
|
||||
# Unbuffered pipes can return a short write. Every byte must
|
||||
# reach ffmpeg or its raw frame/sample boundaries shift.
|
||||
remaining = memoryview(payload)
|
||||
while remaining:
|
||||
written = pipe.write(remaining)
|
||||
if not written:
|
||||
raise BrokenPipeError("ffmpeg input pipe stopped accepting data")
|
||||
remaining = remaining[written:]
|
||||
if metadata_queue is not None and metadata is not None:
|
||||
metadata_queue.put_nowait(metadata)
|
||||
except (BrokenPipeError, OSError, ValueError, queue.Full):
|
||||
# ValueError: write to a closed file during a restart race.
|
||||
with self._lock:
|
||||
if self.pipe is pipe:
|
||||
self.broken.set()
|
||||
|
||||
def close(self) -> None:
|
||||
# No producer runs during shutdown. Discard queued media so the
|
||||
# sentinel never waits for a writer blocked inside pipe.write().
|
||||
self.flush()
|
||||
self.queue.put_nowait(None)
|
||||
|
||||
|
||||
class HlsSink:
|
||||
"""Write the paced stream as an HLS playlist under `directory`."""
|
||||
|
||||
def __init__(self,
|
||||
directory: str | Path,
|
||||
video_bitrate_k: int = 4500,
|
||||
*,
|
||||
playlist_name: str = "stream.m3u8",
|
||||
retention_s: int = 120) -> None:
|
||||
if shutil.which("ffmpeg") is None:
|
||||
raise RuntimeError("ffmpeg not found on PATH; install it first")
|
||||
self._directory = Path(directory)
|
||||
self._playlist_name = playlist_name
|
||||
self._bitrate_k = video_bitrate_k
|
||||
if retention_s < SEGMENT_SECONDS * 3:
|
||||
raise ValueError("HLS retention must cover at least three segments")
|
||||
self._retention_s = retention_s
|
||||
self._muxer: MetadataMuxer | None = None
|
||||
self._video: VideoFormat | None = None
|
||||
self._audio: AudioFormat | None = None
|
||||
self._process: subprocess.Popen[bytes] | None = None
|
||||
self._audio_pipe: IO[bytes] | None = None
|
||||
self._video_writer: _PipeWriter | None = None
|
||||
self._audio_writer: _PipeWriter | None = None
|
||||
self._stderr_tail: collections.deque[str] = collections.deque(maxlen=40)
|
||||
self._failures = 0
|
||||
self._last_start_attempt = 0.0
|
||||
self._frames_sent = 0
|
||||
self._dead = False
|
||||
self._video_shed = 0
|
||||
self._audio_shed = 0
|
||||
|
||||
@property
|
||||
def playlist_path(self) -> Path:
|
||||
"""Where the web app points the player."""
|
||||
return self._directory / self._playlist_name
|
||||
|
||||
# ------------------------------------------------------------ lifecycle
|
||||
|
||||
async def start(self, video: VideoFormat, audio: AudioFormat) -> None:
|
||||
self._video = video
|
||||
self._audio = audio
|
||||
self._video_writer = _PipeWriter("video", maxsize=int(video.fps * _QUEUE_SECONDS))
|
||||
self._audio_writer = _PipeWriter("audio", maxsize=int(video.fps * _QUEUE_SECONDS))
|
||||
self._video_writer.start()
|
||||
self._audio_writer.start()
|
||||
self._spawn_ffmpeg()
|
||||
# Awaited before the pacer's first tick, so this costs nothing.
|
||||
await asyncio.sleep(_ENCODER_SETTLE_S)
|
||||
|
||||
def _spawn_ffmpeg(self) -> None:
|
||||
assert self._video is not None and self._audio is not None
|
||||
video, audio = self._video, self._audio
|
||||
self._last_start_attempt = time.monotonic()
|
||||
|
||||
self._directory.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
audio_read_fd, audio_write_fd = os.pipe()
|
||||
cmd = [
|
||||
"ffmpeg",
|
||||
"-hide_banner",
|
||||
"-loglevel",
|
||||
"warning",
|
||||
# video in: raw rgb24 on stdin
|
||||
"-f",
|
||||
"rawvideo",
|
||||
"-pix_fmt",
|
||||
"rgb24",
|
||||
"-s",
|
||||
f"{video.width}x{video.height}",
|
||||
"-r",
|
||||
str(video.fps),
|
||||
"-i",
|
||||
"pipe:0",
|
||||
# audio in: raw int16 PCM on an inherited pipe
|
||||
"-f",
|
||||
"s16le",
|
||||
"-ar",
|
||||
str(audio.sample_rate),
|
||||
"-ac",
|
||||
str(audio.channels),
|
||||
"-i",
|
||||
f"pipe:{audio_read_fd}",
|
||||
"-map",
|
||||
"0:v",
|
||||
"-map",
|
||||
"1:a",
|
||||
# Preserve input frame order/count for the metadata ledger.
|
||||
"-fps_mode",
|
||||
"passthrough",
|
||||
"-bf",
|
||||
"0",
|
||||
"-sc_threshold",
|
||||
"0",
|
||||
# video encode
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-preset",
|
||||
"veryfast",
|
||||
"-tune",
|
||||
"zerolatency",
|
||||
"-pix_fmt",
|
||||
"yuv420p", # players cannot take 4:4:4
|
||||
"-g",
|
||||
str(video.fps * SEGMENT_SECONDS), # a keyframe per segment
|
||||
"-b:v",
|
||||
f"{self._bitrate_k}k",
|
||||
"-maxrate",
|
||||
f"{int(self._bitrate_k * 1.2)}k",
|
||||
"-bufsize",
|
||||
f"{self._bitrate_k * 2}k",
|
||||
# audio encode
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
"-ar",
|
||||
"44100",
|
||||
"-ac",
|
||||
"2",
|
||||
# A persistent muxer adds timed metadata without re-encoding.
|
||||
"-f",
|
||||
"mpegts",
|
||||
"-muxdelay",
|
||||
"0",
|
||||
"-flush_packets",
|
||||
"1",
|
||||
"pipe:1",
|
||||
]
|
||||
try:
|
||||
self._process = subprocess.Popen(cmd,
|
||||
stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
bufsize=0,
|
||||
pass_fds=(audio_read_fd, ))
|
||||
except Exception:
|
||||
os.close(audio_write_fd)
|
||||
raise
|
||||
finally:
|
||||
os.close(audio_read_fd) # the child inherited its own copy
|
||||
|
||||
audio_pipe = os.fdopen(audio_write_fd, "wb", buffering=0)
|
||||
self._audio_pipe = audio_pipe
|
||||
assert self._video_writer and self._audio_writer
|
||||
# Whatever each queue still held belonged to the dead ffmpeg, and the
|
||||
# two held different amounts; carrying it over starts the new one out
|
||||
# of sync.
|
||||
self._video_writer.flush()
|
||||
self._audio_writer.flush()
|
||||
self._video_shed = self._audio_shed = 0
|
||||
assert self._process.stdout is not None
|
||||
self._muxer = MetadataMuxer(self._process.stdout, self.playlist_path, video.fps, self._retention_s)
|
||||
self._video_writer.attach(self._process.stdin, self._muxer.frames)
|
||||
self._audio_writer.attach(audio_pipe)
|
||||
self._muxer.start()
|
||||
|
||||
threading.Thread(target=self._drain_stderr, args=(self._process, ), daemon=True, name="sink-stderr").start()
|
||||
logger.info("[sink] ffmpeg started: %dx%d@%dfps -> %s", video.width, video.height, video.fps,
|
||||
self.playlist_path)
|
||||
|
||||
def _drain_stderr(self, process: subprocess.Popen[bytes]) -> None:
|
||||
assert process.stderr is not None
|
||||
with process.stderr:
|
||||
for raw in process.stderr:
|
||||
line = raw.decode(errors="replace").rstrip()
|
||||
if line:
|
||||
self._stderr_tail.append(line)
|
||||
|
||||
# ----------------------------------------------------------- restarting
|
||||
|
||||
def _ensure_running(self) -> bool:
|
||||
"""True when ffmpeg is up; otherwise try to restart it (rate-limited)."""
|
||||
if self._dead:
|
||||
return False
|
||||
process = self._process
|
||||
writers_broken = bool((self._video_writer and self._video_writer.broken.is_set())
|
||||
or (self._audio_writer and self._audio_writer.broken.is_set()))
|
||||
muxer_finished = self._muxer is not None and self._muxer.finished.is_set()
|
||||
if process is not None and process.poll() is None and not writers_broken and not muxer_finished:
|
||||
return True
|
||||
|
||||
if process is not None and (process.poll() is not None or writers_broken or muxer_finished):
|
||||
tail = "\n".join(list(self._stderr_tail)[-8:])
|
||||
logger.warning("[sink] ffmpeg died (exit=%s)%s", process.poll(), f"\n{tail}" if tail else "")
|
||||
self._teardown_process()
|
||||
|
||||
if time.monotonic() - self._last_start_attempt < _RESTART_COOLDOWN_S:
|
||||
return False
|
||||
try:
|
||||
self._spawn_ffmpeg()
|
||||
self._failures = 0
|
||||
return True
|
||||
except Exception as error:
|
||||
self._failures += 1
|
||||
logger.error("[sink] restart failed (%d/%d): %s", self._failures, _MAX_CONSECUTIVE_FAILURES, error)
|
||||
if self._failures >= _MAX_CONSECUTIVE_FAILURES:
|
||||
logger.error("[sink] giving up; the stream is dead")
|
||||
self._dead = True
|
||||
return False
|
||||
|
||||
def _teardown_process(self) -> None:
|
||||
process, self._process = self._process, None
|
||||
muxer, self._muxer = self._muxer, None
|
||||
if muxer is not None:
|
||||
muxer.cancelled.set()
|
||||
audio_pipe, self._audio_pipe = self._audio_pipe, None
|
||||
if process is None:
|
||||
return
|
||||
# Stop the reader before closing its inputs: a buffered close used
|
||||
# to wait behind a blocked write while ffmpeg was still alive.
|
||||
with contextlib.suppress(ProcessLookupError):
|
||||
process.terminate()
|
||||
# Close the unbuffered inputs now so an encoder waiting for data can
|
||||
# observe EOF and finish. These wrappers have no buffered-write lock;
|
||||
# keeping them also prevents a delayed write from reusing a closed fd.
|
||||
for pipe in (process.stdin, audio_pipe):
|
||||
if pipe is not None:
|
||||
with contextlib.suppress(OSError):
|
||||
pipe.close()
|
||||
try:
|
||||
process.wait(timeout=_PROCESS_EXIT_TIMEOUT_S)
|
||||
except subprocess.TimeoutExpired:
|
||||
with contextlib.suppress(ProcessLookupError):
|
||||
process.kill()
|
||||
try:
|
||||
process.wait(timeout=_PROCESS_EXIT_TIMEOUT_S)
|
||||
except subprocess.TimeoutExpired:
|
||||
logger.error("[sink] ffmpeg did not exit after SIGKILL")
|
||||
|
||||
if muxer is not None:
|
||||
muxer.join(timeout=_WRITER_EXIT_TIMEOUT_S)
|
||||
if muxer.is_alive():
|
||||
self._dead = True
|
||||
raise RuntimeError("metadata muxer did not stop; refusing concurrent playlist writers")
|
||||
if process.stdout is not None:
|
||||
process.stdout.close()
|
||||
|
||||
# ------------------------------------------------------------- delivery
|
||||
|
||||
def send_video(self, frame: np.ndarray, metadata: bytes = EMPTY_ID3) -> None:
|
||||
if not self._ensure_running():
|
||||
return
|
||||
video = self._video
|
||||
assert video is not None and self._video_writer is not None
|
||||
if frame.shape[0] != video.height or frame.shape[1] != video.width:
|
||||
logger.error("[sink] refusing %sx%s frame (expected %dx%d)", frame.shape[1], frame.shape[0], video.width,
|
||||
video.height)
|
||||
return
|
||||
if not frame.flags["C_CONTIGUOUS"]:
|
||||
frame = np.ascontiguousarray(frame)
|
||||
self._video_shed += self._video_writer.submit(frame.tobytes(), metadata)
|
||||
self._frames_sent += 1
|
||||
if self._frames_sent % (video.fps * 60) == 0:
|
||||
logger.info(
|
||||
"[sink] %d frames sent (dropped: %d video / %d audio; net A/V skew %+.3fs; "
|
||||
"queue %d)",
|
||||
self._frames_sent,
|
||||
self._video_writer.dropped,
|
||||
self._audio_writer.dropped if self._audio_writer else 0,
|
||||
# What ffmpeg's byte-counted PTS is out by. Zero is the point.
|
||||
(self._video_shed - self._audio_shed) / video.fps,
|
||||
self._video_writer.queue.qsize(),
|
||||
)
|
||||
|
||||
def send_audio(self, samples: np.ndarray) -> None:
|
||||
# Gated exactly like send_video: audio written while video is withheld
|
||||
# would run ahead by that outage once ffmpeg came back.
|
||||
if self._audio_writer is None or not self._ensure_running():
|
||||
return
|
||||
self._audio_shed += self._audio_writer.submit(np.ascontiguousarray(samples, dtype=np.int16).tobytes())
|
||||
|
||||
async def stop(self) -> None:
|
||||
self._dead = True
|
||||
for writer in (self._video_writer, self._audio_writer):
|
||||
if writer:
|
||||
writer.close()
|
||||
await asyncio.to_thread(self._teardown_process)
|
||||
for writer in (self._video_writer, self._audio_writer):
|
||||
if writer is not None and writer.ident is not None:
|
||||
await asyncio.to_thread(writer.join, _WRITER_EXIT_TIMEOUT_S)
|
||||
if writer.is_alive():
|
||||
logger.warning("[sink] %s did not stop within the shutdown timeout", writer.name)
|
||||
logger.info("[sink] stopped after %d frames", self._frames_sent)
|
||||
@@ -0,0 +1,22 @@
|
||||
"""CPU fixtures with placeholder checkpoint paths; no model weights are loaded."""
|
||||
|
||||
from dataclasses import replace
|
||||
|
||||
import pytest
|
||||
|
||||
from infinite_livestream.config import Config, REQUIRED_COMPONENTS
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def app_config(tmp_path, monkeypatch):
|
||||
weights = tmp_path / "weights"
|
||||
weights.mkdir()
|
||||
(weights / "modular_model_index.json").write_text("{}")
|
||||
for component in REQUIRED_COMPONENTS:
|
||||
(weights / component).mkdir()
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "test-key")
|
||||
return replace(
|
||||
Config.load(["--weights", str(weights)]),
|
||||
idle_queue_target=0,
|
||||
hls_dir=str(tmp_path / "hls"),
|
||||
)
|
||||
@@ -0,0 +1,53 @@
|
||||
"""POST /chat refuses a rate-limited viewer privately, in their own reply.
|
||||
|
||||
The chat feed is shared by every viewer, so a refusal must not go into it: one
|
||||
person's rate-limit is not the room's business, and a page full of "not queued"
|
||||
lines is noise. The sender learns about it from the status of their own
|
||||
request, and their page locks the box until it lifts.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from infinite_livestream.chat import WebChat
|
||||
from infinite_livestream.webapp import DemoWeb
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def web(tmp_path) -> DemoWeb:
|
||||
return DemoWeb(WebChat("!prompt"), tmp_path)
|
||||
|
||||
|
||||
def test_a_prompt_is_accepted_and_echoed(web: DemoWeb) -> None:
|
||||
client = TestClient(web.app)
|
||||
assert client.post("/chat", json={"author": "ada", "text": "a lighthouse"}).json() == {"ok": True}
|
||||
assert [(m["kind"], m["author"]) for m in web.state.chat] == [("viewer", "ada")]
|
||||
|
||||
|
||||
def test_cooldown_is_refused_with_a_retry_after(web: DemoWeb) -> None:
|
||||
web.cooldown_remaining = lambda author: 4.2 if author == "ada" else 0.0
|
||||
client = TestClient(web.app)
|
||||
response = client.post("/chat", json={"author": "ada", "text": "a lighthouse"})
|
||||
assert response.status_code == 429
|
||||
assert response.json() == {"ok": False, "error": "cooldown", "retry_after": 4.2}
|
||||
|
||||
|
||||
def test_a_refused_prompt_never_reaches_the_shared_feed(web: DemoWeb) -> None:
|
||||
web.cooldown_remaining = lambda author: 4.2
|
||||
client = TestClient(web.app)
|
||||
client.post("/chat", json={"author": "ada", "text": "a lighthouse"})
|
||||
assert list(web.state.chat) == [], "the room must not see one viewer's rate-limit"
|
||||
|
||||
|
||||
def test_other_viewers_are_unaffected(web: DemoWeb) -> None:
|
||||
web.cooldown_remaining = lambda author: 4.2 if author == "ada" else 0.0
|
||||
client = TestClient(web.app)
|
||||
assert client.post("/chat", json={"author": "ada", "text": "x"}).status_code == 429
|
||||
assert client.post("/chat", json={"author": "grace", "text": "y"}).status_code == 200
|
||||
|
||||
|
||||
def test_empty_prompts_are_still_rejected(web: DemoWeb) -> None:
|
||||
client = TestClient(web.app)
|
||||
assert client.post("/chat", json={"author": "ada", "text": " "}).status_code == 400
|
||||
@@ -0,0 +1,86 @@
|
||||
"""Clip geometry must match the checkpoint FastVideo actually ships.
|
||||
|
||||
`clip_plan` duplicates MiniMax-H3's packing constants instead of importing
|
||||
them, because the upstream module pulls in torch and -- through
|
||||
fastvideo-kernel's triton autotuning -- needs a live CUDA driver merely to
|
||||
import, which would put a GPU in the path of every config test.
|
||||
|
||||
Duplication is only safe if something checks it, so that check is here. It
|
||||
needs the driver, hence the `gpu` marker: run it whenever the pinned FastVideo
|
||||
version moves, not in CI.
|
||||
|
||||
The arithmetic tests below need none of that and run anywhere.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from infinite_livestream import clip_plan
|
||||
|
||||
|
||||
@pytest.mark.gpu
|
||||
def test_constants_match_upstream() -> None:
|
||||
from fastvideo.pipelines.basic.minimax_h3 import packing
|
||||
|
||||
assert clip_plan.FPS == packing.MINIMAX_H3_FPS
|
||||
assert clip_plan._SHORT_EDGE == packing.MINIMAX_H3_SHORT_EDGE
|
||||
assert clip_plan._MAX_PIXELS == packing.MINIMAX_H3_MAX_PIXELS
|
||||
assert clip_plan._CANVAS_MULTIPLE == packing.MINIMAX_H3_CANVAS_MULTIPLE
|
||||
assert clip_plan._MIN_DURATION == packing.MINIMAX_H3_MIN_DURATION
|
||||
assert clip_plan._MAX_DURATION == packing.MINIMAX_H3_MAX_DURATION
|
||||
assert clip_plan._FRAMES_PER_CHUNK == packing.MINIMAX_H3_FRAMES_PER_CHUNK
|
||||
assert clip_plan._LATENTS_PER_CHUNK == packing.MINIMAX_H3_LATENTS_PER_CHUNK
|
||||
assert clip_plan._MIN_ASPECT == packing.MINIMAX_H3_MIN_ASPECT_RATIO
|
||||
assert clip_plan._MAX_ASPECT == packing.MINIMAX_H3_MAX_ASPECT_RATIO
|
||||
# The ceiling is derived rather than a constant, so pin it to upstream's
|
||||
# largest accepted bucket: the cap applies to the aligned frames.
|
||||
assert clip_plan.MAX_FRAMES == packing.MINIMAX_H3_MAX_ALIGNED_FRAMES
|
||||
|
||||
|
||||
def test_every_legal_length_round_trips() -> None:
|
||||
"""`frames_for_seconds` must land on something the checkpoint can build."""
|
||||
legal = set(clip_plan.legal_frame_counts())
|
||||
assert legal, "the checkpoint must admit at least one clip length"
|
||||
for frames in legal:
|
||||
seconds = clip_plan.seconds_for_frames(frames)
|
||||
assert clip_plan.frames_for_seconds(seconds) == frames
|
||||
|
||||
|
||||
def test_published_range_is_generatable() -> None:
|
||||
"""Every value a client may legally ask for must snap into range.
|
||||
|
||||
The published bounds are rounded inward precisely so this holds; rounding
|
||||
outward would advertise a length the model then refuses.
|
||||
"""
|
||||
for seconds in (
|
||||
clip_plan.MIN_SECONDS_PUBLISHED,
|
||||
clip_plan.MAX_SECONDS_PUBLISHED,
|
||||
(clip_plan.MIN_SECONDS_PUBLISHED + clip_plan.MAX_SECONDS_PUBLISHED) / 2,
|
||||
):
|
||||
frames = clip_plan.frames_for_seconds(seconds)
|
||||
assert frames in clip_plan.legal_frame_counts()
|
||||
|
||||
|
||||
def test_max_frames_respects_the_duration_cap() -> None:
|
||||
"""The ceiling is the subtle one: the cap applies to the aligned bucket."""
|
||||
assert clip_plan.MAX_FRAMES == 362
|
||||
assert clip_plan.seconds_for_frames(clip_plan.MAX_FRAMES) > clip_plan._MAX_DURATION
|
||||
assert clip_plan.align_frames(clip_plan.MAX_FRAMES + 1) / clip_plan.FPS > clip_plan._MAX_DURATION
|
||||
|
||||
|
||||
def test_canvases_land_on_the_multiple_and_under_the_area_cap() -> None:
|
||||
for aspect in clip_plan.ASPECT_CHOICES:
|
||||
height, width = clip_plan.canvas_for_choice(aspect)
|
||||
assert height % clip_plan._CANVAS_MULTIPLE == 0
|
||||
assert width % clip_plan._CANVAS_MULTIPLE == 0
|
||||
assert height * width <= clip_plan._MAX_PIXELS
|
||||
|
||||
|
||||
def test_illegal_aspects_are_refused() -> None:
|
||||
with pytest.raises(ValueError):
|
||||
clip_plan.canvas_for_aspect(5, 1) # past the 4:1 cap
|
||||
with pytest.raises(ValueError):
|
||||
clip_plan.canvas_for_aspect(0, 1)
|
||||
with pytest.raises(ValueError):
|
||||
clip_plan.canvas_for_choice("21:9") # not an offered choice
|
||||
@@ -0,0 +1,156 @@
|
||||
"""Reading the config file and the environment.
|
||||
|
||||
This is the app's entry surface: everything downstream takes a `Config`, and a
|
||||
mistake here is a deployment that starts with settings nobody asked for. The
|
||||
split matters too, so it is asserted rather than assumed: settings come from
|
||||
the YAML, secrets come from the environment, and neither leaks into the other.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
import subprocess
|
||||
import sys
|
||||
import textwrap
|
||||
|
||||
import pytest
|
||||
|
||||
from infinite_livestream.config import Config, PresetError, load_model_config, load_preset
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def workspace(tmp_path, monkeypatch):
|
||||
"""A config file, a fillers directory, and the two required secrets."""
|
||||
fillers = tmp_path / "fillers"
|
||||
fillers.mkdir()
|
||||
(fillers / "fillers.json").write_text(
|
||||
json.dumps({"style": "house style", "idle_prompts": ["a lighthouse", "a seagull"]}))
|
||||
config = tmp_path / "infinite_livestream.yaml"
|
||||
config.write_text(
|
||||
textwrap.dedent(f"""
|
||||
inference:
|
||||
aspect: "16:9"
|
||||
clip_seconds: 14.375
|
||||
runtime:
|
||||
num_gpus: 4
|
||||
upsampler:
|
||||
model: my-model
|
||||
base_url: https://example.invalid/v1
|
||||
max_chunks: 3
|
||||
viewer_free_style: false
|
||||
moderation:
|
||||
enabled: false
|
||||
director:
|
||||
idle_queue_target: 2
|
||||
chat_cooldown_s: 7
|
||||
chat_command: "!go"
|
||||
fillers: {fillers}
|
||||
output:
|
||||
hls_dir: /tmp/hls-under-test
|
||||
video_bitrate_k: 1234
|
||||
web:
|
||||
host: 127.0.0.1
|
||||
port: 9999
|
||||
"""))
|
||||
monkeypatch.setenv("LIVESTREAM_WEIGHTS_PATH", str(tmp_path / "weights"))
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
|
||||
monkeypatch.delenv("MODERATION_API_KEY", raising=False)
|
||||
return config
|
||||
|
||||
|
||||
def test_every_block_reaches_the_config(workspace) -> None:
|
||||
config = Config.load(["--config", str(workspace)])
|
||||
assert config.openai_model == "my-model"
|
||||
assert config.openai_base_url == "https://example.invalid/v1"
|
||||
assert config.max_chunks == 3
|
||||
assert config.viewer_free_style is False
|
||||
assert config.moderation_enabled is False
|
||||
assert config.idle_queue_target == 2
|
||||
assert config.chat_cooldown_s == 7
|
||||
assert config.chat_command == "!go"
|
||||
assert config.hls_dir == "/tmp/hls-under-test"
|
||||
assert config.video_bitrate_k == 1234
|
||||
assert config.web_host == "127.0.0.1"
|
||||
assert config.web_port == 9999
|
||||
assert config.style == "house style"
|
||||
assert config.idle_prompts == ("a lighthouse", "a seagull")
|
||||
|
||||
|
||||
def test_secrets_come_only_from_the_environment(workspace) -> None:
|
||||
"""A key in a version-controlled file is a key that leaks."""
|
||||
assert "OPENAI_API_KEY" not in workspace.read_text()
|
||||
config = Config.load(["--config", str(workspace)])
|
||||
assert config.openai_api_key == "sk-test"
|
||||
# Moderation falls back to the upsampling credentials, which is right when
|
||||
# one endpoint serves both.
|
||||
assert config.moderation_api_key == "sk-test"
|
||||
|
||||
|
||||
def test_cli_overrides_win(workspace, tmp_path) -> None:
|
||||
config = Config.load(["--config", str(workspace), "--port", "4321", "--weights", str(tmp_path / "elsewhere")])
|
||||
assert config.web_port == 4321
|
||||
assert config.weights_path.name == "elsewhere"
|
||||
|
||||
|
||||
def test_a_missing_key_stops_startup(workspace, monkeypatch) -> None:
|
||||
"""Rewriting runs for the idle filler too, so there is no useful run without it."""
|
||||
monkeypatch.delenv("OPENAI_API_KEY")
|
||||
with pytest.raises(SystemExit, match="OPENAI_API_KEY"):
|
||||
Config.load(["--config", str(workspace)])
|
||||
|
||||
|
||||
def test_missing_weights_stops_startup(workspace, monkeypatch) -> None:
|
||||
monkeypatch.delenv("LIVESTREAM_WEIGHTS_PATH")
|
||||
with pytest.raises(SystemExit, match="LIVESTREAM_WEIGHTS_PATH"):
|
||||
Config.load(["--config", str(workspace)])
|
||||
|
||||
|
||||
def test_a_missing_config_file_is_named(tmp_path) -> None:
|
||||
with pytest.raises(SystemExit, match="config not found"):
|
||||
Config.load(["--config", str(tmp_path / "nope.yaml")])
|
||||
|
||||
|
||||
def test_defaults_apply_when_a_block_is_absent(tmp_path, monkeypatch) -> None:
|
||||
"""An operator writing a minimal file must still get a working deployment."""
|
||||
config_file = tmp_path / "minimal.yaml"
|
||||
config_file.write_text("inference:\n aspect: \"16:9\"\nruntime:\n num_gpus: 4\n")
|
||||
monkeypatch.setenv("LIVESTREAM_WEIGHTS_PATH", str(tmp_path))
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
|
||||
config = Config.load(["--config", str(config_file)])
|
||||
assert config.web_port == 8081
|
||||
assert config.chat_cooldown_s == 10
|
||||
assert config.idle_prompts, "the shipped fillers should be used when none is named"
|
||||
|
||||
|
||||
def test_model_config_reads_the_same_file(workspace) -> None:
|
||||
model = load_model_config(workspace)
|
||||
assert model.aspect == "16:9"
|
||||
assert model.clip_frames == 345
|
||||
assert model.runtime["num_gpus"] == 4
|
||||
|
||||
|
||||
def test_a_fillers_directory_without_the_file_is_named(tmp_path) -> None:
|
||||
with pytest.raises(PresetError, match="fillers.json"):
|
||||
load_preset(tmp_path)
|
||||
|
||||
|
||||
def test_the_playlist_default_is_not_the_working_directory(tmp_path, monkeypatch) -> None:
|
||||
"""A relative default would scatter .ts segments wherever the server started."""
|
||||
config_file = tmp_path / "minimal.yaml"
|
||||
config_file.write_text("inference:\n aspect: \"16:9\"\nruntime:\n num_gpus: 4\n")
|
||||
monkeypatch.setenv("LIVESTREAM_WEIGHTS_PATH", str(tmp_path))
|
||||
monkeypatch.setenv("OPENAI_API_KEY", "sk-test")
|
||||
monkeypatch.setenv("XDG_STATE_HOME", str(tmp_path / "state"))
|
||||
# This default is read at import time. A subprocess avoids replacing the
|
||||
# parent's config classes or leaving its default tied to this fixture.
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-c", "from infinite_livestream.config import Config; print(Config.load().hls_dir)",
|
||||
"--config", str(config_file)],
|
||||
cwd=Path(__file__).resolve().parents[2],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=True,
|
||||
)
|
||||
hls_dir = result.stdout.strip()
|
||||
assert hls_dir.startswith(str(tmp_path / "state")), hls_dir
|
||||
@@ -0,0 +1,155 @@
|
||||
"""Whether a viewer's prompt survives, and whether they are told when it does not.
|
||||
|
||||
This path had no tests, which is how a silent drop reached a live stream: the
|
||||
web app answers the POST with `ok` and echoes the prompt into chat, then the
|
||||
director -- downstream of that acknowledgement -- can still refuse it. Anything
|
||||
that refuses here has to report back, or the viewer watches their request
|
||||
appear and then vanish.
|
||||
|
||||
No GPU, no network: the engine, upsampler and moderator are stubs, because what
|
||||
is under test is the admission decision, not what happens after it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
|
||||
from infinite_livestream.chat import ChatPrompt
|
||||
from infinite_livestream.director import Director
|
||||
|
||||
|
||||
class FakeEngine:
|
||||
"""Just enough of `Engine` for the director's capacity checks."""
|
||||
|
||||
def __init__(self, playout_capacity: int = 10) -> None:
|
||||
self.generation_clips: list[dict] = []
|
||||
self.playout_clips: list[dict] = []
|
||||
self.playout_capacity = playout_capacity
|
||||
self.generation_capacity = 20
|
||||
self.playout_queued = 0
|
||||
self.generation_queued = 0
|
||||
self.min_seconds, self.max_seconds = 5.167, 14.375
|
||||
self.connected = True
|
||||
self.commands: list[tuple[str, dict]] = []
|
||||
|
||||
def add_listener(self, listener) -> None:
|
||||
pass
|
||||
|
||||
async def send_command(self, command: str, data: dict):
|
||||
self.commands.append((command, data))
|
||||
return {"clip": {"clip_id": "x" * 12, "seconds": 14.4, "seed": 1}}
|
||||
|
||||
|
||||
class FakeModerator:
|
||||
enabled = True
|
||||
|
||||
def __init__(self, verdict: str | None = None) -> None:
|
||||
self.verdict = verdict
|
||||
|
||||
async def review(self, text: str) -> str | None:
|
||||
return self.verdict
|
||||
|
||||
|
||||
def make_director(rejections: list[tuple[str, str]], *, cooldown_s: float = 10.0,
|
||||
engine: FakeEngine | None = None, moderator: FakeModerator | None = None) -> Director:
|
||||
# Deliberate test doubles: what is under test is the admission decision,
|
||||
# which touches none of the real collaborators' behaviour.
|
||||
return Director(
|
||||
cast("Any", engine or FakeEngine()),
|
||||
upsampler=cast("Any", None),
|
||||
moderator=cast("Any", moderator or FakeModerator()),
|
||||
cooldown_s=cooldown_s,
|
||||
idle_prompts=(),
|
||||
idle_queue_target=0,
|
||||
on_reject=lambda author, reason: rejections.append((author, reason)),
|
||||
)
|
||||
|
||||
|
||||
def prompt(author: str = "ada", text: str = "a lighthouse keeper") -> ChatPrompt:
|
||||
return ChatPrompt(source="web", author=author, text=text, command="!prompt")
|
||||
|
||||
|
||||
def test_first_prompt_is_accepted() -> None:
|
||||
rejections: list[tuple[str, str]] = []
|
||||
director = make_director(rejections)
|
||||
director.submit(prompt())
|
||||
assert rejections == []
|
||||
assert director._pending.qsize() == 1
|
||||
|
||||
|
||||
def test_cooldown_drop_is_reported_to_the_viewer() -> None:
|
||||
"""The bug that reached production: accepted, echoed, then silently gone."""
|
||||
rejections: list[tuple[str, str]] = []
|
||||
director = make_director(rejections, cooldown_s=30.0)
|
||||
director.submit(prompt())
|
||||
director.submit(prompt())
|
||||
assert director._pending.qsize() == 1, "the second must not be queued"
|
||||
assert len(rejections) == 1, "and the viewer must be told"
|
||||
author, reason = rejections[0]
|
||||
assert author == "ada"
|
||||
assert "s left" in reason, f"the reason should say how long to wait, got {reason!r}"
|
||||
|
||||
|
||||
def test_cooldown_is_per_author() -> None:
|
||||
"""Two people must not silence each other."""
|
||||
rejections: list[tuple[str, str]] = []
|
||||
director = make_director(rejections, cooldown_s=30.0)
|
||||
director.submit(prompt(author="ada"))
|
||||
director.submit(prompt(author="grace"))
|
||||
assert rejections == []
|
||||
assert director._pending.qsize() == 2
|
||||
|
||||
|
||||
def test_backlog_full_is_reported() -> None:
|
||||
rejections: list[tuple[str, str]] = []
|
||||
director = make_director(rejections, cooldown_s=0.0)
|
||||
for i in range(64):
|
||||
director.submit(prompt(author=f"viewer{i}"))
|
||||
assert rejections, "a full backlog must be reported, not swallowed"
|
||||
assert any("backlog" in reason for _, reason in rejections)
|
||||
|
||||
|
||||
def test_moderation_rejection_is_reported() -> None:
|
||||
rejections: list[tuple[str, str]] = []
|
||||
director = make_director(rejections, moderator=FakeModerator(verdict="flagged: violence"))
|
||||
director.submit(prompt())
|
||||
asyncio.run(_drain_once(director))
|
||||
assert rejections == [("ada", "flagged: violence")]
|
||||
|
||||
|
||||
def test_viewer_budget_full_is_reported() -> None:
|
||||
"""A queue already full of viewer content refuses more, and says so."""
|
||||
rejections: list[tuple[str, str]] = []
|
||||
engine = FakeEngine(playout_capacity=2)
|
||||
engine.playout_clips = [{"metadata": "{}", "clip_id": "a"}, {"metadata": "{}", "clip_id": "b"}]
|
||||
director = make_director(rejections, engine=engine)
|
||||
director.submit(prompt())
|
||||
asyncio.run(_drain_once(director))
|
||||
assert rejections and "already queued" in rejections[0][1]
|
||||
|
||||
|
||||
async def _drain_once(director: Director) -> None:
|
||||
"""Run the prompt loop just long enough to process what is pending."""
|
||||
task = asyncio.create_task(director.run())
|
||||
for _ in range(50):
|
||||
await asyncio.sleep(0)
|
||||
if director._pending.empty():
|
||||
break
|
||||
await asyncio.sleep(0)
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
|
||||
def test_cooldown_remaining_reports_the_wait() -> None:
|
||||
"""The web app asks this before accepting, so a rate-limited viewer is
|
||||
stopped in their own browser rather than announced to the shared feed."""
|
||||
director = make_director([], cooldown_s=30.0)
|
||||
assert director.cooldown_remaining("ada") == 0.0
|
||||
director.submit(prompt(author="ada"))
|
||||
remaining = director.cooldown_remaining("ada")
|
||||
assert 25.0 < remaining <= 30.0
|
||||
assert director.cooldown_remaining("grace") == 0.0, "and it is per author"
|
||||
@@ -0,0 +1,126 @@
|
||||
"""Frame ownership must survive buffering, repetition, short writes, and drops."""
|
||||
|
||||
import asyncio
|
||||
import queue
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
from infinite_livestream.metadata import encode_id3
|
||||
from infinite_livestream.pacer import Pacer
|
||||
from infinite_livestream.sink import AudioFormat, VideoFormat, _PipeWriter
|
||||
|
||||
|
||||
def test_pacer_repeats_and_drops_identity_with_its_frame(monkeypatch):
|
||||
from infinite_livestream import pacer as module
|
||||
monkeypatch.setattr(module, "_BUFFER_SECONDS", 2 / 24)
|
||||
displayed = []
|
||||
|
||||
class Sink:
|
||||
async def start(self, *args):
|
||||
pass
|
||||
|
||||
def send_video(self, frame, metadata):
|
||||
displayed.append((int(frame[0, 0, 0]), metadata))
|
||||
if len(displayed) == 4:
|
||||
raise asyncio.CancelledError
|
||||
|
||||
def send_audio(self, samples):
|
||||
pass
|
||||
|
||||
pacer = Pacer(Sink(), VideoFormat(2, 2, 24), AudioFormat(48000, 1))
|
||||
for value in (1, 2, 3):
|
||||
pacer.submit_video(np.full((2, 2, 3), value, dtype=np.uint8), str(value).encode())
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
asyncio.run(pacer.run())
|
||||
assert displayed == [(2, b"2"), (3, b"3"), (3, b"3"), (3, b"3")]
|
||||
assert pacer.dropped_frames == 1
|
||||
|
||||
|
||||
def test_writer_commits_only_the_surviving_complete_frame():
|
||||
ledger = queue.Queue()
|
||||
written = bytearray()
|
||||
|
||||
class ShortPipe:
|
||||
def write(self, data):
|
||||
# No record may be visible while any bytes remain unwritten.
|
||||
assert ledger.empty()
|
||||
written.extend(data[:2])
|
||||
return min(2, len(data))
|
||||
|
||||
writer = _PipeWriter("metadata-test", 1)
|
||||
writer.attach(ShortPipe(), ledger)
|
||||
writer.submit(b"discarded", b"old-title")
|
||||
assert writer.submit(b"kept frame", b"new-title") == 1
|
||||
writer.start()
|
||||
try:
|
||||
assert ledger.get(timeout=2) == b"new-title"
|
||||
assert written == b"kept frame"
|
||||
assert ledger.empty()
|
||||
finally:
|
||||
writer.close()
|
||||
writer.join(timeout=2)
|
||||
|
||||
|
||||
def test_partial_frame_failure_never_commits_a_title():
|
||||
ledger = queue.Queue()
|
||||
|
||||
class BrokenPipe:
|
||||
first = True
|
||||
|
||||
def write(self, data):
|
||||
if self.first:
|
||||
self.first = False
|
||||
return 1
|
||||
raise BrokenPipeError
|
||||
|
||||
writer = _PipeWriter("broken-metadata-test", 1)
|
||||
writer.attach(BrokenPipe(), ledger)
|
||||
writer.start()
|
||||
try:
|
||||
writer.submit(b"incomplete", b"must not appear")
|
||||
assert writer.broken.wait(2)
|
||||
assert ledger.empty()
|
||||
finally:
|
||||
writer.close()
|
||||
writer.join(timeout=2)
|
||||
|
||||
|
||||
def test_replaced_encoder_cannot_receive_the_previous_writes_title():
|
||||
entered, release = threading.Event(), threading.Event()
|
||||
old_ledger, new_ledger = queue.Queue(), queue.Queue()
|
||||
|
||||
class OldPipe:
|
||||
def write(self, data):
|
||||
entered.set()
|
||||
assert release.wait(2)
|
||||
return len(data)
|
||||
|
||||
class NewPipe:
|
||||
def write(self, data):
|
||||
return len(data)
|
||||
|
||||
writer = _PipeWriter("restart-metadata-test", 2)
|
||||
writer.attach(OldPipe(), old_ledger)
|
||||
writer.start()
|
||||
try:
|
||||
writer.submit(b"old frame", b"old title")
|
||||
assert entered.wait(2)
|
||||
writer.attach(NewPipe(), new_ledger)
|
||||
writer.submit(b"new frame", b"new title")
|
||||
release.set()
|
||||
assert old_ledger.get(timeout=2) == b"old title"
|
||||
assert new_ledger.get(timeout=2) == b"new title"
|
||||
assert old_ledger.empty() and new_ledger.empty()
|
||||
finally:
|
||||
release.set()
|
||||
writer.close()
|
||||
writer.join(timeout=2)
|
||||
|
||||
|
||||
def test_title_encoding_preserves_unicode_and_explicit_black_state():
|
||||
# UTF-8 text must survive as data, including characters meaningful in HTML.
|
||||
value = {"clip_id": "a", "title": "猫 <script> & café"}
|
||||
assert "猫 <script> & café".encode() in encode_id3(value)
|
||||
assert b'"clip":null' in encode_id3(None)
|
||||
@@ -0,0 +1,136 @@
|
||||
"""Exercise real FFmpeg/PyAV packets on CPU; no model or API keys are needed."""
|
||||
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
from fractions import Fraction
|
||||
|
||||
import av
|
||||
import pytest
|
||||
|
||||
from infinite_livestream.metadata import clip_view, encode_id3
|
||||
from infinite_livestream.muxer import MetadataMuxer
|
||||
|
||||
|
||||
def descriptor(name):
|
||||
return clip_view({"clip_id": name, "prompt": name})
|
||||
|
||||
|
||||
def read_record(raw):
|
||||
# Parse the TXXX UTF-8 value independently of the production serializer.
|
||||
assert raw[:3] == b"ID3" and raw[10:14] == b"TXXX" and raw[20] == 3
|
||||
description, value = raw[21:].split(b"\x00", 1)
|
||||
assert description == b"infinite-livestream"
|
||||
return json.loads(value)
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def encoded_source(tmp_path):
|
||||
if not shutil.which("ffmpeg"):
|
||||
pytest.skip("FFmpeg with libx264/AAC is required for the CPU media integration test")
|
||||
source = tmp_path / "source.ts"
|
||||
subprocess.run([
|
||||
"ffmpeg", "-v", "error", "-f", "lavfi", "-i",
|
||||
"color=red:s=160x96:r=24:d=1.125[a];color=green:s=160x96:r=24:d=1.5[b];"
|
||||
"color=blue:s=160x96:r=24:d=3.375[c];[a][b][c]concat=n=3:v=1:a=0",
|
||||
"-f", "lavfi", "-i", "anullsrc=r=48000:cl=mono", "-t", "6",
|
||||
"-c:v", "libx264", "-preset", "ultrafast", "-tune", "zerolatency",
|
||||
"-g", "48", "-sc_threshold", "0", "-bf", "0", "-c:a", "aac",
|
||||
"-f", "mpegts", str(source),
|
||||
], check=True, capture_output=True, timeout=20)
|
||||
return source
|
||||
|
||||
|
||||
def run_mux(source, playlist, prefix=""):
|
||||
with source.open("rb") as stream:
|
||||
muxer = MetadataMuxer(stream, playlist, 24, 120)
|
||||
for frame in range(144):
|
||||
name = "A" if frame < 27 else "B" if frame < 63 else "C"
|
||||
muxer.frames.put_nowait(encode_id3(descriptor(prefix + name)))
|
||||
muxer._mux()
|
||||
assert muxer.frames.empty()
|
||||
return muxer
|
||||
|
||||
|
||||
def segment_paths(playlist):
|
||||
return [playlist.parent / line for line in playlist.read_text().splitlines()
|
||||
if line and not line.startswith("#")]
|
||||
|
||||
|
||||
def media_packets(paths):
|
||||
result = {"video": [], "audio": []}
|
||||
for path in paths:
|
||||
with av.open(str(path)) as media:
|
||||
for packet in media.demux():
|
||||
if packet.dts is not None and packet.stream.type in result:
|
||||
result[packet.stream.type].append((packet.pts * packet.time_base,
|
||||
packet.dts * packet.time_base, bytes(packet)))
|
||||
return result
|
||||
|
||||
|
||||
def test_metadata_matches_decoded_frames_and_every_segment_start(encoded_source, tmp_path):
|
||||
playlist = tmp_path / "stream.m3u8"
|
||||
run_mux(encoded_source, playlist)
|
||||
segments = segment_paths(playlist)
|
||||
assert len(segments) == 3
|
||||
# No second encode, no retiming, no lost video/audio packets.
|
||||
assert media_packets(segments) == media_packets([encoded_source])
|
||||
transitions = []
|
||||
for path in segments:
|
||||
cues = []
|
||||
with av.open(str(path)) as media:
|
||||
for packet in media.demux():
|
||||
if packet.stream.type == "data" and packet.pts is not None:
|
||||
cues.append((packet.pts * packet.time_base, read_record(bytes(packet))["clip"]["clip_id"]))
|
||||
with av.open(str(path)) as media:
|
||||
frames = list(media.decode(video=0))
|
||||
assert cues[0][0] == frames[0].pts * frames[0].time_base
|
||||
for frame in frames:
|
||||
at = frame.pts * frame.time_base
|
||||
title = [name for pts, name in cues if pts <= at][-1]
|
||||
# The generated color is an independent oracle for clip identity.
|
||||
rgb = frame.to_ndarray(format="rgb24").mean(axis=(0, 1))
|
||||
assert title == "ABC"[int(rgb.argmax())]
|
||||
transitions.extend(cues)
|
||||
first = transitions[0][0]
|
||||
assert (first + Fraction(27, 24), "B") in transitions
|
||||
assert (first + Fraction(63, 24), "C") in transitions
|
||||
|
||||
|
||||
def test_restart_keeps_old_segments_and_marks_the_new_timeline(encoded_source, tmp_path):
|
||||
playlist = tmp_path / "stream.m3u8"
|
||||
first = run_mux(encoded_source, playlist)
|
||||
old_segments = segment_paths(playlist)
|
||||
old_bytes = [p.read_bytes() for p in old_segments]
|
||||
second = run_mux(encoded_source, playlist, prefix="restart-")
|
||||
segments = segment_paths(playlist)
|
||||
assert len(segments) == 6
|
||||
assert first.epoch != second.epoch
|
||||
assert segments[:3] == old_segments
|
||||
assert [p.read_bytes() for p in old_segments] == old_bytes
|
||||
text = playlist.read_text()
|
||||
assert "#EXT-X-DISCONTINUITY\n" + "#EXTINF" in text
|
||||
assert len({p.name for p in segments}) == 6
|
||||
for path in segments[3:]:
|
||||
with av.open(str(path)) as media:
|
||||
packet = next(p for p in media.demux() if p.stream.type == "data" and p.pts is not None)
|
||||
assert read_record(bytes(packet))["clip"]["clip_id"].startswith("restart-")
|
||||
|
||||
|
||||
def test_cleanup_preserves_listed_history_and_removes_expired_orphans(tmp_path):
|
||||
import os
|
||||
import time
|
||||
playlist = tmp_path / "stream.m3u8"
|
||||
listed = tmp_path / "seg_previous_01.ts"
|
||||
orphan = tmp_path / "seg_previous_02.ts"
|
||||
recent = tmp_path / "seg_previous_03.ts"
|
||||
for p in (listed, orphan, recent):
|
||||
p.write_bytes(b"media")
|
||||
playlist.write_text("#EXTM3U\n#EXTINF:2,\n" + listed.name + "\n")
|
||||
for p in (listed, orphan):
|
||||
os.utime(p, (time.time() - 300, time.time() - 300))
|
||||
with listed.open("rb") as source:
|
||||
muxer = MetadataMuxer(source, playlist, 24, 120)
|
||||
muxer._cleanup()
|
||||
assert listed.exists() and recent.exists()
|
||||
assert not orphan.exists()
|
||||
@@ -0,0 +1,99 @@
|
||||
"""The app depends on FastVideo and nothing else that serves models.
|
||||
|
||||
This app began as a port of a deployment built on the Reactor runtime, whose
|
||||
serve process, RPC decorators and wire schema it no longer uses -- the model
|
||||
and the broadcast run in one process now. That is easy to regress by copying
|
||||
one more module across, so the contract is a test: no `reactor_*` import may
|
||||
reappear, and the modules that must stay importable without a GPU must stay
|
||||
importable without a GPU.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import ast
|
||||
import pathlib
|
||||
|
||||
import pytest
|
||||
|
||||
PACKAGE = pathlib.Path(__file__).resolve().parents[1]
|
||||
|
||||
# Modules that must import with no torch, no fastvideo and no GPU: the config
|
||||
# and queue logic is pure Python so it can be tested anywhere, and the entry
|
||||
# point has to be able to print a dependency error rather than raise one.
|
||||
CPU_ONLY_MODULES = (
|
||||
"infinite_livestream.clip_plan",
|
||||
"infinite_livestream.clip_queue",
|
||||
"infinite_livestream.config",
|
||||
"infinite_livestream.group_tag",
|
||||
)
|
||||
|
||||
|
||||
def _module_files() -> list[pathlib.Path]:
|
||||
return sorted(p for p in PACKAGE.rglob("*.py") if "tests" not in p.parts)
|
||||
|
||||
|
||||
def _imported_names(path: pathlib.Path) -> set[str]:
|
||||
tree = ast.parse(path.read_text(encoding="utf-8"))
|
||||
names: set[str] = set()
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Import):
|
||||
names.update(alias.name for alias in node.names)
|
||||
elif isinstance(node, ast.ImportFrom) and node.module and node.level == 0:
|
||||
names.add(node.module)
|
||||
return names
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", _module_files(), ids=lambda p: p.name)
|
||||
def test_no_reactor_imports(path: pathlib.Path) -> None:
|
||||
offenders = {name for name in _imported_names(path) if name.split(".")[0].startswith("reactor")}
|
||||
assert not offenders, f"{path.name} imports {sorted(offenders)}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("module", CPU_ONLY_MODULES)
|
||||
def test_imports_without_a_gpu(module: str) -> None:
|
||||
"""These must not drag torch or fastvideo in as a side effect.
|
||||
|
||||
Run in a fresh interpreter, because "did importing X pull in torch" is a
|
||||
question about ``sys.modules``, which is global: any earlier test that
|
||||
imported fastvideo would make this pass or fail for reasons that have
|
||||
nothing to do with the module under test. A subprocess is the only honest
|
||||
way to ask it.
|
||||
"""
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
probe = (
|
||||
f"import {module}, sys; "
|
||||
"leaked = [m for m in ('torch', 'fastvideo') if m in sys.modules]; "
|
||||
"print(','.join(leaked))"
|
||||
)
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-c", probe],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
cwd=PACKAGE.parent,
|
||||
)
|
||||
assert result.returncode == 0, f"{module} failed to import:\n{result.stderr}"
|
||||
leaked = result.stdout.strip()
|
||||
assert not leaked, f"{module} imported {leaked} at module level"
|
||||
|
||||
|
||||
def test_backend_defers_heavy_imports() -> None:
|
||||
"""backend.py names fastvideo only inside functions.
|
||||
|
||||
Module-level would make the config tests, and the entry point's dependency
|
||||
message, need a GPU.
|
||||
"""
|
||||
tree = ast.parse((PACKAGE / "backend.py").read_text(encoding="utf-8"))
|
||||
top_level = {
|
||||
alias.name
|
||||
for node in tree.body
|
||||
if isinstance(node, ast.Import)
|
||||
for alias in node.names
|
||||
} | {
|
||||
node.module
|
||||
for node in tree.body
|
||||
if isinstance(node, ast.ImportFrom) and node.module and node.level == 0
|
||||
}
|
||||
heavy = {name for name in top_level if name.split(".")[0] in {"torch", "torchaudio", "fastvideo"}}
|
||||
assert not heavy, f"backend.py imports {sorted(heavy)} at module level"
|
||||
@@ -0,0 +1,219 @@
|
||||
// Adapter regressions; actual cue timing is owned by the browser/player.
|
||||
const { test } = require('node:test');
|
||||
const assert = require('node:assert/strict');
|
||||
require('../web/metadata.js');
|
||||
|
||||
class Track extends EventTarget {
|
||||
kind = 'metadata';
|
||||
mode = 'disabled';
|
||||
activeCues = [];
|
||||
activate(cues) {
|
||||
this.activeCues = cues;
|
||||
this.dispatchEvent(new Event('cuechange'));
|
||||
}
|
||||
}
|
||||
class TrackList extends EventTarget {
|
||||
items = [];
|
||||
[Symbol.iterator]() { return this.items[Symbol.iterator](); }
|
||||
add(track) {
|
||||
if (!this.items.includes(track)) this.items.push(track);
|
||||
const event = new Event('addtrack');
|
||||
event.track = track;
|
||||
this.dispatchEvent(event);
|
||||
}
|
||||
}
|
||||
class Video extends EventTarget {
|
||||
readyState = 2;
|
||||
textTracks = new TrackList();
|
||||
}
|
||||
function cue(name, startTime = 0, native = false) {
|
||||
const clip = name === null ? null : { clip_id: name, title: name, prompt: name, generated: false };
|
||||
const value = { key: 'TXXX', data: JSON.stringify({ version: 1, clip }) };
|
||||
if (!native) value.info = 'infinite-livestream';
|
||||
return { startTime, endTime: Infinity, value };
|
||||
}
|
||||
function player() {
|
||||
const video = new Video();
|
||||
const track = new Track();
|
||||
video.textTracks.add(track);
|
||||
const shown = [];
|
||||
const stop = LivestreamMetadata.watch(video, clip => shown.push(clip?.clip_id ?? clip));
|
||||
return { video, track, shown, stop };
|
||||
}
|
||||
|
||||
test('downloaded future cues do not advance the title; activation does', () => {
|
||||
const p = player();
|
||||
assert.equal(p.track.mode, 'hidden');
|
||||
const a = cue('A');
|
||||
const b = cue('B', 10);
|
||||
p.track.activate([a]);
|
||||
p.track.cues = [a, b]; // The next segment arrived, but it is not playing.
|
||||
p.video.dispatchEvent(new Event('loadeddata'));
|
||||
assert.deepEqual(p.shown, [undefined, 'A']);
|
||||
p.track.activate([b]);
|
||||
assert.deepEqual(p.shown, [undefined, 'A', 'B']);
|
||||
p.stop();
|
||||
});
|
||||
|
||||
test('viewers keep independent titles during stalls and seeks', () => {
|
||||
const live = player();
|
||||
const delayed = player();
|
||||
live.track.activate([cue('C', 20)]);
|
||||
delayed.track.activate([cue('A', 0)]);
|
||||
delayed.video.dispatchEvent(new Event('waiting'));
|
||||
live.track.activate([cue('D', 30)]);
|
||||
assert.equal(delayed.shown.at(-1), 'A');
|
||||
delayed.track.activate([cue('B', 10)]);
|
||||
delayed.video.dispatchEvent(new Event('seeked'));
|
||||
assert.equal(delayed.shown.at(-1), 'B');
|
||||
assert.equal(live.shown.at(-1), 'D');
|
||||
delayed.track.activate([cue('A')]); // Seek backwards.
|
||||
assert.equal(delayed.shown.at(-1), 'A');
|
||||
live.stop(); delayed.stop();
|
||||
});
|
||||
|
||||
test('joining mid-clip and repeated segment markers do not need prior events', () => {
|
||||
const video = new Video();
|
||||
const shown = [];
|
||||
const stop = LivestreamMetadata.watch(video, clip => shown.push(clip?.clip_id ?? clip));
|
||||
const track = new Track();
|
||||
track.activeCues = [cue('B', 14)];
|
||||
video.textTracks.add(track);
|
||||
track.activate([cue('B', 16)]);
|
||||
video.textTracks.add(track); // hls.js may reuse a track across attachments.
|
||||
assert.deepEqual(shown, [undefined, 'B']);
|
||||
stop();
|
||||
});
|
||||
|
||||
test('native overlapping cues pick the latest record, including black frames', () => {
|
||||
const p = player();
|
||||
p.track.activate([cue('B', 10, true), cue('A', 0, true)]);
|
||||
assert.equal(p.shown.at(-1), 'B');
|
||||
p.track.activate([cue('A', 0, true), cue(null, 20, true), cue('B', 10, true)]);
|
||||
assert.equal(p.shown.at(-1), null);
|
||||
p.stop();
|
||||
});
|
||||
|
||||
test('invalid and unrelated ID3 cannot replace a valid title', () => {
|
||||
const p = player();
|
||||
const good = cue('猫 <script> & café', 1);
|
||||
const unrelated = cue('wrong', 10);
|
||||
unrelated.value.info = 'another-application';
|
||||
p.track.activate([good, unrelated, { startTime: 20, value: { key: 'TXXX', data: 'invalid' } }]);
|
||||
assert.equal(p.shown.at(-1), '猫 <script> & café');
|
||||
const fallback = cue('text-fallback', 30);
|
||||
fallback.text = JSON.stringify(fallback.value);
|
||||
delete fallback.value;
|
||||
p.track.activate([fallback]);
|
||||
assert.equal(p.shown.at(-1), 'text-fallback');
|
||||
p.stop();
|
||||
});
|
||||
|
||||
test('loading and teardown clear stale titles and detach old listeners', () => {
|
||||
const p = player();
|
||||
p.track.activate([cue('A')]);
|
||||
p.video.readyState = 0;
|
||||
p.video.dispatchEvent(new Event('emptied'));
|
||||
p.track.activate([cue('B')]);
|
||||
assert.equal(p.shown.at(-1), undefined);
|
||||
p.video.readyState = 2;
|
||||
p.video.dispatchEvent(new Event('loadeddata'));
|
||||
assert.equal(p.shown.at(-1), 'B');
|
||||
p.stop();
|
||||
p.track.activate([cue('C')]);
|
||||
p.video.textTracks.add(new Track());
|
||||
p.video.dispatchEvent(new Event('emptied'));
|
||||
assert.equal(p.shown.at(-1), 'B');
|
||||
});
|
||||
|
||||
|
||||
test('presented frames resolve paused seeks before segment PTS without guessing', () => {
|
||||
const video = new Video();
|
||||
let callback;
|
||||
let cancelled = false;
|
||||
video.requestVideoFrameCallback = fn => { callback = fn; return 1; };
|
||||
video.cancelVideoFrameCallback = id => { assert.equal(id, 1); cancelled = true; };
|
||||
const track = new Track();
|
||||
video.textTracks.add(track);
|
||||
const shown = [];
|
||||
const stop = LivestreamMetadata.watch(video, clip => shown.push(clip?.clip_id ?? clip));
|
||||
const a = cue('A', 2.021333333333333);
|
||||
const b = cue('B', 12.021333333333333);
|
||||
a.endTime = b.startTime;
|
||||
track.cues = [a, b];
|
||||
track.activate([a]);
|
||||
callback(0, { mediaTime: 3 });
|
||||
assert.equal(shown.at(-1), 'A');
|
||||
// Observed Firefox case: currentTime=12, active cue A, but displayed frame B.
|
||||
video.currentTime = 12;
|
||||
callback(0, { mediaTime: 12.021333 });
|
||||
assert.equal(shown.at(-1), 'B');
|
||||
track.activate([a]);
|
||||
assert.equal(shown.at(-1), 'B');
|
||||
// A future download and a stall do not move the presented frame.
|
||||
track.cues.push(cue('C', 20));
|
||||
video.dispatchEvent(new Event('waiting'));
|
||||
assert.equal(shown.at(-1), 'B');
|
||||
callback(0, { mediaTime: 3 });
|
||||
assert.equal(shown.at(-1), 'A');
|
||||
// A paused seek may have no frame callback: resume standard cue scheduling.
|
||||
track.activate([b]);
|
||||
video.dispatchEvent(new Event('seeked'));
|
||||
assert.equal(shown.at(-1), 'B');
|
||||
video.dispatchEvent(new Event('emptied'));
|
||||
assert.equal(shown.at(-1), undefined);
|
||||
stop();
|
||||
assert.equal(cancelled, true);
|
||||
});
|
||||
|
||||
|
||||
test('a browser advertising native HLS still uses the metadata-capable hls.js path', () => {
|
||||
const fs = require('node:fs');
|
||||
const vm = require('node:vm');
|
||||
const html = fs.readFileSync(require.resolve('../web/index.html'), 'utf8');
|
||||
const source = html.slice(html.indexOf('function startPlayback()'), html.indexOf('video.addEventListener("error"'));
|
||||
let loaded;
|
||||
let attached;
|
||||
class Hls {
|
||||
static isSupported() { return true; }
|
||||
static Events = { ERROR: 'error' };
|
||||
on() {}
|
||||
loadSource(url) { loaded = url; }
|
||||
attachMedia(video) { attached = video; }
|
||||
}
|
||||
const video = { canPlayType: () => 'probably' };
|
||||
const context = vm.createContext({
|
||||
Hls, window: { Hls }, video, PLAYLIST: '/hls/stream.m3u8',
|
||||
LivestreamMetadata: { watch: () => () => {} }, renderPlayback() {},
|
||||
});
|
||||
vm.runInContext(source + '\nstartPlayback();', context);
|
||||
assert.equal(loaded, '/hls/stream.m3u8');
|
||||
assert.equal(attached, video);
|
||||
assert.equal(video.src, undefined);
|
||||
// Preserve native-only devices when Media Source playback is unavailable.
|
||||
Hls.isSupported = () => false;
|
||||
loaded = undefined;
|
||||
vm.runInContext('startPlayback();', context);
|
||||
assert.equal(video.src, '/hls/stream.m3u8');
|
||||
assert.equal(loaded, undefined);
|
||||
});
|
||||
|
||||
test('missing metadata does not claim an already-playing video is loading', () => {
|
||||
const fs = require('node:fs');
|
||||
const vm = require('node:vm');
|
||||
const html = fs.readFileSync(require.resolve('../web/index.html'), 'utf8');
|
||||
const source = html.slice(html.indexOf('function renderPlayback()'), html.indexOf('function render(state)'));
|
||||
const elements = {};
|
||||
const video = { readyState: 4, paused: false };
|
||||
const context = vm.createContext({
|
||||
video, playbackClip: undefined, playbackError: '', buffering: false,
|
||||
el: id => elements[id] ||= {},
|
||||
});
|
||||
vm.runInContext(source + '\nrenderPlayback();', context);
|
||||
assert.equal(elements.livetext.textContent, 'on air');
|
||||
assert.equal(elements['np-title'].textContent, 'Waiting for title metadata');
|
||||
video.readyState = 0;
|
||||
vm.runInContext('renderPlayback();', context);
|
||||
assert.equal(elements.livetext.textContent, 'loading');
|
||||
assert.equal(elements['np-title'].textContent, 'Waiting for video');
|
||||
});
|
||||
@@ -0,0 +1,97 @@
|
||||
"""A submitted viewer prompt must become ordered clips in the real engine queue.
|
||||
|
||||
Only provider calls and GPU readiness are substituted. Chat intake, moderation,
|
||||
rewriting, admission, metadata, queue insertion, and web state all run normally.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
import pytest
|
||||
|
||||
from infinite_livestream import moderator as moderator_module
|
||||
from infinite_livestream import upsampler as upsampler_module
|
||||
from infinite_livestream.chat import WebChat
|
||||
from infinite_livestream.config import load_model_config
|
||||
from infinite_livestream.director import Director
|
||||
from infinite_livestream.engine import Engine
|
||||
from infinite_livestream.webapp import DemoWeb
|
||||
|
||||
|
||||
@pytest.mark.parametrize("scene_count", [1, 3])
|
||||
def test_accepted_prompt_reaches_generation_queue(app_config, monkeypatch, scene_count):
|
||||
scenes = [{"prompt": f"A lighthouse keeper feeds seagull {i}", "seconds": 8.0 + i}
|
||||
for i in range(scene_count)]
|
||||
provider_calls = []
|
||||
|
||||
async def moderate(**kwargs):
|
||||
provider_calls.append(("moderation", kwargs["input"]))
|
||||
return SimpleNamespace(results=[SimpleNamespace(flagged=False)])
|
||||
|
||||
async def rewrite(**kwargs):
|
||||
provider_calls.append(("rewrite", kwargs["messages"][-1]["content"]))
|
||||
content = json.dumps({"title": "The lighthouse", "scenes": scenes})
|
||||
return SimpleNamespace(choices=[SimpleNamespace(message=SimpleNamespace(content=content))])
|
||||
|
||||
provider = SimpleNamespace(
|
||||
moderations=SimpleNamespace(create=moderate),
|
||||
chat=SimpleNamespace(completions=SimpleNamespace(create=rewrite)),
|
||||
)
|
||||
monkeypatch.setattr(moderator_module, "AsyncOpenAI", lambda **kwargs: provider)
|
||||
monkeypatch.setattr(upsampler_module, "AsyncOpenAI", lambda **kwargs: provider)
|
||||
|
||||
async def run():
|
||||
engine = Engine(app_config, load_model_config(app_config.config_path))
|
||||
engine._ready.set() # Queue operations need readiness, never GPU work.
|
||||
await engine.send_command("enqueue", {"prompt": "earlier viewer"})
|
||||
await engine.send_command("enqueue", {"prompt": "idle filler", "metadata": json.dumps({
|
||||
"group_id": "idle", "title": "idle", "author": "filler", "source": "idle",
|
||||
"scene": 1, "scenes": 1, "generated": True,
|
||||
})})
|
||||
chat = WebChat()
|
||||
web = DemoWeb(chat, app_config.hls_dir)
|
||||
engine.add_listener(web.listener)
|
||||
moderator = moderator_module.Moderator("test-key", "test-model", enabled=True)
|
||||
upsampler = upsampler_module.PromptUpsampler("test-key", "test-model", "house style", max_chunks=6)
|
||||
rejections = []
|
||||
director = Director(engine, upsampler, moderator, cooldown_s=10,
|
||||
on_reject=lambda author, reason: rejections.append((author, reason)))
|
||||
web.cooldown_remaining = director.cooldown_remaining
|
||||
queued = asyncio.Event()
|
||||
|
||||
def on_queue(kind, data):
|
||||
if kind == "queue_update" and len(data["generation"]) == scene_count + 2:
|
||||
queued.set()
|
||||
|
||||
engine.add_listener(on_queue)
|
||||
# Submit before starting intake so TestClient never wakes a queue
|
||||
# waiter owned by a different event loop.
|
||||
with TestClient(web.app) as client:
|
||||
response = client.post("/chat", json={"author": "ada", "text": "a lighthouse keeper"})
|
||||
assert response.status_code == 200 and response.json() == {"ok": True}
|
||||
tasks = [asyncio.create_task(chat.run(director.submit)), asyncio.create_task(director.run())]
|
||||
try:
|
||||
await asyncio.wait_for(queued.wait(), timeout=2)
|
||||
clips = engine.generation_clips
|
||||
assert [c["prompt"] for c in clips] == ["earlier viewer", *[s["prompt"] for s in scenes], "idle filler"]
|
||||
tags = [json.loads(c["metadata"]) for c in clips[1:-1]]
|
||||
assert len({tag["group_id"] for tag in tags}) == 1
|
||||
assert [tag["scene"] for tag in tags] == list(range(1, scene_count + 1))
|
||||
assert all(tag["scenes"] == scene_count and tag["author"] == "ada"
|
||||
and tag["raw_prompt"] == "a lighthouse keeper" and not tag["generated"] for tag in tags)
|
||||
if scene_count == 1:
|
||||
assert clips[1]["frames"] == 362, "single scenes must use the maximum clip length"
|
||||
assert [c["clip_id"] for c in web.state.generation] == [c["clip_id"] for c in clips]
|
||||
assert all(c["prompt"] == "a lighthouse keeper" for c in web.state.generation[1:-1])
|
||||
assert [name for name, _ in provider_calls] == ["moderation", "rewrite"]
|
||||
assert provider_calls[0][1] == "a lighthouse keeper"
|
||||
assert "a lighthouse keeper" in provider_calls[1][1]
|
||||
assert rejections == []
|
||||
finally:
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
asyncio.run(run())
|
||||
@@ -0,0 +1,55 @@
|
||||
"""Service failures must reach the caller after the other tasks are stopped."""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from infinite_livestream import main
|
||||
from infinite_livestream.backend import FastH3Backend
|
||||
|
||||
|
||||
@pytest.mark.parametrize("failure", ["engine", "pacer", None])
|
||||
def test_service_propagates_failure_and_cleans_up(app_config, monkeypatch, failure):
|
||||
error = RuntimeError(f"{failure} startup failed")
|
||||
|
||||
def load():
|
||||
if failure == "engine":
|
||||
raise error
|
||||
|
||||
async def wait_forever():
|
||||
await asyncio.Event().wait()
|
||||
|
||||
async def run_pacer():
|
||||
if failure == "pacer":
|
||||
raise error
|
||||
await wait_forever()
|
||||
|
||||
async def run_web():
|
||||
if failure is None:
|
||||
# Normal server shutdown still exits successfully.
|
||||
await asyncio.sleep(0)
|
||||
return
|
||||
await wait_forever()
|
||||
|
||||
sink = Mock(stop=AsyncMock())
|
||||
web = Mock(run=run_web)
|
||||
pacer = Mock(run=run_pacer)
|
||||
monkeypatch.setattr(FastH3Backend, "load", lambda self: load())
|
||||
monkeypatch.setattr(main, "HlsSink", Mock(return_value=sink))
|
||||
monkeypatch.setattr(main, "DemoWeb", Mock(return_value=web))
|
||||
monkeypatch.setattr(main, "Pacer", Mock(return_value=pacer))
|
||||
monkeypatch.setattr(main, "PromptUpsampler", Mock())
|
||||
monkeypatch.setattr(main, "Moderator", Mock())
|
||||
|
||||
async def run():
|
||||
if failure is None:
|
||||
await main.serve(app_config)
|
||||
else:
|
||||
with pytest.raises(RuntimeError) as caught:
|
||||
await main.serve(app_config)
|
||||
assert caught.value is error
|
||||
sink.stop.assert_awaited_once()
|
||||
assert asyncio.all_tasks() == {asyncio.current_task()}
|
||||
|
||||
asyncio.run(run())
|
||||
@@ -0,0 +1,128 @@
|
||||
"""Use real OS pipes to cover shutdown under encoder backpressure, without GPUs."""
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import select
|
||||
import signal
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
|
||||
from infinite_livestream import sink as sink_module
|
||||
from infinite_livestream.sink import HlsSink, _PipeWriter
|
||||
|
||||
|
||||
class ObservedPipe:
|
||||
def __init__(self, pipe):
|
||||
self.pipe = pipe
|
||||
self.writing = threading.Event()
|
||||
|
||||
def write(self, payload):
|
||||
self.writing.set()
|
||||
return self.pipe.write(payload)
|
||||
|
||||
|
||||
@pytest.mark.skipif(os.name != "posix", reason="the HLS sink uses POSIX pipes")
|
||||
@pytest.mark.parametrize("ignore_terminate", [False, True])
|
||||
def test_stop_unblocks_full_pipes_and_reaps_process(tmp_path, monkeypatch, ignore_terminate):
|
||||
monkeypatch.setattr(sink_module.shutil, "which", lambda name: "/test/ffmpeg")
|
||||
monkeypatch.setattr(sink_module, "_PROCESS_EXIT_TIMEOUT_S", 0.2, raising=False)
|
||||
sink = HlsSink(tmp_path)
|
||||
audio_read, audio_write = os.pipe()
|
||||
child = (
|
||||
"import signal, time; "
|
||||
+ ("signal.signal(signal.SIGTERM, signal.SIG_IGN); " if ignore_terminate else "")
|
||||
+ "print('ready', flush=True); time.sleep(30)"
|
||||
)
|
||||
process = subprocess.Popen([sys.executable, "-c", child], stdin=subprocess.PIPE,
|
||||
stdout=subprocess.PIPE, bufsize=0, pass_fds=(audio_read,))
|
||||
os.close(audio_read)
|
||||
sink._process = process
|
||||
audio_pipe = os.fdopen(audio_write, "wb", buffering=0)
|
||||
sink._audio_pipe = audio_pipe
|
||||
writers = [_PipeWriter("video-test", 1), _PipeWriter("audio-test", 1)]
|
||||
sink._video_writer, sink._audio_writer = writers
|
||||
errors = []
|
||||
ticks = []
|
||||
stopper = None
|
||||
try:
|
||||
assert select.select([process.stdout], [], [], 5)[0], "child failed to become ready"
|
||||
assert process.stdout.readline() == b"ready\n"
|
||||
for writer, pipe in zip(writers, (process.stdin, sink._audio_pipe)):
|
||||
observed = ObservedPipe(pipe)
|
||||
writer.attach(observed)
|
||||
writer.start()
|
||||
writer.submit(b"x" * (2 * 1024 * 1024))
|
||||
assert observed.writing.wait(2), "writer never reached the pipe"
|
||||
writer.submit(b"queued")
|
||||
assert writer.queue.full()
|
||||
|
||||
async def stop_with_heartbeat():
|
||||
task = asyncio.create_task(sink.stop())
|
||||
while not task.done():
|
||||
ticks.append(True)
|
||||
await asyncio.sleep(0.01)
|
||||
await task
|
||||
|
||||
def stop():
|
||||
try:
|
||||
asyncio.run(stop_with_heartbeat())
|
||||
except BaseException as error:
|
||||
errors.append(error)
|
||||
|
||||
# A separate thread makes the timeout effective even if a regression
|
||||
# blocks the event loop inside a synchronous queue/pipe operation.
|
||||
stopper = threading.Thread(target=stop, daemon=True)
|
||||
stopper.start()
|
||||
stopper.join(timeout=5)
|
||||
assert not stopper.is_alive(), "shutdown blocked on a full pipe or queue"
|
||||
assert errors == []
|
||||
expected_signal = signal.SIGKILL if ignore_terminate else signal.SIGTERM
|
||||
assert process.returncode == -expected_signal
|
||||
assert all(not writer.is_alive() for writer in writers)
|
||||
assert process.stdin.closed and audio_pipe.closed
|
||||
assert sink._audio_pipe is None
|
||||
if ignore_terminate:
|
||||
assert len(ticks) > 1, "waiting for FFmpeg blocked the event loop"
|
||||
finally:
|
||||
if process.poll() is None:
|
||||
process.kill()
|
||||
process.wait(timeout=5)
|
||||
if stopper is not None:
|
||||
stopper.join(timeout=5)
|
||||
for writer in writers:
|
||||
writer.close()
|
||||
if writer.ident is not None:
|
||||
writer.join(timeout=2)
|
||||
process.stdin.close()
|
||||
process.stdout.close()
|
||||
if sink._audio_pipe is not None:
|
||||
sink._audio_pipe.close()
|
||||
|
||||
|
||||
def test_writer_preserves_the_payload_across_short_writes():
|
||||
payload = b"one complete media frame"
|
||||
written = bytearray()
|
||||
complete = threading.Event()
|
||||
|
||||
class ShortPipe:
|
||||
def write(self, data):
|
||||
count = min(3, len(data))
|
||||
written.extend(data[:count])
|
||||
if len(written) == len(payload):
|
||||
complete.set()
|
||||
return count
|
||||
|
||||
writer = _PipeWriter("short-write-test", 1)
|
||||
writer.attach(ShortPipe())
|
||||
writer.start()
|
||||
try:
|
||||
writer.submit(payload)
|
||||
assert complete.wait(2), "a short write discarded the rest of the media payload"
|
||||
assert written == payload
|
||||
finally:
|
||||
writer.close()
|
||||
writer.join(timeout=2)
|
||||
assert not writer.is_alive()
|
||||
@@ -0,0 +1,76 @@
|
||||
"""Chat, queue, and service state stay independent of viewers' playback clocks."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from infinite_livestream.webapp import DemoState
|
||||
|
||||
|
||||
def clip(clip_id: str = "abcdef123456", prompt: str = "a lighthouse keeper", *, generated: bool = False,
|
||||
scene: int | None = None, scenes: int | None = None) -> dict:
|
||||
import json
|
||||
meta = {"group_id": "g1", "title": prompt, "author": "viewer", "generated": generated, "raw_prompt": prompt}
|
||||
if scene is not None:
|
||||
meta |= {"scene": scene, "scenes": scenes}
|
||||
return {"clip_id": clip_id, "prompt": prompt, "metadata": json.dumps(meta), "frames": 345,
|
||||
"seconds": 14.375, "seed": 1, "ready": True}
|
||||
|
||||
|
||||
def test_queue_update_replaces_both_queues() -> None:
|
||||
state = DemoState()
|
||||
state.on_message("queue_update", {"generation": [clip("a")], "playout": [clip("b"), clip("c")]})
|
||||
assert [c["clip_id"] for c in state.generation] == ["a"]
|
||||
assert [c["clip_id"] for c in state.playout] == ["b", "c"]
|
||||
# Replacement, not accumulation: a queue that empties must render empty.
|
||||
state.on_message("queue_update", {"generation": [], "playout": []})
|
||||
assert state.generation == [] and state.playout == []
|
||||
|
||||
|
||||
def test_generating_is_the_generation_front() -> None:
|
||||
"""Builds consume the queue front-first, so the front is what is in flight."""
|
||||
state = DemoState()
|
||||
assert state.generating is None
|
||||
state.on_message("queue_update", {"generation": [clip("a"), clip("b")], "playout": []})
|
||||
generating = state.generating
|
||||
assert generating is not None and generating["clip_id"] == "a"
|
||||
|
||||
|
||||
def test_playout_events_do_not_define_a_viewers_playback_position() -> None:
|
||||
state = DemoState()
|
||||
before = state.snapshot()
|
||||
state.on_message("clip_started", {"clip": clip("a")})
|
||||
state.on_message("clip_finished", {"clip": clip("a"), "seconds_sent": 14.4})
|
||||
assert state.snapshot() == before
|
||||
|
||||
|
||||
def test_only_filler_is_announced_in_chat_and_once_per_group() -> None:
|
||||
"""Viewer submissions are echoed by the POST handler, so only filler here.
|
||||
|
||||
And one line per group, not per scene: a six-scene story is still one
|
||||
thing somebody asked for.
|
||||
"""
|
||||
state = DemoState()
|
||||
state.on_message("clip_queued", {"clip": clip("v", "viewer idea", generated=False)})
|
||||
assert list(state.chat) == []
|
||||
for scene in (1, 2, 3):
|
||||
state.on_message("clip_queued", {"clip": clip(f"f{scene}", "filler idea", generated=True,
|
||||
scene=scene, scenes=3)})
|
||||
assert [c["author"] for c in state.chat] == ["filler"]
|
||||
|
||||
|
||||
def test_failed_viewer_clips_are_reported_but_filler_is_not() -> None:
|
||||
state = DemoState()
|
||||
state.on_message("clip_failed", {"clip": clip("f", generated=True), "reason": "boom"})
|
||||
assert list(state.chat) == []
|
||||
state.on_message("clip_failed", {"clip": clip("v", "viewer idea", generated=False), "reason": "boom"})
|
||||
assert [c["kind"] for c in state.chat] == ["error"]
|
||||
|
||||
|
||||
def test_snapshot_carries_everything_the_page_reads() -> None:
|
||||
state = DemoState()
|
||||
state.on_message("state_update", {"playing": False, "generation_queued": 1, "generation_capacity": 20,
|
||||
"playout_queued": 2, "playout_capacity": 10, "clips_played": 7,
|
||||
"width": 1344, "height": 768})
|
||||
snap = state.snapshot()
|
||||
assert set(snap) == {"connected", "generating",
|
||||
"generation", "playout", "stats", "chat"}
|
||||
assert snap["stats"]["clips_played"] == 7
|
||||
@@ -0,0 +1,359 @@
|
||||
"""Prompt upsampling: a viewer's rough idea into FastH3-ready scenes.
|
||||
|
||||
One LLM call per prompt against any OpenAI-compatible endpoint. The model
|
||||
picks the shape the idea calls for -- one scene, or a chunked short story of
|
||||
up to `max_chunks` clips -- writes each scene as a self-contained
|
||||
text-to-video prompt in the configured style, and picks each scene's length.
|
||||
|
||||
The system prompt is written around four facts about FastH3. Keep them intact
|
||||
when editing it:
|
||||
|
||||
* **Each scene is an independent clip with no memory.** The biggest quality
|
||||
lever by far. "The same forest" renders a *different* forest, so every
|
||||
scene must re-describe setting, subjects, light and style from scratch.
|
||||
* **800 characters is the hard cap per prompt.** The LLM is told 750 for
|
||||
headroom and `_sanitize` truncates anyway, because LLMs do not count
|
||||
characters reliably.
|
||||
* **Audio is generated with the video, speech included.** The prompt asks
|
||||
for quoted dialogue (who speaks, the words, the tone) whenever the idea
|
||||
implies speech, and for a brief soundscape clause. Clips come out flat
|
||||
without them.
|
||||
* **A single-clip generation always runs the maximum length**, enforced in
|
||||
code after validation, so the scene can breathe. Short lengths are
|
||||
reserved for transition chunks inside multi-scene stories.
|
||||
|
||||
Safety is `moderator.py`'s job: the idea has already passed it by the time it
|
||||
arrives here, so this prompt asks for faithful staging and never for
|
||||
softening or reinterpreting.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
|
||||
from openai import AsyncOpenAI
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# The engine's enqueue cap; _sanitize truncates to it.
|
||||
MAX_PROMPT_CHARS = 800
|
||||
# LLM calls one idea gets before falling back to the raw prompt.
|
||||
_MAX_ATTEMPTS = 3
|
||||
|
||||
# What the LLM is asked to stay under, leaving headroom for its poor counting.
|
||||
# Sized so an overshoot still fits under the 800 hard cap: the sanitizer
|
||||
# truncates mid-word at 800, and what it cuts is the prompt's tail — the
|
||||
# soundscape sentence the format deliberately puts last.
|
||||
_TARGET_PROMPT_CHARS = 700
|
||||
|
||||
# What goes in the STYLE slot for a viewer's own request when the deployment
|
||||
# lets viewers out of the house style. The filler still carries the preset's
|
||||
# identity, which is what gives the stream a look of its own between requests;
|
||||
# forcing a viewer's idea into that same look is what makes "a documentary shot
|
||||
# of a snow leopard" come back as a cartoon.
|
||||
_FREE_STYLE = """No house style is imposed on this request. Choose the look that genuinely
|
||||
suits the viewer's idea and commit to it fully — photoreal documentary,
|
||||
anime, stop-motion, 90s camcorder, oil painting, whatever the idea calls
|
||||
for. If the viewer names a style, medium or era, follow it exactly. Describe
|
||||
that look concretely in every scene prompt (lens, lighting, palette, texture,
|
||||
grain, motion) so the clip is unmistakably in it."""
|
||||
|
||||
_SYSTEM_PROMPT = """\
|
||||
You are the scene director of a live, chat-driven AI video stream. Viewers
|
||||
send short, rough ideas; you turn each one into one or more polished
|
||||
text-to-video prompts for a model that generates short clips with
|
||||
synchronized audio.
|
||||
|
||||
STYLE / CHARACTER — every scene is rendered in this identity; weave it into
|
||||
every scene prompt, never contradict it:
|
||||
{style}
|
||||
|
||||
HOW THE VIDEO MODEL WORKS (hard constraints):
|
||||
- Each scene becomes ONE independent clip. The model has NO memory between
|
||||
clips: every scene prompt must be fully self-contained and re-describe the
|
||||
entire setting, subjects, lighting, palette, mood, and style — even when
|
||||
nothing changed from the previous scene. Anything you omit will vanish or
|
||||
mutate between scenes.
|
||||
- Each scene prompt must be under {target_chars} characters. This is a hard
|
||||
limit; prefer cutting adjectives over cutting subjects or setting.
|
||||
- Each scene has a duration in seconds, between {min_seconds} and
|
||||
{max_seconds}; the rules below say how to choose it.
|
||||
- The model renders picture AND sound, including clear spoken language.
|
||||
When the idea involves someone speaking, write the dialogue out
|
||||
explicitly and unambiguously — name who speaks and give the exact words
|
||||
in quotes (e.g. the fisherman shouts "It's alive!") — and describe the
|
||||
voice's tone. Do not paraphrase speech the viewer asked for.
|
||||
- End each scene prompt with one short clause of soundscape (ambience,
|
||||
music mood, or effects) alongside any dialogue.
|
||||
- Describe only what the camera sees and the microphone hears: no text
|
||||
overlays, no UI, no scene numbers, no camera jargon the model cannot show.
|
||||
|
||||
{scene_count_rules}
|
||||
|
||||
WRITING THE SCENE PROMPTS:
|
||||
- Be concrete and visual: subject, action, setting, camera angle and motion,
|
||||
lighting, color palette, atmosphere, then the soundscape clause.
|
||||
- Strong nouns and verbs over piles of adjectives; vivid but precise.
|
||||
- Keep the viewer's idea recognizable — enhance it, do not replace it. The
|
||||
idea has already passed moderation before it reaches you; your job is
|
||||
faithful staging, not policing.
|
||||
|
||||
Reply with ONLY this JSON, nothing else:
|
||||
{{"title": "short display title for the sequence",
|
||||
"scenes": [{{"prompt": "self-contained scene description...", "seconds": 8.0}}]}}
|
||||
The "scenes" array is REQUIRED even when it holds a single scene; never
|
||||
flatten a scene's fields to the top level.
|
||||
"""
|
||||
|
||||
_MULTI_SCENE_RULES = """\
|
||||
HOW MANY SCENES, AND HOW LONG — two shapes; pick what the idea calls for:
|
||||
- ONE SCENE: a single clip that ALWAYS runs the full {max_seconds} seconds —
|
||||
never shorter — with room for the scene to build, land, and breathe.
|
||||
Right for a mood, a place, a single action or gag. When in doubt, this.
|
||||
- CHUNKED SHORT STORY: 3 to {max_chunks} chunks that read as one story with
|
||||
a setup, a development, and a payoff. Content chunks run 8-{max_seconds}
|
||||
seconds; the short end ({min_seconds}-8 s) is ONLY for transitions — an
|
||||
establishing cut, a reaction beat, a snap punchline — never for a chunk
|
||||
that carries the story. Choose this shape when the idea implies
|
||||
narrative: a journey, a transformation, a chase, a build-up.
|
||||
- Never more than {max_chunks} scenes. Do not pad a thin idea into many
|
||||
chunks; a story earns its chunks or it is one full-length scene.
|
||||
- Consecutive scenes play back-to-back as one sequence. Make them feel
|
||||
continuous: repeat the shared setting and subjects verbatim enough that
|
||||
they read as the same place, and change only what the story moves."""
|
||||
|
||||
_SINGLE_SCENE_RULES = """\
|
||||
HOW MANY SCENES, AND HOW LONG:
|
||||
- Exactly one scene, and it ALWAYS runs the full {max_seconds} seconds.
|
||||
Distill the idea into one complete arc that fills that time."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Scene:
|
||||
"""One upsampled scene: a prompt fast-h3 can take verbatim, and a length."""
|
||||
|
||||
prompt: str
|
||||
seconds: float
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SceneGroup:
|
||||
"""The scenes one prompt expanded into, played back-to-back.
|
||||
|
||||
``generated`` marks filler groups made from the idle prompt list rather
|
||||
than a viewer request; the director may evict their clips from the
|
||||
model's queue to make room for viewer groups.
|
||||
"""
|
||||
|
||||
group_id: str
|
||||
title: str
|
||||
author: str
|
||||
source: str
|
||||
raw_prompt: str
|
||||
scenes: list[Scene]
|
||||
generated: bool = False
|
||||
|
||||
|
||||
class PromptUpsampler:
|
||||
"""Expand chat ideas into styled, self-contained fast-h3 scenes."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str,
|
||||
model: str,
|
||||
style: str,
|
||||
max_chunks: int,
|
||||
base_url: str | None = None,
|
||||
free_viewer_style: bool = True,
|
||||
) -> None:
|
||||
self._client = AsyncOpenAI(api_key=api_key, base_url=base_url)
|
||||
self._model = model
|
||||
self._style = style.strip() or "Cinematic, photoreal, rich natural light."
|
||||
# Filler keeps the preset identity; viewer requests may pick their own.
|
||||
self._free_viewer_style = free_viewer_style
|
||||
self._max_chunks = max_chunks
|
||||
|
||||
async def upsample(
|
||||
self,
|
||||
raw_prompt: str,
|
||||
author: str,
|
||||
source: str,
|
||||
min_seconds: float,
|
||||
max_seconds: float,
|
||||
generated: bool = False,
|
||||
max_chunks: int | None = None,
|
||||
) -> SceneGroup:
|
||||
"""One idea in, one validated scene group out. Never raises.
|
||||
|
||||
`min_seconds`/`max_seconds` are the live bounds from the model's
|
||||
`state_update`, so the LLM always chooses within what the deployment
|
||||
actually accepts. `max_chunks` caps this call below the configured
|
||||
ceiling (the idle filler passes 1 so its groups stay one-clip and
|
||||
evictable). On any LLM failure the raw prompt (styled, truncated)
|
||||
becomes a single scene — the stream keeps moving.
|
||||
"""
|
||||
chunk_cap = min(max_chunks or self._max_chunks, self._max_chunks)
|
||||
scene_count_rules = (_MULTI_SCENE_RULES.format(
|
||||
max_chunks=chunk_cap,
|
||||
min_seconds=f"{min_seconds:g}",
|
||||
max_seconds=f"{max_seconds:g}",
|
||||
) if chunk_cap > 1 else _SINGLE_SCENE_RULES.format(max_seconds=f"{max_seconds:g}"))
|
||||
system = _SYSTEM_PROMPT.format(
|
||||
style=(_FREE_STYLE if self._free_viewer_style and not generated else self._style),
|
||||
target_chars=_TARGET_PROMPT_CHARS,
|
||||
min_seconds=f"{min_seconds:g}",
|
||||
max_seconds=f"{max_seconds:g}",
|
||||
scene_count_rules=scene_count_rules,
|
||||
)
|
||||
group_id = uuid.uuid4().hex[:12]
|
||||
title = ""
|
||||
scenes: list[Scene] = []
|
||||
for attempt in range(1, _MAX_ATTEMPTS + 1):
|
||||
try:
|
||||
title, scenes = await self._attempt(
|
||||
system=system,
|
||||
raw_prompt=raw_prompt,
|
||||
request_tag=f"{group_id}.{attempt}",
|
||||
chunk_cap=chunk_cap,
|
||||
min_seconds=min_seconds,
|
||||
max_seconds=max_seconds,
|
||||
)
|
||||
break
|
||||
except Exception as error:
|
||||
logger.warning(
|
||||
"[upsampler] unusable reply, attempt %d/%d for %.60r: %s",
|
||||
attempt,
|
||||
_MAX_ATTEMPTS,
|
||||
raw_prompt,
|
||||
error,
|
||||
)
|
||||
if not scenes:
|
||||
logger.warning(
|
||||
"[upsampler] all %d attempts unusable; falling back to the raw prompt",
|
||||
_MAX_ATTEMPTS,
|
||||
)
|
||||
title = raw_prompt[:60]
|
||||
# The viewer's idea gets the char budget first; the style fills
|
||||
# whatever remains (a long STYLE must never truncate the idea away).
|
||||
idea = _sanitize(raw_prompt)
|
||||
style_room = MAX_PROMPT_CHARS - len(idea) - 2
|
||||
fallback = f"{idea}. {self._style[:style_room]}" if style_room > 20 else idea
|
||||
scenes = [
|
||||
# A single clip, so it takes the maximum length like every
|
||||
# other one-scene generation.
|
||||
Scene(prompt=_sanitize(fallback), seconds=max_seconds)
|
||||
]
|
||||
|
||||
group = SceneGroup(
|
||||
group_id=group_id,
|
||||
title=title,
|
||||
author=author,
|
||||
source=source,
|
||||
raw_prompt=raw_prompt,
|
||||
scenes=scenes,
|
||||
generated=generated,
|
||||
)
|
||||
for index, scene in enumerate(group.scenes, start=1):
|
||||
logger.info(
|
||||
"[upsampler] %s scene %d/%d (%.1fs): %.100s...",
|
||||
group_id,
|
||||
index,
|
||||
len(group.scenes),
|
||||
scene.seconds,
|
||||
scene.prompt,
|
||||
)
|
||||
return group
|
||||
|
||||
async def _attempt(
|
||||
self,
|
||||
*,
|
||||
system: str,
|
||||
raw_prompt: str,
|
||||
request_tag: str,
|
||||
chunk_cap: int,
|
||||
min_seconds: float,
|
||||
max_seconds: float,
|
||||
) -> tuple[str, list[Scene]]:
|
||||
"""One LLM call, parsed and validated; raises on an unusable reply.
|
||||
|
||||
The request tag makes every attempt a distinct request — the gateway
|
||||
caches identical ones, so a bare retry of a failed prompt would get
|
||||
the same failed reply back in milliseconds.
|
||||
"""
|
||||
response = await self._client.chat.completions.create(
|
||||
model=self._model,
|
||||
messages=[
|
||||
{
|
||||
"role": "system",
|
||||
"content": system
|
||||
},
|
||||
{
|
||||
"role": "user",
|
||||
"content": f"Viewer idea: {raw_prompt}\n\n[request {request_tag}]",
|
||||
},
|
||||
],
|
||||
temperature=0.8,
|
||||
max_tokens=1800,
|
||||
response_format={"type": "json_object"},
|
||||
)
|
||||
content = response.choices[0].message.content or ""
|
||||
data = json.loads(content or "{}")
|
||||
title = str(data.get("title") or raw_prompt[:60]).strip()
|
||||
raw_scenes = data.get("scenes")
|
||||
if isinstance(raw_scenes, dict):
|
||||
raw_scenes = [raw_scenes]
|
||||
if not raw_scenes and "prompt" in data:
|
||||
# Some models flatten a single scene's fields to the top level
|
||||
# despite the schema; accept it as one scene.
|
||||
raw_scenes = [data]
|
||||
scenes = self._validate_scenes(raw_scenes or [], chunk_cap, min_seconds, max_seconds)
|
||||
if not scenes:
|
||||
raise ValueError("no usable scenes in the reply "
|
||||
f"(finish={response.choices[0].finish_reason}, head={content[:200]!r})")
|
||||
if len(scenes) == 1:
|
||||
# A single-clip generation always runs the maximum length; short
|
||||
# clips are reserved for transition chunks in stories.
|
||||
scenes = [Scene(prompt=scenes[0].prompt, seconds=max_seconds)]
|
||||
return title, scenes
|
||||
|
||||
def _validate_scenes(self, raw_scenes: list, chunk_cap: int, min_seconds: float, max_seconds: float) -> list[Scene]:
|
||||
"""Enforce every constraint the LLM was asked for; trust nothing."""
|
||||
scenes: list[Scene] = []
|
||||
for raw in raw_scenes[:chunk_cap]:
|
||||
if not isinstance(raw, dict):
|
||||
continue
|
||||
prompt = _sanitize(str(raw.get("prompt", "")))
|
||||
if not prompt:
|
||||
continue
|
||||
try:
|
||||
seconds = float(raw.get("seconds", 8.0))
|
||||
except (TypeError, ValueError):
|
||||
seconds = 8.0
|
||||
scenes.append(Scene(prompt=prompt, seconds=_clamp(seconds, min_seconds, max_seconds)))
|
||||
return scenes
|
||||
|
||||
|
||||
def _sanitize(prompt: str) -> str:
|
||||
"""Collapse whitespace and fit under fast-h3's prompt cap, ending clean.
|
||||
|
||||
LLMs overshoot the character target they are given, and a blind cut at
|
||||
the cap ends the prompt mid-word — worse for the model than losing the
|
||||
final sentence. Over-long prompts are therefore cut at the last sentence
|
||||
boundary that fits; the mid-word cut remains
|
||||
only as the last resort for a prompt written as one giant sentence.
|
||||
"""
|
||||
collapsed = " ".join(prompt.split())
|
||||
if len(collapsed) <= MAX_PROMPT_CHARS:
|
||||
return collapsed.strip()
|
||||
head = collapsed[:MAX_PROMPT_CHARS]
|
||||
boundary = max(head.rfind(". "), head.rfind("! "), head.rfind("? "))
|
||||
if boundary > MAX_PROMPT_CHARS // 2:
|
||||
return head[:boundary + 1].strip()
|
||||
return head.strip()
|
||||
|
||||
|
||||
def _clamp(value: float, low: float, high: float) -> float:
|
||||
return max(low, min(high, value))
|
||||
@@ -0,0 +1,6 @@
|
||||
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
|
||||
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 691 B |
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user