You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何修改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')
)

关键修改说明

  1. 补全列引用:原代码中F.split(F.trim('col2'), '\s+')的'col2'是字符串,不是列对象,必须改为F.col('col2')才能正确读取列值。
  2. 新增长度字段:在结构体中加入val_len(元素的token长度),作为cnt相同时的筛选依据。
  3. 调整排序逻辑:sort_array支持传入排序方向数组,[False, True]表示先按cnt降序(False),再按val_len升序(True),确保相同cnt下,长度最短的元素排在最前面。
  4. 简化匹配逻辑:不需要额外的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.23 11:45:28