Team Ai
Modelpublic

FlameF0X/ShellD

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
5likes30downloads
inference.py794 linesDownload Raw Back to root
1"""
2ShellD (Shell Diffusion) - Standalone Inference
3================================================
4Generate 256x256 images from text prompts using a pre-trained ShellD model.
5
6This file is fully self-contained — it does NOT import from train.py.
7It duplicates only the architecture classes needed for inference, with all
8training-specific code (dataset, optimizers, VAE pretraining, git cloning, etc.) removed.
9
10Usage (load from Hugging Face):
11    from inference import ShellDInference
12
13    pipe = ShellDInference("FlameF0X/ShellD")
14    img = pipe.generate("a mountain lake at sunset")
15    img.save("output.png")
16
17Usage (load from local directory):
18    pipe = ShellDInference("./ShellD_model")
19"""
20
21import json
22import math
23from dataclasses import dataclass, asdict
24from pathlib import Path
25from typing import Dict, Any, List, Optional, Tuple
26
27import torch
28import torch.nn as nn
29from huggingface_hub import snapshot_download
30import torch.nn.functional as F
31from safetensors.torch import load_file
32from sentence_transformers import SentenceTransformer
33from PIL import Image
34import numpy as np
35
36
37# ═══════════════════════════════════════════════════════════
38# Config
39# ═══════════════════════════════════════════════════════════
40
41@dataclass
42class ShellDConfig:
43    """Model hyperparameters. Must match the saved config.json."""
44    image_size: int = 256
45    latent_dim: int = 16
46    ae_hidden_dim: int = 64
47    ae_num_blocks: int = 3
48    hidden_dim: int = 256
49    num_hidden_layers: int = 12
50    num_heads: int = 8
51    patch_size: int = 4
52    text_encoder_name: str = "./all-MiniLM-L6-v2"
53    text_encoder_dim: int = 384
54    num_timesteps: int = 1000
55    beta_start: float = 1e-4
56    beta_end: float = 0.02
57    model_name: str = "ShellD"
58    dropout: float = 0.0               # match train.py; 0 = no dropout (inference)
59    kl_weight: float = 0.1             # VAE KL weight (not used in inference)
60    cfg_dropout_prob: float = 0.15     # CFG dropout prob (not used in inference)
61
62    @classmethod
63    def from_json(cls, path: str) -> "ShellDConfig":
64        with open(path) as f:
65            d = json.load(f)
66        # Keep only fields that exist in the dataclass
67        valid_keys = {f.name for f in cls.__dataclass_fields__.values()}
68        d = {k: v for k, v in d.items() if k in valid_keys}
69        return cls(**d)
70
71
72# ═══════════════════════════════════════════════════════════
73# VAE — Encoder / Decoder
74# ═══════════════════════════════════════════════════════════
75
76class ResidualBlock(nn.Module):
77    def __init__(self, in_ch: int, out_ch: int):
78        super().__init__()
79        self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1)
80        self.norm1 = nn.GroupNorm(8, out_ch)
81        self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)
82        self.norm2 = nn.GroupNorm(8, out_ch)
83        self.skip = nn.Conv2d(in_ch, out_ch, 1) if in_ch != out_ch else nn.Identity()
84
85    def forward(self, x: torch.Tensor) -> torch.Tensor:
86        residual = self.skip(x)
87        x = F.silu(self.norm1(self.conv1(x)))
88        x = self.norm2(self.conv2(x))
89        return F.silu(x + residual)
90
91
92class Encoder(nn.Module):
93    def __init__(self, cfg: ShellDConfig):
94        super().__init__()
95        self.cfg = cfg
96        self.init_conv = nn.Conv2d(3, cfg.ae_hidden_dim, 3, padding=1)
97
98        blocks = []
99        ch = cfg.ae_hidden_dim
100        for _ in range(cfg.ae_num_blocks):
101            out_ch = min(ch * 2, 256)
102            blocks.append(
103                nn.Sequential(
104                    ResidualBlock(ch, out_ch),
105                    ResidualBlock(out_ch, out_ch),
106                    nn.Conv2d(out_ch, out_ch, 4, stride=2, padding=1),  # downsample
107                )
108            )
109            ch = out_ch
110        self.down_blocks = nn.ModuleList(blocks)
111
112        self.mid = nn.Sequential(
113            ResidualBlock(ch, ch),
114            ResidualBlock(ch, ch),
115        )
116        self.out_conv = nn.Conv2d(ch, cfg.latent_dim * 2, 3, padding=1)
117
118    def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
119        x = self.init_conv(x)
120        for block in self.down_blocks:
121            x = block(x)
122        x = self.mid(x)
123        x = self.out_conv(x)
124        mean, logvar = x.chunk(2, dim=1)
125        return mean, logvar
126
127
128class Decoder(nn.Module):
129    def __init__(self, cfg: ShellDConfig):
130        super().__init__()
131        self.cfg = cfg
132        ch = cfg.ae_hidden_dim * (2 ** cfg.ae_num_blocks)
133        self.init_conv = nn.Conv2d(cfg.latent_dim, ch, 3, padding=1)
134
135        self.mid = nn.Sequential(
136            ResidualBlock(ch, ch),
137            ResidualBlock(ch, ch),
138        )
139
140        up_blocks = []
141        for _ in range(cfg.ae_num_blocks):
142            out_ch = ch // 2
143            up_blocks.append(
144                nn.Sequential(
145                    ResidualBlock(ch, out_ch),
146                    ResidualBlock(out_ch, out_ch),
147                    nn.Upsample(scale_factor=2, mode="bilinear", align_corners=False),
148                )
149            )
150            ch = out_ch
151        self.up_blocks = nn.ModuleList(up_blocks)
152
153        self.out_conv = nn.Sequential(
154            ResidualBlock(ch, ch),
155            nn.Conv2d(ch, 3, 3, padding=1),
156            nn.Tanh(),
157        )
158
159    def forward(self, z: torch.Tensor) -> torch.Tensor:
160        x = self.init_conv(z)
161        x = self.mid(x)
162        for block in self.up_blocks:
163            x = block(x)
164        return self.out_conv(x)
165
166
167class VAE(nn.Module):
168    def __init__(self, cfg: ShellDConfig):
169        super().__init__()
170        self.cfg = cfg
171        self.encoder = Encoder(cfg)
172        self.decoder = Decoder(cfg)
173
174    def encode(self, x: torch.Tensor) -> torch.Tensor:
175        mean, logvar = self.encoder(x)
176        logvar = logvar.clamp(-10.0, 10.0)
177        std = torch.exp(0.5 * logvar)
178        eps = torch.randn_like(std)
179        return mean + eps * std
180
181    def decode(self, z: torch.Tensor) -> torch.Tensor:
182        return self.decoder(z)
183
184    def forward(self, x: torch.Tensor):
185        """Full forward (used only for training; included for completeness)."""
186        mean, logvar = self.encoder(x)
187        logvar = logvar.clamp(-10.0, 10.0)
188        std = torch.exp(0.5 * logvar)
189        eps = torch.randn_like(std)
190        z = mean + eps * std
191        return self.decode(z)
192
193
194# ═══════════════════════════════════════════════════════════
195# DiT Backbone
196# ═══════════════════════════════════════════════════════════
197
198class PatchEmbed(nn.Module):
199    def __init__(self, cfg: ShellDConfig):
200        super().__init__()
201        self.cfg = cfg
202        self.latent_size = cfg.image_size // (2 ** cfg.ae_num_blocks)
203        self.num_patches_1d = self.latent_size // cfg.patch_size
204        self.num_patches = self.num_patches_1d ** 2
205        self.patch_dim = cfg.latent_dim * (cfg.patch_size ** 2)
206        self.proj = nn.Linear(self.patch_dim, cfg.hidden_dim)
207
208    def forward(self, z: torch.Tensor) -> torch.Tensor:
209        B, C, H, W = z.shape
210        ps = self.cfg.patch_size
211        n = self.num_patches_1d
212        z = z.reshape(B, C, n, ps, n, ps)
213        z = z.permute(0, 2, 4, 1, 3, 5).reshape(B, self.num_patches, self.patch_dim)
214        return self.proj(z)
215
216
217class DiTBlock(nn.Module):
218    def __init__(self, cfg: ShellDConfig):
219        super().__init__()
220        dim = cfg.hidden_dim
221        drop = cfg.dropout
222
223        self.norm1 = nn.LayerNorm(dim)
224        self.attn = nn.MultiheadAttention(dim, cfg.num_heads, batch_first=True)
225        self.drop_attn = nn.Dropout(drop)
226
227        self.norm_cross = nn.LayerNorm(dim)
228        self.cross_attn = nn.MultiheadAttention(dim, cfg.num_heads, batch_first=True)
229        self.drop_cross = nn.Dropout(drop)
230        self.text_proj = nn.Linear(cfg.text_encoder_dim, dim)
231
232        self.norm2 = nn.LayerNorm(dim)
233        self.mlp = nn.Sequential(
234            nn.Linear(dim, dim * 4),
235            nn.GELU(),
236            nn.Dropout(drop),
237            nn.Linear(dim * 4, dim),
238        )
239        self.drop_mlp = nn.Dropout(drop)
240
241        self.timestep_mlp = nn.Sequential(
242            nn.Linear(dim, dim * 4),
243            nn.SiLU(),
244            nn.Linear(dim * 4, dim),
245        )
246
247    def forward(self, x: torch.Tensor, text_emb: torch.Tensor, t_emb: torch.Tensor):
248        # Self-attention
249        h = self.norm1(x)
250        attn_out, _ = self.attn(h, h, h)
251        x = x + self.drop_attn(attn_out)
252        # Cross-attention with text
253        text_proj = self.text_proj(text_emb)
254        h = self.norm_cross(x)
255        cross_out, _ = self.cross_attn(h, text_proj, text_proj)
256        x = x + self.drop_cross(cross_out)
257        # Timestep conditioning
258        t_proj = self.timestep_mlp(t_emb)
259        x = x + t_proj.unsqueeze(1)
260        # MLP
261        h = self.norm2(x)
262        x = x + self.drop_mlp(self.mlp(h))
263        return x
264
265
266class DiT(nn.Module):
267    def __init__(self, cfg: ShellDConfig):
268        super().__init__()
269        self.cfg = cfg
270        self.latent_size = cfg.image_size // (2 ** cfg.ae_num_blocks)
271        self.num_patches_1d = self.latent_size // cfg.patch_size
272        self.num_patches = self.num_patches_1d ** 2
273
274        self.patch_embed = PatchEmbed(cfg)
275        self.pos_embed = nn.Parameter(torch.randn(1, self.num_patches, cfg.hidden_dim))
276        self.blocks = nn.ModuleList([DiTBlock(cfg) for _ in range(cfg.num_hidden_layers)])
277        self.norm = nn.LayerNorm(cfg.hidden_dim)
278        self.out_proj = nn.Linear(cfg.hidden_dim, cfg.latent_dim * (cfg.patch_size ** 2))
279
280        self.time_mlp = nn.Sequential(
281            nn.Linear(cfg.hidden_dim, cfg.hidden_dim * 4),
282            nn.SiLU(),
283            nn.Linear(cfg.hidden_dim * 4, cfg.hidden_dim),
284        )
285
286    @staticmethod
287    def timestep_embedding(t: torch.Tensor, dim: int, max_period: int = 10000):
288        half = dim // 2
289        freqs = torch.exp(
290            -math.log(max_period) * torch.arange(0, half, dtype=torch.float32) / half
291        ).to(t.device)
292        args = t[:, None].float() * freqs[None]
293        emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
294        if dim % 2 == 1:
295            emb = torch.cat([emb, torch.zeros_like(emb[:, :1])], dim=-1)
296        return emb
297
298    def forward(self, z: torch.Tensor, text_emb: torch.Tensor, t: torch.Tensor):
299        B = z.shape[0]
300        dim = self.cfg.hidden_dim
301        t_emb = self.timestep_embedding(t, dim)
302        t_emb = self.time_mlp(t_emb)
303
304        x = self.patch_embed(z)
305        x = x + self.pos_embed
306
307        for blk in self.blocks:
308            x = blk(x, text_emb, t_emb)
309
310        x = self.norm(x)
311        x = self.out_proj(x)
312
313        ps = self.cfg.patch_size
314        H = W = self.latent_size
315        n = self.num_patches_1d
316        x = x.reshape(B, n, n, self.cfg.latent_dim, ps, ps)
317        x = x.permute(0, 3, 1, 4, 2, 5).reshape(B, self.cfg.latent_dim, H, W)
318        return x
319
320
321# ═══════════════════════════════════════════════════════════
322# Full ShellD Model (Inference-only)
323# ═══════════════════════════════════════════════════════════
324
325class ShellDModel(nn.Module):
326    """ShellD model for inference. Minimal — no training helpers."""
327
328    def __init__(self, cfg: ShellDConfig):
329        super().__init__()
330        self.cfg = cfg
331        self.vae = VAE(cfg)
332        self.dit = DiT(cfg)
333        self.text_encoder = None  # loaded separately
334        # Null text embedding for classifier-free guidance (loaded from checkpoint)
335        self.null_text_embed: Optional[torch.Tensor] = None
336
337    def encode_text(self, prompts: List[str], device: torch.device) -> torch.Tensor:
338        assert self.text_encoder is not None, "Text encoder not loaded"
339        with torch.no_grad():
340            emb = self.text_encoder.encode(prompts, convert_to_tensor=True)
341        return emb.to(device).unsqueeze(1)  # [B, 1, 384]
342
343    def forward(self, z: torch.Tensor, text_emb: torch.Tensor, t: torch.Tensor) -> torch.Tensor:
344        """Predict the noise at timestep t given noisy latent z and text."""
345        return self.dit(z, text_emb, t)
346
347
348# ═══════════════════════════════════════════════════════════
349# Diffusion Schedule (Reverse Process Only)
350# ═══════════════════════════════════════════════════════════
351
352class DiffusionSchedule:
353    """DDPM schedule for the reverse (denoising) process."""
354
355    def __init__(self, cfg: ShellDConfig, device: torch.device):
356        betas = torch.linspace(cfg.beta_start, cfg.beta_end, cfg.num_timesteps, device=device)
357        alphas = 1.0 - betas
358        alphas_cumprod = torch.cumprod(alphas, dim=0)
359
360        self.betas = betas
361        self.alphas = alphas
362        self.alphas_cumprod = alphas_cumprod
363        self.sqrt_alphas_cumprod = alphas_cumprod.sqrt()
364        self.sqrt_one_minus_alphas_cumprod = (1.0 - alphas_cumprod).sqrt()
365
366    @torch.no_grad()
367    def sample(
368        self,
369        model: ShellDModel,
370        text_emb: torch.Tensor,
371        num_steps: Optional[int] = None,
372        cfg_scale: float = 3.0,
373        seed: Optional[int] = None,
374    ) -> torch.Tensor:
375        """
376        DDPM reverse sampling with optional classifier-free guidance.
377
378        Args:
379            model: The ShellD model.
380            text_emb: Text embedding [B, 1, 384].
381            num_steps: Number of denoising steps (default: cfg.num_timesteps).
382            cfg_scale: Classifier-free guidance scale. 1.0 = no guidance.
383            seed: Optional random seed for reproducibility.
384
385        Returns:
386            Denoised latent tensor [B, latent_dim, H', W'].
387        """
388        if seed is not None:
389            torch.manual_seed(seed)
390
391        device = text_emb.device
392        B = text_emb.shape[0]
393        cfg = model.cfg
394
395        num_steps = num_steps or cfg.num_timesteps
396
397        # Latent spatial size
398        H = W = cfg.image_size // (2 ** cfg.ae_num_blocks)
399
400        # Start from random noise
401        z = torch.randn(B, cfg.latent_dim, H, W, device=device)
402
403        # For classifier-free guidance, we need the learned null embedding
404        if cfg_scale != 1.0:
405            if model.null_text_embed is not None:
406                uncond_emb = model.null_text_embed.to(device).expand(B, -1, -1)
407            else:
408                # Fallback: zeros (works if model was trained with zero-dropped text)
409                uncond_emb = torch.zeros_like(text_emb)
410
411        # Time step resampling for faster inference (DDPM-style, evenly spaced)
412        step_indices = torch.linspace(0, cfg.num_timesteps - 1, num_steps, device=device, dtype=torch.long)
413
414        for i in range(num_steps - 1, -1, -1):
415            t = step_indices[i]
416            t_batch = t.expand(B)
417
418            # Predict noise
419            if cfg_scale != 1.0:
420                # Classifier-free guidance: combine conditional and unconditional predictions
421                z_in = torch.cat([z, z], dim=0)
422                t_in = torch.cat([t_batch, t_batch], dim=0)
423                text_in = torch.cat([text_emb, uncond_emb], dim=0)
424                noise_pred = model(z_in, text_in, t_in)
425                noise_cond, noise_uncond = noise_pred.chunk(2, dim=0)
426                noise_pred = noise_uncond + cfg_scale * (noise_cond - noise_uncond)
427            else:
428                noise_pred = model(z, text_emb, t_batch)
429
430            # DDPM update step
431            alpha = self.alphas[t]
432            alpha_cumprod = self.alphas_cumprod[t]
433            beta = self.betas[t]
434
435            # Compute predicted x0 (for logging purposes, not needed for step)
436            # x0_pred = (z - sqrt_one_minus_ac * noise_pred) / sqrt_ac
437
438            # Mean for posterior
439            coef1 = 1.0 / alpha.sqrt()
440            coef2 = beta / (1.0 - alpha_cumprod).sqrt()
441            z_mean = coef1 * (z - coef2 * noise_pred)
442
443            if i > 0:
444                noise = torch.randn_like(z)
445                z = z_mean + beta.sqrt() * noise
446            else:
447                z = z_mean
448
449        return z
450
451
452# ═══════════════════════════════════════════════════════════
453# High-Level Inference Pipeline
454# ═══════════════════════════════════════════════════════════
455
456class ShellDInference:
457    """
458    High-level pipeline for ShellD text-to-image generation.
459
460    Usage:
461        pipe = ShellDInference("./ShellD_model")
462        img = pipe.generate("a cat sitting on a mat")
463        img.save("cat.png")
464    """
465
466    def __init__(
467        self,
468        model_dir: str,
469        device: Optional[str] = None,
470        text_encoder_path: Optional[str] = None,
471    ):
472        """
473        Args:
474            model_dir: Path to the model directory, or a Hugging Face repo ID
475                       (e.g. "FlameF0X/ShellD"). If a repo ID is given, the
476                       weights are automatically downloaded via huggingface_hub.
477            device: Device to run on ('cuda', 'cpu', or None for auto-detect).
478            text_encoder_path: Optional override path for the text encoder model.
479                               Defaults to the path in config.json.
480        """
481        self.device = torch.device(
482            device or ("cuda" if torch.cuda.is_available() else "cpu")
483        )
484        print(f"ShellD inference using device: {self.device}")
485
486        # Resolve model source: Hugging Face repo or local path
487        local_path = Path(model_dir)
488        if local_path.exists():
489            # Local directory
490            self.model_dir = local_path
491        else:
492            # Assume it's a Hugging Face repo ID — download via hub
493            print(f"Downloading model from Hugging Face: {model_dir}")
494            self.model_dir = Path(snapshot_download(repo_id=model_dir))
495            print(f"Model cached at: {self.model_dir}")
496
497        # 1. Load config
498        config_path = self.model_dir / "config.json"
499        if not config_path.exists():
500            raise FileNotFoundError(f"config.json not found at {config_path}")
501        self.cfg = ShellDConfig.from_json(str(config_path))
502
503        # 2. Build model
504        self.model = ShellDModel(self.cfg).to(self.device)
505        self.model.eval()
506
507        # 3. Load weights
508        weights_path = self.model_dir / "model.safetensors"
509        if not weights_path.exists():
510            raise FileNotFoundError(f"model.safetensors not found at {weights_path}")
511        self._load_weights(str(weights_path))
512
513        # 4. Load text encoder
514        encoder_path = text_encoder_path or self.cfg.text_encoder_name
515        print(f"Loading text encoder from: {encoder_path}")
516        self.model.text_encoder = SentenceTransformer(encoder_path)
517
518        # 5. Diffusion schedule
519        self.diffusion = DiffusionSchedule(self.cfg, self.device)
520
521        # Print parameter count
522        total = sum(p.numel() for p in self.model.parameters())
523        trainable = sum(p.numel() for p in self.model.parameters() if p.requires_grad)
524        print(f"Model loaded: {total:,} total params ({trainable:,} trainable)")
525
526    def _load_weights(self, weights_path: str):
527        """Load state dict, mapping prefixes to the correct submodules."""
528        sd = load_file(weights_path)
529
530        vae_sd = {k.replace("vae.", ""): v for k, v in sd.items() if k.startswith("vae.")}
531        dit_sd = {k.replace("dit.", ""): v for k, v in sd.items() if k.startswith("dit.")}
532        txt_sd = {k.replace("text_encoder.", ""): v for k, v in sd.items() if k.startswith("text_encoder.")}
533
534        # Backward compat: old checkpoints (no dropout) have final Linear at mlp.2;
535        # new architecture (with dropout) has it at mlp.3. Remap silently.
536        if any(".mlp.2.weight" in k for k in dit_sd) and not any(".mlp.3.weight" in k for k in dit_sd):
537            remapped = {}
538            for k, v in dit_sd.items():
539                remapped[k.replace(".mlp.2.", ".mlp.3.")] = v
540            dit_sd = remapped
541            print("  ↻ Remapped old DiT checkpoint (no dropout) → new MLP layout.")
542
543        self.model.vae.load_state_dict(vae_sd)
544        self.model.dit.load_state_dict(dit_sd)
545
546        # Restore null text embedding for CFG
547        if "null_text_embed" in sd:
548            self.model.null_text_embed = sd["null_text_embed"].to(self.device)
549            print("Null text embedding loaded for CFG.")
550        else:
551            print("No null_text_embed in checkpoint — CFG will use zeros (may degrade quality).")
552
553        if txt_sd and self.model.text_encoder is not None:
554            self.model.text_encoder.load_state_dict(txt_sd)
555            print("Text encoder weights loaded from safetensors.")
556        else:
557            print("Text encoder loaded from SentenceTransformer cache (weights not in safetensors).")
558
559    @torch.no_grad()
560    def generate(
561        self,
562        prompt: str,
563        num_steps: int = 250,
564        cfg_scale: float = 3.0,
565        seed: Optional[int] = None,
566        output_size: Optional[int] = None,
567    ) -> Image.Image:
568        """
569        Generate an image from a text prompt.
570
571        Args:
572            prompt: Text description of the desired image.
573            num_steps: Number of denoising steps (fewer = faster, lower quality).
574                       Recommended: 250-1000.
575            cfg_scale: Classifier-free guidance scale. Higher = more prompt adherence.
576                       1.0 = no guidance. Typical range: 2.0-5.0.
577            seed: Random seed for reproducibility.
578            output_size: If set, the output image is resized to (output_size, output_size).
579
580        Returns:
581            A PIL Image.
582        """
583        self.model.eval()
584
585        # Encode prompt
586        text_emb = self.model.encode_text([prompt], self.device)  # [1, 1, 384]
587
588        # Sample latent
589        z = self.diffusion.sample(
590            self.model,
591            text_emb,
592            num_steps=num_steps,
593            cfg_scale=cfg_scale,
594            seed=seed,
595        )
596
597        # Decode latent to image (VAE decoder outputs [-1, 1] via Tanh)
598        img_tensor = self.model.vae.decode(z)  # [1, 3, 256, 256], values in [-1, 1]
599
600        # Convert to PIL: [-1,1] → [0,255]
601        img_np = img_tensor[0].permute(1, 2, 0).cpu().numpy()  # [256, 256, 3]
602        img_np = ((img_np + 1.0) * 127.5).clip(0, 255).astype(np.uint8)
603        img = Image.fromarray(img_np)
604
605        if output_size is not None:
606            img = img.resize((output_size, output_size), Image.Resampling.BICUBIC)
607
608        return img
609
610    def _decode_latent(self, z: torch.Tensor) -> Image.Image:
611        """Decode a latent tensor to a PIL Image (helper for streaming)."""
612        img_tensor = self.model.vae.decode(z)  # [B, 3, H, W] in [-1,1]
613        img_np = img_tensor[0].permute(1, 2, 0).cpu().numpy()
614        img_np = ((img_np + 1.0) * 127.5).clip(0, 255).astype(np.uint8)
615        return Image.fromarray(img_np)
616
617    @torch.no_grad()
618    def generate_stream(
619        self,
620        prompt: str,
621        num_steps: int = 250,
622        cfg_scale: float = 3.0,
623        seed: Optional[int] = None,
624        display_every: int = 10,
625    ):
626        """
627        Generator that yields (PIL.Image, str) tuples showing the denoising
628        process unfold step-by-step.  Callers can display each intermediate
629        image as it's produced for a live "diffusion reveal" effect.
630
631        Yields:
632            (image, status_label) — status_label is e.g. "Step 50/250".
633        """
634        self.model.eval()
635        device = self.device
636        cfg = self.model.cfg
637
638        if seed is not None:
639            torch.manual_seed(seed)
640
641        # --- Encode prompt ---
642        text_emb = self.model.encode_text([prompt], device)  # [1, 1, 384]
643        if cfg_scale != 1.0:
644            if self.model.null_text_embed is not None:
645                uncond_emb = self.model.null_text_embed.to(device).expand(1, -1, -1)
646            else:
647                uncond_emb = torch.zeros_like(text_emb)
648        else:
649            uncond_emb = None
650
651        # --- Start from pure noise ---
652        H = W = cfg.image_size // (2 ** cfg.ae_num_blocks)
653        z = torch.randn(1, cfg.latent_dim, H, W, device=device)
654
655        # --- Timestep schedule (evenly spaced) ---
656        step_indices = torch.linspace(
657            0, cfg.num_timesteps - 1, num_steps, device=device, dtype=torch.long
658        )
659
660        # Yield initial noise snapshot
661        yield self._decode_latent(z), f"Step 0/{num_steps} — pure noise"
662
663        # --- DDPM reverse denoising loop ---
664        for i in range(num_steps - 1, -1, -1):
665            t = step_indices[i]
666            t_batch = t.expand(1)
667
668            # Predict noise (with classifier-free guidance)
669            if cfg_scale != 1.0:
670                z_in = torch.cat([z, z], dim=0)
671                t_in = torch.cat([t_batch, t_batch], dim=0)
672                text_in = torch.cat([text_emb, uncond_emb], dim=0)
673                noise_pred = self.model(z_in, text_in, t_in)
674                noise_cond, noise_uncond = noise_pred.chunk(2, dim=0)
675                noise_pred = noise_uncond + cfg_scale * (noise_cond - noise_uncond)
676            else:
677                noise_pred = self.model(z, text_emb, t_batch)
678
679            # DDPM update
680            alpha = self.diffusion.alphas[t]
681            alpha_cumprod = self.diffusion.alphas_cumprod[t]
682            beta = self.diffusion.betas[t]
683
684            coef1 = 1.0 / alpha.sqrt()
685            coef2 = beta / (1.0 - alpha_cumprod).sqrt()
686            z_mean = coef1 * (z - coef2 * noise_pred)
687
688            if i > 0:
689                z = z_mean + beta.sqrt() * torch.randn_like(z)
690            else:
691                z = z_mean
692
693            # Yield intermediate snapshot at display intervals
694            step_no = num_steps - i
695            if step_no % display_every == 0 or i == 0:
696                yield self._decode_latent(z), f"Step {step_no}/{num_steps}"
697
698    @torch.no_grad()
699    def generate_batch(
700        self,
701        prompts: List[str],
702        num_steps: int = 250,
703        cfg_scale: float = 3.0,
704        seed: Optional[int] = None,
705    ) -> List[Image.Image]:
706        """
707        Generate images for multiple prompts efficiently (batched).
708
709        Args:
710            prompts: List of text prompts.
711            num_steps: Number of denoising steps.
712            cfg_scale: Classifier-free guidance scale.
713            seed: Random seed (applied per-prompt).
714
715        Returns:
716            List of PIL Images.
717        """
718        self.model.eval()
719        B = len(prompts)
720
721        # Encode all prompts
722        text_emb = self.model.encode_text(prompts, self.device)  # [B, 1, 384]
723
724        # Sample latents
725        z = self.diffusion.sample(
726            self.model,
727            text_emb,
728            num_steps=num_steps,
729            cfg_scale=cfg_scale,
730            seed=seed,
731        )
732
733        # Decode
734        img_tensor = self.model.vae.decode(z)  # [B, 3, 256, 256] in [-1,1]
735
736        images = []
737        for i in range(B):
738            img_np = img_tensor[i].permute(1, 2, 0).cpu().numpy()
739            img_np = ((img_np + 1.0) * 127.5).clip(0, 255).astype(np.uint8)
740            images.append(Image.fromarray(img_np))
741
742        return images
743
744
745# ═══════════════════════════════════════════════════════════
746# Command-Line Interface
747# ═══════════════════════════════════════════════════════════
748
749def main():
750    import argparse
751
752    parser = argparse.ArgumentParser(description="ShellD - Text-to-Image Generation")
753    parser.add_argument("--model_dir", type=str, default="FlameF0X/ShellD",
754                        help="Path to model directory or Hugging Face repo ID (default: FlameF0X/ShellD)")
755    parser.add_argument("--text_encoder", type=str, default=None,
756                        help="Override path for the text encoder model")
757    parser.add_argument("--prompt", type=str, required=True,
758                        help="Text prompt for image generation")
759    parser.add_argument("--output", type=str, default="output.png",
760                        help="Output image path")
761    parser.add_argument("--steps", type=int, default=250,
762                        help="Number of denoising steps (default: 250)")
763    parser.add_argument("--cfg", type=float, default=3.0,
764                        help="Classifier-free guidance scale (default: 3.0)")
765    parser.add_argument("--seed", type=int, default=None,
766                        help="Random seed for reproducibility")
767    parser.add_argument("--device", type=str, default=None,
768                        help="Device: 'cuda' or 'cpu'")
769    parser.add_argument("--size", type=int, default=None,
770                        help="Output image size (square, resized)")
771
772    args = parser.parse_args()
773
774    pipe = ShellDInference(
775        model_dir=args.model_dir,
776        device=args.device,
777        text_encoder_path=args.text_encoder,
778    )
779
780    img = pipe.generate(
781        prompt=args.prompt,
782        num_steps=args.steps,
783        cfg_scale=args.cfg,
784        seed=args.seed,
785        output_size=args.size,
786    )
787
788    img.save(args.output)
789    print(f"Image saved to {args.output}")
790
791
792if __name__ == "__main__":
793    main()
794