You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于Pandas DataFrame的SetFit零样本分类器训练超时问题咨询

小数据集SetFit训练超时排查与方案优化

问题背景

我计划在标注样本数据集上训练一个零样本分类器,参考教程时发现大多使用自有数据和相同预训练模型,想确认当前方法是否为最优方案。我的数据集仅包含88条记录、11个标签,但在笔记本电脑上训练时长已超500小时仍未结束,需要协助排查问题。

现有代码

数据处理代码

import pandas as pd
from datasets import Dataset
    
# Sample feedback data, it will have 8 samples per label
feedback_dict = [
    {'text': 'The product is great and works well.', 'label': 'Product Performance'},
    {'text': 'I love the design of the product.', 'label': 'Product Design'},
    {'text': 'The product is difficult to use.', 'label': 'Usability'},
    {'text': 'The customer service was very helpful.', 'label': 'Customer Service'},
    {'text': 'The product was delivered on time.', 'label': 'Delivery Time'}
]

# Create a DataFrame with the feedback data
df = pd.DataFrame(feedback_dict)

# convert to Dataset format
df = Dataset.from_pandas(df)

模型训练代码

from setfit import SetFitModel, SetFitTrainer

# Select a model
model = SetFitModel.from_pretrained("sentence-transformers/paraphrase-mpnet-base-v2")

# training with Setfit
trainer = SetFitTrainer(
    model=model,
    train_dataset=df, # to keep the code simple I do not create the df_train
    eval_dataset=df, # to keep the code simple I do not create the df_eval
    column_mapping={"text": "text", "label": "label"} 
)

trainer.train()

排查与优化方案

训练超时原因排查

  • 默认参数冗余:SetFitTrainer的train()方法默认设置num_iterations=20(对比学习迭代次数),对于仅88条数据的小样本场景,过多的迭代会导致大量重复计算,拖慢训练速度。
  • 硬件性能瓶颈:paraphrase-mpnet-base-v2属于参数较大的模型,若仅用笔记本CPU训练,计算效率极低,会导致训练时长大幅增加。
  • 数据集复用异常:训练集与验证集完全相同,可能触发训练过程中的过拟合检测或数据加载逻辑异常,导致训练无法正常终止。

针对性优化措施

  1. 调整训练核心参数
    显式减少对比学习迭代次数和分类头训练轮次,降低不必要的计算量:

    trainer.train(num_iterations=5, num_epochs=1)
    

    小数据集下,num_iterations设置为3-5即可满足需求,num_epochs保持1-2轮足够。

  2. 更换轻量预训练模型
    选用参数更精简的模型,在保证分类效果的前提下提升训练速度,比如:

    model = SetFitModel.from_pretrained("sentence-transformers/all-MiniLM-L6-v2")
    

    该模型参数量仅为paraphrase-mpnet-base-v2的1/6左右,训练速度提升明显。

  3. 拆分训练与验证集
    即使数据量小,也要拆分独立的验证集,避免数据复用导致的逻辑异常:

    df_split = df.train_test_split(test_size=0.2, seed=42)
    trainer = SetFitTrainer(
        model=model,
        train_dataset=df_split['train'],
        eval_dataset=df_split['test'],
        column_mapping={"text": "text", "label": "label"} 
    )
    
  4. 启用GPU加速
    若笔记本搭载NVIDIA GPU,确保安装对应版本的CUDA和PyTorch,让模型在GPU上运行,训练速度可提升数倍至数十倍。

方案合理性确认

你的场景属于少样本分类(有88条标注数据),SetFit是这类场景的优质方案,通过参数调整和硬件优化即可高效完成训练。如果是严格意义上的零样本分类(无标注数据),则建议使用BART-large-mnli、DeBERTa-v3-large-mnli等支持自然语言推理的模型,通过构造假设句实现零样本分类。

内容的提问来源于stack exchange,提问作者PeCaDe

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.18 19:45:00