已合并
Support block sparse attention grad TND GQA #4945
fgd_dragon创建于 5月14日
Support block sparse attention grad TND GQA #4945
已合并
共 4 个文件变更+371-21
| @@ -541,6 +541,9 @@ query 的 head 数 N1 与 key/value 的 head 数 N2,需满足 N1 >= N2 且 N1 | |||
| 541 | block_sparse_mask 必传,且shape必须为[batch, headNum, ceilDiv(maxQS, blockShapeX), ceilDiv(maxKVS, blockShapeY)]. | 541 | block_sparse_mask 必传,且shape必须为[batch, headNum, ceilDiv(maxQS, blockShapeX), ceilDiv(maxKVS, blockShapeY)]. |
| 542 | block_shape 必传,必须包含至少两个元素[blockShapeX, blockShapeY],且值必须大于0;blockShapeY 必须为 128 的倍数。 | 542 | block_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=64 或 128; 反向路径当前支持 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 数. | ||
| 544 | inner_precise 仅支持 0(表示float32 softmax) 或 1(表示float16 softmax);当 query/key/value 为 bfloat16 时,仅支持 0. | 547 | inner_precise 仅支持 0(表示float32 softmax) 或 1(表示float16 softmax);当 query/key/value 为 bfloat16 时,仅支持 0. |
| 545 | 548 | ||
| 546 | 支持的PyTorch版本 | 549 | 支持的PyTorch版本 |
| @@ -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 | ||
| @@ -23,24 +23,40 @@ const int64_t MAX_HEAD_DIM = 128; | |||
| 23 | using npu_preparation = at_npu::native::OpPreparation; | 23 | using npu_preparation = at_npu::native::OpPreparation; |
| 24 | 24 | ||
| 25 | 25 | ||
| 26 | -// 入参检查 | 26 | +// Validate input parameters. |
| 27 | static void check_params(const at::Tensor &query, | 27 | static 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 不能超过 128 | 40 | + // 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. |
| 44 | std::tuple<at::Tensor, at::Tensor, at::Tensor> npu_block_sparse_attention_backward( | 60 | std::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_shape | 85 | ? *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 | } |
| @@ -1,21 +1,20 @@ | |||
| 1 | """ | 1 | """ |
| 2 | npu_block_sparse_attention_backward 反向算子单测。 | 2 | npu_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 的 fp32 | 16 | +attention_out、softmax_lse,确保双方使用同一 P 矩阵。TND/GQA autograd 用例走 NPU 正向并与 CPU |
| 18 | -P 存在差异,会导致梯度对比间歇性超阈值。 | 17 | +backward 标杆对比,验证实际网络中 .backward() 路径可用。 |
| 19 | """ | 18 | """ |
| 20 | 19 | ||
| 21 | import gc | 20 | import gc |
| @@ -31,6 +30,13 @@ DTYPE = torch.float16 | |||
| 31 | B, S, N, D = 2, 32, 8, 128 # head_dim=128,满足 head_dim <= 128 算子限制 | 30 | B, S, N, D = 2, 32, 8, 128 # head_dim=128,满足 head_dim <= 128 算子限制 |
| 32 | NUM_KV_HEADS = 8 | 31 | NUM_KV_HEADS = 8 |
| 33 | BLOCK_SHAPE = [128, 128] | 32 | BLOCK_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 | ||
| 36 | def _softmax_np(x): | 42 | def _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 | + | ||
| 148 | class TestNPUBlockSparseAttentionBackward(TestCase): | 303 | class 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 | 323 | ||
| 165 | 324 | ||
| 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 | + | ||
| 486 | + | ||
| 487 | + | ||
| 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 | + | ||
| 493 | + | ||
| 494 | + | ||
| 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 | + | ||
| 500 | + | ||
| 501 | + | ||
| 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 | + | ||
| 510 | + | ||
| 511 | + | ||
| 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 | + | ||
| 521 | + | ||
| 522 | + | ||
| 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 | + | ||
| 532 | + | ||
| 533 | + | ||
| 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 | + | ||
| 543 | + | ||
| 544 | + | ||
| 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 | + | ||
| 555 | + | ||
| 556 | + | ||
| 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 | + | ||
| 592 | + | ||
| 593 | + | ||
| 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 | 631 | ||
| 302 | 632 | ||
| 303 | def test_npu_block_sparse_attention_autograd_backward(self, device="npu"): | 633 | def test_npu_block_sparse_attention_autograd_backward(self, device="npu"): |