在 speculative mode 下,测试 Attention 在传入 mask_offset 下的功能
(self)
| 408 | ) |
| 409 | |
| 410 | def test_mask_offset(self): |
| 411 | """ |
| 412 | 在 speculative mode 下,测试 Attention 在传入 mask_offset 下的功能 |
| 413 | """ |
| 414 | prefill_len = 8192 |
| 415 | dec_len_q = 5 |
| 416 | total_len = prefill_len + dec_len_q |
| 417 | mask = paddle.tril(paddle.ones((self.bsz, dec_len_q, total_len), dtype="float32"), diagonal=prefill_len) |
| 418 | mask = paddle.where(mask == 1, paddle.zeros_like(mask), paddle.full_like(mask, fill_value=float("-inf"))) |
| 419 | self.run_append_c16_attention(prefill_len, 0, True, use_qknorm=self.use_qknorm) |
| 420 | |
| 421 | mask_offset = paddle.tile( |
| 422 | paddle.tensor( |
| 423 | [0, prefill_len + 1, 0, prefill_len + 2, 0, prefill_len + 3, 0, prefill_len + 4, 0, prefill_len + 5], |
| 424 | dtype="int32", |
| 425 | ), |
| 426 | [self.bsz], |
| 427 | ).astype("int32") |
| 428 | dec_out = self.run_append_c16_attention( |
| 429 | dec_len_q, prefill_len, False, use_qknorm=self.use_qknorm, mask_offset=mask_offset |
| 430 | ) |
| 431 | |
| 432 | ref_out = self.ref_attention(self.CURRENT_Q[0], self.TOTAL_K, self.TOTAL_V, mask, use_qknorm=self.use_qknorm) |
| 433 | np.testing.assert_allclose( |
| 434 | ref_out.astype("float32").numpy(), dec_out.astype("float32").numpy(), rtol=1e-03, atol=5e-03 |
| 435 | ) |
| 436 | |
| 437 | def test_consistency_with_multi_tokens(self): |
| 438 | """ |
nothing calls this directly
no test coverage detected