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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 02:39:03