评价此页

Nested Tensors 入门#

Nested Tensor(嵌套张量)概括了常规稠密张量的形状,允许表示参差不齐(ragged-sized)的数据。

  • 对于常规张量,每个维度都是规则的并具有固定大小

  • 对于嵌套张量,并非所有维度的大小都是规则的;其中一些维度是参差不齐的

嵌套张量是表示各种领域内序列数据的自然解决方案

  • 在 NLP 中,句子可以具有可变长度,因此一批句子可以形成一个嵌套张量

  • 在 CV 中,图像可以具有可变形状,因此一批图像可以形成一个嵌套张量

在本教程中,我们将演示嵌套张量的基本用法,并通过一个实际案例说明它们在处理不同长度序列数据时的实用性。特别是,它们对于构建能够高效处理参差不齐序列输入的 Transformer 模型非常有价值。下面,我们展示了使用嵌套张量实现多头注意力机制的方法,结合 torch.compile 使用,其性能优于在带有填充(padding)的张量上进行简单操作的效果。

嵌套张量目前是一个原型功能,未来可能会发生变化。

嵌套张量初始化#

在 Python 前端,可以通过张量列表创建嵌套张量。我们将 nt[i] 表示为 nestedtensor 的第 i 个张量分量。

通过将所有底层张量填充为相同形状,nestedtensor 可以转换为常规张量。

所有张量都具有一个属性,用于确定它们是否为嵌套张量;

通常从形状不规则的张量批次构建 nestedtensor。即假设维度 0 是批处理维度。索引维度 0 会返回第一个底层张量分量。

# When indexing a nestedtensor's 0th dimension, the result is a regular tensor.

需要注意的一点是,维度 0 的切片操作目前尚不支持。这意味着目前无法构建一个合并了底层张量分量的视图(view)。

嵌套张量操作#

由于每个操作都必须为 nestedtensor 显式实现,因此目前 nestedtensor 的操作覆盖范围比常规张量窄。目前仅涵盖索引、dropout、softmax、transpose、reshape、linear、bmm 等基本操作。不过,覆盖范围正在扩展。如果您需要特定的操作,请提交 issue 以帮助我们确定优先级。

reshape

reshape 操作用于改变张量的形状。其针对常规张量的完整语义可以在此处找到。对于常规张量,在指定新形状时,单个维度可以是 -1,在这种情况下,它会从剩余维度和元素数量中推断出来。

nestedtensor 的语义类似,只是 -1 不再进行推断。相反,它继承了旧的大小(此处 nt[0] 为 2,nt[1] 为 3)。对于参差不齐的维度,-1 是唯一合法的指定大小。

转置

transpose 操作用于交换张量的两个维度。其完整语义可以在此处找到。请注意,对于 nestedtensor,维度 0 是特殊的;它被假定为批处理维度,因此不支持涉及 nestedtensor 维度 0 的转置操作。

其他

其他操作与常规张量具有相同的语义。将操作应用于 nestedtensor 等同于将操作应用于底层的张量分量,结果也是一个 nestedtensor。

为什么选择嵌套张量#

当数据是序列时,通常每个样本的长度不同。例如,在一批句子中,每个句子包含的单词数量不同。处理可变序列的一种常用技术是手动将每个数据张量填充(pad)为相同的形状以形成批次。例如,我们有 2 个长度不同的句子和一个词汇表。为了将其表示为单个张量,我们用 0 填充到该批次中的最大长度。

这种将一批数据填充到最大长度的技术并不理想。填充数据对于计算而言是不必要的,并且通过分配比实际需要更大的张量来浪费内存。此外,并非所有操作在应用于填充数据时都具有相同的语义。对于矩阵乘法,为了忽略填充项,需要填充为 0;而对于 softmax,必须填充为负无穷(-inf)以忽略特定条目。嵌套张量的主要目标是使用标准 PyTorch 张量 UX 促进对参差不齐数据的操作,从而消除低效且复杂的填充和掩码需求。

让我们看一个实际的例子:Transformers 中使用的多头注意力组件。我们可以以一种能够处理填充张量或嵌套张量的方式来实现它。

按照 Transformer 论文设置超参数

除了 dropout 概率外:设置为 0 以进行正确性检查

让我们根据齐普夫定律(Zipf’s law)生成一些真实的模拟数据。

创建嵌套张量批次输入

生成 query、key、value 的填充形式以便比较

构建模型

检查正确性和性能

# padding-specific step: remove output projection bias from padded entries for fair comparison








# warm up compile first...


# ...now benchmark



# warm up compile first...

# ...now benchmark



# padding-specific step: remove output projection bias from padded entries for fair comparison

请注意,如果没有 torch.compile,Python 子类嵌套张量的开销可能使其比在填充张量上执行相同的计算更慢。然而,一旦启用了 torch.compile,在嵌套张量上操作会带来多倍的加速。随着批次中填充比例的增加,避免填充带来的无效计算变得愈发重要。

结论#

在本教程中,我们学习了如何使用嵌套张量执行基本操作,以及如何以避免填充计算的方式实现 Transformer 的多头注意力。有关更多信息,请查看 torch.nested 命名空间的文档。

另请参阅#

# %%%%%%RUNNABLE_CODE_REMOVED%%%%%%

脚本总运行时间:(0 分 0.002 秒)