列表范围索引中的三个点(...)是什么含义?NumPy索引语法疑问
NumPy索引中
...的含义 你看到的三个点是Python内置的Ellipsis对象,在NumPy(以及PyTorch、TensorFlow等支持多维数组的库)的索引语法中,作用是自动补全当前位置到数组末尾所有未显式指定的维度的全选操作,等价于写多个连续的:,不需要手动计算数组维度数量。
用法示例
假设你有一个4维数组arr,形状为(2, 3, 4, 5):
arr[1, ...]等价于arr[1, :, :, :]:选中第一个维度下标为1的元素,剩下3个维度全部选中arr[..., 2]等价于arr[:, :, :, 2]:选中最后一个维度下标为2的元素,前面3个维度全部选中arr[1, ..., 2]等价于arr[1, :, :, 2]:第一个维度选下标1,最后一个维度选下标2,中间2个维度全部选中
对应示例代码解释
pos_encoding = angle_rads[np.newaxis, ...]
这段代码的作用是在angle_rads数组的最前面插入一个新的维度,原数组的所有维度全部保持原样选中。比如angle_rads原本的形状是(序列长度, 隐藏层维度),操作后pos_encoding的形状会变为(1, 序列长度, 隐藏层维度),这类操作常见于Transformer位置编码的实现中,用来适配批量计算时的维度广播规则。
注意:该语法是多维数组的专属索引语法,原生Python列表不支持这种写法。
内容的提问来源于stack exchange,提问作者marlon
相关产品推荐
相关产品推荐

