Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
README.md74 linesDownload Raw Back to root
1---2title: BlockDiffusion API3sdk: gradio4app_file: app.py5---6 7## BlockDiffusion8 9目标:训练一个专用于生成 **16×16《我的世界》纹理** 的小型 Diffusion(UNet)模型(PyTorch),支持 **提示词条件生成** 与 **RGBA 透明通道**。10 11### Hugging Face Spaces(在线API + 伪私有)12- **Space入口文件**:`app.py`13- **伪私有方式**:在 Space 的 `Settings -> Secrets` 配置 `SPACE_API_KEY`(若网页端不允许下划线则用 `SPACEAPIKEY`),所有推理请求必须携带该 Key 才会执行。14- **checkpoint来源(二选一)**:15  - 配置 `BD_CKPT_PATH`(Space仓库内的文件路径)16  - 或配置 `BD_CKPT_REPO` + `BD_CKPT_FILENAME`(从 Hugging Face Hub 下载;私有仓库可再配 `HF_TOKEN`)17 18### 当前能力(已实现)19- **16×16 专用 UNet 扩散**:避免过深/过度下采样(16→8→4)20- **噪声调度**:linear / cosine21- **loss**:基础为 MSE,支持可选 SNR 加权(提升稳定性/质量)22- **训练/采样解耦**:Diffusion 公式与采样逻辑独立于训练器23- **文本条件(提示词)**:24  - 支持 `captions.jsonl/csv`(图文配对)25  - 支持“从文件名自动提取提示词”的弱标注训练(见 `train_oneclick.py`)26  - 支持 CFG(训练丢条件 + 采样 scale)27- **RGB/RGBA**:`channels=3/4`,RGBA 输出 PNG 透明通道28 29### 主要入口30- **一键训练(推荐)**:`train_oneclick.py`31  - **无需命令行参数**,只修改脚本顶部常量32  - 默认从 `DATA_DIR` 扫描图片,并按文件名生成提示词训练33  - 默认支持 **自动续训**:扫描 `outputs/<RUN_NAME>/checkpoints/step_*.pt` 选择最大 step 继续34- **可参数化训练**:`train.py`(保留)35- **采样**:`sample.py`(支持 `--prompt` + CFG)36- **WebUI 采样预览**:`webui.py`(浏览器里快速对比 checkpoint)37 38### 数据与提示词规则39#### 方式 A:文件名弱标注(`train_oneclick.py`)40适配命名:`HASH-tag_tag_tag.png`、`HASH-side.png`、`HASH-top.png` 等。41 42提取规则(实现于 `blockdiffusion/data/filename_captions.py`):43- 去扩展名44- 取第一个 `-` 后的部分45- `_`/`-` 视为分隔符(转空格),转小写46 47#### 方式 B:显式 captions(`train.py --captions`)48- `captions.jsonl`:每行 `{"file":"相对data_dir的路径或绝对路径","text":"提示词"}`49- `captions.csv`:两列 `file,text`(含表头)50 51### 输出与预览52- checkpoints:`outputs/<RUN_NAME>/checkpoints/step_XXXX.pt`53- 训练中采样图:`outputs/<RUN_NAME>/samples/step_XXXX.png`54- 采样脚本输出:默认 `samples/samples.png`55 56### 核心模块导航(给接手者)57- **模型文档**:见 `blockdiffusion/models/README.md`(只讲模型:结构/输入输出/条件/ckpt 格式)58- `blockdiffusion/models/unet16.py`:16×16 UNet(支持条件向量注入)59- `blockdiffusion/models/text_unet16.py`:TextCondUNet16(tokens/mask→cond→UNet)60- `blockdiffusion/text/simple_tokenizer.py`:轻量 tokenizer/vocab(可随 checkpoint 保存)61- `blockdiffusion/text/text_encoder.py`:轻量 TextEncoder(masked mean pooling + MLP)62- `blockdiffusion/diffusion/schedules.py`:beta schedule(linear/cosine)63- `blockdiffusion/diffusion/gaussian_diffusion.py`:训练损失 + DDPM/DDIM 采样 + CFG 训练丢条件64- `blockdiffusion/train/trainer.py`:AMP/EMA/checkpoint/采样、梯度累积、warmup+余弦、channels_last65- `blockdiffusion/utils/ema.py`:EMA(续训时自动对齐 device)66- `blockdiffusion/utils/image_io.py`:保存 RGB/RGBA 网格 PNG67 68### 常用提速/提质开关(修改 `train_oneclick.py` 顶部常量)69- **提速**:`NUM_WORKERS`、`PREFETCH_FACTOR`、`PERSISTENT_WORKERS`、`CHANNELS_LAST`70- **显存/吞吐**:`BATCH_SIZE`、`GRAD_ACCUM_STEPS`71- **质量/稳定**:`SNR_GAMMA`、`LR_WARMUP_STEPS`、`LR_MIN_RATIO`、`COND_DROP_PROB`、`SAMPLE_CFG_SCALE`72 73 74