评价此页

torch.vsplit#

torch.vsplit(input, indices_or_sections) List of Tensors#

根据 indices_or_sections 将输入 input(一个至少有两维的张量)垂直分割成多个张量。每个分割都是 input 的视图。

这等效于调用 torch.tensor_split(input, indices_or_sections, dim=0)(分割维度为 0),唯一的区别是,如果 indices_or_sections 是一个整数,则它必须能整除分割维度,否则将抛出运行时错误。

此函数基于 NumPy 的 numpy.vsplit()

参数

示例

>>> t = torch.arange(16.0).reshape(4,4)
>>> t
tensor([[ 0.,  1.,  2.,  3.],
        [ 4.,  5.,  6.,  7.],
        [ 8.,  9., 10., 11.],
        [12., 13., 14., 15.]])
>>> torch.vsplit(t, 2)
(tensor([[0., 1., 2., 3.],
         [4., 5., 6., 7.]]),
 tensor([[ 8.,  9., 10., 11.],
         [12., 13., 14., 15.]]))
>>> torch.vsplit(t, [3, 6])
(tensor([[ 0.,  1.,  2.,  3.],
         [ 4.,  5.,  6.,  7.],
         [ 8.,  9., 10., 11.]]),
 tensor([[12., 13., 14., 15.]]),
 tensor([], size=(0, 4)))