cuda: keep live MTP correctness fixes

Retain the useful CUDA pieces from cghart's branch without the disabled decode experiments: add the fd-cache model-map guard for base+MTP loads, support the Q4_K routed expert decode path used by the MTP draft head, and use CUDA for non-Apple ds4_test runs.
This commit is contained in:
cghart
2026-05-12 22:18:48 +02:00
committed by antirez
parent 8809b90a1e
commit a5f4f761b7
2 changed files with 181 additions and 11 deletions
+177 -11
View File
@@ -47,6 +47,13 @@ typedef struct {
uint16_t dmin;
} cuda_block_q2_K;
typedef struct {
uint16_t d;
uint16_t dmin;
uint8_t scales[12];
uint8_t qs[CUDA_QK_K / 2];
} cuda_block_q4_K;
typedef struct {
float d;
int8_t qs[CUDA_QK_K];
@@ -68,6 +75,7 @@ static int g_model_device_owned;
static int g_model_range_mapping_supported = 1;
static int g_model_hmm_direct;
static int g_model_fd = -1;
static const void *g_model_fd_host_base;
static int g_model_direct_fd = -1;
static uint64_t g_model_direct_align = 1;
static uint64_t g_model_file_size;
@@ -988,6 +996,7 @@ static const char *cuda_model_range_ptr_from_fd(
uint64_t bytes,
const char *what) {
if (g_model_fd < 0 || bytes == 0) return NULL;
if (g_model_fd_host_base != NULL && model_map != g_model_fd_host_base) return NULL;
const uint64_t limit = cuda_model_cache_limit_bytes();
if (g_model_range_bytes > limit || bytes > limit - g_model_range_bytes) {
if (getenv("DS4_CUDA_WEIGHT_CACHE_VERBOSE")) {
@@ -1369,6 +1378,9 @@ extern "C" int ds4_gpu_set_model_map(const void *model_map, uint64_t model_size)
g_model_range_mapping_supported = 1;
g_model_hmm_direct = 0;
g_model_cache_full = 0;
if (g_model_fd >= 0 && g_model_fd_host_base == NULL) {
g_model_fd_host_base = model_map;
}
const char *copy_env = getenv("DS4_CUDA_COPY_MODEL");
if (copy_env && copy_env[0]) {
@@ -1427,6 +1439,7 @@ extern "C" int ds4_gpu_set_model_map_range(const void *model_map, uint64_t model
extern "C" int ds4_gpu_set_model_fd(int fd) {
g_model_fd = fd;
g_model_fd_host_base = g_model_host_base;
g_model_file_size = 0;
if (g_model_direct_fd >= 0) {
(void)close(g_model_direct_fd);
@@ -7216,6 +7229,47 @@ __device__ static DS4_CUDA_UNUSED void dev_dot_iq2_xxs_q8_K_block8(
for (uint32_t p = 0; p < n; p++) acc[p] += 0.125f * xd * ys[p]->d * (float)bsum[p];
}
__device__ static void dev_q4_K_get_scale_min(
uint32_t j,
const uint8_t *scales,
uint8_t *d_out,
uint8_t *m_out) {
if (j < 4u) {
*d_out = scales[j] & 63u;
*m_out = scales[j + 4u] & 63u;
} else {
*d_out = (scales[j + 4u] & 0x0fu) | ((scales[j - 4u] >> 6u) << 4u);
*m_out = (scales[j + 4u] >> 4u) | ((scales[j] >> 6u) << 4u);
}
}
__device__ __forceinline__ static int32_t dev_dot_q4_32(const uint8_t *qs, const int8_t *q8, int shift) {
int32_t sum = 0;
#pragma unroll
for (uint32_t i = 0; i < 32u; i += 4u) {
const int32_t v = (*(const int32_t *)(qs + i) >> shift) & 0x0f0f0f0f;
sum = __dp4a(v, *(const int32_t *)(q8 + i), sum);
}
return sum;
}
__device__ static float dev_dot_q4_K_q8_K_block(const cuda_block_q4_K *x, const cuda_block_q8_K *y) {
const float xd = dev_f16_to_f32(x->d);
const float xmin = dev_f16_to_f32(x->dmin);
int isum = 0;
int summs = 0;
#pragma unroll
for (uint32_t j = 0; j < 8u; j++) {
uint8_t sc, m;
dev_q4_K_get_scale_min(j, x->scales, &sc, &m);
summs += (int)m * (int)(y->bsums[2u * j] + y->bsums[2u * j + 1u]);
const uint32_t byte_off = (j >> 1u) * 32u;
const int shift = (j & 1u) ? 4 : 0;
isum += (int)sc * dev_dot_q4_32(x->qs + byte_off, y->qs + j * 32u, shift);
}
return y->d * xd * (float)isum - y->d * xmin * (float)summs;
}
__device__ static float dev_dot_q2_K_q8_K_block(const cuda_block_q2_K *x, const cuda_block_q8_K *y) {
const uint8_t *q2 = x->qs;
const int8_t *q8 = y->qs;
@@ -8421,6 +8475,60 @@ __global__ static void moe_down_qwarp32_kernel(
if (lane == 0) down_out[(uint64_t)pair * out_dim + row] = acc;
}
__global__ static void moe_gate_up_mid_decode_q4K_qwarp32_kernel(
float *gate_out,
float *up_out,
float *mid_out,
const char *gate_base,
const char *up_base,
const cuda_block_q8_K *xq,
const int32_t *selected,
const float *weights,
uint64_t gate_expert_bytes,
uint64_t gate_row_bytes,
uint32_t xq_blocks,
uint32_t expert_mid_dim,
uint32_t n_expert,
uint32_t write_aux,
float clamp) {
uint32_t lane = threadIdx.x & 7u;
uint32_t row_lane = threadIdx.x >> 3u;
uint32_t pair = blockIdx.y;
uint32_t tok = pair / n_expert;
uint32_t slot = pair - tok * n_expert;
int32_t expert_i = selected[(uint64_t)tok * n_expert + slot];
if (expert_i < 0) expert_i = 0;
uint32_t expert = (uint32_t)expert_i;
const cuda_block_q8_K *xqb = xq + (uint64_t)tok * xq_blocks;
for (uint32_t rr = 0; rr < 4u; rr++) {
uint32_t row = blockIdx.x * 128u + row_lane + rr * 32u;
if (row >= expert_mid_dim) continue;
const cuda_block_q4_K *gr = (const cuda_block_q4_K *)(gate_base + (uint64_t)expert * gate_expert_bytes + (uint64_t)row * gate_row_bytes);
const cuda_block_q4_K *ur = (const cuda_block_q4_K *)(up_base + (uint64_t)expert * gate_expert_bytes + (uint64_t)row * gate_row_bytes);
float gate = 0.0f;
float up = 0.0f;
for (uint32_t b = lane; b < xq_blocks; b += 8u) {
gate += dev_dot_q4_K_q8_K_block(gr + b, xqb + b);
up += dev_dot_q4_K_q8_K_block(ur + b, xqb + b);
}
gate = quarter_warp_sum_f32(gate, lane);
up = quarter_warp_sum_f32(up, lane);
if (lane == 0) {
if (clamp > 1.0e-6f) {
if (gate > clamp) gate = clamp;
if (up > clamp) up = clamp;
if (up < -clamp) up = -clamp;
}
const uint64_t off = (uint64_t)pair * expert_mid_dim + row;
if (write_aux) {
gate_out[off] = gate;
up_out[off] = up;
}
mid_out[off] = (gate / (1.0f + expf(-gate))) * up * weights[(uint64_t)tok * n_expert + slot];
}
}
}
__global__ static void moe_down_sum6_qwarp32_kernel(
float *out,
const char *down_base,
@@ -8448,6 +8556,33 @@ __global__ static void moe_down_sum6_qwarp32_kernel(
if (lane == 0) out[row] = total;
}
__global__ static void moe_down_q4K_sum6_qwarp32_kernel(
float *out,
const char *down_base,
const cuda_block_q8_K *midq,
const int32_t *selected,
uint64_t down_expert_bytes,
uint64_t down_row_bytes,
uint32_t midq_blocks,
uint32_t out_dim) {
uint32_t lane = threadIdx.x & 7u;
uint32_t row = blockIdx.x * 32u + (threadIdx.x >> 3u);
if (row >= out_dim) return;
float total = 0.0f;
#pragma unroll
for (uint32_t slot = 0; slot < 6u; slot++) {
int32_t expert_i = selected[slot];
if (expert_i < 0) expert_i = 0;
const cuda_block_q4_K *wr = (const cuda_block_q4_K *)(down_base + (uint64_t)(uint32_t)expert_i * down_expert_bytes + (uint64_t)row * down_row_bytes);
const cuda_block_q8_K *xq = midq + (uint64_t)slot * midq_blocks;
float acc = 0.0f;
for (uint32_t b = lane; b < midq_blocks; b += 8u) acc += dev_dot_q4_K_q8_K_block(wr + b, xq + b);
acc = quarter_warp_sum_f32(acc, lane);
if (lane == 0) total += acc;
}
if (lane == 0) out[row] = total;
}
__global__ static void moe_down_sorted_qwarp32_kernel(
float *down_out,
const char *down_base,
@@ -9161,7 +9296,9 @@ static int routed_moe_launch(
out->bytes < (uint64_t)n_tokens * out_dim * sizeof(float)) {
return 0;
}
if (gate_type != 16u || down_type != 10u) return 0;
const int q4k_path = (gate_type == 12u && down_type == 12u);
if (!q4k_path && (gate_type != 16u || down_type != 10u)) return 0;
if (q4k_path && (n_tokens != 1u || n_expert != 6u)) return 0;
const uint64_t gate_bytes = 256ull * gate_expert_bytes;
const uint64_t down_bytes = 256ull * down_expert_bytes;
if (gate_bytes > model_size - gate_offset ||
@@ -9420,7 +9557,24 @@ static int routed_moe_launch(
clamp);
} else if (ok) {
dim3 qgrid((expert_mid_dim + 127u) / 128u, n_tokens * n_expert, 1);
if (use_decode_lut_gate) {
if (use_decode_lut_gate && q4k_path) {
moe_gate_up_mid_decode_q4K_qwarp32_kernel<<<qgrid, 256>>>(
(float *)gate->ptr,
(float *)up->ptr,
(float *)mid->ptr,
gate_w,
up_w,
xq,
(const int32_t *)selected->ptr,
(const float *)weights->ptr,
gate_expert_bytes,
gate_row_bytes,
xq_blocks,
expert_mid_dim,
n_expert,
write_gate_up,
clamp);
} else if (use_decode_lut_gate) {
moe_gate_up_mid_decode_lut_qwarp32_kernel<<<qgrid, 256>>>(
(float *)gate->ptr,
(float *)up->ptr,
@@ -9478,15 +9632,27 @@ static int routed_moe_launch(
}
if (use_direct_down_sum6) {
dim3 sgrid((out_dim + 31u) / 32u, 1, 1);
moe_down_sum6_qwarp32_kernel<<<sgrid, 256>>>(
(float *)out->ptr,
down_w,
midq,
(const int32_t *)selected->ptr,
down_expert_bytes,
down_row_bytes,
midq_blocks,
out_dim);
if (q4k_path) {
moe_down_q4K_sum6_qwarp32_kernel<<<sgrid, 256>>>(
(float *)out->ptr,
down_w,
midq,
(const int32_t *)selected->ptr,
down_expert_bytes,
down_row_bytes,
midq_blocks,
out_dim);
} else {
moe_down_sum6_qwarp32_kernel<<<sgrid, 256>>>(
(float *)out->ptr,
down_w,
midq,
(const int32_t *)selected->ptr,
down_expert_bytes,
down_row_bytes,
midq_blocks,
out_dim);
}
} else if (use_atomic_down) {
uint64_t n = (uint64_t)n_tokens * out_dim;
zero_kernel<<<(n + 255u) / 256u, 256>>>((float *)out->ptr, n);
+4
View File
@@ -19,7 +19,11 @@ static ds4_engine *test_get_engine(bool quality) {
ds4_engine_options opt = {
.model_path = test_model_path(),
#ifdef __APPLE__
.backend = DS4_BACKEND_METAL,
#else
.backend = DS4_BACKEND_CUDA,
#endif
.quality = quality,
};
TEST_ASSERT(ds4_engine_open(slot, &opt) == 0);