PyTorch 的可微量化框架。
项目描述
通过伪量化噪声的可微模型压缩
DiffQ 使用伪量化噪声执行可微量化。它可以自动调整每个权重或权重组使用的位数,以在模型大小和准确性之间实现给定的权衡。
阅读我们的论文了解更多详情。
这是怎么回事?
有关版本的详细信息,请参阅更改日志。
- 2022-08-24:v0.2.3:修复了加载旧量化状态时的错误。
- 2021-11-25:版本 0.2.2:添加对 torchscript 的支持。
要求
DiffQ 需要 Python 3.7 和相当新的 PyTorch 版本(理想情况下是 1.7.1)。要安装 DiffQ,您可以从存储库的根目录运行:
pip install .
你也可以直接从 PyPI 安装pip install diffq.
用法
import torch
from torch.nn import functional as F
import diffq
from diffq import DiffQuantizer
model = MyModel()
optim = ... # The optimizer must be created before the quantizer
quantizer = DiffQuantizer(model)
quantizer.setup_optimizer(optim)
# Distributed data parallel must be created after DiffQuantizer!
dmodel = torch.distributed.DistributedDataParallel(...)
penalty = 1e-3
model.train() # call model.eval() on eval to automatically use true quantized weights.
for batch in loader:
...
optim.zero_grad()
# The `penalty` parameter here will control the tradeoff between model size and model accuracy.
loss = F.mse_loss(x, y) + penalty * quantizer.model_size()
optim.step()
# To get the true model size with when doing proper bit packing.
print(f"Model is {quantizer.true_model_size():.1f} MB")
# When you want to dump your final model:
torch.save(quantizer.get_quantized_state(), "some_file.th")
# You can later load back the model with
model = MyModel()
diffq.restore_quantized_state(model, torch.load("some_file.th"))
# For DiffQ models, we support exporting the model to Torscript with optimal storage.
# Once loaded, the model will be stored in fp32 in memory (int8 support coming up).
from diffq.ts_export import export
export(quantizer, 'quantized.ts')
文档
有关详细文档,请参阅API文档。我们将在下文中介绍几个方面。
量化器对象
量化器在创建时附加到模型。所有 Quantizer 对象都提供相同的基本功能:
- 如果模型处于评估模式,则会自动切换到前向量化权重。
- 前向训练的量化器特定代码(例如,带有 QAT 的 UniformQuantizer 的 STE,DiffQ 的噪声注入)。
- 提供对量化模型大小和状态的访问。
量化大小和状态
该方法quantizer.model_size()提供了可微分的模型大小(对于 DiffQ),同时quantizer.true_model_size()提供了真实的、最佳位压缩的模型大小(不可微分)。quantizer.compressed_model_size()您可以使用gzip. 这实际上可能大于真实模型大小,并揭示了有关特定量化方法的熵使用的有趣信息。
用 获得位压缩量化状态,用quantizer.get_quantized_state()恢复quantizer.restore_quantized_state()。位打包针对速度进行了优化,并且可能会受到一些开销(实际上,Uniform 和 LSQ 不超过 120B,DiffQ 不超过 1kB)。
如果您无法访问原始量化器,例如在推理时,您可以使用diffq.restore_quantized_state(model, quantized_state).
量化器和优化
一些量化器会添加额外的可优化参数(DiffQuantizer 和 LSQ)。这些参数可能需要与主要模型权重不同的优化器或超参数。通常,DiffQ 位参数始终使用 Adam 进行优化。因此,您应该始终 在量化器之前创建主优化器。然后,您可以使用此优化器或其他优化器设置量化器:
model = MyModel(...)
opt = torch.optim.Adam(model.parameters())
quantizer = diffq.DiffQuantizer(model)
quantizer.setup_optimizer(opt, **optim_overrides)
这提供了使用单独超参数的自由。例如,DiffQuantizer
将始终为 bits 参数停用 weight_decay。
如果主要优化器是 SGD,建议为量化器使用第二个 Adam 优化器。
警告:您必须始终DistributedDataParallel
在创建量化器后包装您的模型,否则量化器参数将不会被优化!
火炬脚本支持
目前 TorchScript 支持是实验性的。我们支持使用 TorchScript 将模型保存到具有最佳存储空间的磁盘。加载后,模型将存储在内存中的 FP32 中。我们正在努力在内存中增加对 int8 的支持。请参阅diffq.ts_export.exportAPI 中的函数。
例子
examples/我们在文件夹中提供了三个示例。一种是针对 CIFAR-10/100,使用 Wide-ResNet、ResNet 或 MobileNet 等标准架构。第二个是基于DeiT视觉转换器。第三个是 Wikitext-103 上的语言建模任务,使用Fairseq
DeiT 和 Fairseq 示例在特定提交时作为原始代码库的补丁提供。您可以通过运行初始化 git 子模块并应用补丁
make examples
有关每个示例的更多详细信息,请查看其特定的 README:
开发安装
这将在开发人员模式下安装依赖项和 a diffq(对文件的更改将直接反映),以及运行单元测试的依赖项。
pip install -e '.[dev]'
更新基于补丁的示例
为了更新补丁,首先运行make examples以正确初始化子存储库。然后执行您想要的所有更改,提交它们并运行make patches. 这将更新每个 repo 的补丁。完成此操作后,并检查您所做的所有更改是否正确包含在新的补丁文件中,您可以运行make reset(这将删除您从子模块中所做的所有更改,因此请在调用此之前检查补丁文件)呼叫git add -u .; git commit -m "my changes"和推动。
测试
您可以运行单元测试
make tests
引文
如果您在论文中使用此代码或结果,请引用我们的工作:
@article{defossez2021differentiable,
title={Differentiable Model Compression via Pseudo Quantization Noise},
author={D{\'e}fossez, Alexandre and Adi, Yossi and Synnaeve, Gabriel},
journal={arXiv preprint arXiv:2104.09987},
year={2021}
}
执照
此存储库是在 CC-BY-NC 4.0 下发布的。许可证文件中的
许可证,但以下部分属于 MIT 许可证。这些文件examples/cifar/src/mobilenet.py和examples/cifar/src/src/resnet.py取自kuangliu/pytorch-cifar,作为 MIT 发布。该文件examples/cifar/src/wide_resnet.py取自meliketoy/wide-resnet,作为 MIT 发布。有关详细许可,请参阅每个文件头。
项目详情
下载文件
下载适用于您平台的文件。如果您不确定要选择哪个,请了解有关安装包的更多信息。
源分布
内置发行版
diffq -0.2.3-cp310-cp310-manylinux_2_5_x86_64.manylinux1_x86_64.manylinux_2_12_x86_64.manylinux2010_x86_64.whl 的哈希值
| 算法 | 哈希摘要 | |
|---|---|---|
| SHA256 | f54fde05500cabd48c71c8b6b0a364bf746ecf3edb71d570af6ab9b7736cb3c9 |
|
| MD5 | e59de118c8664dffcc6a20a797287890 |
|
| 布莱克2-256 | b5fbab92b5183732c50866cadaf3cffb64b41f70c946cd67e567d2e67a19a8cb |
diffq -0.2.3-cp310-cp310-manylinux_2_5_i686.manylinux1_i686.manylinux_2_12_i686.manylinux2010_i686.whl 的哈希值
| 算法 | 哈希摘要 | |
|---|---|---|
| SHA256 | 827d0f631863dc9a815ced20bc9376c5467cb81f2ae8c1656ed91a54bf016b5a |
|
| MD5 | 947c5643615226cb1291cad751b14ce5 |
|
| 布莱克2-256 | f262054de9a40d56a89901f2e38021af286acefdf903efb9a0f4195b56403dca |