PyTorch张量操作实战:从view、reshape到维度变换避坑指南
2026/9/15 1:12:18 网站建设 项目流程

说实话,先不管 PyTorch 官网那一堆花哨的教程,也不管你从哪个视频里看到“从零入门 PyTorch”的集数,真正决定你能不能把模型跑起来的,从来都是张量操作。为什么这么说?因为 PyTorch 整个基础框架,说白了就是搭在张量(Tensor)上的一套自动求导系统。你喂给模型的数据是张量,模型里的权重是张量,梯度也是张量。你要是连张量的维度变换都玩不转,哪怕把 YOLO 的源码翻烂了也改不出自己想要的结构。这篇文章就是把“张量操作”这层窗户纸捅破,适合刚装完 PyTorch、还分不清 view 和 reshape 有什么区别的初学者,也适合那些写模型写了不少但老在维度上报错的半新手。看完你能直接照着敲,遇到问题也知道去哪里排查。

1. 张量到底是个什么东西

1.1 它不只是个“多维数组”

很多教程喜欢把张量解释成“多维数组”,这句话没错,但特别容易误导人。NumPy 里也有多维数组,你用 numpy 也能做矩阵乘法、也能做广播、也能做 reshape。那 PyTorch 的张量和 NumPy 的 ndarray 到底差在哪儿?差在两件事:GPU 加速和自动求导。

GPU 加速很好理解,你在 CPU 上算一个 512x512 的矩阵乘法,可能还行;但要是你在训练神经网络,前向传播里随便一层就是几百万次浮点运算,CPU 直接卡死。张量可以放到 GPU 上计算,也就是把数据从内存搬到显存里,运算速度能快几十倍。自动求导更是 PyTorch 的看家本事,你定义一个张量时需要设置requires_grad=True,之后所有对这个张量的操作都会被记录在一个“计算图”里,调用backward()时梯度会自动算出来。你如果用 NumPy,就得自己手动实现反向传播,写两行就想摔键盘。

我见过有些同学用torch.tensor([1, 2, 3])创建了一个张量,然后整天疑惑为什么别人的代码里有.cuda()或者device这种东西,就是因为他没意识到:张量除了“值”本身,还带着dtypedevicerequires_grad这些隐性属性。这三个属性一旦不匹配,就会产生各种莫名其妙的问题。所以学张量操作,第一件事不是背函数,而是建立这个意识:你手里拿的不是一个孤零零的数组,而是一个携带着硬件位置、数值类型、梯度状态的数据容器。

1.2 动手之前,先想清楚三件事

我在带新人的时候,要求他们在创建任何张量之前,先回答三个小问题:

  • 这个数据要放在 CPU 还是 GPU?如果放 GPU,那 CPU 上的数据要先.to(device)搬过去。
  • 数值类型是 float32 还是 int64?模型权重一般都要 float32,标签一般用 int64,索引一般用 int64。
  • 这个张量需要求梯度吗?只有需要训练的参数才需要requires_grad=True,别一股脑全开,浪费内存不说,回头有时候梯度反传到你想不到的地方。

别小看这三个问题。很多人训练 RNN 或者 Transformer 的时候报错“Expected floating point type for target with class probabilities”或者"RuntimeError: expected scalar type Long but found Float",八成就是第二个问题没管好。你不需要把每个类型都背下来,但至少要知道torch.float32torch.long大概对应什么场景,出错了才知道往哪个方向排查。

我记得我第一次用 PyTorch 训练一个简单的全连接网络时,数据是从 DataFrame 里读出来的,里面既有整数又有小数,直接转成torch.from_numpy(df.values),结果报错说数据类型不支持。当时就是没意识到 NumPy 数组默认可能是 int64,但网络参数要求 float32。你光看数值是对的,但类型对不上,框架一样不惯着你。这个坑说起来低级,踩过的人却不少。

2. 创建张量的几类实用方法

2.1 从列表或者 NumPy 转过来

最直接的创建方式就是torch.tensor([1, 2, 3]),但要注意torch.tensor是会复制数据的。如果是从 NumPy 转过来的,我建议用torch.from_numpy(),因为它是共享内存的,转换开销小。这句话什么意思?就是你改 NumPy 数组的值,对应的张量也会变;反之亦然。好处是快,坏处是容易在你不注意的时候发生数据被意外修改的“灵异事件”。

import torch import numpy as np arr = np.array([1.0, 2.0, 3.0]) t = torch.from_numpy(arr) arr[0] = 99.0 print(t) # tensor([99., 2., 3.])

看到了吧?这个共享机制在实际工程里是把双刃剑。如果数据在预处理阶段还要频繁改动,最好用torch.tensor(arr)复制一份,避免后续模型训练时数据被无意中覆盖。如果只是想快速把 NumPy 数据送进网络,而且确定不会再改原数组,用from_numpy是更优解。

从 Python 列表创建目前我还是推荐torch.tensor()。如果你用torch.Tensor([1, 2, 3])这种构造器,它其实等同于torch.FloatTensor,行为上会有一些隐性默认值。torch.tensor()会根据数据自动推断 dtype,而torch.Tensor()永远创建 float32,这个区别经常让人困惑。我的原则是:什么时候都别省那几个字符,用torch.tensor(),让类型显式一点。

2.2 按形状创建:zeros、ones、randn 怎么选

造测试数据的时候,经常需要一个“形状正确”但内容无所谓的张量。这时候torch.zeros()torch.ones()torch.randn()torch.empty()轮流上场。

  • torch.zeros(3, 4):3 行 4 列,全部填 0。适合生成掩码、初始化偏置。
  • torch.ones(3, 4):全部填 1。适合做全连接层的 bias 初始化。
  • torch.randn(3, 4):从标准正态分布里采样,适合测试模型前向传播。
  • torch.empty(3, 4):分配内存,但里面的值是垃圾桶里捡来的,没初始化。只有你确定马上会被覆盖时才用。

实际写模型时,你还要区分torch.rand_like(t)torch.randn_like(t)。前者生成 0 到 1 之间的均匀分布随机数,后者生成标准正态随机数,它们都会保持和t相同的形状。这两个函数在做 Dropout 或噪声注入的测试时特别方便,不用手动传 shape。

下面是一段肉眼可见的实例:

import torch a = torch.zeros(2, 3) b = torch.ones(2, 3) c = torch.randn(2, 3) d = torch.full((2, 3), 7.5) # 全部填 7.5 print("a device:", a.device) print("d shape:", d.shape)

torch.full((2, 3), 7.5)可能很多人不知道,它用来生成“全填充为指定值”的张量,比torch.ones() * 7.5直观得多,而且不会引入额外的乘法运算节点。

2.3 dtype、device、requires_grad 的坑

这三兄弟是真的能联手坑人。我举个最典型的例子:

x = torch.randint(0, 10, (3,)) y = x.float() # 转 float32 z = torch.randn_like(y) z.requires_grad_(True) # 原地设置 requires_grad

第一行是 int64,第二行转成 float32,第三行用randn_like保证了形状和 dtype 都和 y 一致,第四行原地开启梯度记录。这个流程在代码里看起来平平无奇,但如果你没有第二行,直接对 int 型张量计算梯度,PyTorch 会直接告诉你“RuntimeError: Only Tensors of floating point and complex dtype can require gradients”。这也是我在实际问答里见到概率极高的一类报错。

还有一个隐性坑:torch.arange(0, 10)生成的默认 dtype 在某些版本里是 int64,但torch.range(0, 10)是 float32。很多人用range以为它和 Python 的range一样,结果类型对不上。现在 PyTorch 官方其实都建议用torch.arange替代torch.range,因为torch.range的行为太反直觉:它会包含结尾值,而且 dtype 默认不是整数。记住一个小原则:索引类张量用 int64,计算类张量用 float32,别混着用。

3. 形状操作“变形记”:view、reshape、permute、transpose

3.1 view 和 reshape 的区别到底在哪

这是我在各技术论坛上被人问得最多的问题之一。一句话解释:view只对“内存连续”的张量有效,它直接复用底层内存,不复制数据;reshape更聪明,如果张量内存连续,它就和view一样不会复制,如果不连续,它会先拷贝一份让内存连续,再改变视图。

那“内存连续”是什么意思?你可以把张量的内存想象成一排连续的小格子。普通创建一个 3x4 张量,数据在内存里就是按行依次排下来的,连续。但是你一旦执行了t.t()(转置),逻辑上得到的是一个 4x3 张量,可底层内存顺序还是原来的 12 个格子,没法按新形状直接“顺序解读”,这时候内存就不连续了。view看到这种情况直接报错,reshape则会先contiguous()拷贝出连续的内存,再返回新视图。

x = torch.arange(12).reshape(3, 4) y = x.t() # transpose try: z = y.view(4, 3) except RuntimeError as e: print("view 报错:", e) z = y.reshape(4, 3) print(z)

所以我个人建议:新手优先用reshape,因为它更智能,不容易在形状变换时栽跟头。想真正提升性能或者对内存布局很敏感的人,再去深挖viewcontiguous()。遇到“viewsize is not compatible”这类报错,别慌,改成reshape,或者先调用.contiguous()view一般就通了。

3.2 permute 和 transpose 千万别混淆

x.transpose(dim0, dim1)只能交换两个维度,例如二维矩阵转置就是x.transpose(0, 1)x.permute(dims)能一次性把所有维度重新排列,比如x.permute(2, 0, 1)表示把原来第三维挪到最前面。这两个操作都是视图操作,不复制数据,所以改数据会互相影响,这点要记住。

尤其在图像处理里,PyTorch 的默认张量布局是[C, H, W],也就是通道在前。你用 OpenCV 读出来的图片是[H, W, C],也就是通道在最后。想塞进 PyTorch 模型,就需要把通道维度换到最前面。一开始我见过有人写img.permute(2, 0, 1),后来也有人用img.transpose(0, 2)transpose(1, 2),结果自己也绕晕了。我建议只记permute(2, 0, 1)这一个写法,图像从[H, W, C]换成[C, H, W],它是最明确的。

permute 还有一个非常容易踩的坑:变换之后张量通常不连续。你如果接着调用view,十有八九又报“view size is not compatible”。正确的姿势是先permute,再.contiguous(),再view

img = torch.randn(224, 224, 3) # 模拟 HWC img_chw = img.permute(2, 0, 1) # 变成 CHW print(img_chw.shape) # torch.Size([3, 224, 224])

3.3 unsqueeze 和 squeeze 的实战场景

这两个函数名字长得像兄弟,但功能正好相反。unsqueeze是在指定位置插入一个新的维度,长度为 1;squeeze是把长度为 1 的维度去掉。很多人刚开始根本不知道为什么要“多加一个长度为 1 的维度”。

我给你说一个最常见的场景。假设你有一条数据,形状是[4],表示一个样本的 4 个特征。但 PyTorch 的线性层nn.Linear要求输入至少是二维的:第一维是 batch size,第二维是特征数,也就是[1, 4]。所以你得用x.unsqueeze(0)把它变成[1, 4]。如果最后你想把这一个样本的特征向量恢复成[4],就用x.squeeze(0)

还有卷积层,输入要求[B, C, H, W]。如果你做单张图片推理,图片读进来是[C, H, W],也要先unsqueeze(0)变成[1, C, H, W]。这就叫“手动补一个 batch 维度”。

这里要注意.squeeze()默认是去掉所有长度为 1 的维度,但你可以给它传参数,只去除指定位置。比如x.squeeze(dim=2),如果你原以为第二维长度是 1,但实际不是,它不会报错,而是什么都不做。这个“静默不报错”的行为,有时候反而会掩盖问题。我自己写代码时,能显式传dim就尽量显式传,避免无意中把多个长度 1 的维度全删了。

3.4 用形状推导法避免“维度灾难”

张量操作多了以后,我发现自己陷入一个怪圈:不停地加 reshape、permute、unsqueeze,最后把维度搞得连自己都不认识了。后来我养成了一个习惯:每一步关键变换,都手动写下当前的shape,再写下目标shape,用笔推一遍。

举例来说,我有一个 attention 矩阵,原始形状是[batch, heads, seq_len, seq_len],我想去掉中间那个 heads 维度,把它合并到 batch 里去。我会写:

# 当前: [B, H, T, T] attn = attn.permute(0, 2, 3, 1) # 变成 [B, T, T, H] attn = attn.reshape(B, T, T * H) # 合并掉 H

如果我不写注释,一周后回来看这段代码一定得重新推半天。所以这里强烈建议:关键形状变换前后,都写注释。调试时也可以临时加print(x.shape)。这个方法笨,但特别有效,尤其是刚开始学 Transformer 的人,Shape 推导熟练了,各种框架代码读起来都会顺畅很多。

4. 张量的索引、切片与组合拼接

4.1 像 NumPy 一样索引,但小心维度套路

PyTorch 的索引语法和 NumPy 高度一致:x[0]取第一个,x[:, 1]取所有行的第二列,x[1:, :2]切片,x[[0, 2]]按列表取行。这些我都默认你会,真正想提醒的是下面两点。

第一,布尔掩码索引返回的是一维张量。比如x[x > 0],你得到长度等于“满足条件的元素个数”的向量,原来的形状全丢了。这在做损失计算时有用,但你要是想保留矩阵结构,需要另想办法。

第二,索引结果一般情况下会共享数据,也就是返回的是视图。什么意思?你修改y = x[0]中的元素,x也会变。如果你不想影响原张量,就要用y = x[0].clone()

第三,花式索引(fancy indexing)有时候会复制数据,比如x[[0, 1, 2]]这种索引操作和切片不一样,它不保持内存连续性。这个对普通算法影响不大,但如果你做自定义算子,就会碰上性能问题。

x = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) mask = x > 4 print(x[mask]) # 一维:tensor([5, 6, 7, 8, 9]) print(x[x[:, 1] > 2]) # 按某一列过滤行:第2列>2的所有行

4.2 torch.cat 还是 torch.stack,这是个问题

torch.cat沿现有的某个维度拼接,要求除拼接维度外,其余维度完全一致。torch.stack是增加一个新的维度,把所有张量“堆”起来。你可以这么理解:cat是把几块积木沿长度方向接起来,stack是把它们叠成几层。

举例:

a = torch.tensor([[1, 2], [3, 4]]) b = torch.tensor([[5, 6], [7, 8]]) c = torch.cat([a, b], dim=0) # 形状 [4, 2] d = torch.cat([a, b], dim=1) # 形状 [2, 4] s = torch.stack([a, b], dim=0) # 形状 [2, 2, 2]

从形状结果就能看出区别。cat不增加总维度数,stack总是增加一个维度。数据加载的时候,如果你想把多个特征矩阵拼在一起,用cat;如果你想把一批单独的样本“摞”成批次张量,用stack

实际写 dataloader 的时候,我最常用的组合是stack配合unsqueeze。比如 batch 里的每条样本是[seq_len, hidden],我想得到[batch, seq_len, hidden],直接torch.stack(batch_list, dim=0)就行,代码短而且不容易出错。

4.3 split 和 chunk 不是一回事

torch.split可以按“每块大小”拆,也可以按“每块数量”拆,接口比较灵活。torch.chunk则是固定拆成 n 块,如果张量长度不能被 n 整除,最后一块会小一点。这个“整除”问题很容易埋雷。

比如你有一个长度为 10 的张量,想用torch.chunk(x, 3, dim=0)拆成 3 块,PyTorch 给你返回 3 个张量,长度分别是 4、4、2,不是均匀的。如果你用torch.split(x, 3, dim=0),则返回长度为 3、3、3、1 的多个张量。初学者常常忽略这个差别,导致循环里处理每一块时,以为大小一样,结果最后一块尺寸不一样就崩了。

x = torch.arange(10) for piece in torch.split(x, 3): print(piece.shape) # torch.Size([3]) # torch.Size([3]) # torch.Size([3]) # torch.Size([1]) for piece in torch.chunk(x, 3): print(piece.shape) # torch.Size([4]) # torch.Size([4]) # torch.Size([2])

如果你希望风格稳定,就用split,因为它可以指定固定尺寸。如果你只是想把一个 batch 平均分到多张卡上,但还不能保证一定能整除,那么chunk的“最后一块比较小”行为也许更符合预期,但一定记得处理好边界条件。

5. 常见形状错误与排查经验

5.1 我最常碰到的三类张量报错

这些报错几乎每个人都会遇到,我把它们整理成一个速查表,希望能帮你省点时间。

报错信息常见原因解决方向
view size is not compatible操作时张量内存不连续,或新形状元素总量与原形状不一致reshape代替,或先.contiguous()view
Expected scalar type Long but found Float/ 反之dtype 不匹配,如标签用 float 而损失函数期望 long.long().float()转换
mat1 and mat2 shapes cannot be multiplied矩阵乘法中 inner dimensions 不一致打印两个张量的shape,检查维度对应关系
Expected all tensors to be on the same device一部分张量在 CPU,一部分在 GPU统一调用.to(device)
Sizes of tensors must match except in dimensiontorch.cat或广播时,除指定维度外其他维度不一致打印参与拼接的每个张量的shape,对齐形状

这些报错本身不要怕,怕的是你不看报错信息就瞎改。我的建议是:遇到报错,第一件事把print(x.shape, y.shape)加在报错那一行之前,看清楚两个张量的形状再动手。大多数情况下,错误自己就能定位了。

5.2 用几行代码定位维度问题

我在调试程序时,会习惯性地写一个极小的“形状检查脚本”,把关键的张量挨个打印出来。你用不着用什么高级调试器,print 大法最直接。

def describe(name, tensor): print(f"{name}: shape={tuple(tensor.shape)}, dtype={tensor.dtype}, device={tensor.device}") a = torch.randn(2, 3) b = torch.randn(3, 4) describe("a", a) describe("b", b) try: c = torch.matmul(a, b) describe("c", c) except RuntimeError as e: print("matmul error:", e)

这样一眼就能看出a(2, 3)b(3, 4),乘法的最后维度是 3 对 3,没问题。如果写成a(2, 3)b(2, 4),你马上就能发现内维度 3 和 2 不匹配。

5.3 错误排查的三个独门心得

第一,不要盲目用.cpu().numpy()。很多初学者为了打印张量,先.cpu().numpy(),结果遇到“can‘t convert cuda:0 device type tensor to numpy”这种报错。其实你直接print(tensor)就行,张量支持直接输出,不非得转 NumPy。

第二,当你怀疑模型训练出了问题,先检查损失函数的输入。分类问题最常见的就是预测值没经过softmax之前是全连接层输出,形状是[batch, num_classes],标签形状是[batch]。这两个形状对应清楚,损失函数的坑基本就没了。

第三,requires_grad也是排查时的重要指标。如果某个张量本来需要梯度,但你在中间做了一次detach()或者从 NumPy 转过来时没设置requires_grad=True,反向传播到这里梯度就断了。症状往往是“模型不更新”或“某些层梯度为 None”。排查方法很简单:在loss.backward()之前打印一下模型参数的.grad,看看是不是 None。

for name, param in model.named_parameters(): if param.grad is None: print(f"{name} 没有梯度")

写在最后:我的一些小经验

张量操作这种东西,看十篇教程不如自己敲一遍。刚开始练的时候,我也会把viewreshapepermute混着用,现在反而更克制了:能不用高级花活就不用,代码越直白越好。我给自己的规矩是:所有跨维度的变换,一律写注释,标注变换前后的形状;所有拼接操作,先打印参与拼接张量的形状;所有涉及requires_grad的地方,尽量显式声明而不是靠隐式继承。

最后分享一个我一直在用的小技巧:遇到任何张量形状问题,先别急着搜网上的“报错解决方案”,花一分钟在小纸上把(batch, seq_len, hidden)这类中间形状推一遍。很多时候,不是 PyTorch 函数有问题,而是你自己把维度关系搞拧了。把形状推清楚,再看代码其实就顺了。这也是我见过的大多数 PyTorch 高手不约而同的习惯。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询