[Kernel] Coalesce the KDA CuTe DSL decode state transpose: ~3x faster, bit-identical (#39680)
Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5.1
parent
db39b7f961
commit
f447bb7080
@@ -203,8 +203,8 @@ def _define_kernels():
|
|||||||
|
|
||||||
for k_iter in range(NUM_K_ITERS_SMALL):
|
for k_iter in range(NUM_K_ITERS_SMALL):
|
||||||
flat_idx = tidx + k_iter * NUM_THREADS
|
flat_idx = tidx + k_iter * NUM_THREADS
|
||||||
k_load = flat_idx // TILE_V_SMALL
|
k_load = flat_idx % TILE_K
|
||||||
v_load = flat_idx % TILE_V_SMALL
|
v_load = flat_idx // TILE_K
|
||||||
if k_load < TILE_K:
|
if k_load < TILE_K:
|
||||||
v_global_load = v_tile * TILE_V_SMALL + v_load
|
v_global_load = v_tile * TILE_V_SMALL + v_load
|
||||||
h_val = 0.0
|
h_val = 0.0
|
||||||
@@ -262,8 +262,8 @@ def _define_kernels():
|
|||||||
|
|
||||||
for k_iter in range(NUM_K_ITERS_SMALL):
|
for k_iter in range(NUM_K_ITERS_SMALL):
|
||||||
flat_idx = tidx + k_iter * NUM_THREADS
|
flat_idx = tidx + k_iter * NUM_THREADS
|
||||||
k_write = flat_idx // TILE_V_SMALL
|
k_write = flat_idx % TILE_K
|
||||||
v_write = flat_idx % TILE_V_SMALL
|
v_write = flat_idx // TILE_K
|
||||||
if k_write < TILE_K:
|
if k_write < TILE_K:
|
||||||
v_global_write = v_tile * TILE_V_SMALL + v_write
|
v_global_write = v_tile * TILE_V_SMALL + v_write
|
||||||
if v_global_write < v.shape[3]:
|
if v_global_write < v.shape[3]:
|
||||||
@@ -424,8 +424,8 @@ def _define_kernels():
|
|||||||
|
|
||||||
for k_iter in range(NUM_K_ITERS_SMALL):
|
for k_iter in range(NUM_K_ITERS_SMALL):
|
||||||
flat_idx = tidx + k_iter * NUM_THREADS
|
flat_idx = tidx + k_iter * NUM_THREADS
|
||||||
k_load = flat_idx // TILE_V_SMALL
|
k_load = flat_idx % TILE_K
|
||||||
v_load = flat_idx % TILE_V_SMALL
|
v_load = flat_idx // TILE_K
|
||||||
if k_load < TILE_K:
|
if k_load < TILE_K:
|
||||||
v_global_load = v_tile * TILE_V_SMALL + v_load
|
v_global_load = v_tile * TILE_V_SMALL + v_load
|
||||||
h_val = 0.0
|
h_val = 0.0
|
||||||
@@ -483,8 +483,8 @@ def _define_kernels():
|
|||||||
|
|
||||||
for k_iter in range(NUM_K_ITERS_SMALL):
|
for k_iter in range(NUM_K_ITERS_SMALL):
|
||||||
flat_idx = tidx + k_iter * NUM_THREADS
|
flat_idx = tidx + k_iter * NUM_THREADS
|
||||||
k_write = flat_idx // TILE_V_SMALL
|
k_write = flat_idx % TILE_K
|
||||||
v_write = flat_idx % TILE_V_SMALL
|
v_write = flat_idx // TILE_K
|
||||||
if k_write < TILE_K:
|
if k_write < TILE_K:
|
||||||
v_global_write = v_tile * TILE_V_SMALL + v_write
|
v_global_write = v_tile * TILE_V_SMALL + v_write
|
||||||
if v_global_write < v.shape[3]:
|
if v_global_write < v.shape[3]:
|
||||||
@@ -639,8 +639,8 @@ def _define_kernels():
|
|||||||
|
|
||||||
for k_iter in range(NUM_K_ITERS):
|
for k_iter in range(NUM_K_ITERS):
|
||||||
flat_idx = tidx + k_iter * NUM_THREADS_LARGE
|
flat_idx = tidx + k_iter * NUM_THREADS_LARGE
|
||||||
k_load = flat_idx // TILE_V
|
k_load = flat_idx % TILE_K
|
||||||
v_load = flat_idx % TILE_V
|
v_load = flat_idx // TILE_K
|
||||||
if k_load < TILE_K:
|
if k_load < TILE_K:
|
||||||
v_global_load = v_tile * TILE_V + v_load
|
v_global_load = v_tile * TILE_V + v_load
|
||||||
h_val = 0.0
|
h_val = 0.0
|
||||||
@@ -698,8 +698,8 @@ def _define_kernels():
|
|||||||
|
|
||||||
for k_iter in range(NUM_K_ITERS):
|
for k_iter in range(NUM_K_ITERS):
|
||||||
flat_idx = tidx + k_iter * NUM_THREADS_LARGE
|
flat_idx = tidx + k_iter * NUM_THREADS_LARGE
|
||||||
k_write = flat_idx // TILE_V
|
k_write = flat_idx % TILE_K
|
||||||
v_write = flat_idx % TILE_V
|
v_write = flat_idx // TILE_K
|
||||||
if k_write < TILE_K:
|
if k_write < TILE_K:
|
||||||
v_global_write = v_tile * TILE_V + v_write
|
v_global_write = v_tile * TILE_V + v_write
|
||||||
if v_global_write < v.shape[3]:
|
if v_global_write < v.shape[3]:
|
||||||
@@ -854,8 +854,8 @@ def _define_kernels():
|
|||||||
|
|
||||||
for k_iter in range(NUM_K_ITERS):
|
for k_iter in range(NUM_K_ITERS):
|
||||||
flat_idx = tidx + k_iter * NUM_THREADS_LARGE
|
flat_idx = tidx + k_iter * NUM_THREADS_LARGE
|
||||||
k_load = flat_idx // TILE_V
|
k_load = flat_idx % TILE_K
|
||||||
v_load = flat_idx % TILE_V
|
v_load = flat_idx // TILE_K
|
||||||
if k_load < TILE_K:
|
if k_load < TILE_K:
|
||||||
v_global_load = v_tile * TILE_V + v_load
|
v_global_load = v_tile * TILE_V + v_load
|
||||||
h_val = 0.0
|
h_val = 0.0
|
||||||
@@ -913,8 +913,8 @@ def _define_kernels():
|
|||||||
|
|
||||||
for k_iter in range(NUM_K_ITERS):
|
for k_iter in range(NUM_K_ITERS):
|
||||||
flat_idx = tidx + k_iter * NUM_THREADS_LARGE
|
flat_idx = tidx + k_iter * NUM_THREADS_LARGE
|
||||||
k_write = flat_idx // TILE_V
|
k_write = flat_idx % TILE_K
|
||||||
v_write = flat_idx % TILE_V
|
v_write = flat_idx // TILE_K
|
||||||
if k_write < TILE_K:
|
if k_write < TILE_K:
|
||||||
v_global_write = v_tile * TILE_V + v_write
|
v_global_write = v_tile * TILE_V + v_write
|
||||||
if v_global_write < v.shape[3]:
|
if v_global_write < v.shape[3]:
|
||||||
|
|||||||
Reference in New Issue
Block a user