新闻详情

新闻详情

首页 / 资讯中心 / 详情

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

发布时间:2026/9/18 17:36:03来源:尧图网络
Argilla 中基于 PEFT(LoRA)的 Token 分类微调实战:从 ArgillaTrainer 到参数配置全解析
Argilla 中基于 PEFTLoRA的 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 PEFTParameter Efficient Fine-Tuning库的 LoRALow Rank Adaptation实现对 Argilla 中标注好的 Token 分类数据集进行参数高效微调。读完本文你将掌握ArgillaTrainer(frameworkpeft)的完整训练闭环数据加载 → LoRA 配置 → 训练 → 预测并理解LoraConfig、AutoModelForTokenClassification、TrainingArguments三组update_config参数的底层映射原理与默认值来源。一、PEFT 与 LoRA为什么在 Argilla 中微调 Token 分类模型Token 分类Token Classification是命名实体识别NER、词性标注等任务的基础范式通常需要对预训练 Transformer 模型做下游微调。传统全量微调Full Fine-Tuning会为每个下游任务复制一份完整模型权重显存与存储开销巨大而 PEFTParameter Efficient Fine-Tuning只训练极少量额外参数冻结主干网络即可获得接近全量微调的效果。其中 LoRALow Rank Adaptation是 PEFT 中最常用的实现它在冻结的权重矩阵旁注入低秩分解的可训练矩阵秩为r从而把可训练参数量降到极低水平。在 Argilla 中frameworkpeft正是围绕 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( namemy_dataset_name, workspacemy_workspace_name, frameworkpeft, train_size0.8 ) trainer.update_config(lora_alpha8, num_train_epochs3) trainer.train(output_dirtoken-classification) records trainer.predict(The ArgillaTrainer is great!, as_argilla_recordsTrue)这段代码背后的执行链路在源码中非常清晰数据加载与任务识别在 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_size0.8它表示 80% 数据用于训练、20% 用于验证。从 base.py 可以看到只要传入train_size就会触发内部 train/test 切分self._split_applied True这会影响后续评估策略的自动选择详见第五节。三、三组update_config参数LoRA、模型加载与训练超参官方片段的核心价值在于完整列出了 PEFT 框架可用的update_config参数共分三组。这一设计的实现原理是update_config(**kwargs)会把关键字参数通过filter_allowed_argsargilla-v1/src/argilla_v1/training/utils.py按目标函数的形参名白名单过滤后分别写入lora_kwargs、model_kwargs、trainer_kwargs三个字典。也就是说同一个方法参数名决定归属写错参数名会被静默过滤不报错但也不会生效因此务必对照下列清单。3.1peft.LoraConfigLoRA 适配器配置# peft.LoraConfig trainer.update_config( r8, target_modulesNone, lora_alpha16, lora_dropout0.1, fan_in_fan_outFalse, biasnone, inference_modeFalse, modules_to_saveNone, init_lora_weightsTrue )这些参数直接对应 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 类模型为Truebiasnone偏置项训练策略none全部冻结、all或lora_onlyinference_modeFalse是否以推理模式构建适配器modules_to_saveNone除 LoRA 外还需完整微调并保存的模块列表如新增的分类头init_lora_weightsTrue是否使用高斯分布初始化 LoRA 权重PEFT 官方推荐的初始化方式注意lora_alpha与lora_dropout在LoraConfig默认值里分别是16与0.1而官方片段第 24 行trainer.update_config(lora_alpha8, num_train_epochs3)展示了如何在训练前快速覆盖单项配置——lora_alpha8配合默认r8时缩放因子为 1。此外还有一个不可手动覆盖、由源码自动注入的关键字段task_type在 peft.py 的init_model中ArgillaPeftTrainer会根据记录类型自动设置task_type——TextClassificationRecord对应SEQ_CLSTokenClassificationRecord对应TOKEN_CLS本主题Text2TextRecord则暂不支持并抛出NotImplementedError。3.2transformers.AutoModelForTokenClassification预训练模型加载官方片段第二组参数以AutoModelForTextClassification命名其本质是AutoModelForTokenClassification.from_pretrained(...)的入参白名单在 Token 分类场景下本主题底层模型类由 transformers.py 中的_model_class AutoModelForTokenClassification决定# transformers.AutoModelForTokenClassificationfrom_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-casedforce_downloadFalse是否忽略缓存强制重新下载resume_downloadFalse是否续传不完整的下载文件proxiesNone代理设置字典tokenNoneHugging Face Hub 访问令牌私有模型需要cache_dirNone模型缓存目录local_files_onlyFalse是否仅使用本地文件、禁止联网下载除上述可配置项外init_training_argstransformers.py还会自动注入三个与任务强相关的字段num_labelslen(self._label_list)、id2labelself._id2label、label2idself._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_strategyno有验证集时为epoch。同时默认logging_steps1、num_train_epochs1transformers.py。片段中显式给出num_train_epochs3、seed42等即为覆盖这些默认值的典型用法。设备自动探测训练前会根据torch.backends.mps.is_available()与torch.cuda.is_available()自动选择cpu/mps/cudatransformers.py并在train()时通过no_cuda/use_mps_device同步到TrainingArgumentstransformers.py。四、train / predict / save源码级的完整调用链ArgillaTrainer的公开方法只是薄封装真正逻辑都在各框架训练器中train(output_dir)transformers.py先init_model(newTrue)初始化/加载 LoRA 模型再preprocess_datasets()做分词与标签对齐Token 分类使用is_split_into_wordsTrue将非首子词标签置为-100随后构造Trainer并train()若有验证集则调用evaluate()并打印指标最后save(output_dir)并初始化推理 pipeline。predict(text, as_argilla_recordsTrue)peft.pyPEFT 训练器的预测不走 Transformers pipeline而是自行完成分词时开启return_offsets_mappingTrue拿到字符偏移对 logits 做softmax与argmax遍历预测结果遇到非O标签会剥掉B-/I-前缀并把连续的I-实体片段合并为单个实体其 score 取片段内各 token 得分的均值最后若as_argilla_recordsTrue包装成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.9ArgillaPeftTrainer在模块导入时即做版本检查低于 3.9 会直接抛出异常peft.py。必需依赖ArgillaTransformersTrainer在初始化时调用require_dependencies([torch, datasets, transformers, evaluate, seqeval])transformers.pyPEFT 训练器在此基础上额外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(fDataset {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按数据集nameworkspace加载记录的用法而 fine_tune.md 主文档展示了基于FeedbackDatasetTrainingTask的新式用法——两者 API 形态略有差异但update_config/train/predict的核心工作流与本节参数清单是相通的。小结在 Argilla 中使用frameworkpeft微调 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),仅供参考
网站建设高端定制企业官网
RELATED

相关资讯

更多精彩内容,欢迎继续阅读

较早相关资讯

最新相关资讯

微信小程序连续扫码实战:camera组件避坑与性能优化指南 2026/9/18 18:12:10

微信小程序连续扫码实战:camera组件避坑与性能优化指南

1. 从一个真实需求说起:为什么要死磕连续扫码去年接了一个仓储盘点的小程序项目,需求方开口第一句话就是:“我要能一直扫,扫完一个自动接着扫下一个,中间不要让我点任何按钮。”听起来很简单对吧?微信小程序…

阅读更多 →
SCA落地实战:从Gitee集成到选型框架与依赖治理 2026/9/18 18:12:10

SCA落地实战:从Gitee集成到选型框架与依赖治理

1. 先搞清楚:SCA 到底解决什么问题1.1 SCA 的核心价值:不只是查漏洞很多人一听到软件成分分析(Software Composition Analysis,简称 SCA)就以为是“扫描依赖里有没有已知 CVE”。这个理解不能说错,但太窄了…

阅读更多 →
CANN ops-nn 算子开发指南:aclnnHardsigmoid 与 aclnnInplaceHardsigmoid 两段式接口详解 2026/9/18 18:12:10

CANN ops-nn 算子开发指南:aclnnHardsigmoid 与 aclnnInplaceHardsigmoid 两段式接口详解

CANN ops-nn 算子开发指南:aclnnHardsigmoid 与 aclnnInplaceHardsigmoid 两段式接口详解 【免费下载链接】ops-nn 本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。 项目地址: https://gitcode.com/cann/ops-nn HardSigmoid 是…

阅读更多 →
vllm-omni serve 命令详解:基于 Stage 的分进程部署与参数配置指南 2026/9/18 18:12:10

vllm-omni serve 命令详解:基于 Stage 的分进程部署与参数配置指南

vllm-omni serve 命令详解:基于 Stage 的分进程部署与参数配置指南 【免费下载链接】vllm-omni A framework for efficient model inference with omni-modality models 项目地址: https://gitcode.com/GitHub_Trending/vl/vllm-omni 导读 vllm serve 是 vL…

阅读更多 →
Visual Studio中C++多源文件独立运行的三种实操方案 2026/9/18 18:12:10

Visual Studio中C++多源文件独立运行的三种实操方案

1. 项目概述:为什么“多个源文件分开运行”是个伪命题,但却是新手最真实的痛点在 Visual Studio(VS)里点开一个 C 项目,看到七八个.cpp文件堆在解决方案资源管理器里,心里就发毛:“我改了main.c…

阅读更多 →
SaaS订阅订单管理实战:OPM体系如何与网络及服务器协同 2026/9/18 18:09:10

SaaS订阅订单管理实战:OPM体系如何与网络及服务器协同

做SaaS商业化这些年,订阅订单的管理一直是个被低估的环节。很多人觉得订单嘛,支付成功就算完事了,数据同步了,剩下的交给财务开票就行。但真正跑起来才发现,订阅制订单和传统买断制完全是两种生物——买断制是一次性交…

阅读更多 →

今日资讯

本周资讯

本月资讯

看完文章仍有疑问?

联系尧图顾问,获取一对一建站咨询

立即免费咨询 📞 400-888-8888
📞