__global__ static void rms_norm_plain_kernel(float *out, const float *x, uint32_t n, uint32_t rows, float eps) { uint32_t row = blockIdx.x; if (row >= rows) return; const float *xr = x + (uint64_t)row * n; float *orow = out + (uint64_t)row * n; float sum = 0.0f; for (uint32_t i = threadIdx.x; i < n; i += blockDim.x) { float v = xr[i]; sum += v * v; } __shared__ float partial[256]; partial[threadIdx.x] = sum; __syncthreads(); for (uint32_t stride = blockDim.x >> 1; stride > 0; stride >>= 1) { if (threadIdx.x < stride) partial[threadIdx.x] += partial[threadIdx.x + stride]; __syncthreads(); } float scale = rsqrtf(partial[0] / (float)n + eps); for (uint32_t i = threadIdx.x; i < n; i += blockDim.x) { orow[i] = xr[i] * scale; } } __global__ static void rms_norm_weight_kernel(float *out, const float *x, const float *w, uint32_t n, uint32_t rows, float eps) { uint32_t row = blockIdx.x; if (row >= rows) return; const float *xr = x + (uint64_t)row * n; float *orow = out + (uint64_t)row * n; float sum = 0.0f; for (uint32_t i = threadIdx.x; i < n; i += blockDim.x) { float v = xr[i]; sum += v * v; } __shared__ float partial[256]; partial[threadIdx.x] = sum; __syncthreads(); for (uint32_t stride = blockDim.x >> 1; stride > 0; stride >>= 1) { if (threadIdx.x < stride) partial[threadIdx.x] += partial[threadIdx.x + stride]; __syncthreads(); } float scale = rsqrtf(partial[0] / (float)n + eps); for (uint32_t i = threadIdx.x; i < n; i += blockDim.x) { orow[i] = xr[i] * scale * w[i]; } } __global__ static void dsv4_qkv_rms_norm_rows_kernel( float *q_out, const float *q, const float *q_w, uint32_t q_n, float *kv_out, const float *kv, const float *kv_w, uint32_t kv_n, uint32_t rows, float eps) { const uint32_t row = blockIdx.x; const uint32_t which = blockIdx.y; if (row >= rows || which > 1u) return; const uint32_t n = which == 0u ? q_n : kv_n; const float *xr = (which == 0u ? q : kv) + (uint64_t)row * n; float *orow = (which == 0u ? q_out : kv_out) + (uint64_t)row * n; const float *w = which == 0u ? q_w : kv_w; float sum = 0.0f; for (uint32_t i = threadIdx.x; i < n; i += blockDim.x) { const float v = xr[i]; sum += v * v; } __shared__ float partial[256]; partial[threadIdx.x] = sum; __syncthreads(); for (uint32_t stride = blockDim.x >> 1; stride > 0; stride >>= 1) { if (threadIdx.x < stride) partial[threadIdx.x] += partial[threadIdx.x + stride]; __syncthreads(); } const float scale = rsqrtf(partial[0] / (float)n + eps); for (uint32_t i = threadIdx.x; i < n; i += blockDim.x) { orow[i] = xr[i] * scale * w[i]; } } __global__ static void head_rms_norm_kernel(float *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, float eps) { uint32_t row = blockIdx.x; if (row >= n_tok * n_head) return; float *xr = x + (uint64_t)row * head_dim; float sum = 0.0f; for (uint32_t i = threadIdx.x; i < head_dim; i += blockDim.x) { float v = xr[i]; sum += v * v; } __shared__ float partial[256]; partial[threadIdx.x] = sum; __syncthreads(); for (uint32_t stride = blockDim.x >> 1; stride > 0; stride >>= 1) { if (threadIdx.x < stride) partial[threadIdx.x] += partial[threadIdx.x + stride]; __syncthreads(); } float scale = rsqrtf(partial[0] / (float)head_dim + eps); for (uint32_t i = threadIdx.x; i < head_dim; i += blockDim.x) xr[i] *= scale; } __device__ static float rope_yarn_ramp_dev(float low, float high, int i0); __global__ static void head_rms_norm_rope_tail_kernel( float *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, uint32_t n_rot, uint32_t pos0, uint32_t n_ctx_orig, int inverse, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow, float eps) { uint32_t row = blockIdx.x; if (row >= n_tok * n_head) return; uint32_t t = row / n_head; float *xr = x + (uint64_t)row * head_dim; float sum = 0.0f; for (uint32_t i = threadIdx.x; i < head_dim; i += blockDim.x) { float v = xr[i]; sum += v * v; } __shared__ float partial[256]; partial[threadIdx.x] = sum; __syncthreads(); for (uint32_t stride = blockDim.x >> 1; stride > 0; stride >>= 1) { if (threadIdx.x < stride) partial[threadIdx.x] += partial[threadIdx.x + stride]; __syncthreads(); } const float scale = rsqrtf(partial[0] / (float)head_dim + eps); const uint32_t n_nope = head_dim - n_rot; for (uint32_t i = threadIdx.x; i < n_nope; i += blockDim.x) { xr[i] *= scale; } float corr0 = 0.0f, corr1 = 0.0f; if (ext_factor != 0.0f) { float denom = 2.0f * logf(freq_base); corr0 = floorf((float)n_rot * logf((float)n_ctx_orig / (beta_fast * 2.0f * (float)M_PI)) / denom); corr1 = ceilf((float)n_rot * logf((float)n_ctx_orig / (beta_slow * 2.0f * (float)M_PI)) / denom); corr0 = fmaxf(0.0f, corr0); corr1 = fminf((float)(n_rot - 1), corr1); } const float theta_scale = powf(freq_base, -2.0f / (float)n_rot); for (uint32_t pair = threadIdx.x; pair < n_rot / 2; pair += blockDim.x) { uint32_t i = pair * 2u; float theta_extrap = (float)(pos0 + t) * powf(theta_scale, (float)pair); float theta_interp = freq_scale * theta_extrap; float theta = theta_interp; float mscale = attn_factor; if (ext_factor != 0.0f) { float ramp_mix = rope_yarn_ramp_dev(corr0, corr1, (int)i) * ext_factor; theta = theta_interp * (1.0f - ramp_mix) + theta_extrap * ramp_mix; mscale *= 1.0f + 0.1f * logf(1.0f / freq_scale); } float c = cosf(theta) * mscale; float s = sinf(theta) * mscale; if (inverse) s = -s; float *tail = xr + n_nope; float x0 = tail[i] * scale; float x1 = tail[i + 1] * scale; tail[i] = x0 * c - x1 * s; tail[i + 1] = x0 * s + x1 * c; } } __global__ static void head_rms_norm_rope_tail_from_half_kernel( float *out, const __half *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, uint32_t n_rot, uint32_t pos0, uint32_t n_ctx_orig, int inverse, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow, float eps) { uint32_t row = blockIdx.x; if (row >= n_tok * n_head) return; uint32_t t = row / n_head; const __half *xr = x + (uint64_t)row * head_dim; float *orow = out + (uint64_t)row * head_dim; float sum = 0.0f; for (uint32_t i = threadIdx.x; i < head_dim; i += blockDim.x) { float v = __half2float(xr[i]); sum += v * v; } __shared__ float partial[256]; partial[threadIdx.x] = sum; __syncthreads(); for (uint32_t stride = blockDim.x >> 1; stride > 0; stride >>= 1) { if (threadIdx.x < stride) partial[threadIdx.x] += partial[threadIdx.x + stride]; __syncthreads(); } const float scale = rsqrtf(partial[0] / (float)head_dim + eps); const uint32_t n_nope = head_dim - n_rot; for (uint32_t i = threadIdx.x; i < n_nope; i += blockDim.x) { orow[i] = __half2float(xr[i]) * scale; } float corr0 = 0.0f, corr1 = 0.0f; if (ext_factor != 0.0f) { float denom = 2.0f * logf(freq_base); corr0 = floorf((float)n_rot * logf((float)n_ctx_orig / (beta_fast * 2.0f * (float)M_PI)) / denom); corr1 = ceilf((float)n_rot * logf((float)n_ctx_orig / (beta_slow * 2.0f * (float)M_PI)) / denom); corr0 = fmaxf(0.0f, corr0); corr1 = fminf((float)(n_rot - 1), corr1); } const float theta_scale = powf(freq_base, -2.0f / (float)n_rot); const __half *tail = xr + n_nope; float *otail = orow + n_nope; for (uint32_t pair = threadIdx.x; pair < n_rot / 2; pair += blockDim.x) { uint32_t i = pair * 2u; float theta_extrap = (float)(pos0 + t) * powf(theta_scale, (float)pair); float theta_interp = freq_scale * theta_extrap; float theta = theta_interp; float mscale = attn_factor; if (ext_factor != 0.0f) { float ramp_mix = rope_yarn_ramp_dev(corr0, corr1, (int)i) * ext_factor; theta = theta_interp * (1.0f - ramp_mix) + theta_extrap * ramp_mix; mscale *= 1.0f + 0.1f * logf(1.0f / freq_scale); } float c = cosf(theta) * mscale; float s = sinf(theta) * mscale; if (inverse) s = -s; float x0 = __half2float(tail[i]) * scale; float x1 = __half2float(tail[i + 1]) * scale; otail[i] = x0 * c - x1 * s; otail[i + 1] = x0 * s + x1 * c; } } __device__ static float rope_yarn_ramp_dev(float low, float high, int i0) { float y = ((float)(i0 / 2) - low) / fmaxf(0.001f, high - low); return 1.0f - fminf(1.0f, fmaxf(0.0f, y)); } __global__ static void rope_tail_kernel( float *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, uint32_t n_rot, uint32_t pos0, uint32_t pos_stride, uint32_t n_ctx_orig, int inverse, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow) { uint32_t gid = blockIdx.x * blockDim.x + threadIdx.x; uint32_t pairs = n_tok * n_head * (n_rot / 2); if (gid >= pairs) return; uint32_t pair = gid % (n_rot / 2); uint32_t tmp = gid / (n_rot / 2); uint32_t h = tmp % n_head; uint32_t t = tmp / n_head; uint32_t n_nope = head_dim - n_rot; uint32_t i = pair * 2; float corr0 = 0.0f, corr1 = 0.0f; if (ext_factor != 0.0f) { float denom = 2.0f * logf(freq_base); corr0 = floorf((float)n_rot * logf((float)n_ctx_orig / (beta_fast * 2.0f * (float)M_PI)) / denom); corr1 = ceilf((float)n_rot * logf((float)n_ctx_orig / (beta_slow * 2.0f * (float)M_PI)) / denom); corr0 = fmaxf(0.0f, corr0); corr1 = fminf((float)(n_rot - 1), corr1); } const float theta_scale = powf(freq_base, -2.0f / (float)n_rot); float theta_extrap = (float)(pos0 + t * pos_stride) * powf(theta_scale, (float)pair); float theta_interp = freq_scale * theta_extrap; float theta = theta_interp; float mscale = attn_factor; if (ext_factor != 0.0f) { float ramp_mix = rope_yarn_ramp_dev(corr0, corr1, (int)i) * ext_factor; theta = theta_interp * (1.0f - ramp_mix) + theta_extrap * ramp_mix; mscale *= 1.0f + 0.1f * logf(1.0f / freq_scale); } float c = cosf(theta) * mscale; float s = sinf(theta) * mscale; if (inverse) s = -s; float *tail = x + ((uint64_t)t * n_head + h) * head_dim + n_nope; float x0 = tail[i]; float x1 = tail[i + 1]; tail[i] = x0 * c - x1 * s; tail[i + 1] = x0 * s + x1 * c; } __device__ static float dsv4_e4m3fn_value_dev(int i) { int exp = (i >> 3) & 15; int mant = i & 7; if (exp == 0) return (float)mant * 0.001953125f; return (1.0f + (float)mant * 0.125f) * exp2f((float)exp - 7.0f); } __device__ static float dsv4_e4m3fn_dequant_dev(float x) { float sign = x < 0.0f ? -1.0f : 1.0f; float ax = fminf(fabsf(x), 448.0f); int lo = 0, hi = 126; while (lo < hi) { int mid = (lo + hi + 1) >> 1; if (dsv4_e4m3fn_value_dev(mid) <= ax) lo = mid; else hi = mid - 1; } int best = lo; if (best < 126) { float bd = fabsf(ax - dsv4_e4m3fn_value_dev(best)); float nd = fabsf(ax - dsv4_e4m3fn_value_dev(best + 1)); if (nd < bd || (nd == bd && (((best + 1) & 1) == 0) && ((best & 1) != 0))) best++; } return sign * dsv4_e4m3fn_value_dev(best); } __device__ static float dsv4_e2m1fn_value_dev(int i) { switch (i & 7) { case 0: return 0.0f; case 1: return 0.5f; case 2: return 1.0f; case 3: return 1.5f; case 4: return 2.0f; case 5: return 3.0f; case 6: return 4.0f; default: return 6.0f; } } __device__ static float dsv4_e2m1fn_dequant_dev(float x) { float sign = x < 0.0f ? -1.0f : 1.0f; float ax = fminf(fabsf(x), 6.0f); int best = 0; float best_diff = fabsf(ax - dsv4_e2m1fn_value_dev(0)); for (int i = 1; i < 8; i++) { float diff = fabsf(ax - dsv4_e2m1fn_value_dev(i)); if (diff < best_diff || (diff == best_diff && ((i & 1) == 0) && ((best & 1) != 0))) { best = i; best_diff = diff; } } return sign * dsv4_e2m1fn_value_dev(best); } __device__ static float model_scalar_dev(const void *base, uint64_t offset, uint32_t type, uint64_t idx) { const char *p = (const char *)base + offset; if (type == 1u) return __half2float(((const __half *)p)[idx]); return ((const float *)p)[idx]; } __device__ static float model_ape_value_dev(const void *base, uint64_t offset, uint32_t type, uint32_t width, uint32_t row, uint32_t col) { const char *p = (const char *)base + offset; if (type == 1u) return __half2float(((const __half *)p)[(uint64_t)row * width + col]); if (type == 8u) { const uint64_t row_bytes = ((uint64_t)width + 31u) / 32u * 34u; const unsigned char *blk = (const unsigned char *)p + (uint64_t)row * row_bytes + (uint64_t)(col >> 5) * 34u; const float d = q8_0_scale_scalar(blk); const int8_t q = ((const int8_t *)(blk + 2u))[col & 31u]; return d * (float)q; } return ((const float *)p)[(uint64_t)row * width + col]; } __device__ static float rope_yarn_ramp_cpu_equiv_dev(float low, float high, int i0) { float y = ((float)(i0 / 2) - low) / fmaxf(0.001f, high - low); return 1.0f - fminf(1.0f, fmaxf(0.0f, y)); } __device__ static DS4_ROCM_UNUSED void rope_tail_one_dev(float *x, uint32_t head_dim, uint32_t n_rot, uint32_t pos, uint32_t n_ctx_orig, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow) { uint32_t n_nope = head_dim - n_rot; float corr0 = 0.0f, corr1 = 0.0f; if (ext_factor != 0.0f) { float denom = 2.0f * logf(freq_base); corr0 = fmaxf(0.0f, floorf((float)n_rot * logf((float)n_ctx_orig / (beta_fast * 2.0f * (float)M_PI)) / denom)); corr1 = fminf((float)(n_rot - 1), ceilf((float)n_rot * logf((float)n_ctx_orig / (beta_slow * 2.0f * (float)M_PI)) / denom)); } for (uint32_t i = 0; i < n_rot; i += 2) { float theta_extrap = (float)pos * powf(freq_base, -((float)i) / (float)n_rot); float theta_interp = freq_scale * theta_extrap; float theta = theta_interp; float mscale = attn_factor; if (ext_factor != 0.0f) { float mix = rope_yarn_ramp_cpu_equiv_dev(corr0, corr1, (int)i) * ext_factor; theta = theta_interp * (1.0f - mix) + theta_extrap * mix; mscale *= 1.0f + 0.1f * logf(1.0f / freq_scale); } float c = cosf(theta) * mscale; float s = sinf(theta) * mscale; float x0 = x[n_nope + i]; float x1 = x[n_nope + i + 1]; x[n_nope + i] = x0 * c - x1 * s; x[n_nope + i + 1] = x0 * s + x1 * c; } } extern "C" int ds4_gpu_rms_norm_plain_tensor(ds4_gpu_tensor *out, const ds4_gpu_tensor *x, uint32_t n, float eps) { if (!cuda_tensor_has_f32(out, n) || !cuda_tensor_has_f32(x, n)) return 0; if (n == 0u) return 1; rms_norm_plain_kernel<<<1, 256>>>((float *)out->ptr, (const float *)x->ptr, n, 1, eps); return cuda_ok(cudaGetLastError(), "rms_norm_plain launch"); } extern "C" int ds4_gpu_rms_norm_plain_rows_tensor(ds4_gpu_tensor *out, const ds4_gpu_tensor *x, uint32_t n, uint32_t rows, float eps) { if (!cuda_tensor_has_elems2(out, n, rows, sizeof(float)) || !cuda_tensor_has_elems2(x, n, rows, sizeof(float))) return 0; if (n == 0u || rows == 0u) return 1; rms_norm_plain_kernel<<>>((float *)out->ptr, (const float *)x->ptr, n, rows, eps); return cuda_ok(cudaGetLastError(), "rms_norm_plain launch"); } extern "C" int ds4_gpu_rms_norm_weight_tensor(ds4_gpu_tensor *out, const ds4_gpu_tensor *x, const void *model_map, uint64_t model_size, uint64_t weight_offset, uint32_t n, float eps) { uint64_t weight_bytes = 0; if (!model_map || !cuda_u64_mul_checked(n, sizeof(float), &weight_bytes) || !cuda_model_range_fits(model_size, weight_offset, weight_bytes) || !cuda_tensor_has_f32(out, n) || !cuda_tensor_has_f32(x, n)) return 0; if (n == 0u) return 1; const char *wptr = cuda_model_range_ptr(model_map, weight_offset, weight_bytes, "rms_weight"); if (!wptr) return 0; const float *w = (const float *)wptr; rms_norm_weight_kernel<<<1, 256>>>((float *)out->ptr, (const float *)x->ptr, w, n, 1, eps); return cuda_ok(cudaGetLastError(), "rms_norm_weight launch"); } extern "C" int ds4_gpu_rms_norm_weight_rows_tensor(ds4_gpu_tensor *out, const ds4_gpu_tensor *x, const void *model_map, uint64_t model_size, uint64_t weight_offset, uint32_t n, uint32_t rows, float eps) { uint64_t weight_bytes = 0; if (!model_map || !cuda_u64_mul_checked(n, sizeof(float), &weight_bytes) || !cuda_model_range_fits(model_size, weight_offset, weight_bytes) || !cuda_tensor_has_elems2(out, n, rows, sizeof(float)) || !cuda_tensor_has_elems2(x, n, rows, sizeof(float))) return 0; if (n == 0u || rows == 0u) return 1; const char *wptr = cuda_model_range_ptr(model_map, weight_offset, weight_bytes, "rms_weight"); if (!wptr) return 0; const float *w = (const float *)wptr; rms_norm_weight_kernel<<>>((float *)out->ptr, (const float *)x->ptr, w, n, rows, eps); return cuda_ok(cudaGetLastError(), "rms_norm_weight launch"); } extern "C" int ds4_gpu_dsv4_qkv_rms_norm_rows_tensor( ds4_gpu_tensor *q_out, const ds4_gpu_tensor *q, const void *model_map, uint64_t model_size, uint64_t q_weight_offset, uint32_t q_n, ds4_gpu_tensor *kv_out, const ds4_gpu_tensor *kv, uint64_t kv_weight_offset, uint32_t kv_n, uint32_t rows, float eps) { uint64_t q_weight_bytes = 0, kv_weight_bytes = 0; if (!model_map || !cuda_u64_mul_checked(q_n, sizeof(float), &q_weight_bytes) || !cuda_u64_mul_checked(kv_n, sizeof(float), &kv_weight_bytes) || !cuda_model_range_fits(model_size, q_weight_offset, q_weight_bytes) || !cuda_model_range_fits(model_size, kv_weight_offset, kv_weight_bytes) || !cuda_tensor_has_elems2(q_out, q_n, rows, sizeof(float)) || !cuda_tensor_has_elems2(q, q_n, rows, sizeof(float)) || !cuda_tensor_has_elems2(kv_out, kv_n, rows, sizeof(float)) || !cuda_tensor_has_elems2(kv, kv_n, rows, sizeof(float))) { return 0; } if ((q_n == 0u && kv_n == 0u) || rows == 0u) return 1; const float *q_w = (const float *)cuda_model_range_ptr(model_map, q_weight_offset, q_weight_bytes, "q_rms_weight"); const float *kv_w = (const float *)cuda_model_range_ptr(model_map, kv_weight_offset, kv_weight_bytes, "kv_rms_weight"); if (!q_w || !kv_w) return 0; dim3 grid(rows, 2u, 1u); dsv4_qkv_rms_norm_rows_kernel<<>>( (float *)q_out->ptr, (const float *)q->ptr, q_w, q_n, (float *)kv_out->ptr, (const float *)kv->ptr, kv_w, kv_n, rows, eps); return cuda_ok(cudaGetLastError(), "dsv4 qkv rms norm rows launch"); } extern "C" int ds4_gpu_head_rms_norm_tensor(ds4_gpu_tensor *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, float eps) { uint64_t rows64 = 0; if (!cuda_u64_mul_checked(n_tok, n_head, &rows64) || rows64 > UINT32_MAX || !cuda_tensor_has_elems3(x, n_tok, n_head, head_dim, sizeof(float))) return 0; if (rows64 == 0u || head_dim == 0u) return 1; head_rms_norm_kernel<<<(uint32_t)rows64, 256>>>((float *)x->ptr, n_tok, n_head, head_dim, eps); return cuda_ok(cudaGetLastError(), "head_rms_norm launch"); } extern "C" int ds4_gpu_head_rms_norm_rope_tail_tensor(ds4_gpu_tensor *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, uint32_t n_rot, uint32_t pos0, uint32_t n_ctx_orig, bool inverse, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow, float eps) { uint64_t rows64 = 0; if (n_rot > head_dim || (n_rot & 1u) || !cuda_u64_mul_checked(n_tok, n_head, &rows64) || rows64 > UINT32_MAX || !cuda_tensor_has_elems3(x, n_tok, n_head, head_dim, sizeof(float))) return 0; if (rows64 == 0u || head_dim == 0u) return 1; head_rms_norm_rope_tail_kernel<<<(uint32_t)rows64, 256>>>((float *)x->ptr, n_tok, n_head, head_dim, n_rot, pos0, n_ctx_orig, inverse ? 1 : 0, freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow, eps); return cuda_ok(cudaGetLastError(), "head_rms_norm_rope_tail launch"); } extern "C" int ds4_gpu_attn_q_b_f16_head_rms_rope_tail_tensor( ds4_gpu_tensor *out, ds4_gpu_tensor *q_half, const void *model_map, uint64_t model_size, uint64_t weight_offset, uint64_t in_dim, uint64_t out_dim, const ds4_gpu_tensor *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, uint32_t n_rot, uint32_t pos0, uint32_t n_ctx_orig, bool inverse, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow, float eps) { if (!g_cublas_ready || !out || !q_half || !x || !model_map || n_tok == 0 || n_rot > head_dim || (n_rot & 1u) || out_dim != (uint64_t)n_head * head_dim || x->bytes < (uint64_t)n_tok * in_dim * sizeof(float) || out->bytes < (uint64_t)n_tok * out_dim * sizeof(float) || q_half->bytes < (uint64_t)n_tok * out_dim * sizeof(__half)) return 0; const uint64_t blocks = (in_dim + 31u) / 32u; if (weight_offset > model_size || out_dim > UINT64_MAX / (blocks * 34u)) return 0; const uint64_t weight_bytes = out_dim * blocks * 34u; if (weight_bytes > model_size - weight_offset) return 0; const __half *w_f16 = cuda_q8_f16_ptr(model_map, weight_offset, weight_bytes, in_dim, out_dim, "attn_q_b"); if (!w_f16) return 0; const uint64_t xh_count = (uint64_t)n_tok * in_dim; __half *xh = (__half *)cuda_tmp_alloc(xh_count * sizeof(__half), "attn q_b f16 activations"); if (!xh) return 0; f32_to_f16_kernel<<<(xh_count + 255u) / 256u, 256>>>(xh, (const float *)x->ptr, xh_count); if (!cuda_ok(cudaGetLastError(), "attn q_b f16 activation convert launch")) return 0; const float alpha = 1.0f; const float beta = 0.0f; cublasStatus_t st = cublasGemmEx(g_cublas, CUBLAS_OP_T, CUBLAS_OP_N, (int)out_dim, (int)n_tok, (int)in_dim, &alpha, w_f16, CUDA_R_16F, (int)in_dim, xh, CUDA_R_16F, (int)in_dim, &beta, q_half->ptr, CUDA_R_16F, (int)out_dim, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT); if (st != CUBLAS_STATUS_SUCCESS) { fprintf(stderr, "ds4: " DS4_GPU_BLAS_NAME " attn q_b f16-out matmul failed: status %d\n", (int)st); return 0; } head_rms_norm_rope_tail_from_half_kernel<<>>( (float *)out->ptr, (const __half *)q_half->ptr, n_tok, n_head, head_dim, n_rot, pos0, n_ctx_orig, inverse ? 1 : 0, freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow, eps); return cuda_ok(cudaGetLastError(), "attn q_b f16-out head_rms_norm_rope launch"); } static int cuda_rope_tail_stride_tensor(ds4_gpu_tensor *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, uint32_t n_rot, uint32_t pos0, uint32_t pos_stride, uint32_t n_ctx_orig, bool inverse, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow) { if (!x || n_rot > head_dim || (n_rot & 1) || x->bytes < (uint64_t)n_tok * n_head * head_dim * sizeof(float)) return 0; uint32_t pairs = n_tok * n_head * (n_rot / 2); rope_tail_kernel<<<(pairs + 255) / 256, 256>>>((float *)x->ptr, n_tok, n_head, head_dim, n_rot, pos0, pos_stride, n_ctx_orig, inverse ? 1 : 0, freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow); return cuda_ok(cudaGetLastError(), "rope_tail launch"); } extern "C" int ds4_gpu_rope_tail_tensor(ds4_gpu_tensor *x, uint32_t n_tok, uint32_t n_head, uint32_t head_dim, uint32_t n_rot, uint32_t pos0, uint32_t n_ctx_orig, bool inverse, float freq_base, float freq_scale, float ext_factor, float attn_factor, float beta_fast, float beta_slow) { return cuda_rope_tail_stride_tensor(x, n_tok, n_head, head_dim, n_rot, pos0, 1u, n_ctx_orig, inverse, freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow); }