PySpark求单行DataFrame与多行DataFrame列的交集方法
单行与多行DataFrame的列交集高效实现方案
需求说明
现有两个DataFrame:
- 单行DataFrame:
+-----+--------------------+ | col1| col2| +-----+--------------------+ | A | [B, C, D] | +-----+--------------------+
- 多行DataFrame:
+----------+--------------------+ | col1| col2| +----------+--------------------+ | F |[A, B, C] | | G |[J, K, B] | | H |[C, H, D] | +----------+--------------------+
需要计算多行DataFrame中每行col2与单行DataFrame的col2的交集,得到如下结果:
+----------+--------------------+ | col1| col2| +----------+--------------------+ | F |[B, C] | | G |[B] | | H |[C, D] | +----------+--------------------+
高效实现方法
1. PySpark 版本
Spark内置的array_intersect函数可直接计算数组交集,性能优异,适合大数据场景:
from pyspark.sql import SparkSession from pyspark.sql.functions import array_intersect, lit # 初始化Spark会话 spark = SparkSession.builder.appName("array_intersection").getOrCreate() # 构造示例数据 single_df = spark.createDataFrame([("A", ["B", "C", "D"])], ["col1", "col2"]) multi_df = spark.createDataFrame( [("F", ["A", "B", "C"]), ("G", ["J", "K", "B"]), ("H", ["C", "H", "D"])], ["col1", "col2"] ) # 提取单行DataFrame中的基准数组 base_array = single_df.select("col2").collect()[0][0] # 计算每行的数组交集 result_df = multi_df.withColumn("col2", array_intersect(multi_df["col2"], lit(base_array))) # 查看结果 result_df.show(truncate=False)
2. Pandas 版本
针对小数据量可用直观的apply方法,大数据量推荐向量化的展开-筛选-聚合流程:
方法A:apply + 集合交集(简单直观)
import pandas as pd # 构造示例数据 single_df = pd.DataFrame({"col1": ["A"], "col2": [["B", "C", "D"]]}) multi_df = pd.DataFrame({ "col1": ["F", "G", "H"], "col2": [["A", "B", "C"], ["J", "K", "B"], ["C", "H", "D"]] }) # 转换为集合提升交集计算效率 base_set = set(single_df["col2"].iloc[0]) # 逐行计算交集 multi_df["col2"] = multi_df["col2"].apply(lambda x: list(set(x) & base_set)) print(multi_df)
方法B:向量化操作(大数据量高效)
避免apply的循环开销,通过向量化操作实现:
import pandas as pd # 构造示例数据 single_df = pd.DataFrame({"col1": ["A"], "col2": [["B", "C", "D"]]}) multi_df = pd.DataFrame({ "col1": ["F", "G", "H"], "col2": [["A", "B", "C"], ["J", "K", "B"], ["C", "H", "D"]] }) # 提取基准值 base_values = single_df["col2"].explode().unique() # 展开多行DataFrame的数组列 exploded_multi = multi_df.explode("col2") # 筛选出属于基准集合的值 filtered = exploded_multi[exploded_multi["col2"].isin(base_values)] # 重新聚合为数组 result_df = filtered.groupby("col1")["col2"].agg(list).reset_index() print(result_df)
内容的提问来源于stack exchange,提问作者A.M.
相关产品推荐
相关产品推荐

