Download generate_gallery.py from solintellegence/SolPix: direct link, hf CLI and curl.
- Browser
- Download file 6.28 kB
-
https://huggingface.co/solintellegence/SolPix/resolve/main/generate_gallery.py
- Command line
-
hf download hf://solintellegence/SolPix/generate_gallery.py
-
curl -L -o generate_gallery.py https://huggingface.co/solintellegence/SolPix/resolve/main/generate_gallery.py
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() | |