基于TensorFlow训练随机森林回归模型时遇维度错误求助
TensorFlow随机森林回归训练报错修复方案
问题场景
我在尝试用TensorFlow训练随机森林回归模型处理连续数值数据时,调用estimator.fit()后先输出了森林参数日志,随后抛出了这个维度错误:
ValueError: Shape must be at least rank 2 but is rank 1 for 'concat' (op: 'ConcatV2') with input shapes: [?], [?], [?], [] and with computed input tensors: input[3] = <1>.
错误原因分析
排查代码后,发现两个核心问题:
- 参数拼写错误:定义
ForestHParams时,把feature_columns拼成了feature_colums(少了一个字母n),导致特征列没有被正确传递给模型,模型无法识别输入特征的结构,进而引发维度拼接错误。 - 特征数量不匹配:
num_features参数设置为2,但实际用到了3个特征(Average_Score、lat、lng),参数和实际特征数量不一致,打乱了模型对输入维度的预期。
修复后的完整代码
import tensorflow as tf from tensorflow.contrib.tensor_forest.python import tensor_forest from tensorflow.contrib.tensor_forest.client import random_forest from tensorflow.python.estimator.inputs import numpy_io import pandas as pd import numpy as np def getFeatures(): Average_Score = tf.feature_column.numeric_column('Average_Score') lat = tf.feature_column.numeric_column('lat') lng = tf.feature_column.numeric_column('lng') return [Average_Score, lat, lng] # 导入并预处理酒店数据 Hotel_Reviews = pd.read_csv("./DataMining/Hotel_Reviews.csv") # 优化过滤逻辑:保留经纬度都非空的有效样本 Hotel_Reviews_Filtered = Hotel_Reviews[(Hotel_Reviews.lat.notnull()) & (Hotel_Reviews.lng.notnull())] Hotel_Reviews_Filtered_Target = Hotel_Reviews_Filtered[["Reviewer_Score"]] Hotel_Reviews_Filtered_Features = Hotel_Reviews_Filtered[["Average_Score","lat","lng"]] # 转换为模型兼容的输入格式 x = {key: np.array(values) for key, values in Hotel_Reviews_Filtered_Features.to_dict('list').items()} y = Hotel_Reviews_Filtered_Target.values # 修正参数拼写和特征数量,调用fill()补全默认参数 params = tf.contrib.tensor_forest.python.tensor_forest.ForestHParams( feature_columns=getFeatures(), # 修复拼写错误 num_classes=1, num_features=3, # 匹配实际的3个特征 regression=True, num_trees=10, max_nodes=1000 ).fill() # 构建随机森林估算器 iest = random_forest.TensorForestEstimator( params, graph_builder_class=tensor_forest.RandomForestGraphs ) # 定义训练输入函数 train_input_fn = numpy_io.numpy_input_fn( x=x, y=y, batch_size=1000, num_epochs=1, shuffle=True ) # 启动训练 iest.fit(input_fn=train_input_fn, steps=500)
额外优化提示
- 数据过滤逻辑:原代码用
|(或)保留样本,会混入只有单一坐标的无效数据,改成&(且)能确保每个样本都有完整的经纬度信息,提升模型效果。 - 参数补全:调用
ForestHParams的fill()方法可以自动补全TensorFlow随机森林的默认参数,避免遗漏必要配置。
内容的提问来源于stack exchange,提问作者a_parida
相关产品推荐
相关产品推荐

