From 2145525a4081d66ff1a87cf43ef809f95a85ac0c Mon Sep 17 00:00:00 2001 From: Gaurav Garg Date: Sat, 26 Sep 2026 07:59:15 -0700 Subject: [PATCH] Revert "Change max context length for auto-fitting with unified KV (#28849)" (#29437) This reverts commit b04d4e567cd2fb8d2ded6e17d38dbbcfafe29063. --- common/fit.cpp | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/common/fit.cpp b/common/fit.cpp index faa595f84f..7a03008295 100644 --- a/common/fit.cpp +++ b/common/fit.cpp @@ -192,9 +192,9 @@ static void common_params_fit_impl( uint32_t hp_nct = 0; // hparams.n_ctx_train uint32_t hp_nex = 0; // hparams.n_expert - // size the context for all sequences, but keep minimums and alignment per KV stream - const uint32_t n_seq_max = std::max(1, cparams->n_seq_max); - const uint32_t n_streams = cparams->kv_unified ? 1 : n_seq_max; + // with non-unified kv, we need to take into account n_streams + // for example, if memory can hold more than model's trained context size, we must extend the n_ctx to hold enough n_streams + const uint32_t n_streams = cparams->kv_unified ? 1 : std::max(1, cparams->n_seq_max); const bool n_ctx_auto = cparams->n_ctx == 0; dmds_t dmds_extra; // memory of the extra model, laid out on the devices of the main model @@ -264,15 +264,15 @@ static void common_params_fit_impl( dmds_t dmds_full = common_get_device_memory_data_impl(path_model, mparams, cparams, devs, hp_ngl, hp_nct, hp_nex, log_level); // saturate instead of overflowing, this also preserves the UINT32_MAX sentinel of n_ctx_min: - const uint32_t n_ctx_max = (uint32_t) std::min(uint64_t(hp_nct) * n_seq_max, UINT32_MAX); + const uint32_t n_ctx_max = (uint32_t) std::min(uint64_t(hp_nct) * n_streams, UINT32_MAX); const uint32_t n_ctx_min_total = (uint32_t) std::min(uint64_t(n_ctx_min) * n_streams, UINT32_MAX); // llama_context would use only hp_nct in total for n_ctx == 0, resolve the context before measuring anything else: if (n_ctx_auto) { cparams->n_ctx = n_ctx_max; - if (n_seq_max > 1) { - LOG_TRC("%s: context size unset -> using %" PRIu32 " for %" PRIu32 " sequences:\n", - __func__, n_ctx_max, n_seq_max); + if (n_streams > 1) { + LOG_TRC("%s: context size unset and KV cache not unified -> using %" PRIu32 " for %" PRIu32 " sequences:\n", + __func__, n_ctx_max, n_streams); dmds_full = common_get_device_memory_data_impl(path_model, mparams, cparams, devs, hp_ngl, hp_nct, hp_nex, log_level); } }