IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /max/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. /max/get-started.md).

Mojo function

logsoftmax

def logsoftmax[dtype: DType, simd_width: Int, rank: Int, input_fn: def[_simd_width: Int](Coord[**?]) capturing thin -> SIMD[dtype, _simd_width], target: StringSlice[ImmStaticOrigin] = StringSlice("cpu"), has_prologue_fusion: Bool = True](shape: Coord, output: TileTensor[dtype, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type], axis: Int, context: Optional[DeviceContext] = None)

Computes log-softmax over the given axis using a caller-supplied input lambda.

Delegates to softmax with logsoftmax=True, which applies an elementwise log to the normalized outputs.

Parameters:

  • ​dtype (DType): The dtype of the input and output buffers.
  • ​simd_width (Int): The simd_width to use in vectorization.
  • ​rank (Int): The rank of the input and output tensors.
  • ​input_fn (def[_simd_width: Int](Coord[**?]) capturing thin -> SIMD[dtype, _simd_width]): The elementwise input lambda.
  • ​target (StringSlice[ImmStaticOrigin]): The target device ("cpu" or "gpu").
  • ​has_prologue_fusion (Bool): Whether the input lambda supports prologue fusion.

Args:

def logsoftmax[dtype: DType, simd_width: Int, rank: Int, target: StringSlice[ImmStaticOrigin] = StringSlice("cpu")](input: TileTensor[dtype, Storage=input.Storage, address_space=input.address_space, linear_idx_type=input.linear_idx_type], output: TileTensor[dtype, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type], axis: Int, context: Optional[DeviceContext] = None)

Computes log-softmax over the given axis of input and stores the result in output.

Wraps input with a load lambda and delegates to softmax with logsoftmax=True.

Parameters:

  • ​dtype (DType): The dtype of the input and output buffers.
  • ​simd_width (Int): The simd_width to use in vectorization.
  • ​rank (Int): The rank of the input and output tensors.
  • ​target (StringSlice[ImmStaticOrigin]): The target device ("cpu" or "gpu").

Args: