Compare commits
10
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
afbacb5dfa | ||
|
|
c73f91d50e | ||
|
|
7b3cd1cd49 | ||
|
|
649af98bf7 | ||
|
|
7cfa512123 | ||
|
|
0878e4c9e8 | ||
|
|
e82951f1c4 | ||
|
|
d572a4bb84 | ||
|
|
a2fd9dd6e2 | ||
|
|
bb440bf0b7 |
@@ -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()
|
||||
@@ -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
|
||||
@@ -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",
|
||||
]
|
||||
Reference in New Issue
Block a user