[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):
|
||||
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]:
|
||||
|
||||
Reference in New Issue
Block a user