SolPix / generate_gallery.py
j0no12's picture
Finalize SolPix v1.0.0 under Apache-2.0
a5318e5 verified
Raw History Blame Contribute Delete
6.28 kB
from __future__ import annotations
import argparse
import hashlib
import json
import os
from pathlib import Path
import torch
from PIL import Image
from transformers import AutoTokenizer, T5EncoderModel
from solpix import AutoencoderDCSol, SolPixConfig, SolPixTransformer2D
TEXT_MODEL = "google/flan-t5-base"
PROMPTS = [
"Three Black men sharing french fries at a neighborhood diner, candid documentary photography.",
"A red fox standing in fresh snow beneath pine trees at winter dawn, wildlife photography.",
"A glass greenhouse filled with ferns after rain, soft natural light, botanical photograph.",
"A handmade cobalt blue teapot on a pale stone table, clean studio product photograph.",
"A white sailboat crossing a calm blue bay at golden hour, fine art landscape photograph.",
"An orange cat curled on a wooden chair in a sunlit bookshop, cozy editorial photograph.",
"A small street cafe reflected in wet pavement at night, warm window light, city photograph.",
"A wooden lighthouse on a rocky coast under a cloudy sky, atmospheric landscape photograph.",
"A bowl of ripe peaches on a kitchen counter, morning light, natural still life photograph.",
"A snow-covered cabin among tall pine trees at blue hour, quiet winter landscape photograph.",
"A baker placing fresh bread on a cooling rack in a bright kitchen, documentary photograph.",
"A goldfinch perched on a thin branch among spring blossoms, close-up wildlife photograph.",
"A red bicycle leaning against a brick wall on a leafy neighborhood street, lifestyle photograph.",
"A lemon cake with a slice cut out on a ceramic plate, bright tabletop food photograph.",
"A small observatory beneath a clear star-filled sky, distant mountains, night landscape photograph.",
]
def encode(text: str, tokenizer, encoder, device: torch.device):
tokens = tokenizer(text, return_tensors="pt", truncation=True, max_length=96)
ids = tokens["input_ids"].to(device)
mask = tokens["attention_mask"].to(device).bool()
with torch.inference_mode():
embeddings = encoder(input_ids=ids, attention_mask=mask).last_hidden_state
return embeddings.float(), mask
def main() -> None:
parser = argparse.ArgumentParser(description="Create individual SolPix release samples.")
parser.add_argument("--checkpoint", required=True)
parser.add_argument("--output-dir", required=True)
parser.add_argument("--seed", type=int, default=260926)
parser.add_argument("--steps", type=int, default=40)
parser.add_argument("--guidance-scale", type=float, default=3.5)
args = parser.parse_args()
root = Path(__file__).resolve().parent
os.environ.setdefault("HF_HOME", str(root / "hf_cache"))
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
dtype = torch.bfloat16 if device.type == "cuda" and torch.cuda.is_bf16_supported() else torch.float32
checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False)
model = SolPixTransformer2D(SolPixConfig(**checkpoint["model_config"])).to(device).eval()
model_keys = model.state_dict()
weights = checkpoint.get("ema") or checkpoint["model"]
weights = {name: value for name, value in weights.items() if name in model_keys}
missing = set(model_keys) - set(weights)
if missing:
raise ValueError(f"Checkpoint is missing model weights: {sorted(missing)[:8]}")
model.load_state_dict(weights, strict=True)
tokenizer = AutoTokenizer.from_pretrained(TEXT_MODEL)
text_encoder = T5EncoderModel.from_pretrained(TEXT_MODEL, torch_dtype=dtype).to(device).eval()
text_encoder.requires_grad_(False)
empty, empty_mask = encode("", tokenizer, text_encoder, device)
vae = AutoencoderDCSol.from_pretrained(torch_dtype=dtype).to(device).eval()
scale = float(vae.config.scaling_factor or 0.41407)
time_grid = torch.linspace(1.0, 0.0, args.steps + 1, device=device)
records = []
for index, prompt in enumerate(PROMPTS, start=1):
conditional, conditional_mask = encode(prompt, tokenizer, text_encoder, device)
sample_seed = args.seed + index - 1
generator = torch.Generator(device="cpu").manual_seed(sample_seed)
latents = torch.randn(1, 32, 16, 16, generator=generator).to(device)
with torch.inference_mode():
for current, following in zip(time_grid[:-1], time_grid[1:]):
time = current.expand(1)
with torch.autocast(device_type=device.type, dtype=dtype, enabled=device.type == "cuda"):
v_cond = model(latents, time, conditional, conditional_mask)
v_uncond = model(latents, time, empty, empty_mask)
velocity = v_uncond + args.guidance_scale * (v_cond - v_uncond)
latents = latents + (following - current) * velocity.float()
decoded = vae.decode((latents / scale).to(dtype=dtype), return_dict=False)[0]
pixels = (decoded.float().squeeze(0).permute(1, 2, 0).cpu().clamp(-1, 1) * 127.5 + 127.5)
image = Image.fromarray(pixels.to(torch.uint8).numpy(), mode="RGB")
image_name = f"sample_{index:02d}.png"
image_path = output_dir / image_name
image.save(image_path, format="PNG", optimize=True)
record = {
"index": index,
"image": image_name,
"checkpoint_step": int(checkpoint.get("step", -1)),
"prompt": prompt,
"seed": sample_seed,
"sampling_steps": args.steps,
"guidance_scale": args.guidance_scale,
"text_encoder": TEXT_MODEL,
"decoder": AutoencoderDCSol.model_id,
"decoder_revision": AutoencoderDCSol.revision,
"sha256": hashlib.sha256(image_path.read_bytes()).hexdigest(),
}
(output_dir / f"sample_{index:02d}.json").write_text(json.dumps(record, indent=2) + "\n")
records.append(record)
print(f"[{index}/{len(PROMPTS)}] checkpoint={record['checkpoint_step']} {image_path}", flush=True)
(output_dir / "manifest.json").write_text(json.dumps(records, indent=2) + "\n")
if __name__ == "__main__":
main()