Compare commits
55
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
93a375c654 | ||
|
|
83de8caab4 | ||
|
|
7540d1181f | ||
|
|
9758d931d3 | ||
|
|
7219b3b743 | ||
|
|
7d8604b4fa | ||
|
|
f04e675088 | ||
|
|
9363caf64e | ||
|
|
bb1e8935ee | ||
|
|
58256b1282 | ||
|
|
91d8fb85e6 | ||
|
|
11d55fb5a2 | ||
|
|
d2c5395132 | ||
|
|
704e56674f | ||
|
|
555a600f3f | ||
|
|
c6128ab765 | ||
|
|
6a6ebf1ee3 | ||
|
|
f1f7ac0738 | ||
|
|
47335f09ee | ||
|
|
9e7a30acbc | ||
|
|
022aedb0e1 | ||
|
|
874e4e2f9c | ||
|
|
eca1444101 | ||
|
|
000d48b74d | ||
|
|
2cf9aa0d09 | ||
|
|
be0bacbe35 | ||
|
|
42c2fe6a62 | ||
|
|
8bfcb66447 | ||
|
|
0cc0478463 | ||
|
|
a312d7c6e1 | ||
|
|
b7a923c0bb | ||
|
|
0b4e9764b6 | ||
|
|
df88af316b | ||
|
|
033b75a662 | ||
|
|
dcd82f93e8 | ||
|
|
a22da778a6 | ||
|
|
38513d896c | ||
|
|
375b944b58 | ||
|
|
f4d1b59b9e | ||
|
|
ee4021e5cb | ||
|
|
269544c09b | ||
|
|
2d3bb7c795 | ||
|
|
a2dd0b6398 | ||
|
|
3b9ecb3485 | ||
|
|
87dbb78002 | ||
|
|
06b6c43d3f | ||
|
|
9973307b11 | ||
|
|
ace421bc2b | ||
|
|
fc04d01930 | ||
|
|
c2a86bef4c | ||
|
|
9521077f21 | ||
|
|
e73521d04e | ||
|
|
0c56a4aa53 | ||
|
|
eb33639cbd | ||
|
|
7e70f8598a |
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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.
|
||||
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
@@ -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",
|
||||
]
|
||||
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user