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

SGD文本分类模型训练报错:y应为1维数组/类别数需大于1

报错原因分析
  • 标签维度错误:你通过np.array('Low')生成的是0维标量数组,不符合sklearn要求的1D输入格式,触发y should be a 1d array报错
  • 增量训练方法误用:SGD分类器的fit方法是重置模型从头训练,不是增量学习。你每次仅传入单条样本,样本中只有1个类别,触发The number of classes has to be greater than one报错
  • 特征提取器使用错误:retrain函数中调用cv.fit_transform(X)会覆盖预训练阶段得到的词汇表,新生成的特征维度和预训练模型要求的输入维度不匹配,即使解决前两个报错也无法正常推理
  • 冗余代码:你定义的Tokenizer全程没有参与特征提取流程,属于无效代码,可以直接删除
修复方案

1. 调整标签生成逻辑

将所有生成newy的代码从0维数组改为1维数组:

# 原来的写法
newy=np.array(output)
# 修改为
newy = np.array([output])

# 同理,手动指定标签的地方也修改
newy=np.array('Low') → newy = np.array(['Low'])
newy=np.array('Medium') → newy = np.array(['Medium'])

2. 修正retrain函数逻辑

不要重新拟合CountVectorizer,改用partial_fit做增量训练,同时传入预训练模型已有的类别信息:

def retrain(X,y):
    X=preprocess_text(X)
    X=X.lower()
    X=[X]
    # 不要重新fit CountVectorizer,仅做转换
    X=cv.transform(X)
    # 用partial_fit做增量训练,传入已有的类别避免类别数不足报错
    sgd.partial_fit(X,y, classes=sgd.classes_)
    with open('sgd.pickle', 'wb') as f:
        pickle.dump(sgd, f)
    print("Model trained on new data")

3. 补充持久化优化(可选但推荐)

为了避免重启程序后CountVectorizer的词汇表丢失,建议初始训练完成后将模型和特征提取器一起保存:

# 初始训练完成后保存
with open('model_package.pickle', 'wb') as f:
    pickle.dump((sgd, cv), f)

# 加载时同时读取两个对象
with open('../model_package.pickle', 'rb') as f:
    sgd, cv = pickle.load(f)

4. 训练优化建议

单条样本增量训练容易导致模型效果波动大,建议攒3-5条包含不同类别的样本后再统一做增量训练,效果会更稳定。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 15:45:03