[JIT Kernel][DSv4] Optimize epilogue of c128 (#26671)
Co-authored-by: Yeu Jaung <jpppppppppppppppppppppppp@users.noreply.github.com> Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Yeu Jaung
Claude Opus 5
parent
accc51c6db
commit
7331287c1c
@@ -37,17 +37,37 @@ using PlanD = device::compress::DecodePlan;
|
|||||||
using PlanC = device::compress::CompressPlan;
|
using PlanC = device::compress::CompressPlan;
|
||||||
using PlanW = device::compress::WritePlan;
|
using PlanW = device::compress::WritePlan;
|
||||||
|
|
||||||
/// \brief Each thread will handle this many elements (split along head_dim)
|
/// \brief Block-level configuration shared by all C128 kernels (independent of head_dim).
|
||||||
constexpr int32_t kTileElements = 2;
|
struct C128Config {
|
||||||
/// \brief Each warp will handle this many elements (split along 128)
|
/// \brief Each thread loads/stores this many elements (split along head_dim).
|
||||||
constexpr int32_t kElementsPerWarp = 8;
|
static constexpr int32_t kTileElements = 2;
|
||||||
constexpr uint32_t kNumWarps = 128 / kElementsPerWarp;
|
/// \brief Each warp handles this many elements (split along the softmax dim of 128).
|
||||||
constexpr uint32_t kBlockSize = device::kWarpThreads * kNumWarps;
|
static constexpr int32_t kElementsPerWarp = 8;
|
||||||
constexpr uint32_t kWriteBlockSize = 128; // one warp per write
|
static constexpr uint32_t kNumWarps = 128 / kElementsPerWarp;
|
||||||
|
static constexpr uint32_t kBlockSize = device::kWarpThreads * kNumWarps;
|
||||||
|
/// \brief Block size used by the prefill write kernel (one warp per write plan tile).
|
||||||
|
static constexpr uint32_t kWriteBlockSize = 128;
|
||||||
|
static constexpr uint32_t kNumWriteWarps = kWriteBlockSize / device::kWarpThreads;
|
||||||
|
/// \brief Per-warp scratch buffer used to stage partial softmax results before the
|
||||||
|
/// final block-level reduction. Padded to avoid bank conflicts.
|
||||||
|
using SharedStorage = device::AlignedVector<float, kTileElements>;
|
||||||
|
using SharedBuffer = SharedStorage[kNumWarps][device::kWarpThreads];
|
||||||
|
};
|
||||||
|
|
||||||
|
template <int64_t kHeadDim_>
|
||||||
|
struct C128Trait : public C128Config {
|
||||||
|
static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64
|
||||||
|
static constexpr int64_t kHeadDim = kHeadDim_;
|
||||||
|
static constexpr int64_t kScoreOffset = kHeadDim;
|
||||||
|
static constexpr int64_t kElementSize = kHeadDim * 2;
|
||||||
|
static constexpr int64_t kPageElementSize = 128 * kElementSize; // page size = 128
|
||||||
|
static constexpr uint32_t kNumSplit = kHeadDim / kTileDim;
|
||||||
|
static_assert(kHeadDim % kTileDim == 0);
|
||||||
|
};
|
||||||
|
|
||||||
/// \brief Need to reduce register usage to increase occupancy
|
/// \brief Need to reduce register usage to increase occupancy
|
||||||
#define C128_KERNEL __global__ __launch_bounds__(kBlockSize, 2)
|
#define C128_KERNEL __global__ __launch_bounds__(C128Config::kBlockSize, 2)
|
||||||
#define WRITE_KERNEL __global__ __launch_bounds__(kWriteBlockSize, 16)
|
#define WRITE_KERNEL __global__ __launch_bounds__(C128Config::kWriteBlockSize, 16)
|
||||||
|
|
||||||
struct Compress128DecodeParams {
|
struct Compress128DecodeParams {
|
||||||
void* __restrict__ kv_buffer;
|
void* __restrict__ kv_buffer;
|
||||||
@@ -69,28 +89,6 @@ struct Compress128PrefillParams {
|
|||||||
uint32_t num_write;
|
uint32_t num_write;
|
||||||
};
|
};
|
||||||
|
|
||||||
struct Compress128SharedBuffer {
|
|
||||||
using Storage = device::AlignedVector<float, kTileElements>;
|
|
||||||
Storage data[kNumWarps][device::kWarpThreads + 1]; // padding to avoid bank conflict
|
|
||||||
SGL_DEVICE Storage& operator()(uint32_t warp_id, uint32_t lane_id) {
|
|
||||||
return data[warp_id][lane_id];
|
|
||||||
}
|
|
||||||
SGL_DEVICE float& operator()(uint32_t warp_id, uint32_t lane_id, uint32_t tile_id) {
|
|
||||||
return data[warp_id][lane_id][tile_id];
|
|
||||||
}
|
|
||||||
};
|
|
||||||
|
|
||||||
template <int64_t kHeadDim_>
|
|
||||||
struct C128Trait {
|
|
||||||
static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64
|
|
||||||
static constexpr int64_t kHeadDim = kHeadDim_;
|
|
||||||
static constexpr int64_t kScoreOffset = kHeadDim;
|
|
||||||
static constexpr int64_t kElementSize = kHeadDim * 2;
|
|
||||||
static constexpr int64_t kPageElementSize = 128 * kElementSize; // page size = 128
|
|
||||||
static constexpr uint32_t kNumSplit = kHeadDim / kTileDim;
|
|
||||||
static_assert(kHeadDim % kTileDim == 0);
|
|
||||||
};
|
|
||||||
|
|
||||||
template <typename Trait, bool kUsePDL, typename BufferFloat, typename InputFloat, typename OutFloat>
|
template <typename Trait, bool kUsePDL, typename BufferFloat, typename InputFloat, typename OutFloat>
|
||||||
SGL_DEVICE void c128_forward(
|
SGL_DEVICE void c128_forward(
|
||||||
const BufferFloat* kv_buf, // [128n, 128n + 127]
|
const BufferFloat* kv_buf, // [128n, 128n + 127]
|
||||||
@@ -100,6 +98,10 @@ SGL_DEVICE void c128_forward(
|
|||||||
const int32_t buffer_len) {
|
const int32_t buffer_len) {
|
||||||
using namespace device;
|
using namespace device;
|
||||||
|
|
||||||
|
constexpr uint32_t kTileElements = Trait::kTileElements;
|
||||||
|
constexpr uint32_t kElementsPerWarp = Trait::kElementsPerWarp;
|
||||||
|
constexpr uint32_t kNumWarps = Trait::kNumWarps;
|
||||||
|
|
||||||
const auto warp_id = threadIdx.x / kWarpThreads;
|
const auto warp_id = threadIdx.x / kWarpThreads;
|
||||||
const auto lane_id = threadIdx.x % kWarpThreads;
|
const auto lane_id = threadIdx.x % kWarpThreads;
|
||||||
|
|
||||||
@@ -112,7 +114,7 @@ SGL_DEVICE void c128_forward(
|
|||||||
const int32_t warp_offset = warp_id * kElementsPerWarp;
|
const int32_t warp_offset = warp_id * kElementsPerWarp;
|
||||||
|
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int32_t i = 0; i < 8; ++i) {
|
for (int32_t i = 0; i < kElementsPerWarp; ++i) {
|
||||||
const int32_t j = i + warp_offset;
|
const int32_t j = i + warp_offset;
|
||||||
bias[i] = gmem_in.load(score_bias + j * Trait::kHeadDim);
|
bias[i] = gmem_in.load(score_bias + j * Trait::kHeadDim);
|
||||||
}
|
}
|
||||||
@@ -129,7 +131,7 @@ SGL_DEVICE void c128_forward(
|
|||||||
score[i] = gmem_in.load(src + j * Trait::kElementSize + Trait::kScoreOffset);
|
score[i] = gmem_in.load(src + j * Trait::kElementSize + Trait::kScoreOffset);
|
||||||
}
|
}
|
||||||
} else { // mixed dtype
|
} else { // mixed dtype
|
||||||
using StorageBuffer = AlignedVector<BufferFloat, kTileElements>;
|
using StorageBuffer = AlignedVector<BufferFloat, Trait::kTileElements>;
|
||||||
const auto gmem_buffer = tile::Memory<StorageBuffer>{lane_id, kWarpThreads};
|
const auto gmem_buffer = tile::Memory<StorageBuffer>{lane_id, kWarpThreads};
|
||||||
|
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
@@ -153,11 +155,12 @@ SGL_DEVICE void c128_forward(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/// NOTE: part 2: safe online softmax + weighted sum
|
/// NOTE: part 2: per-warp partial softmax (online softmax stats + weighted sum)
|
||||||
using TmpStorage = typename Compress128SharedBuffer::Storage;
|
using SharedBuffer = typename Trait::SharedBuffer;
|
||||||
__shared__ Compress128SharedBuffer s_local_val_max;
|
using TmpStorage = typename Trait::SharedStorage;
|
||||||
__shared__ Compress128SharedBuffer s_local_exp_sum;
|
__shared__ SharedBuffer s_local_val_max;
|
||||||
__shared__ Compress128SharedBuffer s_local_product;
|
__shared__ SharedBuffer s_local_exp_sum;
|
||||||
|
__shared__ SharedBuffer s_local_product;
|
||||||
|
|
||||||
TmpStorage tmp_val_max;
|
TmpStorage tmp_val_max;
|
||||||
TmpStorage tmp_exp_sum;
|
TmpStorage tmp_exp_sum;
|
||||||
@@ -178,19 +181,16 @@ SGL_DEVICE void c128_forward(
|
|||||||
for (int32_t i = 0; i < kTileElements; ++i) {
|
for (int32_t i = 0; i < kTileElements; ++i) {
|
||||||
const auto& score = score_fp32[i];
|
const auto& score = score_fp32[i];
|
||||||
float max_value = score[0];
|
float max_value = score[0];
|
||||||
float sum_exp_value = 0.0f;
|
|
||||||
|
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int32_t j = 1; j < kElementsPerWarp; ++j) {
|
for (int32_t j = 1; j < kElementsPerWarp; ++j) {
|
||||||
const auto fp32_score = score[j];
|
max_value = fmaxf(max_value, score[j]);
|
||||||
max_value = fmaxf(max_value, fp32_score);
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
float sum_exp_value = 0.0f;
|
||||||
float sum_product = 0.0f;
|
float sum_product = 0.0f;
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int32_t j = 0; j < 8; ++j) {
|
for (int32_t j = 0; j < kElementsPerWarp; ++j) {
|
||||||
const auto fp32_score = score[j];
|
const auto exp_score = expf(score[j] - max_value);
|
||||||
const auto exp_score = expf(fp32_score - max_value);
|
|
||||||
sum_product += cast<float>(kv[j][i]) * exp_score;
|
sum_product += cast<float>(kv[j][i]) * exp_score;
|
||||||
sum_exp_value += exp_score;
|
sum_exp_value += exp_score;
|
||||||
}
|
}
|
||||||
@@ -201,45 +201,48 @@ SGL_DEVICE void c128_forward(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// naturally aligned, so no bank conflict
|
// naturally aligned, so no bank conflict
|
||||||
s_local_val_max(warp_id, lane_id) = tmp_val_max;
|
s_local_val_max[warp_id][lane_id] = tmp_val_max;
|
||||||
s_local_exp_sum(warp_id, lane_id) = tmp_exp_sum;
|
s_local_exp_sum[warp_id][lane_id] = tmp_exp_sum;
|
||||||
s_local_product(warp_id, lane_id) = tmp_product;
|
s_local_product[warp_id][lane_id] = tmp_product;
|
||||||
|
|
||||||
__syncthreads();
|
__syncthreads();
|
||||||
|
|
||||||
/// NOTE: part 3: online softmax
|
|
||||||
/// NOTE: We have `kTileElements * kWarpThreads * kNumWarps` values to reduce
|
|
||||||
/// each reduce will consume `kNumWarps` threads (use partial warp reduction)
|
|
||||||
constexpr uint32_t kReductionCount = kTileElements * kWarpThreads * kNumWarps;
|
|
||||||
constexpr uint32_t kIteration = kReductionCount / kBlockSize;
|
|
||||||
|
|
||||||
PDLTriggerSecondary<kUsePDL>();
|
PDLTriggerSecondary<kUsePDL>();
|
||||||
|
|
||||||
|
/// NOTE: part 3: final reduction + write-back.
|
||||||
|
/// Only the first `kTileElements` warps participate; each thread reduces over
|
||||||
|
/// `kNumWarps` partial values entirely in registers and writes one output element.
|
||||||
|
/// The remaining warps exit early, freeing issue slots and avoiding redundant writes.
|
||||||
|
if (warp_id < kTileElements) {
|
||||||
|
const uint32_t tx = threadIdx.x;
|
||||||
|
const uint32_t local_lane_id = tx / kTileElements; // [0, kWarpThreads)
|
||||||
|
const uint32_t local_tile_id = tx % kTileElements; // [0, kTileElements)
|
||||||
|
|
||||||
|
float local_val_max[kNumWarps];
|
||||||
|
float local_exp_sum[kNumWarps];
|
||||||
|
float local_product[kNumWarps];
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (uint32_t i = 0; i < kIteration; ++i) {
|
for (uint32_t i = 0; i < kNumWarps; ++i) {
|
||||||
/// NOTE: Range `[0, kTileElements * kWarpThreads * kNumWarps)`
|
local_val_max[i] = s_local_val_max[i][local_lane_id][local_tile_id];
|
||||||
const uint32_t j = i * kBlockSize + warp_id * kWarpThreads + lane_id;
|
local_exp_sum[i] = s_local_exp_sum[i][local_lane_id][local_tile_id];
|
||||||
/// NOTE: Range `[0, kNumWarps)`
|
local_product[i] = s_local_product[i][local_lane_id][local_tile_id];
|
||||||
const uint32_t local_warp_id = j % kNumWarps;
|
}
|
||||||
/// NOTE: Range `[0, kTileElements * kWarpThreads)`
|
|
||||||
const uint32_t local_elem_id = j / kNumWarps;
|
float global_max = local_val_max[0];
|
||||||
/// NOTE: Range `[0, kTileElements)`
|
#pragma unroll
|
||||||
const uint32_t local_tile_id = local_elem_id % kTileElements;
|
for (uint32_t i = 1; i < kNumWarps; ++i) {
|
||||||
/// NOTE: Range `[0, kWarpThreads)`
|
global_max = fmaxf(global_max, local_val_max[i]);
|
||||||
const uint32_t local_lane_id = local_elem_id / kTileElements;
|
}
|
||||||
/// NOTE: each warp will access the whole tile (all `kTileElements`)
|
|
||||||
/// and for different lanes, the memory access only differ in `local_warp_id`
|
float global_exp_sum = 0.0f;
|
||||||
/// so there's no bank conflict in shared memory access.
|
float global_product = 0.0f;
|
||||||
static_assert(kTileElements * kNumWarps == kWarpThreads, "TODO: support other configs");
|
#pragma unroll
|
||||||
const auto local_val_max = s_local_val_max(local_warp_id, local_lane_id, local_tile_id);
|
for (uint32_t i = 0; i < kNumWarps; ++i) {
|
||||||
const auto local_exp_sum = s_local_exp_sum(local_warp_id, local_lane_id, local_tile_id);
|
const auto exp_val = expf(local_val_max[i] - global_max);
|
||||||
const auto local_product = s_local_product(local_warp_id, local_lane_id, local_tile_id);
|
global_exp_sum += local_exp_sum[i] * exp_val;
|
||||||
const auto global_val_max = warp::reduce_max<kNumWarps>(local_val_max);
|
global_product += local_product[i] * exp_val;
|
||||||
const auto rescale = expf(local_val_max - global_val_max);
|
}
|
||||||
const auto global_exp_sum = warp::reduce_sum<kNumWarps>(local_exp_sum * rescale);
|
kv_out[tx] = cast<OutFloat>(global_product / global_exp_sum);
|
||||||
const auto final_scale = rescale / global_exp_sum;
|
|
||||||
const auto global_product = warp::reduce_sum<kNumWarps>(local_product * final_scale);
|
|
||||||
kv_out[local_elem_id] = cast<OutFloat>(global_product);
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -247,7 +250,7 @@ template <typename Trait, typename BufferFloat, typename InputFloat>
|
|||||||
SGL_DEVICE void c128_write_decode(BufferFloat* kv_buf, const InputFloat* kv_src) {
|
SGL_DEVICE void c128_write_decode(BufferFloat* kv_buf, const InputFloat* kv_src) {
|
||||||
using namespace device;
|
using namespace device;
|
||||||
|
|
||||||
using StorageInput = AlignedVector<InputFloat, kTileElements>;
|
using StorageInput = AlignedVector<InputFloat, Trait::kTileElements>;
|
||||||
const auto gmem_input = tile::Memory<StorageInput>::warp();
|
const auto gmem_input = tile::Memory<StorageInput>::warp();
|
||||||
|
|
||||||
StorageInput data[2];
|
StorageInput data[2];
|
||||||
@@ -262,14 +265,14 @@ SGL_DEVICE void c128_write_decode(BufferFloat* kv_buf, const InputFloat* kv_src)
|
|||||||
gmem_input.store(kv_buf + Trait::kHeadDim * i, data[i]);
|
gmem_input.store(kv_buf + Trait::kHeadDim * i, data[i]);
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
using StorageBuffer = AlignedVector<BufferFloat, kTileElements>;
|
using StorageBuffer = AlignedVector<BufferFloat, Trait::kTileElements>;
|
||||||
const auto gmem_buffer = tile::Memory<StorageBuffer>::warp();
|
const auto gmem_buffer = tile::Memory<StorageBuffer>::warp();
|
||||||
|
|
||||||
StorageBuffer data_cast[2];
|
StorageBuffer data_cast[2];
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int32_t i = 0; i < 2; ++i) {
|
for (int32_t i = 0; i < 2; ++i) {
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int32_t j = 0; j < kTileElements; ++j) {
|
for (int32_t j = 0; j < Trait::kTileElements; ++j) {
|
||||||
data_cast[i][j] = cast<BufferFloat>(data[i][j]);
|
data_cast[i][j] = cast<BufferFloat>(data[i][j]);
|
||||||
}
|
}
|
||||||
gmem_buffer.store(kv_buf + Trait::kHeadDim * i, data_cast[i]);
|
gmem_buffer.store(kv_buf + Trait::kHeadDim * i, data_cast[i]);
|
||||||
@@ -277,6 +280,7 @@ SGL_DEVICE void c128_write_decode(BufferFloat* kv_buf, const InputFloat* kv_src)
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// \brief Need to reduce register usage to increase occupancy.
|
||||||
template <int64_t kHeadDim, typename BufferFloat, typename InputFloat, typename OutFloat, bool kUsePDL>
|
template <int64_t kHeadDim, typename BufferFloat, typename InputFloat, typename OutFloat, bool kUsePDL>
|
||||||
C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodeParams params) {
|
C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodeParams params) {
|
||||||
using namespace device;
|
using namespace device;
|
||||||
@@ -301,7 +305,7 @@ C128_KERNEL void flash_c128_decode(const __grid_constant__ Compress128DecodePara
|
|||||||
|
|
||||||
PDLWaitPrimary<kUsePDL>();
|
PDLWaitPrimary<kUsePDL>();
|
||||||
// the write warp must match the load warp in the following `c128_forward`
|
// the write warp must match the load warp in the following `c128_forward`
|
||||||
if (warp_id == kNumWarps - 1) {
|
if (warp_id == Trait::kNumWarps - 1) {
|
||||||
c128_write_decode<Trait, BufferFloat, InputFloat>(kv_dst, kv_src);
|
c128_write_decode<Trait, BufferFloat, InputFloat>(kv_dst, kv_src);
|
||||||
}
|
}
|
||||||
if (plan.write_loc % 128 == 127) {
|
if (plan.write_loc % 128 == 127) {
|
||||||
@@ -339,7 +343,7 @@ template <int64_t kHeadDim, typename BufferFloat, typename InputFloat, typename
|
|||||||
WRITE_KERNEL void write_c128_prefill(const __grid_constant__ Compress128PrefillParams params) {
|
WRITE_KERNEL void write_c128_prefill(const __grid_constant__ Compress128PrefillParams params) {
|
||||||
using namespace device;
|
using namespace device;
|
||||||
using Trait = C128Trait<kHeadDim>;
|
using Trait = C128Trait<kHeadDim>;
|
||||||
using StorageInput = AlignedVector<InputFloat, kTileElements>;
|
using StorageInput = AlignedVector<InputFloat, Trait::kTileElements>;
|
||||||
|
|
||||||
const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x;
|
const uint32_t global_tid = blockIdx.x * blockDim.x + threadIdx.x;
|
||||||
const uint32_t global_wid = global_tid / kWarpThreads; // warp id
|
const uint32_t global_wid = global_tid / kWarpThreads; // warp id
|
||||||
@@ -374,14 +378,14 @@ WRITE_KERNEL void write_c128_prefill(const __grid_constant__ Compress128PrefillP
|
|||||||
gmem_input.store(kv_buf, data[i], i);
|
gmem_input.store(kv_buf, data[i], i);
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
using StorageBuffer = AlignedVector<BufferFloat, kTileElements>;
|
using StorageBuffer = AlignedVector<BufferFloat, Trait::kTileElements>;
|
||||||
const auto gmem_buffer = tile::Memory<StorageBuffer>::warp();
|
const auto gmem_buffer = tile::Memory<StorageBuffer>::warp();
|
||||||
|
|
||||||
StorageBuffer data_cast[2];
|
StorageBuffer data_cast[2];
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int32_t i = 0; i < 2; ++i) {
|
for (int32_t i = 0; i < 2; ++i) {
|
||||||
#pragma unroll
|
#pragma unroll
|
||||||
for (int32_t j = 0; j < kTileElements; ++j) {
|
for (int32_t j = 0; j < Trait::kTileElements; ++j) {
|
||||||
data_cast[i][j] = cast<BufferFloat>(data[i][j]);
|
data_cast[i][j] = cast<BufferFloat>(data[i][j]);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -395,12 +399,10 @@ WRITE_KERNEL void write_c128_prefill(const __grid_constant__ Compress128PrefillP
|
|||||||
|
|
||||||
template <int64_t kHeadDim, typename BufferFloat, typename InputFloat, typename OutFloat, bool kUsePDL>
|
template <int64_t kHeadDim, typename BufferFloat, typename InputFloat, typename OutFloat, bool kUsePDL>
|
||||||
struct FlashCompress128Kernel {
|
struct FlashCompress128Kernel {
|
||||||
|
using Trait = C128Trait<kHeadDim>;
|
||||||
static constexpr auto decode_kernel = flash_c128_decode<kHeadDim, BufferFloat, InputFloat, OutFloat, kUsePDL>;
|
static constexpr auto decode_kernel = flash_c128_decode<kHeadDim, BufferFloat, InputFloat, OutFloat, kUsePDL>;
|
||||||
static constexpr auto prefill_c_kernel = flash_c128_prefill<kHeadDim, BufferFloat, InputFloat, OutFloat, kUsePDL>;
|
static constexpr auto prefill_c_kernel = flash_c128_prefill<kHeadDim, BufferFloat, InputFloat, OutFloat, kUsePDL>;
|
||||||
static constexpr auto prefill_w_kernel = write_c128_prefill<kHeadDim, BufferFloat, InputFloat, OutFloat, kUsePDL>;
|
static constexpr auto prefill_w_kernel = write_c128_prefill<kHeadDim, BufferFloat, InputFloat, OutFloat, kUsePDL>;
|
||||||
static constexpr int64_t kTileDim = kTileElements * device::kWarpThreads; // 64
|
|
||||||
static constexpr uint32_t kNumSplit = kHeadDim / kTileDim;
|
|
||||||
using Trait = C128Trait<kHeadDim>;
|
|
||||||
|
|
||||||
static void run_decode(
|
static void run_decode(
|
||||||
const tvm::ffi::TensorView kv_buffer,
|
const tvm::ffi::TensorView kv_buffer,
|
||||||
@@ -441,8 +443,8 @@ struct FlashCompress128Kernel {
|
|||||||
.plan_d = plan_d,
|
.plan_d = plan_d,
|
||||||
.batch_size = batch_size,
|
.batch_size = batch_size,
|
||||||
};
|
};
|
||||||
const uint32_t num_blocks = batch_size * kNumSplit;
|
const uint32_t num_blocks = batch_size * Trait::kNumSplit;
|
||||||
LaunchKernel(num_blocks, kBlockSize, device_.unwrap()) //
|
LaunchKernel(num_blocks, Trait::kBlockSize, device_.unwrap()) //
|
||||||
.enable_pdl(kUsePDL)(decode_kernel, params);
|
.enable_pdl(kUsePDL)(decode_kernel, params);
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -495,15 +497,12 @@ struct FlashCompress128Kernel {
|
|||||||
.num_write = num_w,
|
.num_write = num_w,
|
||||||
};
|
};
|
||||||
RuntimeCheck(num_q_tokens >= num_w, "invalid prefill plan: num_q < num_w");
|
RuntimeCheck(num_q_tokens >= num_w, "invalid prefill plan: num_q < num_w");
|
||||||
if (const auto num_c_blocks = num_c * kNumSplit) {
|
if (const auto num_c_blocks = num_c * Trait::kNumSplit) {
|
||||||
constexpr auto kBlockSize_C = kBlockSize;
|
LaunchKernel(num_c_blocks, Trait::kBlockSize, device) //
|
||||||
LaunchKernel(num_c_blocks, kBlockSize_C, device) //
|
|
||||||
.enable_pdl(kUsePDL)(prefill_c_kernel, params);
|
.enable_pdl(kUsePDL)(prefill_c_kernel, params);
|
||||||
}
|
}
|
||||||
constexpr uint32_t kWarpsPerWriteBlock = kWriteBlockSize / device::kWarpThreads;
|
if (const auto num_w_blocks = div_ceil(num_w * Trait::kNumSplit, Trait::kNumWriteWarps)) {
|
||||||
if (const auto num_w_blocks = div_ceil(num_w * kNumSplit, kWarpsPerWriteBlock)) {
|
LaunchKernel(num_w_blocks, Trait::kWriteBlockSize, device) //
|
||||||
constexpr auto kBlockSize_W = kWriteBlockSize;
|
|
||||||
LaunchKernel(num_w_blocks, kBlockSize_W, device) //
|
|
||||||
.enable_pdl(kUsePDL)(prefill_w_kernel, params);
|
.enable_pdl(kUsePDL)(prefill_w_kernel, params);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user