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
54 changed files with 11088 additions and 32 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.
+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,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
+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.
+5 -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
@@ -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"]
+25 -2
View File
@@ -1,6 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
"""RL training methods."""
from fastvideo.train.methods.rl.diffusion_nft import DiffusionNFTMethod
from __future__ import annotations
__all__ = ["DiffusionNFTMethod"]
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)
@@ -10,6 +10,10 @@ 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,
@@ -19,10 +23,12 @@ from fastvideo.train.methods.rl.common.validation import (
__all__ = [
"DiffusionSampler",
"GRPOLossResult",
"KRepeatSample",
"RLValidationConfig",
"SamplingConfig",
"SamplingResult",
"compute_grpo_loss",
"distributed_k_repeat_indices",
"media_to_video_array",
"validation_caption",
+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,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",
]
+103 -8
View File
@@ -1,16 +1,44 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable reward models for training methods."""
from fastvideo.train.methods.rl.rewards.frame_rewards import (
ClipScoreScorer,
PickScoreScorer,
)
from fastvideo.train.methods.rl.rewards.media import (
MultiRewardScorer,
RewardScorer,
select_first_frame,
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,
@@ -18,6 +46,12 @@ def build_multi_reward_scorer(
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 = {
@@ -27,11 +61,72 @@ def build_multi_reward_scorer(
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,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",
]
+19 -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,18 @@ 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,
@@ -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
+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 ---
@@ -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,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