import torch import torch.nn as nn import torch.nn.functional as F from diffusers import AutoencoderKL import math import os import gradio as gr import torchvision.transforms as transforms from contextlib import nullcontext from huggingface_hub import hf_hub_download # Configuration device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Using device: {device}") T = 1000 IMG_SIZE = 32 # latent spatial dim: 256px / 8 (VAE) = 32 CHANNELS = 4 # SD VAE latent channels TIME_DIM = 320 # sinusoidal embedding dim TIME_EMB_DIM = 1280 # projected time embedding dim LATENT_SCALE = 0.18215 HF_MODEL_REPO_ID = "keysun89/face_ldm_ckpt" CKPT_FILENAME = "ckpt_epoch_0095.pt" # DDPM Noise Schedule betas = torch.linspace(1e-4, 0.02, T, device=device) alphas = 1.0 - betas alpha_bars = torch.cumprod(alphas, dim=0) sqrt_one_minus_ab = torch.sqrt(1.0 - alpha_bars) sqrt_recip_a = torch.sqrt(1.0 / alphas) betas_tilde = betas.clone() betas_tilde[1:] = betas[1:] * (1.0 - alpha_bars[:-1]) / (1.0 - alpha_bars[1:]) # Model Architecture Components class TimeEmbedding(nn.Module): def __init__(self, dim=TIME_DIM, time_dim=TIME_EMB_DIM): super().__init__() self.dim = dim self.mlp = nn.Sequential( nn.Linear(dim, time_dim), nn.SiLU(), nn.Linear(time_dim, time_dim), ) def forward(self, t): half = self.dim // 2 freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / (half - 1)) args = t[:, None].float() * freqs[None, :] emb = torch.cat([torch.sin(args), torch.cos(args)], dim=1) return self.mlp(emb) class ResidualBlock(nn.Module): def __init__(self, in_ch, out_ch, time_dim=TIME_EMB_DIM): super().__init__() self.norm1 = nn.GroupNorm(min(32, in_ch), in_ch) self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1) self.norm2 = nn.GroupNorm(min(32, out_ch), out_ch) self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1) self.time_proj = nn.Linear(time_dim, out_ch) self.skip = nn.Identity() if in_ch == out_ch else nn.Conv2d(in_ch, out_ch, 1) def forward(self, x, time_emb): h = self.conv1(F.silu(self.norm1(x))) h = h + self.time_proj(F.silu(time_emb))[:, :, None, None] h = self.conv2(F.silu(self.norm2(h))) return h + self.skip(x) class SelfAttention(nn.Module): def __init__(self, ch, n_heads=8): super().__init__() self.norm = nn.GroupNorm(min(32, ch), ch) self.qkv = nn.Linear(ch, 3 * ch) self.out = nn.Linear(ch, ch) self.n_heads = n_heads self.d_head = ch // n_heads def forward(self, x): B, C, H, W = x.shape h = self.norm(x).view(B, C, H * W).transpose(1, 2) qkv = self.qkv(h).chunk(3, dim=-1) q, k, v = [t.view(B, H * W, self.n_heads, self.d_head).transpose(1, 2) for t in qkv] scale = math.sqrt(self.d_head) attn = (q @ k.transpose(-2, -1) / scale).softmax(dim=-1) out = (attn @ v).transpose(1, 2).reshape(B, H * W, C) out = self.out(out).transpose(1, 2).view(B, C, H, W) return x + out class AttentionBlock(nn.Module): def __init__(self, ch, n_heads=8): super().__init__() self.attn = SelfAttention(ch, n_heads) self.norm = nn.LayerNorm(ch) self.ff1 = nn.Linear(ch, 4 * ch * 2) self.ff2 = nn.Linear(4 * ch, ch) def forward(self, x): x = self.attn(x) B, C, H, W = x.shape h = x.view(B, C, H * W).transpose(1, 2) gate, val = self.ff1(self.norm(h)).chunk(2, dim=-1) h = self.ff2(gate * F.gelu(val)) + h return h.transpose(1, 2).view(B, C, H, W) class Downsample(nn.Module): def __init__(self, ch): super().__init__() self.conv = nn.Conv2d(ch, ch, 3, stride=2, padding=1) def forward(self, x): return self.conv(x) class Upsample(nn.Module): def __init__(self, ch): super().__init__() self.conv = nn.Conv2d(ch, ch, 3, padding=1) def forward(self, x): return self.conv(F.interpolate(x, scale_factor=2.0, mode='nearest')) class UNet(nn.Module): def __init__(self, in_ch=CHANNELS, base_ch=128): super().__init__() ch = base_ch self.stem = nn.Conv2d(in_ch, ch, 3, padding=1) self.enc0_r1 = ResidualBlock(ch, ch) self.enc0_r2 = ResidualBlock(ch, ch) self.enc0_dn = Downsample(ch) self.enc1_r1 = ResidualBlock(ch, ch * 2) self.enc1_a1 = AttentionBlock(ch * 2) self.enc1_r2 = ResidualBlock(ch * 2, ch * 2) self.enc1_a2 = AttentionBlock(ch * 2) self.enc1_dn = Downsample(ch * 2) ch2 = ch * 2 self.enc2_r1 = ResidualBlock(ch2, ch2 * 2) self.enc2_a1 = AttentionBlock(ch2 * 2) self.enc2_r2 = ResidualBlock(ch2 * 2, ch2 * 2) self.enc2_a2 = AttentionBlock(ch2 * 2) self.enc2_dn = Downsample(ch2 * 2) ch4 = ch2 * 2 self.enc3_r1 = ResidualBlock(ch4, ch4) self.enc3_a1 = AttentionBlock(ch4) self.enc3_r2 = ResidualBlock(ch4, ch4) self.enc3_a2 = AttentionBlock(ch4) self.mid_r1 = ResidualBlock(ch4, ch4) self.mid_a = AttentionBlock(ch4) self.mid_r2 = ResidualBlock(ch4, ch4) self.dec3_r1 = ResidualBlock(ch4 * 2, ch4) self.dec3_a1 = AttentionBlock(ch4) self.dec3_r2 = ResidualBlock(ch4 * 2, ch4) self.dec3_a2 = AttentionBlock(ch4) self.dec3_r3 = ResidualBlock(ch4 * 2, ch4) self.dec3_a3 = AttentionBlock(ch4) self.dec3_up = Upsample(ch4) self.dec2_r1 = ResidualBlock(ch4 + ch4, ch4) self.dec2_a1 = AttentionBlock(ch4) self.dec2_r2 = ResidualBlock(ch4 + ch4, ch2) self.dec2_a2 = AttentionBlock(ch2) self.dec2_r3 = ResidualBlock(ch2 + ch2, ch2) self.dec2_up = Upsample(ch2) self.dec1_r1 = ResidualBlock(ch2 + ch2, ch2) self.dec1_a1 = AttentionBlock(ch2) self.dec1_r2 = ResidualBlock(ch2 + ch2, ch) self.dec1_a2 = AttentionBlock(ch) self.dec1_r3 = ResidualBlock(ch + ch, ch) self.dec1_up = Upsample(ch) self.dec0_r1 = ResidualBlock(ch + ch, ch) self.dec0_r2 = ResidualBlock(ch + ch, ch) self.out = nn.Sequential( nn.GroupNorm(32, ch), nn.SiLU(), nn.Conv2d(ch, in_ch, 3, padding=1), ) self._init_weights() def _init_weights(self): nn.init.zeros_(self.out[-1].weight) nn.init.zeros_(self.out[-1].bias) def forward(self, x, temb): h = self.stem(x) s00 = self.enc0_r1(h, temb); s01 = self.enc0_r2(s00, temb) h = self.enc0_dn(s01); s0d = h h = self.enc1_r1(h, temb); s10 = self.enc1_a1(h) h = self.enc1_r2(s10, temb); s11 = self.enc1_a2(h) h = self.enc1_dn(s11); s1d = h h = self.enc2_r1(h, temb); s20 = self.enc2_a1(h) h = self.enc2_r2(s20, temb); s21 = self.enc2_a2(h) h = self.enc2_dn(s21); s2d = h h = self.enc3_r1(h, temb); s30 = self.enc3_a1(h) h = self.enc3_r2(s30, temb); s31 = self.enc3_a2(h) h = self.mid_r1(h, temb) h = self.mid_a(h) h = self.mid_r2(h, temb) h = self.dec3_r1(torch.cat([h, s31], dim=1), temb); h = self.dec3_a1(h) h = self.dec3_r2(torch.cat([h, s30], dim=1), temb); h = self.dec3_a2(h) h = self.dec3_r3(torch.cat([h, s2d], dim=1), temb); h = self.dec3_a3(h) h = self.dec3_up(h) h = self.dec2_r1(torch.cat([h, s21], dim=1), temb); h = self.dec2_a1(h) h = self.dec2_r2(torch.cat([h, s20], dim=1), temb); h = self.dec2_a2(h) h = self.dec2_r3(torch.cat([h, s1d], dim=1), temb) h = self.dec2_up(h) h = self.dec1_r1(torch.cat([h, s11], dim=1), temb); h = self.dec1_a1(h) h = self.dec1_r2(torch.cat([h, s10], dim=1), temb); h = self.dec1_a2(h) h = self.dec1_r3(torch.cat([h, s0d], dim=1), temb) h = self.dec1_up(h) h = self.dec0_r1(torch.cat([h, s01], dim=1), temb) h = self.dec0_r2(torch.cat([h, s00], dim=1), temb) return self.out(h) class Diffusion(nn.Module): def __init__(self): super().__init__() self.time_embedding = TimeEmbedding(TIME_DIM, TIME_EMB_DIM) self.unet = UNet() def forward(self, x, t): temb = self.time_embedding(t) return self.unet(x, temb) # Setup Models (Cached for Gradio) print("Loading VAE decoder...") vae = AutoencoderKL.from_pretrained( "CompVis/stable-diffusion-v1-4", subfolder="vae" ).to(device).eval() for p in vae.parameters(): p.requires_grad_(False) print("Loading Diffusion model skeleton...") model = Diffusion().to(device) # Safely fetch weights dynamically from Hugging Face Hub Model Repo try: print(f"Downloading checkpoint '{CKPT_FILENAME}' from repository '{HF_MODEL_REPO_ID}'...") downloaded_ckpt_path = hf_hub_download( repo_id=HF_MODEL_REPO_ID, filename=CKPT_FILENAME ) checkpoint = torch.load(downloaded_ckpt_path, map_location=device, weights_only=True) model.load_state_dict(checkpoint['model']) print("Checkpoint loaded successfully into U-Net!") except Exception as e: print(f"WARNING: Could not pull or load checkpoint from HF Hub. Error details: {e}") print("Using random/untrained weights for structure verification instead.") model.eval() # Adaptive precision context depending on CPU vs GPU runtime availability autocast_context = torch.amp.autocast('cuda') if device.type == 'cuda' else nullcontext() # Inference Functions (Optimized with Fast DDIM Sampler) @torch.no_grad() def decode_latents(latents): latents = latents.to(device) / LATENT_SCALE with autocast_context: imgs = vae.decode(latents).sample return (imgs / 2 + 0.5).clamp(0, 1) @torch.no_grad() def generate_single_image(steps=50, progress=gr.Progress()): """ Generates an image using a fast, deterministic DDIM sampling strategy. This reduces the required iterations from 1,000 down to 25-50 steps. """ x = torch.randn(1, CHANNELS, IMG_SIZE, IMG_SIZE, device=device) # Subsample 1000 steps evenly into your chosen number of steps (e.g., 50) timesteps = torch.linspace(0, T - 1, steps, dtype=torch.long, device=device) timesteps = list(reversed(timesteps.tolist())) for i in progress.tqdm(range(len(timesteps)), desc=f"Running DDIM ({steps} steps)", unit="steps"): t_val = timesteps[i] t_batch = torch.full((1,), t_val, device=device, dtype=torch.long) with autocast_context: eps = model(x, t_batch) # Get alpha_bar for the current step alpha_bar_curr = alpha_bars[t_val] # Get alpha_bar for the next step down in our sequence if i + 1 < len(timesteps): alpha_bar_prev = alpha_bars[timesteps[i + 1]] else: alpha_bar_prev = torch.tensor(1.0, device=device) # Fully denoised boundary # Deterministic DDIM step calculation (sigma = 0) pred_x0 = (x - torch.sqrt(1.0 - alpha_bar_curr) * eps) / torch.sqrt(alpha_bar_curr) dir_xt = torch.sqrt(1.0 - alpha_bar_prev) * eps x = torch.sqrt(alpha_bar_prev) * pred_x0 + dir_xt final_image_tensor = decode_latents(x) # Process raw tensor directly into structured PIL matrix for UI rendering final_image_tensor = final_image_tensor.squeeze(0).cpu() to_pil = transforms.ToPILImage() return to_pil(final_image_tensor) # Gradio Interface Implementation with gr.Blocks(title="Face Latent Diffusion Interface") as demo: gr.Markdown("# Face Latents Diffusion Pipeline (Optimized)") gr.Markdown("Using a fast DDIM sampling architecture to skip steps and generate images significantly faster.") with gr.Row(): with gr.Column(): # Added a step selection slider to let you balance speed vs quality steps_slider = gr.Slider( minimum=10, maximum=100, value=50, step=5, label="Inference Steps (Lower = Faster)" ) gen_btn = gr.Button("Generate Sample", variant="primary") with gr.Column(): output_img = gr.Image(label="Decoded Latent Space Result") gen_btn.click( fn=generate_single_image, inputs=[steps_slider], outputs=[output_img] ) if __name__ == "__main__": demo.launch()