keysun89's picture
Update app.py
e4793f5 verified
Raw
History Blame Contribute Delete
12.8 kB
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()