如何将单个.npz文件作为样本加载到TensorFlow Dataset中
加载大NPZ数据集时的TensorFlow Dataset问题
之前用pandas DataFrame加载数据集正常,但数据量增大后内存占用过高导致训练终止,换成TensorFlow Dataset类尝试懒加载:通过加载文件路径,再用map逐个读取NPZ文件内容,但加载单个文件失败。每个NPZ文件是形状为(1, x, x, z)的numpy数组,按标签分类存放在对应文件夹中。
原代码
数据集加载方法
IMAGE_SUPPORTED_EXTENSIONS = ('.jpg', '.jpeg', '.png') def load_dataset(self): data = [] for label in self.labels: folder = self.main_folder / label file_paths = [str(file_path) for file_path in folder.glob('*') if file_path.suffix in TENSOR_SUPPORTED_EXTENSIONS] latenst_spaces = [DatasetLoader.load_tensor(file_path) for file_path in folder.glob('*') if file_path.suffix in TENSOR_SUPPORTED_EXTENSIONS] dataset = tf.data.Dataset.from_tensor_slices(file_paths) # 绑定路径与标签 dataset = dataset.map(lambda x: (x, label)) dataset = dataset.map(map_function) data.append(dataset) # 合并不同标签的数据集 dataset = data[0] for i in range(1, len(data)): dataset = dataset.concatenate(data[i]) return dataset
Map函数与加载工具
def map_function(element): file_path, label = element npz_data = DatasetLoader.load_tensor(file_path) return (npz_data, label) class DatasetLoader: @staticmethod def load_tensor(file_path): file_path = tf.get_static_value(tf_tensor) file_path = Path(file_path) if file_path.suffix not in ('.npy', '.npz'): raise ValueError(f"Extension {file_path.suffix} not suppported.") try: with np.load(file_path) as tensor: if file_path.suffix == ".npz": for _, item in tensor.items(): tensor = item return np.array(tensor).squeeze() except Exception as e: print(f"Error loading {file_path.stem} file: {str(e)}.", "\nFile path: ", file_path) raise RuntimeError(f"Error loading {file_path.stem} file: {str(e)}.") from e
问题分析与修正代码
核心问题
latenst_spaces提前加载所有NPZ文件,违背了Dataset懒加载的初衷,依然会占用大量内存load_tensor中变量名错误(tf_tensor未定义),且tf.get_static_value无法处理动态张量(Dataset输出的路径是动态张量)- 直接在TensorFlow的
map中调用Python文件读取逻辑,未用tf.py_function包装,导致TensorFlow无法识别
修正后的完整代码
# 定义支持的张量扩展名 TENSOR_SUPPORTED_EXTENSIONS = ('.npy', '.npz') def load_dataset(self): data = [] for label in self.labels: folder = self.main_folder / label # 仅收集文件路径,不提前加载数据 file_paths = [str(file_path) for file_path in folder.glob('*') if file_path.suffix in TENSOR_SUPPORTED_EXTENSIONS] dataset = tf.data.Dataset.from_tensor_slices(file_paths) # 将标签转为Tensor,保证类型统一 dataset = dataset.map(lambda x: (x, tf.constant(label))) # 用tf.py_function包装Python侧的文件读取逻辑 dataset = dataset.map(lambda x, y: tf.py_function( map_function, inp=[x, y], Tout=[tf.float32, tf.string] )) # 恢复张量形状(py_function会丢失形状信息) dataset = dataset.map(lambda x, y: (tf.ensure_shape(x, (None, None, None)), y)) data.append(dataset) # 合并所有数据集 dataset = data[0] for ds in data[1:]: dataset = dataset.concatenate(ds) return dataset def map_function(file_path, label): # 将Tensor路径转为Python字符串 file_path_str = file_path.numpy().decode('utf-8') npz_data = DatasetLoader.load_tensor(file_path_str) return npz_data, label class DatasetLoader: @staticmethod def load_tensor(file_path): file_path = Path(file_path) if file_path.suffix not in TENSOR_SUPPORTED_EXTENSIONS: raise ValueError(f"不支持的扩展名:{file_path.suffix}") try: with np.load(file_path) as tensor: if file_path.suffix == ".npz": # 取NPZ中的第一个数组(假设每个文件仅存一个数组) for _, item in tensor.items(): tensor = item break # 压缩维度为1的轴,(1, x, x, z) → (x, x, z) return np.array(tensor).squeeze().astype(np.float32) except Exception as e: print(f"加载文件失败 {file_path.stem}: {str(e)}\n文件路径: {file_path}") raise RuntimeError(f"加载文件 {file_path.stem} 失败: {str(e)}") from e
关键修改说明
- 移除提前加载所有NPZ数据的代码,严格遵循懒加载逻辑,降低内存占用
- 使用
tf.py_function包装Python文件读取操作,适配TensorFlow的计算图模式 - 修正
load_tensor中的变量名错误,调整NPZ读取逻辑避免不必要的循环 - 为标签创建Tensor并恢复张量形状,保证数据集输出的类型与形状一致性
内容的提问来源于stack exchange,提问作者Lorenzo Sibi
相关产品推荐
相关产品推荐

