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.randtorch.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. 从最后一个维度开始向前比较
  2. 维度大小相等或其中一个为1时兼容
  3. 缺失维度视为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)  # 显式扩展维度

调试技巧

  1. 频繁使用.shape检查维度
  2. 对复杂操作分步执行
  3. 小规模测试后再应用到大数据

常见错误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)。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐