已合并
[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
已合并
共 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 | ||
| 94 | def triton_lighting_indexer_k_tiled_interface( | 95 | def triton_lighting_indexer_k_tiled_interface( |
| 95 | q, kv, weights, cu_seqlen_ks, cu_seqlen_ke | 96 | 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 = 1 | 102 | BQ = 1 |
| 102 | BK = 64 | 103 | BK = 64 |
| 103 | TK = 2048 | 104 | TK = 2048 |
| @@ -1,6 +1,7 @@ | |||
| 1 | # ruff: noqa | 1 | # ruff: noqa |
| 2 | import torch | 2 | import torch |
| 3 | 3 | ||
| 4 | +device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | ||
| 4 | 5 | ||
| 5 | def ref_fp8_mqa_logits( | 6 | def 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_hi | 24 | 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 | ||
| 225 | def generate_random_cu_seqlens( | 226 | def generate_random_cu_seqlens( |
| 226 | per_cp_seqlen, cp_size=4, cp_rank=3, kv_stride=1, average_q_len=512 | 227 | 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_size | 229 | 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 | ||