Instructions to use 8BitStudio/Aniimage-2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use 8BitStudio/Aniimage-2 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("8BitStudio/Aniimage-2", torch_dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
- DiffusionBee
| """ | |
| 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] | |
| 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) | |
| 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 | |
| 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() |