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

使用Dask与PyTorch时遭遇递归深度超限问题及优化需求

问题描述

我使用PyTorch DataLoader从HDF5文件读取温度数组,借助num_workers参数在__getitem__方法返回前完成温度数据预处理。数据按y、x、year索引,每组对应一个数组;由于数据量极大,采用Dask避免全量加载内存,目标生成[y, x, year, band, day]结构的HDF5文件(band0为修正温度,band1为读数日期)。但保存Dask文件时触发以下错误:

RecursionError: maximum recursion depth exceeded while calling a Python object

尝试过将temperature_array转为Dask数组、调整compute调用位置、修改batch size等操作。目前通过sys.setrecursionlimit(3000)临时解决,但担心无法适配更大数据集,寻求替代方案。

相关代码

class TemperatureDataset(Dataset):
    '''
    Creates accumulated degree day timeseries for the specified base.

    Requires some sort of temperature related data to accumulate, and the dates of each image accordingly.
    '''
    # 省略部分代码

    def __getitem__(self, idx):
        # 获取年份和点位,累积该点位对应年份的数据
        temperature_file = h5py.File(self.temperature_path, 'r')
        temperature_data = da.from_array(temperature_file['/data'])
        
        date_file = h5py.File(self.date_path, 'r')
        dates = date_file['images_date'][:].astype(str)
        
        # 调整索引适配多维数组
        shape_of_array = (len(temperature_data[0,:,0]), len(temperature_data[0,0,:]), number_of_years) 
        index = np.unravel_index(idx, shape_of_array)

        y = index[0]
        x = index[1]
        year = int(index[2] + 2001)

        acc_temperature = []
        start = 0
        # 遍历日期找到对应年份的起始位置
        for i in range(len(temperature_data[:,0,0])):
            if int(dates[i][:4]) < year:
                continue
            if int(dates[i][:4]) == year:
                start = i
                # 计算该点位对应年份的累积温度
                acc_temperature = np.cumsum(temperature_data[i:(i + self.T), index[0], index[1]].compute())
                break

        # 处理日期(可忽略具体逻辑)
        date_int_list = []
        for i in range(start, start + NUMBER_OF_DATA_DAYS):
            datetime_obj = datetime.strptime(dates[i], '%Y-%m-%d')
            date_int = int(datetime_obj.strftime("%Y%m%d"))
            date_int_list.append(date_int)

        return acc_temperature, torch.tensor(date_int_list), index[0], index[1], year


if __name__ == '__main__':
    data = TemperatureDataset(path_to_temperature_data, path_to_dates)
    dataloader = DataLoader(data, batch_size = batch_size, shuffle = False, num_workers = num_workers)

    temperature_file = h5py.File(path_to_temperature_data, 'r')
    temperature_data = da.from_array(temperature_file['/data'])

    # 创建空Dask数组用于存储结果
    loaded_data = da.empty(shape=(len(temperature_data[0,:,0]), len(temperature_data[0,0,:]), number_of_years, 2, NUMBER_OF_DATA_DAYS), dtype = np.float32)
    
    for batch in tqdm(dataloader):
        for i in range(batch_size):
            if i >= len(batch[0]):
                break
            temperature_array = da.from_array(batch[0][i].numpy())
            dates = batch[1][i].numpy()
            y = int(batch[2][i])
            x = int(batch[3][i])
            year = int(batch[4][i] - 2001)

            # 将结果赋值到Dask数组对应位置
            loaded_data[y, x, year, 0, :] = temperature_array
            loaded_data[y, x, year, 1, :] = dates

    # 保存到HDF5
    progress_bar = ProgressBar()
    with progress_bar:
        da.to_hdf5(f'accumulated_degree_day_data_per_x_y_year_base_{base}.hdf5', {'/data': loaded_data})
问题根源与替代方案

核心问题

你当前的实现是在循环中逐个修改Dask数组loaded_data的切片,每一次赋值都会生成新的Dask任务图。随着循环次数增加,任务图的嵌套层级不断加深,最终触发Python递归深度超限错误。调大递归限制只是临时缓解——数据量越大,任务图嵌套越深,迟早会再次触发问题。

替代方案

1. 直接用Dask完成全流程预处理(推荐)

绕开PyTorch DataLoader,用Dask的分块处理能力直接完成数据加载、预处理和写入,避免生成嵌套任务图:

import dask.array as da
import h5py
import datetime

# 加载原数据并设置合理分块
with h5py.File(path_to_temperature_data, 'r') as temp_f:
    # 根据内存情况调整chunks大小,比如(100, 100, 100)
    temp_da = da.from_array(temp_f['/data'], chunks=(100, temp_f['/data'].shape[1], temp_f['/data'].shape[2]))
with h5py.File(path_to_dates, 'r') as date_f:
    dates = date_f['images_date'][:].astype(str)

# 预处理单块数据的函数
def process_chunk(temp_chunk, dates):
    # 按年份分组处理
    processed = []
    for year_offset in range(number_of_years):
        year = 2001 + year_offset
        # 筛选对应年份的日期索引
        year_indices = [i for i, d in enumerate(dates) if int(d[:4]) == year]
        if not year_indices:
            continue
        # 提取对应年份的温度数据并计算累积值
        year_temp = temp_chunk[year_indices]
        acc_temp = da.cumsum(year_temp, axis=0)
        # 处理日期格式
        year_dates = da.array([int(datetime.datetime.strptime(dates[i], '%Y-%m-%d').strftime("%Y%m%d")) for i in year_indices])
        # 堆叠温度和日期到band维度
        year_data = da.stack([acc_temp, year_dates], axis=-2)
        processed.append(year_data)
    # 按year维度拼接结果
    return da.concatenate(processed, axis=-3)

# 构建目标Dask数组
loaded_data = da.map_blocks(
    process_chunk,
    temp_da,
    chunks=(temp_da.chunks[1], temp_da.chunks[2], number_of_years, 2, NUMBER_OF_DATA_DAYS),
    dtype=np.float32,
    args=(dates,)
)

# 写入HDF5文件
da.to_hdf5(f'accumulated_degree_day_data_per_x_y_year_base_{base}.hdf5', {'/data': loaded_data})

2. 若必须用DataLoader,直接写入HDF5而非Dask数组

放弃用Dask数组临时存储结果,直接在循环中将batch数据写入HDF5文件的对应位置,彻底避免任务图嵌套:

import h5py
from tqdm import tqdm

# 提前创建HDF5文件并定义数据集
with h5py.File(f'accumulated_degree_day_data_per_x_y_year_base_{base}.hdf5', 'w') as f:
    y_size = len(temperature_data[0,:,0])
    x_size = len(temperature_data[0,0,:])
    shape = (y_size, x_size, number_of_years, 2, NUMBER_OF_DATA_DAYS)
    # 可添加compression参数压缩数据,比如compression='gzip'
    dset = f.create_dataset('/data', shape=shape, dtype=np.float32)
    
    # 遍历DataLoader写入数据
    for batch in tqdm(dataloader):
        for i in range(len(batch[0])):
            temp_arr = batch[0][i].numpy()
            date_arr = batch[1][i].numpy()
            y = int(batch[2][i])
            x = int(batch[3][i])
            year_idx = int(batch[4][i] - 2001)
            
            # 写入对应位置
            dset[y, x, year_idx, 0, :] = temp_arr
            dset[y, x, year_idx, 1, :] = date_arr

3. 优化Dataset实现,减少资源浪费

当前__getitem__中每次都打开HDF5文件,多worker环境下会重复打开,既浪费资源又可能引发问题:

class TemperatureDataset(Dataset):
    def __init__(self, temperature_path, date_path):
        self.temperature_path = temperature_path
        self.date_path = date_path
        # 在__init__中提前打开文件(只读模式多进程安全)
        self.temp_file = h5py.File(temperature_path, 'r')
        self.temp_data = da.from_array(self.temp_file['/data'])
        # 提前加载日期数据到内存
        with h5py.File(date_path, 'r') as f:
            self.dates = f['images_date'][:].astype(str)
        self.shape_of_array = (len(self.temp_data[0,:,0]), len(self.temp_data[0,0,:]), number_of_years)

    def __getitem__(self, idx):
        index = np.unravel_index(idx, self.shape_of_array)
        y, x, year_offset = index
        year = 2001 + year_offset

        acc_temperature = []
        start = 0
        # 遍历日期找到对应年份起始位置
        for i in range(len(self.temp_data[:,0,0])):
            if int(self.dates[i][:4]) < year:
                continue
            if int(self.dates[i][:4]) == year:
                start = i
                acc_temperature = np.cumsum(self.temp_data[i:(i + self.T), y, x].compute())
                break

        # 处理日期(逻辑不变)
        date_int_list = []
        for i in range(start, start + NUMBER_OF_DATA_DAYS):
            datetime_obj = datetime.strptime(self.dates[i], '%Y-%m-%d')
            date_int = int(datetime_obj.strftime("%Y%m%d"))
            date_int_list.append(date_int)

        return acc_temperature, torch.tensor(date_int_list), y, x, year

    def __del__(self):
        # 确保文件关闭
        self.temp_file.close()

内容的提问来源于stack exchange,提问作者Joel Hempel

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 07:47:46