當前位置: 首頁>>代碼示例 >>用法及示例精選 >>正文


Python PyTorch Tensor.index_copy_用法及代碼示例


本文簡要介紹python語言中 torch.Tensor.index_copy_ 的用法。

用法:

Tensor.index_copy_(dim, index, tensor) → Tensor

參數

  • dim(int) -索引的維度

  • index(LongTensor) - tensor 的索引可供選擇

  • tensor(Tensor) -包含要複製的值的張量

通過按 index 中給出的順序選擇索引,將 tensor 的元素複製到 self 張量中。例如,如果 dim == 0index[i] == j ,則將 tensor 的第 i 行複製到 self 的第 j 行。

tensor 的第 dim 維度必須與 index 的長度(必須是向量)具有相同的大小,並且所有其他維度必須匹配 self ,否則將引發錯誤。

注意

如果 index 包含重複條目,則 tensor 中的多個元素將被複製到 self 的同一索引中。結果是不確定的,因為它取決於最後出現的副本。

例子:

>>> x = torch.zeros(5, 3)
>>> t = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=torch.float)
>>> index = torch.tensor([0, 4, 2])
>>> x.index_copy_(0, index, t)
tensor([[ 1.,  2.,  3.],
        [ 0.,  0.,  0.],
        [ 7.,  8.,  9.],
        [ 0.,  0.,  0.],
        [ 4.,  5.,  6.]])

相關用法


注:本文由純淨天空篩選整理自pytorch.org大神的英文原創作品 torch.Tensor.index_copy_。非經特殊聲明,原始代碼版權歸原作者所有,本譯文未經允許或授權,請勿轉載或複製。