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 配置 → 训练 → 预测),并理解LoraConfig、AutoModelForTokenClassification、TrainingArguments三组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支持transformers、setfit、spacy、peft、span_marker、trl、openai等框架;其中 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)这段代码背后的执行链路在源码中非常清晰:
- 数据加载与任务识别:在 argilla-v1/src/argilla_v1/training/base.py 中,
ArgillaTrainer.__init__会根据name/workspace调用active_client(),先加载 1 条记录快照以自动识别数据集类型(DatasetForTextClassification、DatasetForTokenClassification或DatasetForText2Text),再调用prepare_for_training(framework=..., settings=..., train_size=..., seed=...)完成数据切分与格式化。 - 训练器分发:当
framework is Framework.PEFT时(base.py),内部实例化ArgillaPeftTrainer,并传入record_class、已准备的dataset、multi_label、settings、seed、model等上下文。 - 训练与预测:
update_config、train、predict都是对内部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_kwargs、model_kwargs、trainer_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中被硬编码,含义如下:
| 参数 | 默认值 | 作用 |
|---|---|---|
r | 8 | LoRA 低秩矩阵的秩,控制可训练参数规模;r越大适配能力越强,但参数量与过拟合风险也上升 |
target_modules | None | 指定注入 LoRA 的模块(如["q_lin", "v_lin"]);None时由 PEFT 按模型类型自动推断 |
lora_alpha | 16 | LoRA 缩放因子,实际缩放为lora_alpha / r,影响适配器更新幅度 |
lora_dropout | 0.1 | LoRA 层的 dropout 概率,用于正则化 |
fan_in_fan_out | False | 权重矩阵是否以 fan_in/fan_out 方式存储(部分 GPT 类模型为True) |
bias | "none" | 偏置项训练策略:"none"(全部冻结)、"all"或"lora_only" |
inference_mode | False | 是否以推理模式构建适配器 |
modules_to_save | None | 除 LoRA 外还需完整微调并保存的模块列表(如新增的分类头) |
init_lora_weights | True | 是否使用高斯分布初始化 LoRA 权重(PEFT 官方推荐的初始化方式) |
注意:lora_alpha与lora_dropout在LoraConfig默认值里分别是16与0.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_path | 由model决定 | 预训练模型 ID 或本地目录;ArgillaTrainer未显式传model时,transformers.py 会回退到默认的"bert-base-cased" |
force_download | False | 是否忽略缓存强制重新下载 |
resume_download | False | 是否续传不完整的下载文件 |
proxies | None | 代理设置字典 |
token | None | Hugging Face Hub 访问令牌(私有模型需要) |
cache_dir | None | 模型缓存目录 |
local_files_only | False | 是否仅使用本地文件、禁止联网下载 |
除上述可配置项外,init_training_args(transformers.py)还会自动注入三个与任务强相关的字段:num_labels=len(self._label_list)、id2label=self._id2label、label2id=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.Trainer的TrainingArguments。需要说明三点:
- 参数来源:
trainer_kwargs初始值由get_default_args(TrainingArguments.__init__)通过内省inspect.getfullargspec自动抓取(utils.py),因此上面的数值本质上是对 Transformers 库默认值的显式复述;传入的自定义值会覆盖对应默认值。 - 自动调整的默认项:当没有
train_size(即无验证集)时,evaluation_strategy="no";有验证集时为"epoch"。同时默认logging_steps=1、num_train_epochs=1(transformers.py)。片段中显式给出num_train_epochs=3、seed=42等即为覆盖这些默认值的典型用法。 - 设备自动探测:训练前会根据
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),随后构造Trainer并train();若有验证集则调用evaluate()并打印指标;最后save(output_dir)并初始化推理 pipeline。predict(text, as_argilla_records=True)(peft.py):PEFT 训练器的预测不走 Transformers pipeline,而是自行完成:分词时开启return_offsets_mapping=True拿到字符偏移;对 logits 做softmax与argmax;遍历预测结果,遇到非"O"标签会剥掉B-/I-前缀,并把连续的I-实体片段合并为单个实体,其 score 取片段内各 token 得分的均值;最后若as_argilla_records=True,包装成TokenClassificationRecord(包含entity_group、score、word、start、end字符级跨度)。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填充标签后计算并打印precision、recall、f1、accuracy四项总体指标(overall_*)。因此建议在ArgillaTrainer中始终保留train_size(如0.8),以获得可复现、可对比的验证集评估结果。
六、运行环境与前置依赖
结合仓库代码,PEFT 训练链路对运行环境有以下硬性要求:
- Python ≥ 3.9:
ArgillaPeftTrainer在模块导入时即做版本检查,低于 3.9 会直接抛出异常(peft.py)。 - 必需依赖:
ArgillaTransformersTrainer在初始化时调用require_dependencies(["torch", "datasets", "transformers", "evaluate", "seqeval"])(transformers.py),PEFT 训练器在此基础上额外require_dependencies("peft")(peft.py)。因此至少需要安装peft、torch、transformers、datasets、evaluate、seqeval。 - 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.training(argilla-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),仅供参考