ARTICLE DETAIL

资讯详情

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

Argilla 与 Hugging Face AutoTrain 结合:使用 ArgillaTrainer 零代码微调文本分类模型

Argilla 与 Hugging Face AutoTrain 结合:使用 ArgillaTrainer 零代码微调文本分类模型 Argilla 与 Hugging Face AutoTrain 结合使用 ArgillaTrainer 零代码微调文本分类模型【免费下载链接】argillaArgilla is a collaboration tool for AI engineers and domain experts to build high-quality datasets项目地址: https://gitcode.com/GitHub_Trending/ar/argillaArgilla 的ArgillaTrainer是连接 Argilla 标注数据集与各类主流 NLP 训练框架的统一训练入口。本文将聚焦frameworkautotrain这一路径讲解如何将 Argilla 数据集无缝接入 Hugging Face AutoTrain Advanced 完成文本分类模型微调涵盖环境准备、训练器初始化、update_config参数体系、底层实现原理与已知限制让你既能照抄上手也能理解其内部工作机制。一、为什么用 ArgillaTrainer 接 AutoTrainAutoTrain Advanced 是 Hugging Face 推出的自动化模型训练方案它把数据预处理、模型搜索、超参数调优、分布式训练等环节高度封装用户只需准备好数据并声明任务类型即可在 Hugging Face 平台上自动完成训练。而 Argilla 的核心价值在于高质量数据集的构建与人工标注流程。将两者结合便形成了一条完整的产线在 Argilla 中标注/校验文本分类数据 → 通过ArgillaTrainer一键把数据集格式化为 AutoTrain 所需的结构 → 在 Hugging Face AutoTrain 平台自动训练模型。整个过程无需手写数据转换代码这正是本指南所依托的 autotrain_code.md 片段的核心场景。ArgillaTrainer的通用封装思路在 base.py 中体现得十分清晰它对外提供update_config调整框架参数、train启动训练、predict推理三个统一接口对内则根据framework参数分派到不同的 Trainer 子类。AutoTrain 对应的子类ArgillaAutoTrainTrainer位于 autotrain_advanced.py。二、环境准备依赖与认证在使用 AutoTrain 框架前需要满足以下前提条件1. 安装依赖ArgillaAutoTrainTrainer在类定义处通过require_dependencies([autotrain-advanced, datasets])声明了必需依赖见 autotrain_advanced.py缺失时会直接报错提示安装pip install autotrain-advanced datasets2. 配置环境变量AutoTrain 训练是托管在 Hugging Face 平台上的因此需要认证信息。源码在类加载阶段就强制校验两个环境变量缺少任何一个都会抛出KeyErrorAUTOTRAIN_USERNAME os.environ[AUTOTRAIN_USERNAME] HF_TOKEN os.environ[HF_AUTH_TOKEN]AUTOTRAIN_USERNAME你的 Hugging Face 用户名训练项目会创建在你的账号下HF_AUTH_TOKENHugging Face 访问令牌用于平台 API 认证与模型上传。3. 数据集准备训练数据来自 Argilla 中已有的标注数据集。ArgillaTrainer初始化时会调用argilla.load(name..., workspace...)拉取数据集如果数据集为空会直接抛出ValueError。同时它会根据记录类型自动识别任务文本分类数据集TextClassificationRecord即可用于本指南的文本分类训练。三、快速上手训练一个文本分类模型以下是 autotrain_code.md 给出的完整最小示例from argilla.training import ArgillaTrainer trainer ArgillaTrainer( namemy_dataset_name, workspacemy_workspace_name, frameworkautotrain, train_size0.8 ) trainer.update_config(modelroberta-base, hub_model[{learning_rate: 0.0002}, {learning_rate: 0.0003}]) trainer.train(output_dirtext-classification) records trainer.predict(The ArgillaTrainer is great!, as_argilla_recordsTrue)逐行拆解步骤说明ArgillaTrainer(...)以数据集名称、工作区名称、框架autotrain初始化训练器train_size0.8表示 80% 数据用于训练、20% 用于验证。train_size会在内部触发prepare_for_training的 train/test 切分见 base.py 中self._split_applied True的逻辑update_config(modelroberta-base, ...)指定使用 Hugging Face Hub 上的roberta-base作为基础模型并传入两个不同的hub_model参数组learning_rate分别为 0.0002 与 0.0003表示会运行两次不同超参的调优任务train(output_dirtext-classification)提交训练任务。output_dir是训练项目的输出目录名predict(...)对输入文本做预测。注意AutoTrain 框架本身不提供本地推理predict会打印错误日志并提示改用frameworktransformers加载训练好的模型进行推理详见下文限制与注意事项train_size的切分行为从源码可以确认在 autotrain_advanced.py 的构造函数中若传入的是DatasetDict即已切分的数据集则自动取[train]与[test]作为训练/验证集否则全部数据作为训练集、验证集为空。四、深度配置update_config的两种模式AutoTrain 的配置体系分为两个互斥的入口文档中的第二个代码块展示了完整的配置写法trainer.update_config( model autotrain, # hub models like roberta-base autotrain [{ source_language: en, num_models: 5 }], hub_model [{ learning_rate: 0.001, optimizer: adam, scheduler: linear, train_batch_size: 8, epochs: 10, percentage_warmup: 0.1, gradient_accumulation_steps: 1, weight_decay: 0.1, tasks: text_binary_classification, # this is inferred from the dataset }] )4.1 模式一modelautotrain自动模型搜索将model设置为字符串autotrain会触发 AutoTrain 的自动化模型搜索模式。此时配置项放在autotrain列表中参数含义source_language源语言如en影响模型搜索候选集num_models尝试训练的模型数量例如 5 表示自动搜索并训练 5 个候选模型从源码看autotrain模式下只允许一个任务参数组——get_job_params中len(job_params) 1会直接抛出ValueError(Only one job parameter is allowed for AutoTrain.)。同时get_project_cost在计算项目成本时使用的是trainer_kwargs[autotrain][0][num_models]来估算模型数量见 autotrain_advanced.py。4.2 模式二modelhub模型名指定模型调优将model设置为具体的 Hugging Face Hub 模型名如roberta-base配置项放在hub_model列表中。该列表可以包含多个字典每个字典对应一次独立的调优任务从而实现多组超参对比实验。快速上手示例中传入两个不同learning_rate的字典正是这种用法。hub_model支持的主要超参数参数默认值由 AutoTrain 参数定义说明learning_rate0.001学习率optimizeradam优化器schedulerlinear学习率调度策略train_batch_size8训练批大小epochs10训练轮数percentage_warmup0.1warmup 步数比例gradient_accumulation_steps1梯度累积步数weight_decay0.1权重衰减tasks自动推断任务类型通常无需手动指定4.3 参数如何被填充与校验update_config的内部逻辑见 autotrain_advanced.pymodel若为autotrain会被归一化为小写形式hub_model与autotrain都必须是list[dict]否则抛出ValueError调用init_training_args()从autotrain.params.Params按task与training_type取出全部参数定义对用户未显式提供的键填入默认值——这也是上表默认值一栏的来源调用initialize_project()基于数据集、模型与 job 参数构建 AutoTrainProject并打印项目信息与预估成本。五、任务类型推断与数据集流转AutoTrain 的任务类型不需要手动指定源码会根据数据集自动推断。在 autotrain_advanced.py 的构造函数中若记录类型为TextClassificationRecord且为单标签分类标签类别数 ≤ 2 → 任务为text_binary_classification标签类别数 2 → 任务为text_multi_class_classification推断出的任务会通过job_params[i].update({task: self.task})注入每个任务参数组这正是update_config中tasks注释this is inferred from the dataset的底层含义。数据集流转路径如下ArgillaTrainer通过argilla.load()加载 Argilla 数据集prepare_for_training()将记录转换为datasets库的Dataset/DatasetDict含 train/test 切分AutoTrainMixin.prepare_dataset()基于AutoTrainDataset构造数据集对象内部使用column_mapping{text: text, label: label}完成列映射并统计_num_samples用于成本估算initialize_project()构建Project将dataset、hub_model与job_params组装为可提交的训练项目。训练提交则由train()完成三步project.dataset.prepare()预处理数据→project.create()创建项目→project.approve(project_id)批准并启动训练。训练完成后模型会保存到你的 Hugging Face 账号下save()方法会提示Models are saved on the HuggingFace Hub。六、限制与注意事项结合源码autotrain框架存在以下明确限制使用前需知悉不支持多标签文本分类TextClassification的multi_labelTrue会抛出NotImplementedError不支持 Token 分类与 Text2Text 任务源码中对TokenClassificationRecord直接raise NotImplementedErrorpredict()不提供本地推理调用后仅打印错误日志提示改用ArgillaTrainer(..., frameworktransformers, modelmy_model_name)加载 AutoTrain 产出的模型进行推理文档中的records trainer.predict(...)只是 API 形态的示意不支持seed设置 seed 会被忽略并固定为 42autotrain-advanced不支持该参数save()不可用模型由平台托管在 Hugging Face Hub本地保存不受支持autotrain模式每次仅允许一组参数多组参数会触发ValueError认证严格AUTOTRAIN_USERNAME与HF_AUTH_TOKEN环境变量缺失会导致导入时报错。七、训练完成后的推理衔接由于 AutoTrain 框架本身不承载本地推理官方建议的训练后推理路径是在 AutoTrain 平台或训练日志中找到保存的模型再用 Transformers 框架的ArgillaTrainer重新加载from argilla.training import ArgillaTrainer trainer ArgillaTrainer( namemy_dataset_name, workspacemy_workspace_name, frameworktransformers, modelyour_hf_username/output_dir_name # AutoTrain 产出的模型 ) records trainer.predict(The ArgillaTrainer is great!, as_argilla_recordsTrue)此时predict(..., as_argilla_recordsTrue)返回的将是 Argilla 记录格式可直接回写标注数据集形成训练 → 推理 → 再标注 → 再训练的数据闭环。八、进一步探索完整的训练配置总览AutoTrain、Transformers、SetFit、spaCy、PEFT 等各框架的update_config参数对照train_update_config.mdArgillaTrainer的通用架构与训练工作流fine_tune.mdAutoTrain Trainer 的完整源码实现任务推断、项目构建、成本计算等autotrain_advanced.pyArgillaTrainer基类与框架分派逻辑base.py支持的框架枚举定义models.py如需查看 Token 分类场景的对应代码片段可参考 autotrain_code.md注意源码中 Token 分类尚未被autotrain-advanced支持此片段属于文档规划的 API 形态展示。【免费下载链接】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),仅供参考
返回列表