torch.Tensor.masked_scatter_
-
Tensor.masked_scatter_(mask, source) -
Копирует элементы из
sourceв тензорselfв позициях, гдеmaskимеет значение True. Элементы изsourceкопируются вselfначиная с позиции 0 вsourceи продолжаются в порядке один за другим для каждого случая, когдаmaskравно True. Формаmaskдолжна быть совместима с трансляцией с формой базового тензора.sourceдолжен иметь не меньше элементов, чем количество единиц вmask.- Параметры
-
- mask (BoolTensor) – булевый массив
- source (Tensor) – тензор для копирования
Примечание
Данная операция
maskвыполняется над тензоромself, а не над переданным тензоромsource.Пример
>>> self = torch.tensor([[0, 0, 0, 0, 0], [0, 0, 0, 0, 0]]) >>> mask = torch.tensor([[0, 0, 0, 1, 1], [1, 1, 0, 1, 1]]) >>> source = torch.tensor([[0, 1, 2, 3, 4], [5, 6, 7, 8, 9]]) >>> self.masked_scatter_(mask, source) tensor([[0, 0, 0, 0, 1], [2, 3, 0, 4, 5]])
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.Tensor.masked_scatter_.html