PyTorch Data Sparsifier 详解:为 Tensor、参数与 Embedding 数据实现通用稀疏化
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
导读
本文围绕 PyTorch 仓库torch/ao/pruning/_experimental/data_sparsifier中的 Data Sparsifier 模块展开,系统讲解其设计动机、内部实现与完整使用流程。与传统基于模型的 Sparsifier 不同,Data Sparsifier 直接对裸数据张量(torch.Tensor、nn.Parameter、nn.Embedding/nn.EmbeddingBag)进行稀疏化,由稀疏化器自身持有 mask。读完本文,你将掌握BaseDataSparsifier的核心 API(add_data、step、get_data、get_mask、squash_mask、state_dict),能够自定义自己的数据稀疏化算法,并将稀疏化应用到模型 Embedding、训练数据预处理以及训练后稀疏量化等真实场景。
一、Data Sparsifier 是什么
Data Sparsifier 位于 torch/ao/pruning/_experimental/data_sparsifier,它继承自BaseSparsifier(定义于 torch/ao/pruning/sparsifier/base_sparsifier.py),目标是对通用数据张量进行稀疏化,这里的"数据"既包括可训练的(如模型权重参数),也包括不可训练的(如输入特征、中间数据)。
与普通 Sparsifier 的核心区别在于:
- 普通 Sparsifier(如 torch/ao/pruning/sparsifier/base_sparsifier.py)通过
prepare(model, config)接收一个模型和配置,mask 以参数化(parametrization)的形式挂在用户模型的对应层上; - Data Sparsifier不接收模型或层,而是接收
(name, data)数据对。因此 mask由数据稀疏化器自己持有,而不是由用户模型持有。
这一设计的关键实现技巧是:引入一个私有的容器模型(container model),将数据注册为"参数化 buffer"。从 base_data_sparsifier.py 可以看到容器只是一个空的nn.Module子类:
class _Container(nn.Module): passBaseDataSparsifier负责所有的"内务处理"(注册、mask 维护、state 管理),而使用者只需在自定义子类中实现核心的update_mask逻辑即可。
二、支持的数据类型
从 base_data_sparsifier.py 可以看到支持的类型集合:
EMBEDDING_TYPES = { nn.Embedding, nn.EmbeddingBag, } SUPPORTED_TYPES = { torch.Tensor, nn.Parameter, *EMBEDDING_TYPES, }即:
- 普通张量
torch.Tensor - 参数
nn.Parameter - Embedding 与 EmbeddingBag 模块(
nn.Embedding/nn.EmbeddingBag)
在add_data时,若数据类型不在SUPPORTED_TYPES中会直接抛出AssertionError(base_data_sparsifier.py)。
内部通过_extract_weight统一取出真正要稀疏化的张量:对 Tensor 和 Parameter 直接返回自身;对 Embedding / EmbeddingBag 则返回其.weight(base_data_sparsifier.py):
def _extract_weight(self, data): # extract the weight parameter instead of underlying data if type(data) in [torch.Tensor, nn.Parameter]: return data elif type(data) in EMBEDDING_TYPES: return data.weight配套测试 test/ao/sparsity/test_data_sparsifier.py 分别对torch.randn张量、nn.Parameter、nn.Embedding/nn.EmbeddingBag三类数据跑同一套检查(构造、step、squash、state_dict、内存引用等),验证了这三类数据类型都能被完整支持。
三、核心 API 详解
BaseDataSparsifier是抽象基类,抽象方法为update_mask,它负责为所有已注册数据计算新 mask。
3.1 构造与data_list
构造函数(base_data_sparsifier.py)签名如下:
def __init__(self, data_list=None, **defaults):data_list:(name, data)元组的列表,构造时即调用add_data(name, data, **self.defaults)逐条注册;defaults:默认配置字典,所有未在单条数据配置中显式指定的 key 都会落到默认值。
>>> data_list = [('tensor_1', torch.randn(3,3)), ('tensor_2', torch.randn(4,4))] >>> defaults = {'sparsity_level': 0.7} >>> sparsifier = DerivedDataSparsifier(data_list=data_list, **defaults) >>> new_tensor_to_add = {'name': 'tensor_3', 'data': torch.randn(5,5), 'sparsity_level': 0.3} >>> sparsifier.add_data(**new_tensor_to_add) # tensor_1 和 tensor_2 使用 sparsity_level=0.7,tensor_3 使用 sparsity_level=0.33.2add_data:注册数据
data_sparsifier = ImplementedDataSparsifier() data_sparsifier.add_data(name=name, data=data, **some_config)add_data(name, data, reuse_mask=True, **config)(base_data_sparsifier.py)做以下事情:
- 校验类型:数据必须在
SUPPORTED_TYPES内; - 合并配置:
local_args = copy.deepcopy(self.defaults)然后local_args.update(config),实现"默认配置 + 单条数据专属配置"的覆盖关系; - 准备 mask 与参数化类:
mask = local_args.get("mask", torch.ones_like(weight)),param_class = local_args.get("parametrization", utils.FakeSparsity),即默认 mask 全 1、默认参数化方式为FakeSparsity; - 注册到容器:
self._container.register_buffer(name=name, tensor=weight)后调用parametrize.register_parametrization(self._container, name, param_class(mask)),同时把 mask 记入self.state[name]["mask"],把合并后的配置记入self.data_groups[name]。
关于同名数据替换的语义(base_data_sparsifier.py):
- 若
name已存在,会发出UserWarning("Replacing existing data of the same name..."); - 默认复用旧 mask(
reuse_mask=True),因此新数据形状必须与旧数据一致,否则抛AssertionError; - 默认复用旧配置,除非在
config中显式传入新配置,则用新配置覆盖; - 想丢弃旧 mask 时显式传
reuse_mask=False。
注意:包含
.的 name 不是合法名称。原因在于容器模型内部通过属性访问数据,.会被解释为层级分隔。Lightning 工具函数中专门用_get_valid_name把.替换为_(见 lightning/callbacks/_data_sparstity_utils.py)。
add_data会返回getattr(self._container, name),即注册后的数据对象。
3.3step:计算 mask
data_sparsifier.step()step()(base_data_sparsifier.py)在torch.no_grad()下遍历self.data_groups,对每个(name, config)取出未稀疏化的原始数据(self.get_data(name)),然后调用self.update_mask(name, data, **config):
def step(self): if not self.enable_mask_update: return with torch.no_grad(): for name, config in self.data_groups.items(): # get non-sparsified data data = self.get_data(name) self.update_mask(name, data, **config)enable_mask_update置为False可以临时停用 mask 更新(该开关继承自BaseSparsifier,见 base_sparsifier.py)。
3.4get_mask与get_data:读取 mask 与数据
get_mask(name)直接返回self.state[name]["mask"](base_data_sparsifier.py),name 不存在时抛ValueError。
get_data(name, return_original=True)(base_data_sparsifier.py):
return_original=True:返回未应用 mask的原始数据(通过parametrizations.<name>.original取得)。若 mask 已被 squash 掉,会抛ValueError(因为原始值已不存在);return_original=False:返回应用了 mask(参数化)之后的稀疏化数据,即data * mask。
original_data = data_sparsifier.get_data(name=name, return_original=True) # 返回未应用 mask 的数据 sparsified_data = data_sparsifier.get_data(name=name, return_original=False) # 返回 data * mask3.5squash_mask:落地并移除 mask
data_sparsifier.squash_mask()squash_mask(*args, leave_parametrized=True, names=None, **kwargs)(base_data_sparsifier.py)移除数据上的参数化:
names:字符串列表,指定只对哪些数据 squash;为None时对data_groups中所有 key 执行;leave_parametrized=True:移除参数化前先把 mask 应用到数据上(即数据变为data * mask);leave_parametrized=False:移除参数化但不应用 mask(数据保持原始值)。
底层通过parametrize.remove_parametrizations(self._container, name, leave_parametrized=leave_parametrized)实现。
3.6state_dict/load_state_dict:序列化与恢复
state_dict()(base_data_sparsifier.py)返回可序列化字典,包含三部分:
state:name -> mask的映射(mask 默认转为sparse COO格式存储以节省空间,由_convert_mask完成);data_groups:所有稀疏化配置分组,key 为数据名;_container:内部容器模型的 state dict。
load_state_dict(state_dict, strict=True)(base_data_sparsifier.py):
strict=True:先重置容器再精确恢复到state_dict状态;strict=False:不重置,已有数据保留,state_dict中的内容与现有状态合并。
_load_container_from_state会依据容器 state dict 中是否存在parametrizations.<name>.original键来判断该数据当时是否处于参数化状态,从而决定恢复时是否重新注册参数化(base_data_sparsifier.py)。
配套测试 test/ao/sparsity/test_data_sparsifier.py 验证了:序列化后 mask 为稀疏格式,加载后转回稠密并与原 mask 完全一致,同时data_groups与容器参数化状态也保持一致。
四、自定义数据稀疏化器
自定义数据稀疏化器只需两步:
- 继承
BaseDataSparsifier; - 实现
update_mask(self, name, data, **kwargs)。
以下示例来自 README.md:将所有绝对值小于阈值的条目置零。
class ImplementedDataSparsifier(BaseDataSparsifier): def __init__(self, threshold): super().__init__(threshold=threshold) def update_mask(self, name, data, threshold): mask = self.get_mask(name) mask[torch.abs(data) < threshold] = 0.04.1update_mask的编写约束
README 特别强调了两条硬性规则(在step的实现中也可以印证):
- 何时调用由
BaseDataSparsifier负责:用户无需手动触发update_mask,只需调用step(),基类会遍历全部数据并自动调用; - mask 必须原地(inplace)修改:
step中通过self.get_mask(name)取到的是state里持有引用,因此必须原地修改才能生效。
合法的原地操作示例:
mask[:10] = torch.zeros(10) # 修改 mask 的一部分 mask *= another_mask # 使用原地运算符 mask.data = torch.zeros_like(mask) # 替换底层数据非原地操作会引入 bug,例如:
mask = torch.zeros_like(mask) # 重新赋值,外部引用不会更新 mask = mask * another_mask # 非原地算术运算从DataNormSparsifier的实现可以看到原地写法的正确示范:它统一通过mask.data = ...赋值(见 data_norm_sparsifier.py)。
4.2 现成实现:DataNormSparsifier
仓库自带一个完整实现DataNormSparsifier(data_norm_sparsifier.py),它按块计算范数并把范数最小的稀疏块置零,由三个超参数控制:
| 参数 | 类型 | 默认值 | 含义 |
|---|---|---|---|
sparsity_level | float | 0.5 | 被置零的稀疏块比例 |
sparse_block_shape | tuple[int, int] | (1, 4) | 稀疏块形状,块从张量零索引处开始划分 |
zeros_per_block | int | None | 块内元素总数 | 每个稀疏块内期望的零的个数;未指定时整块置零 |
norm | str | "L1" | 范数类型,仅支持"L1"/"L2" |
其内部逻辑(data_norm_sparsifier.py)为:
- 校验:
zeros_per_block不能超过块内元素总数、不能为负;当前仅支持 2D 数据(1D 数据会先补成[None, :]); - 计算范数:L1 用
torch.abs(data),L2 用data * data; - 两级 mask:
- 数据级 mask(
__get_data_level_mask):用F.avg_pool2d按块池化得到每块范数,排序后把round(sparsity_level * num_blocks)个最小范数块整体置零; - 块级 mask(
__get_block_level_mask):用F.unfold展开成块,对每块内元素排序后置零最小的zeros_per_block个; - 最终
mask.data = torch.where(data_lvl_mask == 1, data_lvl_mask, block_lvl_mask)合并两级 mask;
- 数据级 mask(
- 边界处理:
sparsity_level <= 0或zeros_per_block == 0时 mask 全 1(不稀疏化);sparsity_level >= 1.0且块全置零时 mask 全 0;张量边缘不整除块大小时做 padding(padding 区域填充 NaN,避免边缘数据被误删)。
测试 test/ao/sparsity/test_data_sparsifier.py 还提供了实际稀疏度的上下界估算公式:实际稀疏度并不严格等于sparsity_level,而是取决于张量形状与数据本身:
number_blocks = ceil(height / block_height) * ceil(width / block_width) values_per_block = block_height * block_width min_values_sparsified = round(number_blocks * sparsity_level) max_values_sparsified = min_values_sparsified * min(values_per_block, zeros_per_block) lower_bound = min_values_sparsified / (height * width) upper_bound = min(1.0, max_values_sparsified / (height * width))此外DataNormSparsifier的构造参数全部属于"默认参数",在add_data阶段传入的配置可以按 name 逐一覆盖。
五、实战场景
5.1 简单示例:稀疏化张量与参数
tensor1 = torch.randn(100, 100) param1 = nn.Parameter(torch.randn(200, 32)) my_sparsifier = ImplementedDataSparsifier(threshold=0.2) my_sparsifier.add_data(name='tensor1', data=tensor1, threshold=0.5) # 单条数据专属配置覆盖默认 threshold my_sparsifier.add_data(name='param1', data=param1) my_sparsifier.step() # 计算 mask my_sparsifier.squash_mask() # 应用并移除 mask5.2 稀疏化模型中的 Embedding / EmbeddingBag
Data Sparsifier 的典型用途是训练后稀疏化模型中的 Embedding 层(推荐系统、NLP 模型中 Embedding 表往往巨大):
class Model(nn.Module): def __init__(self, feature_dim, emb_dim, num_classes): self.emb = nn.EmbeddingBag(feature_dim, emb_dim) self.linear1 = nn.Linear(emb_dim, 32) self.linear2 = nn.Linear(32, num_classes) self.relu = nn.ReLU() def forward(self, x): out = self.emb(x) out = self.relu(self.linear1(out)) out = self.linear2(out) return out model = Model(100, 32, 10) my_sparsifier = ImplementedDataSparsifier(threshold=0.5) my_sparsifier.add_data(name='emb', data=model.emb) # ... 训练模型 ... my_sparsifier.step() # 为 embeddings 创建 mask my_sparsifier.squash_mask() # 应用并移除 mask注意add_data接收的是model.emb这个 EmbeddingBag 模块,内部会通过_extract_weight取其.weight进行稀疏化,而 mask 挂在 sparsifier 自己的容器模型上,不影响原模型的参数化结构。
5.3 训练数据场景:送入模型前稀疏化输入
如果输入数据可以在送入模型之前就被稀疏化,也可以在训练循环中使用 Data Sparsifier。批处理输入需要先挂到 sparsifier 上,再送入模型:
model = SomeModel() data_sparsifier = ImplementedDataSparsifier(threshold=0.2) data_name = 'train_data' for x, y in train_data_loader: x = data_sparsifier.add_data(name=data_name, data=x) # 注册并返回注册后的数据 ... y_out = model(x) ... data_sparsifier.step()每轮迭代对同一name调用add_data会触发同名替换逻辑:默认复用旧 mask 与旧配置,因此要求每批数据形状一致。若不需要复用 mask,可显式传入reuse_mask=False。
5.4 训练后稀疏量化:post_training_sparse_quantize
仓库还提供post_training_sparse_quantize(quantization_utils.py),将稀疏化与量化组合应用于模型中的 Embedding / EmbeddingBag:
post_training_sparse_quantize( model, data_sparsifier_class=DataNormSparsifier, sparsify_first=True, # True:先稀疏化再量化;False:先量化再稀疏化 select_embeddings=None, # None 表示处理模型所有 embedding;也可传入模块列表 sparsity_level=0.8, sparse_block_shape=(1, 1), )sparsify_first=True:先对 embedding 权重做稀疏化并squash_mask,再走torch.ao.quantization.prepare / convert量化;sparsify_first=False:先量化,保存每通道的 scale 与 zero-point,反量化后进行稀疏化与squash_mask,最后用保存的量化参数通过torch.quantize_per_channel重新量化。
测试 test/ao/sparsity/test_data_sparsifier.py 验证了两种顺序下:embedding 都被稀疏化到约sparsity_level(80%),并被转换为torch.ao.nn.quantized.modules.embedding_ops.Embedding / EmbeddingBag,而模型中的nn.Linear不会被量化。
5.5 Lightning 回调集成
数据稀疏化器还提供了 PyTorch Lightning 回调封装(lightning/callbacks/data_sparsity.py):
PostTrainingDataSparsity:在on_fit_end中对训练好的模型副本做一次step()+squash_mask(),通过<callback>.sparsified获取稀疏化后的模型;TrainingAwareDataSparsity:在训练过程中逐 epoch 调用 sparsifier 与配套 scheduler 的step(),并在 epoch 间通过state_dict/load_state_dict传递稀疏化状态,on_train_end时 squash mask。
配套的 lightning/callbacks/README.md 与 lightning/tests/test_callbacks.py 可作进一步参考。
六、设计要点小结
- mask 归属:Data Sparsifier 通过私有
_Container模型 +parametrize参数化机制持有 mask,使用者无需维护任何 mask 状态; - 配置合并规则:默认配置 →
add_data显式配置逐层覆盖,替换同名数据时默认复用旧配置与旧 mask(reuse_mask控制); - 类型支持:Tensor / Parameter / Embedding / EmbeddingBag,统一经
_extract_weight提取权重; - 命名约束:数据名不能包含
.(容器属性访问的限制); - mask 更新:必须原地修改,
step()统一驱动,enable_mask_update可整体暂停; - 序列化友好:
state_dict将 mask 转 sparse COO 存储,load_state_dict支持严格/非严格两种恢复模式; - 生态衔接:配套
DataNormSparsifier开箱即用实现,post_training_sparse_quantize打通稀疏化与量化,Lightning 回调支持训练后与训练中两种接入方式。
对源码细节感兴趣可继续深入阅读 base_data_sparsifier.py、data_norm_sparsifier.py 及单元测试 test/ao/sparsity/test_data_sparsifier.py。
【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考