8BitStudio commited on
Commit
67b7bd0
Β·
verified Β·
1 Parent(s): d94e421

Upload generate_images.py

Browse files
Files changed (1) hide show
  1. generate_images.py +1459 -0
generate_images.py ADDED
@@ -0,0 +1,1459 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Aniimage Generator β€” Generate anime images from text prompts.
3
+ https://huggingface.co/8BitStudio/Aniimage-2
4
+
5
+ Usage:
6
+ pip install -U torch torchvision "diffusers>=0.37.1" "transformers>=4.46,<5" accelerate safetensors pillow huggingface_hub
7
+ python generate_hf_aniimage2_corrected.py
8
+ """
9
+
10
+ import os
11
+ import sys
12
+ import gc
13
+ import json
14
+ import torch
15
+ import numpy as np
16
+ import tkinter as tk
17
+ from tkinter import ttk, simpledialog
18
+ from pathlib import Path
19
+ from PIL import Image, ImageTk
20
+ from threading import Thread
21
+
22
+ # ── Paths ─────────────────────────────────────────────────────────────────────
23
+ SCRIPT_DIR = Path(__file__).resolve().parent
24
+ MODEL_DIR = SCRIPT_DIR / "models"
25
+ OUTPUT_DIR = SCRIPT_DIR / "generated"
26
+
27
+ # ── HuggingFace repo ─────────────────────────────────────────────────────────
28
+ HF_REPO_ID = "8BitStudio/Aniimage-2"
29
+
30
+ # ── Aniimage-2 training configuration fallback ────────────────────────────────
31
+ # The downloaded model_config.json is preferred. These values mirror it so the
32
+ # launcher still behaves correctly if only the UNet files were copied locally.
33
+ UNET_CONFIG = dict(
34
+ sample_size=64,
35
+ in_channels=4,
36
+ out_channels=4,
37
+ block_out_channels=(256, 512, 768, 1024),
38
+ layers_per_block=2,
39
+ cross_attention_dim=768,
40
+ attention_head_dim=8,
41
+ down_block_types=("CrossAttnDownBlock2D", "CrossAttnDownBlock2D",
42
+ "CrossAttnDownBlock2D", "DownBlock2D"),
43
+ up_block_types=("UpBlock2D", "CrossAttnUpBlock2D",
44
+ "CrossAttnUpBlock2D", "CrossAttnUpBlock2D"),
45
+ )
46
+
47
+ # Aniimage-2 was trained with this VAE, not the SD 1.x MSE VAE.
48
+ VAE_ID = "madebyollin/sdxl-vae-fp16-fix"
49
+ CLIP_ID = "openai/clip-vit-large-patch14"
50
+
51
+ SCHEDULER_LIST = [
52
+ "DPM++ 2M Karras",
53
+ "DPM++ SDE Karras",
54
+ "Euler a",
55
+ "Euler",
56
+ "DDIM",
57
+ ]
58
+
59
+ DEFAULT_NEGATIVE = (
60
+ "low quality, ugly, blurry, distorted, deformed, bad anatomy, "
61
+ "bad proportions, extra limbs, missing limbs, watermark, text, "
62
+ "signature, washed out, flat colors, manga panel, disfigured, "
63
+ "poorly drawn, jpeg artifacts, cropped, out of frame"
64
+ )
65
+
66
+
67
+ # ── Model discovery ───────────────────────────────────────────────────────────
68
+
69
+ def _read_json(path: Path):
70
+ """Read a JSON file, returning an empty dict when it is unusable."""
71
+ try:
72
+ data = json.loads(path.read_text(encoding="utf-8"))
73
+ return data if isinstance(data, dict) else {}
74
+ except (OSError, ValueError, TypeError):
75
+ return {}
76
+
77
+
78
+ def _looks_like_unet_config(config: dict) -> bool:
79
+ """Return True when a config contains the core Diffusers UNet fields."""
80
+ required = {
81
+ "in_channels", "out_channels", "block_out_channels",
82
+ "down_block_types", "up_block_types",
83
+ }
84
+ return required.issubset(config)
85
+
86
+
87
+ def _find_model_config(model_dir: Path):
88
+ """Find the Aniimage model_config.json that describes training settings."""
89
+ if not model_dir.exists():
90
+ return None
91
+ candidates = [
92
+ p for p in model_dir.rglob("model_config.json")
93
+ if p.is_file() and ".cache" not in p.parts
94
+ ]
95
+ for path in sorted(candidates, key=lambda p: (len(p.relative_to(model_dir).parts), str(p))):
96
+ config = _read_json(path)
97
+ if isinstance(config.get("unet"), dict):
98
+ return path
99
+ return None
100
+
101
+
102
+ def _find_unet_assets(model_dir: Path):
103
+ """Find Aniimage UNet weights/config, including nested repo folders.
104
+
105
+ Aniimage-2 has been published with an extra ``Aniimage-2/unet`` directory
106
+ inside the repository snapshot. Searching recursively keeps the launcher
107
+ compatible with that layout as well as normal Diffusers layouts.
108
+ """
109
+ if not model_dir.exists():
110
+ return None
111
+
112
+ # Prefer the canonical single-file names, but accept fp16/variant names.
113
+ weight_candidates = []
114
+ for pattern in (
115
+ "diffusion_pytorch_model.safetensors",
116
+ "diffusion_pytorch_model.bin",
117
+ "diffusion_pytorch_model*.safetensors",
118
+ "diffusion_pytorch_model*.bin",
119
+ ):
120
+ weight_candidates.extend(model_dir.rglob(pattern))
121
+
122
+ # Remove duplicates, metadata files, and anything under the HF cache.
123
+ unique_weights = []
124
+ seen = set()
125
+ for path in weight_candidates:
126
+ if not path.is_file() or path.name.endswith(".index.json"):
127
+ continue
128
+ if ".cache" in path.parts:
129
+ continue
130
+ key = str(path.resolve())
131
+ if key not in seen:
132
+ seen.add(key)
133
+ unique_weights.append(path)
134
+
135
+ if unique_weights:
136
+ # Exact canonical filenames first, then the shallowest path.
137
+ def weight_rank(path: Path):
138
+ exact = path.name in {
139
+ "diffusion_pytorch_model.safetensors",
140
+ "diffusion_pytorch_model.bin",
141
+ }
142
+ safe = path.suffix.lower() == ".safetensors"
143
+ return (not exact, not safe, len(path.relative_to(model_dir).parts), str(path))
144
+
145
+ weights_path = sorted(unique_weights, key=weight_rank)[0]
146
+
147
+ # Search from the weights folder upward, then across the model root.
148
+ config_candidates = []
149
+ current = weights_path.parent
150
+ while True:
151
+ config_candidates.extend((current / "config.json", current / "model_config.json"))
152
+ if current == model_dir or model_dir not in current.parents:
153
+ break
154
+ current = current.parent
155
+ config_candidates.extend(model_dir.rglob("config.json"))
156
+ config_candidates.extend(model_dir.rglob("model_config.json"))
157
+
158
+ config_path = None
159
+ seen_configs = set()
160
+ for candidate in config_candidates:
161
+ if not candidate.is_file() or ".cache" in candidate.parts:
162
+ continue
163
+ key = str(candidate.resolve())
164
+ if key in seen_configs:
165
+ continue
166
+ seen_configs.add(key)
167
+ if _looks_like_unet_config(_read_json(candidate)):
168
+ config_path = candidate
169
+ break
170
+
171
+ return {
172
+ "kind": "diffusers_weights",
173
+ "weights": weights_path,
174
+ "config": config_path,
175
+ }
176
+
177
+ # Older local checkpoints are still supported.
178
+ for filename in ("ema_unet.pt", "unet.pt"):
179
+ candidates = [p for p in model_dir.rglob(filename)
180
+ if p.is_file() and ".cache" not in p.parts]
181
+ if candidates:
182
+ checkpoint = sorted(
183
+ candidates,
184
+ key=lambda p: (len(p.relative_to(model_dir).parts), str(p)),
185
+ )[0]
186
+ return {
187
+ "kind": "checkpoint",
188
+ "weights": checkpoint,
189
+ "config": None,
190
+ }
191
+
192
+ return None
193
+
194
+
195
+ def _detect_model_resolution(model_dir: Path) -> str:
196
+ """Infer output resolution from model metadata or the UNet config."""
197
+ model_config_path = _find_model_config(model_dir)
198
+ if model_config_path:
199
+ image_size = _read_json(model_config_path).get("image_size")
200
+ if isinstance(image_size, int) and image_size > 0:
201
+ return str(image_size)
202
+
203
+ assets = _find_unet_assets(model_dir)
204
+ config_path = assets.get("config") if assets else None
205
+ if config_path:
206
+ config = _read_json(config_path)
207
+ sample_size = config.get("sample_size")
208
+ if isinstance(sample_size, (list, tuple)) and sample_size:
209
+ sample_size = sample_size[0]
210
+ if isinstance(sample_size, int) and sample_size > 0:
211
+ return str(sample_size * 8)
212
+
213
+ return "512" if "aniimage-2" in model_dir.name.lower() else "256"
214
+
215
+
216
+ def download_from_hf():
217
+ """Download Aniimage-2 from Hugging Face if it is not already present."""
218
+ try:
219
+ from huggingface_hub import snapshot_download
220
+ except ImportError:
221
+ print("Install huggingface_hub: pip install huggingface_hub")
222
+ return None
223
+
224
+ MODEL_DIR.mkdir(parents=True, exist_ok=True)
225
+ aniimage_dir = MODEL_DIR / "Aniimage-2"
226
+
227
+ existing = _find_unet_assets(aniimage_dir)
228
+ existing_config = _find_model_config(aniimage_dir)
229
+ if existing and existing_config:
230
+ print(f"Aniimage-2 weights already downloaded: {existing['weights']}")
231
+ return aniimage_dir
232
+
233
+ print(f"Downloading Aniimage-2 from {HF_REPO_ID}...")
234
+ aniimage_dir.mkdir(parents=True, exist_ok=True)
235
+
236
+ try:
237
+ snapshot_download(
238
+ repo_id=HF_REPO_ID,
239
+ local_dir=aniimage_dir,
240
+ allow_patterns=[
241
+ "Aniimage-2/model_config.json",
242
+ "Aniimage-2/unet/*",
243
+ ],
244
+ )
245
+ except Exception as exc:
246
+ print(f"Aniimage-2 download failed: {exc}")
247
+ return None
248
+
249
+ assets = _find_unet_assets(aniimage_dir)
250
+ if not assets:
251
+ print(
252
+ "Aniimage-2 repository downloaded, but no supported UNet weights "
253
+ "were found anywhere below:\n"
254
+ f" {aniimage_dir}\n"
255
+ "Expected diffusion_pytorch_model.safetensors or "
256
+ "diffusion_pytorch_model.bin."
257
+ )
258
+ return None
259
+
260
+ print(f"Download complete! Found weights at: {assets['weights']}")
261
+ return aniimage_dir
262
+
263
+
264
+ def find_models():
265
+ """Find models, including checkpoints nested inside repository folders."""
266
+ options = []
267
+ if MODEL_DIR.exists():
268
+ for d in sorted(MODEL_DIR.iterdir()):
269
+ if not d.is_dir():
270
+ continue
271
+ assets = _find_unet_assets(d)
272
+ if not assets:
273
+ continue
274
+ resolution = _detect_model_resolution(d)
275
+ model_kind = (
276
+ "safetensors"
277
+ if assets["weights"].suffix.lower() == ".safetensors"
278
+ else assets["kind"]
279
+ )
280
+ options.append((model_kind, d.name, d, resolution))
281
+ return options
282
+
283
+
284
+ # ── Theme ─────────────────────────────────────────────────────────────────────
285
+
286
+ C = {
287
+ "bg": "#111119",
288
+ "panel": "#1b1b2f",
289
+ "card": "#24243e",
290
+ "card_sel": "#3a3a6e",
291
+ "border": "#2e2e52",
292
+ "accent": "#6c5ce7",
293
+ "accent_h": "#8577ed",
294
+ "red": "#e74c3c",
295
+ "green": "#2ecc71",
296
+ "text": "#eaeaea",
297
+ "text2": "#a0a0b8",
298
+ "text3": "#60607a",
299
+ "input": "#16162a",
300
+ "input_fg": "#dcdcf0",
301
+ }
302
+
303
+
304
+ class Generator:
305
+ def __init__(self, device="cuda"):
306
+ self.device = device if device == "cuda" and torch.cuda.is_available() else "cpu"
307
+ self.dtype = self._select_dtype()
308
+ self.vae = None
309
+ self.text_encoder = None
310
+ self.tokenizer = None
311
+ self.unet = None
312
+ self.scheduler = None
313
+ self.loaded_checkpoint = None
314
+ self.loaded_vae_id = None
315
+ self.model_config = {}
316
+ self._clip_inner = None
317
+ self._clip_full_layers = None
318
+ self.latent_size = 64
319
+ self.output_size = 512
320
+ self.prediction_type = "v_prediction"
321
+ self.zero_terminal_snr = True
322
+ self.timestep_spacing = "trailing"
323
+ self.guidance_rescale = 0.7
324
+ self.num_train_timesteps = 1000
325
+ self.beta_schedule = "scaled_linear"
326
+ self.clip_penultimate = True
327
+ self.vae_id = VAE_ID
328
+ self.scheduler_name = "DPM++ SDE Karras"
329
+ self.cancelled = False
330
+ self._configure_backends()
331
+
332
+ def _select_dtype(self):
333
+ if self.device != "cuda":
334
+ return torch.float32
335
+ bf16_supported = getattr(torch.cuda, "is_bf16_supported", lambda: False)()
336
+ return torch.bfloat16 if bf16_supported else torch.float16
337
+
338
+ def _configure_backends(self):
339
+ if self.device == "cuda":
340
+ torch.backends.cudnn.benchmark = True
341
+ torch.backends.cuda.matmul.allow_tf32 = True
342
+ torch.backends.cudnn.allow_tf32 = True
343
+ if hasattr(torch, "set_float32_matmul_precision"):
344
+ torch.set_float32_matmul_precision("high")
345
+
346
+ def _autocast(self):
347
+ return torch.autocast(
348
+ device_type="cuda",
349
+ dtype=self.dtype,
350
+ enabled=(self.device == "cuda"),
351
+ )
352
+
353
+ def switch_device(self, new_device):
354
+ """Switch device and rebuild the models in the correct precision."""
355
+ new_device = new_device if new_device == "cuda" and torch.cuda.is_available() else "cpu"
356
+ if new_device == self.device:
357
+ return
358
+ self.device = new_device
359
+ self.dtype = self._select_dtype()
360
+ self._configure_backends()
361
+ self.vae = None
362
+ self.text_encoder = None
363
+ self.tokenizer = None
364
+ self.unet = None
365
+ self.scheduler = None
366
+ self.loaded_checkpoint = None
367
+ self.loaded_vae_id = None
368
+ self._clip_inner = None
369
+ self._clip_full_layers = None
370
+ gc.collect()
371
+ if torch.cuda.is_available():
372
+ torch.cuda.empty_cache()
373
+ print(f"Switched to {self.device.upper()} ({self.dtype})")
374
+
375
+ def _load_model_metadata(self, model_path: Path, res_label: str):
376
+ """Load the exact training objective and component IDs for Aniimage-2."""
377
+ config_path = _find_model_config(model_path)
378
+ config = _read_json(config_path) if config_path else {}
379
+ self.model_config = config
380
+
381
+ self.prediction_type = config.get("prediction_type", "v_prediction")
382
+ self.zero_terminal_snr = bool(config.get("zero_terminal_snr", True))
383
+ self.timestep_spacing = config.get(
384
+ "timestep_spacing",
385
+ "trailing" if self.zero_terminal_snr else "leading",
386
+ )
387
+ self.guidance_rescale = float(config.get("guidance_rescale", 0.7))
388
+ self.num_train_timesteps = int(config.get("num_train_timesteps", 1000))
389
+ self.beta_schedule = config.get("beta_schedule", "scaled_linear")
390
+ self.clip_penultimate = bool(config.get("clip_penultimate", True))
391
+ self.vae_id = config.get("vae", VAE_ID)
392
+
393
+ try:
394
+ fallback_size = int(res_label)
395
+ except (TypeError, ValueError):
396
+ fallback_size = 512
397
+ self.output_size = int(config.get("image_size", fallback_size))
398
+ self.latent_size = self.output_size // 8
399
+
400
+ if config_path:
401
+ print(f"Using model metadata: {config_path}")
402
+ else:
403
+ print("model_config.json was not found; using Aniimage-2 defaults.")
404
+
405
+ def _apply_clip_layer_mode(self):
406
+ if self.text_encoder is None:
407
+ return
408
+ self._clip_inner = getattr(self.text_encoder, "text_model", self.text_encoder)
409
+ if self._clip_full_layers is None:
410
+ self._clip_full_layers = self._clip_inner.encoder.layers
411
+ if self.clip_penultimate:
412
+ self._clip_inner.encoder.layers = self._clip_full_layers[:-1]
413
+ print("Text encoder: CLIP penultimate layer (matches training).")
414
+ else:
415
+ self._clip_inner.encoder.layers = self._clip_full_layers
416
+
417
+ def load_shared(self):
418
+ from diffusers import AutoencoderKL
419
+ from transformers import (CLIPConfig, CLIPTextConfig,
420
+ CLIPTextModel, CLIPTokenizer)
421
+
422
+ load_kwargs = {"low_cpu_mem_usage": True}
423
+ if self.device == "cuda":
424
+ load_kwargs["torch_dtype"] = self.dtype
425
+
426
+ if self.vae is None or self.loaded_vae_id != self.vae_id:
427
+ print(f"Loading VAE: {self.vae_id}...")
428
+ self.vae = AutoencoderKL.from_pretrained(
429
+ self.vae_id,
430
+ **load_kwargs,
431
+ ).to(self.device).eval()
432
+ self.vae.requires_grad_(False)
433
+ self.vae.enable_slicing()
434
+ if self.device == "cuda":
435
+ self.vae.to(memory_format=torch.channels_last)
436
+ self.loaded_vae_id = self.vae_id
437
+
438
+ if self.text_encoder is None:
439
+ print(f"Loading CLIP text encoder: {CLIP_ID}...")
440
+ self.tokenizer = CLIPTokenizer.from_pretrained(CLIP_ID)
441
+
442
+ # Explicitly pass the nested text config. This avoids the
443
+ # CLIPConfig.hidden_size crash seen with some Transformers builds.
444
+ clip_config = CLIPConfig.from_pretrained(CLIP_ID)
445
+ text_config = getattr(clip_config, "text_config", None)
446
+ if isinstance(text_config, dict):
447
+ text_config = CLIPTextConfig.from_dict(text_config)
448
+ if not isinstance(text_config, CLIPTextConfig):
449
+ text_config = CLIPTextConfig.from_pretrained(CLIP_ID)
450
+
451
+ self.text_encoder = CLIPTextModel.from_pretrained(
452
+ CLIP_ID,
453
+ config=text_config,
454
+ **load_kwargs,
455
+ ).to(self.device).eval()
456
+ self.text_encoder.requires_grad_(False)
457
+ self._clip_full_layers = None
458
+
459
+ self._apply_clip_layer_mode()
460
+ self.scheduler = self._make_scheduler(self.scheduler_name)
461
+ print("Shared models loaded.")
462
+
463
+ def _make_scheduler(self, name="DPM++ SDE Karras"):
464
+ from diffusers import (DDIMScheduler, DPMSolverMultistepScheduler,
465
+ EulerAncestralDiscreteScheduler,
466
+ EulerDiscreteScheduler)
467
+ base = dict(
468
+ num_train_timesteps=self.num_train_timesteps,
469
+ beta_schedule=self.beta_schedule,
470
+ prediction_type=self.prediction_type,
471
+ rescale_betas_zero_snr=self.zero_terminal_snr,
472
+ timestep_spacing=self.timestep_spacing,
473
+ )
474
+ if name == "DPM++ 2M Karras":
475
+ return DPMSolverMultistepScheduler(
476
+ **base, algorithm_type="dpmsolver++",
477
+ solver_order=2, use_karras_sigmas=True)
478
+ if name == "DPM++ SDE Karras":
479
+ return DPMSolverMultistepScheduler(
480
+ **base, algorithm_type="sde-dpmsolver++",
481
+ solver_order=2, use_karras_sigmas=True)
482
+ if name == "Euler a":
483
+ return EulerAncestralDiscreteScheduler(**base)
484
+ if name == "Euler":
485
+ return EulerDiscreteScheduler(**base)
486
+ return DDIMScheduler(
487
+ **base, clip_sample=False, set_alpha_to_one=False)
488
+
489
+ def set_scheduler(self, name):
490
+ self.scheduler_name = name
491
+ self.scheduler = self._make_scheduler(name)
492
+
493
+ def load_model(self, model_path: Path, res_label: str = "512"):
494
+ if str(model_path) == self.loaded_checkpoint:
495
+ return
496
+ from diffusers import UNet2DConditionModel
497
+
498
+ assets = _find_unet_assets(model_path)
499
+ if not assets:
500
+ raise FileNotFoundError(
501
+ f"No supported UNet weights found anywhere inside {model_path}"
502
+ )
503
+
504
+ self._load_model_metadata(model_path, res_label)
505
+ self.load_shared()
506
+
507
+ weights_path = assets["weights"]
508
+ config_path = assets.get("config")
509
+ suffix = weights_path.suffix.lower()
510
+ same_dir_config = weights_path.parent / "config.json"
511
+
512
+ print(
513
+ f"Loading UNet from {weights_path} "
514
+ f"({self.output_size}px, {self.prediction_type}, {self.dtype})..."
515
+ )
516
+
517
+ self.unet = None
518
+ if torch.cuda.is_available():
519
+ torch.cuda.empty_cache()
520
+
521
+ loaded_directly = False
522
+ if same_dir_config.exists() and _looks_like_unet_config(_read_json(same_dir_config)):
523
+ try:
524
+ kwargs = {"low_cpu_mem_usage": True}
525
+ if self.device == "cuda":
526
+ kwargs["torch_dtype"] = self.dtype
527
+ if suffix == ".safetensors":
528
+ kwargs["use_safetensors"] = True
529
+ elif suffix == ".bin":
530
+ kwargs["use_safetensors"] = False
531
+ self.unet = UNet2DConditionModel.from_pretrained(
532
+ weights_path.parent,
533
+ **kwargs,
534
+ ).to(self.device)
535
+ loaded_directly = True
536
+ print("Loaded the repository UNet config and weights directly.")
537
+ except Exception as exc:
538
+ print(f"Direct Diffusers load failed ({exc}); loading manually.")
539
+
540
+ if not loaded_directly:
541
+ if config_path:
542
+ unet_config = _read_json(config_path)
543
+ print(f"Using UNet config: {config_path}")
544
+ elif isinstance(self.model_config.get("unet"), dict):
545
+ unet_config = dict(self.model_config["unet"])
546
+ print("Using UNet config from model_config.json.")
547
+ else:
548
+ unet_config = dict(UNET_CONFIG)
549
+ print("Using built-in Aniimage-2 UNet config.")
550
+ unet_config["sample_size"] = self.latent_size
551
+ self.unet = UNet2DConditionModel.from_config(unet_config)
552
+
553
+ if suffix == ".safetensors":
554
+ from safetensors.torch import load_file
555
+ state = load_file(str(weights_path), device="cpu")
556
+ else:
557
+ try:
558
+ state = torch.load(weights_path, map_location="cpu", weights_only=True)
559
+ except TypeError:
560
+ state = torch.load(weights_path, map_location="cpu")
561
+
562
+ if weights_path.name == "ema_unet.pt" and isinstance(state, dict) and "shadow_params" in state:
563
+ params = dict(self.unet.named_parameters())
564
+ keys = list(params.keys())
565
+ if len(state["shadow_params"]) != len(keys):
566
+ raise RuntimeError("EMA parameter count does not match the UNet.")
567
+ for key, shadow_param in zip(keys, state["shadow_params"]):
568
+ params[key].data.copy_(shadow_param)
569
+ else:
570
+ if isinstance(state, dict) and "state_dict" in state:
571
+ state = state["state_dict"]
572
+ if isinstance(state, dict) and state and all(
573
+ isinstance(key, str) and key.startswith("module.") for key in state
574
+ ):
575
+ state = {key[7:]: value for key, value in state.items()}
576
+ self.unet.load_state_dict(state, strict=True)
577
+
578
+ if self.device == "cuda":
579
+ self.unet = self.unet.to(device=self.device, dtype=self.dtype)
580
+ else:
581
+ self.unet = self.unet.to(self.device)
582
+
583
+ sample_size = self.unet.config.sample_size
584
+ if isinstance(sample_size, (list, tuple)) and sample_size:
585
+ sample_size = sample_size[0]
586
+ if isinstance(sample_size, int) and sample_size > 0:
587
+ self.latent_size = sample_size
588
+ vae_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
589
+ self.output_size = sample_size * vae_factor
590
+
591
+ self.unet.eval().requires_grad_(False)
592
+ if self.device == "cuda":
593
+ self.unet.to(memory_format=torch.channels_last)
594
+ self.scheduler = self._make_scheduler(self.scheduler_name)
595
+ self.loaded_checkpoint = str(model_path)
596
+ print(
597
+ f"Ready at {self.output_size}x{self.output_size}; "
598
+ f"zero-SNR={self.zero_terminal_snr}, spacing={self.timestep_spacing}, "
599
+ f"CFG rescale={self.guidance_rescale}."
600
+ )
601
+
602
+ def _encode_prompts(self, prompt: str, negative_prompt: str):
603
+ tokens = self.tokenizer(
604
+ [negative_prompt or "", prompt],
605
+ padding="max_length",
606
+ max_length=self.tokenizer.model_max_length,
607
+ truncation=True,
608
+ return_tensors="pt",
609
+ )
610
+ with self._autocast():
611
+ return self.text_encoder(tokens.input_ids.to(self.device))[0]
612
+
613
+ @staticmethod
614
+ def _cfg_rescale(noise_cfg, noise_text, amount):
615
+ if amount <= 0:
616
+ return noise_cfg
617
+ dims = tuple(range(1, noise_cfg.ndim))
618
+ std_text = noise_text.std(dim=dims, keepdim=True)
619
+ std_cfg = noise_cfg.std(dim=dims, keepdim=True).clamp_min(1e-6)
620
+ noise_rescaled = noise_cfg * (std_text / std_cfg)
621
+ return amount * noise_rescaled + (1.0 - amount) * noise_cfg
622
+
623
+ def _decode_latents(self, latents, post_process=False):
624
+ del post_process # Kept for compatibility with the preview callbacks.
625
+ decode_dtype = self.dtype if self.device == "cuda" else torch.float32
626
+ scaled = (latents / self.vae.config.scaling_factor).to(dtype=decode_dtype)
627
+ with self._autocast():
628
+ image = self.vae.decode(scaled).sample
629
+ image = (image.float() / 2 + 0.5).clamp(0, 1)
630
+ image = image[0].cpu().permute(1, 2, 0).numpy()
631
+ image = (image * 255).round().astype("uint8")
632
+ return Image.fromarray(image)
633
+
634
+ @torch.inference_mode()
635
+ def generate(self, prompt: str, negative_prompt: str = "",
636
+ steps: int = 50, guidance_scale: float = 7.5,
637
+ seed: int = -1, preview_callback=None,
638
+ preview_every: int = 5) -> tuple:
639
+
640
+ if seed < 0:
641
+ seed = torch.randint(0, 2**32, (1,)).item()
642
+ generator = torch.Generator(device=self.device).manual_seed(seed)
643
+ embeddings = self._encode_prompts(prompt, negative_prompt)
644
+
645
+ scheduler = self._make_scheduler(self.scheduler_name)
646
+ scheduler.set_timesteps(int(steps), device=self.device)
647
+
648
+ in_channels = int(self.unet.config.in_channels)
649
+ latents = torch.randn(
650
+ (1, in_channels, self.latent_size, self.latent_size),
651
+ generator=generator,
652
+ device=self.device,
653
+ dtype=torch.float32,
654
+ ) * scheduler.init_noise_sigma
655
+
656
+ total_steps = len(scheduler.timesteps)
657
+ preview_interval = max(1, int(preview_every))
658
+
659
+ for step_i, timestep in enumerate(scheduler.timesteps):
660
+ if self.cancelled:
661
+ return None, seed
662
+
663
+ latent_input = torch.cat([latents, latents], dim=0)
664
+ latent_input = scheduler.scale_model_input(latent_input, timestep)
665
+
666
+ with self._autocast():
667
+ prediction = self.unet(
668
+ latent_input,
669
+ timestep,
670
+ encoder_hidden_states=embeddings,
671
+ ).sample
672
+
673
+ pred_negative, pred_text = prediction.chunk(2)
674
+ prediction = pred_negative + float(guidance_scale) * (pred_text - pred_negative)
675
+ prediction = self._cfg_rescale(
676
+ prediction, pred_text, self.guidance_rescale)
677
+ latents = scheduler.step(prediction, timestep, latents).prev_sample
678
+
679
+ if (preview_callback
680
+ and (step_i + 1) % preview_interval == 0
681
+ and step_i < total_steps - 1):
682
+ preview_callback(
683
+ self._decode_latents(latents),
684
+ step_i + 1,
685
+ total_steps,
686
+ )
687
+
688
+ return self._decode_latents(latents), seed
689
+
690
+ @torch.inference_mode()
691
+ def refine(self, source_image: Image.Image, prompt: str,
692
+ negative_prompt: str = "", extra_steps: int = 20,
693
+ strength: float = 0.35, guidance_scale: float = 7.5,
694
+ preview_callback=None, preview_every: int = 5) -> Image.Image:
695
+
696
+ img = source_image.convert("RGB").resize(
697
+ (self.output_size, self.output_size), Image.LANCZOS)
698
+ img_tensor = torch.from_numpy(np.array(img)).float().div(127.5).sub(1.0)
699
+ img_tensor = img_tensor.permute(2, 0, 1).unsqueeze(0).to(self.device)
700
+ img_tensor = img_tensor.to(
701
+ dtype=self.dtype if self.device == "cuda" else torch.float32)
702
+
703
+ with self._autocast():
704
+ latents = self.vae.encode(img_tensor).latent_dist.sample()
705
+ latents = (latents * self.vae.config.scaling_factor).float()
706
+ embeddings = self._encode_prompts(prompt, negative_prompt)
707
+
708
+ scheduler = self._make_scheduler(self.scheduler_name)
709
+ scheduler.set_timesteps(int(extra_steps), device=self.device)
710
+ start_step = max(0, int(len(scheduler.timesteps) * (1.0 - float(strength))))
711
+ timesteps = scheduler.timesteps[start_step:]
712
+ if len(timesteps) == 0:
713
+ return source_image.copy()
714
+
715
+ noise = torch.randn_like(latents)
716
+ latents = scheduler.add_noise(latents, noise, timesteps[:1])
717
+ total_steps = len(timesteps)
718
+ preview_interval = max(1, int(preview_every))
719
+
720
+ for step_i, timestep in enumerate(timesteps):
721
+ if self.cancelled:
722
+ return None
723
+
724
+ latent_input = torch.cat([latents, latents], dim=0)
725
+ latent_input = scheduler.scale_model_input(latent_input, timestep)
726
+ with self._autocast():
727
+ prediction = self.unet(
728
+ latent_input,
729
+ timestep,
730
+ encoder_hidden_states=embeddings,
731
+ ).sample
732
+
733
+ pred_negative, pred_text = prediction.chunk(2)
734
+ prediction = pred_negative + float(guidance_scale) * (pred_text - pred_negative)
735
+ prediction = self._cfg_rescale(
736
+ prediction, pred_text, self.guidance_rescale)
737
+ latents = scheduler.step(prediction, timestep, latents).prev_sample
738
+
739
+ if (preview_callback
740
+ and (step_i + 1) % preview_interval == 0
741
+ and step_i < total_steps - 1):
742
+ preview_callback(
743
+ self._decode_latents(latents),
744
+ step_i + 1,
745
+ total_steps,
746
+ )
747
+
748
+ return self._decode_latents(latents)
749
+
750
+
751
+ # ── GUI ───────────────────────────────────────────────────────────────────────
752
+
753
+ class App:
754
+ def __init__(self):
755
+ self.gen = Generator()
756
+ self.models = find_models()
757
+ self.generated_images = []
758
+ self.generated_seeds = []
759
+ self.photo_refs = []
760
+ self.generating = False
761
+ self.selected_index = None
762
+
763
+ self.root = tk.Tk()
764
+ self.root.title("Aniimage")
765
+ self.root.configure(bg=C["bg"])
766
+ self.root.resizable(True, True)
767
+ self.root.geometry("900x780")
768
+ self.root.minsize(640, 500)
769
+
770
+ self._setup_styles()
771
+ self._build_ui()
772
+
773
+ def _setup_styles(self):
774
+ s = ttk.Style()
775
+ s.theme_use("clam")
776
+
777
+ # Base
778
+ s.configure(".", background=C["bg"], foreground=C["text"], font=("Segoe UI", 10))
779
+ s.configure("TFrame", background=C["bg"])
780
+ s.configure("TLabel", background=C["bg"], foreground=C["text"])
781
+ s.configure("TCheckbutton", background=C["bg"], foreground=C["text"])
782
+
783
+ # Combobox β€” readable text
784
+ s.configure("TCombobox", fieldbackground=C["input"], foreground=C["input_fg"],
785
+ selectbackground=C["accent"], selectforeground="#ffffff",
786
+ arrowcolor=C["text2"], padding=4)
787
+ s.map("TCombobox",
788
+ fieldbackground=[("readonly", C["input"])],
789
+ foreground=[("readonly", C["input_fg"])],
790
+ selectbackground=[("readonly", C["accent"])],
791
+ selectforeground=[("readonly", "#ffffff")])
792
+ # Combobox dropdown list colors
793
+ self.root.option_add("*TCombobox*Listbox.background", C["input"])
794
+ self.root.option_add("*TCombobox*Listbox.foreground", C["input_fg"])
795
+ self.root.option_add("*TCombobox*Listbox.selectBackground", C["accent"])
796
+ self.root.option_add("*TCombobox*Listbox.selectForeground", "#ffffff")
797
+ self.root.option_add("*TCombobox*Listbox.font", ("Segoe UI", 10))
798
+
799
+ # Spinbox
800
+ s.configure("TSpinbox", fieldbackground=C["input"], foreground=C["input_fg"],
801
+ arrowcolor=C["text2"], padding=3)
802
+
803
+ # Buttons
804
+ s.configure("TButton", font=("Segoe UI", 10), padding=(14, 7),
805
+ background=C["card"], foreground=C["text"])
806
+ s.map("TButton", background=[("active", C["card_sel"]), ("disabled", C["bg"])],
807
+ foreground=[("disabled", C["text3"])])
808
+
809
+ s.configure("Go.TButton", font=("Segoe UI", 11, "bold"), padding=(20, 9),
810
+ background=C["accent"], foreground="#ffffff")
811
+ s.map("Go.TButton", background=[("active", C["accent_h"]),
812
+ ("disabled", C["border"])])
813
+
814
+ s.configure("Stop.TButton", font=("Segoe UI", 10, "bold"), padding=(14, 7),
815
+ background=C["red"], foreground="#ffffff")
816
+ s.map("Stop.TButton", background=[("active", "#c0392b"),
817
+ ("disabled", C["border"])])
818
+
819
+ # Labelframe
820
+ s.configure("TLabelframe", background=C["bg"], foreground=C["text2"])
821
+ s.configure("TLabelframe.Label", background=C["bg"],
822
+ foreground=C["text2"], font=("Segoe UI", 9, "bold"))
823
+
824
+ # Scrollbar
825
+ s.configure("Vertical.TScrollbar", background=C["card"],
826
+ troughcolor=C["bg"], arrowcolor=C["text3"])
827
+
828
+ def _make_entry(self, parent, font_size=11, dim=False):
829
+ """Create a styled tk.Entry with readable text."""
830
+ return tk.Entry(parent, font=("Segoe UI", font_size),
831
+ bg=C["input"], fg=C["input_fg"] if not dim else C["text2"],
832
+ insertbackground=C["input_fg"],
833
+ relief="flat", bd=6,
834
+ selectbackground=C["accent"], selectforeground="#ffffff",
835
+ highlightthickness=1, highlightcolor=C["accent"],
836
+ highlightbackground=C["border"])
837
+
838
+ def _build_ui(self):
839
+ # ── Header ────────────────────────────────────────────────────────
840
+ header = tk.Frame(self.root, bg=C["panel"], padx=20, pady=12)
841
+ header.pack(fill=tk.X)
842
+
843
+ tk.Label(header, text="Aniimage", bg=C["panel"], fg=C["accent"],
844
+ font=("Segoe UI", 20, "bold")).pack(side=tk.LEFT)
845
+ tk.Label(header, text="by 8BitStudio", bg=C["panel"], fg=C["text3"],
846
+ font=("Segoe UI", 10)).pack(side=tk.LEFT, padx=(10, 0), pady=(6, 0))
847
+
848
+ # Device switch β€” right side of header
849
+ device_frame = tk.Frame(header, bg=C["panel"])
850
+ device_frame.pack(side=tk.RIGHT)
851
+
852
+ tk.Label(device_frame, text="Device:", bg=C["panel"], fg=C["text2"],
853
+ font=("Segoe UI", 9)).pack(side=tk.LEFT, padx=(0, 5))
854
+
855
+ self.device_var = tk.StringVar(value="GPU" if self.gen.device == "cuda" else "CPU")
856
+ devices = ["GPU", "CPU"] if torch.cuda.is_available() else ["CPU"]
857
+ device_combo = ttk.Combobox(device_frame, textvariable=self.device_var,
858
+ values=devices, state="readonly", width=5)
859
+ device_combo.pack(side=tk.LEFT)
860
+ device_combo.bind("<<ComboboxSelected>>", self._on_device_change)
861
+
862
+ # ── Main content β€” two-column: controls left, images right ────────
863
+ main = tk.Frame(self.root, bg=C["bg"])
864
+ main.pack(fill=tk.BOTH, expand=True, padx=12, pady=(8, 12))
865
+
866
+ # Left panel (controls)
867
+ left = tk.Frame(main, bg=C["panel"], width=340, padx=16, pady=12)
868
+ left.pack(side=tk.LEFT, fill=tk.Y, padx=(0, 8))
869
+ left.pack_propagate(False)
870
+
871
+ # Right panel (image grid)
872
+ right = tk.Frame(main, bg=C["bg"])
873
+ right.pack(side=tk.LEFT, fill=tk.BOTH, expand=True)
874
+
875
+ self._build_controls(left)
876
+ self._build_grid(right)
877
+
878
+ def _build_controls(self, parent):
879
+ # ── Model ─────────────────────────────────────────────────────────
880
+ tk.Label(parent, text="Model", bg=C["panel"], fg=C["text2"],
881
+ font=("Segoe UI", 9, "bold")).pack(anchor=tk.W)
882
+
883
+ self.model_var = tk.StringVar()
884
+ model_names = [m[1] for m in self.models] or ["No models found"]
885
+ self.model_combo = ttk.Combobox(parent, textvariable=self.model_var,
886
+ values=model_names, state="readonly", width=32)
887
+ self.model_combo.pack(fill=tk.X, pady=(3, 12))
888
+ self.model_combo.current(len(model_names) - 1)
889
+
890
+ # ── Prompt ────────────────────────────────────────────────────────
891
+ tk.Label(parent, text="Prompt", bg=C["panel"], fg=C["text2"],
892
+ font=("Segoe UI", 9, "bold")).pack(anchor=tk.W)
893
+ self.prompt_entry = self._make_entry(parent)
894
+ self.prompt_entry.pack(fill=tk.X, pady=(3, 8))
895
+ self.prompt_entry.insert(0, "a smiling anime girl with long blue hair")
896
+ self.prompt_entry.bind("<Return>", lambda e: self.on_generate())
897
+
898
+ # ── Negative prompt ───────────────────────────────────────────────
899
+ tk.Label(parent, text="Negative prompt", bg=C["panel"], fg=C["text3"],
900
+ font=("Segoe UI", 9)).pack(anchor=tk.W)
901
+ self.neg_entry = self._make_entry(parent, font_size=9, dim=True)
902
+ self.neg_entry.pack(fill=tk.X, pady=(3, 12))
903
+ self.neg_entry.insert(0, DEFAULT_NEGATIVE)
904
+
905
+ # ── Settings grid ─────────────────────────────────────────────────
906
+ grid = tk.Frame(parent, bg=C["panel"])
907
+ grid.pack(fill=tk.X, pady=(0, 8))
908
+
909
+ # Row 1: Scheduler
910
+ tk.Label(grid, text="Scheduler", bg=C["panel"], fg=C["text2"],
911
+ font=("Segoe UI", 9)).grid(row=0, column=0, sticky="w", pady=(0, 6))
912
+ self.scheduler_var = tk.StringVar(value="DPM++ SDE Karras")
913
+ sched_combo = ttk.Combobox(grid, textvariable=self.scheduler_var,
914
+ values=SCHEDULER_LIST, state="readonly", width=18)
915
+ sched_combo.grid(row=0, column=1, columnspan=3, sticky="ew", padx=(8, 0), pady=(0, 6))
916
+ sched_combo.bind("<<ComboboxSelected>>", self._on_scheduler_change)
917
+
918
+ # Row 2: Steps, CFG, Count
919
+ tk.Label(grid, text="Steps", bg=C["panel"], fg=C["text2"],
920
+ font=("Segoe UI", 9)).grid(row=1, column=0, sticky="w", pady=(0, 6))
921
+ self.steps_var = tk.StringVar(value="50")
922
+ tk.Entry(grid, textvariable=self.steps_var, width=5, font=("Segoe UI", 10),
923
+ bg=C["input"], fg=C["input_fg"], insertbackground=C["input_fg"],
924
+ relief="flat", bd=4).grid(row=1, column=1, sticky="w", padx=(8, 12), pady=(0, 6))
925
+
926
+ tk.Label(grid, text="CFG", bg=C["panel"], fg=C["text2"],
927
+ font=("Segoe UI", 9)).grid(row=1, column=2, sticky="w", pady=(0, 6))
928
+ self.cfg_var = tk.StringVar(value="7.5")
929
+ tk.Entry(grid, textvariable=self.cfg_var, width=5, font=("Segoe UI", 10),
930
+ bg=C["input"], fg=C["input_fg"], insertbackground=C["input_fg"],
931
+ relief="flat", bd=4).grid(row=1, column=3, sticky="w", padx=(8, 0), pady=(0, 6))
932
+
933
+ # Row 3: Count, Live preview
934
+ tk.Label(grid, text="Count", bg=C["panel"], fg=C["text2"],
935
+ font=("Segoe UI", 9)).grid(row=2, column=0, sticky="w", pady=(0, 6))
936
+ self.count_var = tk.StringVar(value="4")
937
+ ttk.Spinbox(grid, from_=1, to=12, textvariable=self.count_var, width=4,
938
+ font=("Segoe UI", 10)).grid(row=2, column=1, sticky="w", padx=(8, 12), pady=(0, 6))
939
+
940
+ self.live_preview_var = tk.BooleanVar(value=False)
941
+ ttk.Checkbutton(grid, text="Live preview",
942
+ variable=self.live_preview_var).grid(
943
+ row=2, column=2, columnspan=2, sticky="w", pady=(0, 6))
944
+
945
+ grid.columnconfigure(1, weight=1)
946
+ grid.columnconfigure(3, weight=1)
947
+
948
+ # ── Buttons ───────────────────────────────────────────────────────
949
+ btn_frame = tk.Frame(parent, bg=C["panel"])
950
+ btn_frame.pack(fill=tk.X, pady=(0, 10))
951
+
952
+ self.gen_btn = ttk.Button(btn_frame, text="Generate", command=self.on_generate,
953
+ style="Go.TButton")
954
+ self.gen_btn.pack(fill=tk.X, pady=(0, 5))
955
+
956
+ btn_row = tk.Frame(btn_frame, bg=C["panel"])
957
+ btn_row.pack(fill=tk.X)
958
+
959
+ self.stop_btn = ttk.Button(btn_row, text="Stop", command=self.on_stop,
960
+ state=tk.DISABLED, style="Stop.TButton")
961
+ self.stop_btn.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=(0, 3))
962
+
963
+ self.save_btn = ttk.Button(btn_row, text="Save Selected", command=self.on_save,
964
+ state=tk.DISABLED)
965
+ self.save_btn.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=(3, 3))
966
+
967
+ self.save_all_btn = ttk.Button(btn_row, text="Save All", command=self.on_save_all,
968
+ state=tk.DISABLED)
969
+ self.save_all_btn.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=(3, 0))
970
+
971
+ # ── Prompt queue ─────────────────────────────────────────────────
972
+ sep = tk.Frame(parent, height=1, bg=C["border"])
973
+ sep.pack(fill=tk.X, pady=(8, 10))
974
+
975
+ tk.Label(parent, text="Prompt Queue", bg=C["panel"], fg=C["text2"],
976
+ font=("Segoe UI", 9, "bold")).pack(anchor=tk.W)
977
+
978
+ queue_input = tk.Frame(parent, bg=C["panel"])
979
+ queue_input.pack(fill=tk.X, pady=(4, 0))
980
+
981
+ self.queue_entry = self._make_entry(queue_input, font_size=9)
982
+ self.queue_entry.pack(side=tk.LEFT, fill=tk.X, expand=True, padx=(0, 4))
983
+ self.queue_entry.bind("<Return>", lambda e: self._queue_add())
984
+
985
+ ttk.Button(queue_input, text="Add", width=4,
986
+ command=self._queue_add).pack(side=tk.LEFT)
987
+
988
+ self.queue_listbox = tk.Listbox(
989
+ parent, height=4, bg=C["input"], fg=C["input_fg"],
990
+ selectbackground=C["accent"], selectforeground="#fff",
991
+ font=("Segoe UI", 9), activestyle="none",
992
+ relief="flat", bd=4, highlightthickness=0)
993
+ self.queue_listbox.pack(fill=tk.X, pady=(5, 0))
994
+
995
+ queue_btns = tk.Frame(parent, bg=C["panel"])
996
+ queue_btns.pack(fill=tk.X, pady=(4, 0))
997
+
998
+ self.queue_run_btn = ttk.Button(queue_btns, text="Run Queue",
999
+ command=self.on_run_queue, style="Go.TButton")
1000
+ self.queue_run_btn.pack(side=tk.LEFT, padx=(0, 4))
1001
+
1002
+ for txt, cmd in [("Remove", self._queue_remove), ("Clear", self._queue_clear),
1003
+ ("Up", self._queue_move_up), ("Down", self._queue_move_down),
1004
+ ("+ Current", self._queue_add_current)]:
1005
+ ttk.Button(queue_btns, text=txt, command=cmd).pack(side=tk.LEFT, padx=2)
1006
+
1007
+ # ── Status bar ────────────────────────────────────────────────────
1008
+ status_frame = tk.Frame(parent, bg=C["bg"], padx=8, pady=6)
1009
+ status_frame.pack(fill=tk.X, side=tk.BOTTOM)
1010
+
1011
+ self.status_var = tk.StringVar(value="Ready")
1012
+ tk.Label(status_frame, textvariable=self.status_var,
1013
+ bg=C["bg"], fg=C["green"], font=("Segoe UI", 9),
1014
+ anchor="w").pack(fill=tk.X)
1015
+
1016
+ def _build_grid(self, parent):
1017
+ self.canvas = tk.Canvas(parent, bg=C["bg"], highlightthickness=0)
1018
+ scrollbar = ttk.Scrollbar(parent, orient=tk.VERTICAL, command=self.canvas.yview)
1019
+ self.grid_frame = tk.Frame(self.canvas, bg=C["bg"])
1020
+
1021
+ self.grid_frame.bind("<Configure>",
1022
+ lambda e: self.canvas.configure(
1023
+ scrollregion=self.canvas.bbox("all")))
1024
+ self.canvas_window = self.canvas.create_window((0, 0), window=self.grid_frame,
1025
+ anchor="nw")
1026
+ self.canvas.configure(yscrollcommand=scrollbar.set)
1027
+
1028
+ self.canvas.pack(side=tk.LEFT, fill=tk.BOTH, expand=True)
1029
+ scrollbar.pack(side=tk.RIGHT, fill=tk.Y)
1030
+
1031
+ self.canvas.bind("<Configure>", self._on_canvas_resize)
1032
+ self.canvas.bind_all("<MouseWheel>",
1033
+ lambda e: self.canvas.yview_scroll(
1034
+ int(-1 * (e.delta / 120)), "units"))
1035
+
1036
+ self.placeholder = tk.Label(self.grid_frame,
1037
+ text="Generated images\nwill appear here",
1038
+ bg=C["bg"], fg=C["text3"],
1039
+ font=("Segoe UI", 13), justify="center")
1040
+ self.placeholder.grid(row=0, column=0, pady=80)
1041
+
1042
+ # ── Event handlers ────────────────────────────────────────────────────
1043
+
1044
+ def _on_device_change(self, event=None):
1045
+ choice = self.device_var.get()
1046
+ new_dev = "cuda" if choice == "GPU" else "cpu"
1047
+ self.status_var.set(f"Switching to {choice}...")
1048
+ self.root.update()
1049
+ self.gen.switch_device(new_dev)
1050
+ self.status_var.set(f"Now using {choice}")
1051
+
1052
+ def _on_scheduler_change(self, event=None):
1053
+ name = self.scheduler_var.get()
1054
+ self.gen.set_scheduler(name)
1055
+ self.status_var.set(f"Scheduler: {name}")
1056
+
1057
+ def _on_canvas_resize(self, event):
1058
+ self.canvas.itemconfig(self.canvas_window, width=event.width)
1059
+ if self.generated_images:
1060
+ self._layout_grid()
1061
+
1062
+ def _get_grid_cols(self):
1063
+ canvas_w = self.canvas.winfo_width()
1064
+ if canvas_w < 50:
1065
+ canvas_w = 560
1066
+ tile_size = self._get_tile_size()
1067
+ return max(1, canvas_w // (tile_size + 16))
1068
+
1069
+ def _get_tile_size(self):
1070
+ n = len(self.generated_images)
1071
+ if n <= 2: return 260
1072
+ elif n <= 4: return 220
1073
+ elif n <= 6: return 180
1074
+ else: return 160
1075
+
1076
+ def _layout_grid(self):
1077
+ for w in self.grid_frame.winfo_children():
1078
+ w.destroy()
1079
+ self.photo_refs.clear()
1080
+
1081
+ if not self.generated_images:
1082
+ return
1083
+
1084
+ tile_size = self._get_tile_size()
1085
+ cols = self._get_grid_cols()
1086
+
1087
+ for i, (img, seed) in enumerate(zip(self.generated_images, self.generated_seeds)):
1088
+ row, col = divmod(i, cols)
1089
+ is_selected = (i == self.selected_index)
1090
+
1091
+ card_bg = C["accent"] if is_selected else C["card"]
1092
+ card = tk.Frame(self.grid_frame, bg=card_bg, padx=3, pady=3)
1093
+ card.grid(row=row, column=col, padx=5, pady=5, sticky="nsew")
1094
+
1095
+ display = img.resize((tile_size, tile_size), Image.LANCZOS)
1096
+ photo = ImageTk.PhotoImage(display)
1097
+ self.photo_refs.append(photo)
1098
+
1099
+ img_label = tk.Label(card, image=photo, bg=card_bg, bd=0)
1100
+ img_label.pack()
1101
+ img_label.bind("<Button-1>", lambda e, idx=i: self._select_image(idx))
1102
+ img_label.bind("<Button-3>", lambda e, idx=i: self._show_refine_menu(e, idx))
1103
+
1104
+ tk.Label(card, text=f"seed: {seed}", bg=card_bg,
1105
+ fg=C["text3"], font=("Segoe UI", 8)).pack()
1106
+
1107
+ for c in range(cols):
1108
+ self.grid_frame.columnconfigure(c, weight=1)
1109
+
1110
+ def _select_image(self, idx):
1111
+ if idx >= len(self.generated_images):
1112
+ return
1113
+ self.selected_index = idx
1114
+ self.save_btn.configure(state=tk.NORMAL)
1115
+ self.status_var.set(f"Selected image {idx + 1} (seed: {self.generated_seeds[idx]})")
1116
+ self._layout_grid()
1117
+
1118
+ def _show_refine_menu(self, event, idx):
1119
+ if self.generating:
1120
+ return
1121
+ menu = tk.Menu(self.root, tearoff=0, bg=C["card"], fg=C["text"],
1122
+ activebackground=C["accent"], activeforeground="#fff",
1123
+ font=("Segoe UI", 10), bd=0)
1124
+ menu.add_command(label=" Refine (more steps)... ",
1125
+ command=lambda: self._ask_refine(idx))
1126
+ menu.tk_popup(event.x_root, event.y_root)
1127
+
1128
+ def _ask_refine(self, idx):
1129
+ extra = simpledialog.askinteger(
1130
+ "Refine Image", "Extra denoising steps:",
1131
+ initialvalue=20, minvalue=5, maxvalue=200, parent=self.root)
1132
+ if extra is None:
1133
+ return
1134
+ self._select_image(idx)
1135
+ self.generating = True
1136
+ self.gen.cancelled = False
1137
+ self.gen_btn.configure(state=tk.DISABLED)
1138
+ self.stop_btn.configure(state=tk.NORMAL)
1139
+ self.status_var.set(f"Refining image {idx + 1}...")
1140
+ self.root.update()
1141
+ Thread(target=self._refine_thread, args=(idx, extra), daemon=True).start()
1142
+
1143
+ def _refine_thread(self, idx, extra_steps):
1144
+ try:
1145
+ source = self.generated_images[idx]
1146
+ prompt = self.prompt_entry.get().strip()
1147
+ neg = self.neg_entry.get().strip()
1148
+ cfg = float(self.cfg_var.get())
1149
+ callback = self._show_preview if self.live_preview_var.get() else None
1150
+
1151
+ refined = self.gen.refine(
1152
+ source_image=source, prompt=prompt, negative_prompt=neg,
1153
+ extra_steps=extra_steps, guidance_scale=cfg,
1154
+ preview_callback=callback, preview_every=5)
1155
+
1156
+ if refined is not None:
1157
+ self.generated_images[idx] = refined
1158
+ self.generated_seeds[idx] = f"{self.generated_seeds[idx]}+R{extra_steps}"
1159
+ self._layout_grid()
1160
+ self.status_var.set(f"Refined image {idx + 1}")
1161
+ else:
1162
+ self.status_var.set("Refine stopped.")
1163
+ self.root.update()
1164
+ except Exception as e:
1165
+ self.status_var.set(f"Refine error: {e}")
1166
+ import traceback; traceback.print_exc()
1167
+ finally:
1168
+ self.generating = False
1169
+ self.gen.cancelled = False
1170
+ self.gen_btn.configure(state=tk.NORMAL)
1171
+ self.stop_btn.configure(state=tk.DISABLED)
1172
+
1173
+ # ── Queue ─────────────────────────────────────────────────────────────
1174
+
1175
+ def _queue_add(self):
1176
+ text = self.queue_entry.get().strip()
1177
+ if text:
1178
+ self.queue_listbox.insert(tk.END, text)
1179
+ self.queue_entry.delete(0, tk.END)
1180
+
1181
+ def _queue_add_current(self):
1182
+ text = self.prompt_entry.get().strip()
1183
+ if text:
1184
+ self.queue_listbox.insert(tk.END, text)
1185
+
1186
+ def _queue_remove(self):
1187
+ sel = self.queue_listbox.curselection()
1188
+ if sel:
1189
+ self.queue_listbox.delete(sel[0])
1190
+
1191
+ def _queue_clear(self):
1192
+ self.queue_listbox.delete(0, tk.END)
1193
+
1194
+ def _queue_move_up(self):
1195
+ sel = self.queue_listbox.curselection()
1196
+ if sel and sel[0] > 0:
1197
+ idx = sel[0]
1198
+ text = self.queue_listbox.get(idx)
1199
+ self.queue_listbox.delete(idx)
1200
+ self.queue_listbox.insert(idx - 1, text)
1201
+ self.queue_listbox.selection_set(idx - 1)
1202
+
1203
+ def _queue_move_down(self):
1204
+ sel = self.queue_listbox.curselection()
1205
+ if sel and sel[0] < self.queue_listbox.size() - 1:
1206
+ idx = sel[0]
1207
+ text = self.queue_listbox.get(idx)
1208
+ self.queue_listbox.delete(idx)
1209
+ self.queue_listbox.insert(idx + 1, text)
1210
+ self.queue_listbox.selection_set(idx + 1)
1211
+
1212
+ def on_run_queue(self):
1213
+ if self.generating or not self.models:
1214
+ return
1215
+ prompts = list(self.queue_listbox.get(0, tk.END))
1216
+ if not prompts:
1217
+ self.status_var.set("Queue is empty")
1218
+ return
1219
+ self.generating = True
1220
+ self.gen.cancelled = False
1221
+ self.gen_btn.configure(state=tk.DISABLED)
1222
+ self.queue_run_btn.configure(state=tk.DISABLED)
1223
+ self.stop_btn.configure(state=tk.NORMAL)
1224
+ Thread(target=self._queue_thread, args=(prompts,), daemon=True).start()
1225
+
1226
+ def _queue_thread(self, prompts):
1227
+ try:
1228
+ idx = self.model_combo.current()
1229
+ mdl = self.models[idx]
1230
+ self.status_var.set(f"Loading {mdl[1]}...")
1231
+ self.root.update()
1232
+ self.gen.load_model(mdl[2], mdl[3])
1233
+
1234
+ neg = self.neg_entry.get().strip()
1235
+ steps = int(self.steps_var.get())
1236
+ cfg = float(self.cfg_var.get())
1237
+ num_images = max(1, min(12, int(self.count_var.get())))
1238
+ live_preview = self.live_preview_var.get()
1239
+
1240
+ self.generated_images.clear()
1241
+ self.generated_seeds.clear()
1242
+ self.selected_index = None
1243
+ if self.placeholder:
1244
+ self.placeholder.destroy()
1245
+ self.placeholder = None
1246
+
1247
+ for p_idx, prompt in enumerate(prompts):
1248
+ if self.gen.cancelled:
1249
+ break
1250
+ self.queue_listbox.selection_clear(0, tk.END)
1251
+ self.queue_listbox.selection_set(p_idx)
1252
+ self.queue_listbox.see(p_idx)
1253
+
1254
+ for img_i in range(num_images):
1255
+ if self.gen.cancelled:
1256
+ break
1257
+ self.status_var.set(
1258
+ f"[{p_idx + 1}/{len(prompts)}] image {img_i + 1}/{num_images}")
1259
+ self.root.update()
1260
+
1261
+ callback = None
1262
+ if live_preview:
1263
+ self._setup_preview_card()
1264
+ callback = self._show_preview
1265
+
1266
+ image, used_seed = self.gen.generate(
1267
+ prompt=prompt, negative_prompt=neg,
1268
+ steps=steps, guidance_scale=cfg,
1269
+ preview_callback=callback, preview_every=5)
1270
+
1271
+ if image is None:
1272
+ break
1273
+ self.generated_images.append(image)
1274
+ self.generated_seeds.append(used_seed)
1275
+ save_path = self._next_save_path(prompt)
1276
+ image.save(save_path)
1277
+ self._layout_grid()
1278
+ self.root.update()
1279
+
1280
+ if self.gen.cancelled:
1281
+ break
1282
+
1283
+ done = len(self.generated_images)
1284
+ self.status_var.set(
1285
+ f"Queue {'stopped' if self.gen.cancelled else 'done'}! {done} images saved.")
1286
+ if done > 0:
1287
+ self.save_all_btn.configure(state=tk.NORMAL)
1288
+
1289
+ except Exception as e:
1290
+ self.status_var.set(f"Queue error: {e}")
1291
+ import traceback; traceback.print_exc()
1292
+ finally:
1293
+ self.generating = False
1294
+ self.gen.cancelled = False
1295
+ self.gen_btn.configure(state=tk.NORMAL)
1296
+ self.queue_run_btn.configure(state=tk.NORMAL)
1297
+ self.stop_btn.configure(state=tk.DISABLED)
1298
+
1299
+ # ── Generation ────────────────────────────────────────────────────────
1300
+
1301
+ def on_stop(self):
1302
+ if self.generating:
1303
+ self.gen.cancelled = True
1304
+ self.status_var.set("Stopping...")
1305
+ self.root.update()
1306
+
1307
+ def on_generate(self):
1308
+ if self.generating or not self.models:
1309
+ return
1310
+ self.generating = True
1311
+ self.gen.cancelled = False
1312
+ self.gen_btn.configure(state=tk.DISABLED)
1313
+ self.stop_btn.configure(state=tk.NORMAL)
1314
+ self.status_var.set("Loading model...")
1315
+ self.root.update()
1316
+ Thread(target=self._generate_thread, daemon=True).start()
1317
+
1318
+ def _setup_preview_card(self):
1319
+ tile_size = self._get_tile_size()
1320
+ cols = self._get_grid_cols()
1321
+ row, col = divmod(len(self.generated_images), cols)
1322
+ card = tk.Frame(self.grid_frame, bg=C["card"], padx=3, pady=3)
1323
+ card.grid(row=row, column=col, padx=5, pady=5, sticky="nsew")
1324
+ self._preview_label = tk.Label(card, bg=C["card"],
1325
+ width=tile_size, height=tile_size)
1326
+ self._preview_label.pack()
1327
+ self.root.update()
1328
+
1329
+ def _show_preview(self, preview_img, step, total):
1330
+ tile_size = self._get_tile_size()
1331
+ display = preview_img.resize((tile_size, tile_size), Image.LANCZOS)
1332
+ photo = ImageTk.PhotoImage(display)
1333
+ self._preview_photo = photo
1334
+ if hasattr(self, '_preview_label') and self._preview_label.winfo_exists():
1335
+ self._preview_label.configure(image=photo)
1336
+ self.status_var.set(f"Step {step}/{total}")
1337
+ self.root.update()
1338
+
1339
+ def _generate_thread(self):
1340
+ try:
1341
+ idx = self.model_combo.current()
1342
+ mdl = self.models[idx]
1343
+ self.status_var.set(f"Loading {mdl[1]}...")
1344
+ self.root.update()
1345
+ self.gen.load_model(mdl[2], mdl[3])
1346
+
1347
+ prompt = self.prompt_entry.get().strip()
1348
+ neg = self.neg_entry.get().strip()
1349
+ steps = int(self.steps_var.get())
1350
+ cfg = float(self.cfg_var.get())
1351
+ num_images = max(1, min(12, int(self.count_var.get())))
1352
+ live_preview = self.live_preview_var.get()
1353
+
1354
+ self.generated_images.clear()
1355
+ self.generated_seeds.clear()
1356
+ self.selected_index = None
1357
+ if self.placeholder:
1358
+ self.placeholder.destroy()
1359
+ self.placeholder = None
1360
+
1361
+ for i in range(num_images):
1362
+ if self.gen.cancelled:
1363
+ break
1364
+ self.status_var.set(f"Generating {i + 1}/{num_images}...")
1365
+ self.root.update()
1366
+
1367
+ callback = None
1368
+ if live_preview:
1369
+ self._setup_preview_card()
1370
+ callback = self._show_preview
1371
+
1372
+ image, used_seed = self.gen.generate(
1373
+ prompt=prompt, negative_prompt=neg,
1374
+ steps=steps, guidance_scale=cfg,
1375
+ preview_callback=callback, preview_every=5)
1376
+
1377
+ if image is None:
1378
+ break
1379
+ self.generated_images.append(image)
1380
+ self.generated_seeds.append(used_seed)
1381
+ self._layout_grid()
1382
+ self.root.update()
1383
+
1384
+ done = len(self.generated_images)
1385
+ if self.gen.cancelled:
1386
+ self.status_var.set(f"Stopped. {done} image(s) kept.")
1387
+ else:
1388
+ self.status_var.set(f"Done! {done} images. Click to select.")
1389
+ if done > 0:
1390
+ self.save_all_btn.configure(state=tk.NORMAL)
1391
+ self.save_btn.configure(state=tk.DISABLED)
1392
+
1393
+ except Exception as e:
1394
+ self.status_var.set(f"Error: {e}")
1395
+ import traceback; traceback.print_exc()
1396
+ finally:
1397
+ self.generating = False
1398
+ self.gen.cancelled = False
1399
+ self.gen_btn.configure(state=tk.NORMAL)
1400
+ self.stop_btn.configure(state=tk.DISABLED)
1401
+
1402
+ # ── Save ──────────────────────────────────────────────────────────────
1403
+
1404
+ def _next_save_path(self, prompt_text):
1405
+ OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
1406
+ slug = prompt_text.strip()[:50] if prompt_text.strip() else "untitled"
1407
+ base = OUTPUT_DIR / f"{slug}.png"
1408
+ if not base.exists():
1409
+ return base
1410
+ n = 1
1411
+ while True:
1412
+ path = OUTPUT_DIR / f"{slug} {n}.png"
1413
+ if not path.exists():
1414
+ return path
1415
+ n += 1
1416
+
1417
+ def on_save(self):
1418
+ if self.selected_index is None or not self.generated_images:
1419
+ return
1420
+ img = self.generated_images[self.selected_index]
1421
+ path = self._next_save_path(self.prompt_entry.get().strip())
1422
+ img.save(path)
1423
+ self.status_var.set(f"Saved: {path.name}")
1424
+
1425
+ def on_save_all(self):
1426
+ if not self.generated_images:
1427
+ return
1428
+ prompt_text = self.prompt_entry.get().strip()
1429
+ for img in self.generated_images:
1430
+ path = self._next_save_path(prompt_text)
1431
+ img.save(path)
1432
+ self.status_var.set(f"Saved {len(self.generated_images)} images")
1433
+
1434
+ def run(self):
1435
+ self.root.mainloop()
1436
+
1437
+
1438
+ # ── Entry point ───────────────────────────────────────────────────────────────
1439
+
1440
+ if __name__ == "__main__":
1441
+ models = find_models()
1442
+ if not models:
1443
+ print("No models found locally. Downloading from HuggingFace...")
1444
+ result = download_from_hf()
1445
+ if result:
1446
+ models = find_models()
1447
+
1448
+ if not models:
1449
+ print("No models found!")
1450
+ print(f"Place model weights in: {MODEL_DIR}/YourModelName/")
1451
+ print("Expected files: diffusion_pytorch_model.safetensors or ema_unet.pt")
1452
+ sys.exit(1)
1453
+
1454
+ print(f"Found {len(models)} model(s): {', '.join(m[1] for m in models)}")
1455
+ print(f"Device: {'CUDA (GPU)' if torch.cuda.is_available() else 'CPU'}")
1456
+ print("Starting Aniimage...")
1457
+
1458
+ app = App()
1459
+ app.run()