PyTorch 张量操作基础
FreeGuideOnline
最新
2026-07-08
python import torch import numpy as np
从列表创建
a = torch.tensor([1, 2, 3]) b = torch.tensor([[1, 2], [3, 4]])
从 NumPy 数组创建
np_arr = np.array([1, 2, 3]) c = torch.from_numpy(np_arr) # 与 np_arr 共享内存
### 创建特定形状的张量
```python
# 全零张量
zeros = torch.zeros(2, 3)
# 全一张量
ones = torch.ones(2, 3)
# 指定值的张量
full = torch.full((2, 3), 3.14)
# 未初始化的张量(内存中遗留的值)
empty = torch.empty(2, 3)
# 创建与已有张量形状相同的张量
like = torch.ones_like(b) # 全一
like_zeros = torch.zeros_like(b)
生成序列和随机张量
# 等差数列
seq = torch.arange(0, 10, step=2) # tensor([0,2,4,6,8])
# 线性间隔
lin = torch.linspace(0, 1, steps=5) # 0 到 1 之间 5 个点
# 随机张量
rand = torch.rand(2, 3) # 均匀分布 [0,1)
randn = torch.randn(2, 3) # 标准正态分布
randint = torch.randint(0, 10, (2,3)) # [0,10) 任意整数
单位矩阵和对角矩阵
# 单位矩阵
I = torch.eye(3) # 3×3 单位矩阵
# 对角线元素为给定值的矩阵
diag = torch.diag(torch.tensor([1, 2, 3]))
张量的属性
创建张量后,可以通过属性快速获取其基本信息。
t = torch.randn(3, 4, 5)
print(t.shape) # torch.Size([3, 4, 5])
print(t.size()) # 同 shape,返回 torch.Size
print(t.size(0)) # 获取第 0 维的大小,即 3
print(t.dtype) # torch.float32
print(t.device) # cpu
print(t.ndim) # 3(维度数)
shape/size():张量的形状dtype:数据类型,如torch.float32,torch.int64device:张量所在设备(CPU 或 GPU)requires_grad:是否需要自动求导(默认 False)
基本张量操作
索引与切片
张量的索引与 NumPy 风格一致,支持多维索引、切片和高级索引。
x = torch.arange(12).reshape(3, 4)
print(x[0]) # 第一行
print(x[:, 1]) # 第二列
print(x[1, 2]) # 第二行第三列元素
# 切片
print(x[0:2, 1:3]) # 前两行,第2、3列
# 高级索引(掩码)
mask = x > 5
print(x[mask]) # 返回 x 中大于 5 的所有元素
连接张量
将多个张量沿指定维度拼接。
a = torch.ones(2, 3)
b = torch.zeros(2, 3)
# 在第 0 维(行)拼接
cat0 = torch.cat((a, b), dim=0) # 形状 (4,3)
# 在第 1 维(列)拼接
cat1 = torch.cat((a, b), dim=1) # 形状 (2,6)
# 堆叠(扩展出新维度)
stacked = torch.stack((a, b), dim=0) # 形状 (2,2,3)
拆分张量
x = torch.arange(6).reshape(2, 3)
# 按块大小拆分
chunks = torch.split(x, 2, dim=1) # 每块 2 列
# 或指定每块尺寸
splits = torch.split(x, [1, 2], dim=1) # 第一块 1 列,第二块 2 列
形状变换
view 和 reshape
两者都可以改变张量形状,但 view 要求张量在内存中连续,否则会报错;reshape 在不连续时会复制一份。
x = torch.randn(2, 3, 4)
# 使用 view
y = x.view(2, 12) # 扁平化为 2×12
z = x.view(-1, 4) # -1 表示该维度自动推断,结果为 6×4
# 使用 reshape
y2 = x.reshape(2, 12)
转置与维度交换
x = torch.randn(2, 3, 4)
# 转置(仅限二维)
a = torch.randn(3, 4)
a_t = a.T # 形状 (4,3)
# 交换两个维度
b = x.transpose(0, 2) # 交换 0 和 2 维度,形状 (4,3,2)
# 任意维度重排
c = x.permute(2, 0, 1) # 将维度按指定顺序重新排列,形状 (4,2,3)
增加或压缩维度
x = torch.tensor([1, 2, 3]) # 形状 (3)
# 增加维度
a = x.unsqueeze(0) # 形状 (1,3)
b = x.unsqueeze(1) # 形状 (3,1)
# 移除大小为1的维度
c = b.squeeze() # 形状 (3)
d = b.squeeze(0) # 若 dim0 为1则移除,否则原样返回
数学运算
张量支持丰富的数学运算,包括逐元素运算、矩阵乘法和广播。
逐元素运算
a = torch.tensor([1, 2, 3])
b = torch.tensor([4, 5, 6])
add = a + b # 加法
sub = a - b # 减法
mul = a * b # 逐元素乘法
div = a / b # 除法
pow = a ** 2 # 幂
# 三角函数、指数对数等
sin = torch.sin(a)
exp = torch.exp(a)
log = torch.log(a.float()) # 必须浮点类型
矩阵乘法
x = torch.randn(3, 4)
y = torch.randn(4, 5)
# 多种矩阵乘法写法
z1 = x.matmul(y)
z2 = x @ y # 运算符重载
z3 = torch.mm(x, y) # 仅二维
# 批量矩阵乘法
batch1 = torch.randn(10, 3, 4)
batch2 = torch.randn(10, 4, 5)
z4 = torch.bmm(batch1, batch2) # 或 batch1 @ batch2
广播机制
形状不同的张量在某些情况下可以自动扩展维度进行计算,规则为从后向前对齐形状,任一维度相等或为 1 即可。
a = torch.ones(3, 1)
b = torch.ones(1, 4)
c = a + b # 形状 (3,4) 自动广播
统计与聚合
张量可以沿指定维度计算均值、和、最大值等。
x = torch.arange(12, dtype=torch.float32).reshape(3, 4)
print(x.sum()) # 所有元素和
print(x.sum(dim=0)) # 按列求和,形状 (4)
print(x.sum(dim=1)) # 按行求和,形状 (3)
print(x.mean())
print(x.mean(dim=1))
print(x.max()) # 最大值
print(x.max(dim=1)) # 返回 (values, indices)
print(x.argmax(dim=1)) # 最大值的索引
张量的设备转移
张量可以在 CPU 和 GPU 之间移动,以便利用 GPU 加速。
# 创建 CPU 张量
cpu_tensor = torch.tensor([1, 2, 3])
# 移至 GPU(如果可用)
if torch.cuda.is_available():
gpu_tensor = cpu_tensor.to('cuda')
# 或
gpu_tensor = cpu_tensor.cuda()
# 从 GPU 移回 CPU
cpu_tensor_again = gpu_tensor.cpu()
# 注意:对于需要 tracking 的梯度,应使用 .detach() 避免计算图错误
与 NumPy 的互操作
PyTorch 张量和 NumPy 数组可以高效互转,但需注意共享内存问题。
a = torch.ones(3, 4)
b = a.numpy() # 转换为 NumPy 数组,共享内存
c = torch.from_numpy(b) # NumPy 转 Tensor,共享内存
# 若不想共享,使用克隆
b = a.clone().numpy()
c = torch.from_numpy(b.copy())
自动求导简介
张量可以通过设置 requires_grad=True 来跟踪所有操作,以便自动计算梯度。
x = torch.tensor(2.0, requires_grad=True)
y = x ** 2 + 3 * x + 1
y.backward() # 计算梯度
print(x.grad) # dy/dx = 2*x + 3 = 7