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
spatial_merge_kernel
def spatial_merge_kernel[dtype: DType, InputLayoutType: TensorLayout, input_origin: ImmOrigin, InputEngine: TensorEngine, OutputLayoutType: TensorLayout, output_origin: MutOrigin, OutputEngine: TensorEngine, GridThwLayoutType: TensorLayout, grid_thw_origin: ImmOrigin, GridThwEngine: TensorEngine](output: TileTensor[dtype, OutputLayoutType, output_origin, Engine=OutputEngine], input: TileTensor[dtype, InputLayoutType, input_origin, Engine=InputEngine], grid_thw: TileTensor[.int64, GridThwLayoutType, grid_thw_origin, Engine=GridThwEngine], batch_size: Int32, hidden_size: Int32, merge_size: Int32)
Spatial merge kernel.
Grid: 1D over all output patches (one block per output patch). Threads: loop over channels (hidden_size x merge_size^2).
Parameters:
- dtype (
DType): Element type of the input and output tensors. - InputLayoutType (
TensorLayout): Compile-timeTensorLayoutof the input tensor. - input_origin (
ImmOrigin): Immutable origin of the input tensor. - InputEngine (
TensorEngine): Engine of the input tensor. - OutputLayoutType (
TensorLayout): Compile-timeTensorLayoutof the output tensor. - output_origin (
MutOrigin): Mutable origin of the output tensor. - OutputEngine (
TensorEngine): Engine of the output tensor. - GridThwLayoutType (
TensorLayout): Compile-timeTensorLayoutof thegrid_thwtensor. - grid_thw_origin (
ImmOrigin): Immutable origin of thegrid_thwtensor. - GridThwEngine (
TensorEngine): Engine of thegrid_thwtensor.
Args:
- output (
TileTensor[dtype, OutputLayoutType, output_origin, Engine=OutputEngine]): Output tensor. - input (
TileTensor[dtype, InputLayoutType, input_origin, Engine=InputEngine]): Input tensor. - grid_thw (
TileTensor[.int64, GridThwLayoutType, grid_thw_origin, Engine=GridThwEngine]): Grid dimensions tensor (B, 3) containing [t, h, w] for each item. - batch_size (
Int32): Number of items in batch. - hidden_size (
Int32): Hidden dimension size. - merge_size (
Int32): Size of spatial merge blocks.