ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

一张图片如何变成512个离散Token?deep-vector-quantization核心量化层VQVAE原理深度剖析

一张图片如何变成512个离散Token?deep-vector-quantization核心量化层VQVAE原理深度剖析 一张图片如何变成512个离散Tokendeep-vector-quantization核心量化层VQVAE原理深度剖析【免费下载链接】deep-vector-quantizationVQVAEs, GumbelSoftmaxes and friends项目地址: https://gitcode.com/gh_mirrors/de/deep-vector-quantizationdeep-vector-quantization是一个用 PyTorch PyTorch Lightning 实现VQVAEVector Quantized Variational Autoencoder向量量化变分自编码器完整训练流程的开源项目通过带离散潜变量瓶颈的自编码器把一张图片压缩成一串离散 Token让图像可以直接接入 GPT 这类序列建模基础设施这也是 DALL-E 的核心思路。本文面向新手带你从零看懂 VQVAE 核心量化层的工作原理。一句话理解 VQVAE图片的离散表示学习普通自编码器把图片压缩成连续向量VQVAE 则多了一步舍入把压缩后的每个位置强制吸附到一张码本codebook里预先定义的若干离散向量之一。结果就是一张 32×32×3 的 CIFAR-10 图片 → 一个 8×8 的整数索引矩阵每个位置的值 ∈ {0, 1, …, 511}即从512 个码本向量中挑一个这串索引就是图片的Token 序列可以像文本词元一样喂给 GPT 建模 整个流程在 dvq/vqvae.py 中只有三行encoder → quantizer → decoder。解剖三部分结构编码器、量化层、解码器VQVAE 由三部分三明治式拼接而成见 dvq/vqvae.py组件作用实现文件编码器 Encoder图片 → 连续特征图dvq/model/deepmind_enc_dec.py量化层 Quantizer特征图 → 离散索引核心dvq/model/quantize.py解码器 Decoder离散特征 → 重建图片dvq/model/deepmind_enc_dec.py编码器32×32 图片变 8×8 特征DeepMind 风格编码器由两次 stride2 的下采样卷积 两个残差块组成输出8×8×128的特征图stride 4。项目还提供了 OpenAI DALL-E 风格编码器stride 8输出 4×4见 dvq/model/openai_enc_dec.py。量化层512 是怎么来的核心参数在 dvq/vqvae.py--num_embeddings 512码本大小即 512 个候选离散向量词表大小--embedding_dim 64每个码本向量的维度量化步骤dvq/model/quantize.py用 1×1 卷积把 8×8×128 特征图投影成 8×8×64拉平成 64 个 64 维向量逐个计算与 512 个码本向量的欧氏距离取距离最近的码本向量得到一个 8×8 的整数索引矩阵⚠️ 注意区分两个数字512 是词表大小每个位置可取的离散状态数而一张 32×32 图片实际产生8×8 64 个 Token 位置每个位置从 512 个码本向量中选 1 个。直通估计器梯度如何跨越不可导的量化argmin取最近向量是不可导的梯度在这里就断了。VQVAE 的解法是直通梯度估计器Straight-Through Estimator关键就一行dvq/model/quantize.pyz_q z_e (z_q - z_e).detach()前向结果就是量化向量z_q数值不变反向.detach()让梯度直接穿透编码器收到的是解码器传回的原始梯度 防码本坍缩k-means 初始化 commitment 损失VQVAE 训练最大的坑是索引坍缩catastrophic index collapse——模型只用少数几个码本向量512 个词退化成几个。项目用了两招数据驱动的 k-means 初始化训练开始时采样 2 万个特征向量跑一遍kmeans2初始化码本而不是随机初始化dvq/model/quantize.pyCommitment 损失用0.25权重惩罚编码器输出远离码本同时让码本向量朝编码器输出移动dvq/model/quantize.py训练时还能通过困惑度perplexity监控码本健康度perplexity 越接近 512说明 512 个码本向量被用得越均匀dvq/vqvae.py。训练目标重建损失 量化损失总损失 重建损失 潜变量损失dvq/vqvae.py重建损失默认用固定方差的正态分布负对数似然本质是归一化的 MSE实现在 dvq/model/loss.py也支持 DALL-E 的 Logit-Laplace 损失量化损失即上文 commitment 项乘以kld_scale10.0加权数据侧默认加载 CIFAR-10带随机裁剪和水平翻转增强见 dvq/data/cifar10.py。Gumbel-Softmax另一条可微分量化路线除了找最近邻 直通梯度项目还实现了Gumbel-Softmax方案dvq/model/quantize.py把量化建模成在 512 类上的采样用 Gumbel 技巧让采样可微。配套两个退火调度dvq/vqvae.py温度 τ 从 1 余弦退火到 1/16前 15 万步逐步从软采样逼近硬选择KL 权重 β 从 0 缓慢升上去避免早期被正则项主导⚡ 实测提示来自 README.mdGumbel 版本收敛稍慢、重建损失略高超参调起来比较娇气新手建议先用默认--vq_flavor vqvae。快速上手一条命令训练你的第一个 VQVAE环境依赖见 requirements.txttorch、torchvision、pytorch-lightning、scipy然后git clone https://gitcode.com/gh_mirrors/de/deep-vector-quantization cd deep-vector-quantization/dvq python vqvae.py --gpus 1 --data_dir /path/to/cifar10这条命令即可在 CIFAR-10 上复现 DeepMind 原始 VQVAE。其他常用参数--enc_dec_flavor deepmind / openai切换两套编解码器架构--num_embeddings码本大小默认 512--embedding_dim码本向量维度默认 64为什么要把图片Token 化把图片变成离散 Token 序列后图像生成就变成了序列预测问题——这正是 DALL-E 的路线先训 VQVAE 做图片↔Token 互转再用 GPT 学习 Token 序列的分布。项目 README 中也明确提到后续目标是复现 DALL-Elogit-laplace 损失 ImageNet 训练。训练完成后还可以用 visualize.ipynb 可视化重建效果 项目文件导航文件说明dvq/vqvae.py训练入口 VQVAE LightningModule 主体dvq/model/quantize.py⭐ 核心量化层VQVAE / Gumbel 两种dvq/model/deepmind_enc_dec.pyDeepMind 风格编解码器stride 4dvq/model/openai_enc_dec.pyOpenAI DALL-E 风格编解码器stride 8dvq/model/loss.pyNormal / LogitLaplace 重建损失dvq/data/cifar10.pyCIFAR-10 数据加载visualize.ipynb训练结果可视化 notebook小结VQVAE 的精髓就在编码器—量化层—解码器这条流水线里——用 1×1 卷积投影、最近邻查找完成连续到离散的舍入用直通估计器救回梯度用 k-means 初始化和 commitment 损失防止码本坍缩。理解了这个核心量化层你就掌握了 DALL-E 等扩散/生成式 AI 系统的地基。【免费下载链接】deep-vector-quantizationVQVAEs, GumbelSoftmaxes and friends项目地址: https://gitcode.com/gh_mirrors/de/deep-vector-quantization创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表