学深度学习第一步不是模型,而是数据怎么在框架里流动。在 PyTorch 里,几乎所有数据都会被表示成 Tensor(张量)。这一节我按“能写代码就算会”的标准,把最常用的数据操作过一遍。

环境:Python + PyTorch(CPU/GPU 都可)


1. Tensor 是什么:多维数组 + 可上 GPU + 支持自动求导

Tensor 本质上就是一个多维数组,但它比 numpy.ndarray 更适合深度学习:

  • 可以放到 GPU(.to("cuda"))

  • 可以做自动求导(后面反向传播用)

  • 运算 API 更适合神经网络


2. 创建张量:最常用的 5 种方式

import torch

# 1) 直接创建(Python 列表 -> Tensor)
x = torch.tensor([1.0, 2.0, 3.0])

# 2) 指定形状,生成全 0 / 全 1
a = torch.zeros(2, 3)
b = torch.ones(2, 3)

# 3) 生成随机数(均匀分布)
c = torch.rand(2, 3)

# 4) 正态分布随机数
d = torch.randn(2, 3)

# 5) 等差序列
e = torch.arange(0, 10, 2)  # 0,2,4,6,8

经验:

  • 写网络时最常见的是 zeros/ones/randn

  • Debug 数据范围时常用 arange


3. 形状与维度:shape / numel / reshape

X = torch.arange(12)
print(X.shape)     # torch.Size([12])
print(X.numel())   # 12

X = X.reshape(3, 4)
print(X)
print(X.shape)     # torch.Size([3, 4])

关键点:

  • reshape 不改数据本身,只是换“视图/结构”(多数情况下是 O(1))

  • numel() 是元素总数(排错非常好用)


4. 索引与切片:和 Python/Numpy 很像

X = torch.arange(12).reshape(3, 4)

print(X[0])        # 第 0 行
print(X[:, 1])     # 第 1 列
print(X[1, 2])     # 第 1 行第 2 列

print(X[0:2, :])   # 前两行

修改某一块数据:

X[0:2, :] = 100
print(X)

5. 拼接与切分:cat / stack / split(非常常用)

cat:沿某个维度拼接(维度不增加)

A = torch.ones(2, 3)
B = torch.zeros(2, 3)

C0 = torch.cat((A, B), dim=0)  # 行方向拼
C1 = torch.cat((A, B), dim=1)  # 列方向拼

stack:拼起来并新增一个维度

S = torch.stack((A, B), dim=0)  # shape: (2, 2, 3)

记忆法:

  • cat:拼在一起,维度不变

  • stack:摞起来,维度 +1


6. 逐元素运算:加减乘除、指数、比较

X = torch.tensor([1.0, 2.0, 4.0, 8.0])
Y = torch.tensor([2.0, 2.0, 2.0, 2.0])

print(X + Y)
print(X - Y)
print(X * Y)
print(X / Y)
print(X ** Y)      # 幂

比较与生成 mask:

print(X == Y)
print(X > Y)

7. 广播机制(Broadcasting):写深度学习一定会遇到

广播解决的问题:不同形状的张量怎么做运算?

A = torch.arange(3).reshape(3, 1)  # (3,1)
B = torch.arange(2).reshape(1, 2)  # (1,2)

print(A + B)  # (3,2)

解释一下:

  • A 会被“复制扩展”成 (3,2)

  • B 也会被“复制扩展”成 (3,2)

  • 然后逐元素相加

本质规则:从后往前对齐维度:

  • 维度相等 ✅

  • 其中一个维度是 1 ✅(可广播)

  • 否则 ❌ 报错


8. 转 numpy / 转 Python 数值:调试时常用

X = torch.arange(6).reshape(2, 3)

# Tensor -> numpy
npX = X.numpy()

# 只有一个元素的 Tensor -> Python 标量
s = torch.tensor([3.14]).item()


9. 保存与加载:模型训练离不开

保存 tensor:

X = torch.arange(12).reshape(3, 4)
torch.save(X, "X.pt")

加载:

Y = torch.load("X.pt")
print(Y)

(后面保存模型参数也用同样思路)


10. 本节最容易踩坑的点(我自己的总结)

  1. reshape 维度要对:元素总数必须一致,否则直接报错

  2. cat 和 stack 容易混:一个不增维,一个增维

  3. 广播不是“随便加”,是严格规则:不满足会报 shape mismatch

  4. .item() 只能用于单元素 tensor


结语:为什么“数据操作”值得单独学?

因为你后面会不断遇到这些问题:

  • 我的 batch 维是第几维?

  • 这个张量怎么从 (N, C, H, W) 变到 (N, H, W, C)?

  • 为什么 loss 能算,但反向传播报维度错?

  • 为什么加法能跑(广播),但 cat 报错?

更多推荐