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

如何存储形状为(3x512x512)的带标签张量用于CNN训练?

适合存储图像张量与标签的几种实用方案

1. Numpy .npz格式(轻量快捷)

适合小规模数据集,直接把张量和标签打包成压缩文件:

import numpy as np

# 示例数据:100个3x512x512张量 + 对应标签
tensors = [np.random.rand(3, 512, 512) for _ in range(100)]
labels = np.array([0, 1] * 50)

# 保存
np.savez('cnn_data.npz', tensors=tensors, labels=labels)

# 加载
data = np.load('cnn_data.npz')
loaded_tensors = data['tensors']
loaded_labels = data['labels']

2. HDF5格式(适配大规模数据)

用h5py库实现分块读写,不用一次性加载全量数据到内存:

import h5py

# 保存
with h5py.File('cnn_data.h5', 'w') as f:
    f.create_dataset('tensors', data=np.array(tensors), dtype='float32')
    f.create_dataset('labels', data=labels, dtype='int32')

# 加载
with h5py.File('cnn_data.h5', 'r') as f:
    loaded_tensors = f['tensors'][:]
    loaded_labels = f['labels'][:]

3. PyTorch .pt格式(直接适配训练流程)

如果用PyTorch做CNN训练,直接存储PyTorch张量对象,省去格式转换步骤:

import torch

# 转成PyTorch张量
torch_tensors = torch.tensor(np.array(tensors))
torch_labels = torch.tensor(labels)

# 保存
torch.save({'tensors': torch_tensors, 'labels': torch_labels}, 'cnn_data.pt')

# 加载
data = torch.load('cnn_data.pt')
loaded_tensors = data['tensors']
loaded_labels = data['labels']

4. 文件系统+标签清单(超大规模数据集首选)

把每个张量存成独立图像文件(比如PNG),用文本文件记录路径和对应标签,训练时按需加载:

import os
from PIL import Image

# 创建图像存储目录
os.makedirs('cnn_images', exist_ok=True)

# 保存图像与标签清单
with open('label_list.txt', 'w') as f:
    for idx, (tensor, label) in enumerate(zip(tensors, labels)):
        # 把CxHxW格式转成PIL需要的HxWxC
        img = Image.fromarray((tensor.transpose(1,2,0)*255).astype(np.uint8))
        img_path = f'cnn_images/img_{idx}.png'
        img.save(img_path)
        f.write(f'{img_path} {label}\n')

# 训练时用自定义Dataset加载(PyTorch示例)
from torch.utils.data import Dataset

class ImageDataset(Dataset):
    def __init__(self, label_file):
        self.samples = []
        with open(label_file, 'r') as f:
            for line in f:
                path, label = line.strip().split()
                self.samples.append((path, int(label)))
    
    def __len__(self):
        return len(self.samples)
    
    def __getitem__(self, idx):
        path, label = self.samples[idx]
        img = Image.open(path)
        tensor = torch.tensor(np.array(img).transpose(2,0,1)).float() / 255.0
        return tensor, label

内容的提问来源于stack exchange,提问作者Dominos-roadster

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 12:00:13