refactor: separate KV token address resolution
This commit is contained in:
@@ -57,7 +57,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
|
|||||||
int s = i / p.head_dim;
|
int s = i / p.head_dim;
|
||||||
int d_dim = i % p.head_dim;
|
int d_dim = i % p.head_dim;
|
||||||
int kc = chunk_start + s;
|
int kc = chunk_start + s;
|
||||||
KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true);
|
int token = KV::resolve_token(p, kctx, kc, true);
|
||||||
|
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim);
|
||||||
k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
||||||
v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
|
v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -73,7 +73,8 @@ __global__ void attn_decode_split_kv_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
int kc = kv0 + r;
|
int kc = kv0 + r;
|
||||||
bool valid = kc < seq_len;
|
bool valid = kc < seq_len;
|
||||||
KVAddr a = KV::kv_addr(p, kctx, kc, d, valid);
|
int token = KV::resolve_token(p, kctx, kc, valid);
|
||||||
|
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
|
||||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||||
cp_async_16_pred(&dK[off], a.k, a.valid);
|
cp_async_16_pred(&dK[off], a.k, a.valid);
|
||||||
cp_async_16_pred(&dV[off], a.v, a.valid);
|
cp_async_16_pred(&dV[off], a.v, a.valid);
|
||||||
|
|||||||
@@ -104,10 +104,18 @@ struct ContigKV {
|
|||||||
c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride;
|
c.kv_base = batch * p.kv_b_stride + kv_head * p.kv_h_stride;
|
||||||
return c;
|
return c;
|
||||||
}
|
}
|
||||||
HOST_DEV_FORCEINLINE KVAddr kv_addr(
|
HOST_DEV_FORCEINLINE int resolve_token(
|
||||||
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) {
|
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
|
||||||
const int g_off = c.kv_base + kc * p.kv_l_stride + d * p.kv_d_stride;
|
return valid ? kc : -1;
|
||||||
return {&p.k_ptr[g_off], &p.v_ptr[g_off], valid};
|
}
|
||||||
|
HOST_DEV_FORCEINLINE KVAddr kv_addr_from_token(
|
||||||
|
const AttentionParams<bf16>& p, const KVContext& c, int token, int d) {
|
||||||
|
const bool valid = token >= 0;
|
||||||
|
const int safe_token = valid ? token : 0;
|
||||||
|
const int64_t gmem_off = (int64_t)c.kv_base
|
||||||
|
+ (int64_t)safe_token * p.kv_l_stride
|
||||||
|
+ (int64_t)d * p.kv_d_stride;
|
||||||
|
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -175,12 +183,16 @@ struct PagedKV {
|
|||||||
c.head_off = (int64_t)kv_head * HEAD_DIM;
|
c.head_off = (int64_t)kv_head * HEAD_DIM;
|
||||||
return c;
|
return c;
|
||||||
}
|
}
|
||||||
HOST_DEV_FORCEINLINE KVAddr kv_addr(
|
HOST_DEV_FORCEINLINE int resolve_token(
|
||||||
const AttentionParams<bf16>& p, const KVContext& c, int kc, int d, bool valid) {
|
const AttentionParams<bf16>& p, const KVContext& c, int kc, bool valid) {
|
||||||
const int slot = valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : 0;
|
return valid ? p.req_to_token[c.req_idx * c.rtt_stride + kc] : -1;
|
||||||
const bool ok = valid && (slot >= 0);
|
}
|
||||||
const int64_t gmem_off = (int64_t)slot * c.pool_stride + c.head_off + d;
|
HOST_DEV_FORCEINLINE KVAddr kv_addr_from_token(
|
||||||
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], ok};
|
const AttentionParams<bf16>& p, const KVContext& c, int slot, int d) {
|
||||||
|
const bool valid = slot >= 0;
|
||||||
|
const int safe_slot = valid ? slot : 0;
|
||||||
|
const int64_t gmem_off = (int64_t)safe_slot * c.pool_stride + c.head_off + d;
|
||||||
|
return {&p.k_ptr[gmem_off], &p.v_ptr[gmem_off], valid};
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
|
|
||||||
|
|||||||
@@ -90,7 +90,8 @@ __global__ void attn_prefill_split_q_kernel_t(AttentionParams<bf16> p) {
|
|||||||
int s = i / HEAD_DIM;
|
int s = i / HEAD_DIM;
|
||||||
int d_dim = i % HEAD_DIM;
|
int d_dim = i % HEAD_DIM;
|
||||||
int kc = kv0 + s;
|
int kc = kv0 + s;
|
||||||
KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true);
|
int token = KV::resolve_token(p, kctx, kc, true);
|
||||||
|
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d_dim);
|
||||||
sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
sK[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
|
||||||
sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
|
sV[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -83,7 +83,8 @@ __global__ void attn_prefill_split_q_mma_kernel(AttentionParams<bf16> p) {
|
|||||||
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
int r = i / Traits::HEAD_DIM, d = i % Traits::HEAD_DIM;
|
||||||
int kc = kv0 + r;
|
int kc = kv0 + r;
|
||||||
bool valid = kc < seq_len;
|
bool valid = kc < seq_len;
|
||||||
KVAddr a = KV::kv_addr(p, kctx, kc, d, valid);
|
int token = KV::resolve_token(p, kctx, kc, valid);
|
||||||
|
KVAddr a = KV::kv_addr_from_token(p, kctx, token, d);
|
||||||
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
int off = r * Traits::LD + swiz_col(d, r, Traits::SWIZ_MASK);
|
||||||
cp_async_16_pred(&dK[off], a.k, a.valid);
|
cp_async_16_pred(&dK[off], a.k, a.valid);
|
||||||
cp_async_16_pred(&dV[off], a.v, a.valid);
|
cp_async_16_pred(&dV[off], a.v, a.valid);
|
||||||
|
|||||||
Reference in New Issue
Block a user