PyTorch 入门:像玩 “乐高” 一样搞懂深度学习工具!
如果把深度学习比作 “搭积木建房子”,那 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)拼 “整齐的一排”:线性张量
想要按规律排列的积木?用arange或linspace,前者按 “步长” 排,后者按 “个数” 排:
# 从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 入门的 “核心口诀”
- 张量是基础:创建、变形、索引、运算、拼接、互转,这六步掌握了,就能玩转 80% 的场景;
- 微分是核心:贴标签(requires_grad=True)、算梯度(backward ())、清梯度(zero_())、更参数,这是模型训练的 “四步口诀”;
- 形状要匹配:广播机制虽方便,但损失计算等场景必须保证形状一致,避免隐形错误;
- 模型训练四步走:准备数据→搭模型→设损失和优化器→循环训练(前向传播→算损失→反向传播→参数更新),所有深度学习模型都离不开这个流程。
PyTorch 的精髓就在于 “简单直观”,不用死记硬背 API,多动手改改代码(比如尝试调整线性回归的 batch_size 或学习率,观察损失变化),很快就能上手。下次咱们可以进阶学习 CNN、RNN,用 PyTorch 拼更复杂的 “AI 机器人”!
更多推荐



所有评论(0)