PySpark实现列值匹配、复杂关联与列追加需求
PySpark实现列值匹配、复杂关联与列追加操作
需求说明
需要完成以下操作:
- 按
version和cover关联policies与rates表 - 匹配
policies中所有以rf开头的列与rates的RF列 - 匹配
policies的rl列与rates的rfval_1列 - 将匹配到的
rates表中rf_amt、rf_coeff、rf_rate列,按policies里rf列的序号追加到原表,生成目标输出
源表结构与数据
policies表
| pol_no | cover | version | rf1 | rf2 | rl_1 | rl_2 |
|---|---|---|---|---|---|---|
| 123 | a | 3 | apples | pears | 3 | 2 |
rates表
| RF | version | cover | rfval_1 | rfval_2 | rfval_3 | rf_amt | rf_coeff | rf_rate |
|---|---|---|---|---|---|---|---|---|
| apples | 3 | a | 3 | null | null | 0.4 | 0.4 | null |
| pears | 3 | a | 2 | null | null | 0.2 | 0.2 | null |
| oranges | 3 | a | 3 | null | null | 0.3 | 0.3 | null |
期望输出表
| pol_no | cover | version | rf1 | rf2 | rl_1 | rl_2 | rfamt_1 | rfcoeff_1 | rfrate_1 | rf_amt2 | rfcoeff_2 | rf_rate2 |
|---|---|---|---|---|---|---|---|---|---|---|---|---|
| 123 | a | 3 | apples | pears | 3 | 2 | 0.4 | 0.4 | null | 0.2 | 0.2 | null |
实现代码
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.types import StructType, StructField, StringType, IntegerType, FloatType # 初始化SparkSession spark = SparkSession.builder.appName("PolicyRateMatch").getOrCreate() # 1. 创建示例DataFrame # policies表结构与数据 policies_schema = StructType([ StructField("pol_no", StringType(), True), StructField("cover", StringType(), True), StructField("version", IntegerType(), True), StructField("rf1", StringType(), True), StructField("rf2", StringType(), True), StructField("rl_1", IntegerType(), True), StructField("rl_2", IntegerType(), True) ]) policies_data = [("123", "a", 3, "apples", "pears", 3, 2)] policies_df = spark.createDataFrame(policies_data, schema=policies_schema) # rates表结构与数据 rates_schema = StructType([ StructField("RF", StringType(), True), StructField("version", IntegerType(), True), StructField("cover", StringType(), True), StructField("rfval_1", IntegerType(), True), StructField("rfval_2", IntegerType(), True), StructField("rfval_3", IntegerType(), True), StructField("rf_amt", FloatType(), True), StructField("rf_coeff", FloatType(), True), StructField("rf_rate", FloatType(), True) ]) rates_data = [ ("apples", 3, "a", 3, None, None, 0.4, 0.4, None), ("pears", 3, "a", 2, None, None, 0.2, 0.2, None), ("oranges", 3, "a", 3, None, None, 0.3, 0.3, None) ] rates_df = spark.createDataFrame(rates_data, schema=rates_schema) # 2. 处理policies表,将rf和rl列转为长格式(unpivot) # 提取rf列和对应的rl列 rf_cols = [col for col in policies_df.columns if col.startswith("rf")] rl_cols = [col for col in policies_df.columns if col.startswith("rl_")] # 创建unpivot的表达式:将每一组rfX和rl_X转为行 unpivot_expr = """ stack({n}, {stack_items} ) as (rf_col_name, rf_value, rl_col_name, rl_value) """.format( n=len(rf_cols), stack_items=", ".join([f"'{rf}', '{policies_df[rf]}', '{rl}', {policies_df[rl]}" for rf, rl in zip(rf_cols, rl_cols)]) ) policies_unpivot_df = policies_df.select( "pol_no", "cover", "version", F.expr(unpivot_expr) ).withColumn("seq", F.regexp_extract("rf_col_name", r"rf(\d+)", 1)) # 提取序号,比如rf1的序号是1 # 3. 关联rates表,匹配条件:version、cover、RF=rf_value、rfval_1=rl_value joined_df = policies_unpivot_df.join( rates_df, (policies_unpivot_df.version == rates_df.version) & (policies_unpivot_df.cover == rates_df.cover) & (policies_unpivot_df.rf_value == rates_df.RF) & (policies_unpivot_df.rl_value == rates_df.rfval_1), how="left" ).select( "pol_no", "cover", "version", "rf_col_name", "rf_value", "rl_col_name", "rl_value", "seq", "rf_amt", "rf_coeff", "rf_rate" ) # 4. 将关联结果转回宽格式(pivot) pivoted_df = joined_df.groupBy("pol_no", "cover", "version").pivot("seq").agg( F.first("rf_amt").alias("rf_amt"), F.first("rf_coeff").alias("rf_coeff"), F.first("rf_rate").alias("rf_rate") ) # 5. 合并原policies表与pivoted表,并调整列顺序 # 整理目标列顺序 original_cols = policies_df.columns new_cols = [] for seq in sorted([str(i) for i in range(1, len(rf_cols)+1)]): new_cols.extend([f"rf_amt{seq}", f"rf_coeff{seq}", f"rf_rate{seq}"]) # 合并并调整列顺序 final_df = policies_df.join(pivoted_df, on=["pol_no", "cover", "version"], how="left") final_df = final_df.select(original_cols + new_cols) # 查看结果 final_df.show(truncate=False)
代码说明
- 数据初始化:创建示例的
policies和ratesDataFrame,模拟源数据结构。 - Unpivot转换:将
policies中分散的rfX和rl_X列转为长格式,方便后续关联匹配,同时提取列的序号用于后续转宽格式。 - 关联匹配:基于
version、cover、RF值匹配、rfval_1值匹配,关联两张表,获取对应的费率数据。 - Pivot转换:将关联后的长格式数据转回宽格式,按序号生成对应的
rf_amtX、rf_coeffX、rf_rateX列。 - 合并与列调整:将原表数据与转换后的费率列合并,并调整列顺序,得到目标输出。
内容的提问来源于stack exchange,提问作者555codewiz
相关产品推荐
相关产品推荐

