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 subtensorsrc[i]is written todstat position[offsets[0], ..., offsets[dim-1], offsets[dim] + index[i]]. The written subtensor spans dimensionsdim+1, ..., dst.rank-1ofdst.The
indextensor must have an integer dtype (int8,int16,int32,int64).srcanddstmust have the same dtype.- Parameters:
src – The source local tensor with the data to write.
dim – The dimension of
dstused 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
dstfor dimensions0..dim. Must containdim + 1values.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
indexto process. If None, all elements ofindexare processed.
- Raises:
TypeError – If
srcis not a LocalTensor,indexis not a LocalTensor, ordstis not a GlobalTensorRuntimeError – If
indexdoes not have an integer dtypeValueError – If
offsetsdoes not containdim + 1values;dimis out of range fordst.rank;srchas an unexpected rank;indexis not rank 1;srcanddsthave different dtypes;srcanddstshapes are incompatible; ordsthas a dynamic dimension afterdimNotImplementedError – If
dimis the last dimension ofdst
Note
Dimensions
dim+1, ..., dst.rank-1ofdstmust be static.Examples
Update full rows by the outermost dimension.
result_gmhas 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])modifiesdstto: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]