diff --git a/python/sglang/kernels/ops/attention/cutedsl_kda.py b/python/sglang/kernels/ops/attention/cutedsl_kda.py index 312769e9d..01cd888cc 100644 --- a/python/sglang/kernels/ops/attention/cutedsl_kda.py +++ b/python/sglang/kernels/ops/attention/cutedsl_kda.py @@ -203,8 +203,8 @@ def _define_kernels(): for k_iter in range(NUM_K_ITERS_SMALL): flat_idx = tidx + k_iter * NUM_THREADS - k_load = flat_idx // TILE_V_SMALL - v_load = flat_idx % TILE_V_SMALL + k_load = flat_idx % TILE_K + v_load = flat_idx // TILE_K if k_load < TILE_K: v_global_load = v_tile * TILE_V_SMALL + v_load h_val = 0.0 @@ -262,8 +262,8 @@ def _define_kernels(): for k_iter in range(NUM_K_ITERS_SMALL): flat_idx = tidx + k_iter * NUM_THREADS - k_write = flat_idx // TILE_V_SMALL - v_write = flat_idx % TILE_V_SMALL + k_write = flat_idx % TILE_K + v_write = flat_idx // TILE_K if k_write < TILE_K: v_global_write = v_tile * TILE_V_SMALL + v_write if v_global_write < v.shape[3]: @@ -424,8 +424,8 @@ def _define_kernels(): for k_iter in range(NUM_K_ITERS_SMALL): flat_idx = tidx + k_iter * NUM_THREADS - k_load = flat_idx // TILE_V_SMALL - v_load = flat_idx % TILE_V_SMALL + k_load = flat_idx % TILE_K + v_load = flat_idx // TILE_K if k_load < TILE_K: v_global_load = v_tile * TILE_V_SMALL + v_load h_val = 0.0 @@ -483,8 +483,8 @@ def _define_kernels(): for k_iter in range(NUM_K_ITERS_SMALL): flat_idx = tidx + k_iter * NUM_THREADS - k_write = flat_idx // TILE_V_SMALL - v_write = flat_idx % TILE_V_SMALL + k_write = flat_idx % TILE_K + v_write = flat_idx // TILE_K if k_write < TILE_K: v_global_write = v_tile * TILE_V_SMALL + v_write if v_global_write < v.shape[3]: @@ -639,8 +639,8 @@ def _define_kernels(): for k_iter in range(NUM_K_ITERS): flat_idx = tidx + k_iter * NUM_THREADS_LARGE - k_load = flat_idx // TILE_V - v_load = flat_idx % TILE_V + k_load = flat_idx % TILE_K + v_load = flat_idx // TILE_K if k_load < TILE_K: v_global_load = v_tile * TILE_V + v_load h_val = 0.0 @@ -698,8 +698,8 @@ def _define_kernels(): for k_iter in range(NUM_K_ITERS): flat_idx = tidx + k_iter * NUM_THREADS_LARGE - k_write = flat_idx // TILE_V - v_write = flat_idx % TILE_V + k_write = flat_idx % TILE_K + v_write = flat_idx // TILE_K if k_write < TILE_K: v_global_write = v_tile * TILE_V + v_write if v_global_write < v.shape[3]: @@ -854,8 +854,8 @@ def _define_kernels(): for k_iter in range(NUM_K_ITERS): flat_idx = tidx + k_iter * NUM_THREADS_LARGE - k_load = flat_idx // TILE_V - v_load = flat_idx % TILE_V + k_load = flat_idx % TILE_K + v_load = flat_idx // TILE_K if k_load < TILE_K: v_global_load = v_tile * TILE_V + v_load h_val = 0.0 @@ -913,8 +913,8 @@ def _define_kernels(): for k_iter in range(NUM_K_ITERS): flat_idx = tidx + k_iter * NUM_THREADS_LARGE - k_write = flat_idx // TILE_V - v_write = flat_idx % TILE_V + k_write = flat_idx % TILE_K + v_write = flat_idx // TILE_K if k_write < TILE_K: v_global_write = v_tile * TILE_V + v_write if v_global_write < v.shape[3]: