Compare commits

...
Author SHA1 Message Date
SolitaryThinker 93a375c654 [bugfix]: resolve InterleaveThinker review findings 2026-06-28 02:32:34 -07:00
Mac Lee 83de8caab4 [docs]: update InterleaveThinker branch name 2026-06-27 11:28:40 +00:00
Mac Lee 7540d1181f [docs]: record InterleaveThinker training dry-runs 2026-06-27 11:24:26 +00:00
Mac Lee 9758d931d3 [docs]: record role model abstraction cleanup 2026-06-27 07:18:43 +00:00
Mac Lee 7219b3b743 [refactor]: add non-diffusion role model base 2026-06-27 07:18:10 +00:00
Mac Lee 7d8604b4fa [docs]: record InterleaveThinker parity checkpoint 2026-06-22 01:17:15 +00:00
Mac Lee f04e675088 [test]: add InterleaveThinker official parity checks 2026-06-22 01:16:36 +00:00
Mac Lee 9363caf64e [docs] record InterleaveThinker workflow namespace correction 2026-06-21 20:40:47 +00:00
Mac Lee bb1e8935ee [refactor] use existing workflow namespace for InterleaveThinker 2026-06-21 20:37:13 +00:00
Mac Lee 58256b1282 [docs] record InterleaveThinker workflow migration 2026-06-21 19:07:47 +00:00
Mac Lee 91d8fb85e6 [misc] format InterleaveThinker workflow helpers 2026-06-21 19:03:31 +00:00
Mac Lee 11d55fb5a2 [refactor] move InterleaveThinker helpers to workflows 2026-06-21 19:00:11 +00:00
Mac Lee d2c5395132 [docs] update InterleaveThinker handoff instructions 2026-06-21 18:53:43 +00:00
Mac Lee 704e56674f [docs] record InterleaveThinker CLI cleanup validation 2026-06-21 06:38:17 +00:00
Mac Lee 555a600f3f [bugfix] remove InterleaveThinker CLI surface 2026-06-21 06:37:44 +00:00
Mac Lee c6128ab765 [docs] record InterleaveThinker review package push 2026-06-20 04:05:05 +00:00
Mac Lee 6a6ebf1ee3 [docs] add InterleaveThinker review package 2026-06-20 04:04:30 +00:00
Mac Lee f1f7ac0738 [docs] record interleave trace eval validation 2026-06-20 04:01:49 +00:00
Mac Lee 47335f09ee [feat] add InterleaveThinker trace evaluation 2026-06-20 03:56:36 +00:00
Mac Lee 9e7a30acbc [docs] record interleave eval validation 2026-06-20 03:43:22 +00:00
Mac Lee 022aedb0e1 [bugfix] clean up interleave eval mypy 2026-06-20 03:39:44 +00:00
Mac Lee 874e4e2f9c [bugfix] defer interleave eval config loading 2026-06-20 03:36:05 +00:00
Mac Lee eca1444101 [feat] add InterleaveThinker prompt-set eval 2026-06-20 03:32:16 +00:00
Mac Lee 000d48b74d [docs] record InterleaveThinker planner GRPO push 2026-06-19 20:28:54 +00:00
Mac Lee 2cf9aa0d09 [feat] add InterleaveThinker planner GRPO path 2026-06-19 20:28:15 +00:00
Mac Lee be0bacbe35 [docs] record InterleaveThinker reference policy push 2026-06-19 20:07:55 +00:00
Mac Lee 42c2fe6a62 [feat] add InterleaveThinker reference policy KL 2026-06-19 20:06:57 +00:00
Mac Lee 8bfcb66447 [docs] record InterleaveThinker PEFT smoke 2026-06-19 19:55:28 +00:00
Mac Lee 0cc0478463 [feat] add PEFT LoRA for InterleaveThinker actors 2026-06-19 19:54:52 +00:00
Mac Lee a312d7c6e1 [docs] record InterleaveThinker GRPO push 2026-06-19 19:33:58 +00:00
Mac Lee b7a923c0bb [feat] add InterleaveThinker GRPO policy loss 2026-06-19 19:33:20 +00:00
Mac Lee 0b4e9764b6 [docs] record InterleaveThinker SFT push 2026-06-19 19:08:55 +00:00
Mac Lee df88af316b [feat] add InterleaveThinker SFT method 2026-06-19 19:08:10 +00:00
Mac Lee 033b75a662 [docs] record InterleaveThinker data normalizer push 2026-06-19 18:48:15 +00:00
Mac Lee dcd82f93e8 [feat] add InterleaveThinker data normalizers 2026-06-19 18:47:41 +00:00
Mac Lee a22da778a6 [docs] record interleave run generator smoke 2026-06-19 18:38:16 +00:00
Mac Lee 38513d896c [docs] record interleave CLI config fix 2026-06-19 18:32:50 +00:00
Mac Lee 375b944b58 [bugfix] preserve interleave CLI config paths 2026-06-19 18:32:16 +00:00
Mac Lee f4d1b59b9e [docs] record InterleaveThinker run CLI push 2026-06-19 18:23:47 +00:00
Mac Lee ee4021e5cb [feat] add InterleaveThinker run CLI 2026-06-19 18:23:03 +00:00
Mac Lee 269544c09b [docs] record InterleaveThinker provider push 2026-06-19 18:11:50 +00:00
Mac Lee 2d3bb7c795 [feat] wire InterleaveThinker model providers 2026-06-19 18:11:03 +00:00
Mac Lee a2dd0b6398 [docs] update InterleaveThinker planner handoff 2026-06-19 17:55:40 +00:00
Mac Lee 3b9ecb3485 [feat] add InterleaveThinker planner actor 2026-06-19 17:54:56 +00:00
Mac Lee 87dbb78002 [docs] add full InterleaveThinker integration plan 2026-06-19 17:28:07 +00:00
Mac Lee 06b6c43d3f [docs] record InterleaveThinker critic smoke 2026-06-19 16:57:19 +00:00
Mac Lee 9973307b11 [docs] update InterleaveThinker integration handoff 2026-06-19 16:46:15 +00:00
Mac Lee ace421bc2b [feat] harden InterleaveThinker critic backend 2026-06-19 16:45:33 +00:00
Mac Lee fc04d01930 [feat] integrate InterleaveThinker model backends 2026-06-19 05:26:24 +00:00
Mac Lee c2a86bef4c [misc] record InterleaveThinker RL handoff 2026-06-18 19:37:47 +00:00
Mac Lee 9521077f21 [feat] add InterleaveThinker RL training loop 2026-06-18 19:37:13 +00:00
Mac Lee e73521d04e [misc] record InterleaveThinker validation handoff 2026-06-18 02:47:52 +00:00
Mac Lee 0c56a4aa53 [feat] add interleave trace runner 2026-06-18 02:44:01 +00:00
Mac Lee eb33639cbd [feat] add InterleaveThinker compatibility service 2026-06-18 02:40:54 +00:00
Mac Lee 7e70f8598a [misc] track InterleaveThinker integration plan 2026-06-18 02:28:24 +00:00
alexzmsandmergify[bot] 633d393568 [ci] layer-0 grad-norm regression for per-method training tests (5a-ii) (#1396)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-06-12 04:45:07 +00:00
Junda SuandPeiyuan Zhang 5854aec2ce [feat] Add Wan RL DiffusionNFT training (#1450)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2026-06-11 21:18:59 -07:00
90 changed files with 14788 additions and 64 deletions
@@ -0,0 +1,409 @@
# Exploration Log: InterleaveThinker FastVideo Integration
## Status
Draft handoff, shortened on 2026-06-21 and updated on 2026-06-22 after
official InterleaveThinker parity validation.
Current working location:
- Directory: `/home/toolbox/FastVideo`
- Branch: `interleavethinker`
- Latest completed integration checkpoint:
`7219b3b7` (`[refactor]: add non-diffusion role model base`)
- Latest observed branch head before the official parity patch:
`9363caf6` (`[docs] record InterleaveThinker workflow namespace correction`)
This file is the canonical handoff for the InterleaveThinker integration work.
It intentionally summarizes older execution logs; use git history for the full
append-only detail if needed.
## Current Hard Instructions
- Work in `/home/toolbox/FastVideo` on the checked-out branch.
- Do not add standalone InterleaveThinker `fastvideo` subcommands such as
`interleave-run`, `interleave-serve`, or `interleave-eval`.
- Do not add a separate InterleaveThinker HTTP API surface unless strictly
necessary.
- Keep useful additions integrated into existing FastVideo library and training
surfaces.
- Keep this handoff updated before context compaction, interruption, or a major
direction change.
- Make focused commits as frequently as useful; push committed checkpoints when
validation evidence should be durable.
- Do not run tests on the local machine. The local environment is not reliable
for this work because both hardware and software prerequisites are missing.
- Run validation on Modal through `fastvideo/tests/modal/launch_l40s_job.py`.
L40S is the normal target, but H100 or B200-class GPUs may be used when the
task needs more memory or speed. Check Modal availability before relying on a
specific larger GPU type.
- User approval is already granted for all Modal actions needed to finish this
task set, including running jobs and uploading files or uncommitted patches
from `/home/toolbox/FastVideo`.
- Prefer FastVideo's modular `fastvideo/train` stack for new training work.
Do not migrate legacy `fastvideo/training` pipelines unless explicitly asked.
- Do not vendor InterleaveThinker, EasyR1, LLaMA-Factory, or their full training
stacks into FastVideo.
- Planner and critic are Transformers Qwen3-VL `RoleModelBase` wrappers, not native
FastVideo DiT components. A native Qwen3-VL port should happen only if
checkpoint conversion, distribution, or performance requirements justify it.
- Boundaries:
- VLM model details live in planner/critic model wrappers.
- RL algorithms live in `fastvideo/train/methods/rl`.
- Reward parsing/scoring lives under `fastvideo/train/methods/rl/rewards`.
- Interleaved inference helpers live under
`fastvideo/workflow/interleave_thinker`; use the pre-existing singular
`fastvideo/workflow` namespace, not a parallel `fastvideo/workflows`
package.
## Goal
Add a native FastVideo integration surface for InterleaveThinker-style workflows:
- run planner -> generator/edit -> critic loops through reusable Python helpers;
- train/fine-tune planner and critic Qwen3-VL actors through FastVideo YAML
configs and the modular trainer;
- support InterleaveThinker SFT, critic GRPO, planner GRPO, reward parsing, and
trace/evaluation utilities;
- keep tests deterministic with fake backends, and reserve real-checkpoint
validation for Modal.
Out of scope unless explicitly re-opened:
- standalone FastVideo CLI commands dedicated to InterleaveThinker;
- a separate InterleaveThinker HTTP API/server surface;
- full-parameter 8B training as a default path;
- deterministic regression tests against live closed-source services.
## Architecture Snapshot
Implemented and retained surfaces:
- `fastvideo.workflow.interleave_thinker`
- schema objects, generator backend translation, orchestrator, provider
adapters, config/runner helpers, prompt-set evaluation, and trace metrics.
- Standalone command registration and standalone server modules were removed.
- `fastvideo.train.models.interleave_thinker`
- shared Qwen3-VL actor base;
- planner wrapper for `InterleaveThinker/InterleaveThinker-Planner-8B`;
- critic wrapper for `InterleaveThinker/Critic-SFT-8B`;
- dataset normalization for planner SFT/RL and critic SFT/RL files.
- `fastvideo.train.methods.fine_tuning.interleave_thinker_sft`
- response-token-only SFT for planner and critic.
- `fastvideo.train.methods.rl.interleave_thinker`
- GRPO-style managed RL loop with grouped rollouts, old logprobs, optional
frozen reference policy KL, and LoRA-first configs.
- `fastvideo.train.methods.rl.rewards.interleave_thinker`
- critic reward parser/scorer and planner format/plan reward utilities.
- `examples/train/configs/interleave_thinker/`
- planner and critic SFT LoRA configs.
- `examples/train/configs/rl/interleave_thinker/`
- critic and planner GRPO LoRA configs.
- `docs/design/interleave_thinker.md`
- review/design entrypoint that should stay shorter and more reviewer-facing
than this exploration file.
Removed by the API/CLI cleanup:
- Interleave-specific `fastvideo` subcommand registration.
- Standalone Interleave compatibility server.
- command/service-oriented examples and scripts.
- `interleave-api` optional extra.
Namespace integration status:
- Completed correction. The reusable helper layer lives under the pre-existing
singular `fastvideo/workflow/interleave_thinker` package, not under a new
parallel `fastvideo/workflows` package.
- Internal imports, tests, examples, docs, and this handoff now use
`fastvideo.workflow.interleave_thinker`.
- The old `fastvideo.entrypoints.interleave` package remains deleted rather
than kept as a compatibility shim. This branch has not merged, so preserving
the old public path is not required.
- Do not scatter the helper code into unrelated core modules unless a genuinely
generic abstraction emerges. The planner -> generator/edit -> critic loop is
InterleaveThinker-specific workflow code, not `VideoGenerator`, training
method, or reward-parser core behavior.
## Condensed Execution History
- Initial service/orchestration slice added Interleave request/trace schema,
generator request translation, fake-provider tests, and an early compatibility
service. The later cleanup removed the standalone service/CLI surface but kept
reusable Python helpers.
- Critic backend hardening added Gemini/Nano Banana-style API wrappers with lazy
imports, fake-client tests, and no live API calls in CI-style tests.
- Real critic smoke loaded `InterleaveThinker/Critic-SFT-8B` on Modal L40S with
`Qwen/Qwen3-VL-8B-Instruct` and produced a non-empty response.
- Shared actor/planner work added a shared Qwen3-VL actor base,
`InterleaveThinkerPlannerModel`, planner parsing, and real planner/critic
smokes. Commit: `3b9ecb34`.
- Provider adapters wired planner and critic model wrappers into the native
`InterleaveOrchestrator`; real planner + fake generator + real critic smoke
passed. Commit: `2d3bb7c7`.
- Native run/config helpers were added and validated with a real FastVideo
FLUX.2-klein generator smoke. Later cleanup removed dedicated command
registration while keeping reusable helper code. Commits included
`ee4021e5` and `375b944b`.
- Dataset normalization added support for upstream planner SFT, critic SFT,
critic RL, and planner RL formats with image path resolution and clear data
errors. Commit: `dcd82f93`.
- Planner/critic SFT added response-token-only supervised fine-tuning and
LoRA-first configs. Commit: `df88af31`.
- Critic GRPO upgraded from advantage-weighted NLL to response-token logprob
policy loss with PPO/GRPO ratio, clipping, optional KL input, and metrics.
Commit: `b7a923c0`.
- PEFT LoRA was added for Qwen actors after FastVideo's native DiT LoRA wrapper
failed on HF Qwen modules. Real one-step critic RL smoke then passed.
Commit: `0cc04784`.
- Optional frozen reference policy KL was added through `models.reference` and
validated with a real one-step critic RL reference smoke. Commit: `42c2fe6a`.
- Planner GRPO added planner rollouts, planner rewards, `planner_rl` data, and
real one-step planner RL smoke. Commit: `2cf9aa0d`.
- Prompt-set evaluation and trace metrics/report helpers were added as reusable
Python/library surfaces. Commits included `eca14441`, `874e4e2f`,
`022aedb0`, and `47335f09`.
- Review package added `docs/design/interleave_thinker.md` and MkDocs nav.
Commit: `6a6ebf1e`.
- API/CLI cleanup removed standalone InterleaveThinker FastVideo commands and
the separate server, restored normal parser behavior, and rewrote docs toward
library/training integration. Commits: `555a600f`, `704e5667`.
- Handoff instructions were condensed and updated with standing Modal approval
and the no-local-tests rule. Commit: `d2c53951`.
- Namespace integration first moved the helper package to
`fastvideo.workflows.interleave_thinker`, renamed stale tests, updated docs
and examples, and deleted the old entrypoints package. Commits: `11d55fb5`,
`91d8fb85`.
- Follow-up correction requested by the user: move the helper package into the
pre-existing singular `fastvideo.workflow.interleave_thinker` namespace and
remove the parallel `fastvideo.workflows` package. Commit: `bb1e8935`.
- Official parity hardening aligned planner, guidance-planner, and critic
prompt literals with upstream InterleaveThinker; matched the upstream demo
message constructor for text/image interleaving including the `max_pixels`
behavior for five or more images; and adjusted Qwen generation so
official-style single-output inference preserves checkpoint generation config
while multi-output/custom-sampling RL paths still pass sampling controls.
Commit: `f04e6750`.
- Abstraction cleanup introduced `RoleModelBase` as the minimal non-diffusion
training role base, made diffusion `ModelBase` inherit from it, moved
`Qwen3VLActorBase` off the diffusion contract, removed the actor dummy
scheduler and diffusion stubs, and added explicit Interleave SFT/RL actor
protocols. This preserves existing diffusion method contracts while making
planner/critic actors honest non-diffusion role models. Commit: `7219b3b7`.
## Validation Evidence
Representative Modal real-checkpoint or GPU-backed smokes:
- Critic SFT smoke:
- model `InterleaveThinker/Critic-SFT-8B`;
- processor `Qwen/Qwen3-VL-8B-Instruct`;
- backend `Qwen3VLForConditionalGeneration`;
- marker `SMOKE_OK`.
- Planner smoke:
- model `InterleaveThinker/InterleaveThinker-Planner-8B`;
- `max_new_tokens=2048` was needed for the tested prompt;
- parsed `3` execution steps;
- marker `PLANNER_SMOKE_OK`.
- Critic refactor smoke:
- real critic through shared actor base;
- marker `CRITIC_REFACTOR_SMOKE_OK`.
- Provider loop smoke:
- real planner + fake deterministic image generator + real critic;
- marker `INTERLEAVE_PROVIDER_REAL_LOOP_SMOKE_OK`.
- FastVideo generator smoke:
- loaded `black-forest-labs/FLUX.2-klein-4B`;
- generated an image and trace through the reusable interleave runner path.
- The old command entrypoint used for this smoke has since been removed.
- Real critic RL smoke:
- trainable LoRA critic student on `InterleaveThinker/Critic-SFT-8B`;
- `ConstantInterleaveEditScorer`;
- one GRPO update completed;
- marker `INTERLEAVE_CRITIC_RL_SMOKE_OK`.
- Real critic RL reference smoke:
- trainable LoRA critic student plus frozen critic reference;
- old and reference response-token logprobs computed;
- marker `INTERLEAVE_CRITIC_RL_REFERENCE_SMOKE_OK`.
- Real planner RL smoke:
- trainable LoRA planner student plus frozen planner reference;
- `InterleavePlannerRewardScorer`;
- one GRPO update completed;
- marker `INTERLEAVE_PLANNER_RL_SMOKE_OK`.
Latest cleanup validation:
- Local `python -m py_compile` passed for touched Python files.
- Local `git diff --check` passed.
- Local `pre-commit run --files ...` passed for surviving changed files:
yapf, ruff, codespell, PyMarkdown, mypy, filename check, and suggestion hook.
- Focused local Interleave tests passed with a temporary CPU-only
`fastvideo_kernel` import stub:
`62 passed, 16 warnings`.
- Focused Modal Interleave/pre-commit validation passed:
`22 passed, 14 warnings`; pre-commit hooks passed.
- Existing API/CLI regression tests after cleanup passed on Modal:
`42 passed, 14 warnings`.
- Namespace migration validation passed on Modal L40S:
- App URL: `https://modal.com/apps/hao-ai-lab/main/ap-APG1eoMxnajN1wpzPd0S4r`
- Commit: `91d8fb85e6bb36bbeacde5e82aac8ccb22a2c9ee`
- Pytest:
`tests/local_tests/test_interleave_workflow_backend.py`,
`tests/local_tests/test_interleave_model_providers.py`,
`tests/local_tests/test_interleave_workflow_runner.py`,
`tests/local_tests/test_interleave_trace_eval.py`, and
`tests/local_tests/test_interleave_thinker_api_models.py`
-> `22 passed, 14 warnings`.
- Pre-commit on changed docs/examples/workflow/reward/test files passed:
yapf, ruff, codespell, PyMarkdown, mypy, filename check, and suggestion.
- `local_patch_applied=false`; validation used the pushed commit.
- Singular workflow namespace correction validation passed on Modal L40S:
- App URL: `https://modal.com/apps/hao-ai-lab/main/ap-zAYZ80ExlxJbpvSDVWbtTu`
- Commit: `bb1e8935ee37ea1e99896cf96fa1ea4139ff119e`
- Pytest:
`tests/local_tests/test_interleave_workflow_backend.py`,
`tests/local_tests/test_interleave_model_providers.py`,
`tests/local_tests/test_interleave_workflow_runner.py`,
`tests/local_tests/test_interleave_trace_eval.py`, and
`tests/local_tests/test_interleave_thinker_api_models.py`
-> `22 passed, 14 warnings`.
- Pre-commit on changed docs/examples/workflow/reward/test files passed:
yapf, ruff, codespell, PyMarkdown, mypy, filename check, and suggestion.
- `local_patch_applied=false`; validation used the pushed commit.
- Official InterleaveThinker parity validation passed on Modal L40S:
- FastVideo base commit: `9363caf64edfe4013c0525f4092b155987974253` with
local patch applied.
- Official InterleaveThinker reference commit observed before validation:
`93511614902c5e4f0c167951a4b78343bd864122`.
- Passing app URL:
`https://modal.com/apps/hao-ai-lab/main/ap-3UpF5p9UD9wiC8geNSJ1eP`
- Command cloned `https://github.com/zhengdian1/InterleaveThinker.git` inside
the Modal job and ran
`pytest tests/local_tests/test_interleave_thinker_official_parity.py -q -s`
with `INTERLEAVETHINKER_REAL_PARITY=1`.
- Result: `5 passed, 14 warnings`.
- Coverage: official prompt-template parity, official demo message-constructor
parity, fake Qwen API-call parity, and real planner/critic checkpoint
generation parity against upstream `UEval.qwen3_vl_api.predict`.
- Modal emitted the known FlashAttention ABI warning after dev dependency
installation; the real checkpoint tests used `attn_implementation=sdpa`.
- Focused InterleaveThinker regression validation passed on Modal L40S:
- App URL: `https://modal.com/apps/hao-ai-lab/main/ap-qdLoZdm5OkqE9nNqwtBQAz`
- Same FastVideo base commit with local patch applied.
- Pytest covered planner/critic model fakes, providers, workflow backend and
runner, API models, RL method/math, SFT method, rewards, data normalization,
and trace evaluation.
- Result: `62 passed, 14 warnings`.
- Pre-commit validation for the parity patch passed on Modal L40S:
- App URL: `https://modal.com/apps/hao-ai-lab/main/ap-cQkolTKGy7J24OP7ZR5mFh`
- Command:
`pre-commit run --files pyproject.toml fastvideo/train/models/interleave_thinker/planner.py fastvideo/train/models/interleave_thinker/critic.py fastvideo/train/models/interleave_thinker/qwen_actor.py tests/local_tests/test_interleave_thinker_planner_model.py tests/local_tests/test_interleave_thinker_critic_model.py tests/local_tests/test_interleave_thinker_official_parity.py`
- Result: yapf, ruff, codespell, PyMarkdown, actionlint, mypy, filename check,
and suggestion passed or were correctly skipped when no files applied.
- Local validation for the parity patch was limited to syntax and diff hygiene:
- `PYTHONDONTWRITEBYTECODE=1 python -m py_compile ...` passed for touched
Python files.
- `git diff --check` passed.
- No local pytest was run.
- Abstraction cleanup validation:
- Local syntax/diff hygiene only:
`PYTHONDONTWRITEBYTECODE=1 python -m py_compile ...` passed for touched
Python files, and `git diff --check` passed.
- Focused InterleaveThinker regression validation passed on Modal L40S:
app URL `https://modal.com/apps/hao-ai-lab/main/ap-75XUQvgThU5DvYENqoJTek`;
result `62 passed, 14 warnings`.
- Existing modular train/config regression subset passed on Modal L40S:
app URL `https://modal.com/apps/hao-ai-lab/main/ap-k1Dtgm9E4gAOKvkHQCVgjM`;
result `69 passed, 14 warnings`.
- A broader train-method Modal run including Wan single-step tests produced
`69 passed, 2 failed, 15 warnings`; the two failing Wan test targets also
failed on an unpatched branch-head comparison job. Treat those failures as
current Modal image / upstream test-environment issues, not regressions from
the role-model abstraction slice.
- Patched broader run:
`https://modal.com/apps/hao-ai-lab/main/ap-X0UxJwUHohtnEU1RQ8lho6`.
- Unpatched comparison:
`https://modal.com/apps/hao-ai-lab/main/ap-ZdMTgR8UP9KhYPHEXR1PAi`.
- Official InterleaveThinker parity passed on Modal L40S with upstream cloned
in-job and `INTERLEAVETHINKER_REAL_PARITY=1`:
app URL `https://modal.com/apps/hao-ai-lab/main/ap-h2OM9Yst0VuoYDAded9Qll`;
result `5 passed, 14 warnings`.
- Modal pre-commit on touched files passed:
app URL `https://modal.com/apps/hao-ai-lab/main/ap-blSbeLVjX6lkE4MZp9TTmq`;
yapf, ruff, codespell, PyMarkdown, mypy, filename check, and suggestion
passed or were correctly skipped when no files applied.
- No local pytest was run.
- Training pipeline dry-run validation, completed 2026-06-27:
- Goal: exercise the modular training entrypoint
`fastvideo.train.entrypoint.train --dry-run` with the public
InterleaveThinker YAML configs, temporary in-job fixtures, single-GPU
distributed overrides, and SDPA attention overrides.
- Planner and critic SFT dry-runs passed on Modal L40S:
app URL `https://modal.com/apps/hao-ai-lab/main/ap-BFk8R8SB47o1CL7l2SbZHk`.
Both commands loaded the real Qwen3-VL checkpoint, enabled PEFT LoRA, built
the dataloader/method via `build_from_config()`, and printed
`Dry-run: config parsed and build_from_config succeeded.`
- Planner and critic GRPO dry-runs passed on Modal H100:
app URL `https://modal.com/apps/hao-ai-lab/main/ap-UZ80OpsMVwwj2guheCf4pN`.
These commands loaded the real trainable student plus frozen reference
Qwen3-VL checkpoints, enabled PEFT LoRA on the student, built temporary
planner/critic RL dataloaders and `InterleaveThinkerRLMethod`, and printed
the same dry-run success line. H100 was used for this slice because each RL
config instantiates both student and reference checkpoints in one process.
- No training steps were executed in these dry-runs; the entrypoint returns
immediately after `build_from_config()` succeeds. No local pytest was run.
Broad-suite status:
- Local broad `pytest tests/ fastvideo/tests/ -q` is not a reliable signal on
this machine due missing GPU/runtime dependencies, blocked Hugging Face
downloads, missing SSIM references, missing `flashinfer`, and missing GUI
libraries for `cv2`.
- Broad Modal attempts did not produce a clean full-suite result. Known blockers
included `flashinfer` absence in the dev image and collection/import fallout
during combined suite runs. Treat focused Modal suites plus targeted API/CLI
regressions as the current evidence until the broad-suite environment is
repaired.
## Current Risks And Decisions
- One-process memory residency for real planner + real critic + real generator
is still not the recommended default. The validated approach separates heavy
concerns or uses fake/lightweight providers for orchestration tests.
- Live Gemini/Nano Banana behavior can change and may incur cost or rate
limits. Unit tests must use fake clients; live API runs should be recorded as
smoke evidence only.
- HF model/dataset access may require tokens and may change over time. Keep
tiny checked-in fixtures for parser, loader, and reward tests.
- Full-parameter 8B training is unvalidated. LoRA is the supported first path.
- Broad test validation needs a better Modal/dev image or a documented skip
strategy for tests requiring unavailable packages and external downloads.
- Keep the standalone CLI/API cleanup intact unless the user explicitly reverses
that product decision.
## Recommended Next Steps
1. For code work, continue from `/home/toolbox/FastVideo` on
`interleavethinker` and inspect `git status --short --branch`
before editing.
2. Read the relevant per-directory `AGENTS.md` before touching files under
`fastvideo/`, `examples/`, `docs/`, `scripts/`, or tests.
3. The API cleanup, namespace correction, official parity hardening, and
abstraction cleanup are complete. No further implementation step from the
current structural-divergence plan is pending.
4. Validate only on Modal. Local syntax-only commands such as `git diff --check`
are acceptable, but no local pytest or other local test execution should be
used.
5. Good next work items are PR decomposition/review packaging, broad-suite Modal
image repair for the known Wan/DTensor and memory failures, or reward/backend
hardening if product requirements call for it.
## Useful Commands
```bash
git status --short --branch
git log --oneline -12
git diff --check
pre-commit run --files <changed paths>
```
Use Modal for all test execution and authoritative validation.
+32
View File
@@ -0,0 +1,32 @@
---
name: add-reward-model
description: Use when adding reusable reward models under fastvideo/train/methods/rl/rewards for RLHF or online RL training.
---
# Add Reward Model
Use for reward models consumed by RL methods.
## Placement
- Put reusable reward code under `fastvideo/train/methods/rl/rewards/`.
- Expose public builders from `fastvideo/train/methods/rl/rewards/__init__.py`.
- Keep method-specific aggregation or advantage logic out of reward classes.
## Media Inputs
- Reward callables receive decoded media tensors.
- Accept single-frame tensors as `[B, C, H, W]` and multi-frame tensors as `[B, C, T, H, W]` when practical.
- Frame selection is reward-specific. Frame scorers such as PickScore and CLIPScore should explicitly select frame `0`; temporal rewards should inspect whichever frames they need.
- Return one scalar reward per prompt/sample.
## Attribution
- If code is ported or closely adapted from another repo, add a short comment or docstring naming the source file/function.
- Preserve SPDX headers used by FastVideo files.
## Tests
- Unit-test tensor layout handling without loading large reward checkpoints.
- Allow fake scorer injection for multi-reward tests.
- Test weighted reward aggregation and metric keys.
+38
View File
@@ -0,0 +1,38 @@
---
name: add-rl-method
description: Use when adding or modifying an RL/RLHF method under fastvideo/train/methods/rl, including DiffusionNFT-like methods.
---
# Add RL Method
Use for new RL methods in the modular `fastvideo/train` stack.
## Required Shape
- Add the method under `fastvideo/train/methods/rl/`.
- Subclass `TrainingMethod`.
- Keep model-family logic in `ModelBase` wrappers.
- Decode generated latents through `ModelBase.decode_latents`; add that hook to the new model wrapper instead of decoding inside the RL method.
- Use `fastvideo/train/methods/rl/common/sampling.py` for generation unless the method has a documented reason to avoid sampling.
- Use `fastvideo/train/methods/rl/common/prompt_sampling.py` for reusable grouped prompt sampling patterns such as DiffusionNFT K-repeat.
- Use `fastvideo/train/methods/rl/rewards/` for reward models.
## Optimization
- Return `manages_optimization() == True` only when the method must own a nonstandard outer/inner loop.
- If using managed optimization, implement `managed_train_step(data_stream, iteration)`.
- Existing trainer callbacks, checkpointing, tracking, and validation should still work.
## Config
- Put method knobs under `method`.
- Put sampler knobs under `method.sampling`.
- Do not put scheduler or trajectory policy into model configs.
- Do not split a diffusers-style scheduler from its built-in `step()` solver in YAML; use `trajectory` only for higher-level ODE vs re-noise behavior.
- Avoid fixed timestep lists in examples unless reproducing a known baseline; prefer scheduler-generated defaults.
## Tests
- Add fake-model tests for sampler/method behavior.
- Add config parse tests for the public YAML.
- Confirm existing train methods stay on the default Trainer path.
+3
View File
@@ -8,3 +8,6 @@
{"name": "decompose-pipeline-pr", "description": "Decompose an oversized FastVideo pipeline PR into a stack of independently-reviewable PRs. Tiers the diff by blast radius (invisible / dead code / cross-cutting infra / activation), produces a branch graph and worktree bootstrap, drafts the AGENTS.md manifest, flags missing tests on cross-cutting infra changes, and extracts lessons from the PR body. Worked example: PR #1280 daVinci-MagiHuman (9.8k LOC) decomposed into 10 stacked PRs.", "path": "decompose-pipeline-pr/SKILL.md", "status": "tested", "trust": "medium"}
{"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"}
{"name": "add-model", "description": "Add a new model (or variant) to FastVideo: DiT + configs + pipeline + presets + registry + tests. Walks through FastVideo's single stage-based pipeline architecture with exact file paths and registration hooks.", "path": "add-model/SKILL.md", "status": "draft", "trust": "low"}
{"name": "rlhf-training-abstractions", "description": "Use when changing FastVideo RLHF/RL training infrastructure, especially sampler, reward, scheduler trajectory, or method boundaries under fastvideo/train.", "path": "rlhf-training-abstractions/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-rl-method", "description": "Use when adding or modifying an RL/RLHF method under fastvideo/train/methods/rl, including DiffusionNFT-like methods.", "path": "add-rl-method/SKILL.md", "status": "draft", "trust": "low"}
{"name": "add-reward-model", "description": "Use when adding reusable reward models under fastvideo/train/methods/rl/rewards for RLHF or online RL training.", "path": "add-reward-model/SKILL.md", "status": "draft", "trust": "low"}
@@ -0,0 +1,41 @@
---
name: rlhf-training-abstractions
description: Use when changing FastVideo RLHF/RL training infrastructure, especially sampler, reward, scheduler trajectory, or method boundaries under fastvideo/train.
---
# RLHF Training Abstractions
Use this skill before editing RLHF-style training code in `fastvideo/train`.
## Boundaries
- RL methods live under `fastvideo/train/methods/rl/` and own algorithm logic: reward collection, advantage computation, policy loss, KL/reference terms, and optimizer cadence.
- Rewards live under `fastvideo/train/methods/rl/rewards/` and must be reusable across RL methods.
- RL methods pass decoded media to rewards; each reward decides whether to use the first frame, sampled frames, or the full video.
- Sampling lives under `fastvideo/train/methods/rl/common/` and must use `ModelBase` primitives plus scheduler math, not model-family inference pipelines.
- Model wrappers under `fastvideo/train/models/` own model-specific forward details.
- Model wrappers also own model-specific latent decoding via `ModelBase.decode_latents`; RL methods should not reach into VAE normalization internals.
- Shared RL helpers such as K-repeat prompt sampling belong under `fastvideo/train/methods/rl/common/` when they are reusable across RL methods.
## Anti-Patterns
- Do not bind RL methods to inference pipeline classes such as `WanDMDPipeline`.
- Do not hardcode timestep lists in a method when the scheduler can generate them.
- Do not put reward-model code inside one RL method.
- Do not make existing non-RL methods use method-managed optimization unless explicitly requested.
## Sampling Policy
- Prefer YAML-configured `method.sampling` with `scheduler`, `trajectory`, `num_steps`, `timesteps`, and `sigmas`.
- Treat diffusers-style scheduler classes as owning both the timestep schedule and their `step()` update rule; avoid a separate `solver` field unless a new sampler truly implements solver math outside the scheduler object.
- Missing `timesteps` means “ask the scheduler”; explicit `timesteps` or `sigmas` are overrides.
- ODE-style trajectories should not re-noise between denoising steps.
- SDE/re-noise behavior must be explicit in config.
## Validation
- Run focused local tests for sampler config and Trainer opt-in behavior.
- Verify existing train methods still report `manages_optimization() == False`.
- Keep fixed-prompt validation helpers in `fastvideo/train/methods/rl/common/validation.py` so new RL methods can reuse sharding and captions.
- Test distributed prompt grouping helpers separately from heavyweight model loading.
- Run `pre-commit run --files <changed paths>`; respect configured excludes.
+158
View File
@@ -0,0 +1,158 @@
# InterleaveThinker Integration Design
This page summarizes the InterleaveThinker integration branch for reviewers.
The detailed execution log remains in
`.agents/exploration/interleavethinker-fastvideo-integration.md`.
## Scope
The branch adds FastVideo-native support for InterleaveThinker-style workflows
without adding new `fastvideo` CLI commands or HTTP API routes:
- Qwen3-VL planner and critic model wrappers;
- planner and critic SFT configs;
- planner and critic GRPO configs with optional reference-policy KL;
- InterleaveThinker reward parsing and scoring utilities;
- Gemini and Nano Banana wrappers for optional network-backed rewards;
- Python orchestration helpers for planner -> generator -> critic traces.
It does not vendor InterleaveThinker, EasyR1, Verl, LLaMA-Factory, or training
framework internals from those projects.
## Integrated Surfaces
### Training
| Surface | Purpose |
|---------|---------|
| `fastvideo.train.models.interleave_thinker.Qwen3VLActorBase` | Shared Transformers Qwen3-VL runtime for planner and critic actors. |
| `InterleaveThinkerPlannerModel` | FastVideo `RoleModelBase` actor wrapper for `InterleaveThinker/InterleaveThinker-Planner-8B`. |
| `InterleaveThinkerCriticModel` | FastVideo `RoleModelBase` actor wrapper for `InterleaveThinker/Critic-SFT-8B` and `InterleaveThinker/InterleaveThinker-Critic-8B`. |
| `InterleaveThinkerSFTMethod` | Response-token supervised fine-tuning method for planner and critic actors. |
| `InterleaveThinkerRLMethod` | Managed GRPO-style loop for planner and critic actors. |
| `fastvideo.train.methods.rl.common.grpo` | Shared GRPO math helpers. |
| `fastvideo.train.methods.rl.rewards.interleave_thinker` | Format, critic, and planner reward utilities. |
| `fastvideo.train.methods.rl.rewards.interleave_api` | Optional Gemini and Nano Banana API-backed reward wrappers. |
Training examples live under:
- `examples/train/configs/interleave_thinker/planner_sft_lora.yaml`;
- `examples/train/configs/interleave_thinker/critic_sft_lora.yaml`;
- `examples/train/configs/interleave_thinker/planner_smoke.yaml`;
- `examples/train/configs/rl/interleave_thinker/critic_grpo.yaml`;
- `examples/train/configs/rl/interleave_thinker/planner_grpo.yaml`.
### Orchestration Helpers
The Python helper layer under `fastvideo.workflow.interleave_thinker` is intentionally
not registered as a CLI or server contract. It provides reusable dataclasses,
provider adapters, image-backend adapters, trace serialization, prompt-set
execution helpers, and saved-trace metrics for tests, examples, and downstream
integration code that already imports FastVideo as a library.
The runnable example is:
- `examples/interleave/interleave_single_prompt.py`.
## Architecture Boundaries
| Layer | Owner | Notes |
|-------|-------|-------|
| Planner and critic actors | `fastvideo/train/models/interleave_thinker/` | Wrap Transformers Qwen3-VL checkpoints. They are training actors, not diffusion pipeline components. |
| RL/SFT algorithms | `fastvideo/train/methods/` | Own loss, reward aggregation, advantage computation, KL, and optimizer cadence. |
| Rewards and API clients | `fastvideo/train/methods/rl/rewards/` | Offline reward aggregation is separate from network-backed Gemini/Nano Banana clients. |
| Image generation/editing helpers | `fastvideo/workflow/interleave_thinker/generator.py` | Presents a small image backend protocol for FastVideo, Nano Banana, and fake backends. |
| Runtime orchestration helpers | `fastvideo/workflow/interleave_thinker/` | Plans steps, calls generator/edit backends, calls critic providers, records traces. |
| Evaluation helpers | `fastvideo/workflow/interleave_thinker/evaluation.py` and `trace_eval.py` | Prompt-set execution and saved-trace reporting remain outside training methods. |
This keeps the Qwen actor implementation reusable by SFT, planner GRPO, critic
GRPO, and inference providers without coupling those paths to a specific
generator service.
## Validation Matrix
All GPU/model validation below ran on Modal L40S through
`fastvideo/tests/modal/launch_l40s_job.py`.
| Area | Evidence | Modal app |
|------|----------|-----------|
| API-backed model/reward wrappers | `27 passed, 14 warnings`; final pre-commit passed. | `ap-QOKlzapm5bSAo3c21lprwv` |
| Critic backend hardening | `30 passed, 14 warnings`; pre-commit passed. | `ap-DplMFq23YYfBx34e6TcsRc` |
| Real critic checkpoint smoke | Loaded `InterleaveThinker/Critic-SFT-8B`; generated one rollout; printed `SMOKE_OK`. | `ap-hDxj5MhLgdnGq22mRLjgIK` |
| FastVideo RL loop skeleton | `22 passed, 14 warnings`; pre-commit passed. | `ap-2Z2sH2UfhMoPmKolG0KY6t` |
| Shared Qwen actor and planner wrapper | `36 passed, 14 warnings`; pre-commit passed. | `ap-ZapOKZPOmhyMZFxZ0X1fQm` |
| Real planner checkpoint smoke | Loaded `InterleaveThinker/InterleaveThinker-Planner-8B`; parsed 3 steps; printed `PLANNER_SMOKE_OK`. | `ap-BzH7QxVXoc5XFXBah5cJ2H` |
| Real critic refactor smoke | Loaded critic wrapper after Qwen base refactor; printed `CRITIC_REFACTOR_SMOKE_OK`. | `ap-NGxUDBNJFiU30Wef0yAQN1` |
| Planner/critic provider adapters | `20 passed, 14 warnings`; pre-commit passed. | `ap-wfRX2DCt30DN903gETbDpj` |
| Real provider loop smoke | Real planner and critic with fake generator; printed `INTERLEAVE_PROVIDER_REAL_LOOP_SMOKE_OK`. | `ap-ZABadeyKBuGVcfy67LqmXt` |
| Dataset normalization | `17 passed, 14 warnings`; pre-commit passed. | `ap-ISuDU2lwc6Pl5NYDZnnBEb` |
| Planner and critic SFT | `20 passed, 14 warnings`; final pre-commit passed. | `ap-1jmIczO3KwZoP3WtLYOIxc` |
| Critic GRPO policy loss | Broad InterleaveThinker test set: `28 passed, 14 warnings`; pre-commit passed. | `ap-aYBz0F0ZiQ2nTGndudnGaH` |
| Real critic RL smoke | Loaded LoRA critic student; generated rollouts; completed one GRPO update; printed `INTERLEAVE_CRITIC_RL_SMOKE_OK`. | `ap-eXMO3I81OcCyxj53XbPWj9` |
| Reference-policy KL | `16 passed, 14 warnings`; real reference smoke printed `INTERLEAVE_CRITIC_RL_REFERENCE_SMOKE_OK`. | `ap-UQ38OTnymREO9bz0L1QzC5` |
| Planner GRPO | `37 passed, 14 warnings`; real planner GRPO smoke printed `INTERLEAVE_PLANNER_RL_SMOKE_OK`. | `ap-PDBijC8opxsMiMU0Uc064A` |
| Prompt-set evaluation helpers | `15 passed, 14 warnings`; pre-commit passed. | `ap-eeQpAgNQvQGi2H8MB0kJCU` |
| Trace-level evaluation helpers | `19 passed, 14 warnings`; pre-commit passed. | `ap-s7ewT9rDZSTdPNhyYrEYO7` |
## Recommended PR Stack
1. **Python orchestration shell**
- schema and trace dataclasses;
- generator backend protocol;
- provider adapters;
- fake-backend tests.
2. **Qwen3-VL actor wrappers**
- shared Qwen actor base;
- planner and critic wrappers;
- data normalization helpers;
- real checkpoint load smokes.
3. **SFT path**
- `InterleaveThinkerSFTMethod`;
- planner and critic SFT configs;
- response-token masking tests.
4. **Reward and API backend path**
- InterleaveThinker reward parser/scorers;
- Gemini and Nano Banana wrappers;
- fake-client tests.
5. **GRPO path**
- shared GRPO helpers;
- `InterleaveThinkerRLMethod`;
- critic GRPO, reference KL, planner GRPO;
- real one-step LoRA smokes.
6. **Evaluation and docs**
- prompt-set runner;
- trace evaluator and HTML report helpers;
- examples and design docs.
Each PR should keep the handoff updated until it lands or is superseded.
## Remaining Risks
- **Full 8B training memory:** Real one-step LoRA smokes passed. Full-parameter
8B optimizer training and longer distributed runs still need dedicated
hardware validation.
- **Checkpoint/resume:** Configs include checkpoint settings, but planner/critic
SFT and GRPO checkpoint/resume smokes are not yet recorded.
- **Closed-source API drift:** Gemini and Nano Banana wrappers are unit-tested
with fake clients. Live API outputs can change and should not be deterministic
CI baselines.
- **EasyR1 parity:** The FastVideo GRPO path matches the important objective
pieces used here, but it is not a wholesale EasyR1/Verl port. Distributed
rollout semantics and memory strategy should remain explicit in docs.
- **Native Qwen3-VL port:** The branch uses Transformers Qwen3-VL wrappers. A
FastVideo-native Qwen3-VL port should only be considered if conversion,
performance, or distributed execution needs justify it.
- **End-to-end real generator cost:** Real planner/critic and real FastVideo
generator pieces have smoke coverage, but large prompt-set runs with all real
components can be expensive and should be scheduled intentionally.
## Review Checklist
- Confirm no training code imports from the legacy `fastvideo/training/` stack.
- Confirm API clients import optional dependencies lazily.
- Confirm fake-provider tests cover planner, critic, generator, reward, and
trace-evaluation behavior without credentials.
- Confirm real-checkpoint smoke commands document whether they used a pushed
commit or an explicitly approved Modal patch upload.
- Confirm public YAML configs are parseable and clearly state credential,
dataset, and hardware assumptions.
+10 -2
View File
@@ -112,9 +112,17 @@ teacher/critic — no code changes needed.
## Model Abstraction
### `ModelBase` — Standard (Bidirectional) Models
### `RoleModelBase` — Minimal Role Models
Every role gets its own `ModelBase` instance owning a `transformer` and
Every training role gets a role-model instance with role-local trainability,
LoRA setup, a `transformer`, and lifecycle hooks such as
`init_preprocessors()` and `on_train_start()`. Non-diffusion actors can inherit
from this base directly when they do not own a scheduler or diffusion runtime
primitives.
### `ModelBase` — Standard (Bidirectional) Diffusion Models
Diffusion roles inherit `ModelBase`, which extends `RoleModelBase` with a
`noise_scheduler`. The base class defines:
- **`prepare_batch()`** — Convert raw dataloader output into forward-ready
+53
View File
@@ -0,0 +1,53 @@
# FastVideo Interleave Examples
This directory contains a small Python example for the reusable Interleave
orchestration helpers. It does not add FastVideo CLI commands or HTTP routes.
## Single-Prompt Trace
Run:
```bash
FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA \
python examples/interleave/interleave_single_prompt.py \
--model-path black-forest-labs/FLUX.2-klein-4B \
--prompt "a brushed steel espresso machine on a marble counter, morning window light" \
--output-dir outputs/interleave_single_prompt
```
The script uses `VideoGenerator` directly with a fallback single-prompt planner
and accept-all critic. It writes an image plus `trace.json`; the trace records
planner/generator/critic attempts and omits base64 image payloads by default.
## Planner And Critic
The real InterleaveThinker planner and critic wrappers are integrated through
the existing FastVideo training config system:
- `examples/train/configs/interleave_thinker/planner_sft_lora.yaml`
- `examples/train/configs/interleave_thinker/critic_sft_lora.yaml`
- `examples/train/configs/interleave_thinker/planner_smoke.yaml`
- `examples/train/configs/rl/interleave_thinker/critic_grpo.yaml`
- `examples/train/configs/rl/interleave_thinker/planner_grpo.yaml`
## Optional Gemini Backends
The RL reward config can use closed-source Google models through lazy wrappers:
- `fastvideo.train.methods.rl.rewards.GeminiNanoBananaEditScorer` generates
edits with Nano Banana and scores them with Gemini.
- `fastvideo.workflow.interleave_thinker.generator.NanoBananaImageGeneratorBackend`
implements the same image backend protocol as the local FastVideo generator.
Install the optional SDK and provide a key only when using these API backends:
```bash
uv pip install -e ".[eval-judge]"
export GEMINI_API_KEY=...
```
Supported Nano Banana aliases are:
- `nano-banana` -> `gemini-2.5-flash-image`
- `nano-banana-pro` -> `gemini-3-pro-image`
- `nano-banana-2` -> `gemini-3.1-flash-image`
@@ -0,0 +1,107 @@
# SPDX-License-Identifier: Apache-2.0
"""Run a one-step interleaved generation trace through FastVideo.
This is intentionally small: it uses the fallback single-prompt planner and an
accept-all critic, so it exercises the native Interleave helper layer without
requiring InterleaveThinker planner/critic checkpoints.
"""
from __future__ import annotations
import argparse
import os
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.api.schema import (
EngineConfig,
GeneratorConfig,
OffloadConfig,
PipelineSelection,
)
from fastvideo.workflow.interleave_thinker import (
AcceptAllCritic,
FastVideoImageGeneratorBackend,
InterleaveOrchestrator,
SinglePromptPlanner,
save_trace,
)
DEFAULT_PROMPT = "a brushed steel espresso machine on a marble counter, morning window light"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run a one-step FastVideo interleave trace.")
parser.add_argument(
"--model-path",
default="black-forest-labs/FLUX.2-klein-4B",
help="HF id or local diffusers-format image model directory.",
)
parser.add_argument("--prompt", default=DEFAULT_PROMPT, help="Interleaved generation instruction.")
parser.add_argument("--output-dir", default="outputs/interleave_single_prompt", help="Output directory.")
parser.add_argument("--trace-path", default=None, help="Trace JSON path. Defaults under output-dir.")
parser.add_argument("--seed", type=int, default=0, help="Generation seed.")
parser.add_argument("--height", type=int, default=1024, help="Output image height.")
parser.add_argument("--width", type=int, default=1024, help="Output image width.")
parser.add_argument("--steps", type=int, default=4, help="Number of denoising steps.")
parser.add_argument("--num-gpus", type=int, default=1, help="Number of GPUs to use.")
parser.add_argument(
"--backend",
default=None,
help="Set FASTVIDEO_ATTENTION_BACKEND, for example TORCH_SDPA.",
)
return parser.parse_args()
def main() -> None:
args = parse_args()
if args.backend:
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
output_dir = Path(args.output_dir)
trace_path = Path(args.trace_path) if args.trace_path else output_dir / "trace.json"
generator_config = GeneratorConfig(
model_path=args.model_path,
engine=EngineConfig(
num_gpus=args.num_gpus,
use_fsdp_inference=False,
offload=OffloadConfig(
dit=False,
vae=True,
text_encoder=True,
pin_cpu_memory=False,
),
),
pipeline=PipelineSelection(workload_type="t2i"),
)
generator = VideoGenerator.from_config(generator_config)
try:
backend = FastVideoImageGeneratorBackend(
generator,
output_dir=str(output_dir),
)
orchestrator = InterleaveOrchestrator(
planner=SinglePromptPlanner(),
generator=backend,
critic=AcceptAllCritic(),
width=args.width,
height=args.height,
num_inference_steps=args.steps,
guidance_scale=1.0,
seed=args.seed,
)
trace = orchestrator.run(args.prompt)
save_trace(trace, trace_path)
if trace.final_image is None or trace.final_image.file_path is None:
raise RuntimeError("Interleave run completed without a final image path")
print(f"Image: {trace.final_image.file_path}")
print(f"Trace: {trace_path}")
finally:
generator.shutdown()
if __name__ == "__main__":
main()
@@ -0,0 +1,52 @@
{
"data": [
{
"caption": "gold tip pyramid in the night, extremely detailed , rain, stars"
},
{
"caption": "a fruit stacking in the shape of a dog, stock image, shutterstock"
},
{
"caption": "Landscape, By Lee madgwick, by Luis Royo, by Louise nevelson"
},
{
"caption": "A colorful poster that says \"philo is a weird\""
},
{
"caption": "Danish male with blue eyes, realistic, viking"
},
{
"caption": "futuristic, cityscape, flying cars, neon lights, towering skyscrapers, glowing purple sky."
},
{
"caption": "a crow with cameras for eyes, sitting on a mans shoulder, anime, studio ghibli, fantasy, fairytale, sketch, digital art, watercolor, dnd, rustic, professional photograph, medieval, hd, 4k"
},
{
"caption": "a background image mixing the matrix and AI"
},
{
"caption": "Golden sunset, a bright orange and yellow sky is visible, lit up by the setting sun, the horizon is a mix of bright colors and deep shadows"
},
{
"caption": "Grim reaper playing an electric guitar"
},
{
"caption": "an epic view of a demonic Rose-ringed parakeet cyborg inside an ironmaiden robot,wearing a noble robe,large view,a surrealist painting, aralan bean and Philippe Druillet,hiromu arakawa,volumetric lighting,detailed shadows"
},
{
"caption": "Ben Shapiro as the cover of ministry's filth pig album, but covered in milk"
},
{
"caption": "king charles spaniel with , ethereal, midjourney style lighting and shadows, insanely detailed, 8k, photorealistic"
},
{
"caption": "A website for a party resort service"
},
{
"caption": "full shot of a steampunk horse"
},
{
"caption": "60s psycedelic spiritual jazz album art"
}
]
}
+1
View File
@@ -87,6 +87,7 @@ training:
# --- training.data [TYPED] -> DataConfig ---
data:
data_path: data/my_dataset # default: ""
preprocessed_data_type: t2v # default: "t2v" ("text_only" for simulate-only DMD text prompts)
train_batch_size: 1 # default: 1
dataloader_num_workers: 4 # default: 0
training_cfg_rate: 0.1 # default: 0.0
@@ -0,0 +1,68 @@
# InterleaveThinker critic LoRA SFT config.
#
# Expects upstream `critic_sft.json` from InterleaveThinker/Train-Data.
models:
student:
_target_: fastvideo.train.models.interleave_thinker.InterleaveThinkerCriticModel
init_from: InterleaveThinker/Critic-SFT-8B
processor_from: Qwen/Qwen3-VL-8B-Instruct
load_backend: true
trainable: true
dataset_kind: critic_sft
image_dir: data/InterleaveThinker/Train-Data
torch_dtype: auto
attn_implementation: flash_attention_2
freeze_vision_tower: true
freeze_multi_modal_projector: true
enable_gradient_checkpointing: true
max_prompt_length: 16384
max_response_length: 4096
lora:
enable: true
rank: 16
alpha: 32
target_modules:
- q_proj
- k_proj
- v_proj
- o_proj
method:
_target_: fastvideo.train.methods.fine_tuning.InterleaveThinkerSFTMethod
training:
distributed:
num_gpus: 8
sp_size: 1
# Qwen actors use FSDP2/HSDP; tensor parallelism is not implemented.
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/InterleaveThinker/Train-Data/critic_sft.json
preprocessed_data_type: text_only
train_batch_size: 1
dataloader_num_workers: 0
seed: 42
optimizer:
learning_rate: 1.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/interleave_thinker_critic_sft_lora
training_state_checkpointing_steps: 50
checkpoints_total_limit: 3
tracker:
project_name: InterleaveThinker-SFT
run_name: interleave_thinker_critic_sft_lora
@@ -0,0 +1,68 @@
# InterleaveThinker planner LoRA SFT config.
#
# Expects upstream `planner_sft.json` from InterleaveThinker/Train-Data.
models:
student:
_target_: fastvideo.train.models.interleave_thinker.InterleaveThinkerPlannerModel
init_from: InterleaveThinker/InterleaveThinker-Planner-8B
processor_from: Qwen/Qwen3-VL-8B-Instruct
load_backend: true
trainable: true
dataset_kind: planner_sft
image_dir: data/InterleaveThinker/Train-Data
torch_dtype: auto
attn_implementation: flash_attention_2
freeze_vision_tower: true
freeze_multi_modal_projector: true
enable_gradient_checkpointing: true
max_prompt_length: 16384
max_response_length: 4096
lora:
enable: true
rank: 16
alpha: 32
target_modules:
- q_proj
- k_proj
- v_proj
- o_proj
method:
_target_: fastvideo.train.methods.fine_tuning.InterleaveThinkerSFTMethod
training:
distributed:
num_gpus: 8
sp_size: 1
# Qwen actors use FSDP2/HSDP; tensor parallelism is not implemented.
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/InterleaveThinker/Train-Data/planner_sft.json
preprocessed_data_type: text_only
train_batch_size: 1
dataloader_num_workers: 0
seed: 42
optimizer:
learning_rate: 1.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/interleave_thinker_planner_sft_lora
training_state_checkpointing_steps: 50
checkpoints_total_limit: 3
tracker:
project_name: InterleaveThinker-SFT
run_name: interleave_thinker_planner_sft_lora
@@ -0,0 +1,44 @@
# InterleaveThinker planner backend smoke config.
#
# This config is intentionally minimal: it verifies that the planner Qwen3-VL
# actor target and actor-specific SFT method can be assembled with the real
# checkpoint/processor pair. It intentionally omits a dataset so --dry-run
# stops after build_from_config.
models:
student:
_target_: fastvideo.train.models.interleave_thinker.InterleaveThinkerPlannerModel
init_from: InterleaveThinker/InterleaveThinker-Planner-8B
processor_from: Qwen/Qwen3-VL-8B-Instruct
load_backend: true
torch_dtype: auto
attn_implementation: flash_attention_2
freeze_vision_tower: true
freeze_multi_modal_projector: true
enable_gradient_checkpointing: true
max_prompt_length: 16384
max_response_length: 4096
trainable: true
method:
_target_: fastvideo.train.methods.fine_tuning.InterleaveThinkerSFTMethod
training:
data:
data_path: ""
train_batch_size: 1
dataloader_num_workers: 0
seed: 42
optimizer:
learning_rate: 1.0e-5
loop:
max_train_steps: 1
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/interleave_thinker_planner_smoke
tracker:
trackers: []
@@ -0,0 +1,112 @@
# InterleaveThinker critic RL bridge.
#
# Upstream InterleaveThinker trains a Qwen3-VL critic checkpoint with EasyR1
# GRPO. This config expresses the same outer loop through FastVideo's modular
# trainer with a local Qwen3-VL actor wrapper plus API-backed Gemini/Nano
# Banana rewards.
models:
student:
_target_: fastvideo.train.models.interleave_thinker.InterleaveThinkerCriticModel
init_from: InterleaveThinker/Critic-SFT-8B
processor_from: Qwen/Qwen3-VL-8B-Instruct
load_backend: true
dataset_kind: critic_rl
image_dir: data/InterleaveThinker/Train-Data
torch_dtype: auto
attn_implementation: flash_attention_2
freeze_vision_tower: true
freeze_multi_modal_projector: true
enable_gradient_checkpointing: true
max_prompt_length: 16384
max_response_length: 4096
trainable: true
lora:
enable: true
rank: 16
alpha: 32
target_modules:
- q_proj
- k_proj
- v_proj
- o_proj
reference:
_target_: fastvideo.train.models.interleave_thinker.InterleaveThinkerCriticModel
init_from: InterleaveThinker/Critic-SFT-8B
processor_from: Qwen/Qwen3-VL-8B-Instruct
load_backend: true
dataset_kind: critic_rl
image_dir: data/InterleaveThinker/Train-Data
torch_dtype: auto
attn_implementation: flash_attention_2
freeze_vision_tower: true
freeze_multi_modal_projector: true
enable_gradient_checkpointing: false
max_prompt_length: 16384
max_response_length: 4096
trainable: false
method:
_target_: fastvideo.train.methods.rl.interleave_thinker.InterleaveThinkerRLMethod
num_generations: 8
num_batches_per_step: 1
temperature: 1.0
top_p: 1.0
max_new_tokens: 2048
format_weight: 0.5
judge_accuracy_weight: 0.2
semantic_weight: 0.6
quality_weight: 0.2
fallback_edit_reward: 0.5
clip_range: 0.2
kl_coef: 0.01
micro_batch_size_per_device_for_update: 1
edit_scorer:
_target_: fastvideo.train.methods.rl.rewards.GeminiNanoBananaEditScorer
image_model: gemini-3.1-flash-image
judge_model: gemini-2.5-pro
output_dir: outputs/interleave_thinker_critic_grpo/reward_api
width: 1024
height: 1024
num_inference_steps: 4
guidance_scale: 1.0
advantage_eps: 0.0001
advantage_clip: 5.0
max_grad_norm: 1.0
terminal_progress: true
training:
distributed:
num_gpus: 8
sp_size: 1
# Qwen actors use FSDP2/HSDP; tensor parallelism is not implemented.
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/InterleaveThinker/Train-Data/critic_rl.jsonl
preprocessed_data_type: text_only
train_batch_size: 16
dataloader_num_workers: 0
seed: 42
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/interleave_thinker_critic_grpo
training_state_checkpointing_steps: 50
checkpoints_total_limit: 3
tracker:
project_name: EasyR1-qwen3-vl
run_name: interleave_thinker_critic_grpo_fastvideo
@@ -0,0 +1,102 @@
# InterleaveThinker planner GRPO bridge.
#
# This config trains the Qwen3-VL planner policy with FastVideo's managed GRPO
# loop. It uses a lightweight planner reward by default: valid
# <think>...</think><answer>{"execution_plan": [...]}</answer> format, plus an
# optional scalar plan score when present in the dataset.
models:
student:
_target_: fastvideo.train.models.interleave_thinker.InterleaveThinkerPlannerModel
init_from: InterleaveThinker/InterleaveThinker-Planner-8B
processor_from: Qwen/Qwen3-VL-8B-Instruct
load_backend: true
trainable: true
dataset_kind: planner_rl
image_dir: data/InterleaveThinker/Train-Data
torch_dtype: auto
attn_implementation: flash_attention_2
freeze_vision_tower: true
freeze_multi_modal_projector: true
enable_gradient_checkpointing: true
max_prompt_length: 16384
max_response_length: 4096
lora:
enable: true
rank: 16
alpha: 32
target_modules:
- q_proj
- k_proj
- v_proj
- o_proj
reference:
_target_: fastvideo.train.models.interleave_thinker.InterleaveThinkerPlannerModel
init_from: InterleaveThinker/InterleaveThinker-Planner-8B
processor_from: Qwen/Qwen3-VL-8B-Instruct
load_backend: true
trainable: false
dataset_kind: planner_rl
image_dir: data/InterleaveThinker/Train-Data
torch_dtype: auto
attn_implementation: flash_attention_2
freeze_vision_tower: true
freeze_multi_modal_projector: true
enable_gradient_checkpointing: false
max_prompt_length: 16384
max_response_length: 4096
method:
_target_: fastvideo.train.methods.rl.interleave_thinker.InterleaveThinkerRLMethod
reward_scorer:
_target_: fastvideo.train.methods.rl.rewards.InterleavePlannerRewardScorer
format_weight: 1.0
fallback_plan_reward: 0.0
num_generations: 8
num_batches_per_step: 1
temperature: 1.0
top_p: 1.0
max_new_tokens: 2048
clip_range: 0.2
kl_coef: 0.01
micro_batch_size_per_device_for_update: 1
advantage_eps: 0.0001
advantage_clip: 5.0
max_grad_norm: 1.0
terminal_progress: true
training:
distributed:
num_gpus: 8
sp_size: 1
# Qwen actors use FSDP2/HSDP; tensor parallelism is not implemented.
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 8
data:
data_path: data/InterleaveThinker/Train-Data/planner_rl.jsonl
preprocessed_data_type: text_only
train_batch_size: 16
dataloader_num_workers: 0
seed: 42
optimizer:
learning_rate: 2.0e-6
betas: [0.9, 0.999]
weight_decay: 0.0
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 1000
gradient_accumulation_steps: 1
checkpoint:
output_dir: outputs/interleave_thinker_planner_grpo
training_state_checkpointing_steps: 50
checkpoints_total_limit: 3
tracker:
project_name: InterleaveThinker-GRPO
run_name: interleave_thinker_planner_grpo_fastvideo
@@ -0,0 +1,116 @@
# DiffusionNFT multi-reward single-frame RL: Wan 2.1 T2V 1.3B on text-only PickScore prompts.
#
# Single-frame RL is represented as a one-latent-frame Wan run:
# num_latent_t: 1
# num_frames: 1
#
# The method trains the full transformer (no LoRA) and keeps an old-policy
# transformer plus a frozen reference transformer, matching the non-LoRA
# DiffusionNFT loss path.
models:
student:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: true
old:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
reference:
_target_: fastvideo.train.models.wan.WanModel
init_from: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
trainable: false
disable_custom_init_weights: true
method:
_target_: fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod
reward_fn:
pickscore: 1.0
clipscore: 1.0
sampling:
num_steps: 25
scheduler: flow_match_euler
trajectory: ode
flow_shift: inherit
validation:
every_steps: 10
num_steps: 40
num_prompts: 16
batch_size: 16
log_samples: true
seed: 42
# Null reuses training.data.data_path. Override this with a held-out
# preprocessed parquet path when one is available.
data_path:
# DiffusionNFT sd3_multi_reward on 4 GPUs resolves to per-GPU sample batch
# size 6, 48 sample batches per outer epoch, and grad accumulation 48.
sample_train_batch_size: 6
train_batch_size: 6
num_batches_per_epoch: 48
num_video_per_prompt: 24
num_inner_epochs: 1
timestep_fraction: 0.99
beta: 0.1
kl_beta: 0.0001
decay_type: 1
adv_mode: all
adv_clip_max: 5
max_grad_norm: 1.0
ema:
enabled: true
decay: 0.9
update_after_step: 0
validation: true
terminal_progress: true
training:
distributed:
num_gpus: 4
sp_size: 1
tp_size: 1
hsdp_replicate_dim: 1
hsdp_shard_dim: 4
data:
data_path: data/pickscore_text_only_preprocessed
preprocessed_data_type: text_only
dataloader_num_workers: 0
train_batch_size: 1
training_cfg_rate: 0.0
seed: 42
num_latent_t: 1
num_height: 448
num_width: 832
num_frames: 1
optimizer:
learning_rate: 3.0e-5
betas: [0.9, 0.999]
weight_decay: 0.0001
lr_scheduler: constant
lr_warmup_steps: 0
loop:
max_train_steps: 100000
gradient_accumulation_steps: 48
checkpoint:
output_dir: outputs/wan2.1_diffusion_nft_pick_clip
training_state_checkpointing_steps: 30
checkpoints_total_limit: 3
tracker:
project_name: diffusion_nft_wan
run_name: wan2.1_diffusion_nft_pick_clip
model:
enable_gradient_checkpointing_type: full
pipeline:
flow_shift: 8
+29
View File
@@ -286,6 +286,35 @@ def run_train_framework_tests():
)
@app.function(gpu="L40S:1",
image=image,
timeout=1800,
secrets=[
modal.Secret.from_dict(
{"HF_API_KEY": os.environ.get("HF_API_KEY", "")})
],
volumes={"/root/data": model_vol})
def seed_grad_norm_references():
"""Record the per-method grad-norm reference for the **CI GPU (L40S only)**.
Phase 2 / 5a-ii one-off seeding entrypoint. Pinned to ``gpu="L40S:1"`` (the
Modal CI runner), so this function only seeds the ``L40S`` key in
``fastvideo/tests/train/methods/grad_norm_refs.json``.
``FASTVIDEO_GRADNORM_UPDATE=1`` makes ``check_grad_norm_regression`` record
the measured norm instead of asserting; ``-rs`` surfaces the recorded value
in the log so it can be copied into the JSON.
To seed any other device (e.g. our local Blackwell dev box → ``GB200``
key), run the same env-var + pytest invocation directly on that
workstation — see the module docstring of ``grad_norm_regression.py`` for
the local command and the ``_DEVICE_MAPPINGS`` table.
"""
run_test(
"export HF_HOME='/root/data/.cache' && hf auth login --token $HF_API_KEY && FASTVIDEO_GRADNORM_UPDATE=1 pytest ./fastvideo/tests/train/methods -vs -rs"
)
@app.function(gpu="L40S:1",
image=image,
timeout=3600,
@@ -67,6 +67,8 @@ class TestConstructor:
assert cb.num_frames is None
assert cb.sampling_timesteps is None
assert cb.output_dir is None
assert cb.offload_training_state is False
assert cb.unload_pipeline_after_validation is False
# Lazy fields not yet populated.
assert cb._pipeline is None
assert cb._sampling_param is None
@@ -83,12 +85,16 @@ class TestConstructor:
guidance_scale="4.5", # type: ignore[arg-type]
num_frames="77", # type: ignore[arg-type]
sampling_timesteps=["1000", "500"],
offload_training_state="1", # type: ignore[arg-type]
unload_pipeline_after_validation="false", # type: ignore[arg-type]
)
assert cb.every_steps == 50
assert cb.sampling_steps == [20, 40]
assert cb.guidance_scale == 4.5
assert cb.num_frames == 77
assert cb.sampling_timesteps == [1000, 500]
assert cb.offload_training_state is True
assert cb.unload_pipeline_after_validation is False
def test_pipeline_kwargs_collected(self) -> None:
cb = ValidationCallback(
@@ -0,0 +1,10 @@
{
"test_wan_causal_dfsft": {
"GB200": 2.9781,
"L40S": 3.2562
},
"test_wan_finetune": {
"GB200": 1.6486,
"L40S": 1.6467
}
}
@@ -0,0 +1,158 @@
# SPDX-License-Identifier: Apache-2.0
"""Layer-0 grad-norm regression for the per-method training smoke tests.
Phase 2 / 5a-ii: layers a device-keyed grad-norm check on top of the
finite/non-zero grad assertions established in 5a-i. After one
``single_train_step`` + ``backward``, the L2 norm of transformer block 0's
trainable gradients is compared against a reference value pinned per GPU in
``grad_norm_refs.json`` (next to this module).
Determinism: the harness seeds both the global RNG and the method's
``cuda_generator`` via ``method.on_train_start()`` (``training.data.seed`` in the
fixture), and the synthetic ``raw_batch`` is built *after* that call, so the
forward/backward is reproducible within bf16 reduction noise on a given GPU.
Why device-keyed: grad norms differ across GPU architectures (kernels,
accumulation order), so a single golden value can't cover every runner. The
JSON currently carries refs for the two GPUs we actually run on — ``L40S`` (CI)
and ``GB200`` (our Blackwell dev box; ``B200`` maps to the same key).
Seeding a reference for the current device:
- **CI / L40S** — invoke ``modal run`` against ``seed_grad_norm_references`` in
``fastvideo/tests/modal/pr_test.py`` (pinned to ``gpu="L40S:1"``), then copy
the recorded value from the log into ``grad_norm_refs.json``.
- **Local / non-L40S GPUs** — on that workstation::
FASTVIDEO_GRADNORM_UPDATE=1 \\
pytest fastvideo/tests/train/methods -vs -rs
The harness writes the measured norm into ``grad_norm_refs.json`` under the
device's key and skips the assertion for that run. Append a new substring
entry to ``_DEVICE_MAPPINGS`` first for any device not already listed.
"""
from __future__ import annotations
import json
import os
from pathlib import Path
import pytest
import torch
_REFS_PATH = Path(__file__).resolve().parent / "grad_norm_refs.json"
_UPDATE_ENV = "FASTVIDEO_GRADNORM_UPDATE"
# bf16 single-step smoke: catch gross breakage (wrong wiring, dead grads,
# scale regressions), not micro-drift from reduction nondeterminism.
_DEFAULT_RTOL = 0.10
# GPU-name substring -> reference key. First match wins. Only devices with
# seeded references in ``grad_norm_refs.json`` are listed here — to add a new
# GPU, append an entry, then seed the reference (see module docstring).
_DEVICE_MAPPINGS: tuple[tuple[str, str], ...] = (
("L40S", "L40S"),
("GB200", "GB200"),
("B200", "GB200"), # same Blackwell arch as GB200
)
def _device_name() -> str:
if not torch.cuda.is_available():
return "CPU"
return torch.cuda.get_device_name(0)
def resolve_device_key(device_name: str | None = None) -> str | None:
"""Map a CUDA device name to its reference key, or None if unsupported.
The substring match is case-insensitive so it survives driver/environment
differences in how ``torch.cuda.get_device_name`` capitalizes the model.
"""
name = device_name if device_name is not None else _device_name()
name_lower = name.lower()
for pattern, key in _DEVICE_MAPPINGS:
if pattern.lower() in name_lower:
return key
return None
def layer0_grad_norm(transformer) -> float:
"""Global L2 norm of transformer block 0's trainable gradients.
Block 0 is the reference surface 5a-i already isolates: its grad is the
*last* one produced during backprop, so a healthy value implies the whole
forward + chain-rule path is intact.
Accumulates the squared sums on the GPU and does a single CPU-GPU sync
(``.item()``) at the end, rather than one per parameter.
"""
blocks = getattr(transformer, "blocks", None)
assert blocks is not None and len(blocks) > 0, (
"transformer is expected to expose a non-empty ``.blocks``")
grads = [
p.grad for p in blocks[0].parameters()
if p.requires_grad and p.grad is not None
]
if not grads:
return 0.0
sq_sum = torch.zeros((), device=grads[0].device, dtype=torch.float32)
for g in grads:
sq_sum += g.detach().float().pow(2).sum()
return sq_sum.sqrt().item()
def _load_refs() -> dict[str, dict[str, float]]:
if _REFS_PATH.exists():
return json.loads(_REFS_PATH.read_text(encoding="utf-8"))
return {}
def _save_refs(refs: dict[str, dict[str, float]]) -> None:
_REFS_PATH.write_text(
json.dumps(refs, indent=2, sort_keys=True) + "\n",
encoding="utf-8")
def check_grad_norm_regression(
test_name: str,
transformer,
*,
rtol: float = _DEFAULT_RTOL,
) -> None:
"""Assert block-0 grad norm matches the device-keyed reference within rtol.
- Skips when the current GPU has no reference (unsupported device, or not
yet seeded) so a new runner never hard-fails before its golden exists.
- With ``FASTVIDEO_GRADNORM_UPDATE=1`` records/updates the reference for the
current device instead of asserting.
"""
norm = layer0_grad_norm(transformer)
device_key = resolve_device_key()
if os.environ.get(_UPDATE_ENV) == "1":
if device_key is None:
pytest.skip(
f"{_UPDATE_ENV}=1 but GPU '{_device_name()}' has no reference "
"key; add it to _DEVICE_MAPPINGS first")
refs = _load_refs()
refs.setdefault(test_name, {})[device_key] = round(norm, 4)
_save_refs(refs)
pytest.skip(
f"recorded grad-norm reference {test_name}[{device_key}] = "
f"{norm:.4f} (assertion skipped under {_UPDATE_ENV}=1)")
ref = _load_refs().get(test_name, {}).get(device_key) \
if device_key is not None else None
if ref is None:
pytest.skip(
f"no grad-norm reference for {test_name} on '{_device_name()}' "
f"(device_key={device_key}); run with {_UPDATE_ENV}=1 to seed it")
rel = abs(norm - ref) / (abs(ref) + 1e-12)
assert rel <= rtol, (
f"{test_name}[{device_key}] grad-norm regression: got {norm:.4f}, "
f"reference {ref:.4f}, relative error {rel:.3%} exceeds rtol "
f"{rtol:.0%}. If this is an intentional change, refresh the reference "
f"with {_UPDATE_ENV}=1 and explain why in the PR.")
@@ -28,6 +28,8 @@ from fastvideo.train.methods.fine_tuning.dfsft import (
from fastvideo.train.models.wan import WanCausalModel
from fastvideo.train.utils.config import load_run_config
from .grad_norm_regression import check_grad_norm_regression
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
@@ -122,3 +124,7 @@ def test_wan_causal_dfsft_single_train_step(
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_causal_dfsft", model.transformer)
@@ -36,6 +36,8 @@ from fastvideo.train.methods.fine_tuning.finetune import (
from fastvideo.train.models.wan import WanModel
from fastvideo.train.utils.config import load_run_config
from .grad_norm_regression import check_grad_norm_regression
_FIXTURE = str(
Path(__file__).resolve().parent.parent / "fixtures"
@@ -139,3 +141,7 @@ def test_wan_finetune_single_train_step(
assert any_nonzero, (
"all layer-0 grads are exactly zero; backward did not "
"reach the first transformer block")
# 5a-ii: device-keyed grad-norm regression on top of the same harness.
# Skips when the current GPU has no seeded reference.
check_grad_norm_regression("test_wan_finetune", model.transformer)
+2 -2
View File
@@ -18,7 +18,7 @@ train/
│ ├── knowledge_distillation/ # KDMethod, KDCausalMethod
│ └── consistency_model/ # Consistency-model training methods
├── models/
│ ├── base.py # ModelBase / CausalModelBase wrappers
│ ├── base.py # RoleModelBase / ModelBase / CausalModelBase wrappers
│ ├── wan/, hunyuan/, cosmos/ # Per-family training wrappers
├── callbacks/ # callback.py base + ema, grad_clip, validation
└── utils/
@@ -47,7 +47,7 @@ Trainer = Method × Model × [Callback...] × Config
## Adding a New Model Plugin
1. Subclass `ModelBase` (or `CausalModelBase`) in `models/<family>/`.
1. Subclass `ModelBase` (or `CausalModelBase`) in `models/<family>/` for diffusion models. Use `RoleModelBase` only for non-diffusion role actors.
2. Wrap the existing inference DiT from `fastvideo/models/dits/`. Do not
reimplement.
3. Expose `trainable_parameters()` so the optimizer factory can group them.
+186 -8
View File
@@ -68,6 +68,8 @@ class ValidationCallback(Callback):
num_frames: int | None = None,
output_dir: str | None = None,
sampling_timesteps: list[int] | None = None,
offload_training_state: bool = False,
unload_pipeline_after_validation: bool = False,
**pipeline_kwargs: Any,
) -> None:
self.pipeline_target = str(pipeline_target)
@@ -78,6 +80,8 @@ class ValidationCallback(Callback):
self.num_frames = (int(num_frames) if num_frames is not None else None)
self.output_dir = (str(output_dir) if output_dir is not None else None)
self.sampling_timesteps = ([int(s) for s in sampling_timesteps] if sampling_timesteps is not None else None)
self.offload_training_state = self._coerce_bool(offload_training_state)
self.unload_pipeline_after_validation = self._coerce_bool(unload_pipeline_after_validation)
self.pipeline_kwargs = dict(pipeline_kwargs)
# Set after on_train_start.
@@ -88,6 +92,12 @@ class ValidationCallback(Callback):
self.validation_random_generator: (torch.Generator | None) = None
self.seed: int = 0
@staticmethod
def _coerce_bool(value: Any) -> bool:
if isinstance(value, str):
return value.strip().lower() in {"1", "true", "yes", "on"}
return bool(value)
# ----------------------------------------------------------
# Callback hooks
# ----------------------------------------------------------
@@ -140,16 +150,183 @@ class ValidationCallback(Callback):
) -> None:
transformer = method.student.transformer
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
self._run_validation_inner(
try:
with self._validation_memory_context(
method,
validation_transformer=transformer,
):
# Look for an EMA callback to temporarily swap
# EMA weights during validation.
ema_cb = self._find_ema_callback()
ctx = ema_cb.ema_context(transformer) if ema_cb is not None else contextlib.nullcontext(transformer)
with ctx as t:
self._run_validation_inner(
method,
step,
t,
)
finally:
if self.unload_pipeline_after_validation:
self._clear_pipeline_cache()
@contextlib.contextmanager
def _validation_memory_context(
self,
method: TrainingMethod,
*,
validation_transformer: torch.nn.Module,
):
if not self.offload_training_state:
yield
return
optimizer_tensor_records: list[tuple[Any, Any, torch.device]] = []
module_records: list[tuple[str, torch.nn.Module, torch.device]] = []
try:
self._offload_optimizer_states_to_cpu(
method,
step,
t,
optimizer_tensor_records,
)
self._offload_inactive_role_modules_to_cpu(
method,
validation_transformer=validation_transformer,
module_records=module_records,
)
self._empty_cuda_cache()
yield
finally:
self._restore_inactive_role_modules(module_records)
self._restore_optimizer_states(optimizer_tensor_records)
self._empty_cuda_cache()
def _offload_optimizer_states_to_cpu(
self,
method: TrainingMethod,
records: list[tuple[Any, Any, torch.device]],
) -> None:
optimizers = getattr(method, "_optimizer_dict", {})
if not optimizers:
return
moved = 0
for optimizer in optimizers.values():
state = getattr(optimizer, "state", None)
if not isinstance(state, dict):
continue
for param_state in state.values():
moved += self._offload_tensor_container_to_cpu(
param_state,
records,
)
if moved:
logger.info(
"Offloaded %d optimizer state tensors to CPU for validation.",
moved,
)
def _offload_tensor_container_to_cpu(
self,
obj: Any,
records: list[tuple[Any, Any, torch.device]],
) -> int:
moved = 0
if isinstance(obj, dict):
for key, value in list(obj.items()):
if torch.is_tensor(value) and value.device.type == "cuda":
records.append((obj, key, value.device))
obj[key] = value.detach().cpu()
moved += 1
else:
moved += self._offload_tensor_container_to_cpu(value, records)
return moved
if isinstance(obj, list):
for idx, value in enumerate(list(obj)):
if torch.is_tensor(value) and value.device.type == "cuda":
records.append((obj, idx, value.device))
obj[idx] = value.detach().cpu()
moved += 1
else:
moved += self._offload_tensor_container_to_cpu(value, records)
return moved
def _restore_optimizer_states(
self,
records: list[tuple[Any, Any, torch.device]],
) -> None:
for container, key, device in reversed(records):
value = container[key]
if torch.is_tensor(value):
container[key] = value.to(device=device)
if records:
logger.info(
"Restored %d optimizer state tensors after validation.",
len(records),
)
def _offload_inactive_role_modules_to_cpu(
self,
method: TrainingMethod,
*,
validation_transformer: torch.nn.Module,
module_records: list[tuple[str, torch.nn.Module, torch.device]],
) -> None:
role_models = getattr(method, "_role_models", {})
if not isinstance(role_models, dict):
return
for role, model in role_models.items():
module = getattr(model, "transformer", None)
if not isinstance(module, torch.nn.Module):
continue
if module is validation_transformer:
continue
device = self._first_cuda_tensor_device(module)
if device is None:
continue
try:
module.to("cpu")
except Exception as exc:
logger.warning(
"Could not offload role %r transformer to CPU before validation: %s",
role,
exc,
)
continue
module_records.append((str(role), module, device))
logger.info(
"Offloaded role %r transformer from %s to CPU for validation.",
role,
device,
)
def _restore_inactive_role_modules(
self,
module_records: list[tuple[str, torch.nn.Module, torch.device]],
) -> None:
for role, module, device in reversed(module_records):
module.to(device)
logger.info(
"Restored role %r transformer to %s after validation.",
role,
device,
)
@staticmethod
def _first_cuda_tensor_device(module: torch.nn.Module) -> torch.device | None:
for tensor in list(module.parameters(recurse=True)) + list(module.buffers(recurse=True)):
device = getattr(tensor, "device", None)
if isinstance(device, torch.device) and device.type == "cuda":
return device
return None
def _clear_pipeline_cache(self) -> None:
self._pipeline = None
self._pipeline_key = None
self._empty_cuda_cache()
@staticmethod
def _empty_cuda_cache() -> None:
if torch.cuda.is_available():
torch.cuda.empty_cache()
def _find_ema_callback(self) -> Any | None:
"""Find the EMA callback in the callback dict."""
@@ -293,6 +470,7 @@ class ValidationCallback(Callback):
}
if flow_shift is not None:
kwargs["flow_shift"] = float(flow_shift)
kwargs.update(self.pipeline_kwargs)
self._pipeline = PipelineCls.from_pretrained(
tc.model_path,
+45 -5
View File
@@ -3,14 +3,14 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from collections.abc import Sequence
from collections.abc import Mapping, Sequence
from typing import Any, Literal, TypeAlias
import torch
from fastvideo import envs
from fastvideo.logger import init_logger
from fastvideo.train.models.base import ModelBase
from fastvideo.train.models.base import RoleModelBase
from fastvideo.train.utils.checkpoint import _RoleModuleContainer
from fastvideo.training.checkpointing_utils import (
ModelWrapper,
@@ -30,7 +30,7 @@ class TrainingMethod(torch.nn.Module, ABC):
plain attributes and manage optimizers directly — no ``RoleManager``
or ``RoleHandle``.
The constructor receives *role_models* (a ``dict[str, ModelBase]``)
The constructor receives *role_models* (a ``Mapping[str, RoleModelBase]``)
and a *cfg* object. It calls ``init_preprocessors`` on the student
and builds ``self.role_modules`` for FSDP wrapping.
@@ -47,11 +47,11 @@ class TrainingMethod(torch.nn.Module, ABC):
self,
*,
cfg: Any,
role_models: dict[str, ModelBase],
role_models: Mapping[str, RoleModelBase],
) -> None:
super().__init__()
self.tracker: Any | None = None
self._role_models: dict[str, ModelBase] = dict(role_models)
self._role_models: dict[str, RoleModelBase] = dict(role_models)
self.student = role_models["student"]
self.training_config = cfg.training
@@ -197,6 +197,46 @@ class TrainingMethod(torch.nn.Module, ABC):
# -- Shared hooks (override in subclasses as needed) --
def manages_optimization(self) -> bool:
"""Whether the method owns backward/optimizer stepping internally.
Most methods return loss tensors and let :class:`Trainer` handle
gradient accumulation, callbacks, optimizer stepping, and scheduler
stepping. RL-style methods such as DiffusionNFT need to preserve their
own sample-then-inner-train loop, so they can provide a specific
``managed_train_step``.
"""
return False
def managed_train_step(
self,
data_stream: Any,
iteration: int,
) -> tuple[
dict[str, torch.Tensor],
dict[str, Any],
dict[str, LogScalar],
]:
"""Run one method-managed step.
Subclasses that return ``True`` from :meth:`manages_optimization`
should override this. The fallback consumes one dataloader batch and
delegates to ``single_train_step`` so tests can exercise the hook with
tiny fake methods.
"""
return self.single_train_step(next(data_stream), iteration)
def on_validation_begin(self, iteration: int = 0) -> dict[str, LogScalar]:
"""Run method-owned validation, if any.
Pipeline-style validation should remain in callbacks. Methods that
intentionally avoid inference pipelines, such as RL methods with their
own sampler/reward loop, can override this hook and return metrics for
the trainer to log at ``iteration``.
"""
del iteration
return {}
def get_grad_clip_targets(
self,
iteration: int,
@@ -55,6 +55,8 @@ class DMD2Method(TrainingMethod):
raise ValueError("DMD2Method requires critic to be trainable")
self._cfg_uncond = self._parse_cfg_uncond()
self._rollout_mode = self._parse_rollout_mode()
self._validate_preprocessed_data_type()
self._configure_student_negative_conditioning()
self._denoising_step_list: torch.Tensor | None = (None)
# Initialize preprocessors on student.
@@ -206,6 +208,13 @@ class DMD2Method(TrainingMethod):
return targets
def _parse_rollout_mode(self, ) -> Literal["simulate", "data_latent"]:
"""Parse how DMD2 obtains the latent point used for rollout.
``simulate`` starts from fresh noise and lets the student create an
artificial latent trajectory, so it can run with text-only data.
``data_latent`` starts from preprocessed VAE latents and perturbs them
at a sampled denoising timestep.
"""
raw = self.method_config.get("rollout_mode", None)
if raw is None:
raise ValueError("method_config.rollout_mode must be set "
@@ -223,6 +232,34 @@ class DMD2Method(TrainingMethod):
"{simulate, data_latent}, got "
f"{raw!r}")
def _validate_preprocessed_data_type(self) -> None:
data_type = str(getattr(
self.training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
if data_type == "text_only" and self._rollout_mode != "simulate":
raise ValueError("training.data.preprocessed_data_type='text_only' "
"requires method.rollout_mode='simulate'; "
"data_latent rollout requires vae_latent data.")
def _uses_negative_prompt_conditioning(self) -> bool:
if self._cfg_uncond is None:
return True
text_policy = self._cfg_uncond.get("text", None)
if text_policy is None:
return True
return str(text_policy).strip().lower() == "negative_prompt"
def _configure_student_negative_conditioning(self) -> None:
setter = getattr(
self.student,
"set_requires_negative_conditioning",
None,
)
if setter is not None:
setter(self._uses_negative_prompt_conditioning())
def _parse_cfg_uncond(self, ) -> dict[str, Any] | None:
raw = self.method_config.get("cfg_uncond", None)
if raw is None:
@@ -7,10 +7,13 @@ from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.train.methods.fine_tuning.dfsft import DiffusionForcingSFTMethod
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
from fastvideo.train.methods.fine_tuning.interleave_thinker_sft import (
InterleaveThinkerSFTMethod, )
__all__ = [
"DiffusionForcingSFTMethod",
"FineTuneMethod",
"InterleaveThinkerSFTMethod",
]
@@ -25,4 +28,9 @@ def __getattr__(name: str) -> object:
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
return FineTuneMethod
if name == "InterleaveThinkerSFTMethod":
from fastvideo.train.methods.fine_tuning.interleave_thinker_sft import (
InterleaveThinkerSFTMethod, )
return InterleaveThinkerSFTMethod
raise AttributeError(name)
@@ -0,0 +1,145 @@
# SPDX-License-Identifier: Apache-2.0
"""Supervised fine-tuning method for InterleaveThinker Qwen actors."""
from __future__ import annotations
from collections.abc import Mapping
from typing import Any, Protocol, TypeGuard
import torch
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.models.base import RoleModelBase
from fastvideo.train.utils.optimizer import build_optimizer_and_scheduler
class InterleaveSFTActor(Protocol):
"""Role model contract required by ``InterleaveThinkerSFTMethod``."""
transformer: torch.nn.Module
_trainable: bool
def init_preprocessors(self, training_config: Any) -> None:
...
def compute_interleave_sft_loss(self, batch: Mapping[str, Any]) -> Any:
...
class InterleaveThinkerSFTMethod(TrainingMethod):
"""SFT for planner/critic Qwen3-VL actors with response-token masking."""
def __init__(
self,
*,
cfg: Any,
role_models: Mapping[str, RoleModelBase],
) -> None:
super().__init__(cfg=cfg, role_models=role_models)
if "student" not in role_models:
raise ValueError("InterleaveThinkerSFTMethod requires role 'student'")
student = self.student
if not student._trainable:
raise ValueError("InterleaveThinkerSFTMethod requires student to be trainable")
if not _is_interleave_sft_actor(student):
raise TypeError("InterleaveThinkerSFTMethod requires an Interleave SFT actor implementing "
"compute_interleave_sft_loss()")
self._interleave_student = student
self._interleave_student.init_preprocessors(self.training_config)
self._init_optimizer_and_scheduler()
@property
def _optimizer_dict(self) -> dict[str, Any]:
return {"student": self._student_optimizer}
@property
def _lr_scheduler_dict(self) -> dict[str, Any]:
return {"student": self._student_lr_scheduler}
def single_train_step(
self,
batch: dict[str, Any],
iteration: int,
) -> tuple[
dict[str, torch.Tensor],
dict[str, Any],
dict[str, LogScalar],
]:
del iteration
compute_loss = self._interleave_student.compute_interleave_sft_loss
result = compute_loss(batch)
loss_map, metrics = _coerce_sft_result(result)
return loss_map, {}, metrics
def get_optimizers(
self,
iteration: int,
) -> list[torch.optim.Optimizer]:
del iteration
return [self._student_optimizer]
def get_lr_schedulers(
self,
iteration: int,
) -> list[Any]:
del iteration
return [self._student_lr_scheduler]
def _init_optimizer_and_scheduler(self) -> None:
tc = self.training_config
lr = float(tc.optimizer.learning_rate)
if lr <= 0.0:
raise ValueError("training.optimizer.learning_rate must be > 0 for InterleaveThinker SFT")
params = [p for p in self._interleave_student.transformer.parameters() if p.requires_grad]
if not params:
raise ValueError("InterleaveThinkerSFTMethod found no trainable student parameters")
self._student_optimizer, self._student_lr_scheduler = build_optimizer_and_scheduler(
params=params,
optimizer_config=tc.optimizer,
loop_config=tc.loop,
learning_rate=lr,
betas=tc.optimizer.betas,
scheduler_name=str(tc.optimizer.lr_scheduler),
)
def _coerce_sft_result(result: Any) -> tuple[dict[str, torch.Tensor], dict[str, LogScalar]]:
if isinstance(result, tuple) and len(result) == 2:
loss_map_raw, metrics_raw = result
elif isinstance(result, Mapping):
loss_map_raw = result.get("loss_map")
metrics_raw = result.get("metrics", {})
if loss_map_raw is None and "loss" in result:
loss_map_raw = {"total_loss": result["loss"]}
else:
raise TypeError("compute_interleave_sft_loss() must return (loss_map, metrics) or a mapping")
if not isinstance(loss_map_raw, Mapping):
raise TypeError("compute_interleave_sft_loss() result must include a loss_map mapping")
loss_map = {str(key): _coerce_tensor(value) for key, value in loss_map_raw.items()}
if "total_loss" not in loss_map:
if len(loss_map) != 1:
raise ValueError("SFT loss_map must include total_loss when multiple losses are returned")
loss_map["total_loss"] = next(iter(loss_map.values()))
if metrics_raw is None:
metrics_raw = {}
if not isinstance(metrics_raw, Mapping):
raise TypeError("compute_interleave_sft_loss() metrics must be a mapping")
return loss_map, {str(key): value for key, value in metrics_raw.items()}
def _coerce_tensor(value: Any) -> torch.Tensor:
if isinstance(value, torch.Tensor):
return value
return torch.as_tensor(float(value), dtype=torch.float32)
def _is_interleave_sft_actor(model: Any) -> TypeGuard[InterleaveSFTActor]:
return (isinstance(getattr(model, "transformer", None), torch.nn.Module)
and callable(getattr(model, "init_preprocessors", None))
and callable(getattr(model, "compute_interleave_sft_loss", None)))
__all__ = ["InterleaveSFTActor", "InterleaveThinkerSFTMethod"]
+29
View File
@@ -0,0 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
"""RL training methods."""
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
from fastvideo.train.methods.rl.interleave_thinker import (
InterleaveThinkerRLMethod, )
__all__ = [
"DiffusionNFTMethod",
"InterleaveThinkerRLMethod",
]
def __getattr__(name: str) -> object:
if name == "DiffusionNFTMethod":
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
return DiffusionNFTMethod
if name == "InterleaveThinkerRLMethod":
from fastvideo.train.methods.rl.interleave_thinker import (
InterleaveThinkerRLMethod, )
return InterleaveThinkerRLMethod
raise AttributeError(name)
@@ -0,0 +1,36 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable RL training primitives."""
from fastvideo.train.methods.rl.common.sampling import (
DiffusionSampler,
SamplingConfig,
SamplingResult,
)
from fastvideo.train.methods.rl.common.prompt_sampling import (
KRepeatSample,
distributed_k_repeat_indices,
)
from fastvideo.train.methods.rl.common.grpo import (
GRPOLossResult,
compute_grpo_loss,
)
from fastvideo.train.methods.rl.common.validation import (
RLValidationConfig,
media_to_video_array,
validation_caption,
validation_shard_indices,
)
__all__ = [
"DiffusionSampler",
"GRPOLossResult",
"KRepeatSample",
"RLValidationConfig",
"SamplingConfig",
"SamplingResult",
"compute_grpo_loss",
"distributed_k_repeat_indices",
"media_to_video_array",
"validation_caption",
"validation_shard_indices",
]
+120
View File
@@ -0,0 +1,120 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared GRPO/PPO-ratio loss helpers for RL methods."""
from __future__ import annotations
from dataclasses import dataclass
import torch
@dataclass(frozen=True, slots=True)
class GRPOLossResult:
"""Token-masked GRPO objective and scalar diagnostics."""
total_loss: torch.Tensor
policy_loss: torch.Tensor
kl_loss: torch.Tensor
approx_kl: torch.Tensor
clipped_fraction: torch.Tensor
mean_ratio: torch.Tensor
token_count: torch.Tensor
def compute_grpo_loss(
*,
current_logprobs: torch.Tensor,
old_logprobs: torch.Tensor,
advantages: torch.Tensor,
response_mask: torch.Tensor,
clip_range: float = 0.2,
reference_logprobs: torch.Tensor | None = None,
kl_coef: float = 0.0,
) -> GRPOLossResult:
"""Compute a masked GRPO/PPO-ratio loss.
Shapes:
- ``current_logprobs`` and ``old_logprobs``: ``[B, T]``.
- ``advantages``: ``[B]``.
- ``response_mask``: ``[B, T]`` with non-zero values for trainable
response tokens.
"""
current_logprobs = _require_2d("current_logprobs", current_logprobs)
old_logprobs = _require_2d("old_logprobs", old_logprobs).to(current_logprobs.device)
if current_logprobs.shape != old_logprobs.shape:
raise ValueError("current_logprobs and old_logprobs must have the same shape")
mask = _require_2d("response_mask", response_mask).to(
device=current_logprobs.device,
dtype=current_logprobs.dtype,
)
if mask.shape != current_logprobs.shape:
raise ValueError("response_mask must have the same shape as logprobs")
advantages = advantages.to(device=current_logprobs.device, dtype=current_logprobs.dtype)
if advantages.ndim != 1 or int(advantages.shape[0]) != int(current_logprobs.shape[0]):
raise ValueError("advantages must have shape [B] matching logprobs")
token_count = mask.sum()
if float(token_count.detach().cpu()) <= 0.0:
raise ValueError("GRPO loss requires at least one response token")
clip = float(clip_range)
if clip < 0.0:
raise ValueError("clip_range must be non-negative")
log_ratio = current_logprobs - old_logprobs
ratio = torch.exp(log_ratio)
clipped_ratio = torch.clamp(ratio, 1.0 - clip, 1.0 + clip)
expanded_advantages = advantages[:, None]
surrogate = torch.minimum(
ratio * expanded_advantages,
clipped_ratio * expanded_advantages,
)
policy_loss = -_masked_mean(surrogate, mask)
old_policy_kl = _masked_mean((ratio - 1.0) - log_ratio, mask)
clipped_fraction = _masked_mean((ratio - clipped_ratio).abs().gt(1.0e-6).to(mask.dtype), mask)
mean_ratio = _masked_mean(ratio, mask)
if reference_logprobs is None or float(kl_coef) == 0.0:
kl_loss = torch.zeros((), device=current_logprobs.device, dtype=current_logprobs.dtype)
else:
reference_logprobs = _require_2d("reference_logprobs", reference_logprobs).to(current_logprobs.device)
if reference_logprobs.shape != current_logprobs.shape:
raise ValueError("reference_logprobs must have the same shape as current_logprobs")
ref_delta = reference_logprobs - current_logprobs
kl_loss = _masked_mean(torch.exp(ref_delta) - ref_delta - 1.0, mask)
total_loss = policy_loss + float(kl_coef) * kl_loss
return GRPOLossResult(
total_loss=total_loss,
policy_loss=policy_loss,
kl_loss=kl_loss,
approx_kl=old_policy_kl,
clipped_fraction=clipped_fraction,
mean_ratio=mean_ratio,
token_count=token_count.detach(),
)
def _require_2d(
name: str,
value: torch.Tensor,
) -> torch.Tensor:
if not torch.is_tensor(value):
raise TypeError(f"{name} must be a torch.Tensor")
if value.ndim != 2:
raise ValueError(f"{name} must have shape [B, T], got {tuple(value.shape)}")
return value
def _masked_mean(
value: torch.Tensor,
mask: torch.Tensor,
) -> torch.Tensor:
denom = mask.sum().clamp_min(1.0)
return (value * mask).sum() / denom
__all__ = ["GRPOLossResult", "compute_grpo_loss"]
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""Prompt-row sampling helpers for online RL methods.
This module chooses and repeats dataset prompt rows across ranks for RL training
batches. Here, "sampling" means selection, not generator sampling.
"""
from __future__ import annotations
from dataclasses import dataclass
import torch
@dataclass(frozen=True, slots=True)
class KRepeatSample:
"""Local prompt indices for one distributed K-repeat sampling batch."""
local_indices: list[int]
unique_prompt_count: int
def distributed_k_repeat_indices(
*,
dataset_length: int,
batch_size: int,
repeats_per_prompt: int,
world_size: int,
rank: int,
seed: int,
) -> KRepeatSample:
"""Mirror DiffusionNFT's distributed K-repeat prompt sampler.
Adapted from DiffusionNFT's
``scripts/train_nft_sd3.py::DistributedKRepeatSampler``.
"""
dataset_length = int(dataset_length)
batch_size = int(batch_size)
repeats_per_prompt = int(repeats_per_prompt)
world_size = int(world_size)
rank = int(rank)
if dataset_length <= 0:
raise ValueError("dataset_length must be positive")
if batch_size <= 0:
raise ValueError("batch_size must be positive")
if repeats_per_prompt <= 0:
raise ValueError("repeats_per_prompt must be positive")
if world_size <= 0:
raise ValueError("world_size must be positive")
if rank < 0 or rank >= world_size:
raise ValueError(f"rank must be in [0, {world_size}), got {rank}")
total_samples = world_size * batch_size
if total_samples % repeats_per_prompt != 0:
raise ValueError("world_size * batch_size must be divisible by repeats_per_prompt "
f"({world_size} * {batch_size} vs {repeats_per_prompt})")
unique_prompt_count = total_samples // repeats_per_prompt
if unique_prompt_count > dataset_length:
raise ValueError("K-repeat sampling needs at least as many rows as unique prompts "
f"per sampling batch ({dataset_length} < {unique_prompt_count})")
generator = torch.Generator()
generator.manual_seed(int(seed))
indices = torch.randperm(dataset_length, generator=generator)[:unique_prompt_count].tolist()
repeated_indices = [idx for idx in indices for _ in range(repeats_per_prompt)]
shuffled_order = torch.randperm(len(repeated_indices), generator=generator).tolist()
shuffled_samples = [int(repeated_indices[idx]) for idx in shuffled_order]
start = rank * batch_size
end = start + batch_size
return KRepeatSample(
local_indices=shuffled_samples[start:end],
unique_prompt_count=unique_prompt_count,
)
@@ -0,0 +1,223 @@
# SPDX-License-Identifier: Apache-2.0
"""Configurable diffusion samplers for RL training methods."""
from __future__ import annotations
import copy
from dataclasses import dataclass
from typing import Any, Literal
import torch
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler, )
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.models.base import ModelBase
SchedulerName = Literal["flow_match_euler", "model_default"]
TrajectoryName = Literal["ode", "sde_reflow"]
@dataclass(slots=True)
class SamplingConfig:
"""YAML-backed sampling knobs shared by RL methods."""
num_steps: int = 25
scheduler: SchedulerName = "model_default"
trajectory: TrajectoryName = "ode"
flow_shift: float | None = None
timesteps: list[float] | None = None
sigmas: list[float] | None = None
@classmethod
def from_mapping(cls, raw: dict[str, Any] | None) -> SamplingConfig:
if raw is None:
return cls()
if not isinstance(raw, dict):
raise ValueError(f"method.sampling must be a mapping, got {type(raw).__name__}")
supported_keys = {
"flow_shift",
"num_steps",
"scheduler",
"sigmas",
"timesteps",
"trajectory",
}
unsupported_keys = sorted(set(raw) - supported_keys)
if unsupported_keys:
raise ValueError(f"Unsupported method.sampling key(s): {unsupported_keys}. "
f"Supported keys: {sorted(supported_keys)}")
scheduler = str(raw.get("scheduler", "model_default") or "model_default").strip().lower()
if scheduler not in {"flow_match_euler", "model_default"}:
raise ValueError("method.sampling.scheduler must be one of "
"{flow_match_euler, model_default}, got "
f"{raw.get('scheduler')!r}")
trajectory = str(raw.get("trajectory", "ode") or "ode").strip().lower()
if trajectory not in {"ode", "sde_reflow"}:
raise ValueError("method.sampling.trajectory must be one of "
"{ode, sde_reflow}, got "
f"{raw.get('trajectory')!r}")
timesteps = raw.get("timesteps")
sigmas = raw.get("sigmas")
if timesteps is not None:
if not isinstance(timesteps, list) or not timesteps:
raise ValueError("method.sampling.timesteps must be a non-empty list when set")
timesteps = [float(t) for t in timesteps]
if sigmas is not None:
if not isinstance(sigmas, list) or not sigmas:
raise ValueError("method.sampling.sigmas must be a non-empty list when set")
sigmas = [float(s) for s in sigmas]
if timesteps is not None and sigmas is not None and len(timesteps) != len(sigmas):
raise ValueError("method.sampling.timesteps and method.sampling.sigmas must have the same length")
num_steps = int(raw.get("num_steps", 25) or 25)
if num_steps <= 0:
raise ValueError("method.sampling.num_steps must be positive")
return cls(
num_steps=num_steps,
scheduler=scheduler, # type: ignore[arg-type]
trajectory=trajectory, # type: ignore[arg-type]
flow_shift=(None if raw.get("flow_shift", None) in (None, "inherit") else float(raw["flow_shift"])),
timesteps=timesteps,
sigmas=sigmas,
)
@dataclass(slots=True)
class SamplingResult:
latents: torch.Tensor
timesteps: torch.Tensor
sigmas: torch.Tensor
class DiffusionSampler:
"""Thin model/scheduler sampler used by RL methods.
This intentionally does not call FastVideo's full inference pipelines.
RL training needs a reusable sampling primitive that works with
``ModelBase`` wrappers and scheduler math without binding a method to
model-family pipeline classes such as ``WanDMDPipeline``.
"""
def __init__(self, config: SamplingConfig) -> None:
self.config = config
@torch.no_grad()
def sample(
self,
model: ModelBase,
batch: TrainingBatch,
*,
generator: torch.Generator | None,
) -> SamplingResult:
latents = batch.latents
if latents is None:
raise RuntimeError("TrainingBatch.latents is required for RL sampling")
current = torch.randn(
latents.shape,
device=latents.device,
dtype=latents.dtype,
generator=generator,
)
scheduler = self._prepare_scheduler(model, current.device)
timesteps = scheduler.timesteps.to(device=current.device)
sigmas = scheduler.sigmas.to(device=current.device)
original_timesteps = batch.timesteps
try:
if self.config.trajectory == "ode":
pred_clean = current
for timestep in timesteps:
model_timestep = self._model_timestep(timestep, current)
batch.timesteps = model_timestep
pred_noise = model.predict_noise(
current,
model_timestep,
batch,
conditional=True,
attn_kind="dense",
)
current = scheduler.step(
pred_noise.flatten(0, 1),
timestep,
current.flatten(0, 1),
return_dict=False,
)[0].unflatten(0, pred_noise.shape[:2])
pred_clean = current
return SamplingResult(latents=pred_clean, timesteps=timesteps, sigmas=sigmas)
return SamplingResult(
latents=self._sample_sde_reflow(
model,
batch,
current,
timesteps,
generator=generator,
),
timesteps=timesteps,
sigmas=sigmas,
)
finally:
batch.timesteps = original_timesteps
def _prepare_scheduler(
self,
model: ModelBase,
device: torch.device,
) -> Any:
if self.config.scheduler == "flow_match_euler":
shift = self.config.flow_shift
if shift is None:
shift = float(getattr(model.noise_scheduler, "shift", 1.0))
scheduler = FlowMatchEulerDiscreteScheduler(shift=float(shift))
else:
scheduler = copy.deepcopy(model.noise_scheduler)
kwargs: dict[str, Any] = {"device": device}
if self.config.timesteps is not None:
kwargs["timesteps"] = self.config.timesteps
kwargs["num_inference_steps"] = len(self.config.timesteps)
if self.config.sigmas is not None:
kwargs["sigmas"] = self.config.sigmas
kwargs["num_inference_steps"] = len(self.config.sigmas)
if "num_inference_steps" not in kwargs:
kwargs["num_inference_steps"] = self.config.num_steps
scheduler.set_timesteps(**kwargs)
return scheduler
def _sample_sde_reflow(
self,
model: ModelBase,
batch: TrainingBatch,
current: torch.Tensor,
timesteps: torch.Tensor,
*,
generator: torch.Generator | None,
) -> torch.Tensor:
pred_clean = current
for step_idx, timestep in enumerate(timesteps):
timestep_tensor = self._model_timestep(timestep, current)
batch.timesteps = timestep_tensor
pred_clean = model.predict_x0(
current,
timestep_tensor,
batch,
conditional=True,
attn_kind="dense",
)
if step_idx < len(timesteps) - 1:
next_timestep = timesteps[step_idx + 1].reshape(1).to(device=current.device)
noise = torch.randn(
pred_clean.shape,
device=pred_clean.device,
dtype=pred_clean.dtype,
generator=generator,
)
current = model.add_noise(pred_clean, noise, next_timestep)
return pred_clean
@staticmethod
def _model_timestep(
timestep: torch.Tensor,
current: torch.Tensor,
) -> torch.Tensor:
return timestep.reshape(1).to(device=current.device).expand(current.shape[0]).contiguous()
@@ -0,0 +1,81 @@
# SPDX-License-Identifier: Apache-2.0
"""Shared validation helpers for RL training methods."""
from __future__ import annotations
from dataclasses import dataclass
import math
from typing import Any
import torch
@dataclass(slots=True)
class RLValidationConfig:
every_steps: int = 0
num_steps: int = 40 # Reference DiffusionNFT sampling num steps for best visual quality
num_prompts: int = 16
batch_size: int = 16
log_samples: bool = True
seed: int = 42
data_path: str | None = None
sampling: dict[str, Any] | None = None
@classmethod
def from_mapping(cls, raw: dict[str, Any] | None) -> RLValidationConfig:
if raw is None:
return cls()
if not isinstance(raw, dict):
raise ValueError(f"method.validation must be a mapping, got {type(raw).__name__}")
data_path = raw.get("data_path", None)
sampling = raw.get("sampling", None)
if sampling is not None and not isinstance(sampling, dict):
raise ValueError(f"method.validation.sampling must be a mapping, got {type(sampling).__name__}")
return cls(
every_steps=max(0, int(raw.get("every_steps", 0) or 0)),
num_steps=max(1, int(raw.get("num_steps", 40) or 40)),
num_prompts=max(1, int(raw.get("num_prompts", 16) or 16)),
batch_size=max(1, int(raw.get("batch_size", 16) or 16)),
log_samples=bool(raw.get("log_samples", True)),
seed=int(raw.get("seed", 42) or 42),
data_path=(None if data_path in (None, "") else str(data_path)),
sampling=(dict(sampling) if sampling is not None else None),
)
def validation_shard_indices(
num_prompts: int,
*,
rank: int,
world_size: int,
) -> list[tuple[int, bool]]:
"""Return fixed validation prompt indices for one distributed rank."""
num_prompts = max(1, int(num_prompts))
world_size = max(1, int(world_size))
per_rank = int(math.ceil(num_prompts / world_size))
padded_total = per_rank * world_size
return [((idx % num_prompts), idx < num_prompts) for idx in range(rank, padded_total, world_size)]
def validation_caption(
prompt: str,
rewards: dict[str, float],
) -> str:
reward_parts = [f"{key}: {float(rewards[key]):.4f}" for key in sorted(rewards)]
return f"{' | '.join(reward_parts)} | {prompt[:1000]}"
def media_to_video_array(media: torch.Tensor) -> Any:
"""Convert decoded media to a tracker video array.
Accepts ``[C, T, H, W]`` tensors. ``[C, H, W]`` tensors are treated as
``T=1`` media. Output follows the existing tracker convention used
elsewhere in FastVideo: ``[T, C, H, W]`` uint8.
"""
if media.ndim == 3:
media = media.unsqueeze(1)
if media.ndim != 4:
raise ValueError("media must have shape [C, T, H, W] or [C, H, W], "
f"got {tuple(media.shape)}")
video = (media.detach().float().clamp(0, 1) * 255).round().to(torch.uint8)
return video.permute(1, 0, 2, 3).contiguous().cpu().numpy()
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,796 @@
# SPDX-License-Identifier: Apache-2.0
"""InterleaveThinker-style RL method for critic/prompt-refinement actors."""
from __future__ import annotations
from collections import defaultdict
from collections.abc import Callable, Iterator, Mapping, Sequence
from typing import Any, Protocol, TypeGuard
import torch
import torch.distributed as dist
from fastvideo.logger import init_logger
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.methods.rl.rewards import (
InterleaveThinkerRewardScorer, )
from fastvideo.train.models.base import RoleModelBase
from fastvideo.train.utils.config import (
get_optional_float,
get_optional_int,
parse_betas,
)
from fastvideo.train.utils.instantiate import (
instantiate, )
from fastvideo.train.utils.optimizer import (
build_optimizer_and_scheduler, )
logger = init_logger(__name__)
class InterleaveTrainActor(Protocol):
"""Role model contract required for Interleave actor policy updates."""
transformer: torch.nn.Module
_trainable: bool
@property
def device(self) -> torch.device:
...
def init_preprocessors(self, training_config: Any) -> None:
...
def train_interleave_rollouts(self, **kwargs: Any) -> Any:
...
class InterleaveGenerationActor(Protocol):
"""Optional online rollout-generation hook for Interleave actors."""
def generate_interleave_responses(self, batch: dict[str, Any], **kwargs: Any) -> Any:
...
class InterleaveReferenceActor(Protocol):
"""Optional frozen reference-policy logprob hook for Interleave actors."""
def reference_logprobs_for_interleave_rollouts(
self,
rollouts: Sequence[Mapping[str, Any]],
) -> Sequence[Any]:
...
class InterleaveRewardScorer(Protocol):
"""Reward scorer contract consumed by the managed Interleave step."""
def as_tensors(
self,
reward_inputs: Sequence[Mapping[str, Any]],
*,
device: torch.device | str = "cpu",
) -> Mapping[str, Any]:
...
class _CallableRewardScorer:
"""Adapt list-returning reward helpers to the tensor scorer protocol."""
def __init__(
self,
scorer: Callable[[Sequence[Mapping[str, Any]]], Any],
) -> None:
self._scorer = scorer
def as_tensors(
self,
reward_inputs: Sequence[Mapping[str, Any]],
*,
device: torch.device | str = "cpu",
) -> Mapping[str, torch.Tensor]:
return _coerce_reward_output(
self._scorer(reward_inputs),
expected_count=len(reward_inputs),
device=device,
)
class InterleaveThinkerRLMethod(TrainingMethod):
"""GRPO-style InterleaveThinker critic RL.
One ``Trainer`` step performs a complete InterleaveThinker RL outer step:
actor rollout generation, reward scoring, group-normalized advantage
computation, and an actor-owned policy update. The method intentionally
delegates tokenizer/VLM/logprob details to the student model wrapper via two
hooks:
- ``generate_interleave_responses(batch, **kwargs)`` returns rollout dicts.
- ``train_interleave_rollouts(rollouts=..., advantages=..., **kwargs)``
performs the policy update and returns loss/metric dictionaries.
"""
def __init__(
self,
*,
cfg: Any,
role_models: Mapping[str, RoleModelBase],
) -> None:
super().__init__(cfg=cfg, role_models=role_models)
student = self.student
if not student._trainable:
raise ValueError("InterleaveThinkerRLMethod requires a trainable student")
if not _is_interleave_train_actor(student):
raise TypeError("InterleaveThinkerRLMethod requires an Interleave train actor implementing "
"train_interleave_rollouts()")
self._interleave_student = student
self.reference = role_models.get("reference")
self._interleave_reference: InterleaveReferenceActor | None = None
if self.reference is not None:
if self.reference._trainable:
raise ValueError("InterleaveThinkerRLMethod requires models.reference.trainable=false")
if _is_interleave_reference_actor(self.reference):
self._interleave_reference = self.reference
self._freeze_reference_model()
self._interleave_student.init_preprocessors(self.training_config)
self._num_generations = self._read_int("num_generations", 8)
self._num_batches_per_step = self._read_int("num_batches_per_step", 1)
self._max_new_tokens = get_optional_int(
self.method_config,
"max_new_tokens",
where="method.max_new_tokens",
)
self._temperature = self._read_float("temperature", 1.0)
self._top_p = self._read_float("top_p", 1.0)
self._advantage_eps = self._read_float("advantage_eps", 1.0e-4)
self._advantage_clip = get_optional_float(
self.method_config,
"advantage_clip",
where="method.advantage_clip",
)
self._clip_range = self._read_float("clip_range", 0.2)
if self._clip_range < 0.0:
raise ValueError("method.clip_range must be non-negative")
self._kl_coef = self._read_float("kl_coef", 0.0)
if self._kl_coef < 0.0:
raise ValueError("method.kl_coef must be non-negative")
self._update_micro_batch_size = self._read_optional_int_alias(
"micro_batch_size_per_device_for_update",
"update_micro_batch_size",
)
self._max_grad_norm = self._read_float("max_grad_norm", 0.0)
self._terminal_progress = bool(self.method_config.get("terminal_progress", True))
self._reward_scorer = self._build_reward_scorer()
self._student_optimizer: torch.optim.Optimizer | None = None
self._student_lr_scheduler: Any | None = None
self._init_optimizer_and_scheduler()
@property
def _optimizer_dict(self) -> dict[str, Any]:
return {"student": self._student_optimizer} if self._student_optimizer is not None else {}
@property
def _lr_scheduler_dict(self) -> dict[str, Any]:
return {"student": self._student_lr_scheduler} if self._student_lr_scheduler is not None else {}
def manages_optimization(self) -> bool:
return True
def single_train_step(
self,
batch: dict[str, Any],
iteration: int,
) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, LogScalar]]:
del batch, iteration
raise RuntimeError("InterleaveThinkerRLMethod uses managed_train_step()")
def managed_train_step(
self,
data_stream: Iterator[dict[str, Any]],
iteration: int,
) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, LogScalar]]:
self._log_progress(f"[InterleaveThinkerRL] step {iteration}: generating rollouts")
rollouts: list[dict[str, Any]] = []
sample_offset = 0
for batch_idx in range(self._num_batches_per_step):
batch = next(data_stream)
batch_rollouts = self._generate_rollouts(batch, iteration=iteration, batch_idx=batch_idx)
sample_offset = self._namespace_rollout_keys(
batch_rollouts,
batch_idx=batch_idx,
sample_offset=sample_offset,
)
rollouts.extend(batch_rollouts)
if not rollouts:
raise RuntimeError("InterleaveThinkerRLMethod generated no rollouts")
reference_logprob_count = self._attach_reference_logprobs(rollouts)
self._log_progress(f"[InterleaveThinkerRL] step {iteration}: scoring {len(rollouts)} rollouts")
reward_inputs = self._make_reward_inputs(rollouts)
reward_tensors = _coerce_reward_output(
self._reward_scorer.as_tensors(reward_inputs, device=self._interleave_student.device),
expected_count=len(reward_inputs),
device=self._interleave_student.device,
)
self._log_progress(f"[InterleaveThinkerRL] step {iteration}: computing advantages")
group_keys = [self._rollout_group_key(rollout, idx) for idx, rollout in enumerate(rollouts)]
advantages = self._compute_group_advantages(
rewards=reward_tensors["overall"],
group_keys=group_keys,
eps=self._advantage_eps,
clip=self._advantage_clip,
)
self._log_progress(f"[InterleaveThinkerRL] step {iteration}: actor update")
loss_map, train_metrics = self._train_actor(
rollouts,
advantages=advantages,
rewards=reward_tensors,
iteration=iteration,
)
metrics: dict[str, LogScalar] = {}
metrics.update(self._reward_metrics(reward_tensors))
metrics.update(self._advantage_metrics(advantages, group_keys))
metrics.update(train_metrics)
metrics["interleave/num_rollouts"] = float(len(rollouts))
metrics["interleave/num_groups"] = float(len(set(group_keys)))
metrics["interleave/num_batches_per_step"] = float(self._num_batches_per_step)
metrics["interleave/reference_logprob_rollouts"] = float(reference_logprob_count)
return loss_map, {}, metrics
def get_optimizers(
self,
iteration: int,
) -> list[torch.optim.Optimizer]:
del iteration
return [self._student_optimizer] if self._student_optimizer is not None else []
def get_lr_schedulers(
self,
iteration: int,
) -> list[Any]:
del iteration
return [self._student_lr_scheduler] if self._student_lr_scheduler is not None else []
def get_grad_clip_targets(
self,
iteration: int,
) -> dict[str, torch.nn.Module]:
del iteration
transformer = self._interleave_student.transformer
if isinstance(transformer, torch.nn.Module):
return {"student": transformer}
return {}
def on_train_start(self) -> None:
super().on_train_start()
self._freeze_reference_model()
def _init_optimizer_and_scheduler(self) -> None:
transformer = self._interleave_student.transformer
if not isinstance(transformer, torch.nn.Module):
return
params = [p for p in transformer.parameters() if p.requires_grad]
if not params:
return
betas = self.training_config.optimizer.betas
betas_raw = self.method_config.get("betas", None)
if betas_raw is not None:
betas = parse_betas(betas_raw, where="method.betas")
self._student_optimizer, self._student_lr_scheduler = build_optimizer_and_scheduler(
params=params,
optimizer_config=self.training_config.optimizer,
loop_config=self.training_config.loop,
learning_rate=float(self.training_config.optimizer.learning_rate),
betas=betas,
scheduler_name=str(self.training_config.optimizer.lr_scheduler),
)
def _build_edit_scorer(
self,
raw: Any,
) -> Any:
if raw is None:
return None
if callable(raw):
return raw
if isinstance(raw, Mapping):
return instantiate(dict(raw))
raise TypeError("method.edit_scorer must be a callable or a mapping with _target_")
def _build_reward_scorer(self) -> InterleaveRewardScorer:
raw = self.method_config.get("reward_scorer")
if raw is None:
scorer: Any = InterleaveThinkerRewardScorer(
format_weight=self._read_float("format_weight", 0.5),
judge_accuracy_weight=self._read_float("judge_accuracy_weight", 0.2),
semantic_weight=self._read_float("semantic_weight", 0.6),
quality_weight=self._read_float("quality_weight", 0.2),
fallback_edit_reward=self._read_float("fallback_edit_reward", 0.5),
edit_scorer=self._build_edit_scorer(self.method_config.get("edit_scorer")),
)
elif isinstance(raw, Mapping):
scorer = instantiate(dict(raw))
else:
scorer = raw
if _is_interleave_reward_scorer(scorer):
return scorer
if callable(scorer):
return _CallableRewardScorer(scorer)
raise TypeError("method.reward_scorer must implement as_tensors(), be a callable returning per-rollout "
"rewards, or be a mapping with _target_ that constructs one")
def _freeze_reference_model(self) -> None:
if self.reference is None:
return
transformer = getattr(self.reference, "transformer", None)
if isinstance(transformer, torch.nn.Module):
transformer.requires_grad_(False)
transformer.eval()
def _attach_reference_logprobs(
self,
rollouts: Sequence[Mapping[str, Any]],
) -> int:
if self.reference is None:
return 0
pending: list[dict[str, Any]] = []
for rollout in rollouts:
if self._has_reference_logprobs(rollout):
continue
if not isinstance(rollout, dict):
raise TypeError("InterleaveThinker reference logprobs require mutable rollout dictionaries")
pending.append(rollout)
if not pending:
return 0
if self._interleave_reference is None:
raise RuntimeError("models.reference must implement reference_logprobs_for_interleave_rollouts()")
self._log_progress(f"[InterleaveThinkerRL] computing reference logprobs for {len(pending)} rollouts")
self._freeze_reference_model()
with torch.no_grad():
rows = self._interleave_reference.reference_logprobs_for_interleave_rollouts(pending)
if not isinstance(rows, Sequence) or isinstance(rows, str | bytes):
raise TypeError("reference_logprobs_for_interleave_rollouts() must return a sequence")
reference_rows = list(rows)
if len(reference_rows) != len(pending):
raise ValueError("reference logprob row count must match rollout count")
for rollout, row in zip(pending, reference_rows, strict=True):
rollout["reference_logprobs"] = self._coerce_reference_logprobs(row)
return len(pending)
@staticmethod
def _has_reference_logprobs(rollout: Mapping[str, Any]) -> bool:
return rollout.get("reference_logprobs") is not None or rollout.get("ref_logprobs") is not None
@staticmethod
def _coerce_reference_logprobs(row: Any) -> list[float]:
if torch.is_tensor(row):
values = row.detach().cpu().float().flatten().tolist()
elif isinstance(row, Sequence) and not isinstance(row, str | bytes):
values = [float(value) for value in row]
else:
raise TypeError("reference logprob rows must be tensors or sequences of floats")
if not values:
raise ValueError("reference logprob rows must be non-empty")
return values
def _generate_rollouts(
self,
batch: dict[str, Any],
*,
iteration: int,
batch_idx: int,
) -> list[dict[str, Any]]:
if _is_interleave_generation_actor(self._interleave_student):
generated = self._interleave_student.generate_interleave_responses(
batch,
num_generations=self._num_generations,
temperature=self._temperature,
top_p=self._top_p,
max_new_tokens=self._max_new_tokens,
generator=self.cuda_generator,
iteration=iteration,
batch_idx=batch_idx,
)
return self._normalize_generated_rollouts(generated, batch=batch)
return self._offline_rollouts_from_batch(batch)
def _offline_rollouts_from_batch(
self,
batch: Mapping[str, Any],
) -> list[dict[str, Any]]:
items = self._batch_to_items(batch)
rollouts: list[dict[str, Any]] = []
for item_idx, item in enumerate(items):
responses = item["responses"] if "responses" in item else item.get("response")
if responses is None:
raise RuntimeError("Student model must implement generate_interleave_responses(), "
"or batches must contain response/responses for offline rollouts")
if isinstance(responses, str):
response_list = [responses]
elif isinstance(responses, Sequence):
response_list = list(responses)
else:
response_list = [responses]
for response_idx, response in enumerate(response_list):
rollout = dict(item)
for per_response_key in ("edit_score", "edit_scores"):
per_response_value = rollout.get(per_response_key)
if (isinstance(per_response_value, Sequence) and not isinstance(per_response_value, str)
and len(per_response_value) == len(response_list)):
rollout[per_response_key] = per_response_value[response_idx]
rollout.pop("responses", None)
rollout["response"] = str(response)
rollout.setdefault("sample_index", item_idx)
rollout.setdefault("generation_index", response_idx)
rollout.setdefault("group_key", self._rollout_group_key(rollout, item_idx))
rollouts.append(rollout)
return rollouts
def _normalize_generated_rollouts(
self,
generated: Any,
*,
batch: Mapping[str, Any],
) -> list[dict[str, Any]]:
if isinstance(generated, Mapping) and "rollouts" in generated:
generated = generated["rollouts"]
if isinstance(generated, str):
generated = [generated]
if not isinstance(generated, Sequence):
raise TypeError("generate_interleave_responses() must return a sequence of rollout mappings")
batch_items = self._batch_to_items(batch)
rollouts: list[dict[str, Any]] = []
for idx, raw_rollout in enumerate(generated):
if isinstance(raw_rollout, str):
base = dict(batch_items[min(idx // max(1, self._num_generations), len(batch_items) - 1)])
rollout = {**base, "response": raw_rollout}
elif isinstance(raw_rollout, Mapping):
rollout = dict(raw_rollout)
else:
raise TypeError(f"Rollout must be a mapping or string, got {type(raw_rollout).__name__}")
rollout.setdefault("sample_index", idx // max(1, self._num_generations))
rollout.setdefault("generation_index", idx % max(1, self._num_generations))
if "group_key" not in rollout:
sample_index = int(rollout["sample_index"])
if 0 <= sample_index < len(batch_items):
base_item = batch_items[sample_index]
for key, value in base_item.items():
rollout.setdefault(key, value)
rollout["group_key"] = self._rollout_group_key(rollout, idx)
rollouts.append(rollout)
return rollouts
def _make_reward_inputs(
self,
rollouts: Sequence[Mapping[str, Any]],
) -> list[dict[str, Any]]:
reward_inputs: list[dict[str, Any]] = []
for rollout in rollouts:
if "response" not in rollout:
raise ValueError("Each InterleaveThinker rollout must include 'response'")
item = dict(rollout)
item["response"] = str(item["response"])
reward_inputs.append(item)
return reward_inputs
def _train_actor(
self,
rollouts: Sequence[Mapping[str, Any]],
*,
advantages: torch.Tensor,
rewards: Mapping[str, torch.Tensor],
iteration: int,
) -> tuple[dict[str, torch.Tensor], dict[str, LogScalar]]:
result = self._interleave_student.train_interleave_rollouts(
rollouts=rollouts,
advantages=advantages,
rewards=rewards,
iteration=iteration,
optimizer=self._student_optimizer,
lr_scheduler=self._student_lr_scheduler,
gradient_accumulation_steps=max(1, int(self.training_config.loop.gradient_accumulation_steps or 1)),
clip_range=self._clip_range,
kl_coef=self._kl_coef,
update_micro_batch_size=self._update_micro_batch_size,
max_grad_norm=self._max_grad_norm,
)
return self._coerce_train_result(result)
def _coerce_train_result(
self,
result: Any,
) -> tuple[dict[str, torch.Tensor], dict[str, LogScalar]]:
if isinstance(result, tuple) and len(result) == 2:
loss_map_raw, metrics_raw = result
elif isinstance(result, Mapping):
loss_map_raw = result.get("loss_map")
metrics_raw = result.get("metrics", {})
if loss_map_raw is None and "loss" in result:
loss_map_raw = {"total_loss": result["loss"]}
else:
raise TypeError("train_interleave_rollouts() must return (loss_map, metrics) or a mapping")
if not isinstance(loss_map_raw, Mapping):
raise TypeError("train_interleave_rollouts() result must include a loss_map mapping")
if not isinstance(metrics_raw, Mapping):
raise TypeError("train_interleave_rollouts() metrics must be a mapping")
loss_map = {str(k): self._coerce_tensor(v) for k, v in loss_map_raw.items()}
if "total_loss" not in loss_map:
if len(loss_map) != 1:
raise ValueError("loss_map must include total_loss when multiple losses are returned")
only_value = next(iter(loss_map.values()))
loss_map["total_loss"] = only_value
metrics = {str(k): self._coerce_log_scalar(v) for k, v in metrics_raw.items()}
return loss_map, metrics
def _coerce_tensor(self, value: Any) -> torch.Tensor:
if torch.is_tensor(value):
return value
return torch.tensor(float(value), device=self._interleave_student.device, dtype=torch.float32)
def _coerce_log_scalar(self, value: Any) -> LogScalar:
if torch.is_tensor(value):
if value.numel() != 1:
raise ValueError(f"Expected scalar metric tensor, got shape={tuple(value.shape)}")
return value.detach()
if isinstance(value, float | int):
return float(value)
raise TypeError(f"Expected scalar metric, got {type(value).__name__}")
@staticmethod
def _batch_to_items(batch: Mapping[str, Any], ) -> list[dict[str, Any]]:
if "items" in batch and isinstance(batch["items"], Sequence):
return [dict(item) for item in batch["items"]]
batch_size = 1
for value in batch.values():
if isinstance(value, list | tuple):
batch_size = len(value)
break
if torch.is_tensor(value) and value.ndim > 0:
batch_size = int(value.shape[0])
break
items: list[dict[str, Any]] = []
for idx in range(batch_size):
item: dict[str, Any] = {}
for key, value in batch.items():
if isinstance(value, list | tuple) and len(value) == batch_size or torch.is_tensor(
value) and value.ndim > 0 and int(value.shape[0]) == batch_size:
item[key] = value[idx]
else:
item[key] = value
items.append(item)
return items
def _namespace_rollout_keys(
self,
rollouts: Sequence[dict[str, Any]],
*,
batch_idx: int,
sample_offset: int,
) -> int:
"""Keep separately sampled input batches out of the same GRPO group."""
if self._num_batches_per_step <= 1:
return sample_offset
local_samples: dict[str, int] = {}
for rollout_idx, rollout in enumerate(rollouts):
local_group_key = self._rollout_group_key(rollout, rollout_idx)
local_sample_index = rollout.get("sample_index", rollout_idx)
local_sample_key = f"{type(local_sample_index).__name__}:{local_sample_index!r}"
if local_sample_key not in local_samples:
local_samples[local_sample_key] = sample_offset + len(local_samples)
rollout.setdefault("local_group_key", local_group_key)
rollout.setdefault("local_sample_index", local_sample_index)
rollout["outer_batch_index"] = batch_idx
rollout["sample_index"] = local_samples[local_sample_key]
rollout["group_key"] = f"batch:{batch_idx}:{local_group_key}"
return sample_offset + len(local_samples)
@staticmethod
def _rollout_group_key(
rollout: Mapping[str, Any],
index: int,
) -> str:
for key in ("group_key", "problem_id", "sample_index", "origin_prompt", "prompt"):
value = rollout.get(key)
if value is not None:
return str(value)
return str(index)
@staticmethod
def _compute_group_advantages(
*,
rewards: torch.Tensor,
group_keys: Sequence[str],
eps: float,
clip: float | None,
) -> torch.Tensor:
if rewards.ndim != 1:
raise ValueError(f"rewards must have shape [N], got {tuple(rewards.shape)}")
if int(rewards.shape[0]) != len(group_keys):
raise ValueError("reward count must match group key count")
advantages = torch.empty_like(rewards, dtype=torch.float32)
groups: dict[str, list[int]] = defaultdict(list)
for idx, group_key in enumerate(group_keys):
groups[str(group_key)].append(idx)
for indices in groups.values():
index_tensor = torch.tensor(indices, device=rewards.device, dtype=torch.long)
group_rewards = rewards[index_tensor].detach().float()
group_std = group_rewards.std(unbiased=False)
advantages[index_tensor] = (group_rewards - group_rewards.mean()) / (group_std + float(eps))
if clip is not None:
advantages = advantages.clamp(-float(clip), float(clip))
return advantages
def _reward_metrics(
self,
rewards: Mapping[str, torch.Tensor],
) -> dict[str, LogScalar]:
metrics: dict[str, LogScalar] = {}
for key, value in rewards.items():
if value.numel() > 0:
metrics[f"interleave/reward/{key}"] = value.detach().float().mean()
return metrics
def _advantage_metrics(
self,
advantages: torch.Tensor,
group_keys: Sequence[str],
) -> dict[str, LogScalar]:
group_sizes: dict[str, int] = defaultdict(int)
for key in group_keys:
group_sizes[key] += 1
group_size_values = torch.tensor(list(group_sizes.values()), device=advantages.device, dtype=torch.float32)
return {
"interleave/advantage_mean": advantages.detach().float().mean(),
"interleave/advantage_std": advantages.detach().float().std(unbiased=False),
"interleave/group_size_mean": group_size_values.mean(),
}
def _read_int(
self,
key: str,
default: int,
) -> int:
value = get_optional_int(self.method_config, key, where=f"method.{key}")
if value is None:
value = default
if value <= 0:
raise ValueError(f"method.{key} must be a positive integer")
return int(value)
def _read_optional_int_alias(
self,
*keys: str,
) -> int | None:
for key in keys:
value = get_optional_int(self.method_config, key, where=f"method.{key}")
if value is None:
continue
if value <= 0:
raise ValueError(f"method.{key} must be a positive integer")
return int(value)
return None
def _read_float(
self,
key: str,
default: float,
) -> float:
value = get_optional_float(self.method_config, key, where=f"method.{key}")
if value is None:
value = default
return float(value)
def _log_progress(self, message: str) -> None:
if not self._terminal_progress:
return
rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else 0
if rank == 0:
logger.info(message)
def _coerce_reward_output(
raw: Any,
*,
expected_count: int,
device: torch.device | str,
) -> dict[str, torch.Tensor]:
"""Normalize supported reward return shapes to one tensor per metric."""
if isinstance(raw, Mapping):
columns: Mapping[Any, Any] = raw
elif torch.is_tensor(raw):
columns = {"overall": raw}
elif isinstance(raw, Sequence) and not isinstance(raw, str | bytes):
rows = list(raw)
row_mappings = [_reward_row_mapping(row) for row in rows]
if any(row is not None for row in row_mappings):
if not all(row is not None for row in row_mappings):
raise TypeError("reward scorer returned a mixture of mapping and scalar rows")
normalized_rows = [row for row in row_mappings if row is not None]
if not normalized_rows:
columns = {"overall": []}
else:
metric_keys = tuple(normalized_rows[0].keys())
expected_keys = set(metric_keys)
for row in normalized_rows[1:]:
if set(row.keys()) != expected_keys:
raise ValueError("reward scorer mapping rows must contain the same metric keys")
columns = {key: [row[key] for row in normalized_rows] for key in metric_keys}
else:
columns = {"overall": rows}
else:
raise TypeError("reward scorer must return a metric mapping, a reward tensor, or one result per rollout")
tensors: dict[str, torch.Tensor] = {}
for raw_key, value in columns.items():
key = str(raw_key)
try:
if torch.is_tensor(value):
tensor = value.detach().to(device=device, dtype=torch.float32)
else:
tensor = torch.as_tensor(value, device=device, dtype=torch.float32)
except (TypeError, ValueError, RuntimeError) as exc:
raise TypeError(f"reward metric {key!r} must contain numeric values") from exc
if tensor.ndim == 0:
tensor = tensor.reshape(1)
if tensor.ndim != 1:
raise ValueError(f"reward metric {key!r} must have shape [N], got {tuple(tensor.shape)}")
if tensor.numel() != expected_count:
raise ValueError(f"reward metric {key!r} has {tensor.numel()} values for {expected_count} rollouts")
tensors[key] = tensor
if "overall" not in tensors:
raise ValueError("reward scorer output must include an 'overall' metric")
tensors.setdefault("avg", tensors["overall"])
return tensors
def _reward_row_mapping(row: Any) -> Mapping[Any, Any] | None:
if isinstance(row, Mapping):
return row
as_dict = getattr(row, "as_dict", None)
if not callable(as_dict):
return None
result = as_dict()
if not isinstance(result, Mapping):
raise TypeError("reward result as_dict() must return a mapping")
return result
def _is_interleave_train_actor(model: Any) -> TypeGuard[InterleaveTrainActor]:
return (isinstance(getattr(model, "transformer", None), torch.nn.Module)
and callable(getattr(model, "init_preprocessors", None))
and callable(getattr(model, "train_interleave_rollouts", None)))
def _is_interleave_generation_actor(model: Any) -> TypeGuard[InterleaveGenerationActor]:
return callable(getattr(model, "generate_interleave_responses", None))
def _is_interleave_reference_actor(model: Any) -> TypeGuard[InterleaveReferenceActor]:
return callable(getattr(model, "reference_logprobs_for_interleave_rollouts", None))
def _is_interleave_reward_scorer(scorer: Any) -> TypeGuard[InterleaveRewardScorer]:
return callable(getattr(scorer, "as_tensors", None))
__all__ = [
"InterleaveGenerationActor",
"InterleaveReferenceActor",
"InterleaveRewardScorer",
"InterleaveThinkerRLMethod",
"InterleaveTrainActor",
]
@@ -0,0 +1,132 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable reward models for training methods."""
from __future__ import annotations
from typing import TYPE_CHECKING
from fastvideo.train.methods.rl.rewards.interleave_thinker import (
EditScoreProvider,
InterleavePlannerRewardResult,
InterleavePlannerRewardScorer,
InterleaveThinkerAnswer,
InterleaveThinkerEditRequest,
InterleaveThinkerEditScore,
InterleaveThinkerRewardResult,
InterleaveThinkerRewardScorer,
extract_interleave_answer,
extract_interleave_plan_payload,
interleave_format_reward,
interleave_planner_format_reward,
interleave_judge_accuracy_reward,
normalize_interleave_response,
score_interleave_planner_rewards,
score_interleave_thinker_rewards,
)
if TYPE_CHECKING:
from fastvideo.train.methods.rl.rewards.frame_rewards import (
ClipScoreScorer,
PickScoreScorer,
)
from fastvideo.train.methods.rl.rewards.interleave_api import (
ConstantInterleaveEditScorer,
GeminiInterleaveImageScorer,
GeminiNanoBananaEditScorer,
)
from fastvideo.train.methods.rl.rewards.media import (
MultiRewardScorer,
RewardScorer,
)
def build_multi_reward_scorer(
reward_weights,
*,
device="cuda",
scorers: dict[str, RewardScorer] | None = None,
) -> MultiRewardScorer:
from fastvideo.train.methods.rl.rewards.frame_rewards import (
ClipScoreScorer,
PickScoreScorer,
)
from fastvideo.train.methods.rl.rewards.media import MultiRewardScorer
available: dict[str, RewardScorer] = dict(scorers or {})
if not available:
available = {
"pickscore": PickScoreScorer(device=device),
"clipscore": ClipScoreScorer(device=device),
}
return MultiRewardScorer(reward_weights, scorers=available)
def __getattr__(name: str) -> object:
if name in {"ClipScoreScorer", "PickScoreScorer"}:
from fastvideo.train.methods.rl.rewards.frame_rewards import (
ClipScoreScorer,
PickScoreScorer,
)
return {
"ClipScoreScorer": ClipScoreScorer,
"PickScoreScorer": PickScoreScorer,
}[name]
if name in {
"ConstantInterleaveEditScorer",
"GeminiInterleaveImageScorer",
"GeminiNanoBananaEditScorer",
}:
from fastvideo.train.methods.rl.rewards.interleave_api import (
ConstantInterleaveEditScorer,
GeminiInterleaveImageScorer,
GeminiNanoBananaEditScorer,
)
return {
"ConstantInterleaveEditScorer": ConstantInterleaveEditScorer,
"GeminiInterleaveImageScorer": GeminiInterleaveImageScorer,
"GeminiNanoBananaEditScorer": GeminiNanoBananaEditScorer,
}[name]
if name in {"MultiRewardScorer", "RewardScorer", "select_first_frame"}:
from fastvideo.train.methods.rl.rewards.media import (
MultiRewardScorer,
RewardScorer,
select_first_frame,
)
return {
"MultiRewardScorer": MultiRewardScorer,
"RewardScorer": RewardScorer,
"select_first_frame": select_first_frame,
}[name]
raise AttributeError(name)
__all__ = [
"ClipScoreScorer",
"ConstantInterleaveEditScorer",
"EditScoreProvider",
"GeminiInterleaveImageScorer",
"GeminiNanoBananaEditScorer",
"InterleavePlannerRewardResult",
"InterleavePlannerRewardScorer",
"InterleaveThinkerAnswer",
"InterleaveThinkerEditRequest",
"InterleaveThinkerEditScore",
"InterleaveThinkerRewardResult",
"InterleaveThinkerRewardScorer",
"MultiRewardScorer",
"PickScoreScorer",
"RewardScorer",
"build_multi_reward_scorer",
"extract_interleave_answer",
"extract_interleave_plan_payload",
"interleave_format_reward",
"interleave_planner_format_reward",
"interleave_judge_accuracy_reward",
"normalize_interleave_response",
"score_interleave_planner_rewards",
"score_interleave_thinker_rewards",
"select_first_frame",
]
@@ -0,0 +1,130 @@
# SPDX-License-Identifier: Apache-2.0
"""Frame-based reward scorers used by RL training methods."""
from __future__ import annotations
from collections.abc import Mapping, Sequence
from typing import Any
from PIL import Image
import torch
from fastvideo.train.methods.rl.rewards.media import select_first_frame
class PickScoreScorer(torch.nn.Module):
"""PickScore reward, matching DiffusionNFT normalization.
Ported from DiffusionNFT's ``flow_grpo/pickscore_scorer.py``.
"""
def __init__(
self,
*,
device: torch.device | str = "cuda",
dtype: torch.dtype = torch.float32,
) -> None:
super().__init__()
from transformers import AutoModel, AutoProcessor
processor_path = "laion/CLIP-ViT-H-14-laion2B-s32B-b79K"
model_path = "yuvalkirstain/PickScore_v1"
self.device = torch.device(device)
self.dtype = dtype
self.processor = AutoProcessor.from_pretrained(processor_path)
self.model = AutoModel.from_pretrained(model_path).eval().to(self.device)
self.model = self.model.to(dtype=dtype)
@torch.no_grad()
def forward(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
frame_tensor = select_first_frame(media)
frame_np = (frame_tensor.detach().float().clamp(0, 1) * 255).round()
frame_np = frame_np.to(torch.uint8).cpu().numpy().transpose(0, 2, 3, 1)
pil_frames = [Image.fromarray(frame) for frame in frame_np]
frame_inputs = self.processor(
images=pil_frames,
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
)
frame_inputs = {k: v.to(device=self.device) for k, v in frame_inputs.items()}
text_inputs = self.processor(
text=list(prompts),
padding=True,
truncation=True,
max_length=77,
return_tensors="pt",
)
text_inputs = {k: v.to(device=self.device) for k, v in text_inputs.items()}
text_embs = self.model.get_text_features(**text_inputs)
text_embs = text_embs / text_embs.norm(p=2, dim=-1, keepdim=True)
frame_embs = self.model.get_image_features(**frame_inputs)
frame_embs = frame_embs / frame_embs.norm(p=2, dim=-1, keepdim=True)
scores = self.model.logit_scale.exp() * (text_embs @ frame_embs.T)
return scores.diag().float() / 26.0
class ClipScoreScorer(torch.nn.Module):
"""CLIPScore reward, matching DiffusionNFT normalization.
Ported from DiffusionNFT's ``flow_grpo/clip_scorer.py``.
"""
def __init__(
self,
*,
device: torch.device | str = "cuda",
) -> None:
super().__init__()
import torch.nn as nn
import torchvision.transforms as T
from transformers import CLIPModel, CLIPProcessor
def get_size(size: Any) -> Any:
if isinstance(size, int):
return (size, size)
if isinstance(size, Mapping) and "height" in size and "width" in size:
return (size["height"], size["width"])
if isinstance(size, Mapping) and "shortest_edge" in size:
return size["shortest_edge"]
raise ValueError(f"Invalid processor size: {size!r}")
def get_frame_transform(processor: Any) -> torch.nn.Module:
config = processor.to_dict()
resize = T.Resize(get_size(config.get("size"))) if config.get("do_resize") else nn.Identity()
crop = T.CenterCrop(get_size(config.get("crop_size"))) if config.get("do_center_crop") else nn.Identity()
normalize = (T.Normalize(mean=processor.image_mean, std=processor.image_std)
if config.get("do_normalize") else nn.Identity())
return T.Compose([resize, crop, normalize])
self.device = torch.device(device)
self.model = CLIPModel.from_pretrained("openai/clip-vit-large-patch14").to(self.device).eval()
self.processor = CLIPProcessor.from_pretrained("openai/clip-vit-large-patch14")
self.transform = get_frame_transform(self.processor.image_processor)
@torch.no_grad()
def forward(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> torch.Tensor:
frame_tensor = select_first_frame(media).detach().float().clamp(0, 1)
texts = self.processor(
text=list(prompts),
padding="max_length",
truncation=True,
return_tensors="pt",
).to(self.device)
pixels = self.transform(frame_tensor).to(device=self.device, dtype=frame_tensor.dtype)
outputs = self.model(pixel_values=pixels, **texts)
return outputs.logits_per_image.diagonal().float() / 100.0
@@ -0,0 +1,343 @@
# SPDX-License-Identifier: Apache-2.0
"""API-backed InterleaveThinker reward services.
These classes wrap the closed-source models used by InterleaveThinker without
making them mandatory FastVideo dependencies:
- Nano Banana / Gemini native-image models generate or edit images.
- Gemini VLM models score semantic alignment and perceptual quality.
The reusable reward parser in ``interleave_thinker.py`` remains pure; this file
contains the optional networked scorer that can be configured from YAML.
"""
from __future__ import annotations
import base64
import io
import json
import os
from pathlib import Path
import time
import uuid
from typing import Any
from fastvideo.workflow.interleave_thinker.generator import (
NanoBananaImageGeneratorBackend,
encode_file_to_base64,
)
from fastvideo.workflow.interleave_thinker.schema import (
InterleaveEditRequest as InterleaveAPIEditRequest, )
from fastvideo.train.methods.rl.rewards.interleave_thinker import (
InterleaveThinkerEditRequest,
InterleaveThinkerEditScore,
)
_SCORE_SCHEMA = {
"type": "object",
"properties": {
"semantic_score": {
"type": "number"
},
"quality_score": {
"type": "number"
},
"semantic_analysis": {
"type": "string"
},
"quality_analysis": {
"type": "string"
},
},
"required": ["semantic_score", "quality_score"],
}
_SCORE_SYSTEM_PROMPT = ("You are a strict image editing evaluator. Return only JSON. Scores must be "
"numbers from 0 to 10, where 10 is best.")
_SCORE_USER_PROMPT = """\
Evaluate the edited image for an iterative image generation/editing step.
Instruction:
{instruction}
Return:
- semantic_score: how well the edited image satisfies the instruction and preserves required content.
- quality_score: visual quality, realism/coherence, artifact absence, and logical consistency.
Use the full 0-10 range. Do not reward unrelated changes.
"""
_IMAGE_MIME_TYPES = {
".bmp": "image/bmp",
".gif": "image/gif",
".jpeg": "image/jpeg",
".jpg": "image/jpeg",
".png": "image/png",
".tif": "image/tiff",
".tiff": "image/tiff",
".webp": "image/webp",
}
class GeminiInterleaveImageScorer:
"""Gemini VLM scorer returning InterleaveThinker semantic/quality scores."""
def __init__(
self,
*,
model: str = "gemini-2.5-pro",
api_key: str | None = None,
base_url: str | None = None,
max_attempts: int = 6,
retry_delay_s: float = 2.0,
temperature: float = 0.2,
) -> None:
self.model = model
self.api_key = api_key
self.base_url = base_url
self.max_attempts = max(1, int(max_attempts))
self.retry_delay_s = float(retry_delay_s)
self.temperature = float(temperature)
self._client: Any | None = None
def score_images(
self,
*,
edited_image_path: str,
instruction: str,
reference_image_path: str | None = None,
) -> InterleaveThinkerEditScore:
prompt = _SCORE_USER_PROMPT.format(instruction=(instruction or "").strip())
contents: list[Any] = [prompt]
if reference_image_path:
contents.extend(["Reference image before this step:", self._image_part(reference_image_path)])
contents.extend(["Edited/generated image to score:", self._image_part(edited_image_path)])
last_exc: Exception | None = None
for attempt in range(self.max_attempts):
try:
response = self._client_instance().models.generate_content(
model=self.model,
contents=contents,
config=self._make_score_config(),
)
payload = _load_jsonish(getattr(response, "text", "") or "")
return InterleaveThinkerEditScore(
semantic_score=float(payload["semantic_score"]),
quality_score=float(payload["quality_score"]),
)
except Exception as exc: # noqa: BLE001 - remote API errors vary by SDK version
last_exc = exc
if attempt + 1 < self.max_attempts:
time.sleep(self.retry_delay_s)
raise RuntimeError(f"Gemini image scoring failed after {self.max_attempts} attempts: {last_exc}") from last_exc
def _client_instance(self) -> Any:
if self._client is not None:
return self._client
genai, _ = _import_google_genai()
kwargs: dict[str, Any] = {"api_key": _resolve_google_api_key(self.api_key)}
if self.base_url:
kwargs["http_options"] = {"base_url": self.base_url}
self._client = genai.Client(**kwargs)
return self._client
def _make_score_config(self) -> Any:
_, types = _import_google_genai()
return types.GenerateContentConfig(
system_instruction=_SCORE_SYSTEM_PROMPT,
response_mime_type="application/json",
response_schema=_SCORE_SCHEMA,
temperature=self.temperature,
)
def _image_part(self, path: str) -> Any:
_, types = _import_google_genai()
image_path = Path(path)
return types.Part.from_bytes(
data=image_path.read_bytes(),
mime_type=_image_mime_type(image_path),
)
class GeminiNanoBananaEditScorer:
"""Generate an edit with Nano Banana and score it with Gemini.
This callable matches ``EditScoreProvider`` from
``interleave_thinker.py`` and can be passed to
``InterleaveThinkerRewardScorer``.
"""
def __init__(
self,
*,
image_model: str = "gemini-3.1-flash-image",
judge_model: str = "gemini-2.5-pro",
api_key: str | None = None,
base_url: str | None = None,
output_dir: str = "outputs/interleave_thinker_reward_api",
width: int = 1024,
height: int = 1024,
num_inference_steps: int = 4,
guidance_scale: float = 1.0,
aspect_ratio: str | None = "1:1",
image_size: str | None = None,
max_attempts: int = 3,
retry_delay_s: float = 2.0,
treat_white_canvas_as_text_to_image: bool = True,
) -> None:
self.output_dir = output_dir
self.width = int(width)
self.height = int(height)
self.num_inference_steps = int(num_inference_steps)
self.guidance_scale = float(guidance_scale)
self.treat_white_canvas_as_text_to_image = bool(treat_white_canvas_as_text_to_image)
self.editor = NanoBananaImageGeneratorBackend(
model=image_model,
api_key=api_key,
base_url=base_url,
output_dir=output_dir,
aspect_ratio=aspect_ratio,
image_size=image_size,
max_attempts=max_attempts,
retry_delay_s=retry_delay_s,
)
self.scorer = GeminiInterleaveImageScorer(
model=judge_model,
api_key=api_key,
base_url=base_url,
max_attempts=max_attempts,
retry_delay_s=retry_delay_s,
)
def __call__(
self,
request: InterleaveThinkerEditRequest,
) -> InterleaveThinkerEditScore | None:
input_image_path = request.previous_image_path or request.origin_image_path
image_base64 = None
if input_image_path and not self._is_text_to_image_canvas(input_image_path):
image_base64 = encode_file_to_base64(input_image_path)
edit_request = InterleaveAPIEditRequest(
prompt=request.refine_prompt,
image=image_base64,
width=self.width,
height=self.height,
num_inference_steps=self.num_inference_steps,
guidance_scale=self.guidance_scale,
output_format="png",
enhance_prompt=False,
)
generated = self.editor.generate(
edit_request,
request_id=f"reward_{request.index}_{uuid.uuid4().hex}",
)
if not generated.file_path:
return None
return self.scorer.score_images(
reference_image_path=input_image_path if input_image_path else None,
edited_image_path=generated.file_path,
instruction=request.origin_prompt or request.refine_prompt,
)
def _is_text_to_image_canvas(self, path: str) -> bool:
if not self.treat_white_canvas_as_text_to_image:
return False
normalized = path.replace(os.sep, "/")
return normalized.endswith("data/interleave/white.png") or Path(path).name == "white.png"
class ConstantInterleaveEditScorer:
"""Small deterministic scorer for smoke tests and offline debugging."""
def __init__(
self,
*,
semantic_reward: float = 0.5,
quality_reward: float = 0.5,
) -> None:
self.semantic_reward = float(semantic_reward)
self.quality_reward = float(quality_reward)
def __call__(
self,
request: InterleaveThinkerEditRequest,
) -> InterleaveThinkerEditScore:
del request
return InterleaveThinkerEditScore(
semantic_reward=self.semantic_reward,
quality_reward=self.quality_reward,
)
def _resolve_google_api_key(explicit: str | None = None) -> str:
if explicit:
return explicit.strip()
for env_name in ("GEMINI_API_KEY", "GOOGLE_API_KEY"):
value = os.environ.get(env_name)
if value:
return value.strip()
token_path = Path("~/.gemini_token").expanduser()
if token_path.is_file():
return token_path.read_text().strip()
raise ValueError("Gemini API wrappers require GEMINI_API_KEY, GOOGLE_API_KEY, "
"an explicit api_key, or ~/.gemini_token.")
def _image_mime_type(path: str | Path) -> str:
suffix = Path(path).suffix.lower()
mime_type = _IMAGE_MIME_TYPES.get(suffix)
if mime_type is None:
supported = ", ".join(sorted(_IMAGE_MIME_TYPES))
raise ValueError(f"Unsupported InterleaveThinker image extension {suffix!r}; expected one of {supported}")
return mime_type
def _import_google_genai() -> tuple[Any, Any]:
try:
from google import genai
from google.genai import types
except ImportError as exc:
raise RuntimeError("Gemini / Nano Banana API wrappers require google-genai. "
"Install google-genai directly or with `uv pip install -e '.[eval-judge]'`.") from exc
return genai, types
def _load_jsonish(raw: str) -> dict[str, Any]:
text = (raw or "").strip()
if text.startswith("```"):
text = text.strip("`")
if text.startswith("json"):
text = text[4:].strip()
try:
payload = json.loads(text)
except json.JSONDecodeError:
start = text.find("{")
end = text.rfind("}")
if start < 0 or end <= start:
raise
payload = json.loads(text[start:end + 1])
if not isinstance(payload, dict):
raise ValueError(f"Gemini scorer returned non-object JSON: {type(payload).__name__}")
return payload
def image_bytes_to_base64(data: bytes) -> str:
return base64.b64encode(data).decode("utf-8")
def image_to_png_base64(image: Any) -> str:
buffer = io.BytesIO()
image.save(buffer, format="PNG")
return image_bytes_to_base64(buffer.getvalue())
__all__ = [
"ConstantInterleaveEditScorer",
"GeminiInterleaveImageScorer",
"GeminiNanoBananaEditScorer",
"image_bytes_to_base64",
"image_to_png_base64",
]
@@ -0,0 +1,517 @@
# SPDX-License-Identifier: Apache-2.0
"""InterleaveThinker critic reward utilities.
Adapted from InterleaveThinker's
``train/EasyR1/verl/reward_function/interleave_thinker_reward.py``. The
networked edit API and Gemini scorer are intentionally left outside this module
so FastVideo training methods can inject those services without making reward
parsing depend on external credentials.
"""
from __future__ import annotations
import ast
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
import json
import re
from typing import Any
import torch
_ANSWER_PATTERN = re.compile(r"<answer>\s*(.*?)\s*</answer>", re.DOTALL | re.IGNORECASE)
_THINK_PATTERN = re.compile(r"<think\b[^>]*>.*?</think>", re.DOTALL | re.IGNORECASE)
_TAG_SPACE_PATTERN = re.compile(r"\s*(<|>|/)\s*")
@dataclass(frozen=True, slots=True)
class InterleaveThinkerAnswer:
"""Parsed critic answer payload."""
previous_step_success: bool
refine_prompt: str
@dataclass(frozen=True, slots=True)
class InterleaveThinkerEditRequest:
"""Metadata needed by an external image-edit scorer."""
index: int
origin_prompt: str
previous_prompt: str
refine_prompt: str
origin_image_path: str | None
previous_image_path: str | None
previous_step_success: bool
previous_semantic_score: float
previous_quality_score: float
@dataclass(frozen=True, slots=True)
class InterleaveThinkerEditScore:
"""Edit scorer output.
``semantic_score`` and ``quality_score`` are absolute post-edit scores on
the same 0-10 scale used by InterleaveThinker. ``semantic_reward`` and
``quality_reward`` are normalized reward components in [0, 1]. Callers can
provide either absolute scores or normalized rewards.
"""
semantic_score: float | None = None
quality_score: float | None = None
semantic_reward: float | None = None
quality_reward: float | None = None
@dataclass(frozen=True, slots=True)
class InterleaveThinkerRewardResult:
"""One scored InterleaveThinker rollout."""
overall: float
format_reward: float
judge_accuracy_reward: float
edited_image_reward_semantic: float
edited_image_reward_quality: float
predicted_previous_step_success: bool | None
refine_prompt: str
index: int
def as_dict(self) -> dict[str, float]:
return {
"overall": float(self.overall),
"format_reward": float(self.format_reward),
"judge_accuracy_reward": float(self.judge_accuracy_reward),
"edited_image_reward_semantic": float(self.edited_image_reward_semantic),
"edited_image_reward_quality": float(self.edited_image_reward_quality),
"idx": float(self.index),
}
@dataclass(frozen=True, slots=True)
class InterleavePlannerRewardResult:
"""One scored InterleaveThinker planner rollout."""
overall: float
format_reward: float
planner_score: float
index: int
def as_dict(self) -> dict[str, float]:
return {
"overall": float(self.overall),
"format_reward": float(self.format_reward),
"planner_score": float(self.planner_score),
"idx": float(self.index),
}
EditScoreProvider = Callable[[InterleaveThinkerEditRequest], InterleaveThinkerEditScore | Mapping[str, Any] | None]
def normalize_interleave_response(response: str) -> str:
"""Match upstream tag-spacing normalization before parsing."""
return _TAG_SPACE_PATTERN.sub(r"\1", str(response or ""))
def _jsonish_loads(raw: str) -> dict[str, Any] | None:
raw = raw.strip()
if not raw:
return None
try:
parsed = json.loads(raw)
except json.JSONDecodeError:
try:
parsed = ast.literal_eval(raw)
except (ValueError, SyntaxError):
return None
if isinstance(parsed, dict):
return parsed
return None
def extract_interleave_answer(response: str) -> InterleaveThinkerAnswer | None:
"""Extract the ``<answer>`` JSON object expected by InterleaveThinker."""
match = _ANSWER_PATTERN.search(normalize_interleave_response(response))
if not match:
return None
payload = _jsonish_loads(match.group(1))
if payload is None:
return None
previous_step_success = payload.get("previous_step_success")
refine_prompt = payload.get("refine_prompt")
if not isinstance(previous_step_success, bool) or not isinstance(refine_prompt, str):
return None
return InterleaveThinkerAnswer(
previous_step_success=previous_step_success,
refine_prompt=refine_prompt,
)
def interleave_format_reward(response: str) -> float:
"""Return 1 when response has valid ``<think>`` then ``<answer>`` format."""
normalized = normalize_interleave_response(response)
think_match = _THINK_PATTERN.search(normalized)
answer_match = _ANSWER_PATTERN.search(normalized)
if think_match is None or answer_match is None:
return 0.0
if think_match.end() > answer_match.start():
return 0.0
return 1.0 if extract_interleave_answer(normalized) is not None else 0.0
def extract_interleave_plan_payload(response: str) -> dict[str, Any] | None:
"""Extract the planner ``execution_plan`` answer payload."""
match = _ANSWER_PATTERN.search(normalize_interleave_response(response))
if not match:
return None
payload = _jsonish_loads(match.group(1))
if payload is None:
return None
raw_steps = payload.get("execution_plan")
if not isinstance(raw_steps, Sequence) or isinstance(raw_steps, str | bytes) or not raw_steps:
return None
for step in raw_steps:
if not isinstance(step, Mapping):
return None
has_prompt = _non_empty_optional_text(step.get("prompt")) is not None
has_instruction = _non_empty_optional_text(step.get("instruction")) is not None
has_auxiliary = _non_empty_optional_text(step.get("auxiliary_text")) is not None
if not (has_prompt or has_instruction or has_auxiliary):
return None
return payload
def interleave_planner_format_reward(response: str) -> float:
"""Return 1 when a planner response has valid reasoning and plan JSON."""
normalized = normalize_interleave_response(response)
think_match = _THINK_PATTERN.search(normalized)
answer_match = _ANSWER_PATTERN.search(normalized)
if think_match is None or answer_match is None:
return 0.0
if think_match.end() > answer_match.start():
return 0.0
return 1.0 if extract_interleave_plan_payload(normalized) is not None else 0.0
def interleave_judge_accuracy_reward(
predicted_previous_step_success: bool | None,
ground_truth_previous_step_success: bool | None,
) -> float:
"""Reward correct previous-step success prediction."""
if predicted_previous_step_success is None or ground_truth_previous_step_success is None:
return 0.0
return 1.0 if bool(predicted_previous_step_success) == bool(ground_truth_previous_step_success) else 0.0
def _coerce_float(value: Any, *, default: float = 0.0) -> float:
if value is None:
return default
if torch.is_tensor(value):
if value.numel() != 1:
raise ValueError(f"Expected scalar tensor, got shape={tuple(value.shape)}")
return float(value.detach().cpu())
return float(value)
def _non_empty_optional_text(value: Any) -> str | None:
if value is None:
return None
text = str(value).strip()
if not text or text.lower() in {"none", "null"}:
return None
return text
def _mapping_get_bool(mapping: Mapping[str, Any], *keys: str) -> bool | None:
for key in keys:
if key in mapping:
value = mapping[key]
if isinstance(value, bool):
return bool(value)
return None
def _ground_truth(mapping: Mapping[str, Any]) -> tuple[bool | None, float, float]:
raw = mapping.get("ground_truth", mapping.get("evaluation", {}))
if isinstance(raw, str):
raw = _jsonish_loads(raw) or {}
if not isinstance(raw, Mapping):
return None, 0.0, 0.0
success = _mapping_get_bool(raw, "success", "previous_step_success")
return (
success,
_coerce_float(raw.get("semantics", raw.get("semantic_score")), default=0.0),
_coerce_float(raw.get("quality", raw.get("quality_score")), default=0.0),
)
def _make_edit_request(
item: Mapping[str, Any],
*,
index: int,
answer: InterleaveThinkerAnswer | None,
previous_success: bool,
previous_semantics: float,
previous_quality: float,
) -> InterleaveThinkerEditRequest:
previous_prompt = str(item.get("previous_prompt", item.get("rewritten_prompt", "")) or "")
refine_prompt = answer.refine_prompt if answer is not None and answer.refine_prompt else previous_prompt
return InterleaveThinkerEditRequest(
index=index,
origin_prompt=str(item.get("origin_prompt", item.get("prompt", "")) or ""),
previous_prompt=previous_prompt,
refine_prompt=refine_prompt,
origin_image_path=item.get("origin_image_path"),
previous_image_path=item.get("previous_image_path", item.get("edited_image_path")),
previous_step_success=previous_success,
previous_semantic_score=previous_semantics,
previous_quality_score=previous_quality,
)
def _coerce_edit_score(raw: InterleaveThinkerEditScore | Mapping[str, Any] | None) -> InterleaveThinkerEditScore | None:
if raw is None:
return None
if isinstance(raw, InterleaveThinkerEditScore):
return raw
if not isinstance(raw, Mapping):
raise TypeError(f"edit score must be a mapping or InterleaveThinkerEditScore, got {type(raw).__name__}")
return InterleaveThinkerEditScore(
semantic_score=raw.get("semantic_score", raw.get("semantics")),
quality_score=raw.get("quality_score", raw.get("quality")),
semantic_reward=raw.get("semantic_reward", raw.get("edited_image_reward_semantic")),
quality_reward=raw.get("quality_reward", raw.get("edited_image_reward_quality")),
)
def _normalized_edit_rewards(
score: InterleaveThinkerEditScore | None,
*,
previous_semantics: float,
previous_quality: float,
fallback: float,
) -> tuple[float, float]:
if score is None:
return fallback, fallback
semantic_reward = score.semantic_reward
quality_reward = score.quality_reward
if semantic_reward is None and score.semantic_score is not None:
semantic_reward = ((_coerce_float(score.semantic_score) - previous_semantics) / 10.0 + 1.0) / 2.0
if quality_reward is None and score.quality_score is not None:
quality_reward = ((_coerce_float(score.quality_score) - previous_quality) / 10.0 + 1.0) / 2.0
return (
_coerce_float(semantic_reward, default=fallback),
_coerce_float(quality_reward, default=fallback),
)
def _planner_score_from_item(
item: Mapping[str, Any],
*,
fallback: float,
) -> float:
for key in ("planner_score", "plan_score", "reward", "score"):
if key in item:
return _coerce_float(item.get(key), default=fallback)
raw = item.get("ground_truth", item.get("evaluation", {}))
if isinstance(raw, str):
raw = _jsonish_loads(raw) or {}
if isinstance(raw, Mapping):
for key in ("planner_score", "plan_score", "reward", "score"):
if key in raw:
return _coerce_float(raw.get(key), default=fallback)
return fallback
class InterleaveThinkerRewardScorer:
"""Batch scorer for InterleaveThinker critic outputs.
The default weights match upstream InterleaveThinker's ``compute_score``:
``0.5 * format + 0.5 * (0.2 * judge + 0.6 * semantic + 0.2 * quality)``.
"""
def __init__(
self,
*,
format_weight: float = 0.5,
judge_accuracy_weight: float = 0.2,
semantic_weight: float = 0.6,
quality_weight: float = 0.2,
fallback_edit_reward: float = 0.5,
edit_scorer: EditScoreProvider | None = None,
) -> None:
self.format_weight = float(format_weight)
self.judge_accuracy_weight = float(judge_accuracy_weight)
self.semantic_weight = float(semantic_weight)
self.quality_weight = float(quality_weight)
self.fallback_edit_reward = float(fallback_edit_reward)
self.edit_scorer = edit_scorer
if not 0.0 <= self.format_weight <= 1.0:
raise ValueError("format_weight must be in [0, 1]")
inner_total = self.judge_accuracy_weight + self.semantic_weight + self.quality_weight
if inner_total <= 0.0:
raise ValueError("At least one non-format reward weight must be positive")
def __call__(
self,
reward_inputs: Sequence[Mapping[str, Any]],
) -> list[InterleaveThinkerRewardResult]:
results: list[InterleaveThinkerRewardResult] = []
inner_total = self.judge_accuracy_weight + self.semantic_weight + self.quality_weight
for index, item in enumerate(reward_inputs):
response = str(item.get("response", "") or "")
answer = extract_interleave_answer(response)
format_reward = interleave_format_reward(response)
previous_success, previous_semantics, previous_quality = _ground_truth(item)
predicted_success = answer.previous_step_success if answer is not None else None
judge_reward = interleave_judge_accuracy_reward(predicted_success, previous_success)
edit_request = _make_edit_request(
item,
index=index,
answer=answer,
previous_success=bool(previous_success),
previous_semantics=previous_semantics,
previous_quality=previous_quality,
)
edit_score = _coerce_edit_score(item.get("edit_score", item.get("edit_scores")))
if edit_score is None and self.edit_scorer is not None:
edit_score = _coerce_edit_score(self.edit_scorer(edit_request))
semantic_reward, quality_reward = _normalized_edit_rewards(
edit_score,
previous_semantics=previous_semantics,
previous_quality=previous_quality,
fallback=self.fallback_edit_reward,
)
non_format_reward = (self.judge_accuracy_weight * judge_reward + self.semantic_weight * semantic_reward +
self.quality_weight * quality_reward) / inner_total
overall = self.format_weight * format_reward + (1.0 - self.format_weight) * non_format_reward
results.append(
InterleaveThinkerRewardResult(
overall=float(overall),
format_reward=float(format_reward),
judge_accuracy_reward=float(judge_reward),
edited_image_reward_semantic=float(semantic_reward),
edited_image_reward_quality=float(quality_reward),
predicted_previous_step_success=predicted_success,
refine_prompt=edit_request.refine_prompt,
index=index,
))
return results
def as_tensors(
self,
reward_inputs: Sequence[Mapping[str, Any]],
*,
device: torch.device | str = "cpu",
) -> dict[str, torch.Tensor]:
results = self(reward_inputs)
names = [
"overall",
"format_reward",
"judge_accuracy_reward",
"edited_image_reward_semantic",
"edited_image_reward_quality",
]
tensors: dict[str, torch.Tensor] = {}
for name in names:
tensors[name] = torch.tensor([getattr(result, name) for result in results],
device=device,
dtype=torch.float32)
tensors["avg"] = tensors["overall"]
return tensors
class InterleavePlannerRewardScorer:
"""Batch scorer for InterleaveThinker planner outputs.
This scorer is intentionally lightweight: by default it rewards valid
planner response format and can blend in an externally supplied scalar plan
score from each rollout for stronger supervision.
"""
def __init__(
self,
*,
format_weight: float = 1.0,
fallback_plan_reward: float = 0.0,
) -> None:
self.format_weight = float(format_weight)
self.fallback_plan_reward = float(fallback_plan_reward)
if not 0.0 <= self.format_weight <= 1.0:
raise ValueError("format_weight must be in [0, 1]")
def __call__(
self,
reward_inputs: Sequence[Mapping[str, Any]],
) -> list[InterleavePlannerRewardResult]:
results: list[InterleavePlannerRewardResult] = []
for index, item in enumerate(reward_inputs):
response = str(item.get("response", "") or "")
format_reward = interleave_planner_format_reward(response)
planner_score = _planner_score_from_item(
item,
fallback=self.fallback_plan_reward,
)
overall = self.format_weight * format_reward + (1.0 - self.format_weight) * planner_score
results.append(
InterleavePlannerRewardResult(
overall=float(overall),
format_reward=float(format_reward),
planner_score=float(planner_score),
index=index,
))
return results
def as_tensors(
self,
reward_inputs: Sequence[Mapping[str, Any]],
*,
device: torch.device | str = "cpu",
) -> dict[str, torch.Tensor]:
results = self(reward_inputs)
names = ["overall", "format_reward", "planner_score"]
tensors: dict[str, torch.Tensor] = {}
for name in names:
tensors[name] = torch.tensor([getattr(result, name) for result in results],
device=device,
dtype=torch.float32)
tensors["avg"] = tensors["overall"]
return tensors
def score_interleave_thinker_rewards(
reward_inputs: Sequence[Mapping[str, Any]],
**kwargs: Any,
) -> list[dict[str, float]]:
"""Convenience wrapper matching upstream's list-of-dicts return shape."""
scorer = InterleaveThinkerRewardScorer(**kwargs)
return [result.as_dict() for result in scorer(reward_inputs)]
def score_interleave_planner_rewards(
reward_inputs: Sequence[Mapping[str, Any]],
**kwargs: Any,
) -> list[dict[str, float]]:
"""Convenience wrapper for planner rollout rewards."""
scorer = InterleavePlannerRewardScorer(**kwargs)
return [result.as_dict() for result in scorer(reward_inputs)]
__all__ = [
"EditScoreProvider",
"InterleavePlannerRewardResult",
"InterleaveThinkerAnswer",
"InterleaveThinkerEditRequest",
"InterleaveThinkerEditScore",
"InterleaveThinkerRewardResult",
"InterleavePlannerRewardScorer",
"InterleaveThinkerRewardScorer",
"extract_interleave_answer",
"extract_interleave_plan_payload",
"interleave_format_reward",
"interleave_planner_format_reward",
"interleave_judge_accuracy_reward",
"normalize_interleave_response",
"score_interleave_planner_rewards",
"score_interleave_thinker_rewards",
]
@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
"""Generic media reward composition utilities."""
from __future__ import annotations
from collections.abc import Callable, Mapping, Sequence
import torch
RewardScorer = Callable[[torch.Tensor, Sequence[str]], torch.Tensor]
def select_first_frame(media: torch.Tensor) -> torch.Tensor:
"""Return first-frame media as ``[B, C, H, W]``.
This is a helper for reward models that are intrinsically frame-based
(for example PickScore and CLIPScore). Video-aware rewards should inspect
the full ``[B, C, T, H, W]`` tensor themselves.
"""
if not torch.is_tensor(media):
raise TypeError(f"media must be a torch.Tensor, got {type(media).__name__}")
if media.ndim == 5:
return media[:, :, 0]
if media.ndim == 4:
return media
raise ValueError("media must have shape [B, C, H, W] or [B, C, T, H, W], "
f"got {tuple(media.shape)}")
class MultiRewardScorer:
"""Weighted sum of reusable media reward scorers.
Mirrors DiffusionNFT's ``flow_grpo/rewards.py::multi_score`` behavior,
while leaving frame selection to each concrete reward.
"""
def __init__(
self,
reward_weights: Mapping[str, float],
*,
scorers: Mapping[str, RewardScorer],
) -> None:
self.reward_weights = {str(k): float(v) for k, v in reward_weights.items()}
if not self.reward_weights:
raise ValueError("reward_weights must contain at least one reward")
self.scorers = dict(scorers)
unsupported = sorted(set(self.reward_weights) - set(self.scorers))
if unsupported:
raise ValueError(f"Unsupported reward(s): {unsupported}. "
f"Available rewards: {sorted(self.scorers)}")
@torch.no_grad()
def __call__(
self,
media: torch.Tensor,
prompts: Sequence[str],
) -> dict[str, torch.Tensor]:
prompt_count = len(prompts)
if media.shape[0] != prompt_count:
raise ValueError(f"media batch size ({media.shape[0]}) must match prompt count ({prompt_count})")
total: torch.Tensor | None = None
details: dict[str, torch.Tensor] = {}
for name, weight in self.reward_weights.items():
scores = self.scorers[name](media, prompts).detach().float()
if scores.ndim != 1 or int(scores.shape[0]) != prompt_count:
raise ValueError(f"Reward {name!r} must return shape [{prompt_count}], got {tuple(scores.shape)}")
details[name] = scores
weighted = scores * float(weight)
total = weighted if total is None else total.to(weighted.device) + weighted
assert total is not None
details["avg"] = total
return details
+30 -8
View File
@@ -17,18 +17,17 @@ if TYPE_CHECKING:
from fastvideo.pipelines import TrainingBatch
class ModelBase(ABC):
"""Per-role model instance.
class RoleModelBase(ABC):
"""Minimal per-role model instance.
Every role (student, teacher, critic, …) gets its own ``ModelBase``
instance. Each instance owns its own ``transformer`` and
``noise_scheduler``. Heavyweight resources (VAE, dataloader, RNG
seeds) are loaded lazily via :meth:`init_preprocessors`, which the
method calls **only on the student**.
Every training role (student, teacher, critic, reference, …) gets its own
role-model instance. Each instance owns its role-local ``transformer`` and
trainability policy. Heavyweight resources such as dataloaders are loaded
lazily via :meth:`init_preprocessors`, which methods usually call only on
the student.
"""
transformer: torch.nn.Module
noise_scheduler: Any
_trainable: bool
def __init__(
@@ -91,6 +90,29 @@ class ModelBase(ABC):
def on_train_start(self) -> None: # noqa: B027
"""Called once before the training loop begins."""
class ModelBase(RoleModelBase):
"""Diffusion per-role model instance.
Diffusion models additionally own a ``noise_scheduler`` and implement the
latent/noise runtime primitives used by standard FastVideo training
methods. Non-diffusion actor models should inherit from
:class:`RoleModelBase` directly.
"""
noise_scheduler: Any
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
"""Decode ``[B, T, C, H, W]`` latents to ``[B, C, T, H, W]`` media.
RL reward methods call this hook instead of reaching into
model-specific VAE normalization details.
"""
raise NotImplementedError(f"{type(self).__name__} does not implement decode_latents()")
# ------------------------------------------------------------------
# Timestep helpers
# ------------------------------------------------------------------
@@ -0,0 +1,64 @@
# SPDX-License-Identifier: Apache-2.0
"""InterleaveThinker training model adapters."""
from fastvideo.train.models.interleave_thinker.critic import (
INTERLEAVE_CRITIC_PROMPT,
InterleaveThinkerCriticModel,
)
from fastvideo.train.models.interleave_thinker.data import (
DEFAULT_FILENAMES,
IMAGE_EXTENSIONS,
IMAGE_LIST_KEYS,
IMAGE_PATH_KEYS,
InterleaveDatasetKind,
load_critic_rl_records,
load_critic_sft_records,
load_interleave_dataset,
load_planner_rl_records,
load_planner_sft_records,
normalize_critic_rl_record,
normalize_critic_sft_record,
normalize_ground_truth,
normalize_interleave_dataset_record,
normalize_planner_rl_record,
normalize_planner_sft_record,
resolve_interleave_image_path,
validate_image_path,
)
from fastvideo.train.models.interleave_thinker.planner import (
INTERLEAVE_GUIDANCE_PLANNER_PROMPT,
INTERLEAVE_PLANNER_PROMPT,
InterleavePlannerOutput,
InterleavePlannerStep,
InterleaveThinkerPlannerModel,
extract_interleave_plan,
)
__all__ = [
"INTERLEAVE_CRITIC_PROMPT",
"INTERLEAVE_GUIDANCE_PLANNER_PROMPT",
"INTERLEAVE_PLANNER_PROMPT",
"DEFAULT_FILENAMES",
"IMAGE_EXTENSIONS",
"IMAGE_LIST_KEYS",
"IMAGE_PATH_KEYS",
"InterleaveDatasetKind",
"InterleavePlannerOutput",
"InterleavePlannerStep",
"InterleaveThinkerCriticModel",
"InterleaveThinkerPlannerModel",
"extract_interleave_plan",
"load_critic_rl_records",
"load_critic_sft_records",
"load_interleave_dataset",
"load_planner_rl_records",
"load_planner_sft_records",
"normalize_critic_rl_record",
"normalize_critic_sft_record",
"normalize_ground_truth",
"normalize_interleave_dataset_record",
"normalize_planner_rl_record",
"normalize_planner_sft_record",
"resolve_interleave_image_path",
"validate_image_path",
]
@@ -0,0 +1,272 @@
# SPDX-License-Identifier: Apache-2.0
"""InterleaveThinker Qwen3-VL critic actor adapter."""
from __future__ import annotations
import atexit
from collections.abc import Mapping
from contextlib import suppress
from functools import lru_cache
import os
import tempfile
from typing import Any, TYPE_CHECKING
from PIL import Image
import torch
from fastvideo.train.models.interleave_thinker.qwen_actor import (
Qwen3VLActorBase,
_PlaceholderActorModule as _PlaceholderActorModule,
batch_to_items,
first_string,
rollout_group_key,
)
from fastvideo.train.models.interleave_thinker.data import InterleaveDatasetKind, looks_like_uri
if TYPE_CHECKING:
from fastvideo.train.utils.lora import LoraConfig
from fastvideo.train.utils.training_config import TrainingConfig
INTERLEAVE_CRITIC_PROMPT = """<image><image>
# Generation/Edit Evaluation and Prompt Refinement System
You are an expert image editing evaluator and prompt engineer. Your task is to:
1. Evaluate the edited image and output the result in boolean format (True/False).
2. If you think the edited image is not good enough (False), generate an optimized rewritten prompt that addresses the original shortcomings; if you think it is good enough (True), output the [Original Rewritten Prompt].
## Input Information
You have been presented with two images in sequence:
- Original Image: The input image before editing. (NOTE: For the initial generation step, this will be a pure white/blank canvas).
- Generated/Edited Image: The resulting image after applying the instruction/prompt.
Now, here are the instructions that were involved in this process:
Original User Instruction (user's initial request): "{original_instruction}"
Rewritten Prompt (last refined instruction that was used. **NOTE: If this is empty, you must base your evaluation and refinement entirely on the Original User Instruction**): "{rewritten_prompt}"
## Evaluation Instructions
**Evaluate Previous Step (Strict 2-Part Check)**: Carefully compare the **Before Image** and the **After Image**. You must evaluate based on two strict criteria. If the image fails *either* criteria, the step is a FAILURE.
1. **Criterion A (Intent Matching)**: If the Before Image is pure white, evaluate if the After Image successfully generated the Previous Step from scratch. Otherwise, observe the delta (differences). Did the changes match the key meaning and necessary details of the Previous Step?
2. **Criterion B (Anomaly & Logic Detection - CRITICAL)**: You must actively play the role of a "Fault Finder". Do NOT just check if the requested object exists; you MUST check HOW it exists. Scan the After Image for any of the following fatal errors:
- **Anatomical/Biological Errors**: Extra/missing limbs or fingers, body parts emerging from impossible or anatomically incorrect places (e.g., a hand growing out of a chest, stomach, or a wall), distorted faces.
- **Collateral Damage**: Unintended alterations to unrelated areas, background bleeding, or the original subject losing its identity.
## Prompt Refinement Strategy (if NOT GOOD ENOUGH, False)
When generating a new rewritten prompt, analyze:
1. **What went wrong?**
- Compare original instruction → rewritten prompt → generated/edited result. *(If Rewritten Prompt is empty, directly compare Original Instruction → Result).*
- Identify gaps between intent and execution
- Determine if the issue is clarity, specificity, or contradiction
2. **Refinement Approaches:**
**If this is an Initial Generation task (Before image was blank):**
- **Establish Foundation:** Translate the raw user instruction into a comprehensive Text-to-Image prompt.
- **Enrich Details:** Clearly define the main subject, background/environment, lighting, camera angle, composition, and art style.
- **Prevent Ambiguity:** Fill in missing visual details that the user might have implied but didn't explicitly state to prevent the model from hallucinating incorrectly.
- **Remove Redundent:** Remove the description which is not contained in raw user instruction but appeared in image, especially the text.
**If the rewritten prompt was too vague:**
- Add more specific descriptors (exact colors, positions, sizes)
- Include spatial relationships and context
- Specify interaction with existing elements
**If the rewritten prompt was contradictory:**
- Resolve conflicts between requirements
- Prioritize core intent over secondary details
- Simplify complex multi-part instructions
**If important details were lost:**
- Explicitly state preservation requirements
- Add "maintain [aspect]" or "preserve [feature]" clauses
- Reference specific elements from the original image
**If positioning/scale was wrong:**
- Use more precise spatial descriptors
- Add relative size/scale indicators
- Specify foreground/midground/background placement
**If style/appearance was incorrect:**
- Use more specific visual vocabulary
- Add reference to original image's style elements
- Include material/texture/lighting specifications
**If the edit was over/under-processed:**
- Add modifiers like "subtle", "gentle", "dramatic", "significant"
- Specify degree of change more clearly
- Balance enhancement with naturalness
3. **Leverage All Information:**
- Reference what's visible in the original image
- Learn from what the previous rewritten prompt missed
- Use the edited image as feedback on what went wrong
- Maintain what worked, fix what didn't
## Output
The output consists of three parts:
1. A Statement - Analysis process and reasoning;
2. A Boolean - Judge whether the edited images is good enough;
3. A prompt — either the optimized rewritten prompt or the original rewritten prompt.
Here is a output example:
<think>
Detailed explanation of evaluation and new rewritten prompt. If edited image is good enough, explain why it meets requirements. If not good enough, explain specific shortcomings.
</think>
<answer>
{{
'previous_step_success': 'boolean (True ONLY IF the Intent Check is successful AND the Anomaly Check finds ZERO errors. If ANY anomaly is detected, this MUST be False.)',
'refine_prompt': '[Improved rewritten prompt that addresses identified issues and enhances clarity, specificity, and preservation requirements] if NOT GOOD ENOUGH (False), [original rewritten prompt] if GOOD ENOUGH (True)'
}}
</answer>
"""
_DEFAULT_BLANK_CANVAS_SIZE = (1024, 1024)
_BLANK_CANVAS_PATHS: set[str] = set()
class InterleaveThinkerCriticModel(Qwen3VLActorBase):
"""Qwen3-VL actor wrapper for InterleaveThinker critic RL."""
def __init__(
self,
*,
init_from: str = "InterleaveThinker/Critic-SFT-8B",
processor_from: str | None = None,
training_config: TrainingConfig | None = None,
trainable: bool = True,
load_backend: bool = True,
image_dir: str = "",
torch_dtype: str = "auto",
device_map: str | dict[str, Any] | None = None,
attn_implementation: str | None = None,
trust_remote_code: bool = False,
use_cache: bool = False,
freeze_vision_tower: bool = True,
freeze_multi_modal_projector: bool = True,
enable_gradient_checkpointing: bool = True,
max_prompt_length: int = 16384,
max_response_length: int = 4096,
dataset_kind: InterleaveDatasetKind | None = None,
prompt_template: str = INTERLEAVE_CRITIC_PROMPT,
lora: LoraConfig | dict[str, Any] | None = None,
**kwargs: Any,
) -> None:
self.prompt_template = prompt_template
super().__init__(
init_from=init_from,
processor_from=processor_from,
training_config=training_config,
trainable=trainable,
load_backend=load_backend,
image_dir=image_dir,
torch_dtype=torch_dtype,
device_map=device_map,
attn_implementation=attn_implementation,
trust_remote_code=trust_remote_code,
use_cache=use_cache,
freeze_vision_tower=freeze_vision_tower,
freeze_multi_modal_projector=freeze_multi_modal_projector,
enable_gradient_checkpointing=enable_gradient_checkpointing,
max_prompt_length=max_prompt_length,
max_response_length=max_response_length,
dataset_kind=dataset_kind,
lora=lora,
**kwargs,
)
@torch.no_grad()
def generate_interleave_responses(
self,
batch: dict[str, Any],
**kwargs: Any,
) -> list[dict[str, Any]]:
num_generations = max(1, int(kwargs.get("num_generations", 1) or 1))
temperature_value = kwargs.get("temperature", 1.0)
top_p_value = kwargs.get("top_p", 1.0)
temperature = 1.0 if temperature_value is None else float(temperature_value)
top_p = 1.0 if top_p_value is None else float(top_p_value)
max_new_tokens = int(kwargs.get("max_new_tokens") or self.max_response_length)
rollouts: list[dict[str, Any]] = []
for item_idx, item in enumerate(batch_to_items(batch)):
decoded = self.generate_qwen_responses(
self.build_messages(item),
num_generations=num_generations,
temperature=temperature,
top_p=top_p,
max_new_tokens=max_new_tokens,
)
for generation_idx, response in enumerate(decoded):
rollout = dict(item)
rollout["response"] = response
rollout.setdefault("sample_index", item_idx)
rollout.setdefault("generation_index", generation_idx)
rollout.setdefault("group_key", rollout_group_key(rollout, item_idx))
old_logprobs, response_mask = self.response_logprobs_from_messages(
self.build_messages(item),
response,
)
rollout["old_logprobs"] = old_logprobs.detach().cpu().tolist()
rollout["response_mask"] = response_mask.detach().cpu().tolist()
rollouts.append(rollout)
return rollouts
def build_messages(
self,
item: Mapping[str, Any],
) -> list[dict[str, Any]]:
prompt = self.prompt_template.format(
original_instruction=str(item.get("origin_prompt", item.get("prompt", "")) or ""),
rewritten_prompt=str(item.get("previous_prompt", item.get("rewritten_prompt", "")) or ""),
)
return self.build_text_image_messages(prompt, _item_image_paths(item))
def _item_image_paths(item: Mapping[str, Any], ) -> list[str]:
before = first_string(item, "previous_image_path", "origin_image_path", "input_image_path")
after = first_string(item, "edited_image_path", "generated_image_path", "output_image_path")
if not before and after:
before = _materialize_blank_canvas(_image_size(after))
return [value for value in (before, after) if value]
def _image_size(image_path: str) -> tuple[int, int]:
if looks_like_uri(image_path):
return _DEFAULT_BLANK_CANVAS_SIZE
try:
with Image.open(image_path) as image:
return image.size
except (OSError, ValueError):
return _DEFAULT_BLANK_CANVAS_SIZE
@lru_cache(maxsize=16)
def _materialize_blank_canvas(size: tuple[int, int]) -> str:
width, height = size
fd, path = tempfile.mkstemp(
prefix=f"fastvideo-interleave-blank-{width}x{height}-",
suffix=".png",
)
os.close(fd)
try:
Image.new("RGB", size, color="white").save(path, format="PNG")
except Exception:
os.unlink(path)
raise
_BLANK_CANVAS_PATHS.add(path)
return path
def _cleanup_blank_canvases() -> None:
for path in _BLANK_CANVAS_PATHS:
with suppress(FileNotFoundError):
os.unlink(path)
atexit.register(_cleanup_blank_canvases)
__all__ = ["INTERLEAVE_CRITIC_PROMPT", "InterleaveThinkerCriticModel", "_PlaceholderActorModule"]
@@ -0,0 +1,433 @@
# SPDX-License-Identifier: Apache-2.0
"""Dataset normalization utilities for InterleaveThinker training files."""
from __future__ import annotations
from collections.abc import Mapping, Sequence
import json
import os
from pathlib import Path
from typing import Any, Literal
InterleaveDatasetKind = Literal["planner_sft", "planner_rl", "critic_sft", "critic_rl"]
DEFAULT_FILENAMES: dict[InterleaveDatasetKind, str] = {
"planner_sft": "planner_sft.json",
"planner_rl": "planner_rl.jsonl",
"critic_sft": "critic_sft.json",
"critic_rl": "critic_rl.jsonl",
}
IMAGE_PATH_KEYS = (
"origin_image_path",
"previous_image_path",
"edited_image_path",
"generated_image_path",
"input_image_path",
"output_image_path",
"image_path",
"target_img",
)
IMAGE_LIST_KEYS = (
"images",
"input_image_paths",
"image_paths",
)
IMAGE_EXTENSIONS = frozenset({
".bmp",
".gif",
".jpeg",
".jpg",
".png",
".tif",
".tiff",
".webp",
})
def load_interleave_dataset(
data_path: str | os.PathLike[str],
*,
kind: InterleaveDatasetKind,
image_dir: str | os.PathLike[str] = "",
validate_image_files: bool = False,
) -> list[dict[str, Any]]:
records: list[dict[str, Any]] = []
for file_path in _resolve_dataset_files(data_path, kind):
for raw in _load_json_records(file_path):
records.append(
normalize_interleave_dataset_record(
raw,
kind=kind,
image_dir=image_dir,
validate_image_files=validate_image_files,
))
if not records:
raise ValueError(f"No {kind} records found at {data_path!s}")
return records
def load_planner_sft_records(
data_path: str | os.PathLike[str],
*,
image_dir: str | os.PathLike[str] = "",
validate_image_files: bool = False,
) -> list[dict[str, Any]]:
return load_interleave_dataset(
data_path,
kind="planner_sft",
image_dir=image_dir,
validate_image_files=validate_image_files,
)
def load_planner_rl_records(
data_path: str | os.PathLike[str],
*,
image_dir: str | os.PathLike[str] = "",
validate_image_files: bool = False,
) -> list[dict[str, Any]]:
return load_interleave_dataset(
data_path,
kind="planner_rl",
image_dir=image_dir,
validate_image_files=validate_image_files,
)
def load_critic_sft_records(
data_path: str | os.PathLike[str],
*,
image_dir: str | os.PathLike[str] = "",
validate_image_files: bool = False,
) -> list[dict[str, Any]]:
return load_interleave_dataset(
data_path,
kind="critic_sft",
image_dir=image_dir,
validate_image_files=validate_image_files,
)
def load_critic_rl_records(
data_path: str | os.PathLike[str],
*,
image_dir: str | os.PathLike[str] = "",
validate_image_files: bool = False,
) -> list[dict[str, Any]]:
return load_interleave_dataset(
data_path,
kind="critic_rl",
image_dir=image_dir,
validate_image_files=validate_image_files,
)
def normalize_interleave_dataset_record(
record: Mapping[str, Any],
*,
kind: InterleaveDatasetKind,
image_dir: str | os.PathLike[str] = "",
validate_image_files: bool = False,
) -> dict[str, Any]:
normalized = _resolve_record_image_paths(
record,
image_dir=image_dir,
validate_image_files=validate_image_files,
)
if kind == "planner_sft":
return normalize_planner_sft_record(normalized)
if kind == "planner_rl":
return normalize_planner_rl_record(normalized)
if kind == "critic_sft":
return normalize_critic_sft_record(normalized)
if kind == "critic_rl":
return normalize_critic_rl_record(normalized)
raise ValueError(f"Unsupported InterleaveThinker dataset kind: {kind!r}")
def normalize_planner_sft_record(record: Mapping[str, Any]) -> dict[str, Any]:
normalized = dict(record)
messages = _require_messages(normalized, where="planner_sft")
user_message = _first_message_content(messages, "user")
assistant_message = _first_message_content(messages, "assistant")
if not user_message:
raise ValueError("planner_sft record requires a user message")
if not assistant_message:
raise ValueError("planner_sft record requires an assistant message")
normalized.setdefault("instruction", user_message)
normalized.setdefault("response", assistant_message)
images = _string_sequence(normalized.get("images"))
if images:
normalized.setdefault("input_image_paths", images)
return normalized
def normalize_planner_rl_record(record: Mapping[str, Any]) -> dict[str, Any]:
normalized = dict(record)
if "messages" in normalized:
messages = _require_messages(normalized, where="planner_rl")
normalized.setdefault("instruction", _first_message_content(messages, "user"))
instruction = _first_text(
normalized.get("instruction"),
normalized.get("text_input"),
normalized.get("origin_prompt"),
normalized.get("prompt"),
)
if not instruction:
raise ValueError("planner_rl record requires instruction, text_input, origin_prompt, or prompt")
normalized["instruction"] = instruction
images = _string_sequence(normalized.get("images"))
if images:
normalized.setdefault("input_image_paths", images)
return normalized
def normalize_critic_sft_record(record: Mapping[str, Any]) -> dict[str, Any]:
normalized = dict(record)
if "messages" in normalized:
messages = _require_messages(normalized, where="critic_sft")
normalized.setdefault("instruction", _first_message_content(messages, "user"))
normalized.setdefault("response", _first_message_content(messages, "assistant"))
images = _string_sequence(normalized.get("images"))
if images:
if len(images) < 2:
raise ValueError("critic_sft ShareGPT record requires at least two images")
normalized.setdefault("origin_image_path", images[0])
normalized.setdefault("edited_image_path", images[1])
_require_text(normalized, "origin_image_path", "critic_sft")
_require_text(normalized, "edited_image_path", "critic_sft")
normalized.setdefault("previous_image_path", normalized["origin_image_path"])
normalized.setdefault("generated_image_path", normalized["edited_image_path"])
if "previous_prompt" not in normalized and "rewritten_prompt" in normalized:
normalized["previous_prompt"] = normalized["rewritten_prompt"]
if "rewritten_prompt" not in normalized and "previous_prompt" in normalized:
normalized["rewritten_prompt"] = normalized["previous_prompt"]
return normalized
def normalize_critic_rl_record(record: Mapping[str, Any]) -> dict[str, Any]:
normalized = dict(record)
_require_text(normalized, "origin_prompt", "critic_rl")
_require_text(normalized, "origin_image_path", "critic_rl")
_require_text(normalized, "edited_image_path", "critic_rl")
previous_prompt = _first_text(
normalized.get("previous_prompt"),
normalized.get("rewritten_prompt"),
normalized.get("refine_prompt"),
)
if not previous_prompt:
raise ValueError("critic_rl record requires previous_prompt or rewritten_prompt")
normalized["previous_prompt"] = previous_prompt
normalized.setdefault("rewritten_prompt", previous_prompt)
normalized.setdefault("previous_image_path", normalized["origin_image_path"])
normalized.setdefault("generated_image_path", normalized["edited_image_path"])
normalized["ground_truth"] = normalize_ground_truth(normalized)
return normalized
def normalize_ground_truth(record: Mapping[str, Any]) -> dict[str, Any]:
raw = record.get("ground_truth", record.get("evaluation", {}))
if isinstance(raw, str):
try:
raw = json.loads(raw)
except json.JSONDecodeError as exc:
raise ValueError(f"critic_rl ground_truth is not valid JSON: {exc}") from exc
if not isinstance(raw, Mapping):
raw = {}
success = _optional_bool(raw.get("success", raw.get("previous_step_success", record.get("previous_step_success"))))
if success is None:
raise ValueError("critic_rl ground_truth requires boolean success or previous_step_success")
return {
"success": success,
"semantics": _optional_float(raw.get("semantics", raw.get("semantic_score")), default=0.0),
"quality": _optional_float(raw.get("quality", raw.get("quality_score")), default=0.0),
}
def resolve_interleave_image_path(
value: str,
*,
image_dir: str | os.PathLike[str] = "",
validate_image_files: bool = False,
) -> str:
expanded = os.path.expanduser(value)
if not image_dir or os.path.isabs(expanded) or looks_like_uri(expanded):
resolved = expanded
else:
resolved = str(Path(os.path.expanduser(str(image_dir))) / expanded)
validate_image_path(resolved, validate_exists=validate_image_files)
return resolved
def validate_image_path(
value: str,
*,
validate_exists: bool = False,
) -> None:
if looks_like_uri(value):
return
suffix = Path(value).suffix.lower()
if suffix not in IMAGE_EXTENSIONS:
raise ValueError(f"Unsupported image file extension for InterleaveThinker path {value!r}")
if validate_exists and not Path(value).is_file():
raise FileNotFoundError(f"InterleaveThinker image path does not exist: {value}")
def looks_like_uri(value: str) -> bool:
return "://" in value or value.startswith("data:")
def _resolve_dataset_files(
data_path: str | os.PathLike[str],
kind: InterleaveDatasetKind,
) -> list[Path]:
path = Path(os.path.expanduser(str(data_path)))
if path.is_dir():
path = path / DEFAULT_FILENAMES[kind]
if not path.exists():
raise FileNotFoundError(f"InterleaveThinker {kind} data file not found: {path}")
if path.suffix not in {".json", ".jsonl"}:
raise ValueError(f"Unsupported InterleaveThinker data file: {path}")
return [path]
def _load_json_records(file_path: Path) -> list[dict[str, Any]]:
try:
if file_path.suffix == ".jsonl":
records: list[dict[str, Any]] = []
for line_number, line in enumerate(file_path.read_text(encoding="utf-8").splitlines(), start=1):
if not line.strip():
continue
raw_line = json.loads(line)
if not isinstance(raw_line, Mapping):
raise ValueError(f"{file_path}:{line_number} must contain a JSON object")
records.append(dict(raw_line))
return records
raw = json.loads(file_path.read_text(encoding="utf-8"))
except json.JSONDecodeError as exc:
raise ValueError(f"Malformed JSON in InterleaveThinker data file {file_path}: {exc}") from exc
if isinstance(raw, list):
if not all(isinstance(item, Mapping) for item in raw):
raise ValueError(f"{file_path} must contain only JSON objects")
return [dict(item) for item in raw]
if isinstance(raw, Mapping):
for key in ("data", "records"):
value = raw.get(key)
if isinstance(value, list):
if not all(isinstance(item, Mapping) for item in value):
raise ValueError(f"{file_path}.{key} must contain only JSON objects")
return [dict(item) for item in value]
return [dict(raw)]
raise ValueError(f"{file_path} must contain a JSON object or list of objects")
def _resolve_record_image_paths(
record: Mapping[str, Any],
*,
image_dir: str | os.PathLike[str],
validate_image_files: bool,
) -> dict[str, Any]:
normalized = dict(record)
for key in IMAGE_PATH_KEYS:
value = normalized.get(key)
if isinstance(value, str) and value:
normalized[key] = resolve_interleave_image_path(
value,
image_dir=image_dir,
validate_image_files=validate_image_files,
)
for key in IMAGE_LIST_KEYS:
value = normalized.get(key)
if isinstance(value, Sequence) and not isinstance(value, str | bytes):
normalized[key] = [
resolve_interleave_image_path(
item,
image_dir=image_dir,
validate_image_files=validate_image_files,
) if isinstance(item, str) and item else item for item in value
]
return normalized
def _require_messages(record: Mapping[str, Any], *, where: str) -> list[dict[str, Any]]:
raw = record.get("messages")
if not isinstance(raw, Sequence) or isinstance(raw, str | bytes):
raise ValueError(f"{where} record requires ShareGPT-style messages")
messages = [dict(item) for item in raw if isinstance(item, Mapping)]
if len(messages) != len(raw):
raise ValueError(f"{where} messages must be JSON objects")
return messages
def _first_message_content(messages: Sequence[Mapping[str, Any]], role: str) -> str:
for message in messages:
if message.get("role") == role:
content = message.get("content")
if isinstance(content, str):
return content
return ""
def _require_text(record: Mapping[str, Any], key: str, where: str) -> str:
value = record.get(key)
if not isinstance(value, str) or not value:
raise ValueError(f"{where} record requires {key}")
return value
def _string_sequence(value: Any) -> list[str]:
if not isinstance(value, Sequence) or isinstance(value, str | bytes):
return []
return [item for item in value if isinstance(item, str)]
def _first_text(*values: Any) -> str:
for value in values:
if isinstance(value, str) and value:
return value
return ""
def _optional_bool(value: Any) -> bool | None:
if isinstance(value, bool):
return value
return None
def _optional_float(value: Any, *, default: float) -> float:
if value is None:
return default
try:
return float(value)
except (TypeError, ValueError):
return default
__all__ = [
"DEFAULT_FILENAMES",
"IMAGE_EXTENSIONS",
"IMAGE_LIST_KEYS",
"IMAGE_PATH_KEYS",
"InterleaveDatasetKind",
"load_critic_rl_records",
"load_critic_sft_records",
"load_interleave_dataset",
"load_planner_rl_records",
"load_planner_sft_records",
"looks_like_uri",
"normalize_critic_rl_record",
"normalize_critic_sft_record",
"normalize_ground_truth",
"normalize_interleave_dataset_record",
"normalize_planner_rl_record",
"normalize_planner_sft_record",
"resolve_interleave_image_path",
"validate_image_path",
]
@@ -0,0 +1,393 @@
# SPDX-License-Identifier: Apache-2.0
"""InterleaveThinker Qwen3-VL planner adapter."""
from __future__ import annotations
import ast
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
import json
import re
from typing import Any, TYPE_CHECKING
import torch
from fastvideo.train.models.interleave_thinker.qwen_actor import (
Qwen3VLActorBase,
batch_to_items,
rollout_group_key,
)
from fastvideo.train.models.interleave_thinker.data import InterleaveDatasetKind
if TYPE_CHECKING:
from fastvideo.train.utils.lora import LoraConfig
from fastvideo.train.utils.training_config import TrainingConfig
INTERLEAVE_PLANNER_PROMPT = """
# Task Planner, Orchestrator, and Prompt Engineer System
You are an expert **Task Planner, Orchestrator, and Prompt Engineer**.
Your goal is to analyze a user's request, generate a structured execution plan, and optimize EVERY step's instruction into a highly effective Text-to-Image (T2I) prompt or Image Editing instruction.
## Input Information
Here are the instructions that were involved in this process:
Original User Instruction (user's request): "{text_input}"
## Execution Plan Instructions
1. **Dynamic Step Count (Image Operations Only)**: Determine the necessary number of steps. Every step in your execution plan MUST represent an actual image generation or image editing action. **DO NOT** create separate steps solely for generating text, captions, or summaries.
2. **Complete & Polished Output**: Always aim for a fully realized final product. For visual or creative tasks, the final step MUST result in a fully colored, detailed, and polished output. Do not stop at a draft, outline, or uncolored sketch unless the user explicitly requests it.
3. **Text Generation & Auxiliary Text Rule**:
- If the user specifically asks to render or draw text *inside* the image, include this requirement within the `instruction` field.
- If the user explicitly asks for a *separate* text response (e.g., a caption, summary, explanation, or knowledge grounding) to accompany the image, generate this text and place it in the `auxiliary_text` field of the corresponding image generation step.
- If the user does not explicitly request any separate text or caption, you MUST set `auxiliary_text` to `null`.
## Optimize Prompt Instructions
1. **Prompt Optimization for All Steps**: Convert the `instruction` of EVERY step into a highly effective prompt in the `prompt` field.
- **Step 1 (Generation)**: Create a highly detailed T2I prompt representing the foundational stage. Focus *only* on the Step 1 instruction. Do NOT hallucinate unmentioned details or future elements.
- **Subsequent Steps (Editing)**: Create clear, actionable image editing instructions (e.g., "add a red hat", "change the background to a cyberpunk city") based on the current step's goal.
2. **CRITICAL**: The `prompt` field MUST contain ONLY the pure text prompt or editing instruction. DO NOT include meta-text, prefixes (such as "Step 1:", "Prompt:", "Edit:"), or conversational filler. It must be directly usable by the generation/editing API.
## Output
The output consists of two parts:
1. A Statement - Analysis process and reasoning;
2. A JSON — Planing each step and rewrite the instruction to prompt suitable for generation/editing.
Here is a output example
<think>
Part 1: Planning analysis explaining the execution plan. Part 2: Analysis of how the instructions were translated into visual keywords for the T2I prompt and editing instructions.
</think>
<answer>
{
'execution_plan':
[
{'step_number': 1, 'step_name': 'Short name for the step', 'instruction': 'Detailed instruction for this image generation step.', 'prompt': "The optimized, pure T2I prompt suitable for the image generation model. (No 'Step 1:' prefix)", 'auxiliary_text': 'The required caption, summary, or text explanation. Output null if no separate text is explicitly requested.'},
{'step_number': 2, 'step_name': 'Short name for the step', 'instruction': 'Detailed instruction for this image editing step.', 'prompt': "The optimized, pure instruction suitable for the image editing model. (No 'Step 2:' prefix)", 'auxiliary_text': None}
]
}
</answer>
"""
INTERLEAVE_GUIDANCE_PLANNER_PROMPT = """
You are an expert **Multimodal Sequence Planner and Orchestrator**.
Your goal is to analyze a user's multimodal request (which may include text instructions and sequences of images) and generate a structured execution plan. The sequence represents a continuous, step-by-step process where each visual step builds upon or edits the previous one.
## Input Information
You have been presented with a text-images sequence: "{text_input}"
### Instructions
1. **Task Identification & Modality Routing**: Carefully analyze the input to determine the task type.
- **Task A (General Text Response / Problem Solving / Image-to-Text)**: If the user provides a complete sequence of images and asks for text responses for each step (e.g., describing the images, solving a problem, explaining a process, or answering questions), you must write your complete response entirely within the `auxiliary_text` field. You MUST set BOTH the `instruction` and `prompt` fields to `null` for these steps.
- **Task B (Sequence Continuation / Sequential Editing)**: If the user provides a partial sequence and asks to predict/generate the remaining steps, you must generate both the text instruction and the editing prompt. The `prompt` field must contain an optimized instruction specifically tailored for an **image editing model** to modify the previous step's image into the new state.
2. **Strict Step Count & NO Prefix Rule**:
- **Step Count**: Determine the logical number of steps. **CRITICAL**: If the user's input explicitly specifies the number of steps required, you MUST strictly output exactly that number of steps to fulfill the requirement. If continuing a sequence (Task B), your `step_number` MUST start exactly from where the user's input left off.
- **NO Prefixes**: BOTH the `instruction` and `prompt` fields MUST NOT contain any step prefixes, numbers, or bullet points (e.g., DO NOT write "(3)", "Step 3:", or "Step 3: Plant the seeds". Just write "Plant the seeds").
3. **Field Definitions & Usage**:
- `instruction`: The detailed, pure text content or action for the editing step (Task B). You MUST set this to `null` for Task A. (Strictly NO step prefixes).
- `prompt`: The optimized, pure instruction suitable for the **image editing model** to execute the change based on the previous image (Task B). You MUST set this to `null` for Task A. (Strictly NO step prefixes).
- `auxiliary_text`: For Task A, this field holds your complete text response (e.g., descriptions, problem-solving steps, or answers). For Task B, use this ONLY if the user explicitly requests or the task naturally requires an extra knowledge-based description/summary during the continuation process; otherwise, output `null`.
4. **Complete Output**: Ensure the final step achieves a complete resolution of the user's goal based on the sequence context.
## Output
The output consists of two parts:
1. A Statement - Just an dummy reasoning;
2. A JSON — Planing each step and rewrite the instruction to prompt suitable for generation/editing.
Here is a output example
<think>
</think>
<answer>
{
'execution_plan':
[
{'step_number': i, 'step_name': 'Short name for the step', 'instruction': "Detailed instruction for this step (Task B). Output null if this is Task A. Strictly NO prefixes like 'Step i:' or '(i)'.", 'prompt': "The optimized instruction suitable for the image editing model (Task B). Output null if this is Task A. Strictly NO prefixes like 'Step i:' or '(i)'.", 'auxiliary_text': 'The complete text answer/solution for Task A, OR the extra knowledge explanation for Task B. Output null if not needed.'},
{'step_number': i+1, 'step_name': 'Short name for the step', 'instruction': "Detailed instruction for this step (Task B). Output null if this is Task A. Strictly NO prefixes like 'Step i+1:' or '(i+1)'.", 'prompt': "The optimized instruction suitable for the image editing model (Task B). Output null if this is Task A. Strictly NO prefixes like 'Step i+1:' or '(i+1)'.", 'auxiliary_text': 'The complete text answer/solution for Task A, OR the extra knowledge explanation for Task B. Output null if not needed.'}
]
}
</answer>
"""
@dataclass(frozen=True, slots=True)
class InterleavePlannerStep:
step_number: int | None
step_name: str | None
instruction: str | None
prompt: str | None
auxiliary_text: str | None
@dataclass(frozen=True, slots=True)
class InterleavePlannerOutput:
raw_response: str
raw_answer: str
steps: tuple[InterleavePlannerStep, ...]
class InterleaveThinkerPlannerModel(Qwen3VLActorBase):
"""Qwen3-VL actor wrapper for InterleaveThinker planning."""
def __init__(
self,
*,
init_from: str = "InterleaveThinker/InterleaveThinker-Planner-8B",
processor_from: str | None = "Qwen/Qwen3-VL-8B-Instruct",
training_config: TrainingConfig | None = None,
trainable: bool = True,
load_backend: bool = True,
image_dir: str = "",
torch_dtype: str = "auto",
device_map: str | dict[str, Any] | None = None,
attn_implementation: str | None = None,
trust_remote_code: bool = False,
use_cache: bool = False,
freeze_vision_tower: bool = True,
freeze_multi_modal_projector: bool = True,
enable_gradient_checkpointing: bool = True,
max_prompt_length: int = 16384,
max_response_length: int = 4096,
dataset_kind: InterleaveDatasetKind | None = None,
prompt_template: str = INTERLEAVE_PLANNER_PROMPT,
guidance_prompt_template: str = INTERLEAVE_GUIDANCE_PLANNER_PROMPT,
lora: LoraConfig | dict[str, Any] | None = None,
**kwargs: Any,
) -> None:
self.prompt_template = prompt_template
self.guidance_prompt_template = guidance_prompt_template
super().__init__(
init_from=init_from,
processor_from=processor_from,
training_config=training_config,
trainable=trainable,
load_backend=load_backend,
image_dir=image_dir,
torch_dtype=torch_dtype,
device_map=device_map,
attn_implementation=attn_implementation,
trust_remote_code=trust_remote_code,
use_cache=use_cache,
freeze_vision_tower=freeze_vision_tower,
freeze_multi_modal_projector=freeze_multi_modal_projector,
enable_gradient_checkpointing=enable_gradient_checkpointing,
max_prompt_length=max_prompt_length,
max_response_length=max_response_length,
dataset_kind=dataset_kind,
lora=lora,
**kwargs,
)
@torch.no_grad()
def generate_interleave_plans(
self,
batch: dict[str, Any],
**kwargs: Any,
) -> list[dict[str, Any]]:
num_generations = max(1, int(kwargs.get("num_generations", 1) or 1))
temperature_value = kwargs.get("temperature", 1.0)
top_p_value = kwargs.get("top_p", 1.0)
temperature = 1.0 if temperature_value is None else float(temperature_value)
top_p = 1.0 if top_p_value is None else float(top_p_value)
max_new_tokens = int(kwargs.get("max_new_tokens") or self.max_response_length)
outputs: list[dict[str, Any]] = []
for item_idx, item in enumerate(batch_to_items(batch)):
decoded = self.generate_qwen_responses(
self.build_messages(item),
num_generations=num_generations,
temperature=temperature,
top_p=top_p,
max_new_tokens=max_new_tokens,
)
for generation_idx, response in enumerate(decoded):
parsed = extract_interleave_plan(response)
outputs.append({
"item": dict(item),
"response": response,
"plan": parsed,
"steps": list(parsed.steps) if parsed is not None else [],
"sample_index": item_idx,
"generation_index": generation_idx,
})
return outputs
@torch.no_grad()
def generate_interleave_responses(
self,
batch: dict[str, Any],
**kwargs: Any,
) -> list[dict[str, Any]]:
num_generations = max(1, int(kwargs.get("num_generations", 1) or 1))
temperature_value = kwargs.get("temperature", 1.0)
top_p_value = kwargs.get("top_p", 1.0)
temperature = 1.0 if temperature_value is None else float(temperature_value)
top_p = 1.0 if top_p_value is None else float(top_p_value)
max_new_tokens = int(kwargs.get("max_new_tokens") or self.max_response_length)
rollouts: list[dict[str, Any]] = []
for item_idx, item in enumerate(batch_to_items(batch)):
messages = self.build_messages(item)
decoded = self.generate_qwen_responses(
messages,
num_generations=num_generations,
temperature=temperature,
top_p=top_p,
max_new_tokens=max_new_tokens,
)
for generation_idx, response in enumerate(decoded):
parsed = extract_interleave_plan(response)
rollout = dict(item)
rollout["response"] = response
rollout["plan"] = parsed
rollout["steps"] = list(parsed.steps) if parsed is not None else []
rollout.setdefault("sample_index", item_idx)
rollout.setdefault("generation_index", generation_idx)
rollout.setdefault("group_key", _planner_group_key(rollout, item_idx))
old_logprobs, response_mask = self.response_logprobs_from_messages(
messages,
response,
)
rollout["old_logprobs"] = old_logprobs.detach().cpu().tolist()
rollout["response_mask"] = response_mask.detach().cpu().tolist()
rollouts.append(rollout)
return rollouts
def generate_interleave_plan(
self,
instruction: str,
*,
input_image_paths: Sequence[str] | None = None,
**kwargs: Any,
) -> dict[str, Any]:
batch = {
"items": [{
"instruction": instruction,
"input_image_paths": list(input_image_paths or []),
}]
}
plans = self.generate_interleave_plans(batch, **kwargs)
if not plans:
raise RuntimeError("InterleaveThinker planner returned no plan generations")
return plans[0]
def build_messages(
self,
item: Mapping[str, Any],
) -> list[dict[str, Any]]:
image_paths = _planner_image_paths(item)
template = self.guidance_prompt_template if image_paths else self.prompt_template
prompt = template.replace("{text_input}", _planner_instruction(item))
return self.build_text_image_messages(prompt, image_paths)
def extract_interleave_plan(response: str) -> InterleavePlannerOutput | None:
raw_answer = _extract_answer_block(response)
if not raw_answer:
return None
payload = _load_answer_mapping(raw_answer)
if not isinstance(payload, Mapping):
return None
raw_steps = payload.get("execution_plan")
if not isinstance(raw_steps, Sequence) or isinstance(raw_steps, str | bytes):
return None
steps = tuple(_coerce_planner_step(step) for step in raw_steps if isinstance(step, Mapping))
return InterleavePlannerOutput(
raw_response=response,
raw_answer=raw_answer,
steps=steps,
)
def _coerce_planner_step(raw: Mapping[str, Any]) -> InterleavePlannerStep:
return InterleavePlannerStep(
step_number=_optional_int(raw.get("step_number")),
step_name=_optional_text(raw.get("step_name")),
instruction=_optional_text(raw.get("instruction")),
prompt=_optional_text(raw.get("prompt")),
auxiliary_text=_optional_text(raw.get("auxiliary_text")),
)
def _extract_answer_block(response: str) -> str:
match = re.search(r"<answer>\s*(.*?)\s*</answer>", response, flags=re.DOTALL | re.IGNORECASE)
if match:
return match.group(1).strip()
return response.strip()
def _load_answer_mapping(raw_answer: str) -> Any:
try:
return json.loads(raw_answer)
except json.JSONDecodeError:
pass
normalized = re.sub(r"\bnull\b", "None", raw_answer)
normalized = re.sub(r"\btrue\b", "True", normalized)
normalized = re.sub(r"\bfalse\b", "False", normalized)
try:
return ast.literal_eval(normalized)
except (ValueError, SyntaxError):
return None
def _optional_int(value: Any) -> int | None:
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return None
def _optional_text(value: Any) -> str | None:
if value is None:
return None
text = str(value).strip()
if not text or text.lower() in {"none", "null"}:
return None
return text
def _planner_instruction(item: Mapping[str, Any]) -> str:
for key in ("instruction", "text_input", "origin_prompt", "prompt"):
value = item.get(key)
if isinstance(value, str) and value:
return value
return ""
def _planner_image_paths(item: Mapping[str, Any]) -> list[str]:
for key in ("input_image_paths", "image_paths", "images"):
value = item.get(key)
if isinstance(value, Sequence) and not isinstance(value, str | bytes):
paths = [str(path) for path in value if path]
if paths:
return paths
for key in ("input_image_path", "image_path", "origin_image_path"):
value = item.get(key)
if isinstance(value, str) and value:
return [value]
return []
def _planner_group_key(
rollout: Mapping[str, Any],
index: int,
) -> str:
for key in ("group_key", "problem_id", "instruction", "text_input", "origin_prompt", "prompt"):
value = rollout.get(key)
if value is not None:
return str(value)
return rollout_group_key(rollout, index)
__all__ = [
"INTERLEAVE_GUIDANCE_PLANNER_PROMPT",
"INTERLEAVE_PLANNER_PROMPT",
"InterleavePlannerOutput",
"InterleavePlannerStep",
"InterleaveThinkerPlannerModel",
"extract_interleave_plan",
]
File diff suppressed because it is too large Load Diff
+47 -5
View File
@@ -101,6 +101,7 @@ class WanModel(ModelBase):
self.negative_prompt_embeds: (torch.Tensor | None) = None
self.negative_prompt_attention_mask: (torch.Tensor | None) = None
self._requires_negative_conditioning = True
# Timestep mechanics.
self.timestep_shift: float = float(flow_shift)
@@ -160,17 +161,31 @@ class WanModel(ModelBase):
self._init_timestep_mechanics()
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_t2v, )
pyarrow_schema_t2v,
pyarrow_schema_text_only,
)
from fastvideo.train.utils.dataloader import (
build_parquet_t2v_train_dataloader, )
preprocessed_data_type = str(getattr(
training_config.data,
"preprocessed_data_type",
"t2v",
)).strip().lower()
parquet_schema = pyarrow_schema_t2v
if preprocessed_data_type == "text_only":
parquet_schema = pyarrow_schema_text_only
elif preprocessed_data_type != "t2v":
raise ValueError("Unsupported Wan preprocessed_data_type: "
f"{preprocessed_data_type!r}")
text_len = (
training_config.pipeline_config.text_encoder_configs[ # type: ignore[union-attr]
0].arch_config.text_len)
self.dataloader = build_parquet_t2v_train_dataloader(
training_config.data,
text_len=int(text_len),
parquet_schema=pyarrow_schema_t2v,
parquet_schema=parquet_schema,
)
self.start_step = 0
@@ -178,6 +193,9 @@ class WanModel(ModelBase):
def num_train_timesteps(self) -> int:
return int(self.num_train_timestep)
def set_requires_negative_conditioning(self, requires: bool) -> None:
self._requires_negative_conditioning = bool(requires)
def shift_and_clamp_timestep(self, timestep: torch.Tensor) -> torch.Tensor:
timestep = shift_timestep(
timestep,
@@ -187,7 +205,25 @@ class WanModel(ModelBase):
return timestep.clamp(self.min_timestep, self.max_timestep)
def on_train_start(self) -> None:
self.ensure_negative_conditioning()
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
@torch.no_grad()
def decode_latents(
self,
latents_b_t_c_h_w: torch.Tensor,
) -> torch.Tensor:
if self.vae is None:
raise RuntimeError("Wan VAE is not initialized")
latents = latents_b_t_c_h_w.permute(0, 2, 1, 3, 4).float()
if bool(getattr(self.vae, "handles_latent_denorm", False)):
denorm = latents
else:
mean = torch.tensor(self.vae.latents_mean, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
std = torch.tensor(self.vae.latents_std, device=latents.device, dtype=latents.dtype).view(1, -1, 1, 1, 1)
denorm = latents * std + mean
media = self.vae.to(latents.device).decode(denorm)
return (media / 2 + 0.5).clamp(0, 1)
# ------------------------------------------------------------------
# Runtime primitives
@@ -200,7 +236,8 @@ class WanModel(ModelBase):
generator: torch.Generator,
latents_source: Literal["data", "zeros"] = "data",
) -> TrainingBatch:
self.ensure_negative_conditioning()
if self._requires_negative_conditioning:
self.ensure_negative_conditioning()
assert self.training_config is not None
tc = self.training_config
@@ -285,7 +322,7 @@ class WanModel(ModelBase):
attn_kind: Literal["dense", "vsa"] = "dense",
) -> torch.Tensor:
device_type = self.device.type
dtype = noisy_latents.dtype
dtype = self._get_training_dtype()
if conditional:
text_dict = batch.conditional_dict
if text_dict is None:
@@ -301,6 +338,11 @@ class WanModel(ModelBase):
else:
raise ValueError(f"Unknown attn_kind: {attn_kind!r}")
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
# Keep Wan training autocast tied to the model's training dtype, not
# to caller-created intermediates that may accidentally be fp32.
with torch.autocast(device_type, dtype=dtype), set_forward_context(
current_timestep=batch.timesteps,
attn_metadata=attn_metadata,
+3 -1
View File
@@ -128,7 +128,9 @@ class WanCausalModel(WanModel, CausalModelBase):
}
device_type = self.device.type
dtype = noisy_latents.dtype
dtype = self._get_training_dtype()
if noisy_latents.is_floating_point():
noisy_latents = noisy_latents.to(dtype=dtype)
if conditional:
text_dict = batch.conditional_dict
+67 -25
View File
@@ -12,7 +12,7 @@ from tqdm.auto import tqdm
from fastvideo.distributed import get_sp_group, get_world_group
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.methods.base import TrainingMethod
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.utils.tracking import build_tracker
if TYPE_CHECKING:
@@ -82,6 +82,22 @@ class Trainer:
batch = next(data_iter)
yield batch
def _run_method_validation(
self,
method: TrainingMethod,
iteration: int,
) -> None:
hook = getattr(method, "on_validation_begin", None)
if hook is None:
return
validation_metrics: dict[str, LogScalar] = hook(iteration)
validation_metrics = {
k: float(_coerce_log_scalar(v, where=(f"method.on_validation_begin().metrics[{k!r}]")))
for k, v in validation_metrics.items()
}
if self.global_rank == 0 and validation_metrics:
self.tracker.log(validation_metrics, iteration)
def run(
self,
method: TrainingMethod,
@@ -115,6 +131,7 @@ class Trainer:
method,
iteration=start_step,
)
self._run_method_validation(method, start_step)
method.optimizers_zero_grad(start_step)
data_stream = self._iter_dataloader(dataloader)
@@ -130,6 +147,8 @@ class Trainer:
desc="Steps",
disable=self.local_rank > 0,
)
# Allow method-specific optimization flow (e.g. DiffusionNFT).
method_manages_optimization = bool(method.manages_optimization())
for step in progress:
t0 = time.perf_counter()
@@ -137,47 +156,69 @@ class Trainer:
# to CPU once per step right before logging.
loss_sums: dict[str, float | torch.Tensor] = {}
metric_sums: dict[str, float | torch.Tensor] = {}
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
if method_manages_optimization:
loss_map, outputs, step_metrics = method.managed_train_step(
data_stream,
step,
))
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
loss_sums[k] = v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
metric_sums[k] = _coerce_log_scalar(
v,
where=("method.single_train_step()"
where=("method.managed_train_step()"
f".metrics[{k!r}]"),
)
else:
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
step,
))
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
v,
where=("method.single_train_step()"
f".metrics[{k!r}]"),
))
if not method_manages_optimization:
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
# Single CPU sync point: materialise GPU tensors
# to float right before logging.
metrics = {k: float(v) / grad_accum for k, v in loss_sums.items()}
metrics.update({k: float(v) / grad_accum for k, v in metric_sums.items()})
divisor = 1 if method_manages_optimization else grad_accum
metrics = {k: float(v) / divisor for k, v in loss_sums.items()}
metrics.update({k: float(v) / divisor for k, v in metric_sums.items()})
metrics["step_time_sec"] = (time.perf_counter() - t0)
metrics["vsa_sparsity"] = float(tc.vsa_sparsity)
if self.global_rank == 0 and metrics:
@@ -196,6 +237,7 @@ class Trainer:
method,
iteration=step,
)
self._run_method_validation(method, step)
self.callbacks.on_validation_end(
method,
iteration=step,
+4 -4
View File
@@ -25,15 +25,15 @@ def build_from_config(cfg: RunConfig, ) -> tuple[TrainingConfig, TrainingMethod,
and construct it with ``(cfg=cfg, role_models=...)``.
3. Return ``(training_args, method, dataloader, start_step)``.
"""
from fastvideo.train.models.base import ModelBase
from fastvideo.train.models.base import RoleModelBase
# --- 1. Build role model instances ---
role_models: dict[str, ModelBase] = {}
role_models: dict[str, RoleModelBase] = {}
for role, model_cfg in cfg.models.items():
model = instantiate(model_cfg, training_config=cfg.training)
if not isinstance(model, ModelBase):
if not isinstance(model, RoleModelBase):
raise TypeError(f"models.{role}._target_ must resolve to a "
f"ModelBase subclass, got {type(model).__name__}")
f"RoleModelBase subclass, got {type(model).__name__}")
role_models[role] = model
# --- 2. Build method ---
+7
View File
@@ -343,6 +343,12 @@ def _build_training_config(
if init_from is not None:
model_path = str(init_from)
preprocessed_data_type = str(da.get("preprocessed_data_type", "t2v") or "t2v").strip().lower()
if preprocessed_data_type not in {"t2v", "text_only"}:
raise ValueError("training.data.preprocessed_data_type must be one of "
"{'t2v', 'text_only'}, got "
f"{preprocessed_data_type!r}")
return TrainingConfig(
distributed=DistributedConfig(
num_gpus=num_gpus,
@@ -354,6 +360,7 @@ def _build_training_config(
),
data=DataConfig(
data_path=str(da.get("data_path", "") or ""),
preprocessed_data_type=preprocessed_data_type,
train_batch_size=int(da.get("train_batch_size", 1) or 1),
dataloader_num_workers=int(da.get("dataloader_num_workers", 0) or 0),
training_cfg_rate=float(da.get("training_cfg_rate", 0.0) or 0.0),
+1
View File
@@ -23,6 +23,7 @@ class DistributedConfig:
@dataclass(slots=True)
class DataConfig:
data_path: str = ""
preprocessed_data_type: str = "t2v"
train_batch_size: int = 1
dataloader_num_workers: int = 0
training_cfg_rate: float = 0.0
+1 -1
View File
@@ -112,7 +112,7 @@ class BaseTracker:
self._timed_metrics = {}
def log_artifacts(self, artifacts: dict[str, Any], step: int) -> None:
"""Log artifacts such as videos or images.
"""Log tracker artifacts such as sampled media.
By default this is treated the same as :meth:`log`.
"""
+2 -2
View File
@@ -1694,7 +1694,7 @@ class EMA_FSDP:
if p_local.numel() == 0:
# Nothing to swap on this rank for this param
continue
self.saved[name] = p_local.clone().to(device=p_local.device, dtype=p_local.dtype)
self.saved[name] = p_local.clone().to("cpu")
if name in self.ema.shadow:
ema_cpu = self.ema.shadow[name]
if ema_cpu.numel() != p_local.numel():
@@ -1714,7 +1714,7 @@ class EMA_FSDP:
saved_local = self.saved[name]
if saved_local.numel() != p_local.numel():
continue
p_local.copy_(saved_local)
p_local.copy_(saved_local.to(dtype=p_local.dtype, device=p_local.device))
self.saved.clear()
return False
@@ -0,0 +1,114 @@
# SPDX-License-Identifier: Apache-2.0
"""InterleaveThinker workflow helpers for FastVideo."""
from fastvideo.workflow.interleave_thinker.config import (
InterleaveCriticConfig,
InterleaveImageBackendConfig,
InterleavePlannerConfig,
InterleaveRunConfig,
InterleaveRunStateConfig,
load_interleave_run_config,
resolve_interleave_instruction,
)
from fastvideo.workflow.interleave_thinker.evaluation import (
InterleavePromptItem,
InterleavePromptResult,
InterleavePromptSetSummary,
load_interleave_prompt_set,
prompt_set_summary_to_dict,
run_interleave_prompt_set,
run_interleave_prompt_set_config,
save_prompt_set_summary,
)
from fastvideo.workflow.interleave_thinker.generator import (
FastVideoImageGeneratorBackend,
ImageGeneratorBackend,
)
from fastvideo.workflow.interleave_thinker.orchestrator import (
AcceptAllCritic,
CriticProvider,
InterleaveOrchestrator,
PlannerProvider,
SinglePromptPlanner,
)
from fastvideo.workflow.interleave_thinker.providers import (
InterleaveThinkerCriticProvider,
InterleaveThinkerPlannerProvider,
)
from fastvideo.workflow.interleave_thinker.runner import (
InterleaveRunResult,
run_interleave_config,
)
from fastvideo.workflow.interleave_thinker.schema import (
CriticDecision,
CriticInput,
GeneratedImage,
InterleaveAttempt,
InterleaveEditRequest,
InterleaveEditResponse,
InterleaveTrace,
PlannedInterleaveStep,
PlannerInput,
)
from fastvideo.workflow.interleave_thinker.trace import (
save_trace,
trace_to_dict,
)
from fastvideo.workflow.interleave_thinker.trace_eval import (
InterleaveTraceEvaluationSummary,
InterleaveTraceMetrics,
discover_interleave_trace_paths,
evaluate_interleave_traces,
interleave_trace_evaluation_to_dict,
load_interleave_trace_metrics,
write_interleave_trace_evaluation,
write_interleave_trace_html_report,
)
__all__ = [
"AcceptAllCritic",
"CriticDecision",
"CriticInput",
"CriticProvider",
"FastVideoImageGeneratorBackend",
"GeneratedImage",
"ImageGeneratorBackend",
"InterleaveAttempt",
"InterleaveCriticConfig",
"InterleaveEditRequest",
"InterleaveEditResponse",
"InterleaveImageBackendConfig",
"InterleaveOrchestrator",
"InterleavePlannerConfig",
"InterleavePromptItem",
"InterleavePromptResult",
"InterleavePromptSetSummary",
"InterleaveRunConfig",
"InterleaveRunResult",
"InterleaveRunStateConfig",
"InterleaveThinkerCriticProvider",
"InterleaveThinkerPlannerProvider",
"InterleaveTrace",
"InterleaveTraceEvaluationSummary",
"InterleaveTraceMetrics",
"PlannedInterleaveStep",
"PlannerInput",
"PlannerProvider",
"SinglePromptPlanner",
"discover_interleave_trace_paths",
"evaluate_interleave_traces",
"interleave_trace_evaluation_to_dict",
"load_interleave_prompt_set",
"load_interleave_run_config",
"load_interleave_trace_metrics",
"prompt_set_summary_to_dict",
"resolve_interleave_instruction",
"run_interleave_prompt_set",
"run_interleave_prompt_set_config",
"run_interleave_config",
"save_prompt_set_summary",
"save_trace",
"trace_to_dict",
"write_interleave_trace_evaluation",
"write_interleave_trace_html_report",
]
@@ -0,0 +1,218 @@
# SPDX-License-Identifier: Apache-2.0
"""Typed config for native interleaved generation workflows."""
from __future__ import annotations
from collections.abc import Mapping
from copy import deepcopy
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Literal
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides as parse_dotted_overrides
from fastvideo.api.parser import load_raw_config, parse_config
from fastvideo.api.request_metadata import bind_generation_request_raw
from fastvideo.api.schema import GenerationRequest, GeneratorConfig
@dataclass
class InterleaveRunStateConfig:
instruction: str | None = None
initial_image_path: str | None = None
output_dir: str = "outputs/interleave_run"
trace_path: str | None = None
include_images_in_trace: bool = False
@dataclass
class InterleaveImageBackendConfig:
kind: Literal["fastvideo", "nano_banana"] = "fastvideo"
output_dir: str | None = None
model: str = "gemini-3.1-flash-image"
api_key: str | None = None
base_url: str | None = None
aspect_ratio: str | None = None
image_size: str | None = None
max_attempts: int = 3
retry_delay_s: float = 2.0
@dataclass
class InterleavePlannerConfig:
kind: Literal["single_prompt", "interleave_thinker"] = "single_prompt"
init_from: str | None = None
processor_from: str | None = None
load_backend: bool = True
trainable: bool = False
image_dir: str = ""
torch_dtype: str = "auto"
device_map: Any | None = None
attn_implementation: str | None = None
trust_remote_code: bool = False
use_cache: bool = False
max_prompt_length: int = 16384
max_response_length: int = 4096
lora: dict[str, Any] | None = None
num_generations: int = 1
temperature: float = 0.0
top_p: float = 1.0
max_new_tokens: int | None = None
max_attempts_per_step: int = 2
@dataclass
class InterleaveCriticConfig:
kind: Literal["none", "accept_all", "interleave_thinker"] = "accept_all"
init_from: str | None = None
processor_from: str | None = None
load_backend: bool = True
trainable: bool = False
image_dir: str = ""
torch_dtype: str = "auto"
device_map: Any | None = None
attn_implementation: str | None = None
trust_remote_code: bool = False
use_cache: bool = False
max_prompt_length: int = 16384
max_response_length: int = 4096
lora: dict[str, Any] | None = None
num_generations: int = 1
temperature: float = 0.0
top_p: float = 1.0
max_new_tokens: int | None = None
@dataclass
class InterleaveRunConfig:
interleave: InterleaveRunStateConfig = field(default_factory=InterleaveRunStateConfig)
image_backend: InterleaveImageBackendConfig = field(default_factory=InterleaveImageBackendConfig)
planner: InterleavePlannerConfig = field(default_factory=InterleavePlannerConfig)
critic: InterleaveCriticConfig = field(default_factory=InterleaveCriticConfig)
request: GenerationRequest = field(default_factory=GenerationRequest)
generator: GeneratorConfig | None = None
_INTERLEAVE_RUN_OVERRIDE_PREFIXES = (
"interleave.",
"image_backend.",
"planner.",
"critic.",
"request.",
"generator.",
)
def load_interleave_run_config(
path: str | Path,
*,
overrides: list[str] | None = None,
prompt: str | None = None,
input_image: str | None = None,
output_dir: str | None = None,
trace_path: str | None = None,
require_instruction: bool = True,
) -> InterleaveRunConfig:
raw = load_raw_config(path)
raw = _apply_interleave_runtime_fields(
raw,
prompt=prompt,
input_image=input_image,
output_dir=output_dir,
trace_path=trace_path,
)
raw = _apply_interleave_overrides(raw, overrides)
config = parse_config(InterleaveRunConfig, raw)
bind_generation_request_raw(
config.request,
raw.get("request") if isinstance(raw.get("request"), Mapping) else {},
)
validate_interleave_run_config(
config,
require_instruction=require_instruction,
)
return config
def resolve_interleave_instruction(config: InterleaveRunConfig) -> str:
if config.interleave.instruction:
return config.interleave.instruction
if isinstance(config.request.prompt, str) and config.request.prompt:
return config.request.prompt
if isinstance(config.request.prompt, list) and len(config.request.prompt) == 1:
prompt = config.request.prompt[0]
if isinstance(prompt, str) and prompt:
return prompt
raise ValueError("Interleave config requires interleave.instruction or a single request.prompt")
def validate_interleave_run_config(
config: InterleaveRunConfig,
*,
require_instruction: bool = True,
) -> None:
if require_instruction:
resolve_interleave_instruction(config)
if config.image_backend.kind == "fastvideo" and config.generator is None:
raise ValueError("Interleave config with image_backend.kind=fastvideo requires a generator config")
if config.planner.kind == "interleave_thinker" and config.planner.max_new_tokens is not None:
_require_positive_int(config.planner.max_new_tokens, "planner.max_new_tokens")
_require_positive_int(config.planner.max_attempts_per_step, "planner.max_attempts_per_step")
if config.critic.kind == "interleave_thinker" and config.critic.max_new_tokens is not None:
_require_positive_int(config.critic.max_new_tokens, "critic.max_new_tokens")
def _apply_interleave_runtime_fields(
raw: Mapping[str, Any],
*,
prompt: str | None,
input_image: str | None,
output_dir: str | None,
trace_path: str | None,
) -> dict[str, Any]:
merged = deepcopy(dict(raw))
interleave = merged.setdefault("interleave", {})
if not isinstance(interleave, dict):
raise ValueError("interleave must be a mapping")
if prompt is not None:
interleave["instruction"] = prompt
if input_image is not None:
interleave["initial_image_path"] = input_image
if output_dir is not None:
interleave["output_dir"] = output_dir
if trace_path is not None:
interleave["trace_path"] = trace_path
return merged
def _apply_interleave_overrides(
raw: Mapping[str, Any],
overrides: list[str] | None,
) -> dict[str, Any]:
if not overrides:
return deepcopy(dict(raw))
parsed = parse_dotted_overrides(overrides)
for key in parsed:
if "." not in key:
raise ValueError("Overrides must use dotted config paths like --request.sampling.seed 42")
if not key.startswith(_INTERLEAVE_RUN_OVERRIDE_PREFIXES):
allowed = ", ".join(_INTERLEAVE_RUN_OVERRIDE_PREFIXES)
raise ValueError(f"Unsupported override path {key!r}. Allowed prefixes: {allowed}")
return apply_overrides(raw, parsed)
def _require_positive_int(value: int, path: str) -> None:
if value <= 0:
raise ValueError(f"{path} must be > 0; got {value}")
__all__ = [
"InterleaveCriticConfig",
"InterleaveImageBackendConfig",
"InterleavePlannerConfig",
"InterleaveRunConfig",
"InterleaveRunStateConfig",
"load_interleave_run_config",
"resolve_interleave_instruction",
"validate_interleave_run_config",
]
@@ -0,0 +1,417 @@
# SPDX-License-Identifier: Apache-2.0
"""Prompt-set runner and summary metrics for native Interleave workflows."""
from __future__ import annotations
import json
import re
from collections.abc import Callable, Mapping, Sequence
from copy import deepcopy
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from fastvideo.workflow.interleave_thinker.generator import ImageGeneratorBackend
from fastvideo.workflow.interleave_thinker.runner import (
build_critic,
build_image_backend,
build_planner,
)
from fastvideo.workflow.interleave_thinker.schema import InterleaveTrace
from fastvideo.workflow.interleave_thinker.trace import save_trace
@dataclass(frozen=True)
class InterleavePromptItem:
"""One prompt-set row for end-to-end Interleave evaluation."""
sample_id: str
instruction: str
initial_image_path: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(frozen=True)
class InterleavePromptResult:
sample_id: str
instruction: str
trace_path: str
success: bool
attempts: int
final_image_path: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
resumed: bool = False
@dataclass(frozen=True)
class InterleavePromptSetSummary:
output_dir: str
summary_path: str
num_samples: int
num_success: int
success_rate: float
total_attempts: int
average_attempts: float
num_resumed: int
results: list[InterleavePromptResult]
def load_interleave_prompt_set(path: str | Path) -> list[InterleavePromptItem]:
"""Load prompt rows from JSONL, JSON, or plain text files."""
prompt_path = Path(path)
if not prompt_path.exists():
raise FileNotFoundError(f"Prompt set not found: {prompt_path}")
suffix = prompt_path.suffix.lower()
if suffix == ".jsonl":
raw_items = _load_jsonl(prompt_path)
elif suffix == ".json":
raw_items = _load_json(prompt_path)
elif suffix in {".txt", ".prompts"}:
raw_items = [line.strip() for line in prompt_path.read_text(encoding="utf-8").splitlines() if line.strip()]
else:
raise ValueError(f"Unsupported prompt-set file format: {prompt_path}")
items = [_coerce_prompt_item(raw, index) for index, raw in enumerate(raw_items)]
if not items:
raise ValueError(f"Prompt set is empty: {prompt_path}")
return items
def run_interleave_prompt_set_config(
config: Any,
prompt_set_path: str | Path,
*,
output_dir: str | None = None,
summary_path: str | None = None,
limit: int | None = None,
resume: bool = False,
image_backend: ImageGeneratorBackend | None = None,
) -> InterleavePromptSetSummary:
"""Run a typed Interleave config over a prompt-set file."""
prompt_items = load_interleave_prompt_set(prompt_set_path)
if limit is not None:
if limit <= 0:
raise ValueError(f"limit must be > 0; got {limit}")
prompt_items = prompt_items[:limit]
return run_interleave_prompt_set(
config,
prompt_items,
output_dir=output_dir,
summary_path=summary_path,
resume=resume,
image_backend=image_backend,
)
def run_interleave_prompt_set(
config: Any,
prompt_items: Sequence[InterleavePromptItem],
*,
output_dir: str | None = None,
summary_path: str | None = None,
resume: bool = False,
image_backend: ImageGeneratorBackend | None = None,
) -> InterleavePromptSetSummary:
"""Run multiple Interleave traces while reusing planner/generator/critic backends."""
if not prompt_items:
raise ValueError("prompt_items must not be empty")
run_config = deepcopy(config)
root = Path(output_dir or run_config.interleave.output_dir)
root.mkdir(parents=True, exist_ok=True)
run_config.interleave.output_dir = str(root)
planned_rows = _planned_trace_rows(prompt_items, root)
if resume and all(trace_path.exists() for _, item, trace_path in planned_rows):
resumed_results = [_result_from_saved_trace(item, trace_path) for _, item, trace_path in planned_rows]
resolved_summary_path = Path(summary_path or root / "summary.json")
summary = _build_summary(
resumed_results,
output_dir=root,
summary_path=resolved_summary_path,
)
save_prompt_set_summary(summary, resolved_summary_path)
return summary
cleanup: Callable[[], None] = _noop_cleanup
if image_backend is None:
image_backend, cleanup = build_image_backend(run_config)
try:
orchestrator = _build_prompt_set_orchestrator(run_config, image_backend)
results: list[InterleavePromptResult] = []
for index, item, trace_path in planned_rows:
if resume and trace_path.exists():
results.append(_result_from_saved_trace(item, trace_path))
continue
trace = orchestrator.run(
item.instruction,
initial_image_path=item.initial_image_path or run_config.interleave.initial_image_path,
metadata=_trace_metadata(item, index),
)
trace.metadata.update(_trace_metadata(item, index))
save_trace(
trace,
trace_path,
include_images=run_config.interleave.include_images_in_trace,
)
results.append(_result_from_trace(item, trace, trace_path))
resolved_summary_path = Path(summary_path or root / "summary.json")
summary = _build_summary(
results,
output_dir=root,
summary_path=resolved_summary_path,
)
save_prompt_set_summary(summary, resolved_summary_path)
return summary
finally:
cleanup()
def save_prompt_set_summary(
summary: InterleavePromptSetSummary,
path: str | Path | None = None,
) -> None:
output_path = Path(path or summary.summary_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(
json.dumps(
prompt_set_summary_to_dict(summary),
indent=2,
sort_keys=True,
) + "\n",
encoding="utf-8",
)
def prompt_set_summary_to_dict(summary: InterleavePromptSetSummary) -> dict[str, Any]:
return {
"output_dir": summary.output_dir,
"summary_path": summary.summary_path,
"num_samples": summary.num_samples,
"num_success": summary.num_success,
"success_rate": summary.success_rate,
"total_attempts": summary.total_attempts,
"average_attempts": summary.average_attempts,
"num_resumed": summary.num_resumed,
"results": [_prompt_result_to_dict(result) for result in summary.results],
}
def _build_prompt_set_orchestrator(
config: Any,
image_backend: ImageGeneratorBackend,
) -> Any:
from fastvideo.workflow.interleave_thinker.orchestrator import InterleaveOrchestrator
return InterleaveOrchestrator(
planner=build_planner(config.planner),
generator=image_backend,
critic=build_critic(config.critic),
)
def _noop_cleanup() -> None:
pass
def _load_jsonl(path: Path) -> list[Any]:
rows: list[Any] = []
for line_number, line in enumerate(path.read_text(encoding="utf-8").splitlines(), start=1):
if not line.strip():
continue
try:
rows.append(json.loads(line))
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid JSONL row in {path}:{line_number}: {exc}") from exc
return rows
def _load_json(path: Path) -> list[Any]:
raw = json.loads(path.read_text(encoding="utf-8"))
if isinstance(raw, list):
return raw
if isinstance(raw, Mapping):
for key in ("items", "prompts", "samples"):
value = raw.get(key)
if isinstance(value, list):
return value
return [raw]
raise ValueError(f"{path} must contain a prompt list or mapping")
def _coerce_prompt_item(raw: Any, index: int) -> InterleavePromptItem:
if isinstance(raw, str):
return InterleavePromptItem(
sample_id=f"sample_{index:05d}",
instruction=raw,
)
if not isinstance(raw, Mapping):
raise ValueError(f"Prompt row {index} must be a mapping or string")
instruction = _first_text(raw, "instruction", "prompt", "text")
if not instruction:
raise ValueError(f"Prompt row {index} requires instruction, prompt, or text")
sample_id = _first_text(raw, "id", "sample_id", "name") or f"sample_{index:05d}"
initial_image_path = _first_text(raw, "initial_image_path", "input_image", "image_path", "image")
metadata: dict[str, Any] = {}
raw_metadata = raw.get("metadata")
if isinstance(raw_metadata, Mapping):
metadata.update(dict(raw_metadata))
reserved = {
"id",
"sample_id",
"name",
"instruction",
"prompt",
"text",
"initial_image_path",
"input_image",
"image_path",
"image",
"metadata",
}
for key, value in raw.items():
if key not in reserved:
metadata[str(key)] = value
return InterleavePromptItem(
sample_id=str(sample_id),
instruction=str(instruction),
initial_image_path=str(initial_image_path) if initial_image_path else None,
metadata=metadata,
)
def _first_text(row: Mapping[str, Any], *keys: str) -> str | None:
for key in keys:
value = row.get(key)
if isinstance(value, str) and value:
return value
return None
def _trace_metadata(item: InterleavePromptItem, index: int) -> dict[str, Any]:
return {
"prompt_set_id": item.sample_id,
"prompt_set_index": index,
"prompt_set_metadata": dict(item.metadata),
}
def _result_from_trace(
item: InterleavePromptItem,
trace: InterleaveTrace,
trace_path: Path,
) -> InterleavePromptResult:
return InterleavePromptResult(
sample_id=item.sample_id,
instruction=item.instruction,
trace_path=str(trace_path),
success=trace.success,
attempts=len(trace.attempts),
final_image_path=(trace.final_image.file_path if trace.final_image is not None else None),
metadata=dict(item.metadata),
)
def _result_from_saved_trace(
item: InterleavePromptItem,
trace_path: Path,
) -> InterleavePromptResult:
payload = json.loads(trace_path.read_text(encoding="utf-8"))
final_image = payload.get("final_image")
return InterleavePromptResult(
sample_id=item.sample_id,
instruction=item.instruction,
trace_path=str(trace_path),
success=bool(payload.get("success")),
attempts=len(payload.get("attempts") or []),
final_image_path=(final_image.get("file_path") if isinstance(final_image, Mapping) else None),
metadata=dict(item.metadata),
resumed=True,
)
def _build_summary(
results: Sequence[InterleavePromptResult],
*,
output_dir: Path,
summary_path: Path,
) -> InterleavePromptSetSummary:
num_samples = len(results)
num_success = sum(1 for result in results if result.success)
total_attempts = sum(result.attempts for result in results)
return InterleavePromptSetSummary(
output_dir=str(output_dir),
summary_path=str(summary_path),
num_samples=num_samples,
num_success=num_success,
success_rate=(num_success / num_samples if num_samples else 0.0),
total_attempts=total_attempts,
average_attempts=(total_attempts / num_samples if num_samples else 0.0),
num_resumed=sum(1 for result in results if result.resumed),
results=list(results),
)
def _prompt_result_to_dict(result: InterleavePromptResult) -> dict[str, Any]:
return {
"sample_id": result.sample_id,
"instruction": result.instruction,
"trace_path": result.trace_path,
"success": result.success,
"attempts": result.attempts,
"final_image_path": result.final_image_path,
"metadata": dict(result.metadata),
"resumed": result.resumed,
}
def _planned_trace_rows(
prompt_items: Sequence[InterleavePromptItem],
root: Path,
) -> list[tuple[int, InterleavePromptItem, Path]]:
seen_ids: dict[str, int] = {}
rows: list[tuple[int, InterleavePromptItem, Path]] = []
for index, item in enumerate(prompt_items):
sample_dir = root / _unique_sample_dir_name(item.sample_id, index, seen_ids)
rows.append((index, item, sample_dir / "trace.json"))
return rows
def _unique_sample_dir_name(
sample_id: str,
index: int,
seen_ids: dict[str, int],
) -> str:
base = _safe_sample_id(sample_id) or f"sample_{index:05d}"
count = seen_ids.get(base, 0)
seen_ids[base] = count + 1
if count:
return f"{base}_{count + 1}"
return base
def _safe_sample_id(sample_id: str) -> str:
sanitized = re.sub(r"[^A-Za-z0-9_.-]+", "_", str(sample_id)).strip("._-")
return sanitized[:96]
__all__ = [
"InterleavePromptItem",
"InterleavePromptResult",
"InterleavePromptSetSummary",
"load_interleave_prompt_set",
"prompt_set_summary_to_dict",
"run_interleave_prompt_set",
"run_interleave_prompt_set_config",
"save_prompt_set_summary",
]
@@ -0,0 +1,336 @@
# SPDX-License-Identifier: Apache-2.0
"""FastVideo generator adapter for InterleaveThinker-style image calls."""
from __future__ import annotations
import base64
import io
import os
import time
import uuid
from collections.abc import Mapping
from pathlib import Path
from typing import Any, Protocol
from fastvideo.api.compat import (
explicit_request_updates,
legacy_generate_call_to_request,
normalize_generation_request,
)
from fastvideo.api.results import GenerationResult
from fastvideo.api.schema import GenerationRequest
from fastvideo.workflow.interleave_thinker.schema import (
GeneratedImage,
InterleaveEditRequest,
)
_NANO_BANANA_MODEL_ALIASES = {
"nano-banana": "gemini-2.5-flash-image",
"nano-banana-pro": "gemini-3-pro-image",
"nano-banana-2": "gemini-3.1-flash-image",
}
class ImageGeneratorBackend(Protocol):
"""Minimal image-generation backend used by the Interleave app layer."""
def generate(
self,
request: InterleaveEditRequest,
*,
request_id: str | None = None,
) -> GeneratedImage:
...
class FastVideoImageGeneratorBackend:
"""Translate InterleaveThinker image requests into ``VideoGenerator`` calls."""
def __init__(
self,
generator: Any,
*,
output_dir: str,
default_request: GenerationRequest | Mapping[str, Any] | None = None,
) -> None:
self.generator = generator
self.output_dir = output_dir
self.default_request = normalize_generation_request(default_request) if default_request is not None else None
def generate(
self,
request: InterleaveEditRequest,
*,
request_id: str | None = None,
) -> GeneratedImage:
request_id = request_id or uuid.uuid4().hex
request_output_dir = os.path.join(self.output_dir, "interleave")
upload_dir = os.path.join(self.output_dir, "uploads")
os.makedirs(request_output_dir, exist_ok=True)
input_path = None
if request.image:
os.makedirs(upload_dir, exist_ok=True)
input_path = decode_base64_image_to_path(
request.image,
os.path.join(upload_dir, f"{request_id}_input.png"),
)
output_path = os.path.join(request_output_dir, f"{request_id}.png")
generation_request = self._build_generation_request(
request,
output_path=output_path,
input_image_path=input_path,
)
start = time.perf_counter()
result = self.generator.generate(generation_request)
elapsed = time.perf_counter() - start
result = _first_generation_result(result)
file_path = result.video_path or output_path
if not file_path or not os.path.exists(file_path):
raise RuntimeError(f"FastVideo generation did not produce an image at {file_path!r}")
return GeneratedImage(
prompt=request.prompt,
image_base64=encode_file_to_base64(file_path),
file_path=os.path.abspath(file_path),
inference_time_s=result.generation_time or elapsed,
metadata={
"request_id": request_id,
"input_image_path": input_path,
"peak_memory_mb": result.peak_memory_mb,
},
)
def _build_generation_request(
self,
request: InterleaveEditRequest,
*,
output_path: str,
input_image_path: str | None,
) -> GenerationRequest:
kwargs = {}
if self.default_request is not None:
kwargs.update(_safe_explicit_request_updates(self.default_request))
kwargs.update({
"num_frames": 1,
"fps": 1,
"save_video": True,
"return_frames": False,
"output_path": output_path,
})
if input_image_path is not None:
kwargs["image_path"] = input_image_path
if request.width is not None:
kwargs["width"] = int(request.width)
if request.height is not None:
kwargs["height"] = int(request.height)
if request.seed is not None:
kwargs["seed"] = int(request.seed)
if request.resolved_num_inference_steps() is not None:
kwargs["num_inference_steps"] = int(request.resolved_num_inference_steps())
if request.guidance_scale is not None:
kwargs["guidance_scale"] = float(request.guidance_scale)
if request.true_cfg_scale is not None:
kwargs["true_cfg_scale"] = float(request.true_cfg_scale)
if request.negative_prompt is not None:
kwargs["negative_prompt"] = request.negative_prompt
return legacy_generate_call_to_request(
request.prompt,
None,
legacy_kwargs=kwargs,
)
class NanoBananaImageGeneratorBackend:
"""Google Gemini API image backend for Nano Banana models.
This wraps the closed-source Gemini native-image API behind the same
``ImageGeneratorBackend`` protocol used by Interleave orchestration. The SDK
import and API-key validation are intentionally lazy so
installing FastVideo does not require ``google-genai`` unless this backend is
configured.
"""
def __init__(
self,
*,
model: str = "gemini-3.1-flash-image",
api_key: str | None = None,
base_url: str | None = None,
output_dir: str = "outputs/nano_banana",
aspect_ratio: str | None = None,
image_size: str | None = None,
max_attempts: int = 3,
retry_delay_s: float = 2.0,
) -> None:
self.model = _NANO_BANANA_MODEL_ALIASES.get(model, model)
self.api_key = api_key
self.base_url = base_url
self.output_dir = output_dir
self.aspect_ratio = aspect_ratio
self.image_size = image_size
self.max_attempts = max(1, int(max_attempts))
self.retry_delay_s = float(retry_delay_s)
self._client: Any | None = None
def generate(
self,
request: InterleaveEditRequest,
*,
request_id: str | None = None,
) -> GeneratedImage:
request_id = request_id or uuid.uuid4().hex
output_format = (request.output_format or "png").lower()
if output_format == "jpg":
output_format = "jpeg"
output_path = Path(self.output_dir) / "interleave" / f"{request_id}.{output_format}"
output_path.parent.mkdir(parents=True, exist_ok=True)
contents: list[Any] = [request.prompt]
if request.image:
contents.append(_decode_base64_to_pil(request.image))
last_exc: Exception | None = None
start = time.perf_counter()
for attempt in range(self.max_attempts):
try:
response = self._client_instance().models.generate_content(
model=self.model,
contents=contents,
config=self._make_generate_config(),
)
image = _extract_first_response_image(response)
image.save(output_path)
return GeneratedImage(
prompt=request.prompt,
image_base64=encode_file_to_base64(output_path),
file_path=str(output_path.resolve()),
inference_time_s=time.perf_counter() - start,
metadata={
"request_id": request_id,
"model": self.model,
"attempt": attempt + 1,
},
)
except Exception as exc: # noqa: BLE001 - remote API errors vary by SDK version
last_exc = exc
if attempt + 1 < self.max_attempts:
time.sleep(self.retry_delay_s)
raise RuntimeError(
f"Nano Banana generation failed after {self.max_attempts} attempts: {last_exc}") from last_exc
def _client_instance(self) -> Any:
if self._client is not None:
return self._client
genai, _ = _import_google_genai()
kwargs: dict[str, Any] = {"api_key": _resolve_google_api_key(self.api_key)}
if self.base_url:
kwargs["http_options"] = {"base_url": self.base_url}
self._client = genai.Client(**kwargs)
return self._client
def _make_generate_config(self) -> Any:
_, types = _import_google_genai()
kwargs: dict[str, Any] = {"response_modalities": ["TEXT", "IMAGE"]}
if self.aspect_ratio or self.image_size:
image_kwargs: dict[str, Any] = {}
if self.aspect_ratio:
image_kwargs["aspect_ratio"] = self.aspect_ratio
if self.image_size:
image_kwargs["image_size"] = self.image_size
kwargs["image_config"] = types.ImageConfig(**image_kwargs)
return types.GenerateContentConfig(**kwargs)
def _safe_explicit_request_updates(request: GenerationRequest) -> dict[str, Any]:
try:
return explicit_request_updates(request)
except AssertionError:
return explicit_request_updates(normalize_generation_request(request))
def _first_generation_result(result: GenerationResult | list[GenerationResult]) -> GenerationResult:
if isinstance(result, list):
if not result:
raise RuntimeError("FastVideo generation returned an empty result list")
return result[0]
return result
def encode_file_to_base64(path: str | os.PathLike[str]) -> str:
with open(path, "rb") as handle:
return base64.b64encode(handle.read()).decode("utf-8")
def decode_base64_image_to_path(
image_base64: str,
output_path: str | os.PathLike[str],
) -> str:
payload = image_base64.strip()
if payload.startswith("data:") and "," in payload:
payload = payload.split(",", 1)[1]
data = base64.b64decode(payload)
path = Path(output_path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_bytes(data)
return str(path)
def _resolve_google_api_key(explicit: str | None = None) -> str:
if explicit:
return explicit.strip()
for env_name in ("GEMINI_API_KEY", "GOOGLE_API_KEY"):
value = os.environ.get(env_name)
if value:
return value.strip()
token_path = Path("~/.gemini_token").expanduser()
if token_path.is_file():
return token_path.read_text().strip()
raise ValueError("Google Gemini API access requires GEMINI_API_KEY, GOOGLE_API_KEY, "
"an explicit api_key, or ~/.gemini_token.")
def _import_google_genai() -> tuple[Any, Any]:
try:
from google import genai
from google.genai import types
except ImportError as exc:
raise RuntimeError("Nano Banana API backend requires google-genai. "
"Install google-genai directly or with `uv pip install -e '.[eval-judge]'`.") from exc
return genai, types
def _decode_base64_to_pil(image_base64: str) -> Any:
from PIL import Image
payload = image_base64.strip()
if payload.startswith("data:") and "," in payload:
payload = payload.split(",", 1)[1]
return Image.open(io.BytesIO(base64.b64decode(payload))).convert("RGB")
def _extract_first_response_image(response: Any) -> Any:
from PIL import Image
parts = getattr(response, "parts", None)
if parts is None:
candidates = getattr(response, "candidates", None) or []
if candidates:
parts = getattr(getattr(candidates[0], "content", None), "parts", None)
for part in parts or []:
as_image = getattr(part, "as_image", None)
if callable(as_image):
image = as_image()
if isinstance(image, Image.Image):
return image
inline_data = getattr(part, "inline_data", None) or getattr(part, "inlineData", None)
data = getattr(inline_data, "data", None)
if data:
if isinstance(data, str):
data = base64.b64decode(data)
return Image.open(io.BytesIO(data)).convert("RGB")
raise RuntimeError("Gemini image response did not include an image part")
@@ -0,0 +1,185 @@
# SPDX-License-Identifier: Apache-2.0
"""Provider-based interleaved generation orchestration."""
from __future__ import annotations
from collections.abc import Sequence
from typing import Any, Protocol
from fastvideo.workflow.interleave_thinker.generator import (
ImageGeneratorBackend,
encode_file_to_base64,
)
from fastvideo.workflow.interleave_thinker.schema import (
CriticDecision,
CriticInput,
GeneratedImage,
InterleaveAttempt,
InterleaveEditRequest,
InterleaveTrace,
PlannedInterleaveStep,
PlannerInput,
)
class PlannerProvider(Protocol):
"""Plans a user instruction into concrete generator calls."""
def plan(self, request: PlannerInput) -> Sequence[PlannedInterleaveStep]:
...
class CriticProvider(Protocol):
"""Reviews one generated step and optionally proposes a refined prompt."""
def review(self, request: CriticInput) -> CriticDecision:
...
class SinglePromptPlanner:
"""Fallback planner that runs the instruction as one generator prompt."""
def __init__(self, *, max_attempts: int = 1) -> None:
self.max_attempts = max(1, int(max_attempts))
def plan(self, request: PlannerInput) -> Sequence[PlannedInterleaveStep]:
return [
PlannedInterleaveStep(
prompt=request.instruction,
input_image_path=request.initial_image_path,
max_attempts=self.max_attempts,
)
]
class AcceptAllCritic:
"""Fallback critic for smoke tests and simple generation flows."""
def review(self, request: CriticInput) -> CriticDecision:
del request
return CriticDecision(success=True)
class InterleaveOrchestrator:
"""Run planner -> generator -> critic loops for interleaved workflows."""
def __init__(
self,
*,
planner: PlannerProvider,
generator: ImageGeneratorBackend,
critic: CriticProvider | None = None,
width: int | None = None,
height: int | None = None,
num_inference_steps: int | None = None,
guidance_scale: float | None = None,
seed: int | None = None,
) -> None:
self.planner = planner
self.generator = generator
self.critic = critic
self.width = width
self.height = height
self.num_inference_steps = num_inference_steps
self.guidance_scale = guidance_scale
self.seed = seed
def run(
self,
instruction: str,
*,
initial_image_path: str | None = None,
metadata: dict[str, Any] | None = None,
) -> InterleaveTrace:
planner_input = PlannerInput(
instruction=instruction,
initial_image_path=initial_image_path,
metadata=dict(metadata or {}),
)
planned_steps = list(self.planner.plan(planner_input))
attempts: list[InterleaveAttempt] = []
previous_image_path = initial_image_path
final_image: GeneratedImage | None = None
if not planned_steps:
return InterleaveTrace(
instruction=instruction,
attempts=[],
final_image=None,
success=False,
metadata={"error": "planner returned no steps"},
)
for step_index, step in enumerate(planned_steps):
accepted = False
prompt = step.prompt
step_input_path = step.input_image_path or previous_image_path
max_attempts = max(1, int(step.max_attempts))
for attempt_index in range(max_attempts):
request = self._build_generation_request(
prompt,
input_image_path=step_input_path,
)
generated = self.generator.generate(request)
decision = None
if self.critic is not None:
decision = self.critic.review(
CriticInput(
step=step,
attempt_index=attempt_index,
generated=generated,
previous_image_path=step_input_path,
metadata=dict(step.metadata),
))
attempts.append(
InterleaveAttempt(
step_index=step_index,
attempt_index=attempt_index,
prompt=prompt,
generated=generated,
decision=decision,
))
if decision is None or decision.success:
accepted = True
final_image = generated
previous_image_path = generated.file_path or previous_image_path
break
if decision.refine_prompt:
prompt = decision.refine_prompt
if not accepted:
return InterleaveTrace(
instruction=instruction,
attempts=attempts,
final_image=final_image,
success=False,
metadata={"failed_step_index": step_index},
)
return InterleaveTrace(
instruction=instruction,
attempts=attempts,
final_image=final_image,
success=True,
metadata=dict(metadata or {}),
)
def _build_generation_request(
self,
prompt: str,
*,
input_image_path: str | None,
) -> InterleaveEditRequest:
return InterleaveEditRequest(
prompt=prompt,
image=(encode_file_to_base64(input_image_path) if input_image_path else None),
width=self.width,
height=self.height,
seed=self.seed,
num_inference_steps=self.num_inference_steps,
guidance_scale=self.guidance_scale,
)
@@ -0,0 +1,167 @@
# SPDX-License-Identifier: Apache-2.0
"""Model-backed planner and critic providers for Interleave orchestration."""
from __future__ import annotations
from typing import Any
from fastvideo.workflow.interleave_thinker.schema import (
CriticDecision,
CriticInput,
PlannedInterleaveStep,
PlannerInput,
)
from fastvideo.train.methods.rl.rewards import extract_interleave_answer
from fastvideo.train.models.interleave_thinker import (
InterleavePlannerStep,
InterleaveThinkerCriticModel,
InterleaveThinkerPlannerModel,
)
class InterleaveThinkerPlannerProvider:
"""Adapter from ``InterleaveThinkerPlannerModel`` to ``PlannerProvider``."""
def __init__(
self,
model: InterleaveThinkerPlannerModel,
*,
num_generations: int = 1,
temperature: float = 0.0,
top_p: float = 1.0,
max_new_tokens: int = 2048,
max_attempts_per_step: int = 2,
) -> None:
self.model = model
self.num_generations = int(num_generations)
self.temperature = float(temperature)
self.top_p = float(top_p)
self.max_new_tokens = int(max_new_tokens)
self.max_attempts_per_step = int(max_attempts_per_step)
def plan(
self,
request: PlannerInput,
) -> list[PlannedInterleaveStep]:
image_paths = [request.initial_image_path] if request.initial_image_path else []
raw_plan = self.model.generate_interleave_plan(
request.instruction,
input_image_paths=image_paths,
num_generations=self.num_generations,
temperature=self.temperature,
top_p=self.top_p,
max_new_tokens=self.max_new_tokens,
)
steps = raw_plan.get("steps") or []
planned_steps: list[PlannedInterleaveStep] = []
for idx, step in enumerate(steps):
if not isinstance(step, InterleavePlannerStep):
continue
# ``auxiliary_text`` is a text response channel, not an image prompt.
# Guidance-planner Task A intentionally leaves both image fields
# unset, so skip those unsupported text-only steps instead of
# sending their answer to the image generator.
prompt = step.prompt or step.instruction or ""
if not prompt:
continue
planned_steps.append(
PlannedInterleaveStep(
prompt=prompt,
name=step.step_name,
input_image_path=request.initial_image_path if idx == 0 else None,
max_attempts=max(1, self.max_attempts_per_step),
metadata={
"planner_step_number": step.step_number,
"planner_instruction": step.instruction,
"planner_prompt": step.prompt,
"planner_auxiliary_text": step.auxiliary_text,
"planner_generation_index": raw_plan.get("generation_index"),
},
))
return planned_steps
class InterleaveThinkerCriticProvider:
"""Adapter from ``InterleaveThinkerCriticModel`` to ``CriticProvider``."""
def __init__(
self,
model: InterleaveThinkerCriticModel,
*,
num_generations: int = 1,
temperature: float = 0.0,
top_p: float = 1.0,
max_new_tokens: int = 512,
) -> None:
self.model = model
self.num_generations = int(num_generations)
self.temperature = float(temperature)
self.top_p = float(top_p)
self.max_new_tokens = int(max_new_tokens)
def review(
self,
request: CriticInput,
) -> CriticDecision:
if not request.generated.file_path:
return CriticDecision(
success=False,
reason="InterleaveThinker critic requires generated.file_path",
)
item = _critic_item_from_request(request)
rollouts = self.model.generate_interleave_responses(
{"items": [item]},
num_generations=self.num_generations,
temperature=self.temperature,
top_p=self.top_p,
max_new_tokens=self.max_new_tokens,
)
if not rollouts:
return CriticDecision(
success=False,
reason="InterleaveThinker critic returned no rollouts",
)
response = str(rollouts[0].get("response", "") or "")
parsed = extract_interleave_answer(response)
if parsed is None:
return CriticDecision(
success=False,
reason="InterleaveThinker critic response did not parse",
metadata={"critic_response": response},
)
return CriticDecision(
success=parsed.previous_step_success,
refine_prompt=parsed.refine_prompt,
metadata={"critic_response": response},
)
def _critic_item_from_request(request: CriticInput) -> dict[str, Any]:
instruction = _first_text(
request.step.metadata.get("planner_instruction"),
request.step.metadata.get("planner_prompt"),
request.step.prompt,
)
return {
"origin_prompt": instruction,
"previous_prompt": request.generated.prompt,
"previous_image_path": request.previous_image_path,
"edited_image_path": request.generated.file_path,
"generated_image_path": request.generated.file_path,
"attempt_index": request.attempt_index,
"step_name": request.step.name,
"step_metadata": dict(request.step.metadata),
}
def _first_text(*values: Any) -> str:
for value in values:
if isinstance(value, str) and value:
return value
return ""
__all__ = [
"InterleaveThinkerCriticProvider",
"InterleaveThinkerPlannerProvider",
]
@@ -0,0 +1,199 @@
# SPDX-License-Identifier: Apache-2.0
"""Config-driven native InterleaveThinker orchestration runner."""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from fastvideo.workflow.interleave_thinker.config import (
InterleaveCriticConfig,
InterleaveImageBackendConfig,
InterleavePlannerConfig,
InterleaveRunConfig,
resolve_interleave_instruction,
)
from fastvideo.workflow.interleave_thinker.generator import (
FastVideoImageGeneratorBackend,
ImageGeneratorBackend,
NanoBananaImageGeneratorBackend,
)
from fastvideo.workflow.interleave_thinker.orchestrator import (
AcceptAllCritic,
CriticProvider,
InterleaveOrchestrator,
PlannerProvider,
SinglePromptPlanner,
)
from fastvideo.workflow.interleave_thinker.providers import (
InterleaveThinkerCriticProvider,
InterleaveThinkerPlannerProvider,
)
from fastvideo.workflow.interleave_thinker.schema import InterleaveTrace
from fastvideo.workflow.interleave_thinker.trace import save_trace
from fastvideo.train.models.interleave_thinker import (
InterleaveThinkerCriticModel,
InterleaveThinkerPlannerModel,
)
@dataclass(frozen=True)
class InterleaveRunResult:
trace: InterleaveTrace
trace_path: str
def run_interleave_config(
config: InterleaveRunConfig,
*,
image_backend: ImageGeneratorBackend | None = None,
) -> InterleaveRunResult:
"""Run one native interleaved generation trace from a typed config."""
cleanup: Callable[[], None] = lambda: None
if image_backend is None:
image_backend, cleanup = build_image_backend(config)
try:
orchestrator = InterleaveOrchestrator(
planner=build_planner(config.planner),
generator=image_backend,
critic=build_critic(config.critic),
)
trace = orchestrator.run(
resolve_interleave_instruction(config),
initial_image_path=config.interleave.initial_image_path,
metadata={
"image_backend": config.image_backend.kind,
"planner": config.planner.kind,
"critic": config.critic.kind,
},
)
trace_path = resolve_trace_path(config)
save_trace(
trace,
trace_path,
include_images=config.interleave.include_images_in_trace,
)
return InterleaveRunResult(
trace=trace,
trace_path=str(trace_path),
)
finally:
cleanup()
def resolve_trace_path(config: InterleaveRunConfig) -> Path:
if config.interleave.trace_path:
return Path(config.interleave.trace_path)
return Path(config.interleave.output_dir) / "trace.json"
def build_planner(config: InterleavePlannerConfig) -> PlannerProvider:
if config.kind == "single_prompt":
return SinglePromptPlanner(max_attempts=config.max_attempts_per_step)
model = InterleaveThinkerPlannerModel(**_actor_model_kwargs(config), )
return InterleaveThinkerPlannerProvider(
model,
num_generations=config.num_generations,
temperature=config.temperature,
top_p=config.top_p,
max_new_tokens=config.max_new_tokens or config.max_response_length,
max_attempts_per_step=config.max_attempts_per_step,
)
def build_critic(config: InterleaveCriticConfig) -> CriticProvider | None:
if config.kind == "none":
return None
if config.kind == "accept_all":
return AcceptAllCritic()
model = InterleaveThinkerCriticModel(**_actor_model_kwargs(config), )
return InterleaveThinkerCriticProvider(
model,
num_generations=config.num_generations,
temperature=config.temperature,
top_p=config.top_p,
max_new_tokens=config.max_new_tokens or config.max_response_length,
)
def build_image_backend(config: InterleaveRunConfig) -> tuple[ImageGeneratorBackend, Callable[[], None]]:
image_config = config.image_backend
output_dir = image_config.output_dir or config.interleave.output_dir
if image_config.kind == "nano_banana":
return (
NanoBananaImageGeneratorBackend(
model=image_config.model,
api_key=image_config.api_key,
base_url=image_config.base_url,
output_dir=output_dir,
aspect_ratio=image_config.aspect_ratio,
image_size=image_config.image_size,
max_attempts=image_config.max_attempts,
retry_delay_s=image_config.retry_delay_s,
),
lambda: None,
)
return _build_fastvideo_image_backend(config, image_config, output_dir)
def _build_fastvideo_image_backend(
config: InterleaveRunConfig,
image_config: InterleaveImageBackendConfig,
output_dir: str,
) -> tuple[ImageGeneratorBackend, Callable[[], None]]:
del image_config
if config.generator is None:
raise ValueError("FastVideo image backend requires config.generator")
from fastvideo import VideoGenerator
generator = VideoGenerator.from_config(config.generator)
def cleanup() -> None:
generator.shutdown()
return (
FastVideoImageGeneratorBackend(
generator,
output_dir=output_dir,
default_request=config.request,
),
cleanup,
)
def _actor_model_kwargs(config: InterleavePlannerConfig | InterleaveCriticConfig) -> dict[str, Any]:
kwargs: dict[str, Any] = {
"load_backend": config.load_backend,
"trainable": config.trainable,
"image_dir": config.image_dir,
"torch_dtype": config.torch_dtype,
"device_map": config.device_map,
"attn_implementation": config.attn_implementation,
"trust_remote_code": config.trust_remote_code,
"use_cache": config.use_cache,
"max_prompt_length": config.max_prompt_length,
"max_response_length": config.max_response_length,
"lora": config.lora,
}
if config.init_from is not None:
kwargs["init_from"] = config.init_from
if config.processor_from is not None:
kwargs["processor_from"] = config.processor_from
return kwargs
__all__ = [
"InterleaveRunResult",
"build_critic",
"build_image_backend",
"build_planner",
"resolve_trace_path",
"run_interleave_config",
]
@@ -0,0 +1,107 @@
# SPDX-License-Identifier: Apache-2.0
"""Schemas for InterleaveThinker-style orchestration."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Literal
from pydantic import BaseModel, ConfigDict, Field
class InterleaveEditRequest(BaseModel):
"""Image edit/generation request used by Interleave orchestration backends.
InterleaveThinker sends `num_inference_step` while FastVideo uses
`num_inference_steps`; accept both and let the plural form win when both are
provided. Unknown fields are tolerated so model-specific knobs can pass
through without forcing every backend to implement them immediately.
"""
model_config = ConfigDict(extra="allow")
prompt: str
image: str | None = None
negative_prompt: str | None = None
width: int | None = None
height: int | None = None
seed: int | None = None
num_inference_step: int | None = None
num_inference_steps: int | None = None
guidance_scale: float | None = None
true_cfg_scale: float | None = None
output_format: Literal["png", "jpeg", "jpg", "webp"] | None = "png"
enhance_prompt: bool | None = None
def resolved_num_inference_steps(self) -> int | None:
return self.num_inference_steps if self.num_inference_steps is not None else self.num_inference_step
class InterleaveEditResponse(BaseModel):
success: bool
edited_image: str | None = None
file_path: str | None = None
prompt: str | None = None
inference_time_s: float | None = None
error: str | None = None
metadata: dict[str, Any] = Field(default_factory=dict)
@dataclass(slots=True)
class GeneratedImage:
prompt: str
image_base64: str | None = None
file_path: str | None = None
inference_time_s: float | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class PlannedInterleaveStep:
prompt: str
name: str | None = None
input_image_path: str | None = None
max_attempts: int = 2
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class PlannerInput:
instruction: str
initial_image_path: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class CriticInput:
step: PlannedInterleaveStep
attempt_index: int
generated: GeneratedImage
previous_image_path: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class CriticDecision:
success: bool
refine_prompt: str | None = None
reason: str | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(slots=True)
class InterleaveAttempt:
step_index: int
attempt_index: int
prompt: str
generated: GeneratedImage
decision: CriticDecision | None = None
@dataclass(slots=True)
class InterleaveTrace:
instruction: str
attempts: list[InterleaveAttempt]
final_image: GeneratedImage | None
success: bool
metadata: dict[str, Any] = field(default_factory=dict)
@@ -0,0 +1,102 @@
# SPDX-License-Identifier: Apache-2.0
"""Serialization helpers for interleaved generation traces."""
from __future__ import annotations
import json
from pathlib import Path
from typing import Any
from fastvideo.workflow.interleave_thinker.schema import (
CriticDecision,
GeneratedImage,
InterleaveAttempt,
InterleaveTrace,
)
def trace_to_dict(
trace: InterleaveTrace,
*,
include_images: bool = False,
) -> dict[str, Any]:
return {
"instruction": trace.instruction,
"success": trace.success,
"final_image": _generated_image_to_dict(
trace.final_image,
include_images=include_images,
),
"attempts": [_attempt_to_dict(
attempt,
include_images=include_images,
) for attempt in trace.attempts],
"metadata": dict(trace.metadata),
}
def save_trace(
trace: InterleaveTrace,
path: str | Path,
*,
include_images: bool = False,
) -> None:
output_path = Path(path)
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(
json.dumps(
trace_to_dict(
trace,
include_images=include_images,
),
indent=2,
sort_keys=True,
) + "\n",
encoding="utf-8",
)
def _attempt_to_dict(
attempt: InterleaveAttempt,
*,
include_images: bool,
) -> dict[str, Any]:
return {
"step_index": attempt.step_index,
"attempt_index": attempt.attempt_index,
"prompt": attempt.prompt,
"generated": _generated_image_to_dict(
attempt.generated,
include_images=include_images,
),
"decision": _critic_decision_to_dict(attempt.decision),
}
def _generated_image_to_dict(
image: GeneratedImage | None,
*,
include_images: bool,
) -> dict[str, Any] | None:
if image is None:
return None
result = {
"prompt": image.prompt,
"file_path": image.file_path,
"inference_time_s": image.inference_time_s,
"metadata": dict(image.metadata),
}
if include_images:
result["image_base64"] = image.image_base64
return result
def _critic_decision_to_dict(decision: CriticDecision | None) -> dict[str, Any] | None:
if decision is None:
return None
return {
"success": decision.success,
"refine_prompt": decision.refine_prompt,
"reason": decision.reason,
"metadata": dict(decision.metadata),
}
@@ -0,0 +1,464 @@
# SPDX-License-Identifier: Apache-2.0
"""Trace-level evaluation helpers for Interleave prompt-set outputs."""
from __future__ import annotations
import html
import json
from collections import Counter
from collections.abc import Mapping, Sequence
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, cast
@dataclass(frozen=True)
class InterleaveTraceMetrics:
trace_path: str
instruction: str
success: bool
attempts: int
steps: int
retry_attempts: int
failed_step_index: int | None = None
failure_reason: str | None = None
final_image_path: str | None = None
final_prompt: str | None = None
total_inference_time_s: float | None = None
prompt_set_id: str | None = None
prompt_set_index: int | None = None
metadata: dict[str, Any] = field(default_factory=dict)
@dataclass(frozen=True)
class InterleaveTraceEvaluationSummary:
input_paths: list[str]
num_traces: int
num_success: int
success_rate: float
total_attempts: int
average_attempts: float
total_retry_attempts: int
average_retry_attempts: float
traces_with_final_image: int
total_inference_time_s: float | None
average_inference_time_s: float | None
failure_reasons: dict[str, int]
success_by_category: dict[str, dict[str, float]]
traces: list[InterleaveTraceMetrics]
def discover_interleave_trace_paths(paths: Sequence[str | Path]) -> list[Path]:
"""Discover trace JSON files from trace files, summaries, or output dirs."""
if not paths:
raise ValueError("At least one trace, summary, or output directory is required")
discovered: list[Path] = []
for raw_path in paths:
path = Path(raw_path)
if not path.exists():
raise FileNotFoundError(f"Trace input not found: {path}")
if path.is_dir():
summary_path = path / "summary.json"
if summary_path.is_file():
discovered.extend(_trace_paths_from_summary(summary_path))
else:
discovered.extend(sorted(path.rglob("trace.json")))
continue
if path.name == "summary.json":
discovered.extend(_trace_paths_from_summary(path))
continue
discovered.append(path)
unique: list[Path] = []
seen: set[Path] = set()
for trace_path in discovered:
resolved = trace_path.resolve()
if resolved in seen:
continue
seen.add(resolved)
unique.append(trace_path)
if not unique:
raise ValueError(f"No trace files found in inputs: {[str(path) for path in paths]}")
return unique
def evaluate_interleave_traces(paths: Sequence[str | Path]) -> InterleaveTraceEvaluationSummary:
"""Evaluate saved Interleave traces and return aggregate metrics."""
trace_paths = discover_interleave_trace_paths(paths)
traces = [load_interleave_trace_metrics(path) for path in trace_paths]
num_traces = len(traces)
num_success = sum(1 for trace in traces if trace.success)
total_attempts = sum(trace.attempts for trace in traces)
total_retry_attempts = sum(trace.retry_attempts for trace in traces)
inference_times = [trace.total_inference_time_s for trace in traces if trace.total_inference_time_s is not None]
return InterleaveTraceEvaluationSummary(
input_paths=[str(path) for path in paths],
num_traces=num_traces,
num_success=num_success,
success_rate=(num_success / num_traces if num_traces else 0.0),
total_attempts=total_attempts,
average_attempts=(total_attempts / num_traces if num_traces else 0.0),
total_retry_attempts=total_retry_attempts,
average_retry_attempts=(total_retry_attempts / num_traces if num_traces else 0.0),
traces_with_final_image=sum(1 for trace in traces if trace.final_image_path),
total_inference_time_s=(sum(inference_times) if inference_times else None),
average_inference_time_s=((sum(inference_times) / len(inference_times)) if inference_times else None),
failure_reasons=_failure_reason_counts(traces),
success_by_category=_success_by_category(traces),
traces=traces,
)
def load_interleave_trace_metrics(path: str | Path) -> InterleaveTraceMetrics:
trace_path = Path(path)
payload = _load_json_mapping(trace_path)
attempts = _mapping_list(payload.get("attempts"))
metadata = _string_mapping(payload.get("metadata"))
final_image = _optional_mapping(payload.get("final_image"))
final_image_path = _string_value(final_image.get("file_path")) if final_image is not None else None
final_prompt = _string_value(final_image.get("prompt")) if final_image is not None else None
total_time = _sum_attempt_inference_time(attempts)
return InterleaveTraceMetrics(
trace_path=str(trace_path),
instruction=_string_value(payload.get("instruction")) or "",
success=bool(payload.get("success")),
attempts=len(attempts),
steps=_count_steps(attempts),
retry_attempts=sum(1 for attempt in attempts if _int_value(attempt.get("attempt_index")) not in (None, 0)),
failed_step_index=_int_value(metadata.get("failed_step_index")),
failure_reason=_failure_reason(payload, attempts, metadata),
final_image_path=final_image_path,
final_prompt=final_prompt,
total_inference_time_s=total_time,
prompt_set_id=_string_value(metadata.get("prompt_set_id")),
prompt_set_index=_int_value(metadata.get("prompt_set_index")),
metadata=dict(metadata),
)
def interleave_trace_evaluation_to_dict(summary: InterleaveTraceEvaluationSummary) -> dict[str, Any]:
return {
"input_paths": list(summary.input_paths),
"num_traces": summary.num_traces,
"num_success": summary.num_success,
"success_rate": summary.success_rate,
"total_attempts": summary.total_attempts,
"average_attempts": summary.average_attempts,
"total_retry_attempts": summary.total_retry_attempts,
"average_retry_attempts": summary.average_retry_attempts,
"traces_with_final_image": summary.traces_with_final_image,
"total_inference_time_s": summary.total_inference_time_s,
"average_inference_time_s": summary.average_inference_time_s,
"failure_reasons": dict(summary.failure_reasons),
"success_by_category": dict(summary.success_by_category),
"traces": [_trace_metrics_to_dict(trace) for trace in summary.traces],
}
def write_interleave_trace_evaluation(
summary: InterleaveTraceEvaluationSummary,
output_path: str | Path,
) -> None:
path = Path(output_path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
json.dumps(
interleave_trace_evaluation_to_dict(summary),
indent=2,
sort_keys=True,
) + "\n",
encoding="utf-8",
)
def write_interleave_trace_html_report(
summary: InterleaveTraceEvaluationSummary,
output_path: str | Path,
*,
title: str = "Interleave Trace Evaluation",
) -> None:
path = Path(output_path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(_render_html_report(summary, path.parent, title=title), encoding="utf-8")
def _trace_paths_from_summary(summary_path: Path) -> list[Path]:
payload = _load_json_mapping(summary_path)
results = _mapping_list(payload.get("results"))
trace_paths: list[Path] = []
for result in results:
raw_trace_path = _string_value(result.get("trace_path"))
if not raw_trace_path:
continue
candidate = Path(raw_trace_path)
if not candidate.is_absolute() and not candidate.exists():
candidate = summary_path.parent / candidate
if candidate.is_file():
trace_paths.append(candidate)
return trace_paths
def _load_json_mapping(path: Path) -> Mapping[str, Any]:
payload = json.loads(path.read_text(encoding="utf-8"))
if not isinstance(payload, Mapping):
raise ValueError(f"{path} must contain a JSON object")
return cast(Mapping[str, Any], payload)
def _mapping_list(value: Any) -> list[Mapping[str, Any]]:
if not isinstance(value, list):
return []
rows: list[Mapping[str, Any]] = []
for item in value:
if isinstance(item, Mapping):
rows.append(cast(Mapping[str, Any], item))
return rows
def _optional_mapping(value: Any) -> Mapping[str, Any] | None:
if isinstance(value, Mapping):
return cast(Mapping[str, Any], value)
return None
def _string_mapping(value: Any) -> Mapping[str, Any]:
if isinstance(value, Mapping):
return cast(Mapping[str, Any], value)
return {}
def _string_value(value: Any) -> str | None:
if isinstance(value, str) and value:
return value
return None
def _int_value(value: Any) -> int | None:
if isinstance(value, bool):
return None
if isinstance(value, int):
return value
return None
def _float_value(value: Any) -> float | None:
if isinstance(value, bool):
return None
if isinstance(value, int | float):
return float(value)
return None
def _count_steps(attempts: Sequence[Mapping[str, Any]]) -> int:
step_indices = {_int_value(attempt.get("step_index")) for attempt in attempts}
step_indices.discard(None)
return len(step_indices)
def _sum_attempt_inference_time(attempts: Sequence[Mapping[str, Any]]) -> float | None:
total = 0.0
found = False
for attempt in attempts:
generated = _optional_mapping(attempt.get("generated"))
if generated is None:
continue
value = _float_value(generated.get("inference_time_s"))
if value is None:
continue
total += value
found = True
return total if found else None
def _failure_reason(
payload: Mapping[str, Any],
attempts: Sequence[Mapping[str, Any]],
metadata: Mapping[str, Any],
) -> str | None:
if bool(payload.get("success")):
return None
explicit_error = _string_value(metadata.get("error"))
if explicit_error:
return explicit_error
for attempt in reversed(attempts):
decision = _optional_mapping(attempt.get("decision"))
if decision is None:
continue
reason = _string_value(decision.get("reason"))
if reason:
return reason
failed_step = _int_value(metadata.get("failed_step_index"))
if failed_step is not None:
return f"failed_step_{failed_step}"
return "unknown"
def _failure_reason_counts(traces: Sequence[InterleaveTraceMetrics]) -> dict[str, int]:
counts: Counter[str] = Counter()
for trace in traces:
if trace.success:
continue
counts[trace.failure_reason or "unknown"] += 1
return dict(sorted(counts.items()))
def _success_by_category(traces: Sequence[InterleaveTraceMetrics]) -> dict[str, dict[str, float]]:
grouped: dict[str, list[InterleaveTraceMetrics]] = {}
for trace in traces:
category = _metadata_category(trace.metadata)
if category is None:
continue
grouped.setdefault(category, []).append(trace)
result: dict[str, dict[str, float]] = {}
for category, category_traces in sorted(grouped.items()):
total = len(category_traces)
success = sum(1 for trace in category_traces if trace.success)
result[category] = {
"num_traces": float(total),
"num_success": float(success),
"success_rate": success / total if total else 0.0,
}
return result
def _metadata_category(metadata: Mapping[str, Any]) -> str | None:
prompt_metadata = _optional_mapping(metadata.get("prompt_set_metadata"))
if prompt_metadata is None:
return None
return _string_value(prompt_metadata.get("category"))
def _trace_metrics_to_dict(trace: InterleaveTraceMetrics) -> dict[str, Any]:
return {
"trace_path": trace.trace_path,
"instruction": trace.instruction,
"success": trace.success,
"attempts": trace.attempts,
"steps": trace.steps,
"retry_attempts": trace.retry_attempts,
"failed_step_index": trace.failed_step_index,
"failure_reason": trace.failure_reason,
"final_image_path": trace.final_image_path,
"final_prompt": trace.final_prompt,
"total_inference_time_s": trace.total_inference_time_s,
"prompt_set_id": trace.prompt_set_id,
"prompt_set_index": trace.prompt_set_index,
"metadata": dict(trace.metadata),
}
def _render_html_report(
summary: InterleaveTraceEvaluationSummary,
html_dir: Path,
*,
title: str,
) -> str:
rows = "\n".join(_render_trace_row(trace, html_dir) for trace in summary.traces)
failure_rows = "\n".join(f"<li>{html.escape(reason)}: {count}</li>"
for reason, count in sorted(summary.failure_reasons.items()))
if not failure_rows:
failure_rows = "<li>None</li>"
return f"""<!doctype html>
<html lang="en">
<head>
<meta charset="utf-8">
<title>{html.escape(title)}</title>
<style>
body {{ font-family: system-ui, sans-serif; margin: 24px; color: #1f2933; }}
table {{ border-collapse: collapse; width: 100%; }}
th, td {{ border-bottom: 1px solid #d9e2ec; padding: 8px; text-align: left; vertical-align: top; }}
th {{ background: #f0f4f8; }}
img {{ max-width: 180px; max-height: 120px; object-fit: contain; border: 1px solid #bcccdc; }}
.ok {{ color: #1f7a4d; font-weight: 600; }}
.fail {{ color: #b42318; font-weight: 600; }}
.summary {{ display: flex; gap: 24px; flex-wrap: wrap; margin-bottom: 16px; }}
.metric {{ background: #f8fafc; border: 1px solid #d9e2ec; padding: 10px 12px; }}
</style>
</head>
<body>
<h1>{html.escape(title)}</h1>
<section class="summary">
<div class="metric">Traces: {summary.num_traces}</div>
<div class="metric">Success: {summary.num_success}</div>
<div class="metric">Success rate: {summary.success_rate:.4f}</div>
<div class="metric">Avg attempts: {summary.average_attempts:.2f}</div>
<div class="metric">Avg retries: {summary.average_retry_attempts:.2f}</div>
</section>
<h2>Failure Reasons</h2>
<ul>{failure_rows}</ul>
<h2>Traces</h2>
<table>
<thead>
<tr>
<th>Sample</th>
<th>Status</th>
<th>Attempts</th>
<th>Instruction</th>
<th>Final image</th>
<th>Trace</th>
</tr>
</thead>
<tbody>
{rows}
</tbody>
</table>
</body>
</html>
"""
def _render_trace_row(trace: InterleaveTraceMetrics, html_dir: Path) -> str:
sample = trace.prompt_set_id or Path(trace.trace_path).parent.name
status_class = "ok" if trace.success else "fail"
status_text = "success" if trace.success else f"failed: {trace.failure_reason or 'unknown'}"
image_html = _image_html(trace.final_image_path, html_dir)
trace_link = _path_link(trace.trace_path, html_dir)
return (" <tr>"
f"<td>{html.escape(sample)}</td>"
f"<td class=\"{status_class}\">{html.escape(status_text)}</td>"
f"<td>{trace.attempts} ({trace.retry_attempts} retries)</td>"
f"<td>{html.escape(trace.instruction)}</td>"
f"<td>{image_html}</td>"
f"<td>{trace_link}</td>"
"</tr>")
def _image_html(image_path: str | None, html_dir: Path) -> str:
if not image_path:
return ""
path = Path(image_path)
href = _relative_or_raw_path(path, html_dir)
return f"<a href=\"{html.escape(href)}\"><img src=\"{html.escape(href)}\" alt=\"final image\"></a>"
def _path_link(raw_path: str, html_dir: Path) -> str:
href = _relative_or_raw_path(Path(raw_path), html_dir)
return f"<a href=\"{html.escape(href)}\">trace</a>"
def _relative_or_raw_path(path: Path, base_dir: Path) -> str:
try:
return str(path.resolve().relative_to(base_dir.resolve()))
except ValueError:
try:
return str(path.resolve().relative_to(Path.cwd().resolve()))
except ValueError:
return str(path)
__all__ = [
"InterleaveTraceEvaluationSummary",
"InterleaveTraceMetrics",
"discover_interleave_trace_paths",
"evaluate_interleave_traces",
"interleave_trace_evaluation_to_dict",
"load_interleave_trace_metrics",
"write_interleave_trace_evaluation",
"write_interleave_trace_html_report",
]
+1
View File
@@ -161,6 +161,7 @@ nav:
- Debugging: utilities/debugging.md
- Design:
- Overview: design/overview.md
- InterleaveThinker Integration: design/interleave_thinker.md
- Training Architecture: design/training_architecture.md
- Server Contracts:
- Overview: design/server_contracts/index.md
+3 -1
View File
@@ -224,7 +224,9 @@ skip = "./data,./wandb,ui/package-lock.json,*/_vendored/*"
# "tread" matches daVinci-MagiHuman's acronym "TReAD" (Token Routing and
# Early Drop). codespell lowercases ignore-words entries, so the single
# lowercase form silences all case variants.
ignore-words-list = "tread,passt"
# "redundent" is an upstream InterleaveThinker prompt typo preserved for
# official reference parity.
ignore-words-list = "tread,passt,redundent"
[tool.ruff]
# Allow lines to be as long as 120.
@@ -0,0 +1,180 @@
#!/usr/bin/env bash
set -euo pipefail
# Build a text-only Parquet dataset from a one-prompt-per-line file. The script
# shards prompts across GPU_NUM single-GPU torchrun workers, runs
# v1_preprocess.py with --preprocess_task text_only, and writes prompt
# embeddings/captions under OUTPUT_DIR for DMD2/DiffusionNFT text-only runs.
INPUT_FILE="${1:-train.txt}"
OUTPUT_DIR="${OUTPUT_DIR:-data/train_text_only_dmd_preprocessed}"
MODEL_PATH="${MODEL_PATH:-Wan-AI/Wan2.1-T2V-1.3B-Diffusers}"
GPU_NUM="${GPU_NUM:-2}"
BATCH_SIZE="${BATCH_SIZE:-1}"
SAMPLES_PER_FILE="${SAMPLES_PER_FILE:-8}"
FLUSH_FREQUENCY="${FLUSH_FREQUENCY:-8}"
TEXT_MAX_LENGTH="${TEXT_MAX_LENGTH:-512}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
MIN_FREE_GPU_MB="${MIN_FREE_GPU_MB:-22000}"
if [[ ! -f "$INPUT_FILE" ]]; then
echo "Input text file not found: $INPUT_FILE" >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
visible_gpus=$(nvidia-smi --query-gpu=index --format=csv,noheader 2>/dev/null | wc -l | tr -d ' ')
if [[ "$visible_gpus" -lt "$GPU_NUM" ]]; then
echo "Expected at least $GPU_NUM GPUs, found $visible_gpus" >&2
exit 1
fi
for gpu_id in $(seq 0 $((GPU_NUM - 1))); do
free_mb=$(nvidia-smi --id="$gpu_id" --query-gpu=memory.free --format=csv,noheader,nounits | tr -d ' ')
if [[ "$free_mb" -lt "$MIN_FREE_GPU_MB" ]]; then
echo "GPU $gpu_id has only ${free_mb} MiB free; text-only Wan preprocessing needs" \
"about ${MIN_FREE_GPU_MB} MiB." >&2
echo "Free the GPU or lower MIN_FREE_GPU_MB if you know this run will fit." >&2
exit 1
fi
done
MARKER="$OUTPUT_DIR/.fastvideo_text_only_dmd_output"
if [[ -d "$OUTPUT_DIR" && ! -f "$MARKER" ]]; then
echo "Refusing to overwrite existing non-script output directory: $OUTPUT_DIR" >&2
echo "Set OUTPUT_DIR to a new path or remove the directory manually." >&2
exit 1
fi
rm -rf "$OUTPUT_DIR"
mkdir -p "$OUTPUT_DIR"
touch "$MARKER"
echo "Text-only DMD preprocessing config:"
echo " input: $INPUT_FILE"
echo " output: $OUTPUT_DIR"
echo " model: $MODEL_PATH"
echo " gpus: $GPU_NUM"
echo " batch size per GPU: $BATCH_SIZE"
echo " text max length: $TEXT_MAX_LENGTH"
echo " samples per parquet file: $SAMPLES_PER_FILE"
echo " flush frequency: $FLUSH_FREQUENCY"
SHARD_DIR="$OUTPUT_DIR/_text_shards"
mkdir -p "$SHARD_DIR"
python - "$INPUT_FILE" "$SHARD_DIR" "$GPU_NUM" <<'PY'
from pathlib import Path
import sys
input_path = Path(sys.argv[1])
shard_dir = Path(sys.argv[2])
num_shards = int(sys.argv[3])
prompts = [line.rstrip("\n") for line in input_path.read_text(encoding="utf-8").splitlines() if line.strip()]
if not prompts:
raise SystemExit(f"No non-empty prompts found in {input_path}")
for shard_idx in range(num_shards):
shard_prompts = prompts[shard_idx::num_shards]
shard_path = shard_dir / f"train_text_shard_{shard_idx}.txt"
shard_path.write_text("\n".join(shard_prompts) + "\n", encoding="utf-8")
print(f"Wrote {len(shard_prompts)} prompts to {shard_path}")
PY
run_preprocess_worker() {
local gpu_id="$1"
local shard_file="$2"
local shard_output="$3"
local log_file="$4"
local master_port="$5"
local -a cmd=(
torchrun
--nnodes=1
--nproc_per_node=1
--master_port "$master_port"
fastvideo/pipelines/preprocess/v1_preprocess.py
--model_path "$MODEL_PATH"
--data_merge_path "$shard_file"
--preprocess_video_batch_size "$BATCH_SIZE"
--seed 42
--max_height 448
--max_width 832
--num_frames 77
--dataloader_num_workers 0
--output_dir "$shard_output"
--train_fps 16
--samples_per_file "$SAMPLES_PER_FILE"
--flush_frequency "$FLUSH_FREQUENCY"
--text_max_length "$TEXT_MAX_LENGTH"
--video_length_tolerance_range 5
--preprocess_task text_only
)
{
echo "[gpu${gpu_id}] log file: $log_file"
echo "[gpu${gpu_id}] command: CUDA_VISIBLE_DEVICES=${gpu_id} ${cmd[*]}"
} | tee "$log_file"
CUDA_VISIBLE_DEVICES="$gpu_id" "${cmd[@]}" 2>&1 \
| sed -u "s/^/[gpu${gpu_id}] /" \
| tee -a "$log_file"
local status=${PIPESTATUS[0]}
if [[ "$status" -ne 0 ]]; then
echo "[gpu${gpu_id}] preprocessing failed with exit code $status" | tee -a "$log_file"
fi
return "$status"
}
pids=()
for gpu_id in $(seq 0 $((GPU_NUM - 1))); do
shard_file="$SHARD_DIR/train_text_shard_${gpu_id}.txt"
shard_output="$OUTPUT_DIR/shard_${gpu_id}"
mkdir -p "$shard_output"
log_file="$OUTPUT_DIR/preprocess_gpu_${gpu_id}.log"
echo "Launching text-only preprocessing on GPU ${gpu_id}: ${shard_file}"
run_preprocess_worker "$gpu_id" "$shard_file" "$shard_output" "$log_file" "$((29610 + gpu_id))" &
pids+=("$!")
done
failed=0
for pid in "${pids[@]}"; do
if ! wait "$pid"; then
failed=1
fi
done
if [[ "$failed" -ne 0 ]]; then
echo "One or more preprocessing workers failed. Check logs under $OUTPUT_DIR." >&2
exit 1
fi
num_parquet=$(find "$OUTPUT_DIR" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files were produced under $OUTPUT_DIR" >&2
exit 1
fi
echo "Text-only preprocessing complete."
echo "Parquet files: $num_parquet"
echo "Use this training data_path:"
echo "$OUTPUT_DIR"
+118
View File
@@ -0,0 +1,118 @@
#!/usr/bin/env bash
set -euo pipefail
CONFIG="${CONFIG:-examples/train/configs/rl/wan/diffusion_nft_pick_clip.yaml}"
DATA_PATH="${DATA_PATH:-data/pickscore_text_only_preprocessed}"
OUTPUT_DIR="${OUTPUT_DIR:-outputs/wan2.1_diffusion_nft_pick_clip}"
NUM_GPUS="${NUM_GPUS:-4}"
NNODES="${NNODES:-1}"
NODE_RANK="${NODE_RANK:-0}"
MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
MASTER_PORT="${MASTER_PORT:-29531}"
SP_SIZE="${SP_SIZE:-1}"
TP_SIZE="${TP_SIZE:-1}"
HSDP_REPLICATE_DIM="${HSDP_REPLICATE_DIM:-1}"
HSDP_SHARD_DIM="${HSDP_SHARD_DIM:-$NUM_GPUS}"
DATALOADER_NUM_WORKERS="${DATALOADER_NUM_WORKERS:-0}"
NUM_FRAMES="${NUM_FRAMES:-1}"
NUM_LATENT_T="${NUM_LATENT_T:-1}"
PROJECT_NAME="${PROJECT_NAME:-diffusion_nft_wan}"
RUN_NAME="${RUN_NAME:-wan2.1_diffusion_nft_pick_clip}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
LOG_DIR="${LOG_DIR:-logs/train}"
if [[ ! -f "$CONFIG" ]]; then
echo "Training config not found: $CONFIG" >&2
exit 1
fi
if [[ ! -d "$DATA_PATH" ]]; then
echo "Preprocessed dataset directory not found: $DATA_PATH" >&2
echo "Run preprocessing first, for example:" >&2
echo " GPU_NUM=4 BATCH_SIZE=1 OUTPUT_DIR=$DATA_PATH \\" >&2
echo " bash scripts/preprocess/preprocess_train_text_only_dmd.sh DiffusionNFT/dataset/pickscore/train.txt" >&2
exit 1
fi
num_parquet=$(find "$DATA_PATH" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files found under $DATA_PATH" >&2
echo "Run preprocessing first, for example:" >&2
echo " GPU_NUM=4 BATCH_SIZE=1 OUTPUT_DIR=$DATA_PATH \\" >&2
echo " bash scripts/preprocess/preprocess_train_text_only_dmd.sh DiffusionNFT/dataset/pickscore/train.txt" >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export WANDB_MODE="${WANDB_MODE:-online}"
export WANDB_API_KEY="${WANDB_API_KEY:-}"
export WANDB_BASE_URL="${WANDB_BASE_URL:-https://api.wandb.ai}"
export FASTVIDEO_ATTENTION_BACKEND="${FASTVIDEO_ATTENTION_BACKEND:-FLASH_ATTN}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-/tmp/triton_cache_diffusion_nft_wan}"
export PYTORCH_CUDA_ALLOC_CONF="${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}"
mkdir -p "$LOG_DIR" "$OUTPUT_DIR"
timestamp="$(date +%Y%m%d_%H%M%S)"
log_file="$LOG_DIR/diffusion_nft_wan_pick_clip_${timestamp}.log"
cmd=(
torchrun
--nnodes "$NNODES"
--node_rank "$NODE_RANK"
--nproc_per_node "$NUM_GPUS"
--master_addr "$MASTER_ADDR"
--master_port "$MASTER_PORT"
-m fastvideo.train.entrypoint.train
--config "$CONFIG"
--training.data.data_path "$DATA_PATH"
--training.data.preprocessed_data_type text_only
--training.data.dataloader_num_workers "$DATALOADER_NUM_WORKERS"
--training.data.num_frames "$NUM_FRAMES"
--training.data.num_latent_t "$NUM_LATENT_T"
--training.distributed.num_gpus "$NUM_GPUS"
--training.distributed.sp_size "$SP_SIZE"
--training.distributed.tp_size "$TP_SIZE"
--training.distributed.hsdp_replicate_dim "$HSDP_REPLICATE_DIM"
--training.distributed.hsdp_shard_dim "$HSDP_SHARD_DIM"
--training.checkpoint.output_dir "$OUTPUT_DIR"
--training.tracker.project_name "$PROJECT_NAME"
--training.tracker.run_name "$RUN_NAME"
)
echo "DiffusionNFT Wan single-frame RL training config:"
echo " config: $CONFIG"
echo " data path: $DATA_PATH"
echo " parquet files: $num_parquet"
echo " output dir: $OUTPUT_DIR"
echo " frames / latent T: $NUM_FRAMES / $NUM_LATENT_T"
echo " rewards: pickscore + clipscore"
echo " learning rate: 3e-5"
echo " GPUs: $NUM_GPUS"
echo " SP/TP: $SP_SIZE/$TP_SIZE"
echo " HSDP replicate/shard: $HSDP_REPLICATE_DIM/$HSDP_SHARD_DIM"
echo " W&B mode: $WANDB_MODE"
echo " log file: $log_file"
echo "Command:"
printf ' %q' "${cmd[@]}" "$@"
echo
"${cmd[@]}" "$@" 2>&1 | tee "$log_file"
+170
View File
@@ -0,0 +1,170 @@
#!/usr/bin/env bash
set -euo pipefail
CONFIG="${CONFIG:-examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml}"
DATA_PATH="${DATA_PATH:-data/train_text_only_dmd_preprocessed}"
OUTPUT_DIR="${OUTPUT_DIR:-outputs/wan2.1_dmd2_text_only}"
NUM_GPUS="${NUM_GPUS:-2}"
NNODES="${NNODES:-1}"
NODE_RANK="${NODE_RANK:-0}"
MASTER_ADDR="${MASTER_ADDR:-127.0.0.1}"
MASTER_PORT="${MASTER_PORT:-29521}"
SP_SIZE="${SP_SIZE:-1}"
TP_SIZE="${TP_SIZE:-1}"
HSDP_REPLICATE_DIM="${HSDP_REPLICATE_DIM:-1}"
HSDP_SHARD_DIM="${HSDP_SHARD_DIM:-$NUM_GPUS}"
DATALOADER_NUM_WORKERS="${DATALOADER_NUM_WORKERS:-0}"
NUM_FRAMES="${NUM_FRAMES:-1}"
NUM_LATENT_T="${NUM_LATENT_T:-1}"
VALIDATION_NUM_FRAMES="${VALIDATION_NUM_FRAMES:-$NUM_FRAMES}"
VALIDATION_PROMPT_FILE="${VALIDATION_PROMPT_FILE:-}"
VALIDATION_FILE="${VALIDATION_FILE:-examples/train/configs/distribution_matching/wan/dmd2_text_only_validation.json}"
VALIDATION_OFFLOAD_TRAINING_STATE="${VALIDATION_OFFLOAD_TRAINING_STATE:-true}"
VALIDATION_UNLOAD_PIPELINE_AFTER="${VALIDATION_UNLOAD_PIPELINE_AFTER:-true}"
CFG_UNCOND_TEXT="${CFG_UNCOND_TEXT:-zero}"
CFG_UNCOND_ON_MISSING="${CFG_UNCOND_ON_MISSING:-ignore}"
PROJECT_NAME="${PROJECT_NAME:-distillation_wan_text_only}"
RUN_NAME="${RUN_NAME:-wan2.1_dmd2_text_only}"
CONDA_ROOT="${CONDA_ROOT:-/root/miniconda3}"
CONDA_ENV="${CONDA_ENV:-fastvideo}"
LOG_DIR="${LOG_DIR:-logs/train}"
if [[ ! -f "$CONFIG" ]]; then
echo "Training config not found: $CONFIG" >&2
exit 1
fi
if [[ ! -d "$DATA_PATH" ]]; then
echo "Preprocessed dataset directory not found: $DATA_PATH" >&2
echo "Run scripts/preprocess/preprocess_train_text_only_dmd.sh first." >&2
exit 1
fi
num_parquet=$(find "$DATA_PATH" -name '*.parquet' | wc -l | tr -d ' ')
if [[ "$num_parquet" -eq 0 ]]; then
echo "No parquet files found under $DATA_PATH" >&2
echo "Run scripts/preprocess/preprocess_train_text_only_dmd.sh first." >&2
exit 1
fi
if [[ ! -f "$CONDA_ROOT/etc/profile.d/conda.sh" ]]; then
echo "Conda activation script not found under $CONDA_ROOT" >&2
exit 1
fi
# shellcheck source=/dev/null
source "$CONDA_ROOT/etc/profile.d/conda.sh"
conda activate "$CONDA_ENV"
if [[ "${HF_HUB_ENABLE_HF_TRANSFER:-0}" == "1" ]]; then
if ! python -c "import hf_transfer" >/dev/null 2>&1; then
echo "HF_HUB_ENABLE_HF_TRANSFER=1 but hf_transfer is not installed; disabling fast transfer."
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
else
export HF_HUB_ENABLE_HF_TRANSFER=0
fi
export TOKENIZERS_PARALLELISM="${TOKENIZERS_PARALLELISM:-false}"
export WANDB_MODE="${WANDB_MODE:-offline}"
export WANDB_API_KEY="${WANDB_API_KEY:-}"
export WANDB_BASE_URL="${WANDB_BASE_URL:-https://api.wandb.ai}"
export FASTVIDEO_ATTENTION_BACKEND="${FASTVIDEO_ATTENTION_BACKEND:-FLASH_ATTN}"
export TRITON_CACHE_DIR="${TRITON_CACHE_DIR:-/tmp/triton_cache_dmd2_text_only}"
export PYTORCH_CUDA_ALLOC_CONF="${PYTORCH_CUDA_ALLOC_CONF:-expandable_segments:True}"
mkdir -p "$LOG_DIR" "$OUTPUT_DIR"
if [[ -n "$VALIDATION_PROMPT_FILE" ]]; then
if [[ ! -f "$VALIDATION_PROMPT_FILE" ]]; then
echo "Validation prompt file not found: $VALIDATION_PROMPT_FILE" >&2
echo "Set VALIDATION_PROMPT_FILE to a text file, or leave it empty and set VALIDATION_FILE=<validation.json>." >&2
exit 1
fi
python - "$VALIDATION_PROMPT_FILE" "$VALIDATION_FILE" <<'PY'
import json
import os
import sys
prompt_file, validation_file = sys.argv[1:3]
with open(prompt_file, encoding="utf-8") as f:
prompts = [line.strip() for line in f if line.strip()]
if not prompts:
raise SystemExit(f"No validation prompts found in {prompt_file}")
validation_dir = os.path.dirname(os.path.abspath(validation_file))
if validation_dir:
os.makedirs(validation_dir, exist_ok=True)
with open(validation_file, "w", encoding="utf-8") as f:
json.dump(
{"data": [{"caption": prompt} for prompt in prompts]},
f,
indent=2,
ensure_ascii=False,
)
f.write("\n")
print(f"Wrote {len(prompts)} validation prompts to {validation_file}")
PY
elif [[ ! -f "$VALIDATION_FILE" ]]; then
echo "Validation dataset file not found: $VALIDATION_FILE" >&2
exit 1
fi
timestamp="$(date +%Y%m%d_%H%M%S)"
log_file="$LOG_DIR/dmd2_t2v_text_only_${timestamp}.log"
cmd=(
torchrun
--nnodes "$NNODES"
--node_rank "$NODE_RANK"
--nproc_per_node "$NUM_GPUS"
--master_addr "$MASTER_ADDR"
--master_port "$MASTER_PORT"
-m fastvideo.train.entrypoint.train
--config "$CONFIG"
--training.data.data_path "$DATA_PATH"
--training.data.preprocessed_data_type text_only
--training.data.dataloader_num_workers "$DATALOADER_NUM_WORKERS"
--training.data.num_frames "$NUM_FRAMES"
--training.data.num_latent_t "$NUM_LATENT_T"
--callbacks.validation.dataset_file "$VALIDATION_FILE"
--callbacks.validation.num_frames "$VALIDATION_NUM_FRAMES"
--callbacks.validation.offload_training_state "$VALIDATION_OFFLOAD_TRAINING_STATE"
--callbacks.validation.unload_pipeline_after_validation "$VALIDATION_UNLOAD_PIPELINE_AFTER"
--method.cfg_uncond.text "$CFG_UNCOND_TEXT"
--method.cfg_uncond.on_missing "$CFG_UNCOND_ON_MISSING"
--training.distributed.num_gpus "$NUM_GPUS"
--training.distributed.sp_size "$SP_SIZE"
--training.distributed.tp_size "$TP_SIZE"
--training.distributed.hsdp_replicate_dim "$HSDP_REPLICATE_DIM"
--training.distributed.hsdp_shard_dim "$HSDP_SHARD_DIM"
--training.checkpoint.output_dir "$OUTPUT_DIR"
--training.tracker.project_name "$PROJECT_NAME"
--training.tracker.run_name "$RUN_NAME"
)
echo "DMD2 T2V text-only training config:"
echo " config: $CONFIG"
echo " data path: $DATA_PATH"
echo " parquet files: $num_parquet"
echo " output dir: $OUTPUT_DIR"
echo " frames / latent T: $NUM_FRAMES / $NUM_LATENT_T"
echo " validation prompt file: ${VALIDATION_PROMPT_FILE:-<none>}"
echo " validation dataset: $VALIDATION_FILE"
echo " validation frames: $VALIDATION_NUM_FRAMES"
echo " validation offload training state: $VALIDATION_OFFLOAD_TRAINING_STATE"
echo " validation unload pipeline after: $VALIDATION_UNLOAD_PIPELINE_AFTER"
echo " GPUs: $NUM_GPUS"
echo " SP/TP: $SP_SIZE/$TP_SIZE"
echo " HSDP replicate/shard: $HSDP_REPLICATE_DIM/$HSDP_SHARD_DIM"
echo " CFG uncond text/on_missing: $CFG_UNCOND_TEXT/$CFG_UNCOND_ON_MISSING"
echo " W&B mode: $WANDB_MODE"
echo " log file: $log_file"
echo "Command:"
printf ' %q' "${cmd[@]}" "$@"
echo
"${cmd[@]}" "$@" 2>&1 | tee "$log_file"
@@ -0,0 +1,99 @@
from types import SimpleNamespace
import torch
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
class _FakeEMA:
def __init__(self):
self.updates = 0
def update(self, module):
del module
self.updates += 1
def test_reward_diagnostic_metrics_match_per_prompt_groups():
method = object.__new__(DiffusionNFTMethod)
method._trained_prompt_hashes = set()
sample_items = [{
"prompts": ["a", "a"],
}, {
"prompts": ["b", "b"],
}]
rewards = {"avg": torch.tensor([1.0, 3.0, 2.0, 6.0])}
metrics = method._reward_diagnostic_metrics(sample_items, rewards)
assert metrics["group_size"] == 2.0
assert metrics["trained_prompt_num"] == 2.0
assert torch.isclose(metrics["zero_std_ratio"], torch.tensor(0.0))
assert torch.isclose(metrics["reward_std_mean"], torch.tensor(1.5))
assert torch.isclose(metrics["mean_reward_100"], torch.tensor(3.0))
assert torch.isclose(metrics["mean_reward_50"], torch.tensor(4.5))
method._reward_diagnostic_metrics(sample_items, rewards)
assert len(method._trained_prompt_hashes) == 2
def test_update_ema_honors_update_after_step():
method = object.__new__(DiffusionNFTMethod)
method._ema_enabled = True
method._student_ema = _FakeEMA()
method._ema_update_count = 0
method._ema_update_after_step = 1
method.student = SimpleNamespace(transformer=object())
method._update_ema()
assert method._student_ema.updates == 0
assert method._ema_update_count == 1
method._update_ema()
assert method._student_ema.updates == 1
assert method._ema_update_count == 2
def test_num_train_timesteps_uses_explicit_schedule_length():
method = object.__new__(DiffusionNFTMethod)
method._sample_steps = 25
method._timestep_fraction = 0.5
method._sampling_config = SimpleNamespace(
timesteps=[900, 800, 700, 600, 500, 400, 300, 200, 100, 10],
sigmas=None,
)
assert method._num_train_timesteps() == 5
def test_checkpoint_state_saves_frozen_old_policy_weights():
student = torch.nn.Linear(2, 2)
old = torch.nn.Linear(2, 2)
for param in old.parameters():
param.requires_grad_(False)
method = object.__new__(DiffusionNFTMethod)
method._role_models = {
"student": SimpleNamespace(
transformer=student,
_trainable=True,
),
"old": SimpleNamespace(
transformer=old,
_trainable=False,
),
}
method.student = method._role_models["student"]
method.old = method._role_models["old"]
method._student_optimizer = None
method._student_lr_scheduler = None
method._ema_enabled = False
states = method.checkpoint_state()
old_state = states["roles.old.transformer"].state_dict()
assert "weight" in old_state
assert "bias" in old_state
assert torch.equal(old_state["weight"], old.weight)
@@ -0,0 +1,205 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import base64
from pathlib import Path
from fastvideo.workflow.interleave_thinker.orchestrator import InterleaveOrchestrator
from fastvideo.workflow.interleave_thinker.providers import (
InterleaveThinkerCriticProvider,
InterleaveThinkerPlannerProvider,
)
from fastvideo.workflow.interleave_thinker.schema import (
CriticInput,
GeneratedImage,
PlannedInterleaveStep,
PlannerInput,
)
from fastvideo.train.models.interleave_thinker import InterleavePlannerStep
class _FakePlannerModel:
def __init__(self) -> None:
self.calls = []
def generate_interleave_plan(self, instruction, *, input_image_paths=None, **kwargs):
self.calls.append((instruction, input_image_paths, kwargs))
return {
"generation_index": 0,
"response": "raw planner response",
"steps": [
InterleavePlannerStep(
step_number=1,
step_name="Base",
instruction="Draw the base cat shapes",
prompt="simple cat base shapes",
auxiliary_text=None,
),
InterleavePlannerStep(
step_number=2,
step_name="Color",
instruction="Color the cat",
prompt="color the cat orange",
auxiliary_text="Use warm colors.",
),
],
}
class _FakeAuxiliaryOnlyPlannerModel:
def generate_interleave_plan(self, instruction, *, input_image_paths=None, **kwargs):
del instruction, input_image_paths, kwargs
return {
"generation_index": 0,
"response": "raw auxiliary-only planner response",
"steps": [
InterleavePlannerStep(
step_number=1,
step_name="Answer",
instruction=None,
prompt=None,
auxiliary_text="The requested textual answer.",
),
],
}
class _FakeCriticModel:
def __init__(self, responses):
self.responses = list(responses)
self.calls = []
def generate_interleave_responses(self, batch, **kwargs):
self.calls.append((batch, kwargs))
response = self.responses.pop(0)
return [{"response": response}]
class _FakeGenerator:
def __init__(self) -> None:
self.prompts = []
def generate(self, request, *, request_id=None):
del request_id
self.prompts.append(request.prompt)
file_path = f"/tmp/generated-{len(self.prompts)}.png"
Path(file_path).write_bytes(request.prompt.encode("utf-8"))
return GeneratedImage(
prompt=request.prompt,
image_base64=base64.b64encode(request.prompt.encode("utf-8")).decode("utf-8"),
file_path=file_path,
)
def _critic_response(success: bool, refine_prompt: str = "refined prompt") -> str:
success_text = "true" if success else "false"
return (
"<think>review</think>"
f'<answer>{{"previous_step_success": {success_text}, "refine_prompt": "{refine_prompt}"}}</answer>')
def test_interleave_thinker_planner_provider_converts_model_steps():
model = _FakePlannerModel()
provider = InterleaveThinkerPlannerProvider(model, max_attempts_per_step=3)
steps = provider.plan(PlannerInput(instruction="draw a cat", initial_image_path="/tmp/input.png"))
assert [step.prompt for step in steps] == ["simple cat base shapes", "color the cat orange"]
assert steps[0].input_image_path == "/tmp/input.png"
assert steps[1].input_image_path is None
assert steps[0].max_attempts == 3
assert steps[0].metadata["planner_step_number"] == 1
assert steps[1].metadata["planner_auxiliary_text"] == "Use warm colors."
assert model.calls[0][0] == "draw a cat"
assert model.calls[0][1] == ["/tmp/input.png"]
def test_interleave_thinker_planner_provider_does_not_generate_from_auxiliary_text():
provider = InterleaveThinkerPlannerProvider(_FakeAuxiliaryOnlyPlannerModel())
generator = _FakeGenerator()
steps = provider.plan(PlannerInput(instruction="answer this question"))
trace = InterleaveOrchestrator(
planner=provider,
generator=generator,
).run("answer this question")
assert steps == []
assert trace.success is False
assert trace.metadata["error"] == "planner returned no steps"
assert generator.prompts == []
def test_interleave_thinker_critic_provider_converts_answer_to_decision():
model = _FakeCriticModel([_critic_response(False, "make it clearer")])
provider = InterleaveThinkerCriticProvider(model, max_new_tokens=64)
decision = provider.review(
CriticInput(
step=PlannedInterleaveStep(
prompt="simple cat base shapes",
name="Base",
metadata={"planner_instruction": "Draw the base cat shapes"},
),
attempt_index=0,
generated=GeneratedImage(prompt="refined cat base shapes", file_path="/tmp/after.png"),
previous_image_path="/tmp/before.png",
))
assert decision.success is False
assert decision.refine_prompt == "make it clearer"
batch, kwargs = model.calls[0]
item = batch["items"][0]
assert item["origin_prompt"] == "Draw the base cat shapes"
assert item["previous_prompt"] == "refined cat base shapes"
assert item["previous_image_path"] == "/tmp/before.png"
assert item["edited_image_path"] == "/tmp/after.png"
assert kwargs["max_new_tokens"] == 64
def test_interleave_thinker_critic_provider_handles_unparseable_response():
provider = InterleaveThinkerCriticProvider(_FakeCriticModel(["not parseable"]))
decision = provider.review(
CriticInput(
step=PlannedInterleaveStep(prompt="prompt"),
attempt_index=0,
generated=GeneratedImage(prompt="prompt", file_path="/tmp/after.png"),
))
assert decision.success is False
assert decision.reason == "InterleaveThinker critic response did not parse"
assert decision.metadata["critic_response"] == "not parseable"
def test_interleave_orchestrator_runs_through_model_providers():
planner = InterleaveThinkerPlannerProvider(_FakePlannerModel(), max_attempts_per_step=2)
critic_model = _FakeCriticModel([
_critic_response(False, "better cat base"),
_critic_response(True, "better cat base"),
_critic_response(True, "color the cat orange"),
])
critic = InterleaveThinkerCriticProvider(critic_model)
generator = _FakeGenerator()
orchestrator = InterleaveOrchestrator(
planner=planner,
generator=generator,
critic=critic,
)
trace = orchestrator.run("draw a cat")
assert trace.success is True
assert generator.prompts == ["simple cat base shapes", "better cat base", "color the cat orange"]
assert len(trace.attempts) == 3
assert trace.attempts[0].decision is not None
assert trace.attempts[0].decision.success is False
assert trace.final_image is not None
assert trace.final_image.prompt == "color the cat orange"
critic_prompts = [call[0]["items"][0]["previous_prompt"] for call in critic_model.calls]
assert critic_prompts == ["simple cat base shapes", "better cat base", "color the cat orange"]
@@ -0,0 +1,179 @@
import base64
import io
import sys
import types
from PIL import Image
import pytest
from fastvideo.workflow.interleave_thinker.generator import (
NanoBananaImageGeneratorBackend,
)
from fastvideo.workflow.interleave_thinker.schema import InterleaveEditRequest
from fastvideo.train.methods.rl.rewards.interleave_api import (
GeminiInterleaveImageScorer,
GeminiNanoBananaEditScorer,
)
from fastvideo.train.methods.rl.rewards.interleave_thinker import (
InterleaveThinkerEditRequest,
)
from fastvideo.train.models.interleave_thinker.data import IMAGE_EXTENSIONS
def _png_base64(color="red"):
image = Image.new("RGB", (8, 8), color)
buffer = io.BytesIO()
image.save(buffer, format="PNG")
return base64.b64encode(buffer.getvalue()).decode("utf-8")
def _install_fake_google_genai(monkeypatch):
captured = {
"calls": [],
"configs": [],
"parts": [],
}
class FakeImagePart:
def __init__(self, image):
self._image = image
def as_image(self):
return self._image
class FakePartFactory:
@staticmethod
def from_bytes(data, mime_type):
captured["parts"].append((data, mime_type))
return {
"data": data,
"mime_type": mime_type,
}
class FakeImageConfig:
def __init__(self, **kwargs):
self.kwargs = kwargs
class FakeGenerateContentConfig:
def __init__(self, **kwargs):
self.kwargs = kwargs
captured["configs"].append(kwargs)
class FakeTypes:
GenerateContentConfig = FakeGenerateContentConfig
ImageConfig = FakeImageConfig
Part = FakePartFactory
class FakeModels:
def generate_content(self, **kwargs):
captured["calls"].append(kwargs)
model = kwargs["model"]
if "image" in model:
return types.SimpleNamespace(parts=[FakeImagePart(Image.new("RGB", (8, 8), "blue"))])
return types.SimpleNamespace(text='{"semantic_score": 8.0, "quality_score": 7.0}')
class FakeClient:
def __init__(self, **kwargs):
captured["client_kwargs"] = kwargs
self.models = FakeModels()
google_mod = types.ModuleType("google")
genai_mod = types.ModuleType("google.genai")
genai_mod.Client = FakeClient
genai_mod.types = FakeTypes
google_mod.genai = genai_mod
monkeypatch.setitem(sys.modules, "google", google_mod)
monkeypatch.setitem(sys.modules, "google.genai", genai_mod)
return captured
def test_nano_banana_backend_wraps_google_genai_image_api(monkeypatch, tmp_path):
captured = _install_fake_google_genai(monkeypatch)
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
backend = NanoBananaImageGeneratorBackend(
model="nano-banana-2",
output_dir=str(tmp_path),
aspect_ratio="1:1",
image_size="1K",
)
result = backend.generate(
InterleaveEditRequest(
prompt="make a blue square",
image=_png_base64("red"),
),
request_id="abc",
)
assert result.file_path is not None
assert Image.open(result.file_path).getpixel((0, 0)) == (0, 0, 255)
assert captured["client_kwargs"]["api_key"] == "test-key"
assert captured["calls"][0]["model"] == "gemini-3.1-flash-image"
assert captured["configs"][0]["image_config"].kwargs == {
"aspect_ratio": "1:1",
"image_size": "1K",
}
def test_gemini_nano_banana_edit_scorer_generates_and_scores(monkeypatch, tmp_path):
captured = _install_fake_google_genai(monkeypatch)
monkeypatch.setenv("GEMINI_API_KEY", "test-key")
source = tmp_path / "source.png"
Image.new("RGB", (8, 8), "white").save(source)
scorer = GeminiNanoBananaEditScorer(
image_model="nano-banana",
judge_model="gemini-2.5-pro",
output_dir=str(tmp_path / "reward"),
max_attempts=1,
)
score = scorer(
InterleaveThinkerEditRequest(
index=0,
origin_prompt="draw a blue square",
previous_prompt="blue square",
refine_prompt="make it a clean blue square",
origin_image_path=str(source),
previous_image_path=str(source),
previous_step_success=False,
previous_semantic_score=4.0,
previous_quality_score=5.0,
))
assert score is not None
assert score.semantic_score == pytest.approx(8.0)
assert score.quality_score == pytest.approx(7.0)
assert captured["calls"][0]["model"] == "gemini-2.5-flash-image"
assert captured["calls"][1]["model"] == "gemini-2.5-pro"
assert len(captured["parts"]) == 2
def test_gemini_image_scorer_uses_dataset_supported_mime_types(monkeypatch, tmp_path):
captured = _install_fake_google_genai(monkeypatch)
expected_mime_types = {
".bmp": "image/bmp",
".gif": "image/gif",
".jpeg": "image/jpeg",
".jpg": "image/jpeg",
".png": "image/png",
".tif": "image/tiff",
".tiff": "image/tiff",
".webp": "image/webp",
}
assert set(expected_mime_types) == IMAGE_EXTENSIONS
scorer = GeminiInterleaveImageScorer(max_attempts=1)
for suffix, expected_mime_type in expected_mime_types.items():
image_path = tmp_path / f"image{suffix}"
image_path.write_bytes(b"image-bytes")
part = scorer._image_part(str(image_path))
assert part["mime_type"] == expected_mime_type
assert [mime_type for _, mime_type in captured["parts"]] == list(expected_mime_types.values())
@@ -0,0 +1,270 @@
import json
from pathlib import Path
from types import SimpleNamespace
from PIL import Image
import torch
from fastvideo.train.models.interleave_thinker.critic import (
_PlaceholderActorModule,
)
from fastvideo.train.models.interleave_thinker import (
INTERLEAVE_CRITIC_PROMPT,
InterleaveThinkerCriticModel,
)
from fastvideo.train.utils.training_config import (
DataConfig,
TrainingConfig,
)
class _FakeBackendCritic(InterleaveThinkerCriticModel):
@property
def device(self):
return torch.device("cpu")
class _FakeProcessor:
def __init__(self):
self.tokenizer = SimpleNamespace(pad_token_id=0, eos_token_id=2)
def apply_chat_template(
self,
messages,
*,
tokenize,
add_generation_prompt,
return_dict,
return_tensors,
):
del tokenize, add_generation_prompt, return_dict, return_tensors
has_assistant = any(message["role"] == "assistant" for message in messages)
length = 5 if has_assistant else 3
return {
"input_ids": torch.arange(1, length + 1).unsqueeze(0),
"attention_mask": torch.ones(1, length, dtype=torch.long),
}
def batch_decode(self, sequences, **kwargs):
del kwargs
return [
'<think>ok</think><answer>{"previous_step_success": true, "refine_prompt": "better"}</answer>'
for _ in sequences
]
class _FakeQwen(torch.nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.nn.Parameter(torch.tensor(1.0))
self.generate_kwargs = None
self.last_labels = None
def generate(self, **kwargs):
self.generate_kwargs = kwargs
input_ids = kwargs["input_ids"]
num_return_sequences = int(kwargs.get("num_return_sequences", 1))
suffix = torch.tensor([[9, 10]], dtype=input_ids.dtype)
return torch.cat([input_ids.repeat(num_return_sequences, 1), suffix.repeat(num_return_sequences, 1)], dim=1)
def forward(self, **kwargs):
input_ids = kwargs["input_ids"]
vocab_size = max(16, int(input_ids.max().detach().cpu()) + 1)
logits = torch.zeros(
*input_ids.shape,
vocab_size,
dtype=self.weight.dtype,
device=input_ids.device,
)
for idx in range(input_ids.shape[1] - 1):
next_token = input_ids[:, idx + 1]
logits[:, idx].scatter_(1, next_token[:, None], self.weight.expand(input_ids.shape[0], 1))
labels = kwargs.get("labels")
if labels is not None:
self.last_labels = labels.detach().clone()
trainable_fraction = (labels != -100).float().mean()
return SimpleNamespace(
loss=self.weight.pow(2).sum() * trainable_fraction,
logits=logits,
)
return SimpleNamespace(logits=logits)
def test_interleave_thinker_critic_builds_qwen_vl_messages_without_loading_backend():
model = InterleaveThinkerCriticModel(load_backend=False)
messages = model.build_messages({
"origin_prompt": "draw a vase",
"previous_prompt": "a vase on a table",
"origin_image_path": "before.png",
"edited_image_path": "after.png",
})
content = messages[0]["content"]
images = [part for part in content if part["type"] == "image"]
text = "\n".join(part["text"] for part in content if part["type"] == "text")
assert messages[0]["role"] == "user"
assert images == [{
"type": "image",
"image": "before.png",
}, {
"type": "image",
"image": "after.png",
}]
assert "draw a vase" in text
assert "a vase on a table" in text
def test_interleave_thinker_critic_prefers_previous_image_as_before_image():
model = InterleaveThinkerCriticModel(load_backend=False)
messages = model.build_messages({
"origin_prompt": "draw a vase",
"previous_prompt": "a vase on a table",
"origin_image_path": "original.png",
"previous_image_path": "before.png",
"edited_image_path": "after.png",
})
images = [part for part in messages[0]["content"] if part["type"] == "image"]
assert images == [{
"type": "image",
"image": "before.png",
}, {
"type": "image",
"image": "after.png",
}]
def test_interleave_thinker_critic_materializes_blank_canvas_for_initial_generation(tmp_path):
generated_path = tmp_path / "generated.png"
Image.new("RGB", (12, 8), color="blue").save(generated_path)
model = InterleaveThinkerCriticModel(load_backend=False)
messages = model.build_messages({
"origin_prompt": "draw a vase",
"previous_prompt": "a vase on a table",
"edited_image_path": str(generated_path),
})
images = [part["image"] for part in messages[0]["content"] if part["type"] == "image"]
assert len(images) == 2
assert images[1] == str(generated_path)
assert images[0] != images[1]
blank_canvas_path = Path(images[0])
assert blank_canvas_path.is_file()
with Image.open(blank_canvas_path) as blank_canvas:
assert blank_canvas.mode == "RGB"
assert blank_canvas.size == (12, 8)
assert blank_canvas.getextrema() == ((255, 255), (255, 255), (255, 255))
def test_interleave_thinker_critic_initializes_jsonl_dataloader(tmp_path):
data_path = tmp_path / "critic_rl.jsonl"
data_path.write_text(
json.dumps({
"origin_prompt": "draw a chair",
"rewritten_prompt": "a wooden chair",
"origin_image_path": "interleave/before.png",
"edited_image_path": "interleave/after.png",
"evaluation": {
"success": True,
"semantics": 7.0,
"quality": 8.0,
},
}) + "\n")
image_dir = tmp_path / "images"
model = InterleaveThinkerCriticModel(load_backend=False, image_dir=str(image_dir))
model.init_preprocessors(TrainingConfig(data=DataConfig(data_path=str(data_path), train_batch_size=1)))
batch = next(iter(model.dataloader))
assert batch["items"][0]["origin_prompt"] == "draw a chair"
assert batch["items"][0]["origin_image_path"] == str(image_dir / "interleave/before.png")
assert batch["items"][0]["edited_image_path"] == str(image_dir / "interleave/after.png")
assert batch["items"][0]["evaluation"]["success"] is True
assert "{original_instruction}" in INTERLEAVE_CRITIC_PROMPT
def test_interleave_thinker_critic_fake_backend_generates_rollouts():
model = _FakeBackendCritic(load_backend=False)
model.processor = _FakeProcessor()
model.transformer = _FakeQwen()
rollouts = model.generate_interleave_responses(
{
"items": [{
"origin_prompt": "draw a chair",
"previous_prompt": "a wooden chair",
"origin_image_path": "before.png",
"edited_image_path": "after.png",
}]
},
num_generations=2,
temperature=0.7,
top_p=0.9,
max_new_tokens=8,
)
assert len(rollouts) == 2
assert all("previous_step_success" in rollout["response"] for rollout in rollouts)
assert rollouts[0]["group_key"] == rollouts[1]["group_key"]
assert rollouts[0]["sample_index"] == 0
assert len(rollouts[0]["old_logprobs"]) == 2
assert rollouts[0]["response_mask"] == [1.0, 1.0]
assert model.transformer.generate_kwargs["num_return_sequences"] == 2
assert model.transformer.generate_kwargs["max_new_tokens"] == 8
def test_interleave_thinker_critic_fake_backend_trains_response_tokens_only():
model = _FakeBackendCritic(load_backend=False)
assert isinstance(model.transformer, _PlaceholderActorModule)
model.processor = _FakeProcessor()
model.transformer = _FakeQwen()
optimizer = torch.optim.SGD(model.transformer.parameters(), lr=0.1)
before = float(model.transformer.weight.detach())
loss_map, metrics = model.train_interleave_rollouts(
rollouts=[{
"origin_prompt": "draw a chair",
"previous_prompt": "a wooden chair",
"origin_image_path": "before.png",
"edited_image_path": "after.png",
"response": '<think>ok</think><answer>{"previous_step_success": true, "refine_prompt": "better"}</answer>',
}],
advantages=torch.tensor([1.0]),
rewards={},
optimizer=optimizer,
gradient_accumulation_steps=1,
max_grad_norm=0.0,
)
assert "total_loss" in loss_map
assert metrics["actor/policy_loss"] == -1.0
assert metrics["actor/clipped_fraction"] == 0.0
assert metrics["actor/mean_ratio"] == 1.0
assert metrics["actor/response_tokens"] == 2.0
assert float(model.transformer.weight.detach()) > before
def test_interleave_thinker_critic_reference_logprob_hook_restores_training_state():
model = _FakeBackendCritic(load_backend=False)
model.processor = _FakeProcessor()
model.transformer = _FakeQwen()
model.transformer.train()
rows = model.reference_logprobs_for_interleave_rollouts([{
"origin_prompt": "draw a chair",
"previous_prompt": "a wooden chair",
"origin_image_path": "before.png",
"edited_image_path": "after.png",
"response": '<think>ok</think><answer>{"previous_step_success": true, "refine_prompt": "better"}</answer>',
}])
assert len(rows) == 1
assert len(rows[0]) == 2
assert all(isinstance(value, float) for value in rows[0])
assert model.transformer.training is True
@@ -0,0 +1,165 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
from pathlib import Path
import pytest
from fastvideo.train.models.interleave_thinker import (
load_critic_rl_records,
load_critic_sft_records,
load_interleave_dataset,
load_planner_rl_records,
load_planner_sft_records,
normalize_critic_rl_record,
resolve_interleave_image_path,
)
def test_load_planner_sft_records_normalizes_sharegpt_images(tmp_path: Path) -> None:
image_dir = tmp_path / "images"
data_dir = tmp_path / "data"
data_dir.mkdir()
(data_dir / "planner_sft.json").write_text(
json.dumps([{
"messages": [
{
"role": "user",
"content": "draw a cat step by step"
},
{
"role": "assistant",
"content": "<answer>{\"execution_plan\": []}</answer>"
},
],
"images": ["planner/cat_step.png"],
}]),
encoding="utf-8",
)
records = load_planner_sft_records(data_dir, image_dir=image_dir)
assert len(records) == 1
assert records[0]["instruction"] == "draw a cat step by step"
assert records[0]["response"] == '<answer>{"execution_plan": []}</answer>'
assert records[0]["images"] == [str(image_dir / "planner/cat_step.png")]
assert records[0]["input_image_paths"] == [str(image_dir / "planner/cat_step.png")]
def test_load_planner_rl_records_accepts_prompt_only_rows(tmp_path: Path) -> None:
image_dir = tmp_path / "images"
data_path = tmp_path / "planner_rl.jsonl"
data_path.write_text(
json.dumps({
"text_input": "draw a cat in three clear steps",
"images": ["planner/start.png"],
"plan_score": 0.75,
}) + "\n",
encoding="utf-8",
)
records = load_planner_rl_records(data_path, image_dir=image_dir)
assert records[0]["instruction"] == "draw a cat in three clear steps"
assert records[0]["images"] == [str(image_dir / "planner/start.png")]
assert records[0]["input_image_paths"] == [str(image_dir / "planner/start.png")]
assert records[0]["plan_score"] == 0.75
def test_load_critic_sft_records_adds_image_pair_aliases(tmp_path: Path) -> None:
image_dir = tmp_path / "images"
data_path = tmp_path / "critic_sft.json"
data_path.write_text(
json.dumps({
"records": [{
"messages": [
{
"role": "user",
"content": "evaluate this edit"
},
{
"role": "assistant",
"content": "<answer>{\"previous_step_success\": true, \"refine_prompt\": \"keep it\"}</answer>"
},
],
"images": ["before.png", "after.webp"],
"rewritten_prompt": "make the cat orange",
}]
}),
encoding="utf-8",
)
records = load_critic_sft_records(data_path, image_dir=image_dir)
assert records[0]["origin_image_path"] == str(image_dir / "before.png")
assert records[0]["edited_image_path"] == str(image_dir / "after.webp")
assert records[0]["previous_image_path"] == str(image_dir / "before.png")
assert records[0]["generated_image_path"] == str(image_dir / "after.webp")
assert records[0]["previous_prompt"] == "make the cat orange"
assert records[0]["response"].startswith("<answer>")
def test_load_critic_rl_records_normalizes_reward_fields_and_aliases(tmp_path: Path) -> None:
image_dir = tmp_path / "images"
data_path = tmp_path / "critic_rl.jsonl"
data_path.write_text(
json.dumps({
"origin_prompt": "draw a chair",
"rewritten_prompt": "a wooden chair by a window",
"origin_image_path": "chairs/original.jpg",
"edited_image_path": "chairs/edited.png",
"evaluation": {
"success": False,
"semantic_score": "6.5",
"quality_score": 7,
},
"responses": ["<answer>{}</answer>"],
}) + "\n",
encoding="utf-8",
)
records = load_critic_rl_records(data_path, image_dir=image_dir)
assert records[0]["origin_prompt"] == "draw a chair"
assert records[0]["previous_prompt"] == "a wooden chair by a window"
assert records[0]["rewritten_prompt"] == "a wooden chair by a window"
assert records[0]["origin_image_path"] == str(image_dir / "chairs/original.jpg")
assert records[0]["edited_image_path"] == str(image_dir / "chairs/edited.png")
assert records[0]["previous_image_path"] == str(image_dir / "chairs/original.jpg")
assert records[0]["generated_image_path"] == str(image_dir / "chairs/edited.png")
assert records[0]["ground_truth"] == {
"success": False,
"semantics": 6.5,
"quality": 7.0,
}
assert records[0]["responses"] == ["<answer>{}</answer>"]
def test_load_interleave_dataset_rejects_empty_dataset(tmp_path: Path) -> None:
data_path = tmp_path / "critic_rl.jsonl"
data_path.write_text("\n", encoding="utf-8")
with pytest.raises(ValueError, match="No critic_rl records"):
load_interleave_dataset(data_path, kind="critic_rl")
def test_normalize_critic_rl_record_requires_ground_truth_success() -> None:
with pytest.raises(ValueError, match="ground_truth requires boolean success"):
normalize_critic_rl_record({
"origin_prompt": "draw a chair",
"previous_prompt": "a wooden chair",
"origin_image_path": "before.png",
"edited_image_path": "after.png",
"ground_truth": {
"semantics": 5
},
})
def test_resolve_interleave_image_path_validates_extension(tmp_path: Path) -> None:
assert resolve_interleave_image_path("image.png", image_dir=tmp_path) == str(tmp_path / "image.png")
with pytest.raises(ValueError, match="Unsupported image file extension"):
resolve_interleave_image_path("not-an-image.txt", image_dir=tmp_path)
@@ -0,0 +1,115 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from types import SimpleNamespace
import torch
from fastvideo.train.methods.rl.common.grpo import compute_grpo_loss
from fastvideo.train.models.interleave_thinker import InterleaveThinkerCriticModel
class _FakeBackendCritic(InterleaveThinkerCriticModel):
@property
def device(self):
return torch.device("cpu")
class _FakeProcessor:
def apply_chat_template(
self,
messages,
*,
tokenize,
add_generation_prompt,
return_dict,
return_tensors,
):
del tokenize, add_generation_prompt, return_dict, return_tensors
has_assistant = any(message["role"] == "assistant" for message in messages)
input_ids = [3, 4, 1, 2] if has_assistant else [3, 4]
return {
"input_ids": torch.tensor([input_ids], dtype=torch.long),
"attention_mask": torch.ones(1, len(input_ids), dtype=torch.long),
}
class _FakeQwenPolicy(torch.nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.nn.Parameter(torch.tensor(0.0))
def forward(self, **kwargs):
input_ids = kwargs["input_ids"]
batch, seq_len = input_ids.shape
logits = torch.zeros(batch, seq_len, 8, dtype=self.weight.dtype, device=input_ids.device)
for idx in range(seq_len - 1):
next_token = input_ids[:, idx + 1]
logits[:, idx].scatter_(1, next_token[:, None], self.weight.expand(batch, 1))
return SimpleNamespace(logits=logits)
def test_compute_grpo_loss_clips_ratios_and_masks_tokens():
result = compute_grpo_loss(
current_logprobs=torch.log(torch.tensor([[1.5, 1.0], [0.5, 2.0]])),
old_logprobs=torch.zeros(2, 2),
advantages=torch.tensor([1.0, -1.0]),
response_mask=torch.tensor([[1.0, 1.0], [1.0, 0.0]]),
clip_range=0.2,
)
assert torch.isclose(result.policy_loss, torch.tensor(-1.4 / 3.0), atol=1.0e-6)
assert torch.isclose(result.clipped_fraction, torch.tensor(2.0 / 3.0), atol=1.0e-6)
assert torch.isclose(result.mean_ratio, torch.tensor(1.0), atol=1.0e-6)
assert result.token_count.item() == 3.0
def test_compute_grpo_loss_adds_optional_reference_kl():
reference_logprobs = torch.log(torch.tensor([[0.5, 2.0]]))
result = compute_grpo_loss(
current_logprobs=torch.zeros(1, 2),
old_logprobs=torch.zeros(1, 2),
advantages=torch.zeros(1),
response_mask=torch.ones(1, 2),
reference_logprobs=reference_logprobs,
kl_coef=0.25,
)
expected_kl = ((0.5 - torch.log(torch.tensor(0.5)) - 1.0) +
(2.0 - torch.log(torch.tensor(2.0)) - 1.0)) / 2.0
assert torch.isclose(result.kl_loss, expected_kl, atol=1.0e-6)
assert torch.isclose(result.total_loss, 0.25 * expected_kl, atol=1.0e-6)
def test_critic_grpo_update_uses_response_logprobs_for_positive_advantage():
model = _FakeBackendCritic(load_backend=False, trainable=True)
model.processor = _FakeProcessor()
model.transformer = _FakeQwenPolicy()
optimizer = torch.optim.SGD(model.transformer.parameters(), lr=0.5)
before = float(model.transformer.weight.detach())
loss_map, metrics = model.train_interleave_rollouts(
rollouts=[{
"origin_prompt": "draw a chair",
"previous_prompt": "a wooden chair",
"response": '<answer>{"previous_step_success": true, "refine_prompt": "ok"}</answer>',
}],
advantages=torch.tensor([1.0]),
optimizer=optimizer,
lr_scheduler=None,
gradient_accumulation_steps=1,
clip_range=0.2,
kl_coef=0.0,
update_micro_batch_size=1,
)
assert torch.isclose(loss_map["total_loss"], torch.tensor(-1.0), atol=1.0e-6)
assert metrics["actor/policy_loss"] == -1.0
assert metrics["actor/clipped_fraction"] == 0.0
assert metrics["actor/mean_ratio"] == 1.0
assert metrics["actor/response_tokens"] == 2.0
assert float(model.transformer.weight.detach()) > before
@@ -0,0 +1,324 @@
from types import SimpleNamespace
import pytest
import torch
from fastvideo.train.methods.fine_tuning.finetune import FineTuneMethod
from fastvideo.train.methods.rl.interleave_thinker import InterleaveThinkerRLMethod
from fastvideo.train.methods.rl.rewards import score_interleave_thinker_rewards
from fastvideo.train.models.base import RoleModelBase
from fastvideo.train.utils.config import load_run_config
from fastvideo.train.utils.training_config import (
CheckpointConfig,
DataConfig,
DistributedConfig,
ModelTrainingConfig,
OptimizerConfig,
TrackerConfig,
TrainingConfig,
TrainingLoopConfig,
)
def _response(success=True):
success_text = "true" if success else "false"
return f"""
<think>reasoning</think>
<answer>{{"previous_step_success": {success_text}, "refine_prompt": "improve prompt"}}</answer>
"""
class _FakeInterleaveActor(RoleModelBase):
def __init__(self, *, trainable=True):
super().__init__(trainable=trainable)
self.transformer = torch.nn.Linear(1, 1)
self.generate_calls = []
self.train_calls = []
@property
def device(self):
return torch.device("cpu")
def init_preprocessors(self, training_config):
self.training_config = training_config
def generate_interleave_responses(self, batch, **kwargs):
self.generate_calls.append((batch, kwargs))
rollouts = []
rewards = {
"prompt-a": [0.0, 1.0],
"prompt-b": [0.25, 0.75],
}
for prompt, values in rewards.items():
for generation_idx, reward_value in enumerate(values):
rollouts.append({
"group_key": prompt,
"origin_prompt": prompt,
"previous_prompt": f"{prompt} previous",
"response": _response(success=True),
"ground_truth": {
"success": True,
"semantics": 0.0,
"quality": 0.0,
},
"edit_scores": {
"edited_image_reward_semantic": reward_value,
"edited_image_reward_quality": 0.0,
},
"generation_index": generation_idx,
})
return rollouts
def train_interleave_rollouts(self, **kwargs):
self.train_calls.append(kwargs)
advantages = kwargs["advantages"]
return (
{
"total_loss": advantages.pow(2).mean()
},
{
"actor/updates": 1.0
},
)
class _FakeReferenceActor(_FakeInterleaveActor):
def __init__(self):
super().__init__(trainable=False)
self.reference_calls = []
self.transformer.train()
def reference_logprobs_for_interleave_rollouts(self, rollouts):
self.reference_calls.append([dict(rollout) for rollout in rollouts])
return [[-0.1, -0.2] for _ in rollouts]
def _cfg(method_overrides=None):
method = {
"num_generations": 2,
"num_batches_per_step": 1,
"format_weight": 0.0,
"judge_accuracy_weight": 0.0,
"semantic_weight": 1.0,
"quality_weight": 0.0,
"terminal_progress": False,
}
method.update(method_overrides or {})
return SimpleNamespace(
method=method,
validation={},
training=TrainingConfig(
distributed=DistributedConfig(),
data=DataConfig(seed=123, train_batch_size=1),
optimizer=OptimizerConfig(learning_rate=0.0),
loop=TrainingLoopConfig(max_train_steps=10, gradient_accumulation_steps=3),
checkpoint=CheckpointConfig(),
tracker=TrackerConfig(trackers=[]),
model=ModelTrainingConfig(),
),
)
def test_interleave_thinker_managed_step_scores_advantages_and_calls_actor_update():
actor = _FakeInterleaveActor()
method = InterleaveThinkerRLMethod(
cfg=_cfg({
"clip_range": 0.15,
"kl_coef": 0.05,
"micro_batch_size_per_device_for_update": 2,
}),
role_models={"student": actor},
)
batch = {"origin_prompt": ["prompt-a", "prompt-b"]}
loss_map, outputs, metrics = method.managed_train_step(iter([batch]), iteration=7)
assert outputs == {}
assert torch.isclose(loss_map["total_loss"], torch.tensor(0.9996), atol=1.0e-3)
assert metrics["actor/updates"] == 1.0
assert metrics["interleave/num_rollouts"] == 4.0
assert metrics["interleave/num_groups"] == 2.0
assert torch.isclose(metrics["interleave/reward/overall"], torch.tensor(0.5))
assert len(actor.generate_calls) == 1
_, generate_kwargs = actor.generate_calls[0]
assert generate_kwargs["num_generations"] == 2
assert generate_kwargs["temperature"] == 1.0
assert generate_kwargs["top_p"] == 1.0
train_kwargs = actor.train_calls[0]
advantages = train_kwargs["advantages"]
assert advantages.shape == (4,)
assert torch.isclose(advantages[:2].sum(), torch.tensor(0.0), atol=1.0e-5)
assert torch.isclose(advantages[2:].sum(), torch.tensor(0.0), atol=1.0e-5)
assert train_kwargs["gradient_accumulation_steps"] == 3
assert train_kwargs["clip_range"] == 0.15
assert train_kwargs["kl_coef"] == 0.05
assert train_kwargs["update_micro_batch_size"] == 2
assert train_kwargs["rollouts"][0].get("reference_logprobs") is None
def test_interleave_thinker_namespaces_groups_and_samples_across_input_batches():
actor = _FakeInterleaveActor()
method = InterleaveThinkerRLMethod(
cfg=_cfg({
"num_batches_per_step": 2,
}),
role_models={"student": actor},
)
batch = {"origin_prompt": ["prompt-a", "prompt-b"]}
_, _, metrics = method.managed_train_step(iter([batch, batch]), iteration=4)
assert metrics["interleave/num_rollouts"] == 8.0
assert metrics["interleave/num_groups"] == 4.0
train_rollouts = actor.train_calls[0]["rollouts"]
assert [rollout["sample_index"] for rollout in train_rollouts] == [0, 0, 1, 1, 2, 2, 3, 3]
assert [rollout["group_key"] for rollout in train_rollouts] == [
"batch:0:prompt-a",
"batch:0:prompt-a",
"batch:0:prompt-b",
"batch:0:prompt-b",
"batch:1:prompt-a",
"batch:1:prompt-a",
"batch:1:prompt-b",
"batch:1:prompt-b",
]
def test_interleave_thinker_adapts_exported_callable_reward_rows_to_tensors():
actor = _FakeInterleaveActor()
method = InterleaveThinkerRLMethod(
cfg=_cfg({
"reward_scorer": score_interleave_thinker_rewards,
}),
role_models={"student": actor},
)
_, _, metrics = method.managed_train_step(
iter([{
"origin_prompt": ["prompt-a", "prompt-b"]
}]),
iteration=5,
)
assert torch.isclose(metrics["interleave/reward/overall"], torch.tensor(0.75))
assert torch.isclose(metrics["interleave/reward/format_reward"], torch.tensor(1.0))
def test_interleave_thinker_method_attaches_reference_logprobs_to_rollouts():
actor = _FakeInterleaveActor()
reference = _FakeReferenceActor()
method = InterleaveThinkerRLMethod(
cfg=_cfg({
"kl_coef": 0.01
}),
role_models={
"student": actor,
"reference": reference,
},
)
batch = {"origin_prompt": ["prompt-a", "prompt-b"]}
_, _, metrics = method.managed_train_step(iter([batch]), iteration=3)
assert metrics["interleave/reference_logprob_rollouts"] == 4.0
assert len(reference.reference_calls) == 1
assert len(reference.reference_calls[0]) == 4
assert reference.transformer.training is False
assert all(not param.requires_grad for param in reference.transformer.parameters())
train_rollouts = actor.train_calls[0]["rollouts"]
assert [rollout["reference_logprobs"] for rollout in train_rollouts] == [[-0.1, -0.2]] * 4
def test_interleave_thinker_method_requires_frozen_reference_model():
with pytest.raises(ValueError, match="models.reference.trainable=false"):
InterleaveThinkerRLMethod(
cfg=_cfg(),
role_models={
"student": _FakeInterleaveActor(),
"reference": _FakeInterleaveActor(trainable=True),
},
)
def test_interleave_thinker_method_can_use_offline_response_batches():
actor = _FakeInterleaveActor()
actor.generate_interleave_responses = None
method = InterleaveThinkerRLMethod(cfg=_cfg(), role_models={"student": actor})
batch = {
"origin_prompt": ["prompt-a"],
"ground_truth": [{
"success": True
}],
"responses": [[_response(True), _response(True)]],
"edit_scores": [[{
"edited_image_reward_semantic": 0.0
}, {
"edited_image_reward_semantic": 1.0
}]],
}
loss_map, _, metrics = method.managed_train_step(iter([batch]), iteration=1)
assert "total_loss" in loss_map
assert metrics["interleave/num_rollouts"] == 2.0
assert len(actor.train_calls) == 1
def test_interleave_thinker_method_instantiates_configured_edit_scorer():
actor = _FakeInterleaveActor()
actor.generate_interleave_responses = None
method = InterleaveThinkerRLMethod(
cfg=_cfg({
"edit_scorer": {
"_target_": "fastvideo.train.methods.rl.rewards.ConstantInterleaveEditScorer",
"semantic_reward": 0.9,
"quality_reward": 0.1,
}
}),
role_models={"student": actor},
)
batch = {
"origin_prompt": ["prompt-a"],
"previous_prompt": ["prompt-a previous"],
"ground_truth": [{
"success": True
}],
"response": [_response(True)],
}
_, _, metrics = method.managed_train_step(iter([batch]), iteration=2)
assert torch.isclose(metrics["interleave/reward/edited_image_reward_semantic"], torch.tensor(0.9))
assert torch.isclose(metrics["interleave/reward/edited_image_reward_quality"], torch.tensor(0.1))
def test_interleave_thinker_config_parses_public_yaml():
cfg = load_run_config("examples/train/configs/rl/interleave_thinker/critic_grpo.yaml")
assert cfg.models["student"]["_target_"] == (
"fastvideo.train.models.interleave_thinker.InterleaveThinkerCriticModel")
assert cfg.models["student"]["init_from"] == "InterleaveThinker/Critic-SFT-8B"
assert cfg.models["student"]["dataset_kind"] == "critic_rl"
assert cfg.models["student"]["image_dir"] == "data/InterleaveThinker/Train-Data"
assert cfg.models["student"]["lora"]["enable"] is True
assert cfg.models["reference"]["_target_"] == (
"fastvideo.train.models.interleave_thinker.InterleaveThinkerCriticModel")
assert cfg.models["reference"]["init_from"] == "InterleaveThinker/Critic-SFT-8B"
assert cfg.models["reference"]["trainable"] is False
assert cfg.method["_target_"] == "fastvideo.train.methods.rl.interleave_thinker.InterleaveThinkerRLMethod"
assert cfg.method["edit_scorer"]["_target_"] == (
"fastvideo.train.methods.rl.rewards.GeminiNanoBananaEditScorer")
assert cfg.method["num_generations"] == 8
assert cfg.method["clip_range"] == 0.2
assert cfg.method["kl_coef"] == 0.01
assert cfg.method["micro_batch_size_per_device_for_update"] == 1
assert cfg.training.data.data_path.endswith("critic_rl.jsonl")
assert cfg.training.optimizer.learning_rate == 2.0e-6
def test_existing_finetune_method_uses_default_optimizer_path():
assert FineTuneMethod.manages_optimization(FineTuneMethod.__new__(FineTuneMethod)) is False
@@ -0,0 +1,287 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import gc
import os
from pathlib import Path
import sys
from types import SimpleNamespace
import pytest
_UPSTREAM_ENV = "INTERLEAVETHINKER_UPSTREAM_REPO"
_REAL_PARITY_ENV = "INTERLEAVETHINKER_REAL_PARITY"
pytestmark = pytest.mark.skipif(
not os.environ.get(_UPSTREAM_ENV),
reason=f"{_UPSTREAM_ENV} must point to an official InterleaveThinker checkout",
)
def _upstream_root() -> Path:
root = Path(os.environ[_UPSTREAM_ENV]).expanduser().resolve()
if not (root / "demo_klein.py").is_file():
raise RuntimeError(f"{_UPSTREAM_ENV} does not look like InterleaveThinker: {root}")
if str(root) not in sys.path:
sys.path.insert(0, str(root))
return root
def _strip_line_end_whitespace(text: str) -> str:
return "\n".join(line.rstrip() for line in text.splitlines())
def _reset_torch_rng(seed: int = 0) -> None:
import torch
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def test_prompt_templates_match_official_reference() -> None:
_upstream_root()
from UEval.system import ( # type: ignore[import-not-found]
GUIDANCE_GLOBAL_PROMPT_JSON,
Iterative_T2I_PROMPT_QWEN,
NARRATIVE_PROMPT_JSON,
)
from fastvideo.train.models.interleave_thinker import (
INTERLEAVE_GUIDANCE_PLANNER_PROMPT,
INTERLEAVE_PLANNER_PROMPT,
)
from fastvideo.train.models.interleave_thinker.critic import INTERLEAVE_CRITIC_PROMPT
assert _strip_line_end_whitespace(INTERLEAVE_PLANNER_PROMPT) == _strip_line_end_whitespace(
NARRATIVE_PROMPT_JSON)
assert _strip_line_end_whitespace(INTERLEAVE_GUIDANCE_PLANNER_PROMPT) == _strip_line_end_whitespace(
GUIDANCE_GLOBAL_PROMPT_JSON)
original_instruction = "draw a blue square"
rewritten_prompt = "a crisp blue square centered on a white background"
upstream_critic_prompt = ("<image><image>\n" + Iterative_T2I_PROMPT_QWEN).replace(
"{original_instruction}",
original_instruction,
).replace(
"{rewritten_prompt}",
rewritten_prompt,
)
fastvideo_critic_prompt = INTERLEAVE_CRITIC_PROMPT.format(
original_instruction=original_instruction,
rewritten_prompt=rewritten_prompt,
)
assert _strip_line_end_whitespace(fastvideo_critic_prompt) == _strip_line_end_whitespace(upstream_critic_prompt)
def test_text_image_messages_match_official_demo_constructor() -> None:
_upstream_root()
sys.modules.setdefault("json_repair", SimpleNamespace(loads=lambda text: text))
from demo_klein import SingleSampleGenerator # type: ignore[import-not-found]
from fastvideo.train.models.interleave_thinker import InterleaveThinkerPlannerModel
upstream = SingleSampleGenerator.__new__(SingleSampleGenerator)
fastvideo = InterleaveThinkerPlannerModel(load_backend=False)
cases = [
("plain text prompt", None),
("<image> edit this image", ["one.png"]),
("before <image> middle <image> after", ["one.png", "two.png"]),
("<image><image><image><image><image> summarize the sequence", [
"one.png",
"two.png",
"three.png",
"four.png",
"five.png",
]),
]
for prompt, images in cases:
assert fastvideo.build_text_image_messages(prompt, images) == upstream.construct_msgs(prompt, images)
def test_qwen_response_generation_matches_official_predict_with_fake_backend() -> None:
_upstream_root()
import torch
from UEval.qwen3_vl_api import predict as upstream_predict # type: ignore[import-not-found]
from fastvideo.train.models.interleave_thinker.qwen_actor import Qwen3VLActorBase
class Batch(dict):
def __getattr__(self, name: str):
return self[name]
def to(self, device):
self["device"] = str(device)
return self
class FakeProcessor:
def __init__(self) -> None:
self.tokenizer = SimpleNamespace(pad_token_id=0, eos_token_id=2)
def apply_chat_template(
self,
messages,
*,
tokenize,
add_generation_prompt,
return_dict,
return_tensors,
):
del messages, tokenize, add_generation_prompt, return_dict, return_tensors
return Batch({
"input_ids": torch.tensor([[1, 2, 3]], dtype=torch.long),
"attention_mask": torch.ones(1, 3, dtype=torch.long),
})
def batch_decode(self, sequences, **kwargs):
del kwargs
return [" ".join(str(int(token)) for token in sequence.flatten()) for sequence in sequences]
class FakeQwen(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.device = torch.device("cpu")
def generate(self, **kwargs):
input_ids = kwargs["input_ids"]
num_return_sequences = int(kwargs.get("num_return_sequences", 1))
suffix = torch.tensor([[9, 10]], dtype=input_ids.dtype)
return torch.cat(
[
input_ids.repeat(num_return_sequences, 1),
suffix.repeat(num_return_sequences, 1),
],
dim=1,
)
class FakeActor(Qwen3VLActorBase):
@property
def device(self):
return torch.device("cpu")
messages = [{"role": "user", "content": [{"type": "text", "text": "hello"}]}]
processor = FakeProcessor()
qwen = FakeQwen()
upstream_response = upstream_predict(qwen, processor, messages, max_new_tokens=7)
actor = FakeActor(init_from="unused", load_backend=False, trainable=False)
actor.processor = processor
actor.transformer = qwen
fastvideo_response = actor.generate_qwen_responses(
messages,
num_generations=1,
temperature=0.0,
top_p=1.0,
max_new_tokens=7,
)[0]
assert fastvideo_response == upstream_response
@pytest.mark.skipif(
os.environ.get(_REAL_PARITY_ENV) != "1",
reason=f"{_REAL_PARITY_ENV}=1 is required for real checkpoint parity",
)
def test_real_planner_generation_matches_official_predict() -> None:
_upstream_root()
import torch
from UEval.qwen3_vl_api import predict as upstream_predict # type: ignore[import-not-found]
from fastvideo.train.models.interleave_thinker import InterleaveThinkerPlannerModel
model = InterleaveThinkerPlannerModel(
init_from=os.environ.get("INTERLEAVETHINKER_PLANNER_CKPT", "InterleaveThinker/InterleaveThinker-Planner-8B"),
processor_from=os.environ.get("INTERLEAVETHINKER_PROCESSOR_CKPT", "Qwen/Qwen3-VL-8B-Instruct"),
trainable=False,
torch_dtype=os.environ.get("INTERLEAVETHINKER_TORCH_DTYPE", "auto"),
device_map=os.environ.get("INTERLEAVETHINKER_DEVICE_MAP", "cuda:0"),
attn_implementation=os.environ.get("INTERLEAVETHINKER_ATTN_IMPL", "sdpa"),
)
try:
messages = model.build_messages({"instruction": "How to draw a cat step by step?"})
_reset_torch_rng()
upstream_response = upstream_predict(
model.transformer,
model.processor,
messages,
max_new_tokens=64,
)
_reset_torch_rng()
fastvideo_response = model.generate_qwen_responses(
messages,
num_generations=1,
temperature=0.0,
top_p=1.0,
max_new_tokens=64,
)[0]
assert fastvideo_response == upstream_response
finally:
del model
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
@pytest.mark.skipif(
os.environ.get(_REAL_PARITY_ENV) != "1",
reason=f"{_REAL_PARITY_ENV}=1 is required for real checkpoint parity",
)
def test_real_critic_generation_matches_official_predict(tmp_path: Path) -> None:
_upstream_root()
import torch
from PIL import Image
from UEval.qwen3_vl_api import predict as upstream_predict # type: ignore[import-not-found]
from fastvideo.train.models.interleave_thinker import InterleaveThinkerCriticModel
before_path = tmp_path / "before.png"
after_path = tmp_path / "after.png"
Image.new("RGB", (128, 128), "white").save(before_path)
Image.new("RGB", (128, 128), "blue").save(after_path)
model = InterleaveThinkerCriticModel(
init_from=os.environ.get("INTERLEAVETHINKER_CRITIC_CKPT", "InterleaveThinker/Critic-SFT-8B"),
processor_from=os.environ.get("INTERLEAVETHINKER_PROCESSOR_CKPT", "Qwen/Qwen3-VL-8B-Instruct"),
trainable=False,
torch_dtype=os.environ.get("INTERLEAVETHINKER_TORCH_DTYPE", "auto"),
device_map=os.environ.get("INTERLEAVETHINKER_DEVICE_MAP", "cuda:0"),
attn_implementation=os.environ.get("INTERLEAVETHINKER_ATTN_IMPL", "sdpa"),
)
try:
messages = model.build_messages({
"origin_prompt": "draw a blue square",
"previous_prompt": "a crisp blue square centered on a white background",
"previous_image_path": str(before_path),
"edited_image_path": str(after_path),
})
_reset_torch_rng()
upstream_response = upstream_predict(
model.transformer,
model.processor,
messages,
max_new_tokens=64,
)
_reset_torch_rng()
fastvideo_response = model.generate_qwen_responses(
messages,
num_generations=1,
temperature=0.0,
top_p=1.0,
max_new_tokens=64,
)[0]
assert fastvideo_response == upstream_response
finally:
del model
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
@@ -0,0 +1,277 @@
from types import SimpleNamespace
import torch
from fastvideo.train.methods.fine_tuning import InterleaveThinkerSFTMethod
from fastvideo.train.models.interleave_thinker import (
INTERLEAVE_GUIDANCE_PLANNER_PROMPT,
INTERLEAVE_PLANNER_PROMPT,
InterleavePlannerStep,
InterleaveThinkerPlannerModel,
extract_interleave_plan,
)
from fastvideo.train.models.interleave_thinker.qwen_actor import (
_PlaceholderActorModule,
)
from fastvideo.train.models.base import ModelBase, RoleModelBase
from fastvideo.train.utils.builder import build_from_config
from fastvideo.train.utils.config import load_run_config
class _FakeBackendPlanner(InterleaveThinkerPlannerModel):
@property
def device(self):
return torch.device("cpu")
class _FakeProcessor:
def __init__(self):
self.tokenizer = SimpleNamespace(pad_token_id=0, eos_token_id=2)
def apply_chat_template(
self,
messages,
*,
tokenize,
add_generation_prompt,
return_dict,
return_tensors,
):
del tokenize, add_generation_prompt, return_dict, return_tensors
assert messages[0]["role"] == "user"
has_assistant = any(message["role"] == "assistant" for message in messages)
length = 5 if has_assistant else 3
return {
"input_ids": torch.arange(1, length + 1).unsqueeze(0),
"attention_mask": torch.ones(1, length, dtype=torch.long),
}
def batch_decode(self, sequences, **kwargs):
del kwargs
return [
"""
<think>plan</think>
<answer>
{"execution_plan": [
{"step_number": 1, "step_name": "Sketch", "instruction": "Draw a cat", "prompt": "a clean cat sketch", "auxiliary_text": null}
]}
</answer>
"""
for _ in sequences
]
class _FakeQwen(torch.nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.nn.Parameter(torch.tensor(1.0))
self.generate_kwargs = None
def generate(self, **kwargs):
self.generate_kwargs = kwargs
input_ids = kwargs["input_ids"]
num_return_sequences = int(kwargs.get("num_return_sequences", 1))
suffix = torch.tensor([[9, 10]], dtype=input_ids.dtype)
return torch.cat([input_ids.repeat(num_return_sequences, 1), suffix.repeat(num_return_sequences, 1)], dim=1)
def forward(self, **kwargs):
input_ids = kwargs["input_ids"]
vocab_size = max(16, int(input_ids.max().detach().cpu()) + 1)
logits = torch.zeros(
*input_ids.shape,
vocab_size,
dtype=self.weight.dtype,
device=input_ids.device,
)
for idx in range(input_ids.shape[1] - 1):
next_token = input_ids[:, idx + 1]
logits[:, idx].scatter_(1, next_token[:, None], self.weight.expand(input_ids.shape[0], 1))
return SimpleNamespace(logits=logits)
def test_extract_interleave_plan_accepts_json_answer_block():
parsed = extract_interleave_plan("""
<think>ok</think>
<answer>
{"execution_plan": [
{"step_number": 1, "step_name": "Base", "instruction": "Draw a cube", "prompt": "a blue cube", "auxiliary_text": null}
]}
</answer>
""")
assert parsed is not None
assert parsed.steps == (InterleavePlannerStep(
step_number=1,
step_name="Base",
instruction="Draw a cube",
prompt="a blue cube",
auxiliary_text=None,
), )
def test_extract_interleave_plan_accepts_upstream_python_literal_answer_block():
parsed = extract_interleave_plan("""
<answer>
{'execution_plan': [
{'step_number': '2', 'step_name': 'Color', 'instruction': 'Color the cube', 'prompt': 'make the cube red', 'auxiliary_text': None}
]}
</answer>
""")
assert parsed is not None
assert parsed.steps[0].step_number == 2
assert parsed.steps[0].step_name == "Color"
assert parsed.steps[0].auxiliary_text is None
def test_interleave_thinker_planner_builds_text_only_messages_without_backend():
model = InterleaveThinkerPlannerModel(load_backend=False)
assert isinstance(model, RoleModelBase)
assert not isinstance(model, ModelBase)
messages = model.build_messages({"instruction": "Show how to draw a cat step by step."})
content = messages[0]["content"]
images = [part for part in content if part["type"] == "image"]
text = "\n".join(part["text"] for part in content if part["type"] == "text")
assert images == []
assert "Show how to draw a cat step by step." in text
assert "execution_plan" in text
assert "{text_input}" in INTERLEAVE_PLANNER_PROMPT
def test_interleave_thinker_planner_builds_image_conditioned_messages_without_backend():
model = InterleaveThinkerPlannerModel(load_backend=False)
messages = model.build_messages({
"instruction": "Continue this process in two steps.",
"input_image_paths": ["step1.png", "step2.png"],
})
content = messages[0]["content"]
images = [part for part in content if part["type"] == "image"]
text = "\n".join(part["text"] for part in content if part["type"] == "text")
assert images == [{
"type": "image",
"image": "step1.png",
}, {
"type": "image",
"image": "step2.png",
}]
assert "Continue this process in two steps." in text
assert "Multimodal Sequence Planner" in text
assert "{text_input}" in INTERLEAVE_GUIDANCE_PLANNER_PROMPT
def test_interleave_thinker_planner_does_not_concatenate_image_aliases():
model = InterleaveThinkerPlannerModel(load_backend=False)
messages = model.build_messages({
"instruction": "Continue this process in two steps.",
"input_image_paths": ["step1.png", "step2.png"],
"image_paths": ["step1.png", "step2.png"],
"images": ["step1.png", "step2.png"],
})
images = [part["image"] for part in messages[0]["content"] if part["type"] == "image"]
assert images == ["step1.png", "step2.png"]
def test_interleave_thinker_planner_fake_backend_generates_parseable_plan():
model = _FakeBackendPlanner(load_backend=False)
assert isinstance(model.transformer, _PlaceholderActorModule)
model.processor = _FakeProcessor()
model.transformer = _FakeQwen()
plans = model.generate_interleave_plans(
{
"items": [{
"instruction": "Show how to draw a cat step by step."
}]
},
num_generations=1,
temperature=0.7,
top_p=0.9,
max_new_tokens=12,
)
assert len(plans) == 1
assert plans[0]["plan"] is not None
assert plans[0]["steps"][0].prompt == "a clean cat sketch"
assert model.transformer.generate_kwargs["max_new_tokens"] == 12
assert model.transformer.generate_kwargs.get("num_return_sequences", 1) == 1
def test_interleave_thinker_planner_fake_backend_generates_rl_rollouts():
model = _FakeBackendPlanner(load_backend=False)
model.processor = _FakeProcessor()
model.transformer = _FakeQwen()
rollouts = model.generate_interleave_responses(
{
"items": [{
"instruction": "Show how to draw a cat step by step."
}]
},
num_generations=2,
temperature=0.7,
top_p=0.9,
max_new_tokens=12,
)
assert len(rollouts) == 2
assert all(rollout["plan"] is not None for rollout in rollouts)
assert rollouts[0]["group_key"] == rollouts[1]["group_key"]
assert rollouts[0]["group_key"] == "Show how to draw a cat step by step."
assert len(rollouts[0]["old_logprobs"]) == 2
assert rollouts[0]["response_mask"] == [1.0, 1.0]
assert model.transformer.generate_kwargs["num_return_sequences"] == 2
def test_interleave_thinker_planner_config_parses_public_yaml():
cfg = load_run_config("examples/train/configs/interleave_thinker/planner_smoke.yaml")
assert cfg.models["student"]["_target_"] == (
"fastvideo.train.models.interleave_thinker.InterleaveThinkerPlannerModel")
assert cfg.models["student"]["init_from"] == "InterleaveThinker/InterleaveThinker-Planner-8B"
assert cfg.models["student"]["processor_from"] == "Qwen/Qwen3-VL-8B-Instruct"
assert cfg.models["student"]["trainable"] is True
assert cfg.method["_target_"] == "fastvideo.train.methods.fine_tuning.InterleaveThinkerSFTMethod"
assert cfg.training.optimizer.learning_rate == 1.0e-5
def test_interleave_thinker_planner_smoke_config_builds_actor_method(monkeypatch):
def fake_load_backend(self, **kwargs):
del self, kwargs
return _FakeProcessor(), torch.nn.Linear(1, 1)
monkeypatch.setattr(InterleaveThinkerPlannerModel, "_load_backend", fake_load_backend)
cfg = load_run_config("examples/train/configs/interleave_thinker/planner_smoke.yaml")
_, method, dataloader, start_step = build_from_config(cfg)
assert isinstance(method, InterleaveThinkerSFTMethod)
assert dataloader is None
assert start_step == 0
def test_interleave_thinker_planner_grpo_config_parses_public_yaml():
cfg = load_run_config("examples/train/configs/rl/interleave_thinker/planner_grpo.yaml")
assert cfg.models["student"]["_target_"] == (
"fastvideo.train.models.interleave_thinker.InterleaveThinkerPlannerModel")
assert cfg.models["student"]["dataset_kind"] == "planner_rl"
assert cfg.models["student"]["lora"]["enable"] is True
assert cfg.models["reference"]["_target_"] == (
"fastvideo.train.models.interleave_thinker.InterleaveThinkerPlannerModel")
assert cfg.models["reference"]["trainable"] is False
assert cfg.method["_target_"] == "fastvideo.train.methods.rl.interleave_thinker.InterleaveThinkerRLMethod"
assert cfg.method["reward_scorer"]["_target_"] == (
"fastvideo.train.methods.rl.rewards.InterleavePlannerRewardScorer")
assert cfg.method["kl_coef"] == 0.01
assert cfg.training.data.data_path.endswith("planner_rl.jsonl")
@@ -0,0 +1,379 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
from types import SimpleNamespace
import pytest
import torch
from fastvideo.train.models.interleave_thinker import qwen_actor as qwen_actor_module
from fastvideo.train.models.interleave_thinker.qwen_actor import (
Qwen3VLActorBase,
_distributed_token_gradient_scale,
_qwen_sharding_root,
_qwen_transformer_block_condition,
_rank_independent_rng,
_resolve_hsdp_dimensions,
)
from fastvideo.train.utils.config import load_run_config
from fastvideo.train.utils.training_config import (
DataConfig,
DistributedConfig,
TrainingConfig,
)
class _FakeActor(Qwen3VLActorBase):
@property
def device(self):
return torch.device("cpu")
def build_messages(self, item):
return [{
"role": "user",
"content": [{
"type": "text",
"text": str(item.get("prompt", "prompt")),
}],
}]
class _FakeProcessor:
def __init__(self, *, prompt_length=2, empty_full_length=3, full_length=4):
self.prompt_length = prompt_length
self.empty_full_length = empty_full_length
self.full_length = full_length
self.tokenizer = SimpleNamespace(pad_token_id=0, eos_token_id=2)
def apply_chat_template(
self,
messages,
*,
tokenize,
add_generation_prompt,
return_dict,
return_tensors,
):
del tokenize, add_generation_prompt, return_dict, return_tensors
assistant_messages = [message for message in messages if message["role"] == "assistant"]
if not assistant_messages:
length = self.prompt_length
else:
assistant_text = "".join(
str(part.get("text", "")) for part in assistant_messages[-1].get("content", [])
if isinstance(part, dict))
length = self.empty_full_length if not assistant_text else self.full_length
return {
"input_ids": torch.arange(1, length + 1).unsqueeze(0),
"attention_mask": torch.ones(1, length, dtype=torch.long),
}
def batch_decode(self, sequences, **kwargs):
del kwargs
return ["response" for _ in sequences]
class _FakeQwenPolicy(torch.nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.nn.Parameter(torch.tensor(0.0))
self.generate_kwargs = None
def forward(self, **kwargs):
input_ids = kwargs["input_ids"]
batch, seq_len = input_ids.shape
logits = torch.zeros(batch, seq_len, 16, dtype=self.weight.dtype)
for idx in range(seq_len - 1):
next_token = input_ids[:, idx + 1]
logits[:, idx].scatter_(1, next_token[:, None], self.weight.expand(batch, 1))
return SimpleNamespace(logits=logits)
def generate(self, **kwargs):
self.generate_kwargs = kwargs
input_ids = kwargs["input_ids"]
suffix = torch.tensor([[9]], dtype=input_ids.dtype)
return torch.cat([input_ids, suffix], dim=1)
def _make_fake_actor(*, max_prompt_length=8, max_response_length=4):
actor = _FakeActor(
init_from="unused",
load_backend=False,
max_prompt_length=max_prompt_length,
max_response_length=max_response_length,
)
actor.processor = _FakeProcessor()
actor.transformer = _FakeQwenPolicy()
return actor
def test_grpo_update_is_invariant_to_micro_batch_partitioning():
rollouts = [{
"prompt": f"prompt-{index}",
"response": "response",
} for index in range(5)]
advantages = torch.ones(5)
weights = []
for micro_batch_size in (2, 5):
actor = _make_fake_actor()
optimizer = torch.optim.SGD(actor.transformer.parameters(), lr=0.5)
actor.train_interleave_rollouts(
rollouts=rollouts,
advantages=advantages,
optimizer=optimizer,
update_micro_batch_size=micro_batch_size,
gradient_accumulation_steps=1,
)
weights.append(actor.transformer.weight.detach().clone())
assert torch.allclose(weights[0], weights[1], atol=1.0e-6)
def test_actor_rejects_tokenized_prompts_over_configured_limit():
actor = _make_fake_actor(max_prompt_length=1)
with pytest.raises(ValueError, match="max_prompt_length=1"):
actor.generate_qwen_responses(actor.build_messages({"prompt": "too long"}))
def test_actor_caps_generation_at_configured_response_limit():
actor = _make_fake_actor(max_response_length=3)
actor.generate_qwen_responses(
actor.build_messages({"prompt": "prompt"}),
max_new_tokens=99,
temperature=0.0,
)
assert actor.transformer.generate_kwargs["max_new_tokens"] == 3
def test_actor_rejects_tokenized_responses_over_configured_limit():
actor = _make_fake_actor(max_prompt_length=4, max_response_length=1)
actor.processor = _FakeProcessor(full_length=5)
with pytest.raises(ValueError, match="max_response_length=1"):
actor.response_logprobs_from_messages(actor.build_messages({"prompt": "prompt"}), "response")
def test_actor_accepts_response_content_at_exact_configured_limit():
actor = _make_fake_actor(max_prompt_length=4, max_response_length=1)
logprobs, mask = actor.response_logprobs_from_messages(
actor.build_messages({"prompt": "prompt"}),
"response",
)
assert logprobs.numel() == 2
assert mask.tolist() == [1.0, 1.0]
def test_actor_enables_synced_generation_for_distributed_fsdp(monkeypatch):
actor = _make_fake_actor()
monkeypatch.setattr(qwen_actor_module, "_distributed_actor_world_size", lambda: 2)
actor.generate_qwen_responses(actor.build_messages({"prompt": "prompt"}), temperature=0.0)
assert actor.transformer.generate_kwargs["synced_gpus"] is True
def test_interleave_dataloader_restores_shuffle_position(tmp_path):
data_path = tmp_path / "records.jsonl"
data_path.write_text("".join(json.dumps({"prompt": f"prompt-{index}"}) + "\n" for index in range(8)))
training_config = TrainingConfig(
data=DataConfig(
data_path=str(data_path),
train_batch_size=2,
dataloader_num_workers=0,
seed=17,
))
actor = _FakeActor(init_from="unused", load_backend=False)
actor.init_preprocessors(training_config)
iterator = iter(actor.dataloader)
next(iterator)
state = actor.dataloader.state_dict()
expected = next(iterator)
resumed_actor = _FakeActor(init_from="unused", load_backend=False)
resumed_actor.init_preprocessors(training_config)
resumed_actor.dataloader.load_state_dict(state)
assert next(iter(resumed_actor.dataloader)) == expected
def test_interleave_dataloader_shards_records_across_distributed_ranks(monkeypatch, tmp_path):
data_path = tmp_path / "records.jsonl"
data_path.write_text("".join(json.dumps({"prompt": f"prompt-{index}"}) + "\n" for index in range(8)))
training_config = TrainingConfig(
data=DataConfig(
data_path=str(data_path),
train_batch_size=2,
dataloader_num_workers=0,
seed=23,
))
monkeypatch.setattr(qwen_actor_module.dist, "is_available", lambda: True)
monkeypatch.setattr(qwen_actor_module.dist, "is_initialized", lambda: True)
monkeypatch.setattr(qwen_actor_module.dist, "get_world_size", lambda: 2)
rank_records = []
epoch_orders = []
for rank in (0, 1):
monkeypatch.setattr(qwen_actor_module.dist, "get_rank", lambda rank=rank: rank)
actor = _FakeActor(init_from="unused", load_backend=False)
actor.init_preprocessors(training_config)
epoch_orders.append((list(actor.dataloader.sampler), list(actor.dataloader.sampler)))
rank_records.append({
item["prompt"]
for batch in actor.dataloader
for item in batch["items"]
})
assert rank_records[0].isdisjoint(rank_records[1])
assert rank_records[0] | rank_records[1] == {f"prompt-{index}" for index in range(8)}
assert all(first_epoch != second_epoch for first_epoch, second_epoch in epoch_orders)
def test_distributed_sampler_restores_rank_local_epoch_position(monkeypatch, tmp_path):
data_path = tmp_path / "records.jsonl"
data_path.write_text("".join(json.dumps({"prompt": f"prompt-{index}"}) + "\n" for index in range(8)))
training_config = TrainingConfig(data=DataConfig(data_path=str(data_path), train_batch_size=2, seed=31))
monkeypatch.setattr(qwen_actor_module.dist, "is_available", lambda: True)
monkeypatch.setattr(qwen_actor_module.dist, "is_initialized", lambda: True)
monkeypatch.setattr(qwen_actor_module.dist, "get_world_size", lambda: 2)
monkeypatch.setattr(qwen_actor_module.dist, "get_rank", lambda: 1)
actor = _FakeActor(init_from="unused", load_backend=False)
actor.init_preprocessors(training_config)
sampler_iterator = iter(actor.dataloader.sampler)
next(sampler_iterator)
state = actor.dataloader.sampler.state_dict()
expected = next(sampler_iterator)
resumed_actor = _FakeActor(init_from="unused", load_backend=False)
resumed_actor.init_preprocessors(training_config)
resumed_actor.dataloader.sampler.load_state_dict(state)
assert next(iter(resumed_actor.dataloader.sampler)) == expected
def test_distributed_actor_rejects_unsupported_tensor_parallelism():
actor = _FakeActor(
init_from="unused",
load_backend=False,
training_config=TrainingConfig(
distributed=DistributedConfig(
num_gpus=8,
tp_size=4,
sp_size=1,
hsdp_replicate_dim=1,
hsdp_shard_dim=8,
)),
)
with pytest.raises(ValueError, match="do not implement tensor parallelism"):
actor._validate_distributed_config(device_map=None)
def test_hsdp_auto_shard_dimension_uses_configured_gpu_count():
distributed = DistributedConfig(
num_gpus=8,
tp_size=1,
sp_size=1,
hsdp_replicate_dim=1,
hsdp_shard_dim=-1,
)
assert _resolve_hsdp_dimensions(distributed, num_gpus=8) == (1, 8)
def test_distributed_token_gradient_scale_uses_global_denominator(monkeypatch):
monkeypatch.setattr(qwen_actor_module.dist, "is_available", lambda: True)
monkeypatch.setattr(qwen_actor_module.dist, "is_initialized", lambda: True)
monkeypatch.setattr(qwen_actor_module.dist, "get_world_size", lambda: 2)
def fake_all_reduce(value, op):
del op
value.fill_(12.0)
monkeypatch.setattr(qwen_actor_module.dist, "all_reduce", fake_all_reduce)
rank_zero_scale = _distributed_token_gradient_scale(2.0, device=torch.device("cpu"))
rank_one_scale = _distributed_token_gradient_scale(10.0, device=torch.device("cpu"))
assert rank_zero_scale == rank_one_scale == 1.0 / 6.0
def test_rank_independent_rng_repeats_adapter_initialization_and_restores_state():
torch.manual_seed(101)
before = torch.rand(1)
with _rank_independent_rng(7, device=torch.device("cpu")):
first_adapter = torch.rand(4)
after = torch.rand(1)
torch.manual_seed(999)
with _rank_independent_rng(7, device=torch.device("cpu")):
second_adapter = torch.rand(4)
torch.manual_seed(101)
assert torch.equal(before, torch.rand(1))
assert torch.equal(after, torch.rand(1))
assert torch.equal(first_adapter, second_adapter)
@pytest.mark.parametrize(
"config_path",
[
"examples/train/configs/interleave_thinker/planner_sft_lora.yaml",
"examples/train/configs/interleave_thinker/critic_sft_lora.yaml",
"examples/train/configs/rl/interleave_thinker/planner_grpo.yaml",
"examples/train/configs/rl/interleave_thinker/critic_grpo.yaml",
],
)
def test_public_distributed_actor_configs_use_supported_hsdp(config_path):
distributed = load_run_config(config_path).training.distributed
assert distributed.num_gpus == 8
assert distributed.tp_size == 1
assert distributed.sp_size == 1
assert distributed.hsdp_replicate_dim == 1
assert distributed.hsdp_shard_dim == 8
def test_qwen_shard_condition_matches_module_list_blocks():
class VisionBlock(torch.nn.Linear):
pass
model = torch.nn.ModuleDict({
"language": torch.nn.ModuleList([torch.nn.Linear(2, 2)]),
"visual": torch.nn.ModuleList([VisionBlock(2, 2)]),
})
condition = _qwen_transformer_block_condition(model)
assert condition("language.0", model["language"][0]) is True
assert condition("language", model["language"]) is False
assert condition("visual.0", model["visual"][0]) is False
def test_qwen_sharding_root_uses_peft_forward_and_generate_base():
base_model = torch.nn.Sequential(torch.nn.ModuleList([torch.nn.Linear(2, 2)]))
class FakePeftModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.base = base_model
def get_base_model(self):
return self.base
wrapper = FakePeftModel()
assert _qwen_sharding_root(wrapper) is base_model
assert _qwen_sharding_root(base_model) is base_model
@@ -0,0 +1,156 @@
import pytest
from fastvideo.train.methods.rl.rewards import (
InterleavePlannerRewardScorer,
InterleaveThinkerEditScore,
InterleaveThinkerRewardScorer,
extract_interleave_answer,
extract_interleave_plan_payload,
interleave_format_reward,
interleave_planner_format_reward,
score_interleave_planner_rewards,
score_interleave_thinker_rewards,
)
def _response(success=True, refine_prompt="make it sharper"):
return f"""
<think>
The generated image needs a stricter prompt.
</think>
<answer>
{{'previous_step_success': {success}, 'refine_prompt': {refine_prompt!r}}}
</answer>
"""
def test_extract_answer_accepts_upstream_single_quote_jsonish_payload():
answer = extract_interleave_answer(_response(False, "add red highlights"))
assert answer is not None
assert answer.previous_step_success is False
assert answer.refine_prompt == "add red highlights"
def test_format_reward_requires_think_before_valid_answer():
assert interleave_format_reward(_response(True)) == 1.0
missing_think_close = """
<think>unfinished reasoning
<answer>{"previous_step_success": true, "refine_prompt": "ok"}</answer>
"""
assert interleave_format_reward(missing_think_close) == 0.0
answer_inside_think = """
<think>reasoning
<answer>{"previous_step_success": true, "refine_prompt": "ok"}</answer>
</think>
"""
assert interleave_format_reward(answer_inside_think) == 0.0
answer_first = """
<answer>{"previous_step_success": true, "refine_prompt": "ok"}</answer>
<think>late reasoning</think>
"""
assert interleave_format_reward(answer_first) == 0.0
invalid_answer = """
<think>reasoning</think>
<answer>{"previous_step_success": "true", "refine_prompt": "ok"}</answer>
"""
assert interleave_format_reward(invalid_answer) == 0.0
def test_reward_scorer_matches_upstream_default_weighting_with_absolute_edit_scores():
scorer = InterleaveThinkerRewardScorer()
result = scorer([{
"response": _response(False),
"ground_truth": {
"success": False,
"semantics": 6.0,
"quality": 8.0,
},
"edit_score": {
"semantics": 8.0,
"quality": 7.0,
},
}])[0]
assert result.format_reward == 1.0
assert result.judge_accuracy_reward == 1.0
assert result.edited_image_reward_semantic == pytest.approx(0.6)
assert result.edited_image_reward_quality == pytest.approx(0.45)
assert result.overall == pytest.approx(0.825)
def test_reward_scorer_uses_injected_edit_scorer_and_json_string_ground_truth():
requests = []
def edit_scorer(request):
requests.append(request)
return InterleaveThinkerEditScore(semantic_reward=0.75, quality_reward=0.25)
scorer = InterleaveThinkerRewardScorer(format_weight=0.0, edit_scorer=edit_scorer)
results = score_interleave_thinker_rewards(
[{
"response": _response(True, ""),
"ground_truth": '{"success": true, "semantics": 4, "quality": 4}',
"origin_prompt": "draw a glass vase",
"previous_prompt": "a vase on a table",
"origin_image_path": "origin.png",
}],
format_weight=0.0,
edit_scorer=edit_scorer,
)
assert requests[0].refine_prompt == "a vase on a table"
assert requests[0].previous_step_success is True
assert results[0]["overall"] == pytest.approx(0.2 * 1.0 + 0.6 * 0.75 + 0.2 * 0.25)
assert scorer([{
"response": _response(True, ""),
"ground_truth": {
"success": True
},
}])[0].overall >= 0.0
def _planner_response(prompt="a clean cat sketch"):
prompt = prompt.replace('"', '\\"')
return f"""
<think>Plan the sequence.</think>
<answer>
{{"execution_plan": [
{{"step_number": 1, "step_name": "Sketch", "instruction": "Draw a cat", "prompt": "{prompt}", "auxiliary_text": null}}
]}}
</answer>
"""
def test_planner_reward_accepts_execution_plan_answer_block():
payload = extract_interleave_plan_payload(_planner_response())
assert payload is not None
assert interleave_planner_format_reward(_planner_response()) == 1.0
assert interleave_planner_format_reward("<answer>{}</answer>") == 0.0
assert interleave_planner_format_reward(_planner_response().replace("</think>", "")) == 0.0
def test_planner_reward_scorer_blends_format_and_scalar_plan_score():
scorer = InterleavePlannerRewardScorer(format_weight=0.25, fallback_plan_reward=0.1)
result = scorer([{
"response": _planner_response(),
"plan_score": 0.9,
}])[0]
wrapped = score_interleave_planner_rewards([{
"response": _planner_response(),
"ground_truth": {
"score": 0.5
},
}], format_weight=0.5)[0]
assert result.format_reward == 1.0
assert result.planner_score == pytest.approx(0.9)
assert result.overall == pytest.approx(0.25 * 1.0 + 0.75 * 0.9)
assert wrapped["overall"] == pytest.approx(0.5 * 1.0 + 0.5 * 0.5)
@@ -0,0 +1,151 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
from types import SimpleNamespace
import torch
from fastvideo.train.methods.fine_tuning import InterleaveThinkerSFTMethod
from fastvideo.train.models.interleave_thinker import (
InterleaveThinkerCriticModel,
InterleaveThinkerPlannerModel,
)
from fastvideo.train.utils.config import load_run_config
from fastvideo.train.utils.training_config import (
DataConfig,
OptimizerConfig,
TrainingConfig,
TrainingLoopConfig,
)
class _FakeBackendCritic(InterleaveThinkerCriticModel):
@property
def device(self):
return torch.device("cpu")
class _FakeProcessor:
def __init__(self):
self.tokenizer = SimpleNamespace(pad_token_id=0, eos_token_id=2)
def apply_chat_template(
self,
messages,
*,
tokenize,
add_generation_prompt,
return_dict,
return_tensors,
):
del tokenize, add_generation_prompt, return_dict, return_tensors
has_assistant = any(message["role"] == "assistant" for message in messages)
length = 5 if has_assistant else 3
return {
"input_ids": torch.arange(1, length + 1).unsqueeze(0),
"attention_mask": torch.ones(1, length, dtype=torch.long),
}
class _FakeQwen(torch.nn.Module):
def __init__(self):
super().__init__()
self.weight = torch.nn.Parameter(torch.tensor(1.0))
self.last_labels = None
def forward(self, **kwargs):
labels = kwargs["labels"]
self.last_labels = labels.detach().clone()
trainable_fraction = (labels != -100).float().mean()
return SimpleNamespace(loss=self.weight.pow(2).sum() * trainable_fraction)
def test_interleave_thinker_sft_method_trains_response_tokens_only():
model = _FakeBackendCritic(load_backend=False, trainable=True)
model.processor = _FakeProcessor()
model.transformer = _FakeQwen()
method = InterleaveThinkerSFTMethod(
cfg=SimpleNamespace(
method={},
validation={},
training=TrainingConfig(
data=DataConfig(data_path="", train_batch_size=1),
optimizer=OptimizerConfig(learning_rate=0.1),
loop=TrainingLoopConfig(max_train_steps=1),
),
),
role_models={"student": model},
)
before = float(model.transformer.weight.detach())
loss_map, outputs, metrics = method.single_train_step(
{
"items": [{
"origin_prompt": "draw a chair",
"previous_prompt": "a wooden chair",
"origin_image_path": "before.png",
"edited_image_path": "after.png",
"response": '<answer>{"previous_step_success": true, "refine_prompt": "ok"}</answer>',
}]
},
iteration=0,
)
method.backward(loss_map, outputs, grad_accum_rounds=1)
method.optimizers_schedulers_step(0)
assert set(loss_map) == {"total_loss", "sft_loss"}
assert metrics["sft/response_tokens"] == 2.0
assert metrics["sft/num_items"] == 1.0
assert model.transformer.last_labels.tolist()[0][:3] == [-100, -100, -100]
assert all(label != -100 for label in model.transformer.last_labels.tolist()[0][3:])
assert float(model.transformer.weight.detach()) < before
def test_qwen_actor_dataset_kind_uses_planner_sft_normalizer(tmp_path):
image_dir = tmp_path / "images"
data_dir = tmp_path / "data"
data_dir.mkdir()
(data_dir / "planner_sft.json").write_text(
json.dumps([{
"messages": [
{
"role": "user",
"content": "draw a cat"
},
{
"role": "assistant",
"content": "<answer>{\"execution_plan\": []}</answer>"
},
],
"images": ["planner/cat.png"],
}]),
encoding="utf-8",
)
model = InterleaveThinkerPlannerModel(
load_backend=False,
dataset_kind="planner_sft",
image_dir=str(image_dir),
)
model.init_preprocessors(TrainingConfig(data=DataConfig(data_path=str(data_dir), train_batch_size=1)))
batch = next(iter(model.dataloader))
assert batch["items"][0]["instruction"] == "draw a cat"
assert batch["items"][0]["response"] == '<answer>{"execution_plan": []}</answer>'
assert batch["items"][0]["images"] == [str(image_dir / "planner/cat.png")]
def test_interleave_thinker_sft_configs_parse_public_yaml():
planner_cfg = load_run_config("examples/train/configs/interleave_thinker/planner_sft_lora.yaml")
critic_cfg = load_run_config("examples/train/configs/interleave_thinker/critic_sft_lora.yaml")
assert planner_cfg.models["student"]["dataset_kind"] == "planner_sft"
assert planner_cfg.method["_target_"] == "fastvideo.train.methods.fine_tuning.InterleaveThinkerSFTMethod"
assert planner_cfg.models["student"]["lora"]["enable"] is True
assert critic_cfg.models["student"]["dataset_kind"] == "critic_sft"
assert critic_cfg.method["_target_"] == "fastvideo.train.methods.fine_tuning.InterleaveThinkerSFTMethod"
@@ -0,0 +1,181 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
from pathlib import Path
from fastvideo.workflow.interleave_thinker import (
discover_interleave_trace_paths,
evaluate_interleave_traces,
interleave_trace_evaluation_to_dict,
write_interleave_trace_html_report,
)
def test_evaluate_interleave_traces_from_summary(tmp_path: Path) -> None:
output_dir = _write_trace_fixture(tmp_path)
summary = evaluate_interleave_traces([output_dir / "summary.json"])
assert summary.num_traces == 2
assert summary.num_success == 1
assert summary.success_rate == 0.5
assert summary.total_attempts == 3
assert summary.average_attempts == 1.5
assert summary.total_retry_attempts == 1
assert summary.traces_with_final_image == 1
assert summary.total_inference_time_s == 1.5
assert summary.failure_reasons == {"critic rejected final attempt": 1}
assert summary.success_by_category["product"] == {
"num_traces": 1.0,
"num_success": 1.0,
"success_rate": 1.0,
}
payload = interleave_trace_evaluation_to_dict(summary)
assert payload["traces"][0]["prompt_set_id"] == "mug"
assert payload["traces"][1]["failure_reason"] == "critic rejected final attempt"
def test_discover_interleave_trace_paths_from_directory(tmp_path: Path) -> None:
output_dir = _write_trace_fixture(tmp_path)
paths = discover_interleave_trace_paths([output_dir])
assert [path.name for path in paths] == ["trace.json", "trace.json"]
assert {path.parent.name for path in paths} == {"mug", "poster"}
def test_write_interleave_trace_html_report(tmp_path: Path) -> None:
output_dir = _write_trace_fixture(tmp_path)
summary = evaluate_interleave_traces([output_dir])
html_path = tmp_path / "report.html"
write_interleave_trace_html_report(summary, html_path, title="Smoke Report")
html_text = html_path.read_text(encoding="utf-8")
assert "Smoke Report" in html_text
assert "mug" in html_text
assert "critic rejected final attempt" in html_text
assert "<img" in html_text
def _write_trace_fixture(tmp_path: Path) -> Path:
output_dir = tmp_path / "eval"
image_path = output_dir / "mug" / "final.png"
image_path.parent.mkdir(parents=True, exist_ok=True)
image_path.write_bytes(b"fake-image")
_write_json(
output_dir / "mug" / "trace.json",
{
"instruction": "draw a mug",
"success": True,
"metadata": {
"prompt_set_id": "mug",
"prompt_set_index": 0,
"prompt_set_metadata": {
"category": "product",
},
},
"final_image": {
"prompt": "refined mug",
"file_path": str(image_path),
"inference_time_s": 0.4,
"metadata": {},
},
"attempts": [
{
"step_index": 0,
"attempt_index": 0,
"prompt": "draw a mug",
"generated": {
"prompt": "draw a mug",
"file_path": str(output_dir / "mug" / "attempt0.png"),
"inference_time_s": 0.5,
"metadata": {},
},
"decision": {
"success": False,
"refine_prompt": "refined mug",
"reason": "needs refinement",
"metadata": {},
},
},
{
"step_index": 0,
"attempt_index": 1,
"prompt": "refined mug",
"generated": {
"prompt": "refined mug",
"file_path": str(image_path),
"inference_time_s": 0.7,
"metadata": {},
},
"decision": {
"success": True,
"refine_prompt": None,
"reason": None,
"metadata": {},
},
},
],
},
)
_write_json(
output_dir / "poster" / "trace.json",
{
"instruction": "draw a poster",
"success": False,
"metadata": {
"prompt_set_id": "poster",
"prompt_set_index": 1,
"failed_step_index": 0,
"prompt_set_metadata": {
"category": "poster",
},
},
"final_image": None,
"attempts": [
{
"step_index": 0,
"attempt_index": 0,
"prompt": "draw a poster",
"generated": {
"prompt": "draw a poster",
"file_path": str(output_dir / "poster" / "attempt0.png"),
"inference_time_s": 0.3,
"metadata": {},
},
"decision": {
"success": False,
"refine_prompt": None,
"reason": "critic rejected final attempt",
"metadata": {},
},
},
],
},
)
_write_json(
output_dir / "summary.json",
{
"num_samples": 2,
"num_success": 1,
"results": [
{
"sample_id": "mug",
"trace_path": str(output_dir / "mug" / "trace.json"),
},
{
"sample_id": "poster",
"trace_path": str(output_dir / "poster" / "trace.json"),
},
],
},
)
return output_dir
def _write_json(path: Path, payload: dict) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
@@ -0,0 +1,197 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import base64
from pathlib import Path
from fastvideo.api.compat import (
explicit_request_updates,
legacy_generate_call_to_request,
)
from fastvideo.api.results import GenerationResult
from fastvideo.workflow.interleave_thinker.generator import (
FastVideoImageGeneratorBackend,
)
from fastvideo.workflow.interleave_thinker.orchestrator import InterleaveOrchestrator
from fastvideo.workflow.interleave_thinker.schema import (
CriticDecision,
GeneratedImage,
InterleaveEditRequest,
PlannedInterleaveStep,
)
from fastvideo.workflow.interleave_thinker.trace import (
save_trace,
trace_to_dict,
)
def test_interleave_edit_request_accepts_singular_step_field() -> None:
request = InterleaveEditRequest(
prompt="a ceramic cup on a table",
num_inference_step=4,
)
assert request.resolved_num_inference_steps() == 4
plural_request = InterleaveEditRequest(
prompt="a ceramic cup on a table",
num_inference_step=4,
num_inference_steps=8,
)
assert plural_request.resolved_num_inference_steps() == 8
def test_fastvideo_backend_translates_edit_request(tmp_path: Path) -> None:
class FakeGenerator:
def __init__(self) -> None:
self.requests = []
def generate(self, request):
self.requests.append(request)
updates = explicit_request_updates(request)
output_path = Path(updates["output_path"])
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_bytes(b"fake-png")
return GenerationResult(
prompt=request.prompt,
video_path=str(output_path),
generation_time=0.25,
)
default_request = legacy_generate_call_to_request(
"unused",
None,
legacy_kwargs={
"height": 512,
"width": 768,
"num_inference_steps": 4,
},
)
fake = FakeGenerator()
backend = FastVideoImageGeneratorBackend(
fake,
output_dir=str(tmp_path),
default_request=default_request,
)
input_b64 = base64.b64encode(b"input-image").decode("utf-8")
generated = backend.generate(
InterleaveEditRequest(
prompt="turn it into a watercolor",
image=input_b64,
width=1024,
seed=7,
),
request_id="abc123",
)
assert generated.prompt == "turn it into a watercolor"
assert generated.image_base64 == base64.b64encode(b"fake-png").decode("utf-8")
assert generated.file_path is not None
assert generated.file_path.endswith("abc123.png")
updates = explicit_request_updates(fake.requests[0])
assert updates["num_frames"] == 1
assert updates["fps"] == 1
assert updates["height"] == 512
assert updates["width"] == 1024
assert updates["num_inference_steps"] == 4
assert updates["seed"] == 7
assert updates["save_video"] is True
assert updates["return_frames"] is False
assert Path(updates["image_path"]).read_bytes() == b"input-image"
def test_interleave_orchestrator_retries_with_refined_prompt() -> None:
class FakePlanner:
def plan(self, request):
return [
PlannedInterleaveStep(
prompt=request.instruction,
max_attempts=2,
)
]
class FakeGenerator:
def __init__(self) -> None:
self.prompts = []
def generate(self, request, *, request_id=None):
del request_id
self.prompts.append(request.prompt)
return GeneratedImage(
prompt=request.prompt,
image_base64=base64.b64encode(request.prompt.encode("utf-8")).decode("utf-8"),
file_path=f"/tmp/{len(self.prompts)}.png",
)
class RefiningCritic:
def __init__(self) -> None:
self.calls = 0
def review(self, request):
self.calls += 1
if self.calls == 1:
return CriticDecision(
success=False,
refine_prompt="refined prompt",
reason="first attempt missed the instruction",
)
return CriticDecision(success=True)
generator = FakeGenerator()
orchestrator = InterleaveOrchestrator(
planner=FakePlanner(),
generator=generator,
critic=RefiningCritic(),
)
trace = orchestrator.run("initial prompt")
assert trace.success is True
assert generator.prompts == ["initial prompt", "refined prompt"]
assert len(trace.attempts) == 2
assert trace.final_image is not None
assert trace.final_image.prompt == "refined prompt"
def test_trace_serialization_omits_images_by_default(tmp_path: Path) -> None:
generated = GeneratedImage(
prompt="final",
image_base64="large-payload",
file_path="/tmp/final.png",
inference_time_s=0.5,
)
trace = InterleaveOrchestrator(
planner=FakeSingleStepPlanner(),
generator=FakeSingleStepGenerator(generated),
).run("final")
payload = trace_to_dict(trace)
assert payload["success"] is True
assert payload["final_image"]["file_path"] == "/tmp/final.png"
assert "image_base64" not in payload["final_image"]
assert "image_base64" not in payload["attempts"][0]["generated"]
trace_path = tmp_path / "trace.json"
save_trace(trace, trace_path)
assert "large-payload" not in trace_path.read_text(encoding="utf-8")
payload_with_images = trace_to_dict(trace, include_images=True)
assert payload_with_images["final_image"]["image_base64"] == "large-payload"
class FakeSingleStepPlanner:
def plan(self, request):
return [PlannedInterleaveStep(prompt=request.instruction)]
class FakeSingleStepGenerator:
def __init__(self, generated: GeneratedImage) -> None:
self.generated = generated
def generate(self, request, *, request_id=None):
del request, request_id
return self.generated
@@ -0,0 +1,247 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import base64
from pathlib import Path
import pytest
from fastvideo.workflow.interleave_thinker import (
GeneratedImage,
InterleaveEditRequest,
InterleavePromptItem,
load_interleave_prompt_set,
load_interleave_run_config,
resolve_interleave_instruction,
run_interleave_prompt_set,
run_interleave_config,
)
from fastvideo.workflow.interleave_thinker.orchestrator import SinglePromptPlanner
from fastvideo.workflow.interleave_thinker.schema import PlannerInput
def test_interleave_run_config_loads_prompt_and_request_defaults(tmp_path: Path) -> None:
config_path = _write_run_config(tmp_path)
config = load_interleave_run_config(str(config_path))
assert resolve_interleave_instruction(config) == "draw a red mug"
assert config.generator is not None
assert config.generator.model_path == "black-forest-labs/FLUX.2-klein-4B"
assert config.request.sampling.width == 512
assert config.request.sampling.num_inference_steps == 4
def test_interleave_run_config_accepts_runtime_fields_and_dotted_overrides(tmp_path: Path) -> None:
config_path = _write_run_config(tmp_path)
config = load_interleave_run_config(
str(config_path),
prompt="draw a blue mug",
output_dir=str(tmp_path / "override_outputs"),
trace_path=str(tmp_path / "trace_override.json"),
overrides=[
"--request.sampling.seed",
"99",
"--planner.max-attempts-per-step",
"3",
],
)
assert resolve_interleave_instruction(config) == "draw a blue mug"
assert config.interleave.output_dir == str(tmp_path / "override_outputs")
assert config.interleave.trace_path == str(tmp_path / "trace_override.json")
assert config.request.sampling.seed == 99
assert config.planner.max_attempts_per_step == 3
def test_interleave_run_config_rejects_unknown_override_prefix(tmp_path: Path) -> None:
config_path = _write_run_config(tmp_path)
with pytest.raises(ValueError, match="Unsupported override path"):
load_interleave_run_config(
str(config_path),
overrides=["--server.port", "9000"],
)
def test_interleave_eval_config_can_omit_single_instruction(tmp_path: Path) -> None:
config_path = _write_run_config(tmp_path, include_prompt=False)
config = load_interleave_run_config(
str(config_path),
require_instruction=False,
)
with pytest.raises(ValueError, match="requires interleave.instruction"):
resolve_interleave_instruction(config)
def test_run_interleave_config_with_injected_backend_writes_trace(tmp_path: Path) -> None:
config_path = _write_run_config(tmp_path)
config = load_interleave_run_config(str(config_path))
backend = _FakeImageBackend(tmp_path / "generated.png")
result = run_interleave_config(config, image_backend=backend)
assert result.trace.success is True
assert result.trace.final_image is not None
assert result.trace.final_image.file_path == str(tmp_path / "generated.png")
assert Path(result.trace_path).exists()
trace_text = Path(result.trace_path).read_text(encoding="utf-8")
assert "draw a red mug" in trace_text
assert "image_base64" not in trace_text
assert backend.requests[0].prompt == "draw a red mug"
def test_load_interleave_prompt_set_accepts_jsonl_rows(tmp_path: Path) -> None:
prompt_path = tmp_path / "prompts.jsonl"
prompt_path.write_text(
'{"id": "mug", "prompt": "draw a mug", "metadata": {"split": "smoke"}, "difficulty": "easy"}\n'
'"draw a kettle"\n',
encoding="utf-8",
)
items = load_interleave_prompt_set(prompt_path)
assert [item.sample_id for item in items] == ["mug", "sample_00001"]
assert items[0].instruction == "draw a mug"
assert items[0].metadata == {
"split": "smoke",
"difficulty": "easy",
}
assert items[1].instruction == "draw a kettle"
def test_run_interleave_prompt_set_writes_traces_and_summary(tmp_path: Path) -> None:
config_path = _write_run_config(tmp_path, include_prompt=False)
config = load_interleave_run_config(
str(config_path),
require_instruction=False,
)
backend = _FakeImageBackend(tmp_path / "generated.png")
items = [
InterleavePromptItem(sample_id="red/mug", instruction="draw a red mug"),
InterleavePromptItem(sample_id="blue mug", instruction="draw a blue mug"),
]
summary = run_interleave_prompt_set(
config,
items,
output_dir=str(tmp_path / "eval"),
image_backend=backend,
)
assert summary.num_samples == 2
assert summary.num_success == 2
assert summary.success_rate == 1.0
assert summary.total_attempts == 2
assert [request.prompt for request in backend.requests] == ["draw a red mug", "draw a blue mug"]
assert Path(summary.summary_path).exists()
assert Path(summary.results[0].trace_path).exists()
assert "red_mug" in summary.results[0].trace_path
summary_text = Path(summary.summary_path).read_text(encoding="utf-8")
assert "draw a blue mug" in summary_text
def test_run_interleave_prompt_set_resume_uses_existing_trace(tmp_path: Path) -> None:
config_path = _write_run_config(tmp_path, include_prompt=False)
config = load_interleave_run_config(
str(config_path),
require_instruction=False,
)
item = InterleavePromptItem(sample_id="mug", instruction="draw a mug")
backend = _FakeImageBackend(tmp_path / "generated.png")
first = run_interleave_prompt_set(
config,
[item],
output_dir=str(tmp_path / "eval"),
image_backend=backend,
)
resumed = run_interleave_prompt_set(
config,
[item],
output_dir=str(tmp_path / "eval"),
image_backend=_FakeImageBackend(tmp_path / "unused.png"),
resume=True,
)
assert first.num_resumed == 0
assert resumed.num_resumed == 1
assert resumed.results[0].resumed is True
assert resumed.results[0].trace_path == first.results[0].trace_path
def test_single_prompt_planner_uses_configured_attempt_count() -> None:
planner = SinglePromptPlanner(max_attempts=4)
steps = list(planner.plan(PlannerInput(instruction="draw a red mug")))
assert len(steps) == 1
assert steps[0].max_attempts == 4
class _FakeImageBackend:
def __init__(self, output_path: Path) -> None:
self.requests: list[InterleaveEditRequest] = []
self.output_path = output_path
def generate(
self,
request: InterleaveEditRequest,
*,
request_id: str | None = None,
) -> GeneratedImage:
del request_id
self.requests.append(request)
self.output_path.write_bytes(b"fake-image")
return GeneratedImage(
prompt=request.prompt,
image_base64=base64.b64encode(b"fake-image").decode("utf-8"),
file_path=str(self.output_path),
metadata={"backend": "fake"},
)
def _write_run_config(tmp_path: Path, *, include_prompt: bool = True) -> Path:
output_path = tmp_path / "generated.png"
config_path = tmp_path / "interleave_run.yaml"
request_prompt = " prompt: draw a red mug\n" if include_prompt else ""
config_path.write_text(
f"""
generator:
model_path: black-forest-labs/FLUX.2-klein-4B
engine:
num_gpus: 1
pipeline:
workload_type: t2i
image_backend:
kind: fastvideo
planner:
kind: single_prompt
max_attempts_per_step: 1
critic:
kind: accept_all
interleave:
output_dir: {tmp_path}
trace_path: {tmp_path / "trace.json"}
request:
{request_prompt} extensions:
test_output_path: {output_path}
sampling:
width: 512
height: 512
seed: 7
num_inference_steps: 4
""",
encoding="utf-8",
)
return config_path
+57
View File
@@ -0,0 +1,57 @@
import torch
import pytest
from fastvideo.train.methods.rl.rewards import MultiRewardScorer, select_first_frame
def test_select_first_frame_for_video_tensor():
video = torch.arange(2 * 3 * 4 * 5 * 6).reshape(2, 3, 4, 5, 6)
frame = select_first_frame(video)
assert frame.shape == (2, 3, 5, 6)
torch.testing.assert_close(frame, video[:, :, 0])
def test_select_first_frame_keeps_frame_tensor():
frame = torch.randn(2, 3, 5, 6)
selected = select_first_frame(frame)
assert selected is frame
def test_multi_reward_weighted_sum_with_injected_scorers():
def pickscore(media, prompts):
assert media.shape == (2, 3, 4, 5, 6)
assert prompts == ["a", "b"]
return torch.tensor([1.0, 2.0])
def clipscore(media, prompts):
assert media.shape == (2, 3, 4, 5, 6)
assert prompts == ["a", "b"]
return torch.tensor([0.5, 1.5])
scorer = MultiRewardScorer(
{"pickscore": 2.0, "clipscore": 3.0},
scorers={
"pickscore": pickscore,
"clipscore": clipscore,
},
)
scores = scorer(torch.zeros(2, 3, 4, 5, 6), ["a", "b"])
torch.testing.assert_close(scores["pickscore"], torch.tensor([1.0, 2.0]))
torch.testing.assert_close(scores["clipscore"], torch.tensor([0.5, 1.5]))
torch.testing.assert_close(scores["avg"], torch.tensor([3.5, 8.5]))
def test_multi_reward_validates_score_shape():
scorer = MultiRewardScorer(
{"pickscore": 1.0},
scorers={"pickscore": lambda media, prompts: torch.tensor([[1.0], [2.0]])},
)
with pytest.raises(ValueError, match="must return shape"):
scorer(torch.zeros(2, 3, 4, 5, 6), ["a", "b"])
+238
View File
@@ -0,0 +1,238 @@
import torch
import pytest
from fastvideo.pipelines import TrainingBatch
from fastvideo.train.methods.rl.common import (
DiffusionSampler,
SamplingConfig,
distributed_k_repeat_indices,
media_to_video_array,
validation_caption,
validation_shard_indices,
)
from fastvideo.train.utils.config import load_run_config
class _FakeScheduler:
def __init__(self):
self.num_train_timesteps = 1000
self.set_timesteps_calls = []
self.timesteps = torch.empty(0)
self.sigmas = torch.empty(0)
self.step_calls = 0
def set_timesteps(self, num_inference_steps=None, device=None, timesteps=None, sigmas=None):
self.set_timesteps_calls.append({
"num_inference_steps": num_inference_steps,
"timesteps": timesteps,
"sigmas": sigmas,
})
if timesteps is not None:
self.timesteps = torch.tensor(timesteps, device=device, dtype=torch.float32)
else:
self.timesteps = torch.linspace(1000, 0, int(num_inference_steps), device=device)
if sigmas is not None:
self.sigmas = torch.tensor(sigmas, device=device, dtype=torch.float32)
else:
self.sigmas = torch.cat([self.timesteps / 1000.0, torch.zeros(1, device=device)])
def step(self, model_output, timestep, sample, return_dict=False):
del timestep
self.step_calls += 1
prev = sample + model_output
return (prev, ) if not return_dict else {"prev_sample": prev}
class _FakeModel:
def __init__(self):
self.noise_scheduler = _FakeScheduler()
self.add_noise_calls = 0
self.timestep_shapes = []
def predict_noise(self, noisy_latents, timestep, batch, *, conditional, attn_kind):
del conditional, attn_kind
self.timestep_shapes.append(tuple(timestep.shape))
assert batch.timesteps is timestep
return torch.zeros_like(noisy_latents)
def predict_x0(self, noisy_latents, timestep, batch, *, conditional, attn_kind):
del conditional, attn_kind
self.timestep_shapes.append(tuple(timestep.shape))
assert batch.timesteps is timestep
return noisy_latents
def add_noise(self, clean_latents, noise, timestep):
del timestep
self.add_noise_calls += 1
return clean_latents + noise
def _batch():
batch = TrainingBatch()
batch.latents = torch.zeros(2, 1, 3, 4, 4)
return batch
def test_sampler_preserves_latent_dtype():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
batch = _batch()
batch.latents = batch.latents.to(torch.bfloat16)
result = sampler.sample(model, batch, generator=torch.Generator().manual_seed(0))
assert result.latents.dtype is torch.bfloat16
def test_sampler_uses_scheduler_generated_timesteps_by_default():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [1000.0, 666.6666259765625, 333.3333435058594, 0.0]
def test_sampler_honors_explicit_timestep_override():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=3, timesteps=[900, 300, 10]))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [900.0, 300.0, 10.0]
def test_sampler_honors_explicit_timesteps_without_matching_num_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(timesteps=[900, 300, 10]))
result = sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert result.timesteps.tolist() == [900.0, 300.0, 10.0]
assert model.noise_scheduler.set_timesteps_calls == []
def test_sampling_config_rejects_unknown_keys():
with pytest.raises(ValueError, match="Unsupported method.sampling key"):
SamplingConfig.from_mapping({"solver": "dpm2"})
def test_sampler_restores_original_batch_timestep_after_sampling():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=2))
batch = _batch()
original_timesteps = torch.tensor([123.0])
batch.timesteps = original_timesteps
sampler.sample(batch=batch, model=model, generator=torch.Generator().manual_seed(0))
assert batch.timesteps is original_timesteps
def test_euler_sampler_does_not_renoise_between_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4))
sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert model.add_noise_calls == 0
assert model.timestep_shapes == [(2,), (2,), (2,), (2,)]
def test_sde_reflow_sampler_renoises_between_steps():
model = _FakeModel()
sampler = DiffusionSampler(SamplingConfig(num_steps=4, trajectory="sde_reflow"))
sampler.sample(model, _batch(), generator=torch.Generator().manual_seed(0))
assert model.add_noise_calls == 3
def test_diffusion_nft_config_uses_rl_sampler_not_dmd_pipeline():
config_path = "examples/train/configs/rl/wan/diffusion_nft_pick_clip.yaml"
cfg = load_run_config(config_path)
raw_text = open(config_path, encoding="utf-8").read()
assert cfg.method["_target_"] == "fastvideo.train.methods.rl.diffusion_nft.DiffusionNFTMethod"
assert cfg.training.optimizer.learning_rate == 3.0e-5
assert cfg.training.data.num_latent_t == 1
assert cfg.training.data.num_frames == 1
assert "sampling_timesteps" not in raw_text
assert "WanDMDPipeline" not in raw_text
assert "solver" not in cfg.method["sampling"]
assert cfg.method["sampling"]["scheduler"] == "flow_match_euler"
assert cfg.method["sampling"]["trajectory"] == "ode"
assert cfg.method["sampling"]["flow_shift"] == "inherit"
assert "deterministic" not in cfg.method["sampling"]
assert "noise_level" not in cfg.method["sampling"]
assert cfg.method["validation"]["every_steps"] == 10
assert cfg.method["validation"]["num_steps"] == 40
assert cfg.method["validation"]["num_prompts"] == 16
assert cfg.method["validation"]["log_samples"] is True
def test_validation_shard_indices_are_stable_and_padded():
rank0 = validation_shard_indices(5, rank=0, world_size=2)
rank1 = validation_shard_indices(5, rank=1, world_size=2)
assert rank0 == [(0, True), (2, True), (4, True)]
assert rank1 == [(1, True), (3, True), (0, False)]
def test_distributed_k_repeat_indices_repeats_prompts_globally():
rank0 = distributed_k_repeat_indices(
dataset_length=100,
batch_size=6,
repeats_per_prompt=24,
world_size=4,
rank=0,
seed=123,
)
all_indices = []
for rank in range(4):
sample = distributed_k_repeat_indices(
dataset_length=100,
batch_size=6,
repeats_per_prompt=24,
world_size=4,
rank=rank,
seed=123,
)
all_indices.extend(sample.local_indices)
assert rank0.unique_prompt_count == 1
assert len(all_indices) == 24
assert len(set(all_indices)) == 1
def test_validation_caption_puts_rewards_first():
caption = validation_caption(
"a small blue cube",
{
"avg": 0.75,
"pickscore": 0.5,
},
)
assert caption.startswith("avg: 0.7500 | pickscore: 0.5000 | ")
assert caption.endswith("a small blue cube")
def test_media_to_video_array_treats_frame_as_single_frame_video():
frame = torch.ones(3, 4, 5)
video = media_to_video_array(frame)
assert video.shape == (1, 3, 4, 5)
assert video.dtype.name == "uint8"
def test_media_to_video_array_preserves_video_frames():
media = torch.ones(3, 2, 4, 5)
video = media_to_video_array(media)
assert video.shape == (2, 3, 4, 5)
@@ -0,0 +1,157 @@
from dataclasses import dataclass, field
import torch
from fastvideo.train.trainer import Trainer
from fastvideo.train.utils.training_config import (
CheckpointConfig,
DistributedConfig,
ModelTrainingConfig,
OptimizerConfig,
TrackerConfig,
TrainingConfig,
TrainingLoopConfig,
)
class _Tracker:
def __init__(self):
self.logged = []
self.finished = False
def log(self, metrics, step):
self.logged.append((step, metrics))
def finish(self):
self.finished = True
class _Callbacks:
def __init__(self):
self.before_optimizer_steps = 0
self.training_step_ends = 0
def on_train_start(self, method, iteration=0):
pass
def on_before_optimizer_step(self, method, iteration=0):
self.before_optimizer_steps += 1
def on_training_step_end(self, method, metrics, iteration=0):
self.training_step_ends += 1
def on_validation_begin(self, method, iteration=0):
pass
def on_validation_end(self, method, iteration=0):
pass
def on_train_end(self, method, iteration=0):
pass
class _World:
rank = 0
local_rank = 0
class _Method(torch.nn.Module):
def __init__(self):
super().__init__()
self.calls = 0
self.backward_calls = 0
self.optimizer_steps = 0
self.tracker = None
def set_tracker(self, tracker):
self.tracker = tracker
def on_train_start(self):
pass
def manages_optimization(self):
return True
def managed_train_step(self, data_stream, iteration):
batch = next(data_stream)
self.calls += 1
return (
{"total_loss": torch.tensor(float(batch["x"]))},
{},
{"managed_metric": float(iteration)},
)
def backward(self, *args, **kwargs):
self.backward_calls += 1
def optimizers_schedulers_step(self, iteration):
self.optimizer_steps += 1
def optimizers_zero_grad(self, iteration):
pass
class _MethodWithValidation(_Method):
def __init__(self):
super().__init__()
self.validation_iterations = []
def on_validation_begin(self, iteration=0):
self.validation_iterations.append(iteration)
return {"validation/fake": float(iteration)}
def test_trainer_skips_default_optimizer_path_for_managed_methods(monkeypatch):
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.build_tracker", lambda *args, **kwargs: _Tracker())
cfg = TrainingConfig(
distributed=DistributedConfig(),
optimizer=OptimizerConfig(),
loop=TrainingLoopConfig(max_train_steps=1, gradient_accumulation_steps=3),
checkpoint=CheckpointConfig(),
tracker=TrackerConfig(trackers=[]),
model=ModelTrainingConfig(),
)
trainer = Trainer(cfg)
trainer.callbacks = _Callbacks()
method = _Method()
dataloader = [{"x": 2}]
trainer.run(method, dataloader=dataloader, max_steps=1)
assert method.calls == 1
assert method.backward_calls == 0
assert method.optimizer_steps == 0
assert trainer.callbacks.before_optimizer_steps == 0
assert trainer.callbacks.training_step_ends == 1
assert trainer.tracker.logged[0][1]["total_loss"] == 2.0
def test_trainer_logs_method_validation_at_step_zero(monkeypatch):
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: _World())
monkeypatch.setattr("fastvideo.train.trainer.build_tracker", lambda *args, **kwargs: _Tracker())
cfg = TrainingConfig(
distributed=DistributedConfig(),
optimizer=OptimizerConfig(),
loop=TrainingLoopConfig(max_train_steps=1, gradient_accumulation_steps=1),
checkpoint=CheckpointConfig(),
tracker=TrackerConfig(trackers=[]),
model=ModelTrainingConfig(),
)
trainer = Trainer(cfg)
trainer.callbacks = _Callbacks()
method = _MethodWithValidation()
dataloader = [{"x": 2}]
trainer.run(method, dataloader=dataloader, max_steps=1)
assert method.validation_iterations == [0, 1]
assert trainer.tracker.logged[0] == (0, {"validation/fake": 0.0})
@@ -0,0 +1,62 @@
import torch
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.train.models.wan.wan import WanModel
class _CPUWanModel(WanModel):
@property
def device(self) -> torch.device:
return torch.device("cpu")
class _AutocastProbe(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.autocast_enabled: bool | None = None
self.autocast_dtype: torch.dtype | None = None
def forward(
self,
*,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor,
timestep: torch.Tensor,
return_dict: bool,
) -> torch.Tensor:
del encoder_hidden_states, encoder_attention_mask, timestep, return_dict
self.autocast_enabled = torch.is_autocast_enabled("cpu")
self.autocast_dtype = torch.get_autocast_dtype("cpu")
self.hidden_states_dtype = hidden_states.dtype
return hidden_states
def test_wan_predict_noise_uses_training_dtype_autocast_for_fp32_inputs():
model = object.__new__(_CPUWanModel)
model.transformer = _AutocastProbe()
batch = TrainingBatch()
batch.timesteps = torch.tensor([1], dtype=torch.long)
batch.conditional_dict = {
"encoder_hidden_states": torch.randn(1, 4, 8, dtype=torch.float32),
"encoder_attention_mask": torch.ones(1, 4, dtype=torch.float32),
}
noisy_latents = torch.randn(1, 1, 2, 4, 4, dtype=torch.float32)
timestep = torch.tensor([1], dtype=torch.long)
pred_noise = model.predict_noise(
noisy_latents,
timestep,
batch,
conditional=True,
)
assert pred_noise.shape == noisy_latents.shape
assert pred_noise.dtype is torch.bfloat16
assert model.transformer.hidden_states_dtype is torch.bfloat16
assert model.transformer.autocast_enabled is True
assert model.transformer.autocast_dtype is torch.bfloat16