Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
text_encoder.py65 linesDownload Raw Back to text
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