news 2026/9/11 3:18:14

PyTorch Data Sparsifier 详解:为 Tensor、参数与 Embedding 数据实现通用稀疏化

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch Data Sparsifier 详解:为 Tensor、参数与 Embedding 数据实现通用稀疏化

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.Tensornn.Parameternn.Embedding/nn.EmbeddingBag)进行稀疏化,由稀疏化器自身持有 mask。读完本文,你将掌握BaseDataSparsifier的核心 API(add_datastepget_dataget_masksquash_maskstate_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): pass

BaseDataSparsifier负责所有的"内务处理"(注册、mask 维护、state 管理),而使用者只需在自定义子类中实现核心的update_mask逻辑即可。

二、支持的数据类型

从 base_data_sparsifier.py 可以看到支持的类型集合:

EMBEDDING_TYPES = { nn.Embedding, nn.EmbeddingBag, } SUPPORTED_TYPES = { torch.Tensor, nn.Parameter, *EMBEDDING_TYPES, }

即:

  1. 普通张量torch.Tensor
  2. 参数nn.Parameter
  3. 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.Parameternn.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.3

3.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)做以下事情:

  1. 校验类型:数据必须在SUPPORTED_TYPES内;
  2. 合并配置local_args = copy.deepcopy(self.defaults)然后local_args.update(config),实现"默认配置 + 单条数据专属配置"的覆盖关系;
  3. 准备 mask 与参数化类mask = local_args.get("mask", torch.ones_like(weight))param_class = local_args.get("parametrization", utils.FakeSparsity),即默认 mask 全 1、默认参数化方式为FakeSparsity
  4. 注册到容器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...");
  • 默认复用旧 maskreuse_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_maskget_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 * mask

3.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)返回可序列化字典,包含三部分:

  • statename -> 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与容器参数化状态也保持一致。

四、自定义数据稀疏化器

自定义数据稀疏化器只需两步:

  1. 继承BaseDataSparsifier
  2. 实现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.0

4.1update_mask的编写约束

README 特别强调了两条硬性规则(在step的实现中也可以印证):

  1. 何时调用由BaseDataSparsifier负责:用户无需手动触发update_mask,只需调用step(),基类会遍历全部数据并自动调用;
  2. 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_levelfloat0.5被置零的稀疏块比例
sparse_block_shapetuple[int, int](1, 4)稀疏块形状,块从张量零索引处开始划分
zeros_per_blockint | None块内元素总数每个稀疏块内期望的零的个数;未指定时整块置零
normstr"L1"范数类型,仅支持"L1"/"L2"

其内部逻辑(data_norm_sparsifier.py)为:

  1. 校验:zeros_per_block不能超过块内元素总数、不能为负;当前仅支持 2D 数据(1D 数据会先补成[None, :]);
  2. 计算范数:L1 用torch.abs(data),L2 用data * data
  3. 两级 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;
  4. 边界处理:sparsity_level <= 0zeros_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() # 应用并移除 mask

5.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 可作进一步参考。

六、设计要点小结

  1. mask 归属:Data Sparsifier 通过私有_Container模型 +parametrize参数化机制持有 mask,使用者无需维护任何 mask 状态;
  2. 配置合并规则:默认配置 →add_data显式配置逐层覆盖,替换同名数据时默认复用旧配置与旧 mask(reuse_mask控制);
  3. 类型支持:Tensor / Parameter / Embedding / EmbeddingBag,统一经_extract_weight提取权重;
  4. 命名约束:数据名不能包含.(容器属性访问的限制);
  5. mask 更新:必须原地修改,step()统一驱动,enable_mask_update可整体暂停;
  6. 序列化友好state_dict将 mask 转 sparse COO 存储,load_state_dict支持严格/非严格两种恢复模式;
  7. 生态衔接:配套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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/11 3:16:58

工业PCB视觉检测:YOLO26+大模型融合落地实践

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 3:15:43

虚拟机忘记密码?VMware、Linux、Windows重置全攻略

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 3:14:58

YOLOv8-seg中文车牌识别:12类全制式实战部署指南

简介&#xff1a;本资源是一套基于YOLOv8的高兼容性中文车牌识别系统&#xff0c;面向人工智能、自动化、电子信息等专业的高校学生及初阶开发者&#xff0c;解决多类型车牌&#xff08;含单双层蓝牌、新能源绿牌、警用车牌、军牌等12类&#xff09;的端到端检测与识别问题&…

作者头像 李华
网站建设 2026/9/11 3:14:32

Redis数据安全加固实战:访问控制、持久化与分布式锁

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 3:14:00

容器里跑安卓模拟器:一条命令完成 Docker 部署

容器里跑安卓模拟器&#xff1a;一条命令完成 Docker 部署 【免费下载链接】docker-android Android in docker solution with noVNC supported, video recording and mcp server 项目地址: https://gitcode.com/GitHub_Trending/do/docker-android 不想在宿主机装整套 …

作者头像 李华
网站建设 2026/9/11 3:12:37

ArduPilot 完整指南:如何用开源飞控快速上手无人机自主飞行

ArduPilot 完整指南&#xff1a;如何用开源飞控快速上手无人机自主飞行 【免费下载链接】ardupilot ArduPlane, ArduCopter, ArduRover, ArduSub source 项目地址: https://gitcode.com/GitHub_Trending/ar/ardupilot ArduPilot 是 DIY Drones 社区发起并维护的开源飞行…

作者头像 李华