使用Xarray+TensorFlow时,洗牌后数据管道性能骤降的问题排查
Xarray与TensorFlow结合时洗牌后性能骤降的问题分析
问题描述
我正尝试将Xarray与TensorFlow结合使用。每个epoch结束后,会沿time维度对数据进行洗牌。NetCDF文件的维度为(time,lat,lon,level),time维度约有400k个值,我将该维度的chunk大小设为1024。遇到的问题是:对DataArray进行洗牌后,数据管道的性能大幅下降。
我尝试在洗牌后重新设置chunk大小,但并无效果。反之,若将洗牌后的数据写入磁盘,再以chunk大小1024重新读取,性能就能恢复到未洗牌时的水平。请问重新分块操作中我忽略了什么,导致其性能远不及从磁盘重新读取?
原代码
def on_epoch_end(self): "Shuffle dataset at the end of epoch" if self.shuffle == True: # Get the Dask array containing the data values dask_data = self.data.data # Create a shuffled index array along the 'time' dimension shuffled_indices = da.random.permutation( dask_data.shape[0] ) # Use Dask delayed computation to perform the shuffling shuffled_data = dask_data[shuffled_indices, :, :, :] # shuffled_data = da.rechunk(shuffled_data, chunks={0: 1024}) # Create a new DataArray with the shuffled data self.data = xr.DataArray( shuffled_data, coords=self.data.coords, dims=self.data.dims ) self.data = self.data.chunk({"time": 1024}) ## save data to file and read it agan #self.data.to_netcdf("save.nc") #ds = xr.open_dataset( # "save.nc", # chunks={"time": 1024}, #) #first_variable_name = list(ds.variables)[4] #self.data = ds[first_variable_name]
问题原因
核心问题在于Dask数组的逻辑分块和物理存储的不匹配:
- 用
da.random.permutation生成的打乱索引彻底破坏了原数据的连续性。此时调用rechunk或data.chunk()只是做了逻辑上的分块定义,但每个新的time chunk背后对应的是原NetCDF文件中分散的、不连续的多个原始chunk。TensorFlow读取时,每个新chunk都需要从磁盘读取大量零散片段,随机IO的开销直接导致性能暴跌。 - 而将洗牌后的数据写入磁盘再重新读取时,洗牌后的数据被物理上重新组织成了连续的1024 size的time chunk,每个新chunk对应磁盘上连续的存储空间,读取时是高效的顺序IO,性能自然恢复。
优化方案
不需要依赖手动存读NetCDF的方式,可以用更适合Dask的Zarr格式来临时存储洗牌后的数据,减少IO开销:
import zarr def on_epoch_end(self): "Shuffle dataset at the end of epoch" if self.shuffle == True: dask_data = self.data.data shuffled_indices = da.random.permutation(dask_data.shape[0]) shuffled_data = dask_data[shuffled_indices, :, :, :] # 使用Zarr临时存储洗牌后的数据,并行写入效率更高 temp_store = zarr.TempStore() # 显式触发计算,将洗牌后的数据写入临时存储 da.store(shuffled_data, temp_store, compute=True) # 从临时存储重新加载为带指定chunk的DataArray self.data = xr.DataArray( da.from_zarr(temp_store, chunks={"time": 1024}), coords=self.data.coords, dims=self.data.dims )
如果坚持使用NetCDF,也可以优化原有的存读逻辑,确保调用self.data.to_netcdf时指定compute=True,待数据完全写入后再读取。
内容的提问来源于stack exchange,提问作者danny
相关产品推荐
相关产品推荐

