Aniimage-2 / generate_images.py
8BitStudio's picture
Upload generate_images.py
67b7bd0 verified
Raw
History Blame Contribute Delete
61.7 kB
"""
Aniimage Generator β€” Generate anime images from text prompts.
https://huggingface.co/8BitStudio/Aniimage-2
Usage:
pip install -U torch torchvision "diffusers>=0.37.1" "transformers>=4.46,<5" accelerate safetensors pillow huggingface_hub
python generate_hf_aniimage2_corrected.py
"""
import os
import sys
import gc
import json
import torch
import numpy as np
import tkinter as tk
from tkinter import ttk, simpledialog
from pathlib import Path
from PIL import Image, ImageTk
from threading import Thread
# ── Paths ─────────────────────────────────────────────────────────────────────
SCRIPT_DIR = Path(__file__).resolve().parent
MODEL_DIR = SCRIPT_DIR / "models"
OUTPUT_DIR = SCRIPT_DIR / "generated"
# ── HuggingFace repo ─────────────────────────────────────────────────────────
HF_REPO_ID = "8BitStudio/Aniimage-2"
# ── Aniimage-2 training configuration fallback ────────────────────────────────
# The downloaded model_config.json is preferred. These values mirror it so the
# launcher still behaves correctly if only the UNet files were copied locally.
UNET_CONFIG = dict(
sample_size=64,
in_channels=4,
out_channels=4,
block_out_channels=(256, 512, 768, 1024),
layers_per_block=2,
cross_attention_dim=768,
attention_head_dim=8,
down_block_types=("CrossAttnDownBlock2D", "CrossAttnDownBlock2D",
"CrossAttnDownBlock2D", "DownBlock2D"),
up_block_types=("UpBlock2D", "CrossAttnUpBlock2D",
"CrossAttnUpBlock2D", "CrossAttnUpBlock2D"),
)
# Aniimage-2 was trained with this VAE, not the SD 1.x MSE VAE.
VAE_ID = "madebyollin/sdxl-vae-fp16-fix"
CLIP_ID = "openai/clip-vit-large-patch14"
SCHEDULER_LIST = [
"DPM++ 2M Karras",
"DPM++ SDE Karras",
"Euler a",
"Euler",
"DDIM",
]
DEFAULT_NEGATIVE = (
"low quality, ugly, blurry, distorted, deformed, bad anatomy, "
"bad proportions, extra limbs, missing limbs, watermark, text, "
"signature, washed out, flat colors, manga panel, disfigured, "
"poorly drawn, jpeg artifacts, cropped, out of frame"
)
# ── Model discovery ───────────────────────────────────────────────────────────
def _read_json(path: Path):
"""Read a JSON file, returning an empty dict when it is unusable."""
try:
data = json.loads(path.read_text(encoding="utf-8"))
return data if isinstance(data, dict) else {}
except (OSError, ValueError, TypeError):
return {}
def _looks_like_unet_config(config: dict) -> bool:
"""Return True when a config contains the core Diffusers UNet fields."""
required = {
"in_channels", "out_channels", "block_out_channels",
"down_block_types", "up_block_types",
}
return required.issubset(config)
def _find_model_config(model_dir: Path):
"""Find the Aniimage model_config.json that describes training settings."""
if not model_dir.exists():
return None
candidates = [
p for p in model_dir.rglob("model_config.json")
if p.is_file() and ".cache" not in p.parts
]
for path in sorted(candidates, key=lambda p: (len(p.relative_to(model_dir).parts), str(p))):
config = _read_json(path)
if isinstance(config.get("unet"), dict):
return path
return None
def _find_unet_assets(model_dir: Path):
"""Find Aniimage UNet weights/config, including nested repo folders.
Aniimage-2 has been published with an extra ``Aniimage-2/unet`` directory
inside the repository snapshot. Searching recursively keeps the launcher
compatible with that layout as well as normal Diffusers layouts.
"""
if not model_dir.exists():
return None
# Prefer the canonical single-file names, but accept fp16/variant names.
weight_candidates = []
for pattern in (
"diffusion_pytorch_model.safetensors",
"diffusion_pytorch_model.bin",
"diffusion_pytorch_model*.safetensors",
"diffusion_pytorch_model*.bin",
):
weight_candidates.extend(model_dir.rglob(pattern))
# Remove duplicates, metadata files, and anything under the HF cache.
unique_weights = []
seen = set()
for path in weight_candidates:
if not path.is_file() or path.name.endswith(".index.json"):
continue
if ".cache" in path.parts:
continue
key = str(path.resolve())
if key not in seen:
seen.add(key)
unique_weights.append(path)
if unique_weights:
# Exact canonical filenames first, then the shallowest path.
def weight_rank(path: Path):
exact = path.name in {
"diffusion_pytorch_model.safetensors",
"diffusion_pytorch_model.bin",
}
safe = path.suffix.lower() == ".safetensors"
return (not exact, not safe, len(path.relative_to(model_dir).parts), str(path))
weights_path = sorted(unique_weights, key=weight_rank)[0]
# Search from the weights folder upward, then across the model root.
config_candidates = []
current = weights_path.parent
while True:
config_candidates.extend((current / "config.json", current / "model_config.json"))
if current == model_dir or model_dir not in current.parents:
break
current = current.parent
config_candidates.extend(model_dir.rglob("config.json"))
config_candidates.extend(model_dir.rglob("model_config.json"))
config_path = None
seen_configs = set()
for candidate in config_candidates:
if not candidate.is_file() or ".cache" in candidate.parts:
continue
key = str(candidate.resolve())
if key in seen_configs:
continue
seen_configs.add(key)
if _looks_like_unet_config(_read_json(candidate)):
config_path = candidate
break
return {
"kind": "diffusers_weights",
"weights": weights_path,
"config": config_path,
}
# Older local checkpoints are still supported.
for filename in ("ema_unet.pt", "unet.pt"):
candidates = [p for p in model_dir.rglob(filename)
if p.is_file() and ".cache" not in p.parts]
if candidates:
checkpoint = sorted(
candidates,
key=lambda p: (len(p.relative_to(model_dir).parts), str(p)),
)[0]
return {
"kind": "checkpoint",
"weights": checkpoint,
"config": None,
}
return None
def _detect_model_resolution(model_dir: Path) -> str:
"""Infer output resolution from model metadata or the UNet config."""
model_config_path = _find_model_config(model_dir)
if model_config_path:
image_size = _read_json(model_config_path).get("image_size")
if isinstance(image_size, int) and image_size > 0:
return str(image_size)
assets = _find_unet_assets(model_dir)
config_path = assets.get("config") if assets else None
if config_path:
config = _read_json(config_path)
sample_size = config.get("sample_size")
if isinstance(sample_size, (list, tuple)) and sample_size:
sample_size = sample_size[0]
if isinstance(sample_size, int) and sample_size > 0:
return str(sample_size * 8)
return "512" if "aniimage-2" in model_dir.name.lower() else "256"
def download_from_hf():
"""Download Aniimage-2 from Hugging Face if it is not already present."""
try:
from huggingface_hub import snapshot_download
except ImportError:
print("Install huggingface_hub: pip install huggingface_hub")
return None
MODEL_DIR.mkdir(parents=True, exist_ok=True)
aniimage_dir = MODEL_DIR / "Aniimage-2"
existing = _find_unet_assets(aniimage_dir)
existing_config = _find_model_config(aniimage_dir)
if existing and existing_config:
print(f"Aniimage-2 weights already downloaded: {existing['weights']}")
return aniimage_dir
print(f"Downloading Aniimage-2 from {HF_REPO_ID}...")
aniimage_dir.mkdir(parents=True, exist_ok=True)
try:
snapshot_download(
repo_id=HF_REPO_ID,
local_dir=aniimage_dir,
allow_patterns=[
"Aniimage-2/model_config.json",
"Aniimage-2/unet/*",
],
)
except Exception as exc:
print(f"Aniimage-2 download failed: {exc}")
return None
assets = _find_unet_assets(aniimage_dir)
if not assets:
print(
"Aniimage-2 repository downloaded, but no supported UNet weights "
"were found anywhere below:\n"
f" {aniimage_dir}\n"
"Expected diffusion_pytorch_model.safetensors or "
"diffusion_pytorch_model.bin."
)
return None
print(f"Download complete! Found weights at: {assets['weights']}")
return aniimage_dir
def find_models():
"""Find models, including checkpoints nested inside repository folders."""
options = []
if MODEL_DIR.exists():
for d in sorted(MODEL_DIR.iterdir()):
if not d.is_dir():
continue
assets = _find_unet_assets(d)
if not assets:
continue
resolution = _detect_model_resolution(d)
model_kind = (
"safetensors"
if assets["weights"].suffix.lower() == ".safetensors"
else assets["kind"]
)
options.append((model_kind, d.name, d, resolution))
return options
# ── Theme ─────────────────────────────────────────────────────────────────────
C = {
"bg": "#111119",
"panel": "#1b1b2f",
"card": "#24243e",
"card_sel": "#3a3a6e",
"border": "#2e2e52",
"accent": "#6c5ce7",
"accent_h": "#8577ed",
"red": "#e74c3c",
"green": "#2ecc71",
"text": "#eaeaea",
"text2": "#a0a0b8",
"text3": "#60607a",
"input": "#16162a",
"input_fg": "#dcdcf0",
}
class Generator:
def __init__(self, device="cuda"):
self.device = device if device == "cuda" and torch.cuda.is_available() else "cpu"
self.dtype = self._select_dtype()
self.vae = None
self.text_encoder = None
self.tokenizer = None
self.unet = None
self.scheduler = None
self.loaded_checkpoint = None
self.loaded_vae_id = None
self.model_config = {}
self._clip_inner = None
self._clip_full_layers = None
self.latent_size = 64
self.output_size = 512
self.prediction_type = "v_prediction"
self.zero_terminal_snr = True
self.timestep_spacing = "trailing"
self.guidance_rescale = 0.7
self.num_train_timesteps = 1000
self.beta_schedule = "scaled_linear"
self.clip_penultimate = True
self.vae_id = VAE_ID
self.scheduler_name = "DPM++ SDE Karras"
self.cancelled = False
self._configure_backends()
def _select_dtype(self):
if self.device != "cuda":
return torch.float32
bf16_supported = getattr(torch.cuda, "is_bf16_supported", lambda: False)()
return torch.bfloat16 if bf16_supported else torch.float16
def _configure_backends(self):
if self.device == "cuda":
torch.backends.cudnn.benchmark = True
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
if hasattr(torch, "set_float32_matmul_precision"):
torch.set_float32_matmul_precision("high")
def _autocast(self):
return torch.autocast(
device_type="cuda",
dtype=self.dtype,
enabled=(self.device == "cuda"),
)
def switch_device(self, new_device):
"""Switch device and rebuild the models in the correct precision."""
new_device = new_device if new_device == "cuda" and torch.cuda.is_available() else "cpu"
if new_device == self.device:
return
self.device = new_device
self.dtype = self._select_dtype()
self._configure_backends()
self.vae = None
self.text_encoder = None
self.tokenizer = None
self.unet = None
self.scheduler = None
self.loaded_checkpoint = None
self.loaded_vae_id = None
self._clip_inner = None
self._clip_full_layers = None
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
print(f"Switched to {self.device.upper()} ({self.dtype})")
def _load_model_metadata(self, model_path: Path, res_label: str):
"""Load the exact training objective and component IDs for Aniimage-2."""
config_path = _find_model_config(model_path)
config = _read_json(config_path) if config_path else {}
self.model_config = config
self.prediction_type = config.get("prediction_type", "v_prediction")
self.zero_terminal_snr = bool(config.get("zero_terminal_snr", True))
self.timestep_spacing = config.get(
"timestep_spacing",
"trailing" if self.zero_terminal_snr else "leading",
)
self.guidance_rescale = float(config.get("guidance_rescale", 0.7))
self.num_train_timesteps = int(config.get("num_train_timesteps", 1000))
self.beta_schedule = config.get("beta_schedule", "scaled_linear")
self.clip_penultimate = bool(config.get("clip_penultimate", True))
self.vae_id = config.get("vae", VAE_ID)
try:
fallback_size = int(res_label)
except (TypeError, ValueError):
fallback_size = 512
self.output_size = int(config.get("image_size", fallback_size))
self.latent_size = self.output_size // 8
if config_path:
print(f"Using model metadata: {config_path}")
else:
print("model_config.json was not found; using Aniimage-2 defaults.")
def _apply_clip_layer_mode(self):
if self.text_encoder is None:
return
self._clip_inner = getattr(self.text_encoder, "text_model", self.text_encoder)
if self._clip_full_layers is None:
self._clip_full_layers = self._clip_inner.encoder.layers
if self.clip_penultimate:
self._clip_inner.encoder.layers = self._clip_full_layers[:-1]
print("Text encoder: CLIP penultimate layer (matches training).")
else:
self._clip_inner.encoder.layers = self._clip_full_layers
def load_shared(self):
from diffusers import AutoencoderKL
from transformers import (CLIPConfig, CLIPTextConfig,
CLIPTextModel, CLIPTokenizer)
load_kwargs = {"low_cpu_mem_usage": True}
if self.device == "cuda":
load_kwargs["torch_dtype"] = self.dtype
if self.vae is None or self.loaded_vae_id != self.vae_id:
print(f"Loading VAE: {self.vae_id}...")
self.vae = AutoencoderKL.from_pretrained(
self.vae_id,
**load_kwargs,
).to(self.device).eval()
self.vae.requires_grad_(False)
self.vae.enable_slicing()
if self.device == "cuda":
self.vae.to(memory_format=torch.channels_last)
self.loaded_vae_id = self.vae_id
if self.text_encoder is None:
print(f"Loading CLIP text encoder: {CLIP_ID}...")
self.tokenizer = CLIPTokenizer.from_pretrained(CLIP_ID)
# Explicitly pass the nested text config. This avoids the
# CLIPConfig.hidden_size crash seen with some Transformers builds.
clip_config = CLIPConfig.from_pretrained(CLIP_ID)
text_config = getattr(clip_config, "text_config", None)
if isinstance(text_config, dict):
text_config = CLIPTextConfig.from_dict(text_config)
if not isinstance(text_config, CLIPTextConfig):
text_config = CLIPTextConfig.from_pretrained(CLIP_ID)
self.text_encoder = CLIPTextModel.from_pretrained(
CLIP_ID,
config=text_config,
**load_kwargs,
).to(self.device).eval()
self.text_encoder.requires_grad_(False)
self._clip_full_layers = None
self._apply_clip_layer_mode()
self.scheduler = self._make_scheduler(self.scheduler_name)
print("Shared models loaded.")
def _make_scheduler(self, name="DPM++ SDE Karras"):
from diffusers import (DDIMScheduler, DPMSolverMultistepScheduler,
EulerAncestralDiscreteScheduler,
EulerDiscreteScheduler)
base = dict(
num_train_timesteps=self.num_train_timesteps,
beta_schedule=self.beta_schedule,
prediction_type=self.prediction_type,
rescale_betas_zero_snr=self.zero_terminal_snr,
timestep_spacing=self.timestep_spacing,
)
if name == "DPM++ 2M Karras":
return DPMSolverMultistepScheduler(
**base, algorithm_type="dpmsolver++",
solver_order=2, use_karras_sigmas=True)
if name == "DPM++ SDE Karras":
return DPMSolverMultistepScheduler(
**base, algorithm_type="sde-dpmsolver++",
solver_order=2, use_karras_sigmas=True)
if name == "Euler a":
return EulerAncestralDiscreteScheduler(**base)
if name == "Euler":
return EulerDiscreteScheduler(**base)
return DDIMScheduler(
**base, clip_sample=False, set_alpha_to_one=False)
def set_scheduler(self, name):
self.scheduler_name = name
self.scheduler = self._make_scheduler(name)
def load_model(self, model_path: Path, res_label: str = "512"):
if str(model_path) == self.loaded_checkpoint:
return
from diffusers import UNet2DConditionModel
assets = _find_unet_assets(model_path)
if not assets:
raise FileNotFoundError(
f"No supported UNet weights found anywhere inside {model_path}"
)
self._load_model_metadata(model_path, res_label)
self.load_shared()
weights_path = assets["weights"]
config_path = assets.get("config")
suffix = weights_path.suffix.lower()
same_dir_config = weights_path.parent / "config.json"
print(
f"Loading UNet from {weights_path} "
f"({self.output_size}px, {self.prediction_type}, {self.dtype})..."
)
self.unet = None
if torch.cuda.is_available():
torch.cuda.empty_cache()
loaded_directly = False
if same_dir_config.exists() and _looks_like_unet_config(_read_json(same_dir_config)):
try:
kwargs = {"low_cpu_mem_usage": True}
if self.device == "cuda":
kwargs["torch_dtype"] = self.dtype
if suffix == ".safetensors":
kwargs["use_safetensors"] = True
elif suffix == ".bin":
kwargs["use_safetensors"] = False
self.unet = UNet2DConditionModel.from_pretrained(
weights_path.parent,
**kwargs,
).to(self.device)
loaded_directly = True
print("Loaded the repository UNet config and weights directly.")
except Exception as exc:
print(f"Direct Diffusers load failed ({exc}); loading manually.")
if not loaded_directly:
if config_path:
unet_config = _read_json(config_path)
print(f"Using UNet config: {config_path}")
elif isinstance(self.model_config.get("unet"), dict):
unet_config = dict(self.model_config["unet"])
print("Using UNet config from model_config.json.")
else:
unet_config = dict(UNET_CONFIG)
print("Using built-in Aniimage-2 UNet config.")
unet_config["sample_size"] = self.latent_size
self.unet = UNet2DConditionModel.from_config(unet_config)
if suffix == ".safetensors":
from safetensors.torch import load_file
state = load_file(str(weights_path), device="cpu")
else:
try:
state = torch.load(weights_path, map_location="cpu", weights_only=True)
except TypeError:
state = torch.load(weights_path, map_location="cpu")
if weights_path.name == "ema_unet.pt" and isinstance(state, dict) and "shadow_params" in state:
params = dict(self.unet.named_parameters())
keys = list(params.keys())
if len(state["shadow_params"]) != len(keys):
raise RuntimeError("EMA parameter count does not match the UNet.")
for key, shadow_param in zip(keys, state["shadow_params"]):
params[key].data.copy_(shadow_param)
else:
if isinstance(state, dict) and "state_dict" in state:
state = state["state_dict"]
if isinstance(state, dict) and state and all(
isinstance(key, str) and key.startswith("module.") for key in state
):
state = {key[7:]: value for key, value in state.items()}
self.unet.load_state_dict(state, strict=True)
if self.device == "cuda":
self.unet = self.unet.to(device=self.device, dtype=self.dtype)
else:
self.unet = self.unet.to(self.device)
sample_size = self.unet.config.sample_size
if isinstance(sample_size, (list, tuple)) and sample_size:
sample_size = sample_size[0]
if isinstance(sample_size, int) and sample_size > 0:
self.latent_size = sample_size
vae_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
self.output_size = sample_size * vae_factor
self.unet.eval().requires_grad_(False)
if self.device == "cuda":
self.unet.to(memory_format=torch.channels_last)
self.scheduler = self._make_scheduler(self.scheduler_name)
self.loaded_checkpoint = str(model_path)
print(
f"Ready at {self.output_size}x{self.output_size}; "
f"zero-SNR={self.zero_terminal_snr}, spacing={self.timestep_spacing}, "
f"CFG rescale={self.guidance_rescale}."
)
def _encode_prompts(self, prompt: str, negative_prompt: str):
tokens = self.tokenizer(
[negative_prompt or "", prompt],
padding="max_length",
max_length=self.tokenizer.model_max_length,
truncation=True,
return_tensors="pt",
)
with self._autocast():
return self.text_encoder(tokens.input_ids.to(self.device))[0]
@staticmethod
def _cfg_rescale(noise_cfg, noise_text, amount):
if amount <= 0:
return noise_cfg
dims = tuple(range(1, noise_cfg.ndim))
std_text = noise_text.std(dim=dims, keepdim=True)
std_cfg = noise_cfg.std(dim=dims, keepdim=True).clamp_min(1e-6)
noise_rescaled = noise_cfg * (std_text / std_cfg)
return amount * noise_rescaled + (1.0 - amount) * noise_cfg
def _decode_latents(self, latents, post_process=False):
del post_process # Kept for compatibility with the preview callbacks.
decode_dtype = self.dtype if self.device == "cuda" else torch.float32
scaled = (latents / self.vae.config.scaling_factor).to(dtype=decode_dtype)
with self._autocast():
image = self.vae.decode(scaled).sample
image = (image.float() / 2 + 0.5).clamp(0, 1)
image = image[0].cpu().permute(1, 2, 0).numpy()
image = (image * 255).round().astype("uint8")
return Image.fromarray(image)
@torch.inference_mode()
def generate(self, prompt: str, negative_prompt: str = "",
steps: int = 50, guidance_scale: float = 7.5,
seed: int = -1, preview_callback=None,
preview_every: int = 5) -> tuple:
if seed < 0:
seed = torch.randint(0, 2**32, (1,)).item()
generator = torch.Generator(device=self.device).manual_seed(seed)
embeddings = self._encode_prompts(prompt, negative_prompt)
scheduler = self._make_scheduler(self.scheduler_name)
scheduler.set_timesteps(int(steps), device=self.device)
in_channels = int(self.unet.config.in_channels)
latents = torch.randn(
(1, in_channels, self.latent_size, self.latent_size),
generator=generator,
device=self.device,
dtype=torch.float32,
) * scheduler.init_noise_sigma
total_steps = len(scheduler.timesteps)
preview_interval = max(1, int(preview_every))
for step_i, timestep in enumerate(scheduler.timesteps):
if self.cancelled:
return None, seed
latent_input = torch.cat([latents, latents], dim=0)
latent_input = scheduler.scale_model_input(latent_input, timestep)
with self._autocast():
prediction = self.unet(
latent_input,
timestep,
encoder_hidden_states=embeddings,
).sample
pred_negative, pred_text = prediction.chunk(2)
prediction = pred_negative + float(guidance_scale) * (pred_text - pred_negative)
prediction = self._cfg_rescale(
prediction, pred_text, self.guidance_rescale)
latents = scheduler.step(prediction, timestep, latents).prev_sample
if (preview_callback
and (step_i + 1) % preview_interval == 0
and step_i < total_steps - 1):
preview_callback(
self._decode_latents(latents),
step_i + 1,
total_steps,
)
return self._decode_latents(latents), seed
@torch.inference_mode()
def refine(self, source_image: Image.Image, prompt: str,
negative_prompt: str = "", extra_steps: int = 20,
strength: float = 0.35, guidance_scale: float = 7.5,
preview_callback=None, preview_every: int = 5) -> Image.Image:
img = source_image.convert("RGB").resize(
(self.output_size, self.output_size), Image.LANCZOS)
img_tensor = torch.from_numpy(np.array(img)).float().div(127.5).sub(1.0)
img_tensor = img_tensor.permute(2, 0, 1).unsqueeze(0).to(self.device)
img_tensor = img_tensor.to(
dtype=self.dtype if self.device == "cuda" else torch.float32)
with self._autocast():
latents = self.vae.encode(img_tensor).latent_dist.sample()
latents = (latents * self.vae.config.scaling_factor).float()
embeddings = self._encode_prompts(prompt, negative_prompt)
scheduler = self._make_scheduler(self.scheduler_name)
scheduler.set_timesteps(int(extra_steps), device=self.device)
start_step = max(0, int(len(scheduler.timesteps) * (1.0 - float(strength))))
timesteps = scheduler.timesteps[start_step:]
if len(timesteps) == 0:
return source_image.copy()
noise = torch.randn_like(latents)
latents = scheduler.add_noise(latents, noise, timesteps[:1])
total_steps = len(timesteps)
preview_interval = max(1, int(preview_every))
for step_i, timestep in enumerate(timesteps):
if self.cancelled:
return None
latent_input = torch.cat([latents, latents], dim=0)
latent_input = scheduler.scale_model_input(latent_input, timestep)
with self._autocast():
prediction = self.unet(
latent_input,
timestep,
encoder_hidden_states=embeddings,
).sample
pred_negative, pred_text = prediction.chunk(2)
prediction = pred_negative + float(guidance_scale) * (pred_text - pred_negative)
prediction = self._cfg_rescale(
prediction, pred_text, self.guidance_rescale)
latents = scheduler.step(prediction, timestep, latents).prev_sample
if (preview_callback
and (step_i + 1) % preview_interval == 0
and step_i < total_steps - 1):
preview_callback(
self._decode_latents(latents),
step_i + 1,
total_steps,
)
return self._decode_latents(latents)
# ── GUI ───────────────────────────────────────────────────────────────────────
class App:
def __init__(self):
self.gen = Generator()
self.models = find_models()
self.generated_images = []
self.generated_seeds = []
self.photo_refs = []
self.generating = False
self.selected_index = None
self.root = tk.Tk()
self.root.title("Aniimage")
self.root.configure(bg=C["bg"])
self.root.resizable(True, True)
self.root.geometry("900x780")
self.root.minsize(640, 500)
self._setup_styles()
self._build_ui()
def _setup_styles(self):
s = ttk.Style()
s.theme_use("clam")
# Base
s.configure(".", background=C["bg"], foreground=C["text"], font=("Segoe UI", 10))
s.configure("TFrame", background=C["bg"])
s.configure("TLabel", background=C["bg"], foreground=C["text"])
s.configure("TCheckbutton", background=C["bg"], foreground=C["text"])
# Combobox β€” readable text
s.configure("TCombobox", fieldbackground=C["input"], foreground=C["input_fg"],
selectbackground=C["accent"], selectforeground="#ffffff",
arrowcolor=C["text2"], padding=4)
s.map("TCombobox",
fieldbackground=[("readonly", C["input"])],
foreground=[("readonly", C["input_fg"])],
selectbackground=[("readonly", C["accent"])],
selectforeground=[("readonly", "#ffffff")])
# Combobox dropdown list colors
self.root.option_add("*TCombobox*Listbox.background", C["input"])
self.root.option_add("*TCombobox*Listbox.foreground", C["input_fg"])
self.root.option_add("*TCombobox*Listbox.selectBackground", C["accent"])
self.root.option_add("*TCombobox*Listbox.selectForeground", "#ffffff")
self.root.option_add("*TCombobox*Listbox.font", ("Segoe UI", 10))
# Spinbox
s.configure("TSpinbox", fieldbackground=C["input"], foreground=C["input_fg"],
arrowcolor=C["text2"], padding=3)
# Buttons
s.configure("TButton", font=("Segoe UI", 10), padding=(14, 7),
background=C["card"], foreground=C["text"])
s.map("TButton", background=[("active", C["card_sel"]), ("disabled", C["bg"])],
foreground=[("disabled", C["text3"])])
s.configure("Go.TButton", font=("Segoe UI", 11, "bold"), padding=(20, 9),
background=C["accent"], foreground="#ffffff")
s.map("Go.TButton", background=[("active", C["accent_h"]),
("disabled", C["border"])])
s.configure("Stop.TButton", font=("Segoe UI", 10, "bold"), padding=(14, 7),
background=C["red"], foreground="#ffffff")
s.map("Stop.TButton", background=[("active", "#c0392b"),
("disabled", C["border"])])
# Labelframe
s.configure("TLabelframe", background=C["bg"], foreground=C["text2"])
s.configure("TLabelframe.Label", background=C["bg"],
foreground=C["text2"], font=("Segoe UI", 9, "bold"))
# Scrollbar
s.configure("Vertical.TScrollbar", background=C["card"],
troughcolor=C["bg"], arrowcolor=C["text3"])
def _make_entry(self, parent, font_size=11, dim=False):
"""Create a styled tk.Entry with readable text."""
return tk.Entry(parent, font=("Segoe UI", font_size),
bg=C["input"], fg=C["input_fg"] if not dim else C["text2"],
insertbackground=C["input_fg"],
relief="flat", bd=6,
selectbackground=C["accent"], selectforeground="#ffffff",
highlightthickness=1, highlightcolor=C["accent"],
highlightbackground=C["border"])
def _build_ui(self):
# ── Header ────────────────────────────────────────────────────────
header = tk.Frame(self.root, bg=C["panel"], padx=20, pady=12)
header.pack(fill=tk.X)
tk.Label(header, text="Aniimage", bg=C["panel"], fg=C["accent"],
font=("Segoe UI", 20, "bold")).pack(side=tk.LEFT)
tk.Label(header, text="by 8BitStudio", bg=C["panel"], fg=C["text3"],
font=("Segoe UI", 10)).pack(side=tk.LEFT, padx=(10, 0), pady=(6, 0))
# Device switch β€” right side of header
device_frame = tk.Frame(header, bg=C["panel"])
device_frame.pack(side=tk.RIGHT)
tk.Label(device_frame, text="Device:", bg=C["panel"], fg=C["text2"],
font=("Segoe UI", 9)).pack(side=tk.LEFT, padx=(0, 5))
self.device_var = tk.StringVar(value="GPU" if self.gen.device == "cuda" else "CPU")
devices = ["GPU", "CPU"] if torch.cuda.is_available() else ["CPU"]
device_combo = ttk.Combobox(device_frame, textvariable=self.device_var,
values=devices, state="readonly", width=5)
device_combo.pack(side=tk.LEFT)
device_combo.bind("<<ComboboxSelected>>", self._on_device_change)
# ── Main content β€” two-column: controls left, images right ────────
main = tk.Frame(self.root, bg=C["bg"])
main.pack(fill=tk.BOTH, expand=True, padx=12, pady=(8, 12))
# Left panel (controls)
left = tk.Frame(main, bg=C["panel"], width=340, padx=16, pady=12)
left.pack(side=tk.LEFT, fill=tk.Y, padx=(0, 8))
left.pack_propagate(False)
# Right panel (image grid)
right = tk.Frame(main, bg=C["bg"])
right.pack(side=tk.LEFT, fill=tk.BOTH, expand=True)
self._build_controls(left)
self._build_grid(right)
def _build_controls(self, parent):
# ── Model ─────────────────────────────────────────────────────────
tk.Label(parent, text="Model", bg=C["panel"], fg=C["text2"],
font=("Segoe UI", 9, "bold")).pack(anchor=tk.W)
self.model_var = tk.StringVar()
model_names = [m[1] for m in self.models] or ["No models found"]
self.model_combo = ttk.Combobox(parent, textvariable=self.model_var,
values=model_names, state="readonly", width=32)
self.model_combo.pack(fill=tk.X, pady=(3, 12))
self.model_combo.current(len(model_names) - 1)
# ── Prompt ────────────────────────────────────────────────────────
tk.Label(parent, text="Prompt", bg=C["panel"], fg=C["text2"],
font=("Segoe UI", 9, "bold")).pack(anchor=tk.W)
self.prompt_entry = self._make_entry(parent)
self.prompt_entry.pack(fill=tk.X, pady=(3, 8))
self.prompt_entry.insert(0, "a smiling anime girl with long blue hair")
self.prompt_entry.bind("<Return>", lambda e: self.on_generate())
# ── Negative prompt ───────────────────────────────────────────────
tk.Label(parent, text="Negative prompt", bg=C["panel"], fg=C["text3"],
font=("Segoe UI", 9)).pack(anchor=tk.W)
self.neg_entry = self._make_entry(parent, font_size=9, dim=True)
self.neg_entry.pack(fill=tk.X, pady=(3, 12))
self.neg_entry.insert(0, DEFAULT_NEGATIVE)
# ── Settings grid ─────────────────────────────────────────────────
grid = tk.Frame(parent, bg=C["panel"])
grid.pack(fill=tk.X, pady=(0, 8))
# Row 1: Scheduler
tk.Label(grid, text="Scheduler", bg=C["panel"], fg=C["text2"],
font=("Segoe UI", 9)).grid(row=0, column=0, sticky="w", pady=(0, 6))
self.scheduler_var = tk.StringVar(value="DPM++ SDE Karras")
sched_combo = ttk.Combobox(grid, textvariable=self.scheduler_var,
values=SCHEDULER_LIST, state="readonly", width=18)
sched_combo.grid(row=0, column=1, columnspan=3, sticky="ew", padx=(8, 0), pady=(0, 6))
sched_combo.bind("<<ComboboxSelected>>", self._on_scheduler_change)
# Row 2: Steps, CFG, Count
tk.Label(grid, text="Steps", bg=C["panel"], fg=C["text2"],
font=("Segoe UI", 9)).grid(row=1, column=0, sticky="w", pady=(0, 6))
self.steps_var = tk.StringVar(value="50")
tk.Entry(grid, textvariable=self.steps_var, width=5, font=("Segoe UI", 10),
bg=C["input"], fg=C["input_fg"], insertbackground=C["input_fg"],
relief="flat", bd=4).grid(row=1, column=1, sticky="w", padx=(8, 12), pady=(0, 6))
tk.Label(grid, text="CFG", bg=C["panel"], fg=C["text2"],
font=("Segoe UI", 9)).grid(row=1, column=2, sticky="w", pady=(0, 6))
self.cfg_var = tk.StringVar(value="7.5")
tk.Entry(grid, textvariable=self.cfg_var, width=5, font=("Segoe UI", 10),
bg=C["input"], fg=C["input_fg"], insertbackground=C["input_fg"],
relief="flat", bd=4).grid(row=1, column=3, sticky="w", padx=(8, 0), pady=(0, 6))
# Row 3: Count, Live preview
tk.Label(grid, text="Count", bg=C["panel"], fg=C["text2"],
font=("Segoe UI", 9)).grid(row=2, column=0, sticky="w", pady=(0, 6))
self.count_var = tk.StringVar(value="4")
ttk.Spinbox(grid, from_=1, to=12, textvariable=self.count_var, width=4,
font=("Segoe UI", 10)).grid(row=2, column=1, sticky="w", padx=(8, 12), pady=(0, 6))
self.live_preview_var = tk.BooleanVar(value=False)
ttk.Checkbutton(grid, text="Live preview",
variable=self.live_preview_var).grid(
row=2, column=2, columnspan=2, sticky="w", pady=(0, 6))
grid.columnconfigure(1, weight=1)
grid.columnconfigure(3, weight=1)
# ── Buttons ───────────────────────────────────────────────────────
btn_frame = tk.Frame(parent, bg=C["panel"])
btn_frame.pack(fill=tk.X, pady=(0, 10))
self.gen_btn = ttk.Button(btn_frame, text="Generate", command=self.on_generate,
style="Go.TButton")
self.gen_btn.pack(fill=tk.X, pady=(0, 5))
btn_row = tk.Frame(btn_frame, bg=C["panel"])
btn_row.pack(fill=tk.X)
self.stop_btn = ttk.Button(btn_row, text="Stop", command=self.on_stop,
state=tk.DISABLED, style="Stop.TButton")
self.stop_btn.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=(0, 3))
self.save_btn = ttk.Button(btn_row, text="Save Selected", command=self.on_save,
state=tk.DISABLED)
self.save_btn.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=(3, 3))
self.save_all_btn = ttk.Button(btn_row, text="Save All", command=self.on_save_all,
state=tk.DISABLED)
self.save_all_btn.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=(3, 0))
# ── Prompt queue ─────────────────────────────────────────────────
sep = tk.Frame(parent, height=1, bg=C["border"])
sep.pack(fill=tk.X, pady=(8, 10))
tk.Label(parent, text="Prompt Queue", bg=C["panel"], fg=C["text2"],
font=("Segoe UI", 9, "bold")).pack(anchor=tk.W)
queue_input = tk.Frame(parent, bg=C["panel"])
queue_input.pack(fill=tk.X, pady=(4, 0))
self.queue_entry = self._make_entry(queue_input, font_size=9)
self.queue_entry.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=(0, 4))
self.queue_entry.bind("<Return>", lambda e: self._queue_add())
ttk.Button(queue_input, text="Add", width=4,
command=self._queue_add).pack(side=tk.LEFT)
self.queue_listbox = tk.Listbox(
parent, height=4, bg=C["input"], fg=C["input_fg"],
selectbackground=C["accent"], selectforeground="#fff",
font=("Segoe UI", 9), activestyle="none",
relief="flat", bd=4, highlightthickness=0)
self.queue_listbox.pack(fill=tk.X, pady=(5, 0))
queue_btns = tk.Frame(parent, bg=C["panel"])
queue_btns.pack(fill=tk.X, pady=(4, 0))
self.queue_run_btn = ttk.Button(queue_btns, text="Run Queue",
command=self.on_run_queue, style="Go.TButton")
self.queue_run_btn.pack(side=tk.LEFT, padx=(0, 4))
for txt, cmd in [("Remove", self._queue_remove), ("Clear", self._queue_clear),
("Up", self._queue_move_up), ("Down", self._queue_move_down),
("+ Current", self._queue_add_current)]:
ttk.Button(queue_btns, text=txt, command=cmd).pack(side=tk.LEFT, padx=2)
# ── Status bar ────────────────────────────────────────────────────
status_frame = tk.Frame(parent, bg=C["bg"], padx=8, pady=6)
status_frame.pack(fill=tk.X, side=tk.BOTTOM)
self.status_var = tk.StringVar(value="Ready")
tk.Label(status_frame, textvariable=self.status_var,
bg=C["bg"], fg=C["green"], font=("Segoe UI", 9),
anchor="w").pack(fill=tk.X)
def _build_grid(self, parent):
self.canvas = tk.Canvas(parent, bg=C["bg"], highlightthickness=0)
scrollbar = ttk.Scrollbar(parent, orient=tk.VERTICAL, command=self.canvas.yview)
self.grid_frame = tk.Frame(self.canvas, bg=C["bg"])
self.grid_frame.bind("<Configure>",
lambda e: self.canvas.configure(
scrollregion=self.canvas.bbox("all")))
self.canvas_window = self.canvas.create_window((0, 0), window=self.grid_frame,
anchor="nw")
self.canvas.configure(yscrollcommand=scrollbar.set)
self.canvas.pack(side=tk.LEFT, fill=tk.BOTH, expand=True)
scrollbar.pack(side=tk.RIGHT, fill=tk.Y)
self.canvas.bind("<Configure>", self._on_canvas_resize)
self.canvas.bind_all("<MouseWheel>",
lambda e: self.canvas.yview_scroll(
int(-1 * (e.delta / 120)), "units"))
self.placeholder = tk.Label(self.grid_frame,
text="Generated images\nwill appear here",
bg=C["bg"], fg=C["text3"],
font=("Segoe UI", 13), justify="center")
self.placeholder.grid(row=0, column=0, pady=80)
# ── Event handlers ────────────────────────────────────────────────────
def _on_device_change(self, event=None):
choice = self.device_var.get()
new_dev = "cuda" if choice == "GPU" else "cpu"
self.status_var.set(f"Switching to {choice}...")
self.root.update()
self.gen.switch_device(new_dev)
self.status_var.set(f"Now using {choice}")
def _on_scheduler_change(self, event=None):
name = self.scheduler_var.get()
self.gen.set_scheduler(name)
self.status_var.set(f"Scheduler: {name}")
def _on_canvas_resize(self, event):
self.canvas.itemconfig(self.canvas_window, width=event.width)
if self.generated_images:
self._layout_grid()
def _get_grid_cols(self):
canvas_w = self.canvas.winfo_width()
if canvas_w < 50:
canvas_w = 560
tile_size = self._get_tile_size()
return max(1, canvas_w // (tile_size + 16))
def _get_tile_size(self):
n = len(self.generated_images)
if n <= 2: return 260
elif n <= 4: return 220
elif n <= 6: return 180
else: return 160
def _layout_grid(self):
for w in self.grid_frame.winfo_children():
w.destroy()
self.photo_refs.clear()
if not self.generated_images:
return
tile_size = self._get_tile_size()
cols = self._get_grid_cols()
for i, (img, seed) in enumerate(zip(self.generated_images, self.generated_seeds)):
row, col = divmod(i, cols)
is_selected = (i == self.selected_index)
card_bg = C["accent"] if is_selected else C["card"]
card = tk.Frame(self.grid_frame, bg=card_bg, padx=3, pady=3)
card.grid(row=row, column=col, padx=5, pady=5, sticky="nsew")
display = img.resize((tile_size, tile_size), Image.LANCZOS)
photo = ImageTk.PhotoImage(display)
self.photo_refs.append(photo)
img_label = tk.Label(card, image=photo, bg=card_bg, bd=0)
img_label.pack()
img_label.bind("<Button-1>", lambda e, idx=i: self._select_image(idx))
img_label.bind("<Button-3>", lambda e, idx=i: self._show_refine_menu(e, idx))
tk.Label(card, text=f"seed: {seed}", bg=card_bg,
fg=C["text3"], font=("Segoe UI", 8)).pack()
for c in range(cols):
self.grid_frame.columnconfigure(c, weight=1)
def _select_image(self, idx):
if idx >= len(self.generated_images):
return
self.selected_index = idx
self.save_btn.configure(state=tk.NORMAL)
self.status_var.set(f"Selected image {idx + 1} (seed: {self.generated_seeds[idx]})")
self._layout_grid()
def _show_refine_menu(self, event, idx):
if self.generating:
return
menu = tk.Menu(self.root, tearoff=0, bg=C["card"], fg=C["text"],
activebackground=C["accent"], activeforeground="#fff",
font=("Segoe UI", 10), bd=0)
menu.add_command(label=" Refine (more steps)... ",
command=lambda: self._ask_refine(idx))
menu.tk_popup(event.x_root, event.y_root)
def _ask_refine(self, idx):
extra = simpledialog.askinteger(
"Refine Image", "Extra denoising steps:",
initialvalue=20, minvalue=5, maxvalue=200, parent=self.root)
if extra is None:
return
self._select_image(idx)
self.generating = True
self.gen.cancelled = False
self.gen_btn.configure(state=tk.DISABLED)
self.stop_btn.configure(state=tk.NORMAL)
self.status_var.set(f"Refining image {idx + 1}...")
self.root.update()
Thread(target=self._refine_thread, args=(idx, extra), daemon=True).start()
def _refine_thread(self, idx, extra_steps):
try:
source = self.generated_images[idx]
prompt = self.prompt_entry.get().strip()
neg = self.neg_entry.get().strip()
cfg = float(self.cfg_var.get())
callback = self._show_preview if self.live_preview_var.get() else None
refined = self.gen.refine(
source_image=source, prompt=prompt, negative_prompt=neg,
extra_steps=extra_steps, guidance_scale=cfg,
preview_callback=callback, preview_every=5)
if refined is not None:
self.generated_images[idx] = refined
self.generated_seeds[idx] = f"{self.generated_seeds[idx]}+R{extra_steps}"
self._layout_grid()
self.status_var.set(f"Refined image {idx + 1}")
else:
self.status_var.set("Refine stopped.")
self.root.update()
except Exception as e:
self.status_var.set(f"Refine error: {e}")
import traceback; traceback.print_exc()
finally:
self.generating = False
self.gen.cancelled = False
self.gen_btn.configure(state=tk.NORMAL)
self.stop_btn.configure(state=tk.DISABLED)
# ── Queue ─────────────────────────────────────────────────────────────
def _queue_add(self):
text = self.queue_entry.get().strip()
if text:
self.queue_listbox.insert(tk.END, text)
self.queue_entry.delete(0, tk.END)
def _queue_add_current(self):
text = self.prompt_entry.get().strip()
if text:
self.queue_listbox.insert(tk.END, text)
def _queue_remove(self):
sel = self.queue_listbox.curselection()
if sel:
self.queue_listbox.delete(sel[0])
def _queue_clear(self):
self.queue_listbox.delete(0, tk.END)
def _queue_move_up(self):
sel = self.queue_listbox.curselection()
if sel and sel[0] > 0:
idx = sel[0]
text = self.queue_listbox.get(idx)
self.queue_listbox.delete(idx)
self.queue_listbox.insert(idx - 1, text)
self.queue_listbox.selection_set(idx - 1)
def _queue_move_down(self):
sel = self.queue_listbox.curselection()
if sel and sel[0] < self.queue_listbox.size() - 1:
idx = sel[0]
text = self.queue_listbox.get(idx)
self.queue_listbox.delete(idx)
self.queue_listbox.insert(idx + 1, text)
self.queue_listbox.selection_set(idx + 1)
def on_run_queue(self):
if self.generating or not self.models:
return
prompts = list(self.queue_listbox.get(0, tk.END))
if not prompts:
self.status_var.set("Queue is empty")
return
self.generating = True
self.gen.cancelled = False
self.gen_btn.configure(state=tk.DISABLED)
self.queue_run_btn.configure(state=tk.DISABLED)
self.stop_btn.configure(state=tk.NORMAL)
Thread(target=self._queue_thread, args=(prompts,), daemon=True).start()
def _queue_thread(self, prompts):
try:
idx = self.model_combo.current()
mdl = self.models[idx]
self.status_var.set(f"Loading {mdl[1]}...")
self.root.update()
self.gen.load_model(mdl[2], mdl[3])
neg = self.neg_entry.get().strip()
steps = int(self.steps_var.get())
cfg = float(self.cfg_var.get())
num_images = max(1, min(12, int(self.count_var.get())))
live_preview = self.live_preview_var.get()
self.generated_images.clear()
self.generated_seeds.clear()
self.selected_index = None
if self.placeholder:
self.placeholder.destroy()
self.placeholder = None
for p_idx, prompt in enumerate(prompts):
if self.gen.cancelled:
break
self.queue_listbox.selection_clear(0, tk.END)
self.queue_listbox.selection_set(p_idx)
self.queue_listbox.see(p_idx)
for img_i in range(num_images):
if self.gen.cancelled:
break
self.status_var.set(
f"[{p_idx + 1}/{len(prompts)}] image {img_i + 1}/{num_images}")
self.root.update()
callback = None
if live_preview:
self._setup_preview_card()
callback = self._show_preview
image, used_seed = self.gen.generate(
prompt=prompt, negative_prompt=neg,
steps=steps, guidance_scale=cfg,
preview_callback=callback, preview_every=5)
if image is None:
break
self.generated_images.append(image)
self.generated_seeds.append(used_seed)
save_path = self._next_save_path(prompt)
image.save(save_path)
self._layout_grid()
self.root.update()
if self.gen.cancelled:
break
done = len(self.generated_images)
self.status_var.set(
f"Queue {'stopped' if self.gen.cancelled else 'done'}! {done} images saved.")
if done > 0:
self.save_all_btn.configure(state=tk.NORMAL)
except Exception as e:
self.status_var.set(f"Queue error: {e}")
import traceback; traceback.print_exc()
finally:
self.generating = False
self.gen.cancelled = False
self.gen_btn.configure(state=tk.NORMAL)
self.queue_run_btn.configure(state=tk.NORMAL)
self.stop_btn.configure(state=tk.DISABLED)
# ── Generation ────────────────────────────────────────────────────────
def on_stop(self):
if self.generating:
self.gen.cancelled = True
self.status_var.set("Stopping...")
self.root.update()
def on_generate(self):
if self.generating or not self.models:
return
self.generating = True
self.gen.cancelled = False
self.gen_btn.configure(state=tk.DISABLED)
self.stop_btn.configure(state=tk.NORMAL)
self.status_var.set("Loading model...")
self.root.update()
Thread(target=self._generate_thread, daemon=True).start()
def _setup_preview_card(self):
tile_size = self._get_tile_size()
cols = self._get_grid_cols()
row, col = divmod(len(self.generated_images), cols)
card = tk.Frame(self.grid_frame, bg=C["card"], padx=3, pady=3)
card.grid(row=row, column=col, padx=5, pady=5, sticky="nsew")
self._preview_label = tk.Label(card, bg=C["card"],
width=tile_size, height=tile_size)
self._preview_label.pack()
self.root.update()
def _show_preview(self, preview_img, step, total):
tile_size = self._get_tile_size()
display = preview_img.resize((tile_size, tile_size), Image.LANCZOS)
photo = ImageTk.PhotoImage(display)
self._preview_photo = photo
if hasattr(self, '_preview_label') and self._preview_label.winfo_exists():
self._preview_label.configure(image=photo)
self.status_var.set(f"Step {step}/{total}")
self.root.update()
def _generate_thread(self):
try:
idx = self.model_combo.current()
mdl = self.models[idx]
self.status_var.set(f"Loading {mdl[1]}...")
self.root.update()
self.gen.load_model(mdl[2], mdl[3])
prompt = self.prompt_entry.get().strip()
neg = self.neg_entry.get().strip()
steps = int(self.steps_var.get())
cfg = float(self.cfg_var.get())
num_images = max(1, min(12, int(self.count_var.get())))
live_preview = self.live_preview_var.get()
self.generated_images.clear()
self.generated_seeds.clear()
self.selected_index = None
if self.placeholder:
self.placeholder.destroy()
self.placeholder = None
for i in range(num_images):
if self.gen.cancelled:
break
self.status_var.set(f"Generating {i + 1}/{num_images}...")
self.root.update()
callback = None
if live_preview:
self._setup_preview_card()
callback = self._show_preview
image, used_seed = self.gen.generate(
prompt=prompt, negative_prompt=neg,
steps=steps, guidance_scale=cfg,
preview_callback=callback, preview_every=5)
if image is None:
break
self.generated_images.append(image)
self.generated_seeds.append(used_seed)
self._layout_grid()
self.root.update()
done = len(self.generated_images)
if self.gen.cancelled:
self.status_var.set(f"Stopped. {done} image(s) kept.")
else:
self.status_var.set(f"Done! {done} images. Click to select.")
if done > 0:
self.save_all_btn.configure(state=tk.NORMAL)
self.save_btn.configure(state=tk.DISABLED)
except Exception as e:
self.status_var.set(f"Error: {e}")
import traceback; traceback.print_exc()
finally:
self.generating = False
self.gen.cancelled = False
self.gen_btn.configure(state=tk.NORMAL)
self.stop_btn.configure(state=tk.DISABLED)
# ── Save ──────────────────────────────────────────────────────────────
def _next_save_path(self, prompt_text):
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
slug = prompt_text.strip()[:50] if prompt_text.strip() else "untitled"
base = OUTPUT_DIR / f"{slug}.png"
if not base.exists():
return base
n = 1
while True:
path = OUTPUT_DIR / f"{slug} {n}.png"
if not path.exists():
return path
n += 1
def on_save(self):
if self.selected_index is None or not self.generated_images:
return
img = self.generated_images[self.selected_index]
path = self._next_save_path(self.prompt_entry.get().strip())
img.save(path)
self.status_var.set(f"Saved: {path.name}")
def on_save_all(self):
if not self.generated_images:
return
prompt_text = self.prompt_entry.get().strip()
for img in self.generated_images:
path = self._next_save_path(prompt_text)
img.save(path)
self.status_var.set(f"Saved {len(self.generated_images)} images")
def run(self):
self.root.mainloop()
# ── Entry point ───────────────────────────────────────────────────────────────
if __name__ == "__main__":
models = find_models()
if not models:
print("No models found locally. Downloading from HuggingFace...")
result = download_from_hf()
if result:
models = find_models()
if not models:
print("No models found!")
print(f"Place model weights in: {MODEL_DIR}/YourModelName/")
print("Expected files: diffusion_pytorch_model.safetensors or ema_unet.pt")
sys.exit(1)
print(f"Found {len(models)} model(s): {', '.join(m[1] for m in models)}")
print(f"Device: {'CUDA (GPU)' if torch.cuda.is_available() else 'CPU'}")
print("Starting Aniimage...")
app = App()
app.run()