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
def spatial_merge[dtype: DType](output: TileTensor[dtype, Storage=output.Storage, linear_idx_type=output.linear_idx_type], input: TileTensor[dtype, Storage=input.Storage, linear_idx_type=input.linear_idx_type], grid_thw: TileTensor[DType.int64, Storage=grid_thw.Storage, linear_idx_type=grid_thw.linear_idx_type], hidden_size: Int, merge_size: Int, ctx: DeviceContext)
Launches the spatial merge kernel that compresses vision token grids by merging spatial blocks before attention.
Parameters:
- βdtype (
DType): Element type of the input and output tensors.
Args:
- βoutput (
TileTensor[dtype, Storage=output.Storage, linear_idx_type=output.linear_idx_type]): Output tensor holding the merged patches. - βinput (
TileTensor[dtype, Storage=input.Storage, linear_idx_type=input.linear_idx_type]): Input tensor holding the original patch grid. - βgrid_thw (
TileTensor[DType.int64, Storage=grid_thw.Storage, linear_idx_type=grid_thw.linear_idx_type]): Grid dimensions tensor of shape(batch_size, 3)with[t, h, w]per item. - βhidden_size (
Int): Hidden dimension size of each patch. - βmerge_size (
Int): Size of the spatial merge blocks. - βctx (
DeviceContext): Device context used to enqueue the kernel.