fix: synchronize final decode async copy
- wait for the final split-KV tile before reading shared memory - cover long decode with production context capacity
This commit is contained in:
@@ -121,7 +121,10 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
load_tile(ti_begin + i, i);
|
load_tile(ti_begin + i, i);
|
||||||
|
|
||||||
for (int it = 0; it < ntiles; it++) {
|
for (int it = 0; it < ntiles; it++) {
|
||||||
cp_async_wait_group<STAGES - 1>();
|
if (it + 1 == ntiles)
|
||||||
|
cp_async_wait_group<0>();
|
||||||
|
else
|
||||||
|
cp_async_wait_group<STAGES - 1>();
|
||||||
__syncwarp();
|
__syncwarp();
|
||||||
process_tile(it, it & (STAGES - 1));
|
process_tile(it, it & (STAGES - 1));
|
||||||
__syncwarp();
|
__syncwarp();
|
||||||
|
|||||||
@@ -127,14 +127,15 @@ inline void print_paged_row(const char* cfg, float max_err, bool pass) {
|
|||||||
// ======================================================================
|
// ======================================================================
|
||||||
template <int HEAD_DIM>
|
template <int HEAD_DIM>
|
||||||
static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
|
static int run_decode_test(int B, int Hq, int Hkv, int max_seq,
|
||||||
int causal, int seed) {
|
int causal, int seed, int context_capacity = 0,
|
||||||
|
int fixed_seq_len = 0) {
|
||||||
// Variable seq_lens per request
|
// Variable seq_lens per request
|
||||||
srand(seed);
|
srand(seed);
|
||||||
std::vector<int> seq_lens(B);
|
std::vector<int> seq_lens(B);
|
||||||
for (int b = 0; b < B; b++)
|
for (int b = 0; b < B; b++)
|
||||||
seq_lens[b] = 8 + rand() % (max_seq - 8);
|
seq_lens[b] = fixed_seq_len ? fixed_seq_len : 8 + rand() % (max_seq - 8);
|
||||||
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
|
int max_sl = *std::max_element(seq_lens.begin(), seq_lens.end());
|
||||||
int max_ctx = max_sl + 16;
|
int max_ctx = context_capacity ? context_capacity : max_sl + 16;
|
||||||
|
|
||||||
int pool_size = B * max_ctx;
|
int pool_size = B * max_ctx;
|
||||||
int num_reqs = B + 4;
|
int num_reqs = B + 4;
|
||||||
@@ -844,6 +845,9 @@ int main() {
|
|||||||
fail += run_decode_test<256>(1, 2, 1, 256, 0, 9);
|
fail += run_decode_test<256>(1, 2, 1, 256, 0, 9);
|
||||||
fail += run_decode_test<128>(16, 32, 4, 2048, 0, 10);
|
fail += run_decode_test<128>(16, 32, 4, 2048, 0, 10);
|
||||||
fail += run_decode_test<128>(32, 32, 4, 1024, 0, 11);
|
fail += run_decode_test<128>(32, 32, 4, 1024, 0, 11);
|
||||||
|
// Production keeps a fixed 32768-wide request table. This forces 32
|
||||||
|
// splits, so seq_len > 512 gives each split multiple cp.async tiles.
|
||||||
|
fail += run_decode_test<64>(1, 24, 4, 1100, 0, 12, 32768, 1100);
|
||||||
|
|
||||||
// Decode with 2D mask (regression: mixed seq_lens + HasMask)
|
// Decode with 2D mask (regression: mixed seq_lens + HasMask)
|
||||||
fail += run_decode_mask_test<128>(2, 8, 2, 256, 30);
|
fail += run_decode_mask_test<128>(2, 8, 2, 256, 30);
|
||||||
|
|||||||
Reference in New Issue
Block a user