如何对PyTorch张量沿指定维度拼接以调整DataLoader输出形状
PyTorch DataLoader输出维度转换解决方案
你遇到的现象是PyTorch DataLoader的默认行为导致的:默认collate_fn会将批次内的所有样本张量在新增的第0维堆叠,因此单个形状为(250, 150)的样本,在batch_size=10时会得到形状为(10, 250, 150)的输出。你可以通过以下两种方式得到形状为(2500, 150)的张量:
方案1:拿到批次后直接转换
如果仅部分场景需要拼接后的维度,直接对DataLoader的输出做维度调整即可,有两种等效写法:
- 用
flatten合并前两维:
for batch in dataloader: # batch.shape = (10, 250, 150) batch = batch.flatten(start_dim=0, end_dim=1) # 转换后batch.shape = (2500, 150)
- 用
reshape直接重设维度:
for batch in dataloader: batch = batch.reshape(-1, 150)
注意:该方案仅适用于所有样本第0维长度固定为250的场景。
方案2:自定义DataLoader的collate_fn
如果所有批次都需要拼接后的维度,直接修改DataLoader的拼接逻辑更方便,无需每轮迭代额外处理:
- 先定义自定义拼接函数:
import torch def concat_collate(batch): # batch是长度等于batch_size的列表,每个元素为__getitem__返回的(250,150)张量 return torch.cat(batch, dim=0)
- 初始化DataLoader时传入自定义函数:
from torch.utils.data import DataLoader dataloader = DataLoader(your_dataset, batch_size=10, collate_fn=concat_collate)
该方案同时兼容样本第0维长度不固定的场景,避免默认collate_fn堆叠时报错。
内容的提问来源于stack exchange,提问作者DiveIntoML
相关产品推荐
相关产品推荐

