From cfc7424080a336ef0a04a07431a102f9de3a4da3 Mon Sep 17 00:00:00 2001 From: ciru-ai Date: Mon, 13 Jul 2026 03:49:47 -0400 Subject: [PATCH 1/4] feat(rocmfpx): add optimized 2.5 bpw ROCmFP2 kernels --- ggml/include/ggml.h | 4 +- ggml/rocmfpx/rocmfp2_reference.c | 483 +++++++++++++++ ggml/rocmfpx/rocmfp2_reference.h | 151 +++++ ggml/rocmfpx/rocmfpx.c | 196 ++++++ ggml/rocmfpx/rocmfpx.h | 18 + ggml/rocmfpx/test_rocmfp2_reference.c | 565 ++++++++++++++++++ ggml/src/ggml-cpu/ggml-cpu.c | 42 ++ ggml/src/ggml-cpu/ops.cpp | 4 +- ggml/src/ggml-cuda/common.cuh | 7 + ggml/src/ggml-cuda/convert.cu | 10 + ggml/src/ggml-cuda/dequantize.cuh | 18 + ggml/src/ggml-cuda/getrows.cu | 4 + ggml/src/ggml-cuda/ggml-cuda.cu | 4 +- ggml/src/ggml-cuda/mmq.cu | 4 + ggml/src/ggml-cuda/mmq.cuh | 70 +++ ggml/src/ggml-cuda/mmvq.cu | 165 ++++- .../template-instances/generate_cu_files.py | 2 +- .../mmq-instance-q2_0_rocmfpx.cu | 5 + ggml/src/ggml-cuda/vecdotq.cuh | 71 +++ ggml/src/ggml-quants.c | 7 + ggml/src/ggml.c | 13 + gguf-py/gguf/constants.py | 3 + include/llama.h | 1 + scripts/check-rocmfp2-reference.sh | 33 + src/llama-model-loader.cpp | 2 + src/llama-quant.cpp | 13 + tests/test-backend-ops.cpp | 28 +- tests/test-quantize-fns.cpp | 6 +- tools/quantize/quantize.cpp | 1 + 29 files changed, 1895 insertions(+), 35 deletions(-) create mode 100644 ggml/rocmfpx/rocmfp2_reference.c create mode 100644 ggml/rocmfpx/rocmfp2_reference.h create mode 100644 ggml/rocmfpx/test_rocmfp2_reference.c create mode 100644 ggml/src/ggml-cuda/template-instances/mmq-instance-q2_0_rocmfpx.cu create mode 100755 scripts/check-rocmfp2-reference.sh diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 35fb96ab9..071168700 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -436,7 +436,8 @@ extern "C" { GGML_TYPE_Q3_0_ROCMFPX = 104, // ROCmFPx experimental 3-bit UE4M3-scale reference layout GGML_TYPE_TURBO3_0 = 105, // TurboQuant 3-bit KV-cache (3.5 bpw) GGML_TYPE_TURBO4_0 = 106, // TurboQuant 4-bit KV-cache (4.5 bpw) - GGML_TYPE_COUNT = 107, + GGML_TYPE_Q2_0_ROCMFPX = 107, // ROCmFPx experimental 2-bit S40 codebook + dual UE4M3 scales + GGML_TYPE_COUNT = 108, }; // precision @@ -490,6 +491,7 @@ extern "C" { GGML_FTYPE_MOSTLY_Q6_0_ROCMFPX = 110, // ROCmFPx experimental 6-bit reference layout GGML_FTYPE_MOSTLY_Q8_0_ROCMFPX = 111, // ROCmFPx experimental 8-bit reference layout GGML_FTYPE_MOSTLY_Q3_0_ROCMFPX = 112, // ROCmFPx experimental 3-bit reference layout + GGML_FTYPE_MOSTLY_Q2_0_ROCMFPX = 113, // ROCmFPx experimental 2-bit S40 codebook layout }; // available tensor operations: diff --git a/ggml/rocmfpx/rocmfp2_reference.c b/ggml/rocmfpx/rocmfp2_reference.c new file mode 100644 index 000000000..87881cf64 --- /dev/null +++ b/ggml/rocmfpx/rocmfp2_reference.c @@ -0,0 +1,483 @@ +#include "rocmfp2_reference.h" + +#include +#include +#include + +static bool rocmfp2_p1_mapping_is_valid(rocmfp2_p1_mapping mapping) { + return mapping == ROCMFP2_P1_MAPPING_MORD || mapping == ROCMFP2_P1_MAPPING_MSM; +} + +static bool rocmfp2_p1_codebook_is_valid(const rocmfp2_p1_codebook * codebook) { + return codebook != NULL && codebook->inner > 0 && codebook->outer > codebook->inner && codebook->outer <= 127; +} + +static double rocmfp2_p1_decode_scale_valid(uint8_t scale_byte) { + const int exponent = scale_byte >> 3; + const int mantissa = scale_byte & 7; + + if (exponent == 0) { + return ldexp((double) mantissa, -10); + } + + return ldexp((double) (8 + mantissa), exponent - 11); +} + +static int8_t rocmfp2_p1_decode_code_valid( + uint8_t code, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping) { + const int inner = (int) codebook->inner; + const int outer = (int) codebook->outer; + + if (mapping == ROCMFP2_P1_MAPPING_MORD) { + static const int sign[4] = { -1, -1, 1, 1 }; + const int magnitude = (code == 0 || code == 3) ? outer : inner; + return (int8_t) (sign[code] * magnitude); + } + + static const int sign[4] = { 1, 1, -1, -1 }; + const int magnitude = (code == 1 || code == 3) ? outer : inner; + return (int8_t) (sign[code] * magnitude); +} + +static uint8_t rocmfp2_p1_encode_semantic_valid( + bool negative, + bool outer, + rocmfp2_p1_mapping mapping) { + if (mapping == ROCMFP2_P1_MAPPING_MORD) { + if (negative) { + return outer ? 0u : 1u; + } + return outer ? 3u : 2u; + } + + if (negative) { + return outer ? 3u : 2u; + } + return outer ? 1u : 0u; +} + +/* Lower ranks implement the frozen semantic tie rules. */ +static int rocmfp2_p1_code_tie_rank( + uint8_t code, + bool source_negative, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping) { + const int value = (int) rocmfp2_p1_decode_code_valid(code, codebook, mapping); + const bool code_negative = value < 0; + const bool outer = abs(value) == (int) codebook->outer; + + /* Inner magnitude wins first; input sign wins a remaining sign tie. */ + return (outer ? 2 : 0) + (code_negative != source_negative ? 1 : 0); +} + +static uint8_t rocmfp2_p1_select_code_ref_valid( + double source, + double scale, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping) { + /* Every code reconstructs exact zero; P0 canonicalizes this tie to 00. */ + if (scale == 0.0) { + return 0; + } + + const bool source_negative = signbit(source) != 0; + double best_error = INFINITY; + int best_rank = 0; + uint8_t best_code = 0; + bool have_best = false; + + for (uint8_t code = 0; code < 4; ++code) { + const double reconstructed = (double) rocmfp2_p1_decode_code_valid(code, codebook, mapping) * scale; + const double delta = source - reconstructed; + const double error = delta * delta; + const int rank = rocmfp2_p1_code_tie_rank(code, source_negative, codebook, mapping); + + if (!have_best || error < best_error || (error == best_error && rank < best_rank)) { + best_error = error; + best_rank = rank; + best_code = code; + have_best = true; + } + } + + return best_code; +} + +static uint8_t rocmfp2_p1_select_code_optimized_valid( + double source, + double scale, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping) { + /* Keep serialization identical to the exhaustive reference at scale 0. */ + if (scale == 0.0) { + return 0; + } + + const double magnitude = fabs(source); + const double inner_reconstructed = (double) codebook->inner * scale; + const double outer_reconstructed = (double) codebook->outer * scale; + const double inner_delta = magnitude - inner_reconstructed; + const double outer_delta = magnitude - outer_reconstructed; + + /* Equality is the exact inner/outer midpoint tie and keeps inner. */ + const bool outer = outer_delta * outer_delta < inner_delta * inner_delta; + return rocmfp2_p1_encode_semantic_valid(signbit(source) != 0, outer, mapping); +} + +const char * rocmfp2_p1_status_name(rocmfp2_p1_status status) { + switch (status) { + case ROCMFP2_P1_OK: return "OK"; + case ROCMFP2_P1_NONFINITE_SOURCE: return "NONFINITE_SOURCE"; + case ROCMFP2_P1_INVALID_ARGUMENT: return "INVALID_ARGUMENT"; + case ROCMFP2_P1_INVALID_CODEBOOK: return "INVALID_CODEBOOK"; + case ROCMFP2_P1_INVALID_MAPPING: return "INVALID_MAPPING"; + case ROCMFP2_P1_INVALID_SCALE_METADATA: return "INVALID_SCALE_METADATA"; + case ROCMFP2_P1_INVALID_CODE: return "INVALID_CODE"; + } + + return "UNKNOWN_STATUS"; +} + +bool rocmfp2_p1_scale_is_valid(uint8_t scale_byte) { + return scale_byte <= 0x7e; +} + +double rocmfp2_p1_ue4m3_to_binary64(uint8_t scale_byte) { + if (!rocmfp2_p1_scale_is_valid(scale_byte)) { + return NAN; + } + + return rocmfp2_p1_decode_scale_valid(scale_byte); +} + +uint8_t rocmfp2_p1_nearest_ue4m3(double target) { + if (isnan(target) || target < 0.0) { + return 0xff; + } + if (target == 0.0) { + return 0; + } + + const double maximum = rocmfp2_p1_decode_scale_valid(0x7e); + if (!isfinite(target) || target >= maximum) { + return 0x7e; + } + + for (uint8_t upper_byte = 1; upper_byte <= 0x7e; ++upper_byte) { + const double upper = rocmfp2_p1_decode_scale_valid(upper_byte); + if (target <= upper) { + const uint8_t lower_byte = (uint8_t) (upper_byte - 1); + const double lower = rocmfp2_p1_decode_scale_valid(lower_byte); + + /* Exact midpoint ties select the smaller scale byte. */ + return target - lower <= upper - target ? lower_byte : upper_byte; + } + } + + return 0x7e; +} + +rocmfp2_p1_status rocmfp2_p1_decode_code( + uint8_t code, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + int8_t * value) { + if (value == NULL) { + return ROCMFP2_P1_INVALID_ARGUMENT; + } + if (!rocmfp2_p1_codebook_is_valid(codebook)) { + return ROCMFP2_P1_INVALID_CODEBOOK; + } + if (!rocmfp2_p1_mapping_is_valid(mapping)) { + return ROCMFP2_P1_INVALID_MAPPING; + } + if (code > 3) { + return ROCMFP2_P1_INVALID_CODE; + } + + *value = rocmfp2_p1_decode_code_valid(code, codebook, mapping); + return ROCMFP2_P1_OK; +} + +rocmfp2_p1_status rocmfp2_p1_decode_packed_byte( + uint8_t packed, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + int8_t values[4]) { + if (values == NULL) { + return ROCMFP2_P1_INVALID_ARGUMENT; + } + if (!rocmfp2_p1_codebook_is_valid(codebook)) { + return ROCMFP2_P1_INVALID_CODEBOOK; + } + if (!rocmfp2_p1_mapping_is_valid(mapping)) { + return ROCMFP2_P1_INVALID_MAPPING; + } + + for (int lane = 0; lane < 4; ++lane) { + const uint8_t code = (uint8_t) ((packed >> (2 * lane)) & 3u); + values[lane] = rocmfp2_p1_decode_code_valid(code, codebook, mapping); + } + + return ROCMFP2_P1_OK; +} + +rocmfp2_p1_status rocmfp2_p1_pack_codes( + const uint8_t codes[ROCMFP2_P1_BLOCK_WEIGHTS], + const uint8_t scales[ROCMFP2_P1_SCALE_BYTES], + rocmfp2_p1_block * block) { + if (codes == NULL || scales == NULL || block == NULL) { + return ROCMFP2_P1_INVALID_ARGUMENT; + } + if (!rocmfp2_p1_scale_is_valid(scales[0]) || !rocmfp2_p1_scale_is_valid(scales[1])) { + return ROCMFP2_P1_INVALID_SCALE_METADATA; + } + + rocmfp2_p1_block temporary; + for (int packed_index = 0; packed_index < ROCMFP2_P1_DATA_BYTES; ++packed_index) { + uint8_t packed = 0; + for (int lane = 0; lane < 4; ++lane) { + const uint8_t code = codes[4 * packed_index + lane]; + if (code > 3) { + return ROCMFP2_P1_INVALID_CODE; + } + packed |= (uint8_t) (code << (2 * lane)); + } + temporary.d[packed_index] = packed; + } + temporary.s[0] = scales[0]; + temporary.s[1] = scales[1]; + + memcpy(block, &temporary, sizeof(temporary)); + return ROCMFP2_P1_OK; +} + +void rocmfp2_p1_unpack_codes( + const rocmfp2_p1_block * block, + uint8_t codes[ROCMFP2_P1_BLOCK_WEIGHTS]) { + for (int packed_index = 0; packed_index < ROCMFP2_P1_DATA_BYTES; ++packed_index) { + const uint8_t packed = block->d[packed_index]; + for (int lane = 0; lane < 4; ++lane) { + codes[4 * packed_index + lane] = (uint8_t) ((packed >> (2 * lane)) & 3u); + } + } +} + +bool rocmfp2_p1_validate_block(const rocmfp2_p1_block * block) { + return block != NULL && rocmfp2_p1_scale_is_valid(block->s[0]) && rocmfp2_p1_scale_is_valid(block->s[1]); +} + +bool rocmfp2_p1_validate_serialized(const void * data, size_t nbytes) { + if (nbytes == 0) { + return true; + } + if (data == NULL || nbytes % ROCMFP2_P1_BLOCK_BYTES != 0) { + return false; + } + + const uint8_t * bytes = (const uint8_t *) data; + const size_t blocks = nbytes / ROCMFP2_P1_BLOCK_BYTES; + for (size_t block_index = 0; block_index < blocks; ++block_index) { + const size_t base = block_index * ROCMFP2_P1_BLOCK_BYTES; + if (!rocmfp2_p1_scale_is_valid(bytes[base + 8]) || !rocmfp2_p1_scale_is_valid(bytes[base + 9])) { + return false; + } + } + + return true; +} + +rocmfp2_p1_status rocmfp2_p1_select_code_ref( + double source, + uint8_t scale_byte, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + uint8_t * code) { + if (code == NULL) { + return ROCMFP2_P1_INVALID_ARGUMENT; + } + if (!rocmfp2_p1_codebook_is_valid(codebook)) { + return ROCMFP2_P1_INVALID_CODEBOOK; + } + if (!rocmfp2_p1_mapping_is_valid(mapping)) { + return ROCMFP2_P1_INVALID_MAPPING; + } + if (!rocmfp2_p1_scale_is_valid(scale_byte)) { + return ROCMFP2_P1_INVALID_SCALE_METADATA; + } + if (!isfinite(source)) { + return ROCMFP2_P1_NONFINITE_SOURCE; + } + + *code = rocmfp2_p1_select_code_ref_valid( + source, rocmfp2_p1_decode_scale_valid(scale_byte), codebook, mapping); + return ROCMFP2_P1_OK; +} + +rocmfp2_p1_status rocmfp2_p1_select_code_optimized( + double source, + uint8_t scale_byte, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + uint8_t * code) { + if (code == NULL) { + return ROCMFP2_P1_INVALID_ARGUMENT; + } + if (!rocmfp2_p1_codebook_is_valid(codebook)) { + return ROCMFP2_P1_INVALID_CODEBOOK; + } + if (!rocmfp2_p1_mapping_is_valid(mapping)) { + return ROCMFP2_P1_INVALID_MAPPING; + } + if (!rocmfp2_p1_scale_is_valid(scale_byte)) { + return ROCMFP2_P1_INVALID_SCALE_METADATA; + } + if (!isfinite(source)) { + return ROCMFP2_P1_NONFINITE_SOURCE; + } + + *code = rocmfp2_p1_select_code_optimized_valid( + source, rocmfp2_p1_decode_scale_valid(scale_byte), codebook, mapping); + return ROCMFP2_P1_OK; +} + +typedef uint8_t (*rocmfp2_p1_selector)( + double source, + double scale, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping); + +static rocmfp2_p1_status rocmfp2_p1_quantize_block_impl( + const double source[ROCMFP2_P1_BLOCK_WEIGHTS], + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + rocmfp2_p1_block * block, + double group_sse[ROCMFP2_P1_SCALE_BYTES], + rocmfp2_p1_selector selector) { + if (source == NULL || block == NULL) { + return ROCMFP2_P1_INVALID_ARGUMENT; + } + if (!rocmfp2_p1_codebook_is_valid(codebook)) { + return ROCMFP2_P1_INVALID_CODEBOOK; + } + if (!rocmfp2_p1_mapping_is_valid(mapping)) { + return ROCMFP2_P1_INVALID_MAPPING; + } + + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + if (!isfinite(source[i])) { + return ROCMFP2_P1_NONFINITE_SOURCE; + } + } + + uint8_t best_codes[ROCMFP2_P1_BLOCK_WEIGHTS]; + uint8_t best_scales[ROCMFP2_P1_SCALE_BYTES]; + double best_group_sse[ROCMFP2_P1_SCALE_BYTES]; + + for (int group = 0; group < ROCMFP2_P1_SCALE_BYTES; ++group) { + const int source_offset = group * ROCMFP2_P1_GROUP_WEIGHTS; + double best_sse = INFINITY; + uint8_t best_scale = 0; + uint8_t candidate_codes[ROCMFP2_P1_GROUP_WEIGHTS]; + uint8_t saved_codes[ROCMFP2_P1_GROUP_WEIGHTS]; + bool have_best = false; + + /* Exact exhaustive search: every legal finite unsigned UE4M3 byte. */ + for (int scale_index = 0; scale_index <= 0x7e; ++scale_index) { + const uint8_t scale_byte = (uint8_t) scale_index; + const double scale = rocmfp2_p1_decode_scale_valid(scale_byte); + double candidate_sse = 0.0; + + for (int i = 0; i < ROCMFP2_P1_GROUP_WEIGHTS; ++i) { + const double value = source[source_offset + i]; + const uint8_t code = selector(value, scale, codebook, mapping); + const double reconstructed = (double) rocmfp2_p1_decode_code_valid(code, codebook, mapping) * scale; + const double delta = value - reconstructed; + + candidate_codes[i] = code; + candidate_sse += delta * delta; + } + + if (!have_best || candidate_sse < best_sse || + (candidate_sse == best_sse && scale_byte < best_scale)) { + best_sse = candidate_sse; + best_scale = scale_byte; + memcpy(saved_codes, candidate_codes, sizeof(saved_codes)); + have_best = true; + } + } + + memcpy(best_codes + source_offset, saved_codes, sizeof(saved_codes)); + best_scales[group] = best_scale; + best_group_sse[group] = best_sse; + } + + rocmfp2_p1_block temporary; + const rocmfp2_p1_status pack_status = rocmfp2_p1_pack_codes(best_codes, best_scales, &temporary); + if (pack_status != ROCMFP2_P1_OK) { + return pack_status; + } + + memcpy(block, &temporary, sizeof(temporary)); + if (group_sse != NULL) { + group_sse[0] = best_group_sse[0]; + group_sse[1] = best_group_sse[1]; + } + + return ROCMFP2_P1_OK; +} + +rocmfp2_p1_status rocmfp2_p1_quantize_block_ref( + const double source[ROCMFP2_P1_BLOCK_WEIGHTS], + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + rocmfp2_p1_block * block, + double group_sse[ROCMFP2_P1_SCALE_BYTES]) { + return rocmfp2_p1_quantize_block_impl( + source, codebook, mapping, block, group_sse, rocmfp2_p1_select_code_ref_valid); +} + +rocmfp2_p1_status rocmfp2_p1_quantize_block_optimized( + const double source[ROCMFP2_P1_BLOCK_WEIGHTS], + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + rocmfp2_p1_block * block, + double group_sse[ROCMFP2_P1_SCALE_BYTES]) { + return rocmfp2_p1_quantize_block_impl( + source, codebook, mapping, block, group_sse, rocmfp2_p1_select_code_optimized_valid); +} + +rocmfp2_p1_status rocmfp2_p1_dequantize_block( + const rocmfp2_p1_block * block, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + double output[ROCMFP2_P1_BLOCK_WEIGHTS]) { + if (block == NULL || output == NULL) { + return ROCMFP2_P1_INVALID_ARGUMENT; + } + if (!rocmfp2_p1_codebook_is_valid(codebook)) { + return ROCMFP2_P1_INVALID_CODEBOOK; + } + if (!rocmfp2_p1_mapping_is_valid(mapping)) { + return ROCMFP2_P1_INVALID_MAPPING; + } + if (!rocmfp2_p1_validate_block(block)) { + return ROCMFP2_P1_INVALID_SCALE_METADATA; + } + + uint8_t codes[ROCMFP2_P1_BLOCK_WEIGHTS]; + double temporary[ROCMFP2_P1_BLOCK_WEIGHTS]; + rocmfp2_p1_unpack_codes(block, codes); + + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + const int group = i / ROCMFP2_P1_GROUP_WEIGHTS; + const double scale = rocmfp2_p1_decode_scale_valid(block->s[group]); + const int8_t value = rocmfp2_p1_decode_code_valid(codes[i], codebook, mapping); + temporary[i] = (double) value * scale; + } + + memcpy(output, temporary, sizeof(temporary)); + return ROCMFP2_P1_OK; +} diff --git a/ggml/rocmfpx/rocmfp2_reference.h b/ggml/rocmfpx/rocmfp2_reference.h new file mode 100644 index 000000000..2bc72579b --- /dev/null +++ b/ggml/rocmfpx/rocmfp2_reference.h @@ -0,0 +1,151 @@ +#ifndef ROCMFP2_REFERENCE_H +#define ROCMFP2_REFERENCE_H + +#include +#include +#include +#include + +#ifdef __cplusplus +extern "C" { +#endif + +enum { + ROCMFP2_P1_BLOCK_WEIGHTS = 32, + ROCMFP2_P1_GROUP_WEIGHTS = 16, + ROCMFP2_P1_DATA_BYTES = 8, + ROCMFP2_P1_SCALE_BYTES = 2, + ROCMFP2_P1_BLOCK_BYTES = 10, +}; + +/* + * Frozen Phase-1 P0_AOS serialization: + * + * d[0] ... d[7], s[0], s[1] + * + * Weight 4*j+i occupies bits 2*i+1:2*i of d[j]. Scale s[0] + * applies to weights 0..15 and s[1] to weights 16..31. This is a + * byte-canonical format; host word endianness is not part of serialization. + */ +typedef struct { + uint8_t d[ROCMFP2_P1_DATA_BYTES]; + uint8_t s[ROCMFP2_P1_SCALE_BYTES]; +} rocmfp2_p1_block; + +typedef struct { + uint8_t inner; + uint8_t outer; +} rocmfp2_p1_codebook; + +typedef enum { + ROCMFP2_P1_MAPPING_MORD = 0, + ROCMFP2_P1_MAPPING_MSM = 1, +} rocmfp2_p1_mapping; + +typedef enum { + ROCMFP2_P1_OK = 0, + ROCMFP2_P1_NONFINITE_SOURCE, + ROCMFP2_P1_INVALID_ARGUMENT, + ROCMFP2_P1_INVALID_CODEBOOK, + ROCMFP2_P1_INVALID_MAPPING, + ROCMFP2_P1_INVALID_SCALE_METADATA, + ROCMFP2_P1_INVALID_CODE, +} rocmfp2_p1_status; + +#if defined(__cplusplus) +static_assert(FLT_RADIX == 2 && DBL_MANT_DIG == 53 && DBL_MAX_EXP == 1024, + "ROCmFP2 Phase-1 reference requires IEEE-754 binary64 double"); +static_assert(sizeof(rocmfp2_p1_block) == ROCMFP2_P1_BLOCK_BYTES, + "ROCmFP2 P0 block has padding"); +static_assert(offsetof(rocmfp2_p1_block, d) == 0, "ROCmFP2 P0 data offset changed"); +static_assert(offsetof(rocmfp2_p1_block, s) == 8, "ROCmFP2 P0 scale offset changed"); +#else +_Static_assert(FLT_RADIX == 2 && DBL_MANT_DIG == 53 && DBL_MAX_EXP == 1024, + "ROCmFP2 Phase-1 reference requires IEEE-754 binary64 double"); +_Static_assert(sizeof(rocmfp2_p1_block) == ROCMFP2_P1_BLOCK_BYTES, + "ROCmFP2 P0 block has padding"); +_Static_assert(offsetof(rocmfp2_p1_block, d) == 0, "ROCmFP2 P0 data offset changed"); +_Static_assert(offsetof(rocmfp2_p1_block, s) == 8, "ROCmFP2 P0 scale offset changed"); +#endif + +const char * rocmfp2_p1_status_name(rocmfp2_p1_status status); + +bool rocmfp2_p1_scale_is_valid(uint8_t scale_byte); +double rocmfp2_p1_ue4m3_to_binary64(uint8_t scale_byte); +/* Returns invalid byte 0xff for NaN or a negative target. */ +uint8_t rocmfp2_p1_nearest_ue4m3(double target); + +rocmfp2_p1_status rocmfp2_p1_decode_code( + uint8_t code, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + int8_t * value); + +rocmfp2_p1_status rocmfp2_p1_decode_packed_byte( + uint8_t packed, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + int8_t values[4]); + +rocmfp2_p1_status rocmfp2_p1_pack_codes( + const uint8_t codes[ROCMFP2_P1_BLOCK_WEIGHTS], + const uint8_t scales[ROCMFP2_P1_SCALE_BYTES], + rocmfp2_p1_block * block); + +void rocmfp2_p1_unpack_codes( + const rocmfp2_p1_block * block, + uint8_t codes[ROCMFP2_P1_BLOCK_WEIGHTS]); + +bool rocmfp2_p1_validate_block(const rocmfp2_p1_block * block); +bool rocmfp2_p1_validate_serialized(const void * data, size_t nbytes); + +/* + * The code selectors are exposed so signed-zero and exact-midpoint behavior + * can be tested independently of scale selection. The reference selector + * exhaustively evaluates all four codes. The optimized selector uses the + * symmetric sign/magnitude structure. Both use binary64 arithmetic. + */ +rocmfp2_p1_status rocmfp2_p1_select_code_ref( + double source, + uint8_t scale_byte, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + uint8_t * code); + +rocmfp2_p1_status rocmfp2_p1_select_code_optimized( + double source, + uint8_t scale_byte, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + uint8_t * code); + +/* + * Both quantizers search every legal UE4M3 byte (0x00..0x7e) independently + * for each 16-weight group and minimize unweighted SSE in binary64. Exact + * scale ties select the lower byte. Outputs are committed only on success. + */ +rocmfp2_p1_status rocmfp2_p1_quantize_block_ref( + const double source[ROCMFP2_P1_BLOCK_WEIGHTS], + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + rocmfp2_p1_block * block, + double group_sse[ROCMFP2_P1_SCALE_BYTES]); + +rocmfp2_p1_status rocmfp2_p1_quantize_block_optimized( + const double source[ROCMFP2_P1_BLOCK_WEIGHTS], + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + rocmfp2_p1_block * block, + double group_sse[ROCMFP2_P1_SCALE_BYTES]); + +rocmfp2_p1_status rocmfp2_p1_dequantize_block( + const rocmfp2_p1_block * block, + const rocmfp2_p1_codebook * codebook, + rocmfp2_p1_mapping mapping, + double output[ROCMFP2_P1_BLOCK_WEIGHTS]); + +#ifdef __cplusplus +} +#endif + +#endif diff --git a/ggml/rocmfpx/rocmfpx.c b/ggml/rocmfpx/rocmfpx.c index cc57f1fdf..12b011265 100644 --- a/ggml/rocmfpx/rocmfpx.c +++ b/ggml/rocmfpx/rocmfpx.c @@ -175,6 +175,11 @@ size_t rocmfpx_row_size_fp3(int64_t k) { return (size_t) (k / QK_ROCMFP3) * sizeof(block_rocmfp3); } +size_t rocmfpx_row_size_fp2(int64_t k) { + assert(k % QK_ROCMFP2 == 0); + return (size_t) (k / QK_ROCMFP2) * sizeof(block_rocmfp2); +} + size_t rocmfpx_row_size_fp6(int64_t k) { assert(k % QK_ROCMFP6 == 0); return (size_t) (k / QK_ROCMFP6) * sizeof(block_rocmfp6); @@ -232,6 +237,182 @@ static void rocmfpx_prepare_mse_weights( } } +// ROCmFP2 S40 uses the frozen MORD code order {-4, -1, +1, +4}. +static inline int rocmfpx_decode_fp2_code(uint8_t code) { + static const int8_t values[4] = { -4, -1, 1, 4 }; + return values[code & 3u]; +} + +static inline uint8_t rocmfpx_quantize_fp2_code(float x, float inv_scale) { + if (!isfinite(x) || !(inv_scale > 0.0f)) { + return 0; + } + + const float magnitude = fabsf(x * inv_scale); + const bool outer = magnitude > 2.5f; + if (signbit(x)) { + return outer ? 0u : 1u; + } + return outer ? 3u : 2u; +} + +static float rocmfpx_fp2_group_mse_for_scale( + const float * x, const float * mse_weights, int n, uint8_t e, float best_err) { + const float scale = rocmfpx_scale_lookup(e); + const float inv_scale = scale > 0.0f ? 1.0f / scale : 0.0f; + float err = 0.0f; + + for (int i = 0; i < n; ++i) { + if (!isfinite(x[i])) { + continue; + } + const float reconstructed = (float) rocmfpx_decode_fp2_code(rocmfpx_quantize_fp2_code(x[i], inv_scale)) * scale; + const float delta = x[i] - reconstructed; + err += (mse_weights ? mse_weights[i] : 1.0f) * delta * delta; + if (err > best_err) { + return err; + } + } + return err; +} + +static uint8_t rocmfpx_choose_scale_fp2_mse( + const float * x, int n, const float * quant_weights, float sigma2) { + float mse_weights[QK_ROCMFP2/2]; + float max_abs = 0.0f; + float max_abs_weight = 0.0f; + bool all_finite = true; + + if (quant_weights) { + rocmfpx_prepare_mse_weights( + mse_weights, x, n, quant_weights, sigma2, + &max_abs, &max_abs_weight, &all_finite); + } else { + max_abs = rocmfpx_max_abs(x, n); + max_abs_weight = 1.0f; + } + GGML_UNUSED(all_finite); + + if (!(max_abs > 0.0f) || !isfinite(max_abs)) { + return 0; + } + + const float * weights = quant_weights ? mse_weights : NULL; + const uint8_t start_e = rocmfpx_nearest_scale_ue4m3(max_abs / 4.0f); + uint8_t best_e = start_e; + float best_err = INFINITY; + bool lower_done = false; + + for (int delta = 0; delta <= 125; ++delta) { + const int e0 = (int) start_e - delta; + if (!lower_done && e0 >= 1 && e0 <= 126) { + const float scale = rocmfpx_scale_lookup((uint8_t) e0); + const float clip_delta = max_abs - 4.0f * scale; + const float clip_err = max_abs_weight * clip_delta * clip_delta; + if (clip_delta > 0.0f && clip_err > best_err) { + lower_done = true; + } else { + const float err = rocmfpx_fp2_group_mse_for_scale(x, weights, n, (uint8_t) e0, best_err); + if (err < best_err || (err == best_err && e0 < best_e)) { + best_err = err; + best_e = (uint8_t) e0; + } + } + } + + const int e1 = (int) start_e + delta; + if (delta != 0 && e1 >= 1 && e1 <= 126) { + const float err = rocmfpx_fp2_group_mse_for_scale(x, weights, n, (uint8_t) e1, best_err); + if (err < best_err || (err == best_err && e1 < best_e)) { + best_err = err; + best_e = (uint8_t) e1; + } + } + + if ((lower_done || e0 <= 1) && e1 >= 126) { + break; + } + } + return best_e; +} + +static void rocmfpx_quantize_row_fp2_impl( + const float * GGML_RESTRICT x, block_rocmfp2 * GGML_RESTRICT y, + int64_t k, const float * GGML_RESTRICT quant_weights) { + assert(k % QK_ROCMFP2 == 0); + + float sum_x2 = 0.0f; + for (int64_t i = 0; i < k; ++i) { + sum_x2 += isfinite(x[i]) ? x[i] * x[i] : 0.0f; + } + const float sigma2 = sum_x2 / (float) k; + + const int64_t nb = k / QK_ROCMFP2; + for (int64_t ib = 0; ib < nb; ++ib) { + const float * xb = x + ib * QK_ROCMFP2; + const float * qw = quant_weights ? quant_weights + ib * QK_ROCMFP2 : NULL; + block_rocmfp2 * yb = y + ib; + + for (int half = 0; half < 2; ++half) { + const int half_off = half * (QK_ROCMFP2 / 2); + const float * xh = xb + half_off; + const float * qh = qw ? qw + half_off : NULL; + yb->e[half] = rocmfpx_choose_scale_fp2_mse(xh, QK_ROCMFP2 / 2, qh, sigma2); + + const float scale = rocmfpx_scale_lookup(yb->e[half]); + const float inv_scale = scale > 0.0f ? 1.0f / scale : 0.0f; + for (int packed = 0; packed < 4; ++packed) { + uint8_t byte = 0; + for (int lane = 0; lane < 4; ++lane) { + const int j = 4 * packed + lane; + byte |= (uint8_t) (rocmfpx_quantize_fp2_code(xh[j], inv_scale) << (2 * lane)); + } + yb->qs[half * 4 + packed] = byte; + } + } + } +} + +void rocmfpx_quantize_row_fp2_ref( + const float * GGML_RESTRICT x, block_rocmfp2 * GGML_RESTRICT y, int64_t k) { + rocmfpx_quantize_row_fp2_impl(x, y, k, NULL); +} + +void rocmfpx_dequantize_row_fp2( + const block_rocmfp2 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { + assert(k % QK_ROCMFP2 == 0); + const int64_t nb = k / QK_ROCMFP2; + for (int64_t ib = 0; ib < nb; ++ib) { + const block_rocmfp2 * xb = x + ib; + float * yb = y + ib * QK_ROCMFP2; + for (int half = 0; half < 2; ++half) { + const float scale = rocmfpx_scale_lookup(xb->e[half]); + for (int j = 0; j < QK_ROCMFP2 / 2; ++j) { + const uint8_t code = (xb->qs[half * 4 + j / 4] >> (2 * (j % 4))) & 3u; + yb[half * (QK_ROCMFP2 / 2) + j] = (float) rocmfpx_decode_fp2_code(code) * scale; + } + } + } +} + +void rocmfpx_quantize_row_fp2( + const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { + rocmfpx_quantize_row_fp2_ref(x, (block_rocmfp2 *) y, k); +} + +size_t rocmfpx_quantize_fp2( + const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, + int64_t nrows, int64_t n_per_row, const float * imatrix) { + const size_t row_size = rocmfpx_row_size_fp2(n_per_row); + char * qrow = (char *) dst; + for (int64_t row = 0; row < nrows; ++row) { + rocmfpx_quantize_row_fp2_impl( + src + row * n_per_row, (block_rocmfp2 *) qrow, n_per_row, imatrix); + qrow += row_size; + } + return (size_t) nrows * row_size; +} + // --------------------------------------------------------------------------- // FP3 group pack/unpack: 8 × 3-bit codes → 3 bytes (24 bits) // @@ -1111,6 +1292,21 @@ size_t rocmfpx_quantize_fp8(const float * GGML_RESTRICT src, void * GGML_RESTRIC return (size_t) nrows * row_size; } +bool rocmfpx_validate_row_data_fp2(const void * data, size_t nbytes) { + if (nbytes % sizeof(block_rocmfp2) != 0) { + return false; + } + + const block_rocmfp2 * blocks = (const block_rocmfp2 *) data; + const size_t nb = nbytes / sizeof(block_rocmfp2); + for (size_t i = 0; i < nb; ++i) { + if (!rocmfpx_scale_is_valid(blocks[i].e[0]) || !rocmfpx_scale_is_valid(blocks[i].e[1])) { + return false; + } + } + return true; +} + bool rocmfpx_validate_row_data_fp3(const void * data, size_t nbytes) { if (nbytes % sizeof(block_rocmfp3) != 0) { return false; diff --git a/ggml/rocmfpx/rocmfpx.h b/ggml/rocmfpx/rocmfpx.h index 9951a9d87..579d7e448 100644 --- a/ggml/rocmfpx/rocmfpx.h +++ b/ggml/rocmfpx/rocmfpx.h @@ -12,10 +12,14 @@ extern "C" { #define QK_ROCMFPX 32 +#define QK_ROCMFP2 QK_ROCMFPX #define QK_ROCMFP3 QK_ROCMFPX #define QK_ROCMFP6 QK_ROCMFPX #define QK_ROCMFP8 QK_ROCMFPX +#define QS_ROCMFP2 ((QK_ROCMFP2 * 2) / 8) +#define QR_ROCMFP2 1 +#define QI_ROCMFP2 (QK_ROCMFP2 / (4 * QR_ROCMFP2)) #define QS_ROCMFP3 ((QK_ROCMFP3 * 3) / 8) #define QS_ROCMFP6 ((QK_ROCMFP6 * 6) / 8) #define QS_ROCMFP8 QK_ROCMFP8 @@ -31,6 +35,11 @@ extern "C" { // AMD-native experimental family layouts. The GGUF types are registered, but // the layouts stay isolated from the promoted ROCmFP4 formats while evaluated. +typedef struct { + uint8_t qs[QS_ROCMFP2]; + uint8_t e[2]; +} block_rocmfp2; + typedef struct { uint8_t qs[QS_ROCMFP3]; uint8_t e[2]; @@ -47,10 +56,12 @@ typedef struct { } block_rocmfp8; #if defined(__cplusplus) +static_assert(sizeof(block_rocmfp2) == QS_ROCMFP2 + 2*sizeof(uint8_t), "wrong rocmfp2 block size/padding"); static_assert(sizeof(block_rocmfp3) == QS_ROCMFP3 + 2*sizeof(uint8_t), "wrong rocmfp3 block size/padding"); static_assert(sizeof(block_rocmfp6) == QS_ROCMFP6 + 2*sizeof(uint8_t), "wrong rocmfp6 block size/padding"); static_assert(sizeof(block_rocmfp8) == QS_ROCMFP8 + sizeof(uint8_t), "wrong rocmfp8 block size/padding"); #else +_Static_assert(sizeof(block_rocmfp2) == QS_ROCMFP2 + 2*sizeof(uint8_t), "wrong rocmfp2 block size/padding"); _Static_assert(sizeof(block_rocmfp3) == QS_ROCMFP3 + 2*sizeof(uint8_t), "wrong rocmfp3 block size/padding"); _Static_assert(sizeof(block_rocmfp6) == QS_ROCMFP6 + 2*sizeof(uint8_t), "wrong rocmfp6 block size/padding"); _Static_assert(sizeof(block_rocmfp8) == QS_ROCMFP8 + sizeof(uint8_t), "wrong rocmfp8 block size/padding"); @@ -58,10 +69,16 @@ _Static_assert(sizeof(block_rocmfp8) == QS_ROCMFP8 + sizeof(uint8_t), "wrong roc GGML_API float rocmfpx_ue4m3_to_fp32(uint8_t e); GGML_API bool rocmfpx_scale_is_valid(uint8_t e); +GGML_API size_t rocmfpx_row_size_fp2(int64_t k); GGML_API size_t rocmfpx_row_size_fp3(int64_t k); GGML_API size_t rocmfpx_row_size_fp6(int64_t k); GGML_API size_t rocmfpx_row_size_fp8(int64_t k); +GGML_API void rocmfpx_quantize_row_fp2_ref(const float * GGML_RESTRICT x, block_rocmfp2 * GGML_RESTRICT y, int64_t k); +GGML_API void rocmfpx_dequantize_row_fp2(const block_rocmfp2 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); +GGML_API void rocmfpx_quantize_row_fp2(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); +GGML_API size_t rocmfpx_quantize_fp2(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); + GGML_API void rocmfpx_quantize_row_fp3_ref(const float * GGML_RESTRICT x, block_rocmfp3 * GGML_RESTRICT y, int64_t k); GGML_API void rocmfpx_dequantize_row_fp3(const block_rocmfp3 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); GGML_API void rocmfpx_quantize_row_fp3(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); @@ -77,6 +94,7 @@ GGML_API void rocmfpx_dequantize_row_fp8(const block_rocmfp8 * GGML_RESTRICT x GGML_API void rocmfpx_quantize_row_fp8(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); GGML_API size_t rocmfpx_quantize_fp8(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); +GGML_API bool rocmfpx_validate_row_data_fp2(const void * data, size_t nbytes); GGML_API bool rocmfpx_validate_row_data_fp3(const void * data, size_t nbytes); GGML_API bool rocmfpx_validate_row_data_fp6(const void * data, size_t nbytes); GGML_API bool rocmfpx_validate_row_data_fp8(const void * data, size_t nbytes); diff --git a/ggml/rocmfpx/test_rocmfp2_reference.c b/ggml/rocmfpx/test_rocmfp2_reference.c new file mode 100644 index 000000000..96201fc1e --- /dev/null +++ b/ggml/rocmfpx/test_rocmfp2_reference.c @@ -0,0 +1,565 @@ +#include "rocmfp2_reference.h" + +#include +#include +#include +#include + +static unsigned long long test_checks = 0; + +#define CHECK(condition) do { \ + ++test_checks; \ + if (!(condition)) { \ + fprintf(stderr, "FAIL %s:%d: %s\n", __FILE__, __LINE__, #condition); \ + return false; \ + } \ +} while (0) + +#define CHECK_STATUS(expression, expected) do { \ + const rocmfp2_p1_status actual_status_ = (expression); \ + ++test_checks; \ + if (actual_status_ != (expected)) { \ + fprintf(stderr, "FAIL %s:%d: %s returned %s, expected %s\n", \ + __FILE__, __LINE__, #expression, rocmfp2_p1_status_name(actual_status_), \ + rocmfp2_p1_status_name(expected)); \ + return false; \ + } \ +} while (0) + +static int8_t expected_code_value( + uint8_t code, + rocmfp2_p1_codebook codebook, + rocmfp2_p1_mapping mapping) { + if (mapping == ROCMFP2_P1_MAPPING_MORD) { + const int values[4] = { + -(int) codebook.outer, + -(int) codebook.inner, + (int) codebook.inner, + (int) codebook.outer, + }; + return (int8_t) values[code]; + } + + const int values[4] = { + (int) codebook.inner, + (int) codebook.outer, + -(int) codebook.inner, + -(int) codebook.outer, + }; + return (int8_t) values[code]; +} + +static uint8_t expected_semantic_code(bool negative, bool outer, rocmfp2_p1_mapping mapping) { + if (mapping == ROCMFP2_P1_MAPPING_MORD) { + if (negative) { + return outer ? 0u : 1u; + } + return outer ? 3u : 2u; + } + + if (negative) { + return outer ? 3u : 2u; + } + return outer ? 1u : 0u; +} + +static bool test_layout_and_p0_packing(void) { + CHECK(sizeof(rocmfp2_p1_block) == 10); + CHECK(offsetof(rocmfp2_p1_block, d) == 0); + CHECK(offsetof(rocmfp2_p1_block, s) == 8); + + uint8_t codes[ROCMFP2_P1_BLOCK_WEIGHTS]; + uint8_t unpacked[ROCMFP2_P1_BLOCK_WEIGHTS]; + const uint8_t scales[2] = { 0x12, 0x7e }; + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + codes[i] = (uint8_t) (i & 3); + } + + rocmfp2_p1_block block; + memset(&block, 0xa5, sizeof(block)); + CHECK_STATUS(rocmfp2_p1_pack_codes(codes, scales, &block), ROCMFP2_P1_OK); + + const uint8_t expected_bytes[ROCMFP2_P1_BLOCK_BYTES] = { + 0xe4, 0xe4, 0xe4, 0xe4, 0xe4, 0xe4, 0xe4, 0xe4, 0x12, 0x7e, + }; + CHECK(memcmp(&block, expected_bytes, sizeof(expected_bytes)) == 0); + + rocmfp2_p1_unpack_codes(&block, unpacked); + CHECK(memcmp(codes, unpacked, sizeof(codes)) == 0); + + rocmfp2_p1_block unchanged; + rocmfp2_p1_block before; + memset(&unchanged, 0x5a, sizeof(unchanged)); + memcpy(&before, &unchanged, sizeof(before)); + codes[31] = 4; + CHECK_STATUS(rocmfp2_p1_pack_codes(codes, scales, &unchanged), ROCMFP2_P1_INVALID_CODE); + CHECK(memcmp(&unchanged, &before, sizeof(before)) == 0); + codes[31] = 3; + + const uint8_t invalid_scales[2] = { 0x7f, 0x00 }; + CHECK_STATUS(rocmfp2_p1_pack_codes(codes, invalid_scales, &unchanged), + ROCMFP2_P1_INVALID_SCALE_METADATA); + CHECK(memcmp(&unchanged, &before, sizeof(before)) == 0); + + puts("PASS layout_p0: block=10 data=8 scales=2 bpw=2.50 golden=e4e4e4e4e4e4e4e4127e"); + return true; +} + +static bool test_all_packed_bytes(void) { + const rocmfp2_p1_codebook codebook = { 3, 10 }; + + for (int mapping_index = 0; mapping_index < 2; ++mapping_index) { + const rocmfp2_p1_mapping mapping = (rocmfp2_p1_mapping) mapping_index; + for (int packed_value = 0; packed_value <= 0xff; ++packed_value) { + int8_t decoded[4]; + CHECK_STATUS(rocmfp2_p1_decode_packed_byte( + (uint8_t) packed_value, &codebook, mapping, decoded), ROCMFP2_P1_OK); + + for (int lane = 0; lane < 4; ++lane) { + const uint8_t code = (uint8_t) ((packed_value >> (2 * lane)) & 3); + CHECK(decoded[lane] == expected_code_value(code, codebook, mapping)); + + int8_t scalar = 0; + CHECK_STATUS(rocmfp2_p1_decode_code(code, &codebook, mapping, &scalar), ROCMFP2_P1_OK); + CHECK(scalar == decoded[lane]); + } + } + } + + int8_t value = 0; + CHECK_STATUS(rocmfp2_p1_decode_code(4, &codebook, ROCMFP2_P1_MAPPING_MORD, &value), + ROCMFP2_P1_INVALID_CODE); + + puts("PASS packed_byte_decode: mappings=2 bytes=256 lanes=2048"); + return true; +} + +static double expected_scale(uint8_t scale_byte) { + const int exponent = scale_byte >> 3; + const int mantissa = scale_byte & 7; + return exponent == 0 ? scalbn((double) mantissa, -10) : scalbn((double) (8 + mantissa), exponent - 11); +} + +static bool test_scale_encoding_and_boundaries(void) { + double previous = -1.0; + for (int scale_index = 0; scale_index <= 0xff; ++scale_index) { + const uint8_t scale_byte = (uint8_t) scale_index; + const bool valid = scale_index <= 0x7e; + CHECK(rocmfp2_p1_scale_is_valid(scale_byte) == valid); + + const double decoded = rocmfp2_p1_ue4m3_to_binary64(scale_byte); + if (valid) { + CHECK(decoded == expected_scale(scale_byte)); + CHECK(decoded > previous); + CHECK(rocmfp2_p1_nearest_ue4m3(decoded) == scale_byte); + previous = decoded; + } else { + CHECK(isnan(decoded)); + } + } + + CHECK(rocmfp2_p1_ue4m3_to_binary64(0x00) == 0.0); + CHECK(rocmfp2_p1_ue4m3_to_binary64(0x01) == 0x1p-10); + CHECK(rocmfp2_p1_ue4m3_to_binary64(0x07) == 7.0 * 0x1p-10); + CHECK(rocmfp2_p1_ue4m3_to_binary64(0x08) == 8.0 * 0x1p-10); + CHECK(rocmfp2_p1_ue4m3_to_binary64(0x7e) == 224.0); + + for (int upper_index = 1; upper_index <= 0x7e; ++upper_index) { + const uint8_t upper_byte = (uint8_t) upper_index; + const uint8_t lower_byte = (uint8_t) (upper_index - 1); + const double lower = expected_scale(lower_byte); + const double upper = expected_scale(upper_byte); + const double midpoint = (lower + upper) * 0.5; + + CHECK(rocmfp2_p1_nearest_ue4m3(midpoint) == lower_byte); + CHECK(rocmfp2_p1_nearest_ue4m3(nextafter(midpoint, -INFINITY)) == lower_byte); + CHECK(rocmfp2_p1_nearest_ue4m3(nextafter(midpoint, INFINITY)) == upper_byte); + } + + CHECK(rocmfp2_p1_nearest_ue4m3(-1.0) == 0xff); + CHECK(rocmfp2_p1_nearest_ue4m3(NAN) == 0xff); + CHECK(rocmfp2_p1_nearest_ue4m3(nextafter(224.0, INFINITY)) == 0x7e); + CHECK(rocmfp2_p1_nearest_ue4m3(INFINITY) == 0x7e); + + puts("PASS ue4m3_boundaries: legal=127 invalid=129 adjacent_midpoints=126 max=224"); + return true; +} + +static bool test_metadata_validation(void) { + rocmfp2_p1_block block; + memset(&block, 0, sizeof(block)); + + for (int first = 0; first <= 0x7e; ++first) { + for (int second = 0; second <= 0x7e; ++second) { + block.s[0] = (uint8_t) first; + block.s[1] = (uint8_t) second; + CHECK(rocmfp2_p1_validate_block(&block)); + } + } + + for (int invalid = 0x7f; invalid <= 0xff; ++invalid) { + block.s[0] = (uint8_t) invalid; + block.s[1] = 0; + CHECK(!rocmfp2_p1_validate_block(&block)); + block.s[0] = 0; + block.s[1] = (uint8_t) invalid; + CHECK(!rocmfp2_p1_validate_block(&block)); + } + + uint8_t serialized[2 * ROCMFP2_P1_BLOCK_BYTES]; + memset(serialized, 0, sizeof(serialized)); + serialized[8] = 0x7e; + serialized[9] = 0x00; + serialized[18] = 0x01; + serialized[19] = 0x7e; + CHECK(rocmfp2_p1_validate_serialized(serialized, sizeof(serialized))); + CHECK(rocmfp2_p1_validate_serialized(NULL, 0)); + CHECK(!rocmfp2_p1_validate_serialized(NULL, sizeof(rocmfp2_p1_block))); + + for (size_t length = 1; length < sizeof(serialized); ++length) { + if (length != ROCMFP2_P1_BLOCK_BYTES) { + CHECK(!rocmfp2_p1_validate_serialized(serialized, length)); + } + } + + for (int invalid = 0x7f; invalid <= 0xff; ++invalid) { + serialized[8] = (uint8_t) invalid; + CHECK(!rocmfp2_p1_validate_serialized(serialized, sizeof(serialized))); + serialized[8] = 0x00; + serialized[19] = (uint8_t) invalid; + CHECK(!rocmfp2_p1_validate_serialized(serialized, sizeof(serialized))); + serialized[19] = 0x7e; + } + + const rocmfp2_p1_codebook codebook = { 3, 10 }; + double output[ROCMFP2_P1_BLOCK_WEIGHTS]; + double before[ROCMFP2_P1_BLOCK_WEIGHTS]; + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + output[i] = 1234.0 + (double) i; + } + memcpy(before, output, sizeof(before)); + block.s[0] = 0x7f; + block.s[1] = 0; + CHECK_STATUS(rocmfp2_p1_dequantize_block( + &block, &codebook, ROCMFP2_P1_MAPPING_MORD, output), ROCMFP2_P1_INVALID_SCALE_METADATA); + CHECK(memcmp(output, before, sizeof(output)) == 0); + + puts("PASS metadata_validation: legal_pairs=16129 invalid_bytes=129 malformed_lengths=18"); + return true; +} + +static bool test_zero_midpoint_and_scale_ties(void) { + const rocmfp2_p1_codebook codebook = { 3, 10 }; + const uint8_t scale_byte = 0x31; + const double scale = expected_scale(scale_byte); + const double midpoint = ((double) codebook.inner + (double) codebook.outer) * scale * 0.5; + + for (int mapping_index = 0; mapping_index < 2; ++mapping_index) { + const rocmfp2_p1_mapping mapping = (rocmfp2_p1_mapping) mapping_index; + const double probes[] = { + 0.0, + -0.0, + midpoint, + -midpoint, + nextafter(midpoint, INFINITY), + nextafter(-midpoint, -INFINITY), + }; + const uint8_t expected[] = { + expected_semantic_code(false, false, mapping), + expected_semantic_code(true, false, mapping), + expected_semantic_code(false, false, mapping), + expected_semantic_code(true, false, mapping), + expected_semantic_code(false, true, mapping), + expected_semantic_code(true, true, mapping), + }; + + for (size_t i = 0; i < sizeof(probes) / sizeof(probes[0]); ++i) { + uint8_t reference_code = 0xff; + uint8_t optimized_code = 0xff; + CHECK_STATUS(rocmfp2_p1_select_code_ref( + probes[i], scale_byte, &codebook, mapping, &reference_code), ROCMFP2_P1_OK); + CHECK_STATUS(rocmfp2_p1_select_code_optimized( + probes[i], scale_byte, &codebook, mapping, &optimized_code), ROCMFP2_P1_OK); + CHECK(reference_code == expected[i]); + CHECK(optimized_code == expected[i]); + } + + uint8_t zero_scale_positive = 0xff; + uint8_t zero_scale_negative = 0xff; + CHECK_STATUS(rocmfp2_p1_select_code_ref( + 1.0, 0, &codebook, mapping, &zero_scale_positive), ROCMFP2_P1_OK); + CHECK_STATUS(rocmfp2_p1_select_code_optimized( + -1.0, 0, &codebook, mapping, &zero_scale_negative), ROCMFP2_P1_OK); + CHECK(zero_scale_positive == 0); + CHECK(zero_scale_negative == 0); + } + + const rocmfp2_p1_codebook tie_codebook = { 1, 2 }; + double source[ROCMFP2_P1_BLOCK_WEIGHTS]; + const double group0_value = 2.0 * expected_scale(0x20); + const double group1_value = 2.0 * expected_scale(0x21); + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + const double magnitude = i < ROCMFP2_P1_GROUP_WEIGHTS ? group0_value : group1_value; + source[i] = (i & 1) ? -magnitude : magnitude; + } + + for (int mapping_index = 0; mapping_index < 2; ++mapping_index) { + const rocmfp2_p1_mapping mapping = (rocmfp2_p1_mapping) mapping_index; + rocmfp2_p1_block reference; + rocmfp2_p1_block optimized; + double reference_sse[2]; + double optimized_sse[2]; + CHECK_STATUS(rocmfp2_p1_quantize_block_ref( + source, &tie_codebook, mapping, &reference, reference_sse), ROCMFP2_P1_OK); + CHECK_STATUS(rocmfp2_p1_quantize_block_optimized( + source, &tie_codebook, mapping, &optimized, optimized_sse), ROCMFP2_P1_OK); + CHECK(memcmp(&reference, &optimized, sizeof(reference)) == 0); + CHECK(reference.s[0] == 0x20); + CHECK(reference.s[1] == 0x21); + CHECK(reference_sse[0] == 0.0 && reference_sse[1] == 0.0); + CHECK(optimized_sse[0] == 0.0 && optimized_sse[1] == 0.0); + } + + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + source[i] = (i & 1) ? -0.0 : 0.0; + } + for (int mapping_index = 0; mapping_index < 2; ++mapping_index) { + const rocmfp2_p1_mapping mapping = (rocmfp2_p1_mapping) mapping_index; + rocmfp2_p1_block reference; + rocmfp2_p1_block optimized; + uint8_t codes[ROCMFP2_P1_BLOCK_WEIGHTS]; + CHECK_STATUS(rocmfp2_p1_quantize_block_ref( + source, &codebook, mapping, &reference, NULL), ROCMFP2_P1_OK); + CHECK_STATUS(rocmfp2_p1_quantize_block_optimized( + source, &codebook, mapping, &optimized, NULL), ROCMFP2_P1_OK); + CHECK(memcmp(&reference, &optimized, sizeof(reference)) == 0); + CHECK(reference.s[0] == 0 && reference.s[1] == 0); + rocmfp2_p1_unpack_codes(&reference, codes); + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + CHECK(codes[i] == 0); + } + const uint8_t canonical_zero[ROCMFP2_P1_BLOCK_BYTES] = { 0 }; + CHECK(memcmp(&reference, canonical_zero, sizeof(canonical_zero)) == 0); + } + + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + source[i] = (i & 1) ? -10000.0 : 10000.0; + } + rocmfp2_p1_block saturated; + CHECK_STATUS(rocmfp2_p1_quantize_block_ref( + source, &codebook, ROCMFP2_P1_MAPPING_MORD, &saturated, NULL), ROCMFP2_P1_OK); + CHECK(saturated.s[0] == 0x7e && saturated.s[1] == 0x7e); + + puts("PASS deterministic_ties: signed_zero=2 magnitude_midpoint=4 scale_tuple=2 saturation=0x7e"); + return true; +} + +static bool test_python_golden_vectors(void) { + const rocmfp2_p1_codebook codebook = { 3, 10 }; + double source[ROCMFP2_P1_BLOCK_WEIGHTS]; + static const double values[4] = { -10.0, -3.0, 3.0, 10.0 }; + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + source[i] = values[i & 3] * (i < ROCMFP2_P1_GROUP_WEIGHTS ? 1.0 : 2.0); + } + + static const uint8_t expected[2][ROCMFP2_P1_BLOCK_BYTES] = { + { 0xe4, 0xe4, 0xe4, 0xe4, 0xe4, 0xe4, 0xe4, 0xe4, 0x40, 0x48 }, + { 0x4b, 0x4b, 0x4b, 0x4b, 0x4b, 0x4b, 0x4b, 0x4b, 0x40, 0x48 }, + }; + + for (int mapping_index = 0; mapping_index < 2; ++mapping_index) { + const rocmfp2_p1_mapping mapping = (rocmfp2_p1_mapping) mapping_index; + rocmfp2_p1_block reference; + rocmfp2_p1_block optimized; + double output[ROCMFP2_P1_BLOCK_WEIGHTS]; + CHECK_STATUS(rocmfp2_p1_quantize_block_ref( + source, &codebook, mapping, &reference, NULL), ROCMFP2_P1_OK); + CHECK_STATUS(rocmfp2_p1_quantize_block_optimized( + source, &codebook, mapping, &optimized, NULL), ROCMFP2_P1_OK); + CHECK(memcmp(&reference, expected[mapping_index], sizeof(reference)) == 0); + CHECK(memcmp(&optimized, expected[mapping_index], sizeof(optimized)) == 0); + CHECK_STATUS(rocmfp2_p1_dequantize_block( + &reference, &codebook, mapping, output), ROCMFP2_P1_OK); + CHECK(memcmp(output, source, sizeof(source)) == 0); + } + + puts("PASS python_golden_vectors: vectors=2 mord=e4x8+4048 msm=4bx8+4048"); + return true; +} + +static uint64_t splitmix64_next(uint64_t * state) { + uint64_t value = (*state += UINT64_C(0x9e3779b97f4a7c15)); + value = (value ^ (value >> 30)) * UINT64_C(0xbf58476d1ce4e5b9); + value = (value ^ (value >> 27)) * UINT64_C(0x94d049bb133111eb); + return value ^ (value >> 31); +} + +static uint64_t fnv1a64(uint64_t hash, const void * data, size_t size) { + const uint8_t * bytes = (const uint8_t *) data; + for (size_t i = 0; i < size; ++i) { + hash ^= bytes[i]; + hash *= UINT64_C(1099511628211); + } + return hash; +} + +static void fill_fixed_random_block(double source[ROCMFP2_P1_BLOCK_WEIGHTS], uint64_t * state, int block_index) { + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + const uint64_t bits = splitmix64_next(state); + const double unit = (double) (bits >> 11) * 0x1p-53; + const int exponent = (int) ((bits >> 3) % 18) - 10; + source[i] = scalbn(2.0 * unit - 1.0, exponent); + } + + source[(block_index * 7) & 31] = (block_index & 1) ? -0.0 : 0.0; + source[(block_index * 11 + 3) & 31] *= 13.0; +} + +static bool test_fixed_random_reference_optimized_match(void) { + static const rocmfp2_p1_codebook codebooks[8] = { + { 1, 2 }, { 2, 5 }, { 1, 3 }, { 3, 10 }, + { 2, 7 }, { 1, 4 }, { 3, 13 }, { 1, 5 }, + }; + enum { BLOCKS_PER_CODEBOOK = 24 }; + + uint64_t random_state = UINT64_C(0x726f636d66703231); + uint64_t fingerprint = UINT64_C(1469598103934665603); + + for (int codebook_index = 0; codebook_index < 8; ++codebook_index) { + const rocmfp2_p1_codebook * codebook = &codebooks[codebook_index]; + for (int block_index = 0; block_index < BLOCKS_PER_CODEBOOK; ++block_index) { + double source[ROCMFP2_P1_BLOCK_WEIGHTS]; + rocmfp2_p1_block mapping_blocks[2]; + double mapping_outputs[2][ROCMFP2_P1_BLOCK_WEIGHTS]; + double mapping_sse[2][2]; + fill_fixed_random_block(source, &random_state, block_index); + + for (int mapping_index = 0; mapping_index < 2; ++mapping_index) { + const rocmfp2_p1_mapping mapping = (rocmfp2_p1_mapping) mapping_index; + rocmfp2_p1_block reference; + rocmfp2_p1_block optimized; + rocmfp2_p1_block reference_repeat; + rocmfp2_p1_block optimized_repeat; + double reference_sse[2]; + double optimized_sse[2]; + double repeat_sse[2]; + + CHECK_STATUS(rocmfp2_p1_quantize_block_ref( + source, codebook, mapping, &reference, reference_sse), ROCMFP2_P1_OK); + CHECK_STATUS(rocmfp2_p1_quantize_block_optimized( + source, codebook, mapping, &optimized, optimized_sse), ROCMFP2_P1_OK); + CHECK(memcmp(&reference, &optimized, sizeof(reference)) == 0); + CHECK(reference_sse[0] == optimized_sse[0]); + CHECK(reference_sse[1] == optimized_sse[1]); + + CHECK_STATUS(rocmfp2_p1_quantize_block_ref( + source, codebook, mapping, &reference_repeat, repeat_sse), ROCMFP2_P1_OK); + CHECK(memcmp(&reference, &reference_repeat, sizeof(reference)) == 0); + CHECK(reference_sse[0] == repeat_sse[0] && reference_sse[1] == repeat_sse[1]); + + CHECK_STATUS(rocmfp2_p1_quantize_block_optimized( + source, codebook, mapping, &optimized_repeat, repeat_sse), ROCMFP2_P1_OK); + CHECK(memcmp(&optimized, &optimized_repeat, sizeof(optimized)) == 0); + CHECK(optimized_sse[0] == repeat_sse[0] && optimized_sse[1] == repeat_sse[1]); + CHECK(rocmfp2_p1_validate_block(&reference)); + + CHECK_STATUS(rocmfp2_p1_dequantize_block( + &reference, codebook, mapping, mapping_outputs[mapping_index]), ROCMFP2_P1_OK); + for (int group = 0; group < 2; ++group) { + double recomputed_sse = 0.0; + for (int i = 0; i < ROCMFP2_P1_GROUP_WEIGHTS; ++i) { + const int index = group * ROCMFP2_P1_GROUP_WEIGHTS + i; + const double delta = source[index] - mapping_outputs[mapping_index][index]; + recomputed_sse += delta * delta; + } + CHECK(recomputed_sse == reference_sse[group]); + mapping_sse[mapping_index][group] = reference_sse[group]; + } + + memcpy(&mapping_blocks[mapping_index], &reference, sizeof(reference)); + fingerprint = fnv1a64(fingerprint, &reference, sizeof(reference)); + } + + CHECK(mapping_blocks[0].s[0] == mapping_blocks[1].s[0]); + CHECK(mapping_blocks[0].s[1] == mapping_blocks[1].s[1]); + CHECK(mapping_sse[0][0] == mapping_sse[1][0]); + CHECK(mapping_sse[0][1] == mapping_sse[1][1]); + CHECK(memcmp(mapping_outputs[0], mapping_outputs[1], sizeof(mapping_outputs[0])) == 0); + } + } + + printf("PASS fixed_random: codebooks=8 blocks_each=%d mappings=2 ref_opt_equal=384 fingerprint=%016" PRIx64 "\n", + BLOCKS_PER_CODEBOOK, fingerprint); + return true; +} + +static bool test_nonfinite_rejection_and_output_atomicity(void) { + const rocmfp2_p1_codebook codebook = { 3, 10 }; + double source[ROCMFP2_P1_BLOCK_WEIGHTS]; + for (int i = 0; i < ROCMFP2_P1_BLOCK_WEIGHTS; ++i) { + source[i] = (double) i * 0.125 - 2.0; + } + + const double nonfinite_values[3] = { NAN, INFINITY, -INFINITY }; + for (int mapping_index = 0; mapping_index < 2; ++mapping_index) { + const rocmfp2_p1_mapping mapping = (rocmfp2_p1_mapping) mapping_index; + for (int position = 0; position < ROCMFP2_P1_BLOCK_WEIGHTS; ++position) { + const double saved = source[position]; + for (int variant = 0; variant < 3; ++variant) { + rocmfp2_p1_block reference; + rocmfp2_p1_block optimized; + rocmfp2_p1_block before_reference; + rocmfp2_p1_block before_optimized; + double reference_sse[2] = { 123.0, 456.0 }; + double optimized_sse[2] = { 789.0, 987.0 }; + const double reference_sse_before[2] = { 123.0, 456.0 }; + const double optimized_sse_before[2] = { 789.0, 987.0 }; + + source[position] = nonfinite_values[variant]; + memset(&reference, 0x3c, sizeof(reference)); + memset(&optimized, 0xc3, sizeof(optimized)); + memcpy(&before_reference, &reference, sizeof(reference)); + memcpy(&before_optimized, &optimized, sizeof(optimized)); + + CHECK_STATUS(rocmfp2_p1_quantize_block_ref( + source, &codebook, mapping, &reference, reference_sse), + ROCMFP2_P1_NONFINITE_SOURCE); + CHECK_STATUS(rocmfp2_p1_quantize_block_optimized( + source, &codebook, mapping, &optimized, optimized_sse), + ROCMFP2_P1_NONFINITE_SOURCE); + CHECK(memcmp(&reference, &before_reference, sizeof(reference)) == 0); + CHECK(memcmp(&optimized, &before_optimized, sizeof(optimized)) == 0); + CHECK(memcmp(reference_sse, reference_sse_before, sizeof(reference_sse)) == 0); + CHECK(memcmp(optimized_sse, optimized_sse_before, sizeof(optimized_sse)) == 0); + } + source[position] = saved; + } + } + + uint8_t code = 0; + CHECK_STATUS(rocmfp2_p1_select_code_ref( + NAN, 0x10, &codebook, ROCMFP2_P1_MAPPING_MORD, &code), ROCMFP2_P1_NONFINITE_SOURCE); + CHECK_STATUS(rocmfp2_p1_select_code_optimized( + INFINITY, 0x10, &codebook, ROCMFP2_P1_MAPPING_MSM, &code), ROCMFP2_P1_NONFINITE_SOURCE); + + puts("PASS nonfinite_rejection: positions=32 values=3 mappings=2 paths=2 status=NONFINITE_SOURCE"); + return true; +} + +int main(void) { + if (!test_layout_and_p0_packing() || + !test_all_packed_bytes() || + !test_scale_encoding_and_boundaries() || + !test_metadata_validation() || + !test_zero_midpoint_and_scale_ties() || + !test_python_golden_vectors() || + !test_fixed_random_reference_optimized_match() || + !test_nonfinite_rejection_and_output_atomicity()) { + fprintf(stderr, "ROCmFP2 Phase-1 reference tests FAILED after %llu checks\n", test_checks); + return 1; + } + + printf("PASS rocmfp2_phase1_reference: checks=%llu arithmetic=IEEE754-binary64 search=127-scales exhaustive\n", + test_checks); + return 0; +} diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index f3c5bd6bb..45396a747 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -226,6 +226,42 @@ static inline int ggml_rocmfpx_decode_fp3_cpu(uint32_t code) { return table[code & 7u]; } +static inline int ggml_rocmfpx_decode_fp2_cpu(uint32_t code) { + static const int8_t table[4] = { -4, -1, 1, 4 }; + return table[code & 3u]; +} + +static void ggml_vec_dot_rocmfpx_fp2_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) { + GGML_UNUSED(bs); + GGML_UNUSED(bx); + GGML_UNUSED(by); + assert(nrc == 1); + GGML_UNUSED(nrc); + assert(n % QK_ROCMFP2 == 0); + assert(QK_ROCMFP2 == QK8_0); + + const block_rocmfp2 * GGML_RESTRICT x = (const block_rocmfp2 *) vx; + const block_q8_0 * GGML_RESTRICT y = (const block_q8_0 *) vy; + const int nb = n / QK_ROCMFP2; + float sumf = 0.0f; + + for (int ib = 0; ib < nb; ++ib) { + const float dy = GGML_CPU_FP16_TO_FP32(y[ib].d); + int sumi0 = 0; + int sumi1 = 0; + for (int j = 0; j < QK_ROCMFP2/2; ++j) { + const uint8_t c0 = (x[ib].qs[j/4] >> (2*(j % 4))) & 3u; + const uint8_t c1 = (x[ib].qs[4 + j/4] >> (2*(j % 4))) & 3u; + sumi0 += ggml_rocmfpx_decode_fp2_cpu(c0) * (int) y[ib].qs[j]; + sumi1 += ggml_rocmfpx_decode_fp2_cpu(c1) * (int) y[ib].qs[j + QK_ROCMFP2/2]; + } + sumf += dy * ( + rocmfpx_ue4m3_to_fp32(x[ib].e[0]) * (float) sumi0 + + rocmfpx_ue4m3_to_fp32(x[ib].e[1]) * (float) sumi1); + } + *s = sumf; +} + static inline int ggml_rocmfpx_decode_fp6_cpu(uint32_t code) { const int mag = (int) (code & 31u); return (code & 32u) ? -mag : mag; @@ -379,6 +415,12 @@ static const struct ggml_type_traits_cpu type_traits_cpu[GGML_TYPE_COUNT] = { .vec_dot_type = GGML_TYPE_Q8_0, .nrows = 1, }, + [GGML_TYPE_Q2_0_ROCMFPX] = { + .from_float = rocmfpx_quantize_row_fp2, + .vec_dot = ggml_vec_dot_rocmfpx_fp2_q8_0, + .vec_dot_type = GGML_TYPE_Q8_0, + .nrows = 1, + }, [GGML_TYPE_Q6_0_ROCMFPX] = { .from_float = rocmfpx_quantize_row_fp6, .vec_dot = ggml_vec_dot_rocmfpx_fp6_q8_0, diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index b65d62034..45efba895 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -1256,6 +1256,7 @@ void ggml_compute_forward_acc( case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: case GGML_TYPE_Q2_K: @@ -4917,7 +4918,7 @@ static void ggml_compute_forward_get_rows_f32( void ggml_compute_forward_get_rows( const ggml_compute_params * params, - ggml_tensor * dst) { + ggml_tensor * dst) { const ggml_tensor * src0 = dst->src[0]; @@ -4934,6 +4935,7 @@ void ggml_compute_forward_get_rows( case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: case GGML_TYPE_Q2_K: diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index 43e3ba2f4..59c38c8f7 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -1052,6 +1052,13 @@ struct ggml_cuda_type_traits { static constexpr int qi = QI_ROCMFP3; }; +template<> +struct ggml_cuda_type_traits { + static constexpr int qk = QK_ROCMFP2; + static constexpr int qr = QR_ROCMFP2; + static constexpr int qi = QI_ROCMFP2; +}; + template<> struct ggml_cuda_type_traits { static constexpr int qk = QK_ROCMFP6; diff --git a/ggml/src/ggml-cuda/convert.cu b/ggml/src/ggml-cuda/convert.cu index d59c5b808..7f9b70563 100644 --- a/ggml/src/ggml-cuda/convert.cu +++ b/ggml/src/ggml-cuda/convert.cu @@ -944,6 +944,8 @@ to_fp16_cuda_t ggml_get_to_fp16_cuda(ggml_type type) { return dequantize_row_rocmfp4_hip; case GGML_TYPE_Q4_0_ROCMFP4_FAST: return dequantize_row_rocmfp4_fast_hip; + case GGML_TYPE_Q2_0_ROCMFPX: + return dequantize_block_cont_cuda; case GGML_TYPE_Q3_0_ROCMFPX: return dequantize_block_cont_cuda; case GGML_TYPE_Q6_0_ROCMFPX: @@ -1013,6 +1015,8 @@ to_fp32_cuda_t ggml_get_to_fp32_cuda(ggml_type type) { return dequantize_row_rocmfp4_hip; case GGML_TYPE_Q4_0_ROCMFP4_FAST: return dequantize_row_rocmfp4_fast_hip; + case GGML_TYPE_Q2_0_ROCMFPX: + return dequantize_block_cont_cuda; case GGML_TYPE_Q3_0_ROCMFPX: return dequantize_block_cont_cuda; case GGML_TYPE_Q6_0_ROCMFPX: @@ -1054,6 +1058,8 @@ to_fp16_nc_cuda_t ggml_get_to_fp16_nc_cuda(ggml_type type) { return dequantize_block_cuda; case GGML_TYPE_Q4_0_ROCMFP4_FAST: return dequantize_block_cuda; + case GGML_TYPE_Q2_0_ROCMFPX: + return dequantize_block_cuda; case GGML_TYPE_Q3_0_ROCMFPX: return dequantize_block_cuda; case GGML_TYPE_Q6_0_ROCMFPX: @@ -1091,6 +1097,8 @@ to_bf16_nc_cuda_t ggml_get_to_bf16_nc_cuda(ggml_type type) { return dequantize_block_cuda; case GGML_TYPE_Q4_0_ROCMFP4_FAST: return dequantize_block_cuda; + case GGML_TYPE_Q2_0_ROCMFPX: + return dequantize_block_cuda; case GGML_TYPE_Q3_0_ROCMFPX: return dequantize_block_cuda; case GGML_TYPE_Q6_0_ROCMFPX: @@ -1128,6 +1136,8 @@ to_fp32_nc_cuda_t ggml_get_to_fp32_nc_cuda(ggml_type type) { return dequantize_block_cuda; case GGML_TYPE_Q4_0_ROCMFP4_FAST: return dequantize_block_cuda; + case GGML_TYPE_Q2_0_ROCMFPX: + return dequantize_block_cuda; case GGML_TYPE_Q3_0_ROCMFPX: return dequantize_block_cuda; case GGML_TYPE_Q6_0_ROCMFPX: diff --git a/ggml/src/ggml-cuda/dequantize.cuh b/ggml/src/ggml-cuda/dequantize.cuh index 29b09a06d..1cd06c665 100644 --- a/ggml/src/ggml-cuda/dequantize.cuh +++ b/ggml/src/ggml-cuda/dequantize.cuh @@ -93,6 +93,10 @@ static __device__ __forceinline__ uint32_t rocmfpx_get_fp3_code_cuda(const uint8 return (rocmfpx_load_qs_window_cuda(src, byte_pos) >> shift) & 7u; } +static __device__ __forceinline__ uint32_t rocmfpx_get_fp2_code_cuda(const uint8_t * src, const int i) { + return (src[i >> 2] >> (2 * (i & 3))) & 3u; +} + static __device__ __forceinline__ uint32_t rocmfpx_get_fp6_code_cuda(const uint8_t * src, const int i) { const int bit_pos = i * 6; const int byte_pos = bit_pos >> 3; @@ -106,6 +110,10 @@ static __device__ __forceinline__ int rocmfpx_decode_fp3_code_cuda(const uint32_ return (code & 4u) ? -mag : mag; } +static __device__ __forceinline__ int rocmfpx_decode_fp2_code_cuda(const uint32_t code) { + return code == 0u ? -4 : code == 1u ? -1 : code == 2u ? 1 : 4; +} + static __device__ __forceinline__ int rocmfpx_decode_fp6_code_cuda(const uint32_t code) { const int mag = (int) (code & 31u); return (code & 32u) ? -mag : mag; @@ -123,6 +131,16 @@ static __device__ __forceinline__ void dequantize_rocmfpx_fp3(const void * vx, c v.y = d1 * (float) rocmfpx_decode_fp3_code_cuda(rocmfpx_get_fp3_code_cuda(x[ib].qs, i1)); } +static __device__ __forceinline__ void dequantize_rocmfpx_fp2(const void * vx, const int64_t ib, const int iqs, float2 & v) { + const block_rocmfp2 * x = (const block_rocmfp2 *) vx; + const int i0 = iqs + 0; + const int i1 = iqs + 1; + const float d0 = rocmfpx_ue4m3_to_fp32_finite(x[ib].e[i0 >= QK_ROCMFP2/2]); + const float d1 = rocmfpx_ue4m3_to_fp32_finite(x[ib].e[i1 >= QK_ROCMFP2/2]); + v.x = d0 * (float) rocmfpx_decode_fp2_code_cuda(rocmfpx_get_fp2_code_cuda(x[ib].qs, i0)); + v.y = d1 * (float) rocmfpx_decode_fp2_code_cuda(rocmfpx_get_fp2_code_cuda(x[ib].qs, i1)); +} + static __device__ __forceinline__ void dequantize_rocmfpx_fp6(const void * vx, const int64_t ib, const int iqs, float2 & v) { const block_rocmfp6_device * x = (const block_rocmfp6_device *) vx; diff --git a/ggml/src/ggml-cuda/getrows.cu b/ggml/src/ggml-cuda/getrows.cu index 8cdf11002..7fd6ecec4 100644 --- a/ggml/src/ggml-cuda/getrows.cu +++ b/ggml/src/ggml-cuda/getrows.cu @@ -225,6 +225,10 @@ static void ggml_cuda_get_rows_switch_src0_type( get_rows_cuda_q(src0_d, src1_d, dst_d, ne00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb1, nb2, nb3, stream); break; + case GGML_TYPE_Q2_0_ROCMFPX: + get_rows_cuda_q(src0_d, src1_d, dst_d, + ne00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb1, nb2, nb3, stream); + break; case GGML_TYPE_Q6_0_ROCMFPX: get_rows_cuda_q(src0_d, src1_d, dst_d, ne00, nb01, nb02, nb03, ne10, ne11, ne12, nb10, nb11, nb12, nb1, nb2, nb3, stream); diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 9011fe937..a5aa24ddd 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -5549,6 +5549,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: case GGML_TYPE_NVFP4: @@ -5591,6 +5592,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: return true; @@ -5608,7 +5610,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g op->type == GGML_TYPE_Q4_0 || op->type == GGML_TYPE_Q4_1 || op->type == GGML_TYPE_Q5_0 || op->type == GGML_TYPE_Q5_1 || op->type == GGML_TYPE_Q8_0 || op->type == GGML_TYPE_IQ4_NL || op->type == GGML_TYPE_Q4_0_ROCMFP4 || op->type == GGML_TYPE_Q4_0_ROCMFP4_FAST || - op->type == GGML_TYPE_Q3_0_ROCMFPX || op->type == GGML_TYPE_Q6_0_ROCMFPX || + op->type == GGML_TYPE_Q2_0_ROCMFPX || op->type == GGML_TYPE_Q3_0_ROCMFPX || op->type == GGML_TYPE_Q6_0_ROCMFPX || op->type == GGML_TYPE_Q8_0_ROCMFPX || op->type == GGML_TYPE_TURBO3_0 || op->type == GGML_TYPE_TURBO4_0) && op->src[0]->type == GGML_TYPE_F32 && diff --git a/ggml/src/ggml-cuda/mmq.cu b/ggml/src/ggml-cuda/mmq.cu index 788a7e016..6f6d4bb09 100644 --- a/ggml/src/ggml-cuda/mmq.cu +++ b/ggml/src/ggml-cuda/mmq.cu @@ -35,6 +35,9 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con case GGML_TYPE_Q3_0_ROCMFPX: mul_mat_q_case(ctx, args, stream); break; + case GGML_TYPE_Q2_0_ROCMFPX: + mul_mat_q_case(ctx, args, stream); + break; case GGML_TYPE_Q6_0_ROCMFPX: mul_mat_q_case(ctx, args, stream); break; @@ -297,6 +300,7 @@ bool ggml_cuda_should_use_mmq(enum ggml_type type, int cc, int64_t ne11, int64_t case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: case GGML_TYPE_NVFP4: diff --git a/ggml/src/ggml-cuda/mmq.cuh b/ggml/src/ggml-cuda/mmq.cuh index a18d12f86..426bde95f 100644 --- a/ggml/src/ggml-cuda/mmq.cuh +++ b/ggml/src/ggml-cuda/mmq.cuh @@ -75,6 +75,7 @@ static mmq_q8_1_ds_layout mmq_get_q8_1_ds_layout(const ggml_type type_x) { case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: return MMQ_Q8_1_DS_LAYOUT_D4; @@ -222,6 +223,7 @@ static constexpr __host__ __device__ tile_x_sizes mmq_get_dp4a_tile_x_sizes(ggml case GGML_TYPE_IQ4_XS: return MMQ_DP4A_TXS_Q8_0; case GGML_TYPE_IQ4_NL: return MMQ_DP4A_TXS_Q8_0; case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: return MMQ_DP4A_TXS_Q8_0_16; case GGML_TYPE_Q8_0_ROCMFPX: @@ -281,6 +283,7 @@ static constexpr __host__ __device__ int mmq_get_mma_tile_x_k(ggml_type type) { case GGML_TYPE_IQ4_XS: return MMQ_MMA_TILE_X_K_Q8_0; case GGML_TYPE_IQ4_NL: return MMQ_MMA_TILE_X_K_Q8_0; case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: return MMQ_MMA_TILE_X_K_Q3_K; case GGML_TYPE_Q8_0_ROCMFPX: @@ -1054,6 +1057,64 @@ template static __device__ __forceinline__ void loa } } +template static __device__ __forceinline__ void load_tiles_rocmfpx_fp2( + const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) { + constexpr int nwarps = mmq_get_nwarps_device(); + constexpr int warp_size = ggml_cuda_get_physical_warp_size(); + +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + int * x_qs = (int *) x_tile; + float * x_df = (float *) (x_qs + MMQ_TILE_NE_K*2); +#else + constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q2_0_ROCMFPX, mmq_y); + int * x_qs = (int *) x_tile; + float * x_df = (float *) (x_qs + txs.qs); +#endif + + constexpr int threads_per_row = 32; + constexpr int nrows = warp_size / threads_per_row; + const int txi = warp_size > threads_per_row ? threadIdx.x % threads_per_row : threadIdx.x; + const int kbx = txi / QI_ROCMFP2; + const int kqsx = txi % QI_ROCMFP2; + +#pragma unroll + for (int i0 = 0; i0 < mmq_y; i0 += nrows*nwarps) { + int i = i0 + (nrows == 1 ? threadIdx.y : threadIdx.y*nrows + threadIdx.x/threads_per_row); + if (need_check) { + i = min(i, i_max); + } + const block_rocmfp2 * bxi = (const block_rocmfp2 *) x + kbx0 + i*stride + kbx; + const int v0 = rocmfpx_pack4_fp2_vec_cuda(bxi[0].qs[kqsx]); + const int v1 = rocmfpx_pack4_fp2_vec_cuda(bxi[MMQ_TILE_NE_K/QI_ROCMFP2].qs[kqsx]); +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + x_qs[i*MMQ_MMA_TILE_X_K_Q3_K + 0 + txi] = v0; + x_qs[i*MMQ_MMA_TILE_X_K_Q3_K + MMQ_TILE_NE_K + txi] = v1; +#else + x_qs[i*(2*MMQ_TILE_NE_K + 1) + 0 + txi] = v0; + x_qs[i*(2*MMQ_TILE_NE_K + 1) + MMQ_TILE_NE_K + txi] = v1; +#endif + } + + constexpr int blocks_per_tile_x_row = 2*MMQ_TILE_NE_K / QI_ROCMFP2; + constexpr int rows_per_warp = warp_size / blocks_per_tile_x_row; + const int kbxd = threadIdx.x % blocks_per_tile_x_row; +#pragma unroll + for (int i0 = 0; i0 < mmq_y; i0 += nwarps * rows_per_warp) { + int i = i0 + threadIdx.y * rows_per_warp + threadIdx.x / blocks_per_tile_x_row; + if (need_check) { + i = min(i, i_max); + } + const block_rocmfp2 * bxi = (const block_rocmfp2 *) x + kbx0 + i*stride + kbxd; +#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) + x_df[i*MMQ_MMA_TILE_X_K_Q3_K + 2*kbxd + 0] = rocmfpx_ue4m3_to_fp32_finite(bxi->e[0]); + x_df[i*MMQ_MMA_TILE_X_K_Q3_K + 2*kbxd + 1] = rocmfpx_ue4m3_to_fp32_finite(bxi->e[1]); +#else + x_df[i*(2*MMQ_TILE_NE_K*2/QI8_0) + i/(QI8_0/4) + 2*kbxd + 0] = rocmfpx_ue4m3_to_fp32_finite(bxi->e[0]); + x_df[i*(2*MMQ_TILE_NE_K*2/QI8_0) + i/(QI8_0/4) + 2*kbxd + 1] = rocmfpx_ue4m3_to_fp32_finite(bxi->e[1]); +#endif + } +} + template static __device__ __forceinline__ void load_tiles_rocmfpx_fp3( const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) { constexpr int nwarps = mmq_get_nwarps_device(); @@ -3781,6 +3842,14 @@ struct mmq_type_traits { static constexpr vec_dot_mmq_t vec_dot_dp4a = vec_dot_q8_0_16_q8_1_dp4a; }; +template +struct mmq_type_traits { + static constexpr int vdr = VDR_ROCMFP2_Q8_1_MMQ; + static constexpr load_tiles_mmq_t load_tiles = load_tiles_rocmfpx_fp2; + static constexpr vec_dot_mmq_t vec_dot_mma = vec_dot_q8_0_16_q8_1_mma; + static constexpr vec_dot_mmq_t vec_dot_dp4a = vec_dot_q8_0_16_q8_1_dp4a; +}; + template struct mmq_type_traits { static constexpr int vdr = VDR_ROCMFP6_Q8_1_MMQ; @@ -4631,6 +4700,7 @@ extern DECL_MMQ_CASE(GGML_TYPE_MXFP4); extern DECL_MMQ_CASE(GGML_TYPE_Q4_0_ROCMFP4); extern DECL_MMQ_CASE(GGML_TYPE_Q4_0_ROCMFP4_FAST); extern DECL_MMQ_CASE(GGML_TYPE_Q3_0_ROCMFPX); +extern DECL_MMQ_CASE(GGML_TYPE_Q2_0_ROCMFPX); extern DECL_MMQ_CASE(GGML_TYPE_Q6_0_ROCMFPX); extern DECL_MMQ_CASE(GGML_TYPE_Q8_0_ROCMFPX); extern DECL_MMQ_CASE(GGML_TYPE_NVFP4); diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu index 1c019e47f..8912af67f 100644 --- a/ggml/src/ggml-cuda/mmvq.cu +++ b/ggml/src/ggml-cuda/mmvq.cu @@ -71,6 +71,31 @@ #error "GGML_ROCMFPX_RDNA35_NWARPS_MAX_NCOLS must be between 1 and MMVQ_MAX_BATCH_SIZE" #endif +#ifndef GGML_ROCMFP2_RDNA35_NWARPS +#define GGML_ROCMFP2_RDNA35_NWARPS GGML_ROCMFPX_RDNA35_NWARPS +#endif + +#if GGML_ROCMFP2_RDNA35_NWARPS != 1 && GGML_ROCMFP2_RDNA35_NWARPS != 2 && \ + GGML_ROCMFP2_RDNA35_NWARPS != 4 && GGML_ROCMFP2_RDNA35_NWARPS != 8 +#error "GGML_ROCMFP2_RDNA35_NWARPS must be one of: 1, 2, 4, 8" +#endif + +#ifndef GGML_ROCMFP2_RDNA35_NWARPS_MAX_NCOLS +#define GGML_ROCMFP2_RDNA35_NWARPS_MAX_NCOLS GGML_ROCMFPX_RDNA35_NWARPS_MAX_NCOLS +#endif + +#ifndef GGML_ROCMFP2_RDNA35_NWARPS_MIN_NCOLS +#define GGML_ROCMFP2_RDNA35_NWARPS_MIN_NCOLS 1 +#endif + +#if GGML_ROCMFP2_RDNA35_NWARPS_MIN_NCOLS < 1 || GGML_ROCMFP2_RDNA35_NWARPS_MIN_NCOLS > GGML_ROCMFP2_RDNA35_NWARPS_MAX_NCOLS +#error "GGML_ROCMFP2_RDNA35_NWARPS_MIN_NCOLS must be between 1 and GGML_ROCMFP2_RDNA35_NWARPS_MAX_NCOLS" +#endif + +#if GGML_ROCMFP2_RDNA35_NWARPS_MAX_NCOLS < 1 || GGML_ROCMFP2_RDNA35_NWARPS_MAX_NCOLS > MMVQ_MAX_BATCH_SIZE +#error "GGML_ROCMFP2_RDNA35_NWARPS_MAX_NCOLS must be between 1 and MMVQ_MAX_BATCH_SIZE" +#endif + #ifndef GGML_ROCMFPX_RDNA35_MMID_MAX_BATCH #define GGML_ROCMFPX_RDNA35_MMID_MAX_BATCH MMVQ_MAX_BATCH_SIZE #endif @@ -120,6 +145,8 @@ static constexpr __device__ vec_dot_q_cuda_t get_vec_dot_q_cuda(ggml_type type) return vec_dot_rocmfp4_fast_q8_1; case GGML_TYPE_Q3_0_ROCMFPX: return vec_dot_rocmfpx_fp3_q8_1; + case GGML_TYPE_Q2_0_ROCMFPX: + return vec_dot_rocmfpx_fp2_q8_1; case GGML_TYPE_Q6_0_ROCMFPX: return vec_dot_rocmfpx_fp6_q8_1; case GGML_TYPE_Q8_0_ROCMFPX: @@ -158,6 +185,8 @@ static constexpr __host__ __device__ int get_vdr_mmvq(ggml_type type) { return VDR_ROCMFP4_FAST_Q8_1_MMVQ; case GGML_TYPE_Q3_0_ROCMFPX: return VDR_ROCMFP3_Q8_1_MMVQ; + case GGML_TYPE_Q2_0_ROCMFPX: + return VDR_ROCMFP2_Q8_1_MMVQ; case GGML_TYPE_Q6_0_ROCMFPX: return VDR_ROCMFP6_Q8_1_MMVQ; case GGML_TYPE_Q8_0_ROCMFPX: @@ -349,6 +378,7 @@ static constexpr __host__ __device__ int get_mmvq_mmid_max_batch_rdna3_5(ggml_ty case GGML_TYPE_Q4_0_ROCMFP4_FAST: return GGML_ROCMFP4_RDNA35_MMID_MAX_BATCH; case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: return GGML_ROCMFPX_RDNA35_MMID_MAX_BATCH; @@ -550,28 +580,28 @@ static constexpr __host__ __device__ int calc_nwarps(ggml_type type, int ncols_d return 1; } if (table_id == MMVQ_PARAMETERS_RDNA3_5) { - if (ncols_dst >= 1 && ncols_dst <= GGML_ROCMFP4_RDNA35_NWARPS_MAX_NCOLS) { - switch (type) { - // NOTE: giving stock Q4_0 the ROCmFP4 RDNA3.5 config (nwarps=2) - // was tried and measured ~2.3% SLOWER tg on gfx1151 (Q4_0's - // access pattern differs), so Q4_0 is intentionally left at the - // default nwarps=1 here. Run Google QAT (Q4_0) models natively; - // do not add Q4_0 to this switch without a fresh gfx1151 A/B. - case GGML_TYPE_Q4_0_ROCMFP4: - case GGML_TYPE_Q4_0_ROCMFP4_FAST: - return GGML_ROCMFP4_RDNA35_NWARPS; - case GGML_TYPE_Q3_0_ROCMFPX: - case GGML_TYPE_Q6_0_ROCMFPX: - case GGML_TYPE_Q8_0_ROCMFPX: - if (ncols_dst <= GGML_ROCMFPX_RDNA35_NWARPS_MAX_NCOLS) { - return GGML_ROCMFPX_RDNA35_NWARPS; - } - return 1; - default: - return 1; - } + if (ncols_dst < 1) { + return 1; + } + switch (type) { + // NOTE: giving stock Q4_0 the ROCmFP4 RDNA3.5 config (nwarps=2) + // was tried and measured ~2.3% SLOWER tg on gfx1151 (Q4_0's + // access pattern differs), so Q4_0 is intentionally left at the + // default nwarps=1 here. Run Google QAT (Q4_0) models natively; + // do not add Q4_0 to this switch without a fresh gfx1151 A/B. + case GGML_TYPE_Q4_0_ROCMFP4: + case GGML_TYPE_Q4_0_ROCMFP4_FAST: + return ncols_dst <= GGML_ROCMFP4_RDNA35_NWARPS_MAX_NCOLS ? GGML_ROCMFP4_RDNA35_NWARPS : 1; + case GGML_TYPE_Q2_0_ROCMFPX: + return ncols_dst >= GGML_ROCMFP2_RDNA35_NWARPS_MIN_NCOLS && + ncols_dst <= GGML_ROCMFP2_RDNA35_NWARPS_MAX_NCOLS ? GGML_ROCMFP2_RDNA35_NWARPS : 1; + case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q6_0_ROCMFPX: + case GGML_TYPE_Q8_0_ROCMFPX: + return ncols_dst <= GGML_ROCMFPX_RDNA35_NWARPS_MAX_NCOLS ? GGML_ROCMFPX_RDNA35_NWARPS : 1; + default: + return 1; } - return 1; } if (table_id == MMVQ_PARAMETERS_RDNA3_0) { // RDNA3 (W7900): stricter whitelist than RDNA4. @@ -621,6 +651,7 @@ static constexpr __host__ __device__ int calc_rows_per_block(ggml_type type, int case GGML_TYPE_Q4_0_ROCMFP4_FAST: return GGML_ROCMFP4_RDNA35_RPB_WIDE_FAST; case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: return GGML_ROCMFPX_RDNA35_RPB_WIDE; @@ -632,13 +663,71 @@ static constexpr __host__ __device__ int calc_rows_per_block(ggml_type type, int return 1; } +// FP2 has a very small payload but a non-trivial byte-to-int8 expansion. The +// generic multi-column loop calls vec_dot once per destination column, causing +// that expansion (and the FP2 scale decode) to be inlined once per column. +// MTP verification uses exactly these small multi-column shapes. Expand the +// target weights once and reuse them against every Q8_1 activation column. +template +static __device__ __forceinline__ void vec_dot_rocmfpx_fp2_q8_1_ncols( + const void * __restrict__ vx, + const block_q8_1 * __restrict__ y, + const uint32_t stride_col_y, + const int kbx, + const int kby, + const int kqs, + float (&tmp)[ncols_dst][rows_per_cuda_block], + const int row) { + const block_rocmfp2 * bq2 = (const block_rocmfp2 *) vx + kbx; + + int values[VDR_ROCMFP2_Q8_1_MMVQ]; +#pragma unroll + for (int i = 0; i < VDR_ROCMFP2_Q8_1_MMVQ; ++i) { + values[i] = rocmfpx_pack4_fp2_vec_cuda(bq2->qs[kqs + i]); + } + +#if VDR_ROCMFP2_Q8_1_MMVQ <= 4 + const float dx = rocmfpx_ue4m3_to_fp32_finite(bq2->e[kqs / 4]); +#endif + +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { + const block_q8_1 * bq8 = &y[j*stride_col_y + kby]; + const int * q8 = (const int *) bq8->qs; + +#if VDR_ROCMFP2_Q8_1_MMVQ <= 4 + int sumi = 0; +#pragma unroll + for (int i = 0; i < VDR_ROCMFP2_Q8_1_MMVQ; ++i) { + sumi = ggml_cuda_dp4a(values[i], q8[kqs + i], sumi); + } + tmp[j][row] += dx * __low2float(bq8->ds) * sumi; +#else + int sumi0 = 0; + int sumi1 = 0; +#pragma unroll + for (int i = 0; i < VDR_ROCMFP2_Q8_1_MMVQ; ++i) { + const int group = kqs + i; + if (group < QI_ROCMFP2/2) { + sumi0 = ggml_cuda_dp4a(values[i], q8[group], sumi0); + } else { + sumi1 = ggml_cuda_dp4a(values[i], q8[group], sumi1); + } + } + const float dx0 = rocmfpx_ue4m3_to_fp32_finite(bq2->e[0]); + const float dx1 = rocmfpx_ue4m3_to_fp32_finite(bq2->e[1]); + tmp[j][row] += __low2float(bq8->ds) * (dx0*sumi0 + dx1*sumi1); +#endif + } +} + template static constexpr int calc_moe_mmvq_rows_per_block() { #if defined(GGML_USE_HIP) if constexpr (type == GGML_TYPE_Q4_0_ROCMFP4 || type == GGML_TYPE_Q4_0_ROCMFP4_FAST) { return GGML_ROCMFP4_MOE_MMVQ_ROWS_PER_BLOCK; } - if constexpr (type == GGML_TYPE_Q3_0_ROCMFPX || type == GGML_TYPE_Q6_0_ROCMFPX || type == GGML_TYPE_Q8_0_ROCMFPX) { + if constexpr (type == GGML_TYPE_Q2_0_ROCMFPX || type == GGML_TYPE_Q3_0_ROCMFPX || type == GGML_TYPE_Q6_0_ROCMFPX || type == GGML_TYPE_Q8_0_ROCMFPX) { return GGML_ROCMFPX_MOE_MMVQ_ROWS_PER_BLOCK; } #endif @@ -743,16 +832,30 @@ static __global__ void mul_mat_vec_q( // x block quant index when casting the quants to int const int kqs = vdr * (tid % (qi/vdr)); -#pragma unroll - for (int j = 0; j < ncols_dst; ++j) { + if constexpr (type == GGML_TYPE_Q2_0_ROCMFPX) { #pragma unroll for (int i = 0; i < rows_per_cuda_block; ++i) { - tmp[j][i] += vec_dot_q_cuda( - vx, &y[j*stride_col_y + kby], kbx_offset + i*stride_row_x + kbx, kqs); + vec_dot_rocmfpx_fp2_q8_1_ncols( + vx, y, stride_col_y, kbx_offset + i*stride_row_x + kbx, kby, kqs, tmp, i); if constexpr (has_fusion) { if (use_gate) { - tmp_gate[j][i] += vec_dot_q_cuda( - vgate, &y[j*stride_col_y + kby], kbx_offset + i*stride_row_x + kbx, kqs); + vec_dot_rocmfpx_fp2_q8_1_ncols( + vgate, y, stride_col_y, kbx_offset + i*stride_row_x + kbx, kby, kqs, tmp_gate, i); + } + } + } + } else { +#pragma unroll + for (int j = 0; j < ncols_dst; ++j) { +#pragma unroll + for (int i = 0; i < rows_per_cuda_block; ++i) { + tmp[j][i] += vec_dot_q_cuda( + vx, &y[j*stride_col_y + kby], kbx_offset + i*stride_row_x + kbx, kqs); + if constexpr (has_fusion) { + if (use_gate) { + tmp_gate[j][i] += vec_dot_q_cuda( + vgate, &y[j*stride_col_y + kby], kbx_offset + i*stride_row_x + kbx, kqs); + } } } } @@ -1216,6 +1319,12 @@ static void mul_mat_vec_q_switch_type( nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); break; + case GGML_TYPE_Q2_0_ROCMFPX: + mul_mat_vec_q_switch_ncols_dst + (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, + nchannels_x, nchannels_y, nchannels_dst, stride_channel_x, stride_channel_y, stride_channel_dst, + nsamples_x, nsamples_dst, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride, stream); + break; case GGML_TYPE_Q6_0_ROCMFPX: mul_mat_vec_q_switch_ncols_dst (vx, vy, ids, fusion, dst, ncols_x, nrows_x, ncols_dst, stride_row_x, stride_col_y, stride_col_dst, diff --git a/ggml/src/ggml-cuda/template-instances/generate_cu_files.py b/ggml/src/ggml-cuda/template-instances/generate_cu_files.py index 8820e6301..4d98fe34a 100755 --- a/ggml/src/ggml-cuda/template-instances/generate_cu_files.py +++ b/ggml/src/ggml-cuda/template-instances/generate_cu_files.py @@ -42,7 +42,7 @@ "GGML_TYPE_Q2_K", "GGML_TYPE_Q3_K", "GGML_TYPE_Q4_K", "GGML_TYPE_Q5_K", "GGML_TYPE_Q6_K", "GGML_TYPE_IQ2_XXS", "GGML_TYPE_IQ2_XS", "GGML_TYPE_IQ2_S", "GGML_TYPE_IQ3_XXS", "GGML_TYPE_IQ3_S", "GGML_TYPE_IQ1_S", "GGML_TYPE_IQ4_NL", "GGML_TYPE_IQ4_XS", "GGML_TYPE_MXFP4", "GGML_TYPE_Q4_0_ROCMFP4", "GGML_TYPE_Q4_0_ROCMFP4_FAST", "GGML_TYPE_NVFP4", - "GGML_TYPE_Q3_0_ROCMFPX", "GGML_TYPE_Q6_0_ROCMFPX", "GGML_TYPE_Q8_0_ROCMFPX" + "GGML_TYPE_Q2_0_ROCMFPX", "GGML_TYPE_Q3_0_ROCMFPX", "GGML_TYPE_Q6_0_ROCMFPX", "GGML_TYPE_Q8_0_ROCMFPX" ] SOURCE_MMQ = """// This file has been autogenerated by generate_cu_files.py, do not edit manually. diff --git a/ggml/src/ggml-cuda/template-instances/mmq-instance-q2_0_rocmfpx.cu b/ggml/src/ggml-cuda/template-instances/mmq-instance-q2_0_rocmfpx.cu new file mode 100644 index 000000000..b5c5af14c --- /dev/null +++ b/ggml/src/ggml-cuda/template-instances/mmq-instance-q2_0_rocmfpx.cu @@ -0,0 +1,5 @@ +// This file has been autogenerated by generate_cu_files.py, do not edit manually. + +#include "../mmq.cuh" + +DECL_MMQ_CASE(GGML_TYPE_Q2_0_ROCMFPX); diff --git a/ggml/src/ggml-cuda/vecdotq.cuh b/ggml/src/ggml-cuda/vecdotq.cuh index 4c82c8b0b..68e28291e 100644 --- a/ggml/src/ggml-cuda/vecdotq.cuh +++ b/ggml/src/ggml-cuda/vecdotq.cuh @@ -353,6 +353,10 @@ static __device__ __forceinline__ float vec_dot_mxfp4_q8_1( #define GGML_ROCMFP3_Q8_1_MMVQ_VDR 2 #endif +#ifndef GGML_ROCMFP2_Q8_1_MMVQ_VDR +#define GGML_ROCMFP2_Q8_1_MMVQ_VDR 4 +#endif + #ifndef GGML_ROCMFP6_Q8_1_MMVQ_VDR #define GGML_ROCMFP6_Q8_1_MMVQ_VDR 4 #endif @@ -368,6 +372,13 @@ static __device__ __forceinline__ float vec_dot_mxfp4_q8_1( #error "GGML_ROCMFP3_Q8_1_MMVQ_VDR must be 1, 2, 4, or 8" #endif +#if GGML_ROCMFP2_Q8_1_MMVQ_VDR != 1 && \ + GGML_ROCMFP2_Q8_1_MMVQ_VDR != 2 && \ + GGML_ROCMFP2_Q8_1_MMVQ_VDR != 4 && \ + GGML_ROCMFP2_Q8_1_MMVQ_VDR != 8 +#error "GGML_ROCMFP2_Q8_1_MMVQ_VDR must be 1, 2, 4, or 8" +#endif + #if GGML_ROCMFP6_Q8_1_MMVQ_VDR != 1 && \ GGML_ROCMFP6_Q8_1_MMVQ_VDR != 2 && \ GGML_ROCMFP6_Q8_1_MMVQ_VDR != 4 && \ @@ -395,10 +406,12 @@ static __device__ __forceinline__ float vec_dot_mxfp4_q8_1( #endif #define VDR_ROCMFP3_Q8_1_MMVQ GGML_ROCMFP3_Q8_1_MMVQ_VDR +#define VDR_ROCMFP2_Q8_1_MMVQ GGML_ROCMFP2_Q8_1_MMVQ_VDR #define VDR_ROCMFP6_Q8_1_MMVQ GGML_ROCMFP6_Q8_1_MMVQ_VDR #define VDR_ROCMFP8_Q8_1_MMVQ GGML_ROCMFP8_Q8_1_MMVQ_VDR #define VDR_ROCMFP3_Q8_1_MMQ 4 +#define VDR_ROCMFP2_Q8_1_MMQ 4 #ifndef VDR_ROCMFP6_Q8_1_MMQ #define VDR_ROCMFP6_Q8_1_MMQ 4 #endif @@ -422,6 +435,31 @@ static __device__ __forceinline__ int rocmfpx_decode_fp3_code_vec_cuda(const uin return (code & 4u) ? -mag : mag; } +static __device__ __forceinline__ int rocmfpx_pack4_fp2_vec_cuda(const uint8_t packed) { +#if defined(GGML_USE_HIP) + // Spread the low and high code bits independently into the low bits of + // four byte selectors. Separating the planes prevents carries between + // adjacent two-bit codes during the multiply. + constexpr uint32_t byte_lsb = 0x01010101u; + const uint32_t lo = (((uint32_t) packed & 0x55u) * 0x00041041u) & byte_lsb; + const uint32_t hi = ((((uint32_t) packed >> 1) & 0x55u) * 0x00041041u) & byte_lsb; + const uint32_t selectors = lo | (hi << 1); + + // v_perm_b32 selector bytes 0..3 choose bytes from the second operand. + // Packed little-endian table bytes are {-4, -1, +1, +4}. + return __builtin_amdgcn_perm(0u, 0x0401fffcu, selectors); +#else + int result = 0; +#pragma unroll + for (int lane = 0; lane < 4; ++lane) { + const uint32_t code = (packed >> (2 * lane)) & 3u; + const int value = code == 0u ? -4 : code == 1u ? -1 : code == 2u ? 1 : 4; + result |= ((int) (uint8_t) (int8_t) value) << (8 * lane); + } + return result; +#endif +} + static __device__ __forceinline__ int rocmfpx_decode_fp6_code_vec_cuda(const uint32_t code) { #if GGML_ROCMFP6_FAST_SIGNMAG_PACK const int mag = (int) (code & 31u); @@ -594,6 +632,39 @@ static __device__ __forceinline__ float vec_dot_rocmfpx_fp3_q8_1( return db * (rocmfpx_ue4m3_to_fp32_finite(bq3->e[0]) * sumi0 + rocmfpx_ue4m3_to_fp32_finite(bq3->e[1]) * sumi1); } +static __device__ __forceinline__ float vec_dot_rocmfpx_fp2_q8_1( + const void * __restrict__ vbq, const block_q8_1 * __restrict__ bq8_1, const int & kbx, const int & iqs) { + + const block_rocmfp2 * bq2 = (const block_rocmfp2 *) vbq + kbx; + const int * q8 = (const int *) bq8_1->qs; +#if VDR_ROCMFP2_Q8_1_MMVQ <= 4 + int sumi = 0; +#pragma unroll + for (int i = 0; i < VDR_ROCMFP2_Q8_1_MMVQ; ++i) { + const int group = iqs + i; + const int values = rocmfpx_pack4_fp2_vec_cuda(bq2->qs[group]); + sumi = ggml_cuda_dp4a(values, q8[group], sumi); + } + const float db = __low2float(bq8_1->ds); + return db * rocmfpx_ue4m3_to_fp32_finite(bq2->e[iqs / 4]) * sumi; +#else + int sumi0 = 0; + int sumi1 = 0; +#pragma unroll + for (int i = 0; i < VDR_ROCMFP2_Q8_1_MMVQ; ++i) { + const int group = iqs + i; + const int values = rocmfpx_pack4_fp2_vec_cuda(bq2->qs[group]); + if (group < QI_ROCMFP2/2) { + sumi0 = ggml_cuda_dp4a(values, q8[group], sumi0); + } else { + sumi1 = ggml_cuda_dp4a(values, q8[group], sumi1); + } + } + const float db = __low2float(bq8_1->ds); + return db * (rocmfpx_ue4m3_to_fp32_finite(bq2->e[0]) * sumi0 + rocmfpx_ue4m3_to_fp32_finite(bq2->e[1]) * sumi1); +#endif +} + static __device__ __forceinline__ float vec_dot_rocmfpx_fp6_q8_1( const void * __restrict__ vbq, const block_q8_1 * __restrict__ bq8_1, const int & kbx, const int & iqs) { diff --git a/ggml/src/ggml-quants.c b/ggml/src/ggml-quants.c index 6ea013d3c..e73a82259 100644 --- a/ggml/src/ggml-quants.c +++ b/ggml/src/ggml-quants.c @@ -5724,6 +5724,13 @@ bool ggml_validate_row_data(enum ggml_type type, const void * data, size_t nbyte return false; } } break; + case GGML_TYPE_Q2_0_ROCMFPX: + { + if (!rocmfpx_validate_row_data_fp2(data, nbytes)) { + fprintf(stderr, "%s: invalid ROCmFPx FP2 row data\n", __func__); + return false; + } + } break; case GGML_TYPE_Q6_0_ROCMFPX: { if (!rocmfpx_validate_row_data_fp6(data, nbytes)) { diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index dddf0a5aa..901e3b9b9 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -708,6 +708,14 @@ static const struct ggml_type_traits type_traits[GGML_TYPE_COUNT] = { .to_float = (ggml_to_float_t) rocmfpx_dequantize_row_fp3, .from_float_ref = (ggml_from_float_t) rocmfpx_quantize_row_fp3_ref, }, + [GGML_TYPE_Q2_0_ROCMFPX] = { + .type_name = "q2_0_rocmfpx", + .blck_size = QK_ROCMFP2, + .type_size = sizeof(block_rocmfp2), + .is_quantized = true, + .to_float = (ggml_to_float_t) rocmfpx_dequantize_row_fp2, + .from_float_ref = (ggml_from_float_t) rocmfpx_quantize_row_fp2_ref, + }, [GGML_TYPE_Q6_0_ROCMFPX] = { .type_name = "q6_0_rocmfpx", .blck_size = QK_ROCMFP6, @@ -1502,6 +1510,8 @@ enum ggml_type ggml_ftype_to_ggml_type(enum ggml_ftype ftype) { wtype = GGML_TYPE_Q4_0_ROCMFP4_FAST; break; case GGML_FTYPE_MOSTLY_Q3_0_ROCMFPX: wtype = GGML_TYPE_Q3_0_ROCMFPX; break; + case GGML_FTYPE_MOSTLY_Q2_0_ROCMFPX: + wtype = GGML_TYPE_Q2_0_ROCMFPX; break; case GGML_FTYPE_MOSTLY_Q6_0_ROCMFPX: wtype = GGML_TYPE_Q6_0_ROCMFPX; break; case GGML_FTYPE_MOSTLY_Q8_0_ROCMFPX: @@ -8013,6 +8023,9 @@ size_t ggml_quantize_chunk( case GGML_TYPE_Q3_0_ROCMFPX: result = rocmfpx_quantize_fp3(src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; + case GGML_TYPE_Q2_0_ROCMFPX: + result = rocmfpx_quantize_fp2(src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); + break; case GGML_TYPE_Q6_0_ROCMFPX: result = rocmfpx_quantize_fp6(src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index 6f5886286..d641b3eae 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -4474,6 +4474,7 @@ class GGMLQuantizationType(IntEnum): Q6_0_ROCMFPX = 102 Q8_0_ROCMFPX = 103 Q3_0_ROCMFPX = 104 + Q2_0_ROCMFPX = 107 class ExpertGatingFuncType(IntEnum): @@ -4546,6 +4547,7 @@ class LlamaFileType(IntEnum): MOSTLY_Q8_0_ROCMFPX_AGENT = 115 MOSTLY_Q6_0_ROCMFPX_LEAN = 116 MOSTLY_Q6_0_ROCMFPX_AGENT_LEAN = 117 + MOSTLY_Q2_0_ROCMFPX = 119 # except 1d tensors GUESSED = 1024 # not specified in the model file @@ -4674,6 +4676,7 @@ class VisionProjectorType: GGMLQuantizationType.Q6_0_ROCMFPX: (32, 24 + 2), GGMLQuantizationType.Q8_0_ROCMFPX: (32, 32 + 1), GGMLQuantizationType.Q3_0_ROCMFPX: (32, 12 + 2), + GGMLQuantizationType.Q2_0_ROCMFPX: (32, 8 + 2), } diff --git a/include/llama.h b/include/llama.h index 359f6a68f..4b6a4c563 100644 --- a/include/llama.h +++ b/include/llama.h @@ -170,6 +170,7 @@ extern "C" { LLAMA_FTYPE_MOSTLY_Q8_0_ROCMFPX_AGENT = 115, // ROCmFPx 8-bit agent/tool-call coherent routing LLAMA_FTYPE_MOSTLY_Q6_0_ROCMFPX_LEAN = 116, // ROCmFPx 6-bit size/speed-biased routing LLAMA_FTYPE_MOSTLY_Q6_0_ROCMFPX_AGENT_LEAN = 117, // ROCmFPx 6-bit agent routing without Q8-heavy boosts + LLAMA_FTYPE_MOSTLY_Q2_0_ROCMFPX = 119, // ROCmFPx 2-bit S40 codebook + dual UE4M3 scales LLAMA_FTYPE_GUESSED = 1024, // not specified in the model file }; diff --git a/scripts/check-rocmfp2-reference.sh b/scripts/check-rocmfp2-reference.sh new file mode 100755 index 000000000..9b72ad58e --- /dev/null +++ b/scripts/check-rocmfp2-reference.sh @@ -0,0 +1,33 @@ +#!/usr/bin/env bash +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +ROOT="${ROOT:-$(cd "$SCRIPT_DIR/.." && pwd)}" +BUILD_DIR="${BUILD_DIR:-$ROOT/build-rocmfp2-reference}" +CC_BIN="${CC:-cc}" +EXTRA_CFLAGS="${EXTRA_CFLAGS:-}" +read -r -a EXTRA_CFLAGS_ARRAY <<< "$EXTRA_CFLAGS" + +mkdir -p "$BUILD_DIR" + +echo "ROCmFP2 Phase-1 CPU reference check" +echo "source_root=$ROOT" +echo "compiler=$CC_BIN" + +"$CC_BIN" \ + -std=c11 \ + -O2 \ + -g \ + -Wall \ + -Wextra \ + -Werror \ + -pedantic \ + -ffp-contract=off \ + -fno-fast-math \ + "${EXTRA_CFLAGS_ARRAY[@]}" \ + "$ROOT/ggml/rocmfpx/rocmfp2_reference.c" \ + "$ROOT/ggml/rocmfpx/test_rocmfp2_reference.c" \ + -lm \ + -o "$BUILD_DIR/test-rocmfp2-reference" + +"$BUILD_DIR/test-rocmfp2-reference" diff --git a/src/llama-model-loader.cpp b/src/llama-model-loader.cpp index dd37bd905..e2b5adefe 100644 --- a/src/llama-model-loader.cpp +++ b/src/llama-model-loader.cpp @@ -46,6 +46,7 @@ static std::string llama_model_ftype_name(llama_ftype ftype) { case LLAMA_FTYPE_MOSTLY_Q4_0_ROCMFP4_STRIX: return "Q4_0_ROCMFP4_STRIX"; case LLAMA_FTYPE_MOSTLY_Q4_0_ROCMFP4_STRIX_LEAN: return "Q4_0_ROCMFP4_STRIX_LEAN"; case LLAMA_FTYPE_MOSTLY_Q3_0_ROCMFPX: return "Q3_0_ROCMFPX"; + case LLAMA_FTYPE_MOSTLY_Q2_0_ROCMFPX: return "Q2_0_ROCMFPX"; case LLAMA_FTYPE_MOSTLY_Q6_0_ROCMFPX: return "Q6_0_ROCMFPX"; case LLAMA_FTYPE_MOSTLY_Q8_0_ROCMFPX: return "Q8_0_ROCMFPX"; case LLAMA_FTYPE_MOSTLY_Q3_0_ROCMFPX_AGENT: return "Q3_0_ROCMFPX_AGENT"; @@ -820,6 +821,7 @@ llama_model_loader::llama_model_loader( case GGML_TYPE_Q4_0_ROCMFP4: ftype = LLAMA_FTYPE_MOSTLY_Q4_0_ROCMFP4; break; case GGML_TYPE_Q4_0_ROCMFP4_FAST: ftype = LLAMA_FTYPE_MOSTLY_Q4_0_ROCMFP4_FAST; break; case GGML_TYPE_Q3_0_ROCMFPX: ftype = LLAMA_FTYPE_MOSTLY_Q3_0_ROCMFPX; break; + case GGML_TYPE_Q2_0_ROCMFPX: ftype = LLAMA_FTYPE_MOSTLY_Q2_0_ROCMFPX; break; case GGML_TYPE_Q6_0_ROCMFPX: ftype = LLAMA_FTYPE_MOSTLY_Q6_0_ROCMFPX; break; case GGML_TYPE_Q8_0_ROCMFPX: ftype = LLAMA_FTYPE_MOSTLY_Q8_0_ROCMFPX; break; case GGML_TYPE_Q4_1: ftype = LLAMA_FTYPE_MOSTLY_Q4_1; break; diff --git a/src/llama-quant.cpp b/src/llama-quant.cpp index a3325cdcb..1459ef9fe 100644 --- a/src/llama-quant.cpp +++ b/src/llama-quant.cpp @@ -425,6 +425,7 @@ static ggml_type tensor_type_fallback(quantize_state_impl & qs, const ggml_tenso case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: return_type = GGML_TYPE_Q4_0; break; case GGML_TYPE_Q3_0_ROCMFPX: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: return_type = GGML_TYPE_Q8_0; break; case GGML_TYPE_Q4_K: return_type = GGML_TYPE_Q5_0; break; @@ -1197,6 +1198,7 @@ ggml_type llama_ftype_get_default_type(llama_ftype ftype) { case LLAMA_FTYPE_MOSTLY_Q4_0_ROCMFP4_STRIX: return GGML_TYPE_Q4_0_ROCMFP4_FAST; case LLAMA_FTYPE_MOSTLY_Q4_0_ROCMFP4_STRIX_LEAN: return GGML_TYPE_Q4_0_ROCMFP4_FAST; case LLAMA_FTYPE_MOSTLY_Q3_0_ROCMFPX: return GGML_TYPE_Q3_0_ROCMFPX; + case LLAMA_FTYPE_MOSTLY_Q2_0_ROCMFPX: return GGML_TYPE_Q2_0_ROCMFPX; case LLAMA_FTYPE_MOSTLY_Q6_0_ROCMFPX: return GGML_TYPE_Q6_0_ROCMFPX; case LLAMA_FTYPE_MOSTLY_Q6_0_ROCMFPX_LEAN: return GGML_TYPE_Q6_0_ROCMFPX; case LLAMA_FTYPE_MOSTLY_Q8_0_ROCMFPX: return GGML_TYPE_Q8_0_ROCMFPX; @@ -1590,6 +1592,7 @@ static void llama_model_quantize_impl(const std::string & fname_inp, const std:: const int64_t nelements = ggml_nelements(tensor); const float * imatrix = nullptr; + std::vector neutral_imatrix; if (imatrix_data) { auto it = imatrix_data->find(tm.remapped_imatrix_name); if (it == imatrix_data->end()) { @@ -1612,6 +1615,16 @@ static void llama_model_quantize_impl(const std::string & fname_inp, const std:: } } } + if (!imatrix && params->pure && + (new_type == GGML_TYPE_IQ2_XXS || new_type == GGML_TYPE_IQ2_XS || new_type == GGML_TYPE_IQ2_S)) { + // A literal-pure IQ2 artifact cannot use the normal protected-tensor + // fallback. Treat tensors absent from a partial calibration matrix as + // unweighted by assigning equal importance to every input column. + // Calibrated tensors continue to use their published importance values. + neutral_imatrix.assign((size_t) tensor->ne[0] * tensor->ne[2], 1.0f); + imatrix = neutral_imatrix.data(); + LLAMA_LOG_WARN("Using neutral uniform importance weights for missing pure-IQ2 tensor %s\n", tensor->name); + } if (!imatrix && tm.requires_imatrix) { LLAMA_LOG_ERROR("\n\n============================================================\n"); LLAMA_LOG_ERROR("Missing importance matrix for tensor %s in a very low-bit quantization\n", tensor->name); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index d9659610a..cd20826c0 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -7552,7 +7552,7 @@ static const ggml_type all_types[] = { GGML_TYPE_Q8_0, GGML_TYPE_Q1_0, GGML_TYPE_MXFP4, GGML_TYPE_Q4_0_ROCMFP4, GGML_TYPE_Q4_0_ROCMFP4_FAST, GGML_TYPE_NVFP4, - GGML_TYPE_Q3_0_ROCMFPX, GGML_TYPE_Q6_0_ROCMFPX, GGML_TYPE_Q8_0_ROCMFPX, + GGML_TYPE_Q2_0_ROCMFPX, GGML_TYPE_Q3_0_ROCMFPX, GGML_TYPE_Q6_0_ROCMFPX, GGML_TYPE_Q8_0_ROCMFPX, GGML_TYPE_Q2_K, GGML_TYPE_Q3_K, GGML_TYPE_Q4_K, GGML_TYPE_Q5_K, GGML_TYPE_Q6_K, @@ -7570,7 +7570,7 @@ static const ggml_type base_types[] = { GGML_TYPE_Q4_1, // for I8MM tests GGML_TYPE_Q4_K, GGML_TYPE_MXFP4, GGML_TYPE_Q4_0_ROCMFP4, GGML_TYPE_Q4_0_ROCMFP4_FAST, GGML_TYPE_NVFP4, // TODO: or "other" - GGML_TYPE_Q3_0_ROCMFPX, GGML_TYPE_Q6_0_ROCMFPX, GGML_TYPE_Q8_0_ROCMFPX, + GGML_TYPE_Q2_0_ROCMFPX, GGML_TYPE_Q3_0_ROCMFPX, GGML_TYPE_Q6_0_ROCMFPX, GGML_TYPE_Q8_0_ROCMFPX, GGML_TYPE_IQ2_XXS }; @@ -8490,6 +8490,17 @@ static std::vector> make_test_cases_eval() { // gpt-oss issue with Vulkan mmq_id test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_MXFP4, GGML_TYPE_F32, 32, 2, false, 2880, 32, 2880)); + // HY3 exact routed-expert shapes: 192 experts, top-8, 4096 hidden, 1536 intermediate. + for (int bs : {1, 4, 32, 128}) { + for (ggml_type type_a : { + GGML_TYPE_Q2_0_ROCMFPX, GGML_TYPE_Q3_0_ROCMFPX, + GGML_TYPE_IQ2_XXS, GGML_TYPE_IQ2_XS, GGML_TYPE_IQ2_S, + GGML_TYPE_Q2_K, GGML_TYPE_Q3_K}) { + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 192, 8, false, 1536, bs, 4096)); + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 192, 8, false, 4096, bs, 1536)); + } + } + for (ggml_type type_a : all_types) { test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 4, 2, false, 64, 16, 3*ggml_blck_size(type_a))); } @@ -9281,6 +9292,19 @@ static std::vector> make_test_cases_perf() { } } + // HY3: 192 routed experts, top-8, hidden size 4096, expert intermediate size 1536. + // Cover both the gate/up projection and the reversed down projection at decode + // and prompt-sized batches for the native ROCmFPX low-bit kernels. + for (int bs : {1, 4, 32, 128}) { + for (ggml_type type_a : { + GGML_TYPE_Q2_0_ROCMFPX, GGML_TYPE_Q3_0_ROCMFPX, + GGML_TYPE_IQ2_XXS, GGML_TYPE_IQ2_XS, GGML_TYPE_IQ2_S, + GGML_TYPE_Q2_K, GGML_TYPE_Q3_K}) { + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 192, 8, false, 1536, bs, 4096)); + test_cases.emplace_back(new test_mul_mat_id(type_a, GGML_TYPE_F32, 192, 8, false, 4096, bs, 1536)); + } + } + for (int bs : {1, 4, 8, 32, 64, 128, 256, 512}) { for (ggml_type type_a : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, GGML_TYPE_Q4_K, GGML_TYPE_Q6_K, GGML_TYPE_IQ2_XS}) { for (ggml_type type_b : {GGML_TYPE_F32}) { diff --git a/tests/test-quantize-fns.cpp b/tests/test-quantize-fns.cpp index 1e87af9dd..1f719381e 100644 --- a/tests/test-quantize-fns.cpp +++ b/tests/test-quantize-fns.cpp @@ -24,6 +24,7 @@ constexpr float MAX_QUANTIZATION_TOTAL_ERROR_3BITS_XXS = 0.0050f; constexpr float MAX_QUANTIZATION_TOTAL_ERROR_FP4 = 0.0030f; constexpr float MAX_DOT_PRODUCT_ERROR = 0.02f; constexpr float MAX_DOT_PRODUCT_ERROR_LOWBIT = 0.04f; +constexpr float MAX_DOT_PRODUCT_ERROR_ROCMFP2 = 0.07f; constexpr float MAX_DOT_PRODUCT_ERROR_FP4 = 0.03f; constexpr float MAX_DOT_PRODUCT_ERROR_BINARY = 0.40f; constexpr float MAX_DOT_PRODUCT_ERROR_TERNARY = 0.15f; @@ -152,6 +153,7 @@ int main(int argc, char * argv[]) { type == GGML_TYPE_TQ2_0 ? MAX_QUANTIZATION_TOTAL_ERROR_TERNARY : type == GGML_TYPE_Q2_K ? MAX_QUANTIZATION_TOTAL_ERROR_2BITS : type == GGML_TYPE_IQ2_S ? MAX_QUANTIZATION_TOTAL_ERROR_2BITS : + type == GGML_TYPE_Q2_0_ROCMFPX ? MAX_QUANTIZATION_TOTAL_ERROR_2BITS : type == GGML_TYPE_Q3_K ? MAX_QUANTIZATION_TOTAL_ERROR_3BITS : type == GGML_TYPE_IQ3_S ? MAX_QUANTIZATION_TOTAL_ERROR_3BITS : type == GGML_TYPE_IQ3_XXS ? MAX_QUANTIZATION_TOTAL_ERROR_3BITS_XXS : @@ -171,7 +173,9 @@ int main(int argc, char * argv[]) { } const float vec_dot_error = dot_product_error(qfns, qfns_cpu, test_size, test_data.data(), test_data2.data()); - const float max_allowed_error = type == GGML_TYPE_Q2_K || type == GGML_TYPE_IQ2_XS || type == GGML_TYPE_IQ2_XXS || + const float max_allowed_error = type == GGML_TYPE_Q2_0_ROCMFPX + ? MAX_DOT_PRODUCT_ERROR_ROCMFP2 + : type == GGML_TYPE_Q2_K || type == GGML_TYPE_IQ2_XS || type == GGML_TYPE_IQ2_XXS || type == GGML_TYPE_IQ3_XXS || type == GGML_TYPE_IQ3_S || type == GGML_TYPE_IQ2_S || type == GGML_TYPE_Q3_0_ROCMFPX ? MAX_DOT_PRODUCT_ERROR_LOWBIT diff --git a/tools/quantize/quantize.cpp b/tools/quantize/quantize.cpp index 9195f9f26..b9e40b618 100644 --- a/tools/quantize/quantize.cpp +++ b/tools/quantize/quantize.cpp @@ -44,6 +44,7 @@ static const std::vector QUANT_OPTIONS = { { "Q4_0_ROCMFP4_STRIX", LLAMA_FTYPE_MOSTLY_Q4_0_ROCMFP4_STRIX, " ~4.49 bpw ROCmFP4 Strix Halo attn-K/V quality recipe", }, { "Q4_0_ROCMFP4_STRIX_LEAN", LLAMA_FTYPE_MOSTLY_Q4_0_ROCMFP4_STRIX_LEAN, " ~4.38 bpw ROCmFP4 Strix K/V + Q5_K token embeddings", }, { "Q3_0_ROCMFPX", LLAMA_FTYPE_MOSTLY_Q3_0_ROCMFPX, " 3.50 bpw ROCmFPx experimental, ROCm/Vulkan staging", }, + { "Q2_0_ROCMFPX", LLAMA_FTYPE_MOSTLY_Q2_0_ROCMFPX, " 2.50 bpw ROCmFPx S40 codebook + dual UE4M3 scales", }, { "Q6_0_ROCMFPX", LLAMA_FTYPE_MOSTLY_Q6_0_ROCMFPX, " 6.50 bpw ROCmFPx experimental, ROCm/Vulkan staging", }, { "Q8_0_ROCMFPX", LLAMA_FTYPE_MOSTLY_Q8_0_ROCMFPX, " 8.25 bpw ROCmFPx experimental, ROCm/Vulkan staging", }, { "Q3_0_ROCMFPX_AGENT", LLAMA_FTYPE_MOSTLY_Q3_0_ROCMFPX_AGENT, " agent/tool-call coherent ROCmFPx Q3 routing", }, From bc27766d2b373ce93ef2d0ba8727be3e5e0b352f Mon Sep 17 00:00:00 2001 From: ciru-ai Date: Mon, 13 Jul 2026 06:29:06 -0400 Subject: [PATCH 2/4] perf(rocmfpx): add opt-in HY3 adaptive MoE launch --- ggml/src/ggml-cuda/mmvq.cu | 53 +++++++++++++++++++++++++++++++++++--- 1 file changed, 50 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu index 8912af67f..9d244fdec 100644 --- a/ggml/src/ggml-cuda/mmvq.cu +++ b/ggml/src/ggml-cuda/mmvq.cu @@ -120,6 +120,14 @@ #error "GGML_ROCMFPX_MOE_MMVQ_ROWS_PER_BLOCK must be between 1 and 4" #endif +#ifndef GGML_ROCMFPX_HY3_ADAPTIVE_MOE_RPB +#define GGML_ROCMFPX_HY3_ADAPTIVE_MOE_RPB 0 +#endif + +#if GGML_ROCMFPX_HY3_ADAPTIVE_MOE_RPB != 0 && GGML_ROCMFPX_HY3_ADAPTIVE_MOE_RPB != 1 +#error "GGML_ROCMFPX_HY3_ADAPTIVE_MOE_RPB must be 0 or 1" +#endif + #if GGML_ROCMFP4_RDNA35_RPB_WIDE_DUAL != 1 && GGML_ROCMFP4_RDNA35_RPB_WIDE_DUAL != 2 #error "GGML_ROCMFP4_RDNA35_RPB_WIDE_DUAL must be 1 or 2" #endif @@ -1065,8 +1073,8 @@ static void mul_mat_vec_q_switch_fusion( sample_ratio, stride_sample_x, stride_sample_y, stride_sample_dst, ids_stride); } -template -static void mul_mat_vec_q_moe_launch( +template +static void mul_mat_vec_q_moe_launch_rpb( const void * vx, const void * vy, const int32_t * ids, float * dst, const uint32_t ncols_x, const uint3 nchannels_y, const uint32_t nrows_x, const uint32_t stride_row_x, const uint32_t stride_col_y, const uint32_t stride_col_dst, @@ -1074,7 +1082,6 @@ static void mul_mat_vec_q_moe_launch( const uint32_t ncols_dst, const uint32_t ids_stride, const int warp_size, const int nchannels_dst, cudaStream_t stream) { - constexpr int rows_per_block = calc_moe_mmvq_rows_per_block(); const int64_t nblocks_rows = (nrows_x + rows_per_block - 1) / rows_per_block; const dim3 block_nums(nblocks_rows, nchannels_dst); const dim3 block_dims(warp_size, ncols_dst); @@ -1086,6 +1093,46 @@ static void mul_mat_vec_q_moe_launch( ncols_dst, ids_stride); } +template +static void mul_mat_vec_q_moe_launch( + const void * vx, const void * vy, const int32_t * ids, float * dst, + const uint32_t ncols_x, const uint3 nchannels_y, const uint32_t nrows_x, + const uint32_t stride_row_x, const uint32_t stride_col_y, const uint32_t stride_col_dst, + const uint32_t stride_channel_x, const uint32_t stride_channel_y, const uint32_t stride_channel_dst, + const uint32_t ncols_dst, const uint32_t ids_stride, + const int warp_size, const int nchannels_dst, cudaStream_t stream) { + +#if defined(GGML_USE_HIP) && GGML_ROCMFPX_HY3_ADAPTIVE_MOE_RPB + // HY3's 192-expert matrices favor different launch shapes for decode and + // multi-token verification. Keep this opt-in so established FPX recipes + // retain the generic launch policy. + if constexpr (type == GGML_TYPE_Q2_0_ROCMFPX) { + if (ncols_dst == 1) { + if (nrows_x <= 2048) { + mul_mat_vec_q_moe_launch_rpb(vx, vy, ids, dst, ncols_x, nchannels_y, nrows_x, + stride_row_x, stride_col_y, stride_col_dst, stride_channel_x, stride_channel_y, + stride_channel_dst, ncols_dst, ids_stride, warp_size, nchannels_dst, stream); + } else { + mul_mat_vec_q_moe_launch_rpb(vx, vy, ids, dst, ncols_x, nchannels_y, nrows_x, + stride_row_x, stride_col_y, stride_col_dst, stride_channel_x, stride_channel_y, + stride_channel_dst, ncols_dst, ids_stride, warp_size, nchannels_dst, stream); + } + } else { + mul_mat_vec_q_moe_launch_rpb(vx, vy, ids, dst, ncols_x, nchannels_y, nrows_x, + stride_row_x, stride_col_y, stride_col_dst, stride_channel_x, stride_channel_y, + stride_channel_dst, ncols_dst, ids_stride, warp_size, nchannels_dst, stream); + } + return; + } + +#endif + + constexpr int rows_per_block = calc_moe_mmvq_rows_per_block(); + mul_mat_vec_q_moe_launch_rpb(vx, vy, ids, dst, ncols_x, nchannels_y, nrows_x, + stride_row_x, stride_col_y, stride_col_dst, stride_channel_x, stride_channel_y, + stride_channel_dst, ncols_dst, ids_stride, warp_size, nchannels_dst, stream); +} + template static void mul_mat_vec_q_switch_ncols_dst( const void * vx, const void * vy, const int32_t * ids, const ggml_cuda_mm_fusion_args_device fusion, float * dst, From 70ea6a3f49bf5f2e5a97bdf6031090ea3546c3d4 Mon Sep 17 00:00:00 2001 From: ciru-ai Date: Mon, 13 Jul 2026 17:54:04 -0400 Subject: [PATCH 3/4] server: add SSD prompt cache for MTP --- common/arg.cpp | 24 +- common/common.h | 5 + common/speculative.cpp | 307 ++++- common/speculative.h | 6 +- tools/server/README.md | 4 +- tools/server/server-context.cpp | 98 +- tools/server/server-task.cpp | 1092 ++++++++++++++++- tools/server/server-task.h | 132 +- .../tests/unit/test_prompt_cache_disk.py | 320 +++++ tools/server/tests/utils.py | 11 +- 10 files changed, 1919 insertions(+), 80 deletions(-) create mode 100644 tools/server/tests/unit/test_prompt_cache_disk.py diff --git a/common/arg.cpp b/common/arg.cpp index 53a512827..4e0dd1879 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -1357,6 +1357,28 @@ common_params_context common_params_parser_init(common_params & params, llama_ex params.cache_ram_mib = value; } ).set_env("LLAMA_ARG_CACHE_RAM").set_examples({LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI})); + add_opt(common_arg( + {"--cache-disk"}, "PATH", + "base directory for the automatic SSD-backed prompt cache (default: disabled); " + "the server creates and removes an owner-only run directory below PATH", + [](common_params & params, const std::string & value) { + if (value.empty()) { + throw std::invalid_argument("cache disk path must not be empty"); + } + params.cache_disk_path = value; + } + ).set_env("LLAMA_ARG_CACHE_DISK").set_examples({LLAMA_EXAMPLE_SERVER})); + add_opt(common_arg( + {"--cache-disk-limit"}, "N", + string_format("maximum SSD-backed prompt-cache size in MiB when --cache-disk is set " + "(default: %d, 0 - disable)", params.cache_disk_limit_mib), + [](common_params & params, int value) { + if (value < 0) { + throw std::invalid_argument("cache disk limit must be non-negative"); + } + params.cache_disk_limit_mib = value; + } + ).set_env("LLAMA_ARG_CACHE_DISK_LIMIT").set_examples({LLAMA_EXAMPLE_SERVER})); add_opt(common_arg( {"-kvu", "--kv-unified"}, {"-no-kvu", "--no-kv-unified"}, @@ -1368,7 +1390,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex add_opt(common_arg( {"--cache-idle-slots"}, {"--no-cache-idle-slots"}, - "save idle slots to the prompt cache on new task, and clear them when using unified KV (default: enabled, requires cache-ram)", + "save idle slots to the prompt cache on new task, and clear them when using unified KV (default: enabled, requires cache RAM or disk)", [](common_params & params, bool value) { params.cache_idle_slots = value; } diff --git a/common/common.h b/common/common.h index 9462326f4..34bde03fe 100644 --- a/common/common.h +++ b/common/common.h @@ -610,6 +610,7 @@ struct common_params { int32_t n_ctx_checkpoints = 32; // max number of context checkpoints per slot int32_t checkpoint_every_nt = 8192; // make a checkpoint every n tokens during prefill int32_t cache_ram_mib = 8192; // -1 = no limit, 0 - disable, 1 = 1 MiB, etc. + int32_t cache_disk_limit_mib = 8192; // disk prompt-cache limit when cache_disk_path is set std::string hostname = "127.0.0.1"; std::string public_path = ""; // NOLINT @@ -647,6 +648,10 @@ struct common_params { // enable built-in tools std::vector server_tools; + // Optional base directory for the automatic SSD-backed prompt cache. The + // server creates and owns a private per-process directory below this path. + std::string cache_disk_path; + // router server configs std::string models_dir = ""; // directory containing models for the router server std::string models_preset = ""; // directory containing model presets for the router server diff --git a/common/speculative.cpp b/common/speculative.cpp index 8924a12b0..f66709b23 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -178,7 +178,9 @@ struct common_speculative_impl { // (optional) serialize/restore per-seq internal state (e.g. eagle3's deferred boundary). virtual bool get_state(llama_seq_id /*seq_id*/, std::vector & /*data*/) const { return false; } - virtual void set_state(llama_seq_id /*seq_id*/, const std::vector & /*data*/) {} + virtual bool set_state(llama_seq_id /*seq_id*/, const std::vector & /*data*/) { return true; } + virtual bool state_required() const { return false; } + virtual void shift_state(llama_seq_id /*seq_id*/, llama_pos /*delta*/) {} // true if this implementation requires the target context to extract post-norm embeddings virtual bool need_embd() const = 0; @@ -912,15 +914,23 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl { return true; } - void set_state(llama_seq_id seq_id, const std::vector & data) override { - if (!need_boundary_stash()) { - return; - } + bool set_state(llama_seq_id seq_id, const std::vector & data) override { if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) { - return; + return false; + } + if (data.empty()) { + pending_pos_last[seq_id] = -1; + std::fill(pending_g_last[seq_id].begin(), pending_g_last[seq_id].end(), 0.0f); + verify_g[seq_id].clear(); + verify_pos_first[seq_id] = -1; + verify_g_rows[seq_id] = 0; + return !need_boundary_stash(); + } + if (!need_boundary_stash()) { + return true; } if (data.size() != sizeof(llama_pos) + (size_t) n_embd_dec * sizeof(float)) { - return; + return false; } llama_pos pos = -1; @@ -929,6 +939,24 @@ struct common_speculative_impl_draft_eagle3 : public common_speculative_impl { pending_pos_last[seq_id] = pos; pending_g_last[seq_id].resize(n_embd_dec); std::memcpy(pending_g_last[seq_id].data(), data.data() + sizeof(llama_pos), (size_t) n_embd_dec * sizeof(float)); + return true; + } + + bool state_required() const override { + return need_boundary_stash(); + } + + void shift_state(llama_seq_id seq_id, llama_pos delta) override { + if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq || delta == 0) { + return; + } + + if (pending_pos_last[seq_id] >= 0) { + pending_pos_last[seq_id] += delta; + } + if (verify_pos_first[seq_id] >= 0) { + verify_pos_first[seq_id] += delta; + } } bool need_embd() const override { @@ -1263,6 +1291,18 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { // The last h-row of one process() call needs the first token of the NEXT // call to pair with, so it's stashed here until that next call fires. std::vector> pending_h; // [n_seq][n_embd] + std::vector> pending_h_prev; + std::vector pending_h_valid; + std::vector pending_h_prev_valid; + std::vector pending_h_pos; + std::vector pending_h_prev_pos; + + // Boundary selected for the most recent process() batch. This is needed + // when verification accepts row zero, and when a one-token replay advances + // the current boundary while retaining the prior one. + std::vector> process_boundary_h; + std::vector process_boundary_valid; + std::vector process_boundary_pos; std::vector i_batch_beg; std::vector i_batch_end; @@ -1271,6 +1311,7 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { // Row 0 corresponds to the sampled token, row N to the Nth accepted draft token. std::vector> verify_h; std::vector verify_h_rows; + std::vector verify_pos_first; // Per-seq draft length from the last draft() call, used in accept() to // roll back ctx_dft's recurrent state past the AR draft's redundant @@ -1356,6 +1397,15 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { } pending_h.assign(n_seq, std::vector(n_embd, 0.0f)); + pending_h_prev.assign(n_seq, std::vector(n_embd, 0.0f)); + pending_h_valid.assign(n_seq, 0); + pending_h_prev_valid.assign(n_seq, 0); + pending_h_pos.assign(n_seq, -1); + pending_h_prev_pos.assign(n_seq, -1); + + process_boundary_h.assign(n_seq, std::vector(n_embd, 0.0f)); + process_boundary_valid.assign(n_seq, 0); + process_boundary_pos.assign(n_seq, -1); i_last.assign(n_seq, -1); i_batch_beg.assign(n_seq, -1); @@ -1366,6 +1416,7 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { h.reserve((size_t) std::max(1, this->params.n_max) * n_embd); } verify_h_rows.assign(n_seq, 0); + verify_pos_first.assign(n_seq, -1); last_n_drafted.assign(n_seq, 0); drafting.assign(n_seq, 0); @@ -1423,6 +1474,9 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { std::fill(i_batch_beg.begin(), i_batch_beg.end(), -1); std::fill(i_batch_end.begin(), i_batch_end.end(), -1); std::fill(verify_h_rows.begin(), verify_h_rows.end(), 0); + std::fill(verify_pos_first.begin(), verify_pos_first.end(), -1); + std::fill(process_boundary_valid.begin(), process_boundary_valid.end(), 0); + std::fill(process_boundary_pos.begin(), process_boundary_pos.end(), -1); for (int k = 0; k < n_tokens; ++k) { GGML_ASSERT(batch_in.n_seq_id[k] == 1); @@ -1478,7 +1532,37 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { continue; } - set_h(i_batch_beg[seq_id], pending_h[seq_id].data()); + const llama_pos pos_needed = batch_in.pos[i_batch_beg[seq_id]] - 1; + const float * h_boundary = nullptr; + if (pending_h_valid[seq_id] && pending_h_pos[seq_id] == pos_needed) { + h_boundary = pending_h[seq_id].data(); + } else if (pending_h_prev_valid[seq_id] && pending_h_prev_pos[seq_id] == pos_needed) { + h_boundary = pending_h_prev[seq_id].data(); + } else if (pos_needed < 0) { + // A new request can reach BOS without prompt_clear(), notably + // through an explicit id_slot or with prompt caching disabled. + // Never reuse the prior request's deferred MTP boundary there. + std::fill(pending_h[seq_id].begin(), pending_h[seq_id].end(), 0.0f); + std::fill(pending_h_prev[seq_id].begin(), pending_h_prev[seq_id].end(), 0.0f); + pending_h_valid[seq_id] = 0; + pending_h_prev_valid[seq_id] = 0; + pending_h_pos[seq_id] = -1; + pending_h_prev_pos[seq_id] = -1; + h_boundary = pending_h[seq_id].data(); + } else { + LOG_ERR("%s: missing MTP boundary for seq_id=%d pos=%d (current=%d/%d previous=%d/%d)\n", + __func__, (int) seq_id, (int) pos_needed, + (int) pending_h_pos[seq_id], (int) pending_h_valid[seq_id], + (int) pending_h_prev_pos[seq_id], (int) pending_h_prev_valid[seq_id]); + return false; + } + + set_h(i_batch_beg[seq_id], h_boundary); + if (pos_needed >= 0) { + std::memcpy(process_boundary_h[seq_id].data(), h_boundary, row_bytes); + process_boundary_valid[seq_id] = 1; + process_boundary_pos[seq_id] = pos_needed; + } } auto * mem_dft = llama_get_memory(ctx_dft); @@ -1523,13 +1607,33 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { const int32_t n_rows = i_batch_end[seq_id] - i_batch_beg[seq_id] + 1; const float * h_seq = h_tgt + (size_t) i_batch_beg[seq_id] * n_embd; + const llama_pos pos_last = batch_in.pos[i_batch_end[seq_id]]; + + if (n_rows > 1) { + const float * h_prev = h_seq + (size_t) (n_rows - 2) * n_embd; + std::memcpy(pending_h_prev[seq_id].data(), h_prev, row_bytes); + pending_h_prev_valid[seq_id] = 1; + pending_h_prev_pos[seq_id] = batch_in.pos[i_batch_end[seq_id] - 1]; + } else if (process_boundary_valid[seq_id]) { + std::memcpy(pending_h_prev[seq_id].data(), process_boundary_h[seq_id].data(), row_bytes); + pending_h_prev_valid[seq_id] = 1; + pending_h_prev_pos[seq_id] = process_boundary_pos[seq_id]; + } else { + std::fill(pending_h_prev[seq_id].begin(), pending_h_prev[seq_id].end(), 0.0f); + pending_h_prev_valid[seq_id] = 0; + pending_h_prev_pos[seq_id] = -1; + } + if (last_n_drafted[seq_id] == 0) { const float * h_last = h_seq + (size_t) (n_rows - 1) * n_embd; std::memcpy(pending_h[seq_id].data(), h_last, row_bytes); + pending_h_valid[seq_id] = 1; + pending_h_pos[seq_id] = pos_last; continue; } verify_h_rows[seq_id] = n_rows; + verify_pos_first[seq_id] = batch_in.pos[i_batch_beg[seq_id]]; const size_t n_verify_floats = (size_t) (n_rows - 1) * n_embd; if (verify_h[seq_id].size() < n_verify_floats) { verify_h[seq_id].resize(n_verify_floats); @@ -1540,6 +1644,8 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { const float * h_last = h_seq + (size_t) (n_rows - 1) * n_embd; std::memcpy(pending_h[seq_id].data(), h_last, row_bytes); + pending_h_valid[seq_id] = 1; + pending_h_pos[seq_id] = pos_last; } return true; @@ -1565,6 +1671,15 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { last_n_drafted[seq_id] = 0; + const llama_pos pos_needed = dp.n_past - 1; + if (!pending_h_valid[seq_id] || pending_h_pos[seq_id] != pos_needed) { + LOG_WRN("%s: disabling MTP draft for seq_id=%d: boundary pos=%d/%d, needed=%d\n", + __func__, (int) seq_id, (int) pending_h_pos[seq_id], + (int) pending_h_valid[seq_id], (int) pos_needed); + dp.drafting = false; + continue; + } + n_drafting++; drafting[seq_id] = 1; common_sampler_reset(smpls[seq_id].get()); @@ -1724,13 +1839,151 @@ struct common_speculative_state_draft_mtp : public common_speculative_impl { } const int32_t i_h = std::min(n_accepted, n_rows - 1); + const size_t row_bytes = (size_t) n_embd * sizeof(float); if (i_h != n_rows - 1) { - const size_t row_bytes = (size_t) n_embd * sizeof(float); std::memcpy(pending_h[seq_id].data(), verify_h[seq_id].data() + (size_t) i_h * n_embd, row_bytes); } + pending_h_valid[seq_id] = 1; + pending_h_pos[seq_id] = verify_pos_first[seq_id] + i_h; + + if (i_h == 0) { + if (process_boundary_valid[seq_id]) { + std::memcpy(pending_h_prev[seq_id].data(), process_boundary_h[seq_id].data(), row_bytes); + pending_h_prev_valid[seq_id] = 1; + pending_h_prev_pos[seq_id] = process_boundary_pos[seq_id]; + } else { + std::fill(pending_h_prev[seq_id].begin(), pending_h_prev[seq_id].end(), 0.0f); + pending_h_prev_valid[seq_id] = 0; + pending_h_prev_pos[seq_id] = -1; + } + } else { + std::memcpy(pending_h_prev[seq_id].data(), verify_h[seq_id].data() + (size_t) (i_h - 1) * n_embd, row_bytes); + pending_h_prev_valid[seq_id] = 1; + pending_h_prev_pos[seq_id] = verify_pos_first[seq_id] + i_h - 1; + } last_n_drafted[seq_id] = 0; } + static constexpr uint32_t MTP_STATE_MAGIC = 0x3250544d; // "MTP2" in little-endian byte order + static constexpr uint16_t MTP_STATE_VERSION = 2; + static constexpr uint16_t MTP_STATE_CURRENT = 1u << 0; + static constexpr uint16_t MTP_STATE_PREVIOUS = 1u << 1; + static constexpr size_t MTP_STATE_HEADER = sizeof(uint32_t) + 2*sizeof(uint16_t) + sizeof(uint32_t) + 2*sizeof(llama_pos); + + void reset_seq_state(llama_seq_id seq_id) { + if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) { + return; + } + + std::fill(pending_h[seq_id].begin(), pending_h[seq_id].end(), 0.0f); + std::fill(pending_h_prev[seq_id].begin(), pending_h_prev[seq_id].end(), 0.0f); + pending_h_valid[seq_id] = 0; + pending_h_prev_valid[seq_id] = 0; + pending_h_pos[seq_id] = -1; + pending_h_prev_pos[seq_id] = -1; + std::fill(process_boundary_h[seq_id].begin(), process_boundary_h[seq_id].end(), 0.0f); + process_boundary_valid[seq_id] = 0; + process_boundary_pos[seq_id] = -1; + verify_h[seq_id].clear(); + verify_h_rows[seq_id] = 0; + verify_pos_first[seq_id] = -1; + last_n_drafted[seq_id] = 0; + drafting[seq_id] = 0; + i_last[seq_id] = -1; + if (chain_heads) { + chain_h[seq_id].clear(); + } + } + + bool get_state(llama_seq_id seq_id, std::vector & data) const override { + data.clear(); + if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq || !pending_h_valid[seq_id]) { + return false; + } + + const uint16_t flags = MTP_STATE_CURRENT | + (pending_h_prev_valid[seq_id] ? MTP_STATE_PREVIOUS : 0); + const uint32_t width = (uint32_t) n_embd; + const size_t row_bytes = (size_t) n_embd * sizeof(float); + data.resize(MTP_STATE_HEADER + 2*row_bytes); + + size_t off = 0; + std::memcpy(data.data() + off, &MTP_STATE_MAGIC, sizeof(MTP_STATE_MAGIC)); off += sizeof(MTP_STATE_MAGIC); + std::memcpy(data.data() + off, &MTP_STATE_VERSION, sizeof(MTP_STATE_VERSION)); off += sizeof(MTP_STATE_VERSION); + std::memcpy(data.data() + off, &flags, sizeof(flags)); off += sizeof(flags); + std::memcpy(data.data() + off, &width, sizeof(width)); off += sizeof(width); + std::memcpy(data.data() + off, &pending_h_pos[seq_id], sizeof(llama_pos)); off += sizeof(llama_pos); + std::memcpy(data.data() + off, &pending_h_prev_pos[seq_id], sizeof(llama_pos)); off += sizeof(llama_pos); + std::memcpy(data.data() + off, pending_h[seq_id].data(), row_bytes); off += row_bytes; + std::memcpy(data.data() + off, pending_h_prev[seq_id].data(), row_bytes); + return true; + } + + bool set_state(llama_seq_id seq_id, const std::vector & data) override { + if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq) { + return false; + } + + reset_seq_state(seq_id); + if (data.empty() || data.size() < MTP_STATE_HEADER) { + return false; + } + + uint32_t magic = 0; + uint16_t version = 0; + uint16_t flags = 0; + uint32_t width = 0; + llama_pos pos_current = -1; + llama_pos pos_previous = -1; + size_t off = 0; + std::memcpy(&magic, data.data() + off, sizeof(magic)); off += sizeof(magic); + std::memcpy(&version, data.data() + off, sizeof(version)); off += sizeof(version); + std::memcpy(&flags, data.data() + off, sizeof(flags)); off += sizeof(flags); + std::memcpy(&width, data.data() + off, sizeof(width)); off += sizeof(width); + std::memcpy(&pos_current, data.data() + off, sizeof(pos_current)); off += sizeof(pos_current); + std::memcpy(&pos_previous, data.data() + off, sizeof(pos_previous)); off += sizeof(pos_previous); + + const size_t row_bytes = (size_t) n_embd * sizeof(float); + if (magic != MTP_STATE_MAGIC || version != MTP_STATE_VERSION || + (flags & MTP_STATE_CURRENT) == 0 || (flags & ~(MTP_STATE_CURRENT | MTP_STATE_PREVIOUS)) != 0 || + width != (uint32_t) n_embd || pos_current < 0 || + ((flags & MTP_STATE_PREVIOUS) != 0 && pos_previous < 0) || + data.size() != MTP_STATE_HEADER + 2*row_bytes) { + return false; + } + + std::memcpy(pending_h[seq_id].data(), data.data() + off, row_bytes); off += row_bytes; + std::memcpy(pending_h_prev[seq_id].data(), data.data() + off, row_bytes); + pending_h_valid[seq_id] = 1; + pending_h_prev_valid[seq_id] = (flags & MTP_STATE_PREVIOUS) != 0; + pending_h_pos[seq_id] = pos_current; + pending_h_prev_pos[seq_id] = pending_h_prev_valid[seq_id] ? pos_previous : -1; + return true; + } + + bool state_required() const override { + return true; + } + + void shift_state(llama_seq_id seq_id, llama_pos delta) override { + if (seq_id < 0 || seq_id >= (llama_seq_id) n_seq || delta == 0) { + return; + } + + if (pending_h_valid[seq_id]) { + pending_h_pos[seq_id] += delta; + } + if (pending_h_prev_valid[seq_id]) { + pending_h_prev_pos[seq_id] += delta; + } + if (process_boundary_valid[seq_id]) { + process_boundary_pos[seq_id] += delta; + } + if (verify_h_rows[seq_id] > 0 && verify_pos_first[seq_id] >= 0) { + verify_pos_first[seq_id] += delta; + } + } + bool need_embd() const override { return false; } @@ -2610,13 +2863,45 @@ bool common_speculative_get_state(common_speculative * spec, llama_seq_id seq_id return false; } -void common_speculative_set_state(common_speculative * spec, llama_seq_id seq_id, const std::vector & data) { +bool common_speculative_set_state(common_speculative * spec, llama_seq_id seq_id, const std::vector & data) { if (spec == nullptr) { + return true; + } + + bool ok = true; + for (auto & impl : spec->impls) { + const bool restored = impl->set_state(seq_id, data); + ok = ok && (!impl->state_required() || restored); + } + + if (seq_id >= 0 && seq_id < (llama_seq_id) spec->dparams.size()) { + spec->dparams[seq_id].drafting = false; + } + + return ok; +} + +bool common_speculative_state_required(const common_speculative * spec) { + if (spec == nullptr) { + return false; + } + + for (const auto & impl : spec->impls) { + if (impl->state_required()) { + return true; + } + } + + return false; +} + +void common_speculative_shift_state(common_speculative * spec, llama_seq_id seq_id, llama_pos delta) { + if (spec == nullptr || delta == 0) { return; } for (auto & impl : spec->impls) { - impl->set_state(seq_id, data); + impl->shift_state(seq_id, delta); } } diff --git a/common/speculative.h b/common/speculative.h index 82fb9cc80..6ce82e5af 100644 --- a/common/speculative.h +++ b/common/speculative.h @@ -70,7 +70,11 @@ void common_speculative_accept(common_speculative * spec, llama_seq_id, uint16_t // (optional) get/set internal state bool common_speculative_get_state(common_speculative * spec, llama_seq_id seq_id, std::vector & data); -void common_speculative_set_state(common_speculative * spec, llama_seq_id seq_id, const std::vector & data); +bool common_speculative_set_state(common_speculative * spec, llama_seq_id seq_id, const std::vector & data); +bool common_speculative_state_required(const common_speculative * spec); + +// rebase per-sequence positions after the corresponding target/draft contexts shift +void common_speculative_shift_state(common_speculative * spec, llama_seq_id seq_id, llama_pos delta); // print statistics about the speculative decoding void common_speculative_print_stats(const common_speculative * spec); diff --git a/tools/server/README.md b/tools/server/README.md index 4a2ee6df1..5e10c0d16 100644 --- a/tools/server/README.md +++ b/tools/server/README.md @@ -166,8 +166,10 @@ For the full list of features, please refer to [server's changelog](https://gith | `-ctxcp, --ctx-checkpoints, --swa-checkpoints N` | max number of context checkpoints to create per slot (default: 32)[(more info)](https://github.com/ggml-org/llama.cpp/pull/15293)
(env: LLAMA_ARG_CTX_CHECKPOINTS) | | `-cpent, --checkpoint-every-n-tokens N` | create a checkpoint every n tokens during prefill (processing), -1 to disable (default: 8192)
(env: LLAMA_ARG_CHECKPOINT_EVERY_NT) | | `-cram, --cache-ram N` | set the maximum cache size in MiB (default: 8192, -1 - no limit, 0 - disable)[(more info)](https://github.com/ggml-org/llama.cpp/pull/16391)
(env: LLAMA_ARG_CACHE_RAM) | +| `--cache-disk PATH` | base directory for the automatic SSD-backed prompt cache (default: disabled); target and MTP draft states are streamed to an owner-only per-server run directory and removed at shutdown
(env: LLAMA_ARG_CACHE_DISK) | +| `--cache-disk-limit N` | maximum SSD-backed prompt-cache size in MiB when `--cache-disk` is set (default: 8192, 0 - disable)
(env: LLAMA_ARG_CACHE_DISK_LIMIT) | | `-kvu, --kv-unified, -no-kvu, --no-kv-unified` | use single unified KV buffer shared across all sequences (default: enabled if number of slots is auto)
(env: LLAMA_ARG_KV_UNIFIED) | -| `--cache-idle-slots, --no-cache-idle-slots` | save idle slots to the prompt cache on new task, and clear them when using unified KV (default: enabled, requires cache-ram)
(env: LLAMA_ARG_CACHE_IDLE_SLOTS) | +| `--cache-idle-slots, --no-cache-idle-slots` | save idle slots to the prompt cache on new task, and clear them when using unified KV (default: enabled, requires cache RAM or disk)
(env: LLAMA_ARG_CACHE_IDLE_SLOTS) | | `--context-shift, --no-context-shift` | whether to use context shift on infinite text generation (default: disabled)
(env: LLAMA_ARG_CONTEXT_SHIFT) | | `-r, --reverse-prompt PROMPT` | halt generation at PROMPT, return control in interactive mode | | `-sp, --special` | special tokens output enabled (default: false) | diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index f71afe503..fd36d74fc 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -128,21 +128,41 @@ struct server_slot { SRV_WRN(" - saving prompt with length %d, total state size = %.3f MiB (draft: %.3f MiB)\n", (int) prompt.tokens.size(), cur_size / (1024.0 * 1024.0), cur_size_dft / (1024.0 * 1024.0)); - auto * cur = prompt_cache.alloc(prompt, cur_size_tgt, cur_size_dft); - if (cur == nullptr) { + std::vector state_spec; + const bool spec_state_required = common_speculative_state_required(spec); + const bool have_spec_state = common_speculative_get_state(spec, id, state_spec); + if (spec_state_required && !have_spec_state) { + SLT_WRN(*this, "%s", "skipping prompt cache save: required speculative state is unavailable\n"); return false; } - llama_state_seq_get_data_ext(ctx_tgt, cur->data.main.data(), cur_size_tgt, id, LLAMA_STATE_SEQ_FLAGS_NONE); - if (ctx_dft) { - llama_state_seq_get_data_ext(ctx_dft, cur->data.drft.data(), cur_size_dft, id, LLAMA_STATE_SEQ_FLAGS_NONE); - } - - return true; + return prompt_cache.save(prompt, ctx_tgt, ctx_dft, id, state_spec); } bool prompt_load(server_prompt_cache & prompt_cache, const server_tokens & tokens) { - bool res = prompt_cache.load(prompt, tokens, ctx_tgt, ctx_dft, id); + const bool spec_state_required = common_speculative_state_required(spec); + bool cache_hit = false; + uint64_t disk_entry_id = 0; + bool res = prompt_cache.load(prompt, tokens, ctx_tgt, ctx_dft, id, spec_state_required, &cache_hit, &disk_entry_id); + if (res && cache_hit) { + if (spec_state_required && prompt.data.spec.empty()) { + SLT_WRN(*this, "%s", "failed to load required speculative state from prompt cache\n"); + res = false; + } else if (!prompt.data.spec.empty() && !common_speculative_set_state(spec, id, prompt.data.spec)) { + SLT_WRN(*this, "%s", "failed to validate speculative state from prompt cache\n"); + res = false; + } + + prompt.data.spec.clear(); + prompt.data.spec.shrink_to_fit(); + } + if (disk_entry_id != 0) { + if (res) { + prompt_cache.accept_disk_load(disk_entry_id); + } else { + prompt_cache.reject_disk_load(disk_entry_id, "spec-state-rejected"); + } + } if (!res) { SLT_WRN(*this, "%s", "failed to load prompt from cache\n"); } @@ -161,6 +181,7 @@ struct server_slot { if (ctx_dft) { common_context_seq_rm(ctx_dft, id, -1, -1); } + common_speculative_set_state(spec, id, {}); prompt.tokens.clear(); } @@ -1008,17 +1029,29 @@ struct server_context_impl { batch = llama_batch_init(std::max(n_batch, params_base.n_parallel), 0, 1); } - if (params_base.cache_ram_mib != 0) { + const bool cache_disk_enabled = !params_base.cache_disk_path.empty() && params_base.cache_disk_limit_mib > 0; + + if (params_base.cache_ram_mib != 0 || cache_disk_enabled) { if (params_base.cache_ram_mib < 0) { - SRV_INF("prompt cache is enabled, size limit: %s\n", "no limit"); + SRV_INF("prompt cache RAM enabled: limit=%s\n", "unlimited"); + } else if (params_base.cache_ram_mib > 0) { + SRV_INF("prompt cache RAM enabled: limit_mib=%d\n", params_base.cache_ram_mib); } else { - SRV_INF("prompt cache is enabled, size limit: %d MiB\n", params_base.cache_ram_mib); + SRV_INF("%s", "prompt cache RAM disabled: limit_mib=0\n"); + } + + if (cache_disk_enabled) { + SRV_INF("prompt cache SSD enabled: path=%s limit_mib=%d target_and_draft=true\n", + params_base.cache_disk_path.c_str(), params_base.cache_disk_limit_mib); } - SRV_INF("%s", "use `--cache-ram 0` to disable the prompt cache\n"); - prompt_cache = std::make_unique(params_base.cache_ram_mib, n_ctx); + prompt_cache = std::make_unique( + params_base.cache_ram_mib, + n_ctx, + params_base.cache_disk_path, + params_base.cache_disk_limit_mib); } else { - SRV_INF("%s", "prompt cache is disabled - use `--cache-ram N` to enable it\n"); + SRV_INF("%s", "prompt cache is disabled - use `--cache-ram N` or `--cache-disk PATH` to enable it\n"); } SRV_INF("%s", "for more info see https://github.com/ggml-org/llama.cpp/pull/16391\n"); @@ -1067,8 +1100,8 @@ struct server_context_impl { metrics.init(); if (params_base.cache_idle_slots) { - if (params_base.cache_ram_mib == 0) { - SRV_WRN("%s", "--cache-idle-slots requires --cache-ram, disabling\n"); + if (!prompt_cache) { + SRV_WRN("%s", "--cache-idle-slots requires --cache-ram or --cache-disk, disabling\n"); params_base.cache_idle_slots = false; } else { if (params_base.kv_unified) { @@ -1235,6 +1268,8 @@ struct server_context_impl { if (!ret->prompt_load(*prompt_cache, task.tokens)) { ret->prompt_clear(false); + SRV_INF("prompt cache cold fallback: slot=%d reason=target-draft-restore-rejected target_and_draft_cleared=true\n", + ret->id); } prompt_cache->update(); @@ -1983,14 +2018,17 @@ struct server_context_impl { if (!slot.is_processing()) { SLT_INF(slot, "%s", "saving idle slot to prompt cache\n"); - if (slot.prompt_save(*prompt_cache)) { + const bool safe_to_clear = slot.prompt_save(*prompt_cache); + if (safe_to_clear) { SLT_DBG(slot, "%s", "__TEST_TAG_CACHE_IDLE_SLOT__\n"); prompt_cache->update(); - } - if (params_base.kv_unified) { - // [TAG_IDLE_SLOT_CLEAR] - slot.prompt_clear(false); + if (params_base.kv_unified) { + // [TAG_IDLE_SLOT_CLEAR] + slot.prompt_clear(false); + } + } else if (!slot.prompt.tokens.empty()) { + SLT_WRN(slot, "%s", "preserving idle slot because prompt cache save was not safe\n"); } } } @@ -2287,6 +2325,8 @@ struct server_context_impl { common_context_seq_add(ctx_dft.get(), slot.id, n_keep + n_discard, slot.prompt.tokens.pos_next(), -n_discard); } + common_speculative_shift_state(spec.get(), slot.id, -n_discard); + // add generated tokens to cache // ref: https://github.com/ggml-org/llama.cpp/pull/16818#discussion_r2473269481 { @@ -2633,6 +2673,20 @@ struct server_context_impl { n_past = 0; } + // MTP carries only the endpoint and immediately preceding target hidden + // boundaries. Arbitrary partial-prefix rollback cannot be reconstructed + // from target/draft KV state, so reprocess cold instead of pairing a token + // with the wrong hidden row. Full-prefix extension and the exact-hit + // one-token replay below remain supported. + if (common_speculative_state_required(spec.get()) && + n_past > 0 && n_past < slot.prompt.n_tokens()) { + SLT_INF(slot, + "prompt cache cold fallback: reason=spec-boundary-mismatch lcp=%d cached_tokens=%d request_tokens=%d\n", + n_past, slot.prompt.n_tokens(), slot.task->n_tokens()); + n_past = 0; + common_speculative_set_state(spec.get(), slot.id, {}); + } + llama_pos pos_next = slot.prompt.tokens.pos_next(n_past); // ref: https://github.com/ggml-org/llama.cpp/pull/24110 diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp index 0292614f4..634213fd9 100644 --- a/tools/server/server-task.cpp +++ b/tools/server/server-task.cpp @@ -10,6 +10,22 @@ #include "speculative.h" #include "server-common.h" +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#if !defined(_WIN32) +#include +#include +#include +#endif + using json = nlohmann::ordered_json; // @@ -1962,6 +1978,319 @@ json server_task_result_apply_lora::to_json() { // // server_prompt_cache // +namespace { + +namespace fs = std::filesystem; + +constexpr const char * SERVER_PROMPT_CACHE_DISK_NAMESPACE = ".llama-prompt-cache-v1"; +constexpr const char * SERVER_PROMPT_CACHE_OWNER_MAGIC = "llama.cpp automatic prompt cache v1"; + +static bool server_prompt_cache_disk_owned(const fs::path & path) { + std::ifstream owner(path / ".owner"); + std::string magic; + return owner.good() && std::getline(owner, magic) && magic == SERVER_PROMPT_CACHE_OWNER_MAGIC; +} + +static bool server_prompt_cache_disk_remove_file(const std::string & path) { + if (path.empty()) { + return true; + } + + std::error_code ec; + fs::remove(path, ec); + if (ec) { + SRV_WRN("prompt cache disk cleanup failed: path=%s error=%s\n", path.c_str(), ec.message().c_str()); + return false; + } + + // A missing file already satisfies the desired postcondition. This also + // lets a later retry finish a pair after an earlier partial removal. + return true; +} + +static bool server_prompt_cache_disk_size_exact( + const std::string & path, + size_t expected, + size_t * actual_out = nullptr) { + if (path.empty()) { + if (actual_out != nullptr) { + *actual_out = 0; + } + return expected == 0; + } + + std::error_code ec; + const uintmax_t actual = fs::file_size(path, ec); + if (ec || actual > std::numeric_limits::max()) { + if (actual_out != nullptr) { + *actual_out = 0; + } + return false; + } + + if (actual_out != nullptr) { + *actual_out = (size_t) actual; + } + return (size_t) actual == expected; +} + +// llama_state_seq_save_file() closes the file before returning. Reopen it to +// force dirty pages to stable storage and immediately mark the cold state as +// reclaimable. This avoids replacing anonymous cache pressure with several GiB +// of sticky buffered page cache on UMA systems. +static bool server_prompt_cache_disk_flush_and_drop(const std::string & path, bool durable) { +#if !defined(_WIN32) + const int fd = open(path.c_str(), (durable ? O_RDWR : O_RDONLY) | O_CLOEXEC); + if (fd < 0) { + SRV_ERR("prompt cache disk open failed: path=%s error=%s\n", path.c_str(), std::strerror(errno)); + return false; + } + + bool ok = true; + if (durable && fdatasync(fd) != 0) { + SRV_ERR("prompt cache disk fdatasync failed: path=%s error=%s\n", path.c_str(), std::strerror(errno)); + ok = false; + } + +#if defined(POSIX_FADV_DONTNEED) + const int err = posix_fadvise(fd, 0, 0, POSIX_FADV_DONTNEED); + if (err != 0) { + SRV_WRN("prompt cache disk fadvise failed: path=%s error=%s\n", path.c_str(), std::strerror(err)); + } +#endif + + close(fd); + return ok; +#else + GGML_UNUSED(path); + GGML_UNUSED(durable); + return true; +#endif +} + +static bool server_prompt_cache_disk_sync_dir(const std::string & path) { +#if !defined(_WIN32) + const int fd = open(path.c_str(), O_RDONLY | O_DIRECTORY | O_CLOEXEC); + if (fd < 0) { + SRV_ERR("prompt cache disk directory open failed: path=%s error=%s\n", path.c_str(), std::strerror(errno)); + return false; + } + + const bool ok = fsync(fd) == 0; + if (!ok) { + SRV_ERR("prompt cache disk directory fsync failed: path=%s error=%s\n", path.c_str(), std::strerror(errno)); + } + close(fd); + return ok; +#else + GGML_UNUSED(path); + return true; +#endif +} + +static bool server_prompt_cache_tokens_equal(const server_tokens & expected, const llama_tokens & actual) { + return expected.get_tokens() == actual; +} + +} // namespace + +server_prompt_cache::server_prompt_cache( + int32_t limit_size_mib, + size_t limit_tokens, + const std::string & disk_base_path, + int32_t disk_limit_size_mib) { + ram_enabled = limit_size_mib != 0; + limit_size = 1024ull*1024ull*(limit_size_mib < 0 ? 0 : limit_size_mib); + this->limit_tokens = limit_tokens; + + if (disk_base_path.empty() || disk_limit_size_mib <= 0) { + return; + } + + disk_limit_size = 1024ull*1024ull*disk_limit_size_mib; + + std::error_code ec; + fs::path base = fs::absolute(disk_base_path, ec); + if (ec) { + throw std::runtime_error("unable to resolve prompt cache disk path '" + disk_base_path + "': " + ec.message()); + } + + fs::create_directories(base, ec); + if (ec || !fs::is_directory(base)) { + throw std::runtime_error("unable to create prompt cache disk path '" + base.string() + "': " + ec.message()); + } + + const fs::path cache_root = base / SERVER_PROMPT_CACHE_DISK_NAMESPACE; + fs::create_directories(cache_root, ec); + if (ec || !fs::is_directory(cache_root)) { + throw std::runtime_error("unable to create prompt cache namespace '" + cache_root.string() + "': " + ec.message()); + } + fs::permissions(cache_root, fs::perms::owner_all, fs::perm_options::replace, ec); + if (ec) { + throw std::runtime_error("unable to secure prompt cache namespace '" + cache_root.string() + "': " + ec.message()); + } + + // An OOM/SIGKILL cannot run the destructor. Each run therefore holds an + // advisory lock in a magic-marked directory. A later server removes only + // marked run-* directories whose lock is no longer held. + for (const auto & entry : fs::directory_iterator(cache_root, ec)) { + if (ec) { + break; + } + const auto name = entry.path().filename().string(); + const bool is_run_dir = name.rfind("run-", 0) == 0; + const bool is_deleting_dir = name.rfind(".deleting-run-", 0) == 0; + if (!entry.is_directory() || (!is_run_dir && !is_deleting_dir) || !server_prompt_cache_disk_owned(entry.path())) { + continue; + } + +#if !defined(_WIN32) + const fs::path lock_path = entry.path() / ".lock"; + const int fd = open(lock_path.c_str(), O_RDWR | O_CLOEXEC); + if (fd < 0) { + continue; + } + const bool stale = flock(fd, LOCK_EX | LOCK_NB) == 0; + if (stale) { + flock(fd, LOCK_UN); + } + close(fd); + if (!stale) { + continue; + } +#else + // Without an advisory-lock primitive, preserve old directories rather + // than risk deleting a live cache owned by another process. + continue; +#endif + + const auto stale_path = entry.path().string(); + std::error_code rm_ec; + const auto removed = fs::remove_all(entry.path(), rm_ec); + if (!rm_ec) { + SRV_INF("prompt cache disk stale cleanup: path=%s files=%zu\n", stale_path.c_str(), (size_t) removed); + } + } + + const auto stamp = (uint64_t) std::chrono::high_resolution_clock::now().time_since_epoch().count(); +#if !defined(_WIN32) + const auto pid = (uint64_t) getpid(); +#else + const uint64_t pid = 0; +#endif + + fs::path owned; + for (uint32_t suffix = 0; suffix < 1000; ++suffix) { + owned = cache_root / ("run-" + std::to_string(pid) + "-" + std::to_string(stamp) + "-" + std::to_string(suffix)); + if (fs::create_directory(owned, ec)) { + break; + } + if (ec && ec != std::errc::file_exists) { + throw std::runtime_error("unable to create owned prompt cache directory '" + owned.string() + "': " + ec.message()); + } + ec.clear(); + owned.clear(); + } + if (owned.empty() || !fs::is_directory(owned)) { + throw std::runtime_error("unable to allocate a unique prompt cache run directory below '" + cache_root.string() + "'"); + } + + fs::permissions(owned, fs::perms::owner_all, fs::perm_options::replace, ec); + if (ec) { + fs::remove_all(owned); + throw std::runtime_error("unable to secure owned prompt cache directory '" + owned.string() + "': " + ec.message()); + } + +#if !defined(_WIN32) + // Publish and hold the lock before publishing .owner. Stale cleanup only + // considers magic-marked directories, so another startup can never see an + // owned directory in the window before this process has acquired its lock. + const fs::path lock_path = owned / ".lock"; + disk_lock_fd = open(lock_path.c_str(), O_CREAT | O_RDWR | O_CLOEXEC, 0600); + if (disk_lock_fd < 0 || flock(disk_lock_fd, LOCK_EX | LOCK_NB) != 0) { + if (disk_lock_fd >= 0) { + close(disk_lock_fd); + disk_lock_fd = -1; + } + fs::remove_all(owned); + throw std::runtime_error("unable to lock owned prompt cache directory '" + owned.string() + "'"); + } +#else + { + std::ofstream lock(owned / ".lock", std::ios::out | std::ios::trunc); + if (!lock.good()) { + fs::remove_all(owned); + throw std::runtime_error("unable to create prompt cache lock file in '" + owned.string() + "'"); + } + } +#endif + + { + std::ofstream owner(owned / ".owner", std::ios::out | std::ios::trunc); + owner << SERVER_PROMPT_CACHE_OWNER_MAGIC << '\n' + << "pid=" << pid << '\n' + << "created=" << stamp << '\n'; + owner.flush(); + if (!owner.good()) { +#if !defined(_WIN32) + flock(disk_lock_fd, LOCK_UN); + close(disk_lock_fd); + disk_lock_fd = -1; +#endif + fs::remove_all(owned); + throw std::runtime_error("unable to write prompt cache ownership manifest in '" + owned.string() + "'"); + } + } + fs::permissions(owned / ".owner", fs::perms::owner_read | fs::perms::owner_write, fs::perm_options::replace, ec); + if (ec) { +#if !defined(_WIN32) + flock(disk_lock_fd, LOCK_UN); + close(disk_lock_fd); + disk_lock_fd = -1; +#endif + fs::remove_all(owned); + throw std::runtime_error("unable to secure prompt cache ownership manifest in '" + owned.string() + "': " + ec.message()); + } + + this->disk_base_path = base.string(); + this->disk_owned_path = owned.string(); + + SRV_INF("prompt cache disk enabled: path=%s owned_path=%s limit_mib=%d\n", + this->disk_base_path.c_str(), this->disk_owned_path.c_str(), disk_limit_size_mib); +} + +server_prompt_cache::~server_prompt_cache() { + if (disk_owned_path.empty()) { + return; + } + + SRV_INF("prompt cache disk cleanup: path=%s entries=%zu bytes=%zu saves=%" PRIu64 " loads=%" PRIu64 " evictions=%" PRIu64 "\n", + disk_owned_path.c_str(), disk_states.size(), disk_size_total, disk_saves, disk_loads, disk_evictions); + + fs::path cleanup_path = disk_owned_path; + std::error_code ec; + const fs::path trash_path = cleanup_path.parent_path() / (".deleting-" + cleanup_path.filename().string()); + fs::rename(cleanup_path, trash_path, ec); + if (!ec) { + cleanup_path = trash_path; + } else { + ec.clear(); + } + +#if !defined(_WIN32) + if (disk_lock_fd >= 0) { + flock(disk_lock_fd, LOCK_UN); + close(disk_lock_fd); + disk_lock_fd = -1; + } +#endif + + fs::remove_all(cleanup_path, ec); + if (ec) { + SRV_WRN("prompt cache disk cleanup failed: path=%s error=%s\n", cleanup_path.string().c_str(), ec.message().c_str()); + } +} + size_t server_prompt_cache::size() const { size_t res = 0; @@ -1982,27 +2311,403 @@ size_t server_prompt_cache::n_tokens() const { return res; } -server_prompt * server_prompt_cache::alloc(const server_prompt & prompt, size_t state_size_tgt, size_t state_size_dft) { +size_t server_prompt_cache::disk_size() const { + return disk_size_total; +} + +size_t server_prompt_cache::disk_n_tokens() const { + size_t res = 0; + for (const auto & state : disk_states) { + res += state.n_tokens(); + } + return res; +} + +void server_prompt_cache::disable_disk_saves(const char * reason, const std::string & path) { + disk_save_failures++; + if (disk_save_disabled) { + return; + } + + disk_save_disabled = true; + SRV_ERR("prompt cache disk writes disabled: reason=%s failures=%" PRIu64 " entries=%zu accounted_bytes=%zu path=%s cache_path=%s\n", + reason, disk_save_failures, disk_states.size(), disk_size_total, + path.empty() ? "-" : path.c_str(), disk_owned_path.c_str()); +} + +bool server_prompt_cache::save( + const server_prompt & prompt, + llama_context * ctx_main, + llama_context * ctx_drft, + llama_seq_id id_slot, + const std::vector & state_spec) { + bool saved = false; + + if (!disk_owned_path.empty()) { + saved = save_disk(prompt, ctx_main, ctx_drft, id_slot, state_spec) || saved; + } + + if (!ram_enabled) { + return saved; + } + + const size_t state_size_main = llama_state_seq_get_size_ext(ctx_main, id_slot, LLAMA_STATE_SEQ_FLAGS_NONE); + const size_t state_size_drft = ctx_drft ? llama_state_seq_get_size_ext(ctx_drft, id_slot, LLAMA_STATE_SEQ_FLAGS_NONE) : 0; + + auto * cur = alloc(prompt, state_size_main, state_size_drft, state_spec); + if (cur == nullptr) { + return saved; + } + + const size_t n_main = llama_state_seq_get_data_ext( + ctx_main, cur->data.main.data(), state_size_main, id_slot, LLAMA_STATE_SEQ_FLAGS_NONE); + if (n_main != state_size_main) { + SRV_ERR("failed to save RAM prompt cache target state: expected=%zu saved=%zu\n", state_size_main, n_main); + states.pop_back(); + return saved; + } + + if (ctx_drft) { + const size_t n_drft = llama_state_seq_get_data_ext( + ctx_drft, cur->data.drft.data(), state_size_drft, id_slot, LLAMA_STATE_SEQ_FLAGS_NONE); + if (n_drft != state_size_drft) { + SRV_ERR("failed to save RAM prompt cache draft state: expected=%zu saved=%zu\n", state_size_drft, n_drft); + states.pop_back(); + return saved; + } + } + + return true; +} + +bool server_prompt_cache::save_disk( + const server_prompt & prompt, + llama_context * ctx_main, + llama_context * ctx_drft, + llama_seq_id id_slot, + const std::vector & state_spec) { + if (disk_owned_path.empty() || disk_limit_size == 0 || prompt.tokens.empty()) { + return false; + } + + if (prompt.tokens.has_mtmd) { + SRV_WRN("prompt cache disk skip: reason=multimodal tokens=%zu path=%s\n", + prompt.tokens.size(), disk_owned_path.c_str()); + return false; + } + + // If a usable cached prompt already contains the current stateless prompt, + // retain the more useful state without rewriting the SSD. Stateful MTP + // blobs are valid only at their exact token boundary, so they may touch an + // equal-token entry but never a longer containing entry. + for (auto it = disk_states.begin(); it != disk_states.end();) { + if (!it->usable) { + ++it; + continue; + } + + const int lcp = it->tokens.get_common_prefix(prompt.tokens); + const bool exact_tokens = lcp == (int) prompt.tokens.size() && it->tokens.size() == prompt.tokens.size(); + const bool can_touch = state_spec.empty() + ? lcp == (int) prompt.tokens.size() + : exact_tokens; + if (!can_touch) { + ++it; + continue; + } + + const bool pair_shape_ok = !it->path_main.empty() && it->size_main > 0 && + ((ctx_drft != nullptr) == (!it->path_drft.empty() && it->size_drft > 0)); + const bool spec_shape_ok = state_spec.empty() || !it->spec.empty(); + size_t actual_main = 0; + size_t actual_drft = 0; + const bool files_ok = pair_shape_ok && spec_shape_ok && + server_prompt_cache_disk_size_exact(it->path_main, it->size_main, &actual_main) && + server_prompt_cache_disk_size_exact(it->path_drft, it->size_drft, &actual_drft); + if (!files_ok) { + SRV_WRN("prompt cache disk touch rejected: entry=%" PRIu64 " reason=unusable-pair target_bytes=%zu target_actual=%zu draft_bytes=%zu draft_actual=%zu spec_bytes=%zu path=%s\n", + it->id, it->size_main, actual_main, it->size_drft, actual_drft, it->spec.size(), disk_owned_path.c_str()); + auto bad = it++; + bad->usable = false; + if (!erase_disk_state(bad, false, "touch-unusable")) { + disable_disk_saves("touch-unusable-removal", disk_owned_path); + } + continue; + } + + { + const auto id = it->id; + disk_states.splice(disk_states.end(), disk_states, it); + SRV_INF("prompt cache disk touch: entry=%" PRIu64 " lcp=%d tokens=%zu exact=%s stateful=%s safe_to_clear=true path=%s\n", + id, lcp, prompt.tokens.size(), exact_tokens ? "true" : "false", + state_spec.empty() ? "false" : "true", disk_owned_path.c_str()); + return true; + } + } + + if (disk_save_disabled) { + SRV_DBG("prompt cache disk save skip: reason=circuit-open tokens=%zu path=%s\n", + prompt.tokens.size(), disk_owned_path.c_str()); + return false; + } + + const auto & tokens = prompt.tokens.get_tokens(); + const size_t token_bytes = tokens.size()*sizeof(llama_token); + const size_t file_overhead = 3*sizeof(uint32_t) + token_bytes; + const size_t state_size_main = llama_state_seq_get_size_ext(ctx_main, id_slot, LLAMA_STATE_SEQ_FLAGS_NONE); + const size_t state_size_drft = ctx_drft ? llama_state_seq_get_size_ext(ctx_drft, id_slot, LLAMA_STATE_SEQ_FLAGS_NONE) : 0; + const size_t predicted_main = state_size_main + file_overhead; + const size_t predicted_drft = ctx_drft ? state_size_drft + file_overhead : 0; + const size_t predicted_total = predicted_main + predicted_drft; + + if (predicted_total > disk_limit_size) { + SRV_WRN("prompt cache disk skip: reason=oversize target_bytes=%zu draft_bytes=%zu total_bytes=%zu limit_bytes=%zu tokens=%zu path=%s\n", + predicted_main, predicted_drft, predicted_total, disk_limit_size, tokens.size(), disk_owned_path.c_str()); + return false; + } + + const uint64_t entry_id = disk_next_id++; + const fs::path owned = disk_owned_path; + const std::string stem = "state-" + std::to_string(entry_id); + const fs::path path_main_tmp = owned / (stem + "-target.bin.tmp"); + const fs::path path_main = owned / (stem + "-target.bin"); + const fs::path path_drft_tmp = owned / (stem + "-draft.bin.tmp"); + const fs::path path_drft = owned / (stem + "-draft.bin"); + + const auto cleanup_temps = [&]() -> bool { + const bool main_ok = server_prompt_cache_disk_remove_file(path_main_tmp.string()); + const bool drft_ok = server_prompt_cache_disk_remove_file(path_drft_tmp.string()); + return main_ok && drft_ok; + }; + const auto fail_io = [&](const char * reason, const std::string & path) -> bool { + const bool cleanup_ok = cleanup_temps(); + disable_disk_saves(reason, path); + if (!cleanup_ok) { + disable_disk_saves("temporary-cleanup", disk_owned_path); + } + return false; + }; + + const int64_t t_start = ggml_time_us(); + + const size_t n_main = llama_state_seq_save_file( + ctx_main, path_main_tmp.c_str(), id_slot, tokens.data(), tokens.size()); + size_t actual_main = 0; + if (n_main == 0 || + !server_prompt_cache_disk_size_exact(path_main_tmp.string(), n_main, &actual_main) || + !server_prompt_cache_disk_flush_and_drop(path_main_tmp.string(), true)) { + SRV_ERR("prompt cache disk save failed: entry=%" PRIu64 " component=target path=%s\n", + entry_id, path_main_tmp.c_str()); + return fail_io("target-save", path_main_tmp.string()); + } + + size_t n_drft = 0; + if (ctx_drft) { + n_drft = llama_state_seq_save_file( + ctx_drft, path_drft_tmp.c_str(), id_slot, tokens.data(), tokens.size()); + size_t actual_drft = 0; + if (n_drft == 0 || + !server_prompt_cache_disk_size_exact(path_drft_tmp.string(), n_drft, &actual_drft) || + !server_prompt_cache_disk_flush_and_drop(path_drft_tmp.string(), true)) { + SRV_ERR("prompt cache disk save failed: entry=%" PRIu64 " component=draft path=%s\n", + entry_id, path_drft_tmp.c_str()); + return fail_io("draft-save", path_drft_tmp.string()); + } + } + + const size_t actual_total = n_main + n_drft; + if (actual_total > disk_limit_size) { + const bool cleanup_ok = cleanup_temps(); + SRV_WRN("prompt cache disk skip: reason=actual-oversize entry=%" PRIu64 " target_bytes=%zu draft_bytes=%zu total_bytes=%zu limit_bytes=%zu path=%s\n", + entry_id, n_main, n_drft, actual_total, disk_limit_size, disk_owned_path.c_str()); + if (!cleanup_ok) { + disable_disk_saves("actual-oversize-cleanup", disk_owned_path); + } + return false; + } + + std::error_code ec; + fs::permissions(path_main_tmp, fs::perms::owner_read | fs::perms::owner_write, fs::perm_options::replace, ec); + if (ec) { + SRV_ERR("prompt cache disk permissions failed: entry=%" PRIu64 " component=target path=%s error=%s\n", + entry_id, path_main_tmp.string().c_str(), ec.message().c_str()); + return fail_io("target-permissions", path_main_tmp.string()); + } + if (ctx_drft) { + fs::permissions(path_drft_tmp, fs::perms::owner_read | fs::perms::owner_write, fs::perm_options::replace, ec); + if (ec) { + SRV_ERR("prompt cache disk permissions failed: entry=%" PRIu64 " component=draft path=%s error=%s\n", + entry_id, path_drft_tmp.string().c_str(), ec.message().c_str()); + return fail_io("draft-permissions", path_drft_tmp.string()); + } + } + + // The complete target/draft temporary pair is durable. Commit it before + // touching older entries so a rename or directory-sync failure cannot + // destroy a previously usable cache. This permits one incoming entry of + // transient staging headroom above the configured payload limit. + ec.clear(); + fs::rename(path_main_tmp, path_main, ec); + if (ec) { + SRV_ERR("prompt cache disk atomic rename failed: entry=%" PRIu64 " component=target path=%s error=%s\n", + entry_id, path_main.string().c_str(), ec.message().c_str()); + return fail_io("target-rename", path_main.string()); + } + + if (ctx_drft) { + ec.clear(); + fs::rename(path_drft_tmp, path_drft, ec); + if (ec) { + const bool main_cleanup_ok = server_prompt_cache_disk_remove_file(path_main.string()); + const bool temp_cleanup_ok = cleanup_temps(); + SRV_ERR("prompt cache disk atomic rename failed: entry=%" PRIu64 " component=draft path=%s error=%s\n", + entry_id, path_drft.string().c_str(), ec.message().c_str()); + disable_disk_saves("draft-rename", path_drft.string()); + if (!main_cleanup_ok || !temp_cleanup_ok) { + disable_disk_saves("draft-rename-cleanup", disk_owned_path); + } + return false; + } + } + + if (!server_prompt_cache_disk_sync_dir(disk_owned_path) || + !server_prompt_cache_disk_flush_and_drop(path_main.string(), false) || + (ctx_drft && !server_prompt_cache_disk_flush_and_drop(path_drft.string(), false))) { + const bool main_cleanup_ok = server_prompt_cache_disk_remove_file(path_main.string()); + const bool drft_cleanup_ok = server_prompt_cache_disk_remove_file(path_drft.string()); + disable_disk_saves("commit-sync", disk_owned_path); + if (!main_cleanup_ok || !drft_cleanup_ok) { + disable_disk_saves("commit-sync-cleanup", disk_owned_path); + } + return false; + } + + server_prompt_disk_state state; + state.tokens = prompt.tokens.clone(); + state.path_main = path_main.string(); + state.path_drft = ctx_drft ? path_drft.string() : std::string(); + state.size_main = n_main; + state.size_drft = n_drft; + state.spec = state_spec; + state.id = entry_id; + state.usable = true; + state.checkpoints.reserve(prompt.checkpoints.size()); + for (const auto & ckpt : prompt.checkpoints) { + state.checkpoints.push_back({ckpt.n_tokens, ckpt.pos_min, ckpt.pos_max}); + } + + disk_states.push_back(std::move(state)); + disk_size_total += actual_total; + disk_saves++; + disk_bytes_written += actual_total; + + auto new_entry = std::prev(disk_states.end()); + bool reclaim_ok = true; + + // Stateless entries can supersede shorter prefixes. Stateful MTP blobs + // remain independently useful exact-boundary states. + if (state_spec.empty()) { + for (auto it = disk_states.begin(); it != new_entry;) { + const int lcp = it->tokens.get_common_prefix(prompt.tokens); + if (lcp == (int) it->tokens.size()) { + auto obsolete = it++; + if (!erase_disk_state(obsolete, false, "obsolete-prefix")) { + disable_disk_saves("obsolete-reclaim", disk_owned_path); + reclaim_ok = false; + break; + } + } else { + ++it; + } + } + } + + while (reclaim_ok && disk_size_total > disk_limit_size) { + if (disk_states.begin() == new_entry) { + SRV_ERR("prompt cache disk reclaim failed: entry=%" PRIu64 " reason=no-old-victim accounted_bytes=%zu limit_bytes=%zu path=%s\n", + entry_id, disk_size_total, disk_limit_size, disk_owned_path.c_str()); + disable_disk_saves("room-not-reclaimed", disk_owned_path); + reclaim_ok = false; + break; + } + if (!erase_disk_state(disk_states.begin(), true, "lru-limit")) { + disable_disk_saves("lru-reclaim", disk_owned_path); + reclaim_ok = false; + break; + } + } + + if (!reclaim_ok) { + SRV_WRN("prompt cache disk committed over limit: entry=%" PRIu64 " accounted_bytes=%zu limit_bytes=%zu save_disabled=true path=%s\n", + entry_id, disk_size_total, disk_limit_size, disk_owned_path.c_str()); + } + + const double t_ms = (ggml_time_us() - t_start)/1000.0; + SRV_INF("prompt cache disk save: entry=%" PRIu64 " tokens=%zu checkpoints=%zu target_bytes=%zu draft_bytes=%zu spec_bytes=%zu total_bytes=%zu save_ms=%.2f path=%s\n", + entry_id, tokens.size(), prompt.checkpoints.size(), n_main, n_drft, state_spec.size(), actual_total, t_ms, disk_owned_path.c_str()); + log_disk_state(); + + return true; +} + +server_prompt * server_prompt_cache::alloc( + const server_prompt & prompt, + size_t state_size_tgt, + size_t state_size_dft, + const std::vector & state_spec) { // first check if the current state is contained fully in the cache for (auto it = states.begin(); it != states.end(); ++it) { const int cur_lcp_len = it->tokens.get_common_prefix(prompt.tokens); + const bool exact_tokens = cur_lcp_len == (int) prompt.tokens.size() && + it->tokens.size() == prompt.tokens.size(); + const bool cached_boundary = state_spec.empty() + ? cur_lcp_len == (int) prompt.tokens.size() + : exact_tokens && !it->data.spec.empty(); - if (cur_lcp_len == (int) prompt.tokens.size()) { + if (cached_boundary) { SRV_INF("%s", " - prompt is already in the cache, skipping\n"); return nullptr; } } - // next, remove any cached prompts that are fully contained in the current prompt - for (auto it = states.begin(); it != states.end();) { - const int len = it->tokens.get_common_prefix(prompt.tokens); + // calculate checkpoints size to see if it will fit with the prompt + size_t checkpoints_size = 0; + for (const auto & ckpt : prompt.checkpoints) { + checkpoints_size += ckpt.size(); + } - if (len == (int) it->tokens.size()) { - SRV_WRN(" - removing obsolete cached prompt with length %d\n", len); + const size_t state_size_new = state_size_tgt + state_size_dft + state_spec.size() + checkpoints_size; - it = states.erase(it); - } else { - ++it; + // skip over-limit entries to avoid disturbing the cache + if (limit_size > 0 && state_size_new > limit_size) { + SRV_WRN(" - prompt state size %.3f MiB exceeds cache size limit %.3f MiB, skipping\n", + state_size_new / (1024.0 * 1024.0), limit_size / (1024.0 * 1024.0)); + return nullptr; + } + + // Stateful speculative blobs are exact-boundary states. Keep shorter + // boundaries instead of treating them as obsolete prefixes. + if (state_spec.empty()) { + for (auto it = states.begin(); it != states.end();) { + const int len = it->tokens.get_common_prefix(prompt.tokens); + + if (len == (int) it->tokens.size()) { + SRV_WRN(" - removing obsolete cached prompt with length %d\n", len); + + it = states.erase(it); + } else { + ++it; + } + } + } + + if (limit_size > 0) { + // make room before allocating the new vectors to avoid breaching the limit + while (!states.empty() && size() + state_size_new > limit_size) { + SRV_WRN(" - making room for prompt cache entry, removing oldest entry (size = %.3f MiB)\n", + states.front().size() / (1024.0 * 1024.0)); + + states.pop_front(); } } @@ -2030,6 +2735,7 @@ server_prompt * server_prompt_cache::alloc(const server_prompt & prompt, size_t /*.data =*/ { /*.main =*/ std::move(state_data_tgt), /*.drft =*/ std::move(state_data_dft), + /*.spec =*/ state_spec, }, /*.checkpoints =*/ prompt.checkpoints, }); @@ -2037,41 +2743,357 @@ server_prompt * server_prompt_cache::alloc(const server_prompt & prompt, size_t return &states.back(); } -bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tokens_new, llama_context * ctx_tgt, llama_context * ctx_dft, int32_t id_slot) { +bool server_prompt_cache::load_disk( + std::list::iterator it, + server_prompt & prompt, + llama_context * ctx_tgt, + llama_context * ctx_dft, + llama_seq_id id_slot, + size_t lcp, + uint64_t * entry_id_out) { + if (entry_id_out != nullptr) { + *entry_id_out = 0; + } + + const uint64_t entry_id = it->id; + const size_t target_bytes = it->size_main; + const size_t draft_bytes = it->size_drft; + const size_t spec_bytes = it->spec.size(); + const size_t total_bytes = it->size(); + const size_t n_tokens_expected = it->tokens.size(); + const size_t n_checkpoints = it->checkpoints.size(); + const std::string path_main = it->path_main; + const std::string path_drft = it->path_drft; + + const auto reject_entry = [&](const char * reason) -> bool { + it->usable = false; + if (!erase_disk_state(it, false, reason)) { + disable_disk_saves("invalid-entry-removal", disk_owned_path); + } + log_disk_state(); + return false; + }; + + // Validate the entire pair before mutating either context. + size_t actual_main = 0; + size_t actual_drft = 0; + if (path_main.empty() || target_bytes == 0 || + !server_prompt_cache_disk_size_exact(path_main, target_bytes, &actual_main)) { + SRV_ERR("prompt cache disk load failed: entry=%" PRIu64 " component=target reason=size-mismatch expected_bytes=%zu actual_bytes=%zu path=%s\n", + entry_id, target_bytes, actual_main, path_main.c_str()); + return reject_entry("target-size-mismatch"); + } + if (!path_drft.empty()) { + if (ctx_dft == nullptr || draft_bytes == 0 || + !server_prompt_cache_disk_size_exact(path_drft, draft_bytes, &actual_drft)) { + SRV_ERR("prompt cache disk load failed: entry=%" PRIu64 " component=draft reason=size-mismatch expected_bytes=%zu actual_bytes=%zu path=%s\n", + entry_id, draft_bytes, actual_drft, path_drft.c_str()); + return reject_entry("draft-size-mismatch"); + } + } else if (ctx_dft != nullptr || draft_bytes != 0) { + SRV_ERR("prompt cache disk load failed: entry=%" PRIu64 " component=draft reason=missing-draft-file expected_bytes=%zu path=%s\n", + entry_id, draft_bytes, disk_owned_path.c_str()); + return reject_entry("missing-draft-file"); + } + + const int64_t t_start = ggml_time_us(); + + llama_tokens tokens_main(n_tokens_expected); + size_t n_tokens_main = 0; + const size_t nread_main = llama_state_seq_load_file( + ctx_tgt, path_main.c_str(), id_slot, + tokens_main.data(), tokens_main.size(), &n_tokens_main); + tokens_main.resize(n_tokens_main); + server_prompt_cache_disk_flush_and_drop(path_main, false); + + if (nread_main != target_bytes || !server_prompt_cache_tokens_equal(it->tokens, tokens_main)) { + SRV_ERR("prompt cache disk load failed: entry=%" PRIu64 " component=target expected_bytes=%zu read_bytes=%zu expected_tokens=%zu restored_tokens=%zu path=%s\n", + entry_id, target_bytes, nread_main, n_tokens_expected, n_tokens_main, path_main.c_str()); + return reject_entry("corrupt-target"); + } + + size_t nread_drft = 0; + if (!path_drft.empty()) { + llama_tokens tokens_drft(n_tokens_expected); + size_t n_tokens_drft = 0; + nread_drft = llama_state_seq_load_file( + ctx_dft, path_drft.c_str(), id_slot, + tokens_drft.data(), tokens_drft.size(), &n_tokens_drft); + tokens_drft.resize(n_tokens_drft); + server_prompt_cache_disk_flush_and_drop(path_drft, false); + + if (nread_drft != draft_bytes || + !server_prompt_cache_tokens_equal(it->tokens, tokens_drft) || + tokens_drft != tokens_main) { + SRV_ERR("prompt cache disk load failed: entry=%" PRIu64 " component=draft expected_bytes=%zu read_bytes=%zu expected_tokens=%zu restored_tokens=%zu path=%s\n", + entry_id, draft_bytes, nread_drft, n_tokens_expected, n_tokens_drft, path_drft.c_str()); + return reject_entry("corrupt-draft"); + } + } + + server_prompt restored; + restored.tokens = it->tokens.clone(); + restored.data.spec = it->spec; + // Intentionally do not recreate common_prompt_checkpoint payloads. The + // disk entry retained only their small positions, not cloned host/device + // state. Fresh checkpoints are created as processing continues. + prompt = std::move(restored); + + disk_bytes_read += nread_main + nread_drft; + + const double t_ms = (ggml_time_us() - t_start)/1000.0; + SRV_INF("prompt cache disk load: entry=%" PRIu64 " lcp=%zu tokens=%zu checkpoints=%zu target_bytes=%zu draft_bytes=%zu spec_bytes=%zu total_bytes=%zu read_bytes=%zu load_ms=%.2f path=%s\n", + entry_id, lcp, n_tokens_expected, n_checkpoints, target_bytes, draft_bytes, spec_bytes, total_bytes, + nread_main + nread_drft, t_ms, disk_owned_path.c_str()); + + if (entry_id_out != nullptr) { + *entry_id_out = entry_id; + } + return true; +} + +bool server_prompt_cache::erase_disk_state( + std::list::iterator it, + bool eviction, + const char * reason) { + const uint64_t entry_id = it->id; + const size_t target_bytes = it->size_main; + const size_t draft_bytes = it->size_drft; + const size_t spec_bytes = it->spec.size(); + const size_t total_bytes = it->size(); + const size_t tokens = it->tokens.size(); + const std::string path_main = it->path_main; + const std::string path_drft = it->path_drft; + + // Quarantine before touching either component. If only one unlink works, + // retain the full conservative accounting and metadata for a later retry. + it->usable = false; + const bool main_ok = server_prompt_cache_disk_remove_file(path_main); + const bool drft_ok = server_prompt_cache_disk_remove_file(path_drft); + if (!main_ok || !drft_ok) { + SRV_ERR("prompt cache disk removal failed: entry=%" PRIu64 " reason=%s target_removed=%s draft_removed=%s accounted_bytes=%zu path=%s\n", + entry_id, reason, main_ok ? "true" : "false", drft_ok ? "true" : "false", + disk_size_total, disk_owned_path.c_str()); + return false; + } + + if (total_bytes > disk_size_total) { + SRV_ERR("prompt cache disk accounting invariant failed: entry=%" PRIu64 " entry_bytes=%zu accounted_bytes=%zu path=%s\n", + entry_id, total_bytes, disk_size_total, disk_owned_path.c_str()); + return false; + } + disk_size_total -= total_bytes; + + if (eviction) { + disk_evictions++; + disk_bytes_evicted += total_bytes; + SRV_INF("prompt cache disk eviction: entry=%" PRIu64 " reason=%s tokens=%zu target_bytes=%zu draft_bytes=%zu spec_bytes=%zu total_bytes=%zu remaining_bytes=%zu path=%s\n", + entry_id, reason, tokens, target_bytes, draft_bytes, spec_bytes, total_bytes, disk_size_total, disk_owned_path.c_str()); + } else { + SRV_INF("prompt cache disk remove: entry=%" PRIu64 " reason=%s tokens=%zu target_bytes=%zu draft_bytes=%zu spec_bytes=%zu total_bytes=%zu remaining_bytes=%zu path=%s\n", + entry_id, reason, tokens, target_bytes, draft_bytes, spec_bytes, total_bytes, disk_size_total, disk_owned_path.c_str()); + } + + disk_states.erase(it); + return true; +} + +void server_prompt_cache::accept_disk_load(uint64_t entry_id) { + if (entry_id == 0) { + return; + } + + for (auto it = disk_states.begin(); it != disk_states.end(); ++it) { + if (it->id != entry_id || !it->usable) { + continue; + } + + disk_loads++; + disk_states.splice(disk_states.end(), disk_states, it); + SRV_INF("prompt cache disk load accepted: entry=%" PRIu64 " reusable=true path=%s\n", + entry_id, disk_owned_path.c_str()); + log_disk_state(); + return; + } +} + +void server_prompt_cache::reject_disk_load(uint64_t entry_id, const char * reason) { + if (entry_id == 0) { + return; + } + + for (auto it = disk_states.begin(); it != disk_states.end(); ++it) { + if (it->id != entry_id) { + continue; + } + + it->usable = false; + SRV_WRN("prompt cache disk load rejected: entry=%" PRIu64 " reason=%s reusable=false path=%s\n", + entry_id, reason, disk_owned_path.c_str()); + if (!erase_disk_state(it, false, reason)) { + disable_disk_saves("rejected-load-removal", disk_owned_path); + } + log_disk_state(); + return; + } +} + +void server_prompt_cache::update_disk() { + while (!disk_states.empty() && disk_size_total > disk_limit_size) { + if (!erase_disk_state(disk_states.begin(), true, "lru-update-limit")) { + disable_disk_saves("update-limit-removal", disk_owned_path); + break; + } + } + + log_disk_state(); +} + +void server_prompt_cache::log_disk_state() const { + if (disk_owned_path.empty()) { + return; + } + + const size_t unusable = std::count_if(disk_states.begin(), disk_states.end(), + [](const server_prompt_disk_state & state) { return !state.usable; }); + SRV_INF("prompt cache disk state: entries=%zu unusable=%zu bytes=%zu limit_bytes=%zu over_limit=%s tokens=%zu saves=%" PRIu64 " loads=%" PRIu64 " evictions=%" PRIu64 " save_disabled=%s save_failures=%" PRIu64 " bytes_written=%" PRIu64 " bytes_read=%" PRIu64 " bytes_evicted=%" PRIu64 " path=%s\n", + disk_states.size(), unusable, disk_size_total, disk_limit_size, + disk_size_total > disk_limit_size ? "true" : "false", disk_n_tokens(), + disk_saves, disk_loads, disk_evictions, disk_save_disabled ? "true" : "false", disk_save_failures, + disk_bytes_written, disk_bytes_read, disk_bytes_evicted, + disk_owned_path.c_str()); +} + +bool server_prompt_cache::load( + server_prompt & prompt, + const server_tokens & tokens_new, + llama_context * ctx_tgt, + llama_context * ctx_dft, + int32_t id_slot, + bool spec_state_required, + bool * cache_hit, + uint64_t * disk_entry_id) { + if (cache_hit != nullptr) { + *cache_hit = false; + } + if (disk_entry_id != nullptr) { + *disk_entry_id = 0; + } + const int lcp_best = prompt.tokens.get_common_prefix(tokens_new); - float f_keep_best = prompt.tokens.size() > 0 ? float(lcp_best) / prompt.tokens.size() : -1.0f; // empty slot: any cache entry wins - float sim_best = float(lcp_best) / tokens_new.size(); + const bool base_boundary_valid = !spec_state_required || + lcp_best == (int) prompt.tokens.size(); + float f_keep_best = base_boundary_valid && prompt.tokens.size() > 0 ? float(lcp_best) / prompt.tokens.size() : -1.0f; // empty slot: any cache entry wins + float sim_best = base_boundary_valid ? float(lcp_best) / std::max(1, tokens_new.size()) : -1.0f; + + if (spec_state_required && !prompt.tokens.empty() && !base_boundary_valid) { + SRV_INF("prompt cache skip: reason=spec-boundary-mismatch source=slot lcp=%d cached_tokens=%zu request_tokens=%zu\n", + lcp_best, prompt.tokens.size(), tokens_new.size()); + } SRV_INF(" - looking for better prompt, base f_keep = %.3f, sim = %.3f\n", f_keep_best, sim_best); - auto it_best = states.end(); + auto it_best_ram = states.end(); + auto it_best_disk = disk_states.end(); + size_t lcp_selected = 0; + size_t spec_boundary_best = base_boundary_valid ? prompt.tokens.size() : 0; + bool ram_loaded = false; - // find the most similar cached prompt, that would also preserve the most context + // Find the most similar RAM prompt first. On an equal match, the hot RAM + // copy wins and avoids SSD I/O. for (auto it = states.begin(); it != states.end(); ++it) { const int lcp_cur = it->tokens.get_common_prefix(tokens_new); - const float f_keep_cur = float(lcp_cur) / it->tokens.size(); - const float sim_cur = float(lcp_cur) / tokens_new.size(); + if (spec_state_required && + lcp_cur != (int) it->tokens.size()) { + SRV_INF("prompt cache skip: reason=spec-boundary-mismatch source=ram lcp=%d cached_tokens=%zu request_tokens=%zu spec_bytes=%zu\n", + lcp_cur, it->tokens.size(), tokens_new.size(), it->data.spec.size()); + continue; + } + if (spec_state_required && it->data.spec.empty()) { + SRV_INF("prompt cache skip: reason=spec-state-missing source=ram cached_tokens=%zu request_tokens=%zu\n", + it->tokens.size(), tokens_new.size()); + continue; + } + + const float f_keep_cur = float(lcp_cur) / std::max(1, it->tokens.size()); + const float sim_cur = float(lcp_cur) / std::max(1, tokens_new.size()); // don't trash large prompts if (f_keep_cur < 0.25f) { continue; } - if (f_keep_best < f_keep_cur && sim_best < sim_cur) { + const bool is_better = spec_state_required + ? it->tokens.size() > spec_boundary_best + : f_keep_best < f_keep_cur && sim_best < sim_cur; + if (is_better) { f_keep_best = f_keep_cur; sim_best = sim_cur; + spec_boundary_best = it->tokens.size(); - it_best = it; + it_best_ram = it; + it_best_disk = disk_states.end(); + lcp_selected = lcp_cur; } } - if (it_best != states.end()) { + for (auto it = disk_states.begin(); it != disk_states.end(); ++it) { + if (!it->usable) { + continue; + } + + const int lcp_cur = it->tokens.get_common_prefix(tokens_new); + + if (spec_state_required && + lcp_cur != (int) it->tokens.size()) { + SRV_INF("prompt cache skip: reason=spec-boundary-mismatch source=disk entry=%" PRIu64 " lcp=%d cached_tokens=%zu request_tokens=%zu spec_bytes=%zu\n", + it->id, lcp_cur, it->tokens.size(), tokens_new.size(), it->spec.size()); + continue; + } + if (spec_state_required && it->spec.empty()) { + SRV_INF("prompt cache skip: reason=spec-state-missing source=disk entry=%" PRIu64 " cached_tokens=%zu request_tokens=%zu\n", + it->id, it->tokens.size(), tokens_new.size()); + continue; + } + + const float f_keep_cur = float(lcp_cur) / std::max(1, it->tokens.size()); + const float sim_cur = float(lcp_cur) / std::max(1, tokens_new.size()); + + if (f_keep_cur < 0.25f) { + continue; + } + + const bool is_better = spec_state_required + ? it->tokens.size() > spec_boundary_best + : f_keep_best < f_keep_cur && sim_best < sim_cur; + if (is_better) { + f_keep_best = f_keep_cur; + sim_best = sim_cur; + spec_boundary_best = it->tokens.size(); + + it_best_ram = states.end(); + it_best_disk = it; + lcp_selected = lcp_cur; + } + } + + if (it_best_disk != disk_states.end()) { + SRV_INF(" - found better disk prompt with f_keep = %.3f, sim = %.3f, lcp = %zu\n", + f_keep_best, sim_best, lcp_selected); + const bool loaded = load_disk(it_best_disk, prompt, ctx_tgt, ctx_dft, id_slot, lcp_selected, disk_entry_id); + if (loaded && cache_hit != nullptr) { + *cache_hit = true; + } + return loaded; + } + + if (it_best_ram != states.end()) { SRV_INF(" - found better prompt with f_keep = %.3f, sim = %.3f\n", f_keep_best, sim_best); { - auto & data = it_best->data.main; + auto & data = it_best_ram->data.main; const size_t size = data.size(); const size_t n = llama_state_seq_set_data_ext(ctx_tgt, data.data(), size, id_slot, 0); @@ -2086,7 +3108,7 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok } { - auto & data = it_best->data.drft; + auto & data = it_best_ram->data.drft; if (!data.empty()) { GGML_ASSERT(ctx_dft); @@ -2104,22 +3126,22 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok } } - prompt = std::move(*it_best); + prompt = std::move(*it_best_ram); + + states.erase(it_best_ram); - states.erase(it_best); + if (cache_hit != nullptr) { + *cache_hit = true; + } + ram_loaded = true; } - return true; + return base_boundary_valid || ram_loaded; } void server_prompt_cache::update() { if (limit_size > 0) { - // always keep at least one state, regardless of the limits - while (states.size() > 1 && size() > limit_size) { - if (states.empty()) { - break; - } - + while (!states.empty() && size() > limit_size) { SRV_WRN(" - cache size limit reached, removing oldest entry (size = %.3f MiB)\n", states.front().size() / (1024.0 * 1024.0)); states.pop_front(); @@ -2133,11 +3155,7 @@ void server_prompt_cache::update() { const size_t limit_tokens_cur = limit_size > 0 ? std::max(limit_tokens, limit_size/size_per_token) : limit_tokens; if (limit_tokens > 0) { - while (states.size() > 1 && n_tokens() > limit_tokens_cur) { - if (states.empty()) { - break; - } - + while (!states.empty() && n_tokens() > limit_tokens_cur) { SRV_WRN(" - cache token limit (%zu, est: %zu) reached, removing oldest entry (size = %.3f MiB)\n", limit_tokens, limit_tokens_cur, states.front().size() / (1024.0 * 1024.0)); @@ -2152,4 +3170,6 @@ void server_prompt_cache::update() { SRV_INF(" - prompt %p: %7d tokens, checkpoints: %2zu, %9.3f MiB\n", (const void *)&state, state.n_tokens(), state.checkpoints.size(), state.size() / (1024.0 * 1024.0)); } + + update_disk(); } diff --git a/tools/server/server-task.h b/tools/server/server-task.h index 64bdecd79..a3ebc1ef5 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -568,9 +568,10 @@ struct server_task_result_apply_lora : server_task_result { struct server_prompt_data { std::vector main; std::vector drft; + std::vector spec; size_t size() const { - return main.size() + drft.size(); + return main.size() + drft.size() + spec.size(); } }; @@ -606,27 +607,144 @@ struct server_prompt { } }; -struct server_prompt_cache { - server_prompt_cache(int32_t limit_size_mib, size_t limit_tokens) { - this->limit_size = 1024ull*1024ull*(limit_size_mib < 0 ? 0 : limit_size_mib); - this->limit_tokens = limit_tokens; +// The context checkpoint payloads may contain large host vectors and cloned +// backend buffers. Disk-cache entries keep only the small scheduling metadata; +// a restored disk state starts with an empty live checkpoint list and creates +// fresh checkpoints as prompt processing continues. +struct server_prompt_checkpoint_meta { + int64_t n_tokens = 0; + llama_pos pos_min = 0; + llama_pos pos_max = 0; +}; + +struct server_prompt_disk_state { + server_tokens tokens; + std::vector checkpoints; + std::vector spec; + + std::string path_main; + std::string path_drft; + + size_t size_main = 0; + size_t size_drft = 0; + + uint64_t id = 0; + bool usable = true; + + size_t size() const { + return size_main + size_drft; + } + + int n_tokens() const { + return tokens.size(); } +}; + +struct server_prompt_cache { + server_prompt_cache( + int32_t limit_size_mib, + size_t limit_tokens, + const std::string & disk_base_path = {}, + int32_t disk_limit_size_mib = 0); + + ~server_prompt_cache(); std::list states; + // Cold automatic cache. Entries own only token/checkpoint metadata in RAM; + // target and draft context payloads live in owner-only files. + std::list disk_states; + + bool ram_enabled = false; + // in bytes, 0 = no limit size_t limit_size = 0; // in tokens, 0 = no limit size_t limit_tokens = 0; + // Disk fields are disabled when disk_owned_path is empty. + std::string disk_base_path; + std::string disk_owned_path; + size_t disk_limit_size = 0; + size_t disk_size_total = 0; + int disk_lock_fd = -1; + + uint64_t disk_next_id = 1; + uint64_t disk_saves = 0; + uint64_t disk_loads = 0; + uint64_t disk_evictions = 0; + uint64_t disk_bytes_written = 0; + uint64_t disk_bytes_read = 0; + uint64_t disk_bytes_evicted = 0; + uint64_t disk_save_failures = 0; + + // A durable write/removal failure opens this run-level circuit breaker. + // Existing valid entries remain readable, but no further state files are + // created for this server process. + bool disk_save_disabled = false; + size_t size() const; size_t n_tokens() const; - server_prompt * alloc(const server_prompt & prompt, size_t state_size_main, size_t state_size_drft); + size_t disk_size() const; + + size_t disk_n_tokens() const; + + bool save( + const server_prompt & prompt, + llama_context * ctx_main, + llama_context * ctx_drft, + llama_seq_id id_slot, + const std::vector & state_spec); + + server_prompt * alloc( + const server_prompt & prompt, + size_t state_size_main, + size_t state_size_drft, + const std::vector & state_spec); + + bool load( + server_prompt & prompt, + const server_tokens & tokens_new, + llama_context * ctx_main, + llama_context * ctx_drft, + int32_t id_slot, + bool spec_state_required, + bool * cache_hit, + uint64_t * disk_entry_id); - bool load(server_prompt & prompt, const server_tokens & tokens_new, llama_context * ctx_main, llama_context * ctx_drft, int32_t id_slot); + // Called when a stateful speculative implementation rejects a blob after + // the target/draft files themselves restored successfully. + void accept_disk_load(uint64_t entry_id); + + void reject_disk_load(uint64_t entry_id, const char * reason); void update(); + +private: + bool save_disk( + const server_prompt & prompt, + llama_context * ctx_main, + llama_context * ctx_drft, + llama_seq_id id_slot, + const std::vector & state_spec); + + bool load_disk( + std::list::iterator it, + server_prompt & prompt, + llama_context * ctx_main, + llama_context * ctx_drft, + llama_seq_id id_slot, + size_t lcp, + uint64_t * entry_id_out); + + bool erase_disk_state(std::list::iterator it, bool eviction, const char * reason); + + void disable_disk_saves(const char * reason, const std::string & path); + + void update_disk(); + + void log_disk_state() const; }; diff --git a/tools/server/tests/unit/test_prompt_cache_disk.py b/tools/server/tests/unit/test_prompt_cache_disk.py new file mode 100644 index 000000000..c4c783229 --- /dev/null +++ b/tools/server/tests/unit/test_prompt_cache_disk.py @@ -0,0 +1,320 @@ +import math +import os +import re +import socket +import tempfile +from pathlib import Path + +import pytest + +from utils import * + + +LONG_PROMPT = ( + "Once upon a time in a land far away, there lived a brave knight " + "who traveled across mountains and rivers to find the legendary " + "golden sword hidden deep within the enchanted forest of whispers. " + "He met many creatures along the way including dragons and fairies " + "and wizards who helped him on his noble quest to save the kingdom." +) + +MODEL_DRAFT_FILE_URL = "https://huggingface.co/ggml-org/tiny-llamas/resolve/main/stories15M-q4_0.gguf" +MODEL_TARGET_FILE_URL = "https://huggingface.co/ggml-org/test-model-stories260K/resolve/main/stories260K-f32.gguf" +MODEL_TARGET_DRAFT_PAIR_FILE_URL = "https://huggingface.co/ggml-org/tiny-llamas/resolve/main/stories15M.gguf" + +server = ServerPreset.tinyllama2() + + +# This module uses two explicit tiny local files. Do not invoke the parent +# conftest's all-preset Hugging Face preload, which is unrelated to this test +# and prevents offline/no-HTTPS server builds from reaching the assertions. +@pytest.fixture(scope="module", autouse=True) +def do_something(): + yield + + +class LogReader: + def __init__(self, path): + self.path = path + self.pos = 0 + + def drain(self): + with open(self.path) as f: + f.seek(self.pos) + content = f.read() + self.pos = f.tell() + return content + + +def configure_disk_server(cache_dir, limit_mib=64, draft=False): + global server + server = ServerPreset.tinyllama2() + server.model_file = download_file( + MODEL_TARGET_DRAFT_PAIR_FILE_URL if draft else MODEL_TARGET_FILE_URL + ) + server.model_hf_repo = None + server.model_hf_file = None + # Keep the test isolated from a developer shell that already exports a + # llama-server API key. + server.api_key = os.environ.get("LLAMA_API_KEY") + with socket.socket() as sock: + sock.bind((server.server_host, 0)) + server.server_port = sock.getsockname()[1] + server.n_slots = 2 + server.n_ctx = 512 + server.n_gpu_layer = 0 + server.n_gpu_layer_draft = 0 if draft else None + server.n_predict = 1 + server.temperature = 0.0 + server.server_slots = True + server.cache_ram = 0 + server.cache_disk = cache_dir + server.cache_disk_limit = limit_mib + server.kv_unified = True + server.debug = True + if draft: + server.model_draft = download_file(MODEL_DRAFT_FILE_URL) + server.spec_draft_n_min = 1 + server.spec_draft_n_max = 4 + server.fa = "off" + fd, server.log_path = tempfile.mkstemp(suffix=".log") + os.close(fd) + return server + + +def complete(prompt, id_slot=None): + data = { + "prompt": prompt, + "cache_prompt": True, + "n_predict": 1, + "temperature": 0.0, + } + if id_slot is not None: + data["id_slot"] = id_slot + headers = {"Authorization": f"Bearer {server.api_key}"} if server.api_key else None + res = server.make_request("POST", "/completion", data=data, headers=headers) + assert res.status_code == 200 + return res + + +def prime_and_displace(prompt=LONG_PROMPT): + original = complete(prompt, 0) + complete("The quick brown fox checks a different cache slot.", 1) + return original + + +def test_disk_only_parse_restore_and_owned_cleanup(tmp_path): + configure_disk_server(str(tmp_path), limit_mib=64) + server.start() + log = LogReader(server.log_path) + + startup = log.drain() + assert "prompt cache RAM disabled: limit_mib=0" in startup + assert "prompt cache SSD enabled:" in startup + assert "__TEST_TAG_CACHE_IDLE_SLOTS_ENABLED__" in startup + + original = prime_and_displace() + saved = log.drain() + assert re.search(r"prompt cache disk save: .*target_bytes=[1-9][0-9]* draft_bytes=0", saved) + assert re.search(r"cache state: 0 prompts,", saved) + + restored = complete(LONG_PROMPT) + loaded = log.drain() + assert "prompt cache disk load:" in loaded + assert "draft_bytes=0" in loaded + assert restored.body["timings"]["cache_n"] > 0 + assert restored.body["timings"]["prompt_n"] < original.body["timings"]["prompt_n"] + + # A successful load remains a reusable MRU entry. Saving the restored idle + # slot should touch it, then a second restore should hit the same entry. + first_entry = re.search(r"prompt cache disk load: entry=([0-9]+)", loaded) + assert first_entry is not None + complete("A third prompt displaces the restored slot safely.", 1) + touched = log.drain() + assert f"prompt cache disk touch: entry={first_entry.group(1)}" in touched + assert "safe_to_clear=true" in touched + + restored_again = complete(LONG_PROMPT) + loaded_again = log.drain() + assert f"prompt cache disk load: entry={first_entry.group(1)}" in loaded_again + assert "prompt cache disk load accepted:" in loaded_again + assert restored_again.body["timings"]["cache_n"] > 0 + + namespace = tmp_path / ".llama-prompt-cache-v1" + owned = list(namespace.glob("run-*")) + assert len(owned) == 1 + assert not list(owned[0].glob("*.tmp")) + + server.stop() + assert not list(namespace.glob("run-*")) + assert not list(namespace.glob(".deleting-run-*")) + + +def test_disk_lru_enforces_mib_limit(tmp_path): + # Measure this fixture's streamed state size first, then restart with a + # model-independent limit that fits a few entries and must evict under load. + measure_dir = tmp_path / "measure" + configure_disk_server(str(measure_dir), limit_mib=64) + server.start() + log = LogReader(server.log_path) + log.drain() + prime_and_displace() + measured_log = log.drain() + match = re.search(r"prompt cache disk save: .*total_bytes=([1-9][0-9]*)", measured_log) + assert match is not None + entry_bytes = int(match.group(1)) + server.stop() + + limit_mib = max(1, math.ceil((entry_bytes * 3) / (1024 * 1024))) + cache_dir = tmp_path / "bounded" + configure_disk_server(str(cache_dir), limit_mib=limit_mib) + server.start() + log = LogReader(server.log_path) + log.drain() + + # Equal-length, non-prefix prompts produce similarly sized independent LRU + # entries. Twenty-four turns is deliberately above the measured capacity. + for i in range(24): + complete((f"Cache lane {i:02d} unique marker. " * 18), i % 2) + + bounded_log = log.drain() + assert "prompt cache disk eviction:" in bounded_log + state_lines = [line for line in bounded_log.splitlines() if "prompt cache disk state:" in line] + assert state_lines + final_state = state_lines[-1] + bytes_match = re.search(r" bytes=([0-9]+) limit_bytes=([0-9]+) ", final_state) + assert bytes_match is not None + assert int(bytes_match.group(1)) <= int(bytes_match.group(2)) == limit_mib * 1024 * 1024 + + run_dirs = list((cache_dir / ".llama-prompt-cache-v1").glob("run-*")) + assert len(run_dirs) == 1 + payload_bytes = sum(path.stat().st_size for path in run_dirs[0].glob("state-*.bin")) + assert payload_bytes <= limit_mib * 1024 * 1024 + + +def test_disk_cache_round_trips_target_and_draft(tmp_path): + configure_disk_server(str(tmp_path), limit_mib=64, draft=True) + server.start() + log = LogReader(server.log_path) + log.drain() + + original = prime_and_displace() + saved = log.drain() + save_match = re.search( + r"prompt cache disk save: .*target_bytes=([1-9][0-9]*) draft_bytes=([1-9][0-9]*)", + saved, + ) + assert save_match is not None + + restored = complete(LONG_PROMPT) + loaded = log.drain() + load_match = re.search( + r"prompt cache disk load: .*target_bytes=([1-9][0-9]*) draft_bytes=([1-9][0-9]*)", + loaded, + ) + assert load_match is not None + assert "prompt cache cold fallback:" not in loaded + assert restored.body["timings"]["cache_n"] > 0 + assert restored.body["timings"]["prompt_n"] < original.body["timings"]["prompt_n"] + + +def test_disk_cache_rejects_partial_target_draft_pair(tmp_path): + configure_disk_server(str(tmp_path), limit_mib=64, draft=True) + server.start() + log = LogReader(server.log_path) + log.drain() + + original = prime_and_displace() + saved = log.drain() + assert re.search( + r"prompt cache disk save: .*target_bytes=[1-9][0-9]* draft_bytes=[1-9][0-9]*", + saved, + ) + assert re.search(r"cache state: 0 prompts,", saved) + + run_dirs = list((tmp_path / ".llama-prompt-cache-v1").glob("run-*")) + assert len(run_dirs) == 1 + draft_files = list(run_dirs[0].glob("state-*-draft.bin")) + assert draft_files + draft_files[0].unlink() + + restored = complete(LONG_PROMPT) + rejected = log.drain() + assert "prompt cache disk load failed:" in rejected + assert "component=draft" in rejected + assert "prompt cache cold fallback:" in rejected + assert "target_and_draft_cleared=true" in rejected + assert restored.body["timings"]["cache_n"] == 0 + assert restored.body["timings"]["prompt_n"] == original.body["timings"]["prompt_n"] + + # A rejected pair must not poison the slot or the server. + healthy = complete("The server remains healthy after a rejected cache pair.") + assert healthy.status_code == 200 + + +def test_disk_save_failure_opens_breaker_and_preserves_idle_slot(tmp_path): + configure_disk_server(str(tmp_path), limit_mib=64) + server.start() + log = LogReader(server.log_path) + log.drain() + + original = complete(LONG_PROMPT, 0) + run_dirs = list((tmp_path / ".llama-prompt-cache-v1").glob("run-*")) + assert len(run_dirs) == 1 + + # Remove directory write permission so the first target temp cannot be + # created. The slot must remain live because no durable cache exists. + run_dirs[0].chmod(0o500) + try: + complete("This request forces an idle-slot save failure.", 1) + finally: + run_dirs[0].chmod(0o700) + + failed = log.drain() + assert "prompt cache disk writes disabled:" in failed + assert "reason=target-save" in failed + assert "preserving idle slot because prompt cache save was not safe" in failed + assert "safe_to_clear=true" not in failed + + reused_live_slot = complete(LONG_PROMPT, 0) + after_breaker = log.drain() + assert "reason=circuit-open" in after_breaker + assert reused_live_slot.body["timings"]["cache_n"] > 0 + assert reused_live_slot.body["timings"]["prompt_n"] < original.body["timings"]["prompt_n"] + + +def test_failed_corrupt_entry_removal_keeps_conservative_accounting(tmp_path): + configure_disk_server(str(tmp_path), limit_mib=64) + server.start() + log = LogReader(server.log_path) + log.drain() + + original = prime_and_displace() + saved = log.drain() + size_match = re.search(r"prompt cache disk save: .*total_bytes=([1-9][0-9]*)", saved) + assert size_match is not None + accounted = int(size_match.group(1)) + + run_dirs = list((tmp_path / ".llama-prompt-cache-v1").glob("run-*")) + assert len(run_dirs) == 1 + target_files = list(run_dirs[0].glob("state-*-target.bin")) + assert target_files + with target_files[0].open("r+b") as f: + f.truncate(max(1, accounted // 2)) + + # Prevent quarantine cleanup. The entry must remain fully accounted and + # unusable rather than being reported as freed after unlink failure. + run_dirs[0].chmod(0o500) + try: + restored = complete(LONG_PROMPT) + finally: + run_dirs[0].chmod(0o700) + + rejected = log.drain() + assert "reason=size-mismatch" in rejected + assert "prompt cache disk removal failed:" in rejected + assert f"accounted_bytes={accounted}" in rejected + assert re.search(rf"prompt cache disk state: entries=1 unusable=1 bytes={accounted} ", rejected) + assert "save_disabled=true" in rejected + assert restored.body["timings"]["cache_n"] == 0 + assert restored.body["timings"]["prompt_n"] == original.body["timings"]["prompt_n"] diff --git a/tools/server/tests/utils.py b/tools/server/tests/utils.py index c5dba1c13..8a7e8fc47 100644 --- a/tools/server/tests/utils.py +++ b/tools/server/tests/utils.py @@ -64,6 +64,7 @@ class ServerProcess: model_draft: str | None = None n_threads: int | None = None n_gpu_layer: int | None = None + n_gpu_layer_draft: int | None = None n_batch: int | None = None n_ubatch: int | None = None n_ctx: int | None = None @@ -105,6 +106,8 @@ class ServerProcess: media_path: str | None = None sleep_idle_seconds: int | None = None cache_ram: int | None = None + cache_disk: str | None = None + cache_disk_limit: int | None = None no_cache_idle_slots: bool = False log_path: str | None = None webui_mcp_proxy: bool = False @@ -172,8 +175,10 @@ def start(self, timeout_seconds: int = DEFAULT_HTTP_TIMEOUT) -> None: server_args.extend(["--ubatch-size", self.n_ubatch]) if self.n_threads: server_args.extend(["--threads", self.n_threads]) - if self.n_gpu_layer: + if self.n_gpu_layer is not None: server_args.extend(["--n-gpu-layers", self.n_gpu_layer]) + if self.n_gpu_layer_draft is not None: + server_args.extend(["--n-gpu-layers-draft", self.n_gpu_layer_draft]) if self.server_continuous_batching: server_args.append("--cont-batching") if self.server_embeddings: @@ -249,6 +254,10 @@ def start(self, timeout_seconds: int = DEFAULT_HTTP_TIMEOUT) -> None: server_args.extend(["--sleep-idle-seconds", self.sleep_idle_seconds]) if self.cache_ram is not None: server_args.extend(["--cache-ram", self.cache_ram]) + if self.cache_disk is not None: + server_args.extend(["--cache-disk", self.cache_disk]) + if self.cache_disk_limit is not None: + server_args.extend(["--cache-disk-limit", self.cache_disk_limit]) if self.no_cache_idle_slots: server_args.append("--no-cache-idle-slots") if self.webui_mcp_proxy: From 23d4c1f49e885b78b56c84124bb524e4a1fe032e Mon Sep 17 00:00:00 2001 From: ciru-ai Date: Mon, 13 Jul 2026 19:53:22 -0400 Subject: [PATCH 4/4] fix(rocmfpx): classify IFP2 in CPU clamp dispatch --- ggml/src/ggml-cpu/ops.cpp | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index 45efba895..96c11ffa2 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -1255,8 +1255,8 @@ void ggml_compute_forward_acc( case GGML_TYPE_NVFP4: case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: - case GGML_TYPE_Q3_0_ROCMFPX: case GGML_TYPE_Q2_0_ROCMFPX: + case GGML_TYPE_Q3_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: case GGML_TYPE_Q2_K: @@ -5666,6 +5666,7 @@ void ggml_compute_forward_clamp( case GGML_TYPE_NVFP4: case GGML_TYPE_Q4_0_ROCMFP4: case GGML_TYPE_Q4_0_ROCMFP4_FAST: + case GGML_TYPE_Q2_0_ROCMFPX: case GGML_TYPE_Q3_0_ROCMFPX: case GGML_TYPE_Q6_0_ROCMFPX: case GGML_TYPE_Q8_0_ROCMFPX: