""" 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("<>", 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("", 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("<>", 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("", 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("", 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("", self._on_canvas_resize) self.canvas.bind_all("", 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("", lambda e, idx=i: self._select_image(idx)) img_label.bind("", 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()