资讯动态

使用 Argilla + OpenAI 微调文本分类模型:ArgillaOpenAITrainer 实战指南

发布时间:2026/9/18 23:18:13 来源:尧图企业网站定制
使用 Argilla OpenAI 微调文本分类模型ArgillaOpenAITrainer 实战指南【免费下载链接】argillaArgilla is a collaboration tool for AI engineers and domain experts to build high-quality datasets项目地址: https://gitcode.com/GitHub_Trending/ar/argilla本文围绕 Argilla 的ArgillaOpenAITrainer讲解如何以编程方式把 Argilla 中标注好的文本分类数据集接入 OpenAI 微调fine-tuning流程。通过ArgillaTrainer统一入口你只需指定数据集、工作空间与frameworkopenai即可自动完成数据格式转换、训练文件上传、微调任务提交与预测回写。读完本文你将掌握完整的 OpenAI 微调参数体系、底层实现原理与常见坑位规避方法。一、背景为什么用 Argilla 驱动 OpenAI 微调在 NLP 文本分类场景中高质量标注数据与模型微调往往割裂为两个独立环节数据散落在标注工具里模型训练则要手工整理 JSONL、对接各家 API。Argilla 的核心价值在于把数据标注与模型训练串成一条自动化流水线。ArgillaOpenAITrainer正是这条流水线中连接 Argilla 与 OpenAI 的桥梁。其官方定位描述为The ArgillaOpenAITrainer leverages the features of OpenAI to fine-tune programmatically with Argilla.也就是说你不需要手动导出数据集、手工构造 OpenAI 微调文件Trainer 会替你完成数据集读取、prompt/completion 格式化、JSONL 序列化、文件上传、微调任务创建与结果检索的全部流程。从源码结构看该 Trainer 位于 argilla-v1/src/argilla_v1/training/openai.py继承自通用骨架ArgillaTrainerSkeleton见 base.py与setfit、peft、spacy、transformers等框架并列通过frameworkopenai即可切换属于 Argilla v1 训练体系中的一员。二、前置条件与初始化约束2.1 依赖与密钥ArgillaOpenAITrainer在初始化时会做两件硬性校验依赖检查调用require_dependencies(openai0.27.10)要求环境中安装不低于 0.27.10 版本的 OpenAI Python SDK。环境变量检查源码第 37-39 行明确要求OPENAI_API_KEY必须存在于环境变量中否则在导入/初始化阶段直接抛出ValueError(OPENAI_API_KEY not found in environment variables.)。export OPENAI_API_KEYsk-... pip install openai0.27.102.2 任务类型约束ArgillaOpenAITrainer当前只支持单标签文本分类不支持以下两类任务源码第 45-50 行TokenClassificationRecordToken 级分类直接抛出NotImplementedError多标签文本分类multi_labelTrue同样抛出NotImplementedError。因此若你的数据集包含多标签或多任务设置需要先通过 Argilla 的prepare_for_training流程转换为单标签文本分类或改用其他框架如setfit、transformers。三、最小可用示例三步完成微调关联文档给出了一段可直接运行的代码骨架这是整个 OpenAI 微调流程的最小闭环from argilla.training import ArgillaTrainer trainer ArgillaTrainer( namemy_dataset_name, workspacemy_workspace_name, frameworkopenai, train_size0.8 ) trainer.update_config(num_iterations10) trainer.train(output_dirtext-classification) records trainer.predict(The ArgillaTrainer is great!, as_argilla_recordsTrue)注意上述片段中的update_config(num_iterations10)取自通用文档模板OpenAI 微调并不存在num_iterations参数它是 SetFit 等框架的配置项。针对 OpenAI 框架正确的超参数见下一节的update_config说明。ArgillaTrainer内部会将不同框架的update_config分发给对应 Trainer见 base.py。各步骤的实际行为如下ArgillaTrainer(...)根据frameworkopenai在 base.py 中实例化ArgillaOpenAITrainer。train_size0.8表示从数据集中切分 80% 作为训练集剩余 20% 作为验证集ArgillaTrainer内部会调用数据集的prepare_for_training完成切分与格式转换。trainer.train(output_dirtext-classification)将训练集序列化为 JSONL 上传到 OpenAI创建微调任务并轮询等待任务提交成功。trainer.predict(...)使用微调后的模型对新文本做分类预测as_argilla_recordsTrue表示把预测结果包装成 Argilla 记录对象方便回写与后续标注管理。四、微调超参数体系update_config全参数解析关联文档给出了完整的 OpenAI 微调配置示例trainer.update_config( training_file None, validation_file None, model curie, n_epochs 4, batch_size None, learning_rate_multiplier 0.1, prompt_loss_weight 0.1, compute_classification_metrics False, classification_n_classes None, classification_positive_class None, classification_betas None, suffix None )这些参数与源码 openai.py 中init_training_args的签名一一对应。下表给出每个参数的作用与注意事项参数类型默认值含义与说明training_filestrNone训练文件的 OpenAI File ID。为None时train()会自动把训练集上传生成文件名为data_train.jsonlvalidation_filestrNone验证文件 ID。仅当提供了验证集且开启compute_classification_metrics时才会自动上传modelstrcurie微调基座模型名。源码中默认模型为curie若构造ArgillaTrainer时未指定model则回退为gpt-3.5-turbon_epochsint4分类/ 2其他训练轮数。旧版 API 下文本分类默认 4其他任务默认 2batch_sizeintNone微调批次大小None时由 OpenAI 自动选择learning_rate_multiplierfloat0.1学习率倍率prompt_loss_weightfloat0.1prompt 部分损失的权重compute_classification_metricsboolFalse是否计算分类指标。当提供验证集时会自动置为Trueclassification_n_classesintNone分类类别数多分类时由标签 schema 推导classification_positive_classstrNone二分类的正类标签classification_betaslistNone分类 F-beta 指标中的 beta 值suffixstrNone微调模型名称后缀。train(output_dir...)时会把output_dir写入该参数4.1 参数在下层如何生效update_config的底层机制源码第 147-166 行分三步将传入 kwargs 合并进trainer_kwargs字典调用filter_allowed_args(self.init_training_args, **self.trainer_kwargs)过滤掉init_training_args签名之外的非法参数剔除所有值为None的键这些参数由 OpenAI 服务端使用默认值并把model同步回self._model。这套白名单过滤 None 剔除机制意味着你可以放心传入通用参数Trainer 会自动丢弃不适用于 OpenAI 的项。4.2 旧版 API 与新版 API 的分流根据源码第 119-145 行Trainer 会根据基座模型是否为旧版模型走不同分支旧版legacy模型OPENAI_LEGACY_MODELS定义为[babbage, davinci, curie, ada]见 argilla-v1/src/argilla_v1/_constants.py。此时n_epochs、batch_size、learning_rate_multiplier、prompt_loss_weight、compute_classification_metrics、classification_*系列参数会被原样写入trainer_kwargs微调调用openai.FineTune.create。新版模型如gpt-3.5-turbo上述传统参数被折叠进hyperparameters字典微调调用openai.FineTuningJob.create。默认仅写入hyperparameters[n_epochs] n_epochs or 1。也就是说文档示例中model curie属于旧版模型分支其中compute_classification_metrics、classification_n_classes、classification_positive_class、classification_betas只有在提供验证集时才会真正发挥作用。4.3 自动分类配置推导当使用旧版模型且提供验证集时源码第 132-139 行Trainer 会根据标签数量自动推导分类配置标签数为 2classification_positive_class label_schema[0]、compute_classification_metrics True走二分类指标计算标签数 2classification_n_classes len(label_schema)、compute_classification_metrics True走多分类指标计算。五、数据上传与微调任务的底层实现5.1 数据格式转换从 Argilla 记录到 Chat 格式新版 API 下训练数据会被转换为 OpenAI Chat 微调所需的messages结构源码第 85-96 行{ messages: [ {role: user, content: fClassify the following text: {entry[prompt]}}, {role: assistant, content: entry[completion]}, ] }即每条样本的用户消息固定为Classify the following text: 文本助手消息为标注的类别标签。这一转换发生在__init__阶段源码第 78-81 行因此训练集与验证集在进入训练前就已格式化完毕。5.2 JSONL 序列化与上传upload_dataset_to_openai源码第 179-200 行的实现要点移除记录中的id字段将每条样本json.dumps(item) \n编码为 UTF-8 字节流构成JSONL每行一条 JSON格式调用openai.File.create(file..., purposefine-tune)上传返回 OpenAI 侧的文件 ID。训练文件固定命名为data_train.jsonl验证文件为data_test.jsonl。5.3 任务提交与重试机制train方法源码第 202-248 行的完整流程若传入了output_dir将其写入suffix即微调模型的后缀名若training_file为空自动上传训练集若存在验证集且compute_classification_metrics为真自动上传验证集进入while not started_training循环调用FineTune.create旧版或FineTuningJob.create新版创建任务失败则记录 warning 并每 10 秒sleep_timer 10重试直到成功记录任务 ID 到self.finetune_id并提示用openai.FineTuningJob.retrieve(id)新版查询训练进度。值得注意训练是异步的train()返回时微调未必完成OpenAI 会在任务完成时发送邮件通知。之后可通过init_model()源码第 250-266 行拉取微调完成的模型 IDresponse.fine_tuned_model此时self._model会被替换为微调后的模型若任务仍在进行则给出 Fine-tuning is still in progress 警告。六、预测把模型能力接回 Argillapredict方法源码第 268-337 行目前仅支持旧版模型走openai.Completion.create新版 API 的 Chat 预测尚未实现对应分支会抛出NotImplementedError并提示参考 OpenAI Chat 文档。针对文本分类任务predict 会自动注入一组合理的推理参数kwargs[logprobs] len(self._settings.label_schema) # 输出所有类别的对数概率 kwargs[max_tokens] 1 # 只生成 1 个 token即类别 kwargs[temperature] 0 # 贪婪解码保证可复现 kwargs[n] 1随后对每个输入构造prompt f{entry.strip()}{self._separator}其中_separator为\n\n###\n\n、_end_token为 END、_whitespace为 见 argilla-v1/src/argilla_v1/_constants.py调用 Completion 接口后将 logprobs 经np.exp还原为概率并把标签-概率对组装成 Argilla 的TextClassificationRecordpredictionlist(zip(keys, values))。因此predict(..., as_argilla_recordsTrue)返回的是可直接用于 Argilla 管理的预测记录records trainer.predict(The ArgillaTrainer is great!, as_argilla_recordsTrue) # 传入字符串时返回单个记录传入列表时返回记录列表七、模型保存suffix 即保存与其他框架不同OpenAI 微调结果托管在云端save方法源码第 339-347 行并不落盘而是给出明确提示Saving is not supported for OpenAI and is passed via thesuffixargument intrain.即保存模型的语义被映射为微调模型名称后缀。你在train(output_dirtext-classification)中传入的输出目录名会成为微调模型的标识后缀最终通过init_model()拿到fine_tuned_model完整 ID 后即可在 OpenAI 侧管理该模型。八、实践建议与注意事项汇总基于上述源码行为给出几条落地建议密钥先行确保OPENAI_API_KEY在启动进程前已注入环境变量否则导入即失败模型选择决定参数分支curie/davinci/babbage/ada走旧版FineTuneAPI 且支持完整分类指标参数gpt-3.5-turbo等新模型走FineTuningJobAPI分类指标类参数会被忽略验证集建议开启提供验证集且标签数明确时Trainer 会自动开启分类指标计算compute_classification_metrics微调过程将输出可量化的评估结果预测仅限旧版模型当前predict对新版 Chat 模型尚未支持落地前请确认所选基座模型落在 legacy 分支或自行基于model...调用 OpenAI Chat API 完成推理训练为异步任务train()提交成功后即可轮询finetune_idOpenAI 完成时会邮件通知init_model()会拉取最终微调模型 ID。九、延伸阅读Trainer 统一入口与框架分发逻辑argilla-v1/src/argilla_v1/training/base.pyOpenAI Trainer 完整实现argilla-v1/src/argilla_v1/training/openai.pyprompt/completion 分隔符与 legacy 模型常量argilla-v1/src/argilla_v1/_constants.py文本分类任务的其他框架示例setfit / peft / spacy / transformersdocs/_source/_common/snippets/training/text-classification/使用prepare_for_training(frameworkopenai, train_size...)在训练前准备数据docs/_source/_common/tabs/train_prepare_for_training.md文档版快速上手docs/_source/getting_started/quickstart.md【免费下载链接】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),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价