资讯详情

昇思 MindSpore 大模型:属性过滤

📅 2026/9/30 20:22:36 | 华诺云谱 👁 阅读
昇思 MindSpore 大模型:属性过滤
大模型训练、微调、推理流程中经常需要对张量、网络参数、检查点权重、数据集样本进行属性过滤。典型场景筛选指定精度参数、过滤冻结权重、按设备属性筛选算子、过滤无效训练样本、加载 Checkpoint 时按需筛选权重。MindSpore 提供参数属性标记、张量属性、网络 Cell 属性、数据集过滤 API。属性过滤可以实现冻结部分层、选择性加载权重、动态路由算子、清洗训练数据减少显存占用加速训练微调。本文基于 MindSpore 2.3围绕网络参数过滤、Checkpoint 权重过滤、数据集样本属性过滤提供完整代码。一、核心原理MindSpore 中支持多种属性载体Cell网络层可增加自定义attr属性Parameter权重参数支持requires_grad、dtype、自定义标签Dataset样本可携带标签属性通过filter算子过滤CheckpointDict加载权重时基于参数名、形状、数据类型过滤。属性过滤通用流程标记属性 → 定义过滤条件 → 遍历筛选 → 执行后续逻辑冻结 / 加载 / 丢弃。二、场景 1网络 Parameter 属性过滤分层冻结微调最常用场景大模型微调通过属性筛选参数冻结 Backbone仅训练 Head 层。import mindspore as ms from mindspore import nn ms.set_context(modems.GRAPH_MODE, device_targetAscend) class LLMBackbone(nn.Cell): def __init__(self): super().__init__() self.embedding nn.Embedding(vocab_size32000, embedding_size512) self.transformer_layer nn.Dense(512, 512) # 自定义属性标记backbone层 self.embedding.attr {group: backbone} self.transformer_layer.attr {group: backbone} class LLMHead(nn.Cell): def __init__(self): super().__init__() self.lm_head nn.Dense(512, 32000) self.lm_head.attr {group: head} class LLModel(nn.Cell): def __init__(self): super().__init__() self.backbone LLMBackbone() self.head LLMHead() def construct(self, x): emb self.backbone.embedding(x) fea self.backbone.transformer_layer(emb) logits self.head.lm_head(fea) return logits # ----------------属性过滤函数---------------- def filter_params_by_attr(network: nn.Cell, target_group: str): 根据自定义attr属性筛选参数 selected_params [] for cell in network.cells(): if hasattr(cell, attr) and cell.attr.get(group) target_group: for param in cell.trainable_params(): selected_params.append(param) return selected_params if __name__ __main__: model LLModel() # 筛选head参数只训练headbackbone冻结 train_params filter_params_by_attr(model, target_grouphead) optimizer nn.Adam(train_params, learning_rate1e-4) # 冻结其余参数 for param in model.trainable_params(): if param not in train_params: param.requires_grad False代码说明给不同网络模块绑定自定义属性通过属性过滤快速划分训练 / 冻结参数相比字符串匹配参数名属性标记可读性更强适配大模型复杂层级结构。三、场景 2Checkpoint 权重加载属性过滤加载预训练权重时过滤指定属性、精度、名称的权重跳过不匹配参数常用于增量预训练、迁移学习。import mindspore as ms def load_ckpt_with_filter(ckpt_path, network, dtype_filter: ms.Type None): 加载权重支持属性过滤 :param ckpt_path: checkpoint文件路径 :param network: 目标网络 :param dtype_filter: 过滤指定数据类型参数 param_dict ms.load_checkpoint(ckpt_path) filtered_param {} for name, tensor in param_dict.items(): # 条件1精度属性过滤 if dtype_filter is not None and tensor.dtype ! dtype_filter: continue # 条件2过滤Embedding层参数示例 if embedding in name: continue filtered_param[name] tensor # 加载过滤后的权重 ms.load_param_into_net(network, filtered_param, strict_loadFalse) print(fFiltered ckpt params, remain {len(filtered_param)} params) # 使用示例 # load_ckpt_with_filter(pretrain.ckpt, model, dtype_filterms.float32)四、场景 3数据集样本属性过滤SFT 指令微调指令微调数据集每条样本携带属性难度、领域、是否有效通过dataset.filter实现属性过滤清洗脏数据。import mindspore.dataset as ds import numpy as np # 模拟数据集每条样本包含text、label、attr字典 class SFTDataSet: def __init__(self): self.data [ {text:指令1,label:回答1,attr:{domain:general,valid:True}}, {text:指令2,label:回答2,attr:{domain:finance,valid:False}}, {text:指令3,label:回答3,attr:{domain:general,valid:True}}, ] def __getitem__(self, idx): item self.data[idx] return item[text], item[label], item[attr] def __len__(self): return len(self.data) # 属性过滤回调函数 def filter_func(text, label, attr): # 过滤条件有效样本 通用领域 return attr[valid] and attr[domain] general if __name__ __main__: dataset ds.GeneratorDataset(SFTDataSet(), column_names[text,label,attr]) # 属性过滤 filtered_ds dataset.filter(predicatefilter_func) print(after filter data count:, filtered_ds.get_dataset_size())适用场景SFT 数据清洗、领域自适应训练快速筛选指定领域样本。五、场景 4高阶通用封装统一属性过滤工具类工程化封装同时支持网络参数、权重字典过滤统一接口class AttrFilter: staticmethod def filter_trainable_params(net:nn.Cell, filter_func): 通用参数过滤 filter_func(param) - bool return [p for p in net.trainable_params() if filter_func(p)] # 使用示例过滤FP16参数 if __name__ __main__: model LLModel() fp16_params AttrFilter.filter_trainable_params( model, lambda p: p.dtype ms.float16 )六、工程优化要点优先自定义 attr 属性避免硬编码参数名大模型参数名称易随版本改动自定义模块属性更稳定过滤操作放置在图编译前图模式下不要在 construct 内部动态过滤参数会引发编译异常Checkpoint 过滤开启 strict_loadFalse过滤后参数不完整关闭严格加载避免报错大规模数据集过滤启用多线程dataset.filter(num_parallel_workers4)提升数据管道处理速度多维并行场景下参数过滤全局同步MindSpore 自动并行场景过滤参数需要所有 rank 保持一致防止通信异常。七、总结属性过滤是 MindSpore 大模型微调、数据预处理、权重加载的通用基础能力。依靠网络 Cell 自定义属性、Parameter 固有属性、数据集样本标签、Checkpoint 张量信息可以灵活实现参数筛选、权重按需加载、训练样本清洗。本文代码覆盖三大高频场景基于模块属性冻结网络层、Checkpoint 权重条件过滤、SFT 数据集属性筛选。在大模型轻量化微调、领域适配、增量预训练场景中属性过滤能够精准控制训练范围减少无效算力消耗降低显存占用。相比于硬编码字符串匹配参数名基于属性标记的过滤方案具备更好的可维护性适配 MindSpore 大模型 MindFormers 训练生态可直接迁移至 LLaMA、Qwen 等主流大模型微调工程。
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑