wkea/blockdiffusion-api
0
1from __future__ import annotations2 3import math4from dataclasses import dataclass5from typing import Iterable, List, Sequence, Tuple6 7import torch8from torch import nn9from torch.nn import functional as F10 11 12def _num_groups(channels: int, max_groups: int = 32) -> int:13 """14 GroupNorm 组数必须能整除 channels;这里自动选择合适组数。15 """16 for g in range(min(max_groups, channels), 0, -1):17 if channels % g == 0:18 return g19 return 120 21 22class SinusoidalTimeEmbedding(nn.Module):23 """24 将离散 timestep 映射到连续 embedding(sin/cos)。25 """26 27 def __init__(self, dim: int):28 super().__init__()29 self.dim = dim30 31 def forward(self, t: torch.Tensor) -> torch.Tensor:32 half = self.dim // 233 # 频率:exp(-log(10000) * i/(half-1))34 freqs = torch.exp(35 -math.log(10000.0) * torch.arange(0, half, device=t.device, dtype=torch.float32) / max(half - 1, 1)36 )37 args = t.float().unsqueeze(1) * freqs.unsqueeze(0)38 emb = torch.cat([torch.sin(args), torch.cos(args)], dim=1)39 if self.dim % 2 == 1:40 emb = F.pad(emb, (0, 1))41 return emb42 43 44class ResBlock(nn.Module):45 """46 ResNet block(带 time embedding 注入),避免对 16×16 不友好的过深/过度下采样。47 """48 49 def __init__(self, in_ch: int, out_ch: int, time_emb_dim: int, dropout: float = 0.0):50 super().__init__()51 self.in_ch = in_ch52 self.out_ch = out_ch53 54 self.norm1 = nn.GroupNorm(_num_groups(in_ch), in_ch)55 self.conv1 = nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1)56 57 self.time_proj = nn.Sequential(58 nn.SiLU(),59 nn.Linear(time_emb_dim, out_ch),60 )61 62 self.norm2 = nn.GroupNorm(_num_groups(out_ch), out_ch)63 self.dropout = nn.Dropout(dropout) if dropout > 0 else nn.Identity()64 self.conv2 = nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1)65 66 self.skip = nn.Identity() if in_ch == out_ch else nn.Conv2d(in_ch, out_ch, kernel_size=1)67 68 def forward(self, x: torch.Tensor, t_emb: torch.Tensor) -> torch.Tensor:69 h = self.conv1(F.silu(self.norm1(x)))70 # time embedding 以 bias 的形式注入到 feature map71 h = h + self.time_proj(t_emb).unsqueeze(-1).unsqueeze(-1)72 h = self.conv2(self.dropout(F.silu(self.norm2(h))))73 return h + self.skip(x)74 75 76class AttentionBlock(nn.Module):77 """78 轻量自注意力(用于 8×8 / 4×4 等小分辨率特征图)。79 """80 81 def __init__(self, channels: int, num_heads: int = 4):82 super().__init__()83 if channels % num_heads != 0:84 # 保守处理:无法整除就退化为 1 head,避免运行时报错85 num_heads = 186 self.channels = channels87 self.num_heads = num_heads88 self.head_dim = channels // num_heads89 90 self.norm = nn.GroupNorm(_num_groups(channels), channels)91 self.qkv = nn.Conv2d(channels, channels * 3, kernel_size=1)92 self.proj = nn.Conv2d(channels, channels, kernel_size=1)93 94 def forward(self, x: torch.Tensor) -> torch.Tensor:95 b, c, h, w = x.shape96 n = h * w97 qkv = self.qkv(self.norm(x))98 q, k, v = qkv.chunk(3, dim=1)99 100 # (b, heads, head_dim, n)101 q = q.reshape(b, self.num_heads, self.head_dim, n)102 k = k.reshape(b, self.num_heads, self.head_dim, n)103 v = v.reshape(b, self.num_heads, self.head_dim, n)104 105 # 注意力:softmax(q^T k / sqrt(d))106 scale = 1.0 / math.sqrt(self.head_dim)107 attn = torch.einsum("bhdn,bhdm->bhnm", q * scale, k) # (b, heads, n, n)108 attn = torch.softmax(attn, dim=-1)109 out = torch.einsum("bhnm,bhdm->bhdn", attn, v) # (b, heads, head_dim, n)110 out = out.reshape(b, c, h, w)111 out = self.proj(out)112 return x + out113 114 115class Downsample(nn.Module):116 """117 stride=2 卷积下采样(对小分辨率更友好、可学习)。118 """119 120 def __init__(self, channels: int):121 super().__init__()122 self.conv = nn.Conv2d(channels, channels, kernel_size=3, stride=2, padding=1)123 124 def forward(self, x: torch.Tensor) -> torch.Tensor:125 return self.conv(x)126 127 128class Upsample(nn.Module):129 """130 最近邻上采样 + 卷积。131 """132 133 def __init__(self, channels: int):134 super().__init__()135 self.conv = nn.Conv2d(channels, channels, kernel_size=3, padding=1)136 137 def forward(self, x: torch.Tensor) -> torch.Tensor:138 x = F.interpolate(x, scale_factor=2.0, mode="nearest")139 return self.conv(x)140 141 142@dataclass(frozen=True)143class UNet16Config:144 """145 16×16 专用 UNet 配置。146 """147 148 in_channels: int = 3149 out_channels: int = 3 # 预测 eps,通道数与输入一致150 base_channels: int = 64151 channel_mults: Tuple[int, ...] = (1, 2, 2) # 16->8->4,避免过度下采样152 num_res_blocks: int = 2153 attn_resolutions: Tuple[int, ...] = (8, 4) # 小分辨率加注意力更划算154 dropout: float = 0.0155 time_emb_dim: int = 256156 # 文本条件向量维度(与 time_emb_dim 对齐最简单)157 cond_dim: int = 256158 159 160class UNet16(nn.Module):161 """162 小型 UNet(epsilon-predictor),适配 16×16 输入。163 """164 165 def __init__(self, cfg: UNet16Config):166 super().__init__()167 self.cfg = cfg168 169 self.time_embed = nn.Sequential(170 SinusoidalTimeEmbedding(cfg.time_emb_dim),171 nn.Linear(cfg.time_emb_dim, cfg.time_emb_dim * 4),172 nn.SiLU(),173 nn.Linear(cfg.time_emb_dim * 4, cfg.time_emb_dim),174 )175 self.cond_proj = nn.Identity() if cfg.cond_dim == cfg.time_emb_dim else nn.Linear(cfg.cond_dim, cfg.time_emb_dim)176 177 self.in_conv = nn.Conv2d(cfg.in_channels, cfg.base_channels, kernel_size=3, padding=1)178 179 # Down path(按 level 显式组织,避免小分辨率下采样/上采样位置出错)180 ch = cfg.base_channels181 self.down_resblocks = nn.ModuleList()182 self.down_attn = nn.ModuleList()183 self.downsamples = nn.ModuleList()184 185 skip_chs: List[int] = []186 resolution = 16 # 纹理固定为 16×16187 for level, mult in enumerate(cfg.channel_mults):188 out_ch = cfg.base_channels * mult189 level_res = nn.ModuleList()190 level_attn = nn.ModuleList()191 for _ in range(cfg.num_res_blocks):192 level_res.append(ResBlock(ch, out_ch, cfg.time_emb_dim, dropout=cfg.dropout))193 ch = out_ch194 level_attn.append(AttentionBlock(ch) if resolution in cfg.attn_resolutions else nn.Identity())195 skip_chs.append(ch)196 self.down_resblocks.append(level_res)197 self.down_attn.append(level_attn)198 if level != len(cfg.channel_mults) - 1:199 self.downsamples.append(Downsample(ch))200 resolution //= 2201 202 # Mid203 self.mid_block1 = ResBlock(ch, ch, cfg.time_emb_dim, dropout=cfg.dropout)204 self.mid_attn = AttentionBlock(ch)205 self.mid_block2 = ResBlock(ch, ch, cfg.time_emb_dim, dropout=cfg.dropout)206 207 # Up path(从最深层开始逐级上采样)208 self.up_resblocks = nn.ModuleList()209 self.up_attn = nn.ModuleList()210 self.upsamples = nn.ModuleList()211 212 # 当前 resolution 对应最深层(例如 4×4)213 # down 部分循环结束后,resolution 已被更新到最深层214 for level, mult in reversed(list(enumerate(cfg.channel_mults))):215 out_ch = cfg.base_channels * mult216 level_res = nn.ModuleList()217 level_attn = nn.ModuleList()218 for _ in range(cfg.num_res_blocks):219 skip_ch = skip_chs.pop()220 level_res.append(ResBlock(ch + skip_ch, out_ch, cfg.time_emb_dim, dropout=cfg.dropout))221 ch = out_ch222 level_attn.append(AttentionBlock(ch) if resolution in cfg.attn_resolutions else nn.Identity())223 self.up_resblocks.append(level_res)224 self.up_attn.append(level_attn)225 if level != 0:226 self.upsamples.append(Upsample(ch))227 resolution *= 2228 229 self.out_norm = nn.GroupNorm(_num_groups(ch), ch)230 self.out_conv = nn.Conv2d(ch, cfg.out_channels, kernel_size=3, padding=1)231 232 def forward(self, x: torch.Tensor, t: torch.Tensor, cond: torch.Tensor | None = None) -> torch.Tensor:233 """234 输入:235 - x: (B, C, 16, 16) 的噪声图236 - t: (B,) 的 timestep(int64)237 - cond: (B, cond_dim) 的条件向量(可选;为空则无条件)238 输出:239 - eps_pred: (B, C, 16, 16)240 """241 t_emb = self.time_embed(t)242 if cond is not None:243 t_emb = t_emb + self.cond_proj(cond)244 h = self.in_conv(x)245 246 skips: List[torch.Tensor] = []247 for level in range(len(self.cfg.channel_mults)):248 for i in range(self.cfg.num_res_blocks):249 h = self.down_resblocks[level][i](h, t_emb)250 h = self.down_attn[level][i](h)251 skips.append(h)252 if level < len(self.downsamples):253 h = self.downsamples[level](h)254 255 h = self.mid_block1(h, t_emb)256 h = self.mid_attn(h)257 h = self.mid_block2(h, t_emb)258 259 # up_resblocks/up_attn 的顺序是从最深层到最浅层260 for level in range(len(self.cfg.channel_mults)):261 for i in range(self.cfg.num_res_blocks):262 skip = skips.pop()263 h = torch.cat([h, skip], dim=1)264 h = self.up_resblocks[level][i](h, t_emb)265 h = self.up_attn[level][i](h)266 if level < len(self.upsamples):267 h = self.upsamples[level](h)268 269 h = self.out_conv(F.silu(self.out_norm(h)))270 return h271 272 273 