asc.experimental.asctile.scatter

asc.experimental.asctile.scatter(src: LocalTensor, dim: int, index: LocalTensor, dst: GlobalTensor, offsets: Iterable[PlainValue | int], check_bounds: bool = True, num_indices: PlainValue | int | None = None) → None

Scatter subtensors from a local tensor into a global tensor at positions given by an index tensor.

For each index value index[i], the subtensor src[i] is written to dst at position [offsets[0], ..., offsets[dim-1], offsets[dim] + index[i]]. The written subtensor spans dimensions dim+1, ..., dst.rank-1 of dst.

The index tensor must have an integer dtype (int8, int16, int32, int64). src and dst must have the same dtype.

Parameters:
  • src – The source local tensor with the data to write.

  • dim – The dimension of dst used for indexing.

  • index – The index tensor. Must be a rank-1 tensor with an integer dtype.

  • dst – The destination global tensor.

  • offsets – The offsets into dst for dimensions 0..dim. Must contain dim + 1 values.

  • check_bounds – If True, out-of-bounds indices are skipped during writes. If False, no bounds checking is performed and the caller must guarantee all indices are valid. Default is True.

  • num_indices – The number of indices in index to process. If None, all elements of index are processed.

Raises:
  • TypeError – If src is not a LocalTensor, index is not a LocalTensor, or dst is not a GlobalTensor

  • RuntimeError – If index does not have an integer dtype

  • ValueError – If offsets does not contain dim + 1 values; dim is out of range for dst.rank; src has an unexpected rank; index is not rank 1; src and dst have different dtypes; src and dst shapes are incompatible; or dst has a dynamic dimension after dim

  • NotImplementedError – If dim is the last dimension of dst

Note

Dimensions dim+1, ..., dst.rank-1 of dst must be static.

Examples

Update full rows by the outermost dimension. result_gm has shape [1024, 128]:

index = asctile.copy_in(index_gm, [0], [256])
data = asctile.copy_in(changes_gm, [0, 0], [256, 128])
asctile.scatter(data, 0, index, result_gm, [0])

Write every other row. If the inputs have the following contents:

dst = [[0, 1, 2, 3, 4, 5, 6, 7],
       [8, 9, 10, 11, 12, 13, 14, 15],
       [16, 17, 18, 19, 20, 21, 22, 23],
       [24, 25, 26, 27, 28, 29, 30, 31],
       ...,
       [248, 249, 250, 251, 252, 253, 254, 255]]  # shape [32, 8]
index = [0, 2, 4, 6, ..., 30]
src = [[0, 0, 0, 0, 0, 0, 0, 0],
       [1, 1, 1, 1, 1, 1, 1, 1],
       ...,
       [15, 15, 15, 15, 15, 15, 15, 15]]  # shape [16, 8]

Then asctile.scatter(src, 0, index, dst, [0]) modifies dst to:

dst = [[0, 0, 0, 0, 0, 0, 0, 0],
       [8, 9, 10, 11, 12, 13, 14, 15],
       [1, 1, 1, 1, 1, 1, 1, 1],
       ...,
       [15, 15, 15, 15, 15, 15, 15, 15],
       [248, 249, 250, 251, 252, 253, 254, 255]]  # shape [32, 8]