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

如何对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的拼接逻辑更方便,无需每轮迭代额外处理:

  1. 先定义自定义拼接函数:
import torch
def concat_collate(batch):
    # batch是长度等于batch_size的列表,每个元素为__getitem__返回的(250,150)张量
    return torch.cat(batch, dim=0)
  1. 初始化DataLoader时传入自定义函数:
from torch.utils.data import DataLoader
dataloader = DataLoader(your_dataset, batch_size=10, collate_fn=concat_collate)

该方案同时兼容样本第0维长度不固定的场景,避免默认collate_fn堆叠时报错。

内容的提问来源于stack exchange,提问作者DiveIntoML

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 00:27:03