已合并
Support block sparse attention grad TND GQA #4945
Support block sparse attention grad TND GQA #4945
已合并
fgd_dragon创建于 5月14日
4 个文件变更+371-21
Mcodegen/templates/_op_plugin_docs.py+3-0
@@ -541,6 +541,9 @@ query 的 head 数 N1 与 key/value 的 head 数 N2,需满足 N1 >= N2 且 N1
541block_sparse_mask 必传,且shape必须为[batch, headNum, ceilDiv(maxQS, blockShapeX), ceilDiv(maxKVS, blockShapeY)]. 541block_sparse_mask 必传,且shape必须为[batch, headNum, ceilDiv(maxQS, blockShapeX), ceilDiv(maxKVS, blockShapeY)].
542block_shape 必传,必须包含至少两个元素[blockShapeX, blockShapeY],且值必须大于0;blockShapeY 必须为 128 的倍数。542block_shape 必传,必须包含至少两个元素[blockShapeX, blockShapeY],且值必须大于0;blockShapeY 必须为 128 的倍数。
543当 q_input_layout 为 "TND" 时 actual_seq_lengths 必选; 当 kv_input_layout 为 "TND" 时 actual_seq_lengths_kv 必选. 543当 q_input_layout 为 "TND" 时 actual_seq_lengths 必选; 当 kv_input_layout 为 "TND" 时 actual_seq_lengths_kv 必选.
544+actual_seq_lengths 与 actual_seq_lengths_kv 当前必须同时配置或同时不配置,仅配置其中之一会被算子拦截.
545+正向路径当前支持 headDim=64128; 反向路径当前支持 headDim=128.
546+反向路径支持 q_input_layout 和 kv_input_layout 同为 "BNSD" 或同为 "TND",并支持 MHA/GQA 场景. MHA 场景下 N1 == N2, GQA 场景下需满足 N1 > N2 且 N1 % N2 == 0, 其中 N1 为 query 的 head 数, N2 为 key/value 的 head 数.
544inner_precise 仅支持 0(表示float32 softmax) 或 1(表示float16 softmax);当 query/key/value 为 bfloat16 时,仅支持 0.547inner_precise 仅支持 0(表示float32 softmax) 或 1(表示float16 softmax);当 query/key/value 为 bfloat16 时,仅支持 0.
545 548 
546支持的PyTorch版本549支持的PyTorch版本
Mdocs/zh/custom_APIs/torch_npu/torch_npu-npu_block_sparse_attention.md+3-1
@@ -75,8 +75,10 @@ torch_npu.npu_block_sparse_attention(query, key, value, block_sparse_mask, block
75 75 
76- `query``key``value`数据类型必须一致,且为`float16``bfloat16`76- `query``key``value`数据类型必须一致,且为`float16``bfloat16`
77- `query`的head数$N1$与`key`/`value`的head数$N2$需满足$N1 ≥ N2$且$N1 \% N2 = 0$。77- `query`的head数$N1$与`key`/`value`的head数$N2$需满足$N1 ≥ N2$且$N1 \% N2 = 0$。
78+- `actual_seq_lengths``actual_seq_lengths_kv`当前必须同时配置或同时不配置,仅配置其中之一会被算子拦截。
78- 序列长度不需要被`block_shape`整除,分块数按向上取整计算。79- 序列长度不需要被`block_shape`整除,分块数按向上取整计算。
79-- 当前版本下,当且仅当`q_input_layout`和`kv_input_layout`为`"BNSD"`、MHA场景(`query`的head数$N1$与`key`/`value`的head数$N2$相等),且headDim=128时,支持反向计算80+- 正向路径当前支持headDim=64或128反向路径当前支持headDim=128
81+- 反向路径支持`q_input_layout``kv_input_layout`同为`"BNSD"`或同为`"TND"`,并支持MHA/GQA场景。MHA场景下$N1 = N2$,GQA场景下需满足$N1 > N2$且$N1 \% N2 = 0$,其中$N1$为`query`的head数,$N2$为`key`/`value`的head数。
80 82 
81## 调用示例83## 调用示例
82 84 
Mop_plugin/ops/opapi/BlockSparseAttentionBackwardKernelNpuOpApi.cpp+27-12
@@ -23,24 +23,40 @@ const int64_t MAX_HEAD_DIM = 128;
23using npu_preparation = at_npu::native::OpPreparation;23using npu_preparation = at_npu::native::OpPreparation;
24 24 
25 25 
26-// 入参检查26+// Validate input parameters.
27static void check_params(const at::Tensor &query,27static void check_params(const at::Tensor &query,
28 const at::Tensor &key,28 const at::Tensor &key,
29- const at::Tensor &value)29+ const at::Tensor &value,
30+ const c10::OptionalIntArrayRef actual_seq_lengths,
31+ const c10::OptionalIntArrayRef actual_seq_lengths_kv,
32+ c10::string_view q_input_layout,
33+ c10::string_view kv_input_layout)
30{34{
31- // Q/K/V 数据类型必须一致35+ // Q/K/V must use the same dtype.
32 TORCH_CHECK(query.scalar_type() == key.scalar_type() && key.scalar_type() == value.scalar_type(),36 TORCH_CHECK(query.scalar_type() == key.scalar_type() && key.scalar_type() == value.scalar_type(),
33 "query, key, value must have the same dtype, got query=", query.scalar_type(),37 "query, key, value must have the same dtype, got query=", query.scalar_type(),
34 ", key=", key.scalar_type(), ", value=", value.scalar_type(), OPS_ERROR(ErrCode::PARAM));38 ", key=", key.scalar_type(), ", value=", value.scalar_type(), OPS_ERROR(ErrCode::PARAM));
35 39 
36- // head_dim 不能超过 12840+ // The kernel supports head_dim up to 128.
37 int64_t head_dim = query.size(-1);41 int64_t head_dim = query.size(-1);
38 TORCH_CHECK(head_dim <= MAX_HEAD_DIM,42 TORCH_CHECK(head_dim <= MAX_HEAD_DIM,
39 "head_dim must be <= ", MAX_HEAD_DIM, ", but got ", head_dim, OPS_ERROR(ErrCode::PARAM));43 "head_dim must be <= ", MAX_HEAD_DIM, ", but got ", head_dim, OPS_ERROR(ErrCode::PARAM));
44+ 
45+ // TND inputs require non-empty per-batch actual sequence lengths.
46+ if (q_input_layout == "TND") {
47+ TORCH_CHECK(actual_seq_lengths.has_value() && actual_seq_lengths->size() > 0,
48+ "actual_seq_lengths must be specified when q_input_layout is TND",
49+ OPS_ERROR(ErrCode::PARAM));
50+ }
51+ if (kv_input_layout == "TND") {
52+ TORCH_CHECK(actual_seq_lengths_kv.has_value() && actual_seq_lengths_kv->size() > 0,
53+ "actual_seq_lengths_kv must be specified when kv_input_layout is TND",
54+ OPS_ERROR(ErrCode::PARAM));
55+ }
40}56}
41 57 
42 58 
43-// PTA 接口实现59+// PTA API implementation.
44std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_block_sparse_attention_backward(60std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_block_sparse_attention_backward(
45 const at::Tensor &d_out,61 const at::Tensor &d_out,
46 const at::Tensor &query,62 const at::Tensor &query,
@@ -57,30 +73,29 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_block_sparse_attention_backwa
57 int64_t num_key_value_heads,73 int64_t num_key_value_heads,
58 double scale_value)74 double scale_value)
59{75{
60- check_params(query, key, value);76+ check_params(query, key, value, actual_seq_lengths, actual_seq_lengths_kv, q_input_layout, kv_input_layout);
61 77 
62- // 分配输出 Tensor
63 at::Tensor d_query = npu_preparation::apply_tensor_without_format(query);78 at::Tensor d_query = npu_preparation::apply_tensor_without_format(query);
64 at::Tensor d_key = npu_preparation::apply_tensor_without_format(key);79 at::Tensor d_key = npu_preparation::apply_tensor_without_format(key);
65 at::Tensor d_value = npu_preparation::apply_tensor_without_format(value);80 at::Tensor d_value = npu_preparation::apply_tensor_without_format(value);
66 81 
67- // blockShape 非空,未传时使用默认 [128, 128]82+ // Use the default block shape [128, 128] when block_shape is not specified.
68 static const int64_t kDefaultBlockShape[2] = {128, 128};83 static const int64_t kDefaultBlockShape[2] = {128, 128};
69 const at::IntArrayRef block_shape_value = (block_shape.has_value() && block_shape->size() >= 2)84 const at::IntArrayRef block_shape_value = (block_shape.has_value() && block_shape->size() >= 2)
70 ? *block_shape85 ? *block_shape
71 : at::IntArrayRef(kDefaultBlockShape, 2);86 : at::IntArrayRef(kDefaultBlockShape, 2);
72 87 
73- // 初始化 aclnn 中的暂不支持参数88+ // Initialize aclnn parameters that are not exposed by this PTA API.
74 const at::Tensor atten_mask{nullptr};89 const at::Tensor atten_mask{nullptr};
75 const int64_t mask_type = 0;90 const int64_t mask_type = 0;
76 const int64_t pre_tokens = 2147483647;91 const int64_t pre_tokens = 2147483647;
77 const int64_t next_tokens = 2147483647;92 const int64_t next_tokens = 2147483647;
78 93 
79- // 获取到 layout 的指针,直接传给 aclnn 接口,供其获取字符串94+ // Pass layout strings through to aclnn. aclnn owns the final validation of supported layouts.
80 char *q_input_layout_ptr = const_cast<char *>(q_input_layout.data());95 char *q_input_layout_ptr = const_cast<char *>(q_input_layout.data());
81 char *kv_input_layout_ptr = const_cast<char *>(kv_input_layout.data());96 char *kv_input_layout_ptr = const_cast<char *>(kv_input_layout.data());
82 97 
83- // 调用alcnn接口98+ // Call aclnn API.
84 EXEC_NPU_NO_FORMAT_CHECK_CMD(99 EXEC_NPU_NO_FORMAT_CHECK_CMD(
85 aclnnBlockSparseAttentionGrad,100 aclnnBlockSparseAttentionGrad,
86 d_out, query, key, value,101 d_out, query, key, value,
@@ -92,7 +107,7 @@ std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_block_sparse_attention_backwa
92 pre_tokens, next_tokens,107 pre_tokens, next_tokens,
93 d_query, d_key, d_value);108 d_query, d_key, d_value);
94 109 
95- // 返回结果110+ // Return gradients.
96 return std::make_tuple(d_query, d_key, d_value);111 return std::make_tuple(d_query, d_key, d_value);
97}112}
98}113}
Mtest/test_custom_ops/test_npu_block_sparse_attention_backward.py+338-8
@@ -1,21 +1,20 @@
1"""1"""
2npu_block_sparse_attention_backward 反向算子单测。2npu_block_sparse_attention_backward 反向算子单测。
3 3 
4-当前反向算子仅支持 BNSD MHA 场景(BNSD 布局,num_heads == num_kv_heads)。
5-与正向解耦用例:attention_out、softmax_lse 由 CPU 正向标杆生成,仅反向在 NPU 执行并与 CPU 反向标杆对比。
6- 
7测试场景覆盖:4测试场景覆盖:
8- BNSD MHA(num_heads == num_kv_heads)5- BNSD MHA(num_heads == num_kv_heads)
6+- TND GQA(num_heads > num_kv_heads,多个 Q head 共享同一个 KV head)
7+- TND 变长序列(NPU 接口传每 batch 实际长度,CPU 标杆使用累计 offset 切分 TND)
9- head_dim=128(算子限制 head_dim <= 128,用例覆盖边界)8- head_dim=128(算子限制 head_dim <= 128,用例覆盖边界)
10- float16、bfloat16 数据类型9- float16、bfloat16 数据类型
11- 多块稀疏(block_shape=[8,128])10- 多块稀疏(block_shape=[8,128])
12-- 稀疏掩码(每个 q_block 仅 attend 一个 kv_block)11+- 稀疏掩码(每个 q_block 仅 attend 部分 kv_block)
13- 正反向算子同时调用(forward 输出作为 backward 输入)12- 正反向算子同时调用(forward 输出作为 backward 输入)
14-- autograd 前反向绑定13+- autograd 端到端前反向绑定
15 14 
16-精度说明:为保障 NPU 与 CPU 标杆公平对比,所有与 CPU 对比用例使用 CPU 正向标杆生成15+精度说明:为保障 NPU 与 CPU 标杆公平对比,显式 backward CPU 对比用例使用 CPU 正向标杆生成
17-attention_out、softmax_lse,确保双方使用同一 P 矩阵。若使用 NPU 正向输出,NPU 的 fp16 与 CPU 的 fp3216+attention_out、softmax_lse,确保双方使用同一 P 矩阵。TND/GQA autograd 例走 NPU 正向与 CPU
18-P 存在差异,会导致梯度对比间歇性超阈值17+backward 标杆对比,验证实际网络中 .backward() 路径可用
19"""18"""
20 19 
21import gc20import gc
@@ -31,6 +30,13 @@ DTYPE = torch.float16
31B, S, N, D = 2, 32, 8, 128 # head_dim=128,满足 head_dim <= 128 算子限制30B, S, N, D = 2, 32, 8, 128 # head_dim=128,满足 head_dim <= 128 算子限制
32NUM_KV_HEADS = 831NUM_KV_HEADS = 8
33BLOCK_SHAPE = [128, 128]32BLOCK_SHAPE = [128, 128]
33+TND_GQA_B = 13
34+TND_GQA_NUM_HEADS = 12
35+TND_GQA_NUM_KV_HEADS = 3
36+TND_GQA_HEAD_DIM = 128
37+TND_GQA_BLOCK_SHAPE = [128, 128]
38+TND_GQA_Q_LENGTHS = [355, 17, 41, 83, 129, 211, 7, 53, 97, 151, 233, 301, 19]
39+TND_GQA_KV_LENGTHS = [533, 23, 67, 101, 173, 257, 11, 89, 137, 199, 281, 349, 31]
34 40 
35 41 
36def _softmax_np(x):42def _softmax_np(x):
@@ -145,6 +151,155 @@ def cpu_block_sparse_attention_backward_bnsd(
145 )151 )
146 152 
147 153 
154+def _make_cumulative_seq_lengths(lengths):
155+ seq_lengths = [0]
156+ for length in lengths:
157+ seq_lengths.append(seq_lengths[-1] + length)
158+ return seq_lengths
159+ 
160+ 
161+def _make_tnd_gqa_case(
162+ full_mask=True,
163+ num_heads=TND_GQA_NUM_HEADS,
164+ num_kv_heads=TND_GQA_NUM_KV_HEADS,
165+ head_dim=TND_GQA_HEAD_DIM,
166+ block_shape=TND_GQA_BLOCK_SHAPE,
167+ q_lengths=TND_GQA_Q_LENGTHS,
168+ kv_lengths=TND_GQA_KV_LENGTHS,
169+):
170+ batch = len(q_lengths)
171+ assert batch == len(kv_lengths)
172+ assert num_heads % num_kv_heads == 0
173+ scale_value = 1.0 / math.sqrt(head_dim)
174+ actual_seq_offsets = _make_cumulative_seq_lengths(q_lengths)
175+ actual_seq_offsets_kv = _make_cumulative_seq_lengths(kv_lengths)
176+ total_q = actual_seq_offsets[-1]
177+ total_kv = actual_seq_offsets_kv[-1]
178+ 
179+ query = torch.randn(total_q, num_heads, head_dim, dtype=DTYPE)
180+ key = torch.randn(total_kv, num_kv_heads, head_dim, dtype=DTYPE)
181+ value = torch.randn(total_kv, num_kv_heads, head_dim, dtype=DTYPE)
182+ d_out = torch.randn(total_q, num_heads, head_dim, dtype=DTYPE)
183+ 
184+ max_q = max(q_lengths)
185+ max_kv = max(kv_lengths)
186+ ceil_q = (max_q + block_shape[0] - 1) // block_shape[0]
187+ ceil_kv = (max_kv + block_shape[1] - 1) // block_shape[1]
188+ if full_mask:
189+ block_sparse_mask = torch.ones(batch, num_heads, ceil_q, ceil_kv, dtype=torch.int8)
190+ else:
191+ block_sparse_mask = torch.zeros(batch, num_heads, ceil_q, ceil_kv, dtype=torch.int8)
192+ for b in range(batch):
193+ valid_ceil_kv = (kv_lengths[b] + block_shape[1] - 1) // block_shape[1]
194+ for n in range(num_heads):
195+ for q_block in range(ceil_q):
196+ block_sparse_mask[b, n, q_block, (q_block + b + n) % valid_ceil_kv] = 1
197+ if valid_ceil_kv > 1 and (q_block + n) % 2 == 0:
198+ block_sparse_mask[b, n, q_block, (q_block + b + n + 1) % valid_ceil_kv] = 1
199+ 
200+ return {
201+ "query": query,
202+ "key": key,
203+ "value": value,
204+ "d_out": d_out,
205+ "block_sparse_mask": block_sparse_mask,
206+ "block_shape": block_shape,
207+ "actual_seq_lengths": q_lengths,
208+ "actual_seq_lengths_kv": kv_lengths,
209+ "actual_seq_offsets": actual_seq_offsets,
210+ "actual_seq_offsets_kv": actual_seq_offsets_kv,
211+ "num_kv_heads": num_kv_heads,
212+ "scale_value": scale_value,
213+ }
214+ 
215+ 
216+def _expand_block_sparse_mask(mask, block_shape, q_len, kv_len):
217+ block_x, block_y = int(block_shape[0]), int(block_shape[1])
218+ return mask.repeat_interleave(block_x, dim=0).repeat_interleave(block_y, dim=1)[:q_len, :kv_len].bool()
219+ 
220+ 
221+def cpu_block_sparse_attention_tnd_gqa_with_lse(
222+ query, key, value, block_sparse_mask, block_shape, scale_value, actual_seq_offsets, actual_seq_offsets_kv
223+):
224+ query_f = query.cpu().to(torch.float32)
225+ key_f = key.cpu().to(torch.float32)
226+ value_f = value.cpu().to(torch.float32)
227+ mask = block_sparse_mask.cpu()
228+ total_q, num_heads, head_dim = query_f.shape
229+ num_kv_heads = key_f.shape[1]
230+ group_size = num_heads // num_kv_heads
231+ attention_out = torch.zeros(total_q, num_heads, head_dim, dtype=torch.float32)
232+ softmax_lse = torch.zeros(total_q, num_heads, 1, dtype=torch.float32)
233+ 
234+ for b in range(len(actual_seq_offsets) - 1):
235+ q_start, q_end = actual_seq_offsets[b], actual_seq_offsets[b + 1]
236+ kv_start, kv_end = actual_seq_offsets_kv[b], actual_seq_offsets_kv[b + 1]
237+ q_len = q_end - q_start
238+ kv_len = kv_end - kv_start
239+ for n in range(num_heads):
240+ kv_head = n // group_size
241+ q_b = query_f[q_start:q_end, n, :]
242+ k_b = key_f[kv_start:kv_end, kv_head, :]
243+ v_b = value_f[kv_start:kv_end, kv_head, :]
244+ scores = torch.matmul(q_b, k_b.transpose(0, 1)) * float(scale_value)
245+ full_mask = _expand_block_sparse_mask(mask[b, n], block_shape, q_len, kv_len)
246+ scores = scores.masked_fill(~full_mask, -1e10)
247+ probs = torch.softmax(scores, dim=-1)
248+ attention_out[q_start:q_end, n, :] = torch.matmul(probs, v_b)
249+ softmax_lse[q_start:q_end, n, 0] = torch.logsumexp(scores, dim=-1)
250+ 
251+ return attention_out.to(query.dtype), softmax_lse
252+ 
253+ 
254+def cpu_block_sparse_attention_backward_tnd_gqa(
255+ query, key, value, d_out, block_sparse_mask, block_shape, scale_value, actual_seq_offsets, actual_seq_offsets_kv
256+):
257+ query_f = query.cpu().to(torch.float32)
258+ key_f = key.cpu().to(torch.float32)
259+ value_f = value.cpu().to(torch.float32)
260+ d_out_f = d_out.cpu().to(torch.float32)
261+ mask = block_sparse_mask.cpu()
262+ total_q, num_heads, head_dim = query_f.shape
263+ total_kv, num_kv_heads, _ = key_f.shape
264+ group_size = num_heads // num_kv_heads
265+ d_query = torch.zeros(total_q, num_heads, head_dim, dtype=torch.float32)
266+ d_key = torch.zeros(total_kv, num_kv_heads, head_dim, dtype=torch.float32)
267+ d_value = torch.zeros(total_kv, num_kv_heads, head_dim, dtype=torch.float32)
268+ 
269+ for b in range(len(actual_seq_offsets) - 1):
270+ q_start, q_end = actual_seq_offsets[b], actual_seq_offsets[b + 1]
271+ kv_start, kv_end = actual_seq_offsets_kv[b], actual_seq_offsets_kv[b + 1]
272+ q_len = q_end - q_start
273+ kv_len = kv_end - kv_start
274+ for n in range(num_heads):
275+ kv_head = n // group_size
276+ q_b = query_f[q_start:q_end, n, :]
277+ k_b = key_f[kv_start:kv_end, kv_head, :]
278+ v_b = value_f[kv_start:kv_end, kv_head, :]
279+ dout_b = d_out_f[q_start:q_end, n, :]
280+ scores = torch.matmul(q_b, k_b.transpose(0, 1)) * float(scale_value)
281+ full_mask = _expand_block_sparse_mask(mask[b, n], block_shape, q_len, kv_len)
282+ scores = scores.masked_fill(~full_mask, -1e10)
283+ probs = torch.softmax(scores, dim=-1)
284+ d_p = torch.matmul(dout_b, v_b.transpose(0, 1))
285+ d_s = (d_p - (d_p * probs).sum(dim=-1, keepdim=True)) * probs * float(scale_value)
286+ d_query[q_start:q_end, n, :] = torch.matmul(d_s, k_b)
287+ d_key[kv_start:kv_end, kv_head, :] += torch.matmul(d_s.transpose(0, 1), q_b)
288+ d_value[kv_start:kv_end, kv_head, :] += torch.matmul(probs.transpose(0, 1), dout_b)
289+ 
290+ return d_query.to(query.dtype), d_key.to(key.dtype), d_value.to(value.dtype)
291+ 
292+ 
293+def _copy_tnd_gqa_case_to_device(case, device):
294+ return {
295+ "query": case["query"].to(device),
296+ "key": case["key"].to(device),
297+ "value": case["value"].to(device),
298+ "d_out": case["d_out"].to(device),
299+ "block_sparse_mask": case["block_sparse_mask"].to(device),
300+ }
301+ 
302+ 
148class TestNPUBlockSparseAttentionBackward(TestCase):303class TestNPUBlockSparseAttentionBackward(TestCase):
149 """Test npu_block_sparse_attention_backward,与 CPU 反向标杆对比."""304 """Test npu_block_sparse_attention_backward,与 CPU 反向标杆对比."""
150 305 
@@ -161,6 +316,10 @@ class TestNPUBlockSparseAttentionBackward(TestCase):
161 torch.npu.empty_cache()316 torch.npu.empty_cache()
162 super().tearDown()317 super().tearDown()
163 318 
319+ def _assert_tnd_gqa_grads_equal(self, cpu_grads, npu_grads):
320+ for cpu_grad, npu_grad in zip(cpu_grads, npu_grads):
321+ self.assertRtolEqual(cpu_grad.cpu().float(), npu_grad.cpu().float(), prec=0.02, prec16=0.02)
322+ 
164 @SkipIfNotGteCANNVersion("9.0.0")323 @SkipIfNotGteCANNVersion("9.0.0")
165 @SupportedDevices(['Ascend910B'])324 @SupportedDevices(['Ascend910B'])
166 def test_npu_block_sparse_attention_backward_bnsd_cpu_compare(self, device="npu"):325 def test_npu_block_sparse_attention_backward_bnsd_cpu_compare(self, device="npu"):
@@ -298,6 +457,177 @@ class TestNPUBlockSparseAttentionBackward(TestCase):
298 self.assertRtolEqual(dk_cpu.cpu().float(), d_key.cpu().float(), prec=0.01, prec16=0.01)457 self.assertRtolEqual(dk_cpu.cpu().float(), d_key.cpu().float(), prec=0.01, prec16=0.01)
299 self.assertRtolEqual(dv_cpu.cpu().float(), d_value.cpu().float(), prec=0.01, prec16=0.01)458 self.assertRtolEqual(dv_cpu.cpu().float(), d_value.cpu().float(), prec=0.01, prec16=0.01)
300 459 
460+ def _run_tnd_gqa_backward_cpu_compare(self, device, full_mask, **case_kwargs):
461+ torch.npu.empty_cache()
462+ case = _make_tnd_gqa_case(full_mask=full_mask, **case_kwargs)
463+ dq_cpu, dk_cpu, dv_cpu = cpu_block_sparse_attention_backward_tnd_gqa(
464+ case["query"], case["key"], case["value"], case["d_out"], case["block_sparse_mask"],
465+ case["block_shape"], case["scale_value"], case["actual_seq_offsets"], case["actual_seq_offsets_kv"])
466+ attention_out_cpu, softmax_lse_cpu = cpu_block_sparse_attention_tnd_gqa_with_lse(
467+ case["query"], case["key"], case["value"], case["block_sparse_mask"],
468+ case["block_shape"], case["scale_value"], case["actual_seq_offsets"], case["actual_seq_offsets_kv"])
469+ 
470+ npu_case = _copy_tnd_gqa_case_to_device(case, device)
471+ 
472+ d_query, d_key, d_value = torch_npu.npu_block_sparse_attention_backward(
473+ npu_case["d_out"], npu_case["query"], npu_case["key"], npu_case["value"],
474+ attention_out_cpu.to(device), softmax_lse_cpu.to(device), npu_case["block_sparse_mask"],
475+ block_shape=case["block_shape"],
476+ actual_seq_lengths=case["actual_seq_lengths"],
477+ actual_seq_lengths_kv=case["actual_seq_lengths_kv"],
478+ q_input_layout="TND", kv_input_layout="TND",
479+ num_key_value_heads=case["num_kv_heads"],
480+ scale_value=case["scale_value"],
481+ )
482+ torch.npu.synchronize()
483+ self._assert_tnd_gqa_grads_equal((dq_cpu, dk_cpu, dv_cpu), (d_query, d_key, d_value))
484+ 
485+ @SkipIfNotGteCANNVersion("9.0.0")
486+ @SupportedDevices(['Ascend910B'])
487+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
488+ def test_npu_block_sparse_attention_backward_tnd_gqa_full_mask_cpu_compare(self, device="npu"):
489+ """Backward TND GQA with full block sparse mask, compared with CPU golden."""
490+ self._run_tnd_gqa_backward_cpu_compare(device, full_mask=True)
491+ 
492+ @SkipIfNotGteCANNVersion("9.0.0")
493+ @SupportedDevices(['Ascend910B'])
494+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
495+ def test_npu_block_sparse_attention_backward_tnd_gqa_sparse_mask_cpu_compare(self, device="npu"):
496+ """Backward TND GQA with sparse block mask, compared with CPU golden."""
497+ self._run_tnd_gqa_backward_cpu_compare(device, full_mask=False)
498+ 
499+ @SkipIfNotGteCANNVersion("9.0.0")
500+ @SupportedDevices(['Ascend910B'])
501+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
502+ def test_npu_block_sparse_attention_backward_tnd_gqa_single_batch_cpu_compare(self, device="npu"):
503+ """Backward TND GQA with a single batch, compared with CPU golden."""
504+ self._run_tnd_gqa_backward_cpu_compare(
505+ device, full_mask=True,
506+ num_heads=4, num_kv_heads=2,
507+ q_lengths=[127], kv_lengths=[191])
508+ 
509+ @SkipIfNotGteCANNVersion("9.0.0")
510+ @SupportedDevices(['Ascend910B'])
511+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
512+ def test_npu_block_sparse_attention_backward_tnd_gqa_uneven_seq_lengths_cpu_compare(self, device="npu"):
513+ """Backward TND GQA with uneven variable lengths, compared with CPU golden."""
514+ self._run_tnd_gqa_backward_cpu_compare(
515+ device, full_mask=False,
516+ num_heads=4, num_kv_heads=2,
517+ q_lengths=[1, 64, 129, 17],
518+ kv_lengths=[128, 3, 257, 65])
519+ 
520+ @SkipIfNotGteCANNVersion("9.0.0")
521+ @SupportedDevices(['Ascend910B'])
522+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
523+ def test_npu_block_sparse_attention_backward_tnd_gqa_group_size_4_cpu_compare(self, device="npu"):
524+ """Backward TND GQA with group size 4, compared with CPU golden."""
525+ self._run_tnd_gqa_backward_cpu_compare(
526+ device, full_mask=False,
527+ num_heads=8, num_kv_heads=2,
528+ q_lengths=[96, 137],
529+ kv_lengths=[111, 259])
530+ 
531+ @SkipIfNotGteCANNVersion("9.0.0")
532+ @SupportedDevices(['Ascend910B'])
533+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
534+ def test_npu_block_sparse_attention_backward_tnd_mqa_cpu_compare(self, device="npu"):
535+ """Backward TND MQA with one shared KV head, compared with CPU golden."""
536+ self._run_tnd_gqa_backward_cpu_compare(
537+ device, full_mask=False,
538+ num_heads=8, num_kv_heads=1,
539+ q_lengths=[65, 130],
540+ kv_lengths=[129, 33])
541+ 
542+ @SkipIfNotGteCANNVersion("9.0.0")
543+ @SupportedDevices(['Ascend910B'])
544+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
545+ def test_npu_block_sparse_attention_backward_tnd_gqa_non_128_tail_block_cpu_compare(self, device="npu"):
546+ """Backward TND GQA with non-default Q block and tail blocks, compared with CPU golden."""
547+ self._run_tnd_gqa_backward_cpu_compare(
548+ device, full_mask=False,
549+ num_heads=4, num_kv_heads=2,
550+ block_shape=[64, 128],
551+ q_lengths=[65, 127, 3],
552+ kv_lengths=[129, 11, 64])
553+ 
554+ @SkipIfNotGteCANNVersion("9.0.0")
555+ @SupportedDevices(['Ascend910B'])
556+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
557+ def test_npu_block_sparse_attention_backward_tnd_actual_seq_lengths_required(self, device="npu"):
558+ """TND backward requires corresponding actual sequence lengths."""
559+ case = _make_tnd_gqa_case(
560+ full_mask=True,
561+ num_heads=4, num_kv_heads=2,
562+ q_lengths=[16], kv_lengths=[16])
563+ attention_out_cpu, softmax_lse_cpu = cpu_block_sparse_attention_tnd_gqa_with_lse(
564+ case["query"], case["key"], case["value"], case["block_sparse_mask"],
565+ case["block_shape"], case["scale_value"], case["actual_seq_offsets"], case["actual_seq_offsets_kv"])
566+ npu_case = _copy_tnd_gqa_case_to_device(case, device)
567+ 
568+ common_args = (
569+ npu_case["d_out"], npu_case["query"], npu_case["key"], npu_case["value"],
570+ attention_out_cpu.to(device), softmax_lse_cpu.to(device), npu_case["block_sparse_mask"])
571+ common_kwargs = {
572+ "block_shape": case["block_shape"],
573+ "q_input_layout": "TND",
574+ "kv_input_layout": "TND",
575+ "num_key_value_heads": case["num_kv_heads"],
576+ "scale_value": case["scale_value"],
577+ }
578+ with self.assertRaisesRegex(RuntimeError, "actual_seq_lengths must be specified"):
579+ torch_npu.npu_block_sparse_attention_backward(
580+ *common_args,
581+ actual_seq_lengths=None,
582+ actual_seq_lengths_kv=case["actual_seq_lengths_kv"],
583+ **common_kwargs)
584+ with self.assertRaisesRegex(RuntimeError, "actual_seq_lengths_kv must be specified"):
585+ torch_npu.npu_block_sparse_attention_backward(
586+ *common_args,
587+ actual_seq_lengths=case["actual_seq_lengths"],
588+ actual_seq_lengths_kv=None,
589+ **common_kwargs)
590+ 
591+ @SkipIfNotGteCANNVersion("9.0.0")
592+ @SupportedDevices(['Ascend910B'])
593+ @unittest.skip("Skip until gate CANN version supports BlockSparseAttention TND/GQA backward.")
594+ def test_npu_block_sparse_attention_backward_tnd_gqa_autograd_cpu_compare(self, device="npu"):
595+ """End-to-end autograd TND GQA path, compared with CPU backward golden."""
596+ torch.npu.empty_cache()
597+ case = _make_tnd_gqa_case(full_mask=False)
598+ dq_cpu, dk_cpu, dv_cpu = cpu_block_sparse_attention_backward_tnd_gqa(
599+ case["query"], case["key"], case["value"], case["d_out"], case["block_sparse_mask"],
600+ case["block_shape"], case["scale_value"], case["actual_seq_offsets"], case["actual_seq_offsets_kv"])
601+ 
602+ npu_case = _copy_tnd_gqa_case_to_device(case, device)
603+ query = npu_case["query"]
604+ key = npu_case["key"]
605+ value = npu_case["value"]
606+ query.requires_grad = True
607+ key.requires_grad = True
608+ value.requires_grad = True
609+ 
610+ attention_out, _ = torch_npu.npu_block_sparse_attention(
611+ query, key, value, npu_case["block_sparse_mask"], case["block_shape"],
612+ q_input_layout="TND", kv_input_layout="TND",
613+ num_key_value_heads=case["num_kv_heads"],
614+ scale_value=case["scale_value"],
615+ inner_precise=1,
616+ actual_seq_lengths=case["actual_seq_lengths"],
617+ actual_seq_lengths_kv=case["actual_seq_lengths_kv"],
618+ softmax_lse_flag=1,
619+ )
620+ attention_out.backward(gradient=npu_case["d_out"])
621+ torch.npu.synchronize()
622+ 
623+ self.assertIsNotNone(query.grad)
624+ self.assertIsNotNone(key.grad)
625+ self.assertIsNotNone(value.grad)
626+ self.assertEqual(query.grad.shape, query.shape)
627+ self.assertEqual(key.grad.shape, key.shape)
628+ self.assertEqual(value.grad.shape, value.shape)
629+ self._assert_tnd_gqa_grads_equal((dq_cpu, dk_cpu, dv_cpu), (query.grad, key.grad, value.grad))
630+ 
301 @SkipIfNotGteCANNVersion("9.0.0")631 @SkipIfNotGteCANNVersion("9.0.0")
302 @SupportedDevices(['Ascend910B'])632 @SupportedDevices(['Ascend910B'])
303 def test_npu_block_sparse_attention_autograd_backward(self, device="npu"):633 def test_npu_block_sparse_attention_autograd_backward(self, device="npu"):