diff --git a/common/arg.cpp b/common/arg.cpp index bf2b62460..c3c312736 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -28,6 +28,7 @@ #include #include #include +#include #include #include #include @@ -3452,9 +3453,14 @@ common_params_context common_params_parser_init(common_params & params, llama_ex ).set_env("LLAMA_ARG_LOG_FILE")); add_opt(common_arg( {"--log-prompts-dir"}, "PATH", - "Log prompts to directory (only used for debugging, default: disabled)", + "Log prompts to directory (auto-created if not present; only used for debugging, default: disabled)", [](common_params & params, const std::string & value) { params.path_prompts_log_dir = value; + std::error_code ec; + std::filesystem::create_directories(value, ec); + if (ec) { + fprintf(stderr, "warning: failed to create prompts-log-dir '%s': %s\n", value.c_str(), ec.message().c_str()); + } } ).set_examples({LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI})); add_opt(common_arg( diff --git a/common/common.h b/common/common.h index 7939d3b07..996ec511d 100644 --- a/common/common.h +++ b/common/common.h @@ -15,6 +15,7 @@ #include #include #include +#include #if defined(_WIN32) && !defined(_WIN32_WINNT) #define _WIN32_WINNT 0x0A00 diff --git a/common/speculative.cpp b/common/speculative.cpp index 4e0439216..b61d12bce 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -2221,6 +2221,112 @@ int32_t common_speculative_n_max(const common_params_speculative * spec) { return n_max; } +common_params common_base_params_to_speculative(const common_params & params) { + const bool has_draft = params.speculative.has_dft(); + + const auto & params_spec = params.speculative.draft; + common_params result = params; + + if (has_draft) { + result.devices = params_spec.devices; + result.model = params_spec.mparams; + result.n_gpu_layers = params_spec.n_gpu_layers; + result.tensor_buft_overrides = params_spec.tensor_buft_overrides; + + if (params_spec.cpuparams.n_threads > 0) { + result.cpuparams.n_threads = params_spec.cpuparams.n_threads; + result.cpuparams_batch.n_threads = params_spec.cpuparams_batch.n_threads; + } + } + + result.cache_type_k = params_spec.cache_type_k; + result.cache_type_v = params_spec.cache_type_v; + result.n_outputs_max = params.n_parallel; + + return result; +} + +struct common_speculative_init_result::impl { + impl() = default; + ~impl() = default; + + // note: the order in which model, context, etc. are declared matters because their destructors will be called bottom-to-top + llama_model_ptr model; + llama_context_ptr context; +}; + +common_speculative_init_result::common_speculative_init_result( + common_params & params, + llama_model * model_tgt, + llama_context * ctx_tgt) : + pimpl(new impl{}) { + const bool has_draft = params.speculative.has_dft(); + const bool spec_mtp = std::find(params.speculative.types.begin(), + params.speculative.types.end(), + COMMON_SPECULATIVE_TYPE_DRAFT_MTP) != params.speculative.types.end(); + GGML_ASSERT(has_draft || spec_mtp); + + auto mparams = common_model_params_to_llama(params); + auto cparams = common_context_params_to_llama(params); + + if (spec_mtp) { + cparams.ctx_type = LLAMA_CONTEXT_TYPE_MTP; + } + + // note: for small models maybe we can set this to the maximum possible draft from all speculative types + // the extra memory for small models is likely negligible? + cparams.n_rs_seq = 0; + cparams.ctx_other = ctx_tgt; + + std::string model_path; + if (has_draft) { + model_path = params.speculative.draft.mparams.path; + LOG_TRC("%s: loading draft model '%s'\n", __func__, model_path.c_str()); + + llama_model * model_dft = llama_model_load_from_file(params.model.path.c_str(), mparams); + if (model_dft == NULL) { + LOG_ERR("%s: failed to load draft model, '%s'\n", __func__, model_path.c_str()); + return; + } + + pimpl->model.reset(model_dft); + + llama_context * ctx_dft = llama_init_from_model(model_dft, cparams); + if (ctx_dft == nullptr) { + LOG_ERR("%s: failed to create MTP context\n", __func__); + return; + } + + pimpl->context.reset(ctx_dft); + } else if (spec_mtp) { + model_path = params.model.path; + + LOG_TRC("%s: creating MTP draft context against the target model '%s'\n", __func__, model_path.c_str()); + + llama_context * ctx_dft = llama_init_from_model(model_tgt, cparams); + if (ctx_dft == nullptr) { + LOG_ERR("%s: failed to create MTP context\n", __func__); + return; + } + + pimpl->context.reset(ctx_dft); + } +} + +common_speculative_init_result::~common_speculative_init_result() = default; + +llama_model * common_speculative_init_result::model() { + return pimpl->model.get(); +} + +llama_context * common_speculative_init_result::context() { + return pimpl->context.get(); +} + +common_speculative_init_result_ptr common_speculative_init_from_params(common_params & params, llama_model * model_tgt, llama_context * ctx_tgt) { + return std::make_unique(params, model_tgt, ctx_tgt); +} + // initialization of the speculative decoding system // common_speculative * common_speculative_init(common_params_speculative & params, uint32_t n_seq) { diff --git a/common/speculative.h b/common/speculative.h index c58fac3cc..062bf2093 100644 --- a/common/speculative.h +++ b/common/speculative.h @@ -23,6 +23,8 @@ std::string common_speculative_type_to_str(enum common_speculative_type type); // return the max number of draft tokens based on the speculative parameters int32_t common_speculative_n_max(const common_params_speculative * spec); +common_params common_base_params_to_speculative(const common_params & params); + common_speculative * common_speculative_init(common_params_speculative & params, uint32_t n_seq); void common_speculative_free(common_speculative * spec); @@ -80,3 +82,19 @@ struct common_speculative_deleter { }; typedef std::unique_ptr common_speculative_ptr; + +struct common_speculative_init_result { + common_speculative_init_result(common_params & params, llama_model * model_tgt, llama_context * ctx_tgt); + ~common_speculative_init_result(); + + llama_model * model(); + llama_context * context(); + +private: + struct impl; + std::unique_ptr pimpl; +}; + +using common_speculative_init_result_ptr = std::unique_ptr; + +common_speculative_init_result_ptr common_speculative_init_from_params(common_params & params, llama_model * model_tgt, llama_context * ctx_tgt); diff --git a/examples/llama-eval/llama-eval.py b/examples/llama-eval/llama-eval.py index 4bdd239c0..61bdbddd8 100644 --- a/examples/llama-eval/llama-eval.py +++ b/examples/llama-eval/llama-eval.py @@ -362,7 +362,7 @@ class EvalState: case = cases.get(task_id, {}) status = case.get("status", "pending") expected = case.get("expected", "") - answer = case.get("answer", "") if status == "ok" else "" + answer = case.get("answer") or "" if status == "ok" else "" is_correct = case.get("correct", False) if status == "ok" else False response = case.get("response", "") or "" prompt = case.get("prompt", "") or "" @@ -647,7 +647,7 @@ class EvalState: question, prompt, expected = self.get_case(i) case = cases.get(task_id, {}) status = case.get("status", "pending") - answer = case.get("answer", "N/A") if status == "ok" else "N/A" + answer = case.get("answer") or "N/A" if status == "ok" else "N/A" tokens = case.get("tokens") tokens_str = str(tokens) if tokens is not None else "N/A" tps_gen = case.get("tps_gen") diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 22d535ecf..d0ca0676a 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -435,7 +435,8 @@ extern "C" { GGML_TYPE_MXFP4 = 39, // MXFP4 (1 block) GGML_TYPE_NVFP4 = 40, // NVFP4 (4 blocks, E4M3 scale) GGML_TYPE_Q1_0 = 41, - GGML_TYPE_COUNT = 42, + GGML_TYPE_Q2_0 = 42, + GGML_TYPE_COUNT = 43, }; // precision @@ -479,6 +480,7 @@ extern "C" { GGML_FTYPE_MOSTLY_MXFP4 = 25, // except 1d tensors GGML_FTYPE_MOSTLY_NVFP4 = 26, // except 1d tensors GGML_FTYPE_MOSTLY_Q1_0 = 27, // except 1d tensors + GGML_FTYPE_MOSTLY_Q2_0 = 28, // except 1d tensors }; // available tensor operations: diff --git a/ggml/src/ggml-common.h b/ggml/src/ggml-common.h index a51b37dcd..83f9118da 100644 --- a/ggml/src/ggml-common.h +++ b/ggml/src/ggml-common.h @@ -96,6 +96,9 @@ typedef sycl::half2 ggml_half2; #define QI1_0 (QK1_0 / 32) #define QR1_0 1 +#define QI2_0 (QK2_0 / 32) +#define QR2_0 1 + #define QI4_0 (QK4_0 / (4 * QR4_0)) #define QR4_0 2 @@ -181,6 +184,13 @@ typedef struct { } block_q1_0; static_assert(sizeof(block_q1_0) == sizeof(ggml_half) + QK1_0 / 8, "wrong q1_0 block size/padding"); +#define QK2_0 64 +typedef struct { + ggml_half d; // delta (scale) + uint8_t qs[QK2_0 / 4]; // 2 bits per element +} block_q2_0; +static_assert(sizeof(block_q2_0) == sizeof(ggml_half) + QK2_0 / 4, "wrong q2_0 block size/padding"); + #define QK4_0 32 typedef struct { ggml_half d; // delta diff --git a/ggml/src/ggml-cpu/arch-fallback.h b/ggml/src/ggml-cpu/arch-fallback.h index 167ff47f5..8af7be52b 100644 --- a/ggml/src/ggml-cpu/arch-fallback.h +++ b/ggml/src/ggml-cpu/arch-fallback.h @@ -17,6 +17,7 @@ #define ggml_vec_dot_mxfp4_q8_0_generic ggml_vec_dot_mxfp4_q8_0 #define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0 #define ggml_vec_dot_q1_0_q8_0_generic ggml_vec_dot_q1_0_q8_0 +#define ggml_vec_dot_q2_0_q8_0_generic ggml_vec_dot_q2_0_q8_0 #define ggml_vec_dot_tq1_0_q8_K_generic ggml_vec_dot_tq1_0_q8_K #define ggml_vec_dot_tq2_0_q8_K_generic ggml_vec_dot_tq2_0_q8_K #define ggml_vec_dot_q2_K_q8_K_generic ggml_vec_dot_q2_K_q8_K @@ -72,6 +73,7 @@ #define ggml_gemm_q2_K_8x8_q8_K_generic ggml_gemm_q2_K_8x8_q8_K #elif defined(__x86_64__) || defined(__i386__) || defined(_M_IX86) || defined(_M_X64) // quants.c +#define ggml_vec_dot_q2_0_q8_0_generic ggml_vec_dot_q2_0_q8_0 // repack.cpp #define ggml_quantize_mat_q8_0_4x4_generic ggml_quantize_mat_q8_0_4x4 #define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4 @@ -101,6 +103,7 @@ #define quantize_row_q8_K_generic quantize_row_q8_K #define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0 #define ggml_vec_dot_q1_0_q8_0_generic ggml_vec_dot_q1_0_q8_0 +#define ggml_vec_dot_q2_0_q8_0_generic ggml_vec_dot_q2_0_q8_0 #define ggml_vec_dot_tq1_0_q8_K_generic ggml_vec_dot_tq1_0_q8_K #define ggml_vec_dot_tq2_0_q8_K_generic ggml_vec_dot_tq2_0_q8_K #define ggml_vec_dot_iq1_m_q8_K_generic ggml_vec_dot_iq1_m_q8_K @@ -144,6 +147,7 @@ #define ggml_vec_dot_mxfp4_q8_0_generic ggml_vec_dot_mxfp4_q8_0 #define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0 #define ggml_vec_dot_q1_0_q8_0_generic ggml_vec_dot_q1_0_q8_0 +#define ggml_vec_dot_q2_0_q8_0_generic ggml_vec_dot_q2_0_q8_0 // repack.cpp #define ggml_quantize_mat_q8_0_4x4_generic ggml_quantize_mat_q8_0_4x4 #define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8 @@ -178,6 +182,7 @@ #elif defined(__riscv) // quants.c #define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0 +#define ggml_vec_dot_q2_0_q8_0_generic ggml_vec_dot_q2_0_q8_0 // repack.cpp #define ggml_quantize_mat_q8_0_4x4_generic ggml_quantize_mat_q8_0_4x4 #define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8 @@ -212,6 +217,7 @@ #define quantize_row_q8_K_generic quantize_row_q8_K #define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0 #define ggml_vec_dot_q1_0_q8_0_generic ggml_vec_dot_q1_0_q8_0 +#define ggml_vec_dot_q2_0_q8_0_generic ggml_vec_dot_q2_0_q8_0 #define ggml_vec_dot_tq1_0_q8_K_generic ggml_vec_dot_tq1_0_q8_K #define ggml_vec_dot_tq2_0_q8_K_generic ggml_vec_dot_tq2_0_q8_K #define ggml_vec_dot_q2_K_q8_K_generic ggml_vec_dot_q2_K_q8_K @@ -269,6 +275,7 @@ #define ggml_vec_dot_mxfp4_q8_0_generic ggml_vec_dot_mxfp4_q8_0 #define ggml_vec_dot_nvfp4_q8_0_generic ggml_vec_dot_nvfp4_q8_0 #define ggml_vec_dot_q1_0_q8_0_generic ggml_vec_dot_q1_0_q8_0 +#define ggml_vec_dot_q2_0_q8_0_generic ggml_vec_dot_q2_0_q8_0 // repack.cpp #define ggml_quantize_mat_q8_0_4x4_generic ggml_quantize_mat_q8_0_4x4 #define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8 diff --git a/ggml/src/ggml-cpu/arch/arm/quants.c b/ggml/src/ggml-cpu/arch/arm/quants.c index 445366797..636d7be12 100644 --- a/ggml/src/ggml-cpu/arch/arm/quants.c +++ b/ggml/src/ggml-cpu/arch/arm/quants.c @@ -219,6 +219,80 @@ void ggml_vec_dot_q1_0_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const voi #endif } +void ggml_vec_dot_q2_0_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) { + const int qk = QK2_0; + const int nb = n / qk; + + assert(n % qk == 0); + assert(nrc == 1); + UNUSED(nrc); + UNUSED(bx); + UNUSED(by); + UNUSED(bs); + + const block_q2_0 * GGML_RESTRICT x = vx; + const block_q8_0 * GGML_RESTRICT y = vy; + + float sumf = 0.0f; + +#if defined(__ARM_NEON) + // Replicate pattern: each byte repeated 4 times + static const uint8_t tbl_idx_lo[16] = {0,0,0,0, 1,1,1,1, 2,2,2,2, 3,3,3,3}; + static const uint8_t tbl_idx_hi[16] = {4,4,4,4, 5,5,5,5, 6,6,6,6, 7,7,7,7}; + // Right-shift amounts: 0,2,4,6 repeated for each group of 4 + static const int8_t shift_vals[16] = {0,-2,-4,-6, 0,-2,-4,-6, 0,-2,-4,-6, 0,-2,-4,-6}; + + const uint8x16_t idx_lo = vld1q_u8(tbl_idx_lo); + const uint8x16_t idx_hi = vld1q_u8(tbl_idx_hi); + const int8x16_t shifts = vld1q_s8(shift_vals); + const uint8x16_t mask2 = vdupq_n_u8(0x03); + const int8x16_t one = vdupq_n_s8(1); + + float32x4_t sumv = vdupq_n_f32(0.0f); + + for (int i = 0; i < nb; i++) { + const float d0 = GGML_CPU_FP16_TO_FP32(x[i].d); + + // group 64: one Q2_0 block (64 weights) maps to two Q8_0 blocks (2 * 32 = 64) + for (int k = 0; k < 2; k++) { + const block_q8_0 * GGML_RESTRICT yb = &y[i * 2 + k]; + const float d1 = GGML_CPU_FP16_TO_FP32(yb->d); + + // Load 8 bytes of packed 2-bit values + const uint8x8_t raw = vld1_u8(&x[i].qs[k * 8]); + const uint8x16_t raw16 = vcombine_u8(raw, raw); + + // First 16 elements: replicate bytes 0-3, shift, mask, subtract 1 + uint8x16_t bytes0 = vqtbl1q_u8(raw16, idx_lo); + int8x16_t qv0 = vsubq_s8( + vreinterpretq_s8_u8(vandq_u8(vshlq_u8(bytes0, shifts), mask2)), + one); + + // Second 16 elements: replicate bytes 4-7, shift, mask, subtract 1 + uint8x16_t bytes1 = vqtbl1q_u8(raw16, idx_hi); + int8x16_t qv1 = vsubq_s8( + vreinterpretq_s8_u8(vandq_u8(vshlq_u8(bytes1, shifts), mask2)), + one); + + // Load Q8_0 values and dot product + const int8x16_t y0 = vld1q_s8(yb->qs); + const int8x16_t y1 = vld1q_s8(yb->qs + 16); + + int32x4_t p0 = ggml_vdotq_s32(vdupq_n_s32(0), qv0, y0); + int32x4_t p1 = ggml_vdotq_s32(p0, qv1, y1); + + sumv = vmlaq_n_f32(sumv, vcvtq_f32_s32(p1), d0 * d1); + } + } + + sumf = vaddvq_f32(sumv); +#else + ggml_vec_dot_q2_0_q8_0_generic(n, s, bs, vx, bx, vy, by, nrc); + return; +#endif + + *s = sumf; +} void ggml_vec_dot_q4_0_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) { const int qk = QK8_0; diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index 0d110b87d..8ab0abf69 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -231,6 +231,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] = { + .from_float = quantize_row_q2_0, + .vec_dot = ggml_vec_dot_q2_0_q8_0, + .vec_dot_type = GGML_TYPE_Q8_0, + .nrows = 1, + }, [GGML_TYPE_Q4_0] = { .from_float = quantize_row_q4_0, .vec_dot = ggml_vec_dot_q4_0_q8_0, diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index 168f1a89d..0dfdacf66 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -665,6 +665,7 @@ void ggml_compute_forward_add( ggml_compute_forward_add_non_quantized(params, dst); } break; case GGML_TYPE_Q1_0: + case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: @@ -1115,6 +1116,7 @@ void ggml_compute_forward_add1( } } break; case GGML_TYPE_Q1_0: + case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: @@ -1245,6 +1247,7 @@ void ggml_compute_forward_acc( case GGML_TYPE_F16: case GGML_TYPE_BF16: case GGML_TYPE_Q1_0: + case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: @@ -4454,6 +4457,7 @@ void ggml_compute_forward_out_prod( switch (src0->type) { case GGML_TYPE_Q1_0: + case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: @@ -4730,6 +4734,7 @@ void ggml_compute_forward_set( case GGML_TYPE_F16: case GGML_TYPE_BF16: case GGML_TYPE_Q1_0: + case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: @@ -4954,6 +4959,7 @@ void ggml_compute_forward_get_rows( switch (src0->type) { case GGML_TYPE_Q1_0: + case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: @@ -5019,8 +5025,8 @@ void ggml_compute_forward_get_rows( //} } -template -static void ggml_compute_forward_set_rows_f32( +template +static void ggml_compute_forward_set_rows_impl( const ggml_compute_params * params, ggml_tensor * dst) { @@ -5035,7 +5041,7 @@ static void ggml_compute_forward_set_rows_f32( assert(ne0 == nc); assert(ne2 == ne02); assert(ne3 == ne03); - assert(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src0->type == GGML_TYPE_F32 || (src0->type == GGML_TYPE_F16 && dst->type == GGML_TYPE_F16)); assert(ne02 % ne11 == 0); assert(ne03 % ne12 == 0); @@ -5049,6 +5055,8 @@ static void ggml_compute_forward_set_rows_f32( const int64_t ir0 = dr*ith; const int64_t ir1 = std::min(ir0 + dr, nr); + const size_t rs = ggml_row_size(src0->type, nc); + ggml_from_float_t const from_float = ggml_get_type_traits_cpu(dst->type)->from_float; for (int64_t i03 = 0; i03 < ne03; ++i03) { @@ -5062,9 +5070,18 @@ static void ggml_compute_forward_set_rows_f32( GGML_ASSERT(i1 >= 0 && i1 < ne1); - from_float( - (const float *) ((char *) src0->data + i*nb01 + i02*nb02 + i03*nb03), - ((char *) dst->data + i1*nb1 + i02*nb2 + i03*nb3), nc); + if constexpr (std::is_same_v) { + from_float( + (const float *) ((char *) src0->data + i*nb01 + i02*nb02 + i03*nb03), + ((char *) dst->data + i1*nb1 + i02*nb2 + i03*nb3), nc); + } else if constexpr (std::is_same_v) { + memcpy( + ((char *) dst->data + i1*nb1 + i02*nb2 + i03*nb3), + ((char *) src0->data + i*nb01 + i02*nb02 + i03*nb03), + rs); + } else { + GGML_ABORT("src0->type = %d (%s) not supported", src0->type, ggml_type_name(src0->type)); + } } } } @@ -5081,13 +5098,27 @@ void ggml_compute_forward_set_rows( case GGML_TYPE_F32: { if (src1->type == GGML_TYPE_I64) { - ggml_compute_forward_set_rows_f32(params, dst); + ggml_compute_forward_set_rows_impl(params, dst); } else if (src1->type == GGML_TYPE_I32) { - ggml_compute_forward_set_rows_f32(params, dst); + ggml_compute_forward_set_rows_impl(params, dst); } else { GGML_ABORT("src1->type = %d (%s) not supported", src1->type, ggml_type_name(src1->type)); } } break; + case GGML_TYPE_F16: + { + if (dst->type == GGML_TYPE_F16) { + if (src1->type == GGML_TYPE_I64) { + ggml_compute_forward_set_rows_impl(params, dst); + } else if (src1->type == GGML_TYPE_I32) { + ggml_compute_forward_set_rows_impl(params, dst); + } else { + GGML_ABORT("src1->type = %d (%s) not supported", src1->type, ggml_type_name(src1->type)); + } + } else { + GGML_ABORT("dst->type = %d (%s) not supported with src0->type = %d (%s)", dst->type, ggml_type_name(dst->type), src0->type, ggml_type_name(src0->type)); + } + } break; default: { GGML_ABORT("src0->type = %d (%s) not supported", src0->type, ggml_type_name(src0->type)); @@ -5680,6 +5711,7 @@ void ggml_compute_forward_clamp( } break; case GGML_TYPE_BF16: case GGML_TYPE_Q1_0: + case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: diff --git a/ggml/src/ggml-cpu/quants.c b/ggml/src/ggml-cpu/quants.c index e5f9a4083..5e36459f8 100644 --- a/ggml/src/ggml-cpu/quants.c +++ b/ggml/src/ggml-cpu/quants.c @@ -26,6 +26,10 @@ void quantize_row_q1_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, in quantize_row_q1_0_ref(x, y, k); } +void quantize_row_q2_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { + quantize_row_q2_0_ref(x, y, k); +} + void quantize_row_q4_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k) { quantize_row_q4_0_ref(x, y, k); } @@ -170,6 +174,53 @@ void ggml_vec_dot_q1_0_q8_0_generic(int n, float * GGML_RESTRICT s, size_t bs, c *s = sumf; } +void ggml_vec_dot_q2_0_q8_0_generic(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) { + const int qk = QK2_0; + const int nb = n / qk; + + assert(n % qk == 0); + assert(nrc == 1); + UNUSED(nrc); + UNUSED(bx); + UNUSED(by); + UNUSED(bs); + + const block_q2_0 * GGML_RESTRICT x = vx; + const block_q8_0 * GGML_RESTRICT y = vy; + + float sumf = 0.0f; + + for (int i = 0; i < nb; i++) { + const float d0 = GGML_CPU_FP16_TO_FP32(x[i].d); + + float sumi = 0.0f; + + // group 64: one Q2_0 block (64 weights) maps to two Q8_0 blocks (2 * 32 = 64) + for (int k = 0; k < 2; k++) { + const block_q8_0 * GGML_RESTRICT yb = &y[i * 2 + k]; + const float d1 = GGML_CPU_FP16_TO_FP32(yb->d); + int sumi_block = 0; + + const uint8_t * GGML_RESTRICT qs = &x[i].qs[k * 8]; + const int8_t * GGML_RESTRICT qy = yb->qs; + + for (int b = 0; b < 8; ++b) { + const uint8_t byte = qs[b]; + // Extract 4 two-bit values, map {0,1,2,3} -> {-1,0,1,2} + sumi_block += ((int)((byte >> 0) & 3) - 1) * qy[b*4 + 0]; + sumi_block += ((int)((byte >> 2) & 3) - 1) * qy[b*4 + 1]; + sumi_block += ((int)((byte >> 4) & 3) - 1) * qy[b*4 + 2]; + sumi_block += ((int)((byte >> 6) & 3) - 1) * qy[b*4 + 3]; + } + + sumi += d1 * sumi_block; + } + + sumf += d0 * sumi; + } + + *s = sumf; +} void ggml_vec_dot_q4_0_q8_0_generic(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) { const int qk = QK8_0; diff --git a/ggml/src/ggml-cpu/quants.h b/ggml/src/ggml-cpu/quants.h index d4bc87a1c..93ea7eeff 100644 --- a/ggml/src/ggml-cpu/quants.h +++ b/ggml/src/ggml-cpu/quants.h @@ -13,6 +13,7 @@ extern "C" { // Quantization void quantize_row_q1_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); +void quantize_row_q2_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); void quantize_row_q4_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); void quantize_row_q4_1(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); void quantize_row_q5_0(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); @@ -38,6 +39,7 @@ void quantize_row_iq4_xs (const float * GGML_RESTRICT x, void * GGML_RESTRICT y, // Dot product void ggml_vec_dot_q1_0_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); +void ggml_vec_dot_q2_0_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); void ggml_vec_dot_q4_0_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); void ggml_vec_dot_q4_1_q8_1(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); void ggml_vec_dot_q5_0_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); @@ -71,6 +73,7 @@ void quantize_row_q8_0_generic(const float * GGML_RESTRICT x, void * GGML_RESTRI void quantize_row_q8_1_generic(const float * GGML_RESTRICT x, void * GGML_RESTRICT vy, int64_t k); void quantize_row_q8_K_generic(const float * GGML_RESTRICT x, void * GGML_RESTRICT y, int64_t k); void ggml_vec_dot_q1_0_q8_0_generic(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); +void ggml_vec_dot_q2_0_q8_0_generic(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); void ggml_vec_dot_q4_0_q8_0_generic(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); void ggml_vec_dot_q4_1_q8_1_generic(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); void ggml_vec_dot_q5_0_q8_0_generic(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); diff --git a/ggml/src/ggml-cpu/simd-gemm.h b/ggml/src/ggml-cpu/simd-gemm.h index 4119d04f8..2ebd10051 100644 --- a/ggml/src/ggml-cpu/simd-gemm.h +++ b/ggml/src/ggml-cpu/simd-gemm.h @@ -78,7 +78,7 @@ static void simd_gemm( for (int64_t i = 0; i < GEMM_RM; i++) { float a = C[i * N + jj]; for (int64_t kk = 0; kk < K; kk++) { - a += A[i + kk] * B[kk * N + jj]; + a += A[i * K + kk] * B[kk * N + jj]; } C[i * N + jj] = a; } diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index daa4fe3c4..2a5c188fa 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -1512,12 +1512,16 @@ struct ggml_cuda_mm_fusion_args_host { const ggml_tensor * x_bias = nullptr; const ggml_tensor * gate = nullptr; const ggml_tensor * gate_bias = nullptr; + const ggml_tensor * x_scale = nullptr; + const ggml_tensor * gate_scale = nullptr; ggml_glu_op glu_op; }; struct ggml_cuda_mm_fusion_args_device { const void * x_bias = nullptr; const void * gate = nullptr; const void * gate_bias = nullptr; + const void * x_scale = nullptr; + const void * gate_scale = nullptr; ggml_glu_op glu_op; }; diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index ba6ec782c..928c1965d 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -1582,12 +1582,18 @@ static bool ggml_cuda_should_fuse_mul_mat(const ggml_tensor * ffn_up, const ggml_tensor * ffn_gate, const ggml_tensor * glu, const ggml_tensor * ffn_up_bias = nullptr, - const ggml_tensor * ffn_gate_bias = nullptr) { + const ggml_tensor * ffn_gate_bias = nullptr, + const ggml_tensor * ffn_up_scale = nullptr, + const ggml_tensor * ffn_gate_scale = nullptr) { const bool has_bias = ffn_up_bias != nullptr || ffn_gate_bias != nullptr; + const bool has_scale = ffn_up_scale != nullptr || ffn_gate_scale != nullptr; if (has_bias && (!ffn_up_bias || !ffn_gate_bias)) { return false; } + if (has_scale && (!ffn_up_scale || !ffn_gate_scale)) { + return false; + } const bool is_mul_mat = ffn_up->op == GGML_OP_MUL_MAT && ffn_gate->op == GGML_OP_MUL_MAT && glu->op == GGML_OP_GLU; const bool is_mul_mat_id = ffn_up->op == GGML_OP_MUL_MAT_ID && ffn_gate->op == GGML_OP_MUL_MAT_ID && glu->op == GGML_OP_GLU; @@ -1599,34 +1605,45 @@ static bool ggml_cuda_should_fuse_mul_mat(const ggml_tensor * ffn_up, } const ggml_op expected_bias_op = is_mul_mat ? GGML_OP_ADD : GGML_OP_ADD_ID; + const ggml_tensor * ffn_up_bias_src = has_scale ? ffn_up_scale : ffn_up; + const ggml_tensor * ffn_gate_bias_src = has_scale ? ffn_gate_scale : ffn_gate; + const ggml_tensor * ffn_up_out = has_bias ? ffn_up_bias : ffn_up_bias_src; + const ggml_tensor * ffn_gate_out = has_bias ? ffn_gate_bias : ffn_gate_bias_src; + + if (glu->src[0] != ffn_gate_out || glu->src[1] != ffn_up_out) { + return false; + } + + if (has_scale) { + if (ffn_up_scale->op != GGML_OP_MUL || ffn_gate_scale->op != GGML_OP_MUL) { + return false; + } + const bool up_has_mm = ffn_up_scale->src[0] == ffn_up || ffn_up_scale->src[1] == ffn_up; + const bool gate_has_mm = ffn_gate_scale->src[0] == ffn_gate || ffn_gate_scale->src[1] == ffn_gate; + if (!up_has_mm || !gate_has_mm) { + return false; + } + } if (has_bias) { if (ffn_up_bias->op != expected_bias_op || ffn_gate_bias->op != expected_bias_op) { return false; } - if (glu->src[0] != ffn_gate_bias || glu->src[1] != ffn_up_bias) { - return false; - } - if (expected_bias_op == GGML_OP_ADD) { - const bool up_has_mul = ffn_up_bias->src[0] == ffn_up || ffn_up_bias->src[1] == ffn_up; - const bool gate_has_mul = ffn_gate_bias->src[0] == ffn_gate || ffn_gate_bias->src[1] == ffn_gate; + const bool up_has_mul = ffn_up_bias->src[0] == ffn_up_bias_src || ffn_up_bias->src[1] == ffn_up_bias_src; + const bool gate_has_mul = ffn_gate_bias->src[0] == ffn_gate_bias_src || ffn_gate_bias->src[1] == ffn_gate_bias_src; if (!up_has_mul || !gate_has_mul) { return false; } } else { // GGML_OP_ADD_ID - if (ffn_up_bias->src[0] != ffn_up || ffn_gate_bias->src[0] != ffn_gate) { + if (ffn_up_bias->src[0] != ffn_up_bias_src || ffn_gate_bias->src[0] != ffn_gate_bias_src) { return false; } if (ffn_up_bias->src[2] != ffn_up->src[2] || ffn_gate_bias->src[2] != ffn_gate->src[2]) { return false; } } - } else { - if (glu->src[0] != ffn_gate && glu->src[1] != ffn_up) { - return false; - } } if (ffn_up->src[0]->type != ffn_gate->src[0]->type || !ggml_are_same_shape(ffn_up->src[0], ffn_gate->src[0]) || @@ -1638,7 +1655,7 @@ static bool ggml_cuda_should_fuse_mul_mat(const ggml_tensor * ffn_up, return false; } - if (ffn_up->src[2] && (ffn_up->src[2] != ffn_gate->src[2])) { + if (is_mul_mat_id && ffn_up->src[2] != ffn_gate->src[2]) { return false; } @@ -3212,10 +3229,240 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph bool fused_mul_mat_vec = false; int fused_node_count = 0; - // gate + glu + up + auto get_mul_mat_scale = [](const ggml_tensor * scale_node, const ggml_tensor * mm_node) -> const ggml_tensor * { + const bool scale_lhs_mm = scale_node->src[0] == mm_node; + const bool scale_rhs_mm = scale_node->src[1] == mm_node; + if (!scale_lhs_mm && !scale_rhs_mm) { + return nullptr; + } + + const ggml_tensor * scale = scale_lhs_mm ? scale_node->src[1] : scale_node->src[0]; + if (mm_node->src[0]->type != GGML_TYPE_NVFP4 || scale_node->type != GGML_TYPE_F32 || + scale->type != GGML_TYPE_F32 || !ggml_is_contiguous(scale) || ggml_nelements(scale) != 1 || + !ggml_are_same_shape(scale_node, mm_node)) { + return nullptr; + } + + return scale; + }; + + auto get_mul_mat_id_scale = [](const ggml_tensor * reshape, const ggml_tensor * repeat, const ggml_tensor * getrows, + const ggml_tensor * scale_node, const ggml_tensor * mm_node) -> const ggml_tensor * { + if (repeat->src[0] != reshape || getrows->src[0] != repeat || getrows->src[1] != mm_node->src[2]) { + return nullptr; + } + if (!((scale_node->src[0] == mm_node && scale_node->src[1] == getrows) || + (scale_node->src[0] == getrows && scale_node->src[1] == mm_node))) { + return nullptr; + } + + const ggml_tensor * scale = reshape->src[0]; + if (mm_node->src[0]->type != GGML_TYPE_NVFP4 || scale_node->type != GGML_TYPE_F32 || + scale->type != GGML_TYPE_F32 || !ggml_is_contiguous(scale) || ggml_nelements(scale) != mm_node->src[0]->ne[2] || + !ggml_are_same_shape(scale_node, mm_node)) { + return nullptr; + } + + return scale; + }; + + auto get_bias_tensor = [](const ggml_tensor * bias_node, const ggml_tensor * mul_node, ggml_op op_bias) -> const ggml_tensor * { + if (op_bias == GGML_OP_ADD) { + if (bias_node->src[0] == mul_node) { + return bias_node->src[1]; + } + if (bias_node->src[1] == mul_node) { + return bias_node->src[0]; + } + return nullptr; + } + GGML_ASSERT(op_bias == GGML_OP_ADD_ID); + GGML_ASSERT(bias_node->src[0] == mul_node); + return bias_node->src[1]; + }; + + // gate + glu + up, with optional scale/bias on both lanes. for (ggml_op op : { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT_ID }) { const ggml_op bias_op = op == GGML_OP_MUL_MAT ? GGML_OP_ADD : GGML_OP_ADD_ID; + if (op == GGML_OP_MUL_MAT) { + for (const bool with_bias : { false, true }) { + const int gate_idx = i; + const int gate_scale_idx = i + 1; + const int gate_bias_idx = with_bias ? i + 2 : -1; + const int up_idx = with_bias ? i + 3 : i + 2; + const int up_scale_idx = up_idx + 1; + const int up_bias_idx = with_bias ? up_idx + 2 : -1; + const int glu_idx = with_bias ? up_idx + 3 : up_idx + 2; + + const int out_nodes[] = { glu_idx }; + ggml_op ops[7]; + if (with_bias) { + ops[0] = op; + ops[1] = GGML_OP_MUL; + ops[2] = bias_op; + ops[3] = op; + ops[4] = GGML_OP_MUL; + ops[5] = bias_op; + ops[6] = GGML_OP_GLU; + } else { + ops[0] = op; + ops[1] = GGML_OP_MUL; + ops[2] = op; + ops[3] = GGML_OP_MUL; + ops[4] = GGML_OP_GLU; + } + const int n_ops = with_bias ? 7 : 5; + + if (!ggml_can_fuse_subgraph(cgraph, i, n_ops, ops, out_nodes, 1) || + !ggml_cuda_check_fusion_memory_ranges(cgraph, i, n_ops, out_nodes, 1)) { + continue; + } + + ggml_tensor * gate_n = cgraph->nodes[gate_idx]; + ggml_tensor * gate_scale_n = cgraph->nodes[gate_scale_idx]; + ggml_tensor * gate_out_n = with_bias ? cgraph->nodes[gate_bias_idx] : gate_scale_n; + ggml_tensor * up_n = cgraph->nodes[up_idx]; + ggml_tensor * up_scale_n = cgraph->nodes[up_scale_idx]; + ggml_tensor * up_out_n = with_bias ? cgraph->nodes[up_bias_idx] : up_scale_n; + const ggml_tensor * glu = cgraph->nodes[glu_idx]; + + if (!ggml_cuda_should_fuse_mul_mat(up_n, gate_n, glu, + with_bias ? up_out_n : nullptr, with_bias ? gate_out_n : nullptr, up_scale_n, gate_scale_n)) { + continue; + } + + const ggml_tensor * gate_scale = get_mul_mat_scale(gate_scale_n, gate_n); + const ggml_tensor * up_scale = get_mul_mat_scale(up_scale_n, up_n); + if (!gate_scale || !up_scale) { + continue; + } + + const ggml_tensor * up_bias = with_bias ? get_bias_tensor(up_out_n, up_scale_n, bias_op) : nullptr; + const ggml_tensor * gate_bias = with_bias ? get_bias_tensor(gate_out_n, gate_scale_n, bias_op) : nullptr; + if (with_bias && (!ggml_are_same_shape(gate_out_n->src[0], gate_out_n->src[1]) || + !ggml_are_same_shape(up_out_n->src[0], up_out_n->src[1]))) { + continue; + } + + const ggml_tensor * src0 = up_n->src[0]; + const ggml_tensor * src1 = up_n->src[1]; + const ggml_tensor * ids = up_n->src[2]; + + ggml_cuda_mm_fusion_args_host fusion_data{}; + fusion_data.gate = gate_n->src[0]; + fusion_data.x_bias = up_bias; + fusion_data.gate_bias = gate_bias; + fusion_data.x_scale = up_scale; + fusion_data.gate_scale = gate_scale; + fusion_data.glu_op = ggml_get_glu_op(glu); + + if (ggml_cuda_should_fuse_mul_mat_vec_q(up_n)) { + ggml_cuda_mul_mat_vec_q(*cuda_ctx, src0, src1, ids, cgraph->nodes[glu_idx], &fusion_data); + fused_mul_mat_vec = true; + fused_node_count = n_ops; + break; + } + } + + if (fused_mul_mat_vec) { + break; + } + } else { + for (const bool with_bias : { false, true }) { + const int gate_idx = i; + const int gate_scale_idx = i + 4; + const int gate_bias_idx = with_bias ? i + 5 : -1; + const int up_idx = with_bias ? i + 6 : i + 5; + const int up_scale_idx = up_idx + 4; + const int up_bias_idx = with_bias ? up_idx + 5 : -1; + const int glu_idx = with_bias ? up_idx + 6 : up_idx + 5; + + const int out_nodes[] = { glu_idx }; + ggml_op ops[13]; + if (with_bias) { + ops[0] = op; + ops[1] = GGML_OP_RESHAPE; + ops[2] = GGML_OP_REPEAT; + ops[3] = GGML_OP_GET_ROWS; + ops[4] = GGML_OP_MUL; + ops[5] = bias_op; + ops[6] = op; + ops[7] = GGML_OP_RESHAPE; + ops[8] = GGML_OP_REPEAT; + ops[9] = GGML_OP_GET_ROWS; + ops[10] = GGML_OP_MUL; + ops[11] = bias_op; + ops[12] = GGML_OP_GLU; + } else { + ops[0] = op; + ops[1] = GGML_OP_RESHAPE; + ops[2] = GGML_OP_REPEAT; + ops[3] = GGML_OP_GET_ROWS; + ops[4] = GGML_OP_MUL; + ops[5] = op; + ops[6] = GGML_OP_RESHAPE; + ops[7] = GGML_OP_REPEAT; + ops[8] = GGML_OP_GET_ROWS; + ops[9] = GGML_OP_MUL; + ops[10] = GGML_OP_GLU; + } + const int n_ops = with_bias ? 13 : 11; + + if (!ggml_can_fuse_subgraph(cgraph, i, n_ops, ops, out_nodes, 1) || + !ggml_cuda_check_fusion_memory_ranges(cgraph, i, n_ops, out_nodes, 1)) { + continue; + } + + ggml_tensor * gate_n = cgraph->nodes[gate_idx]; + ggml_tensor * gate_scale_n = cgraph->nodes[gate_scale_idx]; + ggml_tensor * gate_out_n = with_bias ? cgraph->nodes[gate_bias_idx] : gate_scale_n; + ggml_tensor * up_n = cgraph->nodes[up_idx]; + ggml_tensor * up_scale_n = cgraph->nodes[up_scale_idx]; + ggml_tensor * up_out_n = with_bias ? cgraph->nodes[up_bias_idx] : up_scale_n; + const ggml_tensor * glu = cgraph->nodes[glu_idx]; + + if (!ggml_cuda_should_fuse_mul_mat(up_n, gate_n, glu, + with_bias ? up_out_n : nullptr, with_bias ? gate_out_n : nullptr, up_scale_n, gate_scale_n)) { + continue; + } + + const ggml_tensor * gate_scale = get_mul_mat_id_scale(cgraph->nodes[gate_idx + 1], cgraph->nodes[gate_idx + 2], + cgraph->nodes[gate_idx + 3], gate_scale_n, gate_n); + const ggml_tensor * up_scale = get_mul_mat_id_scale(cgraph->nodes[up_idx + 1], cgraph->nodes[up_idx + 2], + cgraph->nodes[up_idx + 3], up_scale_n, up_n); + if (!gate_scale || !up_scale) { + continue; + } + + const ggml_tensor * up_bias = with_bias ? get_bias_tensor(up_out_n, up_scale_n, bias_op) : nullptr; + const ggml_tensor * gate_bias = with_bias ? get_bias_tensor(gate_out_n, gate_scale_n, bias_op) : nullptr; + + const ggml_tensor * src0 = up_n->src[0]; + const ggml_tensor * src1 = up_n->src[1]; + const ggml_tensor * ids = up_n->src[2]; + + ggml_cuda_mm_fusion_args_host fusion_data{}; + fusion_data.gate = gate_n->src[0]; + fusion_data.x_bias = up_bias; + fusion_data.gate_bias = gate_bias; + fusion_data.x_scale = up_scale; + fusion_data.gate_scale = gate_scale; + fusion_data.glu_op = ggml_get_glu_op(glu); + + if (ggml_cuda_should_fuse_mul_mat_vec_q(up_n)) { + ggml_cuda_mul_mat_vec_q(*cuda_ctx, src0, src1, ids, cgraph->nodes[glu_idx], &fusion_data); + fused_mul_mat_vec = true; + fused_node_count = n_ops; + break; + } + } + + if (fused_mul_mat_vec) { + break; + } + } + if (ggml_cuda_can_fuse(cgraph, i, { op, bias_op, op, bias_op, GGML_OP_GLU }, {})) { ggml_tensor * glu = cgraph->nodes[i + 4]; ggml_tensor * gate_bias_n = glu->src[0]; @@ -3235,23 +3482,8 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph continue; } - auto get_bias_tensor = [](const ggml_tensor * bias_node, const ggml_tensor * mul_node, ggml_op op_bias) { - if (op_bias == GGML_OP_ADD) { - if (bias_node->src[0] == mul_node) { - return bias_node->src[1]; - } - if (bias_node->src[1] == mul_node) { - return bias_node->src[0]; - } - return (ggml_tensor *) nullptr; - } - GGML_ASSERT(op_bias == GGML_OP_ADD_ID); - GGML_ASSERT(bias_node->src[0] == mul_node); - return bias_node->src[1]; - }; - - ggml_tensor * up_bias_tensor = get_bias_tensor(up_bias_n, up_n, bias_op); - ggml_tensor * gate_bias_tensor = get_bias_tensor(gate_bias_n, gate_n, bias_op); + const ggml_tensor * up_bias_tensor = get_bias_tensor(up_bias_n, up_n, bias_op); + const ggml_tensor * gate_bias_tensor = get_bias_tensor(gate_bias_n, gate_n, bias_op); if (!up_bias_tensor || !gate_bias_tensor) { continue; @@ -3339,7 +3571,95 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph fused_mul_mat_vec = false; fused_node_count = 0; - // gate + add + glu + up + add + // mul_mat + scale + optional bias + for (ggml_op op : { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT_ID }) { + const ggml_op bias_op = op == GGML_OP_MUL_MAT ? GGML_OP_ADD : GGML_OP_ADD_ID; + + for (const bool with_bias : { false, true }) { + const int n_ops = op == GGML_OP_MUL_MAT ? (with_bias ? 3 : 2) : (with_bias ? 6 : 5); + const int out_nodes[] = { i + n_ops - 1 }; + ggml_op ops[6]; + if (op == GGML_OP_MUL_MAT) { + if (with_bias) { + ops[0] = op; + ops[1] = GGML_OP_MUL; + ops[2] = bias_op; + } else { + ops[0] = op; + ops[1] = GGML_OP_MUL; + } + } else { + if (with_bias) { + ops[0] = op; + ops[1] = GGML_OP_RESHAPE; + ops[2] = GGML_OP_REPEAT; + ops[3] = GGML_OP_GET_ROWS; + ops[4] = GGML_OP_MUL; + ops[5] = bias_op; + } else { + ops[0] = op; + ops[1] = GGML_OP_RESHAPE; + ops[2] = GGML_OP_REPEAT; + ops[3] = GGML_OP_GET_ROWS; + ops[4] = GGML_OP_MUL; + } + } + + if (!ggml_can_fuse_subgraph(cgraph, i, n_ops, ops, out_nodes, 1) || + !ggml_cuda_check_fusion_memory_ranges(cgraph, i, n_ops, out_nodes, 1)) { + continue; + } + + ggml_tensor * mm_node = cgraph->nodes[i]; + ggml_tensor * scale_node = op == GGML_OP_MUL_MAT ? cgraph->nodes[i + 1] : cgraph->nodes[i + 4]; + ggml_tensor * out_node = with_bias ? cgraph->nodes[i + n_ops - 1] : scale_node; + + const ggml_tensor * scale = nullptr; + if (op == GGML_OP_MUL_MAT) { + scale = get_mul_mat_scale(scale_node, mm_node); + } else { + scale = get_mul_mat_id_scale(cgraph->nodes[i + 1], cgraph->nodes[i + 2], cgraph->nodes[i + 3], scale_node, mm_node); + } + if (!scale) { + continue; + } + + const ggml_tensor * bias = with_bias ? get_bias_tensor(out_node, scale_node, bias_op) : nullptr; + if (with_bias && !bias) { + continue; + } + if (with_bias && bias_op == GGML_OP_ADD && !ggml_are_same_shape(out_node->src[0], out_node->src[1])) { + continue; + } + if (with_bias && bias_op == GGML_OP_ADD_ID && out_node->src[2] != mm_node->src[2]) { + continue; + } + + const ggml_tensor * src0 = mm_node->src[0]; + const ggml_tensor * src1 = mm_node->src[1]; + const ggml_tensor * ids = mm_node->src[2]; + + ggml_cuda_mm_fusion_args_host fusion_data{}; + fusion_data.x_bias = bias; + fusion_data.x_scale = scale; + + if (ggml_cuda_should_fuse_mul_mat_vec_q(mm_node)) { + ggml_cuda_mul_mat_vec_q(*cuda_ctx, src0, src1, ids, out_node, &fusion_data); + fused_mul_mat_vec = true; + fused_node_count = n_ops; + break; + } + } + if (fused_mul_mat_vec) { + break; + } + } + + if (fused_mul_mat_vec) { + return fused_node_count - 1; + } + + // mul_mat + add for (ggml_op op : { GGML_OP_MUL_MAT, GGML_OP_MUL_MAT_ID }) { const ggml_op bias_op = op == GGML_OP_MUL_MAT ? GGML_OP_ADD : GGML_OP_ADD_ID; @@ -3570,12 +3890,6 @@ static void ggml_cuda_graph_evaluate_and_capture(ggml_backend_cuda_context * cud } } -#ifdef GGML_CUDA_DEBUG - const int nodes_fused = i - prev_i - 1; - if (nodes_fused > 0) { - GGML_LOG_INFO("nodes_fused: %d\n", nodes_fused); - } -#endif prev_i = i; if (ggml_cuda_is_view_or_noop(node)) { @@ -3589,6 +3903,12 @@ static void ggml_cuda_graph_evaluate_and_capture(ggml_backend_cuda_context * cud int nodes_to_skip = ggml_cuda_try_fuse(cuda_ctx, cgraph, i); if (nodes_to_skip != 0) { +#ifdef GGML_CUDA_DEBUG + const int last_fused = i + nodes_to_skip; + GGML_LOG_INFO("nodes_fused: %d, first: %s (%s), last: %s (%s)\n", + nodes_to_skip + 1, ggml_op_name(node->op), node->name, + ggml_op_name(cgraph->nodes[last_fused]->op), cgraph->nodes[last_fused]->name); +#endif i += nodes_to_skip; continue; } diff --git a/ggml/src/ggml-cuda/mmvq.cu b/ggml/src/ggml-cuda/mmvq.cu index 2278c7d9b..a48cc48b2 100644 --- a/ggml/src/ggml-cuda/mmvq.cu +++ b/ggml/src/ggml-cuda/mmvq.cu @@ -521,9 +521,13 @@ static __global__ void mul_mat_vec_q( bool use_gate = false; bool use_bias = false; bool use_gate_bias = false; + bool use_scale = false; + bool use_gate_scale = false; [[maybe_unused]] const void * vgate = nullptr; const float * x_bias = nullptr; const float * gate_bias = nullptr; + const float * x_scale = nullptr; + const float * gate_scale = nullptr; ggml_glu_op active_glu; if constexpr (has_fusion) { @@ -534,34 +538,47 @@ static __global__ void mul_mat_vec_q( x_bias = (const float *) fusion.x_bias; gate_bias = (const float *) fusion.gate_bias; active_glu = fusion.glu_op; + if constexpr (type == GGML_TYPE_NVFP4) { + use_scale = fusion.x_scale != nullptr; + use_gate_scale = fusion.gate_scale != nullptr && use_gate; + x_scale = (const float *) fusion.x_scale; + gate_scale = (const float *) fusion.gate_scale; + } } [[maybe_unused]] float x_biases[ncols_dst] = { 0.0f }; [[maybe_unused]] float gate_biases[ncols_dst] = { 0.0f }; + [[maybe_unused]] float x_scales; + [[maybe_unused]] float gate_scales; if constexpr (has_fusion) { + // 1. Hide latency by prefetching bias, gates and scales here + // 2. load only on threads that won't die after partial sum calculation const uint32_t channel_bias = ids ? channel_x : channel_dst; - if (use_bias) { - x_bias = x_bias + sample_dst*stride_sample_dst + channel_bias*stride_channel_dst + row0; - // 1. Hide latency by prefetching bias and gate here - // 2. load only on threads that won't die after partial sum calculation - if (threadIdx.x < rows_per_cuda_block && threadIdx.y == 0 && - (rows_per_cuda_block == 1 || uint32_t(row0 + threadIdx.x) < stride_col_dst)) { + if (threadIdx.x < rows_per_cuda_block && threadIdx.y == 0 && + (rows_per_cuda_block == 1 || uint32_t(row0 + threadIdx.x) < stride_col_dst)) { + if (use_bias) { + x_bias = x_bias + sample_dst * stride_sample_dst + channel_bias * stride_channel_dst + row0; #pragma unroll for (int j = 0; j < ncols_dst; ++j) { x_biases[j] = x_bias[j * stride_col_dst + threadIdx.x]; } } - } - if (use_gate_bias) { - gate_bias = gate_bias + sample_dst*stride_sample_dst + channel_bias*stride_channel_dst + row0; - if (threadIdx.x < rows_per_cuda_block && threadIdx.y == 0 && - (rows_per_cuda_block == 1 || uint32_t(row0 + threadIdx.x) < stride_col_dst)) { + if (use_gate_bias) { + gate_bias = gate_bias + sample_dst * stride_sample_dst + channel_bias * stride_channel_dst + row0; #pragma unroll for (int j = 0; j < ncols_dst; ++j) { gate_biases[j] = gate_bias[j * stride_col_dst + threadIdx.x]; } } + if constexpr (type == GGML_TYPE_NVFP4) { + if (use_scale) { + x_scales = x_scale[ids ? channel_x : 0]; + } + if (use_gate_scale) { + gate_scales = gate_scale[ids ? channel_x : 0]; + } + } } } @@ -643,11 +660,21 @@ static __global__ void mul_mat_vec_q( if (threadIdx.x < rows_per_cuda_block && (rows_per_cuda_block == 1 || uint32_t(row0 + threadIdx.x) < stride_col_dst)) { float result = tmp[j][threadIdx.x]; if constexpr (has_fusion) { + if constexpr (type == GGML_TYPE_NVFP4) { + if (use_scale) { + result *= x_scales; + } + } if (use_bias) { result += x_biases[j]; } if (use_gate) { float gate_value = tmp_gate[j][threadIdx.x]; + if constexpr (type == GGML_TYPE_NVFP4) { + if (use_gate_scale) { + gate_value *= gate_scales; + } + } if (use_gate_bias) { gate_value += gate_biases[j]; } @@ -673,7 +700,10 @@ static __global__ void mul_mat_vec_q( } if constexpr (!has_fusion) { - GGML_UNUSED_VARS(use_gate, use_bias, use_gate_bias, active_glu, gate_bias, x_bias, tmp_gate); + GGML_UNUSED_VARS(use_gate, use_bias, use_gate_bias, use_scale, use_gate_scale, active_glu, gate_bias, x_bias, x_scale, gate_scale, tmp_gate); + } + if constexpr (type != GGML_TYPE_NVFP4) { + GGML_UNUSED_VARS(use_scale, use_gate_scale, x_scale, gate_scale, x_scales, gate_scales); } } @@ -769,7 +799,8 @@ static void mul_mat_vec_q_switch_fusion( const dim3 & block_nums, const dim3 & block_dims, const int nbytes_shared, const uint32_t ids_stride, cudaStream_t stream) { - const bool has_fusion = fusion.gate != nullptr || fusion.x_bias != nullptr || fusion.gate_bias != nullptr; + const bool has_fusion = fusion.gate != nullptr || fusion.x_bias != nullptr || fusion.gate_bias != nullptr || + fusion.x_scale != nullptr || fusion.gate_scale != nullptr; if constexpr (c_ncols_dst == 1) { if (has_fusion) { const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(block_nums, block_dims, nbytes_shared, stream); @@ -834,7 +865,6 @@ static void mul_mat_vec_q_switch_ncols_dst( const int warp_size = ggml_cuda_info().devices[device].warp_size; const mmvq_parameter_table_id table_id = get_device_table_id(cc); - const bool has_fusion = fusion.gate != nullptr || fusion.x_bias != nullptr || fusion.gate_bias != nullptr; const bool has_ids = ids != nullptr; const auto should_use_small_k = [&](int c_ncols_dst) { @@ -973,8 +1003,6 @@ static void mul_mat_vec_q_switch_ncols_dst( GGML_ABORT("fatal error"); break; } - - GGML_UNUSED(has_fusion); } static void mul_mat_vec_q_switch_type( const void * vx, const ggml_type type_x, const void * vy, const int32_t * ids, const ggml_cuda_mm_fusion_args_device fusion, float * dst, @@ -1154,6 +1182,9 @@ void ggml_cuda_mul_mat_vec_q( if (fusion) { GGML_ASSERT( !ids || dst->ne[2] == 1); GGML_ASSERT( ids || dst->ne[1] == 1); + // Scale fusion is only allowed for NVFP4 currently as the cost of checking this at run-time in the prologue is + // non-negligible for some models such as gpt-oss-20b + GGML_ASSERT((fusion->x_scale == nullptr && fusion->gate_scale == nullptr) || src0->type == GGML_TYPE_NVFP4); if (fusion->x_bias) { GGML_ASSERT(fusion->x_bias->type == GGML_TYPE_F32); @@ -1171,6 +1202,18 @@ void ggml_cuda_mul_mat_vec_q( GGML_ASSERT(!ids || fusion->gate_bias->ne[1] == src0->ne[2]); fusion_local.gate_bias = fusion->gate_bias->data; } + if (fusion->x_scale) { + GGML_ASSERT(fusion->x_scale->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(fusion->x_scale)); + GGML_ASSERT(ggml_nelements(fusion->x_scale) == (ids ? src0->ne[2] : 1)); + fusion_local.x_scale = fusion->x_scale->data; + } + if (fusion->gate_scale) { + GGML_ASSERT(fusion->gate_scale->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(fusion->gate_scale)); + GGML_ASSERT(ggml_nelements(fusion->gate_scale) == (ids ? src0->ne[2] : 1)); + fusion_local.gate_scale = fusion->gate_scale->data; + } fusion_local.glu_op = fusion->glu_op; } diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index b2478fe97..e7ac21a2b 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -160,11 +160,15 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_get_rows(ggml_me return res; } -ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_set_rows(ggml_metal_library_t lib, ggml_type tidx, ggml_type tdst) { +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_set_rows(ggml_metal_library_t lib, const ggml_tensor * op) { char base[256]; char name[256]; - snprintf(base, 256, "kernel_set_rows_%s_%s", ggml_type_name(tdst), ggml_type_name(tidx)); + const auto tsrc = op->src[0]->type; + const auto tidx = op->src[1]->type; + const auto tdst = op->type; + + snprintf(base, 256, "kernel_set_rows_%s_%s_%s", ggml_type_name(tsrc), ggml_type_name(tidx), ggml_type_name(tdst)); snprintf(name, 256, "%s", base); ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index 7c8bde362..dc75a34b2 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -112,7 +112,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cpy struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_pool_1d (ggml_metal_library_t lib, const struct ggml_tensor * op, enum ggml_op_pool op_pool); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_pool_2d (ggml_metal_library_t lib, const struct ggml_tensor * op, enum ggml_op_pool op_pool); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_get_rows (ggml_metal_library_t lib, enum ggml_type tsrc); -struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_set_rows (ggml_metal_library_t lib, enum ggml_type tidx, enum ggml_type tdst); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_set_rows (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_diag (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_repeat (ggml_metal_library_t lib, enum ggml_type tsrc); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_concat (ggml_metal_library_t lib, enum ggml_type tsrc); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 338bf47d5..d51f30b13 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1340,7 +1340,7 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te return op->src[0]->type != GGML_TYPE_NVFP4; case GGML_OP_SET_ROWS: { - if (op->src[0]->type != GGML_TYPE_F32) { + if (op->src[0]->type != GGML_TYPE_F32 && op->src[0]->type != GGML_TYPE_F16) { return false; } diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index d2bc1254a..d196bae4f 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -1202,7 +1202,7 @@ int ggml_metal_op_set_rows(ggml_metal_op_t ctx, int idx) { GGML_TENSOR_LOCALS( int32_t, ne, op, ne); GGML_TENSOR_LOCALS(uint64_t, nb, op, nb); - auto pipeline = ggml_metal_library_get_pipeline_set_rows(lib, op->src[1]->type, op->type); + auto pipeline = ggml_metal_library_get_pipeline_set_rows(lib, op); const int32_t nk0 = ne0/ggml_blck_size(op->type); diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal index ad50372f4..dcb6803f5 100644 --- a/ggml/src/ggml-metal/ggml-metal.metal +++ b/ggml/src/ggml-metal/ggml-metal.metal @@ -42,6 +42,8 @@ typedef matrix bfloat4x4; typedef matrix bfloat2x4; #endif +#define QK_NL 16 + constexpr constant static float kvalues_iq4nl_f[16] = { -127.f, -104.f, -83.f, -65.f, -49.f, -35.f, -22.f, -10.f, 1.f, 13.f, 25.f, 38.f, 53.f, 69.f, 89.f, 113.f }; @@ -9386,7 +9388,40 @@ kernel void kernel_get_rows_f( } } -template +typedef decltype(kernel_get_rows_f) get_rows_f_t; + +template [[host_name("kernel_get_rows_f32")]] kernel get_rows_f_t kernel_get_rows_f; +template [[host_name("kernel_get_rows_f16")]] kernel get_rows_f_t kernel_get_rows_f; +template [[host_name("kernel_get_rows_i32")]] kernel get_rows_f_t kernel_get_rows_f; +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_get_rows_bf16")]] kernel get_rows_f_t kernel_get_rows_f; +#endif + +typedef decltype(kernel_get_rows_q) get_rows_q_t; + +template [[host_name("kernel_get_rows_q1_0")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q4_0")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q4_1")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q5_0")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q5_1")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q8_0")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_mxfp4")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q2_K")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q3_K")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q4_K")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q5_K")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_q6_K")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq2_xxs")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq2_xs")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq3_xxs")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq3_s")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq2_s")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq1_s")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq1_m")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq4_nl")]] kernel get_rows_q_t kernel_get_rows_q; +template [[host_name("kernel_get_rows_iq4_xs")]] kernel get_rows_q_t kernel_get_rows_q; + +template kernel void kernel_set_rows_q32( constant ggml_metal_kargs_set_rows & args, device const void * src0, @@ -9410,14 +9445,14 @@ kernel void kernel_set_rows_q32( const TI i1 = ((const device TI *) ((const device char *) src1 + i10*args.nb10 + i11*args.nb11 + i12*args.nb12))[0]; device block_q * dst_row = ( device block_q *) (( device char *) dst + i1*args.nb1 + i02*args.nb2 + i03*args.nb3); - const device float * src_row = (const device float *) ((const device char *) src0 + i01*args.nb01 + i02*args.nb02 + i03*args.nb03); + const device TS * src_row = (const device TS *) ((const device char *) src0 + i01*args.nb01 + i02*args.nb02 + i03*args.nb03); for (int ind = tiitg%tptg.x; ind < args.nk0; ind += tptg.x) { quantize_func(src_row + 32*ind, dst_row[ind]); } } -template +template kernel void kernel_set_rows_f( constant ggml_metal_kargs_set_rows & args, device const void * src0, @@ -9440,14 +9475,47 @@ kernel void kernel_set_rows_f( const int32_t i10 = i01; const TI i1 = ((const device TI *) ((const device char *) src1 + i10*args.nb10 + i11*args.nb11 + i12*args.nb12))[0]; - device T * dst_row = ( device T *) (( device char *) dst + i1*args.nb1 + i02*args.nb2 + i03*args.nb3); - const device float * src_row = (const device float *) ((const device char *) src0 + i01*args.nb01 + i02*args.nb02 + i03*args.nb03); + device TD * dst_row = ( device TD *) (( device char *) dst + i1*args.nb1 + i02*args.nb2 + i03*args.nb3); + const device TS * src_row = (const device TS *) ((const device char *) src0 + i01*args.nb01 + i02*args.nb02 + i03*args.nb03); for (int ind = tiitg%tptg.x; ind < args.nk0; ind += tptg.x) { - dst_row[ind] = (T) src_row[ind]; + dst_row[ind] = (TD) src_row[ind]; } } +typedef decltype(kernel_set_rows_f) set_rows_f_t; + +template [[host_name("kernel_set_rows_f32_i64_f32")]] kernel set_rows_f_t kernel_set_rows_f; +template [[host_name("kernel_set_rows_f32_i32_f32")]] kernel set_rows_f_t kernel_set_rows_f; +template [[host_name("kernel_set_rows_f32_i64_f16")]] kernel set_rows_f_t kernel_set_rows_f; +template [[host_name("kernel_set_rows_f32_i32_f16")]] kernel set_rows_f_t kernel_set_rows_f; +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_set_rows_f32_i64_bf16")]] kernel set_rows_f_t kernel_set_rows_f; +template [[host_name("kernel_set_rows_f32_i32_bf16")]] kernel set_rows_f_t kernel_set_rows_f; +#endif + +template [[host_name("kernel_set_rows_f16_i64_f16")]] kernel set_rows_f_t kernel_set_rows_f; +template [[host_name("kernel_set_rows_f16_i32_f16")]] kernel set_rows_f_t kernel_set_rows_f; +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_set_rows_bf16_i64_bf16")]] kernel set_rows_f_t kernel_set_rows_f; +template [[host_name("kernel_set_rows_bf16_i32_bf16")]] kernel set_rows_f_t kernel_set_rows_f; +#endif + +typedef decltype(kernel_set_rows_q32) set_rows_q32_t; + +template [[host_name("kernel_set_rows_f32_i64_q8_0")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i32_q8_0")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i64_q4_0")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i32_q4_0")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i64_q4_1")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i32_q4_1")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i64_q5_0")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i32_q5_0")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i64_q5_1")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i32_q5_1")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i64_iq4_nl")]] kernel set_rows_q32_t kernel_set_rows_q32; +template [[host_name("kernel_set_rows_f32_i32_iq4_nl")]] kernel set_rows_q32_t kernel_set_rows_q32; + kernel void kernel_diag_f32( constant ggml_metal_kargs_diag & args, device const char * src0, @@ -10190,75 +10258,6 @@ kernel void kernel_mul_mm_id( } } -#define QK_NL 16 - -// -// get rows -// - -typedef decltype(kernel_get_rows_f) get_rows_f_t; - -template [[host_name("kernel_get_rows_f32")]] kernel get_rows_f_t kernel_get_rows_f; -template [[host_name("kernel_get_rows_f16")]] kernel get_rows_f_t kernel_get_rows_f; -template [[host_name("kernel_get_rows_i32")]] kernel get_rows_f_t kernel_get_rows_f; -#if defined(GGML_METAL_HAS_BF16) -template [[host_name("kernel_get_rows_bf16")]] kernel get_rows_f_t kernel_get_rows_f; -#endif - -typedef decltype(kernel_get_rows_q) get_rows_q_t; - -template [[host_name("kernel_get_rows_q1_0")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q4_0")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q4_1")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q5_0")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q5_1")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q8_0")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_mxfp4")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q2_K")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q3_K")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q4_K")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q5_K")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_q6_K")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq2_xxs")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq2_xs")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq3_xxs")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq3_s")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq2_s")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq1_s")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq1_m")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq4_nl")]] kernel get_rows_q_t kernel_get_rows_q; -template [[host_name("kernel_get_rows_iq4_xs")]] kernel get_rows_q_t kernel_get_rows_q; - -// -// set rows -// - -typedef decltype(kernel_set_rows_f) set_rows_f_t; - -template [[host_name("kernel_set_rows_f32_i64")]] kernel set_rows_f_t kernel_set_rows_f; -template [[host_name("kernel_set_rows_f32_i32")]] kernel set_rows_f_t kernel_set_rows_f; -template [[host_name("kernel_set_rows_f16_i64")]] kernel set_rows_f_t kernel_set_rows_f; -template [[host_name("kernel_set_rows_f16_i32")]] kernel set_rows_f_t kernel_set_rows_f; -#if defined(GGML_METAL_HAS_BF16) -template [[host_name("kernel_set_rows_bf16_i64")]] kernel set_rows_f_t kernel_set_rows_f; -template [[host_name("kernel_set_rows_bf16_i32")]] kernel set_rows_f_t kernel_set_rows_f; -#endif - -typedef decltype(kernel_set_rows_q32) set_rows_q32_t; - -template [[host_name("kernel_set_rows_q8_0_i64")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q8_0_i32")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q4_0_i64")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q4_0_i32")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q4_1_i64")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q4_1_i32")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q5_0_i64")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q5_0_i32")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q5_1_i64")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_q5_1_i32")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_iq4_nl_i64")]] kernel set_rows_q32_t kernel_set_rows_q32; -template [[host_name("kernel_set_rows_iq4_nl_i32")]] kernel set_rows_q32_t kernel_set_rows_q32; - // // matrix-matrix multiplication // diff --git a/ggml/src/ggml-quants.c b/ggml/src/ggml-quants.c index 15d231f70..1ebc50a76 100644 --- a/ggml/src/ggml-quants.c +++ b/ggml/src/ggml-quants.c @@ -71,6 +71,44 @@ void quantize_row_q1_0_ref(const float * GGML_RESTRICT x, block_q1_0 * GGML_REST } } +void quantize_row_q2_0_ref(const float * GGML_RESTRICT x, block_q2_0 * GGML_RESTRICT y, int64_t k) { + static const int qk = QK2_0; + + assert(k % qk == 0); + + const int nb = k / qk; + + for (int i = 0; i < nb; i++) { + // Compute scale as max absolute value in the block + float amax = 0.0f; + for (int j = 0; j < qk; j++) { + const float a = fabsf(x[i*qk + j]); + if (a > amax) amax = a; + } + const float d = amax; + const float id = d > 0.0f ? 1.0f / d : 0.0f; + + y[i].d = GGML_FP32_TO_FP16(d); + + // Clear quant bytes + for (int j = 0; j < qk / 4; ++j) { + y[i].qs[j] = 0; + } + + // Encode 2-bit values: round(w/d) clamped to [-1, 2], then add 1 + // 00 (-1) = -scale, 01 (0) = 0, 10 (+1) = +scale, 11 (+2) = 2*scale + for (int j = 0; j < qk; ++j) { + const float w = x[i*qk + j]; + int q = (int)roundf(w * id) + 1; + if (q < 0) q = 0; + if (q > 3) q = 3; + const int byte_index = j / 4; + const int bit_offset = (j % 4) * 2; + y[i].qs[byte_index] |= ((uint8_t)q << bit_offset); + } + } +} + // reference implementation for deterministic creation of model files void quantize_row_q4_0_ref(const float * GGML_RESTRICT x, block_q4_0 * GGML_RESTRICT y, int64_t k) { static const int qk = QK4_0; @@ -398,6 +436,26 @@ void dequantize_row_q1_0(const block_q1_0 * GGML_RESTRICT x, float * GGML_RESTRI } } +void dequantize_row_q2_0(const block_q2_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { + static const int qk = QK2_0; + + assert(k % qk == 0); + + const int nb = k / qk; + + for (int i = 0; i < nb; i++) { + const float d = GGML_FP16_TO_FP32(x[i].d); + + for (int j = 0; j < qk; ++j) { + const int byte_index = j / 4; + const int bit_offset = (j % 4) * 2; + const uint8_t q = (x[i].qs[byte_index] >> bit_offset) & 0x03; + // 00=-1, 01=0, 10=+1, 11=+2 + y[i*qk + j] = ((int)q - 1) * d; + } + } +} + void dequantize_row_q4_0(const block_q4_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k) { static const int qk = QK4_0; @@ -2052,6 +2110,20 @@ size_t quantize_q1_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, return nrow * row_size; } +size_t quantize_q2_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrow, int64_t n_per_row, const float * quant_weights) { + if (!quant_weights) { + quantize_row_q2_0_ref(src, dst, (int64_t)nrow*n_per_row); + return nrow * ggml_row_size(GGML_TYPE_Q2_0, n_per_row); + } + size_t row_size = ggml_row_size(GGML_TYPE_Q2_0, n_per_row); + char * qrow = (char *)dst; + for (int64_t row = 0; row < nrow; ++row) { + quantize_row_q2_0_ref(src, (block_q2_0*)qrow, n_per_row); + src += n_per_row; + qrow += row_size; + } + return nrow * row_size; +} size_t quantize_q4_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrow, int64_t n_per_row, const float * quant_weights) { if (!quant_weights) { @@ -5461,6 +5533,10 @@ bool ggml_validate_row_data(enum ggml_type type, const void * data, size_t nbyte { VALIDATE_ROW_DATA_D_F16_IMPL(block_q1_0, data, nb); } break; + case GGML_TYPE_Q2_0: + { + VALIDATE_ROW_DATA_D_F16_IMPL(block_q2_0, data, nb); + } break; case GGML_TYPE_Q4_0: { VALIDATE_ROW_DATA_D_F16_IMPL(block_q4_0, data, nb); diff --git a/ggml/src/ggml-quants.h b/ggml/src/ggml-quants.h index d56c86da8..75188f1af 100644 --- a/ggml/src/ggml-quants.h +++ b/ggml/src/ggml-quants.h @@ -15,6 +15,7 @@ extern "C" { // Quantization GGML_API void quantize_row_q1_0_ref(const float * GGML_RESTRICT x, block_q1_0 * GGML_RESTRICT y, int64_t k); +GGML_API void quantize_row_q2_0_ref(const float * GGML_RESTRICT x, block_q2_0 * GGML_RESTRICT y, int64_t k); GGML_API void quantize_row_q4_0_ref(const float * GGML_RESTRICT x, block_q4_0 * GGML_RESTRICT y, int64_t k); GGML_API void quantize_row_q4_1_ref(const float * GGML_RESTRICT x, block_q4_1 * GGML_RESTRICT y, int64_t k); GGML_API void quantize_row_q5_0_ref(const float * GGML_RESTRICT x, block_q5_0 * GGML_RESTRICT y, int64_t k); @@ -43,6 +44,7 @@ GGML_API void quantize_row_iq2_s_ref (const float * GGML_RESTRICT x, block_iq2_ // Dequantization GGML_API void dequantize_row_q1_0(const block_q1_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); +GGML_API void dequantize_row_q2_0(const block_q2_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); GGML_API void dequantize_row_q4_0(const block_q4_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); GGML_API void dequantize_row_q4_1(const block_q4_1 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); GGML_API void dequantize_row_q5_0(const block_q5_0 * GGML_RESTRICT x, float * GGML_RESTRICT y, int64_t k); @@ -93,6 +95,7 @@ GGML_API size_t quantize_q4_K(const float * GGML_RESTRICT src, void * GGML_RESTR GGML_API size_t quantize_q5_K(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); GGML_API size_t quantize_q6_K(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); GGML_API size_t quantize_q1_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); +GGML_API size_t quantize_q2_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); GGML_API size_t quantize_q4_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); GGML_API size_t quantize_q4_1(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); GGML_API size_t quantize_q5_0(const float * GGML_RESTRICT src, void * GGML_RESTRICT dst, int64_t nrows, int64_t n_per_row, const float * imatrix); diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index ae740c1be..6ef6e0d1e 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -685,6 +685,14 @@ static const struct ggml_type_traits type_traits[GGML_TYPE_COUNT] = { .to_float = (ggml_to_float_t) dequantize_row_q1_0, .from_float_ref = (ggml_from_float_t) quantize_row_q1_0_ref, }, + [GGML_TYPE_Q2_0] = { + .type_name = "q2_0", + .blck_size = QK2_0, + .type_size = sizeof(block_q2_0), + .is_quantized = true, + .to_float = (ggml_to_float_t) dequantize_row_q2_0, + .from_float_ref = (ggml_from_float_t) quantize_row_q2_0_ref, + }, [GGML_TYPE_Q4_0] = { .type_name = "q4_0", .blck_size = QK4_0, @@ -1433,6 +1441,7 @@ enum ggml_type ggml_ftype_to_ggml_type(enum ggml_ftype ftype) { case GGML_FTYPE_MOSTLY_Q4_0: wtype = GGML_TYPE_Q4_0; break; case GGML_FTYPE_MOSTLY_Q4_1: wtype = GGML_TYPE_Q4_1; break; case GGML_FTYPE_MOSTLY_Q1_0: wtype = GGML_TYPE_Q1_0; break; + case GGML_FTYPE_MOSTLY_Q2_0: wtype = GGML_TYPE_Q2_0; break; case GGML_FTYPE_MOSTLY_Q5_0: wtype = GGML_TYPE_Q5_0; break; case GGML_FTYPE_MOSTLY_Q5_1: wtype = GGML_TYPE_Q5_1; break; case GGML_FTYPE_MOSTLY_Q8_0: wtype = GGML_TYPE_Q8_0; break; @@ -3933,7 +3942,7 @@ struct ggml_tensor * ggml_set_rows( GGML_ASSERT(b->ne[2] % c->ne[1] == 0); GGML_ASSERT(b->ne[3] % c->ne[2] == 0); GGML_ASSERT(c->ne[3] == 1); - GGML_ASSERT(b->type == GGML_TYPE_F32); + GGML_ASSERT(b->type == GGML_TYPE_F32 || b->type == GGML_TYPE_F16); GGML_ASSERT(c->type == GGML_TYPE_I64 || c->type == GGML_TYPE_I32); GGML_ASSERT(ggml_is_contiguous_rows(a)); @@ -7435,6 +7444,10 @@ static int ggml_node_list_find_tensor(const struct ggml_cgraph * cgraph, return -1; } +static bool ggml_is_constant(const struct ggml_tensor * tensor) { + return tensor->buffer != NULL && ggml_backend_buffer_get_usage(tensor->buffer) == GGML_BACKEND_BUFFER_USAGE_WEIGHTS && (tensor->flags & GGML_TENSOR_FLAG_PARAM) == 0; +} + bool ggml_can_fuse_subgraph_ext(const struct ggml_cgraph * cgraph, const int * node_idxs, int count, @@ -7480,10 +7493,11 @@ bool ggml_can_fuse_subgraph_ext(const struct ggml_cgraph * cgraph, return false; } - // if node is a view, check if the view_src and all it's parent view_srcs are within the subgraph + // if node is a view, check if the view_src and all its parent view_srcs are within the subgraph. + // external view sources are allowed only for weight tensors, which are constant for this graph execution. struct ggml_tensor * view_src = node->view_src; while (view_src) { - if (ggml_node_list_find_tensor(cgraph, node_idxs, count, view_src) == -1) { + if (ggml_node_list_find_tensor(cgraph, node_idxs, count, view_src) == -1 && !ggml_is_constant(view_src)) { return false; } view_src = view_src->view_src; @@ -7755,6 +7769,7 @@ size_t ggml_quantize_chunk( switch (type) { case GGML_TYPE_Q1_0: result = quantize_q1_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; + case GGML_TYPE_Q2_0: result = quantize_q2_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; case GGML_TYPE_Q4_0: result = quantize_q4_0 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; case GGML_TYPE_Q4_1: result = quantize_q4_1 (src + start, (char *) dst + start_row * row_size, nrows, n_per_row, imatrix); break; case GGML_TYPE_Q5_0: result = quantize_q5_0 (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 cd4cdef89..869e436ac 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -4533,6 +4533,7 @@ class GGMLQuantizationType(IntEnum): MXFP4 = 39 NVFP4 = 40 Q1_0 = 41 + Q2_0 = 42 class ExpertGatingFuncType(IntEnum): @@ -4588,6 +4589,7 @@ class LlamaFileType(IntEnum): MOSTLY_MXFP4_MOE = 38 # except 1d tensors MOSTLY_NVFP4 = 39 # except 1d tensors MOSTLY_Q1_0 = 40 # except 1d tensors + MOSTLY_Q2_0 = 41 # except 1d tensors GUESSED = 1024 # not specified in the model file @@ -4713,6 +4715,7 @@ GGML_QUANT_SIZES: dict[GGMLQuantizationType, tuple[int, int]] = { GGMLQuantizationType.MXFP4: (32, 1 + 16), GGMLQuantizationType.NVFP4: (64, 4 + 32), GGMLQuantizationType.Q1_0: (128, 2 + 16), + GGMLQuantizationType.Q2_0: (64, 2 + 16), } diff --git a/include/llama.h b/include/llama.h index 0df6fdf3f..c4e3a1bf1 100644 --- a/include/llama.h +++ b/include/llama.h @@ -158,6 +158,7 @@ extern "C" { LLAMA_FTYPE_MOSTLY_MXFP4_MOE = 38, // except 1d tensors LLAMA_FTYPE_MOSTLY_NVFP4 = 39, // except 1d tensors LLAMA_FTYPE_MOSTLY_Q1_0 = 40, // except 1d tensors + LLAMA_FTYPE_MOSTLY_Q2_0 = 41, // except 1d tensors LLAMA_FTYPE_GUESSED = 1024, // not specified in the model file }; diff --git a/src/llama-batch.cpp b/src/llama-batch.cpp index 6bf76939c..5436717c4 100644 --- a/src/llama-batch.cpp +++ b/src/llama-batch.cpp @@ -505,7 +505,7 @@ llama_ubatch llama_batch_allocr::split_simple(uint32_t n_ubatch) { return ubatch_add(idxs, idxs.size(), false); } -llama_ubatch llama_batch_allocr::split_equal(uint32_t n_ubatch, bool sequential) { +llama_ubatch llama_batch_allocr::split_equal(uint32_t n_ubatch, bool sequential, uint32_t n_keep_tail) { if (sequential && has_cpl) { LLAMA_LOG_ERROR("%s: sequential split is not supported when there are coupled sequences in the input batch (you may need to use the -kvu flag)\n", __func__); @@ -548,7 +548,7 @@ llama_ubatch llama_batch_allocr::split_equal(uint32_t n_ubatch, bool sequential) } } - const uint32_t n_seqs = cur_seq_set.size(); + uint32_t n_seqs = cur_seq_set.size(); // we are done if (n_seqs == 0) { @@ -569,7 +569,7 @@ llama_ubatch llama_batch_allocr::split_equal(uint32_t n_ubatch, bool sequential) std::vector idxs_per_seq(n_seqs); while (true) { - // we can only add new n_seq_tokens tokens if all the sequence sets have at least one more unused token and + // we can only add new n_seq_tokens tokens if all the sequence sets have at least 1 more unused tokens and // if we haven't reached n_ubatch bool can_expand = true; @@ -600,6 +600,72 @@ llama_ubatch llama_batch_allocr::split_equal(uint32_t n_ubatch, bool sequential) } } + // if n_keep_tail > 0, keep only the seqs that either finish in this ubatch or have at least + // n_keep_tail tokens remaining for a future ubatch, so that the trailing n_keep_tail tokens + // of each seq are never split across ubatches + if (n_keep_tail > 0) { + GGML_ASSERT(n_ubatch > n_keep_tail); + + auto n_remaining = [&](uint32_t s) { + return (uint32_t) (seq_set_map[cur_seq_set[s]].size() - cur_idx[s]); + }; + + // keep the longest prefix of seqs that satisfy the constraint, to preserve sequential seq ids + uint32_t n_keep = 0; + while (n_keep < n_seqs) { + const uint32_t remaining = n_remaining(n_keep); + + if (remaining != 0 && remaining < n_keep_tail) { + break; + } + + n_keep++; + } + + // all seqs violate the constraint - resolve the first one directly and emit it alone + if (n_keep == 0) { + auto & idxs = idxs_per_seq[0]; + + const auto & seq_idxs = seq_set_map[cur_seq_set[0]]; + + if (idxs.size() + n_remaining(0) <= n_ubatch) { + // extend the seq to completion + while (n_remaining(0) > 0) { + const int32_t idx = seq_idxs[cur_idx[0]]; + + idxs.push_back(idx); + + used[idx] = true; + ++n_used; + + ++cur_idx[0]; + } + } else { + // truncate the seq so that at least n_keep_tail tokens remain + while (n_remaining(0) < n_keep_tail) { + used[idxs.back()] = false; + --n_used; + + idxs.pop_back(); + + --cur_idx[0]; + } + } + + n_keep = 1; + } + + // return the tokens of the deferred seqs back to the pool + for (uint32_t s = n_keep; s < n_seqs; ++s) { + for (const int32_t idx : idxs_per_seq[s]) { + used[idx] = false; + --n_used; + } + } + + n_seqs = n_keep; + } + // concat the per-sequence-set lists std::vector idxs; @@ -814,7 +880,7 @@ void llama_batch_allocr::ubatch_print(const llama_ubatch & ubatch, int debug) { LLAMA_LOG_DEBUG("%s: output = %p\n", __func__, (void *) ubatch.output); LLAMA_LOG_DEBUG("%s: n_outputs = %d\n", __func__, n_outputs); - if (debug > 1) { + if (debug > 0) { int seq_id_max = 0; for (uint32_t i = 0; i < ubatch.n_tokens; ++i) { for (int s = 0; s < ubatch.n_seq_id[i]; ++s) { diff --git a/src/llama-batch.h b/src/llama-batch.h index f77520e86..a3d1889d4 100644 --- a/src/llama-batch.h +++ b/src/llama-batch.h @@ -104,7 +104,8 @@ public: // make ubatches of equal-length sequences sets // if sequential == true, the tokens in the ubatch will have increasing sequential sequence ids - llama_ubatch split_equal(uint32_t n_ubatch, bool sequential); + // n_keep_tail = minimum trailing tokens of a seq that must land in the same ubatch + llama_ubatch split_equal(uint32_t n_ubatch, bool sequential, uint32_t n_keep_tail); // sequence-set-wise split - each ubatch contains a single sequence-set llama_ubatch split_seq(uint32_t n_ubatch); diff --git a/src/llama-kv-cache-dsa.cpp b/src/llama-kv-cache-dsa.cpp index 916ab6537..241c50365 100644 --- a/src/llama-kv-cache-dsa.cpp +++ b/src/llama-kv-cache-dsa.cpp @@ -113,7 +113,7 @@ llama_memory_context_ptr llama_kv_cache_dsa::init_batch( std::vector ubatches; while (true) { - auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true); + auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0); if (ubatch.n_tokens == 0) { break; diff --git a/src/llama-kv-cache-dsv4.cpp b/src/llama-kv-cache-dsv4.cpp index 3a698d719..9fccf347e 100644 --- a/src/llama-kv-cache-dsv4.cpp +++ b/src/llama-kv-cache-dsv4.cpp @@ -1110,7 +1110,7 @@ llama_memory_context_ptr llama_kv_cache_dsv4::init_batch( if (has_coupled) { ubatch = balloc.split_seq(n_ubatch); } else { - ubatch = balloc.split_equal(n_ubatch, raw_per_seq || comp_per_seq); + ubatch = balloc.split_equal(n_ubatch, raw_per_seq || comp_per_seq, 0); } if (ubatch.n_tokens == 0) { diff --git a/src/llama-kv-cache-iswa.cpp b/src/llama-kv-cache-iswa.cpp index 547533cfd..6c5531418 100644 --- a/src/llama-kv-cache-iswa.cpp +++ b/src/llama-kv-cache-iswa.cpp @@ -219,7 +219,7 @@ llama_memory_context_ptr llama_kv_cache_iswa::init_batch(llama_batch_allocr & ba std::vector ubatches; while (true) { - auto ubatch = balloc.split_equal(n_ubatch, !unified); + auto ubatch = balloc.split_equal(n_ubatch, !unified, 0); if (ubatch.n_tokens == 0) { break; diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index 04e7123f6..064912a44 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -706,7 +706,7 @@ llama_memory_context_ptr llama_kv_cache::init_batch( std::vector ubatches; while (true) { - auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true); + auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0); if (ubatch.n_tokens == 0) { break; diff --git a/src/llama-memory-hybrid-iswa.cpp b/src/llama-memory-hybrid-iswa.cpp index c7d4bcd41..06f7fd542 100644 --- a/src/llama-memory-hybrid-iswa.cpp +++ b/src/llama-memory-hybrid-iswa.cpp @@ -77,15 +77,15 @@ llama_memory_context_ptr llama_memory_hybrid_iswa::init_batch(llama_batch_allocr // if all tokens are output, split by sequence ubatch = balloc.split_seq(n_ubatch); } else { - if (mem_recr->n_rs_seq > 0) { - // [TAG_RECURRENT_ROLLBACK_SPLITS] - // TODO: recurrent state rollback does not support equal splits - ubatch = balloc.split_seq(n_ubatch); - } else { - // Use non-sequential split when KV cache is unified (needed for hellaswag/winogrande/multiple-choice) - const bool unified = (mem_attn->get_base()->get_n_stream() == 1); - ubatch = balloc.split_equal(n_ubatch, !unified); - } + // Use non-sequential split when KV cache is unified (needed for hellaswag/winogrande/multiple-choice) + const bool unified = (mem_attn->get_base()->get_n_stream() == 1); + + // [TAG_RECURRENT_ROLLBACK_SPLITS] + // the trailing (1 + n_rs_seq) tokens of each seq must stay in the same ubatch + // so that the rollback snapshots remain valid + const uint32_t n_rs_seq = mem_recr->n_rs_seq; + + ubatch = balloc.split_equal(n_ubatch, !unified, n_rs_seq > 0 ? n_rs_seq + 1 : 0); } if (ubatch.n_tokens == 0) { diff --git a/src/llama-memory-hybrid.cpp b/src/llama-memory-hybrid.cpp index f2d49cbce..42c7381a9 100644 --- a/src/llama-memory-hybrid.cpp +++ b/src/llama-memory-hybrid.cpp @@ -78,15 +78,15 @@ llama_memory_context_ptr llama_memory_hybrid::init_batch(llama_batch_allocr & ba // if all tokens are output, split by sequence ubatch = balloc.split_seq(n_ubatch); } else { - if (mem_recr->n_rs_seq > 0) { - // [TAG_RECURRENT_ROLLBACK_SPLITS] - // TODO: recurrent state rollback does not support equal splits - ubatch = balloc.split_seq(n_ubatch); - } else { - // Use non-sequential split when KV cache is unified (needed for hellaswag/winogrande/multiple-choice) - const bool unified = (mem_attn->get_n_stream() == 1); - ubatch = balloc.split_equal(n_ubatch, !unified); - } + // Use non-sequential split when KV cache is unified (needed for hellaswag/winogrande/multiple-choice) + const bool unified = (mem_attn->get_n_stream() == 1); + + // [TAG_RECURRENT_ROLLBACK_SPLITS] + // the trailing (1 + n_rs_seq) tokens of each seq must stay in the same ubatch + // so that the rollback snapshots remain valid + const uint32_t n_rs_seq = mem_recr->n_rs_seq; + + ubatch = balloc.split_equal(n_ubatch, !unified, n_rs_seq > 0 ? n_rs_seq + 1 : 0); } if (ubatch.n_tokens == 0) { diff --git a/src/llama-memory-recurrent.cpp b/src/llama-memory-recurrent.cpp index 43d0cca48..0d05e64b1 100644 --- a/src/llama-memory-recurrent.cpp +++ b/src/llama-memory-recurrent.cpp @@ -416,15 +416,12 @@ llama_memory_context_ptr llama_memory_recurrent::init_batch(llama_batch_allocr & // if all tokens are output, split by sequence ubatch = balloc.split_seq(n_ubatch); } else { - if (n_rs_seq > 0) { - // [TAG_RECURRENT_ROLLBACK_SPLITS] - // TODO: recurrent state rollback does not support equal splits - ubatch = balloc.split_seq(n_ubatch); - } else { - // TODO: non-sequential equal split can be done if using unified KV cache - // for simplicity, we always use sequential equal split for now - ubatch = balloc.split_equal(n_ubatch, true); - } + // TODO: non-sequential equal split can be done if using unified KV cache + // for simplicity, we always use sequential equal split for now + // [TAG_RECURRENT_ROLLBACK_SPLITS] + // the trailing (1 + n_rs_seq) tokens of each seq must stay in the same ubatch + // so that the rollback snapshots remain valid + ubatch = balloc.split_equal(n_ubatch, true, n_rs_seq > 0 ? n_rs_seq + 1 : 0); } if (ubatch.n_tokens == 0) { diff --git a/src/llama-model-loader.cpp b/src/llama-model-loader.cpp index d05f8def0..6fb252ef9 100644 --- a/src/llama-model-loader.cpp +++ b/src/llama-model-loader.cpp @@ -37,6 +37,7 @@ const char * llama_ftype_name(llama_ftype ftype) { case LLAMA_FTYPE_MOSTLY_F16: name = LLAMA_FTYPE_PREFIX "F16"; break; case LLAMA_FTYPE_MOSTLY_BF16: name = LLAMA_FTYPE_PREFIX "BF16"; break; case LLAMA_FTYPE_MOSTLY_Q1_0: name = LLAMA_FTYPE_PREFIX "Q1_0"; break; + case LLAMA_FTYPE_MOSTLY_Q2_0: name = LLAMA_FTYPE_PREFIX "Q2_0"; break; case LLAMA_FTYPE_MOSTLY_Q4_0: name = LLAMA_FTYPE_PREFIX "Q4_0"; break; case LLAMA_FTYPE_MOSTLY_Q4_1: name = LLAMA_FTYPE_PREFIX "Q4_1"; break; case LLAMA_FTYPE_MOSTLY_Q5_0: name = LLAMA_FTYPE_PREFIX "Q5_0"; break; @@ -768,6 +769,7 @@ llama_model_loader::llama_model_loader( case GGML_TYPE_IQ3_S: ftype = LLAMA_FTYPE_MOSTLY_IQ3_S; break; case GGML_TYPE_NVFP4: ftype = LLAMA_FTYPE_MOSTLY_NVFP4; break; case GGML_TYPE_Q1_0: ftype = LLAMA_FTYPE_MOSTLY_Q1_0; break; + case GGML_TYPE_Q2_0: ftype = LLAMA_FTYPE_MOSTLY_Q2_0; break; default: { LLAMA_LOG_WARN("%s: unknown type %s\n", __func__, ggml_type_name(type_max)); diff --git a/src/llama-quant.cpp b/src/llama-quant.cpp index d61cf8ec3..d30f8d91a 100644 --- a/src/llama-quant.cpp +++ b/src/llama-quant.cpp @@ -380,6 +380,7 @@ static ggml_type tensor_type_fallback(quantize_state_impl & qs, const ggml_tenso case GGML_TYPE_IQ3_XXS: case GGML_TYPE_IQ3_S: // types on the right: block size 32 case GGML_TYPE_IQ4_XS: return_type = GGML_TYPE_IQ4_NL; break; + case GGML_TYPE_Q2_0: case GGML_TYPE_Q2_K: case GGML_TYPE_Q3_K: case GGML_TYPE_TQ1_0: @@ -482,7 +483,7 @@ static ggml_type llama_tensor_get_type_impl(quantize_state_impl & qs, ggml_type else if (ftype == LLAMA_FTYPE_MOSTLY_IQ3_XXS) { new_type = GGML_TYPE_IQ3_S; } - else if (ftype == LLAMA_FTYPE_MOSTLY_TQ1_0 || ftype == LLAMA_FTYPE_MOSTLY_TQ2_0) { + else if (ftype == LLAMA_FTYPE_MOSTLY_TQ1_0 || ftype == LLAMA_FTYPE_MOSTLY_TQ2_0 || ftype == LLAMA_FTYPE_MOSTLY_Q2_0) { new_type = GGML_TYPE_Q4_K; } } @@ -802,6 +803,7 @@ ggml_type llama_ftype_get_default_type(llama_ftype ftype) { case LLAMA_FTYPE_MOSTLY_BF16: return GGML_TYPE_BF16; case LLAMA_FTYPE_ALL_F32: return GGML_TYPE_F32; case LLAMA_FTYPE_MOSTLY_Q1_0: return GGML_TYPE_Q1_0; + case LLAMA_FTYPE_MOSTLY_Q2_0: return GGML_TYPE_Q2_0; case LLAMA_FTYPE_MOSTLY_MXFP4_MOE: return GGML_TYPE_MXFP4; diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index 156346414..5c04673ce 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -1112,9 +1112,6 @@ struct llm_tokenizer_ugm : llm_tokenizer { // blob containing XOR-compressed compact double array (XCDA) entries uint32_t xcda_blob_size = *(const uint32_t *) &precompiled_charsmap[0]; charsmap_offset += sizeof(xcda_blob_size); - if (xcda_blob_size + charsmap_offset >= precompiled_charsmap.size()) { - throw std::runtime_error("Index out of array bounds in precompiled charsmap!"); - } // Next xcda_blob_size bytes contain entries of XOR-compressed compact // double array (XCDA). Each entry is bit-packed into a 32-bit integer. @@ -1430,7 +1427,15 @@ private: throw std::runtime_error("Index out of array bounds in precompiled charsmap!"); } const char * prefix_replacement = &(tokenizer.prefix_replacements)[longest_prefix_offset]; - return { prefix_replacement, strlen(prefix_replacement), longest_prefix_length }; + size_t max_len = tokenizer.prefix_replacements_size - longest_prefix_offset; + size_t repl_len = 0; + while (repl_len < max_len && prefix_replacement[repl_len] != '\0') { + repl_len++; + } + if (repl_len == max_len) { + throw std::runtime_error("Unterminated string in precompiled charsmap!"); + } + return { prefix_replacement, repl_len, longest_prefix_length }; } // check if the input prefix contains a valid sequence of UTF-8 code units @@ -2254,11 +2259,18 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) { const size_t n_precompiled_charsmap = gguf_get_arr_n(ctx, precompiled_charsmap_keyidx); const char * pc = (const char *) gguf_get_arr_data(ctx, precompiled_charsmap_keyidx); precompiled_charsmap.assign(pc, pc + n_precompiled_charsmap); + if (precompiled_charsmap.size() < sizeof(uint32_t)) { + throw std::runtime_error("precompiled_charsmap too small for xcda_blob_size header!"); + } + uint32_t * xcda_blob_size = (uint32_t *) &precompiled_charsmap[0]; +#if defined(__BYTE_ORDER__) && defined(__ORDER_BIG_ENDIAN__) && __BYTE_ORDER__ == __ORDER_BIG_ENDIAN__ + *xcda_blob_size = __builtin_bswap32(*xcda_blob_size); +#endif + if (*xcda_blob_size + sizeof(uint32_t) >= precompiled_charsmap.size()) { + throw std::runtime_error("Index out of array bounds in precompiled charsmap!"); + } #if defined(__BYTE_ORDER__) && defined(__ORDER_BIG_ENDIAN__) && __BYTE_ORDER__ == __ORDER_BIG_ENDIAN__ // correct endianness of data in precompiled_charsmap binary blob - uint32_t * xcda_blob_size = (uint32_t *) &precompiled_charsmap[0]; - *xcda_blob_size = __builtin_bswap32(*xcda_blob_size); - assert(*xcda_blob_size + sizeof(uint32_t) < n_precompiled_charsmap); size_t xcda_array_size = *xcda_blob_size / sizeof(uint32_t); uint32_t * xcda_array = (uint32_t *) &precompiled_charsmap[sizeof(uint32_t)]; for (size_t i = 0; i < xcda_array_size; ++i) { diff --git a/src/models/delta-net-base.cpp b/src/models/delta-net-base.cpp index ad9ce7714..cf5e38095 100644 --- a/src/models/delta-net-base.cpp +++ b/src/models/delta-net-base.cpp @@ -496,8 +496,8 @@ ggml_tensor * llm_build_delta_net_base::build_conv_state( ggml_build_forward_expand(gf, ggml_cpy(ctx0, conv_state_last, conv_state_update)); } else { // [TAG_RECURRENT_ROLLBACK_SPLITS] - // TODO: this logic incorrectly assumes that the last (n_rs_seq + 1) tokens of a sequence in a batch are - // inside the same ubatch. currently with `split_equal()` this is not correct + // this logic assumes that the last (n_rs_seq + 1) tokens of a sequence in a batch are inside + // the same ubatch, which `split_equal()` guarantees via its n_keep_tail argument const int64_t K = (int64_t) cparams.n_rs_seq + 1; diff --git a/tools/quantize/quantize.cpp b/tools/quantize/quantize.cpp index 6c444e08b..7dce2978e 100644 --- a/tools/quantize/quantize.cpp +++ b/tools/quantize/quantize.cpp @@ -32,6 +32,7 @@ struct quant_option { static const std::vector QUANT_OPTIONS = { { "Q1_0", LLAMA_FTYPE_MOSTLY_Q1_0, " 1.125 bpw quantization", }, + { "Q2_0", LLAMA_FTYPE_MOSTLY_Q2_0, " 2.25 bpw quantization (group 64)", }, { "Q4_0", LLAMA_FTYPE_MOSTLY_Q4_0, " 4.34G, +0.4685 ppl @ Llama-3-8B", }, { "Q4_1", LLAMA_FTYPE_MOSTLY_Q4_1, " 4.78G, +0.4511 ppl @ Llama-3-8B", }, { "MXFP4_MOE",LLAMA_FTYPE_MOSTLY_MXFP4_MOE," MXFP4 MoE", }, diff --git a/tools/server/README-dev.md b/tools/server/README-dev.md index dfc9004de..882adca09 100644 --- a/tools/server/README-dev.md +++ b/tools/server/README-dev.md @@ -57,7 +57,7 @@ The core architecture consists of the following components: - `server_tokens`: Unified representation of token sequences (supports both text and multimodal tokens); used by `server_task` and `server_slot`. - `server_prompt_checkpoint`: For recurrent (e.g., RWKV) and SWA models, stores snapshots of KV cache state. Enables reuse when subsequent requests share the same prompt prefix, saving redundant computation. - `server_models`: Standalone component for managing multiple backend instances (used in router mode). It is completely independent of `server_context`. -- `stream_session_manager`: Process wide owner of resumable SSE stream sessions (`g_stream_sessions`), keyed by conversation id. Backs the replay buffer that lets a client reattach to a generation after an HTTP disconnect. See the "Resumable streaming" section below. +- `stream_session_manager`: process wide owner of resumable SSE stream sessions, keyed by conversation id. A file-static singleton inside `server-stream.cpp`, driven through `server_stream_session_manager_start/stop`. Backs the replay buffer that lets a client reattach to a generation after an HTTP disconnect. See the "Resumable streaming" section below. ```mermaid graph TD @@ -127,10 +127,12 @@ It is opt in via the `X-Conversation-Id` header on `POST /v1/chat/completions`. The feature lives entirely in `server-stream.{h,cpp}` and rests on three types: - `stream_session`: a bounded ring buffer (4 MiB cap, oldest bytes drop first) plus a condvar. `append` pushes raw SSE bytes, `read_from` drains from any offset and blocks for live bytes or finalize, `finalize` wakes readers, `cancel` stops the producer. One conv maps to at most one live session. -- `stream_session_manager` (`g_stream_sessions`): owns all sessions keyed by conv id, enforces the one conv one session invariant via `create_or_replace`, and runs a GC thread that drops completed sessions past their TTL. +- `stream_session_manager`: a file-static singleton (`g_stream_sessions`) inside `server-stream.cpp`, owns all sessions keyed by conv id, enforces the one conv one session invariant via `create_or_replace`, and runs a GC thread that drops completed sessions past their TTL. Exposed to main only through `server_stream_session_manager_start/stop`. - `stream_pipe_producer` / `stream_pipe_consumer`: the write and read ends. The producer owns the session lifetime and finalizes it on destruction; the consumer is read only and never finalizes, so a reader detaching cannot kill a running generation. -Producer side: `server_res_generator` attaches a producer pipe when the header is present. The HTTP content provider mirrors every chunk into the ring before writing it to the socket. While a pipe is attached, `stream_aware_should_stop` ignores peer disconnect, so a dropped socket does not stop generation: only an explicit `DELETE` does. When the peer leaves early, `on_complete` calls `close()`, which drains the rest of the generation into the ring on the http worker. +The implementation is hidden in `server-stream.cpp` (pimpl). The header exposes only the route handler factories, `server_stream_session_attach_pipe`, `server_stream_aware_should_stop`, `server_stream_conv_id_from_headers` and the GC lifecycle; the session, manager and consumer types stay in the `.cpp`. + +Producer side: `server_res_generator` attaches a producer pipe when the header is present. The HTTP content provider mirrors every chunk into the ring before writing it to the socket. While a pipe is attached, `server_stream_aware_should_stop` ignores peer disconnect, so a dropped socket does not stop generation: only an explicit `DELETE` does. When the peer leaves early, `on_complete` calls `close()`, which drains the rest of the generation into the ring on the http worker. Lifetime safety: the producer pipe holds a shared `alive` flag also captured by the session cancel hook. `~server_res_generator` calls `cleanup()` to clear that hook while the reader is still alive, so a `cancel` arriving during teardown can never call `stop()` on a freed response. This ordering is the most fragile part of the feature: finalizing or destroying the producer before `cleanup()` runs reintroduces a use after free. @@ -144,7 +146,7 @@ Routes: Router mode binds the same paths to proxy handlers. A `conv_id -> child` map (`conv_models`), populated when a POST is routed, resolves the owning child in one lookup with no polling. The lookup groups ids per child; GET and DELETE proxy straight to the owner. This loopback REST hop is expected to move to a websocket IPC later, swapping only the transport. -Lifecycle: `g_stream_sessions.start_gc()` runs in main after common init, `stop_gc()` runs first in `clean_up()` and finalizes every live session so no reader hangs. Reader blocking and the post drop drain both run on httplib worker threads, which block on a condvar rather than spin. +Lifecycle: `server_stream_session_manager_start()` runs in main after common init, `server_stream_session_manager_stop()` runs first in `clean_up()` and finalizes every live session so no reader hangs. Reader blocking and the post drop drain both run on httplib worker threads, which block on a condvar rather than spin. | Constant | Value | Role | | --- | --- | --- | diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index bb3b91ab5..aa5d0a2ab 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -897,8 +897,10 @@ private: server_batch batch; - llama_model_ptr model_dft; - llama_context_ptr ctx_dft; + llama_model * model_dft = nullptr; + llama_context * ctx_dft = nullptr; + + common_speculative_init_result_ptr spec_init; common_context_seq_rm_type ctx_tgt_seq_rm_type = COMMON_CONTEXT_SEQ_RM_TYPE_NO; common_context_seq_rm_type ctx_dft_seq_rm_type = COMMON_CONTEXT_SEQ_RM_TYPE_NO; @@ -939,8 +941,10 @@ private: void destroy() { spec.reset(); - ctx_dft.reset(); - model_dft.reset(); + spec_init.reset(); + + ctx_dft = nullptr; + model_dft = nullptr; llama_init.reset(); @@ -1084,30 +1088,15 @@ private: // optionally reserve VRAM for the draft / MTP context before fitting the target model if (params_base.fit_params) { if (has_spec) { - common_params params_dft = params_base; - bool measure_model_bytes = true; + // MTP draft context lives on the target model, only context+compute are new + bool measure_model_bytes = has_draft; - if (has_draft) { - const auto & params_spec = params_base.speculative.draft; - params_dft.devices = params_spec.devices; - params_dft.model = params_spec.mparams; - params_dft.n_gpu_layers = params_spec.n_gpu_layers; - params_dft.cache_type_k = params_spec.cache_type_k; - params_dft.cache_type_v = params_spec.cache_type_v; - params_dft.tensor_buft_overrides = params_spec.tensor_buft_overrides; - } else { - // MTP draft context lives on the target model, only context+compute are new - measure_model_bytes = false; - } - - params_dft.n_outputs_max = params_base.n_parallel; + common_params params_dft = common_base_params_to_speculative(params_base); auto mparams_dft = common_model_params_to_llama(params_dft); auto cparams_dft = common_context_params_to_llama(params_dft); if (spec_mtp) { cparams_dft.ctx_type = LLAMA_CONTEXT_TYPE_MTP; - cparams_dft.type_k = params_base.speculative.draft.cache_type_k; - cparams_dft.type_v = params_base.speculative.draft.cache_type_v; } cparams_dft.n_rs_seq = 0; @@ -1175,82 +1164,36 @@ private: add_bos_token = llama_vocab_get_add_bos(vocab); - if (has_draft) { - // TODO speculative: move to common/speculative.cpp? - const auto & params_spec = params_base.speculative.draft; - - SRV_TRC("loading draft model '%s'\n", params_spec.mparams.path.c_str()); - - auto params_dft = params_base; - - params_dft.devices = params_spec.devices; - params_dft.model = params_spec.mparams; - params_dft.n_gpu_layers = params_spec.n_gpu_layers; - params_dft.cache_type_k = params_spec.cache_type_k; - params_dft.cache_type_v = params_spec.cache_type_v; - - if (params_spec.cpuparams.n_threads > 0) { - params_dft.cpuparams.n_threads = params_spec.cpuparams.n_threads; - params_dft.cpuparams_batch.n_threads = params_spec.cpuparams_batch.n_threads; - } - - params_dft.tensor_buft_overrides = params_spec.tensor_buft_overrides; - - auto mparams_dft = common_model_params_to_llama(params_dft); - - // progress callback - mparams_dft.progress_callback = load_progress_callback; - mparams_dft.progress_callback_user_data = &load_progress_spec; - - model_dft.reset(llama_model_load_from_file(params_dft.model.path.c_str(), mparams_dft)); - if (model_dft == nullptr) { - SRV_ERR("failed to load draft model, '%s'\n", params_dft.model.path.c_str()); - return false; - } - - auto cparams = common_context_params_to_llama(params_dft); - - if (spec_mtp) { - cparams.ctx_type = LLAMA_CONTEXT_TYPE_MTP; - } - - // note: for small models maybe we can set this to the maximum possible draft from all speculative types - // the extra memory for small models is likely negligible? - cparams.n_rs_seq = 0; - cparams.ctx_other = ctx_tgt; - - ctx_dft.reset(llama_init_from_model(model_dft.get(), cparams)); - if (ctx_dft == nullptr) { - SRV_ERR("%s", "failed to create draft context\n"); - return false; - } - - params_base.speculative.draft.ctx_tgt = ctx_tgt; - params_base.speculative.draft.ctx_dft = ctx_dft.get(); - } else if (spec_mtp) { - // no new model load, so we simply report 0.0 and 1.0 progress + if (has_spec) { + // spec_mtp doesn't use load a model internally, so we report 0.0 and 1.0 manually load_progress_callback(0.0f, &load_progress_spec); + load_progress_spec.t_last_load_progress_ms = 0; // reset so internal cbs aren't delayed - SRV_TRC("creating MTP draft context against the target model '%s'\n", - params_base.model.path.c_str()); + { + common_params params_dft = common_base_params_to_speculative(params_base); - auto cparams_mtp = common_context_params_to_llama(params_base); - cparams_mtp.ctx_type = LLAMA_CONTEXT_TYPE_MTP; - cparams_mtp.type_k = params_base.speculative.draft.cache_type_k; - cparams_mtp.type_v = params_base.speculative.draft.cache_type_v; - cparams_mtp.n_rs_seq = 0; - cparams_mtp.n_outputs_max = params_base.n_parallel; - cparams_mtp.ctx_other = ctx_tgt; + // progress callback + params_dft.load_progress_callback = load_progress_callback; + params_dft.load_progress_callback_user_data = &load_progress_spec; - ctx_dft.reset(llama_init_from_model(model_tgt, cparams_mtp)); - if (ctx_dft == nullptr) { - SRV_ERR("%s", "failed to create MTP context\n"); - return false; + spec_init = common_speculative_init_from_params(params_dft, model_tgt, ctx_tgt); + model_dft = spec_init->model(); + ctx_dft = spec_init->context(); + + if (has_draft && model_dft == nullptr) { + SRV_ERR("failed to load draft model, '%s'\n", params_dft.model.path.c_str()); + return false; + } + + if (ctx_dft == nullptr) { + SRV_ERR("%s", "failed to create MTP context\n"); + return false; + } + + params_base.speculative.draft.ctx_tgt = ctx_tgt; + params_base.speculative.draft.ctx_dft = ctx_dft; } - params_base.speculative.draft.ctx_tgt = ctx_tgt; - params_base.speculative.draft.ctx_dft = ctx_dft.get(); - load_progress_callback(1.0f, &load_progress_spec); } @@ -1343,13 +1286,15 @@ private: } if (ctx_dft) { - ctx_dft_seq_rm_type = common_context_can_seq_rm(ctx_dft.get()); + ctx_dft_seq_rm_type = common_context_can_seq_rm(ctx_dft); } if (spec) { SRV_TRC("%s", "speculative decoding context initialized\n"); } else { - ctx_dft.reset(); + spec_init.reset(); + ctx_dft = nullptr; + model_dft = nullptr; } for (int i = 0; i < params_base.n_parallel; i++) { @@ -1357,7 +1302,7 @@ private: slot.id = i; slot.ctx_tgt = ctx_tgt; - slot.ctx_dft = ctx_dft.get(); + slot.ctx_dft = ctx_dft; slot.spec = spec.get(); slot.n_ctx = n_ctx_slot; @@ -2362,8 +2307,8 @@ private: // this is not true for SWA models: https://github.com/ggml-org/llama.cpp/pull/24411#issuecomment-4677983225 cur.update_pos(slot.prompt.n_tokens() - n_tokens_cur, pos_min, pos_max); - cur.update_tgt(ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); - cur.update_dft(ctx_dft.get(), slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + cur.update_tgt(ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + cur.update_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); // stash the draft's speculative state with the checkpoint common_speculative_get_state(spec.get(), slot.id, cur.data_spec); @@ -2899,8 +2844,8 @@ private: common_context_seq_add(ctx_tgt, slot.id, n_keep + n_discard, slot.prompt.n_tokens(), -n_discard); if (ctx_dft) { - common_context_seq_rm (ctx_dft.get(), slot.id, n_keep , n_keep + n_discard); - common_context_seq_add(ctx_dft.get(), slot.id, n_keep + n_discard, slot.prompt.tokens.pos_next(), -n_discard); + common_context_seq_rm (ctx_dft, slot.id, n_keep , n_keep + n_discard); + common_context_seq_add(ctx_dft, slot.id, n_keep + n_discard, slot.prompt.tokens.pos_next(), -n_discard); } // add generated tokens to cache @@ -2972,7 +2917,7 @@ private: llama_memory_seq_pos_max(llama_get_memory(ctx_tgt), slot.id)); if (use_ckpt_dft) { - slot.spec_ckpt.update_dft(ctx_dft.get(), slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + slot.spec_ckpt.update_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); } slot.spec_prompt = slot.prompt.tokens.get_text_tokens(); @@ -3009,10 +2954,10 @@ private: if (ctx_dft) { if (use_ckpt_dft) { - ckpt.load_dft(ctx_dft.get(), slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + ckpt.load_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); } - common_context_seq_rm(ctx_dft.get(), slot.id, ckpt.pos_max + 1, -1); + common_context_seq_rm(ctx_dft, slot.id, ckpt.pos_max + 1, -1); } if (!draft.empty()) { @@ -3021,7 +2966,7 @@ private: (ctx_tgt_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_RS && draft.size() > llama_n_rs_seq(ctx_tgt)); const bool use_ckpt_dft = - (ctx_dft_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_RS && draft.size() > llama_n_rs_seq(ctx_dft.get())); + (ctx_dft_seq_rm_type == COMMON_CONTEXT_SEQ_RM_TYPE_RS && draft.size() > llama_n_rs_seq(ctx_dft)); if (use_ckpt_tgt) { //const int64_t t_start = ggml_time_us(); @@ -3038,7 +2983,7 @@ private: } if (use_ckpt_dft) { - ckpt.update_dft(ctx_dft.get(), slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + ckpt.update_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); } } }); @@ -3219,8 +3164,8 @@ private: common_context_seq_add(ctx_tgt, slot.id, head_c, head_c + n_match, kv_shift); if (ctx_dft) { - common_context_seq_rm (ctx_dft.get(), slot.id, head_p, head_c); - common_context_seq_add(ctx_dft.get(), slot.id, head_c, head_c + n_match, kv_shift); + common_context_seq_rm (ctx_dft, slot.id, head_p, head_c); + common_context_seq_add(ctx_dft, slot.id, head_c, head_c + n_match, kv_shift); } for (size_t i = 0; i < n_match; i++) { @@ -3320,8 +3265,8 @@ private: if (!do_reset) { // restore the context checkpoint - it->load_tgt(ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); - it->load_dft(ctx_dft.get(), slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + it->load_tgt(ctx_tgt, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); + it->load_dft(ctx_dft, slot.id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY); // restore the draft's speculative state common_speculative_set_state(spec.get(), slot.id, it->data_spec); @@ -3395,7 +3340,7 @@ private: common_context_seq_rm(ctx_tgt, slot.id, p0, -1); if (ctx_dft) { - common_context_seq_rm(ctx_dft.get(), slot.id, p0, -1); + common_context_seq_rm(ctx_dft, slot.id, p0, -1); } // If using an alora, there may be uncached tokens that come @@ -4243,7 +4188,7 @@ std::unique_ptr server_routes::handle_completions_impl( } }; - auto effective_should_stop = stream_aware_should_stop(res_this, req.should_stop); + auto effective_should_stop = server_stream_aware_should_stop(res_this, req.should_stop); try { if (effective_should_stop()) { @@ -4339,7 +4284,7 @@ std::unique_ptr server_routes::handle_completions_impl( // attach a producer pipe to the response when X-Conversation-Id is present. // the pipe mirrors SSE chunks into the ring buffer and wires up the cancel hook. - stream_session_attach_pipe(*res, req.headers); + server_stream_session_attach_pipe(*res, req.headers); return res; } diff --git a/tools/server/server-models.cpp b/tools/server/server-models.cpp index 6a8eb2a2b..0cbf520af 100644 --- a/tools/server/server-models.cpp +++ b/tools/server/server-models.cpp @@ -1681,7 +1681,7 @@ void server_models_routes::init_routes() { } // remember which child serves this conversation so the stream routes can route straight // to it without polling, keyed on the exact conv id from the header - std::string conv_id = stream_conv_id_from_headers(req.headers); + std::string conv_id = server_stream_conv_id_from_headers(req.headers); if (!conv_id.empty()) { models.conv_models.remember(conv_id, name); } @@ -1896,7 +1896,7 @@ void server_models_routes::init_routes() { if (!from.empty()) { child_path += "?from=" + from; } - SRV_INF("proxying stream resume to model %s on port %d, path=%s\n", + SRV_TRC("proxying stream resume to model %s on port %d, path=%s\n", owner->name.c_str(), owner->port, child_path.c_str()); auto proxy = std::make_unique( "GET", diff --git a/tools/server/server-stream.cpp b/tools/server/server-stream.cpp index c2bba8ec4..553ac26b1 100644 --- a/tools/server/server-stream.cpp +++ b/tools/server/server-stream.cpp @@ -6,6 +6,12 @@ #include #include #include +#include + +enum class stream_read_status { + OK, + OFFSET_LOST, +}; namespace { constexpr int64_t STREAM_SESSION_TTL_SECONDS = 300; @@ -13,7 +19,6 @@ constexpr size_t STREAM_SESSION_MAX_BYTES = 4 * 1024 * 1024; constexpr int64_t STREAM_SESSION_GC_INTERVAL_SECONDS = 60; constexpr int64_t STREAM_READ_WAKE_INTERVAL_MS = 200; -// returns unix time in seconds int64_t now_seconds() { return std::chrono::duration_cast( std::chrono::system_clock::now().time_since_epoch() @@ -21,6 +26,91 @@ int64_t now_seconds() { } } +// owns all live sessions keyed by conversation_id, one conv = at most one live session. +// a periodic GC evicts expired ones +class stream_session_manager { +public: + stream_session_manager(); + ~stream_session_manager(); + + stream_session_manager(const stream_session_manager &) = delete; + stream_session_manager & operator=(const stream_session_manager &) = delete; + + // install a new session, evicting and cancelling any previous one. conversation_id must be non empty + stream_session_ptr create_or_replace(const std::string & conversation_id); + + stream_session_ptr get(const std::string & conversation_id); + + std::vector list_all() const; + + void evict(const std::string & conversation_id); + + void evict_and_cancel(const std::string & conversation_id); + + void start_gc(); + void stop_gc(); + +private: + void gc_loop(); + + mutable std::shared_mutex map_mu; + std::unordered_map sessions; // key: conversation_id + std::thread gc_thread; + bool running; + std::mutex gc_wake_mu; + std::condition_variable gc_wake_cv; +}; + +// process wide manager, lifecycle controlled by llama-server main() via start_gc/stop_gc +static stream_session_manager g_stream_sessions; + +void server_stream_session_manager_start() { + g_stream_sessions.start_gc(); +} + +void server_stream_session_manager_stop() { + g_stream_sessions.stop_gc(); +} + +struct stream_session { + std::string conversation_id; + int64_t started_ts; // unix seconds at construction + + stream_session(std::string conversation_id_, size_t max_bytes_); + stream_session(const stream_session &) = delete; + stream_session & operator=(const stream_session &) = delete; + + bool append(const char * data, size_t len); + + void finalize(); + + // drain from offset into sink, blocking for more bytes or finalize. OFFSET_LOST if offset + // fell below the dropped prefix + stream_read_status read_from(size_t offset, + const std::function & sink, + const std::function & should_stop); + + bool is_done() const; + bool is_cancelled() const; + size_t total_size() const; // bytes that ever entered the session + size_t dropped_prefix() const; // bytes evicted from the front due to cap + int64_t completed_at() const; // 0 while alive, unix seconds after finalize + + void set_stop_producer(std::function fn); + + void cancel(); + +private: + mutable std::mutex mu; + std::condition_variable cv; + std::vector buffer; + size_t prefix_dropped; + size_t cap_bytes; + bool done; + std::atomic cancelled; // polled lock-free by the should_stop closure, no mu + int64_t completed_ts; + std::function stop_producer; +}; stream_session::stream_session(std::string conversation_id_, size_t max_bytes_) : conversation_id(std::move(conversation_id_)) , started_ts(now_seconds()) @@ -38,7 +128,7 @@ bool stream_session::append(const char * data, size_t len) { } { std::lock_guard lock(mu); - if (done.load(std::memory_order_relaxed)) { + if (done) { return false; } if (len >= cap_bytes) { @@ -62,11 +152,14 @@ bool stream_session::append(const char * data, size_t len) { } void stream_session::finalize() { - bool was_done = done.exchange(true, std::memory_order_acq_rel); - if (was_done) { - return; + { + std::lock_guard lock(mu); + if (done) { + return; + } + done = true; + completed_ts = now_seconds(); } - completed_ts.store(now_seconds(), std::memory_order_release); cv.notify_all(); } @@ -96,7 +189,7 @@ stream_read_status stream_session::read_from(size_t offset, lock.lock(); continue; } - if (done.load(std::memory_order_acquire)) { + if (done) { return stream_read_status::OK; } // wait for new bytes, finalize, or a periodic wake to re check should_stop @@ -105,7 +198,8 @@ stream_read_status stream_session::read_from(size_t offset, } bool stream_session::is_done() const { - return done.load(std::memory_order_acquire); + std::lock_guard lock(mu); + return done; } size_t stream_session::total_size() const { @@ -119,7 +213,8 @@ size_t stream_session::dropped_prefix() const { } int64_t stream_session::completed_at() const { - return completed_ts.load(std::memory_order_acquire); + std::lock_guard lock(mu); + return completed_ts; } void stream_session::set_stop_producer(std::function fn) { @@ -128,7 +223,7 @@ void stream_session::set_stop_producer(std::function fn) { } void stream_session::cancel() { - // flip cancelled first so the producer-side stream_aware_should_stop can break out of the + // flip cancelled first so the producer-side server_stream_aware_should_stop can break out of the // recv() wait even if remove_waiting_task_ids does not notify the condvar (the cancel task // posted by rd.stop() will eventually notify, but we do not want to depend on that timing) cancelled.store(true, std::memory_order_release); @@ -237,18 +332,24 @@ void stream_session_manager::evict_and_cancel(const std::string & conversation_i } void stream_session_manager::start_gc() { - if (running.exchange(true)) { - return; + { + std::lock_guard lock(gc_wake_mu); + if (running) { + return; + } + running = true; } gc_thread = std::thread([this] { gc_loop(); }); } void stream_session_manager::stop_gc() { - bool was_running = running.exchange(false); + bool was_running; + { + std::lock_guard lock(gc_wake_mu); + was_running = running; + running = false; + } if (was_running) { - { - std::lock_guard lock(gc_wake_mu); - } gc_wake_cv.notify_all(); if (gc_thread.joinable()) { gc_thread.join(); @@ -270,15 +371,15 @@ void stream_session_manager::stop_gc() { } void stream_session_manager::gc_loop() { - while (running.load(std::memory_order_acquire)) { + while (true) { { std::unique_lock lock(gc_wake_mu); gc_wake_cv.wait_for(lock, std::chrono::seconds(STREAM_SESSION_GC_INTERVAL_SECONDS), - [this] { return !running.load(std::memory_order_acquire); }); - } - if (!running.load(std::memory_order_acquire)) { - return; + [this] { return !running; }); + if (!running) { + return; + } } int64_t cutoff = now_seconds() - STREAM_SESSION_TTL_SECONDS; std::vector to_drop; @@ -301,10 +402,19 @@ void stream_session_manager::gc_loop() { } } -// process wide manager, lifecycle controlled by llama-server main() via start_gc/stop_gc -stream_session_manager g_stream_sessions; +// stream_pipe -// stream_pipe --------------------------------------------------------------------------------- +// consumer end: read-only replay of the ring buffer, the destructor does not finalize the session +struct stream_pipe_consumer : stream_pipe { + stream_read_status read(size_t & offset, + const std::function & sink, + const std::function & should_stop); + + static std::shared_ptr create(stream_session_ptr session); + +private: + explicit stream_pipe_consumer(stream_session_ptr session); +}; stream_pipe::stream_pipe(stream_session_ptr session) : session_(std::move(session)) { @@ -408,12 +518,10 @@ static server_http_res_ptr make_error_response(int status, const std::string & m return res; } -server_http_context::handler_t make_stream_get_handler() { +server_http_context::handler_t server_stream_make_get_handler() { return [](const server_http_req & req) -> server_http_res_ptr { - // GET /v1/stream/?from=N replays the SSE bytes already buffered for the - // session, blocks for more bytes when the session is still running, returns when - // the session is finalized. the body is streamed back as text/event-stream so the - // browser EventSource can attach to it like a fresh request + // GET /v1/stream/?from=N replays buffered SSE bytes then blocks for live + // bytes until the session finalizes, streamed as text/event-stream for EventSource std::string conv_id = req.get_param("conv_id"); if (conv_id.empty()) { return make_error_response(400, "Missing conversation id in path", ERROR_TYPE_INVALID_REQUEST); @@ -459,11 +567,10 @@ server_http_context::handler_t make_stream_get_handler() { }; } -server_http_context::handler_t make_streams_lookup_handler() { +server_http_context::handler_t server_stream_make_lookup_handler() { return [](const server_http_req & req) -> server_http_res_ptr { - // POST /v1/streams/lookup with body {"conversation_ids": ["X", "Y", ...]} returns the - // matching sessions, only for ids the caller already knows. each id matches the exact key - // and any "::" variant, so one lookup covers every per model session for a conv + // POST /v1/streams/lookup returns the matching sessions, only for ids the caller already + // knows. each id matches the exact key and any "::" per model variant std::vector requested; try { json body = json::parse(req.body); @@ -518,11 +625,10 @@ server_http_context::handler_t make_streams_lookup_handler() { }; } -server_http_context::handler_t make_stream_delete_handler() { +server_http_context::handler_t server_stream_make_delete_handler() { return [](const server_http_req & req) -> server_http_res_ptr { - // DELETE /v1/stream/ is the explicit user Stop, cancels the producer hook - // wired by handle_completions_impl and evicts the buffer. idempotent, a session that - // already finalized or was never created returns 204 either way + // DELETE /v1/stream/ is the explicit user Stop, cancels the producer and evicts + // the buffer. idempotent, returns 204 even if the session was already gone std::string conv_id = req.get_param("conv_id"); if (conv_id.empty()) { return make_error_response(400, "Missing conversation id in path", ERROR_TYPE_INVALID_REQUEST); @@ -536,7 +642,7 @@ server_http_context::handler_t make_stream_delete_handler() { }; } -std::string stream_conv_id_from_headers(const std::map & headers) { +std::string server_stream_conv_id_from_headers(const std::map & headers) { // case-insensitive scan for x-conversation-id static constexpr char target[] = "x-conversation-id"; static constexpr size_t target_len = sizeof(target) - 1; @@ -555,8 +661,8 @@ std::string stream_conv_id_from_headers(const std::map return std::string(); } -void stream_session_attach_pipe(server_http_res & res, const std::map & headers) { - std::string conversation_id = stream_conv_id_from_headers(headers); +void server_stream_session_attach_pipe(server_http_res & res, const std::map & headers) { + std::string conversation_id = server_stream_conv_id_from_headers(headers); SRV_TRC("conv_id=%s (empty=%d)\n", conversation_id.c_str(), conversation_id.empty() ? 1 : 0); if (conversation_id.empty()) { return; @@ -565,7 +671,7 @@ void stream_session_attach_pipe(server_http_res & res, const std::map stream_aware_should_stop(server_http_res * res, std::function fallback) { +std::function server_stream_aware_should_stop(server_http_res * res, std::function fallback) { return [res, fallback = std::move(fallback)]() -> bool { if (res->spipe) { return res->spipe->is_cancelled(); diff --git a/tools/server/server-stream.h b/tools/server/server-stream.h index ff363bb4c..c0c3e924f 100644 --- a/tools/server/server-stream.h +++ b/tools/server/server-stream.h @@ -3,81 +3,23 @@ #include "server-http.h" #include -#include #include -#include #include #include -#include -#include #include -#include -#include -#include -enum class stream_read_status { - OK, - OFFSET_LOST, -}; +// streaming buffer for one generation, survives HTTP disconnect. the producer appends SSE bytes, +// readers drain from any offset via read_from. keyed by conversation_id, one conv = one live session -// streaming buffer for one generation, survives HTTP disconnect. the producer appends raw SSE -// bytes, readers drain from any offset via read_from and block until more bytes or finalize. -// keyed by conversation_id: one conv = at most one live session -struct stream_session { - std::string conversation_id; - int64_t started_ts; // unix seconds at construction, used by /v1/streams listing - - stream_session(std::string conversation_id_, size_t max_bytes_); - stream_session(const stream_session &) = delete; - stream_session & operator=(const stream_session &) = delete; - - // append raw bytes, drops from the front if the cap is reached. - // returns false if the session is already finalized - bool append(const char * data, size_t len); - - // mark the session as complete, wakes all pending readers - void finalize(); - - // drain bytes from offset, calling sink for each chunk. blocks until more - // bytes arrive or finalize is called. returns OK on clean exit, OFFSET_LOST - // if offset falls below the dropped prefix - stream_read_status read_from(size_t offset, - const std::function & sink, - const std::function & should_stop); - - bool is_done() const; - bool is_cancelled() const; - size_t total_size() const; // bytes that ever entered the session - size_t dropped_prefix() const; // bytes evicted from the front due to cap - int64_t completed_at() const; // 0 while alive, unix seconds after finalize - - // attach the producer stop hook used to cancel its reader, pass an empty function to detach - void set_stop_producer(std::function fn); - - // signal the producer to abort its inference asap via the stop hook, idempotent - void cancel(); - -private: - mutable std::mutex mu; - std::condition_variable cv; - std::vector buffer; - size_t prefix_dropped; - size_t cap_bytes; - std::atomic done; - std::atomic cancelled; - std::atomic completed_ts; - std::function stop_producer; // protected by mu -}; +struct stream_session; using stream_session_ptr = std::shared_ptr; -// one end of a stream_session pipe. the base holds the session and the shared query, the -// producer and consumer ends derive from it. virtual dtor so each end runs its own teardown: +// base of the producer/consumer pipe ends. virtual dtor so each runs its own teardown: // the producer finalizes the session, the consumer leaves it untouched struct stream_pipe { virtual ~stream_pipe() = default; - // true if the session was cancelled (e.g. via DELETE /v1/stream/) bool is_cancelled() const; protected: @@ -95,7 +37,6 @@ protected: struct stream_pipe_producer : stream_pipe { ~stream_pipe_producer() override; - // append raw bytes to the session's ring buffer, returns false if already finalized bool write(const char * data, size_t len); // mark the natural end on the wire so a later close() is a no-op @@ -121,83 +62,21 @@ private: server_http_res * res_ = nullptr; }; -// consumer end: read-only replay of the ring buffer, the destructor does not finalize the session -struct stream_pipe_consumer : stream_pipe { - // drain bytes from offset, calling sink for each available chunk. blocks until more data - // arrives or the session finalizes. should_stop is polled, returns OFFSET_LOST if offset - // fell below the dropped prefix - stream_read_status read(size_t & offset, - const std::function & sink, - const std::function & should_stop); +void server_stream_session_manager_start(); +void server_stream_session_manager_stop(); - static std::shared_ptr create(stream_session_ptr session); +// route handler factories wired under /v1/stream/* by server.cpp +server_http_context::handler_t server_stream_make_get_handler(); +server_http_context::handler_t server_stream_make_lookup_handler(); +server_http_context::handler_t server_stream_make_delete_handler(); -private: - explicit stream_pipe_consumer(stream_session_ptr session); -}; +// extract the X-Conversation-Id header value (case-insensitive), empty when absent +std::string server_stream_conv_id_from_headers(const std::map & headers); -// owns all live sessions, runs a periodic GC to evict expired ones. -// the map is keyed by conversation_id, so the invariant "one conv = at most one -// live session" is enforced at the type level -class stream_session_manager { -public: - stream_session_manager(); - ~stream_session_manager(); - - stream_session_manager(const stream_session_manager &) = delete; - stream_session_manager & operator=(const stream_session_manager &) = delete; - - // install a new session for this conversation, evicting and cancelling any previous one. - // the conversation_id must be non empty, the caller is responsible for that check. - // returns the new session - stream_session_ptr create_or_replace(const std::string & conversation_id); - - // lookup, returns null if unknown or already evicted - stream_session_ptr get(const std::string & conversation_id); - - // list every live or recently completed session, used by GET /v1/streams without filter - std::vector list_all() const; - - // remove from the map and finalize, wakes any pending readers - void evict(const std::string & conversation_id); - - // signal the producer to cancel asap then evict, used by the explicit user Stop path - void evict_and_cancel(const std::string & conversation_id); - - void start_gc(); - void stop_gc(); - -private: - void gc_loop(); - - mutable std::shared_mutex map_mu; - std::unordered_map sessions; // key: conversation_id - std::thread gc_thread; - std::atomic running; - std::mutex gc_wake_mu; - std::condition_variable gc_wake_cv; -}; - -// process wide manager, linked by both llama-server and llama-cli. llama-server main() drives -// start_gc/stop_gc, llama-cli leaves it idle. the dtor calls stop_gc() unconditionally so exit -// is safe whether or not the GC thread ran -extern stream_session_manager g_stream_sessions; - -// route handler factories operating on g_stream_sessions, wired under /v1/stream/* by server.cpp. -// keeps the resumable stream surface confined to server-stream -server_http_context::handler_t make_stream_get_handler(); -server_http_context::handler_t make_streams_lookup_handler(); -server_http_context::handler_t make_stream_delete_handler(); - -// extract the X-Conversation-Id header value (case-insensitive), empty when absent. exposed so -// the router can track which child serves a forwarded POST -std::string stream_conv_id_from_headers(const std::map & headers); - -// on an X-Conversation-Id header, create or replace the session and attach a producer pipe to -// res. no-op when absent, called from the server_res_generator constructor -void stream_session_attach_pipe(server_http_res & res, const std::map & headers); +// on an X-Conversation-Id header, create or replace the session and attach a producer pipe to res +void server_stream_session_attach_pipe(server_http_res & res, const std::map & headers); // should_stop closure that ignores peer disconnect when a pipe is attached, so only an explicit // DELETE stops the producer and generation keeps flowing into the ring buffer. without a pipe it // delegates to fallback, the legacy non-resumable flow -std::function stream_aware_should_stop(server_http_res * res, std::function fallback); +std::function server_stream_aware_should_stop(server_http_res * res, std::function fallback); diff --git a/tools/server/server-task.cpp b/tools/server/server-task.cpp index 775f50baf..8d611e520 100644 --- a/tools/server/server-task.cpp +++ b/tools/server/server-task.cpp @@ -730,6 +730,10 @@ json server_task_result_cmpl_final::to_json_oaicompat_resp_stream() { }} }); + if (timings.prompt_n >= 0) { + server_sent_events.back().at("data").push_back({"timings", timings.to_json()}); + } + return server_sent_events; } @@ -1016,6 +1020,7 @@ void server_task_result_cmpl_partial::update(task_result_state & state) { thinking_block_started = state.thinking_block_started; text_block_started = state.text_block_started; + oai_resp_created = state.oai_resp_created; oai_resp_id = state.oai_resp_id; oai_resp_reasoning_id = state.oai_resp_reasoning_id; oai_resp_message_id = state.oai_resp_message_id; @@ -1024,6 +1029,10 @@ void server_task_result_cmpl_partial::update(task_result_state & state) { // track if the accumulated message has any reasoning content anthropic_has_reasoning = !state.chat_msg.reasoning_content.empty(); + if (res_type == TASK_RESPONSE_TYPE_OAI_RESP && !state.oai_resp_created && (is_progress || n_decoded == 1)) { + state.oai_resp_created = true; + } + // Pre-compute state updates based on diffs (for next chunk) for (const common_chat_msg_diff & diff : oaicompat_msg_diffs) { if (!diff.reasoning_content_delta.empty() && !state.thinking_block_started) { @@ -1181,7 +1190,7 @@ json server_task_result_cmpl_partial::to_json_oaicompat_chat() { json server_task_result_cmpl_partial::to_json_oaicompat_resp() { std::vector events; - if (n_decoded == 1) { + if (!oai_resp_created) { events.push_back(json { {"event", "response.created"}, {"data", json { @@ -1204,6 +1213,18 @@ json server_task_result_cmpl_partial::to_json_oaicompat_resp() { }}, }}, }); + } else if (is_progress) { + events.push_back(json { + {"event", "response.in_progress"}, + {"data", json { + {"type", "response.in_progress"}, + {"response", json { + {"id", oai_resp_id}, + {"object", "response"}, + {"status", "in_progress"}, + }}, + }}, + }); } for (const common_chat_msg_diff & diff : oaicompat_msg_diffs) { @@ -1302,6 +1323,17 @@ json server_task_result_cmpl_partial::to_json_oaicompat_resp() { }); } } + + if (!events.empty()) { + json & data = events.back().at("data"); + if (timings.prompt_n >= 0) { + data.push_back({"timings", timings.to_json()}); + } + if (is_progress) { + data.push_back({"prompt_progress", progress.to_json()}); + } + } + return events; } @@ -1631,7 +1663,22 @@ server_prompt * server_prompt_cache::alloc(const server_prompt & prompt, size_t } } - // next, remove any cached prompts that are fully contained in the current prompt + // 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(); + } + + const size_t state_size_new = state_size_tgt + state_size_dft + checkpoints_size; + + // 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; + } + + // 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); @@ -1644,6 +1691,16 @@ server_prompt * server_prompt_cache::alloc(const server_prompt & prompt, size_t } } + 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(); + } + } + std::vector state_data_tgt; std::vector state_data_dft; @@ -1752,12 +1809,7 @@ bool server_prompt_cache::load(server_prompt & prompt, const server_tokens & tok 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(); @@ -1771,11 +1823,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)); diff --git a/tools/server/server-task.h b/tools/server/server-task.h index 49f62d386..dc6b2dac1 100644 --- a/tools/server/server-task.h +++ b/tools/server/server-task.h @@ -117,6 +117,7 @@ struct task_result_state { bool text_block_started = false; // for OpenAI Responses streaming API + bool oai_resp_created = false; const std::string oai_resp_id; const std::string oai_resp_reasoning_id; const std::string oai_resp_message_id; @@ -440,6 +441,7 @@ struct server_task_result_cmpl_partial : server_task_result { bool text_block_started = false; // for OpenAI Responses API + bool oai_resp_created = false; std::string oai_resp_id; std::string oai_resp_reasoning_id; std::string oai_resp_message_id; diff --git a/tools/server/server.cpp b/tools/server/server.cpp index eafef86ba..9e8603be6 100644 --- a/tools/server/server.cpp +++ b/tools/server/server.cpp @@ -85,7 +85,7 @@ int llama_server(int argc, char ** argv) { // start the stream session manager GC right after common init, before any HTTP route can // touch it. lifecycle is symmetric, stop_gc() runs in clean_up() before backend free - g_stream_sessions.start_gc(); + server_stream_session_manager_start(); if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_SERVER)) { return 1; @@ -245,8 +245,8 @@ int llama_server(int argc, char ** argv) { ctx_http.post("/slots/:id_slot", ex_wrapper(routes.post_slots)); // resumable streaming, the conversation_id is the session identity end to end. router and - // child wire different handlers under the same paths: a child binds the local g_stream_sessions - // backed factories, the router binds proxies that resolve the owning child through the + // child wire different handlers under the same paths: a child binds the local session + // factories, the router binds proxies that resolve the owning child through the // conv_id -> model map server_http_context::handler_t stream_get_h; server_http_context::handler_t streams_lookup_h; @@ -256,9 +256,9 @@ int llama_server(int argc, char ** argv) { streams_lookup_h = models_routes->router_streams_lookup; stream_delete_h = models_routes->router_stream_delete; } else { - stream_get_h = make_stream_get_handler(); - streams_lookup_h = make_streams_lookup_handler(); - stream_delete_h = make_stream_delete_handler(); + stream_get_h = server_stream_make_get_handler(); + streams_lookup_h = server_stream_make_lookup_handler(); + stream_delete_h = server_stream_make_delete_handler(); } ctx_http.get ("/v1/stream/:conv_id", ex_wrapper(stream_get_h)); // POST /v1/streams/lookup with body {"conversation_ids": [...]}. you can only ask for ids @@ -343,7 +343,7 @@ int llama_server(int argc, char ** argv) { clean_up = [&models_routes]() { SRV_INF("%s: cleaning up before exit...\n", __func__); // stop the session GC first, it finalizes live sessions and wakes pending readers - g_stream_sessions.stop_gc(); + server_stream_session_manager_stop(); if (models_routes.has_value()) { models_routes->stopping.store(true); // maybe redundant, but just to be safe models_routes->models.unload_all(); @@ -371,7 +371,7 @@ int llama_server(int argc, char ** argv) { clean_up = [&ctx_http, &ctx_server]() { SRV_INF("%s: cleaning up before exit...\n", __func__); // stop the session GC first, it finalizes live sessions and wakes pending readers - g_stream_sessions.stop_gc(); + server_stream_session_manager_stop(); ctx_http.stop(); ctx_server.terminate(); llama_backend_free(); diff --git a/tools/server/tests/unit/test_compat_oai_responses.py b/tools/server/tests/unit/test_compat_oai_responses.py index 7aab4a8ba..14528b487 100644 --- a/tools/server/tests/unit/test_compat_oai_responses.py +++ b/tools/server/tests/unit/test_compat_oai_responses.py @@ -71,3 +71,44 @@ def test_responses_stream_with_openai_library(): assert r.response.output[0].id.startswith("msg_") assert gathered_text == r.response.output_text assert match_regex("(Suddenly)+", r.response.output_text) + + +def test_responses_stream_with_llama_telemetry(): + global server + server.n_ctx = 256 + server.n_batch = 32 + server.n_slots = 1 + server.start() + + saw_progress = False + saw_delta_timings = False + completed = None + + res = server.make_stream_request("POST", "/responses", data={ + "input": "This is a test" * 10, + "max_output_tokens": 8, + "temperature": 0.8, + "stream": True, + "timings_per_token": True, + "return_progress": True, + }) + + for data in res: + if "prompt_progress" in data: + assert data["type"] == "response.in_progress" + assert data["prompt_progress"]["total"] > 0 + assert data["prompt_progress"]["processed"] >= data["prompt_progress"]["cache"] + saw_progress = True + if "timings" in data: + assert "prompt_per_second" in data["timings"] + assert "predicted_per_second" in data["timings"] + if data["type"] == "response.output_text.delta": + saw_delta_timings = True + if data["type"] == "response.completed": + completed = data + + assert saw_progress + assert saw_delta_timings + assert completed is not None + assert "usage" in completed["response"] + assert "timings" in completed diff --git a/tools/ui/package-lock.json b/tools/ui/package-lock.json index 9dce3a0c9..7216de682 100644 --- a/tools/ui/package-lock.json +++ b/tools/ui/package-lock.json @@ -11,7 +11,7 @@ "@chromatic-com/storybook": "5.0.0", "@eslint/compat": "1.4.1", "@eslint/js": "9.39.2", - "@internationalized/date": "3.10.1", + "@internationalized/date": "3.12.2", "@lucide/svelte": "0.515.0", "@modelcontextprotocol/sdk": "1.26.0", "@playwright/test": "1.56.1", @@ -2981,9 +2981,9 @@ } }, "node_modules/@internationalized/date": { - "version": "3.10.1", - "resolved": "https://registry.npmjs.org/@internationalized/date/-/date-3.10.1.tgz", - "integrity": "sha512-oJrXtQiAXLvT9clCf1K4kxp3eKsQhIaZqxEyowkBcsvZDdZkbWrVmnGknxs5flTD0VGsxrxKgBCZty1EzoiMzA==", + "version": "3.12.2", + "resolved": "https://registry.npmjs.org/@internationalized/date/-/date-3.12.2.tgz", + "integrity": "sha512-FY1Y+H64NDs+HAF6omlnWxm3mEpfgaCSWtL5l551ZZfImA+kGjPFgrnJrGjH6lfmLL0g8Z/mBu1R3kufeCp6Jw==", "dev": true, "license": "Apache-2.0", "dependencies": { diff --git a/tools/ui/package.json b/tools/ui/package.json index bcb4165d1..8b3516a02 100644 --- a/tools/ui/package.json +++ b/tools/ui/package.json @@ -30,7 +30,7 @@ "@chromatic-com/storybook": "5.0.0", "@eslint/compat": "1.4.1", "@eslint/js": "9.39.2", - "@internationalized/date": "3.10.1", + "@internationalized/date": "3.12.2", "@lucide/svelte": "0.515.0", "@modelcontextprotocol/sdk": "1.26.0", "@playwright/test": "1.56.1", diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddDropdown.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddDropdown.svelte index 479540321..433a3662d 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddDropdown.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddDropdown.svelte @@ -11,7 +11,8 @@ } from '$lib/constants'; import { ChatFormActionAddToolsSubmenu, - ChatFormActionAddMcpServersSubmenu + ChatFormActionAddMcpServersSubmenu, + ChatFormActionAddReasoningSubmenu } from '$lib/components/app'; import { useAttachmentMenu } from '$lib/hooks/use-attachment-menu.svelte'; @@ -92,7 +93,11 @@ - + + + + + diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormReasoningToggle.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddReasoningSubmenu.svelte similarity index 63% rename from tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormReasoningToggle.svelte rename to tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddReasoningSubmenu.svelte index f6bcbcb09..070fd3ac6 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormReasoningToggle.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddReasoningSubmenu.svelte @@ -2,7 +2,7 @@ import { Lightbulb, LightbulbOff, Check, Info } from '@lucide/svelte'; import * as DropdownMenu from '$lib/components/ui/dropdown-menu'; import * as Tooltip from '$lib/components/ui/tooltip'; - import { ReasoningEffort, MessageRole } from '$lib/enums'; + import { ReasoningEffort } from '$lib/enums'; import { REASONING_EFFORT_TOKENS } from '$lib/constants/reasoning-effort-tokens'; import { REASONING_EFFORT_LEVELS } from '$lib/constants/reasoning-effort'; import type { ReasoningEffortLevel } from '$lib/types'; @@ -18,31 +18,23 @@ import { isRouterMode } from '$lib/stores/server.svelte'; import type { DatabaseMessage } from '$lib/types/database'; - let thinkingEnabled = $derived(conversationsStore.getThinkingEnabled()); - let currentEffort = $derived(conversationsStore.getReasoningEffort()); - let isOff = $derived(!thinkingEnabled); - let tooltipText = $derived(thinkingEnabled ? `${currentEffort} Reasoning` : 'Disabled Reasoning'); let subOpen = $state(false); - // Get conversation model from message history let conversationModel = $derived( chatStore.getConversationModel(activeMessages() as DatabaseMessage[]) ); - // Fallback: if model props aren't available, check if any assistant messages - // for this model in the active conversation have reasoning content. let modelSupportsThinkingFromMessages = $derived.by(() => { const modelId = isRouterMode() ? modelsStore.selectedModelName || conversationModel : null; if (!modelId) return false; + const messages = conversationsStore.activeMessages; + return messages.some( - (m: DatabaseMessage) => - m.role === MessageRole.ASSISTANT && m.model === modelId && !!m.reasoningContent + (m) => m.role === 'assistant' && m.model === modelId && !!m.reasoningContent ); }); - // Check if model supports thinking. Primary: chat template from /props. - // Fallback: message history (reasoning content in assistant messages). let modelSupportsThinking = $derived.by(() => { loadedModelIds(); propsCacheVersion(); @@ -52,15 +44,15 @@ return checkModelSupportsThinking(modelId ?? '') || modelSupportsThinkingFromMessages; } - // In non-router mode, use the built-in supportsThinking return supportsThinking() || modelSupportsThinkingFromMessages; }); - // Check if current item is selected + let thinkingEnabled = $derived(conversationsStore.getThinkingEnabled()); + let currentEffort = $derived(conversationsStore.getReasoningEffort()); + let isOff = $derived(!thinkingEnabled); + function isSelected(item: ReasoningEffortLevel): boolean { - if (item.isOff) { - return isOff; - } + if (item.isOff) return isOff; return thinkingEnabled && currentEffort === item.value; } @@ -76,39 +68,30 @@ {#if modelSupportsThinking} - - - - - {#if thinkingEnabled} - - {:else} - - {/if} - - + + + {#if thinkingEnabled} + + {:else} + + {/if} - -

{tooltipText}

-
-
+ + Reasoning - + {thinkingEnabled ? currentEffort : 'off'} + + +
+ + -
Reasoning effort
- {#each REASONING_EFFORT_LEVELS as level (level.value)} {/each} -
- + + {/if} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActions.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActions.svelte index a80f00bc6..7be356f2a 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActions.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActions.svelte @@ -7,14 +7,20 @@ ChatFormActionModels, ChatFormActionRecord, ChatFormActionSubmit, - ChatFormReasoningToggle + ChatFormContextGauge } from '$lib/components/app'; - import { FileTypeCategory } from '$lib/enums'; + import { FileTypeCategory, MessageRole } from '$lib/enums'; import { mcpStore } from '$lib/stores/mcp.svelte'; import { config } from '$lib/stores/settings.svelte'; - import { conversationsStore } from '$lib/stores/conversations.svelte'; + import { activeMessages, conversationsStore } from '$lib/stores/conversations.svelte'; + import { + activeProcessingState, + isChatStreaming, + isLoading as chatIsLoading + } from '$lib/stores/chat.svelte'; import { getFileTypeCategory } from '$lib/utils'; import { goto } from '$app/navigation'; + import { page } from '$app/state'; import { ROUTES } from '$lib/constants/routes'; interface Props { @@ -93,6 +99,36 @@ let activeMessage = $derived( conversationsStore.activeMessages[conversationsStore.activeMessages.length - 1] ); + + let hasProcessedTokens = $derived.by(() => { + if (!page.params.id) return false; + + const messages = activeMessages() as DatabaseMessage[]; + let totalHistoricalTokens = 0; + for (const m of messages) { + if (m.role !== MessageRole.ASSISTANT) continue; + const timings = m.timings; + if (!timings) continue; + const agenticLlm = timings.agentic?.llm; + if (agenticLlm?.prompt_n != null || agenticLlm?.predicted_n != null) { + totalHistoricalTokens += (agenticLlm?.prompt_n ?? 0) + (agenticLlm?.predicted_n ?? 0); + } else { + totalHistoricalTokens += (timings.prompt_n ?? 0) + (timings.predicted_n ?? 0); + } + } + if (totalHistoricalTokens > 0) return true; + + if (!chatIsLoading() && !isChatStreaming()) return false; + + const processingState = activeProcessingState(); + if (!processingState) return false; + const livePromptTokens = Math.max( + processingState.promptTokens ?? 0, + processingState.promptProgress?.processed ?? 0 + ); + const liveOutputTokens = processingState.outputTokensUsed ?? 0; + return livePromptTokens > 0 || liveOutputTokens > 0; + });
{#if showAddButton} -
+
{/if} -
- +
+ {#if hasProcessedTokens} + + {/if} {#if showModelSelector} - + {#if thinkingEnabled} {:else} @@ -89,23 +88,15 @@ {/if} - + {#each REASONING_EFFORT_LEVELS as level (level.value)} {/each} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ChatFormContextGauge.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ChatFormContextGauge.svelte new file mode 100644 index 000000000..855cf6ce7 --- /dev/null +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ChatFormContextGauge.svelte @@ -0,0 +1,108 @@ + + + + + + + + +
+
+ Context + · + + {formatParameters(gauge.contextUsed)} + / {gauge.contextTotal !== null ? formatParameters(gauge.contextTotal) : '-'} + +
+ + {#if gauge.activeModelId !== null && !gauge.isActiveModelLoaded} + + {:else if showProgressBar} +
+
+
+ +
+ + {gauge.contextPercent}% used + + + {formatParameters((gauge.contextTotal ?? 0) - gauge.contextUsed)} remaining + +
+ {:else} +
No context info available
+ {/if} + + {#if gauge.hasAnyUsage} + + {/if} +
+
+
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetailRow.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetailRow.svelte new file mode 100644 index 000000000..71c6a33cb --- /dev/null +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetailRow.svelte @@ -0,0 +1,20 @@ + + +
+
+ {label} + {value} +
+ + {#if subtitle} +
{subtitle}
+ {/if} +
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetails.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetails.svelte new file mode 100644 index 000000000..fdec5aca5 --- /dev/null +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetails.svelte @@ -0,0 +1,122 @@ + + + + + Token usage details + + + + + + {#if hasCumulative} +
+

+ Across all turns +

+ +
+ {#if cumulativeRead > 0} + 0 + ? `${cumulativeCacheTotal.toLocaleString()} reused from KV cache` + : undefined} + /> + {/if} + {#if cumulativeOutput > 0} + + {/if} +
+
+ {/if} + + {#if hasCurrent} +
+

+ This turn · KV cache +

+ +
+ {#if currentRead > 0} + 0 + ? `${currentFresh.toLocaleString()} fresh + ${currentCache.toLocaleString()} cached` + : undefined} + /> + {/if} + + {#if currentOutput > 0} + + {/if} + +
+
+ KV cache total + {kvTotal.toLocaleString()} tok +
+
+
+
+ {/if} + + {#if averageTokensPerSecond !== null} +
+ +
+ {/if} + + {#each transientDetails as detail (detail)} +
{detail}
+ {/each} +
+
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDial.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDial.svelte new file mode 100644 index 000000000..6e2616d36 --- /dev/null +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDial.svelte @@ -0,0 +1,43 @@ + + + + + + + diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeLoadModel.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeLoadModel.svelte new file mode 100644 index 000000000..24a67cfdf --- /dev/null +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeLoadModel.svelte @@ -0,0 +1,24 @@ + + +{#if modelId !== null && !isLoading} +
+ Available context size is only visible once the model is loaded. + +
+{:else if isLoading} +
+ + Loading model... +
+{/if} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/context-gauge.ts b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/context-gauge.ts new file mode 100644 index 000000000..5b0010015 --- /dev/null +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/context-gauge.ts @@ -0,0 +1,37 @@ +export type ColorLevel = 'ok' | 'warning' | 'critical' | 'neutral'; + +const WARNING_THRESHOLD = 80; +const CRITICAL_THRESHOLD = 95; + +export function colorLevelFromPercent(percent: number | null): ColorLevel { + if (percent === null) return 'neutral'; + if (percent >= CRITICAL_THRESHOLD) return 'critical'; + if (percent >= WARNING_THRESHOLD) return 'warning'; + return 'ok'; +} + +export function colorLevelTextClass(level: ColorLevel): string { + switch (level) { + case 'critical': + return 'text-red-400'; + case 'warning': + return 'text-amber-400'; + case 'ok': + return 'text-muted-foreground'; + default: + return 'text-muted-foreground'; + } +} + +export function colorLevelBgClass(level: ColorLevel): string { + switch (level) { + case 'critical': + return 'bg-red-500'; + case 'warning': + return 'bg-amber-500'; + case 'ok': + return 'bg-green-500'; + default: + return 'bg-muted'; + } +} diff --git a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte index dadcae0c4..2b4a3f9d3 100644 --- a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte @@ -24,6 +24,8 @@ message: DatabaseMessage; toolMessages?: DatabaseMessage[]; isLastAssistantMessage?: boolean; + isLastUserMessage?: boolean; + nextAssistantMessage?: DatabaseMessage | null; siblingInfo?: ChatMessageSiblingInfo | null; } @@ -32,6 +34,8 @@ message, toolMessages = [], isLastAssistantMessage = false, + isLastUserMessage = false, + nextAssistantMessage = null, siblingInfo = null }: Props = $props(); @@ -359,7 +363,9 @@ (ChatMessageStatsView.GENERATION); - let statsContainerEl: HTMLDivElement | undefined = $state(); - - function getScrollParent(el: HTMLElement): HTMLElement | null { - let parent = el.parentElement; - while (parent) { - const style = getComputedStyle(parent); - if (/(auto|scroll)/.test(style.overflowY)) { - return parent; - } - parent = parent.parentElement; - } - return null; - } - - async function handleStatsViewChange(view: ChatMessageStatsView) { - const el = statsContainerEl; - if (!el) { - activeStatsView = view; - - return; - } - - const scrollParent = getScrollParent(el); - if (!scrollParent) { - activeStatsView = view; - - return; - } - - const yBefore = el.getBoundingClientRect().top; - - activeStatsView = view; - - await tick(); - - const delta = el.getBoundingClientRect().top - yBefore; - if (delta !== 0) { - scrollParent.scrollTop += delta; - } - - // Correct any drift after browser paint - requestAnimationFrame(() => { - const drift = el.getBoundingClientRect().top - yBefore; - - if (Math.abs(drift) > 1) { - scrollParent.scrollTop += drift; - } - }); - } - - let highlightAgenticTurns = $derived( - isAgentic && - (currentConfig.alwaysShowAgenticTurns || activeStatsView === ChatMessageStatsView.SUMMARY) - ); - let displayedModel = $derived(message.model ?? null); // model being switched to while it loads, so the selector bar tracks it @@ -291,7 +234,6 @@ {toolMessages} isStreaming={isChatStreaming()} {isLastAssistantMessage} - highlightTurns={highlightAgenticTurns} /> {/if} {:else} @@ -315,10 +257,7 @@
{#if displayedModel} -
+
{#if isRouter} {:else if isLoading() && currentConfig.showMessageStats} {@const liveStats = processingState.getLiveProcessingStats()} {@const genStats = processingState.getLiveGenerationStats()} - {@const promptProgress = processingState.processingState?.promptProgress} - {@const isStillProcessingPrompt = - promptProgress && promptProgress.processed < promptProgress.total} - {#if liveStats || genStats} + {#if genStats} {/if} {/if} diff --git a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessageUser/ChatMessageUser.svelte b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessageUser/ChatMessageUser.svelte index 80a0183e6..f7590c3a3 100644 --- a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessageUser/ChatMessageUser.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessageUser/ChatMessageUser.svelte @@ -2,10 +2,14 @@ import { ChatMessageActionIcons, ChatMessageEditForm, + ChatMessageStatistics, ChatMessageUserBubble } from '$lib/components/app/chat'; import { getMessageEditContext } from '$lib/contexts'; - import { MessageRole } from '$lib/enums'; + import { useProcessingState } from '$lib/hooks/use-processing-state.svelte'; + import { isLoading } from '$lib/stores/chat.svelte'; + import { MessageRole, ChatMessageStatisticsMode } from '$lib/enums'; + import { config } from '$lib/stores/settings.svelte'; interface Props { class?: string; @@ -17,6 +21,8 @@ assistantMessages: number; messageTypes: string[]; } | null; + isLastUserMessage?: boolean; + nextAssistantMessage?: DatabaseMessage | null; showDeleteDialog: boolean; onEdit: () => void; onDelete: () => void; @@ -32,6 +38,8 @@ message, siblingInfo = null, deletionInfo, + isLastUserMessage = false, + nextAssistantMessage = null, showDeleteDialog, onEdit, onDelete, @@ -44,6 +52,37 @@ // Get contexts const editCtx = getMessageEditContext(); + const processingState = useProcessingState(); + + const currentConfig = $derived(config()); + const isActivelyProcessing = $derived(isLastUserMessage && isLoading()); + + // For agentic turns, prefer the cumulative agentic.llm totals over per-call timings. + let storedReadingStats = $derived.by(() => { + const timings = nextAssistantMessage?.timings; + if (!timings?.prompt_n || !timings?.prompt_ms) return null; + + const agentic = timings.agentic; + + return { + promptTokens: agentic ? agentic.llm.prompt_n : timings.prompt_n, + promptMs: agentic ? agentic.llm.prompt_ms : timings.prompt_ms + }; + }); + + let showStoredReadingStats = $derived( + Boolean(currentConfig.showMessageStats) && storedReadingStats !== null + ); + + let showLiveReadingStats = $derived( + Boolean(currentConfig.showMessageStats) && isActivelyProcessing && storedReadingStats === null + ); + + $effect(() => { + if (showLiveReadingStats) { + processingState.startMonitoring(); + } + });
+ {#if showStoredReadingStats} + +
+
+ +
+
+ {:else if showLiveReadingStats} + {@const liveStats = processingState.getLiveProcessingStats()} + {#if liveStats} +
+
+ +
+
+ {/if} + {/if} + {#if message.timestamp}
content, getExtras: () => extras, @@ -36,9 +32,7 @@
{#if editCtx.isEditing} diff --git a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageAgenticContent.svelte b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageAgenticContent.svelte index 07b0489cc..f0d03f547 100644 --- a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageAgenticContent.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageAgenticContent.svelte @@ -41,15 +41,13 @@ toolMessages?: DatabaseMessage[]; isStreaming?: boolean; isLastAssistantMessage?: boolean; - highlightTurns?: boolean; } let { message, toolMessages = [], isStreaming = false, - isLastAssistantMessage = false, - highlightTurns = false + isLastAssistantMessage = false }: Props = $props(); let expandedStates: Record = $state({}); @@ -57,6 +55,7 @@ const showToolCallInProgress = $derived(config().showToolCallInProgress as boolean); const showThoughtInProgress = $derived(config().showThoughtInProgress as boolean); const renderThinkingAsMarkdown = $derived(config().renderThinkingAsMarkdown as boolean); + const showMessageStats = $derived(config().showMessageStats as boolean); const hasReasoningError = $derived( isLastAssistantMessage ? !!agenticLastError(message.convId) : false @@ -354,16 +353,17 @@ {/snippet}
- {#if highlightTurns && turnGroups.length > 1} + {#if turnGroups.length > 1} {#each turnGroups as turn, turnIndex (turnIndex)} {@const turnStats = message?.timings?.agentic?.perTurn?.[turnIndex]} -
- Turn {turnIndex + 1} + +
{#each turn.sections as section, sIdx (turn.flatIndices[sIdx])} {@render renderSection(section, turn.flatIndices[sIdx])} {/each} - {#if turnStats} -
+ + {#if turnStats && showMessageStats} +
:global(*), + .agentic-turn > :global(*) { + min-width: 0; } .agentic-text { width: 100%; } - .agentic-turn { - position: relative; - border: 1.5px dashed var(--muted-foreground); - border-radius: 0.75rem; - padding: 1rem; - transition: background 0.1s; - } - - .agentic-turn-label { - position: absolute; - top: -1rem; - left: 0.75rem; - padding: 0 0.375rem; - background: var(--background); - font-size: 0.7rem; - font-weight: 500; - color: var(--muted-foreground); - text-transform: uppercase; - letter-spacing: 0.05em; - } - .turn-stats { - margin-top: 0.75rem; - padding-top: 0.5rem; border-top: 1px solid hsl(var(--muted) / 0.5); } diff --git a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageStatistics/ChatMessageStatistics.svelte b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageStatistics/ChatMessageStatistics.svelte index 6906adbb1..7ef73a499 100644 --- a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageStatistics/ChatMessageStatistics.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessageStatistics/ChatMessageStatistics.svelte @@ -2,7 +2,7 @@ import { Clock, Gauge, WholeWord, BookOpenText, Sparkles, Wrench, Layers } from '@lucide/svelte'; import { ChatMessageStatisticsBadge } from '$lib/components/app'; import * as Tooltip from '$lib/components/ui/tooltip'; - import { ChatMessageStatsView } from '$lib/enums'; + import { ChatMessageStatsView, ChatMessageStatisticsMode } from '$lib/enums'; import type { ChatMessageAgenticTimings } from '$lib/types/chat'; import { formatPerformanceTime } from '$lib/utils'; import { MS_PER_SECOND, DEFAULT_PERFORMANCE_TIME } from '$lib/constants'; @@ -19,6 +19,7 @@ agenticTimings?: ChatMessageAgenticTimings; onActiveViewChange?: (view: ChatMessageStatsView) => void; hideSummary?: boolean; + mode?: ChatMessageStatisticsMode; } let { @@ -31,19 +32,30 @@ initialView = ChatMessageStatsView.GENERATION, agenticTimings, onActiveViewChange, - hideSummary = false + hideSummary = false, + mode = ChatMessageStatisticsMode.SWITCHABLE }: Props = $props(); - let activeView: ChatMessageStatsView = $derived(initialView); + let isSwitchable = $derived(mode === ChatMessageStatisticsMode.SWITCHABLE); + + let activeView: ChatMessageStatsView = $derived( + mode === ChatMessageStatisticsMode.READING + ? ChatMessageStatsView.READING + : mode === ChatMessageStatisticsMode.GENERATION + ? ChatMessageStatsView.GENERATION + : initialView + ); let hasAutoSwitchedToGeneration = $state(false); $effect(() => { - onActiveViewChange?.(activeView); + if (isSwitchable) { + onActiveViewChange?.(activeView); + } }); // In live mode: auto-switch to GENERATION tab when prompt processing completes $effect(() => { - if (isLive) { + if (isLive && isSwitchable) { // Auto-switch to generation tab only when prompt processing is done (once) if ( !hasAutoSwitchedToGeneration && @@ -91,8 +103,7 @@ formattedPromptTime !== undefined ); - // In live mode, generation tab is disabled until we have generation stats - let isGenerationDisabled = $derived(isLive && !hasGenerationStats); + let isGenerationDisabled = $derived(isLive && isSwitchable && !hasGenerationStats); let hasAgenticStats = $derived(agenticTimings !== undefined && agenticTimings.toolCallsCount > 0); @@ -153,44 +164,44 @@ {/snippet}
-
- {#if hasPromptStats || isLive} - {@render viewButton({ - view: ChatMessageStatsView.READING, - icon: BookOpenText, - label: 'Reading', - tooltipText: 'Reading (prompt processing)' - })} - {/if} - - {@render viewButton({ - view: ChatMessageStatsView.GENERATION, - icon: Sparkles, - label: 'Generation', - tooltipText: isGenerationDisabled - ? 'Generation (waiting for tokens...)' - : 'Generation (token output)', - disabled: isGenerationDisabled - })} - - {#if hasAgenticStats} - {@render viewButton({ - view: ChatMessageStatsView.TOOLS, - icon: Wrench, - label: 'Tools', - tooltipText: 'Tool calls' - })} - - {#if !hideSummary} + {#if isSwitchable} +
+ {#if hasPromptStats || isLive} {@render viewButton({ - view: ChatMessageStatsView.SUMMARY, - icon: Layers, - label: 'Summary', - tooltipText: 'Agentic summary' + view: ChatMessageStatsView.READING, + icon: BookOpenText, + label: 'Reading', + tooltipText: 'Processing' })} {/if} - {/if} -
+ + {@render viewButton({ + view: ChatMessageStatsView.GENERATION, + icon: Sparkles, + label: 'Generation', + tooltipText: isGenerationDisabled ? 'Waiting for tokens...' : 'Generation', + disabled: isGenerationDisabled + })} + + {#if hasAgenticStats} + {@render viewButton({ + view: ChatMessageStatsView.TOOLS, + icon: Wrench, + label: 'Tools', + tooltipText: 'Tool calls' + })} + + {#if !hideSummary} + {@render viewButton({ + view: ChatMessageStatsView.SUMMARY, + icon: Layers, + label: 'Summary', + tooltipText: 'Agentic summary' + })} + {/if} + {/if} +
+ {/if}
{#if activeView === ChatMessageStatsView.GENERATION && hasGenerationStats} @@ -256,7 +267,7 @@ value={formattedAgenticTotalTime} tooltipLabel="Total time (LLM + tools)" /> - {:else if hasPromptStats} + {:else if hasPromptStats && (mode === ChatMessageStatisticsMode.READING || isSwitchable)} = []; @@ -236,18 +238,36 @@ message: msg, toolMessages, isLastAssistantMessage: false, + isLastUserMessage: false, + nextAssistantMessage: null, siblingInfo }); } - // Mark the last assistant message + let lastAssistantIdx = -1; for (let i = result.length - 1; i >= 0; i--) { if (result[i].message.role === MessageRole.ASSISTANT) { result[i].isLastAssistantMessage = true; + lastAssistantIdx = i; break; } } + if (lastAssistantIdx > 0 && result[lastAssistantIdx - 1].message.role === MessageRole.USER) { + result[lastAssistantIdx - 1].isLastUserMessage = true; + } + + for (let i = 0; i < result.length; i++) { + if (result[i].message.role !== MessageRole.USER) continue; + + for (let j = i + 1; j < result.length; j++) { + if (result[j].message.role === MessageRole.ASSISTANT) { + result[i].nextAssistantMessage = result[j].message; + break; + } + } + } + return result; }); @@ -257,12 +277,14 @@ {isVisible ? 'opacity-100' : 'opacity-0'} {previousRouteId === '/(chat)/chat/[id]' ? '' : 'delay-300'}" > - {#each displayMessages as { message, toolMessages, isLastAssistantMessage, siblingInfo } (message.id)} + {#each displayMessages as { message, toolMessages, isLastAssistantMessage, isLastUserMessage, nextAssistantMessage, siblingInfo } (message.id)} {/each} diff --git a/tools/ui/src/lib/components/app/chat/ChatScreen/ChatScreen.svelte b/tools/ui/src/lib/components/app/chat/ChatScreen/ChatScreen.svelte index d2fba5a93..a0d66c642 100644 --- a/tools/ui/src/lib/components/app/chat/ChatScreen/ChatScreen.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatScreen/ChatScreen.svelte @@ -4,12 +4,10 @@ ChatScreenForm, ChatMessages, ChatScreenDragOverlay, - ChatScreenProcessingInfo, ChatScreenStreamResumeStatus, ServerLoadingSplash, ChatScreenServerError } from '$lib/components/app'; - import { setProcessingInfoContext } from '$lib/contexts'; import { createAutoScrollController } from '$lib/hooks/use-auto-scroll.svelte'; import { useChatScreenActiveModel } from '$lib/hooks/use-chat-screen-active-model.svelte'; import { useChatScreenDragAndDrop } from '$lib/hooks/use-chat-screen-drag-and-drop.svelte'; @@ -23,8 +21,7 @@ errorDialog, isLoading, isChatStreaming, - isEditing, - activeProcessingState + isEditing } from '$lib/stores/chat.svelte'; import { conversationsStore, @@ -42,12 +39,6 @@ let { showCenteredEmpty = false } = $props(); - setProcessingInfoContext({ - get showProcessingInfo() { - return showProcessingInfo; - } - }); - let disableAutoScroll = $derived(Boolean(config().disableAutoScroll) || isMobile.current); let isMobileUserScrolledUp = $state(false); let mobileScrollDownHint = $state(false); @@ -63,11 +54,6 @@ let isServerLoading = $derived(serverLoading()); let hasPropsError = $derived(!!serverError()); let isCurrentConversationLoading = $derived(isLoading() || isChatStreaming()); - let showProcessingInfo = $derived( - isCurrentConversationLoading || - (config().keepStatsVisible && !!page.params.id) || - activeProcessingState() !== null - ); let chatFormBottomPosition = $derived.by(() => { if (!isMobile.current) return '1rem'; if (device.isStandalone) return '1.5rem'; @@ -298,10 +284,6 @@ }} /> {/if} - - {#if showProcessingInfo} - - {/if}
- import { untrack } from 'svelte'; - import { PROCESSING_INFO_TIMEOUT } from '$lib/constants'; - import { useProcessingState } from '$lib/hooks/use-processing-state.svelte'; - import { chatStore, isLoading, isChatStreaming } from '$lib/stores/chat.svelte'; - import { activeMessages, activeConversation } from '$lib/stores/conversations.svelte'; - import { config } from '$lib/stores/settings.svelte'; - - const processingState = useProcessingState(); - - let isCurrentConversationLoading = $derived(isLoading()); - let isStreaming = $derived(isChatStreaming()); - let processingDetails = $derived(processingState.getTechnicalDetails()); - - let processingVisible = $derived(processingDetails.length > 0); - - let { onVisibilityChange }: { onVisibilityChange?: (visible: boolean) => void } = $props(); - - $effect(() => { - onVisibilityChange?.(processingVisible); - }); - - $effect(() => { - const conversation = activeConversation(); - - untrack(() => chatStore.setActiveProcessingConversation(conversation?.id ?? null)); - }); - - $effect(() => { - const keepStatsVisible = config().keepStatsVisible; - const shouldMonitor = keepStatsVisible || isCurrentConversationLoading || isStreaming; - - if (shouldMonitor) { - processingState.startMonitoring(); - } - - if (!isCurrentConversationLoading && !isStreaming && !keepStatsVisible) { - const timeout = setTimeout(() => { - if (!config().keepStatsVisible && !isChatStreaming()) { - processingState.stopMonitoring(); - } - }, PROCESSING_INFO_TIMEOUT); - - return () => clearTimeout(timeout); - } - }); - - $effect(() => { - const conversation = activeConversation(); - const messages = activeMessages() as DatabaseMessage[]; - const keepStatsVisible = config().keepStatsVisible; - - if (keepStatsVisible && conversation) { - if (messages.length === 0) { - untrack(() => chatStore.clearProcessingState(conversation.id)); - return; - } - - if (!isCurrentConversationLoading && !isStreaming) { - untrack(() => chatStore.restoreProcessingStateFromMessages(messages, conversation.id)); - } - } - }); - - - - - diff --git a/tools/ui/src/lib/components/app/chat/index.ts b/tools/ui/src/lib/components/app/chat/index.ts index d7004d3ae..4f826841e 100644 --- a/tools/ui/src/lib/components/app/chat/index.ts +++ b/tools/ui/src/lib/components/app/chat/index.ts @@ -241,13 +241,18 @@ export { default as ChatFormActionAddToolsSubmenu } from './ChatForm/ChatFormAct export { default as ChatFormActionAddMcpServersSubmenu } from './ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddMcpServersSubmenu.svelte'; /** - * **ChatFormReasoningToggle** - Thinking toggle button with effort dropdown + * Dropdown submenu for selecting reasoning effort level. * - * A toggle button with lightbulb icon that indicates thinking status. - * Shows the reasoning effort dropdown when clicked. + * Shows a "Reasoning" sub-menu item with a lightbulb icon indicating + * thinking status, and a nested list of effort levels. * Only visible when the current model supports thinking. */ -export { default as ChatFormReasoningToggle } from './ChatForm/ChatFormActions/ChatFormReasoningToggle.svelte'; +export { default as ChatFormActionAddReasoningSubmenu } from './ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddReasoningSubmenu.svelte'; + +/** + * Compact context-usage gauge with per-turn and cumulative breakdown in the tooltip. + */ +export { default as ChatFormContextGauge } from './ChatForm/ChatFormContextGauge/ChatFormContextGauge.svelte'; /** * Hidden file input element for programmatic file selection. @@ -669,14 +674,6 @@ export { default as ChatScreenDragOverlay } from './ChatScreen/ChatScreenDragOve */ export { default as ChatScreenForm } from './ChatScreen/ChatScreenForm.svelte'; -/** - * Processing info display during generation. Shows real-time statistics: - * tokens per second, prompt/completion token counts, and elapsed time. - * Data sourced from slotsService polling during active generation. - * Only visible when `isCurrentConversationLoading` is true. - */ -export { default as ChatScreenProcessingInfo } from './ChatScreen/ChatScreenProcessingInfo.svelte'; - /** * Server error alert displayed when the server is unreachable. * Shows the error message with a retry button. diff --git a/tools/ui/src/lib/components/app/content/CollapsibleContentBlock.svelte b/tools/ui/src/lib/components/app/content/CollapsibleContentBlock.svelte index 8bab55d19..3875b449a 100644 --- a/tools/ui/src/lib/components/app/content/CollapsibleContentBlock.svelte +++ b/tools/ui/src/lib/components/app/content/CollapsibleContentBlock.svelte @@ -76,7 +76,7 @@ open = value; onToggle?.(); }} - class={className} + class="{className} my-0!" > diff --git a/tools/ui/src/lib/components/app/content/SyntaxHighlightedCode.svelte b/tools/ui/src/lib/components/app/content/SyntaxHighlightedCode.svelte index c4d1706bf..d8dfe0ad9 100644 --- a/tools/ui/src/lib/components/app/content/SyntaxHighlightedCode.svelte +++ b/tools/ui/src/lib/components/app/content/SyntaxHighlightedCode.svelte @@ -72,8 +72,8 @@
{@html highlightedHtml}
diff --git a/tools/ui/src/lib/components/app/dialogs/DialogMcpServerRecommendations.svelte b/tools/ui/src/lib/components/app/dialogs/DialogMcpServerRecommendations.svelte index 9b4489b82..cdbc055ee 100644 --- a/tools/ui/src/lib/components/app/dialogs/DialogMcpServerRecommendations.svelte +++ b/tools/ui/src/lib/components/app/dialogs/DialogMcpServerRecommendations.svelte @@ -4,9 +4,10 @@ import * as Dialog from '$lib/components/ui/dialog'; import { fly } from 'svelte/transition'; import { McpServerCardCompact, McpServerForm } from '$lib/components/app/mcp'; - import { RECOMMENDED_MCP_SERVERS } from '$lib/constants'; + import { RECOMMENDED_MCP_SERVERS, SETTINGS_KEYS } from '$lib/constants'; import { conversationsStore } from '$lib/stores/conversations.svelte'; import { mcpStore } from '$lib/stores/mcp.svelte'; + import { settingsStore } from '$lib/stores/settings.svelte'; import { uuid } from '$lib/utils'; import { MCP_SERVERS_ADDED_TO_CHAT_LOCALSTORAGE_KEY, MCP_SERVER_ID_PREFIX } from '$lib/constants'; import type { MCPServerSettingsEntry } from '$lib/types'; @@ -24,6 +25,22 @@ ); let addedServers = $state([]); + let didAddAny = $state(false); + + let selectedRecommendedCount = $derived.by( + () => RECOMMENDED_MCP_SERVERS.filter((server) => selected[server.id]).length + ); + + let footerLabel = $derived.by(() => { + const recommended = selectedRecommendedCount; + const custom = addedServers.length; + const total = recommended + custom; + + if (total === 0) return 'Continue'; + if (recommended === 0) return custom === 1 ? 'Add server' : `Add ${custom} servers`; + if (custom === 0) return recommended === 1 ? 'Add server' : `Add ${recommended} servers`; + return `Add ${recommended} servers and ${custom} custom`; + }); let showAddForm = $state(false); let newServerUrl = $state(''); @@ -44,9 +61,14 @@ showAddForm = false; newServerUrl = ''; newServerHeaders = ''; - addedServers = []; + + if (!didAddAny) { + settingsStore.updateConfig(SETTINGS_KEYS.MCP_SERVERS, []); + } localStorage.setItem(MCP_SERVERS_ADDED_TO_CHAT_LOCALSTORAGE_KEY, 'true'); + addedServers = []; + didAddAny = false; } open = value; onOpenChange?.(value); @@ -59,6 +81,7 @@ } function enableSelected() { + didAddAny = true; localStorage.setItem(MCP_SERVERS_ADDED_TO_CHAT_LOCALSTORAGE_KEY, 'true'); for (const server of RECOMMENDED_MCP_SERVERS) { @@ -83,6 +106,8 @@ function saveNewServer() { if (newServerUrlError) return; + didAddAny = true; + const newServerId = uuid() ?? `${MCP_SERVER_ID_PREFIX}-${Date.now()}`; localStorage.setItem(MCP_SERVERS_ADDED_TO_CHAT_LOCALSTORAGE_KEY, 'true'); @@ -174,7 +199,12 @@ - + diff --git a/tools/ui/src/lib/components/app/settings/SettingsChat/SettingsChatToolsTab.svelte b/tools/ui/src/lib/components/app/settings/SettingsChat/SettingsChatToolsTab.svelte index b56832496..8bb49a095 100644 --- a/tools/ui/src/lib/components/app/settings/SettingsChat/SettingsChatToolsTab.svelte +++ b/tools/ui/src/lib/components/app/settings/SettingsChat/SettingsChatToolsTab.svelte @@ -39,13 +39,17 @@ {@const faviconUrl = group.serverId ? mcpStore.getServerFavicon(group.serverId) : null} - + {#if group.source === 'mcp'} + + {:else} + + {/if} diff --git a/tools/ui/src/lib/components/ui/hover-card/hover-card-content.svelte b/tools/ui/src/lib/components/ui/hover-card/hover-card-content.svelte new file mode 100644 index 000000000..db1406a7b --- /dev/null +++ b/tools/ui/src/lib/components/ui/hover-card/hover-card-content.svelte @@ -0,0 +1,31 @@ + + + + + diff --git a/tools/ui/src/lib/components/ui/hover-card/hover-card-portal.svelte b/tools/ui/src/lib/components/ui/hover-card/hover-card-portal.svelte new file mode 100644 index 000000000..9dc25827d --- /dev/null +++ b/tools/ui/src/lib/components/ui/hover-card/hover-card-portal.svelte @@ -0,0 +1,7 @@ + + + diff --git a/tools/ui/src/lib/components/ui/hover-card/hover-card-trigger.svelte b/tools/ui/src/lib/components/ui/hover-card/hover-card-trigger.svelte new file mode 100644 index 000000000..2d42f89c3 --- /dev/null +++ b/tools/ui/src/lib/components/ui/hover-card/hover-card-trigger.svelte @@ -0,0 +1,7 @@ + + + diff --git a/tools/ui/src/lib/components/ui/hover-card/hover-card.svelte b/tools/ui/src/lib/components/ui/hover-card/hover-card.svelte new file mode 100644 index 000000000..ffc075dcd --- /dev/null +++ b/tools/ui/src/lib/components/ui/hover-card/hover-card.svelte @@ -0,0 +1,7 @@ + + + diff --git a/tools/ui/src/lib/components/ui/hover-card/index.ts b/tools/ui/src/lib/components/ui/hover-card/index.ts new file mode 100644 index 000000000..098f69176 --- /dev/null +++ b/tools/ui/src/lib/components/ui/hover-card/index.ts @@ -0,0 +1,15 @@ +import Root from './hover-card.svelte'; +import Content from './hover-card-content.svelte'; +import Trigger from './hover-card-trigger.svelte'; +import Portal from './hover-card-portal.svelte'; + +export { + Root, + Content, + Trigger, + Portal, + Root as HoverCard, + Content as HoverCardContent, + Trigger as HoverCardTrigger, + Portal as HoverCardPortal +}; diff --git a/tools/ui/src/lib/constants/context-keys.ts b/tools/ui/src/lib/constants/context-keys.ts index 12de0d0bc..0bd733b37 100644 --- a/tools/ui/src/lib/constants/context-keys.ts +++ b/tools/ui/src/lib/constants/context-keys.ts @@ -1,4 +1,3 @@ export const CONTEXT_KEY_MESSAGE_EDIT = 'chat-message-edit'; export const CONTEXT_KEY_CHAT_ACTIONS = 'chat-actions'; export const CONTEXT_KEY_CHAT_SETTINGS_CONFIG = 'chat-settings-config'; -export const CONTEXT_KEY_PROCESSING_INFO = 'processing-info'; diff --git a/tools/ui/src/lib/constants/css-classes.ts b/tools/ui/src/lib/constants/css-classes.ts index ca5386fcd..009b5c52e 100644 --- a/tools/ui/src/lib/constants/css-classes.ts +++ b/tools/ui/src/lib/constants/css-classes.ts @@ -17,3 +17,4 @@ export const PANEL_CLASSES = ` `; export const CHAT_FORM_POPOVER_MAX_HEIGHT = 'max-h-80'; +export const DIALOG_SUBMENU_CONTENT = 'w-60'; diff --git a/tools/ui/src/lib/constants/reasoning-effort.ts b/tools/ui/src/lib/constants/reasoning-effort.ts index d854e912a..28a24420e 100644 --- a/tools/ui/src/lib/constants/reasoning-effort.ts +++ b/tools/ui/src/lib/constants/reasoning-effort.ts @@ -6,6 +6,7 @@ import type { ReasoningEffortLevel } from '$lib/types'; * Keys match the ReasoningEffort enum values for type-safe lookups. */ export const REASONING_EFFORT_LABELS: Record = { + [ReasoningEffort.OFF]: 'Off', [ReasoningEffort.LOW]: 'Low', [ReasoningEffort.MEDIUM]: 'Medium', [ReasoningEffort.HIGH]: 'High', @@ -13,7 +14,7 @@ export const REASONING_EFFORT_LABELS: Record = { }; export const REASONING_EFFORT_LEVELS: ReasoningEffortLevel[] = [ - { value: 'off', label: 'Off', isOff: true }, + { value: ReasoningEffort.OFF, label: 'Off', isOff: true }, { value: ReasoningEffort.LOW, label: 'Low' }, { value: ReasoningEffort.MEDIUM, label: 'Medium' }, { value: ReasoningEffort.HIGH, label: 'High' }, diff --git a/tools/ui/src/lib/constants/settings-keys.ts b/tools/ui/src/lib/constants/settings-keys.ts index d4d782705..ea8963044 100644 --- a/tools/ui/src/lib/constants/settings-keys.ts +++ b/tools/ui/src/lib/constants/settings-keys.ts @@ -22,7 +22,6 @@ export const SETTINGS_KEYS = { // Display SHOW_MESSAGE_STATS: 'showMessageStats', SHOW_THOUGHT_IN_PROGRESS: 'showThoughtInProgress', - KEEP_STATS_VISIBLE: 'keepStatsVisible', AUTO_MIC_ON_EMPTY: 'autoMicOnEmpty', RENDER_USER_CONTENT_AS_MARKDOWN: 'renderUserContentAsMarkdown', DISABLE_AUTO_SCROLL: 'disableAutoScroll', @@ -61,7 +60,6 @@ export const SETTINGS_KEYS = { MCP_REQUEST_TIMEOUT_SECONDS: 'mcpRequestTimeoutSeconds', MCP_DEFAULT_SERVER_OVERRIDES: 'mcpDefaultServerOverrides', AGENTIC_MAX_TURNS: 'agenticMaxTurns', - ALWAYS_SHOW_AGENTIC_TURNS: 'alwaysShowAgenticTurns', AGENTIC_MAX_TOOL_PREVIEW_LINES: 'agenticMaxToolPreviewLines', SHOW_TOOL_CALL_IN_PROGRESS: 'showToolCallInProgress', // Performance diff --git a/tools/ui/src/lib/constants/settings-registry.ts b/tools/ui/src/lib/constants/settings-registry.ts index c25514c06..347a08dfe 100644 --- a/tools/ui/src/lib/constants/settings-registry.ts +++ b/tools/ui/src/lib/constants/settings-registry.ts @@ -258,18 +258,6 @@ const SETTINGS_REGISTRY: Record = { paramType: SyncableParameterType.BOOLEAN } }, - { - key: SETTINGS_KEYS.KEEP_STATS_VISIBLE, - label: 'Keep stats visible after generation', - help: 'Keep processing statistics visible after generation finishes.', - defaultValue: true, - type: SettingsFieldType.CHECKBOX, - section: SETTINGS_SECTION_SLUGS.DISPLAY, - sync: { - serverKey: SETTINGS_KEYS.KEEP_STATS_VISIBLE, - paramType: SyncableParameterType.BOOLEAN - } - }, { key: SETTINGS_KEYS.AUTO_MIC_ON_EMPTY, label: 'Show microphone on empty input', @@ -379,18 +367,6 @@ const SETTINGS_REGISTRY: Record = { paramType: SyncableParameterType.BOOLEAN } }, - { - key: SETTINGS_KEYS.ALWAYS_SHOW_AGENTIC_TURNS, - label: 'Always show agentic turns in conversation', - help: 'Always expand and display agentic loop turns in conversation messages.', - defaultValue: false, - type: SettingsFieldType.CHECKBOX, - section: SETTINGS_SECTION_SLUGS.DISPLAY, - sync: { - serverKey: SETTINGS_KEYS.ALWAYS_SHOW_AGENTIC_TURNS, - paramType: SyncableParameterType.BOOLEAN - } - }, { key: SETTINGS_KEYS.SHOW_BUILD_VERSION, label: 'Show build version information', diff --git a/tools/ui/src/lib/constants/storage.ts b/tools/ui/src/lib/constants/storage.ts index d0c5f2eff..eca9739ba 100644 --- a/tools/ui/src/lib/constants/storage.ts +++ b/tools/ui/src/lib/constants/storage.ts @@ -21,7 +21,6 @@ export const DISABLED_TOOLS_LOCALSTORAGE_KEY = `${STORAGE_APP_NAME}.disabledTool /** Disabled tools keyed by stable selection identity, no migration from the name based key */ export const DISABLED_TOOL_KEYS_LOCALSTORAGE_KEY = `${STORAGE_APP_NAME}.disabledToolKeys`; export const FAVORITE_MODELS_LOCALSTORAGE_KEY = `${STORAGE_APP_NAME}.favoriteModels`; -export const THINKING_ENABLED_DEFAULT_LOCALSTORAGE_KEY = `${STORAGE_APP_NAME}.thinkingEnabledDefault`; export const REASONING_EFFORT_DEFAULT_LOCALSTORAGE_KEY = `${STORAGE_APP_NAME}.reasoningEffortDefault`; /** Set when user has interacted with the MCP server recommendations dialog (checked servers, added custom server, or dismissed) */ export const MCP_SERVERS_ADDED_TO_CHAT_LOCALSTORAGE_KEY = `${STORAGE_APP_NAME}.mcpServersSetupDone`; diff --git a/tools/ui/src/lib/contexts/index.ts b/tools/ui/src/lib/contexts/index.ts index 01cd1d4b7..c6719fa9e 100644 --- a/tools/ui/src/lib/contexts/index.ts +++ b/tools/ui/src/lib/contexts/index.ts @@ -17,9 +17,3 @@ export { setChatSettingsConfigContext, type ChatSettingsConfigContext } from './chat-settings-config.context'; - -export { - getProcessingInfoContext, - setProcessingInfoContext, - type ProcessingInfoContext -} from './processing-info.context'; diff --git a/tools/ui/src/lib/contexts/processing-info.context.ts b/tools/ui/src/lib/contexts/processing-info.context.ts deleted file mode 100644 index 0cf43336f..000000000 --- a/tools/ui/src/lib/contexts/processing-info.context.ts +++ /dev/null @@ -1,16 +0,0 @@ -import { getContext, setContext } from 'svelte'; -import { CONTEXT_KEY_PROCESSING_INFO } from '$lib/constants'; - -export interface ProcessingInfoContext { - readonly showProcessingInfo: boolean; -} - -const PROCESSING_INFO_KEY = Symbol.for(CONTEXT_KEY_PROCESSING_INFO); - -export function setProcessingInfoContext(ctx: ProcessingInfoContext): ProcessingInfoContext { - return setContext(PROCESSING_INFO_KEY, ctx); -} - -export function getProcessingInfoContext(): ProcessingInfoContext { - return getContext(PROCESSING_INFO_KEY); -} diff --git a/tools/ui/src/lib/enums/chat.enums.ts b/tools/ui/src/lib/enums/chat.enums.ts index 278e4af84..f4994bb8e 100644 --- a/tools/ui/src/lib/enums/chat.enums.ts +++ b/tools/ui/src/lib/enums/chat.enums.ts @@ -5,6 +5,12 @@ export enum ChatMessageStatsView { SUMMARY = 'summary' } +export enum ChatMessageStatisticsMode { + SWITCHABLE = 'switchable', + READING = 'reading', + GENERATION = 'generation' +} + /** * Connection state of a streamed completion, drives the resume status indicator. */ diff --git a/tools/ui/src/lib/enums/index.ts b/tools/ui/src/lib/enums/index.ts index d8e47958b..847014e78 100644 --- a/tools/ui/src/lib/enums/index.ts +++ b/tools/ui/src/lib/enums/index.ts @@ -10,6 +10,7 @@ export { AgenticSectionType, ContinueIntentKind, ToolCallType } from './agentic. export { ChatMessageStatsView, + ChatMessageStatisticsMode, StreamConnectionState, ContentPartType, ConversationSelectionMode, diff --git a/tools/ui/src/lib/enums/reasoning-effort.enums.ts b/tools/ui/src/lib/enums/reasoning-effort.enums.ts index dadb0c726..172f118e7 100644 --- a/tools/ui/src/lib/enums/reasoning-effort.enums.ts +++ b/tools/ui/src/lib/enums/reasoning-effort.enums.ts @@ -3,6 +3,7 @@ * These values are sent to the server and mapped to token budgets. */ export enum ReasoningEffort { + OFF = 'off', LOW = 'low', MEDIUM = 'medium', HIGH = 'high', diff --git a/tools/ui/src/lib/hooks/use-context-gauge.svelte.ts b/tools/ui/src/lib/hooks/use-context-gauge.svelte.ts new file mode 100644 index 000000000..e11e2f7ab --- /dev/null +++ b/tools/ui/src/lib/hooks/use-context-gauge.svelte.ts @@ -0,0 +1,295 @@ +/** + * Reactive state for the context usage gauge: resolves the active model, + * fetches its cached props, parses live server stats, and exposes per-turn + * read / fresh / cache / output and cumulative token counts. + */ + +import { + modelsStore, + modelOptions, + selectedModelId, + singleModelName +} from '$lib/stores/models.svelte'; +import { chatStore } from '$lib/stores/chat.svelte'; +import { activeMessages } from '$lib/stores/conversations.svelte'; +import { isRouterMode } from '$lib/stores/server.svelte'; +import { MessageRole } from '$lib/enums'; +import { STATS_UNITS } from '$lib/constants'; +import type { ChatMessageTimings, DatabaseMessage } from '$lib/types'; +import { useProcessingState } from './use-processing-state.svelte'; +import { + colorLevelFromPercent, + type ColorLevel +} from '$lib/components/app/chat/ChatForm/ChatFormContextGauge/context-gauge'; + +interface LiveStats { + freshTokens: number; + promptTokens: number; + cacheTokens: number; + outputTokens: number; +} + +export interface UseContextGaugeReturn { + readonly activeModelId: string | null; + readonly isActiveModelLoaded: boolean; + readonly isActiveModelLoading: boolean; + readonly contextTotal: number | null; + readonly contextUsed: number; + readonly currentRead: number; + readonly currentFresh: number; + readonly currentCache: number; + readonly currentOutput: number; + readonly kvTotal: number; + readonly cumulativeRead: number; + readonly cumulativeOutput: number; + readonly cumulativeCacheTotal: number; + readonly averageTokensPerSecond: number | null; + readonly contextPercent: number | null; + readonly colorLevel: ColorLevel; + readonly transientDetails: string[]; + readonly hasAnyUsage: boolean; + loadModel(): Promise; + startMonitoring(): void; +} + +function lastAssistantTimings(messages: DatabaseMessage[]): ChatMessageTimings | undefined { + for (let i = messages.length - 1; i >= 0; i--) { + const m = messages[i]; + if (m.role === MessageRole.ASSISTANT && m.timings) return m.timings; + } + return undefined; +} + +function deriveLiveStats( + state: ReturnType['processingState'] +): LiveStats | null { + if (!state || (state.status !== 'preparing' && state.status !== 'generating')) { + return null; + } + const promptTokens = state.promptTokens ?? 0; + const cacheTokens = state.cacheTokens ?? 0; + return { + freshTokens: promptTokens, + promptTokens: promptTokens + cacheTokens, + cacheTokens, + outputTokens: state.outputTokensUsed ?? 0 + }; +} + +const TRANSIENT_DETAILS_EXCLUDED_PREFIXES = ['Context:', 'Output:']; + +function filterTransientDetails(raw: string[]): string[] { + return raw.filter((detail) => { + if (TRANSIENT_DETAILS_EXCLUDED_PREFIXES.some((prefix) => detail.startsWith(prefix))) { + return false; + } + return !detail.includes(STATS_UNITS.TOKENS_PER_SECOND); + }); +} + +export function useContextGauge(): UseContextGaugeReturn { + const processingState = useProcessingState(); + + // Resolve the model the gauge reports context for: explicit selection > + // last assistant model > single-model mode (mirrors useChatScreenActiveModel). + const activeModelId = $derived.by(() => { + if (!isRouterMode()) { + return singleModelName(); + } + + const selectedId = selectedModelId(); + if (selectedId) { + const model = modelOptions().find((m) => m.id === selectedId); + if (model) return model.model; + } + + return chatStore.getConversationModel(activeMessages() as DatabaseMessage[]); + }); + + const isActiveModelLoaded = $derived( + activeModelId !== null && modelsStore.isModelLoaded(activeModelId) + ); + + const isActiveModelLoading = $derived( + activeModelId !== null && modelsStore.isModelOperationInProgress(activeModelId) + ); + + // Pull /props on demand so n_ctx surfaces before the first chat request. + $effect(() => { + if (activeModelId && isActiveModelLoaded) { + const cached = modelsStore.getModelProps(activeModelId); + if (!cached) { + void modelsStore.fetchModelProps(activeModelId); + } + } + }); + + const contextTotal = $derived.by(() => { + void modelsStore.propsCacheVersion; + return activeModelId ? modelsStore.getModelContextSize(activeModelId) : null; + }); + + const liveStats = $derived(deriveLiveStats(processingState.processingState)); + + const currentRead = $derived.by(() => { + const timings = lastAssistantTimings(activeMessages() as DatabaseMessage[]); + let read = 0; + if (timings) { + read = (timings.prompt_n ?? 0) + (timings.cache_n ?? 0); + } + // live.promptTokens is already the combined reading (prompt + cache), + // so do not also add live.cacheTokens. + if (liveStats && liveStats.promptTokens > 0) { + read = Math.max(read, liveStats.promptTokens); + } + return read; + }); + + const currentFresh = $derived.by(() => { + const timings = lastAssistantTimings(activeMessages() as DatabaseMessage[]); + const fresh = timings?.prompt_n ?? 0; + return Math.max(fresh, liveStats?.freshTokens ?? 0); + }); + + const currentCache = $derived.by(() => { + const timings = lastAssistantTimings(activeMessages() as DatabaseMessage[]); + const cached = timings?.cache_n ?? 0; + if (liveStats && liveStats.promptTokens > 0) { + return Math.max(cached, liveStats.cacheTokens); + } + return cached; + }); + + const currentOutput = $derived.by(() => { + if (liveStats && liveStats.outputTokens > 0) return liveStats.outputTokens; + const timings = lastAssistantTimings(activeMessages() as DatabaseMessage[]); + return timings?.predicted_n ?? 0; + }); + + const kvTotal = $derived(currentRead + currentOutput); + const contextUsed = $derived(currentRead + currentOutput); + + const cumulative = $derived.by(() => { + const messages = activeMessages() as DatabaseMessage[]; + + // Agentic sessions stamp the same agentic.llm totals onto every + // assistant message; cache_n is never per-turn so cache_total stays 0. + const agenticMessages = messages.filter( + (m) => m.role === MessageRole.ASSISTANT && m.timings?.agentic?.llm?.predicted_n != null + ); + + if (agenticMessages.length > 0) { + const llm = agenticMessages[agenticMessages.length - 1].timings!.agentic!.llm; + const output = llm.predicted_n ?? 0; + const outputMs = llm.predicted_ms ?? 0; + const averageTokensPerSecond = outputMs > 0 && output > 0 ? (output / outputMs) * 1000 : null; + return { + read: llm.prompt_n ?? 0, + output, + cacheTotal: 0, + averageTokensPerSecond + }; + } + + let read = 0; + let output = 0; + let outputMs = 0; + let cacheTotal = 0; + for (const m of messages) { + if (m.role !== MessageRole.ASSISTANT || !m.timings) continue; + read += m.timings.prompt_n ?? 0; + cacheTotal += m.timings.cache_n ?? 0; + output += m.timings.predicted_n ?? 0; + outputMs += m.timings.predicted_ms ?? 0; + } + const averageTokensPerSecond = outputMs > 0 && output > 0 ? (output / outputMs) * 1000 : null; + return { read, output, cacheTotal, averageTokensPerSecond }; + }); + + const contextPercent = $derived.by(() => { + if (contextTotal === null || contextTotal <= 0) return null; + return Math.round((contextUsed / contextTotal) * 100); + }); + + const colorLevel = $derived(colorLevelFromPercent(contextPercent)); + + // Drop lines the surrounding Context / Output / speed rows already render. + const transientDetails = $derived(filterTransientDetails(processingState.getTechnicalDetails())); + + const hasAnyUsage = $derived( + cumulative.read > 0 || + cumulative.output > 0 || + currentRead > 0 || + currentOutput > 0 || + cumulative.averageTokensPerSecond !== null || + transientDetails.length > 0 + ); + + async function loadModel() { + if (!activeModelId || isActiveModelLoading) return; + try { + await modelsStore.loadModel(activeModelId); + } catch { + // toast already surfaced by modelsStore.loadModel + } + } + + return { + get activeModelId() { + return activeModelId; + }, + get isActiveModelLoaded() { + return isActiveModelLoaded; + }, + get isActiveModelLoading() { + return isActiveModelLoading; + }, + get contextTotal() { + return contextTotal; + }, + get contextUsed() { + return contextUsed; + }, + get currentRead() { + return currentRead; + }, + get currentFresh() { + return currentFresh; + }, + get currentCache() { + return currentCache; + }, + get currentOutput() { + return currentOutput; + }, + get kvTotal() { + return kvTotal; + }, + get cumulativeRead() { + return cumulative.read; + }, + get cumulativeOutput() { + return cumulative.output; + }, + get cumulativeCacheTotal() { + return cumulative.cacheTotal; + }, + get averageTokensPerSecond() { + return cumulative.averageTokensPerSecond; + }, + get contextPercent() { + return contextPercent; + }, + get colorLevel() { + return colorLevel; + }, + get transientDetails() { + return transientDetails; + }, + get hasAnyUsage() { + return hasAnyUsage; + }, + loadModel, + startMonitoring: () => processingState.startMonitoring() + }; +} diff --git a/tools/ui/src/lib/hooks/use-mcp-recommendations.svelte.ts b/tools/ui/src/lib/hooks/use-mcp-recommendations.svelte.ts index c8a85fa8f..4f4c2c782 100644 --- a/tools/ui/src/lib/hooks/use-mcp-recommendations.svelte.ts +++ b/tools/ui/src/lib/hooks/use-mcp-recommendations.svelte.ts @@ -54,11 +54,6 @@ export function useMcpRecommendations() { // effect, and we must not wipe the timeout that was just scheduled. if (checked) return; - if (mcpStore.optedInRecommendationIds.size > 0) { - checked = true; - return; - } - const hasRecommendations = mcpStore .getServers() .some((server) => RECOMMENDED_MCP_SERVER_IDS.has(server.id)); diff --git a/tools/ui/src/lib/hooks/use-processing-state.svelte.ts b/tools/ui/src/lib/hooks/use-processing-state.svelte.ts index f28031972..9fbda75d6 100644 --- a/tools/ui/src/lib/hooks/use-processing-state.svelte.ts +++ b/tools/ui/src/lib/hooks/use-processing-state.svelte.ts @@ -1,5 +1,4 @@ import { activeProcessingState } from '$lib/stores/chat.svelte'; -import { config } from '$lib/stores/settings.svelte'; import { STATS_UNITS } from '$lib/constants'; import type { ApiProcessingState, LiveProcessingStats, LiveGenerationStats } from '$lib/types'; @@ -46,7 +45,6 @@ export function useProcessingState(): UseProcessingStateReturn { return activeProcessingState(); }); - // Track last known state for keepStatsVisible functionality $effect(() => { if (processingState && isMonitoring) { lastKnownState = processingState; @@ -88,14 +86,8 @@ export function useProcessingState(): UseProcessingStateReturn { function stopMonitoring(): void { if (!isMonitoring) return; - isMonitoring = false; - // Only clear last known state if keepStatsVisible is disabled - const currentConfig = config(); - if (!currentConfig.keepStatsVisible) { - lastKnownState = null; - lastKnownProcessingStats = null; - } + isMonitoring = false; } function getProcessingMessage(): string { diff --git a/tools/ui/src/lib/services/chat.service.ts b/tools/ui/src/lib/services/chat.service.ts index 92828dd56..1a24374a2 100644 --- a/tools/ui/src/lib/services/chat.service.ts +++ b/tools/ui/src/lib/services/chat.service.ts @@ -340,6 +340,7 @@ export class ChatService { if (stream && conversationId) { headers['X-Conversation-Id'] = streamIdentity(conversationId, options.model); } + const response = await fetch(API_CHAT.COMPLETIONS, { method: 'POST', headers, @@ -1015,7 +1016,7 @@ export class ChatService { * * @param response - The fetch Response object containing the JSON data * @param onComplete - Optional callback invoked when response is successfully parsed - * @param onError - Optional callback invoked if an error occurs during parsing + * @param onError - Optional callback invoked if an error occurs while parsing * @returns {Promise} Promise that resolves to the generated content string * @throws {Error} if the response cannot be parsed or is malformed */ diff --git a/tools/ui/src/lib/services/migration.service.ts b/tools/ui/src/lib/services/migration.service.ts index d7709bc6b..981283be9 100644 --- a/tools/ui/src/lib/services/migration.service.ts +++ b/tools/ui/src/lib/services/migration.service.ts @@ -564,9 +564,9 @@ const configTypesMigration: Migration = { const config = JSON.parse(configRaw); let changed = false; - // Pre-schema configs persisted booleans as the strings "true"/"false", which the - // strict server schema now rejects. Coerce those back to real booleans. No config - // string field holds exactly "true"/"false", so the match is unambiguous. + // Pre-schema configs persisted booleans as "true"/"false" strings; the strict server + // schema rejects them. No config string field holds exactly "true"/"false", so the + // match is unambiguous. for (const key of Object.keys(config)) { if (config[key] === 'true') { config[key] = true; diff --git a/tools/ui/src/lib/stores/agentic.svelte.ts b/tools/ui/src/lib/stores/agentic.svelte.ts index 27491257a..1a677602f 100644 --- a/tools/ui/src/lib/stores/agentic.svelte.ts +++ b/tools/ui/src/lib/stores/agentic.svelte.ts @@ -477,7 +477,7 @@ class AgenticStore { conversationId: string; messages: ApiChatMessageData[]; options: AgenticFlowOptions; - tools: ReturnType; + tools: ReturnType; agenticConfig: AgenticConfig; callbacks: AgenticFlowCallbacks; signal?: AbortSignal; diff --git a/tools/ui/src/lib/stores/chat.svelte.ts b/tools/ui/src/lib/stores/chat.svelte.ts index 6949047b5..fcd07c4fd 100644 --- a/tools/ui/src/lib/stores/chat.svelte.ts +++ b/tools/ui/src/lib/stores/chat.svelte.ts @@ -60,6 +60,7 @@ import { ErrorDialogType, MessageRole, MessageType, + ReasoningEffort, StreamConnectionState } from '$lib/enums'; @@ -2334,7 +2335,8 @@ class ChatStore { if (currentConfig.excludeReasoningFromContext) apiOptions.excludeReasoningFromContext = true; apiOptions.enableThinking = conversationsStore.getThinkingEnabled(); - apiOptions.reasoningEffort = conversationsStore.getReasoningEffort(); + const effort = conversationsStore.getReasoningEffort(); + if (effort !== ReasoningEffort.OFF) apiOptions.reasoningEffort = effort; if (hasValue(currentConfig.temperature)) apiOptions.temperature = Number(currentConfig.temperature); diff --git a/tools/ui/src/lib/stores/conversations.svelte.ts b/tools/ui/src/lib/stores/conversations.svelte.ts index 486202207..b29a900fe 100644 --- a/tools/ui/src/lib/stores/conversations.svelte.ts +++ b/tools/ui/src/lib/stores/conversations.svelte.ts @@ -47,7 +47,6 @@ import { NON_ALPHANUMERIC_REGEX, MULTIPLE_UNDERSCORE_REGEX, SETTINGS_KEYS, - THINKING_ENABLED_DEFAULT_LOCALSTORAGE_KEY, REASONING_EFFORT_DEFAULT_LOCALSTORAGE_KEY } from '$lib/constants'; @@ -84,11 +83,16 @@ class ConversationsStore { /** Pending MCP server overrides for new conversations (before first message) */ pendingMcpServerOverrides = $state(ConversationsStore.loadMcpDefaults()); - /** Global (non-conversation-specific) thinking toggle default */ - pendingThinkingEnabled = $state(ConversationsStore.loadThinkingDefaults()); + /** Global (non-conversation-specific) thinking toggle default, derived from reasoning effort */ + pendingThinkingEnabled = $state(false); /** Global (non-conversation-specific) reasoning effort default */ - pendingReasoningEffort = $state(ConversationsStore.loadReasoningEffortDefault()); + pendingReasoningEffort = $state( + ConversationsStore.loadReasoningEffortDefault() + ); + + /** Last non-off reasoning effort, restored when re-enabling thinking globally */ + private lastNonOffEffort: ReasoningEffort | null = null; private static loadMcpDefaults(): McpServerOverride[] { const raw = config()[SETTINGS_KEYS.MCP_DEFAULT_SERVER_OVERRIDES]; @@ -112,35 +116,14 @@ class ConversationsStore { settingsStore.updateConfig(SETTINGS_KEYS.MCP_DEFAULT_SERVER_OVERRIDES, JSON.stringify(plain)); } - /** Load thinking-enabled default from localStorage */ - private static loadThinkingDefaults(): boolean { - if (typeof globalThis.localStorage === 'undefined') return true; - try { - const raw = localStorage.getItem(THINKING_ENABLED_DEFAULT_LOCALSTORAGE_KEY); - if (!raw) return true; - return raw === 'true'; - } catch { - return true; - } - } - - /** Persist thinking-enabled default to localStorage */ - private saveThinkingDefaults(): void { - if (typeof globalThis.localStorage === 'undefined') return; - localStorage.setItem( - THINKING_ENABLED_DEFAULT_LOCALSTORAGE_KEY, - this.pendingThinkingEnabled ? 'true' : 'false' - ); - } - /** Load reasoning effort default from localStorage */ - private static loadReasoningEffortDefault(): ReasoningEffort { - if (typeof globalThis.localStorage === 'undefined') return ReasoningEffort.MEDIUM; + private static loadReasoningEffortDefault(): ReasoningEffort | ReasoningEffort.OFF { + if (typeof globalThis.localStorage === 'undefined') return ReasoningEffort.OFF; try { const raw = localStorage.getItem(REASONING_EFFORT_DEFAULT_LOCALSTORAGE_KEY); - return (raw as ReasoningEffort) || ReasoningEffort.MEDIUM; + return (raw as ReasoningEffort | ReasoningEffort.OFF) || ReasoningEffort.OFF; } catch { - return ReasoningEffort.MEDIUM; + return ReasoningEffort.OFF; } } @@ -303,10 +286,17 @@ class ConversationsStore { this.pendingMcpServerOverrides = []; } - // Inherit global thinking default into the new conversation - conversation.thinkingEnabled = this.pendingThinkingEnabled; + // Inherit global thinking/reasoning defaults into the new conversation + const thinkingEnabled = this.getThinkingEnabled(); + conversation.thinkingEnabled = thinkingEnabled; + conversation.reasoningEffort = + this.pendingReasoningEffort === ReasoningEffort.OFF ? undefined : this.pendingReasoningEffort; await DatabaseService.updateConversation(conversation.id, { - thinkingEnabled: this.pendingThinkingEnabled + thinkingEnabled, + reasoningEffort: + this.pendingReasoningEffort === ReasoningEffort.OFF + ? undefined + : this.pendingReasoningEffort }); this.conversations = [conversation, ...this.conversations]; @@ -332,7 +322,6 @@ class ConversationsStore { } this.pendingMcpServerOverrides = []; - this.pendingThinkingEnabled = ConversationsStore.loadThinkingDefaults(); this.activeConversation = conversation; if (conversation.currNode) { @@ -363,7 +352,7 @@ class ConversationsStore { this.activeMessages = []; // reload defaults so new chats inherit persisted state this.pendingMcpServerOverrides = ConversationsStore.loadMcpDefaults(); - this.pendingThinkingEnabled = ConversationsStore.loadThinkingDefaults(); + this.pendingReasoningEffort = ConversationsStore.loadReasoningEffortDefault(); } /** @@ -794,9 +783,11 @@ class ConversationsStore { */ getThinkingEnabled(): boolean { if (this.activeConversation) { - return this.activeConversation.thinkingEnabled ?? this.pendingThinkingEnabled; + if (this.activeConversation.thinkingEnabled !== undefined) { + return this.activeConversation.thinkingEnabled; + } } - return this.pendingThinkingEnabled; + return this.getReasoningEffort() !== ReasoningEffort.OFF; } /** @@ -806,8 +797,17 @@ class ConversationsStore { */ async setThinkingEnabled(enabled: boolean): Promise { if (!this.activeConversation) { - this.pendingThinkingEnabled = enabled; - this.saveThinkingDefaults(); + if (enabled) { + const effort = this.lastNonOffEffort ?? ReasoningEffort.LOW; + this.pendingReasoningEffort = effort; + this.saveReasoningEffortDefaults(); + } else { + if (this.pendingReasoningEffort !== ReasoningEffort.OFF) { + this.lastNonOffEffort = this.pendingReasoningEffort; + } + this.pendingReasoningEffort = ReasoningEffort.OFF; + this.saveReasoningEffortDefaults(); + } return; } @@ -831,7 +831,7 @@ class ConversationsStore { * Gets the effective reasoning effort for the active conversation. * Returns the conversation override if set, otherwise the global default. */ - getReasoningEffort(): ReasoningEffort { + getReasoningEffort(): ReasoningEffort | ReasoningEffort.OFF { if (this.activeConversation) { return this.activeConversation.reasoningEffort ?? this.pendingReasoningEffort; } diff --git a/tools/ui/src/lib/stores/mcp.svelte.ts b/tools/ui/src/lib/stores/mcp.svelte.ts index 37eb56332..a53463ee9 100644 --- a/tools/ui/src/lib/stores/mcp.svelte.ts +++ b/tools/ui/src/lib/stores/mcp.svelte.ts @@ -12,21 +12,22 @@ * - Lifecycle management (initialize, shutdown) * - Multi-server coordination * - Tool name conflict detection and resolution - * - OpenAI-compatible tool definition generation * - Automatic tool-to-server routing * - Health checks * + * MCP connection state and raw `Tool[]` per server are owned here; the + * OpenAI-compatible wire format for those tools is built in `toolsStore` + * (see {@link toolsStore.mcpEntries} / {@link toolsStore.getEnabledToolsForLLM}). + * * @see MCPService in services/mcp.service.ts for protocol operations */ import { browser } from '$app/environment'; -import { SvelteSet } from 'svelte/reactivity'; import { SETTINGS_KEYS } from '$lib/constants'; import { MCPService } from '$lib/services/mcp.service'; import { config, settingsStore } from '$lib/stores/settings.svelte'; import { mcpResourceStore } from '$lib/stores/mcp-resources.svelte'; import { serverStore } from '$lib/stores/server.svelte'; -import { conversationsStore } from '$lib/stores/conversations.svelte'; import { mode } from 'mode-watcher'; import { parseMcpServerSettings, @@ -40,9 +41,7 @@ import { HealthCheckStatus, MCPRefType, ColorMode, - UrlProtocol, - JsonSchemaType, - ToolCallType + UrlProtocol } from '$lib/enums'; import { DEFAULT_CACHE_TTL_MS, @@ -53,12 +52,10 @@ import { MCP_RECONNECT_BACKOFF_MULTIPLIER, MCP_RECONNECT_INITIAL_DELAY, MCP_RECONNECT_MAX_DELAY, - MCP_RECONNECT_ATTEMPT_TIMEOUT_MS, - RECOMMENDED_MCP_SERVER_IDS + MCP_RECONNECT_ATTEMPT_TIMEOUT_MS } from '$lib/constants'; import type { MCPToolCall, - OpenAIToolDefinition, ServerStatus, ToolExecutionResult, MCPClientConfig, @@ -582,30 +579,10 @@ class MCPStore { } /** - * Recommended MCP server IDs the user opted in to via per-chat overrides. - * Single source of truth for "which recommendations has the user accepted", - * shared by the recommendations hook and the visible-servers getter. - */ - get optedInRecommendationIds(): ReadonlySet { - const ids = new SvelteSet(); - for (const override of conversationsStore.pendingMcpServerOverrides) { - if (RECOMMENDED_MCP_SERVER_IDS.has(override.serverId) && override.enabled) { - ids.add(override.serverId); - } - } - return ids; - } - - /** - * MCP servers selectable in chat-add UIs and the settings page: - * enabled in settings and either non-recommended or explicitly opted in. + * MCP servers selectable in chat-add UIs and the settings page. */ get visibleMcpServers(): MCPServerSettingsEntry[] { - const optedIn = this.optedInRecommendationIds; - return this.getServersSorted().filter( - (server) => - server.enabled && (!RECOMMENDED_MCP_SERVER_IDS.has(server.id) || optedIn.has(server.id)) - ); + return this.getServersSorted().filter((server) => server.enabled); } async ensureInitialized(perChatOverrides?: McpServerOverride[]): Promise { @@ -979,73 +956,6 @@ class MCPStore { } } - getToolDefinitionsForLLM(): OpenAIToolDefinition[] { - const tools: OpenAIToolDefinition[] = []; - - for (const connection of this.connections.values()) { - for (const tool of connection.tools) { - const rawSchema = (tool.inputSchema as Record) ?? { - type: JsonSchemaType.OBJECT, - properties: {}, - required: [] - }; - - tools.push({ - type: ToolCallType.FUNCTION as const, - function: { - name: tool.name, - description: tool.description, - parameters: this.normalizeSchemaProperties(rawSchema) - } - }); - } - } - - return tools; - } - - private normalizeSchemaProperties(schema: Record): Record { - if (!schema || typeof schema !== 'object') { - return schema; - } - - const normalized = { ...schema }; - if (normalized.properties && typeof normalized.properties === 'object') { - const props = normalized.properties as Record>; - const normalizedProps: Record> = {}; - for (const [key, prop] of Object.entries(props)) { - if (!prop || typeof prop !== 'object') { - normalizedProps[key] = prop; - continue; - } - const normalizedProp = { ...prop }; - if (!normalizedProp.type && normalizedProp.default !== undefined) { - const defaultVal = normalizedProp.default; - if (typeof defaultVal === 'string') normalizedProp.type = 'string'; - else if (typeof defaultVal === 'number') - normalizedProp.type = Number.isInteger(defaultVal) ? 'integer' : 'number'; - else if (typeof defaultVal === 'boolean') normalizedProp.type = 'boolean'; - else if (Array.isArray(defaultVal)) normalizedProp.type = 'array'; - else if (typeof defaultVal === 'object' && defaultVal !== null) - normalizedProp.type = 'object'; - } - if (normalizedProp.properties) - Object.assign( - normalizedProp, - this.normalizeSchemaProperties(normalizedProp as Record) - ); - if (normalizedProp.items && typeof normalizedProp.items === 'object') - normalizedProp.items = this.normalizeSchemaProperties( - normalizedProp.items as Record - ); - normalizedProps[key] = normalizedProp; - } - normalized.properties = normalizedProps; - } - - return normalized; - } - getToolNames(): string[] { return Array.from(this.toolsIndex.keys()); } diff --git a/tools/ui/src/lib/stores/tools.svelte.ts b/tools/ui/src/lib/stores/tools.svelte.ts index a63781988..dcaab5f42 100644 --- a/tools/ui/src/lib/stores/tools.svelte.ts +++ b/tools/ui/src/lib/stores/tools.svelte.ts @@ -13,33 +13,6 @@ import { import { SvelteMap, SvelteSet } from 'svelte/reactivity'; /** Stable selection identity for a tool, shared by the disabled set and the permission store */ -function toolKey(source: ToolSource, name: string, serverId?: string): string { - switch (source) { - case ToolSource.MCP: - return serverId ? `mcp-${serverId}:${name}` : `mcp:${name}`; - case ToolSource.CUSTOM: - return `custom:${name}`; - case ToolSource.FRONTEND: - return `frontend:${name}`; - default: - return `builtin:${name}`; - } -} - -function mcpDefinition( - name: string, - description: string | undefined, - schema?: Record -): OpenAIToolDefinition { - return { - type: ToolCallType.FUNCTION, - function: { - name, - description, - parameters: schema ?? { type: JsonSchemaType.OBJECT, properties: {}, required: [] } - } - }; -} class ToolsStore { private _builtinTools = $state([]); @@ -77,12 +50,96 @@ class ToolsStore { } } + private toolKey(source: ToolSource, name: string, serverId?: string): string { + switch (source) { + case ToolSource.MCP: + return serverId ? `mcp-${serverId}:${name}` : `mcp:${name}`; + case ToolSource.CUSTOM: + return `custom:${name}`; + case ToolSource.FRONTEND: + return `frontend:${name}`; + default: + return `builtin:${name}`; + } + } + + private inferTypeFromDefault(value: unknown): string | undefined { + if (typeof value === 'string') return 'string'; + if (typeof value === 'boolean') return 'boolean'; + if (typeof value === 'number') return Number.isInteger(value) ? 'integer' : 'number'; + if (Array.isArray(value)) return 'array'; + if (value !== null && typeof value === 'object') return 'object'; + return undefined; + } + + /** + * Recursively normalize a JSON Schema object: infers `type` from `default` + * for properties / items that omit it, and descends into nested `properties` + * and `items`. Returns a new object -- does not mutate the input. + */ + private normalizeJsonSchema(schema: Record): Record { + if (!schema || typeof schema !== 'object') return schema; + + const normalized: Record = { ...schema }; + + if (normalized.properties && typeof normalized.properties === 'object') { + const props = normalized.properties as Record>; + const normalizedProps: Record> = {}; + for (const [key, prop] of Object.entries(props)) { + if (!prop || typeof prop !== 'object') { + normalizedProps[key] = prop; + continue; + } + + const normalizedProp: Record = { ...prop }; + + if (!normalizedProp.type && normalizedProp.default !== undefined) { + const inferred = this.inferTypeFromDefault(normalizedProp.default); + if (inferred) normalizedProp.type = inferred; + } + + if (normalizedProp.properties) { + Object.assign( + normalizedProp, + this.normalizeJsonSchema(normalizedProp as Record) + ); + } + + if (normalizedProp.items && typeof normalizedProp.items === 'object') { + normalizedProp.items = this.normalizeJsonSchema( + normalizedProp.items as Record + ); + } + + normalizedProps[key] = normalizedProp; + } + normalized.properties = normalizedProps; + } + + return normalized; + } + + private mcpDefinition( + name: string, + description: string | undefined, + schema?: Record + ): OpenAIToolDefinition { + return { + type: ToolCallType.FUNCTION, + function: { + name, + description, + parameters: schema ?? { type: JsonSchemaType.OBJECT, properties: {}, required: [] } + } + }; + } + get builtinTools(): OpenAIToolDefinition[] { return this._builtinTools; } get mcpTools(): OpenAIToolDefinition[] { - return mcpStore.getToolDefinitionsForLLM(); + return this.mcpEntries().map((e) => e.definition); } get frontendTools(): OpenAIToolDefinition[] { @@ -124,11 +181,22 @@ class ToolsStore { for (const [serverId, connection] of connections) { const serverName = mcpStore.getServerDisplayName(serverId); for (const tool of connection.tools) { - const schema = (tool.inputSchema as Record) ?? undefined; + const rawSchema = (tool.inputSchema as Record) ?? { + type: JsonSchemaType.OBJECT, + properties: {}, + required: [] + }; out.push({ serverId, serverName, - definition: mcpDefinition(tool.name, tool.description, schema) + definition: { + type: ToolCallType.FUNCTION, + function: { + name: tool.name, + description: tool.description, + parameters: this.normalizeJsonSchema(rawSchema) + } + } }); } } @@ -138,7 +206,7 @@ class ToolsStore { out.push({ serverId, serverName, - definition: mcpDefinition(tool.name, tool.description) + definition: this.mcpDefinition(tool.name, tool.description) }); } } @@ -160,14 +228,18 @@ class ToolsStore { for (const def of this._builtinTools) { const name = def.function.name; - push({ source: ToolSource.BUILTIN, key: toolKey(ToolSource.BUILTIN, name), definition: def }); + push({ + source: ToolSource.BUILTIN, + key: this.toolKey(ToolSource.BUILTIN, name), + definition: def + }); } for (const def of this.frontendTools) { const name = def.function.name; push({ source: ToolSource.FRONTEND, - key: toolKey(ToolSource.FRONTEND, name), + key: this.toolKey(ToolSource.FRONTEND, name), definition: def }); } @@ -178,14 +250,18 @@ class ToolsStore { source: ToolSource.MCP, serverId, serverName, - key: toolKey(ToolSource.MCP, name, serverId), + key: this.toolKey(ToolSource.MCP, name, serverId), definition }); } for (const def of this.customTools) { const name = def.function.name; - push({ source: ToolSource.CUSTOM, key: toolKey(ToolSource.CUSTOM, name), definition: def }); + push({ + source: ToolSource.CUSTOM, + key: this.toolKey(ToolSource.CUSTOM, name), + definition: def + }); } return entries; @@ -233,7 +309,8 @@ class ToolsStore { /** * Enabled tool definitions for sending to the LLM. - * MCP tools keep their normalized schemas from mcpStore. + * MCP tool schemas are normalized here so the wire payload is consistent + * across all four sources (built-in, frontend/sandbox, MCP, custom JSON). * The API identifies tools by name, so a name is sent at most once. */ getEnabledToolsForLLM(): OpenAIToolDefinition[] { @@ -256,7 +333,8 @@ class ToolsStore { for (const def of this._builtinTools) take(def); for (const def of this.frontendTools) take(def); - for (const def of mcpStore.getToolDefinitionsForLLM()) take(def); + // mcpEntries() over mcpStore directly so wire shape stays normalized and aligned with the tools UI. + for (const entry of this.mcpEntries()) take(entry.definition); for (const def of this.customTools) take(def); return result; @@ -308,15 +386,17 @@ class ToolsStore { const connection = mcpStore.getConnections().get(serverId); if (!connection) return; for (const tool of connection.tools) { - this._disabledTools.delete(toolKey(ToolSource.MCP, tool.name, serverId)); + this._disabledTools.delete(this.toolKey(ToolSource.MCP, tool.name, serverId)); } this.persistDisabledTools(); } toggleGroup(group: ToolGroup): void { const allEnabled = group.tools.every((t) => this.isToolEnabled(t.key)); + const target = !allEnabled; for (const tool of group.tools) { - this.setToolEnabled(tool.key, !allEnabled); + if (target) this._disabledTools.delete(tool.key); + else this._disabledTools.add(tool.key); } this.persistDisabledTools(); } @@ -332,7 +412,7 @@ class ToolsStore { tools: { name: string; description?: string }[]; }[] { const result: ReturnType = []; - for (const server of mcpStore.getServersSorted().filter((s) => s.enabled)) { + for (const server of mcpStore.visibleMcpServers) { const health = mcpStore.getHealthCheckState(server.id); if (health.status === HealthCheckStatus.SUCCESS && health.tools.length > 0) { result.push({ diff --git a/tools/ui/src/lib/utils/branching.ts b/tools/ui/src/lib/utils/branching.ts index c40abbdd6..6ff701318 100644 --- a/tools/ui/src/lib/utils/branching.ts +++ b/tools/ui/src/lib/utils/branching.ts @@ -111,15 +111,7 @@ function findLeafNodeInMap( } /** - * Convenience wrapper around {@link findLeafNodeInMap} for callers that only have - * a flat message array. - * - * Finds the leaf node (message with no children) for a given message branch. - * Traverses down the tree following the last child until reaching a leaf. - * - * @param messages - All messages in the conversation - * @param messageId - Starting message ID to find leaf for - * @returns The leaf node ID, or the original messageId if no children + * Convenience wrapper around {@link findLeafNodeInMap} for callers that have a flat message array. */ export function findLeafNode(messages: readonly DatabaseMessage[], messageId: string): string { const nodeMap = new Map(messages.map((msg) => [msg.id, msg] as const)); @@ -225,7 +217,6 @@ export function getMessageSiblings( /** * Builds sibling information for every message in a conversation. - * A single node map is shared across all lookups for O(1) access. * * @param messages - All messages in the conversation * @returns Map of message ID to its sibling information