已合并
[FlagGems] fix AssertionError in test_indexer_k_tiled.py::test_lighting_indexer_forward #61
adpiero创建于 3月28日
[FlagGems] fix AssertionError in test_indexer_k_tiled.py::test_lighting_indexer_forward #61
已合并
adpiero创建于 3月28日
共 3 个文件变更+12-8
@@ -90,6 +90,7 @@ def triton_lighting_indexer_k_tiled(
90 )90 )
91 tl.store(out_ptr, out_blk.to(tl.float16), out_msk)91 tl.store(out_ptr, out_blk.to(tl.float16), out_msk)
92 92 
93+device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
93 94 
94def triton_lighting_indexer_k_tiled_interface(95def triton_lighting_indexer_k_tiled_interface(
95 q, kv, weights, cu_seqlen_ks, cu_seqlen_ke96 q, kv, weights, cu_seqlen_ks, cu_seqlen_ke
@@ -97,7 +98,7 @@ def triton_lighting_indexer_k_tiled_interface(
97 Q, H, D = q.shape[0], q.shape[1], q.shape[2]98 Q, H, D = q.shape[0], q.shape[1], q.shape[2]
98 K = kv.shape[0]99 K = kv.shape[0]
99 CU = cu_seqlen_ks.shape[0]100 CU = cu_seqlen_ks.shape[0]
100- logits = torch.full([Q, K], float("-inf"), device="cuda", dtype=torch.float32)101+ logits = torch.full([Q, K], float("-inf"), device=device, dtype=torch.float32)
101 BQ = 1102 BQ = 1
102 BK = 64103 BK = 64
103 TK = 2048104 TK = 2048
@@ -1,6 +1,7 @@
1# ruff: noqa1# ruff: noqa
2import torch2import torch
3 3 
4+device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
4 5 
5def ref_fp8_mqa_logits(6def ref_fp8_mqa_logits(
6 q: torch.Tensor,7 q: torch.Tensor,
@@ -15,10 +16,10 @@ def ref_fp8_mqa_logits(
15 16 
16 seq_len_kv = kv.shape[0]17 seq_len_kv = kv.shape[0]
17 mask_lo = (18 mask_lo = (
18- torch.arange(0, seq_len_kv, device="cuda")[None, :] >= cu_seqlen_ks[:, None]19+ torch.arange(0, seq_len_kv, device=device)[None, :] >= cu_seqlen_ks[:, None]
19 )20 )
20 mask_hi = (21 mask_hi = (
21- torch.arange(0, seq_len_kv, device="cuda")[None, :] < cu_seqlen_ke[:, None]22+ torch.arange(0, seq_len_kv, device=device)[None, :] < cu_seqlen_ke[:, None]
22 )23 )
23 mask = mask_lo & mask_hi24 mask = mask_lo & mask_hi
24 25 
@@ -221,6 +221,7 @@ def per_custom_dims_cast_to_fp8(
221 x_scaled = (x * (1.0 / sf)).to(torch.float8_e4m3fn)221 x_scaled = (x * (1.0 / sf)).to(torch.float8_e4m3fn)
222 return x_scaled, sf.squeeze()222 return x_scaled, sf.squeeze()
223 223 
224+device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
224 225 
225def generate_random_cu_seqlens(226def generate_random_cu_seqlens(
226 per_cp_seqlen, cp_size=4, cp_rank=3, kv_stride=1, average_q_len=512227 per_cp_seqlen, cp_size=4, cp_rank=3, kv_stride=1, average_q_len=512
@@ -228,20 +229,21 @@ def generate_random_cu_seqlens(
228 total_seqlen = per_cp_seqlen * cp_size229 total_seqlen = per_cp_seqlen * cp_size
229 230 
230 cu_seqlens = torch.randint(231 cu_seqlens = torch.randint(
231- 0, average_q_len * 2, (total_seqlen // average_q_len * 2,)232+ 0, average_q_len * 2, (total_seqlen // average_q_len * 2,),
232- ).cuda()233+ device=device
234+ )
233 last_seq_id = torch.where(cu_seqlens.cumsum(0) >= total_seqlen)[0][0]235 last_seq_id = torch.where(cu_seqlens.cumsum(0) >= total_seqlen)[0][0]
234 cu_seqlens = cu_seqlens[:last_seq_id]236 cu_seqlens = cu_seqlens[:last_seq_id]
235 237 
236 if cu_seqlens.sum() < total_seqlen:238 if cu_seqlens.sum() < total_seqlen:
237 cu_seqlens = torch.cat(239 cu_seqlens = torch.cat(
238- [cu_seqlens, torch.tensor([total_seqlen - cu_seqlens.sum()]).cuda()]240+ [cu_seqlens, torch.tensor([total_seqlen - cu_seqlens.sum()], device=device)]
239 )241 )
240 242 
241 cu_seqlens_cumsum = torch.cumsum(cu_seqlens, dim=0)243 cu_seqlens_cumsum = torch.cumsum(cu_seqlens, dim=0)
242 cu_seqlens_k_cumsum = torch.cumsum(cu_seqlens // kv_stride, dim=0)244 cu_seqlens_k_cumsum = torch.cumsum(cu_seqlens // kv_stride, dim=0)
243- cu_seqlens_qs = torch.cat([torch.tensor([0]).cuda(), cu_seqlens_cumsum[:-1]])245+ cu_seqlens_qs = torch.cat([torch.tensor([0], device=device), cu_seqlens_cumsum[:-1]])
244- cu_seqlens_ks = torch.cat([torch.tensor([0]).cuda(), cu_seqlens_k_cumsum[:-1]])246+ cu_seqlens_ks = torch.cat([torch.tensor([0], device=device), cu_seqlens_k_cumsum[:-1]])
245 cu_seqlens_qe = cu_seqlens_cumsum.clone()247 cu_seqlens_qe = cu_seqlens_cumsum.clone()
246 cu_seqlens_ke = cu_seqlens_k_cumsum.clone()248 cu_seqlens_ke = cu_seqlens_k_cumsum.clone()
247 249