
使用 ArgillaTrainer 与 Transformers 框架进行 Token Classification 微调【免费下载链接】argillaArgilla is a collaboration tool for AI engineers and domain experts to build high-quality datasets项目地址: https://gitcode.com/GitHub_Trending/ar/argillaArgilla 为 AI 工程师和领域专家提供了数据标注与模型训练的一体化协作能力。本文聚焦于 Argilla v1 中ArgillaTrainer在Token Classification词元级分类典型如命名实体识别 NER场景下的 Transformers 微调实践完整覆盖从数据集准备、训练器初始化、超参数配置到预测回灌的全流程。读完本文你将掌握如何使用几行代码将 Argilla 中已标注的 token 级数据集接入 Hugging Facetransformers生态完成微调并理解ArgillaTransformersTrainer底层的数据预处理、评估与推断原理。Token Classification 任务与 ArgillaTrainer 定位Token Classification 是 NLP 中的基础任务为文本中的每个词元token赋予一个标签例如词性标注、命名实体识别人名、组织、地名或情感极性词标注。Argilla 的数据模型使用TokenClassificationRecord承载此类标注texttokens 以(label, start, end)三元组表示的annotation并通过TokenClassificationSettings管理标签集合。ArgillaTrainer是 Argilla 提供的训练门面类它内部封装了数据转换、训练与推断逻辑开发者无需关心 Argilla 标注数据到框架训练格式的转换细节。在 argilla-v1/src/argilla_v1/training/base.py 中ArgillaTrainer.__init__会根据framework参数分发到具体的训练器实现其中frameworktransformers会实例化ArgillaTransformersTrainer定义于 argilla-v1/src/argilla_v1/training/transformers.py。根据 docs/_source/practical_guides/fine_tune.md 中的框架支持矩阵Token Classification 任务支持 spaCy、Transformers、PEFT 等框架本文只讨论 Transformers 路线。环境依赖与数据准备安装依赖ArgillaTransformersTrainer在初始化时会强制校验以下依赖见 transformers.py#L31pip install torch datasets transformers evaluate seqeval其中seqeval专用于 token 级序列标注的评估evaluate用于加载seqeval等指标。将标注数据送入 Argilla训练数据需要先作为TokenClassificationRecord日志到 Argilla 服务端。典型方式是从 Hugging Face 数据集转换后rg.log例如 fine_tune.md 中展示的 conll2003 示例import argilla_v1 as rg from datasets import load_dataset dataset_rg rg.DatasetForTokenClassification.from_datasets( datasetload_dataset(conll2003, splittrain[:100]), tagsner_tags, ) rg.log(dataset_rg, nameconll2003, workspaceadmin)ArgillaTrainer在构造时会调用argilla.load()拉取数据集并通过prepare_for_training()完成格式转换。从 datasets.py#L333-L392 可以看到该方法会剔除未标注记录将 IOB 标签序列转换为整数编码的ner_tags列ClassLabel从而得到可直接喂给 Transformers Trainer 的datasets.Dataset。快速上手最小训练流程以下代码片段完整继承自 docs/_source/_common/snippets/training/token-classification/transformers.md是 Token Classification Transformers 的标准工作流from argilla.training import ArgillaTrainer trainer ArgillaTrainer( namemy_dataset_name, workspacemy_workspace_name, frameworktransformers, train_size0.8 ) trainer.update_config(num_train_epochs10) trainer.train(output_dirtoken-classification) records trainer.predict(The ArgillaTrainer is great!, as_argilla_recordsTrue)各参数含义如下参数说明nameArgilla 服务端中数据集的名称必填workspace数据集所在工作区不传时默认使用当前用户工作区framework训练框架此处固定为transformerstrain_size训练集比例0~1 浮点数其余作为验证集不传则全量用于训练train_size0.8意味着 80% 数据用于训练、20% 用于评估——只有当验证集存在时ArgillaTransformersTrainer才会在训练后执行evaluate()并输出指标见 transformers.py#L523-L527。此外还可通过model参数指定基础模型如modeldistilbert-base-uncased、通过seed固定随机种子、通过gpu_id指定 GPU见 base.py#L38-L50。trainer.train(output_dir...)会在训练结束后自动将模型与 tokenizer 保存到output_dir并初始化推理 pipeline。trainer.predict(text, as_argilla_recordsTrue)返回TokenClassificationRecord列表可直接rg.log()回灌 Argilla 做进一步审查。配置详解一模型加载参数原文档展示了第一组update_config调用注释标明这些参数面向transformers.AutoModelForTextClassification。需要说明的是从源码看Token Classification 分支实际使用的是AutoModelForTokenClassification见 transformers.py#L77-L82文档中的类名注释应为文本分类场景的复用表述。这组参数对应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 或本地路径未指定时ArgillaTransformersTrainer默认使用bert-base-cased见 transformers.py#L55-L57force_downloadFalse即使缓存存在也强制重新下载权重resume_downloadFalse断点续传已下载的部分文件proxiesNone下载代理配置字典tokenNoneHugging Face Hub 鉴权 token访问私有模型时使用cache_dirNone模型缓存目录默认使用 HF 全局缓存local_files_onlyFalse仅使用本地文件不发起网络请求关于基础模型的指定更推荐在构造ArgillaTrainer时通过model参数传入ArgillaTrainer(name..., frameworktransformers, modeldistilbert-base-uncased)update_config中的模型参数会被filter_allowed_args按TrainingArguments签名过滤见 transformers.py#L194-L201真正稳定生效的是下一节的训练超参数。配置详解二TrainingArguments 训练超参数原文档的第二组update_config面向transformers.TrainingArguments是微调阶段最核心的可调参数集合。ArgillaTransformersTrainer.update_config()内部通过filter_allowed_args(TrainingArguments.__init__, **kwargs)仅透传TrainingArguments支持的参数见 transformers.py#L194-L201传入的其余参数会被安全忽略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 )各参数的作用域与取值建议如下参数类型/默认说明per_device_train_batch_sizeint默认 8每个设备GPU/CPU的训练 batch 大小per_device_eval_batch_sizeint默认 8每个设备的评估 batch 大小gradient_accumulation_stepsint默认 1梯度累积步数等效扩大 batch size显存受限时常用learning_ratefloat默认 5e-5AdamW 初始学习率微调常用 1e-5 ~ 5e-5weight_decayfloat默认 0AdamW 权重衰减系数抑制过拟合adam_beta1/adam_beta2float默认 0.9Adam 优化器的一阶/二阶矩衰减系数adam_epsilonfloat默认 1e-8Adam 优化器的数值稳定性项max_grad_normfloat默认 1梯度裁剪范数上限num_train_epochsint默认 3训练轮数小数据集可适当增大max_stepsint默认 0最大训练步数0 表示由 epochs 决定log_levelstr默认passive日志级别passive/info/warning/errorlogging_strategystr默认steps日志记录策略steps/epoch/nosave_strategystr默认steps模型保存策略steps/epoch/nosave_stepsint默认 500每 N 步保存一次 checkpointseedint默认 42随机种子保证可复现push_to_hubbool默认False训练结束后是否推送到 Hugging Face Hubhub_model_idstrHub 上的目标仓库 ID如user_name/output_dir_namehub_strategystr默认every_saveHub 推送策略end/every_save/checkpointhub_tokenstr推送所需的 Hub 鉴权 tokenhub_private_repobool默认False推送的仓库是否设为私有需要注意的是init_training_args已预设部分默认值见 transformers.py#L120-L124存在验证集时evaluation_strategyepoch否则为nologging_steps1num_train_epochs1。update_config传入的同名参数会覆盖这些默认值。源码级原理训练管线与数据对齐理解ArgillaTransformersTrainer.train()的调用链见 transformers.py#L492-L531能帮助你排查数据与配置问题init_model(newTrue)根据model_kwargs中的pretrained_model_name_or_path加载AutoTokenizer与AutoModelForTokenClassification。tokenizer 的padding_side默认rightGPT/OPT/BLOOM 等自回归模型为left并强制add_prefix_spaceTrue若 tokenizer 无pad_token_id则回退用eos_token_idmodel_max_length缺失时置为 512见 transformers.py#L126-L153。preprocess_datasets()对 token 分类执行词元标签对齐。核心逻辑是使用tokenizer(..., is_split_into_wordsTrue)将 Argilla 的tokens直接切分再通过word_ids()把每个词的ner_tags标签只赋给该词切分后的第一个 sub-token其余 sub-token 标记为-100在损失计算中被忽略特殊 token 同样置-100见 transformers.py#L226-L245。数据整理后使用DataCollatorForTokenClassification动态 padding。compute_metrics()Token Classification 使用evaluate.load(seqeval)将预测与真实标签中-100的位置剔除后计算 overall 的precision、recall、f1与accuracy四项指标见 transformers.py#L404-L428。构建 Trainer 并训练将TrainingArguments(**trainer_kwargs)、模型、tokenizer、预处理后的训练/评估数据集、指标函数与 data collator 一并交给transformers.Trainer训练结束后自动save(output_dir)同时保存模型与 tokenizer并初始化推理 pipeline见 transformers.py#L511-L531。设备选择逻辑位于 transformers.py#L45-L49优先mpsApple Silicon其次cuda否则回退cpu默认随机种子为 42见 transformers.py#L51-L53。推理与结果回灌predict(text, as_argilla_recordsTrue)在未训练时也会触发加载基座模型 初始化 pipeline的兜底逻辑此时给出的是基座模型未经微调的结果见 transformers.py#L545-L548。对 Token Classification它使用pipeline(tasktoken-classification, aggregation_strategyfirst)将 sub-token 粒度的预测聚合到词级实体见 transformers.py#L176-L183并通过word_to_chars()还原每个实体的字符区间与原始 token 文本见 transformers.py#L569-L583。返回的预测形如(entity_group, start, end)三元组组装为TokenClassificationRecord后即可回灌 Argilla 平台import argilla_v1 as rg records trainer.predict(The ArgillaTrainer is great!, as_argilla_recordsTrue) rg.log(recordsrecords, namemy_dataset_name, workspacemy_workspace_name)当as_argilla_recordsFalse时返回原始 pipeline 输出实体组、起始/结束位置与置信度分数便于自定义后处理。测试验证与扩展仓库的集成测试 argilla-v1/tests/integration/training/test_transformers.py 使用轻量模型prajjwal1/bert-tiny验证了完整链路test_update_config断言update_config能正确写入trainer_kwargstest_train_tokencat覆盖了 Token Classification 的训练 → predict → 返回TokenClassificationRecord闭环test_predict_wo_training验证了未微调直接推理的兜底路径。这些测试可作为你本地复现与排查问题的最小参考。此外ArgillaTrainer还提供了 CLI 入口python -m argilla train支持通过--framework transformers、--model、--train-size、--seed、--update-config-kwargs等参数在外部机器上执行同样的训练流程见 fine_tune.md。训练完成后模型与 tokenizer 保存在output_dir下可直接被transformers生态加载用于部署或二次微调。【免费下载链接】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),仅供参考