如果把深度学习比作 “搭积木建房子”,那 PyTorch 就是最顺手的积木套装 —— 既有灵活的零件(张量),又有好用的工具(函数),甚至还帮你准备了 “自动组装说明书”(自动微分)。今天咱们用 “玩积木” 的思路,把 PyTorch 的核心用法拆成简单有趣的知识点!

一、先搞懂:PyTorch 是啥?为啥大家都爱用?

在学具体操作前,得先明白 PyTorch 的 “定位”—— 它不是一个 “黑盒子工具”,而是一套 “灵活的深度学习工作台”。

1. PyTorch 的核心身份

  • 基于 Python 的深度学习框架:所有操作都用 Python 写,如果你会用 Python 处理数据(比如用 NumPy),上手 PyTorch 会特别快,不用额外学新语言。
  • 以 “张量” 为核心:就像高达模型的 “通用连接件”,所有数据都要封装成张量才能处理,小到单个数字(0 维张量),大到视频数据(4 维张量: batch× 通道 × 高度 × 宽度),都能装进去。
  • 学术与工业双适配:学术界爱用它做研究(改模型快、调试方便),工业界也能用它做生产(支持部署到手机、服务器),比如 Facebook 的推荐系统、特斯拉的自动驾驶,都有用到 PyTorch。

PyTorch的安装:

pip install torch -i https://pypi.tuna.tsinghua.edu.cn/simple

这边建议就是安装这种包的时候,都使用镜像源来进行安装,用过的都说好,嘎嘎快!!

2. PyTorch 的 “独家优势”(对比其他框架)

很多人会拿它和 TensorFlow 比,咱们用 “组装高达” 来比喻:

特点       PyTorch(动态图)TensorFlow 1.x(静态图)
计算图构建边写代码边建图,像“拼到哪算哪”先画好完整图纸再拼,该一步要画全图
调试体验能像 Python 代码一样断点调试,报错位置超明确报错错常找不到具体位置,像 “图纸错了但不知道哪错”
上手难度接近 Python 思维,新手 1 小时能跑通简单模型需要理解 “会话(Session)” 等概念,门槛高
GPU 加速一行代码data.to('cuda')搞定早期需要手动管理设备,步骤繁琐

简单说:PyTorch 把 “复杂的深度学习逻辑” 拆成了 “像搭积木一样简单的步骤”,特别适合新手入门。

二、核心零件:张量(Tensor)的 “全方位攻略”

张量是PyTorch的“基石”,就像高达的“螺丝、装甲、关节”——所有操作都围绕它展开。咱们从“认识张量”到“玩转张量”,一步步来。

先搞懂:张量到底是什么?

你可以吧张量理解成“NumPy数组的升级版“:

  • 0 维张量:单个数字(比如tensor(10)),像一颗单独的螺丝;
  • 1 维张量:一串数字(比如tensor([1,2,3])),像一根长螺丝杆;
  • 3 维及以上:比如 3 维张量(2,3,4),可以理解成 “2 块 3 行 4 列的装甲板叠在一起”,常用于处理文本( batch× 句子长度 × 单词向量);4 维张量常用于图像( batch× 通道 × 高度 × 宽度)

张量是 PyTorch 中的核心数据抽象,PyTorch 支持各种张量子类型。通常地,一维张量称为向量/矢量(vector),二维张量称为矩阵(matrix)。

这里用一张图详细解释一下,就相当于一个不断添加位面的一个魔方,深度学习其实就是依靠一个一个不断添加的层数,去更好的识别一个物品提高准确率。

和 NumPy 数组最大的区别:张量支持GPU 加速—— 把张量放到 GPU 上,计算速度能比 CPU 快几十倍,这对处理大数据至关重要。

1.张量的 “花式创建法”——造 “积木”

创建张量就像 “制作不同形状的积木”

(1)按 “现成数据” 创建:torch.tensor ()

如果你手里有现成的数字、列表或 NumPy 数组,用这个方法直接 “打包成张量”,就像把零散的珠子串成项链:

# 导包
import torch

#单个数字(0维张量,像一颗单独的珠子)
t1 = torch.tensor(10)
print(t1, t1.dtype,t1.ndim)
# 列表(1维张量,像一串珠子)
t2=torch.tensor([1,2,3])
print(t2, t2.dtype,t2.ndim)
# 二维(2维张量,像一块方形积木)
t3 = torch.tensor([[1, 2, 3], [4, 5, 6]])
print(t3, t3.dtype,t3.ndim)

结果为:

小技巧:还能指定数据类型,比如dtype=torch.int让张量变成整数型,避免后续计算出错。

(2)按 “模具” 拼:torch.Tensor ()

不知道具体数据,只知道形状?用这个方法先造一个 “空积木框”,后续再填内容。比如想要 2 行 3 列的积木:

# 2行3列的空张量(默认是float32类型,值是随机的)
t = torch.Tensor(2, 3)
print(t, t.dtype, t.ndim)
# 也能直接填数据,和torch.tensor()类似
t1= torch.Tensor([10,20,30])
print(t1, t1.dtype)

#直接转换int类型
t2 = torch.IntTensor(data=[1, 2, 3])
print(t2, t2.dtype)

结果为:

(3)指定 “材质” 拼:torch.XXXTensor ()

就像乐高有塑料、金属材质,张量也有不同数据类型。用torch.IntTensor()torch.FloatTensor()等,能直接指定 “积木材质”:

# 2行3列的整数型张量(int32)
t1 = torch.IntTensor(2, 3)  
# 1行4列的浮点型张量(float32)
t2 = torch.FloatTensor([1.1,2.2,3.3,4.4])  

这里结果就不展示了上面包含了有

(4)拼 “整齐的一排”:线性张量

想要按规律排列的积木?用arangelinspace,前者按 “步长” 排,后者按 “个数” 排:

# 从1到10,每步加3(左闭右开,像每隔3厘米放一块积木)
t1 = torch.arange(1, 10, 3)  # 结果:[1,4,7]
# 从1到10,平均分成4块(左闭右闭,像把10厘米的板切成4段)
t2 = torch.linspace(1, 10, 4)  # 结果:[1., 4., 7., 10.]

(5)拼 “随机花纹”:随机张量

训练模型时需要 “打乱的积木”?用这三个函数生成随机张量,还能固定 “花纹样式”(随机种子):

# 固定随机种子,每次生成的随机数都一样(方便复现)
torch.manual_seed(666)  
# 1.torch.rand(size): 创建一个指定形状的随机张量,0到1的随机浮点数(像撒了一把小石子)
t1 = torch.rand(2, 3)  
print(t1)
# 2.torch.randn(size): 创建一个指定形状的随机张量,正态分布的随机数(中间多两边少,像小山丘)
t2 = torch.randn(2, 3)
print(t2)  
# 3.torch.randint(low,high,size): 创建一个指定形状的随机张量,0到10的随机整数(像抽盲盒)
t3 = torch.randint(0, 10, (2, 3))  
print(t3)

结果为

(6)拼 “纯色积木”:全 0、全 1、指定值

需要 “统一颜色的积木”?比如做占位符、初始化权重,用这几个函数:

# 2行3列全0张量(像空白的白色积木)
t1 = torch.zeros(2, 3)  
# 2行3列全1张量(像纯色的红色积木)
t2 = torch.ones(2, 3)  
# 2行3列全是10的张量(像定制的蓝色积木)
t3 = torch.full((2, 3), 10)  

进阶技巧:如果想和已有张量形状一样,加个_like就行。比如torch.zeros_like(t1),会生成和 t1 形状相同的全 0 张量,不用再手动写形状。

比如这样

# 手动创建的一个3行3列的随机张量,要求使用xx_like()方法创建同形状的全0全1全指定值张量
t5 = torch.randint(1, 10, size=(3, 3))
print(t5)
t6 = torch.zeros_like(t5)
print(t6)
t7 = torch.ones_like(t5)
print(t7)
t8 = torch.full_like(t5, fill_value=10)
print(t8)

(7)拼 “特殊图案”:单位矩阵

想要 “对角线是 1、其他是 0” 的积木(比如 Identity Matrix)?用torch.eye()

# 3行3列的单位矩阵(像棋盘的对角线)
t = torch.eye(3, 3)  

2. 第二步:改 “积木”—— 张量的形状操作

拼积木时经常要调整形状,比如把方形积木改成长条,张量也能轻松做到:

(1)reshape:最常用的 “变形术”

把张量改成指定形状,只要总元素个数不变就行,像把一张纸折成不同形状:

# 2行3列的张量(6个元素)
t = torch.tensor([[1,2,3],[4,5,6]])  
# 改成1行6列(把方形纸卷成细条)
t1 = t.reshape(1, 6)  
# 用-1自动算维度(列数固定为2,行数自动算6/2=3)
t2 = t.reshape(-1, 2)  
  • reshape(-1, 2) 的作用:-1 表示 “自动计算该维度的大小”。这里要求列数固定为 2,因此行数由 “总元素数 ÷ 列数” 自动计算: 总元素数是 6,列数是 2,所以行数 = \(6 \div 2 = 3\),最终形状为 (3, 2)

这里加个附加 view:

张量内存连续时,reshape() 和 view() 效果一致

import torch

# 创建一个连续内存的张量
t = torch.tensor([[1, 2, 3], [4, 5, 6]])

# 使用 reshape 改变形状
t_reshape = t.reshape(3, 2)
print("reshape result:")
print(t_reshape)

# 使用 view 改变形状
t_view = t.view(3, 2)
print("view result:")
print(t_view)

张量内存不连续时,reshape() 和 view() 的区别:

import torch

# 创建一个张量并进行转置,使其内存不连续
t = torch.tensor([[1, 2, 3], [4, 5, 6]]).transpose(0, 1)
print("Is t contiguous before view/reshape:", t.is_contiguous())

# 尝试使用 view 改变形状,会报错
try:
    t_view = t.view(6)
except RuntimeError as e:
    print(f"Error when using view: {e}")

# 使用 reshape 改变形状,不会报错
t_reshape = t.reshape(6)
print("reshape result:")
print(t_reshape)
print("Is t_reshape contiguous after reshape:", t_reshape.is_contiguous())

在不连续张量上先使用 contiguous() 再用 view()

import torch

# 创建一个张量并进行转置,使其内存不连续
t = torch.tensor([[1, 2, 3], [4, 5, 6]]).transpose(0, 1)
print("Is t contiguous before contiguous():", t.is_contiguous())

# 使用 contiguous() 使张量内存连续
t_contiguous = t.contiguous()
print("Is t_contiguous contiguous:", t_contiguous.is_contiguous())

# 此时可以成功使用 view 改变形状
t_view = t_contiguous.view(6)
print("view result:")
print(t_view)

在这个示例中,先对不连续的张量 t 使用 contiguous() 方法使其内存连续,然后就可以顺利使用 view() 改变其形状了。

(2)squeeze/unsqueeze:给积木 “加套” 或 “脱套”

  • unsqueeze(dim=维度):给张量加一个 “1 维套子”,比如把 1 维的 “细棍” 放进 1 行 5 列的 “小盒子”(升维);
  • squeeze(dim=维度):把多余的 “套子” 去掉,比如把 1 行 5 列的 “小盒子” 变回 “细棍”(降维)。
t = torch.tensor([1,2,3])  # 1维,shape=[3]
# 加套:变成1行3列(0维度加1)
t1 = t.unsqueeze(0)  # shape=[1,3]
# 加套:变成3行1列(1维度加1)
t2 = t.unsqueeze(1)  # shape=[3,1]
# 脱套:去掉所有1维,变回1维
t3 = t2.squeeze()  # shape=[3]

(3)transpose/permute:交换积木的 “面”

就像把魔方的面转来转去,这两个函数能交换张量的维度:

  • transpose:只能交换两个维度,比如把 3 行 4 列的张量转成 4 行 3 列;
  • permute:能一次交换多个维度,更灵活。
# 3维张量:2个3行4列的“积木块”
t = torch.randint(0,10,(2,3,4))  
# 交换1和2维度(3和4交换,变成2个4行3列)
t1 = t.transpose(1,2)  
# 一次交换0、1、2维度(变成4行3列2个)
t2 = t.permute(2,1,0)  

3. 第三步:取 “积木”—— 张量的索引操作

拼积木时经常需要从一堆积木里挑出特定的几块,张量的索引操作就像 “挑积木”,能精准获取你需要的部分。无论是单个元素、某一行 / 列,还是满足条件的元素,都能轻松提取。

1)单个索引:精准定位 “单块积木”

就像从积木堆里拿出特定位置的一块,通过张量[行,列]的格式能直接获取单个元素或某一行 / 列。对于二维张量(类似表格):

import torch
torch.manual_seed(666)
data = torch.randint(1, 10, size=(4, 4))  # 4行4列的随机整数张量
print("原始张量:\n", data)

# 取第2行所有元素(索引从0开始,第2行即索引1)
print("第2行:", data[1, :])  # 等价于 data[1]
# 取第3列所有元素(第3列即索引2)
print("第3列:", data[:, 2])
# 取第2行第3列的元素
print("第2行第3列:", data[1, 2])

小技巧:表示 “取所有”,比如data[1, :]就是 “第 2 行的所有列”。

2)列表索引:批量挑选 “多块积木”

如果需要同时取多行、多列,或多个分散的位置,用列表索引就像 “一次抓出好几块指定的积木”:

# 同时取第2行和第4行(索引1和3)
print("第2行和第4行:\n", data[[1, 3], :])  # 等价于 data[[1, 3]]
# 同时取第2列和第4列(索引1和3)
print("第2列和第4列:\n", data[:, [1, 3]])
# 同时取(0,0)和(1,2)两个位置的元素
print("(0,0)和(1,2)元素:", data[[0, 1], [0, 2]])
# 取第1行和第2行的第1列和第3列(索引0、1行,0、2列)
print("指定行和列的交叉元素:\n", data[[[0], [1]], [0, 2]])

(3)切片索引:连续抓取 “一整段积木”

想取连续的几行或几列(比如前 3 行、前 3 列),用切片start:end就像 “从积木堆里拿出一整排”:

# 取前3行(索引0到2,左闭右开)
print("前3行:\n", data[0:3, :])  # 等价于 data[0:3]
# 取前3列
print("前3列:\n", data[:, 0:3])
# 取前3行前3列的子张量(左上角3×3区域)
print("前3行前3列:\n", data[0:3, 0:3])

(4)布尔索引:按条件筛选 “符合要求的积木”

如果想按条件(比如元素大于 5)挑选积木,先用条件判断生成 “布尔张量”(True/False),再用它做索引:

# 生成“元素是否大于5”的布尔张量
mask = data > 5
print("布尔掩码:\n", mask)
# 提取所有大于5的元素
print("大于5的元素:", data[mask])

# 也能直接用布尔列表指定行/列,比如取第2行和第4行
print("第2行和第4行(布尔索引):\n", data[[False, True, False, True], :])

(5)多维张量索引:高维积木的 “分层挑选”

三维及以上的张量(比如(4,4,5)的三维张量,可理解为 4 个 4 行 5 列的 “积木层”),索引时需要指定每个维度的位置:

# 创建三维张量(4,4,5):4个“层”,每层4行5列
data_3d = torch.randint(1, 10, size=(4, 4, 5))

# 取0轴(第1个维度)的第2个“层”(索引1)
print("0轴第2个层:\n", data_3d[1, :, :])
# 取1轴(第2个维度)的第3行(索引2)
print("1轴第3行:\n", data_3d[:, 2, :])
# 取2轴(第3个维度)的第3列(索引2)
print("2轴第3列:\n", data_3d[:, :, 2])

记忆法:三维张量的维度可记为(层, 行, 列),索引时按[层, 行, 列]依次指定即可。

4. 第四步:拼 “积木”—— 张量的运算和拼接

积木可以拼在一起,张量也能做运算、组合成更大的结构:

(1)基本运算:加加减减除除很简单

像给积木涂颜色、加装饰,张量的加减乘除直接用符号或函数:

t = torch.tensor([[3,7,4],[0,0,6]])  
# 所有元素加10(不修改原张量,像复制一份再涂色)
t1 = t.add(10)  # 等价于 t+10
# 所有元素加10(直接修改原张量,像直接在积木上涂色)
# add_(), mul_(), div_(), sub_()直接修改原张量
t.add_(10)  # 等价于 t +=10
#除
print(t / 10)
print(t.div(10))
t.div_(2)

注意:带下划线的函数(如add_)会直接修改原张量,不带的会生成新张量,别搞混啦!

(2)两种 “乘法”:别搞混!

张量有两种乘法,用法完全不同,像拼积木的两种方式:

点乘(对应元素乘):两个同样形状的张量,对应位置 “积木块” 相乘,用*torch.mul

t1 = torch.tensor([[1,2],[3,4]])
t2 = torch.tensor([[5,6],[7,8]])
t3 = t1 * t2  # 结果:[[5,12],[21,32]]

矩阵乘法:按矩阵规则相乘(要求第一个张量的列数等于第二个的行数),用@torch.matmul

# 3行2列 × 2行4列 = 3行4列
t1 = torch.tensor([[1,2],[3,4],[5,6]])
t2 = torch.tensor([[5,6],[7,8]])
t3 = t1 @ t2  # 结果:[[19,22],[43,50],[67,78]]
print(t1.matmul(t2))

(3)拼接积木:cat 和 stack

想把多堆积木拼在一起?用这两个函数,区别在于是否加新维度:

① torch.cat:在现有维度上 “接长”,不新增维度

就像把两根乐高棍接成一根长棍,cat要求 “除了拼接的维度,其他维度必须完全相同”,否则会 “接口不匹配”:

# 创建两个张量:data1是2行3列,data2是1行3列
data1 = torch.randint(1, 10, size=(2, 3))
data2 = torch.randint(1, 10, size=(1, 3))

# 沿0轴(行)拼接:2行+1行=3行,列数都是3(匹配),拼接成功
print("cat沿0轴拼接:\n", torch.cat([data1, data2], dim=0))

# 沿1轴(列)拼接:data1列数3,data2列数3,但行数2≠1(不匹配),会报错
# print(torch.cat([data1, data2], dim=1))  # 报错!

对于三维张量,还能用负数表示维度(-3对应第 0 维,-2对应第 1 维,-1对应第 2 维),方便快速定位:

# 两个三维张量(1,2,3)
data1 = torch.randint(1, 10, size=(1, 2, 3))
data2 = torch.randint(1, 10, size=(1, 2, 3))

# 沿0轴(或-3轴)拼接:(1+1,2,3)=(2,2,3)
print("沿0轴拼接:\n", torch.cat([data1, data2], dim=0))
print("沿-3轴拼接(等价于0轴):\n", torch.cat([data1, data2], dim=-3))

② torch.stack:新增维度 “叠起来”,要求所有维度完全相同

如果想把两堆积木 “上下叠成一层新的”(新增维度),stack要求 “所有维度必须完全相同”,否则会 “大小不匹配”:

# 两个形状相同的张量:都是2行3列
data1 = torch.randint(1, 10, size=(2, 3))
data2 = torch.randint(1, 10, size=(2, 3))

# 沿0轴新增维度拼接:结果形状为(2,2,3)(第0维是新维度,包含2个原张量)
t = torch.stack([data1, data2], dim=0)
print("stack沿0轴拼接形状:", t.shape)  # torch.Size([2, 2, 3])

# 如果两个张量形状不同(比如data1是2行3列,data2是1行3列),会报错
data2 = torch.randint(1, 10, size=(1, 3))
# t = torch.stack([data1, data2], dim=0)  # 报错!因为形状不同

总结cat是 “接长”(不增维),stack是 “叠高”(增新维);cat只要求非拼接维度相同,stack要求所有维度都相同。

5. 第五步:换 “积木材质”—— 张量类型转换

不同 “积木材质”(数据类型)不能混用,比如整数型和浮点型一起运算会报错。转换方法有两种,像给积木刷漆:

#生成一个三行俩列的张量,数值全部为10
t = torch.full((2,3), 10)  # 默认int64类型
# 方法1:用type()指定类型
t1 = t.type(torch.FloatTensor)  # 转成float32
# 方法2:用快捷函数(更简单)
t2 = t.float()  # 转成float32
t3 = t.double()  # 转成float64
t4 = t.int()     # 转成int32

避坑提醒:计算均值、方差时,张量必须是浮点类型(float/double),否则会报错!

6. 第六步:和 “其他玩具” 互动 —— 张量与 NumPy 互转

有时候你会用到 NumPy 的 “小工具”(比如数据分析),需要把张量和 NumPy 数组互相转换,像把乐高积木和拼图零件互换:

(1)张量转 NumPy

numpy()方法,但要注意共享内存—— 改一个,另一个也会变,就像两个玩具共用一个零件:

t = torch.tensor([2,3,4])
np_arr = t.numpy()  # 转成NumPy数组
np_arr[0] = 100     # 修改NumPy数组
print(t)  # 结果:tensor([100,  3,  4])(张量也变了!)

想避免共享内存?加个copy()

np_arr = t.numpy().copy()  # 复制一份,互不影响

(2)NumPy 转张量

有两种方法,区别在于是否共享内存:

  • torch.from_numpy():共享内存,改一个另一个变;
  • torch.tensor():不共享内存,改一个另一个不变。
np_arr = np.array([2,3,4])
# 方法1:共享内存
t1 = torch.from_numpy(np_arr)
t1[0] = 100
print(np_arr)  # 结果:[100  3  4](NumPy数组变了)

# 方法2:不共享内存
t2 = torch.tensor(np_arr)
t2[0] = 200
print(np_arr)  # 结果:[100  3  4](NumPy数组没变)

(3)单个 “积木” 提取:张量转数字

如果张量只有一个元素(比如损失值),用item()把它提取成 Python 数字,像把单独的积木从盒子里拿出来:

t = torch.tensor([30])
num = t.item()  # 结果:30(Python的int类型)

三、自动 “拼积木”:微分工具有多香?

拼积木时最怕拼错了要从头拆,而 PyTorch 的自动微分就像 “智能纠错仪”—— 能自动找到拼错的地方,告诉你该怎么调整,不用你手动算导数。

1. 核心原理:给积木贴 “监控标签”

想要自动纠错,得先告诉 PyTorch “哪些积木要监控”。创建张量时加requires_grad=True,就像给积木贴了 “重点监控” 标签:

# 贴标签:这个x要监控,后续运算要记下来
x = torch.tensor(10.0, requires_grad=True)

2. 实战:算梯度,找 “拼错的地方”

比如我们要优化函数y = 2x²,想知道 “x 怎么变,y 才能变小”。自动微分会帮你算 “梯度”(导数),梯度的方向就是 y 下降最快的方向:

# 定义函数(拼积木的步骤)
y = 2 * x ** 2
# 反向传播:算梯度(纠错仪开始工作)
y.backward()
# 查看x的梯度(纠错建议)
print(x.grad)  # 结果:tensor(40.)

意思是 “x 每减少 1,y 会减少 40”—— 这就是模型训练时调整参数的依据!

3. 进阶:用梯度下降 “拼出最优模型”

知道了梯度,就能用 “梯度下降法” 一步步调整参数,让 y 最小化。公式很简单:新参数 = 旧参数 - 学习率×梯度,像根据纠错建议慢慢调整积木位置:

多轮迭代的关键:梯度清零

当进行多轮训练时,梯度会自动累加(就像每次的纠错建议会叠加),必须手动清零才能保证每轮的梯度计算准确:

# 目标:找y = 2x²的最小值(理论上x=0时y=0)
x = torch.tensor(10.0, requires_grad=True)
lr = 0.01  # 学习率
epochs = 180  # 多轮迭代

for epoch in range(epochs):
    y = 2 * x **2
    # 关键:每次迭代前清空梯度(避免累加)
    if x.grad is not None:
        x.grad.zero_()  # 下划线表示原地修改
    # 反向传播算梯度
    y.sum().backward()  # 标量才能求导,sum()确保输出是标量
    # 更新参数
    x.data = x.data - lr * x.grad

print(x)  # 结果接近0,成功找到最小值点

4. 避坑提醒:自动微分的 5 个小陷阱

-** 梯度累加 :每次反向传播后,梯度会存在x.grad里,如果不手动清零(x.grad.zero_()),下次会累加,导致建议出错;

- 数据类型 :只有浮点类型的张量能求导,int 类型会报错;

- detach () 的用法 :带自动微分的张量转 NumPy 需先detach()

- 标量求导 backward()要求被求导的对象必须是标量(单个数值),如果是张量需要先用sum()mean()转换,例如loss.sum().backward()

- 多参数梯度计算 **:当模型有多个参数(如权重 w 和偏置 b)时,backward()会自动计算所有参数的梯度,无需单独处理:

w = torch.tensor([1.0], requires_grad=True)
b = torch.tensor([2.0], requires_grad=True)
y = w * 3 + b  # 同时依赖w和b
y.backward()
print(w.grad)  # tensor([3.]),w的梯度
print(b.grad)  # tensor([1.]),b的梯度

四、形状的 “隐形陷阱”:广播机制详解

拼积木时如果形状不匹配,可能会 “强行拼接” 导致结果错误 ——PyTorch 的广播机制就是这样一种 “隐形拼接规则”,用得好能简化代码,用不好会悄悄出错!

1. 什么是广播?

当两个张量形状不同时,PyTorch 会尝试自动扩展它们的维度,使形状匹配后再计算。比如[3,1][3]会被扩展成[3,3]后再运算,但这可能不是你想要的结果!

2. 实战坑例:损失计算的形状陷阱

用均方误差(MSELoss)计算损失时,形状不匹配会导致广播错误:

import torch
from torch import nn

# 预测值:3行1列([3,1])
y_pred = torch.tensor([[1.0], [2.0], [3.0]])
# 真实值:3个元素([3])
batch_y = torch.tensor([1.0, 2.0, 3.0])

# 错误做法:直接计算,触发广播
loss = nn.MSELoss()(y_pred, batch_y)
print(loss)  # 结果1.333,明显错误!

# 正确做法:将真实值reshape成和预测值相同的形状
loss = nn.MSELoss()(y_pred, batch_y.reshape(-1, 1))  # [3]→[3,1]
print(loss)  # 结果0.0,正确!

为什么错误?

y_pred([3,1])和batch_y([3])会被广播成[3,3],相当于计算了 9 个元素的误差(而不是 3 个),导致结果错误。** 解决办法 **:始终确保预测值和真实值的形状完全一致(如都转成[n,1])。

五、终极实战:拼一个 “线性回归机器人”

学了这么多基础,咱们来拼一个完整的 “AI 模型”—— 线性回归,相当于用 PyTorch 搭一个 “能拟合直线的小机器人”。目标是找到一条直线y = wx + b,让它尽可能贴近散点数据。

步骤1:准备“积木材料”(数据集)

from sklearn.datasets import make_regression
from torch.utils.data import TensorDataset, DataLoader

# 生成数据:100个样本,1个特征,真实斜率coef,真实截距14.5
x, y, coef = make_regression(
    n_samples=100, n_features=1, noise=10, coef=True, bias=14.5, random_state=1
)
# 转成张量
x = torch.tensor(x, dtype=torch.float32)
y = torch.tensor(y, dtype=torch.float32)
# 打包成数据集,再分成批次(batch_size=32)
dataset = TensorDataset(x, y)  # 将x和y绑定成数据集
dataloader = DataLoader(dataset, batch_size=32, shuffle=True)  # 打乱数据并分批次
print(f"数据集总样本数:{len(dataset)},批次数量:{len(dataloader)}")  # 100样本→4批次(32+32+32+4)

步骤 2:搭 “机器人骨架”(定义模型)

import torch.nn as nn

# 输入特征数=1,输出特征数=1(y = wx + b)
model = nn.Linear(in_features=1, out_features=1)
# 查看初始参数(随机初始化,像机器人的初始骨骼)
print("初始权重w:", model.weight.data)
print("初始偏置b:", model.bias.data)

步骤 3:装 “纠错和驱动系统”(损失函数和优化器)

import torch.optim as optim

# 损失函数:均方误差(MSE)
loss_fn = nn.MSELoss()
# 优化器:随机梯度下降(SGD),学习率0.01
optimizer = optim.SGD(model.parameters(), lr=0.01)

步骤 4:启动 “机器人”(模型训练)

epochs = 100  # 训练100轮
loss_list = []  # 记录每轮损失

for epoch in range(epochs):
    total_loss = 0.0  # 累计本轮总损失
    batch_cnt = 0     # 累计本批次数量
    for batch_x, batch_y in dataloader:  # 遍历每个批次
        # 1. 前向传播:用模型预测
        y_pred = model(batch_x)  # 等价于 w*batch_x + b
        # 2. 计算损失(注意形状匹配:将batch_y转成[32,1])
        loss = loss_fn(y_pred, batch_y.reshape(-1, 1))
        # 3. 梯度清零
        optimizer.zero_grad()
        # 4. 反向传播:计算梯度
        loss.backward()
        # 5. 参数更新
        optimizer.step()
        # 累计损失和批次
        total_loss += loss.item()
        batch_cnt += 1
    # 计算本轮平均损失并记录
    epoch_loss = total_loss / batch_cnt
    loss_list.append(epoch_loss)
    print(f"第{epoch+1}轮,平均损失:{epoch_loss:.4f}")

# 保存模型参数(像保存机器人的最终状态)
torch.save(model.state_dict(), "linear_model.pth")
print("模型保存成功!")

步骤 5:看 “机器人表现”(可视化结果)

import matplotlib.pyplot as plt

plt.rcParams['font.sans-serif'] = ['KaiTi']  # 显示中文
plt.rcParams['axes.unicode_minus'] = False  # 显示负号

# 1. 绘制损失变化曲线
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.plot(range(epochs), loss_list)
plt.title("每轮损失变化")
plt.xlabel("轮次")
plt.ylabel("均方误差")
plt.grid()

# 2. 绘制拟合结果
plt.subplot(1, 2, 2)
# 绘制真实散点
plt.scatter(x.numpy(), y.numpy(), label="真实数据", alpha=0.6)
# 绘制真实直线(y = coef*x + 14.5)
x_range = torch.linspace(x.min(), x.max(), 1000)
y_true = coef * x_range + 14.5
plt.plot(x_range, y_true, label="真实直线", color="red")
# 绘制模型拟合直线(y = model.weight*x + model.bias)
y_pred = model(x_range.reshape(-1, 1)).detach().numpy()  # 用detach()脱离自动微分
plt.plot(x_range, y_pred, label="拟合直线", color="green", linestyle="--")
plt.legend()
plt.grid()
plt.show()

六、总结:PyTorch 入门的 “核心口诀”

  1. 张量是基础:创建、变形、索引、运算、拼接、互转,这六步掌握了,就能玩转 80% 的场景;
  2. 微分是核心:贴标签(requires_grad=True)、算梯度(backward ())、清梯度(zero_())、更参数,这是模型训练的 “四步口诀”;
  3. 形状要匹配:广播机制虽方便,但损失计算等场景必须保证形状一致,避免隐形错误;
  4. 模型训练四步走:准备数据→搭模型→设损失和优化器→循环训练(前向传播→算损失→反向传播→参数更新),所有深度学习模型都离不开这个流程。

PyTorch 的精髓就在于 “简单直观”,不用死记硬背 API,多动手改改代码(比如尝试调整线性回归的 batch_size 或学习率,观察损失变化),很快就能上手。下次咱们可以进阶学习 CNN、RNN,用 PyTorch 拼更复杂的 “AI 机器人”!

 

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐