PySpark DataFrame中带排除规则的列值匹配最近ratecodeid
PySpark实现:为offer匹配非排除类最近ratecodeid
数据加载代码
dfamr = spark.read.csv("/C:/Sushant Workspace/Tech/Self/pyspark datasets/amrrates.csv", header="true") dfexclusion = spark.read.csv("/C:/Sushant Workspace/Tech/Self/pyspark datasets/exclusionrates.csv", header="true") dfamr.show()
数据说明
dfamr包含字段:ratecodeid、rate、offer1、offer2(每条记录对应一个ratecode的基准rate,以及两个需要匹配的offer值)dfexclusion为排除列表,包含需要跳过的ratecodeid(如R4、R5)
需求与规则
核心需求
为dfamr的每条记录:
- 匹配与
offer1数值最接近的ratecodeid,保存为offer1CodeId - 匹配与
offer2数值最接近的ratecodeid,保存为offer2CodeId
匹配规则
- 计算offer值与所有ratecode的
rate值的差值绝对值,差值越小越优先 - 若最优先的匹配项属于排除列表(如R4、R5),则跳过该选项,选择次优先的非排除项;若次优先项也在排除列表,继续往后筛选,直到找到符合要求的ratecodeid
示例说明
- 当
ratecodeid=R5、offer1=5.5时:最近匹配是R4(rate=5.4),但R4在排除列表,故选择次近的R3(rate=5.3) - 当
ratecodeid=R7、offer2=5.5时:最近匹配是R4(rate=5.4),R4在排除列表,故选择R3 - 当
ratecodeid=R6、offer1=6时:最近匹配是R5(rate=5.85),R5在排除列表;次近是R4(rate=5.4),也在排除列表,故选择R6自身(rate=6)
实现代码
方案1:使用Spark 3.0+的min_by函数(简洁版)
from pyspark.sql.functions import col, abs, min_by # 1. 转换字段为数值类型 dfamr = dfamr.withColumn("offer1", col("offer1").cast("double")) \ .withColumn("offer2", col("offer2").cast("double")) \ .withColumn("rate", col("rate").cast("double")) # 2. 提取所有ratecode的基准rate数据 rate_ref = dfamr.select("ratecodeid", "rate").alias("rate_ref") # 3. 获取排除列表集合 exclusion_ids = set(row["ratecodeid"] for row in dfexclusion.collect()) # 4. 匹配offer1对应的非排除最近ratecodeid offer1_match = dfamr.alias("main").crossJoin(rate_ref) \ .filter(~col("rate_ref.ratecodeid").isin(exclusion_ids)) \ .withColumn("diff", abs(col("main.offer1") - col("rate_ref.rate"))) \ .groupBy("main.ratecodeid", "main.offer1", "main.offer2", "main.rate") \ .agg(min_by("rate_ref.ratecodeid", "diff").alias("offer1CodeId")) # 5. 匹配offer2对应的非排除最近ratecodeid offer2_match = dfamr.alias("main").crossJoin(rate_ref) \ .filter(~col("rate_ref.ratecodeid").isin(exclusion_ids)) \ .withColumn("diff", abs(col("main.offer2") - col("rate_ref.rate"))) \ .groupBy("main.ratecodeid", "main.offer1", "main.offer2", "main.rate") \ .agg(min_by("rate_ref.ratecodeid", "diff").alias("offer2CodeId")) # 6. 合并两个匹配结果 final_result = offer1_match.join(offer2_match, on=["ratecodeid", "offer1", "offer2", "rate"], how="inner") # 查看最终结果 final_result.show()
方案2:使用窗口函数(兼容Spark 2.x版本)
from pyspark.sql.functions import col, abs, row_number from pyspark.sql.window import Window # 1. 转换字段类型 dfamr = dfamr.withColumn("offer1", col("offer1").cast("double")) \ .withColumn("offer2", col("offer2").cast("double")) \ .withColumn("rate", col("rate").cast("double")) # 2. 提取基准rate数据 rate_ref = dfamr.select("ratecodeid", "rate").alias("rate_ref") # 3. 获取排除列表 exclusion_ids = set(row["ratecodeid"] for row in dfexclusion.collect()) # 处理offer1匹配 window_offer1 = Window.partitionBy("main.ratecodeid", "main.offer1", "main.offer2", "main.rate") \ .orderBy("diff") offer1_match = dfamr.alias("main").crossJoin(rate_ref) \ .filter(~col("rate_ref.ratecodeid").isin(exclusion_ids)) \ .withColumn("diff", abs(col("main.offer1") - col("rate_ref.rate"))) \ .withColumn("rank", row_number().over(window_offer1)) \ .filter(col("rank") == 1) \ .select("main.ratecodeid", "main.offer1", "main.offer2", "main.rate", col("rate_ref.ratecodeid").alias("offer1CodeId")) # 处理offer2匹配 window_offer2 = Window.partitionBy("main.ratecodeid", "main.offer1", "main.offer2", "main.rate") \ .orderBy("diff") offer2_match = dfamr.alias("main").crossJoin(rate_ref) \ .filter(~col("rate_ref.ratecodeid").isin(exclusion_ids)) \ .withColumn("diff", abs(col("main.offer2") - col("rate_ref.rate"))) \ .withColumn("rank", row_number().over(window_offer2)) \ .filter(col("rank") == 1) \ .select("main.ratecodeid", "main.offer1", "main.offer2", "main.rate", col("rate_ref.ratecodeid").alias("offer2CodeId")) # 合并结果 final_result = offer1_match.join(offer2_match, on=["ratecodeid", "offer1", "offer2", "rate"], how="inner") final_result.show()
代码说明
crossJoin:将每条主表记录与所有基准ratecode关联,确保能计算offer与所有rate的差值- 过滤逻辑:直接排除掉在排除列表中的ratecodeid,避免后续无效计算
- 匹配逻辑:通过差值绝对值排序,取最小差值对应的ratecodeid;窗口函数版本通过
row_number()标记排序后的位次,取第一位即为最近匹配项
内容的提问来源于stack exchange,提问作者sushapat
相关产品推荐
相关产品推荐

