Compare commits
16
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b3a9874fc8 | ||
|
|
d657cbbf17 | ||
|
|
48534ef4de | ||
|
|
1116f514be | ||
|
|
d451e61749 | ||
|
|
3ff4a8d2d2 | ||
|
|
9343d4cdf4 | ||
|
|
66fb3d1e79 | ||
|
|
aca850cef2 | ||
|
|
1c79779956 | ||
|
|
eee03527ed | ||
|
|
1eb8541094 | ||
|
|
69c214d13a | ||
|
|
0341481aa7 | ||
|
|
d1c3fdd187 | ||
|
|
980e8d933e |
@@ -119,7 +119,7 @@ FastVideo-WorldModel/
|
||||
## Build & Test Commands
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[dev]" # Editable install
|
||||
uv pip install -e .[dev] # Editable install
|
||||
pre-commit run --all-files # Lint/format/spell
|
||||
pytest tests/ # Top-level tests
|
||||
pytest fastvideo/tests/ -v # Package tests
|
||||
|
||||
@@ -6,5 +6,3 @@
|
||||
{"name": "index-related-work", "description": "Ingest a paper or repository into the related work index", "path": "index-related-work/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "search-related-work", "description": "Query the related work index for relevant papers, repos, or comparisons", "path": "search-related-work/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "seed-ssim-references", "description": "Run a new or updated fastvideo/tests/ssim/ test on Modal, pull generated videos, and upload them to FastVideo/ssim-reference-videos so the test has a regression baseline", "path": "seed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "reseed-ssim-references", "description": "Re-seed (overwrite) HF reference videos for an existing fastvideo/tests/ssim/ test and a single model id on Modal L40S. Always backs up current refs first, regenerates on Modal, pauses for the user to eyeball before-vs-after, then uploads with --force scoped to --model-id. Sister skill to seed-ssim-references; use when intentional code change has invalidated existing refs", "path": "reseed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
|
||||
{"name": "reseed-performance-baseline", "description": "Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, or environment-caused benchmark shift. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline must be advanced by replicating one reviewed shifted source result into three success=true records, or five records when explicitly requested", "path": "reseed-performance-baseline/SKILL.md", "status": "draft", "trust": "low"}
|
||||
|
||||
@@ -12,7 +12,7 @@ automates the boilerplate of setting environment variables, picking the right
|
||||
entrypoint, and applying defaults from the closest example script.
|
||||
|
||||
## Prerequisites
|
||||
- The repo is cloned and `fastvideo` is installed (`uv pip install -e ".[dev]"`).
|
||||
- The repo is cloned and `fastvideo` is installed (`uv pip install -e .[dev]`).
|
||||
- Dataset is preprocessed (see `docs/training/data_preprocess.md`).
|
||||
- `WANDB_API_KEY` is set in the environment (or `WANDB_MODE=offline` for local).
|
||||
- GPU resources are available (multi-GPU requires NCCL).
|
||||
|
||||
@@ -1,426 +0,0 @@
|
||||
---
|
||||
name: reseed-performance-baseline
|
||||
description: Re-seed the HF performance-tracking baseline for an intentional runtime, dependency, or environment-caused benchmark shift. Use when performance CI fails because metrics such as latency, throughput, component time, or peak memory changed for an accepted reason and the rolling median baseline in FastVideo/performance-tracking must be advanced by replicating one reviewed shifted source result into three success=true records, or five records when explicitly requested.
|
||||
---
|
||||
|
||||
# Re-seed Performance Baseline
|
||||
|
||||
## Purpose
|
||||
|
||||
Replace or advance the rolling performance baseline for a single
|
||||
`(model_id, gpu_type)` pair in the HF dataset
|
||||
`FastVideo/performance-tracking`.
|
||||
|
||||
Performance comparison uses the median of up to the last 5 successful records
|
||||
for the same model and GPU. Failed records are useful audit history, but they
|
||||
do not move the future baseline because `compare_baseline.py` loads records
|
||||
with `successful_only=True`.
|
||||
|
||||
For a 5-record median, one shifted record is not enough to move the median if
|
||||
the other four records are from the old runtime. This skill therefore creates
|
||||
3 reviewed `success=true` records from one accepted shifted source result by
|
||||
default. If the user explicitly asks for a full reset, create 5 records.
|
||||
|
||||
These replicated records are an intentional operator-approved baseline reset,
|
||||
not independent measurements. Mark them clearly with provenance fields so the
|
||||
HF history remains auditable.
|
||||
|
||||
Use this skill when a performance test fails for an intentional and reviewed
|
||||
reason, such as a torch/runtime/container upgrade that legitimately increases
|
||||
peak memory or changes timings. This is the performance equivalent of
|
||||
`reseed-ssim-references`: backup first, scope tightly, require explicit human
|
||||
approval, then upload reviewed accepted baseline records.
|
||||
|
||||
## When to use
|
||||
|
||||
- A PR or main run failed the rolling performance comparison by more than the
|
||||
allowed regression threshold, and maintainers agree the shift is caused by
|
||||
an intentional runtime, dependency, hardware image, or benchmark environment
|
||||
change rather than a FastVideo logic regression.
|
||||
- One shifted source result has been reviewed and accepted, and the operator
|
||||
wants to replicate it into 3 successful records so the rolling median moves
|
||||
immediately. Use 5 records only when the user explicitly asks to fully reset
|
||||
the last-5 window.
|
||||
|
||||
## When not to use
|
||||
|
||||
- The benchmark failure might be a real code regression. Fix or investigate
|
||||
the code path first.
|
||||
- The fixed benchmark thresholds in
|
||||
`.buildkite/performance-benchmarks/tests/*.json` are too low. Those are a
|
||||
separate gate from the rolling HF baseline and may need a code review change.
|
||||
- There is no clear source run, commit, and rationale. Baseline history is a
|
||||
production signal; do not edit it without provenance.
|
||||
|
||||
## Inputs
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|-----------|----------|-------------|
|
||||
| `model_id` | Yes | Benchmark id, e.g. `wan-t2v-1.3b-2gpu`. This maps to the HF subdirectory after `sanitize(model_id)`. |
|
||||
| `gpu_type` | Yes | Exact GPU device string from the performance record, e.g. the L40S device name emitted by CI. Baselines are GPU-specific. |
|
||||
| `source_result` | Yes | Path or Buildkite artifact URL for one accepted shifted performance JSON. Prefer the normalized `normalized_perf_*.json` artifact emitted by `compare_baseline.py`. |
|
||||
| `replica_count` | No | Number of success records to create from `source_result`. Default: `3`. Only use `5` if the user explicitly asks for a full reset. |
|
||||
| `intent_rationale` | Yes | One-line explanation for why the baseline shift is legitimate. This is written into provenance and should be reused in the PR. |
|
||||
|
||||
Hardcoded defaults:
|
||||
|
||||
- HF repo: `FastVideo/performance-tracking` (`HF_REPO_ID` override is
|
||||
supported by the code, but use the default unless the user explicitly asks).
|
||||
- Local sync root: `/tmp/perf-tracking` or a timestamped local backup under
|
||||
`performance_reseed_backup/`.
|
||||
- Baseline window: last 5 `success=true` records for the same
|
||||
`(model_id, gpu_type)`.
|
||||
- Default reseed count: 3 replicated `success=true` records from one reviewed
|
||||
source result. Explicit full-reset count: 5.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Validate the target and source result
|
||||
|
||||
If `source_result` is a Buildkite artifact URL, download it first into a
|
||||
local scratch directory such as `performance_reseed_source/` and use that
|
||||
downloaded JSON path for the rest of the workflow. If the agent cannot access
|
||||
the artifact because Buildkite authentication is missing, ask the user to
|
||||
download the artifact manually and provide the local path.
|
||||
|
||||
Prefer the normalized Buildkite artifact emitted by `compare_baseline.py`:
|
||||
|
||||
```text
|
||||
perf_reports/results/normalized_perf_*.json
|
||||
```
|
||||
|
||||
That file is already in the HF tracking schema. Load it directly and confirm
|
||||
it has the expected baseline fields:
|
||||
|
||||
```python
|
||||
import json
|
||||
|
||||
with open(source_result, encoding="utf-8") as f:
|
||||
record = json.load(f)
|
||||
```
|
||||
|
||||
If only the older raw `fastvideo/tests/performance/results/perf_*.json`
|
||||
artifact is available, normalize it with `compare_baseline.py`'s shared helper
|
||||
before continuing. Run this from the repository root with
|
||||
`PYTHONPATH=fastvideo/tests/performance` so the script-local `hf_store` import
|
||||
resolves the same way it does in CI:
|
||||
|
||||
```python
|
||||
import json
|
||||
from compare_baseline import normalize_performance_result
|
||||
|
||||
with open(source_result, encoding="utf-8") as f:
|
||||
record = normalize_performance_result(json.load(f))
|
||||
```
|
||||
|
||||
The raw-to-normalized helper maps:
|
||||
|
||||
- `model_id` comes from `benchmark_id`.
|
||||
- `gpu_type` comes from `device`.
|
||||
- `memory` comes from `max_peak_memory_mb`.
|
||||
- `latency` comes from `avg_generation_time_s`.
|
||||
- `throughput` comes from `throughput_fps`.
|
||||
- component timings come from the raw `text_encoder_time_s`, `dit_time_s`,
|
||||
and `vae_decode_time_s` fields when present. If an older raw artifact lacks
|
||||
those keys, they normalize to `None`; that source can still reseed latency,
|
||||
throughput, and memory, but it cannot move component-time baselines.
|
||||
|
||||
Stop if the normalized record's `model_id` or `gpu_type` does not match the
|
||||
requested `model_id` and `gpu_type`.
|
||||
|
||||
The source record may have `success: false` when it came from a failed rolling
|
||||
baseline comparison. That is expected; only the reviewed reseed replicas become
|
||||
new `success: true` baseline records after explicit approval.
|
||||
|
||||
Set `replica_count` to `3` by default. Set it to `5` only when the user
|
||||
explicitly asks to upload the same shifted source result 5 times for a full
|
||||
last-5 reset. Reject other counts unless the user gives a concrete reason.
|
||||
|
||||
Check that `HF_API_KEY` is exported. The sync path may be public, but the
|
||||
upload path requires write access.
|
||||
|
||||
### 1a. How to obtain `source_result` from CI
|
||||
|
||||
The performance CI exports normalized source results for failed rolling
|
||||
baseline comparisons when `compare_baseline.py` ran. The preferred artifact
|
||||
comes from:
|
||||
|
||||
```text
|
||||
perf_reports/results/normalized_perf_*.json
|
||||
```
|
||||
|
||||
and is uploaded by Buildkite with the performance reports. The normal operator
|
||||
flow is:
|
||||
|
||||
1. Open the failed Buildkite performance job.
|
||||
2. Download the `normalized_perf_*.json` artifact for the failed benchmark.
|
||||
3. Pass the local path or artifact URL as `source_result`.
|
||||
|
||||
Do not scrape the Markdown performance summary to reconstruct the JSON. The
|
||||
normalized JSON artifact is the source of truth for reseed metrics and
|
||||
provenance. If only a raw `fastvideo/tests/performance/results/perf_*.json`
|
||||
artifact is present, normalize it with `normalize_performance_result()` before
|
||||
continuing. If no JSON artifact is present, the benchmark likely failed before
|
||||
writing results, so that run is not a valid source for baseline reseeding.
|
||||
|
||||
### 2. Sync and back up existing HF records
|
||||
|
||||
Use `fastvideo/tests/performance/hf_store.py` helpers directly. Do **not** use
|
||||
`compare_baseline.py` as a sync shortcut; on full main runs it can persist
|
||||
records, while this step must only fetch and back up existing history.
|
||||
|
||||
The sync command pattern is:
|
||||
|
||||
```bash
|
||||
export PERFORMANCE_TRACKING_ROOT="${PERFORMANCE_TRACKING_ROOT:-/tmp/perf-tracking}"
|
||||
export HF_REPO_ID="${HF_REPO_ID:-FastVideo/performance-tracking}"
|
||||
PYTHONPATH=fastvideo/tests/performance python -c 'from hf_store import sync_from_hf; import os; sync_from_hf(os.environ["PERFORMANCE_TRACKING_ROOT"], strict=True)'
|
||||
```
|
||||
|
||||
Then back up only the sanitized model directory:
|
||||
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(python - <<'PY'
|
||||
from fastvideo.tests.performance.hf_store import sanitize
|
||||
print(sanitize("<model_id>"))
|
||||
PY
|
||||
)
|
||||
BACKUP_DIR="performance_reseed_backup/${TIMESTAMP}_${SHORT_COMMIT}_${MODEL_SAFE}"
|
||||
mkdir -p "$BACKUP_DIR"
|
||||
cp -R "${PERFORMANCE_TRACKING_ROOT}/${MODEL_SAFE}" "$BACKUP_DIR/" 2>/dev/null || true
|
||||
```
|
||||
|
||||
Write provenance next to the backup:
|
||||
|
||||
```bash
|
||||
cat > "$BACKUP_DIR/PROVENANCE.txt" <<EOF
|
||||
model_id: <model_id>
|
||||
gpu_type: <gpu_type>
|
||||
source_result: <source_result>
|
||||
replica_count: <3_or_5>
|
||||
head_commit: $(git rev-parse HEAD)
|
||||
timestamp_utc: $(date -u +%FT%TZ)
|
||||
reason: <intent_rationale>
|
||||
EOF
|
||||
```
|
||||
|
||||
If the backup has no prior records, this is not a destructive reseed; it is a
|
||||
first baseline seed. Continue, but report that baseline history was empty.
|
||||
|
||||
### 3. Compute old baseline and candidate shift
|
||||
|
||||
Load the last 5 successful records for the target:
|
||||
|
||||
```python
|
||||
from fastvideo.tests.performance.hf_store import load_records_for_model
|
||||
|
||||
records = load_records_for_model(
|
||||
"/tmp/perf-tracking",
|
||||
"<model_id>",
|
||||
"<gpu_type>",
|
||||
last_n=5,
|
||||
successful_only=True,
|
||||
)
|
||||
```
|
||||
|
||||
Print a small table showing the source result metrics, the replicated
|
||||
candidate median, and the old medians for:
|
||||
|
||||
- `latency`
|
||||
- `throughput`
|
||||
- `memory`
|
||||
- `text_encoder_time_s`
|
||||
- `dit_time_s`
|
||||
- `vae_decode_time_s`
|
||||
|
||||
Also print how many successful old records exist. Make clear:
|
||||
|
||||
- 1 shifted record only seeds audit history and usually does not move the
|
||||
median.
|
||||
- 3 replicated shifted records in a 5-record window move the median
|
||||
immediately.
|
||||
- 5 replicated shifted records fully reset the rolling window to the source
|
||||
result's runtime profile.
|
||||
- Replicated records are not independent measurements; they are an intentional
|
||||
approved baseline reset and must be labeled that way.
|
||||
|
||||
### 4. Confirm intent
|
||||
|
||||
Require an explicit confirmation phrase before preparing the upload:
|
||||
|
||||
> About to RE-SEED performance baseline for `<model_id>` on `<gpu_type>`.
|
||||
> This will upload `<N>` new `success=true` records to
|
||||
> `FastVideo/performance-tracking/<sanitize(model_id)>/`.
|
||||
>
|
||||
> Reason: `<intent_rationale>`
|
||||
> Source result: `<source_result>`
|
||||
> Replica count: `<replica_count>`
|
||||
> Note: these records replicate one reviewed measurement to force the rolling
|
||||
> median to the accepted runtime profile.
|
||||
> HEAD: `<git rev-parse --short=12 HEAD>`
|
||||
> Backup: `<BACKUP_DIR>`
|
||||
>
|
||||
> Reply `confirm performance reseed` to proceed, anything else to abort.
|
||||
|
||||
Do not continue unless the user types exactly `confirm performance reseed`.
|
||||
|
||||
### 5. Create the accepted seed records
|
||||
|
||||
Create `replica_count` normalized records from the single source result. Use
|
||||
an explicit allowlist; do not copy the raw result JSON wholesale.
|
||||
|
||||
Each record must include only these baseline fields plus the reseed provenance
|
||||
fields below:
|
||||
|
||||
- `model_id`
|
||||
- `timestamp`
|
||||
- `commit_sha`
|
||||
- `gpu_type`
|
||||
- `latency`
|
||||
- `throughput`
|
||||
- `memory`
|
||||
- `text_encoder_time_s`
|
||||
- `dit_time_s`
|
||||
- `vae_decode_time_s`
|
||||
- `success: true`
|
||||
|
||||
For normalized `normalized_perf_*.json` sources, these fields already exist.
|
||||
For older raw `perf_*.json` sources, map the raw fields exactly as
|
||||
`normalize_performance_result()` in `compare_baseline.py` does:
|
||||
|
||||
| Normalized field | Raw source field |
|
||||
|------------------|------------------|
|
||||
| `model_id` | `benchmark_id` |
|
||||
| `gpu_type` | `device` |
|
||||
| `latency` | `avg_generation_time_s` |
|
||||
| `throughput` | `throughput_fps` |
|
||||
| `memory` | `max_peak_memory_mb` |
|
||||
| `text_encoder_time_s` | `text_encoder_time_s` |
|
||||
| `dit_time_s` | `dit_time_s` |
|
||||
| `vae_decode_time_s` | `vae_decode_time_s` |
|
||||
| `commit_sha` | `commit` |
|
||||
|
||||
Do not upload raw-only fields such as `model_short_name`, `num_gpus`,
|
||||
`num_warmup_runs`, `num_measurement_runs`, `individual_times_s`,
|
||||
`individual_peak_memories_mb`, `thresholds`, or `pr_number`.
|
||||
|
||||
Optional provenance fields are allowed and useful:
|
||||
|
||||
- `baseline_reseed: true`
|
||||
- `baseline_reseed_reason`
|
||||
- `baseline_reseed_source_result`
|
||||
- `baseline_reseed_source_timestamp`
|
||||
- `baseline_reseed_replicated_source: true`
|
||||
- `baseline_reseed_batch_size`
|
||||
- `baseline_reseed_batch_index`
|
||||
- `baseline_reseed_operator`
|
||||
|
||||
Use a fresh reseed timestamp for each replicated record, not the original
|
||||
source result timestamp. This is required because
|
||||
`load_records_for_model(..., last_n=5)` keeps the last records after loading
|
||||
the model directory; stale filenames/timestamps may not enter the last-5
|
||||
window and therefore may not move the median. Preserve the original source
|
||||
timestamp in `baseline_reseed_source_timestamp`.
|
||||
|
||||
Use the existing filename convention from `_write_tracking_record()`:
|
||||
`<sanitize(timestamp)>_<sanitize(commit_sha)>.json` under the sanitized model
|
||||
directory, but include a deterministic suffix such as `_reseed_01`,
|
||||
`_reseed_02`, and `_reseed_03` before `.json` so the replicated files do not
|
||||
overwrite each other. For a 5-record full reset, continue through
|
||||
`_reseed_05`.
|
||||
|
||||
If the source record already exists on HF with `success=false`, do not edit it
|
||||
in place unless the user explicitly asked for an audit-preserving correction.
|
||||
Prefer uploading new accepted seed records so failed history remains visible.
|
||||
|
||||
### 6. Pause before upload
|
||||
|
||||
Print:
|
||||
|
||||
- Backup directory path.
|
||||
- HF paths that will receive the new records.
|
||||
- Old rolling medians.
|
||||
- Source metrics, replica count, and candidate median.
|
||||
- Rationale.
|
||||
|
||||
Ask the user to reply exactly `upload`. Anything else aborts and leaves the
|
||||
prepared records plus backup on disk.
|
||||
|
||||
### 7. Upload only the scoped records
|
||||
|
||||
Use the shared storage helper so the path and repo type match CI:
|
||||
|
||||
```python
|
||||
from fastvideo.tests.performance.hf_store import upload_record
|
||||
|
||||
upload_record("<local_record_path>", record, strict=True)
|
||||
```
|
||||
|
||||
Run it once per prepared record. Each upload goes to:
|
||||
|
||||
```text
|
||||
FastVideo/performance-tracking/<sanitize(model_id)>/<record_filename>.json
|
||||
```
|
||||
|
||||
Never bulk upload the whole tracking root. Never modify another model's
|
||||
directory in the same operation.
|
||||
|
||||
### 8. Report outcome
|
||||
|
||||
Report:
|
||||
|
||||
- Uploaded HF paths.
|
||||
- Backup directory.
|
||||
- Old baseline window count and medians.
|
||||
- Source metrics, replica count, and candidate median.
|
||||
- Expected effect: 3 replicated shifted records move the 5-record median; 5
|
||||
replicated shifted records fully reset the window to the accepted source
|
||||
result.
|
||||
- Any separate threshold changes still needed in
|
||||
`.buildkite/performance-benchmarks/tests/*.json`.
|
||||
|
||||
Include the `intent_rationale` in the PR or follow-up comment so reviewers can
|
||||
distinguish an accepted baseline shift from a hidden regression.
|
||||
|
||||
## Failure modes and handling
|
||||
|
||||
- **`HF_API_KEY` unset.** Stop before upload. Do not create an untracked
|
||||
process that appears to have reseeded but never reached HF.
|
||||
- **Source result does not match target.** Stop. The wrong benchmark or GPU
|
||||
would poison a separate baseline.
|
||||
- **`replica_count` is 5 but the user did not explicitly ask for a full
|
||||
reset.** Stop and use the default count of 3.
|
||||
- **The source result is noisy or suspicious.** Stop. Replicating one result
|
||||
amplifies that measurement into the baseline, so it must be reviewed first.
|
||||
- **HF sync fails.** Stop for destructive reseeds. A stale or empty sync can
|
||||
make the old baseline look missing.
|
||||
- **Candidate still violates fixed thresholds.** Report that this skill only
|
||||
handles the rolling HF baseline; update benchmark JSON thresholds in code
|
||||
review if maintainers accept the new absolute limit.
|
||||
- **The user aborts at either confirmation.** Leave the backup and prepared
|
||||
records on disk. Nothing should be uploaded.
|
||||
- **A bad seed was uploaded.** Use the backup and HF history to identify the
|
||||
uploaded file, then remove or supersede it with an explicitly reviewed
|
||||
corrective record. Do not silently rewrite unrelated history.
|
||||
|
||||
## References
|
||||
|
||||
- `.agents/skills/reseed-ssim-references/SKILL.md` — safety pattern for
|
||||
intentional baseline replacement.
|
||||
- `fastvideo/tests/performance/compare_baseline.py` — normalization, rolling
|
||||
median comparison, and persistence rules.
|
||||
- `fastvideo/tests/performance/hf_store.py` — HF sync, record loading,
|
||||
`sanitize()`, and `upload_record()`.
|
||||
- `fastvideo/tests/performance/test_inference_performance.py` — source result
|
||||
JSON schema.
|
||||
- `.buildkite/performance-benchmarks/tests/*.json` — fixed absolute benchmark
|
||||
thresholds, separate from rolling baseline comparisons.
|
||||
|
||||
## Changelog
|
||||
|
||||
| Date | Change |
|
||||
|------|--------|
|
||||
| 2026-05-03 | Initial version. Sister workflow to `reseed-ssim-references`, scoped to one performance `(model_id, gpu_type)` baseline seed with backup, confirmation, provenance, and `success=true` upload. |
|
||||
| 2026-05-03 | Current policy: replicate one approved shifted source result into 3 success records by default, or 5 only when explicitly requested. Add provenance marker for replicated-source reseeds. |
|
||||
@@ -1,343 +0,0 @@
|
||||
---
|
||||
name: reseed-ssim-references
|
||||
description: Re-seed HF reference videos for a single existing SSIM test on Modal L40S. Always backs up current refs locally first, regenerates on Modal, pauses for the user to eyeball before-vs-after quality, then overwrites the targeted `<model_id>` subtree on `FastVideo/ssim-reference-videos` with `--force`. Use when an intentional code change (model port fix, attention backend swap, kernel upgrade, hyperparameter change) has invalidated existing refs and they need to be regenerated. Pairs with `seed-ssim-references`, which is for first-time seeding only.
|
||||
---
|
||||
|
||||
# Re-seed SSIM Reference Videos
|
||||
|
||||
## Purpose
|
||||
|
||||
Replace the existing SSIM reference videos for a single `(test_file, model_id)`
|
||||
pair on the HF dataset (`FastVideo/ssim-reference-videos`). This is **destructive**
|
||||
on HF — the old refs are overwritten — so the skill always:
|
||||
|
||||
1. Confirms intent with a one-liner the user has to type.
|
||||
2. Downloads the existing refs as a local, timestamped backup.
|
||||
3. Regenerates on Modal L40S (same code path that CI uses).
|
||||
4. Pauses for a side-by-side eyeball of backup vs new mp4s.
|
||||
5. Uploads with `--force`, scoped to the single `--model-id`.
|
||||
6. Reminds the user to keep the backup until the PR lands.
|
||||
|
||||
Pairs with `seed-ssim-references`, which is the inverse (first-time seeding
|
||||
only, refuses to overwrite). Re-seeding is intentionally a separate, more
|
||||
ceremonial operation because mistakenly clobbering production refs is much
|
||||
harder to recover from than failing closed.
|
||||
|
||||
## When to use
|
||||
|
||||
- An intentional code change (model port fix, kernel upgrade, attention
|
||||
backend swap, hyperparameter change in the test itself) has shifted the
|
||||
expected SSIM output and the existing refs no longer represent the new
|
||||
ground truth.
|
||||
- A test is failing in CI **for the right reason** (the new code is correct,
|
||||
the old refs are stale).
|
||||
|
||||
## When not to use
|
||||
|
||||
- A test is failing for the **wrong** reason (the port is buggy, not the
|
||||
refs). Fix the port; re-seeding hides the bug.
|
||||
- A brand-new test that has no refs on HF yet. Use `seed-ssim-references`.
|
||||
- "Just to clean up drift" without a concrete code change to point at. The
|
||||
PR description has to justify *why* refs changed; without a concrete
|
||||
change, there's nothing to write.
|
||||
|
||||
## Inputs
|
||||
|
||||
| Parameter | Required | Description |
|
||||
|-----------|----------|-------------|
|
||||
| `test_file` | Yes | Path to the SSIM test, e.g. `fastvideo/tests/ssim/test_matrixgame_similarity.py`. Validated against `fastvideo/tests/ssim/test_*_similarity.py`. |
|
||||
| `model_id` | Yes | Single model id from the test's `*_MODEL_TO_PARAMS`, e.g. `Matrix-Game-2.0-Diffusers-Base`. Re-seed runs are **per model**. For multi-model tests, invoke the skill once per model. |
|
||||
| `intent_rationale` | Yes | One-line explanation of *why* refs are being regenerated (e.g. "Relax FA-2 head_size whitelist to include 80 — matrix_game now uses FLASH_ATTN instead of TORCH_SDPA"). Recorded in the backup directory and reused in the PR description. |
|
||||
|
||||
Hardcoded:
|
||||
|
||||
- Modal GPU: **L40S** (matches CI; re-seeding from another SKU produces refs
|
||||
that L40S CI cannot match).
|
||||
- Quality tier: **`default`**. `full_quality` is a separate, deliberate
|
||||
operation.
|
||||
- HF repo: `FastVideo/ssim-reference-videos` (override via
|
||||
`FASTVIDEO_SSIM_REFERENCE_HF_REPO`).
|
||||
- Device folder: `L40S_reference_videos`.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
The user has confirmed:
|
||||
|
||||
- `modal` CLI authenticated.
|
||||
- `hf` CLI authenticated, **and** `HF_API_KEY` (or `HUGGINGFACE_HUB_TOKEN` /
|
||||
`HF_TOKEN`) exported with **write** access to
|
||||
`FastVideo/ssim-reference-videos`.
|
||||
- The current branch's code is the change that motivated the re-seed (i.e.
|
||||
`git rev-parse HEAD` is the commit that intentionally invalidated refs).
|
||||
|
||||
Fail fast if any of these are missing.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Validate inputs and confirm intent
|
||||
|
||||
- Verify `test_file` exists and matches `fastvideo/tests/ssim/test_*_similarity.py`.
|
||||
- Grep the file for `*_MODEL_TO_PARAMS` and assert `model_id` is one of its
|
||||
keys. If the file has only a single hardcoded model, accept that model id
|
||||
as the only valid value.
|
||||
- Print the rationale and ask the user to type **`confirm reseed`** (not just
|
||||
`y` — make it deliberate):
|
||||
|
||||
> About to RE-SEED references for model `<model_id>` from test `<test_file>`.
|
||||
> This will OVERWRITE existing refs on
|
||||
> `FastVideo/ssim-reference-videos/reference_videos/default/L40S_reference_videos/<model_id>/`
|
||||
> after backup + Modal regen + eyeball.
|
||||
>
|
||||
> Reason: `<intent_rationale>`
|
||||
> HEAD: `<git rev-parse --short=12 HEAD>`
|
||||
>
|
||||
> Reply `confirm reseed` to proceed, anything else to abort.
|
||||
|
||||
Stop until the user types exactly `confirm reseed`. Anything else aborts
|
||||
with no side effects.
|
||||
|
||||
### 2. Back up existing refs
|
||||
|
||||
Always required. The backup is the only graceful path back if anything goes
|
||||
wrong later.
|
||||
|
||||
```bash
|
||||
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
|
||||
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
MODEL_SAFE=$(echo "<model_id>" | tr '/' '_')
|
||||
BACKUP_DIR="ssim_reseed_backup/${TIMESTAMP}_${SHORT_COMMIT}_${MODEL_SAFE}"
|
||||
mkdir -p "$BACKUP_DIR"
|
||||
|
||||
hf download \
|
||||
--repo-type dataset FastVideo/ssim-reference-videos \
|
||||
--include "reference_videos/default/L40S_reference_videos/<model_id>/**" \
|
||||
--local-dir "$BACKUP_DIR"
|
||||
|
||||
mp4_count=$(find "$BACKUP_DIR" -name "*.mp4" | wc -l)
|
||||
echo "Backup mp4 count: $mp4_count"
|
||||
[ "$mp4_count" -gt 0 ] || {
|
||||
echo "ERROR: backup is empty for <model_id>. Either the model id is wrong"
|
||||
echo "or there are no existing refs (use seed-ssim-references instead)."
|
||||
exit 1
|
||||
}
|
||||
|
||||
# Provenance — used in the PR description
|
||||
cat > "$BACKUP_DIR/PROVENANCE.txt" <<EOF
|
||||
test_file: <test_file>
|
||||
model_id: <model_id>
|
||||
head_commit: $(git rev-parse HEAD)
|
||||
timestamp_utc: $(date -u +%FT%TZ)
|
||||
reason: <intent_rationale>
|
||||
EOF
|
||||
```
|
||||
|
||||
If the `hf download` produces zero mp4s, abort — the user has either picked a
|
||||
non-existent `model_id` or there are no refs yet (in which case
|
||||
`seed-ssim-references` is the right tool).
|
||||
|
||||
### 3. Regenerate on Modal L40S
|
||||
|
||||
Mirror CI's exact env recipe so the regenerated refs are byte-comparable to
|
||||
what CI will produce on the same commit. Two differences from CI:
|
||||
|
||||
1. **Pass the same env prefix CI uses** (`IMAGE_VERSION`, `BUILDKITE_*`) — see
|
||||
`.buildkite/pipeline.yml:1-3` and `.buildkite/scripts/pr_test.sh:62-83`.
|
||||
Without this, `ssim_test.py:17-18` resolves a different GHCR image tag
|
||||
(default is `latest`, CI is `py3.12-latest`), and `ssim_test.py:38-46`
|
||||
bakes different values into the image's frozen env block. **Mismatched
|
||||
image or env is the most common source of SSIM drift between reseed and
|
||||
CI runs.**
|
||||
2. **Do not pass `--skip-reference-download`**. Letting the test fetch the
|
||||
existing refs and run the full SSIM compare gives "before" SSIM numbers
|
||||
for the PR description, and the test still produces the new mp4s
|
||||
regardless of whether the comparison passes or fails.
|
||||
|
||||
```bash
|
||||
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
|
||||
|
||||
IMAGE_VERSION="py3.12-latest" \
|
||||
BUILDKITE_REPO="$(git config --get remote.origin.url)" \
|
||||
BUILDKITE_COMMIT="$(git rev-parse HEAD)" \
|
||||
BUILDKITE_PULL_REQUEST="${BUILDKITE_PULL_REQUEST:-false}" \
|
||||
modal run fastvideo/tests/modal/ssim_test.py \
|
||||
--git-repo="$(git config --get remote.origin.url)" \
|
||||
--git-commit="$(git rev-parse HEAD)" \
|
||||
--hf-api-key="$HF_API_KEY" \
|
||||
--test-files="<test_file>" \
|
||||
--sync-generated-to-volume \
|
||||
--generated-volume-subdir="$SUBDIR" \
|
||||
--no-fail-fast
|
||||
```
|
||||
|
||||
Capture the printed `modal volume get ...` hint — its `<SUBDIR>` matches
|
||||
`$SUBDIR` and is needed for step 4. Capture the SSIM numbers from the test
|
||||
output (or from the JSON next to the generated mp4) for the PR description.
|
||||
|
||||
### 4. Download generated videos
|
||||
|
||||
```bash
|
||||
modal volume get --force hf-model-weights \
|
||||
ssim_generated_videos/default/"$SUBDIR"/generated_videos \
|
||||
./generated_videos_modal/default
|
||||
```
|
||||
|
||||
After this, the new mp4s live at:
|
||||
|
||||
```
|
||||
./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4
|
||||
```
|
||||
|
||||
`--force` is required when `./generated_videos_modal/default` already exists
|
||||
from a prior run; safe on the first run too.
|
||||
|
||||
### 5. PAUSE — user reviews quality side-by-side
|
||||
|
||||
Print the diff and the comparison:
|
||||
|
||||
```bash
|
||||
echo "=== File list diff (backup vs new) ==="
|
||||
diff -u \
|
||||
<(find "$BACKUP_DIR/reference_videos/default/L40S_reference_videos/<model_id>" -name "*.mp4" \
|
||||
| sed "s|$BACKUP_DIR/reference_videos/default/L40S_reference_videos/||" | sort) \
|
||||
<(find ./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id> -name "*.mp4" \
|
||||
| sed "s|./generated_videos_modal/default/generated_videos/L40S_reference_videos/||" | sort) \
|
||||
|| true
|
||||
|
||||
echo
|
||||
echo "=== SSIM numbers from this run (paste into PR) ==="
|
||||
find ./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id> -name "*_ssim.json" -exec cat {} \;
|
||||
```
|
||||
|
||||
Then stop and tell the user:
|
||||
|
||||
> Old refs backed up to `$BACKUP_DIR`.
|
||||
> New videos in `./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/`.
|
||||
>
|
||||
> Open both in a video player. Confirm the new videos:
|
||||
> 1. Look correct (no obvious artifacts, no black/static frames).
|
||||
> 2. Are *intentionally* different from the backup in the way described
|
||||
> in `<intent_rationale>` (e.g. slight numerical drift only, not a
|
||||
> different scene / different motion / corrupted output).
|
||||
>
|
||||
> Reply **`upload`** to overwrite HF, anything else to abort.
|
||||
> Aborting leaves the backup and new videos on disk for inspection — nothing
|
||||
> on HF changes.
|
||||
|
||||
Do not proceed until the user types exactly `upload`. If they abort, leave
|
||||
everything on disk and stop here.
|
||||
|
||||
### 6. Copy into the local reference layout
|
||||
|
||||
Same as `seed-ssim-references` step 5:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--generated-dir ./generated_videos_modal/default/generated_videos/L40S_reference_videos
|
||||
```
|
||||
|
||||
Result: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
|
||||
### 7. Upload with `--force`, scoped to `--model-id`
|
||||
|
||||
The `--force` flag is what makes this skill different from `seed-ssim-references`.
|
||||
Always pair it with `--model-id` so a typo cannot accidentally overwrite a
|
||||
neighboring model's refs.
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py upload \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id "<model_id>" \
|
||||
--force
|
||||
```
|
||||
|
||||
The CLI's overwrite guard refuses without `--force`; with `--force` it
|
||||
overwrites only files under
|
||||
`reference_videos/default/L40S_reference_videos/<model_id>/`.
|
||||
|
||||
### 8. Report success and retention guidance
|
||||
|
||||
Print:
|
||||
|
||||
- The HF path that was overwritten (`<repo>/reference_videos/default/L40S_reference_videos/<model_id>/`).
|
||||
- The local backup directory path.
|
||||
- The new SSIM numbers from step 5.
|
||||
- This restore command, in case the PR review surfaces a problem after
|
||||
upload:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py upload \
|
||||
--quality-tier default \
|
||||
--device-folder L40S_reference_videos \
|
||||
--model-id "<model_id>" \
|
||||
--reference-dir "$BACKUP_DIR/reference_videos/default/L40S_reference_videos" \
|
||||
--force
|
||||
```
|
||||
|
||||
- This PR-description checklist (see `fastvideo/tests/ssim/AGENTS.md` →
|
||||
*Updating Reference Videos*):
|
||||
1. Source commit that produced the new refs (HEAD at re-seed time).
|
||||
2. Test command and GPU SKU (`L40S`).
|
||||
3. Before/after SSIM numbers.
|
||||
4. The `<intent_rationale>` from step 1.
|
||||
5. A note that the backup lives at `$BACKUP_DIR` and should be retained
|
||||
until CI on the PR is green.
|
||||
|
||||
Do **not** auto-rerun the SSIM test — the user does that as part of the PR.
|
||||
|
||||
## Failure modes and how to handle them
|
||||
|
||||
- **`HF_API_KEY` unset.** Stop before step 2.
|
||||
- **Backup is empty (zero mp4s).** Stop before step 3 — the model id is
|
||||
wrong or the refs don't exist yet (use `seed-ssim-references`).
|
||||
- **Modal run fails before generation.** No mp4s on the volume. Don't
|
||||
upload. Investigate the failure (test crash, OOM, partition exhaustion),
|
||||
fix, then retry from step 3. Backup is still intact.
|
||||
- **Quality regressed (visual or metric).** User aborts at step 5. Backup
|
||||
retained. New videos retained on disk for inspection. Nothing on HF
|
||||
changed. Either fix the underlying code change or abandon the re-seed.
|
||||
- **User confirmed `upload` but later realized the new refs are wrong.**
|
||||
Run the restore command from step 8 with the backup `--reference-dir`.
|
||||
This is exactly why the backup exists.
|
||||
- **Multi-model test, only one model is being re-seeded.** Run the skill
|
||||
once per model id. The `--model-id` scope on upload guarantees the others
|
||||
are untouched.
|
||||
|
||||
## Design notes (for future skill maintainers)
|
||||
|
||||
- Per-`model_id` scope is mandatory. The dataset houses many model subtrees;
|
||||
re-seeding the wrong one is hard to undo without backup.
|
||||
- `default` tier only; `full_quality` is a separate, deliberate operation
|
||||
with different params and ~doubled runtime, and isn't what CI gates on.
|
||||
- The skill deliberately does **not** pass `--skip-reference-download` to
|
||||
Modal so we get pre-reseed SSIM numbers for the PR. The `seed`-skill
|
||||
passes it because no refs exist yet; for re-seed, refs do exist and
|
||||
exposing the comparison is informative.
|
||||
- The two-token confirm (`confirm reseed`, then `upload`) is intentional.
|
||||
Re-seeding is high-blast-radius and should not be one-keystroke.
|
||||
- The backup directory is plain mp4s + `PROVENANCE.txt`. No HF metadata is
|
||||
preserved; the restore path uses `reference_videos_cli.py upload
|
||||
--reference-dir` which doesn't need it.
|
||||
|
||||
## References
|
||||
|
||||
- `.agents/skills/seed-ssim-references/SKILL.md` — the first-time seed
|
||||
skill this one parallels. Read it for the Modal flag rationale shared
|
||||
between the two flows.
|
||||
- `fastvideo/tests/ssim/AGENTS.md` — directory rules, including the PR
|
||||
expectations for any reference-video change (rationale, before/after
|
||||
SSIM, source commit/model/backend).
|
||||
- `fastvideo/tests/ssim/reference_videos_cli.py` — `copy-local`, `upload`
|
||||
(with `--model-id`, `--force`), `download`. The overwrite guard at
|
||||
`upload_reference_videos` is the safety net this skill leans on.
|
||||
- `fastvideo/tests/modal/ssim_test.py` — Modal orchestrator;
|
||||
`--sync-generated-to-volume`, `--generated-volume-subdir`,
|
||||
`--skip-reference-download`, `--no-fail-fast`.
|
||||
|
||||
## Changelog
|
||||
|
||||
| Date | Change |
|
||||
|------|--------|
|
||||
| 2026-05-02 | Initial version. Sister skill to `seed-ssim-references`, scoped to single `(test_file, model_id)` re-seeds, with mandatory backup and two-token confirm. |
|
||||
@@ -1,41 +1,26 @@
|
||||
---
|
||||
name: seed-ssim-references
|
||||
description: Seed HF reference artefacts for a single newly-added SSIM test (pixel `.mp4` for `run_text_to_video_similarity_test`-style tests, or latent `.pt` for `run_text_to_latent_similarity_test`-style tests). Runs the test on Modal L40S, downloads the generated artefacts via `modal volume get`, pauses for the user to verify (visual eyeball for mp4, numerics dump for pt), then uploads only that test's files to `FastVideo/ssim-reference-videos`. Use when a new `fastvideo/tests/ssim/test_*_similarity.py` has just been added and has no references on HF yet.
|
||||
description: Seed HF reference videos for a single newly-added SSIM test. Runs the test on Modal L40S, downloads the generated mp4s via `modal volume get`, pauses for the user to eyeball quality, then uploads only that test's files to `FastVideo/ssim-reference-videos`. Use when a new `fastvideo/tests/ssim/test_*_similarity.py` has just been added and has no references on HF yet.
|
||||
---
|
||||
|
||||
# Seed SSIM Reference Artefacts (mp4 or pt)
|
||||
# Seed SSIM Reference Videos
|
||||
|
||||
## Purpose
|
||||
|
||||
A brand-new SSIM test in `fastvideo/tests/ssim/` fails forever until its
|
||||
reference artefacts exist on the HF dataset
|
||||
(`FastVideo/ssim-reference-videos`). The dataset hosts two kinds of artefacts
|
||||
side-by-side per `(model_id, backend, prompt)`:
|
||||
|
||||
- **`.mp4`** — pixel ground-truth for tests that call
|
||||
`run_text_to_video_similarity_test` / `run_image_to_video_similarity_test`
|
||||
in `inference_similarity_utils.py`. Compared via SSIM.
|
||||
- **`.pt`** — pre-VAE latent bundle (fp16 full latent + fp32 slice +
|
||||
metadata + `slice_spec` + `format_version`) for tests that call
|
||||
`run_text_to_latent_similarity_test` in `latent_similarity_utils.py`.
|
||||
Compared via cosine distance on the slice and the full tensor.
|
||||
|
||||
reference videos exist on the HF dataset (`FastVideo/ssim-reference-videos`).
|
||||
This skill:
|
||||
|
||||
1. Detects which artefact type the test produces (pixel vs latent).
|
||||
2. Runs the test on Modal's L40S pool to generate the artefacts.
|
||||
3. Downloads them to the local repo via `modal volume get`.
|
||||
4. Pauses so the user can verify quality:
|
||||
- **mp4**: visual eyeball in a video player.
|
||||
- **pt**: numerics dump (shape, slice stats, NaN/Inf check, metadata).
|
||||
5. Uploads only the new test's files to HF, with a guard that refuses to
|
||||
1. Runs the test on Modal's L40S pool to generate the videos.
|
||||
2. Downloads them to the local repo via `modal volume get`.
|
||||
3. Pauses so the user can eyeball the mp4s and confirm quality.
|
||||
4. Uploads only the new test's files to HF, with a guard that refuses to
|
||||
overwrite anything already present.
|
||||
|
||||
The skill is run **manually**, once per new test. Before invoking it, the user
|
||||
has already sanity-tested the new test locally — it launches `VideoGenerator`
|
||||
and writes an artefact without crashing (the missing-reference assertion at
|
||||
the end is expected). The skill does not re-test locally; it goes straight
|
||||
to Modal L40S (which is what CI uses).
|
||||
and writes an mp4 without crashing. The skill does not re-test locally; it
|
||||
goes straight to Modal L40S (which is what CI uses).
|
||||
|
||||
## When to use
|
||||
|
||||
@@ -84,7 +69,7 @@ Fail fast if the token env var is missing.
|
||||
|
||||
## Steps
|
||||
|
||||
### 1. Ask for the test file, then detect artefact type
|
||||
### 1. Ask for the test file
|
||||
|
||||
If the user didn't name one, ask: *"Which SSIM test file do you want to seed
|
||||
references for? (e.g. `fastvideo/tests/ssim/test_ltx2_similarity.py`)"*.
|
||||
@@ -95,22 +80,6 @@ Validate:
|
||||
- File defines a `*_MODEL_TO_PARAMS` dict — grep it to extract the set of
|
||||
model ids. Those ids drive step 5.
|
||||
|
||||
Detect artefact type by inspecting the file's imports / helper call:
|
||||
|
||||
- **latent** (`.pt`) — file imports `run_text_to_latent_similarity_test`
|
||||
from `fastvideo.tests.ssim.latent_similarity_utils` (or any other helper
|
||||
that ends with `_latent_similarity_test`).
|
||||
- **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_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`,
|
||||
`copy_generated_to_reference`, and `upload_reference_videos` already walk
|
||||
both `.mp4` and `.pt` (see `reference_videos_cli.py`).
|
||||
|
||||
If either check fails, stop and tell the user what's wrong.
|
||||
|
||||
### 2. Run the test on Modal L40S
|
||||
@@ -123,19 +92,9 @@ TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
|
||||
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
|
||||
```
|
||||
|
||||
Then launch the Modal run. The `IMAGE_VERSION` and `BUILDKITE_*` env-prefix
|
||||
**must** match what CI exports in `.buildkite/scripts/pr_test.sh`, otherwise
|
||||
`fastvideo/tests/modal/ssim_test.py` resolves a different GHCR image tag
|
||||
(default is `latest`, CI is `py3.12-latest`) and bakes different values into
|
||||
the image's frozen env block (`ssim_test.py:17-18, 38-46`). Mismatched image
|
||||
or env produces SSIM drift that doesn't show up until the same commit runs
|
||||
in CI.
|
||||
Then launch the Modal run:
|
||||
|
||||
```bash
|
||||
IMAGE_VERSION="py3.12-latest" \
|
||||
BUILDKITE_REPO="$(git config --get remote.origin.url)" \
|
||||
BUILDKITE_COMMIT="$(git rev-parse HEAD)" \
|
||||
BUILDKITE_PULL_REQUEST="${BUILDKITE_PULL_REQUEST:-false}" \
|
||||
modal run fastvideo/tests/modal/ssim_test.py \
|
||||
--git-repo="$(git config --get remote.origin.url)" \
|
||||
--git-commit="$(git rev-parse HEAD)" \
|
||||
@@ -147,19 +106,6 @@ modal run fastvideo/tests/modal/ssim_test.py \
|
||||
--no-fail-fast
|
||||
```
|
||||
|
||||
Env prefix rationale (parity with CI; see `.buildkite/pipeline.yml:1-3` and
|
||||
`.buildkite/scripts/pr_test.sh:62-83`):
|
||||
- `IMAGE_VERSION=py3.12-latest`: pins the Modal image tag to the same one CI
|
||||
uses. Without this, `ssim_test.py:17` falls back to `latest`, which on
|
||||
GHCR is built from `Dockerfile.python3.10` — different Python, torch, and
|
||||
flash-attn wheel than CI's `py3.12-latest` (`infra-build-image.yml:51-67`,
|
||||
`_template-build-image.yml:65-101`).
|
||||
- `BUILDKITE_REPO`/`BUILDKITE_COMMIT`/`BUILDKITE_PULL_REQUEST`: mirror what
|
||||
Buildkite exports. `ssim_test.py:38-46` bakes these into the image's
|
||||
`.env(...)` block; mismatched values can perturb in-container code paths
|
||||
that branch on PR-vs-non-PR. `false` for `BUILDKITE_PULL_REQUEST` matches
|
||||
Buildkite's "non-PR build" sentinel.
|
||||
|
||||
Flag rationale:
|
||||
- `--skip-reference-download`: no refs exist yet, so conftest must not try to
|
||||
pull them.
|
||||
@@ -197,59 +143,17 @@ get` preserves that trailing `generated_videos/` segment.
|
||||
|
||||
### 4. PAUSE — user reviews quality
|
||||
|
||||
Type-aware verification.
|
||||
|
||||
**For `ARTEFACT_TYPE = pixel`** — list the downloaded mp4s and ask the user to
|
||||
open them in a video player:
|
||||
Print the list of downloaded mp4s and their paths, then stop. Tell the user:
|
||||
|
||||
> "Generated videos downloaded to `./generated_videos_modal/default/generated_videos/L40S_reference_videos/`. Please open them and confirm the quality looks correct. Reply **`upload`** to continue, or anything else to abort."
|
||||
|
||||
**For `ARTEFACT_TYPE = latent`** — `.pt` files are not human-watchable. Print
|
||||
a numerics dump for each `.pt` so the user can sanity-check shape, distribution,
|
||||
and metadata:
|
||||
|
||||
```python
|
||||
import torch
|
||||
from pathlib import Path
|
||||
ROOT = Path("./generated_videos_modal/default/generated_videos/L40S_reference_videos")
|
||||
for p in sorted(ROOT.rglob("*.pt")):
|
||||
d = torch.load(p, map_location="cpu", weights_only=False)
|
||||
s = d["expected_slice"]
|
||||
L = d["latent"].float()
|
||||
print(f"=== {p.relative_to(ROOT)} ===")
|
||||
print(f" format_version: {d['format_version']}")
|
||||
print(f" shape: {d['shape']}")
|
||||
print(f" dtype_original: {d['dtype_original']}")
|
||||
print(f" slice_spec: {d['slice_spec']}")
|
||||
print(f" slice shape={tuple(s.shape)} mean={s.mean():+.4f} std={s.std():.4f} min={s.min():+.4f} max={s.max():+.4f}")
|
||||
print(f" latent shape={tuple(L.shape)} mean={L.mean():+.4f} std={L.std():.4f} min={L.min():+.4f} max={L.max():+.4f}")
|
||||
print(f" finite: latent NaN={torch.isnan(L).any().item()} Inf={torch.isinf(L).any().item()}; "
|
||||
f"slice NaN={torch.isnan(s).any().item()} Inf={torch.isinf(s).any().item()}")
|
||||
print(f" metadata: {d['metadata']}\n")
|
||||
```
|
||||
|
||||
Sanity criteria:
|
||||
- `format_version == 1` (matches `LATENT_REFERENCE_FORMAT_VERSION`).
|
||||
- `shape` matches what the model produces (e.g. LTX-2 distilled =
|
||||
`[1, 128, T_lat, H_lat, W_lat]`; Stable Audio Open 1.0 = `[1, 64, 1024]`).
|
||||
- `slice_spec.kind` matches a registered kind (`corner_3x3_first_frame`
|
||||
for video, `audio_first_8_timesteps` for audio).
|
||||
- No `NaN`/`Inf`. `mean ≈ 0`, `std ≈ 1` (denoised latents stay close to
|
||||
the initial Gaussian distribution; very wide deviations suggest
|
||||
numerical drift).
|
||||
- `metadata.prompt` matches the test's prompt.
|
||||
|
||||
Then ask:
|
||||
|
||||
> "Numerics look right? Reply **`upload`** to continue, or anything else to abort."
|
||||
|
||||
Do not proceed until the user explicitly says `upload`. If they abort, leave
|
||||
everything on disk so they can inspect further — no cleanup.
|
||||
|
||||
### 5. Copy into the local reference layout
|
||||
|
||||
Scoped copy — only the new test's artefacts. Single command works for both
|
||||
artefact types because `_iter_reference_files` walks `.mp4` and `.pt`:
|
||||
Scoped copy — only the new test's mp4s. Loop over each `<model_id>` extracted
|
||||
in step 1:
|
||||
|
||||
```bash
|
||||
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
@@ -259,13 +163,12 @@ python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
|
||||
```
|
||||
|
||||
(The `--generated-dir` points at the device-folder root inside the
|
||||
downloaded tree; `copy-local` walks all `<model>/<backend>/*.{mp4,pt}`
|
||||
downloaded tree; `copy-local` walks all `<model>/<backend>/*.mp4`
|
||||
underneath it. Since the Modal run was scoped to a single test file via
|
||||
`--test-files`, only that test's model(s) are present — so the copy is
|
||||
implicitly per-test.)
|
||||
|
||||
Result for pixel: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
Result for latent: same path with `.pt` extension.
|
||||
Result: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
|
||||
|
||||
### 6. Upload to HF — scoped per model_id, with overwrite guard
|
||||
|
||||
@@ -298,54 +201,33 @@ it will auto-download the refs they just uploaded.
|
||||
## Failure modes and how to handle them
|
||||
|
||||
- **`HF_API_KEY` unset.** Stop before step 2. The Modal run needs it (passed
|
||||
via `--hf-api-key`), and step 6 needs it for upload. If the user
|
||||
ran `hf auth login` instead of exporting an env var, read the cached
|
||||
token via `huggingface_hub.get_token()` and forward it to Modal as
|
||||
`--hf-api-key="$CACHED_TOKEN"`.
|
||||
- **Modal run fails before generation.** No artefacts on the volume — nothing
|
||||
to download. Fix the test locally (`pytest fastvideo/tests/ssim/<test_file>`)
|
||||
via `--hf-api-key`), and step 6 needs it for upload.
|
||||
- **Modal run fails before generation.** No mp4s on the volume — nothing to
|
||||
download. Fix the test locally (`pytest fastvideo/tests/ssim/<test_file>`)
|
||||
and retry from step 2.
|
||||
- **`./generated_videos_modal/default/L40S_reference_videos/` missing after
|
||||
`modal volume get`.** The run didn't produce artefacts (most likely the
|
||||
test crashed before writing, or `REQUIRED_GPUS` exceeded the partition
|
||||
capacity — see Modal logs).
|
||||
- **Latent test crashed with FSDP / inference_mode error
|
||||
(`RuntimeError: Inference tensors do not track version counter`).** The
|
||||
test must pass `init_kwargs_override={"use_fsdp_inference": False}` when
|
||||
`sp_size == 1` — see `test_stable_audio_similarity.py` for the pattern.
|
||||
Fix in the test, push, retry.
|
||||
`modal volume get`.** The run didn't produce videos (most likely the test
|
||||
crashed before writing, or `REQUIRED_GPUS` exceeded the partition capacity
|
||||
— see Modal logs).
|
||||
- **Upload guard fires (files already exist).** The test name / model id
|
||||
collides with something already on HF. Verify the user actually wants to
|
||||
replace existing refs; if so, re-run the upload with `--force`. If not,
|
||||
rename the model id in `*_MODEL_TO_PARAMS` and re-seed.
|
||||
- **Quality looks wrong in step 4.** Abort. The artefacts stay on disk for
|
||||
- **Quality looks wrong in step 4.** Abort. The mp4s stay on disk for
|
||||
inspection. The fix is usually in the test's params (resolution, steps,
|
||||
seed) — edit the test, then re-run the skill.
|
||||
- For latent: also check `slice_spec.kind` matches the latent rank
|
||||
(`corner_3x3_first_frame` requires 5-D, `audio_first_8_timesteps`
|
||||
requires 3-D); a rank/kind mismatch raises in `_extract_expected_slice`.
|
||||
|
||||
## Design notes (for future skill maintainers)
|
||||
|
||||
- The skill deliberately runs on Modal, **not** locally, because the CI
|
||||
runner is L40S. Seeding from a different GPU SKU produces refs that CI's
|
||||
L40S runs can't match (pixel SSIM drifts across SKUs; latent cosine has
|
||||
tighter cross-SKU bf16 drift but the configured tolerances assume
|
||||
same-SKU seed → same-SKU verify).
|
||||
L40S runs can't match (SSIM drifts across SKUs).
|
||||
- The skill is default-tier only. `full_quality` refs are seeded by a
|
||||
separate, deliberate operation — they double runtime and aren't what CI
|
||||
gates on.
|
||||
- The overwrite guard in `reference_videos_cli.py upload` is default-on
|
||||
specifically because this skill exists. Re-seeding is a distinct operation
|
||||
that requires explicit `--force`.
|
||||
- Both artefact types share the same Modal flow: the orchestrator sets
|
||||
`--skip-reference-download` + `--no-fail-fast`, runs pytest, the test's
|
||||
helper writes the artefact (`.mp4` via `imageio` for pixel,
|
||||
`save_latent_reference` → `torch.save` for latent) BEFORE the
|
||||
missing-reference assertion raises. `_sync_generated_videos_to_volume` in
|
||||
`ssim_test.py` does a `shutil.copytree` of the whole `generated_videos/`
|
||||
tree, picking up `.mp4`, `.pt`, and the `*_ssim.json` / `*_latent.json`
|
||||
metric files alongside.
|
||||
|
||||
## References
|
||||
|
||||
@@ -354,17 +236,10 @@ it will auto-download the refs they just uploaded.
|
||||
`--skip-reference-download`, `--no-fail-fast`.
|
||||
- `fastvideo/tests/ssim/reference_videos_cli.py` — `copy-local`, `upload`
|
||||
(with `--model-id`, `--force`), `download`, `ensure` subcommands.
|
||||
Extension allowlist is `REFERENCE_EXTENSIONS = VIDEO_EXTENSIONS +
|
||||
LATENT_EXTENSIONS` (`.pt`).
|
||||
- `fastvideo/tests/ssim/README.md` — reference layout, HF repo conventions.
|
||||
- `fastvideo/tests/ssim/inference_similarity_utils.py` — pixel helpers
|
||||
(`run_text_to_video_similarity_test`,
|
||||
`run_image_to_video_similarity_test`, `build_init_kwargs`).
|
||||
- `fastvideo/tests/ssim/latent_similarity_utils.py` — latent helper
|
||||
(`run_text_to_latent_similarity_test`), slice spec dispatch
|
||||
(`_extract_expected_slice`), reference schema
|
||||
(`save_latent_reference` / `load_latent_reference`),
|
||||
`LATENT_REFERENCE_FORMAT_VERSION`.
|
||||
- `fastvideo/tests/ssim/inference_similarity_utils.py` —
|
||||
`run_text_to_video_similarity_test` + `_build_init_kwargs`: what each test
|
||||
config passes to `VideoGenerator.from_pretrained`.
|
||||
|
||||
## Changelog
|
||||
|
||||
@@ -373,4 +248,3 @@ it will auto-download the refs they just uploaded.
|
||||
| 2026-04-17 | Initial version (Modal sync-to-volume flow). |
|
||||
| 2026-04-21 | Rewrite: single-test scope, explicit user-review pause, per-`model_id` upload, HF overwrite guard. Dropped `scripts/seed_ssim.sh`. |
|
||||
| 2026-04-21 | Post-first-run fixes: `modal volume get` needs `--force` when parent exists; download tree has an extra `generated_videos/` level so `--generated-dir` must reflect it. |
|
||||
| 2026-05-01 | Latent (`*.pt`) artefact support: artefact-type detection in step 1, type-aware verification (visual eyeball for mp4, numerics dump for pt) in step 4, FSDP+inference_mode failure-mode added, design notes for the unified Modal flow. Triggered by PR #1253 (LTX-2 latent migration + Stable Audio latent test). |
|
||||
|
||||
@@ -15,21 +15,8 @@ log "Project root: $PROJECT_ROOT"
|
||||
# Install Modal if not available
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Modal not found, installing..."
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "uv not found, bootstrapping..."
|
||||
if ! curl -LsSf https://astral.sh/uv/install.sh | sh; then
|
||||
log "Error: Failed to bootstrap uv via astral.sh installer."
|
||||
exit 1
|
||||
fi
|
||||
export PATH="$HOME/.local/bin:$PATH"
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "Error: uv still not on PATH after bootstrap."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
# --break-system-packages preserves prior `pip install --user` semantics on PEP 668 agents.
|
||||
uv pip install --system --break-system-packages modal
|
||||
|
||||
python3 -m pip install modal
|
||||
|
||||
# Verify installation
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Error: Failed to install modal. Please install it manually."
|
||||
@@ -95,7 +82,7 @@ upload_performance_artifacts() {
|
||||
|
||||
_upload_dashboard() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "dashboard_${SHORT_SHA}_*" | head -n 1)
|
||||
target=$(find "$LOCAL_DIR" -name "dashboard_*${SHORT_SHA}*" | head -n 1)
|
||||
log "TARGET dashboard: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
@@ -109,7 +96,7 @@ upload_performance_artifacts() {
|
||||
|
||||
_upload_perf_summary() {
|
||||
local target
|
||||
target=$(find "$LOCAL_DIR" -name "perf_${SHORT_SHA}_*" | head -n 1)
|
||||
target=$(find "$LOCAL_DIR" -name "perf_*${SHORT_SHA}*" | head -n 1)
|
||||
log "TARGET perf summary: '$target'"
|
||||
|
||||
if [ -n "$target" ]; then
|
||||
@@ -121,19 +108,6 @@ upload_performance_artifacts() {
|
||||
fi
|
||||
}
|
||||
|
||||
_upload_normalized_perf_results() {
|
||||
local found=0
|
||||
while IFS= read -r -d '' target; do
|
||||
found=1
|
||||
log "Found normalized performance result: $target. Uploading to Buildkite..."
|
||||
buildkite-agent artifact upload "$target"
|
||||
done < <(find "$LOCAL_DIR" -path "*/results/normalized_perf_*.json" -print0)
|
||||
|
||||
if [ "$found" -eq 0 ]; then
|
||||
log "No normalized performance result artifacts found. This is expected when the rolling performance comparison did not run."
|
||||
fi
|
||||
}
|
||||
|
||||
_cleanup_modal_volume() {
|
||||
log "Cleaning up perf_reports/ from Modal Volume..."
|
||||
if modal volume rm hf-model-weights "perf_reports/" --recursive; then
|
||||
@@ -152,7 +126,6 @@ upload_performance_artifacts() {
|
||||
_download_reports || { _cleanup_local; return 1; }
|
||||
_upload_dashboard
|
||||
_upload_perf_summary
|
||||
_upload_normalized_perf_results
|
||||
_cleanup_modal_volume
|
||||
_cleanup_local
|
||||
}
|
||||
|
||||
@@ -13,21 +13,8 @@ log "Project root: $PROJECT_ROOT"
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "pre-commit not found, installing..."
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "uv not found, bootstrapping..."
|
||||
if ! curl -LsSf https://astral.sh/uv/install.sh | sh; then
|
||||
log "Error: Failed to bootstrap uv via astral.sh installer."
|
||||
exit 1
|
||||
fi
|
||||
export PATH="$HOME/.local/bin:$PATH"
|
||||
if ! command -v uv &> /dev/null; then
|
||||
log "Error: uv still not on PATH after bootstrap."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
# --break-system-packages preserves prior `pip install --user` semantics on PEP 668 agents.
|
||||
uv pip install --system --break-system-packages pre-commit==4.0.1
|
||||
|
||||
python3 -m pip install --user pre-commit==4.0.1
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "Error: Failed to install pre-commit."
|
||||
exit 1
|
||||
|
||||
@@ -37,11 +37,10 @@ jobs:
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install dependencies
|
||||
run: uv pip install --system -r requirements-mkdocs.txt
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-mkdocs.txt
|
||||
|
||||
- name: Setup Pages
|
||||
uses: actions/configure-pages@v4
|
||||
|
||||
@@ -56,11 +56,10 @@ jobs:
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install build dependencies
|
||||
run: uv pip install --system build twine wheel
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install build twine wheel
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
|
||||
@@ -131,13 +131,11 @@ jobs:
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
uv pip install --system typing-extensions==4.12.2
|
||||
uv pip install --system --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
pip install --upgrade pip
|
||||
pip install typing-extensions==4.12.2
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
@@ -147,20 +145,20 @@ jobs:
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
uv pip install --system setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
|
||||
pip install setuptools ninja packaging wheel triton scikit-build-core cmake build
|
||||
|
||||
cd fastvideo-kernel
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
# Release builds are produced on GPU-less runners, so force-enable TK and target Hopper.
|
||||
export TORCH_CUDA_ARCH_LIST="9.0a"
|
||||
export CMAKE_ARGS="${CMAKE_ARGS:-} -DFASTVIDEO_KERNEL_BUILD_TK=ON -DCMAKE_CUDA_ARCHITECTURES=90a"
|
||||
|
||||
|
||||
# Build standard wheel (no local version suffix) for PyPI
|
||||
python -m build --wheel --outdir dist
|
||||
|
||||
|
||||
# Fix the wheel to be manylinux compliant
|
||||
uv pip install --system auditwheel
|
||||
pip install auditwheel
|
||||
# Point auditwheel at torch libs, but do not vendor them into the wheel.
|
||||
TORCH_LIB_DIR=$(python - <<'PY'
|
||||
import os
|
||||
@@ -213,13 +211,10 @@ jobs:
|
||||
pattern: 'fastvideo_kernel-py*'
|
||||
merge-multiple: true
|
||||
|
||||
- name: Install uv
|
||||
uses: astral-sh/setup-uv@v3
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
uv pip install --system build scikit-build-core cmake ninja
|
||||
|
||||
pip install build scikit-build-core cmake ninja
|
||||
|
||||
cd fastvideo-kernel
|
||||
# We don't need full CUDA/Torch to just package the source (sdist)
|
||||
python -m build --sdist --outdir dist
|
||||
|
||||
@@ -7,13 +7,21 @@ exclude: |
|
||||
fastvideo-kernel/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
assets/prompts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
examples/.*|
|
||||
\.agents/.*|
|
||||
.github/workflows/publish-fastvideo.yml|
|
||||
.github/workflows/_template-build-image.yml
|
||||
.github/workflows/_template-build-image.yml|
|
||||
docs/source/inference/support_matrix.md
|
||||
)
|
||||
repos:
|
||||
- repo: https://github.com/google/yapf
|
||||
|
||||
@@ -11,7 +11,7 @@
|
||||
- Static assets: `assets/` (including `assets/images/`, `assets/videos/`, and `assets/prompts/`) and `comfyui/assets/`.
|
||||
|
||||
## Build, Test, and Development Commands
|
||||
- `uv pip install -e ".[dev]"`: editable install with lint/test extras.
|
||||
- `uv pip install -e .[dev]`: editable install with lint/test extras.
|
||||
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
|
||||
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
|
||||
- `pytest tests/`: run top-level test suite.
|
||||
@@ -23,8 +23,7 @@
|
||||
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
|
||||
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
|
||||
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
|
||||
- Lint via `pre-commit run --files <changed paths>` (or `pre-commit run --all-files` for a full sweep) before committing. Do not shell out to `yapf`/`ruff`/`codespell`/`mypy` directly — pre-commit chains them with the project's config and respects the `.pre-commit-config.yaml` excludes (e.g. `fastvideo/tests/` is intentionally skipped). If pre-commit reports `(no files to check)` for your paths, that exclude is deliberate — don't bypass it.
|
||||
- Target line length is 120 (configured in `pyproject.toml` for ruff, yapf, and isort).
|
||||
- Target line length is 80.
|
||||
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
|
||||
|
||||
## Testing Guidelines
|
||||
@@ -55,31 +54,3 @@ This repository is agent-friendly. Before doing any work, read:
|
||||
If you are exploring a new procedure that has no existing SOP, document your
|
||||
progress in `.agents/exploration/` and flag it for review at the end of your
|
||||
session.
|
||||
|
||||
## Per-Directory AGENTS.md
|
||||
|
||||
Local guidance lives next to the code. Read the in-scope file before editing:
|
||||
|
||||
| Directory | What it covers |
|
||||
|-----------|----------------|
|
||||
| `fastvideo/AGENTS.md` | Core package map, public API, registry-driven model dispatch |
|
||||
| `fastvideo/configs/AGENTS.md` | Arch + pipeline config dataclasses, `param_names_mapping` |
|
||||
| `fastvideo/models/AGENTS.md` | DiT / VAE / encoder / scheduler / loader layout (pre-commit excluded) |
|
||||
| `fastvideo/layers/AGENTS.md` | Tensor-parallel linear/attention layer rules for ports |
|
||||
| `fastvideo/attention/AGENTS.md` | Backend registry + env-var override |
|
||||
| `fastvideo/pipelines/AGENTS.md` | Stage ABC, `basic/<model>/`, `preprocess/`, presets |
|
||||
| `fastvideo/training/AGENTS.md` | Legacy monolithic pipelines (frozen for existing models) |
|
||||
| `fastvideo/train/AGENTS.md` | New modular trainer (methods × models × callbacks, YAML) |
|
||||
| `fastvideo/tests/AGENTS.md` | Test taxonomy, conftest, pre-commit-excluded path |
|
||||
| `fastvideo/tests/ssim/AGENTS.md` | GPU SSIM regression authoring + reference video sync |
|
||||
| `scripts/checkpoint_conversion/AGENTS.md` | Adding a converter for a new HF/official checkpoint |
|
||||
|
||||
## Critical: Two Training Stacks Coexist
|
||||
|
||||
- `fastvideo/training/` — legacy, monolithic per-model `*_training_pipeline.py` and
|
||||
`*_distillation_pipeline.py`. Still authoritative for shipped models.
|
||||
- `fastvideo/train/` — new modular framework (composable methods × models × callbacks
|
||||
driven by YAML). Preferred for new training work.
|
||||
|
||||
Pick the matching stack before editing. Do not migrate a pipeline between them
|
||||
without an explicit ask — the conventions and config surfaces differ.
|
||||
|
||||
@@ -128,7 +128,7 @@ class CLIPFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self, device: str = 'cuda', model_name: str = "openai/clip-vit-base-patch32"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError("Please install transformers: uv pip install transformers")
|
||||
raise ImportError("Please install transformers: pip install transformers")
|
||||
super().__init__(device)
|
||||
self.processor = CLIPProcessor.from_pretrained(model_name)
|
||||
self.model = CLIPModel.from_pretrained(model_name).to(self.device)
|
||||
@@ -171,7 +171,7 @@ class VideoMAEFeatureExtractor(BaseFeatureExtractor):
|
||||
|
||||
def __init__(self, device: str = 'cuda', model_name: str = "MCG-NJU/videomae-base"):
|
||||
if not TRANSFORMERS_AVAILABLE:
|
||||
raise ImportError("Please install transformers: uv pip install transformers")
|
||||
raise ImportError("Please install transformers: pip install transformers")
|
||||
super().__init__(device)
|
||||
self.model = VideoMAEModel.from_pretrained(model_name).to(self.device)
|
||||
self.model.eval()
|
||||
|
||||
@@ -57,7 +57,7 @@ class I3DFeatureExtractor(nn.Module):
|
||||
except Exception as e:
|
||||
raise RuntimeError(f"Failed to load I3D model from Hugging Face Hub. Error: {e}\n"
|
||||
f"Ensure you have internet connection and huggingface_hub installed:\n"
|
||||
f"uv pip install huggingface_hub") from e
|
||||
f"pip install huggingface_hub") from e
|
||||
|
||||
def preprocess(self, videos: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
uv pip install -q opencv-python-headless transformers huggingface_hub
|
||||
pip install -q opencv-python-headless transformers huggingface_hub
|
||||
|
||||
# 2. Run FVD script
|
||||
python benchmarks/fvd/run_fvd.py
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
#!/bin/bash
|
||||
|
||||
# 1. Install missing dependency
|
||||
uv pip install -q opencv-python-headless
|
||||
pip install -q opencv-python-headless
|
||||
+2
-2
@@ -38,10 +38,10 @@ cp -r /path/to/FastVideo/comfyui /path/to/ComfyUI/custom_nodes/FastVideo
|
||||
|
||||
#### Install dependencies:
|
||||
|
||||
Currently, the only dependency is `fastvideo`, which can be installed with `uv`.
|
||||
Currently, the only dependency is `fastvideo`, which can be installed using pip.
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
#### Install missing custom nodes:
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.10 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp310-cp310-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp310-cp310-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.11 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp311-cp311-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp311-cp311-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,15 +42,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.8.3+cu128torch2.11-cp312-cp312-linux_x86_64.whl
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.7.16/flash_attn-2.8.3+cu128torch2.10-cp312-cp312-linux_x86_64.whl
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -42,7 +42,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir ".[dev]" && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
@@ -50,7 +50,7 @@ COPY . .
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e ".[dev]" && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
@@ -43,7 +43,7 @@ COPY . .
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e ".[rocm]" && \
|
||||
uv pip install --no-cache-dir -e .[rocm] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
+1
-1
@@ -6,7 +6,7 @@ This directory contains the FastVideo documentation built with MkDocs.
|
||||
|
||||
```bash
|
||||
# Install dependencies
|
||||
uv pip install -r requirements-mkdocs.txt
|
||||
pip install -r requirements-mkdocs.txt
|
||||
|
||||
# Serve docs with live reload (recommended for development)
|
||||
mkdocs serve
|
||||
|
||||
@@ -99,7 +99,7 @@ cd /FastVideo
|
||||
**Install the package**
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[dev]"
|
||||
uv pip install -e .[dev]
|
||||
```
|
||||
|
||||
The Docker image already includes Flash Attention and most heavy dependencies, so this is fast.
|
||||
|
||||
@@ -49,7 +49,7 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
Install FastVideo in editable mode and set up hooks:
|
||||
|
||||
```bash
|
||||
uv pip install -e ".[dev]"
|
||||
uv pip install -e .[dev]
|
||||
|
||||
# Optional: FlashAttention (builds native kernels)
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
@@ -306,30 +306,6 @@ surfaces:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
num_frames_per_block:
|
||||
sources: [fastvideo.configs.pipelines.wan.MatrixGameI2V480PConfig]
|
||||
audio_channels:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
audio_end_in_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
audio_start_in_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
max_audio_duration_s:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
sample_size:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
sampling_rate:
|
||||
sources:
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioT2AConfig
|
||||
- fastvideo.configs.pipelines.stable_audio.StableAudioOpenSmallConfig
|
||||
compatibility_only:
|
||||
batch_size: "Gen3C inference-only tuning field pending typed batching design."
|
||||
gradient_checkpointing: "Gen3C inference-only compatibility field pending typed batching design."
|
||||
@@ -404,13 +380,6 @@ surfaces:
|
||||
ltx2_stg_scale_audio: request.extensions.ltx2.stg_scale_audio
|
||||
ltx2_stg_blocks_video: request.extensions.ltx2.stg_blocks_video
|
||||
ltx2_stg_blocks_audio: request.extensions.ltx2.stg_blocks_audio
|
||||
audio_start_in_s: request.extensions.stable_audio.audio_start_in_s
|
||||
audio_end_in_s: request.extensions.stable_audio.audio_end_in_s
|
||||
init_audio: request.extensions.stable_audio.init_audio
|
||||
init_audio_strength: request.extensions.stable_audio.init_audio_strength
|
||||
init_noise_level: request.extensions.stable_audio.init_noise_level
|
||||
inpaint_audio: request.extensions.stable_audio.inpaint_audio
|
||||
inpaint_mask: request.extensions.stable_audio.inpaint_mask
|
||||
internal_only:
|
||||
data_type: "Derived from the request shape and not a public input."
|
||||
|
||||
|
||||
@@ -243,7 +243,7 @@ for step in range(start_step, max_steps):
|
||||
|
||||
```bash
|
||||
# Install
|
||||
uv pip install -e ".[dev]"
|
||||
uv pip install -e .[dev]
|
||||
|
||||
# Run DMD2 distillation on Wan 2.1
|
||||
torchrun --nproc_per_node=8 -m fastvideo.train.entrypoint.train \
|
||||
|
||||
@@ -27,7 +27,7 @@ uv pip install fastvideo
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
|
||||
uv pip install fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
### From source
|
||||
@@ -41,11 +41,11 @@ uv pip install -e .
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
Alternative with Conda environment (still drives installs through `uv`):
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
pip install -e .
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
@@ -58,16 +58,14 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install FlashAttention:
|
||||
|
||||
```bash
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -89,7 +87,7 @@ uv pip install -e .
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
### Optional Dependencies
|
||||
@@ -103,7 +101,7 @@ uv pip install flash-attn --no-build-isolation -v
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install flash-attn --no-build-isolation -v
|
||||
pip install flash-attn --no-build-isolation -v
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
|
||||
@@ -57,10 +57,8 @@ uv pip install fastvideo
|
||||
|
||||
#### With Conda environment (alternative)
|
||||
|
||||
`uv` works inside an active conda env too, so prefer `uv pip` for the actual install:
|
||||
|
||||
```bash
|
||||
uv pip install fastvideo
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
@@ -82,7 +80,7 @@ uv pip install -e .
|
||||
Alternative with Conda environment:
|
||||
|
||||
```bash
|
||||
uv pip install -e .
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
- Install MoGe:
|
||||
|
||||
```bash
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
pip install git+https://github.com/microsoft/MoGe.git
|
||||
```
|
||||
|
||||
- If you hit `ImportError: libGL.so.1` (common on Ubuntu/headless nodes), you can try installing OpenCV runtime libs:
|
||||
|
||||
@@ -54,7 +54,7 @@ FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN python example.py
|
||||
We recommend always installing [Flash Attention 2](https://github.com/Dao-AILab/flash-attention):
|
||||
|
||||
```bash
|
||||
uv pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://github.com/Dao-AILab/flash-attention?tab=readme-ov-file#flashattention-3-beta-release) by compiling it from source (takes about 10 minutes for me):
|
||||
@@ -63,7 +63,7 @@ And if using a Hopper+ GPU (ie H100), installing [Flash Attention 3](https://git
|
||||
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention
|
||||
|
||||
cd hopper
|
||||
uv pip install ninja
|
||||
pip install ninja
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
@@ -98,7 +98,7 @@ To use [SageAttention](https://github.com/thu-ml/SageAttention) 2.1.1, please co
|
||||
```bash
|
||||
git clone https://github.com/thu-ml/SageAttention.git
|
||||
cd sageattention
|
||||
python setup.py install # or uv pip install -e .
|
||||
python setup.py install # or pip install -e .
|
||||
```
|
||||
|
||||
### Sage Attention 3
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.1 T2V 1.3B model using
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
uv pip install vsa
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
uv pip install vsa
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### Data-free Distillation
|
||||
|
||||
@@ -4,7 +4,7 @@ These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
uv pip install vsa
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
|
||||
@@ -7,7 +7,7 @@ and the GEN3C diffusion model.
|
||||
|
||||
Requirements:
|
||||
1. Install MoGe:
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
pip install git+https://github.com/microsoft/MoGe.git
|
||||
If you hit `ImportError: libGL.so.1`, install:
|
||||
sudo apt-get update && sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1
|
||||
2. Download and convert weights:
|
||||
|
||||
@@ -48,7 +48,7 @@ Prerequisites:
|
||||
and export your HF token in the shell:
|
||||
export HF_TOKEN=hf_...
|
||||
2. Install optional inference deps (one-time):
|
||||
uv pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
cmake_minimum_required(VERSION 3.26 FATAL_ERROR)
|
||||
project(fastvideo-kernel LANGUAGES CXX)
|
||||
|
||||
# Prefer environment variable (used by CI or uv pip install git+repo_addr) if CMake var is not explicitly set.
|
||||
# Prefer environment variable (used by CI or pip install git+repo_addr) if CMake var is not explicitly set.
|
||||
if(NOT DEFINED GPU_BACKEND AND DEFINED ENV{GPU_BACKEND})
|
||||
set(GPU_BACKEND "$ENV{GPU_BACKEND}")
|
||||
endif()
|
||||
|
||||
@@ -11,7 +11,7 @@ except ImportError:
|
||||
|
||||
def _unsupported(*args, **kwargs):
|
||||
raise ImportError(
|
||||
"flash-attn is not installed. Please install it, e.g., `uv pip install flash-attn`."
|
||||
"flash-attn is not installed. Please install it, e.g., `pip install flash-attn`."
|
||||
)
|
||||
|
||||
_flash_attn_varlen_forward = _unsupported
|
||||
|
||||
@@ -1,67 +0,0 @@
|
||||
# `fastvideo/` — Core Package
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Inference + training framework for video DiTs. Public API entry: `from fastvideo import VideoGenerator, PipelineConfig, SamplingParam`.
|
||||
|
||||
## Public Surface (`__init__.py`)
|
||||
|
||||
```python
|
||||
VideoGenerator # entrypoints/video_generator.py — high-level inference handle
|
||||
PipelineConfig # configs/pipelines/base.py — pipeline wiring dataclass
|
||||
SamplingParam # api/sampling_param.py — runtime sampling knobs
|
||||
```
|
||||
|
||||
CLI entry: `fastvideo` script → `entrypoints/cli/main.py` (subcommands: `generate`, `serve`, `bench`).
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
fastvideo/
|
||||
├── api/ # Schema + presets for the OpenAI-compatible serving layer
|
||||
├── attention/ # Backends + selector (FlashAttn / SageAttn / SDPA / VSA / VMoBA / SLA)
|
||||
├── configs/ # Per-model arch configs + per-pipeline configs (registry-driven)
|
||||
├── dataset/ # Dataloaders (pre-commit excluded — minimal lint surface)
|
||||
├── distributed/ # SP/TP groups, device communicators, init helpers
|
||||
├── entrypoints/ # cli/, openai/, streaming/, video_generator.py
|
||||
├── hooks/ # Runtime hook system for pipelines
|
||||
├── layers/ # Tensor-parallel linears + attention wrappers (port targets)
|
||||
├── models/ # DiT / VAE / encoder / scheduler / loader (pre-commit excluded)
|
||||
├── pipelines/ # basic/<model>/, preprocess/, stages/, training/
|
||||
├── platforms/ # CUDA/ROCm capability + AttentionBackendEnum
|
||||
├── third_party/ # Vendored externals (lint excluded; do not reformat)
|
||||
├── train/ # NEW modular trainer — methods × models × callbacks
|
||||
├── training/ # LEGACY monolithic *_training/distillation_pipeline.py
|
||||
├── worker/ # Multi-process / Ray executors
|
||||
├── workflow/ # Preprocessing workflow base class
|
||||
├── registry.py # Pipeline-config + model-class lookup (canonical)
|
||||
├── envs.py # Env-var declarations
|
||||
├── fastvideo_args.py# Runtime arg dataclass passed through pipelines
|
||||
└── utils.py # FlexibleArgumentParser, qualname resolver, etc.
|
||||
```
|
||||
|
||||
## Where to Look
|
||||
|
||||
| Task | Location |
|
||||
|------|----------|
|
||||
| Add a new pipeline class | `pipelines/basic/<model>/` + `configs/pipelines/<model>.py` + register in `registry.py` |
|
||||
| Add a new model component | `models/<role>/<model>.py` + `configs/models/<role>/<model>.py` |
|
||||
| Wire an existing model into a new pipeline | `pipelines/basic/<model>/presets.py` + reuse stages from `pipelines/stages/` |
|
||||
| Add a converter | `scripts/checkpoint_conversion/<model>_to_*.py` (separate dir, separate AGENTS.md) |
|
||||
| Add an attention backend | `attention/backends/<name>.py` + register in selector |
|
||||
| Add a runtime CLI flag | `fastvideo_args.py` (avoid `argparse` ad-hoc inside stages) |
|
||||
|
||||
## Conventions Specific Here
|
||||
|
||||
- `PipelineStage` subclasses (`pipelines/stages/`) own one verb each (encode, schedule, denoise, decode). Compose, don't fork.
|
||||
- Every pipeline reads from a `PipelineConfig` subclass and a `SamplingParam`. Never read raw env vars inside a stage — go through `fastvideo.envs`.
|
||||
- Logger setup: `from fastvideo.logger import init_logger; logger = init_logger(__name__)`. Do not call `logging.getLogger` directly.
|
||||
- Imports between `train/` and `training/` are **forbidden** — they are independent stacks.
|
||||
|
||||
## Pre-Commit Exclusions (do not assume linted)
|
||||
|
||||
These dirs are listed in `.pre-commit-config.yaml` `exclude`:
|
||||
|
||||
- `fastvideo/third_party/`, `fastvideo/dataset/`, `fastvideo/models/`
|
||||
|
||||
Editing files there will NOT trigger yapf/ruff/mypy/codespell. Format manually if a sibling file shows clear style; do not introduce new violations.
|
||||
@@ -1,58 +0,0 @@
|
||||
# `fastvideo/attention/` — Attention Backends
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Backend registry + selector wrapping FlashAttn / SageAttn / SageAttn3 / SDPA / VSA / VMoBA / SLA / BSA.
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
attention/
|
||||
├── __init__.py # Exports DistributedAttention, LocalAttention, get_attn_backend
|
||||
├── layer.py # DistributedAttention, DistributedAttention_VSA, LocalAttention
|
||||
├── selector.py # get_attn_backend (cached) + env-var override
|
||||
├── backends/
|
||||
│ ├── abstract.py # AttentionBackend / AttentionMetadata / AttentionMetadataBuilder
|
||||
│ ├── flash_attn.py # FA2/FA3
|
||||
│ ├── sage_attn.py # SageAttention v1
|
||||
│ ├── sage_attn3.py # SageAttention v3
|
||||
│ ├── sdpa.py # torch SDPA fallback
|
||||
│ ├── video_sparse_attn.py # VSA (paper: Video Sparse Attention)
|
||||
│ ├── vmoba.py # Video-MoBA
|
||||
│ ├── sla.py # Sliding-window (STA)
|
||||
│ └── bsa_attn.py # Block-sparse
|
||||
└── utils/
|
||||
├── flash_attn_cute.py
|
||||
└── flash_attn_no_pad.py
|
||||
```
|
||||
|
||||
## Selection Order
|
||||
|
||||
`get_attn_backend()` resolves via:
|
||||
|
||||
1. Env-var override `FASTVIDEO_ATTENTION_BACKEND` (see `STR_BACKEND_ENV_VAR` in `fastvideo/utils.py`).
|
||||
2. Per-platform default from `fastvideo/platforms/`.
|
||||
3. Heuristic fallback to SDPA.
|
||||
|
||||
The result is `@lru_cache`d. Tests that need a specific backend must use the
|
||||
`global_force_attn_backend(...)` context manager from `selector.py`, never set
|
||||
the env var mid-process.
|
||||
|
||||
## Adding a Backend
|
||||
|
||||
1. Subclass `AttentionBackend` in `backends/<name>.py`.
|
||||
2. Implement `AttentionMetadata` + `AttentionMetadataBuilder` for the new path.
|
||||
3. Register the enum value in `fastvideo/platforms/interface.py` (`AttentionBackendEnum`).
|
||||
4. Wire string → class resolution in `selector.py`.
|
||||
5. Verify the new backend works with `DistributedAttention` (sequence parallel)
|
||||
and `LocalAttention` (single-rank). If it cannot support SP, document the
|
||||
gap in the backend file's module docstring.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Calling `torch.nn.functional.scaled_dot_product_attention` directly inside a
|
||||
model's forward — go through `DistributedAttention` / `LocalAttention`.
|
||||
- Reading `os.environ[STR_BACKEND_ENV_VAR]` from arbitrary call sites. Use
|
||||
`get_env_variable_attn_backend()`.
|
||||
- Caching backend instances per-module. The selector cache is process-wide; do
|
||||
not duplicate it.
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
from dataclasses import dataclass
|
||||
|
||||
try:
|
||||
@@ -17,7 +18,6 @@ except ImportError:
|
||||
flash_attn_func = flash_attn_3_func
|
||||
fa_version = "3"
|
||||
except ImportError:
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
flash_attn_func = flash_attn_2_func
|
||||
fa_version = "2"
|
||||
|
||||
|
||||
@@ -405,7 +405,7 @@ class SageSLAAttentionImpl(AttentionImpl, nn.Module):
|
||||
|
||||
if not SAGESLA_ENABLED:
|
||||
raise ImportError("SageSLA requires spas_sage_attn. "
|
||||
"Install with: uv pip install git+https://github.com/thu-ml/SpargeAttn.git")
|
||||
"Install with: pip install git+https://github.com/thu-ml/SpargeAttn.git")
|
||||
|
||||
assert head_size in [64, 128], f"SageSLA requires head_size in [64, 128], got {head_size}"
|
||||
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
# `fastvideo/configs/` — Config-Driven Model Registry
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Two layers of dataclass configs feed every pipeline: **arch configs** (what the model is) and **pipeline configs** (how to run it).
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
configs/
|
||||
├── configs.py # Dataset / loader enums (DatasetType, VideoLoaderType)
|
||||
├── utils.py # update_config_from_args, shallow_asdict helpers
|
||||
├── backend/ # Attention backend defaults
|
||||
├── models/
|
||||
│ ├── base.py # ModelConfig ABC
|
||||
│ ├── dits/ # DiTConfig per model (wanvideo, ltx2, hunyuan, ...)
|
||||
│ ├── vaes/ # VAEConfig per model
|
||||
│ ├── encoders/ # EncoderConfig (t5, clip, llama, qwen2_5, gemma, siglip, ...)
|
||||
│ ├── upsamplers/ # UpsamplerConfig (hunyuan15)
|
||||
│ └── audio/ # Audio-model configs (ltx2_audio_vae, ...)
|
||||
├── pipelines/
|
||||
│ ├── base.py # PipelineConfig ABC + (de)serialization
|
||||
│ └── <model>.py # Concrete configs (HunyuanConfig, WanT2V480PConfig, ...)
|
||||
└── *.json # Frozen reference configs for shipped models
|
||||
```
|
||||
|
||||
## How Configs Hook Into the Registry
|
||||
|
||||
`fastvideo/registry.py` imports every concrete `PipelineConfig` and exposes
|
||||
`get_pipeline_config_cls_from_name(...)`. Adding a new pipeline config requires:
|
||||
|
||||
1. Subclass `PipelineConfig` in `pipelines/<model>.py`.
|
||||
2. Reference its component arch configs (DiT / VAE / encoder / upsampler).
|
||||
3. Add the import + name mapping in `fastvideo/registry.py`.
|
||||
|
||||
Configs that do not appear in `registry.py` are unreachable from `VideoGenerator`.
|
||||
|
||||
## Arch vs Pipeline — Where Does This Field Go?
|
||||
|
||||
| Field type | Lives on |
|
||||
|-----------|----------|
|
||||
| Architecture constants (hidden dim, num heads, layer count) | `configs/models/<role>/<model>.py` |
|
||||
| Default sampling params (steps, cfg, shift, fps) | `configs/pipelines/<model>.py` |
|
||||
| Runtime overrides (precision, sp_size, tp_size, attention backend) | `configs/pipelines/base.py` defaults + CLI flags via `fastvideo_args.py` |
|
||||
| `param_names_mapping` for HF → FastVideo state-dict | Arch config (lives with the model definition) |
|
||||
|
||||
If a knob is tunable per inference call → `SamplingParam`, not `PipelineConfig`.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Hard-coding architecture constants inside model classes — always read from the arch config.
|
||||
- Using `argparse` directly here. Configs deserialize from dicts via `update_config_from_args`.
|
||||
- Importing from `fastvideo.pipelines` here. Configs are the lower layer; the dependency is one-way.
|
||||
@@ -3,8 +3,6 @@
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.cli.generate import cmd_init as generate_cmd_init
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
from fastvideo.entrypoints.cli.router_serve import (
|
||||
cmd_init as router_serve_cmd_init, )
|
||||
from fastvideo.entrypoints.cli.serve import cmd_init as serve_cmd_init
|
||||
from fastvideo.entrypoints.cli.bench import cmd_init as bench_cmd_init
|
||||
|
||||
@@ -14,7 +12,6 @@ def cmd_init() -> list[CLISubcommand]:
|
||||
commands = []
|
||||
commands.extend(generate_cmd_init())
|
||||
commands.extend(serve_cmd_init())
|
||||
commands.extend(router_serve_cmd_init())
|
||||
commands.extend(bench_cmd_init())
|
||||
return commands
|
||||
|
||||
|
||||
@@ -1,115 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""``fastvideo router-serve`` CLI subcommand.
|
||||
|
||||
Launches the streaming router from a YAML config. Separate from
|
||||
``fastvideo serve`` because the router is an orthogonal process: it
|
||||
fronts one or more running servers rather than hosting a generator
|
||||
itself.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from fastvideo.api.parser import load_raw_config
|
||||
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.entrypoints.streaming.router.config import (
|
||||
ReplicaEndpoint,
|
||||
RouterConfig,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class RouterServeSubcommand(CLISubcommand):
|
||||
"""Start the multi-replica WebSocket router."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.name = "router-serve"
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
config = _load_router_config(args.config)
|
||||
logger.info(
|
||||
"router listening on %s:%d (%d replicas, %d primary)",
|
||||
config.host,
|
||||
config.port,
|
||||
len(config.replicas),
|
||||
sum(1 for r in config.replicas if r.primary),
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.router.main import run_router
|
||||
|
||||
run_router(config)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
if not args.config:
|
||||
raise ValueError("fastvideo router-serve requires --config PATH")
|
||||
if not os.path.exists(args.config):
|
||||
raise ValueError(f"Router config file not found: {args.config}")
|
||||
|
||||
def subparser_init(
|
||||
self,
|
||||
subparsers: argparse._SubParsersAction,
|
||||
) -> FlexibleArgumentParser:
|
||||
parser = subparsers.add_parser(
|
||||
"router-serve",
|
||||
help="Start the streaming router (multi-replica load balancer)",
|
||||
usage="fastvideo router-serve --config ROUTER_CONFIG",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default="",
|
||||
required=False,
|
||||
help="Path to a YAML/JSON router config. Required.",
|
||||
)
|
||||
return cast(FlexibleArgumentParser, parser)
|
||||
|
||||
|
||||
def _load_router_config(path: str) -> RouterConfig:
|
||||
raw = load_raw_config(path)
|
||||
router_raw = raw.get("router") if isinstance(raw, dict) else None
|
||||
if not isinstance(router_raw, dict):
|
||||
raise ValueError(f"Router config {path!r} must have a top-level `router:` block")
|
||||
|
||||
replicas_raw = router_raw.get("replicas", [])
|
||||
if not isinstance(replicas_raw, list):
|
||||
raise ValueError(f"router.replicas must be a list, got {type(replicas_raw).__name__}")
|
||||
replicas = []
|
||||
for i, r in enumerate(replicas_raw):
|
||||
if not isinstance(r, dict):
|
||||
raise ValueError(f"router.replicas[{i}] must be a mapping, got {type(r).__name__}")
|
||||
url = r.get("url")
|
||||
if not url:
|
||||
raise ValueError(f"router.replicas[{i}] is missing required key 'url'")
|
||||
replicas.append(
|
||||
ReplicaEndpoint(
|
||||
url=url,
|
||||
name=r.get("name"),
|
||||
primary=bool(r.get("primary", False)),
|
||||
weight=float(r.get("weight", 1.0)),
|
||||
))
|
||||
if not replicas:
|
||||
raise ValueError("Router config must list at least one replica under `router.replicas`")
|
||||
|
||||
health_check = router_raw.get("health_check") or {}
|
||||
return RouterConfig(
|
||||
host=str(router_raw.get("host", "0.0.0.0")),
|
||||
port=int(router_raw.get("port", 9000)),
|
||||
replicas=replicas,
|
||||
health_check_path=str(health_check.get("path", "/health")),
|
||||
health_check_interval_seconds=float(health_check.get("interval_seconds", 5.0)),
|
||||
health_check_timeout_seconds=float(health_check.get("timeout_seconds", 2.0)),
|
||||
failure_threshold=int(health_check.get("failure_threshold", 3)),
|
||||
recovery_threshold=int(health_check.get("recovery_threshold", 2)),
|
||||
)
|
||||
|
||||
|
||||
def cmd_init() -> list[CLISubcommand]:
|
||||
return [RouterServeSubcommand()]
|
||||
|
||||
|
||||
__all__ = ["RouterServeSubcommand", "cmd_init"]
|
||||
@@ -11,28 +11,6 @@ from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.gpu_pool import (
|
||||
GpuPool,
|
||||
InProcessGpuPool,
|
||||
PoolAcquireTimeout,
|
||||
SubprocessGpuPool,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.mock_server import (
|
||||
MockGenerator,
|
||||
build_mock_app,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt import (
|
||||
LLMProvider,
|
||||
PromptEnhancer,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt.safety import (
|
||||
PromptSafetyFilter,
|
||||
SafetyDecision,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_logger import (
|
||||
SessionLogEvent,
|
||||
SessionLogger,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.stream import (
|
||||
FragmentedMP4Chunk,
|
||||
FragmentedMP4Encoder,
|
||||
@@ -42,24 +20,12 @@ __all__ = [
|
||||
"BlobStore",
|
||||
"FragmentedMP4Chunk",
|
||||
"FragmentedMP4Encoder",
|
||||
"GpuPool",
|
||||
"InMemoryBlobStore",
|
||||
"InMemorySessionStore",
|
||||
"InProcessGpuPool",
|
||||
"LLMProvider",
|
||||
"MockGenerator",
|
||||
"PoolAcquireTimeout",
|
||||
"PromptEnhancer",
|
||||
"PromptSafetyFilter",
|
||||
"SafetyDecision",
|
||||
"SessionLogEvent",
|
||||
"SessionLogger",
|
||||
"build_mock_app",
|
||||
"Session",
|
||||
"SessionManager",
|
||||
"SessionState",
|
||||
"SessionStore",
|
||||
"SubprocessGpuPool",
|
||||
"build_app",
|
||||
"run_server",
|
||||
]
|
||||
|
||||
@@ -1,542 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GPU pool manager for the streaming server.
|
||||
|
||||
Replaces the single-generator path in PR 7.5 with a typed pool
|
||||
abstraction. Three implementations ship here:
|
||||
|
||||
* :class:`InProcessGpuPool` — one in-process ``VideoGenerator``; used
|
||||
by tests and single-GPU dev deployments.
|
||||
* :class:`SubprocessGpuPool` — one ``multiprocessing.Process`` per
|
||||
GPU, each running :func:`worker_main` against a ``GeneratorConfig``.
|
||||
Jobs are dispatched via ``multiprocessing.Queue``.
|
||||
* :class:`GpuPool` (abstract) — the interface both use.
|
||||
|
||||
Session-to-GPU binding lives in the pool so continuation state stays
|
||||
on the GPU that generated the previous segment (matching the internal
|
||||
``gpu_pool.py``'s per-GPU cache behavior). Cross-GPU handoff is
|
||||
supported via :class:`SessionStore` snapshot + hydrate, which
|
||||
serializes the state before the migration and rehydrates it on the
|
||||
new worker.
|
||||
|
||||
Typed config: workers start from a :class:`GeneratorConfig` (no flat
|
||||
LTX-2 kwargs), satisfying the PR 6 + PR 7 contracts that the public
|
||||
surface doesn't reintroduce the legacy kwarg bag.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import multiprocessing as mp
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from abc import ABC, abstractmethod
|
||||
from concurrent.futures import Future
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Protocol
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
GeneratorConfig,
|
||||
GenerationRequest,
|
||||
GpuPoolConfig,
|
||||
WarmupConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.worker import worker_main
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public interface
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class _GeneratorLike(Protocol):
|
||||
"""Subset the pool calls on a worker-side generator."""
|
||||
|
||||
def generate(self, request: GenerationRequest) -> Any:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
class PoolAssignment:
|
||||
"""The worker a session is currently bound to."""
|
||||
|
||||
gpu_id: int
|
||||
worker_id: str
|
||||
pinned_at: float = field(default_factory=time.monotonic)
|
||||
|
||||
|
||||
class GpuPool(ABC):
|
||||
"""Abstract GPU pool.
|
||||
|
||||
``acquire`` binds a session to a worker and holds that binding
|
||||
across segments so continuation state can stay hot. ``run`` submits
|
||||
a single ``GenerationRequest`` for a bound session.
|
||||
|
||||
Acquire / release are independent of run — a session can run many
|
||||
segments on one acquired worker, and must release on disconnect.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def acquire(
|
||||
self,
|
||||
session_id: str,
|
||||
*,
|
||||
timeout: float | None = None,
|
||||
) -> PoolAssignment:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def run(
|
||||
self,
|
||||
session_id: str,
|
||||
request: GenerationRequest,
|
||||
) -> Any:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def release(self, session_id: str) -> None:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
async def shutdown(self) -> None:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def health(self) -> PoolHealth:
|
||||
...
|
||||
|
||||
|
||||
@dataclass
|
||||
class PoolHealth:
|
||||
total_workers: int
|
||||
available_workers: int
|
||||
active_sessions: int
|
||||
queued_sessions: int = 0
|
||||
|
||||
|
||||
class PoolAcquireTimeout(RuntimeError):
|
||||
"""Raised when ``acquire`` times out waiting for a free worker."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# In-process implementation (single-worker, test / dev)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class InProcessGpuPool(GpuPool):
|
||||
"""Single-process pool backed by one :class:`_GeneratorLike`.
|
||||
|
||||
This is what PR 7.5's server uses by default; PR 7.6 adds the real
|
||||
``SubprocessGpuPool`` alternative but keeps this one for tests and
|
||||
small deployments.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
generator: _GeneratorLike,
|
||||
*,
|
||||
gpu_id: int = 0,
|
||||
session_store: SessionStore | None = None,
|
||||
) -> None:
|
||||
self._generator = generator
|
||||
self._gpu_id = gpu_id
|
||||
self._worker_id = f"inproc-{uuid.uuid4().hex[:6]}"
|
||||
self._session_store = session_store or InMemorySessionStore()
|
||||
self._active: dict[str, PoolAssignment] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
self._gen_lock = asyncio.Lock()
|
||||
|
||||
async def acquire(
|
||||
self,
|
||||
session_id: str,
|
||||
*,
|
||||
timeout: float | None = None,
|
||||
) -> PoolAssignment:
|
||||
async with self._lock:
|
||||
existing = self._active.get(session_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
assignment = PoolAssignment(gpu_id=self._gpu_id, worker_id=self._worker_id)
|
||||
self._active[session_id] = assignment
|
||||
return assignment
|
||||
|
||||
async def run(
|
||||
self,
|
||||
session_id: str,
|
||||
request: GenerationRequest,
|
||||
) -> Any:
|
||||
if session_id not in self._active:
|
||||
raise RuntimeError(f"session {session_id!r} is not acquired on this pool")
|
||||
# Serialize generator access so one GPU runs one request at a
|
||||
# time, matching the internal gpu_pool's per-GPU lock.
|
||||
async with self._gen_lock:
|
||||
loop = asyncio.get_running_loop()
|
||||
return await loop.run_in_executor(None, self._generator.generate, request)
|
||||
|
||||
async def release(self, session_id: str) -> None:
|
||||
async with self._lock:
|
||||
self._active.pop(session_id, None)
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
self._active.clear()
|
||||
|
||||
def health(self) -> PoolHealth:
|
||||
return PoolHealth(
|
||||
total_workers=1,
|
||||
available_workers=1 if not self._active else 0,
|
||||
active_sessions=len(self._active),
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Subprocess implementation (multi-worker, real deployment)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class _WorkerHandle:
|
||||
process: Any # mp.Process or compatible handle with is_alive / join / kill
|
||||
job_queue: mp.Queue
|
||||
result_queue: mp.Queue
|
||||
gpu_id: int
|
||||
worker_id: str
|
||||
ready: threading.Event
|
||||
# ``ready`` flips on either successful boot or boot failure so the
|
||||
# parent stops waiting; ``boot_ok`` is set only on a real ready
|
||||
# acknowledgement and is what gates pool admission.
|
||||
boot_ok: threading.Event
|
||||
shutdown_event: Any # mp.Event is a factory, not a type — Any keeps mypy sane
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PendingJob:
|
||||
job_id: str
|
||||
future: Future
|
||||
session_id: str
|
||||
worker_id: str
|
||||
|
||||
|
||||
class SubprocessGpuPool(GpuPool):
|
||||
"""One ``multiprocessing.Process`` per GPU.
|
||||
|
||||
Each worker boots :class:`fastvideo.VideoGenerator` from a typed
|
||||
:class:`GeneratorConfig` inside the child process (post-
|
||||
``CUDA_VISIBLE_DEVICES`` setup) and consumes jobs from an mp Queue.
|
||||
|
||||
This is the production shape: the parent process stays CPU-only, and
|
||||
GPU state never crosses process boundaries. Continuation state is
|
||||
serialized through :class:`SessionStore` for cross-GPU handoff.
|
||||
|
||||
PR 7.6 ships this as an opt-in; PR 7.5's in-process pool remains the
|
||||
default until nightly runs validate the subprocess path.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
generator_config: GeneratorConfig,
|
||||
*,
|
||||
pool_config: GpuPoolConfig,
|
||||
warmup_config: WarmupConfig | None = None,
|
||||
session_store: SessionStore | None = None,
|
||||
worker_factory: WorkerFactory | None = None,
|
||||
) -> None:
|
||||
self._generator_config = generator_config
|
||||
self._pool_config = pool_config
|
||||
self._warmup_config = warmup_config or WarmupConfig()
|
||||
self._session_store = session_store or InMemorySessionStore()
|
||||
self._worker_factory = worker_factory or _default_worker_factory
|
||||
self._workers: list[_WorkerHandle] = []
|
||||
self._available: asyncio.Queue[int] = asyncio.Queue()
|
||||
self._assignments: dict[str, PoolAssignment] = {}
|
||||
self._worker_by_id: dict[str, _WorkerHandle] = {}
|
||||
self._pending: dict[str, _PendingJob] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
self._result_reader_tasks: list[asyncio.Task] = []
|
||||
|
||||
async def start(self) -> None:
|
||||
"""Spawn worker processes and wait for each to report ready."""
|
||||
num_workers = self._pool_config.num_workers or 1
|
||||
for gpu_id in range(num_workers):
|
||||
handle = self._worker_factory(
|
||||
gpu_id=gpu_id,
|
||||
generator_config=self._generator_config,
|
||||
warmup_config=self._warmup_config,
|
||||
)
|
||||
self._workers.append(handle)
|
||||
self._worker_by_id[handle.worker_id] = handle
|
||||
|
||||
# Wait for each worker's ready event in a thread to avoid
|
||||
# blocking the event loop.
|
||||
loop = asyncio.get_running_loop()
|
||||
await asyncio.gather(*[
|
||||
loop.run_in_executor(None, handle.ready.wait, self._warmup_config.timeout_seconds)
|
||||
for handle in self._workers
|
||||
])
|
||||
|
||||
# Start background result readers — one task per worker
|
||||
# drains its result queue and resolves futures in _pending.
|
||||
for handle in self._workers:
|
||||
task = asyncio.create_task(self._drain_results(handle))
|
||||
self._result_reader_tasks.append(task)
|
||||
|
||||
# Only admit workers that successfully booted. Anything that
|
||||
# failed boot (timeout, crash, error sentinel) stays out of the
|
||||
# available queue so we never assign a session to it.
|
||||
for idx, handle in enumerate(self._workers):
|
||||
if handle.boot_ok.is_set():
|
||||
await self._available.put(idx)
|
||||
else:
|
||||
logger.error(
|
||||
"pool: worker %s failed to boot; skipping",
|
||||
handle.worker_id,
|
||||
)
|
||||
|
||||
async def acquire(
|
||||
self,
|
||||
session_id: str,
|
||||
*,
|
||||
timeout: float | None = None,
|
||||
) -> PoolAssignment:
|
||||
async with self._lock:
|
||||
existing = self._assignments.get(session_id)
|
||||
if existing is not None:
|
||||
return existing
|
||||
try:
|
||||
idx = await asyncio.wait_for(self._available.get(), timeout=timeout)
|
||||
except asyncio.TimeoutError as exc:
|
||||
raise PoolAcquireTimeout(f"no worker available after {timeout}s") from exc
|
||||
handle = self._workers[idx]
|
||||
assignment = PoolAssignment(gpu_id=handle.gpu_id, worker_id=handle.worker_id)
|
||||
async with self._lock:
|
||||
self._assignments[session_id] = assignment
|
||||
return assignment
|
||||
|
||||
async def run(
|
||||
self,
|
||||
session_id: str,
|
||||
request: GenerationRequest,
|
||||
) -> Any:
|
||||
assignment = self._assignments.get(session_id)
|
||||
if assignment is None:
|
||||
raise RuntimeError(f"session {session_id!r} not acquired on this pool")
|
||||
handle = self._worker_by_id[assignment.worker_id]
|
||||
job_id = uuid.uuid4().hex
|
||||
future: Future = Future()
|
||||
self._pending[job_id] = _PendingJob(
|
||||
job_id=job_id,
|
||||
future=future,
|
||||
session_id=session_id,
|
||||
worker_id=handle.worker_id,
|
||||
)
|
||||
# mp.Queue.put can block if the underlying pipe buffer is full;
|
||||
# offload to a thread so the event loop keeps serving other
|
||||
# sessions. If the put itself fails, drop the pending entry so
|
||||
# _drain_results doesn't dangle a future forever.
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
await loop.run_in_executor(
|
||||
None,
|
||||
handle.job_queue.put,
|
||||
{
|
||||
"job_id": job_id,
|
||||
"request": request
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
self._pending.pop(job_id, None)
|
||||
raise
|
||||
return await asyncio.wrap_future(future)
|
||||
|
||||
async def release(self, session_id: str) -> None:
|
||||
async with self._lock:
|
||||
assignment = self._assignments.pop(session_id, None)
|
||||
if assignment is None:
|
||||
return
|
||||
idx = next((i for i, h in enumerate(self._workers) if h.worker_id == assignment.worker_id), None)
|
||||
if idx is None:
|
||||
return
|
||||
# Don't return a dead worker to the pool; otherwise the next
|
||||
# acquire will hand a session to a process that can't run jobs.
|
||||
if not self._workers[idx].process.is_alive():
|
||||
logger.warning(
|
||||
"pool: worker %s died; not returning to available queue",
|
||||
self._workers[idx].worker_id,
|
||||
)
|
||||
return
|
||||
await self._available.put(idx)
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
# Signal all workers in parallel; .put may block on a full pipe,
|
||||
# so off-load it the same way run() does.
|
||||
async def _signal(handle: _WorkerHandle) -> None:
|
||||
try:
|
||||
handle.shutdown_event.set()
|
||||
await loop.run_in_executor(None, handle.job_queue.put, None)
|
||||
except Exception: # pragma: no cover - best-effort cleanup
|
||||
pass
|
||||
|
||||
await asyncio.gather(*(_signal(h) for h in self._workers))
|
||||
# Join in parallel so total shutdown is bounded by the slowest
|
||||
# worker, not the sum of all timeouts.
|
||||
await asyncio.gather(*(loop.run_in_executor(None, handle.process.join, 5.0) for handle in self._workers))
|
||||
for handle in self._workers:
|
||||
if handle.process.is_alive():
|
||||
handle.process.kill()
|
||||
for task in self._result_reader_tasks:
|
||||
task.cancel()
|
||||
self._result_reader_tasks.clear()
|
||||
self._workers.clear()
|
||||
self._worker_by_id.clear()
|
||||
|
||||
def health(self) -> PoolHealth:
|
||||
return PoolHealth(
|
||||
total_workers=len(self._workers),
|
||||
available_workers=self._available.qsize(),
|
||||
active_sessions=len(self._assignments),
|
||||
)
|
||||
|
||||
async def _drain_results(self, handle: _WorkerHandle) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
try:
|
||||
while not handle.shutdown_event.is_set():
|
||||
try:
|
||||
msg = await loop.run_in_executor(None, _safe_queue_get, handle.result_queue, 0.5)
|
||||
except Exception:
|
||||
logger.exception("pool: worker %s result reader failed", handle.worker_id)
|
||||
return
|
||||
if msg is None:
|
||||
continue
|
||||
job_id = msg.get("job_id")
|
||||
if job_id is None:
|
||||
continue
|
||||
pending = self._pending.pop(job_id, None)
|
||||
if pending is None:
|
||||
continue
|
||||
if msg.get("kind") == "error":
|
||||
pending.future.set_exception(RuntimeError(msg["error"]))
|
||||
else:
|
||||
pending.future.set_result(msg.get("result"))
|
||||
finally:
|
||||
# If we exit for any reason — shutdown, exception, cancel —
|
||||
# surface that to any in-flight jobs on this worker so their
|
||||
# await never hangs on a future no one will resolve.
|
||||
for jid in [jid for jid, job in self._pending.items() if job.worker_id == handle.worker_id]:
|
||||
pending = self._pending.pop(jid, None)
|
||||
if pending is not None and not pending.future.done():
|
||||
pending.future.set_exception(
|
||||
RuntimeError(f"worker {handle.worker_id} result reader exited "
|
||||
"with pending jobs"))
|
||||
|
||||
|
||||
def _safe_queue_get(q: mp.Queue, timeout: float) -> Any | None:
|
||||
try:
|
||||
return q.get(timeout=timeout)
|
||||
except queue.Empty:
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Worker process
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class WorkerFactory(Protocol):
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
*,
|
||||
gpu_id: int,
|
||||
generator_config: GeneratorConfig,
|
||||
warmup_config: WarmupConfig,
|
||||
) -> _WorkerHandle:
|
||||
...
|
||||
|
||||
|
||||
def _default_worker_factory(
|
||||
*,
|
||||
gpu_id: int,
|
||||
generator_config: GeneratorConfig,
|
||||
warmup_config: WarmupConfig,
|
||||
) -> _WorkerHandle:
|
||||
"""Spawn a real multiprocessing worker.
|
||||
|
||||
The child process calls :func:`worker_main` which constructs a
|
||||
:class:`VideoGenerator` from ``generator_config`` and runs a
|
||||
blocking job loop. The ``ready`` event flips after the warmup
|
||||
request completes.
|
||||
"""
|
||||
ctx = mp.get_context("spawn")
|
||||
job_queue: mp.Queue = ctx.Queue()
|
||||
result_queue: mp.Queue = ctx.Queue()
|
||||
ready = threading.Event()
|
||||
boot_ok = threading.Event()
|
||||
shutdown_event = ctx.Event()
|
||||
worker_id = f"gpu{gpu_id}-{uuid.uuid4().hex[:6]}"
|
||||
process = ctx.Process(
|
||||
target=worker_main,
|
||||
kwargs={
|
||||
"gpu_id": gpu_id,
|
||||
"worker_id": worker_id,
|
||||
"generator_config": generator_config,
|
||||
"warmup_config": warmup_config,
|
||||
"job_queue": job_queue,
|
||||
"result_queue": result_queue,
|
||||
"shutdown_event": shutdown_event,
|
||||
},
|
||||
daemon=False,
|
||||
)
|
||||
process.start()
|
||||
|
||||
# Block the parent-side ``ready`` flag until the worker posts a
|
||||
# ready acknowledgement on the result queue. We drain that single
|
||||
# sentinel here; subsequent results belong to jobs. ``boot_ok``
|
||||
# only flips on a real ready; on error we set ``ready`` to unblock
|
||||
# the parent's wait but leave ``boot_ok`` clear so the pool keeps
|
||||
# the worker out of the available queue.
|
||||
def _await_ready() -> None:
|
||||
while not shutdown_event.is_set():
|
||||
try:
|
||||
msg = result_queue.get(timeout=1.0)
|
||||
except queue.Empty:
|
||||
continue
|
||||
if isinstance(msg, dict) and msg.get("kind") == "ready":
|
||||
boot_ok.set()
|
||||
ready.set()
|
||||
return
|
||||
if isinstance(msg, dict) and msg.get("kind") == "error":
|
||||
logger.error("pool: worker %s failed to boot: %s", worker_id, msg.get("error"))
|
||||
ready.set()
|
||||
return
|
||||
|
||||
threading.Thread(target=_await_ready, daemon=True).start()
|
||||
|
||||
return _WorkerHandle(
|
||||
process=process,
|
||||
job_queue=job_queue,
|
||||
result_queue=result_queue,
|
||||
gpu_id=gpu_id,
|
||||
worker_id=worker_id,
|
||||
ready=ready,
|
||||
boot_ok=boot_ok,
|
||||
shutdown_event=shutdown_event,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"GpuPool",
|
||||
"InProcessGpuPool",
|
||||
"PoolAcquireTimeout",
|
||||
"PoolAssignment",
|
||||
"PoolHealth",
|
||||
"SubprocessGpuPool",
|
||||
"WorkerFactory",
|
||||
"worker_main",
|
||||
]
|
||||
@@ -1,122 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Mock streaming server — a frontend dev aid.
|
||||
|
||||
Boots the same FastAPI app the real streaming server uses, but backs
|
||||
it with :class:`InProcessGpuPool` wrapping a synthetic generator that
|
||||
emits pre-baked RGB frames. No GPU or model weights required.
|
||||
|
||||
Use cases:
|
||||
|
||||
* Frontend development without a real model loaded.
|
||||
* Integration tests that exercise the WS protocol end-to-end.
|
||||
* Reproducing protocol bugs locally.
|
||||
|
||||
Launch: ``python -m fastvideo.entrypoints.streaming.mock_server``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
SamplingConfig,
|
||||
ServeConfig,
|
||||
StreamingConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.server import build_app
|
||||
|
||||
|
||||
@dataclass
|
||||
class MockGenerator:
|
||||
"""Generator stand-in that returns synthetic gradient frames.
|
||||
|
||||
Each call produces one segment worth of frames whose pixels vary by
|
||||
a constant derived from the request seed and segment index. Latency
|
||||
is configurable via ``sleep_ms`` so the caller can exercise slow-
|
||||
generate scenarios without spinning a GPU.
|
||||
"""
|
||||
|
||||
sleep_ms: float = 0.0
|
||||
|
||||
def generate(self, request: GenerationRequest) -> dict[str, Any]:
|
||||
if self.sleep_ms:
|
||||
time.sleep(self.sleep_ms / 1000.0)
|
||||
width = max(16, request.sampling.width)
|
||||
height = max(16, request.sampling.height)
|
||||
num_frames = max(1, request.sampling.num_frames)
|
||||
frames = [_gradient_frame(height, width, idx, seed=request.sampling.seed) for idx in range(num_frames)]
|
||||
state = ContinuationState(
|
||||
kind="ltx2.v1",
|
||||
payload={
|
||||
"schema_version": 1,
|
||||
"segment_index": 0,
|
||||
"source_prompt": request.prompt,
|
||||
},
|
||||
)
|
||||
return {
|
||||
"frames": frames,
|
||||
"audio_sample_rate": 24000,
|
||||
"state": state,
|
||||
}
|
||||
|
||||
|
||||
def _gradient_frame(height: int, width: int, idx: int, *, seed: int) -> np.ndarray:
|
||||
base = (idx * 17 + seed * 3) % 256
|
||||
row = np.linspace(base, (base + 64) % 256, width, dtype=np.uint8)
|
||||
frame = np.tile(row, (height, 1))
|
||||
stacked = np.stack([frame, np.roll(frame, 8, axis=1), np.roll(frame, 16, axis=1)], axis=-1)
|
||||
return stacked.astype(np.uint8)
|
||||
|
||||
|
||||
def build_mock_app(*, sleep_ms: float = 0.0):
|
||||
"""Build a FastAPI app backed by :class:`MockGenerator`."""
|
||||
serve_config = ServeConfig(
|
||||
generator=GeneratorConfig(model_path="/models/mock"),
|
||||
streaming=StreamingConfig(
|
||||
session_timeout_seconds=120,
|
||||
generation_segment_cap=6,
|
||||
),
|
||||
)
|
||||
serve_config.default_request.sampling = SamplingConfig(
|
||||
num_frames=24,
|
||||
height=256,
|
||||
width=256,
|
||||
fps=24,
|
||||
num_inference_steps=1,
|
||||
)
|
||||
return build_app(serve_config, MockGenerator(sleep_ms=sleep_ms))
|
||||
|
||||
|
||||
def main() -> None: # pragma: no cover - CLI entry
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--host", default="127.0.0.1")
|
||||
parser.add_argument("--port", type=int, default=8000)
|
||||
parser.add_argument(
|
||||
"--sleep-ms",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Per-segment artificial latency for testing slow paths",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
import uvicorn
|
||||
|
||||
app = build_mock_app(sleep_ms=args.sleep_ms)
|
||||
uvicorn.run(app, host=args.host, port=args.port)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MockGenerator",
|
||||
"build_mock_app",
|
||||
"main",
|
||||
]
|
||||
|
||||
if __name__ == "__main__": # pragma: no cover - CLI entry
|
||||
main()
|
||||
@@ -1,36 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Prompt pipeline for the streaming server.
|
||||
|
||||
* :mod:`providers` — LLM backend abstraction + built-in adapters
|
||||
* :mod:`enhancer` — provider-agnostic enhance / auto-extend / rewrite
|
||||
operations on top of the provider layer
|
||||
|
||||
All of this is optional; the streaming server runs fine without it
|
||||
(PR 7.5's skeleton never invokes the enhancer). When the operator
|
||||
enables ``ServeConfig.streaming.prompt.enabled``, the server routes
|
||||
each ``session_init_v2`` curated prompt through ``enhance`` before the
|
||||
first segment.
|
||||
"""
|
||||
from fastvideo.entrypoints.streaming.prompt.enhancer import (
|
||||
PromptEnhancer,
|
||||
PromptOperation,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMMessage,
|
||||
LLMProvider,
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
LLMTimeoutError,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"LLMMessage",
|
||||
"LLMProvider",
|
||||
"LLMProviderError",
|
||||
"LLMRequest",
|
||||
"LLMResponse",
|
||||
"LLMTimeoutError",
|
||||
"PromptEnhancer",
|
||||
"PromptOperation",
|
||||
]
|
||||
@@ -1,197 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Provider-agnostic prompt orchestration for the streaming server.
|
||||
|
||||
Three operations the streaming server needs:
|
||||
|
||||
* ``enhance`` — polish a user prompt (add cinematic detail, fix syntax)
|
||||
* ``auto_extend`` — generate a follow-on prompt for loop generation
|
||||
* ``rewrite`` — rewrite a seed prompt for a user-directed rewrite flow
|
||||
|
||||
All three share the same orchestration: pick a provider in priority
|
||||
order, submit an ``LLMRequest``, fall back to the next provider on
|
||||
retryable errors, and surface a structured :class:`LLMResponse` back
|
||||
to the caller.
|
||||
|
||||
System prompts are loaded from ``system_prompt_dir`` on construction
|
||||
and can be hot-reloaded via :meth:`PromptEnhancer.reload_system_prompts`.
|
||||
The streaming server's management endpoint calls that method in
|
||||
response to a ``rewrite_seed_prompts_started`` frame.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import os
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, replace
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMMessage,
|
||||
LLMProvider,
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class PromptOperation(enum.Enum):
|
||||
ENHANCE = "enhance"
|
||||
AUTO_EXTEND = "auto_extend"
|
||||
REWRITE = "rewrite"
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SystemPrompts:
|
||||
enhance: str
|
||||
auto_extend: str
|
||||
rewrite: str
|
||||
|
||||
|
||||
_DEFAULT_SYSTEM_PROMPTS = _SystemPrompts(
|
||||
enhance=("You are a prompt enhancer for cinematic video generation. Given "
|
||||
"a user prompt, produce an enhanced prompt that is more vivid, "
|
||||
"specific, and concrete. Keep the subject intact; add lighting, "
|
||||
"camera, and motion detail. Reply with just the enhanced prompt."),
|
||||
auto_extend=("You are a video continuation assistant. Given the current "
|
||||
"sequence of prompts, produce one new prompt that naturally "
|
||||
"continues the sequence. Reply with just the next prompt."),
|
||||
rewrite=("You are a creative prompt rewriter. Given a seed prompt, produce "
|
||||
"a set of alternative prompts that explore different angles, "
|
||||
"styles, and moods. Reply with one prompt per line."),
|
||||
)
|
||||
|
||||
|
||||
class PromptEnhancer:
|
||||
"""Orchestrates prompt operations across a priority-ordered provider
|
||||
list with structured fallback + hot-reloadable system prompts.
|
||||
|
||||
Usage::
|
||||
|
||||
enhancer = PromptEnhancer(
|
||||
providers=[CerebrasProvider(), GroqProvider()],
|
||||
model="gpt-oss-120b",
|
||||
system_prompt_dir="/etc/fastvideo/prompts",
|
||||
)
|
||||
response = await enhancer.enhance("a fox running through snow")
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
providers: Sequence[LLMProvider],
|
||||
model: str,
|
||||
timeout_ms: int = 20000,
|
||||
temperature: float = 0.7,
|
||||
max_tokens: int | None = 256,
|
||||
system_prompt_dir: str | None = None,
|
||||
) -> None:
|
||||
if not providers:
|
||||
raise ValueError("PromptEnhancer requires at least one LLMProvider")
|
||||
self._providers = list(providers)
|
||||
self._model = model
|
||||
self._timeout_ms = timeout_ms
|
||||
self._temperature = temperature
|
||||
self._max_tokens = max_tokens
|
||||
self._system_prompt_dir = system_prompt_dir
|
||||
self._system_prompts = self._load_system_prompts()
|
||||
|
||||
@property
|
||||
def providers(self) -> list[LLMProvider]:
|
||||
return list(self._providers)
|
||||
|
||||
def register_provider(self, provider: LLMProvider, *, priority: int = -1) -> None:
|
||||
"""Insert an additional provider. ``priority=0`` makes it primary;
|
||||
``priority=-1`` (default) appends as a fallback."""
|
||||
if priority < 0:
|
||||
self._providers.append(provider)
|
||||
else:
|
||||
self._providers.insert(priority, provider)
|
||||
|
||||
def reload_system_prompts(self) -> None:
|
||||
"""Re-read the system prompt files from ``system_prompt_dir``.
|
||||
|
||||
The streaming server exposes this via a management endpoint so
|
||||
operators can iterate on prompt templates without restarting
|
||||
workers.
|
||||
"""
|
||||
self._system_prompts = self._load_system_prompts()
|
||||
logger.info("prompt enhancer: reloaded system prompts from %s", self._system_prompt_dir or "defaults")
|
||||
|
||||
async def enhance(self, prompt: str) -> LLMResponse:
|
||||
return await self._run(
|
||||
PromptOperation.ENHANCE,
|
||||
system=self._system_prompts.enhance,
|
||||
user=prompt,
|
||||
)
|
||||
|
||||
async def auto_extend(self, prior_prompts: Sequence[str]) -> LLMResponse:
|
||||
user = "\n".join(prior_prompts)
|
||||
return await self._run(
|
||||
PromptOperation.AUTO_EXTEND,
|
||||
system=self._system_prompts.auto_extend,
|
||||
user=user,
|
||||
)
|
||||
|
||||
async def rewrite(self, seed_prompt: str) -> LLMResponse:
|
||||
return await self._run(
|
||||
PromptOperation.REWRITE,
|
||||
system=self._system_prompts.rewrite,
|
||||
user=seed_prompt,
|
||||
)
|
||||
|
||||
async def _run(
|
||||
self,
|
||||
operation: PromptOperation,
|
||||
*,
|
||||
system: str,
|
||||
user: str,
|
||||
) -> LLMResponse:
|
||||
request = LLMRequest(
|
||||
messages=[
|
||||
LLMMessage(role="system", content=system),
|
||||
LLMMessage(role="user", content=user),
|
||||
],
|
||||
model=self._model,
|
||||
max_tokens=self._max_tokens,
|
||||
temperature=self._temperature,
|
||||
timeout_ms=self._timeout_ms,
|
||||
)
|
||||
last_error: LLMProviderError | None = None
|
||||
for idx, provider in enumerate(self._providers):
|
||||
try:
|
||||
response = await provider.complete(request)
|
||||
if idx > 0:
|
||||
# Mark the fallback flag without losing any other
|
||||
# response fields the provider populated.
|
||||
response = replace(response, fallback_used=True)
|
||||
return response
|
||||
except LLMProviderError as exc:
|
||||
logger.warning("prompt %s: provider %s failed: %s; trying next", operation.value, provider.name, exc)
|
||||
last_error = exc
|
||||
if not exc.retryable:
|
||||
break
|
||||
assert last_error is not None
|
||||
raise last_error
|
||||
|
||||
def _load_system_prompts(self) -> _SystemPrompts:
|
||||
if not self._system_prompt_dir:
|
||||
return _DEFAULT_SYSTEM_PROMPTS
|
||||
return _SystemPrompts(
|
||||
enhance=_read_prompt(self._system_prompt_dir, "enhance.txt", _DEFAULT_SYSTEM_PROMPTS.enhance),
|
||||
auto_extend=_read_prompt(self._system_prompt_dir, "auto_extend.txt", _DEFAULT_SYSTEM_PROMPTS.auto_extend),
|
||||
rewrite=_read_prompt(self._system_prompt_dir, "rewrite.txt", _DEFAULT_SYSTEM_PROMPTS.rewrite),
|
||||
)
|
||||
|
||||
|
||||
def _read_prompt(dirname: str, filename: str, default: str) -> str:
|
||||
path = os.path.join(dirname, filename)
|
||||
if not os.path.exists(path):
|
||||
return default
|
||||
with open(path, encoding="utf-8") as f:
|
||||
content = f.read().strip()
|
||||
return content or default
|
||||
|
||||
|
||||
__all__ = ["PromptEnhancer", "PromptOperation"]
|
||||
@@ -1,24 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LLM provider implementations used by the prompt enhancer."""
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMMessage,
|
||||
LLMProvider,
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
LLMTimeoutError,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.cerebras import (
|
||||
CerebrasProvider, )
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.groq import GroqProvider
|
||||
|
||||
__all__ = [
|
||||
"CerebrasProvider",
|
||||
"GroqProvider",
|
||||
"LLMMessage",
|
||||
"LLMProvider",
|
||||
"LLMProviderError",
|
||||
"LLMRequest",
|
||||
"LLMResponse",
|
||||
"LLMTimeoutError",
|
||||
]
|
||||
@@ -1,101 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Shared HTTP path for OpenAI-compatible ``/chat/completions`` providers.
|
||||
|
||||
Cerebras and Groq both expose the OpenAI chat-completions schema, so
|
||||
the request shape, error mapping, and response decoding are identical
|
||||
between them. This module centralizes that logic; the per-provider
|
||||
modules stay thin (just defaults + env var wiring).
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
LLMTimeoutError,
|
||||
)
|
||||
|
||||
|
||||
async def complete_openai_compatible(
|
||||
*,
|
||||
api_key: str | None,
|
||||
api_key_hint: str,
|
||||
base_url: str,
|
||||
provider_name: str,
|
||||
request: LLMRequest,
|
||||
) -> LLMResponse:
|
||||
"""Issue a chat-completions call and decode the OpenAI response."""
|
||||
if not api_key:
|
||||
raise LLMProviderError(
|
||||
f"{provider_name} provider requires {api_key_hint} "
|
||||
"(or explicit api_key=...)",
|
||||
retryable=False,
|
||||
)
|
||||
try:
|
||||
import httpx
|
||||
except ImportError as exc: # pragma: no cover - optional dep
|
||||
raise LLMProviderError(
|
||||
f"{provider_name} provider requires httpx; install httpx",
|
||||
retryable=False,
|
||||
) from exc
|
||||
|
||||
timeout_s = (request.timeout_ms or 20000) / 1000.0
|
||||
t0 = time.perf_counter()
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=timeout_s) as client:
|
||||
response = await client.post(
|
||||
f"{base_url}/chat/completions",
|
||||
headers={
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
},
|
||||
json={
|
||||
"model": request.model,
|
||||
"messages": [{
|
||||
"role": m.role,
|
||||
"content": m.content
|
||||
} for m in request.messages],
|
||||
"max_tokens": request.max_tokens,
|
||||
"temperature": request.temperature,
|
||||
},
|
||||
)
|
||||
except httpx.TimeoutException as exc:
|
||||
raise LLMTimeoutError(f"{provider_name} timed out after {timeout_s}s") from exc
|
||||
except httpx.HTTPError as exc:
|
||||
raise LLMProviderError(f"{provider_name} HTTP error: {exc}") from exc
|
||||
|
||||
if response.status_code >= 400:
|
||||
# 5xx and 429 (rate-limit) are retryable: another provider may
|
||||
# succeed. 4xx (auth, bad-request, etc.) are client errors —
|
||||
# the enhancer should stop fallback traversal.
|
||||
retryable = (response.status_code >= 500 or response.status_code == 429)
|
||||
raise LLMProviderError(
|
||||
f"{provider_name} returned {response.status_code}: "
|
||||
f"{response.text[:200]}",
|
||||
retryable=retryable,
|
||||
)
|
||||
|
||||
try:
|
||||
data = response.json()
|
||||
except Exception as exc:
|
||||
# Non-JSON body usually means a proxy / load-balancer error
|
||||
# page; leave it retryable so a fallback provider can try.
|
||||
raise LLMProviderError(f"{provider_name} returned non-JSON body: {exc}") from exc
|
||||
|
||||
choices = data.get("choices") or []
|
||||
if not choices:
|
||||
raise LLMProviderError(f"{provider_name} returned no choices")
|
||||
content = choices[0].get("message", {}).get("content") or ""
|
||||
|
||||
latency_ms = (time.perf_counter() - t0) * 1000.0
|
||||
return LLMResponse(
|
||||
content=content.strip(),
|
||||
provider=provider_name,
|
||||
model=request.model,
|
||||
latency_ms=latency_ms,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["complete_openai_compatible"]
|
||||
@@ -1,85 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LLM provider protocol + DTOs used by the prompt enhancer.
|
||||
|
||||
Third-party users add a new provider by implementing
|
||||
:class:`LLMProvider` and registering it with a prompt enhancer
|
||||
instance. The shipped providers live in sibling modules
|
||||
(``cerebras.py``, ``groq.py``) and each is ~100-200 LOC — the
|
||||
provider layer is intentionally thin so the enhancer stays
|
||||
provider-agnostic.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal, Protocol, runtime_checkable
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMMessage:
|
||||
role: Literal["system", "user", "assistant"]
|
||||
content: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMRequest:
|
||||
messages: list[LLMMessage]
|
||||
model: str
|
||||
max_tokens: int | None = None
|
||||
temperature: float | None = None
|
||||
timeout_ms: int | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class LLMResponse:
|
||||
content: str
|
||||
provider: str
|
||||
model: str
|
||||
latency_ms: float
|
||||
fallback_used: bool = False
|
||||
|
||||
|
||||
class LLMProviderError(RuntimeError):
|
||||
"""Raised when an LLM provider fails a request.
|
||||
|
||||
``retryable`` controls whether the enhancer falls back to the next
|
||||
provider. It is settable per-instance so the same exception type
|
||||
can describe retryable transport errors (5xx, 429) and
|
||||
non-retryable client errors (4xx auth/bad-request) without forcing
|
||||
a separate subclass for every status family.
|
||||
"""
|
||||
|
||||
def __init__(self, message: str, *, retryable: bool = True) -> None:
|
||||
super().__init__(message)
|
||||
self.retryable = retryable
|
||||
|
||||
|
||||
class LLMTimeoutError(LLMProviderError):
|
||||
"""Raised when an LLM provider times out — always retryable."""
|
||||
|
||||
def __init__(self, message: str) -> None:
|
||||
super().__init__(message, retryable=True)
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class LLMProvider(Protocol):
|
||||
"""Provider interface every LLM adapter implements.
|
||||
|
||||
Providers are async-first because every built-in implementation
|
||||
talks to an HTTP API. Synchronous providers can wrap their call in
|
||||
``asyncio.to_thread`` internally.
|
||||
"""
|
||||
|
||||
name: str
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
...
|
||||
|
||||
|
||||
__all__ = [
|
||||
"LLMMessage",
|
||||
"LLMProvider",
|
||||
"LLMProviderError",
|
||||
"LLMRequest",
|
||||
"LLMResponse",
|
||||
"LLMTimeoutError",
|
||||
]
|
||||
@@ -1,44 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Cerebras LLM provider (OpenAI-compatible chat endpoint)."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers._openai_compat import (
|
||||
complete_openai_compatible, )
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
)
|
||||
|
||||
_DEFAULT_BASE_URL = "https://api.cerebras.ai/v1"
|
||||
_API_KEY_ENV = "CEREBRAS_API_KEY"
|
||||
|
||||
|
||||
@dataclass
|
||||
class CerebrasProvider:
|
||||
"""Cerebras inference adapter.
|
||||
|
||||
``api_key`` falls back to ``CEREBRAS_API_KEY`` when unset.
|
||||
"""
|
||||
|
||||
api_key: str | None = None
|
||||
base_url: str = _DEFAULT_BASE_URL
|
||||
name: str = "cerebras"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.api_key is None:
|
||||
self.api_key = os.environ.get(_API_KEY_ENV)
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
return await complete_openai_compatible(
|
||||
api_key=self.api_key,
|
||||
api_key_hint=_API_KEY_ENV,
|
||||
base_url=self.base_url,
|
||||
provider_name=self.name,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["CerebrasProvider"]
|
||||
@@ -1,46 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Groq LLM provider (OpenAI-compatible chat endpoint)."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers._openai_compat import (
|
||||
complete_openai_compatible, )
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import (
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
)
|
||||
|
||||
_DEFAULT_BASE_URL = "https://api.groq.com/openai/v1"
|
||||
_API_KEY_ENV = "GROQ_API_KEY"
|
||||
|
||||
|
||||
@dataclass
|
||||
class GroqProvider:
|
||||
"""Groq inference adapter.
|
||||
|
||||
Identical wire format to :class:`CerebrasProvider`; both go through
|
||||
:func:`complete_openai_compatible`. The two providers differ only
|
||||
in base URL, env var, and model id conventions.
|
||||
"""
|
||||
|
||||
api_key: str | None = None
|
||||
base_url: str = _DEFAULT_BASE_URL
|
||||
name: str = "groq"
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.api_key is None:
|
||||
self.api_key = os.environ.get(_API_KEY_ENV)
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
return await complete_openai_compatible(
|
||||
api_key=self.api_key,
|
||||
api_key_hint=_API_KEY_ENV,
|
||||
base_url=self.base_url,
|
||||
provider_name=self.name,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
__all__ = ["GroqProvider"]
|
||||
@@ -1,82 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Rewrite payload builder.
|
||||
|
||||
The UI's "rewrite seed prompts" flow asks the enhancer to produce a
|
||||
batch of alternative prompts given one seed. This module packages the
|
||||
seed + options into the payload the enhancer expects and unpacks the
|
||||
response back into a typed :class:`RewriteResult`.
|
||||
|
||||
Separating this from :mod:`enhancer` keeps the enhancer provider-
|
||||
agnostic; anything UI-specific (how many alternatives to request, how
|
||||
to split the response, temperature) lives here.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.enhancer import PromptEnhancer
|
||||
|
||||
_LEADING_MARKER_RE = re.compile(r"^(?:[-*•]\s*|\d+\s*[.)]\s*)+")
|
||||
|
||||
|
||||
@dataclass
|
||||
class RewriteOptions:
|
||||
count: int = 3
|
||||
"""Number of alternative prompts to request."""
|
||||
temperature: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class RewriteResult:
|
||||
seed_prompt: str
|
||||
alternatives: list[str]
|
||||
provider: str
|
||||
model: str
|
||||
latency_ms: float
|
||||
fallback_used: bool = False
|
||||
|
||||
|
||||
async def build_rewrite(
|
||||
enhancer: PromptEnhancer,
|
||||
seed_prompt: str,
|
||||
*,
|
||||
options: RewriteOptions | None = None,
|
||||
) -> RewriteResult:
|
||||
"""Run a rewrite op through the enhancer and return a typed result."""
|
||||
if not seed_prompt.strip():
|
||||
raise ValueError("rewrite seed prompt must be non-empty")
|
||||
options = options or RewriteOptions()
|
||||
response = await enhancer.rewrite(seed_prompt)
|
||||
alternatives = _split_response(response.content, limit=options.count)
|
||||
return RewriteResult(
|
||||
seed_prompt=seed_prompt,
|
||||
alternatives=alternatives,
|
||||
provider=response.provider,
|
||||
model=response.model,
|
||||
latency_ms=response.latency_ms,
|
||||
fallback_used=response.fallback_used,
|
||||
)
|
||||
|
||||
|
||||
def _split_response(content: str, *, limit: int) -> list[str]:
|
||||
"""Split the LLM response into discrete prompt candidates.
|
||||
|
||||
The shipped system prompt instructs the model to emit one prompt
|
||||
per line; this function is forgiving about numbered lists or
|
||||
leading bullets so user-supplied system prompts don't break it.
|
||||
"""
|
||||
lines = [line.strip() for line in content.splitlines() if line.strip()]
|
||||
cleaned: list[str] = []
|
||||
for line in lines:
|
||||
stripped = _LEADING_MARKER_RE.sub("", line).strip()
|
||||
if stripped:
|
||||
cleaned.append(stripped)
|
||||
return cleaned[:max(1, limit)]
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RewriteOptions",
|
||||
"RewriteResult",
|
||||
"build_rewrite",
|
||||
]
|
||||
@@ -1,146 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Optional prompt safety filter.
|
||||
|
||||
Uses a fastText classifier to score prompts against a banned-content
|
||||
rubric. Only loaded when ``ServeConfig.streaming.safety.enabled`` is
|
||||
True and fastText is installed — users who don't need it see no
|
||||
runtime cost.
|
||||
|
||||
Install: ``pip install fastvideo[prompt-safety]`` (ships fasttext as an
|
||||
optional extra) or install fasttext directly.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import enum
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SafetyDecision(enum.Enum):
|
||||
ALLOW = "allow"
|
||||
BLOCK = "block"
|
||||
UNAVAILABLE = "unavailable"
|
||||
"""Returned when the classifier can't run (not configured, fastText
|
||||
missing). Safety is opt-in; the server treats ``UNAVAILABLE`` as
|
||||
``ALLOW`` but logs it so operators know the filter is off."""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SafetyResult:
|
||||
prompt: str
|
||||
decision: SafetyDecision
|
||||
score: float = 0.0
|
||||
label: str | None = None
|
||||
reason: str | None = None
|
||||
|
||||
|
||||
class PromptSafetyFilter:
|
||||
"""Minimal fastText-backed prompt safety filter.
|
||||
|
||||
Loads the classifier lazily on first use so the streaming server
|
||||
can construct the filter eagerly at startup without paying the
|
||||
model-load cost when safety is disabled.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
classifier_path: str | None,
|
||||
enabled: bool = True,
|
||||
block_threshold: float = 0.5,
|
||||
) -> None:
|
||||
self._classifier_path = classifier_path
|
||||
self._enabled = enabled
|
||||
self._block_threshold = block_threshold
|
||||
self._model: Any | None = None
|
||||
self._load_attempted = False
|
||||
self._load_lock = threading.Lock()
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return self._enabled and self._classifier_path is not None
|
||||
|
||||
def classify(self, prompt: str) -> SafetyResult:
|
||||
if not self.enabled:
|
||||
return SafetyResult(
|
||||
prompt=prompt,
|
||||
decision=SafetyDecision.UNAVAILABLE,
|
||||
reason="safety filter not enabled",
|
||||
)
|
||||
model = self._ensure_loaded()
|
||||
if model is None:
|
||||
return SafetyResult(
|
||||
prompt=prompt,
|
||||
decision=SafetyDecision.UNAVAILABLE,
|
||||
reason="fastText model unavailable",
|
||||
)
|
||||
try:
|
||||
labels, probs = model.predict(prompt.replace("\n", " "), k=1)
|
||||
except Exception as exc: # pragma: no cover - defensive
|
||||
logger.warning("safety: classifier failed: %s", exc)
|
||||
return SafetyResult(
|
||||
prompt=prompt,
|
||||
decision=SafetyDecision.UNAVAILABLE,
|
||||
reason=f"classifier error: {exc}",
|
||||
)
|
||||
label = labels[0].removeprefix("__label__") if labels else None
|
||||
score = float(probs[0]) if len(probs) else 0.0
|
||||
decision = (SafetyDecision.BLOCK if
|
||||
(label == "unsafe" and score >= self._block_threshold) else SafetyDecision.ALLOW)
|
||||
return SafetyResult(
|
||||
prompt=prompt,
|
||||
decision=decision,
|
||||
score=score,
|
||||
label=label,
|
||||
)
|
||||
|
||||
def _ensure_loaded(self) -> Any | None:
|
||||
if self._model is not None:
|
||||
return self._model
|
||||
if self._load_attempted:
|
||||
return None
|
||||
with self._load_lock:
|
||||
if self._model is not None:
|
||||
return self._model
|
||||
if self._load_attempted:
|
||||
return None
|
||||
self._load_attempted = True
|
||||
if self._classifier_path is None:
|
||||
return None
|
||||
try:
|
||||
import fasttext # type: ignore[import-not-found]
|
||||
except ImportError:
|
||||
logger.warning("safety: fasttext not installed; safety filter disabled. "
|
||||
"Install fastvideo[prompt-safety] to enable.")
|
||||
return None
|
||||
try:
|
||||
self._model = fasttext.load_model(self._classifier_path)
|
||||
except Exception as exc: # pragma: no cover - requires real model
|
||||
logger.warning("safety: failed to load %s: %s", self._classifier_path, exc)
|
||||
return None
|
||||
return self._model
|
||||
|
||||
|
||||
def first_blocked(
|
||||
filter_: PromptSafetyFilter,
|
||||
prompts: list[str],
|
||||
) -> SafetyResult | None:
|
||||
"""Return the first prompt the filter blocks, or ``None``."""
|
||||
for prompt in prompts:
|
||||
result = filter_.classify(prompt)
|
||||
if result.decision is SafetyDecision.BLOCK:
|
||||
return result
|
||||
return None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"PromptSafetyFilter",
|
||||
"SafetyDecision",
|
||||
"SafetyResult",
|
||||
"first_blocked",
|
||||
]
|
||||
@@ -1,27 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Multi-replica load balancer + WebSocket proxy for the streaming server.
|
||||
|
||||
Sits in front of one-or-more streaming-server replicas and forwards
|
||||
WebSocket sessions to a healthy primary, with failover to secondaries.
|
||||
Kept in-repo under ``fastvideo/entrypoints/streaming/router/`` per the
|
||||
PR plan's default; the alternative (separate package) is an open
|
||||
question deferred to review.
|
||||
"""
|
||||
from fastvideo.entrypoints.streaming.router.registry import (
|
||||
Replica,
|
||||
ReplicaHealth,
|
||||
ReplicaRegistry,
|
||||
ReplicaStatus,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.router.config import RouterConfig
|
||||
from fastvideo.entrypoints.streaming.router.main import build_router_app, run_router
|
||||
|
||||
__all__ = [
|
||||
"Replica",
|
||||
"ReplicaHealth",
|
||||
"ReplicaRegistry",
|
||||
"ReplicaStatus",
|
||||
"RouterConfig",
|
||||
"build_router_app",
|
||||
"run_router",
|
||||
]
|
||||
@@ -1,88 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Typed router configuration."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from urllib.parse import urlparse
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReplicaEndpoint:
|
||||
"""One backend replica the router can route to."""
|
||||
|
||||
url: str
|
||||
"""HTTP base URL, e.g. ``http://host:8000``. WebSocket URL is
|
||||
derived automatically by replacing the scheme."""
|
||||
name: str | None = None
|
||||
primary: bool = False
|
||||
"""``True`` = prefer this replica over others in steady state."""
|
||||
weight: float = 1.0
|
||||
|
||||
|
||||
@dataclass
|
||||
class RouterConfig:
|
||||
"""Typed router config loaded from a YAML file.
|
||||
|
||||
Example::
|
||||
|
||||
router:
|
||||
host: 0.0.0.0
|
||||
port: 9000
|
||||
replicas:
|
||||
- url: http://streamer-a:8000
|
||||
primary: true
|
||||
- url: http://streamer-b:8000
|
||||
health_check:
|
||||
path: /health
|
||||
interval_seconds: 5
|
||||
failure_threshold: 3
|
||||
|
||||
Validation runs in ``__post_init__``: empty replicas, non-positive
|
||||
intervals/timeouts, thresholds < 1, non-http(s) URLs, and more than
|
||||
one primary all raise ``ValueError`` so misconfigurations surface at
|
||||
load time rather than as confusing runtime failures.
|
||||
"""
|
||||
|
||||
host: str = "0.0.0.0"
|
||||
port: int = 9000
|
||||
replicas: list[ReplicaEndpoint] = field(default_factory=list)
|
||||
health_check_path: str = "/health"
|
||||
health_check_interval_seconds: float = 5.0
|
||||
health_check_timeout_seconds: float = 2.0
|
||||
failure_threshold: int = 3
|
||||
recovery_threshold: int = 2
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.replicas:
|
||||
raise ValueError("RouterConfig.replicas must list at least one replica")
|
||||
if self.health_check_interval_seconds <= 0:
|
||||
raise ValueError(f"health_check_interval_seconds must be > 0, got {self.health_check_interval_seconds}")
|
||||
if self.health_check_timeout_seconds <= 0:
|
||||
raise ValueError(f"health_check_timeout_seconds must be > 0, got {self.health_check_timeout_seconds}")
|
||||
if self.failure_threshold < 1:
|
||||
raise ValueError(f"failure_threshold must be >= 1, got {self.failure_threshold}")
|
||||
if self.recovery_threshold < 1:
|
||||
raise ValueError(f"recovery_threshold must be >= 1, got {self.recovery_threshold}")
|
||||
seen_urls: set[str] = set()
|
||||
for replica in self.replicas:
|
||||
if not replica.url.startswith(("http://", "https://")):
|
||||
raise ValueError(f"ReplicaEndpoint.url must start with http:// or https://, got {replica.url!r}")
|
||||
parsed = urlparse(replica.url)
|
||||
if parsed.path not in ("", "/"):
|
||||
raise ValueError(f"ReplicaEndpoint.url must be a base host[:port] URL without a path; "
|
||||
f"got {replica.url!r} with path {parsed.path!r}. The router appends "
|
||||
"`/health` and `/v1/stream` itself.")
|
||||
if parsed.query or parsed.fragment:
|
||||
raise ValueError(f"ReplicaEndpoint.url must not include query/fragment; got {replica.url!r}")
|
||||
if replica.url in seen_urls:
|
||||
raise ValueError(f"Duplicate ReplicaEndpoint.url {replica.url!r}; "
|
||||
"router selection keys by URL so duplicates would silently collapse")
|
||||
seen_urls.add(replica.url)
|
||||
primaries = sum(1 for r in self.replicas if r.primary)
|
||||
if primaries > 1:
|
||||
raise ValueError(f"RouterConfig allows at most one primary replica; got {primaries}. "
|
||||
"Multi-primary load distribution is deferred — promote one replica to "
|
||||
"primary and treat the rest as secondaries.")
|
||||
|
||||
|
||||
__all__ = ["ReplicaEndpoint", "RouterConfig"]
|
||||
@@ -1,218 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Router FastAPI entry point.
|
||||
|
||||
Exposes the same ``/v1/stream`` WebSocket path the backend servers do,
|
||||
accepts a client, picks a healthy replica from the registry, and
|
||||
proxies frames bidirectionally.
|
||||
|
||||
PR 7.9 ships the minimum-viable shape: explicit replica list, single
|
||||
primary, JSON + binary passthrough in both directions, and a
|
||||
``/status`` endpoint for operators. Sticky-session routing (so a
|
||||
reconnect lands on the same backend) is left for a follow-up.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
from fastvideo.entrypoints.streaming.router.config import RouterConfig
|
||||
from fastvideo.entrypoints.streaming.router.registry import (
|
||||
ReplicaRegistry,
|
||||
run_health_check_loop,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _RouterState:
|
||||
config: RouterConfig
|
||||
registry: ReplicaRegistry
|
||||
stop_event: asyncio.Event
|
||||
health_task: asyncio.Task | None = None
|
||||
|
||||
|
||||
def build_router_app(
|
||||
config: RouterConfig,
|
||||
*,
|
||||
registry: ReplicaRegistry | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the router FastAPI app.
|
||||
|
||||
``registry`` can be injected for tests; defaults to one built from
|
||||
``config.replicas``.
|
||||
"""
|
||||
registry = registry or ReplicaRegistry(config.replicas)
|
||||
state = _RouterState(
|
||||
config=config,
|
||||
registry=registry,
|
||||
stop_event=asyncio.Event(),
|
||||
)
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _lifespan(_app: FastAPI):
|
||||
state.health_task = asyncio.create_task(
|
||||
run_health_check_loop(
|
||||
registry=state.registry,
|
||||
config=state.config,
|
||||
stop_event=state.stop_event,
|
||||
))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
state.stop_event.set()
|
||||
if state.health_task is not None:
|
||||
with contextlib.suppress(asyncio.CancelledError):
|
||||
await state.health_task
|
||||
|
||||
app = FastAPI(title="FastVideo Streaming Router", lifespan=_lifespan)
|
||||
|
||||
@app.get("/status")
|
||||
async def _status() -> JSONResponse:
|
||||
return JSONResponse({
|
||||
"replicas": [{
|
||||
"url": r.url,
|
||||
"primary": r.primary,
|
||||
"status": r.health.status.value,
|
||||
"last_ok_at": r.health.last_ok_at,
|
||||
"last_latency_ms": r.health.last_latency_ms,
|
||||
"consecutive_failures": r.health.consecutive_failures,
|
||||
} for r in state.registry.all()],
|
||||
})
|
||||
|
||||
@app.websocket("/v1/stream")
|
||||
async def _proxy(websocket: WebSocket) -> None:
|
||||
await websocket.accept()
|
||||
replica = state.registry.select()
|
||||
if replica is None:
|
||||
await websocket.send_json({
|
||||
"type": "error",
|
||||
"code": "gpu_unavailable",
|
||||
"message": "router: no healthy replica available",
|
||||
"retryable": True,
|
||||
})
|
||||
await websocket.close(code=1013, reason="no_healthy_replica")
|
||||
return
|
||||
|
||||
ws_url = _websocket_url_for(replica.url)
|
||||
try:
|
||||
await _bridge_session(websocket, ws_url)
|
||||
except WebSocketDisconnect:
|
||||
logger.info("router: client disconnected")
|
||||
except Exception as exc:
|
||||
logger.exception("router: bridge failed: %s", exc)
|
||||
with contextlib.suppress(RuntimeError):
|
||||
await websocket.send_json({
|
||||
"type": "error",
|
||||
"code": "worker_failed",
|
||||
"message": f"router bridge failed: {exc}",
|
||||
"retryable": True,
|
||||
})
|
||||
with contextlib.suppress(RuntimeError):
|
||||
await websocket.close(code=1011)
|
||||
|
||||
app.state.router_state = state
|
||||
return app
|
||||
|
||||
|
||||
def run_router(config: RouterConfig) -> None: # pragma: no cover - CLI
|
||||
import uvicorn
|
||||
|
||||
app = build_router_app(config)
|
||||
uvicorn.run(app, host=config.host, port=config.port)
|
||||
|
||||
|
||||
async def _bridge_session(
|
||||
client_ws: WebSocket,
|
||||
backend_ws_url: str,
|
||||
) -> None:
|
||||
"""Connect to backend and shuttle messages in both directions.
|
||||
|
||||
Uses ``websockets`` for the backend side; imported lazily to keep
|
||||
the router's import graph small for users who only want the server.
|
||||
|
||||
Cancellation: when either direction completes (client disconnect,
|
||||
backend close, exception), the other is cancelled explicitly and
|
||||
both are drained before returning. Unexpected exceptions from the
|
||||
direction that completed first are re-raised; normal disconnect
|
||||
paths (``WebSocketDisconnect``, ``ConnectionClosed``,
|
||||
``CancelledError``) are swallowed.
|
||||
"""
|
||||
try:
|
||||
import websockets
|
||||
except ImportError as exc: # pragma: no cover - optional extra
|
||||
raise RuntimeError("router requires the `websockets` package for backend proxying") from exc
|
||||
|
||||
async with websockets.connect(backend_ws_url + "/v1/stream") as backend_ws:
|
||||
c2b = asyncio.create_task(_forward_client_to_backend(client_ws, backend_ws))
|
||||
b2c = asyncio.create_task(_forward_backend_to_client(backend_ws, client_ws))
|
||||
try:
|
||||
done, _pending = await asyncio.wait(
|
||||
{c2b, b2c},
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
finally:
|
||||
for task in (c2b, b2c):
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
await asyncio.gather(c2b, b2c, return_exceptions=True)
|
||||
for task in done:
|
||||
task_exc = task.exception()
|
||||
if task_exc is not None and not _is_normal_disconnect(task_exc):
|
||||
raise task_exc
|
||||
|
||||
|
||||
def _is_normal_disconnect(exc: BaseException) -> bool:
|
||||
"""Whether ``exc`` is a routine WebSocket teardown vs a real bridge fault."""
|
||||
if isinstance(exc, asyncio.CancelledError | WebSocketDisconnect):
|
||||
return True
|
||||
name = type(exc).__name__
|
||||
# websockets.exceptions.ConnectionClosed{,OK,Error} all subclass
|
||||
# WebSocketException; check by name to avoid the lazy-import dance.
|
||||
return name.startswith("ConnectionClosed")
|
||||
|
||||
|
||||
async def _forward_client_to_backend(client_ws: WebSocket, backend_ws) -> None:
|
||||
try:
|
||||
while True:
|
||||
msg = await client_ws.receive()
|
||||
if msg.get("type") == "websocket.disconnect":
|
||||
break
|
||||
if "text" in msg and msg["text"] is not None:
|
||||
await backend_ws.send(msg["text"])
|
||||
elif "bytes" in msg and msg["bytes"] is not None:
|
||||
await backend_ws.send(msg["bytes"])
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
await backend_ws.close()
|
||||
|
||||
|
||||
async def _forward_backend_to_client(backend_ws, client_ws: WebSocket) -> None:
|
||||
try:
|
||||
async for frame in backend_ws:
|
||||
if isinstance(frame, bytes):
|
||||
await client_ws.send_bytes(frame)
|
||||
else:
|
||||
await client_ws.send_text(frame)
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
await client_ws.close()
|
||||
|
||||
|
||||
def _websocket_url_for(http_url: str) -> str:
|
||||
if http_url.startswith("https://"):
|
||||
return "wss://" + http_url[len("https://"):]
|
||||
if http_url.startswith("http://"):
|
||||
return "ws://" + http_url[len("http://"):]
|
||||
return http_url
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_router_app",
|
||||
"run_router",
|
||||
]
|
||||
@@ -1,268 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Replica registry + health-check loop.
|
||||
|
||||
The registry tracks the set of known backend replicas and their live
|
||||
health. The router consults it for "pick a backend for this session"
|
||||
decisions and a background task updates it from periodic HTTP probes.
|
||||
|
||||
State machine per replica::
|
||||
|
||||
HEALTHY ──(N consecutive failures)──▶ UNHEALTHY
|
||||
▲ │
|
||||
└──────(M consecutive successes)──────┘
|
||||
|
||||
Where N = :attr:`RouterConfig.failure_threshold` and
|
||||
M = :attr:`RouterConfig.recovery_threshold`.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextlib
|
||||
import enum
|
||||
import time
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.entrypoints.streaming.router.config import (
|
||||
ReplicaEndpoint,
|
||||
RouterConfig,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
HttpProbe = Any
|
||||
"""Structural alias for health-probe callables. Concrete signature is
|
||||
``async def __call__(url: str, *, timeout: float) -> tuple[float,
|
||||
str | None]``; typing.Callable cannot express keyword-only parameters,
|
||||
so duck-typing is the pragmatic compromise."""
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ReplicaStatus(enum.Enum):
|
||||
UNKNOWN = "unknown"
|
||||
HEALTHY = "healthy"
|
||||
UNHEALTHY = "unhealthy"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ReplicaHealth:
|
||||
status: ReplicaStatus = ReplicaStatus.UNKNOWN
|
||||
last_ok_at: float | None = None
|
||||
last_failure_at: float | None = None
|
||||
consecutive_failures: int = 0
|
||||
consecutive_successes: int = 0
|
||||
last_latency_ms: float | None = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Replica:
|
||||
endpoint: ReplicaEndpoint
|
||||
health: ReplicaHealth = field(default_factory=ReplicaHealth)
|
||||
|
||||
@property
|
||||
def url(self) -> str:
|
||||
return self.endpoint.url
|
||||
|
||||
@property
|
||||
def primary(self) -> bool:
|
||||
return self.endpoint.primary
|
||||
|
||||
@property
|
||||
def is_healthy(self) -> bool:
|
||||
return self.health.status is ReplicaStatus.HEALTHY
|
||||
|
||||
|
||||
class ReplicaRegistry:
|
||||
"""Stateful map of replica URL → :class:`Replica`.
|
||||
|
||||
Selection favors primary replicas when healthy; otherwise the first
|
||||
healthy non-primary is returned. When none are healthy, the
|
||||
registry returns ``None`` so the router can reject incoming
|
||||
sessions with ``gpu_unavailable``.
|
||||
"""
|
||||
|
||||
def __init__(self, replicas: list[ReplicaEndpoint]) -> None:
|
||||
if not replicas:
|
||||
raise ValueError("ReplicaRegistry requires at least one replica")
|
||||
self._replicas: dict[str, Replica] = {endpoint.url: Replica(endpoint=endpoint) for endpoint in replicas}
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
def all(self) -> list[Replica]:
|
||||
return list(self._replicas.values())
|
||||
|
||||
def get(self, url: str) -> Replica | None:
|
||||
return self._replicas.get(url)
|
||||
|
||||
def primaries(self) -> list[Replica]:
|
||||
return [r for r in self._replicas.values() if r.primary]
|
||||
|
||||
def select(self) -> Replica | None:
|
||||
"""Pick the best healthy replica.
|
||||
|
||||
Priority order:
|
||||
|
||||
1. The first healthy primary (insertion order).
|
||||
2. The first healthy non-primary (insertion order).
|
||||
3. ``None`` when nothing is healthy.
|
||||
|
||||
This MVP picks the first match within each tier; it does NOT
|
||||
load-balance across multiple healthy replicas of the same tier.
|
||||
Round-robin and weighted distribution are deferred until a real
|
||||
N-way active deployment exists.
|
||||
"""
|
||||
healthy_primaries = [r for r in self._replicas.values() if r.primary and r.is_healthy]
|
||||
if healthy_primaries:
|
||||
return healthy_primaries[0]
|
||||
healthy = [r for r in self._replicas.values() if r.is_healthy]
|
||||
if healthy:
|
||||
return healthy[0]
|
||||
return None
|
||||
|
||||
async def record_success(
|
||||
self,
|
||||
replica: Replica,
|
||||
*,
|
||||
recovery_threshold: int,
|
||||
latency_ms: float,
|
||||
) -> None:
|
||||
async with self._lock:
|
||||
h = replica.health
|
||||
h.last_ok_at = time.time()
|
||||
h.last_latency_ms = latency_ms
|
||||
h.consecutive_failures = 0
|
||||
h.consecutive_successes += 1
|
||||
# State machine: UNKNOWN -> HEALTHY is immediate; only the
|
||||
# UNHEALTHY -> HEALTHY transition is gated by recovery_threshold.
|
||||
if h.status is ReplicaStatus.UNKNOWN:
|
||||
logger.info("router: replica %s initial probe ok, marking HEALTHY", replica.url)
|
||||
h.status = ReplicaStatus.HEALTHY
|
||||
h.consecutive_successes = 0
|
||||
elif (h.status is ReplicaStatus.UNHEALTHY and h.consecutive_successes >= recovery_threshold):
|
||||
logger.info("router: replica %s recovered to HEALTHY after %d successes", replica.url,
|
||||
h.consecutive_successes)
|
||||
h.status = ReplicaStatus.HEALTHY
|
||||
h.consecutive_successes = 0
|
||||
|
||||
async def record_failure(
|
||||
self,
|
||||
replica: Replica,
|
||||
*,
|
||||
failure_threshold: int,
|
||||
reason: str,
|
||||
) -> None:
|
||||
async with self._lock:
|
||||
h = replica.health
|
||||
h.last_failure_at = time.time()
|
||||
h.consecutive_successes = 0
|
||||
h.consecutive_failures += 1
|
||||
if (h.status is not ReplicaStatus.UNHEALTHY and h.consecutive_failures >= failure_threshold):
|
||||
logger.warning("router: replica %s marked UNHEALTHY after %d failures: %s", replica.url,
|
||||
h.consecutive_failures, reason)
|
||||
h.status = ReplicaStatus.UNHEALTHY
|
||||
|
||||
|
||||
async def run_health_check_loop(
|
||||
registry: ReplicaRegistry,
|
||||
config: RouterConfig,
|
||||
*,
|
||||
stop_event: asyncio.Event,
|
||||
http_get: HttpProbe | None = None,
|
||||
) -> None:
|
||||
"""Poll all replicas' health endpoints in parallel on a fixed interval.
|
||||
|
||||
``http_get`` is pluggable so unit tests can inject a deterministic
|
||||
probe without hitting the network. The default builds a single
|
||||
``httpx.AsyncClient`` shared across the loop's lifetime so the
|
||||
common case (steady polling against a stable replica set) reuses
|
||||
TCP/TLS connections instead of paying handshake cost per probe.
|
||||
|
||||
Probes within one polling cycle run concurrently via ``asyncio.gather``
|
||||
so a slow replica doesn't push the cycle past
|
||||
``health_check_interval_seconds``.
|
||||
"""
|
||||
if http_get is not None:
|
||||
await _run_loop(registry, config, stop_event, http_get)
|
||||
return
|
||||
async with _build_default_probe(config) as probe:
|
||||
await _run_loop(registry, config, stop_event, probe)
|
||||
|
||||
|
||||
async def _run_loop(
|
||||
registry: ReplicaRegistry,
|
||||
config: RouterConfig,
|
||||
stop_event: asyncio.Event,
|
||||
http_get: Callable[..., Awaitable[tuple[float, str | None]]],
|
||||
) -> None:
|
||||
while not stop_event.is_set():
|
||||
replicas = registry.all()
|
||||
results = await asyncio.gather(
|
||||
*[
|
||||
http_get(replica.url + config.health_check_path, timeout=config.health_check_timeout_seconds)
|
||||
for replica in replicas
|
||||
],
|
||||
return_exceptions=True,
|
||||
)
|
||||
for replica, result in zip(replicas, results, strict=True):
|
||||
if isinstance(result, BaseException):
|
||||
await registry.record_failure(
|
||||
replica,
|
||||
failure_threshold=config.failure_threshold,
|
||||
reason=f"{type(result).__name__}: {result}",
|
||||
)
|
||||
continue
|
||||
status_ms, error = result
|
||||
if error is None:
|
||||
await registry.record_success(
|
||||
replica,
|
||||
recovery_threshold=config.recovery_threshold,
|
||||
latency_ms=status_ms,
|
||||
)
|
||||
else:
|
||||
await registry.record_failure(
|
||||
replica,
|
||||
failure_threshold=config.failure_threshold,
|
||||
reason=error,
|
||||
)
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
stop_event.wait(),
|
||||
timeout=config.health_check_interval_seconds,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def _build_default_probe(
|
||||
config: RouterConfig, ) -> AsyncIterator[Callable[..., Awaitable[tuple[float, str | None]]]]:
|
||||
try:
|
||||
import httpx
|
||||
except ImportError as exc: # pragma: no cover - optional extra
|
||||
raise RuntimeError("router health checks require httpx; install with "
|
||||
"`pip install fastvideo[streaming]` or `pip install httpx`") from exc
|
||||
|
||||
async with httpx.AsyncClient(timeout=config.health_check_timeout_seconds) as client:
|
||||
|
||||
async def probe(url: str, *, timeout: float) -> tuple[float, str | None]:
|
||||
start = time.perf_counter()
|
||||
try:
|
||||
response = await client.get(url, timeout=timeout)
|
||||
except Exception as exc:
|
||||
return 0.0, f"{type(exc).__name__}: {exc}"
|
||||
latency_ms = (time.perf_counter() - start) * 1000.0
|
||||
if response.status_code >= 400:
|
||||
return latency_ms, f"HTTP {response.status_code}"
|
||||
return latency_ms, None
|
||||
|
||||
yield probe
|
||||
|
||||
|
||||
__all__ = [
|
||||
"HttpProbe",
|
||||
"Replica",
|
||||
"ReplicaHealth",
|
||||
"ReplicaRegistry",
|
||||
"ReplicaStatus",
|
||||
"run_health_check_loop",
|
||||
]
|
||||
@@ -51,11 +51,6 @@ from fastvideo.entrypoints.streaming.session import (
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_init_image import (
|
||||
persist_session_init_image, )
|
||||
from fastvideo.entrypoints.streaming.gpu_pool import (
|
||||
GpuPool,
|
||||
InProcessGpuPool,
|
||||
PoolAcquireTimeout,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_store import (
|
||||
InMemorySessionStore,
|
||||
SessionStore,
|
||||
@@ -80,37 +75,25 @@ class _GeneratorProto(Protocol):
|
||||
@dataclass
|
||||
class ServerState:
|
||||
serve_config: ServeConfig
|
||||
pool: GpuPool
|
||||
generator: _GeneratorProto
|
||||
sessions: SessionManager
|
||||
session_store: SessionStore
|
||||
|
||||
|
||||
def build_app(
|
||||
serve_config: ServeConfig,
|
||||
generator: _GeneratorProto | None = None,
|
||||
generator: _GeneratorProto,
|
||||
*,
|
||||
pool: GpuPool | None = None,
|
||||
session_store: SessionStore | None = None,
|
||||
) -> FastAPI:
|
||||
"""Build the FastAPI app used by :func:`run_server`.
|
||||
|
||||
Exposed so tests can drive the WebSocket endpoint in-process via
|
||||
``starlette.testclient.TestClient(app).websocket_connect(...)``.
|
||||
|
||||
Exactly one of ``generator`` (backed by :class:`InProcessGpuPool`)
|
||||
or ``pool`` (for the subprocess-backed production shape) must be
|
||||
given.
|
||||
"""
|
||||
if serve_config.streaming is None:
|
||||
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
|
||||
"server; got None. Add a `streaming:` block to your serve config.")
|
||||
if (generator is None) == (pool is None):
|
||||
raise ValueError("build_app requires exactly one of `generator` or `pool`")
|
||||
|
||||
store = session_store or InMemorySessionStore()
|
||||
if pool is None:
|
||||
assert generator is not None
|
||||
pool = InProcessGpuPool(generator, session_store=store)
|
||||
|
||||
sessions = SessionManager(
|
||||
segment_cap=serve_config.streaming.generation_segment_cap,
|
||||
@@ -118,9 +101,9 @@ def build_app(
|
||||
)
|
||||
state = ServerState(
|
||||
serve_config=serve_config,
|
||||
pool=pool,
|
||||
generator=generator,
|
||||
sessions=sessions,
|
||||
session_store=store,
|
||||
session_store=session_store or InMemorySessionStore(),
|
||||
)
|
||||
|
||||
app = FastAPI(title="FastVideo Streaming")
|
||||
@@ -152,8 +135,6 @@ def build_app(
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
finally:
|
||||
with contextlib.suppress(Exception):
|
||||
await state.pool.release(session.id)
|
||||
_cleanup_session(session, state)
|
||||
|
||||
app.state.server_state = state
|
||||
@@ -197,22 +178,10 @@ async def _handle_session(
|
||||
await _apply_session_init(session, init, state)
|
||||
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
|
||||
session.transition(SessionState.GPU_BINDING)
|
||||
try:
|
||||
assignment = await state.pool.acquire(
|
||||
session.id,
|
||||
timeout=float(state.sessions.session_timeout_seconds),
|
||||
)
|
||||
except PoolAcquireTimeout as exc:
|
||||
await _send_error(websocket, "gpu_unavailable", str(exc), retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.TIMEOUT)
|
||||
return
|
||||
session.gpu_id = assignment.gpu_id
|
||||
await _send_json(websocket,
|
||||
GpuAssigned(
|
||||
gpu_id=assignment.gpu_id,
|
||||
session_timeout=state.sessions.session_timeout_seconds,
|
||||
))
|
||||
await _send_json(websocket, GpuAssigned(
|
||||
gpu_id=0,
|
||||
session_timeout=state.sessions.session_timeout_seconds,
|
||||
))
|
||||
session.transition(SessionState.ACTIVE)
|
||||
await _send_json(websocket, _build_stream_start(session, state))
|
||||
|
||||
@@ -358,13 +327,15 @@ async def _run_segment(
|
||||
))
|
||||
|
||||
start = time.perf_counter()
|
||||
# TODO: pool.run() runs to completion even if the client disconnects
|
||||
# mid-segment. Real cancellation needs the generate_async API.
|
||||
loop = asyncio.get_running_loop()
|
||||
# TODO: executor-wrapped generate() cannot be cancelled, so a
|
||||
# client disconnect mid-segment leaves the GPU work running to
|
||||
# completion. Real cancellation needs the generate_async API.
|
||||
try:
|
||||
result = await state.pool.run(session.id, request)
|
||||
result = await loop.run_in_executor(None, state.generator.generate, request)
|
||||
except Exception as exc:
|
||||
logger.exception("session %s: pool.run failed", session.id[:8])
|
||||
await _send_error(websocket, "worker_failed", f"pool.run failed: {exc}", retryable=True)
|
||||
logger.exception("session %s: generator failed", session.id[:8])
|
||||
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
|
||||
with contextlib.suppress(InvalidSessionTransition):
|
||||
session.transition(SessionState.ERROR)
|
||||
return
|
||||
|
||||
@@ -1,113 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-session JSONL event logger.
|
||||
|
||||
Each session gets its own JSONL file under the configured log root so
|
||||
post-hoc analytics (enhancer latency, GPU assignment, segment timings)
|
||||
can be recovered without a tracing backend. The internal UI uses this
|
||||
format; keeping the same shape makes log tooling portable.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, TextIO
|
||||
|
||||
_FILENAME_SANITIZE_RE = re.compile(r"[^A-Za-z0-9._-]")
|
||||
|
||||
|
||||
@dataclass
|
||||
class SessionLogEvent:
|
||||
"""One line in the session JSONL file."""
|
||||
|
||||
session_id: str
|
||||
event: str
|
||||
payload: dict[str, Any] = field(default_factory=dict)
|
||||
ts: float = field(default_factory=time.time)
|
||||
|
||||
|
||||
class SessionLogger:
|
||||
"""Append-only JSONL logger keyed by session id.
|
||||
|
||||
Thread-safe; the server may be writing from multiple asyncio tasks
|
||||
(fMP4 encoder thread + control-frame handler) for the same session.
|
||||
"""
|
||||
|
||||
def __init__(self, log_dir: str | None) -> None:
|
||||
self._log_dir = log_dir
|
||||
self._files: dict[str, TextIO] = {}
|
||||
self._locks: dict[str, threading.Lock] = {}
|
||||
self._registry_lock = threading.Lock()
|
||||
self._ensure_dir()
|
||||
|
||||
def log(self, event: SessionLogEvent) -> None:
|
||||
if self._log_dir is None:
|
||||
return
|
||||
opened = self._get_file(event.session_id)
|
||||
if opened is None:
|
||||
return
|
||||
handle, lock = opened
|
||||
line = json.dumps({
|
||||
"session_id": event.session_id,
|
||||
"event": event.event,
|
||||
"ts": event.ts,
|
||||
"payload": event.payload,
|
||||
})
|
||||
with lock, contextlib.suppress(ValueError):
|
||||
handle.write(line + "\n")
|
||||
handle.flush()
|
||||
|
||||
def close(self, session_id: str) -> None:
|
||||
with self._registry_lock:
|
||||
handle = self._files.pop(session_id, None)
|
||||
lock = self._locks.pop(session_id, None)
|
||||
if handle is None or lock is None:
|
||||
return
|
||||
with lock, contextlib.suppress(Exception):
|
||||
handle.close()
|
||||
|
||||
def close_all(self) -> None:
|
||||
with self._registry_lock:
|
||||
sids = list(self._files)
|
||||
for sid in sids:
|
||||
self.close(sid)
|
||||
|
||||
def _ensure_dir(self) -> None:
|
||||
if self._log_dir is None:
|
||||
return
|
||||
os.makedirs(self._log_dir, exist_ok=True)
|
||||
|
||||
def _get_file(self, session_id: str) -> tuple[TextIO, threading.Lock] | None:
|
||||
if self._log_dir is None:
|
||||
return None
|
||||
with self._registry_lock:
|
||||
handle = self._files.get(session_id)
|
||||
lock = self._locks.get(session_id)
|
||||
if handle is not None and lock is not None:
|
||||
return handle, lock
|
||||
# Defense-in-depth: session_id is server-generated UUID today,
|
||||
# but sanitize against path traversal in case future code paths
|
||||
# allow client-supplied ids.
|
||||
safe_id = _FILENAME_SANITIZE_RE.sub("_", session_id) or "unknown"
|
||||
path = os.path.join(
|
||||
self._log_dir,
|
||||
f"session-{safe_id}.jsonl",
|
||||
)
|
||||
try:
|
||||
handle = open(path, "a", encoding="utf-8") # noqa: SIM115
|
||||
except OSError:
|
||||
return None
|
||||
lock = threading.Lock()
|
||||
self._files[session_id] = handle
|
||||
self._locks[session_id] = lock
|
||||
return handle, lock
|
||||
|
||||
|
||||
__all__ = [
|
||||
"SessionLogEvent",
|
||||
"SessionLogger",
|
||||
]
|
||||
@@ -1,133 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Per-GPU worker subprocess entry for :class:`SubprocessGpuPool`.
|
||||
|
||||
The pool manages binding, lifecycle, and message dispatch in the parent
|
||||
process. The worker constructs its :class:`VideoGenerator` from a typed
|
||||
:class:`GeneratorConfig`, runs the two-segment warmup so both
|
||||
initial-segment and continuation-branch compile graphs are hot, and
|
||||
then loops on the job queue.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import multiprocessing as mp
|
||||
import queue
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
GeneratorConfig,
|
||||
GenerationRequest,
|
||||
InputConfig,
|
||||
OutputConfig,
|
||||
SamplingConfig,
|
||||
WarmupConfig,
|
||||
)
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Synthetic warmup dimensions: small enough to keep boot fast, big enough
|
||||
# to exercise the real shape-dependent compile paths. Keep in sync with
|
||||
# WarmupConfig if these become user-tunable.
|
||||
_WARMUP_NUM_FRAMES = 8
|
||||
_WARMUP_HEIGHT = 256
|
||||
_WARMUP_WIDTH = 256
|
||||
_WARMUP_NUM_INFERENCE_STEPS = 1
|
||||
|
||||
|
||||
def worker_main(
|
||||
*,
|
||||
gpu_id: int,
|
||||
worker_id: str,
|
||||
generator_config: GeneratorConfig,
|
||||
warmup_config: WarmupConfig,
|
||||
job_queue: mp.Queue,
|
||||
result_queue: mp.Queue,
|
||||
shutdown_event: Any,
|
||||
) -> None: # pragma: no cover - exercised via integration only
|
||||
"""Per-worker subprocess entry.
|
||||
|
||||
Runs inside the child spawned by ``SubprocessGpuPool``. Blocking
|
||||
``VideoGenerator`` construction + generation happens here, not in
|
||||
the parent's event loop.
|
||||
"""
|
||||
import os
|
||||
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = str(gpu_id)
|
||||
try:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
generator = VideoGenerator.from_pretrained(config=generator_config)
|
||||
if warmup_config.enabled:
|
||||
_warmup_worker(generator, warmup_config)
|
||||
result_queue.put({"kind": "ready", "worker_id": worker_id})
|
||||
except Exception as exc:
|
||||
result_queue.put({"kind": "error", "error": repr(exc)})
|
||||
return
|
||||
|
||||
while not shutdown_event.is_set():
|
||||
try:
|
||||
item = job_queue.get(timeout=0.5)
|
||||
except queue.Empty:
|
||||
continue
|
||||
if item is None:
|
||||
break
|
||||
job_id = item["job_id"]
|
||||
request = item["request"]
|
||||
try:
|
||||
result = generator.generate(request)
|
||||
result_queue.put({
|
||||
"kind": "result",
|
||||
"job_id": job_id,
|
||||
"result": result,
|
||||
})
|
||||
except Exception as exc:
|
||||
result_queue.put({
|
||||
"kind": "error",
|
||||
"job_id": job_id,
|
||||
"error": repr(exc),
|
||||
})
|
||||
|
||||
|
||||
def _warmup_worker(
|
||||
generator: Any,
|
||||
warmup_config: WarmupConfig,
|
||||
) -> None:
|
||||
"""Run two synthetic generations so both compile branches are primed.
|
||||
|
||||
Segment 1 is a fresh start (no continuation state) and exercises
|
||||
the initial-segment graph. Segment 2 feeds segment 1's continuation
|
||||
state back in so the conditioning branch is also compiled before
|
||||
the first user request lands.
|
||||
"""
|
||||
sampling = SamplingConfig(
|
||||
num_frames=_WARMUP_NUM_FRAMES,
|
||||
height=_WARMUP_HEIGHT,
|
||||
width=_WARMUP_WIDTH,
|
||||
num_inference_steps=_WARMUP_NUM_INFERENCE_STEPS,
|
||||
)
|
||||
seg1 = GenerationRequest(
|
||||
prompt=warmup_config.prompt,
|
||||
sampling=sampling,
|
||||
inputs=InputConfig(),
|
||||
output=OutputConfig(save_video=False, return_frames=False, return_state=True),
|
||||
)
|
||||
seg1_result = generator.generate(seg1)
|
||||
|
||||
seg2 = GenerationRequest(
|
||||
prompt=warmup_config.prompt,
|
||||
sampling=sampling,
|
||||
inputs=InputConfig(),
|
||||
output=OutputConfig(save_video=False, return_frames=False),
|
||||
state=_extract_continuation_state(seg1_result),
|
||||
)
|
||||
generator.generate(seg2)
|
||||
|
||||
|
||||
def _extract_continuation_state(result: Any) -> Any:
|
||||
state = getattr(result, "state", None)
|
||||
if state is None and isinstance(result, dict):
|
||||
state = result.get("state")
|
||||
return state
|
||||
|
||||
|
||||
__all__ = ["worker_main"]
|
||||
@@ -65,7 +65,6 @@ _FROM_PRETRAINED_CONVENIENCE_KWARGS = frozenset({
|
||||
"pin_cpu_memory",
|
||||
"enable_torch_compile",
|
||||
"torch_compile_kwargs",
|
||||
"output_type",
|
||||
})
|
||||
|
||||
|
||||
@@ -602,20 +601,10 @@ class VideoGenerator:
|
||||
thread = threading.Thread(target=execute_forward_thread)
|
||||
thread.start()
|
||||
latent_batch_size = _infer_latent_batch_size(batch)
|
||||
# When ``output_type == "latent"`` the forward output has latent
|
||||
# shape (e.g. ``[B, C_latent, T_latent, H_latent, W_latent]``)
|
||||
# rather than the pre-allocation's pixel shape. Skip the pinned
|
||||
# ~50 MB buffer entirely; we always fall through to the
|
||||
# ``samples = output_batch.output.cpu()`` branch below in that
|
||||
# mode. ``skip_pixel_prealloc`` also gates the slow-path warning.
|
||||
skip_pixel_prealloc = fastvideo_args.output_type == "latent"
|
||||
if skip_pixel_prealloc:
|
||||
samples = torch.empty(0, device='cpu')
|
||||
else:
|
||||
samples = torch.empty(
|
||||
(latent_batch_size, 3, sampling_param.num_frames, sampling_param.height, sampling_param.width),
|
||||
device='cpu',
|
||||
pin_memory=fastvideo_args.pin_cpu_memory)
|
||||
samples = torch.empty(
|
||||
(latent_batch_size, 3, sampling_param.num_frames, sampling_param.height, sampling_param.width),
|
||||
device='cpu',
|
||||
pin_memory=fastvideo_args.pin_cpu_memory)
|
||||
thread.join()
|
||||
|
||||
if thread_error["error"] is not None:
|
||||
@@ -630,44 +619,29 @@ class VideoGenerator:
|
||||
if output_batch.output.shape == samples.shape:
|
||||
samples.copy_(output_batch.output)
|
||||
else:
|
||||
if not skip_pixel_prealloc:
|
||||
logger.warning("Output shape %s does not match expected shape %s; use slow path",
|
||||
output_batch.output.shape, samples.shape)
|
||||
logger.warning("Output shape %s does not match expected shape %s; use slow path", output_batch.output.shape,
|
||||
samples.shape)
|
||||
samples = output_batch.output.cpu()
|
||||
logging_info = output_batch.logging_info
|
||||
|
||||
gen_time = time.perf_counter() - start_time
|
||||
logger.info("Generated successfully in %.2f seconds", gen_time)
|
||||
|
||||
# Three mutually-exclusive output modes determine whether (a) we
|
||||
# build an RGB frame buffer and (b) what file we write to disk:
|
||||
#
|
||||
# 1. `output_type == "latent"` — VAE is bypassed in DecodingStage
|
||||
# and `samples` holds raw latents (arbitrary channel count).
|
||||
# The RGB grid / uint8 / mp4 / png pipeline below cannot
|
||||
# consume those, so we skip it entirely and let callers work
|
||||
# with the latent tensor directly via `result["samples"]`.
|
||||
# 2. Audio-only workload — `samples` is a 1×3×1×8×8 placeholder
|
||||
# no caller will use; skip the grid loop and save a `.wav`.
|
||||
# 3. Pixel video / image — the historical happy path.
|
||||
is_latent_output = fastvideo_args.output_type == "latent"
|
||||
# Process outputs (skip the make_grid loop for audio-only, where
|
||||
# `samples` is a 1×3×1×8×8 placeholder no caller will use).
|
||||
audio_only = bool(output_batch.extra.get("audio_only"))
|
||||
|
||||
frames: list[np.ndarray] | None
|
||||
if is_latent_output or audio_only:
|
||||
frames = None if is_latent_output else []
|
||||
else:
|
||||
frames: list[np.ndarray] = []
|
||||
if not audio_only:
|
||||
videos = rearrange(samples, "b c t h w -> t b c h w")
|
||||
frames = []
|
||||
for x in videos:
|
||||
x = torchvision.utils.make_grid(x, nrow=6)
|
||||
x = x.permute(1, 2, 0).squeeze(-1)
|
||||
x = (x * 255).to(torch.uint8)
|
||||
frames.append(x.cpu().numpy())
|
||||
|
||||
save_to_disk = batch.save_video and not is_latent_output
|
||||
if save_to_disk:
|
||||
if audio_only:
|
||||
# Save output if requested
|
||||
if batch.save_video:
|
||||
if output_batch.extra.get("audio_only"):
|
||||
# Audio-only workload: write a standalone .wav rather than
|
||||
# muxing the audio into a placeholder mp4 (which forces
|
||||
# ffmpeg to round 8x8 placeholder frames up to 16x16).
|
||||
@@ -680,11 +654,9 @@ class VideoGenerator:
|
||||
logger.info("Saved audio to %s", output_path)
|
||||
elif self._is_image_workload():
|
||||
# Image workloads (t2i, i2i, …): save the first frame as PNG.
|
||||
assert frames is not None # implied by save_to_disk and not audio_only
|
||||
imageio.imwrite(output_path, frames[0])
|
||||
logger.info("Saved image to %s", output_path)
|
||||
else:
|
||||
assert frames is not None # implied by save_to_disk and not audio_only
|
||||
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
|
||||
logger.info("Saved video to %s", output_path)
|
||||
audio = output_batch.extra.get("audio")
|
||||
@@ -708,7 +680,7 @@ class VideoGenerator:
|
||||
"trajectory": output_batch.trajectory_latents,
|
||||
"trajectory_timesteps": output_batch.trajectory_timesteps,
|
||||
"trajectory_decoded": output_batch.trajectory_decoded,
|
||||
"video_path": output_path if save_to_disk else None,
|
||||
"video_path": output_path if batch.save_video else None,
|
||||
"peak_memory_mb": output_batch.extra.get("peak_memory_mb"),
|
||||
}
|
||||
|
||||
@@ -787,7 +759,7 @@ class VideoGenerator:
|
||||
import av
|
||||
except ImportError:
|
||||
logger.warning("PyAV not installed; cannot mux audio. "
|
||||
"Install with: uv pip install av")
|
||||
"Install with: pip install av")
|
||||
return False
|
||||
|
||||
try:
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
# Layer Guidance For Model Ports
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Use this file when adding FastVideo-native model components. Keep it generic:
|
||||
model-specific parameter mappings belong in `scripts/checkpoint_conversion/`, not
|
||||
in this directory.
|
||||
|
||||
## Linear Layers
|
||||
|
||||
- Use `ReplicatedLinear` for DiT and VAE hot paths when the layer is not tensor
|
||||
parallel and should expose a normal `weight`/`bias` state-dict surface.
|
||||
- Use `QKVParallelLinear` for LLM-style fused query/key/value projections when
|
||||
the existing encoder pattern already expects tensor parallel loading.
|
||||
- Use `MergedColumnParallelLinear` for fused MLP gate/up projections that are
|
||||
loaded as packed column shards.
|
||||
- Use `ColumnParallelLinear` and `RowParallelLinear` for tensor-parallel encoder
|
||||
blocks that follow existing `t5.py`, `clip.py`, `llama.py`, or `qwen2_5.py`
|
||||
patterns.
|
||||
- Do not replace a simple official layer with a fused FastVideo layer unless the
|
||||
conversion script explicitly handles the resulting key and tensor layout.
|
||||
|
||||
## Attention Layers
|
||||
|
||||
- Use `DistributedAttention` for standard DiT full-sequence attention when the
|
||||
model should participate in sequence parallel execution.
|
||||
- Use `LocalAttention` for local/window attention or narrow single-GPU parity
|
||||
paths that match existing component style.
|
||||
- Raw `torch.nn.functional.scaled_dot_product_attention` is acceptable for
|
||||
unusual cross-modality flat streams when no FastVideo distributed primitive
|
||||
matches yet. Document the sequence-parallel gap in the owning model file.
|
||||
|
||||
## State-Dict Surface
|
||||
|
||||
- Prototype the native component before writing conversion mappings. The
|
||||
prototype's `state_dict()` is the source of truth for FastVideo target keys and
|
||||
shapes.
|
||||
- Conversion scripts should map official keys into the native state-dict surface;
|
||||
production model code should not be contorted to match checkpoint naming.
|
||||
- Fused and packed FastVideo layers may require tensor split/fuse logic in the
|
||||
converter, especially QKV/KV projections and gated MLP projections.
|
||||
- Record intentional skipped keys in the conversion script with a reason, such
|
||||
as training-only EMA/logvar/optimizer state or dynamically computed buffers.
|
||||
|
||||
## Porting Discipline
|
||||
|
||||
- Match the official layer definition and the official instantiation arguments.
|
||||
A reusable class with different constructor args is not reused.
|
||||
- Keep architecture constants on the component arch config. Runtime sampling,
|
||||
guidance, precision, and pipeline defaults belong on pipeline config or
|
||||
presets.
|
||||
- Prefer small, direct implementations until parity passes. Add helpers only
|
||||
when they serve multiple call sites or make the mapping clearer.
|
||||
@@ -1,55 +0,0 @@
|
||||
# `fastvideo/models/` — Model Implementations
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
DiT / VAE / encoder / scheduler / upsampler / audio model classes. **Pre-commit excludes this directory** — yapf/ruff/mypy do not run on commits here. Match neighboring file style manually.
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
models/
|
||||
├── dits/
|
||||
│ ├── <model>.py # Single-file DiT (wanvideo, ltx2, hunyuanvideo, cosmos, ...)
|
||||
│ ├── hyworld/ # Multi-file DiT family
|
||||
│ ├── lingbotworld/ # ditto
|
||||
│ └── matrixgame/ # ditto
|
||||
├── vaes/ # AutoencoderKL variants per model family
|
||||
├── encoders/ # T5, CLIP, Llama, Qwen2.5, Gemma, SigLIP, Reason1, audio conditioner
|
||||
├── schedulers/ # FlowMatch / EulerDiscrete / DPM custom schedulers
|
||||
├── upsamplers/ # Hunyuan15 super-resolution
|
||||
├── audio/ # Audio-VAE/decoder modules (LTX-2 audio, Stable Audio)
|
||||
├── camera/ # Camera-conditioning modules (Gen3C)
|
||||
└── loader/ # component_loader.py, fsdp_load.py, weight_utils.py
|
||||
```
|
||||
|
||||
`loader/component_loader.py` is the central entry point that the pipeline uses
|
||||
to instantiate model components from a HF directory. New components plug in
|
||||
through `register_*` calls or by extending the `ComponentLoader` mappings.
|
||||
|
||||
## Adding a Model Component (DiT / VAE / Encoder)
|
||||
|
||||
1. Read `fastvideo/layers/AGENTS.md` first — it defines which tensor-parallel
|
||||
linear / attention layer to use. Do not freelance.
|
||||
2. Define the arch in `models/<role>/<model>.py`. Mirror the official reference's
|
||||
constructor args; do not "improve" the layer choices.
|
||||
3. Add the matching arch config in `configs/models/<role>/<model>.py`.
|
||||
4. Expose `param_names_mapping` on the config — it is the **source of truth** for
|
||||
the converter under `scripts/checkpoint_conversion/`.
|
||||
5. Use `init_logger(__name__)`, not stdlib logging.
|
||||
|
||||
## State-Dict Discipline
|
||||
|
||||
- The native component's `state_dict()` defines target keys + shapes.
|
||||
- Conversion scripts (`scripts/checkpoint_conversion/`) bend to the model, not
|
||||
the other way around.
|
||||
- Fused QKV / packed MLP layouts must be documented in the config or the model
|
||||
module — converters need to split/fuse accordingly.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Importing `transformers` / `diffusers` model classes at runtime inside the
|
||||
forward path — these belong in the loader, not the architecture file.
|
||||
- Adding training-only state (EMA buffers, optimizer state) to the inference
|
||||
state-dict surface.
|
||||
- Calling `torch.distributed` directly. Go through `fastvideo.distributed`.
|
||||
- Treating this directory as lint-clean. It isn't (see pre-commit excludes).
|
||||
@@ -183,7 +183,7 @@ def _load_video_with_ffmpeg(
|
||||
except AttributeError as e:
|
||||
raise AttributeError(
|
||||
"Unable to find an ffmpeg installation on your machine. "
|
||||
"Please install via `uv pip install imageio-ffmpeg`") from e
|
||||
"Please install via `pip install imageio-ffmpeg`") from e
|
||||
|
||||
pil_images = []
|
||||
original_fps = None
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
# `fastvideo/pipelines/` — Pipeline Composition
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
Diffusion pipelines are **compositions of `PipelineStage` objects**. Each stage owns one verb (validate / encode / schedule / denoise / decode). Adding a model means assembling stages, not subclassing a megapipeline.
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
pipelines/
|
||||
├── pipeline_batch_info.py # ForwardBatch — the dict passed between stages
|
||||
├── lora_pipeline.py # LoRA-aware base
|
||||
├── composed_pipeline_base.py # Base for stage-composed pipelines
|
||||
├── stages/ # Reusable stage implementations (~30 files)
|
||||
│ ├── base.py # PipelineStage ABC + StageVerificationError
|
||||
│ ├── input_validation.py # Validates ForwardBatch shape/keys
|
||||
│ ├── text_encoding.py # Generic prompt encoder stage
|
||||
│ ├── image_encoding.py # Image conditioning
|
||||
│ ├── latent_preparation.py # Init noise + scheduler
|
||||
│ ├── conditioning.py # CFG / negative prompt fan-out
|
||||
│ ├── denoising.py # Standard diffusion loop
|
||||
│ ├── sd35_conditioning.py # Per-model overrides (named by family)
|
||||
│ ├── longcat_*.py # LongCat I2V/V2V/refine variants
|
||||
│ ├── gen3c_stages.py # Gen3C-specific stages
|
||||
│ ├── gamecraft_denoising.py # GameCraft-specific
|
||||
│ └── matrixgame_denoising.py # MatrixGame-specific
|
||||
├── basic/ # Per-model end-to-end pipelines
|
||||
│ ├── hunyuan/, hunyuan15/, hyworld/, gamecraft/, gen3c/, cosmos/
|
||||
│ ├── wan/, longcat/, ltx2/, lingbotworld/, magi_human/, matrixgame/
|
||||
│ ├── sd35/, stable_audio/, turbodiffusion/
|
||||
│ └── <model>/{<model>_pipeline.py, presets.py, __init__.py}
|
||||
├── preprocess/ # Data preprocessing pipelines (ltx2, wan, matrixgame)
|
||||
└── training/ # Training-time pipeline glue
|
||||
```
|
||||
|
||||
## Stage Authoring Rules
|
||||
|
||||
- Subclass `PipelineStage` from `stages/base.py`. Implement `forward(batch, args) -> ForwardBatch`.
|
||||
- Implement `verify_input` / `verify_output` — both return `VerificationResult`. Failures raise `StageVerificationError`.
|
||||
- Mutate `ForwardBatch` only by reassigning fields you declared in `pipeline_batch_info.py`. New keys → add to the dataclass first.
|
||||
- Stages must be **deterministic given the same `ForwardBatch + FastVideoArgs`**. Side effects (logging, profiling) only.
|
||||
- Read all knobs from the passed-in `FastVideoArgs` / `PipelineConfig`. Never `os.getenv` directly.
|
||||
|
||||
## Per-Model Pipeline Pattern (`basic/<model>/`)
|
||||
|
||||
Every model directory has the same skeleton:
|
||||
|
||||
```
|
||||
basic/<model>/
|
||||
├── __init__.py
|
||||
├── <model>_pipeline.py # Composes stages list
|
||||
├── presets.py # Default PipelineConfig + SamplingParam combos
|
||||
└── (optional) stage_overrides.py, continuation.py, ...
|
||||
```
|
||||
|
||||
`presets.py` is the entry point that `registry.py` imports — it must export the named preset constants used elsewhere in the codebase.
|
||||
|
||||
## Forking vs Reusing a Stage
|
||||
|
||||
Reuse `stages/text_encoding.py` if your model takes text → embeddings via a standard encoder. Fork only when:
|
||||
|
||||
- The model needs a **different ForwardBatch shape** (extra inputs, different output keys).
|
||||
- The denoising loop has structural differences (causal, refine-then-denoise, multi-stream).
|
||||
|
||||
When forking, keep the file name model-prefixed (`longcat_*`, `gamecraft_*`) so the registry stays grep-able.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Putting a full pipeline in a single file under `basic/<model>/` instead of composing stages.
|
||||
- Reading config from globals or env vars inside a stage.
|
||||
- Adding cross-stage state via module-level dicts. Use `ForwardBatch`.
|
||||
@@ -37,7 +37,7 @@ def load_moge_model(
|
||||
from moge.model.v1 import MoGeModel
|
||||
except ImportError as exc:
|
||||
raise ImportError("MoGe is required for GEN3C 3D cache conditioning. "
|
||||
"Install it with: uv pip install git+https://github.com/microsoft/MoGe.git. "
|
||||
"Install it with: pip install git+https://github.com/microsoft/MoGe.git. "
|
||||
"If import fails with libGL.so.1, install system deps: "
|
||||
"sudo apt-get install -y libgl1 libglib2.0-0 libsm6 libxext6 libxrender1") from exc
|
||||
|
||||
|
||||
@@ -33,15 +33,6 @@ class StableAudioDecodingStage(PipelineStage):
|
||||
pc = fastvideo_args.pipeline_config
|
||||
latents = batch.latents
|
||||
|
||||
# Latent regression path: hand back the un-decoded denoised latent
|
||||
# so the LatentSimilarityUtils harness can compare on pre-VAE
|
||||
# numerics. Mirrors the bypass in `pipelines/stages/decoding.py`
|
||||
# used by the video DiTs. Skips the `.to(device)` VAE move so we
|
||||
# don't pay decoder load cost on this shortcut path.
|
||||
if fastvideo_args.output_type == "latent":
|
||||
batch.output = latents.detach().cpu()
|
||||
return batch
|
||||
|
||||
# VAE may be CPU-parked under `vae_cpu_offload=True`.
|
||||
from fastvideo.distributed.parallel_state import get_local_torch_device
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
@@ -195,7 +195,7 @@ class CudaPlatformBase(Platform):
|
||||
except ImportError as e:
|
||||
logger.error("Failed to import SageSLA Attention backend: %s", str(e))
|
||||
raise ImportError("SageSLA Attention backend requires spas_sage_attn. "
|
||||
"Install with: uv pip install git+https://github.com/thu-ml/SpargeAttn.git") from e
|
||||
"Install with: pip install git+https://github.com/thu-ml/SpargeAttn.git") from e
|
||||
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "fastvideo.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
@@ -1,74 +0,0 @@
|
||||
# `fastvideo/tests/` — Package-Level Test Suite
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
> **Pre-commit excludes `fastvideo/tests/`.** Lint/format hooks do not run on
|
||||
> files here. Match style of neighboring tests manually.
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
tests/
|
||||
├── conftest.py # distributed_setup fixture (1×1 SP/TP init+cleanup)
|
||||
├── utils.py # Shared test helpers
|
||||
├── api/ # api/ schema + presets
|
||||
├── attention/ # Backend selector + layer parity
|
||||
├── dataset/ # Dataloader smoke tests
|
||||
├── distributed/ # SP / TP collectives
|
||||
├── encoders/ # Per-encoder parity (t5, clip, llama, qwen, ...)
|
||||
├── entrypoints/ # CLI + streaming + OpenAI-compatible server
|
||||
│ └── streaming/ # Streaming-specific tests
|
||||
├── hooks/ # Runtime hook system
|
||||
├── inference/ # End-to-end inference smoke
|
||||
├── lora_extraction/ # LoRA merge/extract round-trip
|
||||
├── modal/ # Modal CI orchestrators (ssim_test.py, nightly_test.py)
|
||||
├── nightly/ # Long-running suites (gated)
|
||||
├── ops/ # Custom kernel / op tests
|
||||
├── performance/ # Throughput / memory benchmarks (informational)
|
||||
├── ssim/ # GPU SSIM regressions — see ssim/AGENTS.md
|
||||
├── stages/ # Per-stage unit tests
|
||||
├── train/ # New modular trainer tests
|
||||
├── training/ # Legacy training pipeline tests
|
||||
├── transformers/ # transformers-shim tests
|
||||
├── vaes/ # Per-VAE encode/decode parity
|
||||
└── workflow/ # Preprocessing workflow tests
|
||||
```
|
||||
|
||||
## Run Commands (model-domain order)
|
||||
|
||||
```bash
|
||||
pytest fastvideo/tests/ -v # All package tests
|
||||
pytest fastvideo/tests/encoders/ -v # One domain
|
||||
pytest fastvideo/tests/ssim/ -vs # SSIM (GPU-heavy; see ssim/AGENTS.md)
|
||||
pytest fastvideo/tests/ssim/ -vs --ssim-full-quality # Full-quality SSIM params
|
||||
pytest tests/ -v # Top-level repo tests (different scope)
|
||||
modal run fastvideo/tests/modal/ssim_test.py # Orchestrated CI-style SSIM run
|
||||
```
|
||||
|
||||
`tests/local_tests/` (top-level) holds component checks that need a local
|
||||
working tree but are not part of the package suite.
|
||||
|
||||
## Conventions
|
||||
|
||||
- Name files `test_<feature>_<expected_behavior>.py`. Place near the domain
|
||||
(`tests/encoders/test_t5_*.py`, not `tests/test_everything.py`).
|
||||
- Use the `distributed_setup` fixture for any test that touches
|
||||
`fastvideo.distributed.*`. It seeds torch + numpy and tears down the PG.
|
||||
- GPU tests must `pytest.skip(...)` when hardware is missing — never `xfail`.
|
||||
Document `REQUIRED_GPUS` near the top for SSIM tests (see ssim/AGENTS.md).
|
||||
- `nightly/` and `performance/` tests should be guarded by an environment marker
|
||||
so the default `pytest fastvideo/tests/` stays fast.
|
||||
|
||||
## Modal CI
|
||||
|
||||
`tests/modal/ssim_test.py` is the orchestrator that auto-discovers
|
||||
`test_*.py` files under `ssim/` and schedules one subprocess per `*_MODEL_TO_PARAMS`
|
||||
key. New SSIM tests need no CI wiring — just declare `REQUIRED_GPUS = N`.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Adding a hard-coded path to a model checkpoint without a corresponding
|
||||
`pytest.skip` when the path is missing.
|
||||
- Forgetting `cleanup_dist_env_and_memory()` in tests that bypass the
|
||||
`distributed_setup` fixture — tears the next test in the worker.
|
||||
- Putting Modal-specific code outside `tests/modal/`.
|
||||
@@ -1,261 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for PR 7.8 streaming auxiliaries.
|
||||
|
||||
Covers:
|
||||
|
||||
* PromptSafetyFilter gracefully disables when fastText isn't installed
|
||||
* SafetyResult semantics (allow, block, unavailable)
|
||||
* RewriteOptions + _split_response parsing behavior
|
||||
* SessionLogger JSONL append semantics + close lifecycle
|
||||
* MockServer builds an app that drives the WS protocol end-to-end
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.safety import (
|
||||
PromptSafetyFilter,
|
||||
SafetyDecision,
|
||||
first_blocked,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.prompt.rewrite import (
|
||||
RewriteOptions,
|
||||
_split_response,
|
||||
build_rewrite,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.session_logger import (
|
||||
SessionLogEvent,
|
||||
SessionLogger,
|
||||
)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Safety
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPromptSafetyFilter:
|
||||
|
||||
def test_disabled_by_default_when_no_path(self):
|
||||
f = PromptSafetyFilter(classifier_path=None)
|
||||
assert f.enabled is False
|
||||
result = f.classify("hi")
|
||||
assert result.decision is SafetyDecision.UNAVAILABLE
|
||||
|
||||
def test_disabled_when_enabled_false(self):
|
||||
f = PromptSafetyFilter(classifier_path="/tmp/m.bin", enabled=False)
|
||||
assert f.enabled is False
|
||||
|
||||
def test_unavailable_when_fasttext_missing(self, monkeypatch):
|
||||
# Force `import fasttext` inside _ensure_loaded to fail.
|
||||
monkeypatch.setitem(sys.modules, "fasttext", None)
|
||||
f = PromptSafetyFilter(classifier_path="/tmp/m.bin", enabled=True)
|
||||
result = f.classify("hi")
|
||||
assert result.decision is SafetyDecision.UNAVAILABLE
|
||||
|
||||
def test_block_when_classifier_flags_unsafe(self, monkeypatch, tmp_path):
|
||||
fake_model = types.SimpleNamespace(
|
||||
predict=lambda text, k=1: (["__label__unsafe"], [0.95]))
|
||||
stub = types.SimpleNamespace(load_model=lambda _p: fake_model)
|
||||
monkeypatch.setitem(sys.modules, "fasttext", stub)
|
||||
model_path = str(tmp_path / "m.bin")
|
||||
Path(model_path).write_text("")
|
||||
f = PromptSafetyFilter(classifier_path=model_path, enabled=True)
|
||||
result = f.classify("please")
|
||||
assert result.decision is SafetyDecision.BLOCK
|
||||
assert result.label == "unsafe"
|
||||
assert result.score == pytest.approx(0.95)
|
||||
|
||||
def test_allow_when_classifier_flags_safe(self, monkeypatch, tmp_path):
|
||||
fake_model = types.SimpleNamespace(
|
||||
predict=lambda text, k=1: (["__label__safe"], [0.99]))
|
||||
stub = types.SimpleNamespace(load_model=lambda _p: fake_model)
|
||||
monkeypatch.setitem(sys.modules, "fasttext", stub)
|
||||
f = PromptSafetyFilter(classifier_path="ignored", enabled=True)
|
||||
result = f.classify("hello")
|
||||
assert result.decision is SafetyDecision.ALLOW
|
||||
|
||||
def test_below_threshold_allows_even_if_unsafe_label(self, monkeypatch):
|
||||
fake_model = types.SimpleNamespace(
|
||||
predict=lambda text, k=1: (["__label__unsafe"], [0.3]))
|
||||
stub = types.SimpleNamespace(load_model=lambda _p: fake_model)
|
||||
monkeypatch.setitem(sys.modules, "fasttext", stub)
|
||||
f = PromptSafetyFilter(
|
||||
classifier_path="m", enabled=True, block_threshold=0.5)
|
||||
assert f.classify("x").decision is SafetyDecision.ALLOW
|
||||
|
||||
def test_first_blocked_returns_first_hit(self, monkeypatch):
|
||||
responses = iter([
|
||||
(["__label__safe"], [0.9]),
|
||||
(["__label__unsafe"], [0.9]),
|
||||
(["__label__safe"], [0.9]),
|
||||
])
|
||||
fake_model = types.SimpleNamespace(
|
||||
predict=lambda text, k=1: next(responses))
|
||||
stub = types.SimpleNamespace(load_model=lambda _p: fake_model)
|
||||
monkeypatch.setitem(sys.modules, "fasttext", stub)
|
||||
f = PromptSafetyFilter(classifier_path="m", enabled=True)
|
||||
blocked = first_blocked(f, ["ok", "bad", "also ok"])
|
||||
assert blocked is not None
|
||||
assert blocked.prompt == "bad"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Rewrite
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestRewriteSplit:
|
||||
|
||||
def test_plain_lines(self):
|
||||
assert _split_response("one\ntwo\nthree", limit=3) == [
|
||||
"one", "two", "three",
|
||||
]
|
||||
|
||||
def test_numbered_list(self):
|
||||
assert _split_response("1. first\n2. second", limit=3) == [
|
||||
"first", "second",
|
||||
]
|
||||
|
||||
def test_bulleted_list(self):
|
||||
assert _split_response("- one\n* two\n• three", limit=3) == [
|
||||
"one", "two", "three",
|
||||
]
|
||||
|
||||
def test_respects_limit(self):
|
||||
assert _split_response("a\nb\nc\nd", limit=2) == ["a", "b"]
|
||||
|
||||
def test_limit_min_one(self):
|
||||
assert _split_response("only one", limit=0) == ["only one"]
|
||||
|
||||
|
||||
class _StubEnhancer:
|
||||
|
||||
async def rewrite(self, seed):
|
||||
from fastvideo.entrypoints.streaming.prompt.providers.base import LLMResponse
|
||||
|
||||
return LLMResponse(
|
||||
content="1. alpha\n2. beta\n3. gamma",
|
||||
provider="stub",
|
||||
model="m",
|
||||
latency_ms=1.0,
|
||||
)
|
||||
|
||||
|
||||
class TestBuildRewrite:
|
||||
|
||||
def test_empty_seed_rejected(self):
|
||||
import asyncio
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
asyncio.run(build_rewrite(_StubEnhancer(), " "))
|
||||
|
||||
def test_returns_limited_alternatives(self):
|
||||
import asyncio
|
||||
|
||||
result = asyncio.run(build_rewrite(
|
||||
_StubEnhancer(), "seed",
|
||||
options=RewriteOptions(count=2)))
|
||||
assert result.seed_prompt == "seed"
|
||||
assert result.alternatives == ["alpha", "beta"]
|
||||
assert result.provider == "stub"
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Session logger
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestSessionLogger:
|
||||
|
||||
def test_no_log_dir_is_noop(self):
|
||||
logger = SessionLogger(None)
|
||||
logger.log(SessionLogEvent(session_id="s", event="x")) # no raise
|
||||
|
||||
def test_appends_jsonl(self, tmp_path):
|
||||
logger = SessionLogger(str(tmp_path))
|
||||
logger.log(SessionLogEvent(
|
||||
session_id="s1",
|
||||
event="start",
|
||||
payload={"preset": "ltx2"},
|
||||
ts=1.0,
|
||||
))
|
||||
logger.log(SessionLogEvent(
|
||||
session_id="s1",
|
||||
event="segment",
|
||||
payload={"idx": 0},
|
||||
ts=2.0,
|
||||
))
|
||||
logger.close("s1")
|
||||
path = tmp_path / "session-s1.jsonl"
|
||||
lines = path.read_text().splitlines()
|
||||
assert len(lines) == 2
|
||||
first = json.loads(lines[0])
|
||||
assert first["event"] == "start"
|
||||
assert first["payload"]["preset"] == "ltx2"
|
||||
|
||||
def test_separate_files_per_session(self, tmp_path):
|
||||
logger = SessionLogger(str(tmp_path))
|
||||
logger.log(SessionLogEvent(session_id="a", event="e"))
|
||||
logger.log(SessionLogEvent(session_id="b", event="e"))
|
||||
logger.close_all()
|
||||
assert (tmp_path / "session-a.jsonl").exists()
|
||||
assert (tmp_path / "session-b.jsonl").exists()
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Mock server
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestMockServer:
|
||||
|
||||
def test_build_mock_app_returns_fastapi(self):
|
||||
from fastapi import FastAPI
|
||||
|
||||
from fastvideo.entrypoints.streaming.mock_server import build_mock_app
|
||||
|
||||
app = build_mock_app()
|
||||
assert isinstance(app, FastAPI)
|
||||
|
||||
def test_mock_generator_produces_frames(self):
|
||||
from fastvideo.api.schema import GenerationRequest, SamplingConfig
|
||||
from fastvideo.entrypoints.streaming.mock_server import MockGenerator
|
||||
|
||||
gen = MockGenerator()
|
||||
result = gen.generate(GenerationRequest(
|
||||
prompt="x",
|
||||
sampling=SamplingConfig(
|
||||
num_frames=3, height=32, width=32, num_inference_steps=1),
|
||||
))
|
||||
assert len(result["frames"]) == 3
|
||||
assert result["frames"][0].shape == (32, 32, 3)
|
||||
assert result["state"].kind == "ltx2.v1"
|
||||
|
||||
def test_mock_app_health_endpoint(self):
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from fastvideo.entrypoints.streaming.mock_server import build_mock_app
|
||||
|
||||
app = build_mock_app()
|
||||
client = TestClient(app)
|
||||
assert client.get("/health").json()["status"] == "ok"
|
||||
|
||||
def test_mock_app_ws_handshake(self):
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from fastvideo.entrypoints.streaming.mock_server import build_mock_app
|
||||
|
||||
app = build_mock_app()
|
||||
client = TestClient(app)
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
ws.send_json({"type": "session_init_v2"})
|
||||
assert ws.receive_json()["type"] == "queue_status"
|
||||
assert ws.receive_json()["type"] == "gpu_assigned"
|
||||
assert ws.receive_json()["type"] == "ltx2_stream_start"
|
||||
@@ -1,463 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GPU pool tests.
|
||||
|
||||
InProcessGpuPool is exercised end-to-end. SubprocessGpuPool is driven
|
||||
with an injected ``worker_factory`` that stands up a fake worker inside
|
||||
a thread (not a subprocess) so the test suite stays CPU-only.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import multiprocessing as mp
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
GeneratorConfig,
|
||||
GenerationRequest,
|
||||
GpuPoolConfig,
|
||||
WarmupConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.gpu_pool import (
|
||||
GpuPool,
|
||||
InProcessGpuPool,
|
||||
PoolAcquireTimeout,
|
||||
SubprocessGpuPool,
|
||||
_WorkerHandle,
|
||||
)
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# In-process pool
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class _MockGenerator:
|
||||
sleep_s: float = 0.0
|
||||
|
||||
def generate(self, request: GenerationRequest) -> dict[str, Any]:
|
||||
if self.sleep_s:
|
||||
time.sleep(self.sleep_s)
|
||||
return {
|
||||
"frames": [],
|
||||
"prompt_echo": request.prompt,
|
||||
}
|
||||
|
||||
|
||||
class TestInProcessGpuPool:
|
||||
|
||||
def test_is_gpu_pool(self):
|
||||
assert isinstance(
|
||||
InProcessGpuPool(_MockGenerator()), GpuPool)
|
||||
|
||||
def test_acquire_returns_deterministic_assignment(self):
|
||||
pool = InProcessGpuPool(_MockGenerator(), gpu_id=7)
|
||||
|
||||
async def run():
|
||||
a = await pool.acquire("sess-a")
|
||||
return a
|
||||
|
||||
a = asyncio.run(run())
|
||||
assert a.gpu_id == 7
|
||||
assert a.worker_id.startswith("inproc-")
|
||||
|
||||
def test_acquire_is_sticky_across_calls(self):
|
||||
pool = InProcessGpuPool(_MockGenerator())
|
||||
|
||||
async def run():
|
||||
a = await pool.acquire("sess-a")
|
||||
b = await pool.acquire("sess-a")
|
||||
return a, b
|
||||
|
||||
a, b = asyncio.run(run())
|
||||
assert a == b
|
||||
|
||||
def test_run_without_acquire_raises(self):
|
||||
pool = InProcessGpuPool(_MockGenerator())
|
||||
|
||||
async def run():
|
||||
with pytest.raises(RuntimeError):
|
||||
await pool.run(
|
||||
"sess-a", GenerationRequest(prompt="hi"))
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_run_returns_generator_output(self):
|
||||
pool = InProcessGpuPool(_MockGenerator())
|
||||
|
||||
async def run():
|
||||
await pool.acquire("sess-a")
|
||||
return await pool.run(
|
||||
"sess-a", GenerationRequest(prompt="hi"))
|
||||
|
||||
result = asyncio.run(run())
|
||||
assert result["prompt_echo"] == "hi"
|
||||
|
||||
def test_release_frees_binding(self):
|
||||
pool = InProcessGpuPool(_MockGenerator())
|
||||
|
||||
async def run():
|
||||
await pool.acquire("sess-a")
|
||||
await pool.release("sess-a")
|
||||
with pytest.raises(RuntimeError):
|
||||
await pool.run(
|
||||
"sess-a", GenerationRequest(prompt="hi"))
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_health_reports_active_sessions(self):
|
||||
pool = InProcessGpuPool(_MockGenerator())
|
||||
|
||||
async def run():
|
||||
await pool.acquire("sess-a")
|
||||
health = pool.health()
|
||||
assert health.total_workers == 1
|
||||
assert health.active_sessions == 1
|
||||
assert health.available_workers == 0
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Subprocess pool (driven by a thread-backed fake worker factory)
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
class _ThreadWorker:
|
||||
"""Stand-in for a subprocess worker.
|
||||
|
||||
Runs a Python thread that pulls jobs from ``job_queue`` and invokes
|
||||
a supplied mock generator. The control flow matches
|
||||
:func:`worker_main` exactly (ready + result dict shapes) so the
|
||||
parent-side pool under test exercises the same code paths.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
generator: _MockGenerator,
|
||||
*,
|
||||
job_queue: mp.Queue,
|
||||
result_queue: mp.Queue,
|
||||
shutdown_event: threading.Event,
|
||||
warmup_ms: float = 0.0,
|
||||
) -> None:
|
||||
self._generator = generator
|
||||
self._job_queue = job_queue
|
||||
self._result_queue = result_queue
|
||||
self._shutdown_event = shutdown_event
|
||||
self._warmup_ms = warmup_ms
|
||||
self._thread = threading.Thread(target=self._run, daemon=True)
|
||||
|
||||
def start(self) -> None:
|
||||
self._thread.start()
|
||||
|
||||
def join(self, timeout: float | None = None) -> None:
|
||||
self._thread.join(timeout)
|
||||
|
||||
def _run(self) -> None:
|
||||
if self._warmup_ms:
|
||||
time.sleep(self._warmup_ms / 1000.0)
|
||||
self._result_queue.put({"kind": "ready"})
|
||||
while not self._shutdown_event.is_set():
|
||||
try:
|
||||
item = self._job_queue.get(timeout=0.1)
|
||||
except queue.Empty:
|
||||
continue
|
||||
if item is None:
|
||||
return
|
||||
try:
|
||||
result = self._generator.generate(item["request"])
|
||||
self._result_queue.put({
|
||||
"kind": "result",
|
||||
"job_id": item["job_id"],
|
||||
"result": result,
|
||||
})
|
||||
except Exception as exc: # pragma: no cover - defensive
|
||||
self._result_queue.put({
|
||||
"kind": "error",
|
||||
"job_id": item["job_id"],
|
||||
"error": repr(exc),
|
||||
})
|
||||
|
||||
|
||||
def _thread_worker_factory(generator_builder):
|
||||
"""Return a WorkerFactory that uses thread workers instead of procs."""
|
||||
|
||||
def factory(
|
||||
*,
|
||||
gpu_id: int,
|
||||
generator_config: GeneratorConfig,
|
||||
warmup_config: WarmupConfig,
|
||||
) -> _WorkerHandle:
|
||||
ctx = mp.get_context("spawn")
|
||||
job_queue: mp.Queue = ctx.Queue()
|
||||
result_queue: mp.Queue = ctx.Queue()
|
||||
shutdown_event = threading.Event()
|
||||
ready = threading.Event()
|
||||
boot_ok = threading.Event()
|
||||
mp_shutdown = ctx.Event()
|
||||
|
||||
generator = generator_builder(gpu_id)
|
||||
worker = _ThreadWorker(
|
||||
generator,
|
||||
job_queue=job_queue,
|
||||
result_queue=result_queue,
|
||||
shutdown_event=shutdown_event,
|
||||
)
|
||||
worker.start()
|
||||
|
||||
# Drain the ready sentinel from the queue in the same way the
|
||||
# real factory does (thread waiter populates ``ready`` /
|
||||
# ``boot_ok``).
|
||||
def _await_ready() -> None:
|
||||
while True:
|
||||
try:
|
||||
msg = result_queue.get(timeout=1.0)
|
||||
except queue.Empty:
|
||||
if shutdown_event.is_set():
|
||||
return
|
||||
continue
|
||||
if msg.get("kind") == "ready":
|
||||
boot_ok.set()
|
||||
ready.set()
|
||||
return
|
||||
if msg.get("kind") == "error":
|
||||
ready.set()
|
||||
return
|
||||
|
||||
threading.Thread(target=_await_ready, daemon=True).start()
|
||||
|
||||
class _FakeProcess:
|
||||
def __init__(self, stop: threading.Event, worker: _ThreadWorker):
|
||||
self._stop = stop
|
||||
self._worker = worker
|
||||
|
||||
def is_alive(self) -> bool:
|
||||
return self._worker._thread.is_alive()
|
||||
|
||||
def join(self, timeout: float | None = None) -> None:
|
||||
self._stop.set()
|
||||
self._worker.join(timeout)
|
||||
|
||||
def kill(self) -> None:
|
||||
self._stop.set()
|
||||
|
||||
fake_process = _FakeProcess(shutdown_event, worker)
|
||||
|
||||
return _WorkerHandle(
|
||||
process=fake_process, # type: ignore[arg-type]
|
||||
job_queue=job_queue,
|
||||
result_queue=result_queue,
|
||||
gpu_id=gpu_id,
|
||||
worker_id=f"gpu{gpu_id}-fake",
|
||||
ready=ready,
|
||||
boot_ok=boot_ok,
|
||||
shutdown_event=mp_shutdown,
|
||||
)
|
||||
|
||||
return factory
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def pool_factory():
|
||||
"""Provide a SubprocessGpuPool built against thread workers."""
|
||||
|
||||
async def _build(num_workers: int = 2):
|
||||
pool = SubprocessGpuPool(
|
||||
generator_config=GeneratorConfig(model_path="/models/fake"),
|
||||
pool_config=GpuPoolConfig(num_workers=num_workers),
|
||||
warmup_config=WarmupConfig(enabled=False),
|
||||
worker_factory=_thread_worker_factory(
|
||||
lambda gpu_id: _MockGenerator()),
|
||||
)
|
||||
await pool.start()
|
||||
return pool
|
||||
|
||||
return _build
|
||||
|
||||
|
||||
class TestSubprocessGpuPool:
|
||||
|
||||
def test_start_spawns_requested_workers(self, pool_factory):
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=3)
|
||||
try:
|
||||
h = pool.health()
|
||||
assert h.total_workers == 3
|
||||
assert h.available_workers == 3
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_acquire_decrements_available(self, pool_factory):
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=2)
|
||||
try:
|
||||
a = await pool.acquire("sess-a")
|
||||
assert a.worker_id.endswith("-fake")
|
||||
assert pool.health().available_workers == 1
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_acquire_timeout_when_all_busy(self, pool_factory):
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=1)
|
||||
try:
|
||||
await pool.acquire("sess-a")
|
||||
with pytest.raises(PoolAcquireTimeout):
|
||||
await pool.acquire("sess-b", timeout=0.1)
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_run_returns_worker_result(self, pool_factory):
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=1)
|
||||
try:
|
||||
await pool.acquire("sess-a")
|
||||
result = await pool.run(
|
||||
"sess-a", GenerationRequest(prompt="hello"))
|
||||
assert result["prompt_echo"] == "hello"
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_release_returns_worker_to_pool(self, pool_factory):
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=1)
|
||||
try:
|
||||
await pool.acquire("sess-a")
|
||||
await pool.release("sess-a")
|
||||
# Now a second acquire should succeed without timeout.
|
||||
a = await pool.acquire("sess-b", timeout=1.0)
|
||||
assert a.worker_id.endswith("-fake")
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_sticky_binding_across_multiple_runs(self, pool_factory):
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=2)
|
||||
try:
|
||||
a1 = await pool.acquire("sess-a")
|
||||
a2 = await pool.acquire("sess-a")
|
||||
assert a1.worker_id == a2.worker_id
|
||||
# Two runs land on the same worker.
|
||||
await pool.run("sess-a", GenerationRequest(prompt="1"))
|
||||
await pool.run("sess-a", GenerationRequest(prompt="2"))
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_run_without_acquire_raises(self, pool_factory):
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=1)
|
||||
try:
|
||||
with pytest.raises(RuntimeError):
|
||||
await pool.run(
|
||||
"sess-x", GenerationRequest(prompt="x"))
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_shutdown_is_idempotent(self, pool_factory):
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=2)
|
||||
await pool.shutdown()
|
||||
await pool.shutdown() # no raise
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
class TestSubprocessGpuPoolFailureModes:
|
||||
"""Coverage for boot/runtime failures the pool has to absorb."""
|
||||
|
||||
def test_failed_boot_excluded_from_available(self):
|
||||
"""A worker whose factory leaves ``boot_ok`` unset must not be
|
||||
handed out by ``acquire``. Otherwise a session lands on a dead
|
||||
worker and ``run`` blocks forever on the missing result."""
|
||||
|
||||
def factory_with_one_failure(*, gpu_id, generator_config, warmup_config):
|
||||
ctx = mp.get_context("spawn")
|
||||
job_queue: mp.Queue = ctx.Queue()
|
||||
result_queue: mp.Queue = ctx.Queue()
|
||||
ready = threading.Event()
|
||||
boot_ok = threading.Event()
|
||||
mp_shutdown = ctx.Event()
|
||||
ready.set()
|
||||
# gpu_id 0 boots fine; gpu_id 1 fails (boot_ok stays clear).
|
||||
if gpu_id == 0:
|
||||
boot_ok.set()
|
||||
|
||||
class _AliveProcess:
|
||||
def is_alive(self) -> bool:
|
||||
return True
|
||||
def join(self, timeout: float | None = None) -> None:
|
||||
return
|
||||
def kill(self) -> None:
|
||||
return
|
||||
|
||||
return _WorkerHandle(
|
||||
process=_AliveProcess(), # type: ignore[arg-type]
|
||||
job_queue=job_queue,
|
||||
result_queue=result_queue,
|
||||
gpu_id=gpu_id,
|
||||
worker_id=f"gpu{gpu_id}",
|
||||
ready=ready,
|
||||
boot_ok=boot_ok,
|
||||
shutdown_event=mp_shutdown,
|
||||
)
|
||||
|
||||
async def run():
|
||||
pool = SubprocessGpuPool(
|
||||
generator_config=GeneratorConfig(model_path="/m"),
|
||||
pool_config=GpuPoolConfig(num_workers=2),
|
||||
warmup_config=WarmupConfig(enabled=False),
|
||||
worker_factory=factory_with_one_failure,
|
||||
)
|
||||
await pool.start()
|
||||
try:
|
||||
# Only worker 0 booted; only one slot should be available.
|
||||
assert pool.health().available_workers == 1
|
||||
a = await pool.acquire("sess-a", timeout=0.1)
|
||||
assert a.worker_id == "gpu0"
|
||||
with pytest.raises(PoolAcquireTimeout):
|
||||
await pool.acquire("sess-b", timeout=0.1)
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
def test_release_skips_dead_worker(self, pool_factory):
|
||||
"""If a worker died while bound, releasing the session must not
|
||||
return its slot to the available queue — a later acquire would
|
||||
hand the dead slot to a new session."""
|
||||
|
||||
async def run():
|
||||
pool = await pool_factory(num_workers=1)
|
||||
try:
|
||||
await pool.acquire("sess-a")
|
||||
# Simulate the worker dying mid-session.
|
||||
pool._workers[0].process._stop.set()
|
||||
pool._workers[0].process._worker.join(timeout=1.0)
|
||||
await pool.release("sess-a")
|
||||
# Available queue must remain empty.
|
||||
assert pool.health().available_workers == 0
|
||||
finally:
|
||||
await pool.shutdown()
|
||||
|
||||
asyncio.run(run())
|
||||
@@ -1,207 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the provider-agnostic prompt enhancer."""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt import (
|
||||
LLMProvider,
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMResponse,
|
||||
PromptEnhancer,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _StaticProvider:
|
||||
name: str
|
||||
content: str
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
return LLMResponse(
|
||||
content=self.content,
|
||||
provider=self.name,
|
||||
model=request.model,
|
||||
latency_ms=1.0,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _FailingProvider:
|
||||
name: str
|
||||
message: str = "boom"
|
||||
retryable: bool = True
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
raise LLMProviderError(self.message, retryable=self.retryable)
|
||||
|
||||
|
||||
class TestConstruction:
|
||||
|
||||
def test_requires_at_least_one_provider(self):
|
||||
with pytest.raises(ValueError):
|
||||
PromptEnhancer(providers=[], model="m")
|
||||
|
||||
def test_registers_system_prompts_from_defaults(self, tmp_path):
|
||||
enh = PromptEnhancer(
|
||||
providers=[_StaticProvider("p", "ok")],
|
||||
model="m",
|
||||
)
|
||||
# Defaults live inside _DEFAULT_SYSTEM_PROMPTS; no hot-reload file
|
||||
# means the enhancer falls back to the shipped values.
|
||||
assert "prompt enhancer" in enh._system_prompts.enhance.lower()
|
||||
|
||||
def test_reads_override_files_when_present(self, tmp_path):
|
||||
(tmp_path / "enhance.txt").write_text("customized enhance")
|
||||
(tmp_path / "auto_extend.txt").write_text("") # empty -> default
|
||||
enh = PromptEnhancer(
|
||||
providers=[_StaticProvider("p", "ok")],
|
||||
model="m",
|
||||
system_prompt_dir=str(tmp_path),
|
||||
)
|
||||
assert enh._system_prompts.enhance == "customized enhance"
|
||||
# Empty file fell back to the default.
|
||||
assert "continuation assistant" in enh._system_prompts.auto_extend
|
||||
|
||||
|
||||
class TestEnhance:
|
||||
|
||||
def test_enhance_returns_primary_provider_response(self):
|
||||
enh = PromptEnhancer(
|
||||
providers=[_StaticProvider("primary", "enhanced!")],
|
||||
model="m",
|
||||
)
|
||||
response = asyncio.run(enh.enhance("a fox"))
|
||||
assert response.content == "enhanced!"
|
||||
assert response.provider == "primary"
|
||||
assert response.fallback_used is False
|
||||
|
||||
def test_enhance_falls_back_on_provider_error(self):
|
||||
enh = PromptEnhancer(
|
||||
providers=[
|
||||
_FailingProvider("a"),
|
||||
_StaticProvider("b", "fallback!"),
|
||||
],
|
||||
model="m",
|
||||
)
|
||||
response = asyncio.run(enh.enhance("x"))
|
||||
assert response.provider == "b"
|
||||
assert response.content == "fallback!"
|
||||
assert response.fallback_used is True
|
||||
|
||||
def test_enhance_stops_on_non_retryable_error(self):
|
||||
enh = PromptEnhancer(
|
||||
providers=[
|
||||
_FailingProvider("a", message="hard-fail", retryable=False),
|
||||
_StaticProvider("b", "should-not-be-reached"),
|
||||
],
|
||||
model="m",
|
||||
)
|
||||
with pytest.raises(LLMProviderError, match="hard-fail"):
|
||||
asyncio.run(enh.enhance("x"))
|
||||
|
||||
def test_enhance_raises_when_all_providers_fail(self):
|
||||
enh = PromptEnhancer(
|
||||
providers=[
|
||||
_FailingProvider("a"),
|
||||
_FailingProvider("b"),
|
||||
],
|
||||
model="m",
|
||||
)
|
||||
with pytest.raises(LLMProviderError):
|
||||
asyncio.run(enh.enhance("x"))
|
||||
|
||||
|
||||
class TestAutoExtendAndRewrite:
|
||||
|
||||
def test_auto_extend_joins_prior_prompts(self):
|
||||
captured: list[str] = []
|
||||
|
||||
class _Capturer:
|
||||
name = "cap"
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
user = next(m.content for m in request.messages
|
||||
if m.role == "user")
|
||||
captured.append(user)
|
||||
return LLMResponse(
|
||||
content="next",
|
||||
provider=self.name,
|
||||
model=request.model,
|
||||
latency_ms=1.0,
|
||||
)
|
||||
|
||||
enh = PromptEnhancer(providers=[_Capturer()], model="m")
|
||||
asyncio.run(enh.auto_extend(["first", "second"]))
|
||||
assert captured[0] == "first\nsecond"
|
||||
|
||||
def test_rewrite_passes_seed_through(self):
|
||||
captured: list[str] = []
|
||||
|
||||
class _Capturer:
|
||||
name = "cap"
|
||||
|
||||
async def complete(self, request: LLMRequest) -> LLMResponse:
|
||||
user = next(m.content for m in request.messages
|
||||
if m.role == "user")
|
||||
captured.append(user)
|
||||
return LLMResponse(
|
||||
content="one\ntwo\nthree",
|
||||
provider=self.name,
|
||||
model=request.model,
|
||||
latency_ms=1.0,
|
||||
)
|
||||
|
||||
enh = PromptEnhancer(providers=[_Capturer()], model="m")
|
||||
response = asyncio.run(enh.rewrite("seed prompt"))
|
||||
assert captured == ["seed prompt"]
|
||||
assert response.content.splitlines() == ["one", "two", "three"]
|
||||
|
||||
|
||||
class TestRegisterProvider:
|
||||
|
||||
def test_register_appends_by_default(self):
|
||||
enh = PromptEnhancer(
|
||||
providers=[_StaticProvider("primary", "a")],
|
||||
model="m",
|
||||
)
|
||||
enh.register_provider(_StaticProvider("extra", "b"))
|
||||
assert [p.name for p in enh.providers] == ["primary", "extra"]
|
||||
|
||||
def test_register_priority_zero_makes_primary(self):
|
||||
enh = PromptEnhancer(
|
||||
providers=[_StaticProvider("old", "a")],
|
||||
model="m",
|
||||
)
|
||||
enh.register_provider(
|
||||
_StaticProvider("new-primary", "b"), priority=0)
|
||||
assert enh.providers[0].name == "new-primary"
|
||||
|
||||
def test_registered_provider_is_used_in_enhance(self):
|
||||
enh = PromptEnhancer(
|
||||
providers=[_FailingProvider("broken")],
|
||||
model="m",
|
||||
)
|
||||
enh.register_provider(_StaticProvider("fallback", "ok"))
|
||||
response = asyncio.run(enh.enhance("x"))
|
||||
assert response.content == "ok"
|
||||
assert response.fallback_used is True
|
||||
|
||||
|
||||
class TestHotReload:
|
||||
|
||||
def test_reload_picks_up_new_file(self, tmp_path):
|
||||
(tmp_path / "enhance.txt").write_text("first version")
|
||||
enh = PromptEnhancer(
|
||||
providers=[_StaticProvider("p", "ok")],
|
||||
model="m",
|
||||
system_prompt_dir=str(tmp_path),
|
||||
)
|
||||
assert enh._system_prompts.enhance == "first version"
|
||||
(tmp_path / "enhance.txt").write_text("second version")
|
||||
enh.reload_system_prompts()
|
||||
assert enh._system_prompts.enhance == "second version"
|
||||
@@ -1,298 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Tests for the LLM provider protocol + built-in adapters.
|
||||
|
||||
Real Cerebras/Groq HTTP calls are stubbed with a fake ``httpx`` module
|
||||
inserted into ``sys.modules`` so the unit tests don't depend on
|
||||
external API availability or paid keys.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import sys
|
||||
import types
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.prompt.providers import (
|
||||
CerebrasProvider,
|
||||
GroqProvider,
|
||||
LLMMessage,
|
||||
LLMProvider,
|
||||
LLMProviderError,
|
||||
LLMRequest,
|
||||
LLMTimeoutError,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _FakeResponse:
|
||||
status_code: int
|
||||
payload: dict[str, Any]
|
||||
text: str = ""
|
||||
|
||||
def json(self) -> dict[str, Any]:
|
||||
return self.payload
|
||||
|
||||
|
||||
class _FakeAsyncClient:
|
||||
|
||||
def __init__(self, response_or_exc, *, captured: list) -> None:
|
||||
self._response_or_exc = response_or_exc
|
||||
self._captured = captured
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return None
|
||||
|
||||
async def post(self, url, **kwargs):
|
||||
self._captured.append({"url": url, **kwargs})
|
||||
if isinstance(self._response_or_exc, Exception):
|
||||
raise self._response_or_exc
|
||||
return self._response_or_exc
|
||||
|
||||
|
||||
def _install_fake_httpx(
|
||||
monkeypatch,
|
||||
*,
|
||||
response: _FakeResponse | None = None,
|
||||
exception: Exception | None = None,
|
||||
) -> list[dict]:
|
||||
"""Replace ``sys.modules['httpx']`` with a stub exposing the bits
|
||||
providers touch: ``AsyncClient``, ``HTTPError``, ``TimeoutException``.
|
||||
|
||||
Returns a list the test can inspect to see what requests went out.
|
||||
"""
|
||||
captured: list[dict] = []
|
||||
|
||||
class _HTTPError(Exception):
|
||||
pass
|
||||
|
||||
class _TimeoutException(_HTTPError):
|
||||
pass
|
||||
|
||||
payload = response if exception is None else exception
|
||||
|
||||
def _client_factory(*_args, **_kwargs):
|
||||
return _FakeAsyncClient(payload, captured=captured)
|
||||
|
||||
stub = types.SimpleNamespace(
|
||||
AsyncClient=_client_factory,
|
||||
HTTPError=_HTTPError,
|
||||
TimeoutException=_TimeoutException,
|
||||
)
|
||||
monkeypatch.setitem(sys.modules, "httpx", stub)
|
||||
return captured
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Cerebras
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestCerebrasProvider:
|
||||
|
||||
def test_is_llm_provider(self):
|
||||
assert isinstance(CerebrasProvider(api_key="x"), LLMProvider)
|
||||
|
||||
def test_requires_api_key(self, monkeypatch):
|
||||
monkeypatch.delenv("CEREBRAS_API_KEY", raising=False)
|
||||
provider = CerebrasProvider()
|
||||
with pytest.raises(LLMProviderError, match="CEREBRAS_API_KEY"):
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
|
||||
def test_api_key_from_env(self, monkeypatch):
|
||||
monkeypatch.setenv("CEREBRAS_API_KEY", "from-env")
|
||||
provider = CerebrasProvider()
|
||||
assert provider.api_key == "from-env"
|
||||
|
||||
def test_explicit_api_key_wins(self, monkeypatch):
|
||||
monkeypatch.setenv("CEREBRAS_API_KEY", "from-env")
|
||||
provider = CerebrasProvider(api_key="explicit")
|
||||
assert provider.api_key == "explicit"
|
||||
|
||||
def test_success_path(self, monkeypatch):
|
||||
captured = _install_fake_httpx(
|
||||
monkeypatch,
|
||||
response=_FakeResponse(
|
||||
status_code=200,
|
||||
payload={
|
||||
"choices": [{"message": {"content": "enhanced"}}],
|
||||
},
|
||||
),
|
||||
)
|
||||
provider = CerebrasProvider(api_key="secret")
|
||||
result = asyncio.run(provider.complete(
|
||||
LLMRequest(
|
||||
messages=[LLMMessage(role="user", content="a fox")],
|
||||
model="gpt-oss-120b",
|
||||
max_tokens=64,
|
||||
temperature=0.7,
|
||||
)))
|
||||
assert result.content == "enhanced"
|
||||
assert result.provider == "cerebras"
|
||||
assert result.model == "gpt-oss-120b"
|
||||
# Verify the HTTP request was shaped correctly.
|
||||
body = captured[0]["json"]
|
||||
assert body["model"] == "gpt-oss-120b"
|
||||
assert body["messages"] == [{"role": "user", "content": "a fox"}]
|
||||
assert captured[0]["headers"]["Authorization"] == "Bearer secret"
|
||||
|
||||
def test_http_error_wrapped(self, monkeypatch):
|
||||
# Stage the stub first, then raise that stub's own HTTPError so
|
||||
# the provider's ``except httpx.HTTPError`` catches it.
|
||||
stub = types.SimpleNamespace()
|
||||
|
||||
class _HTTPError(Exception):
|
||||
pass
|
||||
|
||||
class _TimeoutException(_HTTPError):
|
||||
pass
|
||||
|
||||
stub.HTTPError = _HTTPError
|
||||
stub.TimeoutException = _TimeoutException
|
||||
stub.AsyncClient = lambda *a, **k: _FakeAsyncClient(
|
||||
_HTTPError("boom"), captured=[])
|
||||
monkeypatch.setitem(sys.modules, "httpx", stub)
|
||||
|
||||
provider = CerebrasProvider(api_key="k")
|
||||
with pytest.raises(LLMProviderError, match="HTTP"):
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
|
||||
def test_timeout_wrapped(self, monkeypatch):
|
||||
_install_fake_httpx(monkeypatch)
|
||||
# Install timeout exception after stubbing.
|
||||
stub = sys.modules["httpx"]
|
||||
|
||||
def raising_factory(*_a, **_kw):
|
||||
raise stub.TimeoutException("timed out") # type: ignore[attr-defined]
|
||||
|
||||
stub.AsyncClient = raising_factory # type: ignore[attr-defined]
|
||||
provider = CerebrasProvider(api_key="k")
|
||||
with pytest.raises(LLMTimeoutError):
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
|
||||
def test_4xx_raises_non_retryable_provider_error(self, monkeypatch):
|
||||
_install_fake_httpx(
|
||||
monkeypatch,
|
||||
response=_FakeResponse(
|
||||
status_code=401,
|
||||
payload={},
|
||||
text="unauthorized",
|
||||
),
|
||||
)
|
||||
provider = CerebrasProvider(api_key="k")
|
||||
with pytest.raises(LLMProviderError, match="401") as excinfo:
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
assert excinfo.value.retryable is False
|
||||
|
||||
def test_429_raises_retryable_provider_error(self, monkeypatch):
|
||||
_install_fake_httpx(
|
||||
monkeypatch,
|
||||
response=_FakeResponse(
|
||||
status_code=429,
|
||||
payload={},
|
||||
text="rate limited",
|
||||
),
|
||||
)
|
||||
provider = CerebrasProvider(api_key="k")
|
||||
with pytest.raises(LLMProviderError, match="429") as excinfo:
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
assert excinfo.value.retryable is True
|
||||
|
||||
def test_5xx_raises_retryable_provider_error(self, monkeypatch):
|
||||
_install_fake_httpx(
|
||||
monkeypatch,
|
||||
response=_FakeResponse(
|
||||
status_code=503,
|
||||
payload={},
|
||||
text="service unavailable",
|
||||
),
|
||||
)
|
||||
provider = CerebrasProvider(api_key="k")
|
||||
with pytest.raises(LLMProviderError, match="503") as excinfo:
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
assert excinfo.value.retryable is True
|
||||
|
||||
def test_non_json_body_raises_provider_error(self, monkeypatch):
|
||||
# Simulate a proxy/load-balancer HTML error page slipping in
|
||||
# with a 200 status: response.json() raises, and the provider
|
||||
# must wrap it in an LLMProviderError instead of bubbling.
|
||||
class _BadJsonResponse:
|
||||
status_code = 200
|
||||
text = "<html>oops</html>"
|
||||
|
||||
def json(self):
|
||||
raise ValueError("Expecting value: line 1 column 1 (char 0)")
|
||||
|
||||
_install_fake_httpx(
|
||||
monkeypatch,
|
||||
response=_BadJsonResponse(), # type: ignore[arg-type]
|
||||
)
|
||||
provider = CerebrasProvider(api_key="k")
|
||||
with pytest.raises(LLMProviderError, match="non-JSON"):
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
|
||||
def test_empty_choices_raises(self, monkeypatch):
|
||||
_install_fake_httpx(
|
||||
monkeypatch,
|
||||
response=_FakeResponse(
|
||||
status_code=200, payload={"choices": []}),
|
||||
)
|
||||
provider = CerebrasProvider(api_key="k")
|
||||
with pytest.raises(LLMProviderError, match="no choices"):
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------
|
||||
# Groq
|
||||
# ----------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestGroqProvider:
|
||||
|
||||
def test_is_llm_provider(self):
|
||||
assert isinstance(GroqProvider(api_key="x"), LLMProvider)
|
||||
|
||||
def test_requires_api_key(self, monkeypatch):
|
||||
monkeypatch.delenv("GROQ_API_KEY", raising=False)
|
||||
provider = GroqProvider()
|
||||
with pytest.raises(LLMProviderError, match="GROQ_API_KEY"):
|
||||
asyncio.run(provider.complete(
|
||||
LLMRequest(messages=[], model="m")))
|
||||
|
||||
def test_api_key_from_env(self, monkeypatch):
|
||||
monkeypatch.setenv("GROQ_API_KEY", "from-env")
|
||||
provider = GroqProvider()
|
||||
assert provider.api_key == "from-env"
|
||||
|
||||
def test_success_path(self, monkeypatch):
|
||||
captured = _install_fake_httpx(
|
||||
monkeypatch,
|
||||
response=_FakeResponse(
|
||||
status_code=200,
|
||||
payload={
|
||||
"choices": [{"message": {"content": "groq out"}}],
|
||||
},
|
||||
),
|
||||
)
|
||||
provider = GroqProvider(api_key="secret")
|
||||
result = asyncio.run(provider.complete(
|
||||
LLMRequest(
|
||||
messages=[LLMMessage(role="user", content="a deer")],
|
||||
model="llama-3.1-70b",
|
||||
)))
|
||||
assert result.content == "groq out"
|
||||
assert result.provider == "groq"
|
||||
assert captured[0]["headers"]["Authorization"] == "Bearer secret"
|
||||
@@ -1,279 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Router tests — registry semantics + health loop behavior.
|
||||
|
||||
Avoids real WebSocket proxying (the bridge requires the ``websockets``
|
||||
package and two running uvicorn processes); test_server.py already
|
||||
covers the end-to-end WS protocol against a direct backend.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.entrypoints.streaming.router.config import (
|
||||
ReplicaEndpoint,
|
||||
RouterConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.router.registry import (
|
||||
ReplicaRegistry,
|
||||
ReplicaStatus,
|
||||
run_health_check_loop,
|
||||
)
|
||||
|
||||
|
||||
def _registry(
|
||||
*,
|
||||
num_primary: int = 1,
|
||||
num_secondary: int = 1,
|
||||
) -> ReplicaRegistry:
|
||||
replicas = [
|
||||
ReplicaEndpoint(
|
||||
url=f"http://primary-{i}:8000",
|
||||
primary=True,
|
||||
)
|
||||
for i in range(num_primary)
|
||||
] + [
|
||||
ReplicaEndpoint(
|
||||
url=f"http://secondary-{i}:8000",
|
||||
primary=False,
|
||||
)
|
||||
for i in range(num_secondary)
|
||||
]
|
||||
return ReplicaRegistry(replicas)
|
||||
|
||||
|
||||
class TestReplicaRegistry:
|
||||
|
||||
def test_requires_replicas(self):
|
||||
with pytest.raises(ValueError):
|
||||
ReplicaRegistry([])
|
||||
|
||||
def test_select_none_when_all_unknown(self):
|
||||
assert _registry().select() is None
|
||||
|
||||
def test_select_prefers_healthy_primary(self):
|
||||
reg = _registry()
|
||||
|
||||
async def promote():
|
||||
primary = reg.primaries()[0]
|
||||
await reg.record_success(
|
||||
primary, recovery_threshold=1, latency_ms=1.0)
|
||||
# Also mark secondary healthy — primary should still win.
|
||||
sec = next(r for r in reg.all() if not r.primary)
|
||||
await reg.record_success(
|
||||
sec, recovery_threshold=1, latency_ms=1.0)
|
||||
return reg.select()
|
||||
|
||||
pick = asyncio.run(promote())
|
||||
assert pick is not None
|
||||
assert pick.primary
|
||||
|
||||
def test_falls_back_to_secondary_when_primary_unhealthy(self):
|
||||
reg = _registry()
|
||||
|
||||
async def run():
|
||||
primary = reg.primaries()[0]
|
||||
sec = next(r for r in reg.all() if not r.primary)
|
||||
# Fail primary past threshold, succeed secondary.
|
||||
for _ in range(3):
|
||||
await reg.record_failure(
|
||||
primary, failure_threshold=3, reason="mock")
|
||||
await reg.record_success(
|
||||
sec, recovery_threshold=1, latency_ms=1.0)
|
||||
return reg.select()
|
||||
|
||||
pick = asyncio.run(run())
|
||||
assert pick is not None
|
||||
assert not pick.primary
|
||||
|
||||
def test_failure_threshold_transitions_to_unhealthy(self):
|
||||
reg = _registry()
|
||||
primary = reg.primaries()[0]
|
||||
|
||||
async def run():
|
||||
for _ in range(2):
|
||||
await reg.record_failure(
|
||||
primary, failure_threshold=3, reason="x")
|
||||
assert primary.health.status is not ReplicaStatus.UNHEALTHY
|
||||
await reg.record_failure(
|
||||
primary, failure_threshold=3, reason="x")
|
||||
return primary
|
||||
|
||||
result = asyncio.run(run())
|
||||
assert result.health.status is ReplicaStatus.UNHEALTHY
|
||||
assert result.health.consecutive_failures == 3
|
||||
|
||||
def test_recovery_threshold_returns_to_healthy(self):
|
||||
reg = _registry()
|
||||
primary = reg.primaries()[0]
|
||||
|
||||
async def run():
|
||||
for _ in range(3):
|
||||
await reg.record_failure(
|
||||
primary, failure_threshold=3, reason="x")
|
||||
assert primary.health.status is ReplicaStatus.UNHEALTHY
|
||||
for _ in range(2):
|
||||
await reg.record_success(
|
||||
primary, recovery_threshold=2, latency_ms=5.0)
|
||||
return primary
|
||||
|
||||
result = asyncio.run(run())
|
||||
assert result.health.status is ReplicaStatus.HEALTHY
|
||||
assert result.health.last_latency_ms == 5.0
|
||||
|
||||
def test_record_success_resets_failure_counter(self):
|
||||
reg = _registry()
|
||||
primary = reg.primaries()[0]
|
||||
|
||||
async def run():
|
||||
await reg.record_failure(
|
||||
primary, failure_threshold=10, reason="x")
|
||||
await reg.record_success(
|
||||
primary, recovery_threshold=1, latency_ms=1.0)
|
||||
return primary
|
||||
|
||||
result = asyncio.run(run())
|
||||
assert result.health.consecutive_failures == 0
|
||||
|
||||
|
||||
class TestHealthCheckLoop:
|
||||
|
||||
def test_loop_transitions_replicas_on_probe_results(self):
|
||||
config = RouterConfig(
|
||||
replicas=[
|
||||
ReplicaEndpoint(url="http://a", primary=True),
|
||||
ReplicaEndpoint(url="http://b"),
|
||||
],
|
||||
health_check_interval_seconds=0.01,
|
||||
failure_threshold=1,
|
||||
recovery_threshold=1,
|
||||
)
|
||||
reg = ReplicaRegistry(config.replicas)
|
||||
stop_event = asyncio.Event()
|
||||
|
||||
calls: list[str] = []
|
||||
|
||||
async def probe(url, *, timeout):
|
||||
calls.append(url)
|
||||
if "http://a" in url:
|
||||
return 1.0, None
|
||||
return 0.0, "mock failure"
|
||||
|
||||
async def run() -> None:
|
||||
task = asyncio.create_task(run_health_check_loop(
|
||||
registry=reg, config=config, stop_event=stop_event,
|
||||
http_get=probe,
|
||||
))
|
||||
await asyncio.sleep(0.05)
|
||||
stop_event.set()
|
||||
await task
|
||||
|
||||
asyncio.run(run())
|
||||
a = reg.get("http://a")
|
||||
b = reg.get("http://b")
|
||||
assert a is not None
|
||||
assert b is not None
|
||||
assert a.health.status is ReplicaStatus.HEALTHY
|
||||
assert b.health.status is ReplicaStatus.UNHEALTHY
|
||||
assert any("/health" in c for c in calls)
|
||||
|
||||
|
||||
class TestRouterApp:
|
||||
|
||||
def test_status_endpoint_lists_replicas(self):
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from fastvideo.entrypoints.streaming.router.main import build_router_app
|
||||
|
||||
config = RouterConfig(
|
||||
replicas=[
|
||||
ReplicaEndpoint(url="http://a", primary=True),
|
||||
ReplicaEndpoint(url="http://b"),
|
||||
],
|
||||
health_check_interval_seconds=60, # don't actually poll
|
||||
)
|
||||
reg = ReplicaRegistry(config.replicas)
|
||||
app = build_router_app(config, registry=reg)
|
||||
client = TestClient(app)
|
||||
response = client.get("/status")
|
||||
body = response.json()
|
||||
urls = {r["url"] for r in body["replicas"]}
|
||||
assert urls == {"http://a", "http://b"}
|
||||
# Initial status is UNKNOWN.
|
||||
assert all(r["status"] == "unknown" for r in body["replicas"])
|
||||
|
||||
def test_ws_rejects_when_no_healthy_replica(self):
|
||||
from starlette.testclient import TestClient
|
||||
|
||||
from fastvideo.entrypoints.streaming.router.main import build_router_app
|
||||
|
||||
config = RouterConfig(
|
||||
replicas=[ReplicaEndpoint(url="http://a", primary=True)],
|
||||
health_check_interval_seconds=60,
|
||||
)
|
||||
reg = ReplicaRegistry(config.replicas)
|
||||
app = build_router_app(config, registry=reg)
|
||||
client = TestClient(app)
|
||||
with client.websocket_connect("/v1/stream") as ws:
|
||||
err = ws.receive_json()
|
||||
assert err["type"] == "error"
|
||||
assert err["code"] == "gpu_unavailable"
|
||||
|
||||
|
||||
class TestUnknownToHealthyImmediate:
|
||||
"""Initial probe must promote UNKNOWN -> HEALTHY without waiting for recovery_threshold."""
|
||||
|
||||
def test_first_success_promotes_unknown(self):
|
||||
reg = _registry(num_primary=1, num_secondary=0)
|
||||
primary = reg.primaries()[0]
|
||||
assert primary.health.status is ReplicaStatus.UNKNOWN
|
||||
|
||||
async def run():
|
||||
await reg.record_success(primary, recovery_threshold=10, latency_ms=1.0)
|
||||
|
||||
asyncio.run(run())
|
||||
assert primary.health.status is ReplicaStatus.HEALTHY
|
||||
|
||||
def test_unhealthy_recovery_still_gated_by_threshold(self):
|
||||
reg = _registry(num_primary=1, num_secondary=0)
|
||||
primary = reg.primaries()[0]
|
||||
|
||||
async def run():
|
||||
for _ in range(3):
|
||||
await reg.record_failure(primary, failure_threshold=3, reason="x")
|
||||
assert primary.health.status is ReplicaStatus.UNHEALTHY
|
||||
await reg.record_success(primary, recovery_threshold=2, latency_ms=1.0)
|
||||
assert primary.health.status is ReplicaStatus.UNHEALTHY # 1/2 successes
|
||||
await reg.record_success(primary, recovery_threshold=2, latency_ms=1.0)
|
||||
assert primary.health.status is ReplicaStatus.HEALTHY # 2/2 successes
|
||||
|
||||
asyncio.run(run())
|
||||
|
||||
|
||||
class TestConfigValidation:
|
||||
"""RouterConfig.__post_init__ rejects malformed configs."""
|
||||
|
||||
def test_rejects_path_in_url(self):
|
||||
with pytest.raises(ValueError, match="without a path"):
|
||||
RouterConfig(replicas=[ReplicaEndpoint(url="http://host:8000/api")])
|
||||
|
||||
def test_rejects_query_in_url(self):
|
||||
with pytest.raises(ValueError, match="query/fragment"):
|
||||
RouterConfig(replicas=[ReplicaEndpoint(url="http://host:8000?x=1")])
|
||||
|
||||
def test_rejects_fragment_in_url(self):
|
||||
with pytest.raises(ValueError, match="query/fragment"):
|
||||
RouterConfig(replicas=[ReplicaEndpoint(url="http://host:8000#frag")])
|
||||
|
||||
def test_rejects_duplicate_urls(self):
|
||||
with pytest.raises(ValueError, match="Duplicate"):
|
||||
RouterConfig(replicas=[
|
||||
ReplicaEndpoint(url="http://host:8000"),
|
||||
ReplicaEndpoint(url="http://host:8000"),
|
||||
])
|
||||
|
||||
def test_accepts_trailing_slash(self):
|
||||
# parsed.path == "/" should be allowed
|
||||
cfg = RouterConfig(replicas=[ReplicaEndpoint(url="http://host:8000/")])
|
||||
assert cfg.replicas[0].url == "http://host:8000/"
|
||||
@@ -1,86 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Unit coverage for :mod:`fastvideo.entrypoints.streaming.worker`.
|
||||
|
||||
The full ``worker_main`` loop runs in a subprocess and is exercised via
|
||||
``test_gpu_pool.py``'s subprocess integration tests. This file covers
|
||||
the in-process pieces:
|
||||
|
||||
* the two-segment warmup feeds segment 1's continuation state into
|
||||
segment 2 so both compile branches are primed before the worker
|
||||
reports ready
|
||||
* result-shape extractors handle both attribute-style and dict-style
|
||||
generator returns
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from fastvideo.api.schema import (
|
||||
ContinuationState,
|
||||
GenerationRequest,
|
||||
WarmupConfig,
|
||||
)
|
||||
from fastvideo.entrypoints.streaming.worker import (
|
||||
_extract_continuation_state,
|
||||
_warmup_worker,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _RecordingGenerator:
|
||||
"""Captures the requests passed to ``generate`` and returns a
|
||||
canned :class:`ContinuationState` on the first call so the warmup
|
||||
can feed it into the second call.
|
||||
"""
|
||||
|
||||
state_to_return: ContinuationState | None = field(default_factory=lambda: ContinuationState(
|
||||
kind="ltx2.v1",
|
||||
payload={"schema_version": 1, "segment_index": 1},
|
||||
))
|
||||
requests: list[GenerationRequest] = field(default_factory=list)
|
||||
|
||||
def generate(self, request: GenerationRequest) -> dict[str, Any]:
|
||||
self.requests.append(request)
|
||||
return {"frames": [], "state": self.state_to_return}
|
||||
|
||||
|
||||
class TestWarmupTwoSegment:
|
||||
|
||||
def test_warmup_runs_segment_one_then_segment_two_with_returned_state(self) -> None:
|
||||
gen = _RecordingGenerator()
|
||||
_warmup_worker(gen, WarmupConfig(enabled=True, prompt="warm"))
|
||||
|
||||
assert len(gen.requests) == 2
|
||||
|
||||
seg1 = gen.requests[0]
|
||||
assert seg1.state is None
|
||||
assert seg1.output.return_state is True
|
||||
|
||||
seg2 = gen.requests[1]
|
||||
assert seg2.state is gen.state_to_return
|
||||
|
||||
def test_warmup_passes_through_when_no_state_returned(self) -> None:
|
||||
gen = _RecordingGenerator(state_to_return=None)
|
||||
_warmup_worker(gen, WarmupConfig(enabled=True, prompt="warm"))
|
||||
|
||||
assert len(gen.requests) == 2
|
||||
assert gen.requests[1].state is None
|
||||
|
||||
|
||||
class TestExtractContinuationState:
|
||||
|
||||
def test_extracts_from_attribute(self) -> None:
|
||||
|
||||
class _R:
|
||||
state = ContinuationState(kind="k", payload={})
|
||||
|
||||
assert _extract_continuation_state(_R()).kind == "k"
|
||||
|
||||
def test_extracts_from_dict(self) -> None:
|
||||
state = ContinuationState(kind="k", payload={})
|
||||
assert _extract_continuation_state({"state": state}) is state
|
||||
|
||||
def test_returns_none_for_missing(self) -> None:
|
||||
assert _extract_continuation_state({}) is None
|
||||
assert _extract_continuation_state(object()) is None
|
||||
@@ -30,8 +30,6 @@ image = (modal.Image.from_registry(
|
||||
os.environ.get("TEST_SCOPE", ""),
|
||||
"IMAGE_VERSION":
|
||||
os.environ.get("IMAGE_VERSION", ""),
|
||||
"HF_REPO_ID":
|
||||
"FastVideo/performance-tracking",
|
||||
}))
|
||||
|
||||
|
||||
@@ -68,19 +66,25 @@ def run_test(pytest_command: str):
|
||||
cd fastvideo-kernel &&
|
||||
./build.sh &&
|
||||
cd .. &&
|
||||
uv pip install -e ".[test]" &&
|
||||
uv pip install -e .[test] &&
|
||||
{pytest_command}
|
||||
"""
|
||||
|
||||
# result = subprocess.run(["/bin/bash", "-c", command],
|
||||
# stdout=sys.stdout,
|
||||
# stderr=sys.stderr,
|
||||
# check=False)
|
||||
|
||||
# sys.exit(result.returncode)
|
||||
|
||||
result = subprocess.run(["/bin/bash", "-c", command],
|
||||
stdout=sys.stdout,
|
||||
stderr=sys.stderr,
|
||||
check=False)
|
||||
|
||||
# Modal containers crash on sys.exit(0); raise on failure, return on success.
|
||||
if result.returncode != 0:
|
||||
raise RuntimeError(
|
||||
f"Test command failed with exit code {result.returncode}")
|
||||
raise RuntimeError(f"Test command failed with exit code {result.returncode}")
|
||||
# On success, just return — don't call sys.exit()
|
||||
|
||||
@app.function(gpu="H100:1",
|
||||
image=image,
|
||||
@@ -236,27 +240,21 @@ def run_lora_extraction_tests():
|
||||
timeout=1800,
|
||||
secrets=[
|
||||
modal.Secret.from_dict(
|
||||
{"HF_API_KEY": os.environ.get("HF_API_KEY", "")})
|
||||
{"HF_API_KEY": os.environ.get("HF_API_KEY", ""),
|
||||
"HF_REPO_ID": "FastVideo/performance-tracking"})
|
||||
],
|
||||
volumes={"/root/data": model_vol})
|
||||
volumes={
|
||||
"/root/data": model_vol,
|
||||
})
|
||||
def run_performance_tests():
|
||||
# compare_baseline.py emits normalized_perf_*.json artifacts for manual
|
||||
# performance-baseline reseeds when the rolling comparison runs.
|
||||
run_test(
|
||||
"export HF_HOME='/root/data/.cache' && "
|
||||
"export PERFORMANCE_TRACKING_ROOT='/tmp/perf-tracking' && "
|
||||
"hf auth login --token $HF_API_KEY && "
|
||||
"pytest ./fastvideo/tests/performance -vs; "
|
||||
"PYTEST_RC=$?; "
|
||||
"PERF_RC=0; "
|
||||
"if [ $PYTEST_RC -eq 0 ]; then "
|
||||
"python ./fastvideo/tests/performance/compare_baseline.py; "
|
||||
"PERF_RC=$?; "
|
||||
"fi; "
|
||||
"python ./fastvideo/tests/performance/dashboard.py || true; "
|
||||
"FINAL_RC=$PYTEST_RC; "
|
||||
"if [ $FINAL_RC -eq 0 ]; then FINAL_RC=$PERF_RC; fi; "
|
||||
"exit $FINAL_RC")
|
||||
"pytest ./fastvideo/tests/performance -vs && "
|
||||
"python ./fastvideo/tests/performance/compare_baseline.py && "
|
||||
"python ./fastvideo/tests/performance/dashboard.py"
|
||||
)
|
||||
|
||||
|
||||
@app.function(gpu="L40S:1",
|
||||
|
||||
@@ -480,15 +480,11 @@ def _prepare_ssim_workspace(
|
||||
{checkout_command}
|
||||
rm -rf fastvideo/tests/ssim/reference_videos
|
||||
git_retry git submodule update --init --recursive
|
||||
uv pip install -e ".[test]"
|
||||
uv pip install -e .[test]
|
||||
cd fastvideo-kernel
|
||||
./build.sh
|
||||
cd ..
|
||||
uv pip install git+https://github.com/microsoft/MoGe.git
|
||||
# Stable Audio Open 1.0 inference deps (optional in basic install,
|
||||
# required by `StableAudioDenoisingStage`; consumed by
|
||||
# `test_stable_audio_similarity.py`).
|
||||
uv pip install k_diffusion einops_exts alias_free_torch torchsde
|
||||
export HF_HOME='/root/data/.cache'
|
||||
hf auth login --token "$HF_API_KEY"
|
||||
"""
|
||||
|
||||
@@ -3,30 +3,26 @@
|
||||
|
||||
This script:
|
||||
1) reads current benchmark results from fastvideo/tests/performance/results,
|
||||
2) syncs the canonical baseline from the configured HF dataset repo,
|
||||
3) compares each current record against the median of up to 5 prior records
|
||||
(filtered by gpu_type, successful only),
|
||||
4) on persist runs (full-suite on main branch), writes the normalized record
|
||||
back to the HF dataset repo,
|
||||
5) exits non-zero if any metric regresses by more than PERF_MAX_REGRESSION
|
||||
(default 5%).
|
||||
2) writes normalized tracking records to the Modal volume path,
|
||||
3) compares each current record against the mean of up to 5 prior records,
|
||||
4) exits non-zero if any metric regresses by more than 15%.
|
||||
"""
|
||||
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import statistics
|
||||
import sys
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from hf_store import (
|
||||
load_records_for_model,
|
||||
safe_float,
|
||||
sanitize,
|
||||
sync_from_hf,
|
||||
upload_record,
|
||||
)
|
||||
from hf_store import sync_from_hf, upload_record, load_records_for_model, sanitize, safe_float
|
||||
|
||||
# Use the env var passed by Modal, fallback to a default if needed
|
||||
HF_REPO_ID = os.environ.get("HF_REPO_ID", "FastVideo/performance-tracking")
|
||||
HF_TOKEN = os.environ.get("HF_API_KEY")
|
||||
|
||||
RESULTS_DIR = os.path.join(
|
||||
os.path.dirname(os.path.abspath(__file__)),
|
||||
@@ -36,14 +32,25 @@ TRACKING_ROOT = os.environ.get(
|
||||
"PERFORMANCE_TRACKING_ROOT",
|
||||
"/tmp/perf-tracking",
|
||||
)
|
||||
PERF_REPORTS_DIR = os.environ.get("PERF_REPORTS_DIR", "/root/data/perf_reports")
|
||||
MAX_REGRESSION = float(os.environ.get("PERF_MAX_REGRESSION", "0.05"))
|
||||
|
||||
|
||||
def _should_persist_tracking() -> bool:
|
||||
test_scope = os.environ.get("TEST_SCOPE", "")
|
||||
branch = os.environ.get("BUILDKITE_BRANCH", "")
|
||||
return test_scope == "full" and branch == "main"
|
||||
# test_scope = os.environ.get("TEST_SCOPE", "")
|
||||
# branch = os.environ.get("BUILDKITE_BRANCH", "")
|
||||
# return test_scope == "full" and branch == "main"
|
||||
return True # only for testing purpose.
|
||||
|
||||
def _sanitize(value: str) -> str:
|
||||
return re.sub(r"[^A-Za-z0-9._-]", "_", value)
|
||||
|
||||
|
||||
def _safe_float(value: Any) -> float | None:
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _load_current_results() -> list[dict[str, Any]]:
|
||||
@@ -55,14 +62,7 @@ def _load_current_results() -> list[dict[str, Any]]:
|
||||
return records
|
||||
|
||||
|
||||
def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Normalize a raw perf_*.json result into the HF tracking schema.
|
||||
|
||||
The Buildkite artifact intentionally keeps the raw benchmark output from
|
||||
test_inference_performance.py. Baseline comparison, main-branch persistence,
|
||||
and manual baseline reseeds should all use this mapping so the stored HF
|
||||
records do not drift from the artifact schema.
|
||||
"""
|
||||
def _normalize_record(result: dict[str, Any]) -> dict[str, Any]:
|
||||
benchmark_id = result.get("benchmark_id", "unknown")
|
||||
model_id = benchmark_id
|
||||
|
||||
@@ -71,9 +71,9 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
timestamp = datetime.now(timezone.utc).isoformat()
|
||||
|
||||
commit_sha = result.get("commit") or os.environ.get("BUILDKITE_COMMIT", "")
|
||||
latency = safe_float(result.get("avg_generation_time_s"))
|
||||
throughput = safe_float(result.get("throughput_fps"))
|
||||
memory = safe_float(result.get("max_peak_memory_mb"))
|
||||
latency = _safe_float(result.get("avg_generation_time_s"))
|
||||
throughput = _safe_float(result.get("throughput_fps"))
|
||||
memory = _safe_float(result.get("max_peak_memory_mb"))
|
||||
|
||||
return {
|
||||
"model_id": model_id,
|
||||
@@ -87,16 +87,12 @@ def normalize_performance_result(result: dict[str, Any]) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _normalize_record(result: dict[str, Any]) -> dict[str, Any]:
|
||||
return normalize_performance_result(result)
|
||||
|
||||
|
||||
def _write_tracking_record(record: dict[str, Any]) -> str:
|
||||
model_dir = os.path.join(TRACKING_ROOT, sanitize(record["model_id"]))
|
||||
model_dir = os.path.join(TRACKING_ROOT, _sanitize(record["model_id"]))
|
||||
os.makedirs(model_dir, exist_ok=True)
|
||||
|
||||
timestamp = sanitize(record["timestamp"])
|
||||
commit = sanitize(record["commit_sha"] or "unknown")
|
||||
timestamp = _sanitize(record["timestamp"])
|
||||
commit = _sanitize(record["commit_sha"] or "unknown")
|
||||
out_path = os.path.join(model_dir, f"{timestamp}_{commit}.json")
|
||||
|
||||
with open(out_path, "w", encoding="utf-8") as f:
|
||||
@@ -104,27 +100,11 @@ def _write_tracking_record(record: dict[str, Any]) -> str:
|
||||
|
||||
return out_path
|
||||
|
||||
|
||||
def _write_normalized_artifact(record: dict[str, Any]) -> None:
|
||||
try:
|
||||
results_dir = os.path.join(PERF_REPORTS_DIR, "results")
|
||||
os.makedirs(results_dir, exist_ok=True)
|
||||
timestamp = sanitize(record["timestamp"])
|
||||
model_id = sanitize(record["model_id"])
|
||||
commit = sanitize(record["commit_sha"] or "unknown")
|
||||
path = os.path.join(
|
||||
results_dir,
|
||||
f"normalized_perf_{model_id}_{timestamp}_{commit}.json",
|
||||
)
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
json.dump(record, f, indent=2)
|
||||
print(f"Normalized performance result written to {path}")
|
||||
except Exception as e:
|
||||
print(f"Failed to write normalized performance result artifact: {e}")
|
||||
|
||||
|
||||
def _baseline_metric(records: list[dict[str, Any]], key: str) -> float | None:
|
||||
values = [safe_float(r.get(key)) for r in records]
|
||||
values = [
|
||||
_safe_float(r.get(key))
|
||||
for r in records
|
||||
]
|
||||
values = [v for v in values if v is not None]
|
||||
if not values:
|
||||
return None
|
||||
@@ -140,23 +120,25 @@ def _check_regressions(
|
||||
|
||||
for metric in ("latency", "memory"):
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
curr = safe_float(current.get(metric))
|
||||
curr = _safe_float(current.get(metric))
|
||||
if baseline is None or curr is None or baseline <= 0:
|
||||
continue
|
||||
regression = (curr - baseline) / baseline
|
||||
if regression > max_regression:
|
||||
failures.append(f"{current['model_id']} {metric} regressed by "
|
||||
f"{regression * 100:.1f}% "
|
||||
f"(current={curr:.3f}, baseline_median={baseline:.3f})")
|
||||
failures.append(
|
||||
f"{current['model_id']} {metric} regressed by {regression * 100:.1f}% "
|
||||
f"(current={curr:.3f}, baseline_median={baseline:.3f})"
|
||||
)
|
||||
|
||||
baseline_tp = _baseline_metric(baseline_records, "throughput")
|
||||
curr_tp = safe_float(current.get("throughput"))
|
||||
curr_tp = _safe_float(current.get("throughput"))
|
||||
if baseline_tp is not None and curr_tp is not None and baseline_tp > 0:
|
||||
regression = (baseline_tp - curr_tp) / baseline_tp
|
||||
if regression > max_regression:
|
||||
failures.append(f"{current['model_id']} throughput regressed by "
|
||||
f"{regression * 100:.1f}% "
|
||||
f"(current={curr_tp:.3f}, baseline_median={baseline_tp:.3f})")
|
||||
failures.append(
|
||||
f"{current['model_id']} throughput regressed by {regression * 100:.1f}% "
|
||||
f"(current={curr_tp:.3f}, baseline_median={baseline_tp:.3f})"
|
||||
)
|
||||
|
||||
return failures
|
||||
|
||||
@@ -166,7 +148,7 @@ def _metric_delta_percent(
|
||||
current: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
) -> float | None:
|
||||
curr = safe_float(current.get(metric))
|
||||
curr = _safe_float(current.get(metric))
|
||||
baseline = _baseline_metric(baseline_records, metric)
|
||||
if curr is None or baseline is None or baseline <= 0:
|
||||
return None
|
||||
@@ -183,23 +165,22 @@ def _compact_value(value: float | None, precision: int = 3) -> str:
|
||||
return "n/a"
|
||||
return f"{value:.{precision}f}"
|
||||
|
||||
|
||||
def _build_summary_row(
|
||||
record: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
has_failed: bool,
|
||||
record: dict[str, Any],
|
||||
baseline_records: list[dict[str, Any]],
|
||||
has_failed: bool
|
||||
) -> dict[str, Any]:
|
||||
"""Format a single benchmark result as a row for the Markdown table."""
|
||||
|
||||
latency_base = _baseline_metric(baseline_records, "latency")
|
||||
throughput_base = _baseline_metric(baseline_records, "throughput")
|
||||
memory_base = _baseline_metric(baseline_records, "memory")
|
||||
|
||||
"""Formats a single benchmark result into a row for the Markdown summary table."""
|
||||
|
||||
latency_base = _safe_float(_baseline_metric(baseline_records, "latency"))
|
||||
throughput_base = _safe_float(_baseline_metric(baseline_records, "throughput"))
|
||||
memory_base = _safe_float(_baseline_metric(baseline_records, "memory"))
|
||||
|
||||
# Calculate percentages for the 'Worst Regression' column
|
||||
latency_reg = _metric_delta_percent("latency", record, baseline_records)
|
||||
throughput_reg = _metric_delta_percent("throughput", record, baseline_records)
|
||||
memory_reg = _metric_delta_percent("memory", record, baseline_records)
|
||||
|
||||
|
||||
regressions = [v for v in (latency_reg, throughput_reg, memory_reg) if v is not None]
|
||||
worst_regression_pct = max(regressions) if regressions else None
|
||||
|
||||
@@ -207,17 +188,16 @@ def _build_summary_row(
|
||||
"model_id": record["model_id"],
|
||||
"gpu_type": record["gpu_type"],
|
||||
"baseline_n": len(baseline_records),
|
||||
"latency_curr": safe_float(record.get("latency")),
|
||||
"latency_curr": _safe_float(record.get("latency")),
|
||||
"latency_base": latency_base,
|
||||
"throughput_curr": safe_float(record.get("throughput")),
|
||||
"throughput_curr": _safe_float(record.get("throughput")),
|
||||
"throughput_base": throughput_base,
|
||||
"memory_curr": safe_float(record.get("memory")),
|
||||
"memory_curr": _safe_float(record.get("memory")),
|
||||
"memory_base": memory_base,
|
||||
"worst_regression_pct": worst_regression_pct,
|
||||
"failed": has_failed,
|
||||
}
|
||||
|
||||
|
||||
def _build_markdown_summary(
|
||||
summary_rows: list[dict[str, Any]],
|
||||
max_regression: float,
|
||||
@@ -227,34 +207,28 @@ def _build_markdown_summary(
|
||||
"",
|
||||
f"Threshold: regressions greater than {max_regression * 100:.1f}% fail",
|
||||
"",
|
||||
("| Model | GPU | Baseline N | Latency (curr/base) | "
|
||||
"Throughput (curr/base) | Memory (curr/base) | "
|
||||
"Worst Regression | Status |"),
|
||||
"| Model | GPU | Baseline N | Latency (curr/base) | Throughput (curr/base) | Memory (curr/base) | Worst Regression | Status |",
|
||||
"|---|---|---:|---|---|---|---:|---|",
|
||||
]
|
||||
|
||||
for row in summary_rows:
|
||||
latency = (f"{_compact_value(row['latency_curr'])} / "
|
||||
f"{_compact_value(row['latency_base'])}")
|
||||
throughput = (f"{_compact_value(row['throughput_curr'])} / "
|
||||
f"{_compact_value(row['throughput_base'])}")
|
||||
memory = (f"{_compact_value(row['memory_curr'], 1)} / "
|
||||
f"{_compact_value(row['memory_base'], 1)}")
|
||||
latency = f"{_compact_value(row['latency_curr'])} / {_compact_value(row['latency_base'])}"
|
||||
throughput = f"{_compact_value(row['throughput_curr'])} / {_compact_value(row['throughput_base'])}"
|
||||
memory = f"{_compact_value(row['memory_curr'], 1)} / {_compact_value(row['memory_base'], 1)}"
|
||||
|
||||
worst_reg = ("n/a" if row["worst_regression_pct"] is None else f"{row['worst_regression_pct']:.1f}%")
|
||||
worst_reg = "n/a" if row["worst_regression_pct"] is None else f"{row['worst_regression_pct']:.1f}%"
|
||||
status = "FAIL" if row["failed"] else "PASS"
|
||||
|
||||
lines.append(f"| {row['model_id']} | {row['gpu_type']} | "
|
||||
f"{row['baseline_n']} | "
|
||||
f"{latency} | {throughput} | {memory} | "
|
||||
f"{worst_reg} | {status} |")
|
||||
lines.append(
|
||||
f"| {row['model_id']} | {row['gpu_type']} | {row['baseline_n']} | "
|
||||
f"{latency} | {throughput} | {memory} | {worst_reg} | {status} |"
|
||||
)
|
||||
|
||||
return "\n".join(lines) + "\n"
|
||||
|
||||
|
||||
def _emit_markdown_summary(markdown: str, commit_sha: str) -> None:
|
||||
print("\n" + markdown)
|
||||
|
||||
|
||||
# 1. Existing GitHub logic (safe to keep)
|
||||
summary_path = os.environ.get("GITHUB_STEP_SUMMARY")
|
||||
if summary_path:
|
||||
@@ -263,68 +237,73 @@ def _emit_markdown_summary(markdown: str, commit_sha: str) -> None:
|
||||
|
||||
# 2. Write to Modal volume for Buildkite to pick up in post-run hook
|
||||
try:
|
||||
os.makedirs(PERF_REPORTS_DIR, exist_ok=True)
|
||||
perf_reports_dir = "/root/data/perf_reports"
|
||||
os.makedirs(perf_reports_dir, exist_ok=True)
|
||||
short_sha = commit_sha[:7] if commit_sha else "unknown"
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
report_path = os.path.join(PERF_REPORTS_DIR, f"perf_{short_sha}_{timestamp}.md")
|
||||
report_path = os.path.join(perf_reports_dir, f"perf_{short_sha}_{timestamp}.md")
|
||||
with open(report_path, "w", encoding="utf-8") as f:
|
||||
f.write(markdown + "\n")
|
||||
print(f"Performance report written to {report_path}")
|
||||
except Exception as e:
|
||||
print(f"Failed to write performance report to Modal volume: {e}")
|
||||
|
||||
|
||||
def main() -> int:
|
||||
persist_tracking = _should_persist_tracking()
|
||||
|
||||
# Strict on persist: a silent sync failure would pollute the baseline.
|
||||
sync_from_hf(TRACKING_ROOT, strict=persist_tracking)
|
||||
# Pull the current state of the world from HF
|
||||
sync_from_hf(TRACKING_ROOT)
|
||||
|
||||
current_results = _load_current_results()
|
||||
if not current_results:
|
||||
print(f"No performance result files found in {RESULTS_DIR}")
|
||||
return 0
|
||||
|
||||
all_failures: list[str] = []
|
||||
summary_rows: list[dict[str, Any]] = []
|
||||
all_failures = []
|
||||
summary_rows = []
|
||||
persist_tracking = _should_persist_tracking()
|
||||
|
||||
if persist_tracking:
|
||||
print("Tracking persistence enabled: full-suite run on main branch")
|
||||
else:
|
||||
print("Tracking persistence disabled: "
|
||||
"only full-suite runs on main branch are persisted")
|
||||
print("Tracking persistence disabled: only full-suite runs on main branch are persisted")
|
||||
|
||||
for raw in current_results:
|
||||
record = _normalize_record(raw)
|
||||
|
||||
baseline_records = load_records_for_model(
|
||||
TRACKING_ROOT,
|
||||
record["model_id"],
|
||||
record["gpu_type"],
|
||||
last_n=5,
|
||||
successful_only=True,
|
||||
TRACKING_ROOT, record["model_id"], record["gpu_type"],
|
||||
last_n=5, successful_only=True
|
||||
)
|
||||
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
|
||||
# Tag the current record based on the failure.
|
||||
if not baseline_records:
|
||||
print(f"No baseline for {record['model_id']} on "
|
||||
f"{record['gpu_type']}. Initializing...")
|
||||
failures: list[str] = []
|
||||
record["success"] = True
|
||||
# INITIALIZATION CASE: First run for this model/GPU
|
||||
print(f"No baseline for {record['model_id']} on {record['gpu_type']}. Initializing...")
|
||||
failures = []
|
||||
record["success"] = True # The first run is always "successful"
|
||||
else:
|
||||
# COMPARISON CASE: Compare against the mean of the last 5 good runs
|
||||
failures = _check_regressions(record, baseline_records, MAX_REGRESSION)
|
||||
record["success"] = not failures
|
||||
all_failures.extend(failures)
|
||||
|
||||
_write_normalized_artifact(record)
|
||||
|
||||
# Strict upload: a silent failure would freeze the rolling baseline.
|
||||
if failures:
|
||||
record["success"] = False
|
||||
all_failures.extend(failures)
|
||||
else:
|
||||
record["success"] = True
|
||||
|
||||
# 5. Persist to HF if we are on main
|
||||
if persist_tracking:
|
||||
# This writes the JSON with the "success" field to /tmp
|
||||
current_path = _write_tracking_record(record)
|
||||
upload_record(current_path, record, strict=True)
|
||||
# This pushes it to the FastVideo/performance-tracking repo
|
||||
upload_record(current_path, record)
|
||||
|
||||
summary_rows.append(_build_summary_row(record, baseline_records, bool(failures)))
|
||||
|
||||
summary_row = _build_summary_row(record, baseline_records, bool(failures))
|
||||
summary_rows.append(summary_row)
|
||||
|
||||
commit_sha = os.environ.get("BUILDKITE_COMMIT", "unknown")[:7]
|
||||
|
||||
markdown = _build_markdown_summary(summary_rows, MAX_REGRESSION)
|
||||
_emit_markdown_summary(markdown, commit_sha)
|
||||
|
||||
@@ -337,6 +316,5 @@ def main() -> int:
|
||||
print("Performance baseline comparison passed")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import shutil
|
||||
import subprocess
|
||||
from datetime import datetime
|
||||
|
||||
import plotly.express as px
|
||||
|
||||
@@ -26,11 +26,11 @@ from huggingface_hub import HfApi, snapshot_download
|
||||
HF_REPO_ID: str = os.environ.get("HF_REPO_ID", "FastVideo/performance-tracking")
|
||||
HF_TOKEN: str | None = os.environ.get("HF_API_KEY")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Low-level helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def sanitize(value: str) -> str:
|
||||
"""Return a filesystem- and HF-path-safe version of *value*."""
|
||||
return re.sub(r"[^A-Za-z0-9._-]", "_", value)
|
||||
@@ -50,23 +50,16 @@ def safe_float(value: Any) -> float | None:
|
||||
# HF I/O
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def sync_from_hf(local_dir: str, *, strict: bool = False) -> str:
|
||||
def sync_from_hf(local_dir: str) -> str:
|
||||
"""Download the HF dataset repo snapshot to *local_dir*.
|
||||
|
||||
Returns *local_dir* so callers can chain: ``load_records(sync_from_hf(...))``.
|
||||
|
||||
By default (``strict=False``) failures are logged and *local_dir* is
|
||||
returned unchanged, so dashboard / PR consumers stay resilient when HF is
|
||||
unavailable. Callers that depend on the sync for correctness (e.g. the
|
||||
main-branch baseline writer) must pass ``strict=True`` so that misconfig
|
||||
or transient HF errors fail loud rather than silently reset the baseline.
|
||||
On failure (empty repo, no credentials, network error) the function logs a
|
||||
warning and returns *local_dir* unchanged so the caller can still work with
|
||||
whatever is already on disk.
|
||||
"""
|
||||
if not HF_REPO_ID:
|
||||
msg = "hf_store: HF_REPO_ID not set"
|
||||
if strict:
|
||||
raise RuntimeError(f"{msg}; cannot sync.")
|
||||
print(f"{msg}, skipping sync.")
|
||||
print("hf_store: HF_REPO_ID not set, skipping sync.")
|
||||
return local_dir
|
||||
|
||||
print(f"hf_store: syncing from {HF_REPO_ID} → {local_dir}")
|
||||
@@ -79,31 +72,19 @@ def sync_from_hf(local_dir: str, *, strict: bool = False) -> str:
|
||||
allow_patterns="*.json",
|
||||
)
|
||||
except Exception as exc:
|
||||
if strict:
|
||||
raise
|
||||
print(f"hf_store: sync skipped — {exc}")
|
||||
|
||||
return local_dir
|
||||
|
||||
|
||||
def upload_record(
|
||||
local_path: str,
|
||||
record: dict[str, Any],
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> None:
|
||||
def upload_record(local_path: str, record: dict[str, Any]) -> None:
|
||||
"""Upload *local_path* to the HF repo under ``<model_id>/<filename>``.
|
||||
|
||||
By default failures (missing token, network errors) are logged and
|
||||
swallowed. Pass ``strict=True`` when the upload is part of a write-path
|
||||
that must not silently lose records — otherwise the rolling baseline can
|
||||
stop advancing without any signal in the build log.
|
||||
Silently skips if HF_TOKEN is absent so local/CI runs without credentials
|
||||
don't crash.
|
||||
"""
|
||||
if not HF_TOKEN:
|
||||
msg = "hf_store: HF_API_KEY not set"
|
||||
if strict:
|
||||
raise RuntimeError(f"{msg}; cannot upload.")
|
||||
print(f"{msg}, skipping upload.")
|
||||
print("hf_store: HF_API_KEY not set, skipping upload.")
|
||||
return
|
||||
|
||||
model_id = record.get("model_id", "unknown")
|
||||
@@ -121,8 +102,6 @@ def upload_record(
|
||||
)
|
||||
print(f"hf_store: uploaded → {HF_REPO_ID}/{path_in_repo}")
|
||||
except Exception as exc:
|
||||
if strict:
|
||||
raise
|
||||
print(f"hf_store: upload failed — {exc}")
|
||||
|
||||
|
||||
@@ -130,7 +109,6 @@ def upload_record(
|
||||
# Record loading
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def load_records(
|
||||
local_dir: str,
|
||||
*,
|
||||
@@ -142,7 +120,7 @@ def load_records(
|
||||
Args:
|
||||
local_dir: Root directory previously populated by :func:`sync_from_hf`.
|
||||
days: When set, discard records whose ``timestamp`` is older than this
|
||||
many days. Records with a missing/unparsable timestamp are kept.
|
||||
many days. Records with a missing/unparseable timestamp are kept.
|
||||
successful_only: When True, only records with ``success=True`` are
|
||||
returned. Useful when building a regression baseline.
|
||||
|
||||
@@ -176,7 +154,7 @@ def load_records(
|
||||
if ts < cutoff:
|
||||
continue
|
||||
except ValueError:
|
||||
pass # keep records with unparsable timestamps
|
||||
pass # keep records with unparseable timestamps
|
||||
|
||||
records.append(data)
|
||||
|
||||
@@ -273,4 +251,4 @@ def load_as_dataframe(
|
||||
return pd.DataFrame()
|
||||
|
||||
df = pd.DataFrame(records)
|
||||
return normalize_dataframe(df)
|
||||
return normalize_dataframe(df)
|
||||
@@ -1,7 +1,5 @@
|
||||
# SSIM Directory Guidelines
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
## Scope
|
||||
These instructions apply to everything under `fastvideo/tests/ssim/`.
|
||||
|
||||
|
||||
@@ -132,7 +132,7 @@ def _assert_similarity(
|
||||
)
|
||||
|
||||
|
||||
def build_init_kwargs(
|
||||
def _build_init_kwargs(
|
||||
base_params: dict[str, object],
|
||||
) -> dict[str, object]:
|
||||
init_kwargs: dict[str, object] = {
|
||||
@@ -167,7 +167,7 @@ def build_init_kwargs(
|
||||
return init_kwargs
|
||||
|
||||
|
||||
def build_generation_kwargs(
|
||||
def _build_generation_kwargs(
|
||||
base_params: dict[str, object],
|
||||
num_inference_steps: int,
|
||||
output_dir: str,
|
||||
@@ -222,11 +222,11 @@ def run_text_to_video_similarity_test(
|
||||
base_params = params_map[model_id]
|
||||
num_inference_steps = int(base_params["num_inference_steps"])
|
||||
|
||||
init_kwargs = build_init_kwargs(base_params)
|
||||
init_kwargs = _build_init_kwargs(base_params)
|
||||
if init_kwargs_override:
|
||||
init_kwargs.update(init_kwargs_override)
|
||||
|
||||
generation_kwargs = build_generation_kwargs(
|
||||
generation_kwargs = _build_generation_kwargs(
|
||||
base_params,
|
||||
num_inference_steps,
|
||||
output_dir,
|
||||
@@ -297,11 +297,11 @@ def run_image_to_video_similarity_test(
|
||||
base_params = params_map[model_id]
|
||||
num_inference_steps = int(base_params["num_inference_steps"])
|
||||
|
||||
init_kwargs = build_init_kwargs(base_params)
|
||||
init_kwargs = _build_init_kwargs(base_params)
|
||||
if init_kwargs_override:
|
||||
init_kwargs.update(init_kwargs_override)
|
||||
|
||||
generation_kwargs = build_generation_kwargs(
|
||||
generation_kwargs = _build_generation_kwargs(
|
||||
base_params,
|
||||
num_inference_steps,
|
||||
output_dir,
|
||||
|
||||
@@ -1,494 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Latent-space regression helpers for numerically fragile SSIM tests.
|
||||
|
||||
Motivation
|
||||
----------
|
||||
Pixel-space SSIM is a poor regression signal for distilled / few-step
|
||||
models (e.g. LTX-2 distilled): a single mis-rounded bf16 accumulator in
|
||||
the VAE decoder can drive mean SSIM from ~0.95 to ~0.50 without any real
|
||||
quality regression.
|
||||
|
||||
Inspired by diffusers' "small signature slice + bounded full-tensor
|
||||
distance" testing philosophy, applied here to the **pre-VAE latent**
|
||||
rather than the decoded pixel/audio output:
|
||||
|
||||
* `tests/pipelines/ltx2/test_ltx2.py` (diffusers) compares
|
||||
``output_type='pt'`` (pixel) slices via
|
||||
``torch.allclose(generated_slice, expected_slice, atol=1e-4)``;
|
||||
* `tests/pipelines/stable_audio/test_stable_audio.py` (diffusers)
|
||||
compares decoded audio samples via
|
||||
``np.abs(expected - actual).max() < 1.5e-3``;
|
||||
* `tests/pipelines/cogvideo/test_cogvideox.py` (diffusers) compares
|
||||
full pixel video tensors via
|
||||
``numpy_cosine_similarity_distance(...) < 1e-3``.
|
||||
|
||||
Diffusers does *not* assert on latents directly — that is a FastVideo
|
||||
adaptation. Distilled few-step pipelines amplify per-step bf16 noise
|
||||
enough that VAE-decoded comparisons are unreliable across our
|
||||
heterogeneous CI pool, so we move the assertion upstream of the VAE.
|
||||
|
||||
Design
|
||||
------
|
||||
* Inference is run with ``output_type='latent'`` so ``DecodingStage``
|
||||
hands back the un-decoded latent on ``result["samples"]``.
|
||||
* The reference artefact is a ``.pt`` bundle (tensor + metadata) hosted
|
||||
on the same HF dataset as the mp4 references, selected by
|
||||
``<GPU>_reference_videos/<model_id>/<backend>/<prompt>.pt``.
|
||||
* Two assertions are performed:
|
||||
1. A small signature slice (default video: ``latent[0, :, 0, :3, :3]``;
|
||||
audio: ``latent[0, :, :8]``) is compared via cosine distance with
|
||||
a loose tolerance. Primary pass/fail gate.
|
||||
2. The full latent is compared via cosine distance with a slightly
|
||||
tighter tolerance, guarding against shape-correct but globally
|
||||
drifted outputs.
|
||||
* Tolerances default to 5e-3 (slice) and 1e-2 (full). diffusers uses
|
||||
``1e-3`` against deterministic CPU dummy components; we relax for
|
||||
cross-GPU-arch bf16 differences on the rented CI pool
|
||||
(A40/L40S/H100/B200).
|
||||
|
||||
The helper intentionally reuses ``build_init_kwargs`` /
|
||||
``build_generation_kwargs`` from :mod:`inference_similarity_utils` so
|
||||
model params (vae tiling, sp_size, flow shift, …) flow through a single
|
||||
source of truth.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from logging import Logger
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch.nn.functional import cosine_similarity
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
attention_backend,
|
||||
build_generation_kwargs,
|
||||
build_init_kwargs,
|
||||
shutdown_executor,
|
||||
)
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
build_generated_output_dir,
|
||||
build_reference_folder_path,
|
||||
select_ssim_params,
|
||||
)
|
||||
|
||||
LATENT_REFERENCE_EXTENSION = ".pt"
|
||||
LATENT_REFERENCE_FORMAT_VERSION = 1
|
||||
|
||||
DEFAULT_SLICE_SPEC: dict[str, Any] = {
|
||||
"kind": "corner_3x3_first_frame",
|
||||
"version": 1,
|
||||
}
|
||||
|
||||
AUDIO_FIRST_8_TIMESTEPS_SPEC: dict[str, Any] = {
|
||||
"kind": "audio_first_8_timesteps",
|
||||
"version": 1,
|
||||
}
|
||||
|
||||
|
||||
def _extract_expected_slice(
|
||||
latent: torch.Tensor,
|
||||
spec: dict[str, Any],
|
||||
) -> torch.Tensor:
|
||||
"""Return a 1-D fp32 signature slice from ``latent`` per ``spec``.
|
||||
|
||||
Dispatch by ``spec["kind"]``:
|
||||
|
||||
* ``"corner_3x3_first_frame"`` — 5-D video latent ``[B, C, T, H, W]``;
|
||||
returns ``latent[0, :, 0, :3, :3]`` flattened (length = C * 9).
|
||||
* ``"audio_first_8_timesteps"`` — 3-D audio latent ``[B, C, T]``;
|
||||
returns ``latent[0, :, :8]`` flattened (length = C * 8).
|
||||
"""
|
||||
kind = spec.get("kind", "corner_3x3_first_frame")
|
||||
if kind == "corner_3x3_first_frame":
|
||||
if latent.dim() != 5:
|
||||
raise ValueError(
|
||||
f"corner_3x3_first_frame requires 5-D latent [B,C,T,H,W]; "
|
||||
f"got shape {tuple(latent.shape)}")
|
||||
_, _, t, h, w = latent.shape
|
||||
if t < 1 or h < 3 or w < 3:
|
||||
raise ValueError(
|
||||
"corner_3x3_first_frame requires T>=1, H>=3, W>=3; got "
|
||||
f"shape {tuple(latent.shape)}")
|
||||
return (latent[0, :, 0, :3, :3]
|
||||
.detach()
|
||||
.to(torch.float32)
|
||||
.reshape(-1)
|
||||
.contiguous())
|
||||
if kind == "audio_first_8_timesteps":
|
||||
if latent.dim() != 3:
|
||||
raise ValueError(
|
||||
f"audio_first_8_timesteps requires 3-D latent [B,C,T]; "
|
||||
f"got shape {tuple(latent.shape)}")
|
||||
_, _, t = latent.shape
|
||||
if t < 8:
|
||||
raise ValueError(
|
||||
"audio_first_8_timesteps requires T>=8; got "
|
||||
f"shape {tuple(latent.shape)}")
|
||||
return (latent[0, :, :8]
|
||||
.detach()
|
||||
.to(torch.float32)
|
||||
.reshape(-1)
|
||||
.contiguous())
|
||||
raise ValueError(f"Unknown slice kind: {kind!r}")
|
||||
|
||||
|
||||
def _cosine_distance(a: torch.Tensor, b: torch.Tensor) -> float:
|
||||
"""Return ``1 - cos(a, b)`` as a Python float, operating in fp32."""
|
||||
a32 = a.detach().to(torch.float32).reshape(-1)
|
||||
b32 = b.detach().to(torch.float32).reshape(-1)
|
||||
if a32.shape != b32.shape:
|
||||
raise ValueError(
|
||||
f"Cosine shape mismatch: {tuple(a32.shape)} vs {tuple(b32.shape)}"
|
||||
)
|
||||
sim = cosine_similarity(a32.unsqueeze(0), b32.unsqueeze(0), dim=1).item()
|
||||
return 1.0 - float(sim)
|
||||
|
||||
|
||||
def save_latent_reference(
|
||||
path: str,
|
||||
latent: torch.Tensor,
|
||||
*,
|
||||
metadata: dict[str, Any],
|
||||
slice_spec: dict[str, Any] | None = None,
|
||||
) -> None:
|
||||
"""Persist a latent bundle to ``path``.
|
||||
|
||||
Storage format (dict pickled via ``torch.save``):
|
||||
|
||||
* ``latent``: full latent as fp16 on cpu
|
||||
* ``shape``: original shape (list)
|
||||
* ``dtype_original``: str
|
||||
* ``expected_slice``: fp32 1-D signature slice
|
||||
* ``slice_spec``: dict describing how the slice was built
|
||||
* ``metadata``: caller-provided context (prompt, backend, steps, …)
|
||||
* ``format_version``: int
|
||||
|
||||
fp16 is lossy but bounded; it keeps ref artefacts small (~a few MB
|
||||
per prompt) while preserving enough dynamic range for cosine-based
|
||||
regression. Slice values stay fp32 because the primary assertion is
|
||||
computed against them.
|
||||
"""
|
||||
spec = slice_spec if slice_spec is not None else DEFAULT_SLICE_SPEC
|
||||
parent = os.path.dirname(path)
|
||||
if parent:
|
||||
os.makedirs(parent, exist_ok=True)
|
||||
|
||||
latent_cpu = latent.detach().to("cpu")
|
||||
payload: dict[str, Any] = {
|
||||
"latent": latent_cpu.to(torch.float16),
|
||||
"shape": list(latent_cpu.shape),
|
||||
"dtype_original": str(latent_cpu.dtype),
|
||||
"expected_slice": _extract_expected_slice(latent_cpu, spec),
|
||||
"slice_spec": spec,
|
||||
"metadata": metadata,
|
||||
"format_version": LATENT_REFERENCE_FORMAT_VERSION,
|
||||
}
|
||||
torch.save(payload, path)
|
||||
|
||||
|
||||
def load_latent_reference(path: str) -> dict[str, Any]:
|
||||
"""Inverse of :func:`save_latent_reference` — always loads to cpu.
|
||||
|
||||
Enforces ``format_version == LATENT_REFERENCE_FORMAT_VERSION`` so a
|
||||
schema change forces a deliberate reseed instead of silently
|
||||
misinterpreting old artefacts.
|
||||
"""
|
||||
# ``weights_only=False`` is required because the payload is a dict of
|
||||
# tensors + plain-Python metadata (slice_spec, prompt, ...). The trust
|
||||
# boundary is the controlled HF dataset configured via
|
||||
# FASTVIDEO_SSIM_REFERENCE_HF_REPO (default
|
||||
# FastVideo/ssim-reference-videos), which is org-write-gated.
|
||||
payload = torch.load(path, map_location="cpu", weights_only=False)
|
||||
fmt = payload.get("format_version") if isinstance(payload, dict) else None
|
||||
if fmt != LATENT_REFERENCE_FORMAT_VERSION:
|
||||
raise ValueError(
|
||||
f"Latent reference at {path!r} has format_version={fmt!r}; "
|
||||
f"expected {LATENT_REFERENCE_FORMAT_VERSION}. Re-seed via the "
|
||||
"test that produces this artefact, then re-upload through "
|
||||
"fastvideo/tests/ssim/reference_videos_cli.py upload.")
|
||||
return payload
|
||||
|
||||
|
||||
def write_latent_similarity_results(
|
||||
output_dir: str,
|
||||
metrics: dict[str, float],
|
||||
*,
|
||||
reference_path: str,
|
||||
generated_path: str,
|
||||
num_inference_steps: int,
|
||||
prompt: str,
|
||||
model_id: str,
|
||||
attention_backend_name: str,
|
||||
slice_spec: dict[str, Any],
|
||||
slice_cosine_threshold: float,
|
||||
full_cosine_threshold: float,
|
||||
passed: bool,
|
||||
) -> bool:
|
||||
"""Persist latent regression metrics next to the generated artefact.
|
||||
|
||||
Mirrors :func:`fastvideo.tests.utils.write_ssim_results` so downstream
|
||||
CI tooling can scrape one schema for both pixel and latent runs.
|
||||
The filename is ``steps{N}_{prompt[:100]}_latent.json``.
|
||||
"""
|
||||
try:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
prompt_prefix = prompt[:100].strip()
|
||||
filename = f"steps{num_inference_steps}_{prompt_prefix}_latent.json"
|
||||
target = os.path.join(output_dir, filename)
|
||||
payload = {
|
||||
"metrics": metrics,
|
||||
"reference_latent": reference_path,
|
||||
"generated_latent": generated_path,
|
||||
"model_id": model_id,
|
||||
"attention_backend": attention_backend_name,
|
||||
"slice_spec": slice_spec,
|
||||
"thresholds": {
|
||||
"slice_cosine": slice_cosine_threshold,
|
||||
"full_cosine": full_cosine_threshold,
|
||||
},
|
||||
"passed": passed,
|
||||
"parameters": {
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"prompt": prompt,
|
||||
},
|
||||
}
|
||||
with open(target, "w", encoding="utf-8") as handle:
|
||||
json.dump(payload, handle, indent=2, sort_keys=True)
|
||||
return True
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
|
||||
def _assert_latent_similarity(
|
||||
*,
|
||||
logger: Logger,
|
||||
gen_latent: torch.Tensor,
|
||||
reference_path: str,
|
||||
slice_cosine_threshold: float,
|
||||
full_cosine_threshold: float,
|
||||
model_id: str,
|
||||
attention_backend_name: str,
|
||||
output_dir: str,
|
||||
generated_path: str,
|
||||
num_inference_steps: int,
|
||||
prompt: str,
|
||||
) -> dict[str, float]:
|
||||
ref = load_latent_reference(reference_path)
|
||||
expected_slice = ref["expected_slice"]
|
||||
ref_full = ref["latent"].to(torch.float32)
|
||||
ref_shape = tuple(ref.get("shape", ref_full.shape))
|
||||
slice_spec = ref.get("slice_spec", DEFAULT_SLICE_SPEC)
|
||||
|
||||
if tuple(gen_latent.shape) != ref_shape:
|
||||
raise AssertionError(
|
||||
f"Generated latent shape {tuple(gen_latent.shape)} does not "
|
||||
f"match reference shape {ref_shape} for {model_id} with "
|
||||
f"backend {attention_backend_name}")
|
||||
|
||||
gen_slice = _extract_expected_slice(gen_latent, slice_spec)
|
||||
slice_cos = _cosine_distance(gen_slice, expected_slice)
|
||||
# The reference full tensor was fp16-quantized at seed time. Round-trip
|
||||
# the generated tensor through fp16 so both sides share the same
|
||||
# quantization floor and the cosine distance is symmetric. This does
|
||||
# NOT affect ``slice_cos`` because the reference slice is persisted in
|
||||
# fp32 (see :func:`save_latent_reference`).
|
||||
gen_full_matched = gen_latent.to(torch.float16).to(torch.float32)
|
||||
full_cos = _cosine_distance(gen_full_matched, ref_full)
|
||||
max_abs_diff = float(
|
||||
(gen_full_matched - ref_full).abs().max().item())
|
||||
|
||||
metrics: dict[str, float] = {
|
||||
"slice_cosine_distance": slice_cos,
|
||||
"full_cosine_distance": full_cos,
|
||||
"max_abs_diff": max_abs_diff,
|
||||
}
|
||||
logger.info(
|
||||
"Latent regression metrics for %s/%s: %s",
|
||||
model_id,
|
||||
attention_backend_name,
|
||||
metrics,
|
||||
)
|
||||
|
||||
failures: list[str] = []
|
||||
if slice_cos > slice_cosine_threshold:
|
||||
failures.append(
|
||||
f"slice cosine {slice_cos:.6e} > threshold "
|
||||
f"{slice_cosine_threshold:.6e}")
|
||||
if full_cos > full_cosine_threshold:
|
||||
failures.append(
|
||||
f"full cosine {full_cos:.6e} > threshold "
|
||||
f"{full_cosine_threshold:.6e}")
|
||||
|
||||
passed = not failures
|
||||
write_latent_similarity_results(
|
||||
output_dir,
|
||||
metrics,
|
||||
reference_path=reference_path,
|
||||
generated_path=generated_path,
|
||||
num_inference_steps=num_inference_steps,
|
||||
prompt=prompt,
|
||||
model_id=model_id,
|
||||
attention_backend_name=attention_backend_name,
|
||||
slice_spec=slice_spec,
|
||||
slice_cosine_threshold=slice_cosine_threshold,
|
||||
full_cosine_threshold=full_cosine_threshold,
|
||||
passed=passed,
|
||||
)
|
||||
|
||||
if failures:
|
||||
raise AssertionError(
|
||||
f"Latent regression exceeded tolerance for {model_id} with "
|
||||
f"backend {attention_backend_name}: {'; '.join(failures)}. "
|
||||
f"Full metrics: {metrics}")
|
||||
|
||||
return metrics
|
||||
|
||||
|
||||
def _extract_latent_from_result(result: Any) -> torch.Tensor:
|
||||
"""Pull a fp32 cpu latent tensor (5-D video or 3-D audio) from
|
||||
``generate_video`` output.
|
||||
"""
|
||||
if not isinstance(result, dict):
|
||||
raise RuntimeError(
|
||||
"VideoGenerator.generate_video returned unexpected payload "
|
||||
f"(type={type(result)!r}); expected dict with 'samples'.")
|
||||
samples = result.get("samples")
|
||||
if samples is None:
|
||||
raise RuntimeError(
|
||||
"VideoGenerator did not return latent samples. Ensure "
|
||||
"output_type='latent' and return_frames=True for this call.")
|
||||
if not isinstance(samples, torch.Tensor):
|
||||
raise RuntimeError(
|
||||
f"Expected torch.Tensor samples; got type={type(samples)!r}.")
|
||||
gen_latent = samples.detach().to(torch.float32).cpu()
|
||||
if gen_latent.dim() not in (3, 5):
|
||||
raise RuntimeError(
|
||||
"Expected 5-D video latent (B,C,T,H,W) or 3-D audio latent "
|
||||
f"(B,C,T); got shape {tuple(gen_latent.shape)}")
|
||||
return gen_latent
|
||||
|
||||
|
||||
def run_text_to_latent_similarity_test(
|
||||
*,
|
||||
logger: Logger,
|
||||
script_dir: str,
|
||||
device_reference_folder: str,
|
||||
prompt: str,
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
default_params_map: dict[str, dict[str, object]],
|
||||
full_quality_params_map: dict[str, dict[str, object]],
|
||||
slice_cosine_threshold: float = 5e-3,
|
||||
full_cosine_threshold: float = 1e-2,
|
||||
init_kwargs_override: dict[str, object] | None = None,
|
||||
generation_kwargs_override: dict[str, object] | None = None,
|
||||
slice_spec: dict[str, Any] | None = None,
|
||||
) -> dict[str, float]:
|
||||
"""Run T2V (or T2A) inference with ``output_type='latent'`` and
|
||||
compare to a reference latent.
|
||||
|
||||
Returns the computed metrics dict on success. Raises
|
||||
``AssertionError`` if any cosine tolerance is exceeded and
|
||||
``FileNotFoundError`` if the reference artefact is missing.
|
||||
"""
|
||||
spec = slice_spec if slice_spec is not None else DEFAULT_SLICE_SPEC
|
||||
with attention_backend(attention_backend_name):
|
||||
output_dir = build_generated_output_dir(
|
||||
script_dir,
|
||||
device_reference_folder,
|
||||
model_id,
|
||||
attention_backend_name,
|
||||
)
|
||||
prompt_prefix = prompt[:100].strip()
|
||||
output_latent_name = f"{prompt_prefix}{LATENT_REFERENCE_EXTENSION}"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
params_map = select_ssim_params(
|
||||
default_params_map,
|
||||
full_quality_params_map,
|
||||
)
|
||||
base_params = params_map[model_id]
|
||||
num_inference_steps = int(base_params["num_inference_steps"])
|
||||
|
||||
init_kwargs = build_init_kwargs(base_params)
|
||||
if init_kwargs_override:
|
||||
init_kwargs.update(init_kwargs_override)
|
||||
# Always wins: the helper exists specifically to compare on latents,
|
||||
# so an override can never silently turn it back into a pixel run.
|
||||
init_kwargs["output_type"] = "latent"
|
||||
|
||||
generation_kwargs = build_generation_kwargs(
|
||||
base_params,
|
||||
num_inference_steps,
|
||||
output_dir,
|
||||
)
|
||||
# We serialize latents ourselves; skip the RGB encoder path.
|
||||
generation_kwargs["save_video"] = False
|
||||
generation_kwargs["return_frames"] = True
|
||||
if generation_kwargs_override:
|
||||
generation_kwargs.update(generation_kwargs_override)
|
||||
|
||||
generator: VideoGenerator | None = None
|
||||
try:
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=base_params["model_path"],
|
||||
**init_kwargs,
|
||||
)
|
||||
result = generator.generate_video(prompt, **generation_kwargs)
|
||||
finally:
|
||||
shutdown_executor(generator)
|
||||
|
||||
gen_latent = _extract_latent_from_result(result)
|
||||
|
||||
generated_latent_path = os.path.join(output_dir, output_latent_name)
|
||||
save_latent_reference(
|
||||
generated_latent_path,
|
||||
gen_latent,
|
||||
metadata={
|
||||
"prompt": prompt,
|
||||
"model_id": model_id,
|
||||
"attention_backend": attention_backend_name,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
},
|
||||
slice_spec=spec,
|
||||
)
|
||||
logger.info("Saved generated latent to %s", generated_latent_path)
|
||||
|
||||
reference_folder = build_reference_folder_path(
|
||||
script_dir,
|
||||
device_reference_folder,
|
||||
model_id,
|
||||
attention_backend_name,
|
||||
)
|
||||
if not os.path.exists(reference_folder):
|
||||
raise FileNotFoundError(
|
||||
f"Reference folder does not exist: {reference_folder}\n"
|
||||
f"To download references, run:\n"
|
||||
f" python fastvideo/tests/ssim/reference_videos_cli.py download")
|
||||
|
||||
reference_latent_path = os.path.join(
|
||||
reference_folder,
|
||||
output_latent_name,
|
||||
)
|
||||
if not os.path.exists(reference_latent_path):
|
||||
raise FileNotFoundError(
|
||||
"Reference latent missing for prompt/backend: "
|
||||
f"{reference_latent_path}")
|
||||
|
||||
return _assert_latent_similarity(
|
||||
logger=logger,
|
||||
gen_latent=gen_latent,
|
||||
reference_path=reference_latent_path,
|
||||
slice_cosine_threshold=slice_cosine_threshold,
|
||||
full_cosine_threshold=full_cosine_threshold,
|
||||
model_id=model_id,
|
||||
attention_backend_name=attention_backend_name,
|
||||
output_dir=output_dir,
|
||||
generated_path=generated_latent_path,
|
||||
num_inference_steps=num_inference_steps,
|
||||
prompt=prompt,
|
||||
)
|
||||
@@ -13,12 +13,6 @@ from pathlib import Path
|
||||
|
||||
|
||||
VIDEO_EXTENSIONS = (".mp4", ".avi", ".mov", ".mkv", ".webm", ".flv")
|
||||
# Additional artefact types stored under the same reference folders. Latent
|
||||
# tensors (`.pt`) back the latent-slice regression tests used by flaky
|
||||
# pixel-space models (e.g. LTX-2 distilled). Keeping them in the same upload
|
||||
# flow means seeding a new test only requires one HF round-trip.
|
||||
LATENT_EXTENSIONS = (".pt",)
|
||||
REFERENCE_EXTENSIONS = VIDEO_EXTENSIONS + LATENT_EXTENSIONS
|
||||
HF_TOKEN_ENV_KEYS = ("HF_API_KEY", "HUGGINGFACE_HUB_TOKEN", "HF_TOKEN")
|
||||
HF_REPO_ENV_KEY = "FASTVIDEO_SSIM_REFERENCE_HF_REPO"
|
||||
HF_REPO_TYPE_ENV_KEY = "FASTVIDEO_SSIM_REFERENCE_HF_REPO_TYPE"
|
||||
@@ -48,14 +42,9 @@ def _default_repo_type() -> str:
|
||||
return os.environ.get(HF_REPO_TYPE_ENV_KEY, DEFAULT_REPO_TYPE)
|
||||
|
||||
|
||||
def _iter_reference_files(root: Path) -> Iterable[Path]:
|
||||
"""Yield video and latent (.pt) references under `root`.
|
||||
|
||||
Used by copy-local and the "has local references" marker probe so that
|
||||
latent-only tests (no mp4) still satisfy readiness checks.
|
||||
"""
|
||||
def _iter_video_files(root: Path) -> Iterable[Path]:
|
||||
for path in root.rglob("*"):
|
||||
if path.is_file() and path.suffix.lower() in REFERENCE_EXTENSIONS:
|
||||
if path.is_file() and path.suffix.lower() in VIDEO_EXTENSIONS:
|
||||
yield path
|
||||
|
||||
|
||||
@@ -150,7 +139,7 @@ def _has_local_reference_videos(base_dir: Path, quality_tier: str) -> bool:
|
||||
return False
|
||||
tier_root = _reference_tier_root(base_dir, quality_tier)
|
||||
for ref_dir in _discover_reference_dirs(tier_root):
|
||||
for _ in _iter_reference_files(ref_dir):
|
||||
for _ in _iter_video_files(ref_dir):
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -169,7 +158,7 @@ def _load_hf_sdk():
|
||||
from huggingface_hub import HfApi, snapshot_download
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
"huggingface_hub is required for download/upload.\nInstall with: uv pip install huggingface_hub"
|
||||
"huggingface_hub is required for download/upload.\nInstall with: pip install huggingface_hub"
|
||||
) from exc
|
||||
return HfApi, snapshot_download
|
||||
|
||||
@@ -184,14 +173,14 @@ def copy_generated_to_reference(
|
||||
raise FileNotFoundError(f"Generated directory not found: {generated_dir}")
|
||||
|
||||
copied = 0
|
||||
for ref_file in _iter_reference_files(generated_dir):
|
||||
rel = ref_file.relative_to(generated_dir)
|
||||
for video_file in _iter_video_files(generated_dir):
|
||||
rel = video_file.relative_to(generated_dir)
|
||||
dst = reference_dir / rel
|
||||
if dry_run:
|
||||
print(f"[dry-run] {ref_file} -> {dst}")
|
||||
print(f"[dry-run] {video_file} -> {dst}")
|
||||
else:
|
||||
dst.parent.mkdir(parents=True, exist_ok=True)
|
||||
shutil.copy2(ref_file, dst)
|
||||
shutil.copy2(video_file, dst)
|
||||
print(f"Copied: {rel}")
|
||||
copied += 1
|
||||
return copied
|
||||
|
||||
@@ -1,25 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Latent-slice regression test for LTX-2 distilled text-to-video.
|
||||
"""SSIM-based similarity test for LTX-2 distilled text-to-video.
|
||||
|
||||
Pixel-space SSIM is not a useful signal for this model: 4 distilled
|
||||
steps + bf16 attention + tiled VAE decode produce outputs that pass
|
||||
visual QA but occupy a very wide region in pixel space.
|
||||
|
||||
Inspired by diffusers' slice-vs-full regression philosophy — see
|
||||
``diffusers/tests/pipelines/ltx2/test_ltx2.py`` (compares pixel slices
|
||||
via ``torch.allclose(..., atol=1e-4)``) and
|
||||
``diffusers/tests/pipelines/cogvideo/test_cogvideox.py`` (full pixel
|
||||
tensors via ``numpy_cosine_similarity_distance(...) < 1e-3``).
|
||||
Diffusers itself does NOT compare latents; we apply the same "small
|
||||
signature slice + bounded full-tensor distance" idea to the **pre-VAE
|
||||
latent** because distilled few-step pipelines amplify per-step bf16
|
||||
noise enough that VAE-decoded comparisons are unreliable.
|
||||
|
||||
Parameters are kept identical to the original SSIM run so that
|
||||
reference artefacts generated on Modal L40S remain bit-compatible with
|
||||
production inference.
|
||||
Parameters derived from examples/inference/basic/basic_ltx2_distilled.py,
|
||||
with resolution + num_inference_steps reduced to keep GPU CI runtime
|
||||
bounded. Full-quality variant (via ``--ssim-full-quality``) falls back
|
||||
to the ``ltx2_distilled`` preset defaults.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
@@ -28,9 +14,7 @@ from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
resolve_inference_device_reference_folder,
|
||||
)
|
||||
from fastvideo.tests.ssim.latent_similarity_utils import (
|
||||
run_text_to_latent_similarity_test,
|
||||
run_text_to_video_similarity_test,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -87,13 +71,6 @@ LTX2_DISTILLED_TEST_PROMPTS = [
|
||||
"deadpan, absurd, and quietly tragic.",
|
||||
]
|
||||
|
||||
# Tolerances chosen on top of diffusers' ``1e-3`` defaults. LTX-2 distilled
|
||||
# amplifies per-step numerical noise, and FastVideo's CI pool spans L40S /
|
||||
# A40 / H100 so cross-architecture bf16 drift must be absorbed. Values can
|
||||
# be tightened after an initial stable window of reference refreshes.
|
||||
SLICE_COSINE_DISTANCE_THRESHOLD = 5e-3
|
||||
FULL_COSINE_DISTANCE_THRESHOLD = 1e-2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prompt", LTX2_DISTILLED_TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["FLASH_ATTN"])
|
||||
@@ -103,7 +80,7 @@ def test_ltx2_distilled_inference_similarity(
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
run_text_to_latent_similarity_test(
|
||||
run_text_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
@@ -112,6 +89,5 @@ def test_ltx2_distilled_inference_similarity(
|
||||
model_id=model_id,
|
||||
default_params_map=LTX2_DISTILLED_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FULL_QUALITY_LTX2_DISTILLED_MODEL_TO_PARAMS,
|
||||
slice_cosine_threshold=SLICE_COSINE_DISTANCE_THRESHOLD,
|
||||
full_cosine_threshold=FULL_COSINE_DISTANCE_THRESHOLD,
|
||||
min_acceptable_ssim=0.60,
|
||||
)
|
||||
|
||||
@@ -1,111 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Latent-slice regression test for Stable Audio Open 1.0 text-to-audio.
|
||||
|
||||
Companion to ``test_ltx2_similarity.py`` — applies the same latent
|
||||
cosine-distance philosophy to 3-D audio latents ``[B, 64, T_latent]``.
|
||||
|
||||
Why latent-space and not waveform-space SSIM:
|
||||
- ``dpmpp-3m-sde`` (k-diffusion) accumulates per-step bf16 noise; the
|
||||
Oobleck VAE then magnifies any residual drift into the time-domain
|
||||
waveform. A few mis-rounded accumulators drive sample-wise diff well
|
||||
past audible thresholds without indicating a real regression.
|
||||
- Diffusers' own
|
||||
``tests/pipelines/stable_audio/test_stable_audio.py`` compares
|
||||
decoded audio samples via
|
||||
``np.abs(expected - actual).max() < 1.5e-3``; that bound holds for
|
||||
CPU dummy components but does not survive cross-architecture bf16
|
||||
on our CI pool (L40S/A40/H100/B200).
|
||||
- Comparing the **pre-VAE latent** moves the assertion upstream of
|
||||
the dominant noise source.
|
||||
|
||||
Slice spec: ``audio_first_8_timesteps`` returns
|
||||
``latent[0, :, :8]`` (= 64 channels × 8 latent timesteps = 512
|
||||
elements). The full latent ``[1, 64, 1024]`` for SA-1.0 is also
|
||||
compared via cosine distance.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
resolve_inference_device_reference_folder,
|
||||
)
|
||||
from fastvideo.tests.ssim.latent_similarity_utils import (
|
||||
AUDIO_FIRST_8_TIMESTEPS_SPEC,
|
||||
run_text_to_latent_similarity_test,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 1
|
||||
|
||||
device_reference_folder = resolve_inference_device_reference_folder(logger)
|
||||
|
||||
STABLE_AUDIO_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": "FastVideo/stable-audio-open-1.0-Diffusers",
|
||||
# 8/8/1 are the SA-1.0 SamplingParam defaults; the values are
|
||||
# placeholders unused by the audio pipeline but must satisfy the
|
||||
# shared InputValidationStage divisible-by-8 check.
|
||||
"height": 8,
|
||||
"width": 8,
|
||||
"num_frames": 1,
|
||||
"num_inference_steps": 25,
|
||||
"guidance_scale": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"fps": 24,
|
||||
}
|
||||
|
||||
_SA_FULL_DEFAULTS = SamplingParam.from_pretrained(STABLE_AUDIO_PARAMS["model_path"])
|
||||
STABLE_AUDIO_FULL_QUALITY_PARAMS = {
|
||||
**STABLE_AUDIO_PARAMS,
|
||||
"num_inference_steps": _SA_FULL_DEFAULTS.num_inference_steps,
|
||||
"guidance_scale": _SA_FULL_DEFAULTS.guidance_scale,
|
||||
"seed": _SA_FULL_DEFAULTS.seed,
|
||||
}
|
||||
|
||||
STABLE_AUDIO_MODEL_TO_PARAMS = {
|
||||
"stable-audio-open-1.0-Diffusers": STABLE_AUDIO_PARAMS,
|
||||
}
|
||||
FULL_QUALITY_STABLE_AUDIO_MODEL_TO_PARAMS = {
|
||||
"stable-audio-open-1.0-Diffusers": STABLE_AUDIO_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
STABLE_AUDIO_TEST_PROMPTS = [
|
||||
"Lo-fi hip hop instrumental with vinyl crackle and gentle piano.",
|
||||
]
|
||||
|
||||
SLICE_COSINE_DISTANCE_THRESHOLD = 5e-3
|
||||
FULL_COSINE_DISTANCE_THRESHOLD = 1e-2
|
||||
|
||||
|
||||
@pytest.mark.parametrize("prompt", STABLE_AUDIO_TEST_PROMPTS)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(STABLE_AUDIO_MODEL_TO_PARAMS.keys()))
|
||||
def test_stable_audio_inference_similarity(
|
||||
prompt: str,
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
run_text_to_latent_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=STABLE_AUDIO_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FULL_QUALITY_STABLE_AUDIO_MODEL_TO_PARAMS,
|
||||
slice_cosine_threshold=SLICE_COSINE_DISTANCE_THRESHOLD,
|
||||
full_cosine_threshold=FULL_COSINE_DISTANCE_THRESHOLD,
|
||||
slice_spec=AUDIO_FIRST_8_TIMESTEPS_SPEC,
|
||||
# FSDP + @torch.inference_mode in StableAudioDenoisingStage hits
|
||||
# "Inference tensors do not track version counter" on single-GPU
|
||||
# unshard. SA-1.0 fits on one B200 anyway.
|
||||
init_kwargs_override={"use_fsdp_inference": False},
|
||||
)
|
||||
@@ -1,76 +0,0 @@
|
||||
# `fastvideo/train/` — Modular Training Framework
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
YAML-driven trainer composed from interchangeable **methods × models × callbacks**. Preferred location for new training code. (See sibling `fastvideo/training/AGENTS.md` for the legacy stack.)
|
||||
|
||||
## Layout
|
||||
|
||||
```
|
||||
train/
|
||||
├── trainer.py # Core training loop coordinator
|
||||
├── README.md # User-facing overview (legacy → new diff)
|
||||
├── entrypoint/ # train.py + dcp_to_diffusers.py CLI entrypoints
|
||||
├── methods/
|
||||
│ ├── base.py # TrainingMethod ABC
|
||||
│ ├── fine_tuning/ # FineTuneMethod, DiffusionForcingSFTMethod
|
||||
│ ├── distribution_matching/ # DMD2Method, SelfForcingMethod
|
||||
│ ├── knowledge_distillation/ # KDMethod, KDCausalMethod
|
||||
│ └── consistency_model/ # Consistency-model training methods
|
||||
├── models/
|
||||
│ ├── base.py # ModelBase / CausalModelBase wrappers
|
||||
│ ├── wan/, hunyuan/, cosmos/ # Per-family training wrappers
|
||||
├── callbacks/ # callback.py base + ema, grad_clip, validation
|
||||
└── utils/
|
||||
├── training_config.py # Hierarchical YAML config dataclasses
|
||||
├── builder.py # Build trainer from config
|
||||
├── checkpoint.py # DCP save/load
|
||||
├── optimizer.py # AdamW / fused optimizer factory
|
||||
├── tracking.py # build_tracker (W&B / TensorBoard)
|
||||
└── dataloader.py # StatefulDataLoader wiring
|
||||
```
|
||||
|
||||
## Composition Model
|
||||
|
||||
```
|
||||
Trainer = Method × Model × [Callback...] × Config
|
||||
```
|
||||
|
||||
- A **Method** owns the loss + optimizer step (`compute_loss`, `step_post_grad`).
|
||||
- A **Model** owns the forward + parameter-grouping (`forward`, `trainable_parameters`).
|
||||
- **Callbacks** subscribe to lifecycle hooks (`on_train_start`,
|
||||
`on_training_step_end`, `on_before_optimizer_step`, `on_validation_begin`,
|
||||
`on_validation_end`, `on_train_end`) and compose freely.
|
||||
- **Config** is a Pydantic-style hierarchical YAML resolved by
|
||||
`utils/training_config.py`. Dotted-key CLI overrides go through
|
||||
`parse_overrides`.
|
||||
|
||||
## Adding a New Model Plugin
|
||||
|
||||
1. Subclass `ModelBase` (or `CausalModelBase`) in `models/<family>/`.
|
||||
2. Wrap the existing inference DiT from `fastvideo/models/dits/`. Do not
|
||||
reimplement.
|
||||
3. Expose `trainable_parameters()` so the optimizer factory can group them.
|
||||
4. Register in `utils/builder.py` if the trainer dispatches by name.
|
||||
|
||||
## Adding a New Method
|
||||
|
||||
1. Subclass `TrainingMethod` in `methods/<family>/`.
|
||||
2. Implement `compute_loss(batch, model_outputs) -> dict`.
|
||||
3. If the method needs a teacher / second model, expose a `build_extras(cfg)`
|
||||
classmethod — never instantiate inside `__init__`.
|
||||
|
||||
## Configs
|
||||
|
||||
YAML lives under `examples/train/`. Schemas in `utils/training_config.py`. Add
|
||||
new fields with explicit defaults; agents and humans both rely on the dataclass
|
||||
to discover knobs.
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Importing from `fastvideo.training.*` (legacy stack). Forbidden cross-import.
|
||||
- Adding a new training method as a fork of an existing pipeline file. Compose
|
||||
via `Method` instead.
|
||||
- Logging via stdlib `logging` — use `init_logger(__name__)`.
|
||||
- Mutating the global state of a model in a callback. Callbacks operate on the
|
||||
trainer state object passed in.
|
||||
@@ -1,58 +0,0 @@
|
||||
# `fastvideo/training/` — Legacy Monolithic Training Pipelines
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
> **Status:** maintenance mode. New training work should go in `fastvideo/train/`.
|
||||
> Existing shipped recipes (Wan, LTX-2, MatrixGame distillation) still live here.
|
||||
|
||||
## What's Here
|
||||
|
||||
```
|
||||
training/
|
||||
├── training_pipeline.py # Base TrainingPipeline ABC
|
||||
├── distillation_pipeline.py # Base DistillationPipeline (~80 KB)
|
||||
├── self_forcing_distillation_pipeline.py # Self-forcing causal distill (~55 KB)
|
||||
├── wan_training_pipeline.py # Wan T2V finetune
|
||||
├── wan_i2v_training_pipeline.py # Wan I2V finetune
|
||||
├── wan_distillation_pipeline.py # Wan T2V distill
|
||||
├── wan_i2v_distillation_pipeline.py # Wan I2V distill
|
||||
├── wan_self_forcing_distillation_pipeline.py
|
||||
├── ltx2_training_pipeline.py # LTX-2 training
|
||||
├── matrixgame_training_pipeline.py # MatrixGame training
|
||||
├── ode_causal_pipeline.py # ODE-causal pipeline
|
||||
├── checkpointing_utils.py # save/load helpers
|
||||
├── activation_checkpoint.py # AC wrapping helpers
|
||||
├── trackers.py # WandbTracker (legacy)
|
||||
└── training_utils.py # Grad clip, FSDP state-dicts (~75 KB)
|
||||
```
|
||||
|
||||
## Style of This Stack
|
||||
|
||||
- One file per (model × method). Subclass an existing pipeline if your config is
|
||||
a near-clone; otherwise fork and rename.
|
||||
- Pipelines are torch-distributed-launched directly (no shared trainer). Each
|
||||
pipeline calls `dist.init_process_group` lifecycle through
|
||||
`fastvideo.distributed`.
|
||||
- Configs are flat argparse flags wired in `fastvideo_args.py` and per-pipeline
|
||||
argparse groups. No YAML.
|
||||
|
||||
## When to Touch a File Here
|
||||
|
||||
- Bugfix or behavior tweak to a shipped recipe (Wan, LTX-2, MatrixGame).
|
||||
- Performance / memory regression in `training_utils.py` (used by both stacks).
|
||||
- New attention backend that needs distillation-time wiring.
|
||||
|
||||
## When to NOT Touch a File Here
|
||||
|
||||
- Adding a new model's training. Build it in `fastvideo/train/` instead.
|
||||
- Adding a new training method (DMD2 variant, new self-forcing variant) for a
|
||||
model that already has a `train/` plugin. Add a method class in
|
||||
`train/methods/`.
|
||||
- Refactoring shared utilities. The new stack owns the modular utilities now.
|
||||
|
||||
## Cross-Stack Rule
|
||||
|
||||
Imports from `fastvideo.train.*` into `fastvideo.training.*` (or vice versa)
|
||||
are **forbidden**. They are independent stacks; mixing them creates cyclic
|
||||
dependency hazards. Share via `fastvideo.training.training_utils` or
|
||||
`fastvideo.distributed` only when both stacks already use the helper.
|
||||
@@ -854,7 +854,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
import av
|
||||
except ImportError:
|
||||
logger.warning("PyAV not installed; cannot mux audio. "
|
||||
"Install with: uv pip install av")
|
||||
"Install with: pip install av")
|
||||
return False
|
||||
|
||||
if torch.is_tensor(audio):
|
||||
|
||||
@@ -57,7 +57,7 @@ def assert_ray_available() -> None:
|
||||
"""Raise an exception if Ray is not available."""
|
||||
if ray is None:
|
||||
raise ValueError(f"Failed to import Ray: {ray_import_err}."
|
||||
"Please install Ray with `uv pip install ray`.")
|
||||
"Please install Ray with `pip install ray`.")
|
||||
|
||||
|
||||
def _verify_bundles(placement_group: "PlacementGroup", fastvideo_args: FastVideoArgs, device_str: str):
|
||||
|
||||
+2
-20
@@ -27,7 +27,7 @@ dependencies = [
|
||||
"timm==1.0.11",
|
||||
"peft>=0.15.0",
|
||||
"diffusers>=0.33.1",
|
||||
"torch==2.11.0",
|
||||
"torch>=2.10.0",
|
||||
"torchvision",
|
||||
"torchaudio",
|
||||
|
||||
@@ -98,10 +98,6 @@ torchvision = [
|
||||
{ index = "pytorch-cpu", marker = "sys_platform != 'linux'" },
|
||||
{ index = "pytorch-cu128", marker = "sys_platform == 'linux'" },
|
||||
]
|
||||
torchaudio = [
|
||||
{ index = "pytorch-cpu", marker = "sys_platform != 'linux'" },
|
||||
{ index = "pytorch-cu128", marker = "sys_platform == 'linux'" },
|
||||
]
|
||||
|
||||
[[tool.uv.index]]
|
||||
name = "pytorch-cpu"
|
||||
@@ -115,7 +111,7 @@ explicit = true
|
||||
|
||||
[project.optional-dependencies]
|
||||
|
||||
# flash-attn: uv pip install flash-attn==2.8.1 --no-cache-dir --no-build-isolation
|
||||
# flash-attn: pip install flash-attn==2.8.1 --no-cache-dir --no-build-isolation
|
||||
|
||||
|
||||
lint = [
|
||||
@@ -134,20 +130,6 @@ test = [
|
||||
|
||||
dev = [ "fastvideo[lint]", "fastvideo[test]", ]
|
||||
|
||||
prompt-safety = [
|
||||
"fasttext",
|
||||
]
|
||||
|
||||
prompt-enhancer = [
|
||||
"httpx",
|
||||
]
|
||||
|
||||
streaming = [
|
||||
"fastvideo[prompt-enhancer]",
|
||||
"fastvideo[prompt-safety]",
|
||||
"websockets",
|
||||
]
|
||||
|
||||
rocm = [
|
||||
"amdsmi",
|
||||
]
|
||||
|
||||
@@ -27,7 +27,7 @@ dependencies = [
|
||||
"timm==1.0.11",
|
||||
"peft>=0.15.0",
|
||||
"diffusers>=0.33.1",
|
||||
"torch==2.11.0",
|
||||
"torch>=2.10.0",
|
||||
"torchvision",
|
||||
|
||||
# Acceleration & Optimization
|
||||
@@ -84,7 +84,7 @@ prerelease = "allow"
|
||||
|
||||
[project.optional-dependencies]
|
||||
|
||||
# flash-attn: uv pip install flash-attn==2.8.1 --no-cache-dir --no-build-isolation
|
||||
# flash-attn: pip install flash-attn==2.8.1 --no-cache-dir --no-build-isolation
|
||||
|
||||
|
||||
lint = [
|
||||
|
||||
@@ -1,69 +0,0 @@
|
||||
# `scripts/checkpoint_conversion/` — Official → FastVideo Converters
|
||||
|
||||
**Generated:** 2026-05-02
|
||||
|
||||
> **Pre-commit excludes `scripts/`.** Format / lint by hand against neighboring
|
||||
> converters before opening a PR.
|
||||
|
||||
## What Lives Here
|
||||
|
||||
```
|
||||
checkpoint_conversion/
|
||||
├── convert_gamecraft_full.py # Combined DiT + VAE
|
||||
├── convert_gamecraft_vae.py # VAE only
|
||||
├── convert_gamecraft_weights.py # DiT only
|
||||
├── convert_gen3c_to_fastvideo.py
|
||||
├── convert_ltx2_weights.py
|
||||
├── convert_turbodiffusion_to_diffusers.py
|
||||
├── convert_turbodiffusion_i2v_to_diffusers.py
|
||||
├── extract_llava_text_encoder.py # Encoder extraction from a multimodal repo
|
||||
├── longcat_to_fastvideo.py
|
||||
├── stable_audio_to_diffusers.py
|
||||
├── wan_to_diffusers.py
|
||||
├── validate_longcat_weights.py # Post-conversion validation
|
||||
├── pt_to_safetensors.py # Generic format flip
|
||||
└── create_hf_repo.py # Push to HF after conversion
|
||||
```
|
||||
|
||||
## Naming Convention
|
||||
|
||||
| Pattern | Use |
|
||||
|---------|-----|
|
||||
| `convert_<model>_*.py` / `<model>_to_<format>.py` | One-shot converter for a model family |
|
||||
| `extract_<role>_*.py` | Pull a sub-component out of a multimodal repo |
|
||||
| `validate_<model>_*.py` | Post-conversion shape / norm sanity checks |
|
||||
| `<format>_to_<format>.py` | Generic format conversion (no model knowledge) |
|
||||
|
||||
## Authoring a New Converter
|
||||
|
||||
1. **Prototype the FastVideo-native component first** under
|
||||
`fastvideo/models/<role>/<model>.py`. Its `state_dict()` is the target
|
||||
surface — never the other way around.
|
||||
2. Mirror the official-checkpoint key pattern → FastVideo-native key pattern in
|
||||
a `param_names_mapping` (declared on the arch config in
|
||||
`fastvideo/configs/models/<role>/<model>.py`).
|
||||
3. The converter's only job: load the official checkpoint, apply the mapping
|
||||
(with split/fuse for QKV / packed MLP), and write a safetensors directory
|
||||
that the `ComponentLoader` (`fastvideo/models/loader/`) can read.
|
||||
4. Record intentionally skipped keys (training-only EMA, optimizer state,
|
||||
logvar, dynamic buffers) in a constant near the top of the converter — with
|
||||
a one-line reason each.
|
||||
5. Add a smoke test under `tests/local_tests/<model>/` or
|
||||
`fastvideo/tests/<role>/` that loads the converted weights and runs a
|
||||
parity assertion against the official reference (1 forward pass).
|
||||
|
||||
## Cross-Reference
|
||||
|
||||
- `fastvideo/layers/AGENTS.md` — which native layer to target (fused QKV,
|
||||
parallel MLP, etc.). Determines the split/fuse logic in the converter.
|
||||
- `fastvideo/models/AGENTS.md` — model-side discipline (lint excluded, native
|
||||
state-dict is the source of truth).
|
||||
|
||||
## Anti-Patterns
|
||||
|
||||
- Hard-coding tensor renames inside the model class instead of the converter.
|
||||
- Converting via a private fork of `transformers` / `diffusers` weight loaders.
|
||||
Read tensors directly with `safetensors` or `torch.load`.
|
||||
- Pushing to HF before the validate / smoke test passes locally.
|
||||
- Skipping the explicit "skipped keys" list — it documents intent and prevents
|
||||
silent precision loss across re-conversions.
|
||||
@@ -378,7 +378,7 @@ def download_checkpoint(
|
||||
) -> Path:
|
||||
"""Download checkpoint from HuggingFace Hub."""
|
||||
if hf_hub_download is None:
|
||||
raise RuntimeError("huggingface_hub is required for --download. Install with: uv pip install huggingface_hub")
|
||||
raise RuntimeError("huggingface_hub is required for --download. Install with: pip install huggingface_hub")
|
||||
|
||||
print(f"Downloading {filename} from {repo_id}...")
|
||||
path = hf_hub_download(
|
||||
@@ -403,7 +403,7 @@ def resolve_model_dir(
|
||||
if snapshot_download is None:
|
||||
raise RuntimeError(
|
||||
"huggingface_hub is required to download component source. "
|
||||
"Install with: uv pip install huggingface_hub")
|
||||
"Install with: pip install huggingface_hub")
|
||||
|
||||
print(f"Downloading component source repo: {model_name_or_path}")
|
||||
downloaded = snapshot_download(
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user