PyTorch 基本使用学习笔记


PyTorch 基本使用学习笔记

1. PyTorch 简介

  • 基于 Python 的深度学习框架,核心数据抽象是 张量(Tensor)

  • 特点:类似 NumPy 的张量计算、自动微分、动态计算图、GPU 加速(CUDA)、跨平台。

  • 安装:

    bash 复制代码
    pip install torch -i https://pypi.tuna.tsinghua.edu.cn/simple

2. 张量的创建

2.1 基本创建方式

方法 说明
torch.tensor(data) 根据指定数据创建张量
torch.Tensor(shape) 根据形状创建张量,也可传入数据
torch.IntTensor / FloatTensor / DoubleTensor 创建指定类型的张量
python 复制代码
import torch
import numpy as np

# 标量
data = torch.tensor(10)

# 从 numpy 数组创建
data = np.random.randn(2, 3)
data = torch.tensor(data)

# 从列表创建,默认 float32
data = [[10., 20., 30.], [40., 50., 60.]]
data = torch.tensor(data)

# 根据形状创建
data = torch.Tensor(2, 3)          # 2行3列,默认 float32
data = torch.Tensor([10])          # tensor([10.])
data = torch.Tensor([10, 20])      # tensor([10., 20.])

# 指定类型
data = torch.IntTensor(2, 3)       # int32
data = torch.IntTensor([2.5, 3.3]) # 类型转换 → tensor([2, 3], dtype=torch.int32)
data = torch.ShortTensor()         # int16
data = torch.LongTensor()          # int64
data = torch.FloatTensor()         # float32
data = torch.DoubleTensor()        # float64

2.2 线性张量与随机张量

python 复制代码
# 线性张量
torch.arange(0, 10, 2)      # tensor([0, 2, 4, 6, 8])
torch.linspace(0, 9, 10)    # 在 [0,9] 上等距生成 10 个数

# 随机种子
torch.initial_seed()        # 查看当前随机种子
torch.manual_seed(100)      # 设置随机种子

# 随机浮点张量
torch.randn(2, 3)           # 标准正态分布

# 随机整数张量
torch.randint(low, high, size=())

2.3 0、1、指定值张量

python 复制代码
torch.zeros(2, 3)           # 全 0
torch.zeros_like(data)      # 按 data 形状创建全 0
torch.ones(2, 3)            # 全 1
torch.ones_like(data)       # 按 data 形状创建全 1
torch.full([2, 3], 10)      # 全为指定值
torch.full_like(data, 20)   # 按 data 形状创建指定值

2.4 张量元素类型转换

python 复制代码
data.type(torch.DoubleTensor)   # 转为 float64
data.half()                     # float16
data.double()                   # float64
data.float()                    # float32
data.short()                    # int16
data.int()                      # int32
data.long()                     # int64

3. 张量与 NumPy 的转换

3.1 张量 → NumPy 数组

  • Tensor.numpy():共享内存,修改一个会影响另一个。
  • Tensor.numpy().copy():避免共享内存。
python 复制代码
data_tensor = torch.tensor([2, 3, 4])
data_numpy = data_tensor.numpy()          # 共享内存
data_numpy = data_tensor.numpy().copy()   # 不共享内存

3.2 NumPy 数组 → 张量

  • torch.from_numpy(ndarray):共享内存。
  • torch.tensor(ndarray):不共享内存。
python 复制代码
data_numpy = np.array([2, 3, 4])
data_tensor = torch.from_numpy(data_numpy)   # 共享内存
data_tensor = torch.tensor(data_numpy)       # 不共享内存

3.3 标量张量 → 数字

python 复制代码
data = torch.tensor(30)
data.item()   # 30

4. 张量数值计算

4.1 基本运算

python 复制代码
data = torch.randint(0, 10, [2, 3])

data.add(10)      # 不修改原数据
data.add_(10)     # 修改原数据(带下划线)
data.sub(100)
data.mul(100)
data.div(100)
data.neg()

4.2 点乘(Hadamard 积)

相同形状张量对应位置相乘:

python 复制代码
data1 = torch.tensor([[1, 2], [3, 4]])
data2 = torch.tensor([[5, 6], [7, 8]])
torch.mul(data1, data2)   # 或 data1 * data2
# tensor([[ 5, 12],
#         [21, 32]])

4.3 矩阵乘法

要求:第一个矩阵 shape (n, m),第二个矩阵 shape (m, p),结果 shape (n, p)

python 复制代码
data1 = torch.tensor([[1, 2], [3, 4], [5, 6]])
data2 = torch.tensor([[5, 6], [7, 8]])

data1 @ data2                # 运算符 @
torch.matmul(data1, data2)   # torch.matmul

5. 张量运算函数

python 复制代码
data = torch.randint(0, 10, [2, 3], dtype=torch.float64)

data.mean()          # 均值
data.mean(dim=0)     # 按列
data.mean(dim=1)     # 按行
data.sum()           # 总和
data.sum(dim=0)
data.sum(dim=1)
torch.pow(data, 2)   # 平方
data.sqrt()          # 平方根
data.exp()           # e^n
data.log()           # 以 e 为底
data.log2()          # 以 2 为底
data.log10()         # 以 10 为底

注意:mean() 要求张量为 Float 或 Double 类型。


6. 张量索引操作

python 复制代码
data = torch.randint(0, 10, [4, 5])

data[0]              # 第 0 行
data[:, 0]           # 第 0 列
data[[0, 1], [1, 2]] # 返回 (0,1)、(1,2) 两个元素
data[[0, 1], [1, 2]] # 返回 0、1 行的 1、2 列共 4 个元素
data[:, :2]          # 前 2 列
data[2:, :2]         # 第 2 行到最后的前 2 列
data[data[:, 2] > 5] # 布尔索引:第三列大于 5 的行
data[:, data[1] > 5] # 布尔索引:第二行大于 5 的列

多维索引:

python 复制代码
data = torch.randint(0, 10, [3, 4, 5])
data[0, :, :]   # 0 轴第一个数据
data[:, 0, :]   # 1 轴第一个数据
data[:, :, 0]   # 2 轴第一个数据

7. 张量形状操作

函数 说明
reshape() 改变形状,数据不变
squeeze() 删除形状为 1 的维度
unsqueeze(dim) 添加形状为 1 的维度
transpose(dim0, dim1) 交换两个维度
permute(dims) 一次交换多个维度
view() 改变形状,要求内存连续
contiguous() 将张量变为连续,配合 view 使用
python 复制代码
data = torch.tensor([[10, 20, 30], [40, 50, 60]])
data.shape, data.size()          # torch.Size([2, 3])

data.reshape(1, 6)               # torch.Size([1, 6])

mydata1 = torch.tensor([1, 2, 3, 4, 5])
mydata1.unsqueeze(dim=0)         # torch.Size([1, 5])
mydata1.unsqueeze(dim=1)         # torch.Size([5, 1])
mydata1.unsqueeze(dim=-1)        # torch.Size([5, 1])
mydata1.unsqueeze(dim=0).squeeze()  # torch.Size([5])

data = torch.tensor(np.random.randint(0, 10, [3, 4, 5]))
torch.transpose(data, 1, 2)      # torch.Size([3, 5, 4])
torch.permute(data, [1, 2, 0])   # torch.Size([4, 5, 3])
data.permute([1, 2, 0])          # torch.Size([4, 5, 3])

data = torch.tensor([[10, 20, 30], [40, 50, 60]])
data.is_contiguous()             # True
data.view(3, 2)                  # 可用

mydata3 = torch.transpose(data, 0, 1)
mydata3.is_contiguous()          # False
mydata3.contiguous().view(2, 3)  # 先用 contiguous 再用 view

8. 张量拼接操作

8.1 torch.cat()

按指定维度拼接,不增加维度。

python 复制代码
data1 = torch.randint(0, 10, [1, 2, 3])
data2 = torch.randint(0, 10, [1, 2, 3])

torch.cat([data1, data2], dim=0)  # torch.Size([2, 2, 3])
torch.cat([data1, data2], dim=1)  # torch.Size([1, 4, 3])
torch.cat([data1, data2], dim=2)  # torch.Size([1, 2, 6])

8.2 torch.stack()

在新维度上拼接,所有输入张量形状必须完全相同。

python 复制代码
data1 = torch.randint(0, 10, [2, 3])
data2 = torch.randint(0, 10, [2, 3])

torch.stack([data1, data2], dim=0)  # torch.Size([2, 2, 3])
torch.stack([data1, data2], dim=1)  # torch.Size([2, 2, 3])
torch.stack([data1, data2], dim=2)  # torch.Size([2, 3, 2])

9. 自动微分模块

9.1 基本概念

  • PyTorch 不支持向量张量对向量张量求导,只支持 标量对向量 求导。
  • requires_grad=True:自动计算梯度并保存到 grad
  • y.backward():反向传播,y 必须是标量。
  • x.grad:获取梯度,会累加历史梯度。
  • x.grad.zero_():清空梯度。
python 复制代码
x = torch.tensor(10, requires_grad=True, dtype=torch.float32)
y = 2 * x ** 2
y.sum().backward()
print(x.grad)   # tensor(40.)

向量张量:

python 复制代码
x = torch.tensor([10, 20], requires_grad=True, dtype=torch.float32)
y = 2 * x ** 2
y.sum().backward()
print(x.grad)   # tensor([40., 80.])

9.2 梯度下降法求最优解

公式:w = w - r * gradr 是学习率)

python 复制代码
x = torch.tensor(10, requires_grad=True, dtype=torch.float32)
y = x ** 2 + 20

for i in range(1, 1001):
    y = x ** 2 + 20
    if x.grad is not None:
        x.grad.zero_()
    y.backward()
    x.data = x.data - 0.01 * x.grad

注意:不能写成 x = x - 0.01 * x.grad,应使用 x.data = x.data - 0.01 * x.grad

9.3 梯度计算注意点

  • 不能将自动微分的张量直接转成 NumPy 数组,会报错。
  • 使用 detach() 产生一个新张量,共享数据但不自动微分。
python 复制代码
x1 = torch.tensor([10, 20], requires_grad=True, dtype=torch.float64)
# x1.numpy()  # 报错
x2 = x1.detach()
print(x1.requires_grad)  # True
print(x2.requires_grad)  # False
print(x1.data, x2.data)  # 共享内存

9.4 自动微分模块应用

python 复制代码
x = torch.ones(2, 5)
y = torch.zeros(2, 3)

w = torch.randn(5, 3, requires_grad=True)
b = torch.randn(3, requires_grad=True)

z = torch.matmul(x, w) + b
loss = torch.nn.MSELoss()
loss = loss(z, y)
loss.backward()

print(w.grad)
print(b.grad)

10. 线性回归案例

10.1 模型构建流程

  1. 准备训练集数据
  2. 构建模型
  3. 设置损失函数和优化器
  4. 模型训练

10.2 使用 API

组件 替代
nn.MSELoss() 平方损失函数
data.DataLoader 数据加载器
optim.SGD 优化器
nn.Linear 假设函数

10.3 代码实现

python 复制代码
import torch
from torch.utils.data import TensorDataset, DataLoader
from torch import nn, optim
from sklearn.datasets import make_regression
import matplotlib.pyplot as plt

plt.rcParams['font.sans-serif'] = ['SimHei']
plt.rcParams['axes.unicode_minus'] = False


def create_dataset():
    x, y, coef = make_regression(
        n_samples=100, n_features=1, noise=10,
        coef=True, bias=14.5, random_state=0
    )
    x = torch.tensor(x)
    y = torch.tensor(y)
    return x, y, coef


if __name__ == "__main__":
    x, y, coef = create_dataset()

    plt.scatter(x, y)
    x_line = torch.linspace(x.min(), x.max(), 1000)
    y_line = torch.tensor([v * coef + 14.5 for v in x_line])
    plt.plot(x_line, y_line, label='real')
    plt.grid()
    plt.legend()
    plt.show()

    dataset = TensorDataset(x, y)
    dataloader = DataLoader(dataset=dataset, batch_size=16, shuffle=True)

    model = nn.Linear(in_features=1, out_features=1)
    criterion = nn.MSELoss()
    optimizer = optim.SGD(params=model.parameters(), lr=1e-2)

    epochs = 100
    epoch_loss = []
    total_loss = 0.0
    train_sample = 0.0

    for _ in range(epochs):
        for train_x, train_y in dataloader:
            y_pred = model(train_x.type(torch.float32))
            loss = criterion(
                y_pred,
                train_y.reshape(-1, 1).type(torch.float32)
            )
            total_loss += loss.item()
            train_sample += 1

            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

        epoch_loss.append(total_loss / train_sample)

    plt.plot(range(epochs), epoch_loss)
    plt.title('损失变化曲线')
    plt.grid()
    plt.show()

    plt.scatter(x, y)
    x_line = torch.linspace(x.min(), x.max(), 1000)
    y1 = torch.tensor([v * model.weight + model.bias for v in x_line])
    y2 = torch.tensor([v * coef + 14.5 for v in x_line])
    plt.plot(x_line, y1, label='训练')
    plt.plot(x_line, y2, label='真实')
    plt.grid()
    plt.legend()
    plt.show()

11. 总结

  • 张量创建torch.tensortorch.TensorIntTensor/FloatTensor/DoubleTensor、线性/随机张量、0/1/指定值张量。
  • 类型转换type()half/double/float/short/int/long()numpy()from_numpy()item()
  • 数值计算add/sub/mul/div/neg、点乘、矩阵乘法 @ / matmul
  • 运算函数sum/mean/sqrt/pow/exp/log 等。
  • 索引:行列索引、列表索引、范围索引、布尔索引、多维索引。
  • 形状操作reshape/squeeze/unsqueeze/transpose/permute/view/contiguous
  • 拼接cat(不增加维度)、stack(新维度拼接)。
  • 自动微分requires_gradbackward()gradzero_()detach()
  • 线性回归流程:数据准备 → 模型构建 → 损失函数与优化器 → 训练循环。
相关推荐
志尊宝1 小时前
Vue3 零基础每日笔记(016):v-for 列表渲染——key 的作用与为什么不能用 index
javascript·vue.js·笔记
ouynagda1 小时前
51 单片机 LED 流水灯 + 数码管动态扫描学习笔记(STC89C52)
笔记·单片机·学习
zzzll11111 小时前
LLM 学习第 24 课:Agent Harness
前端·人工智能·学习
智者知已应修善业1 小时前
【如何将此图变为000到101就返回】2026-4-7
驱动开发·经验分享·笔记·硬件架构·硬件工程
wixzjsh2 小时前
51单片机学习笔记|最小系统、中断、定时器+数码管动态扫描实现0~9999计数器
笔记·学习·51单片机
平头哥AI2 小时前
Day 21 _ error 是个普通值_errors.New 与 fmt.Errorf 造出来,沿调用栈抛到 main 接住
后端·学习·golang·go
snow@li2 小时前
服务器运维: k3s 安装及配置的完整操作笔记
笔记
ECT-OS-JiuHuaShan2 小时前
哲学是迭代学,数学是拓扑学
开发语言·人工智能·学习·算法·机器学习·php·拓扑学
leo_messi942 小时前
Mysql学习(十二) -- SQL执行到底做了什么事?
sql·学习·mysql