ARTICLE DETAIL

资讯详情

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

使用 Argilla 的 ArgillaPeftTrainer 进行 LoRA 文本分类微调:PEFT 框架完整实践指南

使用 Argilla 的 ArgillaPeftTrainer 进行 LoRA 文本分类微调:PEFT 框架完整实践指南 使用 Argilla 的 ArgillaPeftTrainer 进行 LoRA 文本分类微调PEFT 框架完整实践指南【免费下载链接】argillaArgilla is a collaboration tool for AI engineers and domain experts to build high-quality datasets项目地址: https://gitcode.com/GitHub_Trending/ar/argilla本指南以 Argilla 官方文档中 PEFT 训练代码片段peft.md为核心讲解如何通过ArgillaTrainer的peft框架在 Argilla 标注数据集之上使用 Hugging Face PEFT 库的低秩适配LoRA技术微调文本分类模型。读完本文你将掌握从数据加载、LoRA 配置、训练参数调优到推理回写 Argilla 记录的完整实战流程并理解其底层实现原理。ArgillaTrainer 与 PEFT 框架概览ArgillaTrainer是 Argilla v1 提供的一个高阶训练包装器内部封装了众多主流 NLP 训练库。它以统一、直观的 API 屏蔽了从 Argilla 数据集到各框架数据格式之间的转换细节让你可以用几乎相同的代码在不同框架Transformers、SetFit、spaCy、PEFT、TRL 等之间切换训练。从源码看ArgillaTrainer定义于 argilla-v1/src/argilla_v1/training/base.py其构造参数包括参数说明name要加载的 Argilla 数据集名称必填framework训练框架可选transformers、setfit、spacy、spacy-transformers、peft、openai、trl、span_marker、sentence-transformers等workspace数据集所在工作区默认使用当前用户的工作区model基线模型名称或路径未指定时各框架使用默认模型Transformers/PEFT 系默认bert-base-casedtrain_size训练集占比其余部分作为验证集不指定则全量用于训练seed随机种子用于保证数据切分与训练的可复现性gpu_idGPU ID默认-1使用 CPUframeworkpeft对应框架枚举Framework.PEFT定义于 argilla-v1/src/argilla_v1/client/models.py在ArgillaTrainer.__init__中会被路由到ArgillaPeftTrainer见 base.py。ArgillaPeftTrainer的核心设计是继承ArgillaTransformersTrainer的全部基础能力数据预处理、Trainer 训练循环、评估指标、Pipeline 推理仅在模型初始化层叠加 PEFT 的 LoRA 适配器。其类定义位于 argilla-v1/src/argilla_v1/training/peft.py。快速上手三行代码完成 LoRA 微调以下是官方文档给出的最小可用示例原文完整保留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_dirtext-classification) records trainer.predict(The ArgillaTrainer is great!, as_argilla_recordsTrue)执行流程拆解加载数据ArgillaTrainer初始化时会通过active_client().load(name..., limit1)读取数据集的快照自动判断数据集类型DatasetForTextClassification/DatasetForTokenClassification/DatasetForText2Text并识别是否为多标签multi_label任务随后用prepare_for_training按train_size0.8和seed将数据切分为训练集与验证集见 base.py。配置参数trainer.update_config(lora_alpha8, num_train_epochs3)同时向LoraConfiglora_alpha和TrainingArgumentsnum_train_epochs两个命名空间写入参数——这是ArgillaTrainer.update_config的设计特色它会把**kwargs按各框架构造函数的签名自动过滤分发详见下文配置项详解。训练并保存trainer.train(output_dirtext-classification)执行微调并把模型与分词器保存到text-classification目录。推理回写trainer.predict(..., as_argilla_recordsTrue)返回TextClassificationRecord列表可直接用于 Argilla 的评审或导出流程传单条字符串时返回单条记录源码见 peft.py。需要特别说明的是该示例基于 Argilla v1 的 API训练模块位于 argilla-v1/src/argilla_v1/training。若使用 FeedbackDataset 的新版工作流可参考 fine_tune.md 中TrainingTaskArgillaTrainer的用法底层同样会路由到本指南介绍的 PEFT 训练器。update_config 详解三组可调参数ArgillaTrainer.update_config(**kwargs)是调整训练细节的统一入口base.py它会将参数委托给ArgillaPeftTrainer.update_config。其内部实现是先调用父类ArgillaTransformersTrainer.update_config把属于transformers.TrainingArguments的参数写入trainer_kwargs再调用filter_allowed_args(LoraConfig.__init__, **kwargs)把属于LoraConfig的参数写入lora_kwargs见 peft.py。filter_allowed_args定义于 argilla-v1/src/argilla_v1/training/utils.py通过检查func.__code__.co_varnames实现参数白名单过滤——只有目标函数签名中存在的参数才会被接受传入无关关键字不会报错但会被静默丢弃。这与文档中无需一次性传入全部参数未传入时使用默认配置的说明一致。第一组peft.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 )参数默认值作用r8LoRA 低秩矩阵的秩决定新增可训练参数的规模增大可提升模型容量但会带来更多显存与过拟合风险target_modulesNone要注入 LoRA 适配器的模块名列表None时由 PEFT 按任务类型自动选择如注意力层的q_proj、v_projlora_alpha16LoRA 缩放因子最终适配权重按alpha / r缩放控制更新幅度lora_dropout0.1LoRA 层丢弃率用于正则化fan_in_fan_outFalse权重矩阵存储方式标记部分模型如 GPT-2 的 Conv1D需设为Truebiasnone偏置项训练策略可选none/all/lora_onlyinference_modeFalse是否以推理模式创建适配器modules_to_saveNone除 LoRA 层外还需完整微调并保存的模块如分类头init_lora_weightsTrue是否使用 LoRA 论文中的随机初始化策略初始化适配器权重以上默认值在源码 peft.py 的init_training_args中逐项硬编码可直接与文档对照。第二组transformers.AutoModelForTextClassification模型加载参数# transformers.AutoModelForTextClassification 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 )这些参数最终写入model_kwargs在init_model时以self._model_class.from_pretrained(**self.model_kwargs, return_dictTrue)的形式传给AutoModelForSequenceClassification.from_pretrained见 peft.py。pretrained_model_name_or_path基础模型标识Hugging Face Hub 模型 ID 或本地目录。ArgillaTrainer初始化时若未显式传model默认使用bert-base-cased见 transformers.py。force_download/resume_download是否强制重新下载 / 断点续传模型权重。proxies/token网络代理配置与 Hub 访问令牌。cache_dir模型缓存目录。local_files_only是否仅使用本地缓存不访问网络。注意model_kwargs中还会被自动填充num_labels、id2label、label2id以及多标签时的problem_typemulti_label_classification这些来自数据集设置无需手动配置见 transformers.py。第三组transformers.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, learning_rate 5e-5, 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 )参数默认值说明per_device_train_batch_size/per_device_eval_batch_size8每设备训练 / 评估批量大小gradient_accumulation_steps1梯度累积步数等效扩大批量learning_rate5e-5学习率weight_decay0权重衰减adam_beta1/adam_beta2/adam_epsilon0.9/0.9/1e-8Adam 优化器超参数文档示例中两处learning_rate重复属笔误实际以最后一次为准max_grad_norm1梯度裁剪阈值num_train_epochs3训练轮数源码默认1见 transformers.py文档示例通过update_config覆盖为3max_steps0最大训练步数0表示不限制log_levelpassive日志级别logging_strategy/save_strategysteps日志 / 保存策略save_steps500每 N 步保存一次 checkpointseed42随机种子push_to_hubFalse是否在训练时推送模型到 Hubhub_model_iduser_name/output_dir_nameHub 目标仓库 IDhub_strategyevery_saveHub 推送策略hub_token1234Hub 访问令牌示例值请替换为你自己的 tokenhub_private_repoFalse是否创建私有仓库trainer_kwargs的初始值由get_default_args(TrainingArguments.__init__)从transformers反射得到transformers.py因此即便不调用update_config训练参数也会带有 Transformers 库的全部默认值train时还会自动写入output_dir并根据是否有验证集把evaluation_strategy设为epoch或no。源码级原理ArgillaPeftTrainer 如何组装 LoRA 模型ArgillaPeftTrainer.init_modelpeft.py是整个 PEFT 集成的核心其逻辑分为三步第一步确定任务类型。根据记录类型设置task_typeif self._record_class TextClassificationRecord: self.lora_kwargs[task_type] SEQ_CLS elif self._record_class TokenClassificationRecord: self.lora_kwargs[task_type] TOKEN_CLS else: raise NotImplementedError(rg.Text2TextRecord is not supported yet.)即文本分类走SEQ_CLS序列分类命名实体识别等 token 分类走TOKEN_CLS目前不支持 Text2Text 记录。因此本框架的实际应用范围是文本分类与 token 分类两类任务这与 fine_tune.md 中PEFT 支持 Text Classification的支持矩阵一致。第二步加载模型——优先恢复已训练好的 PEFT 模型。代码先尝试PeftConfig.from_pretrained(...)读取配置如果pretrained_model_name_or_path指向一个已保存的 PEFT 模型则加载其base_model_name_or_path作为基础模型再以PeftModel.from_pretrained挂载适配器从而支持加载上次微调结果继续训练若读取失败例如是普通预训练模型 ID则回退为LoraConfig(**self.lora_kwargs)构建 LoRA 配置加载基础模型后调用get_peft_model(model, config)原地注入适配器。第三步初始化分词器。使用AutoTokenizer.from_pretrained(config.base_model_name_or_path, add_prefix_spaceTrue)加载对应分词器并将模型移动到self.devicecuda/mps/cpu由 transformers.py 按可用性自动选择。训练、推理与保存的完整工作流训练trainer.train(output_dir)train方法继承自ArgillaTransformersTrainertransformers.py关键步骤包括把output_dir写入trainer_kwargsinit_model(newTrue)重建带 LoRA 适配器的模型走 PEFT 分支preprocess_datasets()按任务类型执行 tokenization——文本分类用DataCollatorWithPaddingtoken 分类用DataCollatorForTokenClassification并对非首个子词标签置-100compute_metrics()装配评估指标单标签文本分类用 accuracy多标签用 micro-F1 及逐标签 F1token 分类用seqeval的整体 precision / recall / F1 / accuracy构建transformers.Trainer并train()存在验证集时调用evaluate()并打印指标自动save(output_dir)后初始化推理 Pipeline。推理trainer.predict(text, as_argilla_recordsTrue)ArgillaPeftTrainer.predictpeft.py不依赖pipeline而是直接使用 tokenizer 模型前向传播文本分类对输入做截断与最长 paddingtorch.no_grad()下取 logits多标签任务用sigmoid单标签任务用softmax得到各标签概率。token 分类借助return_offsets_mappingTrue获取字符偏移将B-/I-前缀的实体片段合并为完整实体取各 token 分数的均值作为实体分数输出包含entity_group/score/word/start/end的结构化结果。若as_argilla_recordsTrue结果包装为TextClassificationRecord携带prediction与multi_label或TokenClassificationRecord携带tokens与字符级prediction可直接落回 Argilla。传入单个字符串时返回单条记录传入列表时返回记录列表。保存与复用trainer.save(output_dir)def save(self, output_dir: str): self.trainer_model.save_pretrained(output_dir) self.trainer_tokenizer.save_pretrained(output_dir)模型与分词器均以save_pretrained保存到同一目录peft.py。由于保存的是完整 PEFT 模型含 LoRA 适配器权重后续可通过init_model中PeftConfig.from_pretrainedPeftModel.from_pretrained的分支无缝加载续训。也可以调用trainer.get_trainer_model()、get_trainer_tokenizer()、get_trainer()等访问器直接取得底层对象base.py。环境要求与注意事项Python 版本ArgillaPeftTrainer在导入时即检查sys.version_info (3, 9)低于 3.9 会直接抛出异常Must be using Python 3.9 or higher or PEFT wont workpeft.py。依赖项init_training_args会调用require_dependencies(peft)检查peft是否已安装父类还要求torch、datasets、transformers、evaluate、seqevaltransformers.py。启动训练前请确认这些包已就绪例如pip install peft transformers datasets evaluate seqeval torch。默认模型未指定model时默认使用bert-base-cased首次运行会自动从 Hugging Face Hub 下载也可以像第二组参数那样通过pretrained_model_name_or_path换成distilbert-base-uncased等更轻量的模型。数据集要求ArgillaTrainer要求目标数据集非空否则抛出ValueError: Dataset {name} is emptybase.py数据集的标签设置label2id/id2label会自动从 Argilla 加载无需手工维护映射。参数分发机制update_config基于函数签名白名单过滤参数传入与目标框架签名不匹配的关键字会被静默忽略因此调试时可利用trainer.__repr__()打印当前model_kwargs、trainer_kwargs与lora_kwargs的实际生效值peft.py。综上借助 Argilla 的peft框架你可以在不改动任何数据预处理代码的前提下用最少的参数默认仅r8、lora_alpha16、lora_dropout0.1完成 LoRA 微调将标注数据快速转化为可部署的分类模型并通过as_argilla_recordsTrue让模型预测无缝回流到 Argilla 标注流水线形成标注 → 微调 → 推理 → 再标注的闭环。【免费下载链接】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),仅供参考
返回列表