PySpark从Numpy数组创建DataFrame时维度不匹配报错求助
问题解决:从Numpy数组创建3列PySpark DataFrame报错
错误原因
你遇到的ValueError: Shape of passed values is (125, 3), indices imply (125, 2),本质是PySpark对Numpy数组的解析逻辑和预期不符,导致列数识别错误。虽然你的Numpy数组形状是(125,3),但直接传递给spark.createDataFrame时,内部处理可能出现结构误判。
解决方案
方案1:直接使用Pandas DataFrame(最简洁可靠)
无需将Pandas DataFrame转成Numpy数组,直接传递给spark.createDataFrame即可,PySpark能正确识别列结构:
prob_row_schema = StructType([ StructField('field1', IntegerType(), True), StructField('field2', IntegerType(), True), StructField('field3', IntegerType(), True), ]) df1= pd.DataFrame(range(1,5+1), columns=['field1']) df2= pd.DataFrame(range(1,5+1), columns=['field2']) df3= pd.DataFrame(range(1,5+1), columns=['field3']) df= df1.join(df2, how='cross').join(df3, how='cross') # 直接传入Pandas DataFrame,无需转Numpy数组 df_prob = spark.createDataFrame(df, schema=prob_row_schema)
方案2:将Numpy数组转为列表的列表
如果必须使用Numpy数组,显式将其转为嵌套列表,让PySpark能正确识别每一行的元素:
prob_row_schema = StructType([ StructField('field1', IntegerType(), True), StructField('field2', IntegerType(), True), StructField('field3', IntegerType(), True), ]) df1= pd.DataFrame(range(1,5+1), columns=['field1']) df2= pd.DataFrame(range(1,5+1), columns=['field2']) df3= pd.DataFrame(range(1,5+1), columns=['field3']) df= df1.join(df2, how='cross').join(df3, how='cross') x = df.values # 将Numpy数组转为列表的列表 x_list = x.tolist() df_prob = spark.createDataFrame(data=x_list, schema = prob_row_schema)
验证
两种方案都能正确生成包含125行、3列的PySpark DataFrame,你可以通过df_prob.printSchema()和df_prob.count()确认结果。
内容的提问来源于stack exchange,提问作者mikest31
相关产品推荐
相关产品推荐

