Argilla 中基于 PEFT(LoRA)的 Token 分类微调实战:从 ArgillaTrainer 到参数配置全解析
2026/9/18 17:36:03 网站建设 项目流程

Argilla 中基于 PEFT(LoRA)的 Token 分类微调实战:从 ArgillaTrainer 到参数配置全解析

【免费下载链接】argillaArgilla is a collaboration tool for AI engineers and domain experts to build high-quality datasets项目地址: https://gitcode.com/GitHub_Trending/ar/argilla

本文以 Argilla 官方文档片段(docs/_source/_common/snippets/training/token-classification/peft.md)为核心骨架,完整讲解如何用ArgillaTrainer配合 Hugging Face PEFT(Parameter Efficient Fine-Tuning)库的 LoRA(Low Rank Adaptation)实现,对 Argilla 中标注好的 Token 分类数据集进行参数高效微调。读完本文,你将掌握ArgillaTrainer(framework="peft")的完整训练闭环(数据加载 → LoRA 配置 → 训练 → 预测),并理解LoraConfigAutoModelForTokenClassificationTrainingArguments三组update_config参数的底层映射原理与默认值来源。

一、PEFT 与 LoRA:为什么在 Argilla 中微调 Token 分类模型

Token 分类(Token Classification)是命名实体识别(NER)、词性标注等任务的基础范式,通常需要对预训练 Transformer 模型做下游微调。传统全量微调(Full Fine-Tuning)会为每个下游任务复制一份完整模型权重,显存与存储开销巨大;而 PEFT(Parameter Efficient Fine-Tuning)只训练极少量额外参数,冻结主干网络,即可获得接近全量微调的效果。其中 LoRA(Low Rank Adaptation)是 PEFT 中最常用的实现:它在冻结的权重矩阵旁注入低秩分解的可训练矩阵(秩为r),从而把可训练参数量降到极低水平。

在 Argilla 中,framework="peft"正是围绕 LoRA 实现的训练框架入口。在 argilla-v1/src/argilla_v1/client/models.py 的Framework枚举中,PEFT = "peft"被映射为 "PEFT Transformers library",并且ArgillaTrainer支持transformerssetfitspacypeftspan_markertrlopenai等框架;其中 PEFT 与transformers共享底层微调逻辑(见下文源码分析),区别在于额外注入 LoRA 适配器层。

二、最小可用示例:三行代码跑通 LoRA 微调

官方片段给出了 PEFT 框架下最精简的完整流程。它假设 Argilla 中已存在一个 Token 分类数据集(包含ner_tags标注),通过数据集名称与工作区即可直接拉起训练:

from argilla.training import ArgillaTrainer trainer = ArgillaTrainer( name="<my_dataset_name>", workspace="<my_workspace_name>", framework="peft", train_size=0.8 ) trainer.update_config(lora_alpha=8, num_train_epochs=3) trainer.train(output_dir="token-classification") records = trainer.predict("The ArgillaTrainer is great!", as_argilla_records=True)

这段代码背后的执行链路在源码中非常清晰:

  1. 数据加载与任务识别:在 argilla-v1/src/argilla_v1/training/base.py 中,ArgillaTrainer.__init__会根据name/workspace调用active_client(),先加载 1 条记录快照以自动识别数据集类型(DatasetForTextClassificationDatasetForTokenClassificationDatasetForText2Text),再调用prepare_for_training(framework=..., settings=..., train_size=..., seed=...)完成数据切分与格式化。
  2. 训练器分发:当framework is Framework.PEFT时(base.py),内部实例化ArgillaPeftTrainer,并传入record_class、已准备的datasetmulti_labelsettingsseedmodel等上下文。
  3. 训练与预测update_configtrainpredict都是对内部self._trainer的透明代理(见 base.py),因此对外 API 与transformers框架完全一致,切换框架几乎不需要改动调用代码。

关于train_size=0.8:它表示 80% 数据用于训练、20% 用于验证。从 base.py 可以看到,只要传入train_size就会触发内部 train/test 切分(self._split_applied = True),这会影响后续评估策略的自动选择(详见第五节)。

三、三组update_config参数:LoRA、模型加载与训练超参

官方片段的核心价值在于完整列出了 PEFT 框架可用的update_config参数,共分三组。这一设计的实现原理是:update_config(**kwargs)会把关键字参数通过filter_allowed_args(argilla-v1/src/argilla_v1/training/utils.py)按目标函数的形参名白名单过滤后,分别写入lora_kwargsmodel_kwargstrainer_kwargs三个字典。也就是说:同一个方法,参数名决定归属,写错参数名会被静默过滤(不报错但也不会生效),因此务必对照下列清单。

3.1peft.LoraConfig:LoRA 适配器配置

# `peft.LoraConfig` trainer.update_config( r=8, target_modules=None, lora_alpha=16, lora_dropout=0.1, fan_in_fan_out=False, bias="none", inference_mode=False, modules_to_save=None, init_lora_weights=True )

这些参数直接对应 PEFT 库的LoraConfig,其默认值在 argilla-v1/src/argilla_v1/training/peft.py 的init_training_args中被硬编码,含义如下:

参数默认值作用
r8LoRA 低秩矩阵的秩,控制可训练参数规模;r越大适配能力越强,但参数量与过拟合风险也上升
target_modulesNone指定注入 LoRA 的模块(如["q_lin", "v_lin"]);None时由 PEFT 按模型类型自动推断
lora_alpha16LoRA 缩放因子,实际缩放为lora_alpha / r,影响适配器更新幅度
lora_dropout0.1LoRA 层的 dropout 概率,用于正则化
fan_in_fan_outFalse权重矩阵是否以 fan_in/fan_out 方式存储(部分 GPT 类模型为True
bias"none"偏置项训练策略:"none"(全部冻结)、"all""lora_only"
inference_modeFalse是否以推理模式构建适配器
modules_to_saveNone除 LoRA 外还需完整微调并保存的模块列表(如新增的分类头)
init_lora_weightsTrue是否使用高斯分布初始化 LoRA 权重(PEFT 官方推荐的初始化方式)

注意lora_alphalora_dropoutLoraConfig默认值里分别是160.1,而官方片段第 24 行trainer.update_config(lora_alpha=8, num_train_epochs=3)展示了如何在训练前快速覆盖单项配置——lora_alpha=8配合默认r=8时缩放因子为 1。

此外,还有一个不可手动覆盖、由源码自动注入的关键字段task_type:在 peft.py 的init_model中,ArgillaPeftTrainer会根据记录类型自动设置task_type——TextClassificationRecord对应"SEQ_CLS"TokenClassificationRecord对应"TOKEN_CLS"(本主题),Text2TextRecord则暂不支持并抛出NotImplementedError

3.2transformers.AutoModelForTokenClassification:预训练模型加载

官方片段第二组参数以AutoModelForTextClassification命名,其本质是AutoModelForTokenClassification.from_pretrained(...)的入参白名单;在 Token 分类场景下(本主题),底层模型类由 transformers.py 中的_model_class = AutoModelForTokenClassification决定:

# `transformers.AutoModelForTokenClassification`(from_pretrained 参数) trainer.update_config( pretrained_model_name_or_path = "distilbert-base-uncased", force_download = False, resume_download = False, proxies = None, token = None, cache_dir = None, local_files_only = False )
参数默认值作用
pretrained_model_name_or_pathmodel决定预训练模型 ID 或本地目录;ArgillaTrainer未显式传model时,transformers.py 会回退到默认的"bert-base-cased"
force_downloadFalse是否忽略缓存强制重新下载
resume_downloadFalse是否续传不完整的下载文件
proxiesNone代理设置字典
tokenNoneHugging Face Hub 访问令牌(私有模型需要)
cache_dirNone模型缓存目录
local_files_onlyFalse是否仅使用本地文件、禁止联网下载

除上述可配置项外,init_training_args(transformers.py)还会自动注入三个与任务强相关的字段:num_labels=len(self._label_list)id2label=self._id2labellabel2id=self._label2id,这些来自数据集设置(settings.label2id / id2label),保证分类头维度与 Argilla 标注标签一一对应。

3.3transformers.TrainingArguments:训练超参数

# `transformers.TrainingArguments` trainer.update_config( per_device_train_batch_size = 8, per_device_eval_batch_size = 8, gradient_accumulation_steps = 1, learning_rate = 5e-5, weight_decay = 0, adam_beta1 = 0.9, adam_beta2 = 0.9, adam_epsilon = 1e-8, max_grad_norm = 1, num_train_epochs = 3, max_steps = 0, log_level = "passive", logging_strategy = "steps", save_strategy = "steps", save_steps = 500, seed = 42, push_to_hub = False, hub_model_id = "user_name/output_dir_name", hub_strategy = "every_save", hub_token = "1234", hub_private_repo = False )

这一组参数被写入trainer_kwargs并最终传给transformers.TrainerTrainingArguments。需要说明三点:

  1. 参数来源trainer_kwargs初始值由get_default_args(TrainingArguments.__init__)通过内省inspect.getfullargspec自动抓取(utils.py),因此上面的数值本质上是对 Transformers 库默认值的显式复述;传入的自定义值会覆盖对应默认值。
  2. 自动调整的默认项:当没有train_size(即无验证集)时,evaluation_strategy="no";有验证集时为"epoch"。同时默认logging_steps=1num_train_epochs=1(transformers.py)。片段中显式给出num_train_epochs=3seed=42等即为覆盖这些默认值的典型用法。
  3. 设备自动探测:训练前会根据torch.backends.mps.is_available()torch.cuda.is_available()自动选择"cpu"/"mps"/"cuda"(transformers.py),并在train()时通过no_cuda/use_mps_device同步到TrainingArguments(transformers.py)。

四、train / predict / save:源码级的完整调用链

ArgillaTrainer的公开方法只是薄封装,真正逻辑都在各框架训练器中:

  • train(output_dir)(transformers.py):先init_model(new=True)初始化/加载 LoRA 模型,再preprocess_datasets()做分词与标签对齐(Token 分类使用is_split_into_words=True,将非首子词标签置为-100),随后构造Trainertrain();若有验证集则调用evaluate()并打印指标;最后save(output_dir)并初始化推理 pipeline。
  • predict(text, as_argilla_records=True)(peft.py):PEFT 训练器的预测不走 Transformers pipeline,而是自行完成:分词时开启return_offsets_mapping=True拿到字符偏移;对 logits 做softmaxargmax;遍历预测结果,遇到非"O"标签会剥掉B-/I-前缀,并把连续的I-实体片段合并为单个实体,其 score 取片段内各 token 得分的均值;最后若as_argilla_records=True,包装成TokenClassificationRecord(包含entity_groupscorewordstartend字符级跨度)。
  • save(output_dir)(peft.py):对 LoRA 模型执行save_pretrained(output_dir)并同步保存 tokenizer,输出目录中即可直接获得可用于from_pretrained加载的适配器权重。
  • 断点续训支持init_model会先尝试用PeftConfig.from_pretrained(pretrained_model_name_or_path)加载已有 PEFT 配置,若成功则基于base_model_name_or_path重建基座模型并挂载已训练好的 LoRA 适配器;若失败则视为全新训练,用LoraConfig(**self.lora_kwargs)get_peft_model(model, config)从零注入适配器(peft.py)。这意味着把pretrained_model_name_or_path指向一个已保存的 LoRA 输出目录,即可实现增量微调。

五、评估指标:Token 分类的 seqeval 报告

当传入train_size(如0.8)时,训练结束会自动在验证集上评估。Token 分类的评估函数定义在 transformers.py:使用evaluate.load("seqeval"),在去除-100填充标签后计算并打印precisionrecallf1accuracy四项总体指标(overall_*)。因此建议在ArgillaTrainer中始终保留train_size(如0.8),以获得可复现、可对比的验证集评估结果。

六、运行环境与前置依赖

结合仓库代码,PEFT 训练链路对运行环境有以下硬性要求:

  • Python ≥ 3.9ArgillaPeftTrainer在模块导入时即做版本检查,低于 3.9 会直接抛出异常(peft.py)。
  • 必需依赖ArgillaTransformersTrainer在初始化时调用require_dependencies(["torch", "datasets", "transformers", "evaluate", "seqeval"])(transformers.py),PEFT 训练器在此基础上额外require_dependencies("peft")(peft.py)。因此至少需要安装pefttorchtransformersdatasetsevaluateseqeval
  • PyTorch MPS 回退ArgillaTrainer构造时会检查环境变量PYTORCH_ENABLE_MPS_FALLBACK,未设置则自动置为"1"并给出警告(base.py),以提升 Apple Silicon 上的兼容性。
  • 数据集非空:若目标数据集为空,ArgillaTrainer初始化会抛出ValueError(f"Dataset {self._name} is empty")(base.py)。

七、与整体微调指南的关系

本文档是 Argilla 微调指南的 Token 分类 PEFT 示例片段,更完整的背景(TrainingTask定义、FeedbackDataset数据准备、各框架支持矩阵、模型卡生成与 Hugging Face Hub 推送)可参考 docs/_source/practical_guides/fine_tune.md。需要留意的是:本文片段基于argilla.trainingargilla-v1SDK 的ArgillaTrainer,按数据集name+workspace加载记录)的用法;而 fine_tune.md 主文档展示了基于FeedbackDataset+TrainingTask的新式用法——两者 API 形态略有差异,但update_config/train/predict的核心工作流与本节参数清单是相通的。

小结

在 Argilla 中使用framework="peft"微调 Token 分类模型,本质上是"Argilla 数据层 + Transformers 训练层 + LoRA 适配层"的三层协作:数据切分与标签映射由 base.py 完成,LoRA 注入与 task_type 判定由 peft.py 完成,分词预处理、Trainer 组装与 seqeval 评估由 transformers.py 完成。掌握三组update_config参数(LoraConfig/AutoModel.from_pretrained/TrainingArguments)的归属与默认值,即可在不接触底层样板代码的前提下,快速迭代出高质量的 NER 微调模型。

【免费下载链接】argillaArgilla is a collaboration tool for AI engineers and domain experts to build high-quality datasets项目地址: https://gitcode.com/GitHub_Trending/ar/argilla

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询