llm-fishing / src /image_client.py
reuzed
Warm Modal on load, parallelize fish gen, and improve message pacing.
488cf98
Raw
History Blame Contribute Delete
6.04 kB
"""Image generation — mock placeholders or FLUX Klein backend."""
from __future__ import annotations
import base64
import os
import random
from io import BytesIO
from pathlib import Path
from typing import Any
import httpx
from PIL import Image, ImageDraw
from src.image_utils import (
bait_prompt,
chroma_key_green,
fish_edit_prompt,
fish_greenscreen_prompt,
fish_underwater_prompt,
image_to_png_bytes,
trim_transparent,
)
ROOT = Path(__file__).resolve().parent.parent
IMAGE_SIZE = int(os.getenv("IMAGE_SIZE", "384"))
class ImageClient:
def __init__(self) -> None:
self.mode = os.getenv("IMAGE_MODE", "mock").lower()
self.api_base = os.getenv("IMAGE_API_BASE", "").rstrip("/")
self.flux_model = os.getenv("FLUX_MODEL", "black-forest-labs/FLUX.2-klein-9B")
self.size = IMAGE_SIZE
def _placeholder(self, label: str, *, green: bool = False) -> bytes:
rnd = random.Random(label)
body = tuple(rnd.randint(90, 230) for _ in range(3)) + (255,)
fin = tuple(max(0, c - 50) for c in body[:3]) + (255,)
img = Image.new(
"RGBA", (256, 256), (0, 255, 0, 255) if green else (26, 74, 106, 255)
)
draw = ImageDraw.Draw(img)
if not green:
for _ in range(8):
x, y = rnd.randint(0, 255), rnd.randint(0, 255)
draw.ellipse((x, y, x + 6, y + 6), outline=(180, 220, 235, 120))
draw.polygon([(196, 128), (240, 92), (240, 164)], fill=fin)
draw.ellipse((40, 88, 208, 168), fill=body)
draw.polygon([(110, 92), (140, 60), (160, 94)], fill=fin)
draw.ellipse((62, 108, 86, 132), fill=(255, 255, 255, 255))
draw.ellipse((70, 116, 82, 128), fill=(20, 20, 30, 255))
draw.arc((44, 118, 92, 156), 20, 110, fill=(20, 20, 30, 200), width=3)
if green:
img = chroma_key_green(img)
return image_to_png_bytes(img)
def _generate_remote(self, prompt: str) -> Image.Image:
if not self.api_base:
raise RuntimeError("IMAGE_API_BASE required for remote image mode")
with httpx.Client(timeout=300) as client:
url = self.api_base.rstrip("/")
if not url.endswith(".run"):
url = f"{url}/generate"
r = client.post(
url,
json={"prompt": prompt, "width": self.size, "height": self.size},
)
r.raise_for_status()
b64 = r.json()["image_b64"]
return Image.open(BytesIO(base64.b64decode(b64)))
def _generate_local(self, prompt: str) -> Image.Image:
import torch
from diffusers import Flux2KleinPipeline
pipe = Flux2KleinPipeline.from_pretrained(
self.flux_model,
torch_dtype=torch.bfloat16,
)
pipe.enable_model_cpu_offload()
result = pipe(
prompt=prompt,
num_inference_steps=4,
width=self.size,
height=self.size,
)
return result.images[0]
def _generate_modal(self, prompt: str, source_png: bytes | None = None) -> Image.Image:
import modal
app_name = os.getenv("MODAL_APP_NAME", "llm-fishing")
cls_name = os.getenv("MODAL_IMAGE_CLASS", "ImageGenerator")
ImageGenerator = modal.Cls.from_name(app_name, cls_name)
b64 = ImageGenerator().generate.remote(
prompt,
width=self.size,
height=self.size,
image_b64=base64.b64encode(source_png).decode("ascii") if source_png else None,
)
return Image.open(BytesIO(base64.b64decode(b64)))
def placeholder_png(self, label: str, *, green: bool = False) -> bytes:
return self._placeholder(label, green=green)
def _generate(
self, prompt: str, *, greenscreen: bool, source_png: bytes | None = None
) -> bytes:
try:
if self.mode == "mock":
return self._placeholder(prompt[:20], green=greenscreen)
if self.mode == "modal":
img = self._generate_modal(prompt, source_png)
elif self.mode == "remote":
img = self._generate_remote(prompt)
elif self.mode == "local":
img = self._generate_local(prompt)
else:
raise ValueError(f"Unknown IMAGE_MODE: {self.mode}")
if greenscreen:
img = trim_transparent(chroma_key_green(img))
return image_to_png_bytes(img)
except Exception: # noqa: BLE001
return self._placeholder(prompt[:20], green=greenscreen)
def generate_bait(self, bait: str) -> bytes:
return self._generate(bait_prompt(bait), greenscreen=True)
def generate_fish_underwater(self, fish: dict[str, Any], bait: str = "") -> bytes:
return self._generate(fish_underwater_prompt(fish, bait), greenscreen=False)
def generate_fish_sprite(
self, fish: dict[str, Any], underwater_png: bytes | None = None, bait: str = ""
) -> bytes:
if underwater_png and self.mode == "modal":
return self._generate(
fish_edit_prompt(fish, bait),
greenscreen=True,
source_png=underwater_png,
)
return self._generate(fish_greenscreen_prompt(fish, bait), greenscreen=True)
def warmup(self) -> None:
"""Tiny FLUX render to wake the image container."""
if self.mode == "mock":
return
if self.mode == "modal":
import modal
app_name = os.getenv("MODAL_APP_NAME", "llm-fishing")
cls_name = os.getenv("MODAL_IMAGE_CLASS", "ImageGenerator")
ImageGenerator = modal.Cls.from_name(app_name, cls_name)
ImageGenerator().generate.remote(
"solid bright green square pixel art, game asset",
width=128,
height=128,
)
return
self._generate("solid green pixel", greenscreen=True)