如何在PySpark中实现按优先级选取指定数字的自定义函数
实现思路
预设优先级序列为[1, 3, 2013, 154, 147],优先级从左到右递减,逻辑核心是:
- 先把输入的字符串按逗号拆分,逐个清洗提取有效数字(过滤掉非数字内容)
- 按优先级从高到低遍历预设序列,第一个命中输入中存在的数字就是返回结果
- 若输入中没有匹配的优先级数字,可返回默认值(比如null)
代码实现
方式1:UDF实现(逻辑简单易维护)
适合中小数据量场景,代码可读性高
from pyspark.sql import SparkSession from pyspark.sql.functions import udf from pyspark.sql.types import IntegerType # 初始化Spark spark = SparkSession.builder.appName("priority_select").getOrCreate() # 预设优先级序列,顺序对应优先级从高到低 PRIORITY_LIST = [1, 3, 2013, 154, 147] def select_by_priority(input_str): # 处理空输入 if not input_str: return None # 拆分字符串,清洗提取有效数字 input_nums = set() for item in input_str.split(","): item = item.strip() # 仅保留纯数字内容 if item.isdigit(): input_nums.add(int(item)) # 按优先级遍历,返回第一个匹配值 for num in PRIORITY_LIST: if num in input_nums: return num # 无匹配返回None return None # 注册UDF priority_udf = udf(select_by_priority, IntegerType()) # 功能验证 test_data = [ ("1,3,2,14",), ("151, 152, 2013",), ("3, 147, 151",), ("9999, 154, 8777, 45=3",), ("999,888",) ] df = spark.createDataFrame(test_data, ["a"]) df.withColumn("result", priority_udf("a")).show()
测试输出完全匹配需求:
+--------------------+------+ | a|result| +--------------------+------+ | 1,3,2,14 | 1| | 151, 152, 2013 | 2013| | 3, 147, 151 | 3| |9999, 154, 8777,...| 154| | 999,888 | null| +--------------------+------+
方式2:内置函数实现(性能更高)
避免UDF的序列化开销,适合超大数据量场景
from pyspark.sql import functions as F df = df.withColumn("split_a", F.split(F.col("a"), ",")) # 按优先级顺序逐个判断,取第一个匹配值 df = df.withColumn("result", F.when(F.array_contains(F.transform("split_a", lambda x: F.regexp_extract(F.trim(x), r'^\d+$', 0).cast("int")), 1), 1) .when(F.array_contains(F.transform("split_a", lambda x: F.regexp_extract(F.trim(x), r'^\d+$', 0).cast("int")), 3), 3) .when(F.array_contains(F.transform("split_a", lambda x: F.regexp_extract(F.trim(x), r'^\d+$', 0).cast("int")), 2013), 2013) .when(F.array_contains(F.transform("split_a", lambda x: F.regexp_extract(F.trim(x), r'^\d+$', 0).cast("int")), 154), 154) .when(F.array_contains(F.transform("split_a", lambda x: F.regexp_extract(F.trim(x), r'^\d+$', 0).cast("int")), 147), 147) .otherwise(None) ) df.drop("split_a").show()
输出结果和UDF实现完全一致。
内容的提问来源于stack exchange,提问作者practicalGuy
相关产品推荐
相关产品推荐

