PySpark DataFrame按begin_end区间提取text列子串实现方法
PySpark按索引从文本提取子串实现方案
你不需要强制编写UDF,优先使用Spark内置函数实现的性能远高于Python UDF,两种实现方式如下:
方法1:内置函数实现(推荐,性能最优)
Spark内置函数经过Catalyst优化器优化,没有Python进程和JVM之间的数据序列化开销,大数据量下执行效率是UDF的数倍到数十倍。
由于Spark的substring函数字符串位置从1开始计数,和你给出的0-based索引差1位偏移,直接按偏移计算截取长度即可:
from pyspark.sql import functions as F res = s.withColumn( "new_col", F.substring( str = F.col("text"), # 0-based索引转Spark要求的1-based位置 pos = F.element_at(F.col("begin_end"), 1) + 1, # 截取长度对应Python切片text[a:b+1]的长度 len = F.element_at(F.col("begin_end"), 2) - F.element_at(F.col("begin_end"), 1) + 1 ) )
执行后结果和pandas实现完全一致:
+----------+----------------------------------------------------------------------------------------------------------------------------+-------------+ |begin_end |text |new_col | +----------+----------------------------------------------------------------------------------------------------------------------------+-------------+ |[111, 120]|They say that all cats land on their feet, but this does not apply to my cat. He not only often falls, but also jumps badly.|jumps bad | |[20, 31] |Mom called dad, and when he came home, he took moms car and drove to the store |when he came | +----------+----------------------------------------------------------------------------------------------------------------------------+-------------+
方法2:Python UDF实现(写法直观,性能较差)
如果不想处理位置偏移的换算,也可以写和pandas逻辑几乎一致的UDF,但是因为涉及跨进程数据序列化,仅适合小数据量场景使用:
from pyspark.sql import functions as F from pyspark.sql.types import StringType @F.udf(StringType()) def extract_substr(text, idx_range): start, end = idx_range return text[start:end+1] res = s.withColumn("new_col", extract_substr(F.col("text"), F.col("begin_end")))
内容的提问来源于stack exchange,提问作者Rory
相关产品推荐
相关产品推荐

