PyTorch 深度学习笔记(四)张量运算与变形——数值计算与索引进阶

📅 2026/8/1 8:33:29
PyTorch 深度学习笔记(四)张量运算与变形——数值计算与索引进阶
系列导读本系列共 6 篇从 PyTorch 框架入门到实战案例带你系统掌握深度学习开发。上一篇PyTorch 张量基础——从零开始掌握 Tensor一、张量数值计算1.1 点乘运算Hadamard 积点乘是元素级乘积对应位置元素相乘使用mul()或*运算符importtorch data1torch.tensor([[1,3],[2,4]])data2torch.tensor([[2,2],[3,3]])# 方式 1torch.mul()datatorch.mul(data1,data2)print(data)# tensor([[2, 6],# [6, 12]])# 方式 2* 运算符datadata1*data2print(data)1.2 矩阵乘法 ⭐⭐⭐矩阵乘法要求第一个矩阵(n, m)× 第二个矩阵(m, p) 结果(n, p)Atorch.tensor([[1,2,3],[4,5,6]])# shape (2, 3)Btorch.tensor([[7,8],[9,10],[11,12]])# shape (3, 2)# 方式 1 运算符推荐data3A Bprint(A B --,data3)# 方式 2torch.matmul()data4torch.matmul(A,B)print(matmul --,data4)# 方式 3torch.mm()仅限 2Ddatatorch.mm(A,B)print(mm --,data)二、张量统计函数2.1 常用统计函数importtorchimportmath torch.manual_seed(42)xtorch.randint(0,10,(2,3),dtypetorch.float64)print(x)# tensor([[8., 2., 5.],# [1., 6., 7.]])# 1. 平均值print(平均值,x.mean())# 全局平均print(按列平均,x.mean(dim0))# 压缩行按列计算print(按行平均,x.mean(dim1))# 压缩列按行计算# 2. 求和print(求和,x.sum())print(按列求和,x.sum(dim0))print(按行求和,x.sum(dim1))# 3. 最值print(最小值,x.min())print(按行最小值,x.min(dim0))print(按列最小值,x.min(dim1))# 4. 幂次方print(平方,x.pow(2))# 5. 平方根print(平方根,x.sqrt())# 6. 指数 e^xprint(exp,x.exp())# 7. 自然对数 ln(x)print(log,x.log())dim 参数含义dim0按列计算压缩行结果形状去掉第 0 维dim1按行计算压缩列结果形状去掉第 1 维三、张量索引操作 ⭐⭐⭐⭐3.1 基础索引importtorch datatorch.randint(0,10,(3,4))print(data)# 行索引print(data[0])# 第 1 行# 列索引print(data[:,0])# 第 1 列# 单个元素print(data[0,1])# 第 1 行第 2 列3.2 列表索引# 取 (0,1) 和 (1,2) 两个位置的元素print(data[[0,1],[1,2]])# 输出tensor([data[0,1], data[1,2]])3.3 范围索引# 前 3 行的前 2 列print(data[:3,:2])# 第 2 行到最后的前 2 列print(data[2:,:2])3.4 布尔索引 ⭐⭐⭐# 第三列大于 5 的所有行print(data[data[:,2]5])# 第二行大于 5 的列maskdata[1]5resultdata[:,mask]print(result)3.5 高级索引难点# 取 0、1 行的 1、2 列共 4 个元素print(data[[[0],[1]],[1,2]])3.6 多维索引# 三维张量datatorch.randint(0,10,[3,4,5])# 0 轴上的第 1 个数据print(data[0,:,:])# 1 轴上的第 1 个数据print(data[:,0,:])# 2 轴上的第 1 个数据print(data[:,:,0])四、张量形状操作 ⭐⭐⭐⭐4.1 reshape在保持数据不变的前提下改变维度datatorch.tensor([[1,3,5],[2,4,8]])print(data.shape)# torch.Size([2, 3])# reshape 为 (1, 6)data2data.reshape(1,6)print(data2.shape)# torch.Size([1, 6])# reshape 为 (3, 2)data2data.reshape(3,2)print(data2.shape)# torch.Size([3, 2])4.2 squeeze 与 unsqueezedatatorch.randint(0,6,[2,3])print(data.shape)# torch.Size([2, 3])# unsqueeze在指定位置添加维度1升维y1data.unsqueeze(0)# 位置 0 添加维度 → shape (1, 2, 3)y2data.unsqueeze(1)# 位置 1 添加维度 → shape (2, 1, 3)y3data.unsqueeze(2)# 位置 2 添加维度 → shape (2, 3, 1)# squeeze删除维度1 的位置降维y4y3.squeeze(2)# 删除位置 2 的维度 1 → shape (2, 3)# squeeze() 不指定位置则删除所有维度1记忆口诀unsqueeze(0)→ 包成一整叠 →(1, 2, 3)比作书本unsqueeze(1)→ 每行单独包 →(2, 1, 3)比作每页纸unsqueeze(2)→ 每个数字单独包 →(2, 3, 1)比作纸上的字4.3 transpose 与 permutedatatorch.randint(0,6,(2,3,4))print(data.shape)# torch.Size([2, 3, 4])# transpose交换两个维度data2data.transpose(0,1)# 交换 0 和 1 → shape (3, 2, 4)data3data.transpose(0,2)# 交换 0 和 2 → shape (4, 3, 2)# permute一次交换多个维度参数是维度序号data3data.permute((2,0,1))# (2,3,4) → (4,2,3)print(data3.shape)# torch.Size([4, 2, 3])4.4 view 与 contiguous ⭐⭐⭐# view类似于 reshape但要求张量内存连续datatorch.tensor([[1,2,3],[4,5,6]])# 判断张量是否连续print(data.is_contiguous())# True# 连续的可以直接 viewdata2data.view(3,2)# transpose 后可能不连续data3data.transpose(0,1)print(data3.is_contiguous())# False# 不连续张量使用 view 会报错# data4 data3.view(3, 2) # ❌ 报错# 先用 contiguous() 转为连续再用 viewdata4data3.contiguous().view(3,2)print(data4.is_contiguous())# True判断方法id()判断是否是同一个 Python 对象data_ptr()判断是否共享底层存储五、运算与变形总结类别操作方法说明点乘元素级乘积*/torch.mul()对应位置相乘矩阵乘法矩阵乘积/matmul()/mm()线性代数矩阵乘统计平均值/求和/最值mean() / sum() / min() / max()dim 指定压缩维度数学函数幂/根/指数/对数pow() / sqrt() / exp() / log()元素级运算基础索引行/列/范围data[0] / data[:,0] / data[:3,:2]类似 NumPy布尔索引条件筛选data[data 5]返回满足条件的元素多维索引高维张量data[:, 0, :]指定轴取数据reshape改变形状.reshape()不改变数据squeeze降维.squeeze(dim)删除维度1unsqueeze升维.unsqueeze(dim)添加维度1transpose交换维度.transpose(a, b)一次交换两个permute重排维度.permute(dims)一次重排多个view视图变形.view()要求内存连续contiguous内存连续化.contiguous()配合 view 使用六、下一篇预告PyTorch 深度学习笔记五张量拼接与自动微分——构建神经网络基础将详细介绍张量拼接cat/stack、拆分chunk/split以及自动微分模块 autograd 的原理和使用方法。如果这篇文章对你有帮助欢迎点赞、收藏、关注你的支持是我持续创作的动力。