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

如何将单个.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

问题分析与修正代码

核心问题

  1. latenst_spaces提前加载所有NPZ文件,违背了Dataset懒加载的初衷,依然会占用大量内存
  2. load_tensor中变量名错误(tf_tensor未定义),且tf.get_static_value无法处理动态张量(Dataset输出的路径是动态张量)
  3. 直接在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 11:36:14