说实话,先不管 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这种东西,就是因为他没意识到:张量除了“值”本身,还带着dtype、device、requires_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.float32和torch.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,因为它更智能,不容易在形状变换时栽跟头。想真正提升性能或者对内存布局很敏感的人,再去深挖view和contiguous()。遇到“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 dimension | torch.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} 没有梯度")写在最后:我的一些小经验
张量操作这种东西,看十篇教程不如自己敲一遍。刚开始练的时候,我也会把view、reshape、permute混着用,现在反而更克制了:能不用高级花活就不用,代码越直白越好。我给自己的规矩是:所有跨维度的变换,一律写注释,标注变换前后的形状;所有拼接操作,先打印参与拼接张量的形状;所有涉及requires_grad的地方,尽量显式声明而不是靠隐式继承。
最后分享一个我一直在用的小技巧:遇到任何张量形状问题,先别急着搜网上的“报错解决方案”,花一分钟在小纸上把(batch, seq_len, hidden)这类中间形状推一遍。很多时候,不是 PyTorch 函数有问题,而是你自己把维度关系搞拧了。把形状推清楚,再看代码其实就顺了。这也是我见过的大多数 PyTorch 高手不约而同的习惯。