Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
unet16.py273 linesDownload Raw Back to models
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