mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-01 10:27:30 -05:00
fit: take an optional second model into account
Illustrates the alternative discussed on the draft context fix. The memory of a draft or MTP context is currently handed to the fit as a fixed byte margin, which cannot express a memory that grows with the context the fit is still deciding on. common_fit_params now takes an optional second model that shares the devices of the main one. Its context follows the main context and its memory is measured again whenever that context changes, so the reduce path stays exact instead of conservative. A model that cannot be measured on its own, such as a shared cell MTP context, is skipped with a warning and the main model is fitted alone. This drops the reservation block in the server, which no longer has to probe the trained context size of the target to guess an upper bound.
This commit is contained in:
@@ -1040,74 +1040,7 @@ private:
|
||||
}
|
||||
}
|
||||
|
||||
// optionally reserve VRAM for the draft / MTP context before fitting the target model
|
||||
if (params_base.fit_params) {
|
||||
if (has_spec) {
|
||||
// MTP draft context lives on the target model, only context+compute are new
|
||||
bool measure_model_bytes = has_draft;
|
||||
|
||||
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.n_rs_seq = 0;
|
||||
|
||||
std::vector<ggml_backend_dev_t> devs;
|
||||
uint32_t hp_ngl = 0;
|
||||
uint32_t hp_nct = 0;
|
||||
uint32_t hp_nex = 0;
|
||||
try {
|
||||
// the draft context follows the target context, measure it at the largest context the target can take
|
||||
if (cparams_dft.n_ctx == 0) {
|
||||
auto mparams_tgt = common_model_params_to_llama(params_base);
|
||||
auto cparams_tgt = common_context_params_to_llama(params_base);
|
||||
|
||||
common_get_device_memory_data(
|
||||
params_base.model.path.c_str(), &mparams_tgt, &cparams_tgt,
|
||||
devs, hp_ngl, hp_nct, hp_nex, GGML_LOG_LEVEL_ERROR);
|
||||
|
||||
cparams_dft.n_ctx = hp_nct * (params_base.kv_unified ? 1 : params_base.n_parallel);
|
||||
}
|
||||
|
||||
auto dmd = common_get_device_memory_data(
|
||||
params_dft.model.path.c_str(), &mparams_dft, &cparams_dft,
|
||||
devs, hp_ngl, hp_nct, hp_nex, GGML_LOG_LEVEL_ERROR);
|
||||
|
||||
GGML_ASSERT(!params_base.fit_params_target.empty());
|
||||
size_t total = 0;
|
||||
|
||||
std::vector<ggml_backend_dev_t> tgt_devices = params.devices;
|
||||
|
||||
if (tgt_devices.empty()) {
|
||||
for(size_t i = 0; i < ggml_backend_dev_count(); ++i) {
|
||||
tgt_devices.push_back(ggml_backend_dev_get(i));
|
||||
}
|
||||
}
|
||||
|
||||
for (size_t j = 0; j < devs.size(); ++j) {
|
||||
const size_t bytes = (measure_model_bytes ? dmd[j].model : 0) + dmd[j].context + dmd[j].compute;
|
||||
total += bytes;
|
||||
for (size_t i = 0; i < tgt_devices.size(); i++) {
|
||||
if (tgt_devices[i] == devs[j]) {
|
||||
SRV_DBG("[spec] adding %.2f MiB to fit_params_target for device %s\n",
|
||||
bytes / (1024.0 * 1024.0), ggml_backend_dev_name(devs[j]));
|
||||
params_base.fit_params_target[i] += bytes;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
SRV_TRC("[spec] estimated memory usage of %s is %.2f MiB\n",
|
||||
has_draft ? "draft model" : "MTP context",
|
||||
total / (1024.0 * 1024.0));
|
||||
} catch (const std::exception & e) {
|
||||
SRV_WRN("[spec] failed to measure %s memory: %s\n",
|
||||
has_draft ? "draft model" : "MTP context", e.what());
|
||||
}
|
||||
}
|
||||
}
|
||||
// note: the draft / MTP context is fitted together with the target model, see common_fit_extra_model
|
||||
|
||||
// attach a progress callback
|
||||
{
|
||||
|
||||
Reference in New Issue
Block a user