wkea/blockdiffusion-api
0
1from __future__ import annotations
2
3from dataclasses import dataclass
4from typing import Optional
5
6import torch
7from torch import nn
8
9
10@dataclass(frozen=True)
11class TextEncoderConfig:
12 """
13 文本编码器配置:
14 - 目标是轻量、可训练、可随 checkpoint 一起保存,不依赖外部大模型。
15 """
16
17 vocab_size: int
18 max_len: int = 32
19 emb_dim: int = 128
20 out_dim: int = 256
21 dropout: float = 0.0
22
23
24class TextEncoder(nn.Module):
25 """
26 极简 TextEncoder:
27 - token embedding
28 - masked mean pooling
29 - MLP 投影到 out_dim(通常与 UNet time_emb_dim 对齐)
30 """
31
32 def __init__(self, cfg: TextEncoderConfig, pad_id: int = 0):
33 super().__init__()
34 self.cfg = cfg
35 self.pad_id = int(pad_id)
36
37 self.emb = nn.Embedding(cfg.vocab_size, cfg.emb_dim, padding_idx=self.pad_id)
38 self.proj = nn.Sequential(
39 nn.Dropout(cfg.dropout) if cfg.dropout > 0 else nn.Identity(),
40 nn.Linear(cfg.emb_dim, cfg.out_dim),
41 nn.SiLU(),
42 nn.Linear(cfg.out_dim, cfg.out_dim),
43 )
44
45 def forward(self, tokens: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
46 """
47 tokens: (B, L) int64
48 mask: (B, L) bool/0-1;为 None 时自动根据 pad_id 生成
49 返回:
50 - cond: (B, out_dim)
51 """
52 if mask is None:
53 mask = (tokens != self.pad_id)
54 mask_f = mask.to(dtype=self.emb.weight.dtype)
55
56 x = self.emb(tokens) # (B,L,D)
57 # masked mean
58 denom = mask_f.sum(dim=1, keepdim=True).clamp(min=1.0)
59 pooled = (x * mask_f.unsqueeze(-1)).sum(dim=1) / denom
60
61 # 若整句为空(mask 全 0),pooled 会是全 0,等价于无条件
62 return self.proj(pooled)
63
64
65 