Anime_Face / model.py
adityachaubey's picture
Refactor SinusoidalPositionEmbeddings class and update related methods for clarity and consistency
f7f03ea
Raw
History Blame Contribute Delete
7.37 kB
import torch
import torch.nn as nn
import math
#Varicence scheduler
T = 1000
beta_start = 1e-4
beta_end = 0.02
betas = torch.linspace(beta_start, beta_end, steps=T)
alphas = 1. - betas
alphas_cumprod = torch.cumprod(alphas, dim=0)
#SinusoidalPositionEmbedding
class SinusoidalPositionEmbeddings(torch.nn.Module):
def __init__(self, dim:int):
super().__init__()
assert dim % 2 == 0, "Embedding must divisble by 2"
self.dim = dim
def forward(self, time_stamps: torch.Tensor) -> torch.Tensor :
device = time_stamps.device
half_dim = self.dim // 2
scale = math.log(10000)/ (half_dim - 1)
freqs= torch.exp(torch.arange(half_dim, dtype=torch.float32, device=device) * -scale)
embeddings = time_stamps[:, None].float() * freqs[None, :]
embeddings = torch.cat((embeddings.sin(), embeddings.cos()), dim=-1)
return embeddings
#Time Embeddings
class Time_embeddings(torch.nn.Module):
def __init__(self, time_emb_dim: int, out_dim: int):
super().__init__()
self.sinusoidal_emb = SinusoidalPositionEmbeddings(time_emb_dim)
self.mlp = nn.Sequential(
nn.Linear(time_emb_dim, out_dim),
nn.SiLU(),
nn.Linear(out_dim, out_dim)
)
def forward (self, time_steps: torch.Tensor) -> torch.Tensor :
raw_emb = self.sinusoidal_emb(time_steps)
return self.mlp(raw_emb)
#Resblock
class ResBlock(torch.nn.Module):
def __init__(self, in_channels, out_channels, time_emb_dim):
super().__init__()
self.conv1 = nn.Conv2d(in_channels, out_channels, 3, padding=1)
self.conv2 = nn.Conv2d(out_channels, out_channels, 3, padding=1)
self.norm1 = nn.GroupNorm(8, out_channels)
self.norm2 = nn.GroupNorm(8, out_channels)
self.act = nn.SiLU()
self.time_mlp = nn.Linear(time_emb_dim, out_channels)
#Matches in_channels and out_channels for skip connections
self.residual_conv = nn.Conv2d(in_channels, out_channels, 1) if in_channels != out_channels else nn.Identity()
def forward(self, x, time_emb):
h = self.norm1(self.conv1(x))
h = self.act(h)
# inject time embedding β€” shape (B, out_channels) β†’ (B, out_channels, 1, 1)
time_proj = self.act(self.time_mlp(time_emb))[:, :, None, None]
h = h + time_proj
h = self.norm2(self.conv2(h))
h = self.act(h)
return h + self.residual_conv(x)
#Downsmaple and upsmaple
class Down_sample(torch.nn.Module):
def __init__(self, channels):
super().__init__()
self.conv = nn.Conv2d(channels, channels, 4, stride=2, padding=1)
def forward(self, x):
return self.conv(x)
class Up_sample(torch.nn.Module):
def __init__(self, channels):
super().__init__()
self.conv = nn.ConvTranspose2d(channels, channels, 4, stride=2, padding=1)
def forward(self, x):
return self.conv(x)
#Assembling Unet class
class Unet(nn.Module):
def __init__(self, in_channels=3, time_emb_dim=128):
super().__init__()
# time embedding outputs 256-dim vector
self.time_embeddings = Time_embeddings(time_emb_dim, out_dim=256)
self.inc = nn.Conv2d(in_channels, 64, 3, padding=1)
# Encoder
self.down1 = ResBlock(64, 128, time_emb_dim=256)
self.down_sample1 = Down_sample(128)
self.down2 = ResBlock(128, 256, time_emb_dim=256)
self.down_sample2 = Down_sample(256)
# Bottleneck
self.bot1 = ResBlock(256, 256, time_emb_dim=256)
self.bot2 = ResBlock(256, 256, time_emb_dim=256)
# Decoder (input channels doubled due to skip connection concat)
self.up_sample1 = Up_sample(256)
self.up1 = ResBlock(512, 128, time_emb_dim=256) # 256 + 256 = 512
self.up_sample2 = Up_sample(128)
self.up2 = ResBlock(256, 64, time_emb_dim=256) # 128 + 128 = 256
self.outc = nn.Conv2d(64, in_channels, 3, padding=1)
def forward(self, x, t):
t_emb = self.time_embeddings(t) # (B,) β†’ (B, 256)
# Encoder
x1 = self.inc(x) # (B, 64, 64, 64)
x2 = self.down1(x1, t_emb) # (B, 128, 64, 64)
x3 = self.down_sample1(x2) # (B, 128, 32, 32)
x4 = self.down2(x3, t_emb) # (B, 256, 32, 32)
x5 = self.down_sample2(x4) # (B, 256, 16, 16)
# Bottleneck
x5 = self.bot1(x5, t_emb) # (B, 256, 16, 16)
x5 = self.bot2(x5, t_emb) # (B, 256, 16, 16)
# Decoder
x6 = self.up_sample1(x5) # (B, 256, 32, 32)
x6 = torch.cat([x6, x4], dim=1) # (B, 512, 32, 32) β€” skip from x4
x6 = self.up1(x6, t_emb) # (B, 128, 32, 32)
x7 = self.up_sample2(x6) # (B, 128, 64, 64)
x7 = torch.cat([x7, x2], dim=1) # (B, 256, 64, 64) β€” skip from x2
x7 = self.up2(x7, t_emb) # (B, 64, 64, 64)
return self.outc(x7) # (B, 3, 64, 64)
#EMA weights
class EMA():
def __init__(self ,model ,decay=0.9999):
self.model = model
self.decay = decay
self.shadow = {}
self.backup = {}
def register(self):
for name, param in self.model.named_parameters():
if param.requires_grad:
self.shadow[name] = param.data.clone()
def update(self):
for name , param in self.model.named_parameters():
if param.requires_grad:
assert name in self.shadow
new_avg = self.decay*self.shadow[name] + (1. - self.decay)*param.data
self.shadow[name] = new_avg.clone()
def apply_shadow(self):
for name , param in self.model.named_parameters():
if param.requires_grad:
assert name in self.shadow
self.backup[name] = param.data.clone()
param.data.copy_(self.shadow[name])
def restore(self):
for name, param in self.model.named_parameters():
if param.requires_grad:
assert name in self.backup
param.data.copy_(self.backup[name])
self.backup = {}
#DDIM
@torch.no_grad()
def sample_ddim(model, T,img_size, batch_size=16, channels=3,ddim_step=50, device='cpu'):
model.eval()
alphas_cumprod_device = alphas_cumprod.to(device)
timestep = torch.linspace(0, T-1, steps=ddim_step, dtype=torch.long, device=device)
img = torch.randn(batch_size, channels, img_size, img_size).to(device)
for i in reversed(range(len(timestep))):
t_current =timestep[i]
t_preq = timestep[i-1] if i > 0 else -1
t_tensor = torch.full((batch_size,), t_current, dtype=torch.long, device=device)
pred_noise = model(img , t_tensor)
alphas_bar_t = alphas_cumprod[t_current]
alphas_bar_t_preq = alphas_cumprod[t_preq] if t_preq >= 0 else torch.tensor(1.0, device=device)
pred_x0 = (img - torch.sqrt(1 - alphas_bar_t) * pred_noise)/(torch.sqrt(alphas_bar_t))
pred_x0 = torch.clamp(pred_x0, -1.0, 1.0)
pred_dir = torch.sqrt(1 - alphas_bar_t_preq) * pred_noise
img = torch.sqrt(alphas_bar_t_preq) * pred_x0 + pred_dir
return img