Compare commits

...
16 changed files with 2892 additions and 0 deletions
+56
View File
@@ -0,0 +1,56 @@
# FastVideo Interleave Example
This directory contains a small Python example for running an application-level
InterleaveThinker-style image generation trace on top of FastVideo. The helper
code lives under `fastvideo/workflows/interleave_thinker` and is importable as
`fastvideo.workflows.interleave_thinker`.
## 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 through the local FastVideo image
backend. It writes the generated image and a `trace.json` file under the output
directory. The trace records the generation attempt and omits base64 image
payloads by default.
## How The Example Planner Works
`SinglePromptPlanner` is intentionally minimal: it takes the user instruction
and creates one generation step with that exact prompt. It does not decompose
the instruction, call a learned planner, or create multiple edit steps.
`AcceptAllCritic` is equally small. It accepts the first generated image, so the
example can produce a complete trace without requiring a learned critic model.
The orchestration code still records the planner, generator, and critic result
so a trace has the same shape when a richer app-specific planner or critic is
added later.
## Optional Gemini Backends
The app helper can also use closed-source Google image models through lazy
wrappers when you instantiate the backend directly:
- `fastvideo.workflows.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,106 @@
# SPDX-License-Identifier: Apache-2.0
"""Run a one-step interleaved generation trace through FastVideo.
This is intentionally small: it uses the single-prompt planner and an
accept-all critic from the app-level Interleave helper package.
"""
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.workflows.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()
+1
View File
@@ -34,6 +34,7 @@ fastvideo/
├── training/ # LEGACY monolithic *_training/distillation_pipeline.py
├── worker/ # Multi-process / Ray executors
├── workflow/ # Preprocessing workflow base class
├── workflows/ # Reusable multi-step generation workflows (planner→generate→critic loops); orchestration only, no forward-path logic
├── registry.py # Pipeline-config + model-class lookup (canonical)
├── envs.py # Env-var declarations
├── fastvideo_args.py# Runtime arg dataclass passed through pipelines
@@ -0,0 +1,181 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import json
from pathlib import Path
from fastvideo.workflows.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,258 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import base64
from pathlib import Path
import pytest
from fastvideo.api.compat import (
explicit_request_updates,
legacy_generate_call_to_request,
)
from fastvideo.api.results import GenerationResult
from fastvideo.workflows.interleave_thinker.generator import (
FastVideoImageGeneratorBackend,
NanoBananaImageGeneratorBackend,
)
from fastvideo.workflows.interleave_thinker.orchestrator import InterleaveOrchestrator
from fastvideo.workflows.interleave_thinker.schema import (
CriticDecision,
GeneratedImage,
InterleaveEditRequest,
PlannedInterleaveStep,
)
from fastvideo.workflows.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_nano_banana_client_setup_failure_is_not_retried(tmp_path: Path) -> None:
class FailingClientBackend(NanoBananaImageGeneratorBackend):
def __init__(self) -> None:
super().__init__(
api_key="fake-key",
output_dir=str(tmp_path),
max_attempts=3,
retry_delay_s=999.0,
)
self.client_calls = 0
def _client_instance(self):
self.client_calls += 1
raise RuntimeError("missing Gemini SDK")
backend = FailingClientBackend()
with pytest.raises(RuntimeError, match="missing Gemini SDK"):
backend.generate(InterleaveEditRequest(prompt="draw a red mug"), request_id="client-error")
assert backend.client_calls == 1
def test_nano_banana_generate_config_failure_is_not_retried(tmp_path: Path) -> None:
class FailingConfigBackend(NanoBananaImageGeneratorBackend):
def __init__(self) -> None:
super().__init__(
api_key="fake-key",
output_dir=str(tmp_path),
max_attempts=3,
retry_delay_s=999.0,
)
self.client_calls = 0
self.config_calls = 0
def _client_instance(self):
self.client_calls += 1
return object()
def _make_generate_config(self):
self.config_calls += 1
raise ValueError("invalid Gemini config")
backend = FailingConfigBackend()
with pytest.raises(ValueError, match="invalid Gemini config"):
backend.generate(InterleaveEditRequest(prompt="draw a red mug"), request_id="config-error")
assert backend.client_calls == 1
assert backend.config_calls == 1
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 = []
self.seeds = []
def generate(self, request, *, request_id=None):
del request_id
self.prompts.append(request.prompt)
self.seeds.append(request.seed)
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(),
seed=100,
)
trace = orchestrator.run("initial prompt")
assert trace.success is True
assert generator.prompts == ["initial prompt", "refined prompt"]
assert generator.seeds == [100, 101]
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,246 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import base64
from pathlib import Path
import pytest
from fastvideo.workflows.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.workflows.interleave_thinker.orchestrator import SinglePromptPlanner
from fastvideo.workflows.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:
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
+7
View File
@@ -0,0 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
"""Reusable multi-step generation workflows.
EXPERIMENTAL: this package has no API stability promise yet. Interfaces
may change between releases without deprecation while the training-side
counterpart (planner/critic post-training) lands.
"""
@@ -0,0 +1,106 @@
# SPDX-License-Identifier: Apache-2.0
"""InterleaveThinker example app helpers."""
from fastvideo.workflows.interleave_thinker.config import (
InterleaveCriticConfig,
InterleaveImageBackendConfig,
InterleavePlannerConfig,
InterleaveRunConfig,
InterleaveRunStateConfig,
load_interleave_run_config,
resolve_interleave_instruction,
)
from fastvideo.workflows.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.workflows.interleave_thinker.generator import (
FastVideoImageGeneratorBackend,
ImageGeneratorBackend,
)
from fastvideo.workflows.interleave_thinker.orchestrator import (
AcceptAllCritic,
CriticProvider,
InterleaveOrchestrator,
PlannerProvider,
SinglePromptPlanner,
)
from fastvideo.workflows.interleave_thinker.runner import (
InterleaveRunResult,
run_interleave_config,
)
from fastvideo.workflows.interleave_thinker.schema import (
CriticDecision,
CriticInput,
GeneratedImage,
InterleaveAttempt,
InterleaveEditRequest,
InterleaveTrace,
PlannedInterleaveStep,
PlannerInput,
)
from fastvideo.workflows.interleave_thinker.trace import (
save_trace,
trace_to_dict,
)
from fastvideo.workflows.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",
"InterleaveImageBackendConfig",
"InterleaveOrchestrator",
"InterleavePlannerConfig",
"InterleavePromptItem",
"InterleavePromptResult",
"InterleavePromptSetSummary",
"InterleaveRunConfig",
"InterleaveRunResult",
"InterleaveRunStateConfig",
"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,179 @@
# SPDX-License-Identifier: Apache-2.0
"""Typed config for the InterleaveThinker example app."""
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:
max_attempts_per_step: int = 2
@dataclass
class InterleaveCriticConfig:
kind: Literal["none", "accept_all"] = "accept_all"
@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")
_require_positive_int(config.planner.max_attempts_per_step, "planner.max_attempts_per_step")
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.workflows.interleave_thinker.generator import ImageGeneratorBackend
from fastvideo.workflows.interleave_thinker.runner import (
build_critic,
build_image_backend,
build_planner,
)
from fastvideo.workflows.interleave_thinker.schema import InterleaveTrace
from fastvideo.workflows.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.workflows.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,338 @@
# 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.workflows.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()
client = self._client_instance()
generate_config = self._make_generate_config()
for attempt in range(self.max_attempts):
try:
response = client.models.generate_content(
model=self.model,
contents=contents,
config=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,190 @@
# 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.workflows.interleave_thinker.generator import (
ImageGeneratorBackend,
encode_file_to_base64,
)
from fastvideo.workflows.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,
attempt_index=attempt_index,
)
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,
attempt_index: int = 0,
) -> InterleaveEditRequest:
seed = self.seed
if seed is not None:
seed = seed + attempt_index
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=seed,
num_inference_steps=self.num_inference_steps,
guidance_scale=self.guidance_scale,
)
@@ -0,0 +1,148 @@
# SPDX-License-Identifier: Apache-2.0
"""Config-driven runner for the InterleaveThinker example app."""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from pathlib import Path
from fastvideo.workflows.interleave_thinker.config import (
InterleaveCriticConfig,
InterleaveImageBackendConfig,
InterleavePlannerConfig,
InterleaveRunConfig,
resolve_interleave_instruction,
)
from fastvideo.workflows.interleave_thinker.generator import (
FastVideoImageGeneratorBackend,
ImageGeneratorBackend,
NanoBananaImageGeneratorBackend,
)
from fastvideo.workflows.interleave_thinker.orchestrator import (
AcceptAllCritic,
CriticProvider,
InterleaveOrchestrator,
PlannerProvider,
SinglePromptPlanner,
)
from fastvideo.workflows.interleave_thinker.schema import InterleaveTrace
from fastvideo.workflows.interleave_thinker.trace import save_trace
@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": "single_prompt",
"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:
return SinglePromptPlanner(max_attempts=config.max_attempts_per_step)
def build_critic(config: InterleaveCriticConfig) -> CriticProvider | None:
if config.kind == "none":
return None
return AcceptAllCritic()
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,
)
__all__ = [
"InterleaveRunResult",
"build_critic",
"build_image_backend",
"build_planner",
"resolve_trace_path",
"run_interleave_config",
]
@@ -0,0 +1,93 @@
# SPDX-License-Identifier: Apache-2.0
"""Schemas for the InterleaveThinker example app."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Literal
from pydantic import BaseModel
class InterleaveEditRequest(BaseModel):
"""Image edit/generation request used by Interleave orchestration backends.
Some image-edit callers use `num_inference_step` while FastVideo uses
`num_inference_steps`; accept both and let the plural form win when both
are provided.
"""
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"
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
@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.workflows.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",
]