PyTorch张量操作保姆级习题集:从arange到广播,新手避坑指南
PyTorch张量操作保姆级习题集:从arange到广播,新手避坑指南
当你第一次接触PyTorch的张量操作时,是否曾被各种形状不匹配的错误提示搞得晕头转向?是否在广播机制面前感到困惑不解?这篇文章将带你通过实战习题,避开那些让新手抓狂的常见陷阱。
1. 张量创建与基础属性:那些容易忽略的细节
创建张量看似简单,但魔鬼藏在细节中。让我们从最基础的arange开始:
import torch
# 创建包含前12个整数的行向量
x = torch.arange(12)
常见错误1:忘记导入torch就直接使用。这会导致NameError: name 'torch' is not defined。
检查张量形状时,新手常混淆.shape和.size():
print(x.shape) # torch.Size([12])
print(x.size()) # 同样输出torch.Size([12])
实际上,.shape是属性,.size()是方法,但功能相同。PyTorch为了与NumPy兼容保留了这两种形式。
形状转换的坑:
# 将(12,)的行向量转换为(3,4)的矩阵
y = x.reshape(3, 4)
常见错误2:使用不兼容的形状。比如尝试x.reshape(4,4)会报错,因为16≠12。
2. 特殊张量创建:初始化陷阱全解析
创建全0或全1张量时,形状参数的正确传递至关重要:
zeros_tensor = torch.zeros(2, 3, 4) # 正确
zeros_tensor = torch.zeros((2, 3, 4)) # 也正确
常见错误3:忘记形状参数是可变参数还是元组。以下写法都会报错:
torch.zeros(2, 3, 4,) # 多余的逗号
torch.zeros([2, 3, 4]) # 传入列表而非元组
随机张量创建时,正态分布参数容易混淆:
normal_tensor = torch.randn(3, 4) # 均值0,标准差1
常见错误4:误用torch.rand和torch.randn:
rand生成[0,1)均匀分布randn生成标准正态分布
3. 张量运算:操作符重载的暗礁
PyTorch重载了Python运算符,但有些行为可能与预期不同:
x = torch.tensor([1.0, 2, 4, 8])
y = torch.tensor([2, 2, 2, 2])
# 逐元素运算
print(x + y) # tensor([ 3., 4., 6., 10.])
print(x * y) # tensor([ 2., 4., 8., 16.])
常见错误5:误以为*是矩阵乘法。实际上矩阵乘法应使用@或torch.matmul()。
幂运算的优先级问题:
print(x ** 2) # tensor([ 1., 4., 16., 64.])
常见错误6:忘记Python的运算符优先级,导致-x**2被解析为-(x**2)而非(-x)**2。
4. 广播机制:形状兼容的魔法与陷阱
广播机制是PyTorch的强大特性,也是新手困惑的重灾区:
a = torch.arange(3).reshape(3, 1)
b = torch.arange(2).reshape(1, 2)
print(a + b)
输出将是:
tensor([[0, 1],
[1, 2],
[2, 3]])
广播规则总结:
- 从最后一个维度开始向前比较
- 维度大小相等或其中一个为1时兼容
- 缺失维度视为1
常见错误7:不兼容的形状导致广播失败。例如尝试广播(3,4)和(2,3)会报错。
5. 索引与修改:原地操作的风险
张量索引与NumPy类似,但有些特殊行为需要注意:
X = torch.arange(12, dtype=torch.float32).reshape(3,4)
# 获取最后一行
last_row = X[-1]
# 获取第二到第三行(实际是第二行)
rows_1_2 = X[1:3]
常见错误8:误以为索引从1开始。PyTorch和Python一样使用0-based索引。
原地修改需要特别小心:
# 修改第1行第2列(注意是0-based)
X[1, 2] = 9
# 修改第0行和第1行所有元素
X[:2] = 12
常见错误9:忘记PyTorch的某些操作是原地(in-place)的,可能导致意外的副作用。原地操作通常有_后缀,如add_()。
6. 类型转换:数据精度丢失的隐患
PyTorch与NumPy互操作时,类型转换容易出问题:
# Tensor转NumPy
A = X.numpy()
# NumPy转Tensor
B = torch.from_numpy(A)
常见错误10:GPU张量直接转NumPy会报错。需要先.cpu():
gpu_tensor = torch.randn(3,4).cuda()
# 错误做法:gpu_tensor.numpy()
# 正确做法:gpu_tensor.cpu().numpy()
数据类型转换也容易出错:
a = torch.tensor([3.5])
print(a.int()) # tensor(3, dtype=torch.int32)
print(a.float()) # tensor(3.5000)
print(a.char()) # 报错,没有char()方法
常见错误11:误用不存在的类型转换方法。正确方法包括:
.to(torch.int).type(torch.FloatTensor)
7. 张量拼接与比较:维度对齐的艺术
拼接操作需要特别注意维度匹配:
X = torch.arange(12, dtype=torch.float32).reshape(3,4)
Y = torch.tensor([[2.0, 1, 4, 3], [1, 2, 3, 4], [4, 3, 2, 1]])
# 按行拼接(增加行数)
torch.cat((X, Y), dim=0)
# 按列拼接(增加列数)
torch.cat((X, Y), dim=1)
常见错误12:拼接维度不匹配。比如尝试拼接(3,4)和(2,4)按dim=0会失败。
元素比较操作:
print(X == Y) # 逐元素比较
常见错误13:误用==和torch.equal():
==是逐元素比较,返回布尔张量torch.equal()比较所有元素是否完全相同,返回单个布尔值
8. 实战演练:综合应用与调试技巧
让我们通过一个综合练习巩固所学:
# 创建两个张量
A = torch.randn(2, 3)
B = torch.randn(3, 4)
# 矩阵乘法
C = A @ B
# 尝试广播相加
D = A.unsqueeze(2) + B.unsqueeze(0) # 显式扩展维度
调试技巧:
- 频繁使用
.shape检查维度 - 对复杂操作分步执行
- 小规模测试后再应用到大数据
常见错误14:直接对不匹配形状的张量操作而不先检查形状。养成打印.shape的习惯能节省大量调试时间。
当遇到错误时,PyTorch的错误信息通常包含问题张量的形状信息。例如:
RuntimeError: The size of tensor a (3) must match the size of tensor b (4) at non-singleton dimension 1
这明确指出了在维度1上大小不匹配(3 vs 4)。
更多推荐


所有评论(0)