Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
README.md144 linesDownload Raw Back to models
1## BlockDiffusion 模型说明(Model README)
2
3本文件只描述**模型本身**(网络结构、输入输出、条件机制、checkpoint 格式),不描述训练脚本的使用方式。
4
5### 1. 模型定位
6
7- **任务**:生成/还原 **16×16** 的《我的世界》纹理(支持 **RGB 或 RGBA**)。
8- **扩散建模**:epsilon-predictor(预测噪声 \(\epsilon\)),网络输出通道数与输入一致。
9- **条件方式**:可选的**文本提示词条件**(Text Conditioning),并支持 CFG(Classifier-Free Guidance)的“条件丢弃训练 + 采样引导”范式。
10
11### 2. 模型组件与代码位置
12
13- **UNet 主体**:`blockdiffusion/models/unet16.py`
14  - `UNet16Config`:UNet 配置(通道、层数、注意力分辨率等)
15  - `UNet16`:16×16 专用 UNet(epsilon predictor)
16- **文本条件封装**:`blockdiffusion/models/text_unet16.py`
17  - `TextCondConfig`:文本侧配置(长度、embedding 维度等)
18  - `TextCondUNet16`:TextEncoder + UNet16 的组合模型
19- **文本编码器**:`blockdiffusion/text/text_encoder.py`
20  - `TextEncoder`:masked mean pooling + MLP 投影到条件向量
21- **Tokenizer / 词表**:`blockdiffusion/text/simple_tokenizer.py`
22  - `tokenize()`:轻量分词(支持 CJK 单字切分)
23  - `SimpleVocab`:极简词表(含 `<pad>/<unk>`),可序列化到 checkpoint
24- **CFG 采样包装器**:`blockdiffusion/train/guidance.py`
25  - `CFGuidanceModel`:把任意支持 `forward(x, t, **kwargs)` 的模型包装成 CFG 采样模型
26
27### 3. 前向接口(输入/输出约定)
28
29#### 3.1 `UNet16`
30
31- **函数签名(语义)**:`eps = UNet16(x, t, cond=None)`
32- **输入**
33  - `x`:形状 `(B, C, 16, 16)`,取值范围约定为 **[-1, 1]**
34  - `t`:形状 `(B,)`,dtype 通常为 `int64`,表示扩散 timestep
35  - `cond`(可选):形状 `(B, cond_dim)` 的条件向量;`None` 表示无条件
36- **输出**
37  - `eps`:形状 `(B, C, 16, 16)`,预测噪声 \(\epsilon\)
38
39#### 3.2 `TextCondUNet16`
40
41- **函数签名(语义)**:`eps = TextCondUNet16(x, t, tokens=None, mask=None)`
42- **输入**
43  - `x` / `t`:同 `UNet16`
44  - `tokens`(可选):形状 `(B, L)`,dtype 通常为 `int64`
45  - `mask`(可选):形状 `(B, L)`,bool 或 0/1
46    - `mask=None` 时会自动由 `tokens != pad_id` 推导
47    - 当 `tokens=None` 或 `mask` 全 0 时,模型**等价于无条件**
48- **输出**
49  - `eps`:同 `UNet16`
50
51### 4. 条件注入机制(Text Conditioning)
52
53- `TextEncoder` 把 `tokens/mask` 编码为 `cond`(形状 `(B, time_emb_dim)`)。
54- `cond` 通过“**加和**”注入到 UNet 的 time embedding:
55  - `t_emb = time_embed(t)`
56  - `t_emb = t_emb + cond_proj(cond)`(无 `cond` 时不加)
57- 这是一种轻量、易训练、可随 checkpoint 保存的条件方案,不依赖外部大型文本模型。
58
59### 5. CFG(Classifier-Free Guidance)说明
60
61- **训练侧**:通过一定概率把条件 `mask` 置空(等价无条件),让模型同时学到 cond/uncond。
62- **采样侧**:`CFGuidanceModel` 组合两次前向,按以下公式合成:
63
64\[
65\epsilon = \epsilon_{\text{uncond}} + s \cdot (\epsilon_{\text{cond}} - \epsilon_{\text{uncond}})
66\]
67
68其中 \(s\) 为 `cfg_scale`。
69
70### 6. 关键配置项(建议对照 checkpoint 的 meta)
71
72#### 6.1 `UNet16Config`
73
74| 字段 | 含义 | 备注 |
75|---|---|---|
76| `in_channels` | 输入通道数 | 3=RGB,4=RGBA |
77| `out_channels` | 输出通道数 | epsilon 预测,通常与输入一致 |
78| `base_channels` | 基础通道宽度 | 模型容量主开关 |
79| `channel_mults` | 各层通道倍率 | 默认 `(1,2,2)` 对应 16→8→4 |
80| `num_res_blocks` | 每层 ResBlock 数 | 16×16 不建议太大 |
81| `attn_resolutions` | 启用自注意力的分辨率 | 默认在 8/4 上更划算 |
82| `dropout` | ResBlock dropout | 0 表示关闭 |
83| `time_emb_dim` | time embedding 维度 | 也作为条件注入维度基准 |
84| `cond_dim` | 条件向量维度 | 通过 `cond_proj` 对齐到 `time_emb_dim` |
85
86#### 6.2 `TextCondConfig`
87
88| 字段 | 含义 | 备注 |
89|---|---|---|
90| `max_text_len` | 最大 token 长度 \(L\) | 影响 `tokens` 形状与 TextEncoder pooling |
91| `emb_dim` | token embedding 维度 | TextEncoder 内部使用 |
92| `dropout` | TextEncoder dropout | 0 表示关闭 |
93
94> `TextEncoderConfig.out_dim` 会被设置为 `UNet16Config.time_emb_dim`,从而保证条件向量可直接注入 UNet。
95
96### 7. Tokenizer 与词表(`SimpleVocab`)
97
98- **Tokenizer 特性**
99  - 统一转小写
100  - 英文/数字按“词”切分(把 `_`/`-` 视作词内字符)
101  - CJK(中日韩)字符按“单字”切分,适配无空格中文提示词
102- **词表约定**
103  - `id=0`:`<pad>`
104  - `id=1`:`<unk>`
105- **重要限制(续训/推理一致性)**
106  - 续训或推理时必须复用 checkpoint 内保存的词表;否则 `Embedding` 尺寸会不匹配。
107  - 新出现但不在旧词表中的词会被编码成 `<unk>`:可继续运行,但无法学习/表达新词语义。
108
109### 8. Checkpoint(.pt)格式与推荐加载流程
110
111#### 8.1 文件内容(关键键)
112
113训练保存的 checkpoint 是一个 `dict`,关键字段如下:
114
115- `model`:模型 `state_dict`
116- `ema`:EMA `state_dict`(若存在,采样建议使用)
117- `step`:保存时的全局 step
118- `meta`:**重建模型所需的结构化信息**(强烈依赖)
119- `trainer_cfg` / `optimizer` / `scaler`:与训练相关(推理通常不需要)
120
121#### 8.2 `meta` 约定(常用字段)
122
123- `model_type`:`"text_cond"` 或(缺省/其他)无条件
124- `unet_cfg`:`UNet16Config` 的字典形式(`dataclasses.asdict`)
125- `text_cfg`:`TextCondConfig` 的字典形式(仅 text_cond)
126- `tokenizer`:`SimpleVocab.to_dict()` 的结果(仅 text_cond)
127- `diffusion`:`{"timesteps": ..., "beta_schedule": ...}`(采样需一致)
128- `data`:`{"image_size": 16, "channels": 3/4}`(采样 shape 需一致)
129
130#### 8.3 推荐加载流程(文字版)
131
132- **步骤**
133  - 读取 `.pt` 到内存(`torch.load`)
134  - 从 `meta` 判断 `model_type`
135  - 由 `meta["unet_cfg"]` 重建 `UNet16Config`
136  - 若为 `text_cond`:从 `meta["tokenizer"]` 重建 `SimpleVocab`,从 `meta["text_cfg"]` 重建 `TextCondConfig`
137  - 实例化 `UNet16` 或 `TextCondUNet16`
138  - `load_state_dict(ckpt["model"], strict=True)`
139  - 如需更稳的采样:把 `ckpt["ema"]` 应用到模型参数
140
141> 兼容性提示:若旧 checkpoint 不含 `meta/unet_cfg`,项目内的采样入口会回落到默认的 16×16 RGB 无条件配置;但这会丢失“通道数/条件/调度”等信息,不建议用于长期维护。
142
143
144