如何修改PySpark代码,获取交集最多且长度最短的col1元素
问题
现有如下PySpark代码,用于获取col1中与col2拥有最多共同token的元素:
c1_arr = F.col('col1') c2_arr = F.split(F.trim('col2'), '\s+') arr_of_struct = F.transform( c1_arr, lambda x: F.struct( F.size(F.array_intersect(c2_arr, F.split(F.trim(x), '\s+'))).alias('cnt'), x.alias('val'), ) ) top_val = F.sort_array(arr_of_struct, False)[0]
现在需要修改代码,以获取col1中同时满足与col2拥有最多共同token且长度最短的元素。
示例数据
col1 col2 ["come and get", "computer", "come and get more" ] "come for good" ["summer is hot", "summer is too hot", "hot weather"] "hot tea" ["summer is hot", "summer is too hot", "hot weather"] "hot summer"
期望输出
col1 col2 match ["come and get", "computer", "come and get more" ] "come for good" "come and get" ["summer is hot", "summer is too hot", "hot weather"] "hot tea" "hot weather" ["summer is hot", "summer is too hot", "hot weather"] "hot summer" "summer is hot"
我尝试了以下代码,但不确定是否正确或如何优化:
df = df.select( '*', F.when(((top_val['cnt'] > 1) & (F.size(c2_arr) > 1))|((top_val['cnt'] > 0) & (F.size(c2_arr) == 1))|((F.size(F.split(F.trim(top_val['val']),'\s+'))==1) &(top_val['cnt']>0 )),top_val['val']).alias('match'))
上述代码的三个条件含义:
- 当交集数量大于1且
col2的token数量大于1时,选中匹配元素; - 当
col2的token数量等于1时,交集数量需大于0; - 当交集数量为1且
col1的元素token数量为1时,选中该元素。
请问如何修改原代码以正确实现需求?
解决方案
你原来的代码思路存在缺陷:仅靠sort_array按cnt降序排序,无法处理cnt相同但需要筛选长度最短元素的场景。正确的做法是在生成结构体时加入元素的token长度作为第二排序字段,再调整排序逻辑:
修改后的完整代码
from pyspark.sql import functions as F # 定义列引用 c1_arr = F.col('col1') c2_arr = F.split(F.trim(F.col('col2')), '\s+') # 原代码漏了F.col(),需补全列引用 # 生成包含共同token数、元素长度、元素值的结构体数组 arr_of_struct = F.transform( c1_arr, lambda x: F.struct( # 计算与col2的共同token数量 F.size(F.array_intersect(c2_arr, F.split(F.trim(x), '\s+'))).alias('cnt'), # 计算当前元素的token长度,用于二次排序 F.size(F.split(F.trim(x), '\s+')).alias('val_len'), x.alias('val'), ) ) # 排序规则:先按共同token数降序,再按元素长度升序,取第一个元素 top_val = F.sort_array(arr_of_struct, [False, True])[0] # 提取最终匹配结果 df = df.select( '*', top_val['val'].alias('match') )
关键修改说明
- 补全列引用:原代码中
F.split(F.trim('col2'), '\s+')的'col2'是字符串,不是列对象,必须改为F.col('col2')才能正确读取列值。 - 新增长度字段:在结构体中加入
val_len(元素的token长度),作为cnt相同时的筛选依据。 - 调整排序逻辑:
sort_array支持传入排序方向数组,[False, True]表示先按cnt降序(False),再按val_len升序(True),确保相同cnt下,长度最短的元素排在最前面。 - 简化匹配逻辑:不需要额外的
when条件判断,排序后的第一个元素直接满足“最多共同token+最短长度”的要求,逻辑更简洁可靠。
示例验证
- 第一行:
col1中"come and get"和"come and get more"的cnt都是1,但前者长度更短,被选中; - 第二行:
col1中"summer is hot"和"hot weather"的cnt都是1,"hot weather"长度更短,被选中; - 第三行:
"summer is hot"的cnt为2(匹配col2的两个token),是所有元素中最高的,直接被选中。
内容的提问来源于stack exchange,提问作者user15649753
相关产品推荐
相关产品推荐

