跳转到内容
新建笔记

PyTorch Tensor:类型、形状、索引与共享存储

理解张量操作时,先回答四个问题:形状怎样变化、元素的逻辑顺序有没有改变、是否共享数据、数据类型和设备是否合适。只记函数名,容易把“返回另一个张量对象”误认为“复制了一份独立数据”。

本页保留原笔记中的创建、变形、索引、统计、类型转换与设备检查主题,按数据流重新组织。API 依据 PyTorch 2.8 文档;下面的数值示例可以在 CPU 上运行。

1. 环境、类型与设备分别检查

跳转到“1. 环境、类型与设备分别检查”

PyTorch、TorchVision、Ultralytics 是不同的包。检查张量 API 不要求先安装或导入后两者;cuDNN 信息可以通过 PyTorch 后端查询,不需要假定存在一个可 import cudnn 的包。

import torch
from importlib.metadata import version, PackageNotFoundError
print("PyTorch:", torch.__version__)
print("CUDA build version:", torch.version.cuda)
print("CUDA available:", torch.cuda.is_available())
print("cuDNN version:", torch.backends.cudnn.version())
for package in ("torchvision", "ultralytics"):
try:
print(package, version(package))
except PackageNotFoundError:
print(package, "not installed")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
x = torch.tensor([1.0, 2.0], dtype=torch.float32, device=device)
print("shape/dtype/device:", x.shape, x.dtype, x.device)
print("elements/dimensions:", x.numel(), x.ndim)

torch.version.cuda 表示该 PyTorch 构建对应的 CUDA 版本,不是显卡驱动版本,也不证明此进程能使用 GPU。torch.cuda.is_available() 才用于检查当前运行条件。CPU 构建的 CUDA/cuDNN 版本可能为 None。CUDA 检查、后端信息

2. 创建数据时明确初始内容

跳转到“2. 创建数据时明确初始内容”
用途接口要点
从已有数值创建torch.tensor(data)复制输入数据;类型由输入推断或通过 dtype 指定
全零、全一、常量zeros、ones、full尺寸和数值不同;full((2, 3), 5) 默认推断为整数
预留未初始化内容empty不能把其初始元素当作有效样本;先完整写入
等步长序列arange(start, end, step)通常不包含 end,浮点步长需注意舍入
指定点数的序列linspace(start, end, steps)默认包含两端,steps 是点数
随机数rand、randn、randint分别为均匀、标准正态、整数范围采样
按已有张量构造zeros_like、ones_like、empty_like默认沿用形状、类型与设备,可显式覆盖
复用 NumPy 数据from_numpy在支持的 CPU 数组类型上共享底层内存,需要注意共同修改

需要复现实验时设置随机种子,并记录版本与设备。固定种子不承诺不同后端、设备或软件版本之间逐位相同。张量创建

3. 形状变换不等于轴变换

跳转到“3. 形状变换不等于轴变换”

以 [N, C, H, W] 为例,reshape 改变逻辑元素序列的分组方式;要得到 [N, H, W, C] 的轴含义,使用 permute(0, 2, 3, 1)。仅把四个形状数字改顺序,并不会自动把通道移动到正确位置。

操作形状含义数据关系
view(shape)按兼容的步幅解释同一数据,元素总数不变返回视图;步幅不兼容时失败
reshape(shape)保持逻辑元素顺序,重新分组可能返回视图,也可能复制,不依赖哪一种
flatten(start_dim, end_dim)合并指定范围内的轴可能返回原对象、视图或复制
permute、transpose交换轴的位置;t() 主要用于不超过二维的情形对普通稠密张量返回视图
squeeze(dim)、unsqueeze(dim)删除大小为 1 的指定轴、增加大小为 1 的轴返回视图
contiguous()获得指定内存格式的连续表示已满足要求可返回自身,否则复制
split、chunk沿一个轴拆分返回包含视图的元组,不能当作独立副本
cat、stack沿已有轴连接、沿新轴堆叠通常生成新的输出数据
repeat按各轴重复内容复制数据
expand扩展大小为 1 的轴或增加前导轴使用步幅为 0 的视图,不物理重复元素

reshape 并没有被禁止或不推荐;它适合“只关心结果形状”的场景。需要独立数据时明确 clone();需要明确共享且能满足步幅约束时使用 view()。不要把 is_contiguous() 当作所有 view 成败的唯一条件,部分非连续张量仍可形成兼容视图。视图规则、view 步幅条件

无参数 squeeze() 会移除所有大小为 1 的轴,包括大小恰好为 1 的批次轴。chunk(k) 也不保证总能返回恰好 k 个等长块;需要固定段数时检查 tensor_split 的语义。squeeze、chunk

4. 用修改行为理解共享关系

跳转到“4. 用修改行为理解共享关系”
import torch
base = torch.arange(12, dtype=torch.float64).reshape(3, 4)
reshaped_view = base.view(2, 6)
reshaped_view[0, 0] = -1
assert base[0, 0].item() == -1
transposed = base.transpose(0, 1)
assert transposed.shape == (4, 3) and not transposed.is_contiguous()
# 这个具体转置的步幅不能直接合并成一个轴。
try:
transposed.view(-1)
except RuntimeError:
pass
else:
raise AssertionError("expected an incompatible-stride view to fail")
flattened = transposed.reshape(-1)
old_value = base[0, 0].item()
flattened[0] = 123
assert base[0, 0].item() == old_value # 本例reshape实际复制
pieces = base.split(2, dim=0)
assert [tuple(part.shape) for part in pieces] == [(2, 4), (1, 4)]
pieces[0][0, 1] = -2
assert base[0, 1].item() == -2
assert len(torch.arange(13).chunk(6)) == 5
image = torch.arange(16).reshape(4, 4)
expanded = image.expand(3, -1, -1)
repeated = image.repeat(3, 1, 1)
assert expanded.shape == repeated.shape == (3, 4, 4)
assert expanded.stride()[0] == 0
image[0, 0] = 99
assert expanded[2, 0, 0].item() == 99
assert repeated[2, 0, 0].item() == 0
independent = expanded.clone()
independent[0, 0, 0] = -5
assert image[0, 0].item() == 99
batch = torch.zeros(1, 3, 1, 4)
assert batch.squeeze().shape == (3, 4)
assert batch.squeeze(2).shape == (1, 3, 4)
assert batch.flatten(1).shape == (1, 12)
print("形状、视图和独立副本检查通过")

扩展视图可能有多个逻辑位置指向同一内存。若要分别写入各份内容,先克隆;不要在这类重叠视图上依赖逐元素原位写入。上例用整数或不需要梯度的张量说明修改关系,不能由此推断自动求导图中的任意原位操作都合法。

5. 索引、连接与矩阵运算

跳转到“5. 索引、连接与矩阵运算”

基本切片通常共享数据;整数数组和布尔掩码等高级索引读取通常产生副本,但通过索引赋值会修改被赋值的原张量。index_select 用一维索引选取整个切片;gather 的索引与输入具有相同维数,可以为各位置选不同元素。索引规则、gather

import torch
x = torch.arange(12, dtype=torch.float64).reshape(3, 4)
selected_rows = x.index_select(0, torch.tensor([2, 0]))
assert selected_rows.shape == (2, 4)
indices = torch.tensor([[0, 2], [1, 3], [3, 0]])
selected_values = x.gather(1, indices)
torch.testing.assert_close(selected_values,
torch.tensor([[0., 2.], [5., 7.], [11., 8.]], dtype=x.dtype))
masked = torch.masked_select(x, x > 8)
assert masked.tolist() == [9.0, 10.0, 11.0]
# scatter_是Tensor的方法;目标、src的dtype一致,索引为long。
destination = torch.zeros(3, 3)
scatter_index = torch.tensor([[0, 1], [2, 0]], dtype=torch.long)
source = torch.tensor([[1., 2.], [3., 4.]])
destination.scatter_(0, scatter_index, source)
torch.testing.assert_close(destination,
torch.tensor([[1., 4., 0.], [0., 2., 0.], [3., 0., 0.]]))
assert torch.cat([x, x], dim=0).shape == (6, 4)
assert torch.stack([x, x], dim=0).shape == (2, 3, 4)
a = torch.arange(6, dtype=torch.float64).reshape(2, 3)
b = torch.arange(12, dtype=torch.float64).reshape(3, 4)
torch.testing.assert_close(a @ b, torch.mm(a, b))
assert torch.matmul(a[None], b[None]).shape == (1, 2, 4)
assert torch.bmm(a[None], b[None]).shape == (1, 2, 4)
assert torch.outer(torch.tensor([1., 2.]), torch.tensor([3., 4.])).shape == (2, 2)
print("索引、散布、连接和矩阵运算检查通过")

* 是逐元素乘法,mm 是二维矩阵乘法,bmm 对应相同批次数的三维批量矩阵乘法,matmul 支持更多维数并在批次维广播。广播按尾部维度对齐,不能用它代替对形状含义的检查。

原来的 torch.scatter_(...) 应改为 tensor.scatter_(...);函数形式 torch.scatter(...) 返回结果。重复目标索引还涉及覆盖顺序和梯度问题,上例刻意使用无冲突索引。scatter_

6. 统计、逐元素运算与返回值

跳转到“6. 统计、逐元素运算与返回值”
操作族例子需要确认
聚合sum、mean、max、min是否指定 dim、是否保留轴;整数张量求均值应先明确转换类型
累积cumsum(dim=0)、cumprod(dim=0)必须指定累计的轴,原来无 dim 的调用不完整
方差、标准差var、std默认 correction=1;总体统计使用 correction=0,样本数不足时结果可能非有限
排序、前几个sort、topk同时返回数值与索引;并列值不应假定索引顺序稳定
众数、中位数mode、medianmode 返回值和索引;偶数个元素的 median 取中间偏小值,不是两者平均
范数linalg.vector_norm、linalg.matrix_norm明确向量/矩阵、轴和范数类型;新代码不依赖已弃用的 torch.norm
数学函数exp、log、sqrt、sin、cos、abs、tanh、sigmoid通常逐元素作用;exp 不是矩阵指数
比较、逻辑与位运算eq/ne/lt/le/gt/ge、logical_and/or/not、bitwise_and/or/not区分布尔条件与整数位模式,类型支持不同
对角与三角区域diag、tril、triudiag 对一维输入构造矩阵、对二维输入提取对角线
分布采样multinomial输入可以是未归一化非负权重,要求总和为正;是否放回需明确

all()、any() 汇总布尔条件,numel() 返回元素总数。数学函数的定义域也要检查,例如负实数的 sqrt 不会自动把实数张量改成复数。方差、中位数、向量范数

7. 类型转换、自动求导与 NumPy

跳转到“7. 类型转换、自动求导与 NumPy”

常见类型包括 float32/float、float64/double、float16/half、bfloat16,int8、int16/short、int32/int、int64/long,以及 uint8、bool、complex64、complex128。uint16/uint32/uint64 等类型的存在不代表全部算子都支持它们,应按当前版本与后端核对。类型支持

使用 x.to(dtype=torch.float64, device="cpu") 明确类型与设备。若要求本来就满足,to 可以返回原对象;需要强制副本可传 copy=True。.float()、.double()、.long() 是常用简写。PyTorch 2.8 的普通 Tensor 没有 .cast() 这个公开转换接口;旧 .type() 还允许类型名或 Tensor 类,不能概括为“在任何写法下都绝不会改变设备”。to、type

clone() 复制数据但保留可微关系;detach() 切断自动求导关系却仍共享数据。需要两者同时满足时用 detach().clone()。转换成整数会使结果不再具备浮点自动求导的含义。

import numpy as np
import torch
array = np.array([1.0, 2.0, 3.0], dtype=np.float32)
shared = torch.from_numpy(array)
copied = torch.tensor(array)
array[0] = 9.0
assert shared[0].item() == 9.0 and copied[0].item() == 1.0
shared[1] = 8.0
assert array[1] == 8.0
x = torch.tensor([1.0, 2.0], requires_grad=True)
clone = x.clone()
clone.sum().backward()
torch.testing.assert_close(x.grad, torch.ones_like(x))
detached = x.detach()
assert not detached.requires_grad and detached.data_ptr() == x.data_ptr()
independent = x.detach().clone()
independent[0] = -1
assert x[0].item() == 1.0
assert x.to(dtype=x.dtype, device=x.device) is x
assert x.to(copy=True).data_ptr() != x.data_ptr()
assert not x.to(torch.int64).requires_grad
# 明确移动到CPU并脱离计算图,再导出NumPy;它仍可能共享CPU数据。
exported = x.detach().cpu().numpy()
independent_array = exported.copy()
assert independent_array.tolist() == [1.0, 2.0]
print("类型、自动求导和NumPy共享检查通过")

只把只读 NumPy 数组传给 from_numpy 并不能获得可安全写入的独立张量;有独立写入需求时先复制。普通张量与 NumPy 的共享也不自动提供线程同步。from_numpy

进一步把张量连接成模型,见torch.nn;把输出连接到监督目标,见损失函数。