Spark RDD中实现以ARRIVAL_DELAY为key的自定义哈希分区方法
航班RDD按ARRIVAL_DELAY自定义分区实现方案
实现逻辑说明
我们需要先统计ARRIVAL_DELAY列的最大值、最小值确定区间范围,再继承Spark官方的Partitioner类实现自定义分区规则,最后将RDD转换为键值对结构后调用分区方法即可。
完整实现代码
第一步:依赖导入与基础RDD生成
from pyspark import SparkContext, Partitioner from pyspark.sql import Row # 你已有的解析方法直接复用 def import_parse_rdd(data): # create rdd rdd = sc.textFile(data) # remove the header header = rdd.first() rdd = rdd.filter(lambda row: row != header) #filter out header # split by comma split_rdd = rdd.map(lambda line: line.split(',')) row_rdd = split_rdd.map(lambda line: Row( YEAR = int(line[0]),MONTH = int(line[1]),DAY = int(line[2]),DAY_OF_WEEK = int(line[3]) ,AIRLINE = line[4],FLIGHT_NUMBER = int(line[5]), TAIL_NUMBER = line[6],ORIGIN_AIRPORT = line[7],DESTINATION_AIRPORT = line[8], SCHEDULED_DEPARTURE = line[9],DEPARTURE_TIME = line[10],DEPARTURE_DELAY = 0 if "".__eq__(line[11]) else float(line[11]),TAXI_OUT = 0 if "".__eq__(line[12]) else float(line[12]), WHEELS_OFF = line[13],SCHEDULED_TIME = line[14],ELAPSED_TIME = 0 if "".__eq__(line[15]) else float(line[15]),AIR_TIME = 0 if "".__eq__(line[16]) else float(line[16]),DISTANCE = 0 if "".__eq__(line[17]) else float(line[17]),WHEELS_ON = line[18],TAXI_IN = 0 if "".__eq__(line[19]) else float(line[19]), SCHEDULED_ARRIVAL = line[20],ARRIVAL_TIME = line[21],ARRIVAL_DELAY = 0 if "".__eq__(line[22]) else float(line[22]),DIVERTED = line[23],CANCELLED = line[24],CANCELLATION_REASON = line[25],AIR_SYSTEM_DELAY = line[26], SECURITY_DELAY = line[27],AIRLINE_DELAY = line[28],LATE_AIRCRAFT_DELAY = line[29],WEATHER_DELAY = line[30]) ) return row_rdd # 生成基础RDD,*建议缓存避免后续重复计算* flight_rdd = import_parse_rdd("你的航班数据文件路径").cache()
第二步:统计ARRIVAL_DELAY的最大最小值
# 提取所有ARRIVAL_DELAY值 delay_rdd = flight_rdd.map(lambda row: row.ARRIVAL_DELAY) max_delay = delay_rdd.max() min_delay = delay_rdd.min() # 自定义分区数量,可根据实际数据量调整 PARTITION_NUM = 5 # 计算每个分区对应的延迟区间长度 interval = (max_delay - min_delay) / PARTITION_NUM
第三步:实现自定义分区器
class DelayHashPartitioner(Partitioner): def __init__(self, num_partitions, min_delay, interval): self.num_partitions = num_partitions self.min_delay = min_delay self.interval = interval # 返回分区总数 def numPartitions(self): return self.num_partitions # 输入键(这里是ARRIVAL_DELAY值),返回对应的分区id def getPartition(self, key): # 计算当前延迟所属区间 partition_id = int((key - self.min_delay) / self.interval) # 边界处理:等于最大值的记录分到最后一个分区,避免越界 return min(partition_id, self.num_partitions -1)
第四步:执行分区
# partitionBy要求RDD为键值对结构,所以先转换为(ARRIVAL_DELAY, 原Row)格式 key_value_rdd = flight_rdd.map(lambda row: (row.ARRIVAL_DELAY, row)) # 应用自定义分区器 partitioned_rdd = key_value_rdd.partitionBy( PARTITION_NUM, partitionFunc=DelayHashPartitioner(PARTITION_NUM, min_delay, interval) ) # 如果需要恢复为仅保留Row的RDD,提取values即可 final_rdd = partitioned_rdd.values()
效果验证
可以用以下代码查看每个分区的记录数,确认分布符合预期:
partition_count = final_rdd.mapPartitionsWithIndex( lambda idx, iter: [(idx, len(list(iter)))] ).collect() print(partition_count)
内容的提问来源于stack exchange,提问作者jixubi
相关产品推荐
相关产品推荐

