PySpark DataFrame遍历数组生成对应键值对新列的实现问询
问题描述
现有如下PySpark DataFrame:
| name | fruits | apple | banana | orange |
|---|---|---|---|---|
| Alice | ["apple","banana","orange"] | 5 | 8 | 3 |
| Bob | ["apple"] | 2 | 9 | 1 |
需要生成新列new_col,存储键值对(最终可转为JSON格式):键为fruits数组中的元素,值为对应列的数值,最终效果如下:
| name | fruits | apple | banana | orange | new_col |
|---|---|---|---|---|---|
| Alice | ["apple","banana","orange"] | 5 | 8 | 3 | {"apple":5, "banana":8, "orange":3} |
| Bob | ["apple"] | 2 | 9 | 1 | {"apple":2} |
用户尝试用UDF实现,但语法有误,现有代码如下:
from pyspark.sql.functions import udf, col from pyspark.sql.types import MapType, StringType from pyspark.sql import SparkSession # Create a Spark session spark = SparkSession.builder.appName("example").getOrCreate() # Sample data data = [("Alice", ["apple", "banana", "orange"], 5, 8, 3), ("Bob", ["apple"], 2, 9, 1)] # Define the schema schema = ["name", "fruits", "apple", "banana", "orange"] # Create a DataFrame df = spark.createDataFrame(data, schema=schema) # Show the initial DataFrame print("Initial DataFrame:") display(df) # Define a UDF to create a dictionary @udf(MapType(StringType(), StringType())) def json_map(fruits): result = {} for i in fruits: result[i] = col(i) return result # Apply the UDF to the 'fruits' column new_df = df.withColumn('test', json_map(col('fruits'))) # Display the updated DataFrame display(new_df)
解决方案
方法一:使用PySpark原生函数(推荐,避免UDF性能损耗)
利用map_from_arrays和transform函数实现,无需自定义UDF:
from pyspark.sql.functions import map_from_arrays, transform, col, to_json from pyspark.sql import SparkSession spark = SparkSession.builder.appName("example").getOrCreate() data = [("Alice", ["apple", "banana", "orange"], 5, 8, 3), ("Bob", ["apple"], 2, 9, 1)] schema = ["name", "fruits", "apple", "banana", "orange"] df = spark.createDataFrame(data, schema=schema) # 生成Map类型列,若需JSON字符串则用to_json包裹 df = df.withColumn( "new_col", map_from_arrays( col("fruits"), transform(col("fruits"), lambda x: col(x)) ) ) # 若需要JSON格式字符串,替换上面的withColumn为: # df = df.withColumn( # "new_col", # to_json(map_from_arrays(col("fruits"), transform(col("fruits"), lambda x: col(x)))) # ) display(df)
说明
transform(col("fruits"), lambda x: col(x)):遍历fruits数组,取出每个元素对应列的数值,生成数值数组。map_from_arrays:将键数组(fruits)和数值数组配对,生成Map类型列;用to_json可直接转为JSON字符串。
方法二:修正自定义UDF
原UDF的问题是:UDF内部无法直接用col(i)引用Spark列,需将所需列作为参数传入:
from pyspark.sql.functions import udf, col, to_json from pyspark.sql.types import MapType, StringType, IntegerType from pyspark.sql import SparkSession spark = SparkSession.builder.appName("example").getOrCreate() data = [("Alice", ["apple", "banana", "orange"], 5, 8, 3), ("Bob", ["apple"], 2, 9, 1)] schema = ["name", "fruits", "apple", "banana", "orange"] df = spark.createDataFrame(data, schema=schema) @udf(MapType(StringType(), IntegerType())) def json_map(fruits, apple, banana, orange): fruit_dict = {"apple": apple, "banana": banana, "orange": orange} return {fruit: fruit_dict[fruit] for fruit in fruits} # 调用UDF时传入所有水果列 new_df = df.withColumn( "new_col", json_map(col("fruits"), col("apple"), col("banana"), col("orange")) ) # 若需转为JSON字符串: # new_df = new_df.withColumn("new_col", to_json(col("new_col"))) display(new_df)
说明
- UDF需显式接收所有水果列的数值,因为UDF内部处理的是Python原生数据,无法动态引用Spark列。
- 返回类型定义为
MapType(StringType(), IntegerType())更贴合数值类型,后续可通过to_json转为JSON字符串。
内容的提问来源于stack exchange,提问作者buttermilk
相关产品推荐
相关产品推荐

