BareGit
#!/usr/bin/env python3
"""Opt-in local Krea 2 Turbo integration check.

This script only accepts explicit local paths. It never resolves a Hub model
identifier and is intentionally outside ordinary unit-test discovery.
"""

from __future__ import annotations

import argparse
from pathlib import Path
import time

from PIL import Image
import torch

from diffusion_cli.config import (
    GenerationDefaults,
    ImageGenerationRequest,
    ModelProfile,
    UserConfig,
)
from diffusion_cli.generation_service import GenerationService


def buildParser() -> argparse.ArgumentParser:
    """Build arguments for the explicit local-checkpoint integration run."""

    parser = argparse.ArgumentParser()
    parser.add_argument("--diffusion-model", type=Path, required=True)
    parser.add_argument("--text-encoder", type=Path, required=True)
    parser.add_argument("--vae", type=Path, required=True)
    parser.add_argument("--tokenizer", type=Path, required=True)
    parser.add_argument("--output", type=Path, default=Path("krea2-test.png"))
    parser.add_argument("--seed", type=int, default=0)
    return parser


def main() -> int:
    """Generate one fixed-default Turbo image and report basic statistics."""

    args = buildParser().parse_args()
    profile = ModelProfile(
        name="krea2-turbo",
        architecture="krea2",
        variant="turbo",
        diffusion_model=args.diffusion_model,
        text_encoder=args.text_encoder,
        vae=args.vae,
        tokenizer=args.tokenizer,
    )
    user_config = UserConfig(
        models=None,
        generation=GenerationDefaults(
            output=args.output,
            device="cuda",
            dtype="auto",
        ),
        model_profiles={profile.name: profile},
        default_model=profile.name,
    )
    request = ImageGenerationRequest(
        prompt="a fox walking through fresh snow",
        seed=args.seed,
    )
    service = GenerationService(user_config, model_residency="staged")
    if torch.cuda.is_available():
        torch.cuda.reset_peak_memory_stats()
    started = time.perf_counter()
    paths = service.generateToFiles(request)
    elapsed = time.perf_counter() - started
    peak_memory = (
        torch.cuda.max_memory_allocated() / (1024 ** 3)
        if torch.cuda.is_available()
        else 0.0
    )
    for path in paths:
        with Image.open(path) as image:
            extrema = image.convert("RGB").getextrema()
            print(
                f"generated: {path} size={image.size} "
                f"extrema={extrema}"
            )
    print(f"elapsed_seconds: {elapsed:.2f}")
    print(f"peak_cuda_gib: {peak_memory:.2f}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())