如何将PyTorch DataLoader批次合并为单个NumPy数组?
PyTorch DataLoader数据合并方案
直接用列表收集所有批次的输入和标签,最后一次性拼接,完全不用判断首次迭代:
import numpy as np # 遍历DataLoader时收集所有批次数据 inputs_list = [] labels_list = [] for batch in train_loader: inputs, labels = batch # 若数据是PyTorch Tensor,先转成numpy数组 inputs_list.append(inputs.numpy()) labels_list.append(labels.numpy()) # 沿样本维度(第0轴)拼接所有数组 all_inputs = np.concatenate(inputs_list, axis=0) all_labels = np.concatenate(labels_list, axis=0)
这种方式不管批次大小是否能整除总样本数,都能生成形状正确的完整数组。
多维数组通用合并方法
对于形状为(300,28,28,3)和(500,28,28,3)的两个数组,直接用np.concatenate沿第0轴合并即可:
arr1 = np.random.rand(300,28,28,3) arr2 = np.random.rand(500,28,28,3) merged_arr = np.concatenate([arr1, arr2], axis=0) # 合并后数组形状为(800,28,28,3)
如果是仅在第0轴合并,也可以用np.vstack((arr1, arr2)),但np.concatenate更通用,支持指定任意合并轴。
内容的提问来源于stack exchange,提问作者CarinaE
相关产品推荐
相关产品推荐

