PySpark MLP分类器评估报A & B维度不匹配错误
问题根因
报错触发点在多层感知机前向传播的矩阵乘法环节,核心错误是MLP网络层配置维度和实际输入特征维度不匹配,对应栈里的A & B Dimension mismatch报错。MultilayerPerceptronClassifier的layers参数有严格的配置规则:
- 列表第一位:必须和输入特征向量的维度完全一致
- 列表最后一位:必须和分类任务的类别总数完全一致
- 中间位:对应各隐藏层的神经元数量,可按需调整
你当前任务是3分类,最后一层设为3是正确的,但你手动设置的第一层维度11,和Pipeline前置流程输出的features列实际维度不符。你之前对接的其他分类器(如随机森林、逻辑回归)不需要手动声明输入维度,会自动适配特征长度,所以之前流程能正常运行,而MLP要求手动声明全网络层维度,参数写错就会触发矩阵运算维度错误。
修复步骤
- 先获取实际特征维度,不要靠臆测写输入层大小:
from pyspark.ml import Pipeline # 用未加MLP的前置Pipeline处理训练集,拿到真实特征维度 prep_pipe = Pipeline(stages=stages.copy()) prep_model = prep_pipe.fit(train) sample = prep_model.transform(train).first() real_feature_dim = len(sample.features) - 修改MLP的
layers参数,把第一位替换成上一步拿到的real_feature_dim,最后一位保持3即可。隐藏层神经元数可以根据训练效果调整,只要首尾维度匹配就不会再报维度错误。 - 你贴的代码里评估指标计算、打印的逻辑重复写了两遍,直接删掉冗余部分即可,不影响运行但会产生重复计算。
修正后可运行代码
train, test, validation = df.randomSplit([0.7, 0.2, 0.1], 1234) # 先获取真实特征维度 prep_pipe = Pipeline(stages=stages.copy()) prep_model = prep_pipe.fit(train) real_feature_dim = len(prep_model.transform(train).first().features) # 正确配置MLP层维度 mlp = MultilayerPerceptronClassifier( labelCol='label', featuresCol='features', maxIter=100, layers=[real_feature_dim, 6, 4, 3], seed=1234 ) stages.append(mlp) pipeline = Pipeline(stages=stages) model = pipeline.fit(train) pred = model.transform(test) # 初始化评估器 acc_evaluator = MulticlassClassificationEvaluator(labelCol="label", predictionCol="prediction", metricName="accuracy") prec_evaluator = MulticlassClassificationEvaluator(labelCol="label", predictionCol="prediction", metricName="weightedPrecision") recall_evaluator = MulticlassClassificationEvaluator(labelCol="label", predictionCol="prediction", metricName="weightedRecall") f1_evaluator = MulticlassClassificationEvaluator(labelCol="label", predictionCol="prediction", metricName="f1") # 计算指标 accuracy = acc_evaluator.evaluate(pred) precision = prec_evaluator.evaluate(pred) recall = recall_evaluator.evaluate(pred) f1 = f1_evaluator.evaluate(pred) # 打印结果 print("Test Error = %g" % (1.0 - accuracy)) print("Accuracy = %g" % (accuracy)) print("Precision = %g" % (precision)) print("Recall = %g" % (recall)) print("F1 = %g" % (f1))
注意:Spark ML的多分类要求标签值是从0开始的连续整数,3分类对应标签应为0、1、2,如果你的标签编码不符合这个规则,后续可能触发标签范围错误,提前做字符串/离散值到0基索引的转换即可。
内容的提问来源于stack exchange,提问作者Lorenzo Maggio
相关产品推荐
相关产品推荐

