Argilla ArgillaTrainer.update_config 全框架训练配置详解:OpenAI、Transformers、PEFT、TRL 等九大框架参数参考手册
Argilla ArgillaTrainer.update_config 全框架训练配置详解OpenAI、Transformers、PEFT、TRL 等九大框架参数参考手册【免费下载链接】argillaArgilla is a collaboration tool for AI engineers and domain experts to build high-quality datasets项目地址: https://gitcode.com/GitHub_Trending/ar/argilla本指南以 Argilla 的ArgillaTrainer.update_config()方法为线索系统梳理其覆盖的 OpenAI、AutoTrain、SetFit、spaCy、Transformers、PEFTLoRA、SpanMarker、TRL、sentence-transformers 九大训练框架的全部可配置参数及其默认值。读完本文你将理解update_config的底层映射机制能够针对不同任务与框架精准调参直接复制参数示例完成模型微调并知道如何通过 CLI 在外部机器上以 JSON 字符串传递同样的配置。背景为什么需要update_configArgilla 的ArgillaTrainer是对多个主流 NLP 训练库的封装wrapper。它的设计目标是用户只需用一套统一 API 完成「数据集准备 → 训练 → 推理」而无需关心 Argilla 数据到各框架输入格式的转换细节。相关概念在 fine_tune.md 中有完整介绍首先通过TrainingTask.for_*定义任务再初始化ArgillaTrainer并指定framework随后即可调用三个核心方法ArgillaTrainer.update_config—— 修改框架相关的训练参数本文主题ArgillaTrainer.train—— 启动训练ArgillaTrainer.predict—— 执行推理。典型的最小工作流如下摘自 fine_tune.mdfrom argilla.feedback import ArgillaTrainer, FeedbackDataset, TrainingTask dataset FeedbackDataset.from_huggingface(repo_idargilla/emotion) task TrainingTask.for_text_classification( textdataset.field_by_name(text), labeldataset.question_by_name(label), ) trainer ArgillaTrainer(datasetdataset, tasktask, frameworksetfit) trainer.update_config(num_iterations1) trainer.train(output_dirmy_setfit_model) trainer.predict(This is awesome!)update_config的底层工作方式源码级原理update_config并不是一个黑盒赋值方法它在 feedback/training/base.py 中做了参数合法性校验def update_config(self, *args, **kwargs) - None: def get_all_keys(d): keys [] for k, v in d.items(): keys.append(k) if isinstance(v, dict): keys get_all_keys(v) return keys trainer_kwargs self._trainer.get_trainer_kwargs() model_kwargs self._trainer.get_model_kwargs() all_keys get_all_keys({**trainer_kwargs, **model_kwargs}) for kwarg in kwargs: if kwarg not in all_keys: warnings.warn( f{kwarg} is not a valid default argument for {self._trainer.__class__.__name__}. fValid default arguments are: {all_keys}. , UserWarning, stacklevel2, ) return super().update_config(*args, **kwargs)从源码可以归纳出三点关键行为参数来源每个框架 trainer 内部维护model_kwargs模型初始化参数与trainer_kwargs训练过程参数两个字典update_config会把两者全部键收集起来做白名单校验传入不在白名单中的参数会触发UserWarning而不是静默失败。透传底层框架校验通过后配置最终透传给具体框架 trainer 的update_config见 training/base.py。因此本文列出的参数名与默认值直接对应底层库如transformers.TrainingArguments、setfit.SetFitTrainer的构造函数签名。默认值自动推导以 sentence-transformers 框架为例sentence_transformers.pytrainer 初始化时会用get_default_args(self._trainer_cls.__init__)和get_default_args(self._trainer_cls.fit)从底层类的签名中抓取默认参数——这正是下方各框架参数表与底层库保持一致的机制来源。需要强调的是以下列出的参数并不需要全部传入。它们展示的是各框架的默认配置你只需传入想覆盖的少量参数即可。ArgillaTrainer在初始化时通过dataset.prepare_for_training(framework..., task...)完成数据格式化dispatch 逻辑见 feedback/training/base.py你可以在初始化后打印 trainer 对象查看可配置参数全貌。各框架训练配置参考下文逐框架给出update_config可接受的参数及默认值。每个代码块前的注释标注了参数实际归属的底层类。OpenAIOpenAI 框架支持 Chat Completion 等任务可同时使用新版FineTune与 legacy 接口对应 openai.py# OpenAI.FineTune trainer.update_config( training_file None, validation_file None, model gpt-3.5-turbo-0613, hyperparameters {n_epochs: 1}, suffix None ) # OpenAI.FineTune (legacy) trainer.update_config( training_file None, validation_file None, model curie, n_epochs 2, 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 )参数说明参数含义training_file/validation_file训练/验证文件 ID为None时由 Argilla 根据数据集自动准备并上传model基础模型标识新版默认gpt-3.5-turbo-0613legacy 默认curiehyperparameters新版接口的参数字典如{n_epochs: 1}n_epochs、batch_size、learning_rate_multiplier、prompt_loss_weightlegacy 接口的训练超参数None表示使用 OpenAI 侧默认值compute_classification_metrics是否计算分类指标classification_n_classes/classification_positive_class/classification_betas分类任务的类别数、正类与 beta 配置suffix微调模型名称后缀AutoTrainAutoTrain 框架对应AutoTrain.autotrain_advanced实现见 autotrain_advanced.py支持在 Hub 模型上自动搜索训练配置# AutoTrain.autotrain_advanced 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 }] )参数说明model基础模型标识可填autotrain或具体的 Hub 模型名如roberta-baseautotrain自动训练配置列表source_language指定源语言num_models指定参与搜索/集成的模型数量hub_model目标模型超参列表包括learning_rate、optimizer如adam、scheduler如linear、train_batch_size、epochs、percentage_warmup、gradient_accumulation_steps、weight_decay以及tasks如text_binary_classification通常从数据集自动推断。SetFitSetFit 框架同时暴露模型初始化与训练器两组参数对应setfit.SetFitModel与setfit.SetFitTrainer# setfit.SetFitModel trainer.update_config( pretrained_model_name_or_path all-MiniLM-L6-v2, force_download False, resume_download False, proxies None, token None, cache_dir None, local_files_only False ) # setfit.SetFitTrainer trainer.update_config( metric accuracy, num_iterations 20, num_epochs 1, learning_rate 2e-5, batch_size 16, seed 42, use_amp True, warmup_proportion 0.1, distance_metric BatchHardTripletLossDistanceFunction.cosine_distance, margin 0.25, samples_per_label 2 )参数说明模型侧pretrained_model_name_or_path指定预训练句子编码器默认all-MiniLM-L6-v2force_download、resume_download、local_files_only控制下载行为proxies、token、cache_dir处理网络与缓存训练侧metric默认accuracynum_iterations默认 20是 SetFit 特有的「对比学习句子对迭代次数」num_epochs、learning_rate、batch_size、seed为常规训练超参use_amp启用混合精度warmup_proportion为热身比例对比学习相关distance_metric使用余弦距离BatchHardTripletLossDistanceFunction.cosine_distancemargin为三元组损失边界默认 0.25samples_per_label控制每标签采样数。spaCyspaCy 框架的配置面向spacy.training流程对应 spacy.py注意gpu_allocator为 0 表示由系统分配# spacy.training trainer.update_config( dev_corpus corpora.dev, train_corpus corpora.train, seed 42, gpu_allocator 0, accumulate_gradient 1, patience 1600, max_epochs 0, max_steps 20000, eval_frequency 200, frozen_components [], annotating_components [], before_to_disk None, before_update None )参数说明dev_corpus/train_corpus指定验证与训练语料键名seed固定随机数保证可复现gpu_allocator用于 GPU 显存分配accumulate_gradient为梯度累积步数patience为早停耐心值默认 1600max_epochs0 表示不限制与max_steps联合控制训练停止条件eval_frequency控制评估频率frozen_components/annotating_components分别指定冻结与仅做标注的流水线组件before_to_disk/before_update为回调钩子。TransformersTransformers 框架组合了transformers.AutoModelForTextClassification模型初始化与transformers.TrainingArguments训练参数对应 transformers.py# 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 ) # 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 )参数说明模型侧与 SetFit 一致默认编码器为distilbert-base-uncased训练侧per_device_train_batch_size/per_device_eval_batch_size为单卡批大小默认 8gradient_accumulation_steps梯度累积默认 1learning_rate默认5e-5weight_decay默认 0Adam 优化器三参数adam_beta1、adam_beta2、adam_epsilonmax_grad_norm梯度裁剪num_train_epochs训练轮数默认 3max_steps为 0 时不覆盖 epoch 逻辑日志与保存log_levelpassive、logging_strategysteps、save_strategysteps配合save_steps500按步保存 checkpointseed42固定随机种子Hub 集成push_to_hub是否推送、hub_model_id目标仓库名、hub_strategyevery_save每次保存即推送、hub_token与hub_private_repo控制认证与仓库可见性。Peft (LoRA)PEFT 框架在 Transformers 基础上叠加peft.LoraConfig用于参数高效微调对应 peft.py# 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, ) # 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 ) # 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 )LoraConfig参数说明rLoRA 秩默认 8控制新增可训练矩阵的维度target_modules要注入 LoRA 适配器的模块名列表None时由实现自动推断lora_alpha缩放系数默认 16通常与r配合调整lora_dropoutLoRA 层 dropout默认 0.1fan_in_fan_out针对 Conv1D 等权重布局的开关bias可训练偏置策略none/all/lora_onlyinference_mode是否以推理模式创建适配器modules_to_save除 LoRA 外额外保存的模块init_lora_weights是否初始化 LoRA 权重。后两组模型初始化与TrainingArguments与 Transformers 框架完全相同可参照上节说明。SpanMarkerSpanMarker 框架用于 Token ClassificationNER配置包括SpanMarkerConfig与transformers.TrainingArguments对应 span_marker.py# SpanMarkerConfig trainer.update_config( pretrained_model_name_or_path distilbert-base-cased, model_max_length 256, marker_max_length 128, entity_max_length 8, ) # 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 )SpanMarkerConfig参数说明pretrained_model_name_or_path默认distilbert-base-cased注意此处为 cased 版本适合含实体大小写的 NER 场景model_max_length整个输入序列的最大长度默认 256marker_max_lengthmarker实体标记部分的最大长度默认 128entity_max_length单个实体跨度允许的最大 token 数默认 8。TRLTRL 框架覆盖 SFT监督微调、Reward Modeling、PPO、DPO 等 RLHF 任务参数来自trl.RewardTrainer、trl.SFTTrainer、trl.PPOTrainer或trl.DPOTrainer实现见 trl.py其中训练参数统一通过transformers.TrainingArguments传入# Parameters from trl.RewardTrainer, trl.SFTTrainer, trl.PPOTrainer or trl.DPOTrainer. # 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 )这些参数与 Transformers 框架完全同构可参照上文。值得注意的是TRL 场景下部分任务如 PPO还需要通过update_config传入任务特有对象——例如reward_modelreward pipeline、length_sampler_kwargs生成 token 数上下界、generation_kwargs生成策略与configtrl.PPOConfig(...)详见 fine_tune.md 的 PPO 章节。sentence-transformerssentence-transformers 框架参数最丰富分为「模型初始化」「训练过程」「外部参数」三组对应 sentence_transformers.py。默认 Bi-Encoder 为sentence-transformers/all-MiniLM-L6-v2若在ArgillaTrainer中传入framework_kwargs{cross_encoder: True}则默认换为cross-encoder/ms-marco-MiniLM-L-6-v2。# Parameters related to the model initialization from sentence_transformers.SentenceTransformer trainer.update_config( modelsentence-transformers/all-MiniLM-L6-v2, modules False, devicecuda, cache_folderdir/folder, use_auth_tokenTrue ) # and from sentence_transformers.CrossEncoder trainer.update_config( modelcross-encoder/ms-marco-MiniLM-L-6-v2, num_labels2, max_length128, devicecpu, tokenizer_args{}, automodel_args{}, default_activation_functionNone ) # Related to the training procedure from sentence_transformers.SentenceTransformer trainer.update_config( steps_per_epoch 2, checkpoint_path: str None, checkpoint_save_steps: int 500, checkpoint_save_total_limit: int 0 ) # and from sentence_transformers.CrossEncoder trainer.update_config( loss_fct None, activation_fct nn.Identity(), ) # The remaining arguments are common for both procedures trainer.update_config( evaluator: SentenceEvaluator evaluation.EmbeddingSimilarityEvaluator, epochs: int 1, scheduler: str WarmupLinear, warmup_steps: int 10000, optimizer_class: Type[Optimizer] torch.optim.AdamW, optimizer_params : Dict[str, object] {lr: 2e-5}, weight_decay: float 0.01, evaluation_steps: int 0, output_path: str None, save_best_model: bool True, max_grad_norm: float 1, use_amp: bool False, callback: Callable[[float, int, int], None] None, show_progress_bar: bool True, ) # Other parameters that dont correspond to the initialization or the trainer, but # can be set externally. trainer.update_config( batch_size8, # It will be passed to the DataLoader to generate batches during training. loss_clslosses.BatchAllTripletLoss )参数说明模型初始化Bi-Encodermodel指定预训练编码器modules控制加载时是否执行模块变换device默认cudacache_folder指定模型缓存目录use_auth_token访问私有 Hub 模型模型初始化Cross-Encodernum_labels默认 2、max_length默认 128、device默认cpu、tokenizer_args/automodel_args透传给 tokenizer 与 automodel 的额外参数、default_activation_function输出激活函数训练过程Bi-Encodersteps_per_epoch、checkpoint 系列参数checkpoint_path、checkpoint_save_steps500、checkpoint_save_total_limit0训练过程Cross-Encoderloss_fct自定义损失、activation_fct默认nn.Identity()通用训练参数evaluator默认EmbeddingSimilarityEvaluatorepochs1schedulerWarmupLinear与warmup_steps10000optimizer_classtorch.optim.AdamW及optimizer_params{lr: 2e-5}weight_decay0.01evaluation_steps0不单独评估output_path输出目录save_best_modelTrue保存最优模型max_grad_norm1use_ampFalsecallback训练回调show_progress_barTrue外部参数batch_size8会被传入 DataLoader 用于分批loss_cls指定损失类如losses.BatchAllTripletLoss。从源码看sentence_transformers.pyloss_cls会依据「是否有标签」与「样本含 2 句还是 3 句」自动选择两句话有整数标签用ContrastiveLoss、浮点标签用CosineSimilarityLoss三句话triplet整数标签用BatchHardTripletLoss、浮点标签用BatchAllTripletLoss因此默认值由数据形态自动推断。实战要点配置、训练与 CLI组合使用。update_config支持多次调用也可以与其他方法串联trainer ArgillaTrainer( datasetdataset, tasktask, frameworktransformers, train_size0.8, ) trainer.update_config(per_device_train_batch_size16, num_train_epochs5) trainer.train(output_dirmy_model) trainer.predict(This is awesome!)CLI 方式。update_config的参数也可以通过命令行传递CLI 支持在 fine_tune.md 有完整说明python -m argilla train的--update-config-kwargs选项接受一个 JSON 可序列化字符串内部仍调用对应 trainer 的update_config方法适合在外部机器上执行训练。常用 CLI 选项还包括--framework、--model、--train-size、--seed、--device、--output-dir等。同一套机制的另一种数据集形态。本文介绍的是 Feedback Dataset含 LLM/RLHF 任务的扩展框架集对于更早的 v1 数据集TextClassification、TokenClassification、Text2TextArgillaTrainer同样提供update_config其参数清单见配套文档 train_update_config_other_datasets.md两者用法一致仅框架支持矩阵略有差异如新增 SpanMarker 于 Token Classification、OpenAI 于 Text2Text。小结ArgillaTrainer.update_config()的价值在于用一套 Python 方法统一了九大训练框架的调参入口参数名与默认值直接透传底层库签名并带有参数合法性告警。掌握本文的参数手册后你既能为文本分类、句子相似度、NER 等经典任务快速调参也能为 SFT、Reward Modeling、PPO、DPO、Chat Completion 等 LLM 工作流精细控制训练过程。想进一步深入可继续阅读 fine_tune.md 中各任务的完整训练与推理示例或查看 cheatsheet.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),仅供参考