Implementation:InternLM Lmdeploy Gemm CtaMap
| Knowledge Sources | |
|---|---|
| Domains | GPU_Kernels, GEMM |
| Last Updated | 2026-02-07 15:00 GMT |
Overview
Implements CTA-to-tile mapping strategies for GEMM kernels, including static GemmScheduler with swizzled tile assignment and DynamicScheduler for grouped/batched GEMM workloads.
Description
This header provides two tile scheduling strategies:
GemmScheduler (templated on Order) maps CUDA grid blocks to GEMM output tiles using a log-tile swizzling scheme that improves L2 cache locality. It supports split-K partitioning by dividing the K dimension into chunks distributed across the grid's Z dimension. The scheduler computes tile offsets and K-iteration ranges from block indices at kernel launch time.
DynamicScheduler supports grouped GEMM where each group may have different problem dimensions. It reads precomputed tile assignments, GEMM shapes, and K-iteration ranges from a Tape structure, allowing a single kernel launch to process heterogeneous GEMM problems.
Helper functions get_log_tile and get_tiled_shape compute the swizzle factor and tile counts for a given problem size.
Usage
Used by GemmUniversal and KernelImpl to determine which output tile each CTA computes and how the K dimension is partitioned for split-K execution.
Code Reference
Source Location
- Repository: InternLM_Lmdeploy
- File: src/turbomind/kernels/gemm/cta_map.h
Signature
template<Order order_>
class GemmScheduler {
public:
GemmScheduler(int4 gemm_shape, int2 tiled_mn, int splits, int log_tile, int cta_k, int chunk_size);
static int get_log_tile(int2 tiled_mn, int tile_size);
static dim3 get_grid_shape(int4 tiled_shape, int log_tile);
std::true_type init(); // device
int4 tile_offset() const;
int2 iter_k_range() const;
};
template<Order order_>
class DynamicScheduler {
public:
DynamicScheduler(const Tape& tape);
bool init(); // device
int4 tile_offset() const;
int2 iter_k_range() const;
};
Import
#include "src/turbomind/kernels/gemm/cta_map.h"
I/O Contract
Inputs
| Name | Type | Required | Description |
|---|---|---|---|
| gemm_shape | int4 | Yes | (M, N, K, batch) problem dimensions |
| tiled_mn | int2 | Yes | Number of tiles in M and N dimensions |
| splits | int | Yes | Split-K factor |
| log_tile | int | Yes | Log2 of the swizzle tile size |
| tape | Tape | For DynamicScheduler | Precomputed tile assignments for grouped GEMM |
Outputs
| Name | Type | Description |
|---|---|---|
| tile_offset | int4 | (tile_m, tile_n, split_id, group_id) for the current CTA |
| iter_k_range | int2 | (begin, end) K-iteration range for this CTA |
| grid_shape | dim3 | Grid dimensions for kernel launch |
Usage Examples
// Static scheduling with swizzle
GemmScheduler<kColMajor> sched({M, N, K, 1}, tiled_mn, splits, log_tile, CTA_K, chunk_size);
auto grid = sched.get_grid_shape();
// Inside kernel
sched.init(); // reads blockIdx
auto offset = sched.tile_offset();
auto [k_beg, k_end] = sched.iter_k_range();