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
9from blockdiffusion.models.unet16 import UNet16, UNet16Config
10from blockdiffusion.text.simple_tokenizer import SimpleVocab
11from blockdiffusion.text.text_encoder import TextEncoder, TextEncoderConfig
12
13
14@dataclass(frozen=True)
15class TextCondConfig:
16 """
17 文本条件配置(用于 checkpoint/meta 保存)。
18 """
19
20 max_text_len: int = 32
21 emb_dim: int = 128
22 dropout: float = 0.0
23
24
25class TextCondUNet16(nn.Module):
26 """
27 文本条件 UNet:
28 - TextEncoder 将 tokens/mask 编码为 cond 向量
29 - cond 向量注入到 UNet 的 time embedding(加和)
30 - tokens/mask 为空或 mask 全 0 时等价于无条件
31 """
32
33 def __init__(
34 self,
35 unet_cfg: UNet16Config,
36 vocab: SimpleVocab,
37 text_cfg: TextCondConfig,
38 ):
39 super().__init__()
40 self.unet_cfg = unet_cfg
41 self.text_cfg = text_cfg
42
43 self.vocab = vocab # 仅用于保存 meta;不注册为 buffer/module
44
45 te_cfg = TextEncoderConfig(
46 vocab_size=vocab.size,
47 max_len=text_cfg.max_text_len,
48 emb_dim=text_cfg.emb_dim,
49 out_dim=unet_cfg.time_emb_dim,
50 dropout=text_cfg.dropout,
51 )
52 self.text_encoder = TextEncoder(te_cfg, pad_id=vocab.pad_id)
53 self.unet = UNet16(unet_cfg)
54
55 def forward(
56 self,
57 x: torch.Tensor,
58 t: torch.Tensor,
59 tokens: Optional[torch.Tensor] = None,
60 mask: Optional[torch.Tensor] = None,
61 ) -> torch.Tensor:
62 cond = None
63 if tokens is not None:
64 cond = self.text_encoder(tokens, mask=mask)
65 return self.unet(x, t, cond=cond)
66
67
68 