拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

昇思 MindSpore 大模型:属性过滤

昇思 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(mode=ms.GRAPH_MODE, device_target="Ascend") class LLMBackbone(nn.Cell): def __init__(self): super().__init__() self.embedding = nn.Embedding(vocab_size=32000, embedding_size=512) 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参数,只训练head,backbone冻结 train_params = filter_params_by_attr(model, target_group="head") optimizer = nn.Adam(train_params, learning_rate=1e-4) # 冻结其余参数 for param in model.trainable_params(): if param not in train_params: param.requires_grad = False

代码说明:给不同网络模块绑定自定义属性,通过属性过滤快速划分训练 / 冻结参数;相比字符串匹配参数名,属性标记可读性更强,适配大模型复杂层级结构。

三、场景 2:Checkpoint 权重加载属性过滤

加载预训练权重时,过滤指定属性、精度、名称的权重,跳过不匹配参数,常用于增量预训练、迁移学习。

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_load=False) print(f"Filtered ckpt params, remain {len(filtered_param)} params") # 使用示例 # load_ckpt_with_filter("pretrain.ckpt", model, dtype_filter=ms.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(predicate=filter_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_load=False

过滤后参数不完整,关闭严格加载,避免报错;

大规模数据集过滤启用多线程

dataset.filter(num_parallel_workers=4)提升数据管道处理速度;

多维并行场景下参数过滤全局同步

MindSpore 自动并行场景,过滤参数需要所有 rank 保持一致,防止通信异常。

七、总结

属性过滤是 MindSpore 大模型微调、数据预处理、权重加载的通用基础能力。依靠网络 Cell 自定义属性、Parameter 固有属性、数据集样本标签、Checkpoint 张量信息,可以灵活实现参数筛选、权重按需加载、训练样本清洗。

本文代码覆盖三大高频场景:基于模块属性冻结网络层、Checkpoint 权重条件过滤、SFT 数据集属性筛选。在大模型轻量化微调、领域适配、增量预训练场景中,属性过滤能够精准控制训练范围,减少无效算力消耗,降低显存占用。

相比于硬编码字符串匹配参数名,基于属性标记的过滤方案具备更好的可维护性,适配 MindSpore 大模型 MindFormers 训练生态,可直接迁移至 LLaMA、Qwen 等主流大模型微调工程。

返回列表