mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-09-19 01:05:09 +02:00
Merge branch 'upstream' into concedo_experimental
# Conflicts: # .github/workflows/build-cuda-windows.yml # .github/workflows/release.yml # app/llama.cpp # build-xcframework.sh # docs/speculative.md # ggml/src/ggml-opencl/ggml-opencl.cpp # ggml/src/ggml-opencl/kernels/cvt.cl # ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp # ggml/src/ggml-webgpu/ggml-webgpu.cpp # ggml/src/ggml-webgpu/wgsl-shaders/set_rows.wgsl # requirements/requirements-server-bench.txt # scripts/server-bench.py # tests/test-backend-ops.cpp # tests/test-llama-archs.cpp # tools/cli/README.md # tools/completion/README.md # tools/server/README.md
This commit is contained in:
@@ -0,0 +1,22 @@
|
||||
name: "ccache-clear"
|
||||
description: "Delete all GitHub Actions caches matching a key prefix"
|
||||
inputs:
|
||||
key:
|
||||
description: "Cache key prefix to match and delete"
|
||||
required: true
|
||||
|
||||
runs:
|
||||
using: "composite"
|
||||
steps:
|
||||
- name: Clear caches
|
||||
shell: bash
|
||||
run: |
|
||||
CACHES=$(gh cache list --key "ccache-${{ inputs.key }}" --json id,key --jq '.[] | "\(.id) \(.key)"' 2>/dev/null)
|
||||
if [ -z "$CACHES" ]; then
|
||||
echo "No caches found with key prefix: ${{ inputs.key }}"
|
||||
exit 0
|
||||
fi
|
||||
while read -r id key; do
|
||||
echo "Deleting cache: $id ($key)"
|
||||
gh cache delete "$id"
|
||||
done <<< "$CACHES"
|
||||
+4
-1
@@ -1447,6 +1447,9 @@ class TextModel(ModelBase):
|
||||
if chkhsh == "0fe1cf6eda062318a1af7270f3331a85c539a01778ff948e24388e949c5282f4":
|
||||
# ref: https://huggingface.co/evilfreelancer/ruGPT3XL
|
||||
res = "gpt-2"
|
||||
if chkhsh == "9e454714343b69b99b71795c1d27a68c2a1d15dab111f4d353109f966af29da7":
|
||||
# ref: https://huggingface.co/LiquidAI/LFM2.5-8B-A1B
|
||||
res = "lfm2"
|
||||
if chkhsh == "0ef9807a4087ebef797fc749390439009c3b9eda9ad1a097abbe738f486c01e5":
|
||||
# ref: https://huggingface.co/meta-llama/Meta-Llama-3-8B
|
||||
res = "llama-bpe"
|
||||
@@ -1598,7 +1601,7 @@ class TextModel(ModelBase):
|
||||
# ref: https://huggingface.co/K-intelligence/Midm-2.0-Base-Instruct
|
||||
res = "midm-2.0"
|
||||
if chkhsh == "169bf0296a13c4d9b7672313f749eb36501d931022de052aad6e36f2bf34dd51":
|
||||
# ref: https://huggingface.co/LiquidAI/LFM2-Tokenizer
|
||||
# ref: https://huggingface.co/LiquidAI/LFM2.5-350M
|
||||
res = "lfm2"
|
||||
if chkhsh == "2085e1638f6c377a0aa4ead21b27bb4cb941bf800df86ed391011769c1758dfb":
|
||||
# ref: https://huggingface.co/LGAI-EXAONE/EXAONE-4.0-32B
|
||||
|
||||
@@ -139,7 +139,7 @@ models = [
|
||||
{"name": "seed-coder", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/ByteDance-Seed/Seed-Coder-8B-Base", },
|
||||
{"name": "a.x-4.0", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/skt/A.X-4.0", },
|
||||
{"name": "midm-2.0", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/K-intelligence/Midm-2.0-Base-Instruct", },
|
||||
{"name": "lfm2", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/LiquidAI/LFM2-Tokenizer"},
|
||||
{"name": "lfm2", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/LiquidAI/LFM2.5-350M", },
|
||||
{"name": "exaone4", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/LGAI-EXAONE/EXAONE-4.0-32B", },
|
||||
{"name": "mellum", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/JetBrains/Mellum-4b-base", },
|
||||
{"name": "modern-bert", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/answerdotai/ModernBERT-base", },
|
||||
@@ -183,6 +183,8 @@ pre_computed_hashes = [
|
||||
# jina-v2-de variants
|
||||
{"name": "jina-v2-de", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/aari1995/German_Semantic_V3", "chkhsh": "b3d1dd861f1d4c5c0d2569ce36baf3f90fe8a102db3de50dd71ff860d91be3df"},
|
||||
{"name": "gpt-2", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/evilfreelancer/ruGPT3XL", "chkhsh": "0fe1cf6eda062318a1af7270f3331a85c539a01778ff948e24388e949c5282f4"},
|
||||
# lfm2 variants
|
||||
{"name": "lfm2", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/LiquidAI/LFM2.5-8B-A1B", "chkhsh": "9e454714343b69b99b71795c1d27a68c2a1d15dab111f4d353109f966af29da7"},
|
||||
]
|
||||
|
||||
|
||||
|
||||
@@ -977,6 +977,35 @@ void ggml_vec_dot_q8_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
sumf = hsum_float_8(acc);
|
||||
|
||||
*s = sumf;
|
||||
|
||||
#elif defined(__loongarch_sx)
|
||||
|
||||
__m128 acc = (__m128)__lsx_vldi(0);
|
||||
|
||||
for (; ib < nb; ++ib) {
|
||||
const float d = GGML_CPU_FP16_TO_FP32(x[ib].d) * GGML_CPU_FP16_TO_FP32(y[ib].d);
|
||||
const __m128i qx_0 = __lsx_vld((const __m128i *)x[ib].qs, 0);
|
||||
const __m128i qx_1 = __lsx_vld((const __m128i *)x[ib].qs + 1, 0);
|
||||
const __m128i qy_0 = __lsx_vld((const __m128i *)y[ib].qs, 0);
|
||||
const __m128i qy_1 = __lsx_vld((const __m128i *)y[ib].qs + 1, 0);
|
||||
|
||||
const __m128i p16_0 = lsx_maddubs_h(qx_0, qy_0);
|
||||
const __m128i p16_1 = lsx_maddubs_h(qx_1, qy_1);
|
||||
|
||||
// Sum int16 pairs → int32
|
||||
const __m128i s_0 = __lsx_vaddwev_w_h(p16_0, p16_1);
|
||||
const __m128i s_1 = __lsx_vaddwod_w_h(p16_0, p16_1);
|
||||
|
||||
const __m128 q = __lsx_vffint_s_w(__lsx_vadd_w(s_0, s_1));
|
||||
acc = __lsx_vfmadd_s(__lsx_vreplfr2vr_s(d), q, acc);
|
||||
}
|
||||
|
||||
__m128 res = lsx_hadd_s(acc, acc);
|
||||
res = lsx_hadd_s(res, res);
|
||||
sumf = ((v4f32)res)[0];
|
||||
|
||||
*s = sumf;
|
||||
|
||||
#else
|
||||
UNUSED(nb);
|
||||
UNUSED(ib);
|
||||
@@ -1443,6 +1472,99 @@ void ggml_vec_dot_q6_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
|
||||
|
||||
*s = hsum_float_8(acc);
|
||||
|
||||
#elif defined(__loongarch_sx)
|
||||
|
||||
const __m128i m32s = __lsx_vreplgr2vr_b(32);
|
||||
|
||||
__m128 acc_0 = (__m128)__lsx_vldi(0);
|
||||
__m128 acc_1 = (__m128)__lsx_vldi(0);
|
||||
|
||||
for (int i = 0; i < nb; ++i) {
|
||||
|
||||
const float d = y[i].d * GGML_CPU_FP16_TO_FP32(x[i].d);
|
||||
|
||||
const uint8_t * GGML_RESTRICT q4 = x[i].ql;
|
||||
const uint8_t * GGML_RESTRICT qh = x[i].qh;
|
||||
const int8_t * GGML_RESTRICT q8 = y[i].qs;
|
||||
|
||||
const __m128i scale_i8 = __lsx_vld(x[i].scales, 0);
|
||||
const __m128i scales_lo = __lsx_vsllwil_h_b(scale_i8, 0);
|
||||
const __m128i scales_hi = __lsx_vsllwil_h_b(__lsx_vbsrl_v(scale_i8, 8), 0);
|
||||
|
||||
__m128i sumi_0 = __lsx_vldi(0);
|
||||
__m128i sumi_1 = __lsx_vldi(0);
|
||||
|
||||
for (int j = 0; j < QK_K/128; ++j) {
|
||||
|
||||
const __m128i q4bitsH_0 = __lsx_vld((const __m128i*)qh, 0); qh += 16;
|
||||
const __m128i q4bitsH_1 = __lsx_vld((const __m128i*)qh, 0); qh += 16;
|
||||
|
||||
const __m128i q4h_0 = __lsx_vslli_b(__lsx_vandi_b(q4bitsH_0, 3), 4);
|
||||
const __m128i q4h_1 = __lsx_vslli_b(__lsx_vandi_b(q4bitsH_1, 3), 4);
|
||||
const __m128i q4h_2 = __lsx_vslli_b(__lsx_vandi_b(q4bitsH_0, 3 << 2), 2);
|
||||
const __m128i q4h_3 = __lsx_vslli_b(__lsx_vandi_b(q4bitsH_1, 3 << 2), 2);
|
||||
const __m128i q4h_4 = __lsx_vandi_b(q4bitsH_0, 3 << 4);
|
||||
const __m128i q4h_5 = __lsx_vandi_b(q4bitsH_1, 3 << 4);
|
||||
const __m128i q4h_6 = __lsx_vsrli_b(__lsx_vandi_b(q4bitsH_0, 3 << 6), 2);
|
||||
const __m128i q4h_7 = __lsx_vsrli_b(__lsx_vandi_b(q4bitsH_1, 3 << 6), 2);
|
||||
|
||||
const __m128i q4bits1_0 = __lsx_vld((const __m128i*)q4, 0); q4 += 16;
|
||||
const __m128i q4bits1_1 = __lsx_vld((const __m128i*)q4, 0); q4 += 16;
|
||||
const __m128i q4bits2_0 = __lsx_vld((const __m128i*)q4, 0); q4 += 16;
|
||||
const __m128i q4bits2_1 = __lsx_vld((const __m128i*)q4, 0); q4 += 16;
|
||||
|
||||
const __m128i q4_0 = __lsx_vor_v(__lsx_vandi_b(q4bits1_0, 0xf), q4h_0);
|
||||
const __m128i q4_1 = __lsx_vor_v(__lsx_vandi_b(q4bits1_1, 0xf), q4h_1);
|
||||
const __m128i q4_2 = __lsx_vor_v(__lsx_vandi_b(q4bits2_0, 0xf), q4h_2);
|
||||
const __m128i q4_3 = __lsx_vor_v(__lsx_vandi_b(q4bits2_1, 0xf), q4h_3);
|
||||
const __m128i q4_4 = __lsx_vor_v(__lsx_vsrli_b(q4bits1_0, 4), q4h_4);
|
||||
const __m128i q4_5 = __lsx_vor_v(__lsx_vsrli_b(q4bits1_1, 4), q4h_5);
|
||||
const __m128i q4_6 = __lsx_vor_v(__lsx_vsrli_b(q4bits2_0, 4), q4h_6);
|
||||
const __m128i q4_7 = __lsx_vor_v(__lsx_vsrli_b(q4bits2_1, 4), q4h_7);
|
||||
|
||||
const __m128i q8_0 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q8_1 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q8_2 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q8_3 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q8_4 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q8_5 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q8_6 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q8_7 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
|
||||
__m128i p16_0 = lsx_maddubs_h(__lsx_vsub_b(q4_0, m32s), q8_0);
|
||||
__m128i p16_1 = lsx_maddubs_h(__lsx_vsub_b(q4_1, m32s), q8_1);
|
||||
__m128i p16_2 = lsx_maddubs_h(__lsx_vsub_b(q4_2, m32s), q8_2);
|
||||
__m128i p16_3 = lsx_maddubs_h(__lsx_vsub_b(q4_3, m32s), q8_3);
|
||||
__m128i p16_4 = lsx_maddubs_h(__lsx_vsub_b(q4_4, m32s), q8_4);
|
||||
__m128i p16_5 = lsx_maddubs_h(__lsx_vsub_b(q4_5, m32s), q8_5);
|
||||
__m128i p16_6 = lsx_maddubs_h(__lsx_vsub_b(q4_6, m32s), q8_6);
|
||||
__m128i p16_7 = lsx_maddubs_h(__lsx_vsub_b(q4_7, m32s), q8_7);
|
||||
|
||||
const __m128i sc_vec = j == 0 ? scales_lo : scales_hi;
|
||||
|
||||
p16_0 = lsx_madd_h(__lsx_vreplvei_h(sc_vec, 0), p16_0);
|
||||
p16_1 = lsx_madd_h(__lsx_vreplvei_h(sc_vec, 1), p16_1);
|
||||
p16_2 = lsx_madd_h(__lsx_vreplvei_h(sc_vec, 2), p16_2);
|
||||
p16_3 = lsx_madd_h(__lsx_vreplvei_h(sc_vec, 3), p16_3);
|
||||
p16_4 = lsx_madd_h(__lsx_vreplvei_h(sc_vec, 4), p16_4);
|
||||
p16_5 = lsx_madd_h(__lsx_vreplvei_h(sc_vec, 5), p16_5);
|
||||
p16_6 = lsx_madd_h(__lsx_vreplvei_h(sc_vec, 6), p16_6);
|
||||
p16_7 = lsx_madd_h(__lsx_vreplvei_h(sc_vec, 7), p16_7);
|
||||
|
||||
sumi_0 = __lsx_vadd_w(sumi_0, __lsx_vadd_w(p16_0, p16_2));
|
||||
sumi_1 = __lsx_vadd_w(sumi_1, __lsx_vadd_w(p16_1, p16_3));
|
||||
sumi_0 = __lsx_vadd_w(sumi_0, __lsx_vadd_w(p16_4, p16_6));
|
||||
sumi_1 = __lsx_vadd_w(sumi_1, __lsx_vadd_w(p16_5, p16_7));
|
||||
}
|
||||
|
||||
__m128 p_0 = __lsx_vfmul_s(__lsx_vreplfr2vr_s(d), __lsx_vffint_s_w(sumi_0));
|
||||
__m128 p_1 = __lsx_vfmul_s(__lsx_vreplfr2vr_s(d), __lsx_vffint_s_w(sumi_1));
|
||||
acc_0 = __lsx_vfadd_s(p_0, acc_0);
|
||||
acc_1 = __lsx_vfadd_s(p_1, acc_1);
|
||||
}
|
||||
|
||||
*s = hsum_float_4x4(acc_0, acc_1, (__m128)__lsx_vldi(0), (__m128)__lsx_vldi(0));
|
||||
|
||||
#else
|
||||
UNUSED(x);
|
||||
UNUSED(y);
|
||||
@@ -2149,6 +2271,35 @@ void ggml_vec_dot_iq4_xs_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const v
|
||||
|
||||
*s = hsum_float_8(accum);
|
||||
|
||||
#elif defined(__loongarch_sx)
|
||||
|
||||
const __m128i values128 = __lsx_vld((const __m128i*)kvalues_iq4nl, 0);
|
||||
|
||||
__m128 accum = (__m128)__lsx_vldi(0);
|
||||
for (int ibl = 0; ibl < nb; ++ibl) {
|
||||
const uint8_t * qs = x[ibl].qs;
|
||||
const int8_t * q8 = y[ibl].qs;
|
||||
uint16_t sh = x[ibl].scales_h;
|
||||
__m128i sumi = __lsx_vldi(0);
|
||||
for (int ib = 0; ib < QK_K/32; ++ib) {
|
||||
const __m128i q4bits = __lsx_vld((const __m128i*)qs, 0); qs += 16;
|
||||
const __m128i q8b_0 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q8b_1 = __lsx_vld((const __m128i*)q8, 0); q8 += 16;
|
||||
const __m128i q4b_0 = __lsx_vshuf_b(values128, values128, __lsx_vandi_b(q4bits, 0xf));
|
||||
const __m128i q4b_1 = __lsx_vshuf_b(values128, values128, __lsx_vsrli_b(q4bits, 4));
|
||||
const __m128i p16_0 = lsx_maddubs_h(q4b_0, q8b_0);
|
||||
const __m128i p16_1 = lsx_maddubs_h(q4b_1, q8b_1);
|
||||
const int16_t ls = (((x[ibl].scales_l[ib/2] >> ((ib & 1) * 4)) & 0xf) | ((sh & 0x3) << 4)) - 32;
|
||||
sh >>= 2;
|
||||
sumi = __lsx_vadd_w(lsx_madd_h(p16_0, __lsx_vreplgr2vr_h(ls)), sumi);
|
||||
sumi = __lsx_vadd_w(lsx_madd_h(p16_1, __lsx_vreplgr2vr_h(ls)), sumi);
|
||||
}
|
||||
const float ds = GGML_CPU_FP16_TO_FP32(x[ibl].d) * y[ibl].d;
|
||||
accum = __lsx_vfadd_s(__lsx_vfmul_s(__lsx_vreplfr2vr_s(ds), __lsx_vffint_s_w(sumi)), accum);
|
||||
}
|
||||
|
||||
*s = ((v4f32)lsx_hadd_s(lsx_hadd_s(accum, accum), lsx_hadd_s(accum, accum)))[0];
|
||||
|
||||
#else
|
||||
UNUSED(x);
|
||||
UNUSED(y);
|
||||
|
||||
@@ -1125,25 +1125,12 @@ static inline void __lasx_f32cx8_store(ggml_fp16_t * x, __m256 y) {
|
||||
#define GGML_F16_EPR 4
|
||||
|
||||
static inline __m128 __lsx_f16x4_load(const ggml_fp16_t * x) {
|
||||
float tmp[4];
|
||||
|
||||
tmp[0] = GGML_CPU_FP16_TO_FP32(x[0]);
|
||||
tmp[1] = GGML_CPU_FP16_TO_FP32(x[1]);
|
||||
tmp[2] = GGML_CPU_FP16_TO_FP32(x[2]);
|
||||
tmp[3] = GGML_CPU_FP16_TO_FP32(x[3]);
|
||||
|
||||
return (__m128)__lsx_vld(tmp, 0);
|
||||
return __lsx_vfcvtl_s_h(__lsx_vld((const void *)x, 0));
|
||||
}
|
||||
|
||||
static inline void __lsx_f16x4_store(ggml_fp16_t * x, __m128 y) {
|
||||
float arr[4];
|
||||
|
||||
__lsx_vst(y, arr, 0);
|
||||
|
||||
x[0] = GGML_CPU_FP32_TO_FP16(arr[0]);
|
||||
x[1] = GGML_CPU_FP32_TO_FP16(arr[1]);
|
||||
x[2] = GGML_CPU_FP32_TO_FP16(arr[2]);
|
||||
x[3] = GGML_CPU_FP32_TO_FP16(arr[3]);
|
||||
__m128i a = __lsx_vfcvt_h_s(y, y);
|
||||
memcpy(x, &a, sizeof(ggml_fp16_t) * 4);
|
||||
}
|
||||
|
||||
#define GGML_F32Cx4 __m128
|
||||
|
||||
@@ -1732,6 +1732,8 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rope(ggml_metal_
|
||||
ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_im2col(ggml_metal_library_t lib, const ggml_tensor * op) {
|
||||
assert(op->op == GGML_OP_IM2COL);
|
||||
|
||||
GGML_TENSOR_LOCALS(int64_t, ne0, op->src[0], ne);
|
||||
|
||||
GGML_ASSERT(ggml_is_contiguous(op->src[1]));
|
||||
GGML_ASSERT(op->src[1]->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(op->type == GGML_TYPE_F16 || op->type == GGML_TYPE_F32);
|
||||
@@ -1739,7 +1741,11 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_im2col(ggml_meta
|
||||
char base[256];
|
||||
char name[256];
|
||||
|
||||
snprintf(base, 256, "kernel_im2col_%s", ggml_type_name(op->type));
|
||||
if (ne00*ne01 <= 1024) {
|
||||
snprintf(base, 256, "kernel_im2col_%s", ggml_type_name(op->type));
|
||||
} else {
|
||||
snprintf(base, 256, "kernel_im2col_ext_%s", ggml_type_name(op->type));
|
||||
}
|
||||
snprintf(name, 256, "%s", base);
|
||||
|
||||
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
|
||||
|
||||
@@ -3635,16 +3635,26 @@ int ggml_metal_op_im2col(ggml_metal_op_t ctx, int idx) {
|
||||
|
||||
auto pipeline = ggml_metal_library_get_pipeline_im2col(lib, op);
|
||||
|
||||
GGML_ASSERT(KH*KW <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline));
|
||||
if (KH*KW <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)) {
|
||||
const uint64_t ntptg0 = std::min(ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)/(KH*KW), N);
|
||||
|
||||
const uint64_t ntptg0 = std::min(ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)/(KH*KW), N);
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 1);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 2);
|
||||
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 1);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 2);
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, IC, OH, OW, ntptg0, KH, KW);
|
||||
} else {
|
||||
const uint64_t n_threads = std::min(ggml_metal_pipeline_max_theads_per_threadgroup(pipeline), N);
|
||||
const int64_t quotient = N / n_threads + (N % n_threads > 0 ? 1 : 0);
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, IC, OH, OW, ntptg0, KH, KW);
|
||||
ggml_metal_encoder_set_pipeline(enc, pipeline);
|
||||
ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 1);
|
||||
ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 2);
|
||||
|
||||
ggml_metal_encoder_dispatch_threadgroups(enc, quotient * CHW, OH, OW, n_threads, 1, 1);
|
||||
}
|
||||
|
||||
return 1;
|
||||
}
|
||||
|
||||
@@ -4696,59 +4696,59 @@ kernel void kernel_im2col(
|
||||
template [[host_name("kernel_im2col_f32")]] kernel im2col_t kernel_im2col<float>;
|
||||
template [[host_name("kernel_im2col_f16")]] kernel im2col_t kernel_im2col<half>;
|
||||
|
||||
// TODO: obsolete -- remove
|
||||
//typedef void (im2col_ext_t)(
|
||||
// constant ggml_metal_kargs_im2col & args,
|
||||
// device const float * x,
|
||||
// device char * dst,
|
||||
// uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
// uint3 tgpg[[threadgroups_per_grid]],
|
||||
// uint3 tpitg[[thread_position_in_threadgroup]],
|
||||
// uint3 ntg[[threads_per_threadgroup]]);
|
||||
//
|
||||
//template <typename T>
|
||||
//kernel void kernel_im2col_ext(
|
||||
// constant ggml_metal_kargs_im2col & args,
|
||||
// device const float * x,
|
||||
// device char * dst,
|
||||
// uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
// uint3 tgpg[[threadgroups_per_grid]], // tgpg[0] = D x IC x KH x KW, CHW = IC x KH x KW
|
||||
// uint3 tpitg[[thread_position_in_threadgroup]],
|
||||
// uint3 ntg[[threads_per_threadgroup]]) { // [M, 1, 1]
|
||||
// const int64_t KHW = (int64_t)args.KHW;
|
||||
//
|
||||
// const int64_t d = tgpig[0] / args.CHW;
|
||||
// const int64_t chw = tgpig[0] % args.CHW;
|
||||
// const int64_t tgpig_0 = chw / KHW; // 0 ~ (IC - 1)
|
||||
// const int64_t HW = tgpig[0] % KHW;
|
||||
//
|
||||
// const int64_t tpitg_0 = (d * ntg[0]) + tpitg[0];
|
||||
// if (tpitg_0 >= args.N) {
|
||||
// return;
|
||||
// }
|
||||
//
|
||||
// const int64_t tpitg_1 = HW / args.KW;
|
||||
// const int64_t tpitg_2 = HW % args.KW;
|
||||
//
|
||||
// const int64_t iiw = tgpig[2] * args.s0 + tpitg_2 * args.d0 - args.p0;
|
||||
// const int64_t iih = tgpig[1] * args.s1 + tpitg_1 * args.d1 - args.p1;
|
||||
//
|
||||
// const int64_t offset_dst =
|
||||
// (tpitg_0 * tgpg[1] * tgpg[2] + tgpig[1] * tgpg[2] + tgpig[2]) * args.CHW +
|
||||
// (tgpig_0 * KHW + tpitg_1 * args.KW + tpitg_2);
|
||||
//
|
||||
// device T * pdst = (device T *) (dst);
|
||||
//
|
||||
// if (iih < 0 || iih >= args.IH || iiw < 0 || iiw >= args.IW) {
|
||||
// pdst[offset_dst] = 0.0f;
|
||||
// } else {
|
||||
// const int64_t offset_src = tpitg_0 * args.ofs0 + tgpig_0 * args.ofs1;
|
||||
// pdst[offset_dst] = x[offset_src + iih * args.IW + iiw];
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//template [[host_name("kernel_im2col_ext_f32")]] kernel im2col_ext_t kernel_im2col_ext<float>;
|
||||
//template [[host_name("kernel_im2col_ext_f16")]] kernel im2col_ext_t kernel_im2col_ext<half>;
|
||||
// TODO: optimize
|
||||
typedef void (im2col_ext_t)(
|
||||
constant ggml_metal_kargs_im2col & args,
|
||||
device const float * x,
|
||||
device char * dst,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
uint3 tgpg[[threadgroups_per_grid]],
|
||||
uint3 tpitg[[thread_position_in_threadgroup]],
|
||||
uint3 ntg[[threads_per_threadgroup]]);
|
||||
|
||||
template <typename T>
|
||||
kernel void kernel_im2col_ext(
|
||||
constant ggml_metal_kargs_im2col & args,
|
||||
device const float * x,
|
||||
device char * dst,
|
||||
uint3 tgpig[[threadgroup_position_in_grid]],
|
||||
uint3 tgpg[[threadgroups_per_grid]], // tgpg[0] = D x IC x KH x KW, CHW = IC x KH x KW
|
||||
uint3 tpitg[[thread_position_in_threadgroup]],
|
||||
uint3 ntg[[threads_per_threadgroup]]) { // [M, 1, 1]
|
||||
const int64_t KHW = (int64_t)args.KHW;
|
||||
|
||||
const int64_t d = tgpig[0] / args.CHW;
|
||||
const int64_t chw = tgpig[0] % args.CHW;
|
||||
const int64_t tgpig_0 = chw / KHW; // 0 ~ (IC - 1)
|
||||
const int64_t HW = tgpig[0] % KHW;
|
||||
|
||||
const int64_t tpitg_0 = (d * ntg[0]) + tpitg[0];
|
||||
if (tpitg_0 >= args.N) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int64_t tpitg_1 = HW / args.KW;
|
||||
const int64_t tpitg_2 = HW % args.KW;
|
||||
|
||||
const int64_t iiw = tgpig[2] * args.s0 + tpitg_2 * args.d0 - args.p0;
|
||||
const int64_t iih = tgpig[1] * args.s1 + tpitg_1 * args.d1 - args.p1;
|
||||
|
||||
const int64_t offset_dst =
|
||||
(tpitg_0 * tgpg[1] * tgpg[2] + tgpig[1] * tgpg[2] + tgpig[2]) * args.CHW +
|
||||
(tgpig_0 * KHW + tpitg_1 * args.KW + tpitg_2);
|
||||
|
||||
device T * pdst = (device T *) (dst);
|
||||
|
||||
if (iih < 0 || iih >= args.IH || iiw < 0 || iiw >= args.IW) {
|
||||
pdst[offset_dst] = 0.0f;
|
||||
} else {
|
||||
const int64_t offset_src = tpitg_0 * args.ofs0 + tgpig_0 * args.ofs1;
|
||||
pdst[offset_dst] = x[offset_src + iih * args.IW + iiw];
|
||||
}
|
||||
}
|
||||
|
||||
template [[host_name("kernel_im2col_ext_f32")]] kernel im2col_ext_t kernel_im2col_ext<float>;
|
||||
template [[host_name("kernel_im2col_ext_f16")]] kernel im2col_ext_t kernel_im2col_ext<half>;
|
||||
|
||||
template <typename TK>
|
||||
kernel void kernel_conv_2d(
|
||||
|
||||
@@ -697,6 +697,7 @@ struct vk_device_struct {
|
||||
uint32_t coopmat_int_k;
|
||||
|
||||
bool coopmat2;
|
||||
bool coopmat2_bf16_support {};
|
||||
bool coopmat2_decode_vector;
|
||||
|
||||
bool pipeline_executable_properties_support {};
|
||||
@@ -3145,7 +3146,7 @@ struct vk_fa_tuning_params {
|
||||
};
|
||||
|
||||
static bool ggml_vk_flash_attn_scalar_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc, ggml_type k_type, ggml_type v_type);
|
||||
static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc);
|
||||
static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc, ggml_type k_type = GGML_TYPE_F16);
|
||||
|
||||
static vk_fa_tuning_params get_fa_tuning_params_scalar(const vk_device& device, uint32_t hsk, uint32_t hsv, uint32_t n_rows, uint32_t n_kv, ggml_type k_type, ggml_type v_type, bool f32acc) {
|
||||
|
||||
@@ -3285,6 +3286,13 @@ static vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_
|
||||
FaCodePath path = device->coopmat2 ? FA_COOPMAT2 :
|
||||
device->coopmat1_fa_support ? FA_COOPMAT1 : FA_SCALAR;
|
||||
|
||||
if (path == FA_COOPMAT2 && k_type == GGML_TYPE_BF16 && !device->coopmat2_bf16_support) {
|
||||
path = FA_COOPMAT1;
|
||||
}
|
||||
if (path == FA_COOPMAT1 && k_type == GGML_TYPE_BF16 && !device->coopmat_bf16_support) {
|
||||
path = FA_SCALAR;
|
||||
}
|
||||
|
||||
if (path == FA_COOPMAT1 && device->architecture == vk_device_architecture::NVIDIA_TURING) {
|
||||
// Nvidia compiler bug, see https://github.com/ggml-org/llama.cpp/pull/19075#issuecomment-3820716090
|
||||
path = FA_SCALAR;
|
||||
@@ -3294,7 +3302,7 @@ static vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_
|
||||
bool shape_ok = (f32acc && device->coopmat_support_16x16x16_f32acc) ||
|
||||
(!f32acc && device->coopmat_support_16x16x16_f16acc);
|
||||
const vk_fa_tuning_params params = get_fa_tuning_params_coopmat1(device, hsk, hsv, n_rows, n_kv, k_type, v_type, f32acc);
|
||||
bool shmem_ok = ggml_vk_flash_attn_coopmat_shmem_support(device, params, hsk, hsv, f32acc);
|
||||
bool shmem_ok = ggml_vk_flash_attn_coopmat_shmem_support(device, params, hsk, hsv, f32acc, k_type);
|
||||
|
||||
if (!shape_ok || !shmem_ok) {
|
||||
path = FA_SCALAR;
|
||||
@@ -3340,8 +3348,8 @@ static vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const
|
||||
|
||||
static std::vector<uint32_t> get_fa_spec_constants(const vk_fa_pipeline_state& state) {
|
||||
const auto fa_block_bytes = [](ggml_type t) -> uint32_t {
|
||||
// decodeBufF32 uses a block of vec4s for a better memory access pattern.
|
||||
return t == GGML_TYPE_F32 ? 16u : (uint32_t) ggml_type_size(t);
|
||||
if (t == GGML_TYPE_F32) return 16u;
|
||||
return (uint32_t) ggml_type_size(t);
|
||||
};
|
||||
return {
|
||||
/* 0 WorkGroupSize */ state.workgroup_size,
|
||||
@@ -3855,10 +3863,16 @@ static void ggml_vk_load_shaders(vk_device& device) {
|
||||
const uint32_t fa_sgs = fa.first.subgroup_size;
|
||||
const bool fa_ds = fa.first.subgroup_size == 0;
|
||||
|
||||
const bool bf16_kv = fa.first.k_type == GGML_TYPE_BF16;
|
||||
const bool use_mmq = ggml_vk_fa_scalar_uses_mmq(device, fa.first.k_type);
|
||||
const void * spv_data = nullptr;
|
||||
size_t spv_size = 0;
|
||||
if (use_mmq) {
|
||||
const char *name = nullptr;
|
||||
if (bf16_kv) {
|
||||
spv_data = flash_attn_f32_f16_fp32_data;
|
||||
spv_size = flash_attn_f32_f16_fp32_len;
|
||||
name = aligned ? "flash_attn_f32_bf16_aligned" : "flash_attn_f32_bf16";
|
||||
} else if (use_mmq) {
|
||||
#if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT)
|
||||
if (device->fp16) {
|
||||
if (f32acc) { spv_data = flash_attn_f32_f16_int8_data; spv_size = flash_attn_f32_f16_int8_len; }
|
||||
@@ -3868,6 +3882,7 @@ static void ggml_vk_load_shaders(vk_device& device) {
|
||||
spv_size = flash_attn_f32_f16_fp32_int8_len;
|
||||
}
|
||||
#endif
|
||||
name = aligned ? "flash_attn_f32_f16_aligned" : "flash_attn_f32_f16";
|
||||
} else {
|
||||
if (device->fp16) {
|
||||
if (f32acc) { spv_data = flash_attn_f32_f16_data; spv_size = flash_attn_f32_f16_len; }
|
||||
@@ -3876,8 +3891,8 @@ static void ggml_vk_load_shaders(vk_device& device) {
|
||||
spv_data = flash_attn_f32_f16_fp32_data;
|
||||
spv_size = flash_attn_f32_f16_fp32_len;
|
||||
}
|
||||
name = aligned ? "flash_attn_f32_f16_aligned" : "flash_attn_f32_f16";
|
||||
}
|
||||
const char *name = aligned ? "flash_attn_f32_f16_aligned" : "flash_attn_f32_f16";
|
||||
ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7,
|
||||
sizeof(vk_flash_attn_push_constants), {Br, 1, 1},
|
||||
get_fa_spec_constants(fa.first), aligned ? Bc : 1, true,
|
||||
@@ -3895,11 +3910,25 @@ static void ggml_vk_load_shaders(vk_device& device) {
|
||||
const uint32_t fa_sgs = fa.first.subgroup_size;
|
||||
const bool fa_ds = fa.first.subgroup_size == 0;
|
||||
|
||||
const bool bf16_kv = fa.first.k_type == GGML_TYPE_BF16;
|
||||
|
||||
const void * spv_data;
|
||||
size_t spv_size;
|
||||
if (f32acc) { spv_data = flash_attn_f32_f16_cm1_data; spv_size = flash_attn_f32_f16_cm1_len; }
|
||||
else { spv_data = flash_attn_f32_f16_f16acc_cm1_data; spv_size = flash_attn_f32_f16_f16acc_cm1_len; }
|
||||
const char *name = aligned ? "flash_attn_f32_f16_aligned_cm1" : "flash_attn_f32_f16_cm1";
|
||||
const char *name;
|
||||
if (bf16_kv) {
|
||||
#if defined(VK_KHR_shader_bfloat16) && defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT)
|
||||
if (!device->coopmat_bf16_support) continue;
|
||||
spv_data = flash_attn_f32_f16_bf16_cm1_data;
|
||||
spv_size = flash_attn_f32_f16_bf16_cm1_len;
|
||||
name = aligned ? "flash_attn_f32_bf16_aligned_cm1" : "flash_attn_f32_bf16_cm1";
|
||||
#else
|
||||
continue;
|
||||
#endif
|
||||
} else {
|
||||
if (f32acc) { spv_data = flash_attn_f32_f16_cm1_data; spv_size = flash_attn_f32_f16_cm1_len; }
|
||||
else { spv_data = flash_attn_f32_f16_f16acc_cm1_data; spv_size = flash_attn_f32_f16_f16acc_cm1_len; }
|
||||
name = aligned ? "flash_attn_f32_f16_aligned_cm1" : "flash_attn_f32_f16_cm1";
|
||||
}
|
||||
ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7,
|
||||
sizeof(vk_flash_attn_push_constants), {Br, 1, 1},
|
||||
get_fa_spec_constants(fa.first), aligned ? Bc : 1, true,
|
||||
@@ -3917,10 +3946,20 @@ static void ggml_vk_load_shaders(vk_device& device) {
|
||||
const bool aligned = fa.first.aligned;
|
||||
const bool f32acc = fa.first.f32acc;
|
||||
|
||||
const bool bf16_kv = fa.first.k_type == GGML_TYPE_BF16;
|
||||
const void * spv_data;
|
||||
size_t spv_size;
|
||||
const char * name;
|
||||
if (aligned) {
|
||||
if (bf16_kv) {
|
||||
#if defined(VK_KHR_shader_bfloat16) && defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT)
|
||||
if (!device->coopmat2_bf16_support) continue;
|
||||
spv_data = flash_attn_f32_f16_bf16_cm2_data;
|
||||
spv_size = flash_attn_f32_f16_bf16_cm2_len;
|
||||
name = aligned ? "flash_attn_f32_bf16_aligned_cm2" : "flash_attn_f32_bf16_cm2";
|
||||
#else
|
||||
continue;
|
||||
#endif
|
||||
} else if (aligned) {
|
||||
if (f32acc) { spv_data = flash_attn_f32_f16_cm2_data; spv_size = flash_attn_f32_f16_cm2_len; name = "flash_attn_f32_f16_aligned_f32acc_cm2"; }
|
||||
else { spv_data = flash_attn_f32_f16_f16acc_cm2_data; spv_size = flash_attn_f32_f16_f16acc_cm2_len; name = "flash_attn_f32_f16_aligned_f16acc_cm2"; }
|
||||
} else {
|
||||
@@ -5804,46 +5843,72 @@ static vk_device ggml_vk_get_device(size_t idx) {
|
||||
found_fp16_256 = false,
|
||||
found_fp32_128 = false,
|
||||
found_fp32_256 = false;
|
||||
bool found_bf16_128 = false,
|
||||
found_bf16_256 = false;
|
||||
// need to support fp16*fp16 with fp16/fp32 accumulator, for workgroupsize 128
|
||||
// with 32x16x16 and 256 with 32x32x16.
|
||||
for (auto &prop : flexible_dimensions) {
|
||||
if (prop.saturatingAccumulation == VK_FALSE &&
|
||||
prop.scope == VK_SCOPE_WORKGROUP_KHR &&
|
||||
prop.AType == VK_COMPONENT_TYPE_FLOAT16_KHR &&
|
||||
prop.BType == VK_COMPONENT_TYPE_FLOAT16_KHR) {
|
||||
prop.scope == VK_SCOPE_WORKGROUP_KHR) {
|
||||
|
||||
if (prop.workgroupInvocations == 128 &&
|
||||
prop.MGranularity <= 32 &&
|
||||
prop.NGranularity <= 16 &&
|
||||
prop.KGranularity <= 16) {
|
||||
if (prop.CType == VK_COMPONENT_TYPE_FLOAT16_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT16_KHR) {
|
||||
found_fp16_128 = true;
|
||||
if (prop.AType == VK_COMPONENT_TYPE_FLOAT16_KHR &&
|
||||
prop.BType == VK_COMPONENT_TYPE_FLOAT16_KHR) {
|
||||
|
||||
if (prop.workgroupInvocations == 128 &&
|
||||
prop.MGranularity <= 32 &&
|
||||
prop.NGranularity <= 16 &&
|
||||
prop.KGranularity <= 16) {
|
||||
if (prop.CType == VK_COMPONENT_TYPE_FLOAT16_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT16_KHR) {
|
||||
found_fp16_128 = true;
|
||||
}
|
||||
if (prop.CType == VK_COMPONENT_TYPE_FLOAT32_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT32_KHR) {
|
||||
found_fp32_128 = true;
|
||||
}
|
||||
}
|
||||
if (prop.CType == VK_COMPONENT_TYPE_FLOAT32_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT32_KHR) {
|
||||
found_fp32_128 = true;
|
||||
if (prop.workgroupInvocations == 256 &&
|
||||
prop.MGranularity <= 32 &&
|
||||
prop.NGranularity <= 32 &&
|
||||
prop.KGranularity <= 16) {
|
||||
if (prop.CType == VK_COMPONENT_TYPE_FLOAT16_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT16_KHR) {
|
||||
found_fp16_256 = true;
|
||||
}
|
||||
if (prop.CType == VK_COMPONENT_TYPE_FLOAT32_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT32_KHR) {
|
||||
found_fp32_256 = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
if (prop.workgroupInvocations == 256 &&
|
||||
prop.MGranularity <= 32 &&
|
||||
prop.NGranularity <= 32 &&
|
||||
prop.KGranularity <= 16) {
|
||||
if (prop.CType == VK_COMPONENT_TYPE_FLOAT16_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT16_KHR) {
|
||||
found_fp16_256 = true;
|
||||
|
||||
#if defined(VK_KHR_shader_bfloat16) && defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT)
|
||||
if (prop.AType == VK_COMPONENT_TYPE_BFLOAT16_KHR &&
|
||||
prop.BType == VK_COMPONENT_TYPE_BFLOAT16_KHR &&
|
||||
prop.CType == VK_COMPONENT_TYPE_FLOAT32_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT32_KHR) {
|
||||
|
||||
if (prop.workgroupInvocations == 128 &&
|
||||
prop.MGranularity <= 32 &&
|
||||
prop.NGranularity <= 16 &&
|
||||
prop.KGranularity <= 16) {
|
||||
found_bf16_128 = true;
|
||||
}
|
||||
if (prop.CType == VK_COMPONENT_TYPE_FLOAT32_KHR &&
|
||||
prop.ResultType == VK_COMPONENT_TYPE_FLOAT32_KHR) {
|
||||
found_fp32_256 = true;
|
||||
if (prop.workgroupInvocations == 256 &&
|
||||
prop.MGranularity <= 32 &&
|
||||
prop.NGranularity <= 32 &&
|
||||
prop.KGranularity <= 16) {
|
||||
found_bf16_256 = true;
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
}
|
||||
if (found_fp16_128 && found_fp16_256 &&
|
||||
found_fp32_128 && found_fp32_256 &&
|
||||
coopmat2_props.cooperativeMatrixFlexibleDimensionsMaxDimension >= 512) {
|
||||
device->coopmat2 = true;
|
||||
device->coopmat2_bf16_support = found_bf16_128 && found_bf16_256;
|
||||
device->coopmat2_decode_vector = coopmat2_decode_vector_support && coopmat2_decode_vector_features.cooperativeMatrixDecodeVector;
|
||||
}
|
||||
}
|
||||
@@ -9476,7 +9541,8 @@ static bool ggml_vk_flash_attn_scalar_shmem_support(const vk_device& device, con
|
||||
const uint32_t Br = params.block_rows;
|
||||
const uint32_t Bc = params.block_cols;
|
||||
|
||||
const uint32_t float_type_size = device->fp16 ? sizeof(ggml_fp16_t) : sizeof(float);
|
||||
// BF16 uses the fp32 shader (FLOAT_TYPE=float)
|
||||
const uint32_t float_type_size = (device->fp16 && k_type != GGML_TYPE_BF16) ? sizeof(ggml_fp16_t) : sizeof(float);
|
||||
|
||||
const bool mmq = ggml_vk_fa_scalar_uses_mmq(device, k_type);
|
||||
|
||||
@@ -9517,7 +9583,7 @@ static bool ggml_vk_flash_attn_scalar_shmem_support(const vk_device& device, con
|
||||
return supported;
|
||||
}
|
||||
|
||||
static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc) {
|
||||
static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool f32acc, ggml_type k_type) {
|
||||
// Needs to be kept up to date on shader changes
|
||||
const uint32_t Br = params.block_rows;
|
||||
const uint32_t Bc = params.block_cols;
|
||||
@@ -9547,8 +9613,10 @@ static bool ggml_vk_flash_attn_coopmat_shmem_support(const vk_device& device, co
|
||||
const uint32_t vsh_stride = MatBc / 4 * row_split;
|
||||
const uint32_t ksh = ((kvshstride >= vsh_stride) ? (Bc * kvshstride) : (Bc * vsh_stride)) * f16vec4;
|
||||
|
||||
// BF16 PVMat accumulator is f32 (no bf16 accumulator support), so pvsh is vec4 (16 bytes)
|
||||
const uint32_t pvsh_elem_size = (k_type == GGML_TYPE_BF16) ? 16u : f16vec4;
|
||||
const uint32_t osh_stride = params.row_split * MatBr / 4;
|
||||
const uint32_t pvsh = MatBc * osh_stride * f16vec4;
|
||||
const uint32_t pvsh = MatBc * osh_stride * pvsh_elem_size;
|
||||
|
||||
const uint32_t slope = Br * acctype;
|
||||
|
||||
@@ -9617,7 +9685,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
|
||||
uint32_t workgroups_y = (uint32_t)neq2;
|
||||
uint32_t workgroups_z = (uint32_t)neq3;
|
||||
|
||||
const bool f32acc = !ctx->device->fp16 || dst->op_params[3] == GGML_PREC_F32;
|
||||
const bool f32acc = !ctx->device->fp16 || dst->op_params[3] == GGML_PREC_F32 || k->type == GGML_TYPE_BF16;
|
||||
|
||||
// For scalar/coopmat1 FA, we can use the "large" size to accommodate qga.
|
||||
// For coopmat2 FA, we always use the small size (which is still pretty large for gqa).
|
||||
@@ -16428,6 +16496,7 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
|
||||
switch (t) {
|
||||
case GGML_TYPE_F32:
|
||||
case GGML_TYPE_F16:
|
||||
case GGML_TYPE_BF16:
|
||||
case GGML_TYPE_Q8_0:
|
||||
case GGML_TYPE_Q5_1:
|
||||
case GGML_TYPE_Q5_0:
|
||||
@@ -16443,6 +16512,9 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
|
||||
if (!fa_kv_ok(op->src[1]->type) || !fa_kv_ok(op->src[2]->type)) {
|
||||
return false;
|
||||
}
|
||||
if ((op->src[1]->type == GGML_TYPE_BF16) != (op->src[2]->type == GGML_TYPE_BF16)) {
|
||||
return false;
|
||||
}
|
||||
if (!coopmat2 && !(device->subgroup_shuffle && device->subgroup_vote)) {
|
||||
// scalar/coopmat1 FA uses subgroupShuffle/subgroupAll
|
||||
return false;
|
||||
|
||||
@@ -97,8 +97,17 @@ layout (binding = 6) readonly buffer MO {uint32_t data_mask_opt[];};
|
||||
#define FA_TYPE_Q5_0 6u
|
||||
#define FA_TYPE_Q5_1 7u
|
||||
#define FA_TYPE_Q8_0 8u
|
||||
#define FA_TYPE_BF16 30u
|
||||
#define FA_TYPE_Q1_0 41u
|
||||
|
||||
#if defined(BFLOAT16)
|
||||
#define O_TYPE float
|
||||
#define O_TYPEV4 vec4
|
||||
#else
|
||||
#define O_TYPE FLOAT_TYPE
|
||||
#define O_TYPEV4 FLOAT_TYPEV4
|
||||
#endif
|
||||
|
||||
// Number of matrix elements per buffer block, derived from the K/V type spec
|
||||
// constant. F32 is treated as a vec4 "block" of 4 floats. F16 uses block size 1
|
||||
// and bypasses the dequant path entirely. Quants follow their ggml block sizes.
|
||||
@@ -111,6 +120,7 @@ uint fa_block_elems(uint ty) {
|
||||
case FA_TYPE_Q5_0: return uint(QUANT_K_Q5_0);
|
||||
case FA_TYPE_Q5_1: return uint(QUANT_K_Q5_1);
|
||||
case FA_TYPE_Q8_0: return uint(QUANT_K_Q8_0);
|
||||
case FA_TYPE_BF16: return 1u;
|
||||
case FA_TYPE_Q1_0: return uint(QUANT_K_Q1_0); // cm2-only, harmless elsewhere
|
||||
default: return 1u;
|
||||
}
|
||||
@@ -248,7 +258,7 @@ const float FATTN_KQ_MAX_OFFSET = 3.0f*0.6931f;
|
||||
|
||||
// Store the output when doing grouped query attention.
|
||||
// Rows index by Q's dimension 2, and the first N rows are valid.
|
||||
void gqaStore(const in uint32_t r, const in uint32_t c, const in FLOAT_TYPEV4 elems, const in uint32_t o_offset, const in uint32_t iq2, const in uint32_t N)
|
||||
void gqaStore(const in uint32_t r, const in uint32_t c, const in O_TYPEV4 elems, const in uint32_t o_offset, const in uint32_t iq2, const in uint32_t N)
|
||||
{
|
||||
uint32_t offset = (iq2 + r) * HSV / 4 + c;
|
||||
data_ov4[o_offset + offset] = D_TYPEV4(elems);
|
||||
|
||||
@@ -6,6 +6,10 @@
|
||||
#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
|
||||
#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require
|
||||
|
||||
#if defined(BFLOAT16)
|
||||
#extension GL_EXT_bfloat16 : enable
|
||||
#endif
|
||||
|
||||
#extension GL_KHR_shader_subgroup_basic : enable
|
||||
#extension GL_KHR_shader_subgroup_arithmetic : enable
|
||||
#extension GL_KHR_shader_subgroup_vote : enable
|
||||
@@ -14,7 +18,9 @@
|
||||
|
||||
#include "types.glsl"
|
||||
#include "flash_attn_base.glsl"
|
||||
#if !defined(BFLOAT16)
|
||||
#include "flash_attn_dequant.glsl"
|
||||
#endif
|
||||
|
||||
// These need to be supported N,M values for a MatBc x MatBr x 16 coopmatmuladd
|
||||
const uint32_t MatBr = 16;
|
||||
@@ -27,32 +33,32 @@ const uint32_t cols_per_thread = Bc / cols_per_iter;
|
||||
|
||||
layout (binding = 0) readonly buffer Q {float data_q[];};
|
||||
layout (binding = 0) readonly buffer QV4 {vec4 data_qv4[];};
|
||||
layout (binding = 1) readonly buffer K {float16_t data_k[];};
|
||||
layout (binding = 1) readonly buffer KV4 {f16vec4 data_kv4[];};
|
||||
layout (binding = 2) readonly buffer V {float16_t data_v[];};
|
||||
layout (binding = 2) readonly buffer VV4 {f16vec4 data_vv4[];};
|
||||
layout (binding = 1) readonly buffer K {FLOAT_TYPE data_k[];};
|
||||
layout (binding = 1) readonly buffer KV4 {FLOAT_TYPEV4 data_kv4[];};
|
||||
layout (binding = 2) readonly buffer V {FLOAT_TYPE data_v[];};
|
||||
layout (binding = 2) readonly buffer VV4 {FLOAT_TYPEV4 data_vv4[];};
|
||||
layout (binding = 3) readonly buffer M {float16_t data_m[];};
|
||||
|
||||
shared float tmpsh[row_split];
|
||||
|
||||
const uint32_t qstride = HSK_pad / 4 + 2; // in units of f16vec4
|
||||
shared f16vec4 Qf[Br * qstride];
|
||||
const uint32_t qstride = HSK_pad / 4 + 2;
|
||||
shared FLOAT_TYPEV4 Qf[Br * qstride];
|
||||
|
||||
const uint psh_stride = Br / 4 + 2;
|
||||
shared f16vec4 Psh[Bc * psh_stride];
|
||||
shared FLOAT_TYPEV4 Psh[Bc * psh_stride];
|
||||
|
||||
// Avoid padding for hsk==256 to make it fit in 48KB shmem.
|
||||
const uint32_t sfshstride = (HSK <= 128) ? (Br / 4 + 2) : Br / 4;
|
||||
shared ACC_TYPEV4 sfsh[Bc * sfshstride];
|
||||
|
||||
const uint32_t D_pad = HSK_pad > HSV_pad ? HSK_pad : HSV_pad;
|
||||
const uint32_t kvsh_stride = (SHMEM_STAGING != 0 ? D_pad : MatBr) / 4 + 2; // in units of f16vec4
|
||||
const uint32_t kvsh_stride = (SHMEM_STAGING != 0 ? D_pad : MatBr) / 4 + 2;
|
||||
const uint v_cols = MatBc / 4 * row_split; // total cols, 4 vec4s per MatBc * number of subgroups
|
||||
const uint vsh_stride = v_cols;
|
||||
shared f16vec4 kvsh[(kvsh_stride >= vsh_stride) ? (Bc * kvsh_stride) : (Bc * vsh_stride)];
|
||||
shared FLOAT_TYPEV4 kvsh[(kvsh_stride >= vsh_stride) ? (Bc * kvsh_stride) : (Bc * vsh_stride)];
|
||||
|
||||
const uint32_t osh_stride = row_split * MatBr / 4;
|
||||
shared f16vec4 pvsh[MatBc * osh_stride];
|
||||
shared O_TYPEV4 pvsh[MatBc * osh_stride];
|
||||
|
||||
shared ACC_TYPE slope[Br];
|
||||
|
||||
@@ -76,7 +82,7 @@ void main() {
|
||||
if ((HSK % 16) != 0) {
|
||||
[[unroll]] for (uint i = 0; i < Br * qstride; i += gl_WorkGroupSize.x) {
|
||||
if (i + tid < Br * qstride) {
|
||||
Qf[i + tid] = f16vec4(0);
|
||||
Qf[i + tid] = FLOAT_TYPEV4(0);
|
||||
}
|
||||
}
|
||||
barrier();
|
||||
@@ -89,15 +95,15 @@ void main() {
|
||||
uint32_t r = (idx + tid) / (HSK / 4);
|
||||
if (r < Br && d < HSK / 4 &&
|
||||
i * Br + r < N) {
|
||||
Qf[r * qstride + d] = f16vec4(data_qv4[q_offset / 4 + (i * Br + r) * q_stride / 4 + d] * p.scale);
|
||||
Qf[r * qstride + d] = FLOAT_TYPEV4(data_qv4[q_offset / 4 + (i * Br + r) * q_stride / 4 + d] * p.scale);
|
||||
}
|
||||
}
|
||||
barrier();
|
||||
|
||||
f16vec4 Of[rows_per_thread][d_per_thread];
|
||||
O_TYPEV4 Of[rows_per_thread][d_per_thread];
|
||||
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
|
||||
[[unroll]] for (uint32_t d = 0; d < d_per_thread; ++d) {
|
||||
Of[r][d] = f16vec4(0.0);
|
||||
Of[r][d] = O_TYPEV4(0.0);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -222,15 +228,18 @@ void main() {
|
||||
uint32_t d = (idx + tid) % (HSK_pad / 4);
|
||||
uint32_t c = (idx + tid) / (HSK_pad / 4);
|
||||
if (idx + gl_WorkGroupSize.x <= Bc * HSK_pad / 4 || c < Bc) {
|
||||
f16vec4 K_Tf = f16vec4(0);
|
||||
FLOAT_TYPEV4 K_Tf = FLOAT_TYPEV4(0);
|
||||
if ((!KV_bounds_check || j * Bc + c < KV) && (HSK == HSK_pad || d < HSK / 4)) {
|
||||
#if !defined(BFLOAT16)
|
||||
if (USE_DECODE_K) {
|
||||
uint coord = (j * Bc + c) * k_stride * BLOCK_SIZE_K + 4 * d;
|
||||
uint ib = coord / BLOCK_SIZE_K;
|
||||
uint iqs = (coord % BLOCK_SIZE_K);
|
||||
K_Tf = dequantize4(ib, iqs, k_offset, BINDING_IDX_K);
|
||||
} else {
|
||||
K_Tf = f16vec4(data_kv4[k_offset / 4 + (j * Bc + c) * k_stride / 4 + d]);
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + c) * k_stride / 4 + d]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -244,16 +253,16 @@ void main() {
|
||||
// Bc split across workgroup (four subgroups), loop over HSK in chunks of 16: 16 x 16 * 16 x 16 -> 16 x 16
|
||||
// This is written transposed in order to allow for N being 8 if implementations need it
|
||||
coopmat<ACC_TYPE, gl_ScopeSubgroup, MatBc, MatBr, gl_MatrixUseAccumulator> SfMat = coopmat<ACC_TYPE, gl_ScopeSubgroup, MatBc, MatBr, gl_MatrixUseAccumulator>(0);
|
||||
coopmat<float16_t, gl_ScopeSubgroup, MatBc, 16, gl_MatrixUseA> KMat;
|
||||
coopmat<float16_t, gl_ScopeSubgroup, 16, MatBr, gl_MatrixUseB> QMat;
|
||||
coopmat<FLOAT_TYPE, gl_ScopeSubgroup, MatBc, 16, gl_MatrixUseA> KMat;
|
||||
coopmat<FLOAT_TYPE, gl_ScopeSubgroup, 16, MatBr, gl_MatrixUseB> QMat;
|
||||
|
||||
[[unroll]] for (uint32_t d = 0; d < HSK_pad / 16; ++d) {
|
||||
// If SHMEM_STAGING is set, a Bc * HSK_pad size tile of K is loaded to shmem
|
||||
// If not, f16 K is loaded directly from global memory if aligned, otherwise
|
||||
// If not, K is loaded directly from global memory if aligned, otherwise
|
||||
// staged through a Bc * MatBr size staging buffer.
|
||||
// If K is not type f16, then it is always staged for dequantization.
|
||||
// If K is a quant type, then it is always staged for dequantization.
|
||||
if (SHMEM_STAGING == 0) {
|
||||
// For quants we always need to dequant into kvsh; for f16 we can load
|
||||
// For quants we always need to dequant into kvsh; for f16/bf16 we can load
|
||||
// directly from global memory when alignment / bounds allow it.
|
||||
const bool stage_k = USE_DECODE_K || KV_bounds_check || d * 16 + 16 > HSK;
|
||||
if (stage_k) {
|
||||
@@ -262,15 +271,18 @@ void main() {
|
||||
uint32_t col_vec = (idx + tid) % (MatBr / 4);
|
||||
uint32_t row = (idx + tid) / (MatBr / 4);
|
||||
if (idx + tid < Bc * MatBr / 4) {
|
||||
f16vec4 K_Tf = f16vec4(0);
|
||||
FLOAT_TYPEV4 K_Tf = FLOAT_TYPEV4(0);
|
||||
if ((!KV_bounds_check || j * Bc + row < KV) && (HSK == HSK_pad || d * 16 + col_vec * 4 < HSK)) {
|
||||
#if !defined(BFLOAT16)
|
||||
if (USE_DECODE_K) {
|
||||
uint coord = (j * Bc + row) * k_stride * BLOCK_SIZE_K + d * 16 + col_vec * 4;
|
||||
uint ib = coord / BLOCK_SIZE_K;
|
||||
uint iqs = (coord % BLOCK_SIZE_K);
|
||||
K_Tf = dequantize4(ib, iqs, k_offset, BINDING_IDX_K);
|
||||
} else {
|
||||
K_Tf = f16vec4(data_kv4[k_offset / 4 + (j * Bc + row) * k_stride / 4 + d * 16 / 4 + col_vec]);
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
K_Tf = FLOAT_TYPEV4(data_kv4[k_offset / 4 + (j * Bc + row) * k_stride / 4 + d * 16 / 4 + col_vec]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -357,7 +369,7 @@ void main() {
|
||||
[[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) {
|
||||
const uint d_local = d0 / threads_per_rowgroup;
|
||||
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
|
||||
Of[r][d_local] = float16_t(eMf[r]) * Of[r][d_local];
|
||||
Of[r][d_local] = O_TYPE(eMf[r]) * Of[r][d_local];
|
||||
}
|
||||
}
|
||||
|
||||
@@ -368,10 +380,10 @@ void main() {
|
||||
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; r += 4) {
|
||||
const uint row = tile_row(r);
|
||||
if (KV_bounds_check && j * Bc + col >= KV) {
|
||||
Psh[col * psh_stride + row / 4] = f16vec4(0.0f);
|
||||
Psh[col * psh_stride + row / 4] = FLOAT_TYPEV4(0.0f);
|
||||
} else {
|
||||
const vec4 mfvec = vec4(Mf[r], Mf[r + 1], Mf[r + 2], Mf[r + 3]);
|
||||
const f16vec4 Pf = f16vec4(exp(vec4(sfsh[row / 4 + col * sfshstride]) - mfvec));
|
||||
const FLOAT_TYPEV4 Pf = FLOAT_TYPEV4(exp(vec4(sfsh[row / 4 + col * sfshstride]) - mfvec));
|
||||
[[unroll]] for (uint32_t vec_idx = 0; vec_idx < 4; ++vec_idx) {
|
||||
Lf[r + vec_idx] += Pf[vec_idx];
|
||||
}
|
||||
@@ -385,15 +397,18 @@ void main() {
|
||||
uint32_t d = (idx + tid) % (HSV_pad / 4);
|
||||
uint32_t c = (idx + tid) / (HSV_pad / 4);
|
||||
if (idx + gl_WorkGroupSize.x <= Bc * HSV_pad / 4 || c < Bc) {
|
||||
f16vec4 V_Tf = f16vec4(0);
|
||||
FLOAT_TYPEV4 V_Tf = FLOAT_TYPEV4(0);
|
||||
if ((!KV_bounds_check || j * Bc + c < KV) && (HSV == HSV_pad || d < HSV / 4)) {
|
||||
#if !defined(BFLOAT16)
|
||||
if (USE_DECODE_V) {
|
||||
uint coord = (j * Bc + c) * v_stride * BLOCK_SIZE_V + 4 * d;
|
||||
uint ib = coord / BLOCK_SIZE_V;
|
||||
uint iqs = (coord % BLOCK_SIZE_V);
|
||||
V_Tf = dequantize4(ib, iqs, v_offset, BINDING_IDX_V);
|
||||
} else {
|
||||
V_Tf = f16vec4(data_vv4[v_offset / 4 + (j * Bc + c) * v_stride / 4 + d]);
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
V_Tf = FLOAT_TYPEV4(data_vv4[v_offset / 4 + (j * Bc + c) * v_stride / 4 + d]);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -409,7 +424,7 @@ void main() {
|
||||
[[unroll]] for (uint32_t hsv_tile = 0; hsv_tile < num_hsv_tiles; ++hsv_tile) {
|
||||
const uint hsv_offset = (hsv_tile * row_split + gl_SubgroupID) * 16;
|
||||
|
||||
coopmat<float16_t, gl_ScopeSubgroup, MatBc, MatBr, gl_MatrixUseAccumulator> PVMat = coopmat<float16_t, gl_ScopeSubgroup, MatBc, MatBr, gl_MatrixUseAccumulator>(0);
|
||||
coopmat<O_TYPE, gl_ScopeSubgroup, MatBc, MatBr, gl_MatrixUseAccumulator> PVMat = coopmat<O_TYPE, gl_ScopeSubgroup, MatBc, MatBr, gl_MatrixUseAccumulator>(0);
|
||||
|
||||
// Preload V tiles for [Bc, 16 * num subgroups]
|
||||
const uint v_rows = Bc;
|
||||
@@ -417,11 +432,11 @@ void main() {
|
||||
const uint v_loads_per_thread = v_total / gl_WorkGroupSize.x;
|
||||
|
||||
// If SHMEM_STAGING is set, a Bc * HSV_pad size tile of V is loaded to shmem.
|
||||
// If not, f16 V is loaded directly from global memory if aligned, otherwise
|
||||
// If not, V is loaded directly from global memory if aligned, otherwise
|
||||
// staged through a Bc * MatBr size staging buffer.
|
||||
// If V is not type f16, then it is always staged for dequantization.
|
||||
// If V is a quant type, then it is always staged for dequantization.
|
||||
if (SHMEM_STAGING == 0) {
|
||||
// For quants we always preload via kvsh. For f16 we only preload when
|
||||
// For quants we always preload via kvsh. For f16/bf16 we only preload when
|
||||
// alignment / bounds force it (otherwise we coopMatLoad direct from data_vv4).
|
||||
const bool stage_v = USE_DECODE_V || KV_bounds_check;
|
||||
if (stage_v) {
|
||||
@@ -438,13 +453,16 @@ void main() {
|
||||
const uint iqs = coord % BLOCK_SIZE_V;
|
||||
|
||||
if (!KV_bounds_check || (v_row < KV && v_col < HSV)) {
|
||||
#if !defined(BFLOAT16)
|
||||
if (USE_DECODE_V) {
|
||||
kvsh[row * vsh_stride + col] = dequantize4(ib, iqs, v_offset, BINDING_IDX_V);
|
||||
} else {
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
kvsh[row * vsh_stride + col] = data_vv4[(v_offset + v_row * v_stride + v_col) / 4];
|
||||
}
|
||||
} else {
|
||||
kvsh[row * vsh_stride + col] = f16vec4(0.0f);
|
||||
kvsh[row * vsh_stride + col] = FLOAT_TYPEV4(0.0f);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -459,7 +477,7 @@ void main() {
|
||||
|
||||
if (SHMEM_STAGING == 0) {
|
||||
if (!USE_DECODE_V && !KV_bounds_check) {
|
||||
// F16 values can be loaded directly from global memory
|
||||
// F16/BF16 values can be loaded directly from global memory
|
||||
const uint v_tile_row = j * Bc + bc_chunk * MatBc;
|
||||
const uint v_tile_offset = v_offset / 4 + v_tile_row * v_stride / 4 + hsv_offset / 4;
|
||||
coopMatLoad(QMat, data_vv4, v_tile_offset, v_stride / 4, gl_CooperativeMatrixLayoutRowMajor);
|
||||
@@ -573,7 +591,7 @@ void main() {
|
||||
|
||||
[[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) {
|
||||
const uint d_local = d0 / threads_per_rowgroup;
|
||||
Of[r][d_local] *= float16_t(ms);
|
||||
Of[r][d_local] *= O_TYPE(ms);
|
||||
}
|
||||
} else {
|
||||
vs = exp(sink - Mf[r]);
|
||||
@@ -591,7 +609,7 @@ void main() {
|
||||
[[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) {
|
||||
const uint d_local = d0 / threads_per_rowgroup;
|
||||
[[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
|
||||
Of[r][d_local] *= float16_t(Lfrcp[r]);
|
||||
Of[r][d_local] *= O_TYPE(Lfrcp[r]);
|
||||
#if defined(FLOAT_TYPE_MAX)
|
||||
Of[r][d_local] = clamp(Of[r][d_local], -FLOAT_TYPE_MAX, FLOAT_TYPE_MAX);
|
||||
#endif
|
||||
|
||||
@@ -8,6 +8,10 @@
|
||||
#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require
|
||||
#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require
|
||||
|
||||
#if defined(BFLOAT16)
|
||||
#extension GL_EXT_bfloat16 : enable
|
||||
#endif
|
||||
|
||||
#extension GL_KHR_memory_scope_semantics : enable
|
||||
#extension GL_KHR_cooperative_matrix : enable
|
||||
#extension GL_NV_cooperative_matrix2 : enable
|
||||
@@ -21,7 +25,9 @@
|
||||
|
||||
#include "types.glsl"
|
||||
#include "flash_attn_base.glsl"
|
||||
#if !defined(BFLOAT16)
|
||||
#include "dequant_funcs_cm2.glsl"
|
||||
#endif
|
||||
|
||||
// buffer_reference stride = sizeof(struct) = FaBlockBytesK/V.
|
||||
layout(buffer_reference, std430, buffer_reference_align = 1) buffer decodeBufFA_K {
|
||||
@@ -31,6 +37,7 @@ layout(buffer_reference, std430, buffer_reference_align = 1) buffer decodeBufFA_
|
||||
uint8_t raw[FaBlockBytesV];
|
||||
};
|
||||
|
||||
#if !defined(BFLOAT16)
|
||||
float16_t faDecodeK(const decodeBufFA_K bl_in, const uint blockCoords[2], const uint coordInBlock[2]) {
|
||||
switch (FaTypeK) {
|
||||
case FA_TYPE_F32: return dequantFuncF32 (decodeBufF32 (bl_in), blockCoords, coordInBlock);
|
||||
@@ -91,6 +98,7 @@ f16vec4 faDecodeVVector(const decodeBufFA_V bl_in, const uint blockCoords[2], co
|
||||
#define FADECODEK , faDecodeK
|
||||
#define FADECODEV , faDecodeV
|
||||
#endif
|
||||
#endif
|
||||
|
||||
layout (binding = 0) readonly buffer Q {uint8_t data_q[];};
|
||||
layout (binding = 1) readonly buffer K {uint8_t data_k[];};
|
||||
@@ -195,15 +203,15 @@ void main() {
|
||||
tensorLayoutV = setTensorLayoutStrideNV(tensorLayoutV, v_stride, 1);
|
||||
|
||||
coopmat<Q_TYPE, gl_ScopeWorkgroup, Br, HSK_pad, gl_MatrixUseAccumulator> Q;
|
||||
coopmat<float16_t, gl_ScopeWorkgroup, Br, HSK_pad, gl_MatrixUseA> Qf16;
|
||||
coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, Br, HSK_pad, gl_MatrixUseA> Qf16;
|
||||
|
||||
uint32_t q_offset = gqa_iq1*p.nb01*4/*sizeof(float)*/ + iq2*p.nb02+iq3*p.nb03;
|
||||
coopMatLoadTensorNV(Q, data_q, q_offset, sliceTensorLayoutNV(tensorLayoutQ, i * Br, Br, 0, HSK_pad));
|
||||
|
||||
Qf16 = coopmat<float16_t, gl_ScopeWorkgroup, Br, HSK_pad, gl_MatrixUseA>(Q);
|
||||
Qf16 *= float16_t(p.scale);
|
||||
Q *= Q_TYPE(p.scale);
|
||||
Qf16 = coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, Br, HSK_pad, gl_MatrixUseA>(Q);
|
||||
|
||||
coopmat<float16_t, gl_ScopeWorkgroup, Br, HSV_pad, gl_MatrixUseAccumulator> O = coopmat<float16_t, gl_ScopeWorkgroup, Br, HSV_pad, gl_MatrixUseAccumulator>(0);
|
||||
coopmat<O_TYPE, gl_ScopeWorkgroup, Br, HSV_pad, gl_MatrixUseAccumulator> O = coopmat<O_TYPE, gl_ScopeWorkgroup, Br, HSV_pad, gl_MatrixUseAccumulator>(0);
|
||||
|
||||
coopmat<ACC_TYPE, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseAccumulator> L, M;
|
||||
|
||||
@@ -291,16 +299,20 @@ void main() {
|
||||
|
||||
coopmat<ACC_TYPE, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseAccumulator> S = coopmat<ACC_TYPE, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseAccumulator>(0);
|
||||
|
||||
coopmat<float16_t, gl_ScopeWorkgroup, HSK_pad, Bc, gl_MatrixUseB> K_T;
|
||||
coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, HSK_pad, Bc, gl_MatrixUseB> K_T;
|
||||
|
||||
uint32_t k_offset = ik2*p.nb12 + ik3*p.nb13;
|
||||
// F16: bs_k==1 (direct load). F32: bs_k==4 (vec4 / dequantFuncF32). Q4/Q8 family: bs_k==32. Q1_0: bs_k==128.
|
||||
#if defined(BFLOAT16)
|
||||
coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose);
|
||||
#else
|
||||
const bool k_use_decode = (bs_k > 1u);
|
||||
if (k_use_decode) {
|
||||
coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose FADECODEK);
|
||||
} else {
|
||||
coopMatLoadTensorNV(K_T, data_k, k_offset, sliceTensorLayoutNV(tensorLayoutK, j * Bc, Bc, 0, HSK_pad), tensorViewTranspose);
|
||||
}
|
||||
#endif
|
||||
S = coopMatMulAdd(Qf16, K_T, S);
|
||||
|
||||
if (LOGIT_SOFTCAP) {
|
||||
@@ -351,22 +363,26 @@ void main() {
|
||||
coopMatPerElementNV(P, P, replacePadding, ACC_TYPE(0.0), R, C);
|
||||
}
|
||||
|
||||
coopmat<float16_t, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseA> P_A = coopmat<float16_t, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseA>(P);
|
||||
coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseA> P_A = coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseA>(P);
|
||||
|
||||
// compute rowsum by multiplying by matrix of all ones.
|
||||
coopmat<float16_t, gl_ScopeWorkgroup, Bc, Bc, gl_MatrixUseB> One = coopmat<float16_t, gl_ScopeWorkgroup, Bc, Bc, gl_MatrixUseB>(1.0);
|
||||
coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, Bc, Bc, gl_MatrixUseB> One = coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, Bc, Bc, gl_MatrixUseB>(1.0);
|
||||
|
||||
rowsum = coopmat<ACC_TYPE, gl_ScopeWorkgroup, Br, Bc, gl_MatrixUseAccumulator>(0.0);
|
||||
rowsum = coopMatMulAdd(P_A, One, rowsum);
|
||||
|
||||
coopmat<float16_t, gl_ScopeWorkgroup, Bc, HSV_pad, gl_MatrixUseB> V;
|
||||
coopmat<FLOAT_TYPE, gl_ScopeWorkgroup, Bc, HSV_pad, gl_MatrixUseB> V;
|
||||
uint32_t v_offset = iv2*p.nb22 + iv3*p.nb23;
|
||||
#if defined(BFLOAT16)
|
||||
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad));
|
||||
#else
|
||||
const bool v_use_decode = (bs_v > 1u);
|
||||
if (v_use_decode) {
|
||||
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad) FADECODEV);
|
||||
} else {
|
||||
coopMatLoadTensorNV(V, data_v, v_offset, sliceTensorLayoutNV(tensorLayoutV, j * Bc, Bc, 0, HSV_pad));
|
||||
}
|
||||
#endif
|
||||
|
||||
L = eM*L + rowsum;
|
||||
|
||||
@@ -378,7 +394,7 @@ void main() {
|
||||
// resize eM by using smear/reduce
|
||||
coopMatReduceNV(eMdiag, eM, gl_CooperativeMatrixReduceRowNV, smearReduce);
|
||||
|
||||
O *= coopmat<float16_t, gl_ScopeWorkgroup, Br, HSV_pad, gl_MatrixUseAccumulator>(eMdiag);
|
||||
O *= coopmat<O_TYPE, gl_ScopeWorkgroup, Br, HSV_pad, gl_MatrixUseAccumulator>(eMdiag);
|
||||
O = coopMatMulAdd(P_A, V, O);
|
||||
}
|
||||
|
||||
@@ -427,7 +443,7 @@ void main() {
|
||||
if (sink > Mr[i]) {
|
||||
ms = exp(Mr[i] - sink);
|
||||
|
||||
O[i] *= float16_t(ms);
|
||||
O[i] *= O_TYPE(ms);
|
||||
} else {
|
||||
vs = exp(sink - Mr[i]);
|
||||
}
|
||||
|
||||
@@ -28,6 +28,9 @@ layout (binding = 2) readonly buffer V_PACKED_Q5_1 { block_q5_1_packed16 data[];
|
||||
layout (binding = 1) readonly buffer K_PACKED_Q8_0 { block_q8_0_packed16 data[]; } k_packed_q8_0;
|
||||
layout (binding = 2) readonly buffer V_PACKED_Q8_0 { block_q8_0_packed16 data[]; } v_packed_q8_0;
|
||||
|
||||
layout (binding = 1) readonly buffer K_PACKED_BF16 { u16vec4 data[]; } k_packed_bf16;
|
||||
layout (binding = 2) readonly buffer V_PACKED_BF16 { u16vec4 data[]; } v_packed_bf16;
|
||||
|
||||
// Q4_1 and Q5_1 packed32 views: aliased to the same memory as the packed16
|
||||
// views, used by the MMQ K-side hot path for fast 4-uint loads.
|
||||
layout (binding = 1) readonly buffer K_PACKED_Q4_1_P32 { block_q4_1_packed32 data[]; } k_packed_q4_1_p32;
|
||||
@@ -99,6 +102,9 @@ layout (binding = 1) readonly buffer K_PACKED_Q5_1_P32 { block_q5_1_packed32 dat
|
||||
return FLOAT_TYPE(BUF.data[a_offset + ib].d) * FLOAT_TYPEV4(v0.x, v0.y, v1.x, v1.y); \
|
||||
}
|
||||
|
||||
#define FA_DEQUANT4_BF16(BUF) \
|
||||
return FLOAT_TYPEV4(bf16_to_fp32(uvec4(BUF.data[(a_offset + ib) / 4])));
|
||||
|
||||
FLOAT_TYPEV4 dequantize4(uint ib, uint iqs, uint a_offset, uint binding_idx) {
|
||||
if (binding_idx == BINDING_IDX_K) {
|
||||
switch (FaTypeK) {
|
||||
@@ -108,6 +114,7 @@ FLOAT_TYPEV4 dequantize4(uint ib, uint iqs, uint a_offset, uint binding_idx) {
|
||||
case FA_TYPE_Q5_0: FA_DEQUANT4_Q5_0(k_packed_q5_0)
|
||||
case FA_TYPE_Q5_1: FA_DEQUANT4_Q5_1(k_packed_q5_1)
|
||||
case FA_TYPE_Q8_0: FA_DEQUANT4_Q8_0(k_packed_q8_0)
|
||||
case FA_TYPE_BF16: FA_DEQUANT4_BF16(k_packed_bf16)
|
||||
}
|
||||
} else {
|
||||
switch (FaTypeV) {
|
||||
@@ -117,6 +124,7 @@ FLOAT_TYPEV4 dequantize4(uint ib, uint iqs, uint a_offset, uint binding_idx) {
|
||||
case FA_TYPE_Q5_0: FA_DEQUANT4_Q5_0(v_packed_q5_0)
|
||||
case FA_TYPE_Q5_1: FA_DEQUANT4_Q5_1(v_packed_q5_1)
|
||||
case FA_TYPE_Q8_0: FA_DEQUANT4_Q8_0(v_packed_q8_0)
|
||||
case FA_TYPE_BF16: FA_DEQUANT4_BF16(v_packed_bf16)
|
||||
}
|
||||
}
|
||||
return FLOAT_TYPEV4(0);
|
||||
|
||||
@@ -679,6 +679,28 @@ void process_shaders() {
|
||||
}
|
||||
}
|
||||
|
||||
const std::map<std::string, std::string> fa_bf16_dict = {
|
||||
{"FLOAT_TYPE", "bfloat16_t"},
|
||||
{"FLOAT_TYPEV2", "bf16vec2"},
|
||||
{"FLOAT_TYPEV4", "bf16vec4"},
|
||||
{"ACC_TYPE", "float"},
|
||||
{"ACC_TYPEV2", "vec2"},
|
||||
{"ACC_TYPEV4", "vec4"},
|
||||
{"BFLOAT16", "1"},
|
||||
};
|
||||
|
||||
#if defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT) && defined(GGML_VULKAN_COOPMAT_GLSLC_SUPPORT)
|
||||
string_to_spv("flash_attn_f32_f16_bf16", "flash_attn_cm1.comp",
|
||||
merge_maps(fa_bf16_dict, {{"Q_TYPE", "float"}, {"D_TYPE", "float"}, {"D_TYPEV4", "vec4"}, {"COOPMAT", "1"}}),
|
||||
true, true, false, false);
|
||||
#endif
|
||||
|
||||
#if defined(GGML_VULKAN_BFLOAT16_GLSLC_SUPPORT) && defined(GGML_VULKAN_COOPMAT2_GLSLC_SUPPORT)
|
||||
string_to_spv("flash_attn_f32_f16_bf16", "flash_attn_cm2.comp",
|
||||
merge_maps(fa_bf16_dict, {{"Q_TYPE", "float"}, {"D_TYPE", "float"}, {"D_TYPEV4", "vec4"}}),
|
||||
true, false, true, false);
|
||||
#endif
|
||||
|
||||
std::map<std::string, std::string> base_dict = {{"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}};
|
||||
|
||||
for (const auto& tname : type_names) {
|
||||
|
||||
@@ -0,0 +1,224 @@
|
||||
#ifdef DST_Q8_0
|
||||
#define BLOCK_SIZE 32u
|
||||
#define BLOCK_BYTES 34u
|
||||
#define QS_WORDS 8u
|
||||
#elif defined(DST_Q4_0)
|
||||
#define BLOCK_SIZE 32u
|
||||
#define BLOCK_BYTES 18u
|
||||
#define QS_WORDS 4u
|
||||
#endif
|
||||
|
||||
@group(0) @binding(0)
|
||||
var<storage, read_write> src: array<f32>;
|
||||
|
||||
@group(0) @binding(1)
|
||||
var<storage, read_write> idx: array<u32>;
|
||||
|
||||
@group(0) @binding(2)
|
||||
#ifdef PAIR_BLOCKS
|
||||
var<storage, read_write> dst: array<u32>;
|
||||
#else
|
||||
var<storage, read_write> dst: array<atomic<u32>>;
|
||||
#endif
|
||||
|
||||
#ifdef I64_IDX
|
||||
@group(0) @binding(3)
|
||||
var<storage, read_write> error: atomic<u32>;
|
||||
#define PARAMS_BINDING 4
|
||||
#else
|
||||
#define PARAMS_BINDING 3
|
||||
#endif
|
||||
|
||||
struct Params {
|
||||
offset_src: u32, // in elements
|
||||
offset_idx: u32, // in elements
|
||||
offset_dst: u32, // in blocks
|
||||
|
||||
// Strides (in elements / blocks)
|
||||
stride_src1: u32,
|
||||
stride_src2: u32,
|
||||
stride_src3: u32,
|
||||
|
||||
stride_idx0: u32,
|
||||
stride_idx1: u32,
|
||||
stride_idx2: u32,
|
||||
|
||||
stride_dst1: u32,
|
||||
stride_dst2: u32,
|
||||
stride_dst3: u32,
|
||||
|
||||
// Shape of src
|
||||
ne0: u32,
|
||||
n_rows: u32,
|
||||
ne2: u32,
|
||||
ne3: u32,
|
||||
|
||||
// Shape of idx
|
||||
idx1: u32,
|
||||
idx2: u32,
|
||||
};
|
||||
|
||||
@group(0) @binding(PARAMS_BINDING)
|
||||
var<uniform> params: Params;
|
||||
|
||||
// if the quantization type is unaligned and there are an odd number of blocks per row, we need to store atomically
|
||||
#ifndef PAIR_BLOCKS
|
||||
fn merge_store_dst_word(word_idx: u32, mask: u32, bits: u32) {
|
||||
loop {
|
||||
let old = atomicLoad(&dst[word_idx]);
|
||||
let merged = (old & ~mask) | (bits & mask);
|
||||
let result = atomicCompareExchangeWeak(&dst[word_idx], old, merged);
|
||||
if (result.exchanged) {
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
#else
|
||||
fn merge_store_dst_word(word_idx: u32, mask: u32, bits: u32) {
|
||||
let old = dst[word_idx];
|
||||
dst[word_idx] = (old & ~mask) | (bits & mask);
|
||||
}
|
||||
#endif
|
||||
|
||||
fn store_u16(dst_word_idx: u32, block_byte_offset: u32, byte_offset: u32, value: u32) {
|
||||
let total_byte_offset = block_byte_offset + byte_offset;
|
||||
let word_idx = dst_word_idx + total_byte_offset / 4u;
|
||||
let shift = (total_byte_offset & 2u) * 8u;
|
||||
let mask = 0xFFFFu << shift;
|
||||
merge_store_dst_word(word_idx, mask, (value & 0xFFFFu) << shift);
|
||||
}
|
||||
|
||||
fn store_u32(dst_word_idx: u32, block_byte_offset: u32, byte_offset: u32, value: u32) {
|
||||
let total_byte_offset = block_byte_offset + byte_offset;
|
||||
let word_idx = dst_word_idx + total_byte_offset / 4u;
|
||||
let shift = (total_byte_offset & 3u) * 8u;
|
||||
|
||||
if (shift == 0u) {
|
||||
#ifdef PAIR_BLOCKS
|
||||
dst[word_idx] = value;
|
||||
#else
|
||||
atomicStore(&dst[word_idx], value);
|
||||
#endif
|
||||
return;
|
||||
}
|
||||
|
||||
let lo_mask = 0xFFFFFFFFu << shift;
|
||||
let hi_mask = (1u << shift) - 1u;
|
||||
merge_store_dst_word(word_idx, lo_mask, value << shift);
|
||||
merge_store_dst_word(word_idx + 1u, hi_mask, value >> (32u - shift));
|
||||
}
|
||||
|
||||
fn quantize_block_params(src_block: u32) -> vec2<f32> {
|
||||
#ifdef DST_Q8_0
|
||||
var amax = 0.0;
|
||||
for (var j: u32 = 0u; j < BLOCK_SIZE; j++) {
|
||||
amax = max(amax, abs(src[src_block + j]));
|
||||
}
|
||||
|
||||
let d = amax / 127.0;
|
||||
let id = select(0.0, 1.0 / d, d > 0.0);
|
||||
return vec2(d, id);
|
||||
#elif defined(DST_Q4_0)
|
||||
var amax = 0.0;
|
||||
var max_val = 0.0;
|
||||
for (var j: u32 = 0u; j < BLOCK_SIZE; j++) {
|
||||
let v = src[src_block + j];
|
||||
let av = abs(v);
|
||||
if (amax < av) {
|
||||
amax = av;
|
||||
max_val = v;
|
||||
}
|
||||
}
|
||||
|
||||
let d = max_val / -8.0;
|
||||
let id = select(0.0, 1.0 / d, d != 0.0);
|
||||
return vec2(d, id);
|
||||
#endif
|
||||
}
|
||||
|
||||
fn quantize_block_word(src_block: u32, j: u32, id: f32) -> u32 {
|
||||
#ifdef DST_Q8_0
|
||||
let base = src_block + j * 4u;
|
||||
return (u32(i32(round(src[base + 0u] * id)) & 0xFF) << 0u) |
|
||||
(u32(i32(round(src[base + 1u] * id)) & 0xFF) << 8u) |
|
||||
(u32(i32(round(src[base + 2u] * id)) & 0xFF) << 16u) |
|
||||
(u32(i32(round(src[base + 3u] * id)) & 0xFF) << 24u);
|
||||
#elif defined(DST_Q4_0)
|
||||
var packed_q = 0u;
|
||||
for (var k: u32 = 0u; k < 4u; k++) {
|
||||
let x0 = src[src_block + j * 4u + k] * id;
|
||||
let x1 = src[src_block + 16u + j * 4u + k] * id;
|
||||
let q0 = u32(clamp(i32(x0 + 8.5), 0, 15));
|
||||
let q1 = u32(clamp(i32(x1 + 8.5), 0, 15));
|
||||
packed_q |= (q0 & 0xFu) << (8u * k);
|
||||
packed_q |= (q1 & 0xFu) << (8u * k + 4u);
|
||||
}
|
||||
return packed_q;
|
||||
#endif
|
||||
}
|
||||
|
||||
fn quantize_block(src_block: u32, dst_word_idx: u32, block_byte_offset: u32) {
|
||||
let params = quantize_block_params(src_block);
|
||||
let d = params.x;
|
||||
let id = params.y;
|
||||
let packed_d = pack2x16float(vec2(d, 0.0)) & 0xFFFFu;
|
||||
store_u16(dst_word_idx, block_byte_offset, 0u, packed_d);
|
||||
|
||||
for (var j: u32 = 0u; j < QS_WORDS; j++) {
|
||||
store_u32(dst_word_idx, block_byte_offset, 2u + j * 4u, quantize_block_word(src_block, j, id));
|
||||
}
|
||||
}
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
let blocks_per_row = params.ne0 / BLOCK_SIZE;
|
||||
#ifdef PAIR_BLOCKS
|
||||
let blocks_per_invocation = 2u;
|
||||
#else
|
||||
let blocks_per_invocation = 1u;
|
||||
#endif
|
||||
let invocations_per_row = blocks_per_row / blocks_per_invocation;
|
||||
let total_invocations = params.ne3 * params.ne2 * params.n_rows * invocations_per_row;
|
||||
if (gid.x >= total_invocations) {
|
||||
return;
|
||||
}
|
||||
|
||||
var i = gid.x / invocations_per_row;
|
||||
let block_in_row = (gid.x % invocations_per_row) * blocks_per_invocation;
|
||||
|
||||
let i_src3 = i / (params.ne2 * params.n_rows);
|
||||
i = i % (params.ne2 * params.n_rows);
|
||||
let i_src2 = i / params.n_rows;
|
||||
let i_src1 = i % params.n_rows;
|
||||
|
||||
let i_idx2 = i_src3 % params.idx2;
|
||||
let i_idx1 = i_src2 % params.idx1;
|
||||
let i_idx0 = i_src1;
|
||||
|
||||
#ifdef I64_IDX
|
||||
let idx_high = (params.offset_idx + i_idx0 * params.stride_idx0 + i_idx1 * params.stride_idx1 + i_idx2 * params.stride_idx2) * 2u;
|
||||
let idx_val = idx[idx_high];
|
||||
let idx_low_val = idx[idx_high + 1u];
|
||||
|
||||
if (idx_low_val != 0u) {
|
||||
atomicStore(&error, 1u);
|
||||
return;
|
||||
}
|
||||
#else
|
||||
let idx_i = params.offset_idx + i_idx0 * params.stride_idx0 + i_idx1 * params.stride_idx1 + i_idx2 * params.stride_idx2;
|
||||
let idx_val = idx[idx_i];
|
||||
#endif
|
||||
|
||||
let dst_row_blocks = params.offset_dst + idx_val * params.stride_dst1 + i_src2 * params.stride_dst2 + i_src3 * params.stride_dst3;
|
||||
let src_row = params.offset_src + i_src1 * params.stride_src1 + i_src2 * params.stride_src2 + i_src3 * params.stride_src3;
|
||||
let src_block = src_row + block_in_row * BLOCK_SIZE;
|
||||
let dst_block_byte = (dst_row_blocks + block_in_row) * BLOCK_BYTES;
|
||||
|
||||
let dst_word_idx = dst_block_byte / 4u;
|
||||
#ifdef PAIR_BLOCKS
|
||||
quantize_block(src_block, dst_word_idx, 0u);
|
||||
quantize_block(src_block + BLOCK_SIZE, dst_word_idx, BLOCK_BYTES);
|
||||
#else
|
||||
quantize_block(src_block, dst_word_idx, dst_block_byte & 3u);
|
||||
#endif
|
||||
}
|
||||
+7
-3
@@ -2656,14 +2656,18 @@ llm_graph_input_attn_k_dsa * llm_graph_context::build_attn_inp_k_dsa() const {
|
||||
inp->self_k_idxs_mla = mctx_cur->get_mla()->build_input_k_idxs(ctx0, ubatch);
|
||||
|
||||
inp->self_kq_mask_mla = build_attn_inp_kq_mask(ctx0, mctx_cur->get_mla(), ubatch, cparams);
|
||||
inp->self_kq_mask_mla_cnv = cparams.flash_attn ? ggml_cast(ctx0, inp->self_kq_mask_mla, GGML_TYPE_F16) : inp->self_kq_mask_mla;
|
||||
inp->self_kq_mask_mla_cnv = inp->self_kq_mask_mla;
|
||||
}
|
||||
|
||||
{
|
||||
inp->self_k_idxs_lid = mctx_cur->get_lid()->build_input_k_idxs(ctx0, ubatch);
|
||||
|
||||
inp->self_kq_mask_lid = build_attn_inp_kq_mask(ctx0, mctx_cur->get_lid(), ubatch, cparams);
|
||||
inp->self_kq_mask_lid_cnv = cparams.flash_attn ? ggml_cast(ctx0, inp->self_kq_mask_lid, GGML_TYPE_F16) : inp->self_kq_mask_lid;
|
||||
// ensure F32 mask
|
||||
auto cparams_copy = cparams;
|
||||
cparams_copy.flash_attn = false;
|
||||
|
||||
inp->self_kq_mask_lid = build_attn_inp_kq_mask(ctx0, mctx_cur->get_lid(), ubatch, cparams_copy);
|
||||
inp->self_kq_mask_lid_cnv = inp->self_kq_mask_lid;
|
||||
|
||||
inp->self_k_rot_lid = mctx_cur->get_lid()->build_input_k_rot(ctx0);
|
||||
}
|
||||
|
||||
+4
-4
@@ -399,10 +399,10 @@ public:
|
||||
ggml_tensor * self_k_idxs_mla = nullptr; // I64 [n_batch]
|
||||
ggml_tensor * self_k_idxs_lid = nullptr; // I64 [n_batch]
|
||||
|
||||
ggml_tensor * self_kq_mask_mla = nullptr; // F32 [n_kv, n_batch/n_stream, 1, n_stream]
|
||||
ggml_tensor * self_kq_mask_mla_cnv = nullptr; // [n_kv, n_batch/n_stream, 1, n_stream]
|
||||
ggml_tensor * self_kq_mask_lid = nullptr; // F32 [n_kv, n_batch/n_stream, 1, n_stream]
|
||||
ggml_tensor * self_kq_mask_lid_cnv = nullptr; // [n_kv, n_batch/n_stream, 1, n_stream]
|
||||
ggml_tensor * self_kq_mask_mla = nullptr; // F32/F16 [n_kv, n_batch/n_stream, 1, n_stream]
|
||||
ggml_tensor * self_kq_mask_mla_cnv = nullptr; // [n_kv, n_batch/n_stream, 1, n_stream]
|
||||
ggml_tensor * self_kq_mask_lid = nullptr; // F32 [n_kv, n_batch/n_stream, 1, n_stream]
|
||||
ggml_tensor * self_kq_mask_lid_cnv = nullptr; // [n_kv, n_batch/n_stream, 1, n_stream]
|
||||
|
||||
ggml_tensor * self_k_rot_lid = nullptr;
|
||||
|
||||
|
||||
+5
-5
@@ -542,16 +542,16 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
auto get_tensor_config = [&]() -> tensor_config {
|
||||
// standard attention
|
||||
if (std::regex_match(tensor_name, pattern_q_weight) || std::regex_match(tensor_name, pattern_kv_weight)) {
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1, "attn_output.weight");
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1, "attn_output.weight", "ssm_out.weight");
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_q_bias) || std::regex_match(tensor_name, pattern_kv_bias)) {
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_0, "attn_output.weight");
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_0, "attn_output.weight", "ssm_out.weight");
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_qkv_weight)) {
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1);
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1, "attn_output.weight", "ssm_out.weight");
|
||||
}
|
||||
if ( std::regex_match(tensor_name, pattern_qkv_bias)) {
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_0);
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_0, "attn_output.weight", "ssm_out.weight");
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_qk_norm)) {
|
||||
return get_tensor_config_impl(tensor->ne[1] == 1 ? GGML_BACKEND_SPLIT_AXIS_MIRRORED : GGML_BACKEND_SPLIT_AXIS_1, "attn_output.weight");
|
||||
@@ -567,7 +567,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str
|
||||
}
|
||||
|
||||
if (std::regex_match(tensor_name, pattern_attn_gate_weight)) {
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1);
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1, "attn_output.weight", "ssm_out.weight");
|
||||
}
|
||||
if (std::regex_match(tensor_name, pattern_ssm_dt) || std::regex_match(tensor_name, pattern_ssm_a)) {
|
||||
return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_0, "ssm_out.weight");
|
||||
|
||||
+3
-2
@@ -263,8 +263,9 @@ static bool llama_prepare_model_devices(const llama_model_params & params, llama
|
||||
// add GPUs
|
||||
model->devices.insert(model->devices.end(), gpus.begin(), gpus.end());
|
||||
|
||||
// add integrated GPUs only if no other devices were found
|
||||
if (model->devices.empty()) {
|
||||
// add integrated GPUs only if no discrete GPUs were found
|
||||
// (RPC servers do not count, otherwise the local iGPU would be dropped on iGPU+RPC setups)
|
||||
if (gpus.empty()) {
|
||||
model->devices.insert(model->devices.end(), igpus.begin(), igpus.end());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
# SPEED-Bench server benchmark
|
||||
|
||||
A lightweight [SPEED-Bench](https://huggingface.co/datasets/nvidia/SPEED-Bench) client for benchmarking an already-running `llama-server` through its OpenAI-compatible API. It is primarily meant to evaluate speculative decoding (draft model, n-gram, MTP, EAGLE3, ...) by reporting per-category throughput, latency, and draft acceptance.
|
||||
|
||||
The dataset handling follows the [aiperf SPEED-Bench tutorial](https://github.com/ai-dynamo/aiperf/blob/main/docs/tutorials/speed-bench.md), which also documents the dataset layout in more detail.
|
||||
|
||||
## Install
|
||||
|
||||
```bash
|
||||
pip install -r tools/server/bench/speed-bench/requirements.txt
|
||||
```
|
||||
|
||||
## Start a server
|
||||
|
||||
The client does not launch the server, so start `llama-server` yourself first. If you care about throughput numbers, set the client `--concurrency` to the server's slot count (`--np`):
|
||||
|
||||
```bash
|
||||
llama-server \
|
||||
-m target.gguf \
|
||||
-c 8192 \
|
||||
--port 8080 \
|
||||
-ngl 99 -fa on \
|
||||
--np 1 \
|
||||
--jinja
|
||||
```
|
||||
|
||||
For speculative decoding, start the server with the appropriate flags for your setup (e.g. a draft model with `-md`, or `--spec-type ngram-mod`). See the [speculative decoding doc](../../../../docs/speculative.md) for details.
|
||||
|
||||
## Run
|
||||
|
||||
```bash
|
||||
python tools/server/bench/speed-bench/speed_bench.py \
|
||||
--url localhost:8080 \
|
||||
--bench qualitative \
|
||||
--category coding \
|
||||
--osl 1024 \
|
||||
--concurrency 1
|
||||
```
|
||||
|
||||
## Options
|
||||
|
||||
| Option | Default | Description |
|
||||
| --- | --- | --- |
|
||||
| `--url` | `localhost:8080` | Server URL. The scheme and `/v1` are optional and a trailing slash is fine, so `localhost:8080` and `http://localhost:8080/v1/` both work. |
|
||||
| `--model` | none | Optional `model` field sent in each request. |
|
||||
| `--bench` | `qualitative` | SPEED-Bench config, e.g. `qualitative`, `throughput_1k`. See [available dataset variants](https://github.com/ai-dynamo/aiperf/blob/main/docs/tutorials/speed-bench.md#available-dataset-variants). |
|
||||
| `--category` | `all` | Category filter within the bench; comma-separated list or `all`. For `qualitative` the categories are `coding`, `humanities`, `math`, `multilingual`, `qa`, `rag`, `reasoning`, `roleplay`, `stem`, `summarization`, `writing`. For the `throughput_{ISL}` splits they are `high_entropy`, `low_entropy`, `mixed`. |
|
||||
| `--osl` | `1024` | Output sequence length, mapped to `max_tokens`. |
|
||||
| `--extra-inputs` | `{"temperature":0}` | Extra request fields as a JSON object. |
|
||||
| `--concurrency` | `1` | Concurrent client requests; usually match `--np`. |
|
||||
| `--limit` | none | Max samples per category (handy for smoke tests). |
|
||||
| `--timeout` | `600` | Per-request timeout in seconds. |
|
||||
| `--output` | none | Save raw per-request results and the summary to JSON. |
|
||||
|
||||
A few common ones:
|
||||
|
||||
- `--category all` runs every category in the bench.
|
||||
- `--category coding,math` runs just those two.
|
||||
- `--bench throughput_8k` runs a fixed-input-length throughput split.
|
||||
- `--limit 8` keeps at most 8 samples per category, which is enough for a quick check.
|
||||
|
||||
The `throughput_{ISL}` splits use fixed input lengths (1k - 32k), so they are handy for long-context testing and for comparing different `llama-server` batching settings (e.g. sweeping `-ub` / `--ubatch-size`) on prompts of a known size. Make sure the server `-c` is large enough for the chosen split. When raising `-ub`, also raise `-b` to at least the same value, since the physical ubatch cannot exceed the logical batch.
|
||||
|
||||
When `--output` is given, the JSON file holds the run `config`, the `selected_samples` / `completed_samples` / `failed_samples` counts, the per-category `summary` rows, and the per-sample `results`.
|
||||
|
||||
## Metrics
|
||||
|
||||
The summary prints one row per category plus an `overall` row:
|
||||
|
||||
- `samples` - how many samples finished successfully.
|
||||
- `avg_prompt_t/s` - prefill throughput from llama.cpp (`timings.prompt_per_second`), averaged over the category's samples.
|
||||
- `avg_pred_t/s` - decode throughput from llama.cpp (`timings.predicted_per_second`), averaged over the category's samples.
|
||||
- `avg_latency` - average end-to-end request latency seen by the client.
|
||||
- `accept_rate` - `accepted / draft_n` over the category, or `n/a` if nothing was drafted (`draft_n == 0`).
|
||||
|
||||
## Baseline vs speculative decoding
|
||||
|
||||
Save a run from each server with `--output`, then diff the two JSON files with `speed_bench_compare.py`.
|
||||
|
||||
First, start a plain `llama-server` (no speculative decoding) and save a baseline:
|
||||
|
||||
```bash
|
||||
python tools/server/bench/speed-bench/speed_bench.py \
|
||||
--url localhost:8080 \
|
||||
--bench qualitative \
|
||||
--category all \
|
||||
--osl 1024 \
|
||||
--concurrency 1 \
|
||||
--output baseline.json
|
||||
```
|
||||
|
||||
Then restart `llama-server` with speculative decoding enabled and save another run:
|
||||
|
||||
```bash
|
||||
python tools/server/bench/speed-bench/speed_bench.py \
|
||||
--url localhost:8080 \
|
||||
--bench qualitative \
|
||||
--category all \
|
||||
--osl 1024 \
|
||||
--concurrency 1 \
|
||||
--output spec.json
|
||||
```
|
||||
|
||||
Finally compare the two:
|
||||
|
||||
```bash
|
||||
python tools/server/bench/speed-bench/speed_bench_compare.py \
|
||||
--baseline baseline.json \
|
||||
--speculative spec.json
|
||||
```
|
||||
|
||||
The comparison table adds:
|
||||
|
||||
- `decode_speedup = spec_avg_pred_t/s / base_avg_pred_t/s`
|
||||
- `latency_speedup = base_avg_latency / spec_avg_latency`
|
||||
|
||||
Keep `--bench`, `--category`, `--osl`, and `--limit` the same across both runs, otherwise they won't be using the same prompts.
|
||||
@@ -0,0 +1,3 @@
|
||||
datasets
|
||||
requests
|
||||
tqdm
|
||||
@@ -0,0 +1,432 @@
|
||||
#!/usr/bin/env python3
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import concurrent.futures
|
||||
import json
|
||||
import statistics
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import asdict, dataclass
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import requests
|
||||
from datasets import get_dataset_config_names, load_dataset
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
DATASET_REPO = "nvidia/SPEED-Bench"
|
||||
|
||||
@dataclass
|
||||
class Sample:
|
||||
id: str
|
||||
category: str
|
||||
turns: list[str]
|
||||
|
||||
|
||||
@dataclass
|
||||
class RequestResult:
|
||||
id: str
|
||||
category: str
|
||||
ok: bool
|
||||
turns: int
|
||||
latency_s: float
|
||||
prompt_tokens: int
|
||||
completion_tokens: int
|
||||
total_tokens: int
|
||||
finish_reason: str | None
|
||||
draft_n: int
|
||||
draft_n_accepted: int
|
||||
prompt_ms: float | None
|
||||
predicted_ms: float | None
|
||||
prompt_per_second: float | None
|
||||
predicted_per_second: float | None
|
||||
error: str | None
|
||||
|
||||
|
||||
def normalize_base_url(url: str) -> str:
|
||||
url = url.strip().rstrip("/")
|
||||
if not url:
|
||||
raise ValueError("--url cannot be empty")
|
||||
if "://" not in url:
|
||||
url = "http://" + url
|
||||
parsed = urlparse(url)
|
||||
if not parsed.scheme or not parsed.netloc:
|
||||
raise ValueError(f"invalid --url: {url}")
|
||||
if not parsed.path.rstrip("/").endswith("/v1"):
|
||||
url = url + "/v1"
|
||||
return url.rstrip("/")
|
||||
|
||||
|
||||
def parse_extra_inputs(value: str) -> dict[str, Any]:
|
||||
extra = json.loads(value)
|
||||
if not isinstance(extra, dict):
|
||||
raise ValueError("--extra-inputs must be a JSON object")
|
||||
return extra
|
||||
|
||||
|
||||
def extract_turns(row: dict[str, Any]) -> list[str]:
|
||||
turns = row.get("turns")
|
||||
if isinstance(turns, list) and turns:
|
||||
clean_turns = [str(turn).strip() for turn in turns if turn and str(turn).strip()]
|
||||
if clean_turns:
|
||||
return clean_turns
|
||||
raise ValueError("missing or empty turns")
|
||||
|
||||
|
||||
def load_samples(args: argparse.Namespace) -> list[Sample]:
|
||||
bench_names = get_dataset_config_names(DATASET_REPO)
|
||||
if args.bench not in bench_names:
|
||||
raise ValueError(
|
||||
f"unknown --bench {args.bench!r}; available benches: {', '.join(bench_names)}"
|
||||
)
|
||||
|
||||
dataset = load_dataset(DATASET_REPO, name=args.bench, split="test")
|
||||
categories = list(dict.fromkeys(str(category) for category in dataset["category"]))
|
||||
requested_categories = None
|
||||
if args.category != "all":
|
||||
requested_list = [category.strip() for category in args.category.split(",") if category.strip()]
|
||||
if not requested_list:
|
||||
raise ValueError(
|
||||
f"--category must be 'all' or a comma-separated list; available categories: {', '.join(categories)}"
|
||||
)
|
||||
requested_categories = set(requested_list)
|
||||
unknown_categories = [category for category in requested_list if category not in categories]
|
||||
if unknown_categories:
|
||||
unknown = ", ".join(unknown_categories)
|
||||
raise ValueError(
|
||||
f"unknown --category {unknown!r} for bench {args.bench!r}; "
|
||||
f"available categories: all, {', '.join(categories)}"
|
||||
)
|
||||
|
||||
samples: list[Sample] = []
|
||||
samples_per_category: dict[str, int] = {}
|
||||
skipped = 0
|
||||
for index, row_raw in enumerate(dataset):
|
||||
row = dict(row_raw)
|
||||
category_raw = row.get("category")
|
||||
if not isinstance(category_raw, str) or not category_raw.strip():
|
||||
skipped += 1
|
||||
continue
|
||||
category = category_raw.strip()
|
||||
if requested_categories is not None and category not in requested_categories:
|
||||
continue
|
||||
if args.limit is not None and samples_per_category.get(category, 0) >= args.limit:
|
||||
continue
|
||||
|
||||
try:
|
||||
turns = extract_turns(row)
|
||||
except ValueError:
|
||||
skipped += 1
|
||||
continue
|
||||
question_id = row.get("question_id")
|
||||
if not isinstance(question_id, str) or not question_id.strip():
|
||||
skipped += 1
|
||||
continue
|
||||
sample_id = question_id.strip()
|
||||
samples.append(Sample(id=sample_id, category=category, turns=turns))
|
||||
samples_per_category[category] = samples_per_category.get(category, 0) + 1
|
||||
|
||||
if not samples:
|
||||
raise RuntimeError(f"no samples selected from bench={args.bench} category={args.category}")
|
||||
|
||||
if skipped:
|
||||
print(f"speed_bench: skipped {skipped} rows without usable turns")
|
||||
return samples
|
||||
|
||||
|
||||
def parse_completion_response(data: dict[str, Any]) -> tuple[dict[str, Any], dict[str, Any], str | None, str]:
|
||||
usage = data.get("usage") or {}
|
||||
timings = data.get("timings") or {}
|
||||
finish_reason = None
|
||||
content = ""
|
||||
choices = data.get("choices")
|
||||
if isinstance(choices, list) and choices and isinstance(choices[0], dict):
|
||||
choice = choices[0]
|
||||
finish_reason = choice.get("finish_reason")
|
||||
message = choice.get("message")
|
||||
if isinstance(message, dict) and isinstance(message.get("content"), str):
|
||||
content = message["content"]
|
||||
elif isinstance(choice.get("text"), str):
|
||||
content = choice["text"]
|
||||
return usage, timings, finish_reason, content
|
||||
|
||||
|
||||
def run_request(
|
||||
endpoint: str,
|
||||
model: str | None,
|
||||
messages: list[dict[str, str]],
|
||||
osl: int,
|
||||
extra_inputs: dict[str, Any],
|
||||
timeout: float,
|
||||
) -> tuple[dict[str, Any], float]:
|
||||
payload: dict[str, Any] = {
|
||||
"messages": messages,
|
||||
"max_tokens": osl,
|
||||
"stream": False,
|
||||
}
|
||||
if model:
|
||||
payload["model"] = model
|
||||
payload.update(extra_inputs)
|
||||
payload["max_tokens"] = osl
|
||||
|
||||
start = time.perf_counter()
|
||||
response = requests.post(endpoint, json=payload, timeout=timeout)
|
||||
latency_s = time.perf_counter() - start
|
||||
if response.status_code != 200:
|
||||
body = response.text[:500].replace("\n", "\\n")
|
||||
raise RuntimeError(f"HTTP {response.status_code}: {body}")
|
||||
return response.json(), latency_s
|
||||
|
||||
|
||||
def run_one(
|
||||
sample: Sample,
|
||||
endpoint: str,
|
||||
model: str | None,
|
||||
osl: int,
|
||||
extra_inputs: dict[str, Any],
|
||||
timeout: float,
|
||||
) -> RequestResult:
|
||||
selected_turns = sample.turns
|
||||
messages: list[dict[str, str]] = []
|
||||
total_latency_s = 0.0
|
||||
prompt_tokens = 0
|
||||
completion_tokens = 0
|
||||
total_tokens = 0
|
||||
draft_n = 0
|
||||
draft_n_accepted = 0
|
||||
prompt_ms = 0.0
|
||||
predicted_ms = 0.0
|
||||
prompt_per_second = None
|
||||
predicted_per_second = None
|
||||
finish_reason: str | None = None
|
||||
try:
|
||||
for turn in selected_turns:
|
||||
messages.append({"role": "user", "content": turn})
|
||||
data, latency_s = run_request(endpoint, model, messages, osl, extra_inputs, timeout)
|
||||
total_latency_s += latency_s
|
||||
usage, timings, finish_reason, assistant_text = parse_completion_response(data)
|
||||
|
||||
turn_prompt_tokens = int(usage.get("prompt_tokens") or timings.get("prompt_n") or 0)
|
||||
turn_completion_tokens_count = int(usage.get("completion_tokens") or timings.get("predicted_n") or 0)
|
||||
turn_total_tokens_count = int(usage.get("total_tokens") or (turn_prompt_tokens + turn_completion_tokens_count))
|
||||
prompt_tokens += turn_prompt_tokens
|
||||
completion_tokens += turn_completion_tokens_count
|
||||
total_tokens += turn_total_tokens_count
|
||||
draft_n += int(timings.get("draft_n") or 0)
|
||||
draft_n_accepted += int(timings.get("draft_n_accepted") or 0)
|
||||
prompt_ms += float(timings.get("prompt_ms") or 0)
|
||||
predicted_ms += float(timings.get("predicted_ms") or 0)
|
||||
if len(selected_turns) == 1 and isinstance(timings.get("prompt_per_second"), (int, float)):
|
||||
prompt_per_second = float(timings["prompt_per_second"])
|
||||
if len(selected_turns) == 1 and isinstance(timings.get("predicted_per_second"), (int, float)):
|
||||
predicted_per_second = float(timings["predicted_per_second"])
|
||||
|
||||
messages.append({"role": "assistant", "content": assistant_text})
|
||||
|
||||
if total_tokens == 0:
|
||||
total_tokens = prompt_tokens + completion_tokens
|
||||
if len(selected_turns) > 1:
|
||||
prompt_per_second = (prompt_tokens / (prompt_ms / 1000)) if prompt_ms > 0 else None
|
||||
predicted_per_second = (completion_tokens / (predicted_ms / 1000)) if predicted_ms > 0 else None
|
||||
|
||||
return RequestResult(
|
||||
id=sample.id,
|
||||
category=sample.category,
|
||||
ok=True,
|
||||
turns=len(selected_turns),
|
||||
latency_s=total_latency_s,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=total_tokens,
|
||||
finish_reason=finish_reason,
|
||||
draft_n=draft_n,
|
||||
draft_n_accepted=draft_n_accepted,
|
||||
prompt_ms=prompt_ms if prompt_ms > 0 else None,
|
||||
predicted_ms=predicted_ms if predicted_ms > 0 else None,
|
||||
prompt_per_second=prompt_per_second,
|
||||
predicted_per_second=predicted_per_second,
|
||||
error=None,
|
||||
)
|
||||
except Exception as exc:
|
||||
return RequestResult(
|
||||
id=sample.id,
|
||||
category=sample.category,
|
||||
ok=False,
|
||||
turns=len(selected_turns),
|
||||
latency_s=total_latency_s,
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
total_tokens=0,
|
||||
finish_reason=None,
|
||||
draft_n=0,
|
||||
draft_n_accepted=0,
|
||||
prompt_ms=None,
|
||||
predicted_ms=None,
|
||||
prompt_per_second=None,
|
||||
predicted_per_second=None,
|
||||
error=str(exc),
|
||||
)
|
||||
|
||||
|
||||
def summarize_group(category: str, results: list[RequestResult]) -> dict[str, Any]:
|
||||
ok_results = [result for result in results if result.ok]
|
||||
latencies = [result.latency_s for result in ok_results]
|
||||
server_prompt_speeds = [
|
||||
result.prompt_per_second
|
||||
for result in ok_results
|
||||
if result.prompt_per_second is not None
|
||||
]
|
||||
server_completion_speeds = [
|
||||
result.predicted_per_second
|
||||
for result in ok_results
|
||||
if result.predicted_per_second is not None
|
||||
]
|
||||
turns = sum(result.turns for result in ok_results)
|
||||
draft_n = sum(result.draft_n for result in ok_results)
|
||||
accepted = sum(result.draft_n_accepted for result in ok_results)
|
||||
|
||||
return {
|
||||
"category": category,
|
||||
"requests": len(ok_results),
|
||||
"turns": turns,
|
||||
"failed": len(results) - len(ok_results),
|
||||
"avg_prompt_t_s": statistics.mean(server_prompt_speeds) if server_prompt_speeds else None,
|
||||
"avg_pred_t_s": statistics.mean(server_completion_speeds) if server_completion_speeds else None,
|
||||
"avg_latency": statistics.mean(latencies) if latencies else None,
|
||||
"draft_n": draft_n,
|
||||
"accepted": accepted,
|
||||
"accept_rate": (accepted / draft_n) if draft_n > 0 else None,
|
||||
}
|
||||
|
||||
|
||||
def fmt_value(value: Any, kind: str = "") -> str:
|
||||
if value is None:
|
||||
return "n/a"
|
||||
if kind == "int":
|
||||
return str(int(value))
|
||||
if kind == "rate":
|
||||
return f"{float(value):.4f}"
|
||||
if kind == "seconds":
|
||||
return f"{float(value):.3f}s"
|
||||
if kind == "speed":
|
||||
return f"{float(value):.2f}"
|
||||
if kind == "speedup":
|
||||
return f"{float(value):.2f}x"
|
||||
return str(value)
|
||||
|
||||
|
||||
def print_table(rows: list[dict[str, Any]]) -> None:
|
||||
columns = [
|
||||
("category", "category", ""),
|
||||
("samples", "requests", "int"),
|
||||
("avg_prompt_t/s", "avg_prompt_t_s", "speed"),
|
||||
("avg_pred_t/s", "avg_pred_t_s", "speed"),
|
||||
("avg_latency", "avg_latency", "seconds"),
|
||||
("accept_rate", "accept_rate", "rate"),
|
||||
]
|
||||
print_rows(rows, columns)
|
||||
|
||||
|
||||
def print_rows(rows: list[dict[str, Any]], columns: list[tuple[str, str, str]]) -> None:
|
||||
rendered_rows = []
|
||||
for row in rows:
|
||||
rendered_rows.append([fmt_value(row.get(key), kind) for _, key, kind in columns])
|
||||
|
||||
widths = [len(header) for header, _, _ in columns]
|
||||
for rendered in rendered_rows:
|
||||
for i, cell in enumerate(rendered):
|
||||
widths[i] = max(widths[i], len(cell))
|
||||
|
||||
header = " ".join(header.ljust(widths[i]) for i, (header, _, _) in enumerate(columns))
|
||||
print(header)
|
||||
print(" ".join("-" * width for width in widths))
|
||||
for rendered in rendered_rows:
|
||||
print(" ".join(cell.ljust(widths[i]) for i, cell in enumerate(rendered)))
|
||||
|
||||
|
||||
def save_output(path: str, args: argparse.Namespace, samples: list[Sample], results: list[RequestResult], summary: list[dict[str, Any]]) -> None:
|
||||
payload = {
|
||||
"config": {
|
||||
"url": args.url,
|
||||
"model": args.model,
|
||||
"bench": args.bench,
|
||||
"category": args.category,
|
||||
"osl": args.osl,
|
||||
"concurrency": args.concurrency,
|
||||
"extra_inputs": args.extra_inputs,
|
||||
},
|
||||
"selected_samples": len(samples),
|
||||
"completed_samples": sum(1 for result in results if result.ok),
|
||||
"failed_samples": sum(1 for result in results if not result.ok),
|
||||
"summary": summary,
|
||||
"results": [asdict(result) for result in results],
|
||||
}
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
json.dump(payload, f, indent=2, sort_keys=True)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
parser = argparse.ArgumentParser(description="Run SPEED-Bench against an OpenAI-compatible llama-server.")
|
||||
parser.add_argument("--url", default="localhost:8080", help="Server URL, for example localhost:8080 or http://localhost:8080/v1")
|
||||
parser.add_argument("--model", default=None, help="Optional model name to send in OpenAI requests")
|
||||
parser.add_argument("--bench", default="qualitative", help="SPEED-Bench config to run, for example qualitative or throughput_1k")
|
||||
parser.add_argument("--category", default="all", help="Category to run within the selected bench; use all for no category filter")
|
||||
parser.add_argument("--osl", type=int, default=4096, help="Output sequence length, mapped to max_tokens")
|
||||
parser.add_argument("--extra-inputs", default='{"temperature":0}', help="Extra request fields as a JSON object")
|
||||
parser.add_argument("--concurrency", type=int, default=1, help="Concurrent client requests; usually match llama-server --np")
|
||||
parser.add_argument("--limit", type=int, default=None, help="Optional sample limit per category for smoke tests")
|
||||
parser.add_argument("--timeout", type=float, default=600, help="Per-request timeout in seconds")
|
||||
parser.add_argument("--output", default=None, help="Optional path to save raw results JSON")
|
||||
args = parser.parse_args(argv)
|
||||
try:
|
||||
base_url = normalize_base_url(args.url)
|
||||
endpoint = base_url + "/chat/completions"
|
||||
extra_inputs = parse_extra_inputs(args.extra_inputs)
|
||||
args.extra_inputs = extra_inputs
|
||||
samples = load_samples(args)
|
||||
except Exception as exc:
|
||||
print(f"speed_bench: setup failed: {exc}", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
print(f"speed_bench: loaded {len(samples)} samples from bench={args.bench} category={args.category}")
|
||||
|
||||
results: list[RequestResult] = []
|
||||
started = time.perf_counter()
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=args.concurrency) as executor:
|
||||
futures = [
|
||||
executor.submit(run_one, sample, endpoint, args.model, args.osl, extra_inputs, args.timeout)
|
||||
for sample in samples
|
||||
]
|
||||
for future in tqdm(concurrent.futures.as_completed(futures), total=len(futures), desc="speed_bench", unit="sample"):
|
||||
result = future.result()
|
||||
results.append(result)
|
||||
|
||||
elapsed = time.perf_counter() - started
|
||||
categories = list(dict.fromkeys(sample.category for sample in samples))
|
||||
summary = [
|
||||
summarize_group(category, [result for result in results if result.category == category])
|
||||
for category in categories
|
||||
]
|
||||
summary.append(summarize_group("overall", results))
|
||||
print()
|
||||
print(f"Summary (elapsed={elapsed:.2f}s)")
|
||||
print_table(summary)
|
||||
|
||||
if args.output:
|
||||
save_output(args.output, args, samples, results, summary)
|
||||
print(f"\nspeed_bench: wrote {args.output}")
|
||||
|
||||
failed = sum(1 for result in results if not result.ok)
|
||||
if failed:
|
||||
print(f"\nspeed_bench: {failed} samples failed", file=sys.stderr)
|
||||
first_error = next((result.error for result in results if result.error), None)
|
||||
if first_error:
|
||||
print(f"first error: {first_error}", file=sys.stderr)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,84 @@
|
||||
#!/usr/bin/env python3
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
from typing import Any
|
||||
|
||||
from speed_bench import fmt_value, print_rows
|
||||
|
||||
|
||||
def load_summary(path: str) -> list[dict[str, Any]]:
|
||||
with open(path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
summary = data.get("summary")
|
||||
if not isinstance(summary, list):
|
||||
raise ValueError(f"{path} does not contain a summary list")
|
||||
return summary
|
||||
|
||||
|
||||
def compare_rows(baseline: list[dict[str, Any]], speculative: list[dict[str, Any]]) -> list[dict[str, Any]]:
|
||||
baseline_by_category = {row["category"]: row for row in baseline}
|
||||
comparisons = []
|
||||
for row in speculative:
|
||||
base = baseline_by_category.get(row["category"])
|
||||
if not base:
|
||||
continue
|
||||
base_speed = base.get("avg_pred_t_s")
|
||||
spec_speed = row.get("avg_pred_t_s")
|
||||
base_latency = base.get("avg_latency")
|
||||
spec_latency = row.get("avg_latency")
|
||||
comparisons.append(
|
||||
{
|
||||
"category": row["category"],
|
||||
"base_avg_pred_t_s": base_speed,
|
||||
"spec_avg_pred_t_s": spec_speed,
|
||||
"decode_speedup": (spec_speed / base_speed) if base_speed and spec_speed else None,
|
||||
"base_avg_latency": base_latency,
|
||||
"spec_avg_latency": spec_latency,
|
||||
"latency_speedup": (base_latency / spec_latency) if base_latency and spec_latency else None,
|
||||
"accept_rate": row.get("accept_rate"),
|
||||
}
|
||||
)
|
||||
return comparisons
|
||||
|
||||
|
||||
def print_comparison(rows: list[dict[str, Any]]) -> None:
|
||||
if not rows:
|
||||
print("No overlapping categories found for comparison.")
|
||||
return
|
||||
columns = [
|
||||
("category", "category", ""),
|
||||
("base_avg_pred_t/s", "base_avg_pred_t_s", "speed"),
|
||||
("spec_avg_pred_t/s", "spec_avg_pred_t_s", "speed"),
|
||||
("decode_speedup", "decode_speedup", "speedup"),
|
||||
("base_avg_latency", "base_avg_latency", "seconds"),
|
||||
("spec_avg_latency", "spec_avg_latency", "seconds"),
|
||||
("latency_speedup", "latency_speedup", "speedup"),
|
||||
("accept_rate", "accept_rate", "rate"),
|
||||
]
|
||||
print_rows(rows, columns)
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
parser = argparse.ArgumentParser(description="Compare two SPEED-Bench runs (baseline vs speculative).")
|
||||
parser.add_argument("--baseline", required=True, help="Baseline results JSON produced by speed_bench.py --output")
|
||||
parser.add_argument("--speculative", required=True, help="Speculative decoding results JSON produced by speed_bench.py --output")
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
try:
|
||||
baseline = load_summary(args.baseline)
|
||||
speculative = load_summary(args.speculative)
|
||||
except Exception as exc:
|
||||
print(f"speed_bench_compare: failed to load inputs: {exc}", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
comparisons = compare_rows(baseline, speculative)
|
||||
print(f"Comparison: baseline={args.baseline} speculative={args.speculative}")
|
||||
print_comparison(comparisons)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -1,109 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
API_URL="${API_URL:-http://127.0.0.1:8080}"
|
||||
|
||||
CHAT=(
|
||||
"Hello, Assistant."
|
||||
"Hello. How may I help you today?"
|
||||
)
|
||||
|
||||
INSTRUCTION="A chat between a curious human and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the human's questions."
|
||||
|
||||
trim() {
|
||||
shopt -s extglob
|
||||
set -- "${1##+([[:space:]])}"
|
||||
printf "%s" "${1%%+([[:space:]])}"
|
||||
}
|
||||
|
||||
trim_trailing() {
|
||||
shopt -s extglob
|
||||
printf "%s" "${1%%+([[:space:]])}"
|
||||
}
|
||||
|
||||
format_prompt() {
|
||||
if [[ "${#CHAT[@]}" -eq 0 ]]; then
|
||||
echo -n "[INST] <<SYS>>\n${INSTRUCTION}\n<</SYS>>"
|
||||
else
|
||||
LAST_INDEX=$(( ${#CHAT[@]} - 1 ))
|
||||
echo -n "${CHAT[$LAST_INDEX]}\n[INST] $1 [/INST]"
|
||||
fi
|
||||
}
|
||||
|
||||
tokenize() {
|
||||
curl \
|
||||
--silent \
|
||||
--request POST \
|
||||
--url "${API_URL}/tokenize" \
|
||||
--header "Content-Type: application/json" \
|
||||
--data-raw "$(jq -ns --arg content "$1" '{content:$content}')" \
|
||||
| jq '.tokens[]'
|
||||
}
|
||||
|
||||
N_KEEP=$(tokenize "[INST] <<SYS>>\n${INSTRUCTION}\n<</SYS>>" | wc -l)
|
||||
|
||||
chat_completion() {
|
||||
PROMPT="$(trim_trailing "$(format_prompt "$1")")"
|
||||
DATA="$(echo -n "$PROMPT" | jq -Rs --argjson n_keep $N_KEEP '{
|
||||
prompt: .,
|
||||
temperature: 0.2,
|
||||
top_k: 40,
|
||||
top_p: 0.9,
|
||||
n_keep: $n_keep,
|
||||
n_predict: 1024,
|
||||
stop: ["[INST]"],
|
||||
stream: true
|
||||
}')"
|
||||
|
||||
# Create a temporary file to hold the Python output
|
||||
TEMPFILE=$(mktemp)
|
||||
|
||||
exec 3< <(curl \
|
||||
--silent \
|
||||
--no-buffer \
|
||||
--request POST \
|
||||
--url "${API_URL}/completion" \
|
||||
--header "Content-Type: application/json" \
|
||||
--data-raw "${DATA}")
|
||||
|
||||
python -c "
|
||||
import json
|
||||
import sys
|
||||
|
||||
answer = ''
|
||||
while True:
|
||||
line = sys.stdin.readline()
|
||||
if not line:
|
||||
break
|
||||
if line.startswith('data: '):
|
||||
json_content = line[6:].strip()
|
||||
content = json.loads(json_content)['content']
|
||||
sys.stdout.write(content)
|
||||
sys.stdout.flush()
|
||||
answer += content
|
||||
|
||||
answer = answer.rstrip('\n')
|
||||
|
||||
# Write the answer to the temporary file
|
||||
with open('$TEMPFILE', 'w') as f:
|
||||
f.write(answer)
|
||||
" <&3
|
||||
|
||||
exec 3<&-
|
||||
|
||||
# Read the answer from the temporary file
|
||||
ANSWER=$(cat $TEMPFILE)
|
||||
|
||||
# Clean up the temporary file
|
||||
rm $TEMPFILE
|
||||
|
||||
printf "\n"
|
||||
|
||||
CHAT+=("$1" "$(trim "$ANSWER")")
|
||||
}
|
||||
|
||||
while true; do
|
||||
echo -en "\033[0;32m" # Green color
|
||||
read -r -e -p "> " QUESTION
|
||||
echo -en "\033[0m" # Reset color
|
||||
chat_completion "${QUESTION}"
|
||||
done
|
||||
@@ -1,131 +0,0 @@
|
||||
import * as readline from 'node:readline'
|
||||
import { stdin, stdout } from 'node:process'
|
||||
import { readFileSync } from 'node:fs'
|
||||
import { SchemaConverter } from './public_legacy/json-schema-to-grammar.mjs'
|
||||
|
||||
const args = process.argv.slice(2);
|
||||
const grammarJsonSchemaFile = args.find(
|
||||
(_, index) => args[index - 1] === "--grammar-json-schema"
|
||||
);
|
||||
|
||||
const no_cached_prompt = args.find(
|
||||
(_, index) => args[index - 1] === "--no-cache-prompt"
|
||||
) ?? "false";
|
||||
|
||||
const grammarFile = args.find((_, index) => args[index - 1] === "--grammar");
|
||||
|
||||
// Example usage: function,arguments
|
||||
const grammarJsonSchemaPropOrder = args.find(
|
||||
(_, index) => args[index - 1] === "--grammar-json-schema-prop-order"
|
||||
);
|
||||
const propOrder = grammarJsonSchemaPropOrder
|
||||
? grammarJsonSchemaPropOrder
|
||||
.split(",")
|
||||
.reduce((acc, cur, index) => ({ ...acc, [cur]: index }), {})
|
||||
: {};
|
||||
|
||||
let grammar = null
|
||||
if (grammarJsonSchemaFile) {
|
||||
let schema = JSON.parse(readFileSync(grammarJsonSchemaFile, 'utf-8'))
|
||||
const converter = new SchemaConverter({prop_order: propOrder, allow_fetch: true})
|
||||
schema = await converter.resolveRefs(schema, grammarJsonSchemaFile)
|
||||
converter.visit(schema, '')
|
||||
grammar = converter.formatGrammar()
|
||||
}
|
||||
if (grammarFile) {
|
||||
grammar = readFileSync(grammarFile, 'utf-8')
|
||||
}
|
||||
|
||||
// for cached prompt
|
||||
let slot_id = -1;
|
||||
|
||||
const API_URL = 'http://127.0.0.1:8080'
|
||||
|
||||
const chat = [
|
||||
{
|
||||
human: "Hello, Assistant.",
|
||||
assistant: "Hello. How may I help you today?"
|
||||
},
|
||||
{
|
||||
human: "Please tell me the largest city in Europe.",
|
||||
assistant: "Sure. The largest city in Europe is Moscow, the capital of Russia."
|
||||
},
|
||||
]
|
||||
|
||||
const instruction = `A chat between a curious human and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the human's questions.`
|
||||
|
||||
function format_prompt(question) {
|
||||
return `${instruction}\n${
|
||||
chat.map(m =>`### Human: ${m.human}\n### Assistant: ${m.assistant}`).join("\n")
|
||||
}\n### Human: ${question}\n### Assistant:`
|
||||
}
|
||||
|
||||
async function tokenize(content) {
|
||||
const result = await fetch(`${API_URL}/tokenize`, {
|
||||
method: 'POST',
|
||||
body: JSON.stringify({ content })
|
||||
})
|
||||
|
||||
if (!result.ok) {
|
||||
return []
|
||||
}
|
||||
|
||||
return await result.json().tokens
|
||||
}
|
||||
|
||||
const n_keep = await tokenize(instruction).length
|
||||
|
||||
async function chat_completion(question) {
|
||||
const result = await fetch(`${API_URL}/completion`, {
|
||||
method: 'POST',
|
||||
body: JSON.stringify({
|
||||
prompt: format_prompt(question),
|
||||
temperature: 0.2,
|
||||
top_k: 40,
|
||||
top_p: 0.9,
|
||||
n_keep: n_keep,
|
||||
n_predict: 256,
|
||||
cache_prompt: no_cached_prompt === "false",
|
||||
slot_id: slot_id,
|
||||
stop: ["\n### Human:"], // stop completion after generating this
|
||||
grammar,
|
||||
stream: true,
|
||||
})
|
||||
})
|
||||
|
||||
if (!result.ok) {
|
||||
return
|
||||
}
|
||||
|
||||
let answer = ''
|
||||
|
||||
for await (var chunk of result.body) {
|
||||
const t = Buffer.from(chunk).toString('utf8')
|
||||
if (t.startsWith('data: ')) {
|
||||
const message = JSON.parse(t.substring(6))
|
||||
slot_id = message.slot_id
|
||||
answer += message.content
|
||||
process.stdout.write(message.content)
|
||||
if (message.stop) {
|
||||
if (message.truncated) {
|
||||
chat.shift()
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
process.stdout.write('\n')
|
||||
chat.push({ human: question, assistant: answer.trimStart() })
|
||||
}
|
||||
|
||||
const rl = readline.createInterface({ input: stdin, output: stdout });
|
||||
|
||||
const readlineQuestion = (rl, query, options) => new Promise((resolve, reject) => {
|
||||
rl.question(query, options, resolve)
|
||||
});
|
||||
|
||||
while(true) {
|
||||
const question = await readlineQuestion(rl, '> ')
|
||||
await chat_completion(question)
|
||||
}
|
||||
@@ -1,80 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
|
||||
API_URL="${API_URL:-http://127.0.0.1:8080}"
|
||||
|
||||
CHAT=(
|
||||
"Hello, Assistant."
|
||||
"Hello. How may I help you today?"
|
||||
"Please tell me the largest city in Europe."
|
||||
"Sure. The largest city in Europe is Moscow, the capital of Russia."
|
||||
)
|
||||
|
||||
INSTRUCTION="A chat between a curious human and an artificial intelligence assistant. The assistant gives helpful, detailed, and polite answers to the human's questions."
|
||||
|
||||
trim() {
|
||||
shopt -s extglob
|
||||
set -- "${1##+([[:space:]])}"
|
||||
printf "%s" "${1%%+([[:space:]])}"
|
||||
}
|
||||
|
||||
trim_trailing() {
|
||||
shopt -s extglob
|
||||
printf "%s" "${1%%+([[:space:]])}"
|
||||
}
|
||||
|
||||
format_prompt() {
|
||||
echo -n "${INSTRUCTION}"
|
||||
printf "\n### Human: %s\n### Assistant: %s" "${CHAT[@]}" "$1"
|
||||
}
|
||||
|
||||
tokenize() {
|
||||
curl \
|
||||
--silent \
|
||||
--request POST \
|
||||
--url "${API_URL}/tokenize" \
|
||||
--header "Content-Type: application/json" \
|
||||
--data-raw "$(jq -ns --arg content "$1" '{content:$content}')" \
|
||||
| jq '.tokens[]'
|
||||
}
|
||||
|
||||
N_KEEP=$(tokenize "${INSTRUCTION}" | wc -l)
|
||||
|
||||
chat_completion() {
|
||||
PROMPT="$(trim_trailing "$(format_prompt "$1")")"
|
||||
DATA="$(echo -n "$PROMPT" | jq -Rs --argjson n_keep $N_KEEP '{
|
||||
prompt: .,
|
||||
temperature: 0.2,
|
||||
top_k: 40,
|
||||
top_p: 0.9,
|
||||
n_keep: $n_keep,
|
||||
n_predict: 256,
|
||||
cache_prompt: true,
|
||||
stop: ["\n### Human:"],
|
||||
stream: true
|
||||
}')"
|
||||
|
||||
ANSWER=''
|
||||
|
||||
while IFS= read -r LINE; do
|
||||
if [[ $LINE = data:* ]]; then
|
||||
CONTENT="$(echo "${LINE:5}" | jq -r '.content')"
|
||||
printf "%s" "${CONTENT}"
|
||||
ANSWER+="${CONTENT}"
|
||||
fi
|
||||
done < <(curl \
|
||||
--silent \
|
||||
--no-buffer \
|
||||
--request POST \
|
||||
--url "${API_URL}/completion" \
|
||||
--header "Content-Type: application/json" \
|
||||
--data-raw "${DATA}")
|
||||
|
||||
printf "\n"
|
||||
|
||||
CHAT+=("$1" "$(trim "$ANSWER")")
|
||||
}
|
||||
|
||||
while true; do
|
||||
read -r -e -p "> " QUESTION
|
||||
chat_completion "${QUESTION}"
|
||||
done
|
||||
@@ -1734,7 +1734,7 @@ private:
|
||||
return true;
|
||||
}
|
||||
|
||||
void send_partial_response(server_slot & slot, const completion_token_output & tkn, bool is_progress) {
|
||||
void send_partial_response(server_slot & slot, const completion_token_output & tkn, bool is_progress, bool is_begin = false) {
|
||||
auto res = std::make_unique<server_task_result_cmpl_partial>();
|
||||
|
||||
res->id = slot.task->id;
|
||||
@@ -1746,6 +1746,9 @@ private:
|
||||
res->progress.cache = slot.n_prompt_tokens_cache;
|
||||
res->progress.processed = slot.prompt.tokens.size();
|
||||
res->progress.time_ms = (ggml_time_us() - slot.t_start_process_prompt) / 1000;
|
||||
}
|
||||
if (is_begin) {
|
||||
res->is_begin = true;
|
||||
} else {
|
||||
res->content = tkn.text_to_send;
|
||||
res->tokens = { tkn.tok };
|
||||
@@ -2828,10 +2831,15 @@ private:
|
||||
|
||||
slot.prompt.tokens.keep_first(n_past);
|
||||
|
||||
// send initial 0% progress update if needed
|
||||
// this is to signal the client that the request has started processing
|
||||
if (slot.task->params.stream && slot.task->params.return_progress) {
|
||||
send_partial_response(slot, {}, true);
|
||||
if (slot.task->params.stream) {
|
||||
if (slot.task->params.return_progress) {
|
||||
// send initial 0% progress update if needed
|
||||
send_partial_response(slot, {}, true);
|
||||
} else {
|
||||
// otherwise, for streaming without progress, signal HTTP to send the headers (i.e. 200 status)
|
||||
send_partial_response(slot, {}, false, true);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3745,7 +3753,9 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl(
|
||||
// next responses are streamed
|
||||
// to be sent immediately
|
||||
json first_result_json = first_result->to_json();
|
||||
if (res_type == TASK_RESPONSE_TYPE_ANTHROPIC) {
|
||||
if (first_result_json == nullptr) {
|
||||
res->data = ""; // simply send HTTP headers and status code
|
||||
} else if (res_type == TASK_RESPONSE_TYPE_ANTHROPIC) {
|
||||
res->data = format_anthropic_sse(first_result_json);
|
||||
} else if (res_type == TASK_RESPONSE_TYPE_OAI_RESP) {
|
||||
res->data = format_oai_resp_sse(first_result_json);
|
||||
|
||||
@@ -1422,6 +1422,9 @@ void server_task_result_cmpl_partial::update(task_result_state & state) {
|
||||
|
||||
json server_task_result_cmpl_partial::to_json() {
|
||||
GGML_ASSERT(is_updated && "update() must be called before to_json()");
|
||||
if (is_begin) {
|
||||
return nullptr; // simply signal to HTTP handler to send the headers and status code
|
||||
}
|
||||
switch (res_type) {
|
||||
case TASK_RESPONSE_TYPE_NONE:
|
||||
return to_json_non_oaicompat();
|
||||
|
||||
@@ -47,7 +47,7 @@ enum stop_type {
|
||||
};
|
||||
|
||||
struct task_params {
|
||||
bool stream = true;
|
||||
bool stream = false;
|
||||
bool include_usage = false;
|
||||
bool cache_prompt = true; // remember the prompt to avoid reprocessing all prompt
|
||||
bool return_tokens = false;
|
||||
@@ -418,6 +418,8 @@ struct server_task_result_cmpl_partial : server_task_result {
|
||||
|
||||
bool post_sampling_probs;
|
||||
bool is_progress = false;
|
||||
bool is_begin = false; // whether to send 200 status to HTTP client (begin of SSE stream)
|
||||
// ref: https://github.com/ggml-org/llama.cpp/pull/23884
|
||||
completion_token_output prob_output;
|
||||
result_timings timings;
|
||||
result_prompt_progress progress;
|
||||
|
||||
@@ -7,3 +7,9 @@ bun.lockb
|
||||
|
||||
# Miscellaneous
|
||||
/static/
|
||||
|
||||
# Build output
|
||||
/dist/
|
||||
/build/
|
||||
/.svelte-kit/
|
||||
test-results
|
||||
|
||||
@@ -45,8 +45,8 @@ export default ts.config(
|
||||
}
|
||||
},
|
||||
{
|
||||
// Exclude Storybook files from main ESLint rules
|
||||
ignores: ['.storybook/**/*']
|
||||
// Exclude generated build output and Storybook files from ESLint
|
||||
ignores: ['dist/**', 'build/**', '.svelte-kit/**', 'test-results/**', '.storybook/**/*']
|
||||
},
|
||||
storybook.configs['flat/recommended']
|
||||
);
|
||||
|
||||
@@ -186,6 +186,7 @@ export enum MimeTypeAudio {
|
||||
WAVE = 'audio/wave',
|
||||
X_WAV = 'audio/x-wav',
|
||||
X_WAVE = 'audio/x-wave',
|
||||
VND_WAVE = 'audio/vnd.wave',
|
||||
X_PN_WAV = 'audio/x-pn-wav',
|
||||
WEBM = 'audio/webm',
|
||||
WEBM_OPUS = 'audio/webm;codecs=opus'
|
||||
|
||||
@@ -40,6 +40,7 @@ function getAudioInputFormat(mimeType: string): AudioInputFormat {
|
||||
normalizedMimeType === MimeTypeAudio.WAVE ||
|
||||
normalizedMimeType === MimeTypeAudio.X_WAV ||
|
||||
normalizedMimeType === MimeTypeAudio.X_WAVE ||
|
||||
normalizedMimeType === MimeTypeAudio.VND_WAVE ||
|
||||
normalizedMimeType === MimeTypeAudio.X_PN_WAV
|
||||
) {
|
||||
return FileTypeAudio.WAV;
|
||||
|
||||
@@ -40,6 +40,7 @@ export function getFileTypeCategory(mimeType: string): FileTypeCategory | null {
|
||||
case MimeTypeAudio.WAVE:
|
||||
case MimeTypeAudio.X_WAV:
|
||||
case MimeTypeAudio.X_WAVE:
|
||||
case MimeTypeAudio.VND_WAVE:
|
||||
case MimeTypeAudio.X_PN_WAV:
|
||||
case MimeTypeAudio.WEBM:
|
||||
case MimeTypeAudio.WEBM_OPUS:
|
||||
|
||||
Reference in New Issue
Block a user