使用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

