@@ -95,7 +95,7 @@ def _compute(block_kv_start_idx, block_kv_seqlen, o, m_i, l_i):
9595 # Load q: it will stay in L1 throughout. Indices form a matrix because we
9696 # read, compute, and write all in 2d chunks. 1 element ~= 1 CUDA thread index.
9797 # q tile has shape [block_h, head_dim].
98- q = pl . load ( q_ref , ( slice ( None ), slice ( None )), mask = q_mask ) * softmax_scale
98+ q = jnp . where ( q_mask , q_ref [...], 0.0 ) * softmax_scale
9999
100100 mask_indices = jnp .arange (block_k )
101101
@@ -110,13 +110,11 @@ def body(start_k, carry):
110110
111111 def compute ():
112112 curr_k_slice = pl .ds (start_k * block_k , block_k )
113- k = pl . load ( k_ref , ( curr_k_slice , slice ( None )), mask = mask [:, None ], other = 0.0 )
113+ k = jnp . where ( mask [:, None ], k_ref [ curr_k_slice , : ], 0.0 )
114114 k = k .astype (q .dtype )
115115 qk = pl .dot (q , k .T , precision = precision ) # [block_h, block_k]
116116 if bias_ref is not None :
117- qk += pl .load (
118- bias_ref , (slice (None ), curr_k_slice ), mask = mask [None , :], other = 0.0
119- )
117+ qk += jnp .where (mask [None , :], bias_ref [:, curr_k_slice ], 0.0 )
120118
121119 qk = jnp .where (logits_mask [None , :], qk , NEG_INF )
122120
@@ -128,7 +126,7 @@ def compute():
128126 s_curr = jnp .exp (qk - m_next [:, None ])
129127 l_curr = s_curr .sum (axis = - 1 )
130128 l_next = l_prev_corr + l_curr
131- v = pl . load ( v_ref , ( curr_k_slice , slice ( None )), mask = mask [:, None ], other = 0.0 )
129+ v = jnp . where ( mask [:, None ], v_ref [ curr_k_slice , : ], 0.0 )
132130 v = v .astype (q .dtype )
133131 o_curr = pl .dot (s_curr .astype (v .dtype ), v , precision = precision )
134132
@@ -155,7 +153,8 @@ def no_compute():
155153 o = jnp .zeros ((block_h , head_dim ), dtype = jnp .float32 )
156154
157155 block_kv_start_idx = prog_j * split_k_seq_len
158- kv_seq_len = pl .load (kv_seq_len_ref , ())
156+ kv_seq_len = kv_seq_len_ref [...]
157+
159158 block_kv_seqlen = jnp .minimum ((prog_j + 1 ) * split_k_seq_len , kv_seq_len )
160159
161160 # Skip padding in seq dim.
@@ -167,9 +166,9 @@ def no_compute():
167166
168167 # Write output to HBM.
169168 vec_q_mask = q_mask .reshape (- 1 )
170- pl . store ( l_ref , slice ( None ) , l_i , mask = vec_q_mask )
171- pl . store ( m_ref , slice ( None ) , m_i , mask = vec_q_mask )
172- pl . store ( o_ref , ( slice ( None ), slice ( None )) , o , mask = q_mask )
169+ l_ref [...] = jnp . where ( vec_q_mask , l_i , l_ref [...] )
170+ m_ref [...] = jnp . where ( vec_q_mask , m_i , m_ref [...] )
171+ o_ref [...] = jnp . where ( q_mask , o , o_ref [...] )
173172
174173
175174def _get_sm_count () -> int :
0 commit comments