LinearDiscriminantAnalysis调用fit报shrinkage not supported如何解决
报错根因
该报错是scikit-learn的原生设计限制,和版本、运行环境无关:LinearDiscriminantAnalysis的默认求解器为solver='svd',该求解器本身不支持shrinkage参数,只要同时使用默认svd求解器和非空的shrinkage参数就会触发该报错。
修复方案
- 方案1:保留shrinkage参数,显式指定支持该参数的求解器
支持shrinkage的求解器为lsqr和eigen,修改初始化LDA的代码即可,示例如下:
import numpy as np from sklearn.discriminant_analysis import LinearDiscriminantAnalysis # 显式指定solver为lsqr,可替换为eigen lda = LinearDiscriminantAnalysis(solver='lsqr', shrinkage='auto') # 注意原代码中y的生成方式有误,已调整为符合维度要求的标签 lda.fit(np.random.rand(3,2), np.random.randint(0, 2, size=3))
- 方案2:不需要使用shrinkage功能的话,直接删除该参数即可,沿用默认svd求解器:
import numpy as np from sklearn.discriminant_analysis import LinearDiscriminantAnalysis lda = LinearDiscriminantAnalysis() lda.fit(np.random.rand(3,2), np.random.randint(0, 2, size=3))
额外说明:你原有代码中
np.random.randint((1,1,1))的调用方式不符合参数规范,传入元组时会被识别为输出数组的shape参数,无法得到和样本数量匹配的一维标签数组,需要调整为np.random.randint(取值下限, 取值上限, size=样本数)的格式。
内容的提问来源于stack exchange,提问作者Nathane Berrebi
相关产品推荐
相关产品推荐

