IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Mojo function

gather

def gather[dtype: DType, indices_type: DType, //, *, axis: Int, target: StringSpan[ImmStaticOrigin] = StringSpan("cpu")](output: TileTensor[dtype, Engine=output.Engine, address_space=output.address_space, linear_idx_type=output.linear_idx_type], input: TileTensor[dtype, Engine=input.Engine, address_space=input.address_space, linear_idx_type=input.linear_idx_type], indices: TileTensor[indices_type, Engine=indices.Engine, address_space=indices.address_space, linear_idx_type=indices.linear_idx_type], *, context: DeviceContext)

Gather operation as defined in https://github.com/onnx/onnx/blob/main/docs/Operators.md#Gather.

Note that this is NOT the same as the default PyTorch gather (which is equivalent to https://github.com/onnx/onnx/blob/main/docs/Operators.md#gatherelements).

Parameters:

  • ​dtype (DType): Element type of input and output.
  • ​indices_type (DType): Element type of the indices tensor.
  • ​axis (Int): Axis along which to gather from input.
  • ​target (StringSpan[ImmStaticOrigin]): Target backend to execute on, such as "cpu" or "cuda" (defaults to "cpu").

Args:

def gather[dtype: DType, indices_type: DType, InputFnType: def[width: Int, element_alignment: Int](Coord[*?]) -> SIMD[dtype, width] & RegisterPassable & ImplicitlyCopyable, IndicesFnType: def[width: Int](Coord[*?]) -> SIMD[indices_type, width] & RegisterPassable & ImplicitlyCopyable, OutputFnType: def[width: SIMDLength, element_alignment: Int](Coord[*?], SIMD[dtype, width]) -> None & RegisterPassable & ImplicitlyCopyable, *, prefetch_fn: OptionalReg[def(Coord[*?], Coord[*?]) capturing thin -> None] = None, target: StringSpan[ImmStaticOrigin] = StringSpan("cpu")](axis: Axis, input_shape: IndexList[element_type=input_shape.element_type], indices_shape: IndexList[element_type=indices_shape.element_type], output_shape: IndexList[element_type=output_shape.element_type], *, input_fn: InputFnType, indices_fn: IndicesFnType, output_fn: OutputFnType, context: DeviceContext)

Gather operation as defined in https://github.com/onnx/onnx/blob/main/docs/Operators.md#Gather.

Note that this is NOT the same as the default PyTorch gather (which is equivalent to https://github.com/onnx/onnx/blob/main/docs/Operators.md#gatherelements).

Parameters:

  • ​dtype (DType): Element type of the input and output tensors.
  • ​indices_type (DType): Element type of the indices tensor.
  • ​InputFnType (def[width: Int, element_alignment: Int](Coord[*?]) -> SIMD[dtype, width] & RegisterPassable & ImplicitlyCopyable): Function type that loads a SIMD vector from the input tensor at given coordinates.
  • ​IndicesFnType (def[width: Int](Coord[*?]) -> SIMD[indices_type, width] & RegisterPassable & ImplicitlyCopyable): Function type that loads a SIMD vector of indices from the indices buffer at given coordinates.
  • ​OutputFnType (def[width: SIMDLength, element_alignment: Int](Coord[*?], SIMD[dtype, width]) -> None & RegisterPassable & ImplicitlyCopyable): Function type that stores a SIMD vector into the output tensor at given coordinates.
  • ​prefetch_fn (OptionalReg[def(Coord[*?], Coord[*?]) capturing thin -> None]): Optional prefetch callback for software index prefetching (defaults to None).
  • ​target (StringSpan[ImmStaticOrigin]): Target backend to execute on, such as "cpu" or "cuda" (defaults to "cpu").

Args:

Was this page helpful?