如何为PySpark DataFrame新增列:计算数组与列的点积
PySpark计算每行数组与全列value数组的点积
需求说明
现有PySpark DataFrame包含两列:weights(浮点型数组,长度等于DataFrame行数)和value(浮点型)。需新增result列,每行值为该行weights数组与所有行value组成的数组的点积。
示例数据
先创建示例DataFrame:
from pyspark.sql import SparkSession spark = SparkSession.builder.appName("dot_product").getOrCreate() data = [ ([0.0,5.0,4.0,3.0,2.0,1.0,0.0,1.0,2.0,3.0,4.0,5.0],34), ([5.0,0.0,5.0,4.0,3.0,2.0,1.0,0.0,1.0,2.0,3.0,4.0],50), ([4.0,5.0,0.0,5.0,4.0,3.0,2.0,1.0,0.0,1.0,2.0,3.0],56), ([3.0,4.0,5.0,0.0,5.0,4.0,3.0,2.0,1.0,0.0,1.0,2.0],45), ([2.0,3.0,4.0,5.0,0.0,5.0,4.0,3.0,2.0,1.0,0.0,1.0],34), ([1.0,2.0,3.0,4.0,5.0,0.0,5.0,4.0,3.0,2.0,1.0,0.0],36), ([0.0,1.0,2.0,3.0,4.0,5.0,0.0,5.0,4.0,3.0,2.0,1.0],45), ([1.0,0.0,1.0,2.0,3.0,4.0,5.0,0.0,5.0,4.0,3.0,2.0],50), ([2.0,1.0,0.0,1.0,2.0,3.0,4.0,5.0,0.0,5.0,4.0,3.0],57), ([3.0,2.0,1.0,0.0,1.0,2.0,3.0,4.0,5.0,0.0,5.0,4.0],39), ([4.0,3.0,2.0,1.0,0.0,1.0,2.0,3.0,4.0,5.0,0.0,5.0],48), ([5.0,4.0,3.0,2.0,1.0,0.0,1.0,2.0,3.0,4.0,5.0,0.0],39) ] df = spark.createDataFrame(data, ["weights", "value"])
实现步骤
首先收集value列的所有值并广播(分布式场景下广播可减少数据传输开销):
from pyspark.sql import functions as F # 收集value列所有值 value_list = df.select("value").rdd.flatMap(lambda x: x).collect() # 广播数组 broadcast_values = spark.sparkContext.broadcast(value_list)
方法一:使用Spark内置函数(推荐)
无需依赖第三方库,分布式计算效率更高:
df_result = df.withColumn( "result", F.aggregate( # 将weights数组与value数组按位置配对 F.arrays_zip(F.col("weights"), F.array([F.lit(v) for v in broadcast_values.value])), # 初始累加值 F.lit(0.0), # 累加逻辑:每对元素相乘后加到累加器 lambda acc, pair: acc + pair.weights * pair[1], # 返回最终累加结果 lambda acc: acc ) ) # 查看结果 df_result.select("weights", "value", "result").show(truncate=False)
方法二:使用Pandas UDF(代码简洁)
适合熟悉numpy的场景,代码更直观:
from pyspark.sql.functions import pandas_udf import pandas as pd import numpy as np @pandas_udf("double") def calculate_dot_product(weights: pd.Series) -> pd.Series: value_array = np.array(broadcast_values.value) return weights.apply(lambda w_arr: np.dot(w_arr, value_array)) df_result = df.withColumn("result", calculate_dot_product(F.col("weights"))) # 查看结果 df_result.select("weights", "value", "result").show(truncate=False)
结果示例
最终result列的计算结果符合numpy.dot(row['weights'], value_list)的预期,例如第一行的result值为:0.0*34 +5.0*50 +4.0*56 +3.0*45 +2.0*34 +1.0*36 +0.0*45 +1.0*50 +2.0*57 +3.0*39 +4.0*48 +5.0*39 = 1690.0
内容的提问来源于stack exchange,提问作者AndreyF
相关产品推荐
相关产品推荐

