Team Ai
Apppublic

wkea/blockdiffusion-api

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