asc.experimental.asctile.gather

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

Gather subtensors from a global tensor at positions given by an index tensor.

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

The index tensor must have an integer dtype (int8, int16, int32, int64).

Parameters:
  • src – The source global tensor to gather from.

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

  • dim – The dimension of src used for indexing.

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

  • check_bounds – If True, out-of-bounds indices produce pad_value elements. 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.

  • pad_value – The value used to pad out-of-bounds indices and to align the last dimension of the result to 32 bytes. If not specified, 0 is used.

Returns:

A tensor with shape [index.shape[0], src.shape[dim+1], ..., src.shape[src.rank-1]] and the

same dtype as src, located in UB. The last dimension is aligned to 32 bytes and padded with pad_value.

Return type:

LocalTensor

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

  • RuntimeError – If index does not have an integer dtype or is not located in UB

  • ValueError – If offsets does not contain dim + 1 values, dim is out of range for src.rank, index is not rank 1, or src has a dynamic dimension after dim

  • NotImplementedError – If dim is the last dimension of src

Note

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

Examples

Read full rows by outermost dimension. src has shape [1024, 128] and tile has shape [256, 128]:

index = asctile.copy_in(index_gm, [0], [256])
tile = asctile.gather(src, [0], 0, index)

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

src = [[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]

Then tile = asctile.gather(src, [0], 0, index) results in:

tile = [[0, 1, 2, 3, 4, 5, 6, 7],
        [16, 17, 18, 19, 20, 21, 22, 23],
        ...,
        [240, 241, 242, 243, 244, 245, 246, 247]]  # shape [16, 8]