For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /max/get-started.md).
Mojo function
advanced_indexing_setitem_inplace
def advanced_indexing_setitem_inplace[index_rank: Int, updates_rank: Int, input_type: DType, //, start_axis: Int, num_index_tensors: Int, target: StringSlice[ImmStaticOrigin], trace_description: StringSlice[ImmStaticOrigin], UpdatesTensorFn: def[dtype: DType, width: Int](IndexList[updates_rank]) -> SIMD[dtype, width] & ImplicitlyCopyable & RegisterPassable, IndicesFn: def[indices_index: Int](IndexList[index_rank]) -> Int & ImplicitlyCopyable & RegisterPassable](input_tensor: TileTensor[input_type, Storage=input_tensor.Storage, address_space=input_tensor.address_space, linear_idx_type=input_tensor.linear_idx_type], index_tensor_shape: IndexList[index_rank], updates_tensor_strides: IndexList[updates_rank], ctx: DeviceContext, updates_tensor_fn: UpdatesTensorFn, indices_fn: IndicesFn) where (eq UpdatesTensorFn.updates_rank, updates_rank) where (eq IndicesFn.index_rank, index_rank)
Implement basic numpy-style advanced indexing with assignment.
This is designed to be fused with other view-producing operations to implement full numpy-indexing semantics.
This assumes the dimensions in input_tensor not indexed by index tensors
are ":", ie selecting all indices along the slice. For example in numpy:
# rank(indices1) == 2
# rank(indices2) == 2
# rank(updates) == 2
input_tensor[:, :, :, indices1, indices2, :, :] = updatesWe calculate the following for all valid valued indexing variables:
input_tensor[
a, b, c,
indices1[i, j],
indices2[i, j],
d, e
] = updates[i, j]In this example start_axis = 3 and num_index_tensors = 2.
In terms of implementation details, our strategy is to iterate over
all indices over a common iteration range. The idea is we can map
indices in this range to the write location in input_tensor as well
as the data location in updates. An update can illustrate how this is
possible best:
Imagine the input_tensor shape is [A, B, C, D] and we have indexing
tensors I1 and I2 with shape [M, N, K]. Assume I1 and I2 are applied
to dimensions 1 and 2.
I claim an appropriate common iteration range is then (A, M, N, K, D).
Note we expect updates to be the shape [A, M, N, K, D]. We will show
this by providing the mappings into updates and input_tensor:
Consider an arbitrary set of indices in this range (a, m, n, k, d):
- The index into updates is (a, m, n, k, d).
- The index into input_tensor is (a, I1[m, n, k], I2[m, n, k], d).
Note: Currently supports contiguous indexing tensors only; boolean tensor masks, view-fusion, and a unified getitem/setitem interface are not yet implemented.
Parameters:
- βindex_rank (
Int): The rank of the indexing tensors. - βupdates_rank (
Int): The rank of the updates tensor. - βinput_type (
DType): The dtype of the input tensor. - βstart_axis (
Int): The first dimension in input where the indexing tensors are applied. It is assumed the indexing tensors are applied in consecutive dimensions. - βnum_index_tensors (
Int): The number of indexing tensors. - βtarget (
StringSlice[ImmStaticOrigin]): The target architecture to operation on. - βtrace_description (
StringSlice[ImmStaticOrigin]): For profiling, the trace name the operation will appear under. - βUpdatesTensorFn (
def[dtype: DType, width: Int](IndexList[updates_rank]) -> SIMD[dtype, width]&ImplicitlyCopyable&RegisterPassable): The type of the updates-tensor fusion lambda. - βIndicesFn (
def[indices_index: Int](IndexList[index_rank]) -> Int&ImplicitlyCopyable&RegisterPassable): The type of the indices fusion lambda.
Args:
- βinput_tensor (
TileTensor[input_type, Storage=input_tensor.Storage, address_space=input_tensor.address_space, linear_idx_type=input_tensor.linear_idx_type]): The input tensor being indexed into and modified in-place. - βindex_tensor_shape (
IndexList[index_rank]): The shape of each index tensor. - βupdates_tensor_strides (
IndexList[updates_rank]): The strides of the update tensor. - βctx (
DeviceContext): The device context as prepared by the graph compiler. - βupdates_tensor_fn (
UpdatesTensorFn): Fusion lambda for the update tensor. - βindices_fn (
IndicesFn): Fusion lambda for the indices tensors.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!