Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

16 Commits
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

NumberGenerate

在 VAE 潜空间中,用数字条件 Flow Matching 生成 MNIST 手写数字。

预训练权重 · 快速开始 · 可视化结果 · 模型原理

NumberGenerate 是一个面向学习与实验的生成模型项目。它先用无条件 VAE 把二值 MNIST 压缩到潜空间,再训练带时间条件和数字标签条件的 Flow Matching 模型,最后通过 Euler 积分把高斯噪声变成指定数字。

MNIST [N,2,28,28]
    -> VAE Encoder
    -> latent [N,4,14,14]
    -> conditional Flow Matching
    -> VAE Decoder
    -> generated image [N,2,28,28]

数字 0 到 9 在 32 个 Euler 步骤中从噪声逐渐成形
每列对应数字 0–9;从上到下是 32 个 Euler 采样步骤。点击查看原图。

项目特点

  • 完整但紧凑:数据、模型、训练、采样、可视化和测试相互分离。
  • 潜空间生成:Flow 只在 4 × 14 × 14 的 VAE 潜变量上学习速度场。
  • 数字可控:时间嵌入与数字标签嵌入共同调制 8 个条件 ConvNeXt 块。
  • 可解释:内置生成轨迹、逐步流场、特征 PCA、VAE 重建、KL 与能量热力图。
  • 配置集中:默认参数统一放在 config/default.yaml,也支持 YAML 覆盖配置。
  • 可恢复训练:已有 checkpoint 时会恢复模型、优化器、epoch、step 和 loss。

快速开始

1. 安装环境

项目已在 Python 3.12 下验证。CPU 可以运行测试、推理和小规模训练,完整训练建议使用 CUDA。

Windows PowerShell:

git clone https://github.com/1os3/NumberGenerate.git
cd NumberGenerate
python -m venv .venv
.\.venv\Scripts\python.exe -m pip install -r requirements.txt

Linux / macOS:

git clone https://github.com/1os3/NumberGenerate.git
cd NumberGenerate
python3 -m venv .venv
./.venv/bin/python -m pip install -r requirements.txt

主要依赖包括 PyTorch、torchvision、NumPy、PyYAML、Matplotlib、scikit-learn 和 tqdm。

2. 下载预训练权重

预训练权重发布在 Model Release

文件 用途 直接下载
vae.pt VAE 编码器与解码器 下载 vae.pt
flow.pt 数字条件 Flow Matching 模型 下载 flow.pt

下载后放到项目根目录的 ckpt/

NumberGenerate/
└── ckpt/
    ├── vae.pt
    └── flow.pt

也可以在 Windows PowerShell 中直接下载:

New-Item -ItemType Directory -Force ckpt | Out-Null
Invoke-WebRequest `
  -Uri "https://github.com/1os3/NumberGenerate/releases/download/Model/vae.pt" `
  -OutFile "ckpt/vae.pt"
Invoke-WebRequest `
  -Uri "https://github.com/1os3/NumberGenerate/releases/download/Model/flow.pt" `
  -OutFile "ckpt/flow.pt"

3. 一键生成全部可视化

Windows:

.\.venv\Scripts\python.exe -m vis.visualize

Linux / macOS:

./.venv/bin/python -m vis.visualize

首次运行会按需把 MNIST 下载到 datasets/。脚本读取两个 checkpoint,并把 9 张结果图写入 outputs/

outputs/
├── generation_steps.png
├── flow_prediction_steps.png
├── flow_feature_pca_map.png
├── flow_feature_clustering.png
├── vae_reconstruction.png
├── vae_latent_pca.png
├── vae_latent_distribution.png
├── vae_kl_map.png
└── vae_latent_energy_map.png

如果只下载了 vae.pt,可以只生成 VAE 可视化:

.\.venv\Scripts\python.exe -m vis.visualize --mode vae

--only-vae--mode vae 等价。

可视化结果

下面的图片是项目输出的可提交快照;重新运行 vis.visualize 会在 outputs/ 中生成新的随机样本。

从噪声到数字

generation_steps.png 把 0–9 各生成一个样本,并记录全部 32 个历史帧。最开始是潜空间高斯噪声,随后轮廓逐渐聚合,最后由 VAE Decoder 还原为清晰的前景概率图。

MNIST 数字逐步生成轨迹

Flow 每一步在预测什么

flow_prediction_steps.png 跟踪 visual.flow_step_label 指定的数字。每个步骤上方是潜空间预测速度经共享 PCA 投影后的二维流场,下方是最后一个 Flow 主干块的特征经共享 PCA 投影后的 RGB 图。共享投影、色阶与流场幅值范围让不同步骤可以直接比较。

Flow 在 32 个采样步骤中的预测速度和末端特征

单样本末端特征 PCA 图 多时间步骨干中层 GAP 聚类中心距离
Flow 末端特征 PCA RGB 图 Flow 多时间步骨干中层 GAP 特征到聚类中心的距离散点图
  • flow_feature_pca_map.png:在 visual.feature_map_time 指定的时间点,把单个样本的末端特征压到 3 个 PCA 通道并映射为 RGB。
  • flow_feature_clustering.png:固定同一批高斯噪声 z₀ 与 VAE 后验均值 z₁=mu,严格按训练分布中的线性路径 zₜ=(1-t)z₀+t z₁ 构造 t=0、0.25、0.5、0.75、1 五组输入。每个时间步分别取 Flow 骨干中间块(默认 8 块时取第 5 块),通过 GAP 把 [256,14,14] 压成每个样本一个 256 维向量;各维标准化后,K-Means 直接在完整 256 维空间中聚成 10 类,PCA 不参与聚类或评价。

每个子图对应一个时间步。聚类编号先经 Hungarian 匹配对齐到数字 0–9;横轴是对齐后的中心编号,纵轴是样本在该时间步的标准化 256 维空间中到所属中心的欧氏距离,颜色表示真实数字,横向抖动只用于减少重叠。每一列颜色越单一,说明该簇类别纯度越高;点越靠近底部的黑色菱形中心,说明簇内越紧致。accuracy 是对齐后的整体正确率,ARI 和 NMI 是不依赖簇编号的聚类一致性指标。多时间步对比展示条件语义沿 Flow 路径的形成过程;它衡量的是 GAP 后的无监督可分性,而模型生成还会使用未池化的完整空间特征和显式标签条件,因此二者不必完全一致。

VAE 压缩与重建

vae_reconstruction.png 依次展示输入前景、潜变量 PCA 图、重建前景和绝对误差。VAE 本身不接收数字标签,只负责学习稳定的图像压缩与重建。

VAE 输入、潜变量、重建和绝对误差

VAE 潜空间 PCA 前景 / 背景潜变量统计
VAE 潜空间 PCA 散点图 VAE 潜变量均值、标准差和 KL 分布

VAE 是无条件模型,因此 PCA 中不同数字发生混合并不意外;这里更关注潜空间是否连续、数值是否稳定,以及前景和背景位置是否呈现不同统计特征。

信息集中在哪里

空间 KL 强度 潜变量能量 mean(abs(mu))
VAE 空间 KL 热力图 VAE 潜变量能量热力图

两张热力图先把输入前景下采样到 14 × 14,再与潜空间位置对齐:

  • KL 图显示每个位置偏离标准高斯先验的强度。
  • 能量图显示各位置 mu 的平均绝对值。
  • 二者与笔画区域的对应关系可以帮助判断 VAE 是否把容量用在有效结构上。

模型原理

总体流程

flowchart LR
    X["MNIST<br/>[N,2,28,28]"] --> ENC["VAE Encoder"]
    ENC --> POST["mu, logvar<br/>[N,4,14,14]"]
    POST --> Z1["posterior 样本 z₁"]
    Z1 --> DEC["VAE Decoder"]
    DEC --> IMG["重建 / 生成图像"]

    Z0["高斯噪声 z₀"] --> MIX["线性插值 zₜ"]
    Z1 --> MIX
    MIX --> FLOW["FlowModel<br/>8 个条件主干块"]
    T["连续时间 t"] --> FLOW
    Y["数字标签 y"] --> FLOW
    FLOW --> VEL["预测速度 vθ"]
    VEL --> EULER["Euler 积分 × 32"]
    EULER --> DEC
Loading

训练分为两个阶段:先训练 VAE,再冻结 VAE 训练 Flow。采样时不需要输入真实图片,只需要高斯噪声和目标数字标签。

数据表示

MNIST 灰度图先按 data.binarize_threshold 二值化,再转换为二通道 one-hot 概率图:

张量 形状 含义
images [N,2,28,28] 第 0 通道为 background,第 1 通道为 foreground
labels [N] 数字标签,取值 0–9;只传给 Flow

把背景显式表示为一个通道,可以让 VAE 同时重建“有笔画”和“无笔画”的概率,而不是把背景当成缺失值。

VAE

VAE(Variational Autoencoder,变分自编码器)可以直观地理解为一套“带随机性的压缩器 + 解压器”:

  • Encoder 不把图片压成一个固定编码,而是输出潜变量分布的均值 mu 和方差 logvar
  • 从这个分布中采样得到 z,Decoder 再尝试用 z 还原原图。
  • 重建损失要求图片还原准确,KL 损失则让潜空间接近连续、规则的高斯分布,便于后续从中采样和生成。

在本项目中,VAE 不负责区分要生成哪个数字,也不接收数字标签;它只负责把 28 × 28 图片压缩成更小的 4 × 14 × 14 潜变量,并把潜变量解码回图片。

VAE 编码器:从像素到分布

Encoder 的逐层结构如下。LayerNorm2d 在每个空间位置上沿通道维归一化,因此不会混合不同像素的位置统计。

阶段 运算 输出形状 作用
输入 二值 presence channels [N,2,28,28] 显式表示背景与前景
Stem 1×1 Conv: 2 → 32 [N,32,28,28] 把类别通道投影到图像特征空间
残差块 LN → 1×1(32→16) → 3×3 → GELU → 1×1(16→32) + skip [N,32,28,28] 在不改变分辨率的情况下提取局部笔画结构
下采样 2×2 Conv, stride=2: 32 → 4 [N,4,14,14] 同时压缩空间和通道
后验头 LayerNorm2d + 两个独立 1×1 Conv mu, logvar: [N,4,14,14] 参数化对角高斯后验

默认 VAE 共 5,402 个可训练参数。其主要形状变化可简写为:

[N,2,28,28]
  -> 1×1 Conv + ResidualBlock
  -> 2×2 stride-2 Conv + LayerNorm2d
  -> mu, logvar: [N,4,14,14]

重参数化把随机采样写成可求导形式:

sigma = exp(0.5 * logvar)
epsilon ~ N(0,I)
z = mu + sigma * epsilon

VAE 解码器与目标函数

Decoder 先在 [N,4,14,14] 上执行一个同结构残差块,再由 1×1 Conv: 4 → 128 生成 4 组子像素通道;PixelShuffle(2) 将其重排为 [N,32,28,28],最后用 3×3 Conv: 32 → 2 输出背景/前景 logits。训练时直接把 logits 交给 BCEWithLogits 以保持数值稳定,采样和展示时才应用 sigmoid。

训练损失为二通道 BCE 与 KL 正则之和;两项都先逐元素求和,再对 batch 取均值:

L_vae = BCEWithLogits(recon, x) + beta * KL(q(z|x) || N(0,I))
beta  = train.vae_kl_weight = 0.05

其中每个潜变量元素的 KL 为 -0.5 × (1 + logvar - mu² - exp(logvar))。BCE 保留笔画细节,KL 把后验约束到标准高斯附近;beta 控制重建精度与潜空间规则性之间的折中。Flow 使用的是后验采样 z 而不是只用 mu,因此会学习 VAE 编码不确定性所覆盖的数据分布。

条件 Flow Matching(流匹配)

流匹配可以直观地理解为学习一张“随时间变化的导航场”。高斯噪声是起点,真实图片对应的 VAE 潜变量是终点;模型在任意中间时刻看到当前位置后,需要预测下一步应该朝哪个方向、以多快的速度移动。

训练过程可以概括为三步:

  1. 随机取一份噪声 z₀ 和一份真实数据潜变量 z₁,用直线连接它们。
  2. 随机选择时间 t,得到直线上的中间位置 zₜ
  3. 让模型根据 zₜ、时间 t 和数字标签,预测从 z₀ 指向 z₁ 的正确速度。

采样时只有噪声,没有真实终点。模型从 t=0 开始反复预测速度并沿速度方向前进,最终到达一个符合指定数字标签的潜变量,再由 VAE Decoder 还原为图片。数字标签告诉模型“要生成什么”,时间则告诉模型“当前生成到哪一步”。

条件编码与 Flow 主干

时间 t 先编码为 16 组 sin/cos 特征,再经两层 MLP 映射到 128 维;数字标签通过查表也映射到 128 维。相加得到共享条件向量 c = MLP(TimeEmbedding(t)) + Embedding(y)。时间决定速度场处于噪声到数据路径的哪个位置,标签决定路径应进入哪个数字类别。

阶段 默认运算 输出形状 作用
输入投影 1×1 Conv: 4 → 256 [N,256,14,14] 将 VAE 潜变量提升到主干宽度
条件主干 8 个 Conditional ConvNeXt 块 [N,256,14,14] 保持空间分辨率并反复注入时间/标签条件
输出归一化 LayerNorm2d(256) [N,256,14,14] 稳定速度头输入
速度投影 1×1 Conv: 256 → 4 [N,4,14,14] 预测潜变量每个位置的瞬时速度

默认 Flow 共 2,756,868 个可训练参数。单个条件主干块的计算为:

h_dw = DepthwiseConv7×7(h)
(scale, shift) = Linear(c).chunk(2)
h_cond = LayerNorm2d(h_dw) * (1 + scale) + shift
delta = Conv1×1_512→256(GELU(Conv1×1_256→512(h_cond)))
h_next = h + delta

7×7 depthwise 卷积以较低参数量扩大空间感受野,两个 1×1 卷积负责通道混合和 扩展。AdaLN 的调制线性层以零初始化开始,使每个块在训练初期接近普通归一化残差块,再逐渐学会按 t 和数字标签缩放、平移各通道。残差连接则为 8 个块提供稳定的信息与梯度通路。

z_t: [N,4,14,14]
  -> 1×1 Conv, 4 -> 256
  -> 8 × ConditionalDepthwiseSeparableBlock
  -> LayerNorm2d
  -> 1×1 Conv, 256 -> 4
  -> velocity: [N,4,14,14]

Flow 训练使用 VAE posterior 样本作为数据端:

z₁ = mu + eps_post * exp(0.5 * logvar)
z₀ ~ N(0,I)
t  ~ Uniform(0,1)
zₜ = (1-t) * z₀ + t * z₁
target_velocity = z₁ - z₀
L_flow = MSE(FlowModel(zₜ, t, label), target_velocity)

这条路径是条件概率路径 zₜ = (1-t)z₀ + tz₁。对时间求导后得到恒定监督速度 dzₜ/dt = z₁-z₀;模型虽然只看到当前 zₜt 和标签,却通过回归大量随机配对的条件期望,学到可用于生成的连续速度场。它预测的是向量场而不是最终图像,因此同一网络可在任意连续时间被 Euler 求解器重复调用。

从速度场到图像

推理从 z⁽⁰⁾ ~ N(0,I) 开始,默认把 [0,1] 分为 32 段,并用显式 Euler 法更新:

Δt = 1 / 32
z⁽ᵏ⁺¹⁾ = z⁽ᵏ⁾ + Δt * vθ(z⁽ᵏ⁾, kΔt, label)
image = sigmoid(VAE_Decoder(z⁽³²⁾))

步数越多,离散轨迹通常越接近所学常微分方程,但推理成本也线性增加。VAE 将像素空间从 2×28×28 压到 4×14×14,使 Flow 的 32 次网络调用都发生在更小的空间上;代价是最终细节上限受 VAE 重建能力约束。

从头训练

训练 VAE:

.\.venv\Scripts\python.exe -m train.vae_trainer

输出 checkpoint:

ckpt/vae.pt

VAE 训练完成后再训练 Flow:

.\.venv\Scripts\python.exe -m train.flow_trainer

输出 checkpoint:

ckpt/flow.pt

两个训练入口都支持覆盖配置:

.\.venv\Scripts\python.exe -m train.vae_trainer --config path/to/override.yaml
.\.venv\Scripts\python.exe -m train.flow_trainer --config path/to/override.yaml

覆盖文件只需要写要修改的字段。例如快速检查配置可以写成:

data:
  batch_size: 16
  num_workers: 0
train:
  device: cpu
  vae_epochs: 1
  flow_epochs: 1
  max_train_steps: 5

如果目标 checkpoint 已存在,训练器会从下一轮 epoch 继续训练。若要从头开始,请先把旧 checkpoint 改名或移出 ckpt/

关键默认配置

完整配置与中文注释放在 config/default.yaml

配置 默认值 说明
data.batch_size 128 DataLoader batch size
data.binarize_threshold 0.5 MNIST 二值化阈值
train.device auto 自动选择 CUDA 或 CPU
train.vae_epochs 400 VAE 训练轮数
train.flow_epochs 200 Flow 训练轮数
train.vae_kl_weight 0.05 β-VAE 的 KL 权重
model.latent_channels 4 VAE 潜变量与 Flow 输入 / 输出通道数
model.latent_size 14 潜空间边长
model.flow_hidden_channels 256 Flow 主干宽度
model.flow_depth 8 条件主干块数量
model.condition_dim 128 时间 / 标签条件维度
sample.sampling_steps 32 Euler 更新步数
sample.history_steps 32 生成轨迹保存帧数
visual.pca_samples 512 PCA 最多使用的测试样本数

所有默认运行产物都写在项目目录内:数据在 datasets/,权重在 ckpt/,可视化在 outputs/,日志在 logs/

代码结构

NumberGenerate/
├── assets/readme/       README 使用的可视化快照
├── config/              配置加载、类型定义和默认参数
├── data/                MNIST 数据集与 DataLoader
├── model/               LayerNorm2d、AdaLN2d、VAE、FlowModel
├── train/               VAE / Flow 训练、checkpoint 与采样
├── vis/                 生成轨迹、PCA 与空间诊断图
├── tests/               配置、形状、损失和条件注入测试
├── debug/               Flow 诊断脚本
├── Doc/                 开发规范与代码索引
├── datasets/            运行时下载的数据,默认被 Git 忽略
├── ckpt/                训练权重,默认被 Git 忽略
├── outputs/             运行时可视化结果,默认被 Git 忽略
└── logs/                训练日志,默认被 Git 忽略

推荐阅读顺序:

  1. config/default.yaml
  2. data/mnist.py
  3. model/layers.py
  4. model/vae.py
  5. model/flow.py
  6. train/vae_trainer.pytrain/flow_trainer.py
  7. train/sampling.pyvis/plots.py

测试

运行全部单元测试:

.\.venv\Scripts\python.exe -m unittest discover -s tests

运行语法编译检查:

.\.venv\Scripts\python.exe -m compileall -q config data model train vis tests

测试覆盖配置加载、VAE 与 Flow 张量形状、条件注入、损失反向传播、checkpoint、逐步采样轨迹以及可视化入口检查。

License

本项目采用 Apache License 2.0

About

一个流匹配模型教育性项目

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages