maxframe.tensor.dsplit#

maxframe.tensor.dsplit(a, indices_or_sections)[源代码]#

沿第3轴(深度)将张量分割为多个子张量。

请参考 split 文档。dsplit 等价于 axis=2split,只要张量维度大于或等于3,数组总是沿第三个轴进行分割。

参见

split

将张量分割为多个大小相等的子数组。

示例

>>> import maxframe.tensor as mt
>>> x = mt.arange(16.0).reshape(2, 2, 4)
>>> x.execute()
array([[[  0.,   1.,   2.,   3.],
        [  4.,   5.,   6.,   7.]],
       [[  8.,   9.,  10.,  11.],
        [ 12.,  13.,  14.,  15.]]])
>>> mt.dsplit(x, 2).execute()
[array([[[  0.,   1.],
        [  4.,   5.]],
       [[  8.,   9.],
        [ 12.,  13.]]]),
 array([[[  2.,   3.],
        [  6.,   7.]],
       [[ 10.,  11.],
        [ 14.,  15.]]])]
>>> mt.dsplit(x, mt.array([3, 6])).execute()
[array([[[  0.,   1.,   2.],
        [  4.,   5.,   6.]],
       [[  8.,   9.,  10.],
        [ 12.,  13.,  14.]]]),
 array([[[  3.],
        [  7.]],
       [[ 11.],
        [ 15.]]]),
 array([], dtype=float64)]