mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-08-04 02:38:02 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0ef6e55edb | ||
|
|
94bc47f280 | ||
|
|
fe2adf0e72 | ||
|
|
57c092139a | ||
|
|
ee0445c99c | ||
|
|
99111b19ce | ||
|
|
e8e06f78e2 | ||
|
|
dbadb68eec | ||
|
|
39eab74a05 | ||
|
|
c50b34a1e0 | ||
|
|
67d5978bb1 | ||
|
|
563dec81c1 | ||
|
|
96278e39fc |
+20
-5
@@ -27,6 +27,7 @@
|
||||
#include <algorithm>
|
||||
#include <cinttypes>
|
||||
#include <climits>
|
||||
#include <cmath>
|
||||
#include <cstdarg>
|
||||
#include <filesystem>
|
||||
#include <fstream>
|
||||
@@ -2036,7 +2037,13 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
|
||||
{"--repeat-penalty"}, "N",
|
||||
string_format("penalize repeat sequence of tokens (default: %.2f, 1.0 = disabled)", (double)params.sampling.penalty_repeat),
|
||||
[](common_params & params, const std::string & value) {
|
||||
params.sampling.penalty_repeat = std::stof(value);
|
||||
const float penalty_repeat = std::stof(value);
|
||||
if (!std::isfinite(penalty_repeat) ||
|
||||
penalty_repeat <= 0.0f ||
|
||||
!std::isfinite(1.0f/penalty_repeat)) {
|
||||
throw std::runtime_error("error: repeat-penalty must be finite and greater than 0\n");
|
||||
}
|
||||
params.sampling.penalty_repeat = penalty_repeat;
|
||||
params.sampling.user_sampling_config |= common_params_sampling_config::COMMON_PARAMS_SAMPLING_CONFIG_PENALTY_REPEAT;
|
||||
}
|
||||
).set_sampling());
|
||||
@@ -2044,14 +2051,22 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
|
||||
{"--presence-penalty"}, "N",
|
||||
string_format("repeat alpha presence penalty (default: %.2f, 0.0 = disabled)", (double)params.sampling.penalty_present),
|
||||
[](common_params & params, const std::string & value) {
|
||||
params.sampling.penalty_present = std::stof(value);
|
||||
const float penalty_present = std::stof(value);
|
||||
if (!std::isfinite(penalty_present)) {
|
||||
throw std::runtime_error("error: presence-penalty must be finite\n");
|
||||
}
|
||||
params.sampling.penalty_present = penalty_present;
|
||||
}
|
||||
).set_sampling());
|
||||
add_opt(common_arg(
|
||||
{"--frequency-penalty"}, "N",
|
||||
string_format("repeat alpha frequency penalty (default: %.2f, 0.0 = disabled)", (double)params.sampling.penalty_freq),
|
||||
[](common_params & params, const std::string & value) {
|
||||
params.sampling.penalty_freq = std::stof(value);
|
||||
const float penalty_freq = std::stof(value);
|
||||
if (!std::isfinite(penalty_freq)) {
|
||||
throw std::runtime_error("error: frequency-penalty must be finite\n");
|
||||
}
|
||||
params.sampling.penalty_freq = penalty_freq;
|
||||
}
|
||||
).set_sampling());
|
||||
add_opt(common_arg(
|
||||
@@ -2567,7 +2582,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
|
||||
params.mtmd_batch_max_tokens = value;
|
||||
}
|
||||
).set_examples({LLAMA_EXAMPLE_SERVER}).set_env("LLAMA_ARG_MTMD_BATCH_MAX_TOKENS"));
|
||||
if (llama_supports_rpc()) {
|
||||
if (params.is_gen_docs || llama_supports_rpc()) {
|
||||
add_opt(common_arg(
|
||||
{"--rpc"}, "SERVERS",
|
||||
"comma-separated list of RPC servers (host:port)",
|
||||
@@ -3316,7 +3331,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
|
||||
{"--tools"}, "TOOL1,TOOL2,...",
|
||||
"experimental: whether to enable built-in tools for AI agents - do not enable in untrusted environments (default: no tools)\n"
|
||||
"specify \"all\" to enable all tools\n"
|
||||
"available tools: read_file, file_glob_search, grep_search, exec_shell_command, write_file, edit_file, get_datetime\n"
|
||||
"available tools: read_file, file_glob_search, grep_search, exec_shell_command, write_file, edit_file, get_datetime, get_info\n"
|
||||
"note: for security reasons, this will limit --cors-origins to localhost by default",
|
||||
[](common_params & params, const std::string & value) {
|
||||
params.server_tools = parse_csv_row(value);
|
||||
|
||||
+26
-7
@@ -2114,6 +2114,11 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
|
||||
std::optional<json> additional_context;
|
||||
if (is_v4 && has_response_format) {
|
||||
additional_context = json{ { "response_format", inputs.json_schema } };
|
||||
}
|
||||
|
||||
const std::string DSML = "|DSML|";
|
||||
const std::string THINK_START = "<think>";
|
||||
const std::string THINK_END = "</think>";
|
||||
@@ -2125,9 +2130,12 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
||||
const std::string PARAM_START = "<" + DSML + "parameter";
|
||||
const std::string PARAM_END = "</" + DSML + "parameter>";
|
||||
const std::string GEN_PROMPT = "<|Assistant|>";
|
||||
const std::string TC_SEPARATOR = "\n\n";
|
||||
|
||||
data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs, adjusted_messages);
|
||||
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs, adjusted_messages);
|
||||
data.prompt = common_chat_template_direct_apply_impl(
|
||||
tmpl, inputs, adjusted_messages, std::nullopt, additional_context);
|
||||
data.generation_prompt = common_chat_template_generation_prompt_impl(
|
||||
tmpl, inputs, adjusted_messages, std::nullopt, additional_context);
|
||||
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
|
||||
data.supports_thinking = true;
|
||||
data.thinking_start_tag = THINK_START;
|
||||
@@ -2141,9 +2149,16 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
||||
if (inputs.has_continuation()) {
|
||||
const auto & msg = inputs.continue_msg;
|
||||
|
||||
data.generation_prompt = GEN_PROMPT + THINK_START + msg.reasoning_content;
|
||||
if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) {
|
||||
data.generation_prompt += THINK_END + msg.render_content();
|
||||
if (is_v4 && msg.reasoning_content.empty()) {
|
||||
data.generation_prompt = GEN_PROMPT + THINK_END;
|
||||
if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) {
|
||||
data.generation_prompt += msg.render_content();
|
||||
}
|
||||
} else {
|
||||
data.generation_prompt = GEN_PROMPT + THINK_START + msg.reasoning_content;
|
||||
if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) {
|
||||
data.generation_prompt += THINK_END + msg.render_content();
|
||||
}
|
||||
}
|
||||
|
||||
data.prompt += data.generation_prompt;
|
||||
@@ -2242,7 +2257,9 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
||||
|
||||
if (extract_reasoning && inputs.enable_thinking) {
|
||||
reasoning = p.optional(THINK_START + p.reasoning(p.until(THINK_END)) + THINK_END);
|
||||
reasoning_with_tc = THINK_START + p.reasoning(p.until_one_of({ FC_START, THINK_END })) + obligatory_tool_calls;
|
||||
reasoning_with_tc = THINK_START +
|
||||
p.reasoning(p.until_one_of({ TC_SEPARATOR + FC_START, FC_START, THINK_END })) +
|
||||
p.space() + obligatory_tool_calls;
|
||||
allow_reasoning_with_tc = true;
|
||||
} else if (extract_reasoning) {
|
||||
// Thinking disabled but reasoning extraction requested: the generation prompt
|
||||
@@ -2265,7 +2282,9 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha
|
||||
return generation_prompt + reasoning + p.content(p.rest()) + end;
|
||||
}
|
||||
|
||||
auto content_before_tools = p.negate(p.literal(THINK_START)) + p.content(p.until(FC_START));
|
||||
auto content_before_tools = p.negate(p.literal(THINK_START)) +
|
||||
p.content(p.until_one_of({ TC_SEPARATOR + FC_START, FC_START })) +
|
||||
p.space();
|
||||
return allow_reasoning_with_tc ? generation_prompt + (reasoning_with_tc | (reasoning + content_before_tools + tool_calls)) + end :
|
||||
generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||
});
|
||||
|
||||
+30
-12
@@ -998,6 +998,23 @@ bool fs_is_directory(const std::string & path) {
|
||||
return std::filesystem::exists(dir) && std::filesystem::is_directory(dir);
|
||||
}
|
||||
|
||||
std::string common_get_env(const std::string & name) {
|
||||
const char * value = std::getenv(name.c_str());
|
||||
return value == nullptr ? "" : value;
|
||||
}
|
||||
|
||||
void common_set_env(const std::string & name, const std::string & value) {
|
||||
#if defined(_WIN32)
|
||||
_putenv_s(name.c_str(), value.c_str());
|
||||
#else
|
||||
if (value.empty()) {
|
||||
unsetenv(name.c_str());
|
||||
} else {
|
||||
setenv(name.c_str(), value.c_str(), 1);
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
std::string fs_get_cache_directory() {
|
||||
std::string cache_directory = "";
|
||||
auto ensure_trailing_slash = [](std::string p) {
|
||||
@@ -1299,8 +1316,9 @@ common_init_result::common_init_result(common_params & params, bool model_only)
|
||||
pimpl->samplers.resize(cparams.n_seq_max);
|
||||
pimpl->samplers_seq_config.resize(cparams.n_seq_max);
|
||||
|
||||
const int32_t n_ctx = cparams.n_ctx > 0 ? (int32_t) cparams.n_ctx : llama_model_n_ctx_train(model);
|
||||
for (int i = 0; i < (int) cparams.n_seq_max; ++i) {
|
||||
pimpl->samplers[i].reset(common_sampler_init(model, params.sampling));
|
||||
pimpl->samplers[i].reset(common_sampler_init(model, params.sampling, n_ctx));
|
||||
pimpl->samplers_seq_config[i] = { i, common_sampler_get(pimpl->samplers[i].get()) };
|
||||
}
|
||||
|
||||
@@ -1462,18 +1480,18 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode
|
||||
common_init_result::~common_init_result() = default;
|
||||
|
||||
std::string common_get_model_endpoint() {
|
||||
const char * model_endpoint_env = getenv("MODEL_ENDPOINT");
|
||||
// We still respect the use of environment-variable "HF_ENDPOINT" for backward-compatibility.
|
||||
const char * hf_endpoint_env = getenv("HF_ENDPOINT");
|
||||
const char * endpoint_env = model_endpoint_env ? model_endpoint_env : hf_endpoint_env;
|
||||
std::string model_endpoint = "https://huggingface.co/";
|
||||
if (endpoint_env) {
|
||||
model_endpoint = endpoint_env;
|
||||
if (model_endpoint.back() != '/') {
|
||||
model_endpoint += '/';
|
||||
}
|
||||
std::string endpoint = common_get_env("MODEL_ENDPOINT");
|
||||
if (endpoint.empty()) {
|
||||
// the HF_ENDPOINT variable is respected for backward compatibility
|
||||
endpoint = common_get_env("HF_ENDPOINT");
|
||||
}
|
||||
return model_endpoint;
|
||||
if (endpoint.empty()) {
|
||||
return "https://huggingface.co/";
|
||||
}
|
||||
if (endpoint.back() != '/') {
|
||||
endpoint += '/';
|
||||
}
|
||||
return endpoint;
|
||||
}
|
||||
|
||||
char * common_get_model_or_exit(int argc, char * argv[]) {
|
||||
|
||||
@@ -739,6 +739,8 @@ struct common_params {
|
||||
llama_progress_callback load_progress_callback = NULL;
|
||||
void * load_progress_callback_user_data = NULL;
|
||||
bool no_alloc = false; // Don't allocate model buffers
|
||||
|
||||
bool is_gen_docs = false; // whether we are running inside llama-gen-docs
|
||||
};
|
||||
|
||||
// call once at the start of a program if it uses libcommon
|
||||
@@ -863,6 +865,15 @@ std::string string_from(const struct llama_context * ctx, const struct llama_bat
|
||||
|
||||
bool glob_match(const std::string & pattern, const std::string & str);
|
||||
|
||||
//
|
||||
// Environment utils
|
||||
//
|
||||
|
||||
// portable environment access, an unset variable reads as an empty string
|
||||
// and setting an empty value unsets the variable
|
||||
std::string common_get_env(const std::string & name);
|
||||
void common_set_env(const std::string & name, const std::string & value);
|
||||
|
||||
//
|
||||
// Filesystem utils
|
||||
//
|
||||
|
||||
@@ -482,6 +482,7 @@ caps caps_get(jinja::program & prog) {
|
||||
});
|
||||
},
|
||||
[&](context & ctx) {
|
||||
ctx.set_val("enable_thinking", mk_val<value_bool>(true));
|
||||
caps_apply_preserve_reasoning(ctx, true);
|
||||
},
|
||||
nullptr, // tools_fn
|
||||
|
||||
+19
-2
@@ -184,9 +184,26 @@ std::string common_params_sampling::print() const {
|
||||
return std::string(result);
|
||||
}
|
||||
|
||||
struct common_sampler * common_sampler_init(const struct llama_model * model, struct common_params_sampling & params) {
|
||||
const llama_vocab * vocab = llama_model_get_vocab(model);
|
||||
struct common_sampler * common_sampler_init(
|
||||
const struct llama_model * model,
|
||||
struct common_params_sampling & params,
|
||||
int32_t n_ctx) {
|
||||
if (!std::isfinite(params.penalty_repeat) ||
|
||||
params.penalty_repeat <= 0.0f ||
|
||||
!std::isfinite(1.0f/params.penalty_repeat)) {
|
||||
throw std::invalid_argument("penalty_repeat must be finite and greater than 0");
|
||||
}
|
||||
if (!std::isfinite(params.penalty_freq)) {
|
||||
throw std::invalid_argument("penalty_freq must be finite");
|
||||
}
|
||||
if (!std::isfinite(params.penalty_present)) {
|
||||
throw std::invalid_argument("penalty_present must be finite");
|
||||
}
|
||||
if (params.penalty_last_n == -1) {
|
||||
params.penalty_last_n = n_ctx > 0 ? n_ctx : llama_model_n_ctx_train(model);
|
||||
}
|
||||
|
||||
const llama_vocab * vocab = llama_model_get_vocab(model);
|
||||
llama_sampler_chain_params lparams = llama_sampler_chain_default_params();
|
||||
|
||||
lparams.no_perf = params.no_perf;
|
||||
|
||||
+4
-1
@@ -37,7 +37,10 @@ struct common_sampler;
|
||||
// llama_sampler API overloads
|
||||
|
||||
// note: can mutate params in some cases
|
||||
struct common_sampler * common_sampler_init(const struct llama_model * model, struct common_params_sampling & params);
|
||||
struct common_sampler * common_sampler_init(
|
||||
const struct llama_model * model,
|
||||
struct common_params_sampling & params,
|
||||
int32_t n_ctx = 0);
|
||||
|
||||
void common_sampler_free(struct common_sampler * gsmpl);
|
||||
|
||||
|
||||
@@ -535,7 +535,10 @@ class DeepseekV4Model(TextModel):
|
||||
logger.info("Skipping %d DeepSeek-V4 MTP tensor(s) for conversion v0", type(self)._skipped_mtp_tensors)
|
||||
|
||||
# add a default chat template; if the model has a built-in template, it will be overridden later
|
||||
template_path = Path(__file__).parent.parent / "models" / "templates" / "deepseek-ai-DeepSeek-V4.jinja"
|
||||
model_id_hint = self.remote_hf_model_id or self.dir_model.name
|
||||
is_0731 = "0731" in model_id_hint
|
||||
template_name = "deepseek-ai-DeepSeek-V4-Flash-0731.jinja" if is_0731 else "deepseek-ai-DeepSeek-V4.jinja"
|
||||
template_path = Path(__file__).parent.parent / "models" / "templates" / template_name
|
||||
if template_path.is_file():
|
||||
with open(template_path, "r", encoding="utf-8") as f:
|
||||
self.gguf_writer.add_chat_template(f.read())
|
||||
|
||||
@@ -206,10 +206,70 @@ class Glm4MoeModel(TextModel):
|
||||
@ModelBase.register("Glm4MoeLiteForCausalLM")
|
||||
class Glm4MoeLiteModel(DeepseekV2Model):
|
||||
model_arch = gguf.MODEL_ARCH.DEEPSEEK2
|
||||
skip_mtp = False
|
||||
supports_mtp_export = True
|
||||
_n_main_layers: int | None = None
|
||||
|
||||
def set_vocab(self):
|
||||
return self._set_vocab_glm()
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
num_hidden_layers = self.hparams["num_hidden_layers"]
|
||||
self.num_nextn_predict_layers = self.hparams.get("num_nextn_predict_layers", 0)
|
||||
self.skip_mtp = self.no_mtp or self.num_nextn_predict_layers == 0
|
||||
|
||||
if self.skip_mtp:
|
||||
self.block_count = num_hidden_layers
|
||||
else:
|
||||
self.block_count = num_hidden_layers + self.num_nextn_predict_layers
|
||||
|
||||
self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count)
|
||||
|
||||
def set_gguf_parameters(self):
|
||||
super().set_gguf_parameters()
|
||||
|
||||
if self.skip_mtp:
|
||||
return
|
||||
|
||||
self.gguf_writer.add_nextn_predict_layers(self.num_nextn_predict_layers)
|
||||
|
||||
def index_tensors(self, remote_hf_model_id: str | None = None):
|
||||
type(self)._n_main_layers = self.hparams["num_hidden_layers"]
|
||||
return super().index_tensors(remote_hf_model_id=remote_hf_model_id)
|
||||
|
||||
@classmethod
|
||||
def filter_tensors(cls, item):
|
||||
if (titem := super().filter_tensors(item)) is None:
|
||||
return None
|
||||
name, gen = titem
|
||||
|
||||
if cls._n_main_layers is not None:
|
||||
match = re.match(r"model\.layers\.(\d+)\.", name)
|
||||
is_mtp = match is not None and int(match.group(1)) >= cls._n_main_layers
|
||||
if is_mtp and cls.no_mtp:
|
||||
return None
|
||||
if cls.mtp_only and not is_mtp and name not in (
|
||||
"model.embed_tokens.weight", "model.norm.weight", "lm_head.weight",
|
||||
):
|
||||
return None
|
||||
|
||||
return name, gen
|
||||
|
||||
def prepare_metadata(self, vocab_only: bool):
|
||||
from_dir = self.fname_out.is_dir()
|
||||
super().prepare_metadata(vocab_only=vocab_only)
|
||||
|
||||
if not self.mtp_only or not from_dir:
|
||||
return
|
||||
|
||||
output_type: str = self.ftype.name.partition("_")[2]
|
||||
fname_default: str = gguf.naming_convention(
|
||||
self.metadata.name, self.metadata.basename, self.metadata.finetune,
|
||||
self.metadata.version, size_label=None, output_type=output_type, model_type=None)
|
||||
self.fname_out = self.fname_out.parent / f"mtp-{fname_default}.gguf"
|
||||
|
||||
|
||||
@ModelBase.register("GlmMoeDsaForCausalLM")
|
||||
class GlmMoeDsaModel(DeepseekV2Model):
|
||||
|
||||
@@ -70,6 +70,8 @@ static void write_table(std::ostringstream & ss, std::vector<common_arg *> & opt
|
||||
|
||||
static void write_help(std::ostringstream & ss, const md_file & md) {
|
||||
common_params params;
|
||||
params.is_gen_docs = true;
|
||||
|
||||
auto ctx_arg = common_params_parser_init(params, md.ex);
|
||||
|
||||
std::vector<common_arg *> common_options;
|
||||
|
||||
@@ -765,8 +765,9 @@ struct ggml_backend_sched_split {
|
||||
int backend_id;
|
||||
int i_start;
|
||||
int i_end;
|
||||
struct ggml_tensor * inputs[GGML_SCHED_MAX_SPLIT_INPUTS];
|
||||
struct ggml_tensor ** inputs;
|
||||
int n_inputs;
|
||||
int inputs_capacity;
|
||||
// graph view of this split
|
||||
struct ggml_cgraph graph;
|
||||
};
|
||||
@@ -805,8 +806,9 @@ struct ggml_backend_sched {
|
||||
int cur_copy;
|
||||
int next_copy;
|
||||
ggml_backend_event_t events[GGML_SCHED_MAX_BACKENDS][GGML_SCHED_MAX_COPIES];
|
||||
struct ggml_tensor * graph_inputs[GGML_SCHED_MAX_SPLIT_INPUTS];
|
||||
struct ggml_tensor ** graph_inputs;
|
||||
int n_graph_inputs;
|
||||
int graph_inputs_capacity;
|
||||
|
||||
struct ggml_context * ctx;
|
||||
|
||||
@@ -832,6 +834,36 @@ struct ggml_backend_sched {
|
||||
#define tensor_id_copy(id, backend_id, copy_id) sched->hv_tensor_copies[(id) * sched->n_backends * sched->n_copies + (backend_id) * sched->n_copies + (copy_id)]
|
||||
#define tensor_copy(tensor, backend_id, copy_id) tensor_id_copy(hash_id(tensor), backend_id, copy_id)
|
||||
|
||||
static void ggml_backend_sched_split_inputs_grow(struct ggml_backend_sched_split * split) {
|
||||
int new_cap = GGML_SCHED_MAX_SPLIT_INPUTS;
|
||||
if (split->inputs_capacity > 0) {
|
||||
new_cap = 2*split->inputs_capacity;
|
||||
GGML_LOG_WARN("%s: increasing split inputs capacity from %d to %d\n", __func__, split->inputs_capacity, new_cap);
|
||||
}
|
||||
auto * pnew = (struct ggml_tensor **) realloc((void *) split->inputs, new_cap * sizeof(struct ggml_tensor *));
|
||||
if (pnew == NULL) {
|
||||
GGML_LOG_ERROR("%s: failed to allocate %zu bytes\n", __func__, new_cap * sizeof(struct ggml_tensor *));
|
||||
GGML_ABORT("failed to grow split inputs container");
|
||||
}
|
||||
split->inputs = pnew;
|
||||
split->inputs_capacity = new_cap;
|
||||
}
|
||||
|
||||
static void ggml_backend_sched_graph_inputs_grow(ggml_backend_sched_t sched) {
|
||||
int new_cap = GGML_SCHED_MAX_SPLIT_INPUTS;
|
||||
if (sched->graph_inputs_capacity > 0) {
|
||||
new_cap = 2*sched->graph_inputs_capacity;
|
||||
GGML_LOG_WARN("%s: increasing graph inputs capacity from %d to %d\n", __func__, sched->graph_inputs_capacity, new_cap);
|
||||
}
|
||||
auto * pnew = (struct ggml_tensor **) realloc((void *) sched->graph_inputs, new_cap * sizeof(struct ggml_tensor *));
|
||||
if (pnew == NULL) {
|
||||
GGML_LOG_ERROR("%s: failed to allocate %zu bytes\n", __func__, new_cap * sizeof(struct ggml_tensor *));
|
||||
GGML_ABORT("failed to grow graph inputs container");
|
||||
}
|
||||
sched->graph_inputs = pnew;
|
||||
sched->graph_inputs_capacity = new_cap;
|
||||
}
|
||||
|
||||
// returns the priority of the backend, lower id is higher priority
|
||||
static int ggml_backend_sched_backend_id(ggml_backend_sched_t sched, ggml_backend_t backend) {
|
||||
for (int i = 0; i < sched->n_backends; i++) {
|
||||
@@ -1297,7 +1329,7 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra
|
||||
}
|
||||
// check if the split has too many inputs
|
||||
// FIXME: count the number of inputs instead of only checking when full
|
||||
if (split->n_inputs == GGML_SCHED_MAX_SPLIT_INPUTS) {
|
||||
if (split->n_inputs >= split->inputs_capacity) {
|
||||
const size_t id = hash_id(src);
|
||||
int src_backend_id = sched->hv_tensor_backend_ids[id];
|
||||
bool supported = ggml_backend_sched_buffer_supported(sched, src, cur_backend_id);
|
||||
@@ -1313,10 +1345,14 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra
|
||||
split->i_end = i;
|
||||
i_split++;
|
||||
if (i_split >= sched->splits_capacity) {
|
||||
int old_cap = sched->splits_capacity;
|
||||
sched->splits_capacity *= 2;
|
||||
sched->splits = (ggml_backend_sched_split *)
|
||||
realloc(sched->splits, sched->splits_capacity * sizeof(struct ggml_backend_sched_split));
|
||||
GGML_ASSERT(sched->splits != NULL);
|
||||
for (int k = old_cap; k < sched->splits_capacity; k++) {
|
||||
memset(&sched->splits[k], 0, sizeof(struct ggml_backend_sched_split));
|
||||
}
|
||||
}
|
||||
split = &sched->splits[i_split];
|
||||
split->backend_id = node_backend_id;
|
||||
@@ -1353,7 +1389,9 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra
|
||||
SET_CAUSE(tensor_copy, "4.cpy");
|
||||
}
|
||||
int n_graph_inputs = sched->n_graph_inputs++;
|
||||
GGML_ASSERT(n_graph_inputs < GGML_SCHED_MAX_SPLIT_INPUTS);
|
||||
if (n_graph_inputs >= sched->graph_inputs_capacity) {
|
||||
ggml_backend_sched_graph_inputs_grow(sched);
|
||||
}
|
||||
sched->graph_inputs[n_graph_inputs] = src;
|
||||
}
|
||||
}
|
||||
@@ -1373,7 +1411,9 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra
|
||||
SET_CAUSE(tensor_copy, "4.cpy");
|
||||
}
|
||||
int n_inputs = split->n_inputs++;
|
||||
GGML_ASSERT(n_inputs < GGML_SCHED_MAX_SPLIT_INPUTS);
|
||||
if (n_inputs >= split->inputs_capacity) {
|
||||
ggml_backend_sched_split_inputs_grow(split);
|
||||
}
|
||||
split->inputs[n_inputs] = src;
|
||||
}
|
||||
node->src[j] = tensor_id_copy(src_id, cur_backend_id, sched->cur_copy);
|
||||
@@ -1399,7 +1439,11 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra
|
||||
sched->prev_leaf_backend_ids = tmp;
|
||||
}
|
||||
|
||||
int graph_size = std::max(graph->n_nodes, graph->n_leafs) + sched->n_splits*GGML_SCHED_MAX_SPLIT_INPUTS*2*sched->n_copies;
|
||||
int total_inputs = sched->n_graph_inputs;
|
||||
for (int i = 0; i < sched->n_splits; i++) {
|
||||
total_inputs += sched->splits[i].n_inputs;
|
||||
}
|
||||
int graph_size = std::max(graph->n_nodes, graph->n_leafs) + total_inputs * 2 * sched->n_copies;
|
||||
|
||||
// remember the actual graph_size for performing reallocation checks later [GGML_SCHED_DEBUG_REALLOC]
|
||||
sched->debug_prev_graph_size = sched->debug_graph_size;
|
||||
@@ -1782,6 +1826,9 @@ ggml_backend_sched_t ggml_backend_sched_new(
|
||||
sched->splits = (ggml_backend_sched_split *) calloc(initial_splits_capacity, sizeof(sched->splits[0]));
|
||||
sched->splits_capacity = initial_splits_capacity;
|
||||
|
||||
sched->graph_inputs_capacity = GGML_SCHED_MAX_SPLIT_INPUTS;
|
||||
sched->graph_inputs = (struct ggml_tensor **) calloc(sched->graph_inputs_capacity, sizeof(struct ggml_tensor *));
|
||||
|
||||
for (int b = 0; b < n_backends; b++) {
|
||||
sched->backends[b] = backends[b];
|
||||
sched->bufts[b] = bufts ? bufts[b] : ggml_backend_get_default_buffer_type(backends[b]);
|
||||
@@ -1814,7 +1861,11 @@ void ggml_backend_sched_free(ggml_backend_sched_t sched) {
|
||||
ggml_gallocr_free(sched->galloc);
|
||||
ggml_free(sched->ctx);
|
||||
ggml_hash_set_free(&sched->hash_set);
|
||||
for (int i = 0; i < sched->splits_capacity; i++) {
|
||||
free(sched->splits[i].inputs);
|
||||
}
|
||||
free(sched->splits);
|
||||
free(sched->graph_inputs);
|
||||
free(sched->hv_tensor_backend_ids);
|
||||
free(sched->hv_tensor_copies);
|
||||
free(sched->node_backend_ids);
|
||||
|
||||
@@ -7065,7 +7065,7 @@ static inline bool use_flat_gemv_for_large_m_q4_K(const ggml_tensor *tensor) {
|
||||
return tensor->ne[1] >= 32768 && tensor->ne[2] == 1 && tensor->ne[3] == 1;
|
||||
}
|
||||
|
||||
static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_tensor *tensor) {
|
||||
static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) {
|
||||
// gemv_noshuffle variant perf drops for large M, use flat variant for large M.
|
||||
// threshold is well above typical hidden/FFN dims, but below typical vocab sizes.
|
||||
// q6_K flat gemv is worse for smaller K; 2048 seems to be a reasonable threshold.
|
||||
@@ -7083,7 +7083,15 @@ static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_tensor *tensor) {
|
||||
if ((tensor->ne[1] % 128 != 0) && tensor->ne[2] == 1 && tensor->ne[3] == 1) {
|
||||
return true;
|
||||
}
|
||||
return tensor->ne[1] >= 32768 && tensor->ne[0] >= 2048 && tensor->ne[2] == 1 && tensor->ne[3] == 1;
|
||||
|
||||
// The gemv_noshuffle slowdown tracks TOTAL weight size, not ne0 alone; ne0 >= 2048 is a
|
||||
// proxy for "large weight" that misses a narrow-hidden vocab-scale lm_head.
|
||||
// Add a direct size escape so such weights also take the flat path, without changing
|
||||
// which weights ne0 >= 2048 already routes there.
|
||||
// The size escape is not taken on the A7X since its compiler miscompiles the flat K-quant GEMV
|
||||
return tensor->ne[1] >= 32768
|
||||
&& (tensor->ne[0] >= 2048 || (backend_ctx->adreno_gen != ADRENO_GPU_GEN::A7X && ggml_nbytes(tensor) >= (256ull << 20)))
|
||||
&& tensor->ne[2] == 1 && tensor->ne[3] == 1;
|
||||
}
|
||||
|
||||
static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_tensor * op) {
|
||||
@@ -9403,7 +9411,7 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
|
||||
cl_kernel kernel;
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
kernel = backend_ctx->kernel_convert_block_q6_K;
|
||||
if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(tensor)) {
|
||||
if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(backend_ctx, tensor)) {
|
||||
kernel = backend_ctx->kernel_convert_block_q6_K_noshuffle;
|
||||
}
|
||||
#else
|
||||
@@ -9436,7 +9444,7 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
|
||||
tensor->extra = extra;
|
||||
|
||||
#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
|
||||
if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(tensor)) {
|
||||
if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(backend_ctx, tensor)) {
|
||||
cl_int M = tensor->ne[1]; // ne01
|
||||
cl_int K = tensor->ne[0]; // ne00
|
||||
|
||||
@@ -10473,7 +10481,7 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
|
||||
CL_CHECK(clReleaseMemObject(data_device));
|
||||
return;
|
||||
}
|
||||
if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(tensor)) {
|
||||
if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(backend_ctx, tensor)) {
|
||||
static ggml_cl_buffer buf_trans_ql;
|
||||
static ggml_cl_buffer buf_trans_qh;
|
||||
static ggml_cl_buffer buf_trans_s;
|
||||
@@ -18895,7 +18903,7 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co
|
||||
}
|
||||
|
||||
// q6_K x fp32
|
||||
if (src0t == GGML_TYPE_Q6_K && src1t == GGML_TYPE_F32 && !use_flat_gemv_for_large_m_q6_K(src0)) {
|
||||
if (src0t == GGML_TYPE_Q6_K && src1t == GGML_TYPE_F32 && !use_flat_gemv_for_large_m_q6_K(backend_ctx, src0)) {
|
||||
ggml_cl_mul_mat_q6_K_f32_adreno(backend, src0, src1, dst);
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -3221,6 +3221,13 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = {
|
||||
MODEL_TENSOR.FFN_DOWN_SHEXP,
|
||||
MODEL_TENSOR.FFN_UP_SHEXP,
|
||||
MODEL_TENSOR.FFN_EXP_PROBS_B,
|
||||
# NextN/MTP tensors
|
||||
MODEL_TENSOR.NEXTN_EH_PROJ,
|
||||
MODEL_TENSOR.NEXTN_EMBED_TOKENS,
|
||||
MODEL_TENSOR.NEXTN_ENORM,
|
||||
MODEL_TENSOR.NEXTN_HNORM,
|
||||
MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD,
|
||||
MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM,
|
||||
],
|
||||
MODEL_ARCH.DEEPSEEK2OCR: [
|
||||
MODEL_TENSOR.TOKEN_EMBD,
|
||||
|
||||
+4
-3
@@ -1256,6 +1256,7 @@ extern "C" {
|
||||
struct ggml_tensor * probs;
|
||||
struct ggml_tensor * sampled;
|
||||
struct ggml_tensor * candidates;
|
||||
int64_t n_vocab;
|
||||
};
|
||||
|
||||
// user code can implement the interface below in order to create custom llama_sampler
|
||||
@@ -1425,9 +1426,9 @@ extern "C" {
|
||||
/// NOTE: Avoid using on the full vocabulary as searching for repeated tokens can become slow. For example, apply top-k or top-p sampling first.
|
||||
LLAMA_API struct llama_sampler * llama_sampler_init_penalties(
|
||||
int32_t penalty_last_n, // last n tokens to penalize (0 = disable penalty, -1 = context size)
|
||||
float penalty_repeat, // 1.0 = disabled
|
||||
float penalty_freq, // 0.0 = disabled
|
||||
float penalty_present); // 0.0 = disabled
|
||||
float penalty_repeat, // must be > 0.0, 1.0 = disabled
|
||||
float penalty_freq, // must be finite, 0.0 = disabled
|
||||
float penalty_present); // must be finite, 0.0 = disabled
|
||||
|
||||
/// @details DRY sampler, designed by p-e-w, as described in: https://github.com/oobabooga/text-generation-webui/pull/5677, porting Koboldcpp implementation authored by pi6am: https://github.com/LostRuins/koboldcpp/pull/982
|
||||
LLAMA_API struct llama_sampler * llama_sampler_init_dry(
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
{%- if not add_generation_prompt is defined -%}
|
||||
{%- set add_generation_prompt = false -%}
|
||||
{%- endif -%}
|
||||
{%- if not thinking is defined -%}
|
||||
{%- if enable_thinking is defined -%}
|
||||
{%- set thinking = enable_thinking -%}
|
||||
{%- else -%}
|
||||
{%- set thinking = false -%}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- if not drop_thinking is defined -%}
|
||||
{%- set drop_thinking = true -%}
|
||||
{%- endif -%}
|
||||
{%- set dsml_token = '|DSML|' -%}
|
||||
{%- set thinking_start_token = '<think>' -%}
|
||||
{%- set thinking_end_token = '</think>' -%}
|
||||
{%- set reasoning_effort_high = 'Reasoning Effort: Absolute maximum with no shortcuts permitted.\nYou MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.\nExplicitly write out your entire deliberation process, documenting every intermediate step, considered alternative, and rejected hypothesis to ensure absolutely no assumption is left unchecked.\n\n' -%}
|
||||
{%- set reasoning_effort_max = 'Reasoning Effort: Beyond maximum — exhaustive, relentless, and uncompromising.\nYou MUST reason with the utmost depth and rigor, leaving absolutely nothing to chance: exhaustively decompose the problem into its most fundamental components, trace every causal chain to its root, and resolve the underlying cause rather than any surface symptom.\nDo not stop reasoning until you have independently verified the solution from multiple angles and are certain that no assumption remains unchecked and no error remains undiscovered.\n\n' -%}
|
||||
{%- set response_format_template = '## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n' -%}
|
||||
{%- set has_tools = false -%}
|
||||
{%- set tools_header = '## Tools\n\nYou have access to a set of tools to help answer the user\'s question. You can invoke tools by writing a "<' + dsml_token + 'tool_calls>" block like the following:\n\n<' + dsml_token + 'tool_calls>\n<' + dsml_token + 'invoke name="$TOOL_NAME">\n<' + dsml_token + 'parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE</' + dsml_token + 'parameter>\n...\n</' + dsml_token + 'invoke>\n<' + dsml_token + 'invoke name="$TOOL_NAME2">\n...\n</' + dsml_token + 'invoke>\n</' + dsml_token + 'tool_calls>\n\nString parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`.\n\nIf thinking_mode is enabled (triggered by ' + thinking_start_token + '), you MUST output your complete reasoning inside ' + thinking_start_token + '...' + thinking_end_token + ' BEFORE any tool calls or final response.\n\nOtherwise, output directly after ' + thinking_end_token + ' with tool calls or final response.\n\n### Available Tool Schemas\n\n' -%}
|
||||
{%- set tools_footer = '\nYou MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls.\n' -%}
|
||||
{%- set ns = namespace(system_prompt='', is_first_sp=true, has_tool_calls=false) -%}
|
||||
{%- for message in messages -%}
|
||||
{%- if message['role'] == 'system' -%}
|
||||
{%- if ns.is_first_sp -%}
|
||||
{%- set ns.system_prompt = ns.system_prompt + (message['content'] or '') -%}
|
||||
{%- set ns.is_first_sp = false -%}
|
||||
{%- else -%}
|
||||
{%- set ns.system_prompt = ns.system_prompt + '\n\n' + (message['content'] or '') -%}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- if tools is defined and tools -%}
|
||||
{%- set has_tools = true -%}
|
||||
{%- set ts = namespace(schemas='') -%}
|
||||
{%- for tool in tools -%}
|
||||
{%- if tool['type'] == 'function' -%}
|
||||
{%- set ts.schemas = ts.schemas + (tool['function'] | tojson) + '\n' -%}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- if ns.system_prompt -%}
|
||||
{%- set ns.system_prompt = ns.system_prompt + '\n\n' + tools_header + ts.schemas + tools_footer -%}
|
||||
{%- else -%}
|
||||
{%- set ns.system_prompt = tools_header + ts.schemas + tools_footer -%}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- if response_format is defined -%}
|
||||
{%- if ns.system_prompt -%}
|
||||
{%- set ns.system_prompt = ns.system_prompt + '\n\n' -%}
|
||||
{%- endif -%}
|
||||
{%- set ns.system_prompt = ns.system_prompt + response_format_template + (response_format | tojson) -%}
|
||||
{%- endif -%}
|
||||
{{- bos_token -}}
|
||||
{%- if messages and thinking and reasoning_effort is defined and reasoning_effort == 'high' -%}
|
||||
{{- reasoning_effort_high -}}
|
||||
{%- elif messages and thinking and reasoning_effort is defined and reasoning_effort == 'max' -%}
|
||||
{{- reasoning_effort_max -}}
|
||||
{%- endif -%}
|
||||
{{- ns.system_prompt -}}
|
||||
{%- set last_user_idx = namespace(value=-1) -%}
|
||||
{%- for message in messages -%}
|
||||
{%- if message['role'] == 'user' or message['role'] == 'developer' or message['role'] == 'tool' -%}
|
||||
{%- set last_user_idx.value = loop.index0 -%}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- set state = namespace(in_user=false) -%}
|
||||
{%- for message in messages -%}
|
||||
{%- if message['role'] == 'tool' -%}
|
||||
{%- set ns.has_tool_calls = true -%}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- for message in messages -%}
|
||||
{%- if message['role'] == 'user' or message['role'] == 'developer' -%}
|
||||
{%- if state.in_user -%}
|
||||
{{- '\n\n' -}}
|
||||
{%- else -%}
|
||||
{{- '<|User|>' -}}
|
||||
{%- set state.in_user = true -%}
|
||||
{%- endif -%}
|
||||
{{- message['content'] or '' -}}
|
||||
{%- elif message['role'] == 'tool' -%}
|
||||
{%- if state.in_user -%}
|
||||
{{- '\n\n' -}}
|
||||
{%- else -%}
|
||||
{{- '<|User|>' -}}
|
||||
{%- set state.in_user = true -%}
|
||||
{%- endif -%}
|
||||
{{- '<tool_result>' + (message['content'] or '') + '</tool_result>' -}}
|
||||
{%- elif message['role'] == 'assistant' -%}
|
||||
{%- set state.in_user = false -%}
|
||||
{{- '<|Assistant|>' -}}
|
||||
{%- set is_after_last_user = loop.index0 > last_user_idx.value -%}
|
||||
{%- set keep_reasoning = thinking and ((not drop_thinking) or has_tools or is_after_last_user or ns.has_tool_calls) -%}
|
||||
{%- if keep_reasoning -%}
|
||||
{{- thinking_start_token -}}
|
||||
{%- if message['reasoning_content'] is defined and message['reasoning_content'] -%}
|
||||
{{- message['reasoning_content'] -}}
|
||||
{%- endif -%}
|
||||
{{- thinking_end_token -}}
|
||||
{%- else -%}
|
||||
{{- thinking_end_token -}}
|
||||
{%- endif -%}
|
||||
{%- if message['content'] is defined and message['content'] -%}
|
||||
{{- message['content'] -}}
|
||||
{%- endif -%}
|
||||
{%- if message['tool_calls'] -%}
|
||||
{{- '\n\n<' + dsml_token + 'tool_calls>\n' -}}
|
||||
{%- for tool in message['tool_calls'] -%}
|
||||
{%- set func = tool['function'] -%}
|
||||
{{- '<' + dsml_token + 'invoke name="' + func['name'] + '">\n' -}}
|
||||
{%- set args = func['arguments'] -%}
|
||||
{%- if args is string -%}
|
||||
{%- set args = args | from_json -%}
|
||||
{%- endif -%}
|
||||
{%- for key, val in args.items() -%}
|
||||
{%- if val is string -%}
|
||||
{{- '<' + dsml_token + 'parameter name="' + key + '" string="true">' + val + '</' + dsml_token + 'parameter>\n' -}}
|
||||
{%- else -%}
|
||||
{{- '<' + dsml_token + 'parameter name="' + key + '" string="false">' + (val | tojson) + '</' + dsml_token + 'parameter>\n' -}}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- if not args -%}
|
||||
{{- '\n' -}}
|
||||
{%- endif -%}
|
||||
{{- '</' + dsml_token + 'invoke>\n' -}}
|
||||
{%- endfor -%}
|
||||
{{- '</' + dsml_token + 'tool_calls>' -}}
|
||||
{%- endif -%}
|
||||
{{- '<|end▁of▁sentence|>' -}}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- if add_generation_prompt -%}
|
||||
{{- '<|Assistant|>' -}}
|
||||
{%- if thinking -%}
|
||||
{{- thinking_start_token -}}
|
||||
{%- else -%}
|
||||
{{- thinking_end_token -}}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
@@ -9,11 +9,14 @@
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- if not drop_thinking is defined -%}
|
||||
{%- set drop_thinking = false -%}
|
||||
{%- set drop_thinking = true -%}
|
||||
{%- endif -%}
|
||||
{%- set dsml_token = '|DSML|' -%}
|
||||
{%- set thinking_start_token = '<think>' -%}
|
||||
{%- set thinking_end_token = '</think>' -%}
|
||||
{%- set reasoning_effort_max = 'Reasoning Effort: Absolute maximum with no shortcuts permitted.\nYou MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.\nExplicitly write out your entire deliberation process, documenting every intermediate step, considered alternative, and rejected hypothesis to ensure absolutely no assumption is left unchecked.\n\n' -%}
|
||||
{%- set response_format_template = '## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n' -%}
|
||||
{%- set has_tools = false -%}
|
||||
{%- set tools_header = '## Tools\n\nYou have access to a set of tools to help answer the user\'s question. You can invoke tools by writing a "<' + dsml_token + 'tool_calls>" block like the following:\n\n<' + dsml_token + 'tool_calls>\n<' + dsml_token + 'invoke name="$TOOL_NAME">\n<' + dsml_token + 'parameter name="$PARAMETER_NAME" string="true|false">$PARAMETER_VALUE</' + dsml_token + 'parameter>\n...\n</' + dsml_token + 'invoke>\n<' + dsml_token + 'invoke name="$TOOL_NAME2">\n...\n</' + dsml_token + 'invoke>\n</' + dsml_token + 'tool_calls>\n\nString parameters should be specified as is and set `string="true"`. For all other types (numbers, booleans, arrays, objects), pass the value in JSON format and set `string="false"`.\n\nIf thinking_mode is enabled (triggered by ' + thinking_start_token + '), you MUST output your complete reasoning inside ' + thinking_start_token + '...' + thinking_end_token + ' BEFORE any tool calls or final response.\n\nOtherwise, output directly after ' + thinking_end_token + ' with tool calls or final response.\n\n### Available Tool Schemas\n\n' -%}
|
||||
{%- set tools_footer = '\nYou MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls.\n' -%}
|
||||
{%- set ns = namespace(system_prompt='', is_first_sp=true, has_tool_calls=false) -%}
|
||||
@@ -28,6 +31,7 @@
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- if tools is defined and tools -%}
|
||||
{%- set has_tools = true -%}
|
||||
{%- set ts = namespace(schemas='') -%}
|
||||
{%- for tool in tools -%}
|
||||
{%- if tool['type'] == 'function' -%}
|
||||
@@ -40,7 +44,16 @@
|
||||
{%- set ns.system_prompt = tools_header + ts.schemas + tools_footer -%}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- if response_format is defined -%}
|
||||
{%- if ns.system_prompt -%}
|
||||
{%- set ns.system_prompt = ns.system_prompt + '\n\n' -%}
|
||||
{%- endif -%}
|
||||
{%- set ns.system_prompt = ns.system_prompt + response_format_template + (response_format | tojson) -%}
|
||||
{%- endif -%}
|
||||
{{- bos_token -}}
|
||||
{%- if messages and thinking and reasoning_effort is defined and reasoning_effort == 'max' -%}
|
||||
{{- reasoning_effort_max -}}
|
||||
{%- endif -%}
|
||||
{{- ns.system_prompt -}}
|
||||
{%- set last_user_idx = namespace(value=-1) -%}
|
||||
{%- for message in messages -%}
|
||||
@@ -75,8 +88,8 @@
|
||||
{%- set state.in_user = false -%}
|
||||
{{- '<|Assistant|>' -}}
|
||||
{%- set is_after_last_user = loop.index0 > last_user_idx.value -%}
|
||||
{%- set retain_reasoning = (not drop_thinking) or (is_after_last_user or ns.has_tool_calls) -%}
|
||||
{%- if retain_reasoning and thinking -%}
|
||||
{%- set keep_reasoning = thinking and ((not drop_thinking) or has_tools or is_after_last_user or ns.has_tool_calls) -%}
|
||||
{%- if keep_reasoning -%}
|
||||
{{- thinking_start_token -}}
|
||||
{%- if message['reasoning_content'] is defined and message['reasoning_content'] -%}
|
||||
{{- message['reasoning_content'] -}}
|
||||
@@ -104,6 +117,9 @@
|
||||
{{- '<' + dsml_token + 'parameter name="' + key + '" string="false">' + (val | tojson) + '</' + dsml_token + 'parameter>\n' -}}
|
||||
{%- endif -%}
|
||||
{%- endfor -%}
|
||||
{%- if not args -%}
|
||||
{{- '\n' -}}
|
||||
{%- endif -%}
|
||||
{{- '</' + dsml_token + 'invoke>\n' -}}
|
||||
{%- endfor -%}
|
||||
{{- '</' + dsml_token + 'tool_calls>' -}}
|
||||
@@ -118,4 +134,4 @@
|
||||
{%- else -%}
|
||||
{{- thinking_end_token -}}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
{%- endif -%}
|
||||
|
||||
@@ -5,7 +5,7 @@ import os
|
||||
import sys
|
||||
import subprocess
|
||||
|
||||
HTTPLIB_VERSION = "refs/tags/v0.51.0"
|
||||
HTTPLIB_VERSION = "refs/tags/v0.52.0"
|
||||
|
||||
vendor = {
|
||||
"https://github.com/nlohmann/json/releases/latest/download/json.hpp": "vendor/nlohmann/json.hpp",
|
||||
|
||||
@@ -25,6 +25,7 @@ add_library(llama
|
||||
llama-kv-cache.cpp
|
||||
llama-kv-cache-iswa.cpp
|
||||
llama-kv-cache-dsa.cpp
|
||||
llama-kv-cache-msa.cpp
|
||||
llama-kv-cache-dsv4.cpp
|
||||
llama-memory.cpp
|
||||
llama-memory-hybrid.cpp
|
||||
|
||||
@@ -8,6 +8,7 @@
|
||||
#include "llama-kv-cache.h"
|
||||
#include "llama-kv-cache-iswa.h"
|
||||
#include "llama-kv-cache-dsa.h"
|
||||
#include "llama-kv-cache-msa.h"
|
||||
#include "llama-kv-cache-dsv4.h"
|
||||
#include "llama-memory-hybrid.h"
|
||||
#include "llama-memory-hybrid-iswa.h"
|
||||
@@ -518,6 +519,40 @@ bool llm_graph_input_attn_k::can_reuse(const llm_graph_params & params) {
|
||||
return res;
|
||||
}
|
||||
|
||||
llm_graph_input_attn_kv_msa::llm_graph_input_attn_kv_msa(
|
||||
const llama_hparams & hparams,
|
||||
const llama_cparams & cparams,
|
||||
const llama_kv_cache_msa_context * mctx) :
|
||||
llm_graph_input_attn_kv(hparams, cparams, mctx->get_base()),
|
||||
mctx_msa(mctx) {
|
||||
}
|
||||
|
||||
void llm_graph_input_attn_kv_msa::set_input(const llama_ubatch * ubatch) {
|
||||
llm_graph_input_attn_kv::set_input(ubatch);
|
||||
|
||||
if (self_k_idxs_idx) {
|
||||
mctx_msa->get_idx()->set_input_k_idxs(self_k_idxs_idx, ubatch);
|
||||
}
|
||||
}
|
||||
|
||||
bool llm_graph_input_attn_kv_msa::can_reuse(const llm_graph_params & params) {
|
||||
mctx_msa = static_cast<const llama_kv_cache_msa_context *>(params.mctx);
|
||||
|
||||
// the parent class operates on the base cache context
|
||||
this->mctx = mctx_msa->get_base();
|
||||
|
||||
bool res = true;
|
||||
|
||||
res &= self_k_idxs->ne[0] == params.ubatch.n_tokens;
|
||||
if (self_k_idxs_idx) {
|
||||
res &= self_k_idxs_idx->ne[0] == params.ubatch.n_tokens;
|
||||
}
|
||||
|
||||
res &= can_reuse_kq_mask(self_kq_mask, this->mctx, params.ubatch, params.cparams);
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
void llm_graph_input_attn_k_dsa::set_input(const llama_ubatch * ubatch) {
|
||||
mctx->get_mla()->set_input_k_idxs(self_k_idxs_mla, ubatch);
|
||||
|
||||
@@ -3187,6 +3222,34 @@ llm_graph_input_attn_k_dsa * llm_graph_context::build_attn_inp_k_dsa() const {
|
||||
return (llm_graph_input_attn_k_dsa *) res->add_input(std::move(inp));
|
||||
}
|
||||
|
||||
llm_graph_input_attn_kv_msa * llm_graph_context::build_attn_inp_kv_msa(bool msa_enabled) const {
|
||||
const auto * mctx_cur = static_cast<const llama_kv_cache_msa_context *>(mctx);
|
||||
|
||||
auto inp = std::make_unique<llm_graph_input_attn_kv_msa>(hparams, cparams, mctx_cur);
|
||||
|
||||
const auto * mctx_base = mctx_cur->get_base();
|
||||
const auto * mctx_idx = mctx_cur->get_idx();
|
||||
|
||||
{
|
||||
GGML_ASSERT(hparams.swa_type == LLAMA_SWA_TYPE_NONE && "Use llama_kv_cache_iswa for SWA");
|
||||
|
||||
inp->self_k_idxs = mctx_base->build_input_k_idxs(ctx0, ubatch);
|
||||
inp->self_v_idxs = mctx_base->build_input_v_idxs(ctx0, ubatch);
|
||||
|
||||
inp->self_kq_mask = build_attn_inp_kq_mask(ctx0, mctx_base, ubatch, cparams);
|
||||
inp->self_kq_mask_cnv = inp->self_kq_mask;
|
||||
}
|
||||
|
||||
inp->self_k_rot = mctx_base->build_input_k_rot(ctx0);
|
||||
inp->self_v_rot = mctx_base->build_input_v_rot(ctx0);
|
||||
|
||||
if (msa_enabled) {
|
||||
inp->self_k_idxs_idx = mctx_idx->build_input_k_idxs(ctx0, ubatch);
|
||||
}
|
||||
|
||||
return (llm_graph_input_attn_kv_msa *) res->add_input(std::move(inp));
|
||||
}
|
||||
|
||||
// TODO: maybe separate the inner implementation into a separate function
|
||||
// like with the non-sliding window equivalent
|
||||
// once sliding-window hybrid caches are a thing.
|
||||
@@ -3620,6 +3683,7 @@ void llm_graph_context::build_sampling() const {
|
||||
/*.probs =*/ nullptr,
|
||||
/*.sampled =*/ nullptr,
|
||||
/*.candidates =*/ nullptr,
|
||||
/*.n_vocab =*/ logits_seq->ne[0],
|
||||
};
|
||||
|
||||
assert(sampler->iface->backend_apply);
|
||||
|
||||
@@ -23,6 +23,7 @@ struct llama_memory_context_i;
|
||||
|
||||
class llama_kv_cache_context;
|
||||
class llama_kv_cache_dsa_context;
|
||||
class llama_kv_cache_msa_context;
|
||||
class llama_kv_cache_dsv4_raw_context;
|
||||
class llama_kv_cache_dsv4_context;
|
||||
class llama_kv_cache_iswa_context;
|
||||
@@ -425,6 +426,26 @@ public:
|
||||
const llama_kv_cache_dsa_context * mctx;
|
||||
};
|
||||
|
||||
// standard K/V attention input against the base cache, plus destination indices for the indexer key cache
|
||||
class llm_graph_input_attn_kv_msa : public llm_graph_input_attn_kv {
|
||||
public:
|
||||
llm_graph_input_attn_kv_msa(
|
||||
const llama_hparams & hparams,
|
||||
const llama_cparams & cparams,
|
||||
const llama_kv_cache_msa_context * mctx);
|
||||
~llm_graph_input_attn_kv_msa() = default;
|
||||
|
||||
void set_input(const llama_ubatch * ubatch) override;
|
||||
|
||||
bool can_reuse(const llm_graph_params & params) override;
|
||||
|
||||
ggml_tensor * get_k_idxs_idx() const { return self_k_idxs_idx; }
|
||||
|
||||
ggml_tensor * self_k_idxs_idx = nullptr; // I64 [n_batch]
|
||||
|
||||
const llama_kv_cache_msa_context * mctx_msa;
|
||||
};
|
||||
|
||||
class llm_graph_input_attn_kv_iswa : public llm_graph_input_i {
|
||||
public:
|
||||
llm_graph_input_attn_kv_iswa(
|
||||
@@ -1169,6 +1190,8 @@ struct llm_graph_context {
|
||||
|
||||
llm_graph_input_attn_k_dsa * build_attn_inp_k_dsa() const;
|
||||
|
||||
llm_graph_input_attn_kv_msa * build_attn_inp_kv_msa(bool msa_enabled) const;
|
||||
|
||||
ggml_tensor * build_attn(
|
||||
llm_graph_input_attn_k_dsa * inp,
|
||||
ggml_tensor * wo,
|
||||
|
||||
@@ -180,16 +180,6 @@ uint32_t llama_hparams::n_embd_v_gqa_max() const {
|
||||
return val;
|
||||
}
|
||||
|
||||
uint32_t llama_hparams::n_embd_k_idx(uint32_t il) const {
|
||||
if (!indexer_kv || indexer_head_size == 0) {
|
||||
return 0; // arch without a MSA indexer
|
||||
}
|
||||
if (il < n_layer_dense_lead) {
|
||||
return 0; // leading dense layers carry no indexer
|
||||
}
|
||||
return indexer_head_size; // 128
|
||||
}
|
||||
|
||||
uint32_t llama_hparams::n_embd_r() const {
|
||||
if (wkv_head_size != 0) {
|
||||
// for RWKV models
|
||||
|
||||
@@ -230,8 +230,6 @@ struct llama_hparams {
|
||||
// MSA
|
||||
uint32_t indexer_block_size = 0;
|
||||
uint32_t indexer_local_blocks = 0;
|
||||
// MSA stores its indexer keys in the main KV cache (k_idx tensors);
|
||||
bool indexer_kv = false;
|
||||
|
||||
// Indexer is "full" (1) or "shared" (0)
|
||||
// Shared indexers reuse top-k from previous full layer
|
||||
@@ -356,9 +354,6 @@ struct llama_hparams {
|
||||
uint32_t n_embd_k_gqa_max() const;
|
||||
uint32_t n_embd_v_gqa_max() const;
|
||||
|
||||
// dimension of the single-head MSA indexer key stream
|
||||
uint32_t n_embd_k_idx(uint32_t il = 0) const;
|
||||
|
||||
// dimension of the rolling state embeddings
|
||||
// corresponds to Mamba's conv_states size or RWKV's token_shift states size
|
||||
uint32_t n_embd_r() const;
|
||||
|
||||
@@ -23,7 +23,8 @@ llama_kv_cache_dsa::llama_kv_cache_dsa(
|
||||
uint32_t n_pad,
|
||||
uint32_t n_swa,
|
||||
llama_swa_type swa_type,
|
||||
const layer_filter_cb & filter,
|
||||
const layer_filter_cb & filter_mla,
|
||||
const layer_filter_cb & filter_lid,
|
||||
const layer_reuse_cb & reuse) :
|
||||
hparams_lid(model.hparams), n_stream(unified ? 1 : n_seq_max) {
|
||||
|
||||
@@ -32,7 +33,7 @@ llama_kv_cache_dsa::llama_kv_cache_dsa(
|
||||
kv_mla = std::make_unique<llama_kv_cache>(
|
||||
model, model.hparams, type_k, type_v,
|
||||
v_trans, offload, unified, kv_size, n_seq_max, n_pad,
|
||||
n_swa, swa_type, nullptr, filter, reuse, nullptr);
|
||||
n_swa, swa_type, nullptr, filter_mla, reuse, nullptr);
|
||||
|
||||
// we use llama_kv_cache for caching indexer keys
|
||||
// by hand-tweaking some hparams we fool it to create
|
||||
@@ -49,7 +50,7 @@ llama_kv_cache_dsa::llama_kv_cache_dsa(
|
||||
kv_lid = std::make_unique<llama_kv_cache>(
|
||||
model, hparams_lid, type_k, type_v,
|
||||
v_trans, offload, unified, kv_size, n_seq_max, n_pad,
|
||||
n_swa, swa_type, nullptr, filter, reuse, nullptr);
|
||||
n_swa, swa_type, nullptr, filter_lid, reuse, nullptr);
|
||||
}
|
||||
|
||||
void llama_kv_cache_dsa::clear(bool data) {
|
||||
|
||||
@@ -26,7 +26,8 @@ public:
|
||||
uint32_t n_pad,
|
||||
uint32_t n_swa,
|
||||
llama_swa_type swa_type,
|
||||
const layer_filter_cb & filter,
|
||||
const layer_filter_cb & filter_mla,
|
||||
const layer_filter_cb & filter_lid,
|
||||
const layer_reuse_cb & reuse);
|
||||
|
||||
~llama_kv_cache_dsa() = default;
|
||||
|
||||
@@ -0,0 +1,395 @@
|
||||
#include "llama-kv-cache-msa.h"
|
||||
|
||||
#include "llama-impl.h"
|
||||
#include "llama-batch.h"
|
||||
#include "llama-model.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cmath>
|
||||
|
||||
// llama_kv_cache_msa
|
||||
|
||||
llama_kv_cache_msa::llama_kv_cache_msa(
|
||||
const llama_model & model,
|
||||
ggml_type type_k,
|
||||
ggml_type type_v,
|
||||
bool v_trans,
|
||||
bool offload,
|
||||
bool unified,
|
||||
uint32_t kv_size,
|
||||
uint32_t n_seq_max,
|
||||
uint32_t n_pad,
|
||||
uint32_t n_swa,
|
||||
llama_swa_type swa_type,
|
||||
const layer_filter_cb & filter,
|
||||
const layer_filter_cb & filter_idx,
|
||||
const layer_reuse_cb & reuse) :
|
||||
hparams_idx(model.hparams),
|
||||
n_stream(unified ? 1 : n_seq_max), n_seq_max(n_seq_max), n_pad(n_pad),
|
||||
n_swa(n_swa), swa_type(swa_type) {
|
||||
|
||||
LLAMA_LOG_INFO("%s: creating main KV cache, size = %u cells\n", __func__, kv_size);
|
||||
|
||||
kv_base = std::make_unique<llama_kv_cache>(
|
||||
model, model.hparams, type_k, type_v,
|
||||
v_trans, offload, unified, kv_size, n_seq_max, n_pad,
|
||||
n_swa, swa_type, nullptr, filter, reuse, nullptr);
|
||||
|
||||
// the MSA indexer uses a single key head per layer
|
||||
std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1);
|
||||
hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size;
|
||||
// the rope parameters are kept identical to the main cache
|
||||
|
||||
LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size);
|
||||
|
||||
kv_idx = std::make_unique<llama_kv_cache>(
|
||||
model, hparams_idx, type_k, type_v,
|
||||
v_trans, offload, unified, kv_size, n_seq_max, n_pad,
|
||||
n_swa, swa_type, nullptr, filter_idx, reuse, nullptr);
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa::clear(bool data) {
|
||||
kv_base->clear(data);
|
||||
kv_idx ->clear(data);
|
||||
}
|
||||
|
||||
bool llama_kv_cache_msa::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {
|
||||
bool res = true;
|
||||
|
||||
res = res & kv_base->seq_rm(seq_id, p0, p1);
|
||||
res = res & kv_idx ->seq_rm(seq_id, p0, p1);
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) {
|
||||
kv_base->seq_cp(seq_id_src, seq_id_dst, p0, p1);
|
||||
kv_idx ->seq_cp(seq_id_src, seq_id_dst, p0, p1);
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa::seq_keep(llama_seq_id seq_id) {
|
||||
kv_base->seq_keep(seq_id);
|
||||
kv_idx ->seq_keep(seq_id);
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) {
|
||||
kv_base->seq_add(seq_id, p0, p1, shift);
|
||||
kv_idx ->seq_add(seq_id, p0, p1, shift);
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) {
|
||||
kv_base->seq_div(seq_id, p0, p1, d);
|
||||
kv_idx ->seq_div(seq_id, p0, p1, d);
|
||||
}
|
||||
|
||||
llama_pos llama_kv_cache_msa::seq_pos_min(llama_seq_id seq_id) const {
|
||||
return kv_base->seq_pos_min(seq_id);
|
||||
}
|
||||
|
||||
llama_pos llama_kv_cache_msa::seq_pos_max(llama_seq_id seq_id) const {
|
||||
return kv_base->seq_pos_max(seq_id);
|
||||
}
|
||||
|
||||
std::map<ggml_backend_buffer_type_t, size_t> llama_kv_cache_msa::memory_breakdown() const {
|
||||
std::map<ggml_backend_buffer_type_t, size_t> mb = kv_base->memory_breakdown();
|
||||
for (const auto & buft_size : kv_idx->memory_breakdown()) {
|
||||
mb[buft_size.first] += buft_size.second;
|
||||
}
|
||||
return mb;
|
||||
}
|
||||
|
||||
llama_memory_context_ptr llama_kv_cache_msa::init_batch(
|
||||
llama_batch_allocr & balloc,
|
||||
uint32_t n_ubatch,
|
||||
bool embd_all) {
|
||||
GGML_UNUSED(embd_all);
|
||||
|
||||
do {
|
||||
balloc.split_reset();
|
||||
|
||||
std::vector<llama_ubatch> ubatches;
|
||||
while (true) {
|
||||
auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0);
|
||||
|
||||
if (ubatch.n_tokens == 0) {
|
||||
break;
|
||||
}
|
||||
|
||||
ubatches.push_back(std::move(ubatch));
|
||||
}
|
||||
|
||||
if (balloc.get_n_used() < balloc.get_n_tokens()) {
|
||||
// failed to find a suitable split
|
||||
break;
|
||||
}
|
||||
|
||||
auto sinfos_base = kv_base->prepare(ubatches);
|
||||
if (sinfos_base.empty()) {
|
||||
break;
|
||||
}
|
||||
|
||||
auto sinfos_idx = kv_idx->prepare(ubatches);
|
||||
if (sinfos_idx.empty()) {
|
||||
break;
|
||||
}
|
||||
|
||||
assert(sinfos_base.size() == sinfos_idx.size());
|
||||
|
||||
return std::make_unique<llama_kv_cache_msa_context>(
|
||||
this, std::move(sinfos_base), std::move(sinfos_idx), std::move(ubatches));
|
||||
} while (false);
|
||||
|
||||
return std::make_unique<llama_kv_cache_msa_context>(LLAMA_MEMORY_STATUS_FAILED_PREPARE);
|
||||
}
|
||||
|
||||
llama_memory_context_ptr llama_kv_cache_msa::init_full() {
|
||||
return std::make_unique<llama_kv_cache_msa_context>(this);
|
||||
}
|
||||
|
||||
llama_memory_context_ptr llama_kv_cache_msa::init_update(llama_context * lctx, bool optimize) {
|
||||
return std::make_unique<llama_kv_cache_msa_context>(this, lctx, optimize);
|
||||
}
|
||||
|
||||
bool llama_kv_cache_msa::get_can_shift() const {
|
||||
return kv_base->get_can_shift() &&
|
||||
kv_idx ->get_can_shift() &&
|
||||
kv_base->get_size() == kv_idx->get_size();
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const {
|
||||
kv_base->state_write(io, seq_id, flags);
|
||||
kv_idx ->state_write(io, seq_id, flags);
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) {
|
||||
kv_base->state_read(io, seq_id, flags);
|
||||
kv_idx ->state_read(io, seq_id, flags);
|
||||
}
|
||||
|
||||
llama_kv_cache * llama_kv_cache_msa::get_base() const {
|
||||
return kv_base.get();
|
||||
}
|
||||
|
||||
llama_kv_cache * llama_kv_cache_msa::get_idx() const {
|
||||
return kv_idx.get();
|
||||
}
|
||||
|
||||
// llama_kv_cache_msa_context
|
||||
|
||||
llama_kv_cache_msa_context::llama_kv_cache_msa_context(llama_memory_status status) :
|
||||
kv(nullptr), status(status) {}
|
||||
|
||||
llama_kv_cache_msa_context::llama_kv_cache_msa_context(
|
||||
llama_kv_cache_msa * kv) :
|
||||
kv(kv),
|
||||
ctx_base(kv->get_base()->init_full()),
|
||||
ctx_idx (kv->get_idx ()->init_full()),
|
||||
status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {
|
||||
}
|
||||
|
||||
llama_kv_cache_msa_context::llama_kv_cache_msa_context(
|
||||
llama_kv_cache_msa * kv,
|
||||
llama_context * lctx,
|
||||
bool optimize) :
|
||||
kv(kv),
|
||||
ctx_base(kv->get_base()->init_update(lctx, optimize)),
|
||||
ctx_idx (kv->get_idx ()->init_update(lctx, optimize)),
|
||||
status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {
|
||||
}
|
||||
|
||||
llama_kv_cache_msa_context::llama_kv_cache_msa_context(
|
||||
llama_kv_cache_msa * kv,
|
||||
slot_info_vec_t sinfos_base,
|
||||
slot_info_vec_t sinfos_idx,
|
||||
std::vector<llama_ubatch> ubatches) :
|
||||
kv(kv),
|
||||
ubatches(std::move(ubatches)),
|
||||
// here we copy the ubatches. not sure if this is ideal
|
||||
ctx_base(new llama_kv_cache_context(kv->get_base(), std::move(sinfos_base), this->ubatches)),
|
||||
ctx_idx (new llama_kv_cache_context(kv->get_idx (), std::move(sinfos_idx), this->ubatches)),
|
||||
status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) {
|
||||
}
|
||||
|
||||
llama_kv_cache_msa_context::~llama_kv_cache_msa_context() = default;
|
||||
|
||||
bool llama_kv_cache_msa_context::next() {
|
||||
assert(status == LLAMA_MEMORY_STATUS_SUCCESS);
|
||||
|
||||
ctx_base->next();
|
||||
ctx_idx ->next();
|
||||
|
||||
if (++i_next >= ubatches.size()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool llama_kv_cache_msa_context::apply() {
|
||||
assert(!llama_memory_status_is_fail(status));
|
||||
|
||||
bool res = true;
|
||||
|
||||
res = res & ctx_base->apply();
|
||||
res = res & ctx_idx ->apply();
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
llama_memory_status llama_kv_cache_msa_context::get_status() const {
|
||||
return status;
|
||||
}
|
||||
|
||||
const llama_ubatch & llama_kv_cache_msa_context::get_ubatch() const {
|
||||
assert(status == LLAMA_MEMORY_STATUS_SUCCESS);
|
||||
|
||||
return ubatches[i_next];
|
||||
}
|
||||
|
||||
const llama_kv_cache_context * llama_kv_cache_msa_context::get_base() const {
|
||||
assert(status == LLAMA_MEMORY_STATUS_SUCCESS);
|
||||
|
||||
return static_cast<const llama_kv_cache_context *>(ctx_base.get());
|
||||
}
|
||||
|
||||
const llama_kv_cache_context * llama_kv_cache_msa_context::get_idx() const {
|
||||
assert(status == LLAMA_MEMORY_STATUS_SUCCESS);
|
||||
|
||||
return static_cast<const llama_kv_cache_context *>(ctx_idx.get());
|
||||
}
|
||||
|
||||
uint32_t llama_kv_cache_msa_context::get_n_pos() const {
|
||||
// pad the value so that the graph remains constant across batches and can be reused
|
||||
const uint32_t n_pad_cur = std::max(kv->get_n_pad(), 256u);
|
||||
|
||||
llama_pos pos_max = -1;
|
||||
|
||||
for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) kv->get_n_seq_max(); ++seq_id) {
|
||||
pos_max = std::max(pos_max, kv->seq_pos_max(seq_id));
|
||||
}
|
||||
|
||||
return std::max(n_pad_cur, GGML_PAD((uint32_t) (pos_max + 1), n_pad_cur));
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa_context::set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const {
|
||||
GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));
|
||||
GGML_ASSERT(dst->type == GGML_TYPE_I32);
|
||||
GGML_ASSERT(div > 0);
|
||||
|
||||
const int64_t n_tokens = ubatch->n_tokens;
|
||||
const int64_t n_kv = dst->ne[0];
|
||||
const int64_t n_stream_ub = dst->ne[1];
|
||||
|
||||
GGML_ASSERT(n_tokens % n_stream_ub == 0);
|
||||
const int64_t n_tps = n_tokens/n_stream_ub;
|
||||
|
||||
int32_t * data = (int32_t *) dst->data;
|
||||
|
||||
for (int64_t s = 0; s < n_stream_ub; ++s) {
|
||||
const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0];
|
||||
|
||||
const auto & cells = kv->get_base()->get_cells(seq_id);
|
||||
|
||||
for (int64_t j = 0; j < n_kv; ++j) {
|
||||
// the value for empty or other-sequence cells is irrelevant as consumers mask them
|
||||
data[s*n_kv + j] =
|
||||
cells.is_empty(j) || !cells.seq_has(j, seq_id)
|
||||
? 0
|
||||
: (int32_t) (cells.pos_get(j)/div);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa_context::set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const {
|
||||
GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));
|
||||
GGML_ASSERT(dst->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_F32);
|
||||
|
||||
const int64_t n_tokens = ubatch->n_tokens;
|
||||
const int64_t n_pos = dst->ne[0];
|
||||
const int64_t n_stream_ub = dst->ne[1];
|
||||
|
||||
GGML_ASSERT(n_tokens % n_stream_ub == 0);
|
||||
const int64_t n_tps = n_tokens/n_stream_ub;
|
||||
|
||||
for (int64_t s = 0; s < n_stream_ub; ++s) {
|
||||
const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0];
|
||||
|
||||
const auto & cells = kv->get_base()->get_cells(seq_id);
|
||||
|
||||
std::vector<int32_t> map(n_pos, 0);
|
||||
|
||||
for (uint32_t j = 0; j < cells.size(); ++j) {
|
||||
if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const llama_pos p0 = cells.pos_get(j);
|
||||
|
||||
if (p0 < 0 || p0 >= n_pos) {
|
||||
continue;
|
||||
}
|
||||
|
||||
map[p0] = (int32_t) j;
|
||||
}
|
||||
|
||||
if (dst->type == GGML_TYPE_I32) {
|
||||
int32_t * data = (int32_t *) dst->data + s*n_pos;
|
||||
std::copy(map.begin(), map.end(), data);
|
||||
} else {
|
||||
float * data = (float *) dst->data + s*n_pos;
|
||||
for (int64_t p = 0; p < n_pos; ++p) {
|
||||
data[p] = (float) map[p];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void llama_kv_cache_msa_context::set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const {
|
||||
GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer));
|
||||
GGML_ASSERT(dst->type == GGML_TYPE_F32);
|
||||
|
||||
const int64_t n_tokens = ubatch->n_tokens;
|
||||
const int64_t n_pos = dst->ne[0];
|
||||
|
||||
GGML_ASSERT(dst->ne[1] == n_tokens);
|
||||
|
||||
const uint32_t n_swa = kv->get_n_swa();
|
||||
const llama_swa_type swa_type = kv->get_swa_type();
|
||||
|
||||
float * data = (float *) dst->data;
|
||||
|
||||
std::fill(data, data + n_pos*n_tokens, -INFINITY);
|
||||
|
||||
for (int64_t i = 0; i < n_tokens; ++i) {
|
||||
const llama_seq_id seq_id = ubatch->seq_id[i][0];
|
||||
|
||||
const auto & cells = kv->get_base()->get_cells(seq_id);
|
||||
|
||||
const llama_pos p1 = ubatch->pos[i];
|
||||
|
||||
for (uint32_t j = 0; j < cells.size(); ++j) {
|
||||
if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const llama_pos p0 = cells.pos_get(j);
|
||||
|
||||
if (p0 < 0 || p0 >= n_pos) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// causal mask
|
||||
if (p0 > p1) {
|
||||
continue;
|
||||
}
|
||||
|
||||
// apply SWA if any
|
||||
if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
data[i*n_pos + p0] = 0.0f;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,153 @@
|
||||
#pragma once
|
||||
|
||||
#include "llama-kv-cache.h"
|
||||
|
||||
#include <vector>
|
||||
|
||||
// llama_kv_cache_msa
|
||||
|
||||
// uses two instances of llama_kv_cache, one for K/V tensors, and one for the MSA indexer tensors
|
||||
// both receive identical sequence operations and identical ubatches, so their cell layouts stay in synced.
|
||||
// the context also exposes per-ubatch pos - cell translation maps populated from llama_kv_cells via
|
||||
// llama_kv_cache::get_cells(), which the model graph uses to run MSA block selection in position space
|
||||
|
||||
class llama_kv_cache_msa : public llama_memory_i {
|
||||
public:
|
||||
llama_kv_cache_msa(
|
||||
const llama_model & model,
|
||||
ggml_type type_k,
|
||||
ggml_type type_v,
|
||||
bool v_trans,
|
||||
bool offload,
|
||||
bool unified,
|
||||
uint32_t kv_size,
|
||||
uint32_t n_seq_max,
|
||||
uint32_t n_pad,
|
||||
uint32_t n_swa,
|
||||
llama_swa_type swa_type,
|
||||
const layer_filter_cb & filter,
|
||||
const layer_filter_cb & filter_idx,
|
||||
const layer_reuse_cb & reuse);
|
||||
|
||||
~llama_kv_cache_msa() = default;
|
||||
|
||||
// llama_memory_i
|
||||
|
||||
llama_memory_context_ptr init_batch(
|
||||
llama_batch_allocr & balloc,
|
||||
uint32_t n_ubatch,
|
||||
bool embd_all) override;
|
||||
|
||||
llama_memory_context_ptr init_full() override;
|
||||
|
||||
llama_memory_context_ptr init_update(llama_context * lctx, bool optimize) override;
|
||||
|
||||
bool get_can_shift() const override;
|
||||
|
||||
void clear(bool data) override;
|
||||
|
||||
bool seq_rm (llama_seq_id seq_id, llama_pos p0, llama_pos p1) override;
|
||||
void seq_cp (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) override;
|
||||
void seq_keep(llama_seq_id seq_id) override;
|
||||
void seq_add (llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) override;
|
||||
void seq_div (llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) override;
|
||||
|
||||
llama_pos seq_pos_min(llama_seq_id seq_id) const override;
|
||||
llama_pos seq_pos_max(llama_seq_id seq_id) const override;
|
||||
|
||||
std::map<ggml_backend_buffer_type_t, size_t> memory_breakdown() const override;
|
||||
|
||||
// state write/load
|
||||
|
||||
void state_write(llama_io_write_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) const override;
|
||||
void state_read (llama_io_read_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) override;
|
||||
|
||||
// llama_kv_cache_msa specific API
|
||||
|
||||
llama_kv_cache * get_base() const;
|
||||
llama_kv_cache * get_idx () const;
|
||||
|
||||
uint32_t get_n_pad() const { return n_pad; }
|
||||
uint32_t get_n_seq_max() const { return n_seq_max; }
|
||||
uint32_t get_n_swa() const { return n_swa; }
|
||||
llama_swa_type get_swa_type() const { return swa_type; }
|
||||
|
||||
private:
|
||||
// keep the indexer KV cache hparams instance here as llama_kv_cache stores only a reference
|
||||
llama_hparams hparams_idx;
|
||||
|
||||
const uint32_t n_stream = 1;
|
||||
const uint32_t n_seq_max = 1;
|
||||
const uint32_t n_pad = 1;
|
||||
|
||||
const uint32_t n_swa = 0;
|
||||
const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE;
|
||||
|
||||
std::unique_ptr<llama_kv_cache> kv_base;
|
||||
std::unique_ptr<llama_kv_cache> kv_idx;
|
||||
};
|
||||
|
||||
class llama_kv_cache_msa_context : public llama_memory_context_i {
|
||||
public:
|
||||
using slot_info_vec_t = llama_kv_cache::slot_info_vec_t;
|
||||
|
||||
// used for errors
|
||||
llama_kv_cache_msa_context(llama_memory_status status);
|
||||
|
||||
// used to create a full-cache context
|
||||
llama_kv_cache_msa_context(
|
||||
llama_kv_cache_msa * kv);
|
||||
|
||||
// used to create an update context
|
||||
llama_kv_cache_msa_context(
|
||||
llama_kv_cache_msa * kv,
|
||||
llama_context * lctx,
|
||||
bool optimize);
|
||||
|
||||
// used to create a batch processing context from a batch
|
||||
llama_kv_cache_msa_context(
|
||||
llama_kv_cache_msa * kv,
|
||||
slot_info_vec_t sinfos_base,
|
||||
slot_info_vec_t sinfos_idx,
|
||||
std::vector<llama_ubatch> ubatches);
|
||||
|
||||
virtual ~llama_kv_cache_msa_context();
|
||||
|
||||
// llama_memory_context_i
|
||||
|
||||
bool next() override;
|
||||
bool apply() override;
|
||||
|
||||
llama_memory_status get_status() const override;
|
||||
const llama_ubatch & get_ubatch() const override;
|
||||
|
||||
// llama_kv_cache_msa_context specific API
|
||||
|
||||
const llama_kv_cache_context * get_base() const;
|
||||
const llama_kv_cache_context * get_idx () const;
|
||||
|
||||
// max position currently present in the cache plus one, padded MSA blocks are defined over token positions
|
||||
// so the block-selection tensors are sized by this value rather than by the number of cells
|
||||
uint32_t get_n_pos() const;
|
||||
|
||||
// position <-> cell translation maps, populated from the base cache cells
|
||||
// the model graph relates cache contents to token positions only through these per ubatch inputs
|
||||
// value for empty or other-sequence cells is 0 so consumers must mask them
|
||||
void set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const;
|
||||
// positions without a cell map to cell 0, consumers must mask them assumes one sequence per stream
|
||||
void set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const;
|
||||
void set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const;
|
||||
|
||||
private:
|
||||
llama_kv_cache_msa * kv;
|
||||
|
||||
// the index of the next ubatch to process
|
||||
size_t i_next = 0;
|
||||
|
||||
std::vector<llama_ubatch> ubatches;
|
||||
|
||||
const llama_memory_context_ptr ctx_base;
|
||||
const llama_memory_context_ptr ctx_idx;
|
||||
|
||||
const llama_memory_status status;
|
||||
};
|
||||
+20
-278
@@ -112,7 +112,7 @@ llama_kv_cache::llama_kv_cache(
|
||||
auto it = ctx_map.find(buft);
|
||||
if (it == ctx_map.end()) {
|
||||
ggml_init_params params = {
|
||||
/*.mem_size =*/ size_t(3u*(1 + n_stream)*n_layer*ggml_tensor_overhead()), //Reserve tensor metadata for up to 3 tensors per layer (K, V, and optional K_idx), plus one view per tensor per stream.
|
||||
/*.mem_size =*/ size_t(2u*(1 + n_stream)*n_layer*ggml_tensor_overhead()),
|
||||
/*.mem_buffer =*/ NULL,
|
||||
/*.no_alloc =*/ true,
|
||||
};
|
||||
@@ -242,25 +242,9 @@ llama_kv_cache::llama_kv_cache(
|
||||
v_stream.push_back(has_v ? ggml_view_2d(ctx, v, n_embd_v_gqa, kv_size, v->nb[1], s*v->nb[2]) : nullptr);
|
||||
}
|
||||
|
||||
const uint32_t n_embd_k_idx = hparams.n_embd_k_idx(il);
|
||||
ggml_tensor * k_idx = n_embd_k_idx > 0
|
||||
? ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd_k_idx, kv_size, n_stream)
|
||||
: nullptr;
|
||||
if (k_idx) {
|
||||
ggml_format_name(k_idx, "cache_k_idx_l%d", il);
|
||||
msa_strict_slots = (n_stream == n_seq_max);
|
||||
}
|
||||
|
||||
std::vector<ggml_tensor *> k_idx_stream;
|
||||
for (uint32_t s = 0; s < n_stream; ++s) {
|
||||
k_idx_stream.push_back(k_idx
|
||||
? ggml_view_2d(ctx, k_idx, n_embd_k_idx, kv_size, k_idx->nb[1], s*k_idx->nb[2])
|
||||
: nullptr);
|
||||
}
|
||||
|
||||
map_layer_ids[il] = layers.size();
|
||||
|
||||
layers.push_back({ il, k, v, k_idx, k_stream, v_stream, k_idx_stream });
|
||||
layers.push_back({ il, k, v, k_stream, v_stream, });
|
||||
}
|
||||
|
||||
if (reuse) {
|
||||
@@ -309,24 +293,13 @@ llama_kv_cache::llama_kv_cache(
|
||||
}
|
||||
|
||||
{
|
||||
const size_t memory_size_k = size_k_bytes();
|
||||
const size_t memory_size_v = size_v_bytes();
|
||||
const size_t memory_size_k_idx = size_k_idx_bytes();
|
||||
const size_t memory_size_total = memory_size_k + memory_size_v + memory_size_k_idx;
|
||||
const size_t memory_size_k = size_k_bytes();
|
||||
const size_t memory_size_v = size_v_bytes();
|
||||
|
||||
constexpr float mib = 1024.0f * 1024.0f;
|
||||
|
||||
const std::string k_log = format(", K (%s): %7.2f MiB", ggml_type_name(type_k), (float) memory_size_k / mib);
|
||||
const std::string v_log = format(", V (%s): %7.2f MiB", ggml_type_name(type_v), (float) memory_size_v / mib);
|
||||
|
||||
std::string k_idx_log;
|
||||
if (memory_size_k_idx > 0) {
|
||||
k_idx_log = format(", K_idx (%s): %7.2f MiB", ggml_type_name(GGML_TYPE_F32), (float) memory_size_k_idx / mib);
|
||||
}
|
||||
|
||||
LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs)%s%s%s\n", __func__,
|
||||
(float) memory_size_total / mib, kv_size, (int) layers.size(), n_seq_max, n_stream,
|
||||
k_log.c_str(), v_log.c_str(), k_idx_log.c_str());
|
||||
LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs), K (%s): %7.2f MiB, V (%s): %7.2f MiB\n", __func__,
|
||||
(float)(memory_size_k + memory_size_v) / (1024.0f * 1024.0f), kv_size, (int) layers.size(), n_seq_max, n_stream,
|
||||
ggml_type_name(type_k), (float)memory_size_k / (1024.0f * 1024.0f),
|
||||
ggml_type_name(type_v), (float)memory_size_v / (1024.0f * 1024.0f));
|
||||
}
|
||||
|
||||
// TODO: refactor [TAG_KV_CACHE_SHARE_CELLS]
|
||||
@@ -419,39 +392,6 @@ bool llama_kv_cache::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) {
|
||||
p1 = std::numeric_limits<llama_pos>::max();
|
||||
}
|
||||
|
||||
// empty range - nothing to remove
|
||||
if (p0 >= p1) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// MSA anchors block selection to absolute cache slots (slot == position). Tail trim and full removal preserve this invariant, but removing a prefix
|
||||
// or middle range would free slots while later cells survive, desynchronizing the indexer cache. Reject such removals before modifying the cache.
|
||||
if (msa_strict_slots) {
|
||||
for (llama_seq_id sid = 0; sid < (llama_seq_id) seq_to_stream.size(); ++sid) {
|
||||
if (seq_id >= 0 && sid != seq_id) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const auto & cells = v_cells[seq_to_stream[sid]];
|
||||
|
||||
const llama_pos pmin = cells.seq_pos_min(sid);
|
||||
const llama_pos pmax = cells.seq_pos_max(sid);
|
||||
|
||||
if (pmin < 0) {
|
||||
continue; // empty sequence
|
||||
}
|
||||
|
||||
const bool overlaps = p0 <= pmax && p1 > pmin; // the range removes something
|
||||
const bool leaves_tail = p1 <= pmax; // cells beyond the range survive
|
||||
|
||||
if (overlaps && leaves_tail) {
|
||||
LLAMA_LOG_WARN("%s: MSA: partial (non-suffix) removal [%d, %d) for seq %d is not supported "
|
||||
"(block selection is anchored to cache slots) - rejected\n", __func__, p0, p1, sid);
|
||||
return false;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (seq_id >= 0) {
|
||||
auto & cells = v_cells[seq_to_stream[seq_id]];
|
||||
auto & head = v_heads[seq_to_stream[seq_id]];
|
||||
@@ -906,10 +846,6 @@ bool llama_kv_cache::update(llama_context * lctx, bool do_shift, const stream_co
|
||||
if (layer.v_stream[ssrc]) {
|
||||
ggml_backend_tensor_copy(layer.v_stream[ssrc], layer.v_stream[sdst]);
|
||||
}
|
||||
if (layer.k_idx_stream[ssrc]) {
|
||||
GGML_ASSERT(layer.k_idx_stream[sdst]);
|
||||
ggml_backend_tensor_copy(layer.k_idx_stream[ssrc], layer.k_idx_stream[sdst]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1058,44 +994,6 @@ llama_kv_cache::slot_info llama_kv_cache::find_slot(const llama_ubatch & ubatch,
|
||||
|
||||
const auto & cells = v_cells[seq_to_stream[seq_id]];
|
||||
|
||||
if (n_tokens > cells.size()) {
|
||||
LLAMA_LOG_ERROR("%s: n_tokens = %d > size = %u\n", __func__, n_tokens, cells.size());
|
||||
return { };
|
||||
}
|
||||
|
||||
// MSA block selection assumes slot == logical position (append-only streams).
|
||||
if (msa_strict_slots) {
|
||||
for (uint32_t ii = 0; ii < n_tokens; ++ii) {
|
||||
const llama_pos pos = ubatch.pos[s*n_tokens + ii];
|
||||
|
||||
if (pos < 0 || (uint64_t) pos >= cells.size()) {
|
||||
LLAMA_LOG_WARN("%s: MSA: position %d is outside the cache range [0, %u)\n",
|
||||
__func__, pos, cells.size());
|
||||
return { };
|
||||
}
|
||||
|
||||
const uint32_t idx = (uint32_t) pos;
|
||||
|
||||
if (!cells.is_empty(idx)) {
|
||||
LLAMA_LOG_WARN("%s: MSA: required slot %u is already occupied (stream %u)\n",
|
||||
__func__, idx, seq_to_stream[seq_id]);
|
||||
return { };
|
||||
}
|
||||
|
||||
// strictly increasing positions, rules out duplicates and, for contiguous requests, is tightened to exact adjacency
|
||||
if (!res.idxs[s].empty() && (cont ? idx != res.idxs[s].back() + 1
|
||||
: idx <= res.idxs[s].back())) {
|
||||
LLAMA_LOG_WARN("%s: MSA: token positions are not %s within the ubatch\n",
|
||||
__func__, cont ? "contiguous" : "strictly increasing");
|
||||
return { };
|
||||
}
|
||||
|
||||
res.idxs[s].push_back(idx);
|
||||
}
|
||||
|
||||
continue;
|
||||
}
|
||||
|
||||
uint32_t head_cur = v_heads[seq_to_stream[seq_id]];
|
||||
|
||||
// if we have enough unused cells before the current head ->
|
||||
@@ -1104,6 +1002,11 @@ llama_kv_cache::slot_info llama_kv_cache::find_slot(const llama_ubatch & ubatch,
|
||||
head_cur = 0;
|
||||
}
|
||||
|
||||
if (n_tokens > cells.size()) {
|
||||
LLAMA_LOG_ERROR("%s: n_tokens = %d > size = %u\n", __func__, n_tokens, cells.size());
|
||||
return { };
|
||||
}
|
||||
|
||||
uint32_t n_tested = 0;
|
||||
|
||||
// for continuous slots, we test that all tokens in the ubatch fit, starting from the current head
|
||||
@@ -1210,15 +1113,6 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch &
|
||||
|
||||
const auto idx = sinfo.idxs[s][ii];
|
||||
|
||||
if (msa_strict_slots && (llama_pos) idx != ubatch.pos[i]) {
|
||||
LLAMA_LOG_ERROR("%s: MSA slot/position invariant violated: "
|
||||
"writing pos %d into cell %u (stream %u). The indexer cache "
|
||||
"would desync and block selection would silently corrupt. "
|
||||
"This is a bug, please report it with reproduction steps.\n",
|
||||
__func__, ubatch.pos[i], idx, sinfo.strm[s]);
|
||||
GGML_ABORT("MSA: slot != pos");
|
||||
}
|
||||
|
||||
if (!cells.is_empty(idx)) {
|
||||
assert(cells.seq_count(idx) == 1);
|
||||
|
||||
@@ -1262,8 +1156,7 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch &
|
||||
LLAMA_LOG_DEBUG("%s: purging positions [%d, %d] of sequence %d from KV cache\n",
|
||||
__func__, cells.seq_pos_min(s), seq_pos_max_rm[s], s);
|
||||
|
||||
// under MSA strict slots this path should be unreachable, since strict MSA placement never selects occupied cells
|
||||
GGML_ASSERT(seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1));
|
||||
seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1283,12 +1176,6 @@ bool llama_kv_cache::get_can_shift() const {
|
||||
if (hparams.n_pos_per_embd() > 1) {
|
||||
return false;
|
||||
}
|
||||
// shifting would leave k_idx stale
|
||||
for (const auto & layer : layers) {
|
||||
if (layer.k_idx) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -1337,6 +1224,12 @@ ggml_tensor * llama_kv_cache::get_k_storage(int32_t il) const {
|
||||
return layers[ikv].k;
|
||||
}
|
||||
|
||||
const llama_kv_cells & llama_kv_cache::get_cells(llama_seq_id seq_id) const {
|
||||
GGML_ASSERT(seq_id >= 0 && (size_t) seq_id < seq_to_stream.size());
|
||||
|
||||
return v_cells[seq_to_stream[seq_id]];
|
||||
}
|
||||
|
||||
uint32_t llama_kv_cache::get_n_kv(const slot_info & sinfo) const {
|
||||
uint32_t result = 0;
|
||||
|
||||
@@ -1405,23 +1298,6 @@ ggml_tensor * llama_kv_cache::get_v(ggml_context * ctx, int32_t il, uint32_t n_k
|
||||
ggml_row_size(v->type, kv_size*n_embd_v_gqa)*sinfo.s0);
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache::get_k_idx(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const {
|
||||
const int32_t ikv = map_layer_ids.at(il);
|
||||
auto * k_idx = layers[ikv].k_idx;
|
||||
GGML_ASSERT(k_idx);
|
||||
|
||||
const uint64_t kv_size = get_size();
|
||||
const int64_t n_idx = k_idx->ne[0]; // 128
|
||||
const uint32_t ns = sinfo.s1 - sinfo.s0 + 1;
|
||||
|
||||
return ggml_view_4d(ctx, k_idx,
|
||||
n_idx, 1, n_kv, ns,
|
||||
ggml_row_size(k_idx->type, n_idx), // nb1 (single head)
|
||||
ggml_row_size(k_idx->type, n_idx), // nb2 (per cell)
|
||||
ggml_row_size(k_idx->type, n_idx*kv_size), // nb3 (per stream)
|
||||
ggml_row_size(k_idx->type, n_idx*kv_size)*sinfo.s0);
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const {
|
||||
GGML_UNUSED(sinfo);
|
||||
|
||||
@@ -1523,28 +1399,6 @@ ggml_tensor * llama_kv_cache::build_input_k_idxs(ggml_context * ctx, const llama
|
||||
return k_idxs;
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache::cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const {
|
||||
GGML_UNUSED(sinfo);
|
||||
const int32_t ikv = map_layer_ids.at(il);
|
||||
ggml_tensor * k_idx = layers[ikv].k_idx;
|
||||
GGML_ASSERT(k_idx && "cpy_k_idx on a layer with no indexer cache");
|
||||
|
||||
const int64_t n_embd_head = k_idx_cur->ne[0]; // 128
|
||||
const int64_t n_head = k_idx_cur->ne[1]; // 1
|
||||
const int64_t n_tokens = k_idx_cur->ne[2];
|
||||
const int64_t n_embd_gqa = n_embd_head*n_head; // 128
|
||||
|
||||
GGML_ASSERT(ggml_row_size(k_idx_cur->type, n_embd_head) == k_idx_cur->nb[1]);
|
||||
k_idx_cur = ggml_view_2d(ctx, k_idx_cur, n_embd_gqa, n_tokens, k_idx_cur->nb[2], 0);
|
||||
|
||||
const int64_t n_stream = k_idx->ne[2];
|
||||
if (n_stream > 1) {
|
||||
const int64_t kv_size = get_size();
|
||||
k_idx = ggml_reshape_2d(ctx, k_idx, n_embd_gqa, kv_size*n_stream);
|
||||
}
|
||||
return ggml_set_rows(ctx, k_idx, k_idx_cur, k_idxs); // same k_idxs as the K store
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache::build_input_v_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const {
|
||||
const uint32_t n_tokens = ubatch.n_tokens;
|
||||
|
||||
@@ -1979,18 +1833,6 @@ size_t llama_kv_cache::size_v_bytes() const {
|
||||
return size_v_bytes;
|
||||
}
|
||||
|
||||
size_t llama_kv_cache::size_k_idx_bytes() const {
|
||||
size_t size_k_idx_bytes = 0;
|
||||
|
||||
for (const auto & layer : layers) {
|
||||
if (layer.k_idx) {
|
||||
size_k_idx_bytes += ggml_nbytes(layer.k_idx);
|
||||
}
|
||||
}
|
||||
|
||||
return size_k_idx_bytes;
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache::build_rope_shift(
|
||||
const llama_cparams & cparams,
|
||||
ggml_context * ctx,
|
||||
@@ -2303,36 +2145,6 @@ void llama_kv_cache::state_write_data(llama_io_write_i & io, const cell_ranges_t
|
||||
}
|
||||
}
|
||||
|
||||
if (size_k_idx_bytes() > 0) {
|
||||
const uint32_t has_k_idx_u32 = 1;
|
||||
io.write(&has_k_idx_u32, sizeof(has_k_idx_u32));
|
||||
|
||||
for (const auto & layer : layers) {
|
||||
const uint32_t layer_has_k_idx = layer.k_idx ? 1 : 0;
|
||||
io.write(&layer_has_k_idx, sizeof(layer_has_k_idx));
|
||||
|
||||
if (!layer_has_k_idx) {
|
||||
continue;
|
||||
}
|
||||
|
||||
GGML_ASSERT(layer.k_idx_stream[cr.strm]);
|
||||
|
||||
const int32_t k_idx_type_i = (int32_t) layer.k_idx->type;
|
||||
io.write(&k_idx_type_i, sizeof(k_idx_type_i));
|
||||
|
||||
const uint64_t k_idx_size_row = ggml_row_size(layer.k_idx->type, layer.k_idx->ne[0]);
|
||||
io.write(&k_idx_size_row, sizeof(k_idx_size_row));
|
||||
|
||||
for (const auto & range : cr.data) {
|
||||
const size_t range_size = range.second - range.first;
|
||||
const size_t buf_size = range_size * k_idx_size_row;
|
||||
const size_t offset = range.first * k_idx_size_row;
|
||||
|
||||
io.write_tensor(layer.k_idx_stream[cr.strm], offset, buf_size);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (!v_trans) {
|
||||
for (const auto & layer : layers) {
|
||||
const uint32_t il = layer.il;
|
||||
@@ -2581,68 +2393,6 @@ bool llama_kv_cache::state_read_data(llama_io_read_i & io, uint32_t strm, uint32
|
||||
}
|
||||
}
|
||||
|
||||
if (size_k_idx_bytes() > 0) {
|
||||
uint32_t has_k_idx_u32 = 0;
|
||||
io.read(&has_k_idx_u32, sizeof(has_k_idx_u32));
|
||||
|
||||
if (has_k_idx_u32 != 1) {
|
||||
LLAMA_LOG_ERROR("%s: missing k_idx data in KV cache state\n", __func__);
|
||||
return false;
|
||||
}
|
||||
|
||||
for (const auto & layer : layers) {
|
||||
uint32_t layer_has_k_idx = 0;
|
||||
io.read(&layer_has_k_idx, sizeof(layer_has_k_idx));
|
||||
|
||||
const uint32_t expected_layer_has_k_idx = layer.k_idx ? 1 : 0;
|
||||
|
||||
if (layer_has_k_idx != expected_layer_has_k_idx) {
|
||||
LLAMA_LOG_ERROR(
|
||||
"%s: mismatched k_idx state for layer: got %u, expected %u\n",
|
||||
__func__, layer_has_k_idx, expected_layer_has_k_idx);
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!layer_has_k_idx) {
|
||||
continue;
|
||||
}
|
||||
|
||||
GGML_ASSERT(layer.k_idx_stream[strm]);
|
||||
|
||||
int32_t k_idx_type_i = -1;
|
||||
io.read(&k_idx_type_i, sizeof(k_idx_type_i));
|
||||
|
||||
if (k_idx_type_i != (int32_t) layer.k_idx->type) {
|
||||
LLAMA_LOG_ERROR(
|
||||
"%s: mismatched k_idx type: got %d, expected %d\n",
|
||||
__func__, k_idx_type_i, (int32_t) layer.k_idx->type);
|
||||
return false;
|
||||
}
|
||||
|
||||
uint64_t k_idx_size_row = 0;
|
||||
io.read(&k_idx_size_row, sizeof(k_idx_size_row));
|
||||
|
||||
const uint64_t expected_k_idx_size_row = ggml_row_size(layer.k_idx->type, layer.k_idx->ne[0]);
|
||||
|
||||
if (k_idx_size_row != expected_k_idx_size_row) {
|
||||
LLAMA_LOG_ERROR(
|
||||
"%s: mismatched k_idx row size: got %zu, expected %zu\n",
|
||||
__func__, (size_t) k_idx_size_row, (size_t) expected_k_idx_size_row);
|
||||
return false;
|
||||
}
|
||||
|
||||
if (cell_count) {
|
||||
if (sinfo.is_contiguous()) {
|
||||
io.read_tensor(layer.k_idx_stream[strm], sinfo.head() * k_idx_size_row, cell_count * k_idx_size_row);
|
||||
} else {
|
||||
for (uint32_t i = 0; i < cell_count; ++i) {
|
||||
io.read_tensor(layer.k_idx_stream[strm], sinfo.idxs[0][i] * k_idx_size_row, k_idx_size_row);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (!this->v_trans) {
|
||||
for (const auto & layer : layers) {
|
||||
const uint32_t il = layer.il;
|
||||
@@ -2844,10 +2594,6 @@ ggml_tensor * llama_kv_cache_context::get_v(ggml_context * ctx, int32_t il) cons
|
||||
return kv->get_v(ctx, il, n_kv, sinfos[i_cur]);
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache_context::get_k_idx(ggml_context * ctx, int32_t il) const {
|
||||
return kv->get_k_idx(ctx, il, n_kv, sinfos[i_cur]);
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache_context::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const {
|
||||
return kv->cpy_k(ctx, k_cur, k_idxs, il, sinfos[i_cur]);
|
||||
}
|
||||
@@ -2856,10 +2602,6 @@ ggml_tensor * llama_kv_cache_context::cpy_v(ggml_context * ctx, ggml_tensor * v_
|
||||
return kv->cpy_v(ctx, v_cur, v_idxs, il, sinfos[i_cur]);
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache_context::cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il) const {
|
||||
return kv->cpy_k_idx(ctx, k_idx_cur, k_idxs, il, sinfos[i_cur]);
|
||||
}
|
||||
|
||||
ggml_tensor * llama_kv_cache_context::build_input_k_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const {
|
||||
return kv->build_input_k_idxs(ctx, ubatch);
|
||||
}
|
||||
|
||||
+2
-10
@@ -164,6 +164,8 @@ public:
|
||||
std::vector<uint32_t> get_layer_ids() const;
|
||||
ggml_tensor * get_k_storage(int32_t il) const;
|
||||
|
||||
const llama_kv_cells & get_cells(llama_seq_id seq_id) const;
|
||||
|
||||
//
|
||||
// graph_build API
|
||||
//
|
||||
@@ -173,12 +175,10 @@ public:
|
||||
// get views of the current state of the cache
|
||||
ggml_tensor * get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;
|
||||
ggml_tensor * get_v(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;
|
||||
ggml_tensor * get_k_idx(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const;
|
||||
|
||||
// store k_cur and v_cur in the cache based on the provided head location
|
||||
ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const;
|
||||
ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il, const slot_info & sinfo) const;
|
||||
ggml_tensor * cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const;
|
||||
|
||||
//
|
||||
// preparation API
|
||||
@@ -230,11 +230,9 @@ private:
|
||||
|
||||
ggml_tensor * k;
|
||||
ggml_tensor * v;
|
||||
ggml_tensor * k_idx; // MSA single-head indexer keys, F32
|
||||
|
||||
std::vector<ggml_tensor *> k_stream;
|
||||
std::vector<ggml_tensor *> v_stream;
|
||||
std::vector<ggml_tensor *> k_idx_stream;
|
||||
};
|
||||
|
||||
bool v_trans = true; // the value tensor is transposed
|
||||
@@ -263,9 +261,6 @@ private:
|
||||
// env: LLAMA_KV_CACHE_DEBUG
|
||||
int debug = 0;
|
||||
|
||||
// set when a k_idx (indexer) cache exists and the stream layout supports MSA (single seq, or one stream per seq)
|
||||
bool msa_strict_slots = false;
|
||||
|
||||
// this is the SWA type of the cache - not to be confused with the model SWA type
|
||||
const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE;
|
||||
|
||||
@@ -298,7 +293,6 @@ private:
|
||||
|
||||
size_t size_k_bytes() const;
|
||||
size_t size_v_bytes() const;
|
||||
size_t size_k_idx_bytes() const;
|
||||
|
||||
ggml_tensor * build_rope_shift(
|
||||
const llama_cparams & cparams,
|
||||
@@ -378,7 +372,6 @@ public:
|
||||
// get views of the current state of the cache
|
||||
ggml_tensor * get_k(ggml_context * ctx, int32_t il) const;
|
||||
ggml_tensor * get_v(ggml_context * ctx, int32_t il) const;
|
||||
ggml_tensor * get_k_idx(ggml_context * ctx, int32_t il) const;
|
||||
|
||||
// store k_cur and v_cur in the cache based on the provided head location
|
||||
// note: the heads in k_cur and v_cur should be laid out contiguously in memory
|
||||
@@ -388,7 +381,6 @@ public:
|
||||
// - v_idxs [n_tokens] or [n_tokens*n_embd_v_gqa] depending if V cache is transposed
|
||||
ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const;
|
||||
ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il) const;
|
||||
ggml_tensor * cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il) const;
|
||||
|
||||
// create destination indices for each head of the current batch for where it would be written in the KV cache
|
||||
// the indices address the global KV cache (not per stream) - this is not relevant for the user of this API, but
|
||||
|
||||
+28
-3
@@ -11,6 +11,7 @@
|
||||
#include "llama-kv-cache.h"
|
||||
#include "llama-kv-cache-iswa.h"
|
||||
#include "llama-kv-cache-dsa.h"
|
||||
#include "llama-kv-cache-msa.h"
|
||||
#include "llama-kv-cache-dsv4.h"
|
||||
#include "llama-memory-hybrid.h"
|
||||
#include "llama-memory-hybrid-iswa.h"
|
||||
@@ -2071,6 +2072,28 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params,
|
||||
{
|
||||
res = nullptr;
|
||||
} break;
|
||||
case LLM_ARCH_MINIMAX_M3:
|
||||
{
|
||||
// sparse (MSA) layers carry an indexer key cache, but leading dense layers do not
|
||||
llama_kv_cache::layer_filter_cb filter_idx =
|
||||
[&](int32_t il) { return (uint32_t) il >= hparams.n_layer_dense_lead; };
|
||||
|
||||
res = new llama_kv_cache_msa(
|
||||
*this,
|
||||
params.type_k,
|
||||
params.type_v,
|
||||
!cparams.flash_attn,
|
||||
cparams.offload_kqv,
|
||||
cparams.kv_unified,
|
||||
cparams.n_ctx_seq,
|
||||
cparams.n_seq_max,
|
||||
1,
|
||||
hparams.n_swa,
|
||||
hparams.swa_type,
|
||||
nullptr,
|
||||
filter_idx,
|
||||
nullptr);
|
||||
} break;
|
||||
case LLM_ARCH_GLM_DSA:
|
||||
case LLM_ARCH_DEEPSEEK32:
|
||||
{
|
||||
@@ -2101,10 +2124,11 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params,
|
||||
} else {
|
||||
// Main context: DSA cache for the trunk layers only - the nextn
|
||||
// layer(s) are never attended by the trunk graph.
|
||||
llama_kv_cache::layer_filter_cb filter = nullptr;
|
||||
llama_kv_cache::layer_filter_cb filter_mla = nullptr;
|
||||
if (hparams.n_layer_nextn > 0) {
|
||||
filter = [&](uint32_t il) { return il < hparams.n_layer(); };
|
||||
filter_mla = [&](uint32_t il) { return il < hparams.n_layer(); };
|
||||
}
|
||||
llama_kv_cache::layer_filter_cb filter_lid = [&](uint32_t il) { return il < hparams.n_layer() && (arch != LLM_ARCH_GLM_DSA || hparams.is_indexer_full(il)); };
|
||||
|
||||
res = new llama_kv_cache_dsa(
|
||||
*this,
|
||||
@@ -2118,7 +2142,8 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params,
|
||||
1,
|
||||
hparams.n_swa,
|
||||
hparams.swa_type,
|
||||
filter,
|
||||
filter_mla,
|
||||
filter_lid,
|
||||
nullptr);
|
||||
}
|
||||
} break;
|
||||
|
||||
+221
-20
@@ -589,6 +589,7 @@ static bool llama_sampler_backend_support(
|
||||
/*.probs = */ nullptr,
|
||||
/*.sampled = */ nullptr,
|
||||
/*.candidates = */ ggml_new_tensor_1d(ctx, GGML_TYPE_I32, n),
|
||||
/*.n_vocab = */ n,
|
||||
};
|
||||
|
||||
ggml_cgraph * gf = ggml_new_graph(ctx);
|
||||
@@ -2638,7 +2639,7 @@ struct llama_sampler * llama_sampler_init_grammar_lazy_patterns(
|
||||
|
||||
// penalties
|
||||
|
||||
struct llama_sampler_penalties {
|
||||
struct llama_sampler_penalties : public llama_sampler_backend {
|
||||
const int32_t penalty_last_n;
|
||||
const float penalty_repeat;
|
||||
const float penalty_freq;
|
||||
@@ -2648,10 +2649,49 @@ struct llama_sampler_penalties {
|
||||
|
||||
// a frequency map to count token occurrences
|
||||
std::unordered_map<llama_token, int> token_count;
|
||||
|
||||
// backend graph inputs
|
||||
ggml_tensor * inp_token_ids = nullptr;
|
||||
ggml_tensor * inp_counts = nullptr;
|
||||
|
||||
// backend helpers
|
||||
int32_t n_vocab = 0;
|
||||
int32_t n_max = 0;
|
||||
bool has_candidates = false;
|
||||
|
||||
std::vector<int32_t> host_token_ids;
|
||||
std::vector<int32_t> host_counts;
|
||||
|
||||
static bool is_disabled(
|
||||
int32_t penalty_last_n,
|
||||
float penalty_repeat,
|
||||
float penalty_freq,
|
||||
float penalty_present) {
|
||||
return penalty_last_n == 0 ||
|
||||
(penalty_repeat == 1.0f && penalty_freq == 0.0f && penalty_present == 0.0f);
|
||||
}
|
||||
|
||||
bool is_disabled() const {
|
||||
return is_disabled(penalty_last_n, penalty_repeat, penalty_freq, penalty_present);
|
||||
}
|
||||
|
||||
llama_sampler_penalties(
|
||||
int32_t penalty_last_n,
|
||||
float penalty_repeat,
|
||||
float penalty_freq,
|
||||
float penalty_present)
|
||||
: llama_sampler_backend("penalties")
|
||||
, penalty_last_n (penalty_last_n)
|
||||
, penalty_repeat (penalty_repeat)
|
||||
, penalty_freq (penalty_freq)
|
||||
, penalty_present (penalty_present)
|
||||
, prev (penalty_last_n) {
|
||||
}
|
||||
};
|
||||
|
||||
static const char * llama_sampler_penalties_name(const struct llama_sampler * /*smpl*/) {
|
||||
return "penalties";
|
||||
static const char * llama_sampler_penalties_name(const struct llama_sampler * smpl) {
|
||||
auto * ctx = (llama_sampler_penalties *) smpl->ctx;
|
||||
return ctx->get_name();
|
||||
}
|
||||
|
||||
static void llama_sampler_penalties_accept(struct llama_sampler * smpl, llama_token token) {
|
||||
@@ -2688,8 +2728,7 @@ static void llama_sampler_penalties_accept(struct llama_sampler * smpl, llama_to
|
||||
static void llama_sampler_penalties_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) {
|
||||
auto * ctx = (llama_sampler_penalties *) smpl->ctx;
|
||||
|
||||
if ((ctx->penalty_last_n == 0) ||
|
||||
(ctx->penalty_repeat == 1.0f && ctx->penalty_freq == 0.0f && ctx->penalty_present == 0.0f)) {
|
||||
if (ctx->is_disabled()) {
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -2736,7 +2775,8 @@ static struct llama_sampler * llama_sampler_penalties_clone(const struct llama_s
|
||||
{
|
||||
auto * result_ctx = (llama_sampler_penalties *) result->ctx;
|
||||
|
||||
result_ctx->prev = ctx->prev;
|
||||
result_ctx->prev = ctx->prev;
|
||||
result_ctx->token_count = ctx->token_count;
|
||||
}
|
||||
|
||||
return result;
|
||||
@@ -2746,6 +2786,171 @@ static void llama_sampler_penalties_free(struct llama_sampler * smpl) {
|
||||
delete (llama_sampler_penalties *) smpl->ctx;
|
||||
}
|
||||
|
||||
static bool llama_sampler_penalties_backend_init(
|
||||
struct llama_sampler * smpl,
|
||||
ggml_backend_buffer_type_t buft) {
|
||||
auto * sctx = (llama_sampler_penalties *) smpl->ctx;
|
||||
|
||||
const bool res = llama_sampler_backend_support(smpl, buft);
|
||||
|
||||
sctx->init(res);
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
static void llama_sampler_penalties_backend_apply(
|
||||
struct llama_sampler * smpl,
|
||||
struct ggml_context * ctx,
|
||||
struct ggml_cgraph * gf,
|
||||
struct llama_sampler_data * data) {
|
||||
GGML_UNUSED(gf);
|
||||
|
||||
auto * sctx = (llama_sampler_penalties *) smpl->ctx;
|
||||
|
||||
if (sctx->is_disabled()) {
|
||||
return;
|
||||
}
|
||||
|
||||
GGML_ASSERT(data->n_vocab > 0 && data->n_vocab <= INT32_MAX);
|
||||
|
||||
sctx->has_candidates = data->candidates != nullptr;
|
||||
sctx->n_vocab = (int32_t) data->n_vocab;
|
||||
sctx->n_max = std::min(sctx->penalty_last_n, sctx->n_vocab);
|
||||
|
||||
sctx->inp_token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, sctx->n_max);
|
||||
ggml_set_name(sctx->inp_token_ids, "penalties_token_ids");
|
||||
ggml_set_input(sctx->inp_token_ids);
|
||||
|
||||
sctx->inp_counts = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, sctx->n_max);
|
||||
ggml_set_name(sctx->inp_counts, "penalties_counts");
|
||||
ggml_set_input(sctx->inp_counts);
|
||||
|
||||
if ((int32_t) sctx->host_token_ids.size() != sctx->n_max) {
|
||||
sctx->host_token_ids.assign(sctx->n_max, 0);
|
||||
sctx->host_counts.assign(sctx->n_max, 0);
|
||||
}
|
||||
|
||||
// flatten
|
||||
ggml_tensor * logits = ggml_reshape_1d(ctx, data->logits, ggml_nelements(data->logits));
|
||||
ggml_tensor * gathered = logits;
|
||||
ggml_tensor * counts_f32 = ggml_cast(ctx, sctx->inp_counts, GGML_TYPE_F32);
|
||||
|
||||
if (sctx->has_candidates) {
|
||||
ggml_tensor * candidates = ggml_reshape_1d(
|
||||
ctx, data->candidates, ggml_nelements(data->candidates));
|
||||
const int64_t n_candidates = candidates->ne[0];
|
||||
GGML_ASSERT(n_candidates == ggml_nelements(logits));
|
||||
|
||||
ggml_tensor * counts_rows = ggml_fill(
|
||||
ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 1, sctx->n_vocab), 0.0f);
|
||||
ggml_tensor * scatter_rows = ggml_reshape_2d(ctx, counts_f32, 1, sctx->n_max);
|
||||
counts_rows = ggml_set_rows(ctx, counts_rows, scatter_rows, sctx->inp_token_ids);
|
||||
counts_f32 = ggml_get_rows(ctx, counts_rows, candidates);
|
||||
counts_f32 = ggml_reshape_1d(ctx, counts_f32, n_candidates);
|
||||
} else {
|
||||
ggml_tensor * logits_rows = ggml_reshape_2d(ctx, logits, 1, ggml_nelements(logits));
|
||||
gathered = ggml_get_rows(ctx, logits_rows, sctx->inp_token_ids);
|
||||
gathered = ggml_reshape_1d(ctx, gathered, sctx->n_max);
|
||||
}
|
||||
|
||||
ggml_tensor * active_mask = ggml_step(ctx, counts_f32);
|
||||
ggml_tensor * inactive_mask = ggml_sub(ctx, ggml_fill(ctx, active_mask, 1.0f), active_mask);
|
||||
|
||||
ggml_tensor * penalized = gathered;
|
||||
|
||||
if (sctx->penalty_repeat != 1.0f) {
|
||||
ggml_tensor * pos_mask = ggml_step(ctx, penalized);
|
||||
ggml_tensor * neg_mask = ggml_sub(ctx, ggml_fill(ctx, pos_mask, 1.0f), pos_mask);
|
||||
|
||||
ggml_tensor * pos_scale = ggml_scale(ctx, pos_mask, 1.0f/sctx->penalty_repeat);
|
||||
ggml_tensor * neg_scale = ggml_scale(ctx, neg_mask, sctx->penalty_repeat);
|
||||
ggml_tensor * repeat_scale = ggml_add(ctx, pos_scale, neg_scale);
|
||||
|
||||
// scale inactive entries with 1 to avoid -INF * 0 = NaN for values masked by top-p
|
||||
repeat_scale = ggml_mul(ctx, repeat_scale, active_mask);
|
||||
repeat_scale = ggml_add(ctx, repeat_scale, inactive_mask);
|
||||
penalized = ggml_mul(ctx, gathered, repeat_scale);
|
||||
}
|
||||
|
||||
if (sctx->penalty_freq != 0.0f) {
|
||||
ggml_tensor * penalty_freq = ggml_scale(ctx, counts_f32, sctx->penalty_freq);
|
||||
penalized = ggml_sub(ctx, penalized, penalty_freq);
|
||||
}
|
||||
|
||||
if (sctx->penalty_present != 0.0f) {
|
||||
ggml_tensor * penalty_present = ggml_scale(ctx, active_mask, sctx->penalty_present);
|
||||
penalized = ggml_sub(ctx, penalized, penalty_present);
|
||||
}
|
||||
|
||||
if (sctx->has_candidates) {
|
||||
data->logits = penalized;
|
||||
} else {
|
||||
ggml_tensor * logits_rows = ggml_reshape_2d(ctx, logits, 1, ggml_nelements(logits));
|
||||
ggml_tensor * scatter_rows = ggml_reshape_2d(ctx, penalized, 1, sctx->n_max);
|
||||
logits_rows = ggml_set_rows(ctx, logits_rows, scatter_rows, sctx->inp_token_ids);
|
||||
data->logits = ggml_reshape_1d(ctx, logits_rows, ggml_nelements(logits));
|
||||
}
|
||||
}
|
||||
|
||||
static void llama_sampler_penalties_backend_set_input(struct llama_sampler * smpl) {
|
||||
auto * sctx = (llama_sampler_penalties *) smpl->ctx;
|
||||
|
||||
if (!sctx->inp_token_ids || !sctx->inp_counts || sctx->n_max <= 0 || sctx->n_vocab <= 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
if (sctx->is_disabled()) {
|
||||
return;
|
||||
}
|
||||
|
||||
// fill active entries from the map
|
||||
int32_t n_active = 0;
|
||||
|
||||
for (const auto & it : sctx->token_count) {
|
||||
GGML_ASSERT(n_active < sctx->n_max);
|
||||
sctx->host_token_ids[n_active] = it.first;
|
||||
sctx->host_counts [n_active] = it.second;
|
||||
++n_active;
|
||||
}
|
||||
|
||||
// Sorting is required because backend_apply uses ggml_set_rows (a scatter-back operation)
|
||||
std::vector<std::pair<int32_t, int32_t>> entries;
|
||||
entries.reserve(n_active);
|
||||
for (int32_t i = 0; i < n_active; ++i) {
|
||||
entries.emplace_back(sctx->host_token_ids[i], sctx->host_counts[i]);
|
||||
}
|
||||
std::sort(entries.begin(), entries.end(), [](const auto & a, const auto & b) {
|
||||
return a.first < b.first;
|
||||
});
|
||||
for (int32_t i = 0; i < n_active; ++i) {
|
||||
sctx->host_token_ids[i] = entries[i].first;
|
||||
sctx->host_counts [i] = entries[i].second;
|
||||
}
|
||||
|
||||
// Padding: Finds a filler token id that is not present in token_count.
|
||||
// Use it to do padding for the arrays, it avoids resizing every time.
|
||||
// The arrays must always have exactly n_max entries (the GPU tensor is a fixed size).
|
||||
int32_t filler = 0;
|
||||
if (n_active < sctx->n_max) {
|
||||
while (sctx->token_count.find(filler) != sctx->token_count.end()) {
|
||||
++filler;
|
||||
}
|
||||
GGML_ASSERT(filler < sctx->n_vocab);
|
||||
}
|
||||
|
||||
// Fill the rest of the arrays with the filler token id and count 0.
|
||||
// Inactive slots are padded with a unique dummy token ID (count = 0).
|
||||
// The uniqueness matters because ggml_set_rows with duplicate indices can produce non-deterministic or incorrect results.
|
||||
// Using a filler token with count 0 that isn't in the active set is safe, because the active_mask step in backend_apply filters them out via ggml_step(counts_f32)
|
||||
for (int32_t i = n_active; i < sctx->n_max; ++i) {
|
||||
sctx->host_token_ids[i] = filler;
|
||||
sctx->host_counts [i] = 0;
|
||||
}
|
||||
|
||||
ggml_backend_tensor_set(sctx->inp_token_ids, sctx->host_token_ids.data(), 0, sctx->n_max * sizeof(int32_t));
|
||||
ggml_backend_tensor_set(sctx->inp_counts, sctx->host_counts.data(), 0, sctx->n_max * sizeof(int32_t));
|
||||
}
|
||||
|
||||
static struct llama_sampler_i llama_sampler_penalties_i = {
|
||||
/* .name = */ llama_sampler_penalties_name,
|
||||
/* .accept = */ llama_sampler_penalties_accept,
|
||||
@@ -2753,10 +2958,10 @@ static struct llama_sampler_i llama_sampler_penalties_i = {
|
||||
/* .reset = */ llama_sampler_penalties_reset,
|
||||
/* .clone = */ llama_sampler_penalties_clone,
|
||||
/* .free = */ llama_sampler_penalties_free,
|
||||
/* .backend_init = */ nullptr,
|
||||
/* .backend_init = */ llama_sampler_penalties_backend_init,
|
||||
/* .backend_accept = */ nullptr,
|
||||
/* .backend_apply = */ nullptr,
|
||||
/* .backend_set_input = */ nullptr,
|
||||
/* .backend_apply = */ llama_sampler_penalties_backend_apply,
|
||||
/* .backend_set_input = */ llama_sampler_penalties_backend_set_input,
|
||||
};
|
||||
|
||||
struct llama_sampler * llama_sampler_init_penalties(
|
||||
@@ -2766,22 +2971,18 @@ struct llama_sampler * llama_sampler_init_penalties(
|
||||
float penalty_present) {
|
||||
penalty_last_n = std::max(penalty_last_n, 0);
|
||||
|
||||
const bool is_empty = (penalty_last_n == 0 || (penalty_repeat == 1.0f && penalty_freq == 0.0f && penalty_present == 0.0f));
|
||||
|
||||
if (is_empty) {
|
||||
if (llama_sampler_penalties::is_disabled(
|
||||
penalty_last_n, penalty_repeat, penalty_freq, penalty_present)) {
|
||||
return llama_sampler_init_empty("?penalties");
|
||||
}
|
||||
|
||||
return llama_sampler_init(
|
||||
/* .iface = */ &llama_sampler_penalties_i,
|
||||
/* .ctx = */ new llama_sampler_penalties {
|
||||
/* .penalty_last_n = */ penalty_last_n,
|
||||
/* .penalty_repeat = */ penalty_repeat,
|
||||
/* .penalty_freq = */ penalty_freq,
|
||||
/* .penalty_present = */ penalty_present,
|
||||
/* .prev = */ ring_buffer<llama_token>(penalty_last_n),
|
||||
/* .token_count = */ {},
|
||||
}
|
||||
/* .ctx = */ new llama_sampler_penalties(
|
||||
penalty_last_n,
|
||||
penalty_repeat,
|
||||
penalty_freq,
|
||||
penalty_present)
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -2532,6 +2532,12 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) {
|
||||
const std::string & key = kv(std::get<0>(it));
|
||||
int32_t & id = std::get<1>(it);
|
||||
|
||||
if (id >= 0 && static_cast<size_t>(id) >= id_to_token.size()) {
|
||||
LLAMA_LOG_WARN("%s: default special token '%s' = %d out of vocab range, disabling\n",
|
||||
__func__, key.c_str(), id);
|
||||
id = LLAMA_TOKEN_NULL;
|
||||
}
|
||||
|
||||
uint32_t new_id;
|
||||
if (!ml.get_key(std::get<0>(it), new_id, false)) {
|
||||
continue;
|
||||
|
||||
+308
-25
@@ -37,6 +37,11 @@ void llama_model_deepseek2::load_arch_hparams(llama_model_loader & ml) {
|
||||
hparams.rope_yarn_log_mul /= 0.1f;
|
||||
}
|
||||
|
||||
// NextN/MTP
|
||||
ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.n_layer_nextn, false);
|
||||
GGML_ASSERT(hparams.n_layer_nextn == 0 ||
|
||||
hparams.n_layer() + hparams.n_layer_nextn == hparams.n_layer_all);
|
||||
|
||||
// (optional) temperature tuning - used by mistral-large
|
||||
ml.get_key(LLM_KV_ATTENTION_TEMPERATURE_SCALE, hparams.f_attn_temp_scale, false);
|
||||
ml.get_key(LLM_KV_ATTENTION_TEMPERATURE_LENGTH, hparams.n_attn_temp_floor_scale, false); // FIXME why not use temperature_length?
|
||||
@@ -52,10 +57,20 @@ void llama_model_deepseek2::load_arch_hparams(llama_model_loader & ml) {
|
||||
}
|
||||
}
|
||||
|
||||
void llama_model_deepseek2::load_arch_tensors(llama_model_loader &) {
|
||||
void llama_model_deepseek2::load_arch_tensors(llama_model_loader & ml) {
|
||||
LLAMA_LOAD_LOCALS;
|
||||
const int64_t n_expert_shared = hparams.n_expert_shared;
|
||||
|
||||
const bool mtp_only = (hparams.n_layer_nextn > 0) && (ml.get_weight("blk.0.attn_norm.weight") == nullptr);
|
||||
const std::string mtp_probe = "blk." + std::to_string(n_layer) + ".nextn.eh_proj.weight";
|
||||
const bool trunk_only = (hparams.n_layer_nextn > 0) && (ml.get_weight(mtp_probe.c_str()) == nullptr);
|
||||
const int trunk_flags = mtp_only ? TENSOR_NOT_REQUIRED : 0;
|
||||
int mtp_flags = trunk_only ? TENSOR_NOT_REQUIRED : 0;
|
||||
|
||||
if (!ml.load_mtp) {
|
||||
mtp_flags |= TENSOR_SKIP;
|
||||
}
|
||||
|
||||
const bool is_mla = hparams.is_mla();
|
||||
|
||||
// note: these are the actual head sizes you get when treating as MHA or after "decompression" using wv_b for MLA
|
||||
@@ -81,44 +96,45 @@ void llama_model_deepseek2::load_arch_tensors(llama_model_loader &) {
|
||||
output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED);
|
||||
}
|
||||
|
||||
for (int i = 0; i < n_layer; ++i) {
|
||||
for (int i = 0; i < n_layer_all; ++i) {
|
||||
auto & layer = layers[i];
|
||||
const int flags = i < n_layer ? trunk_flags : mtp_flags;
|
||||
|
||||
layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0);
|
||||
layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, flags);
|
||||
if (q_lora_rank > 0) {
|
||||
layer.attn_q_a_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_A_NORM, "weight", i), {q_lora_rank}, 0);
|
||||
layer.attn_q_a_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_A_NORM, "weight", i), {q_lora_rank}, flags);
|
||||
}
|
||||
|
||||
layer.attn_kv_a_norm = create_tensor(tn(LLM_TENSOR_ATTN_KV_A_NORM, "weight", i), {kv_lora_rank}, 0);
|
||||
layer.attn_kv_a_norm = create_tensor(tn(LLM_TENSOR_ATTN_KV_A_NORM, "weight", i), {kv_lora_rank}, flags);
|
||||
|
||||
if (q_lora_rank > 0) {
|
||||
layer.wq_a = create_tensor(tn(LLM_TENSOR_ATTN_Q_A, "weight", i), {n_embd, q_lora_rank}, 0);
|
||||
layer.wq_b = create_tensor(tn(LLM_TENSOR_ATTN_Q_B, "weight", i), {q_lora_rank, n_head * n_embd_head_k_mla}, 0);
|
||||
layer.wq_a = create_tensor(tn(LLM_TENSOR_ATTN_Q_A, "weight", i), {n_embd, q_lora_rank}, flags);
|
||||
layer.wq_b = create_tensor(tn(LLM_TENSOR_ATTN_Q_B, "weight", i), {q_lora_rank, n_head * n_embd_head_k_mla}, flags);
|
||||
} else {
|
||||
layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_head * n_embd_head_k_mla}, 0);
|
||||
layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_head * n_embd_head_k_mla}, flags);
|
||||
}
|
||||
|
||||
layer.wkv_a_mqa = create_tensor(tn(LLM_TENSOR_ATTN_KV_A_MQA, "weight", i), {n_embd, kv_lora_rank + n_embd_head_qk_rope}, 0);
|
||||
layer.wkv_a_mqa = create_tensor(tn(LLM_TENSOR_ATTN_KV_A_MQA, "weight", i), {n_embd, kv_lora_rank + n_embd_head_qk_rope}, flags);
|
||||
|
||||
// note: only old legacy GGUF files will have the unsplit wkv_b tensor in
|
||||
if (is_mla) {
|
||||
layer.wk_b = create_tensor(tn(LLM_TENSOR_ATTN_K_B, "weight", i), {n_embd_head_qk_nope, kv_lora_rank, n_head}, 0);
|
||||
layer.wv_b = create_tensor(tn(LLM_TENSOR_ATTN_V_B, "weight", i), {kv_lora_rank, n_embd_head_v_mla, n_head}, 0);
|
||||
layer.wk_b = create_tensor(tn(LLM_TENSOR_ATTN_K_B, "weight", i), {n_embd_head_qk_nope, kv_lora_rank, n_head}, flags);
|
||||
layer.wv_b = create_tensor(tn(LLM_TENSOR_ATTN_V_B, "weight", i), {kv_lora_rank, n_embd_head_v_mla, n_head}, flags);
|
||||
} else {
|
||||
layer.wkv_b = create_tensor(tn(LLM_TENSOR_ATTN_KV_B, "weight", i), {kv_lora_rank, n_head * (n_embd_head_qk_nope + n_embd_head_v_mla)}, 0);
|
||||
layer.wkv_b = create_tensor(tn(LLM_TENSOR_ATTN_KV_B, "weight", i), {kv_lora_rank, n_head * (n_embd_head_qk_nope + n_embd_head_v_mla)}, flags);
|
||||
}
|
||||
|
||||
layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_head * n_embd_head_v_mla, n_embd}, 0);
|
||||
layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_head * n_embd_head_v_mla, n_embd}, flags);
|
||||
|
||||
layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, 0);
|
||||
layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, flags);
|
||||
|
||||
if (i < (int) hparams.n_layer_dense_lead) {
|
||||
layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd, n_ff}, 0);
|
||||
layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), { n_ff, n_embd}, 0);
|
||||
layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, 0);
|
||||
layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd, n_ff}, flags);
|
||||
layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), { n_ff, n_embd}, flags);
|
||||
layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, flags);
|
||||
} else {
|
||||
layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), {n_embd, n_expert}, 0);
|
||||
layer.ffn_exp_probs_b = create_tensor(tn(LLM_TENSOR_FFN_EXP_PROBS_B, "bias", i), {n_expert}, TENSOR_NOT_REQUIRED);
|
||||
layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), {n_embd, n_expert}, flags);
|
||||
layer.ffn_exp_probs_b = create_tensor(tn(LLM_TENSOR_FFN_EXP_PROBS_B, "bias", i), {n_expert}, TENSOR_NOT_REQUIRED | flags);
|
||||
|
||||
if (n_expert == 0) {
|
||||
throw std::runtime_error("n_expert must be > 0");
|
||||
@@ -128,21 +144,281 @@ void llama_model_deepseek2::load_arch_tensors(llama_model_loader &) {
|
||||
}
|
||||
|
||||
// MoE branch
|
||||
layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), {n_ff_exp, n_embd, n_expert}, 0);
|
||||
create_tensor_gate_up_exps(layer, i, n_embd, n_ff_exp, n_expert, 0);
|
||||
layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), {n_ff_exp, n_embd, n_expert}, flags);
|
||||
create_tensor_gate_up_exps(layer, i, n_embd, n_ff_exp, n_expert, flags);
|
||||
|
||||
// Shared expert branch
|
||||
layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, 0);
|
||||
layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", i), { n_ff_exp * n_expert_shared, n_embd}, 0);
|
||||
layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, 0);
|
||||
layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, flags);
|
||||
layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", i), { n_ff_exp * n_expert_shared, n_embd}, flags);
|
||||
layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, flags);
|
||||
}
|
||||
|
||||
// NextN/MTP tensors
|
||||
if (i >= n_layer) {
|
||||
layer.nextn.eh_proj = create_tensor(tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", i), { 2 * n_embd, n_embd }, mtp_flags);
|
||||
layer.nextn.enorm = create_tensor(tn(LLM_TENSOR_NEXTN_ENORM, "weight", i), { n_embd }, mtp_flags);
|
||||
layer.nextn.hnorm = create_tensor(tn(LLM_TENSOR_NEXTN_HNORM, "weight", i), { n_embd }, mtp_flags);
|
||||
layer.nextn.embed_tokens = create_tensor(tn(LLM_TENSOR_NEXTN_EMBED_TOKENS, "weight", i), { n_embd, n_vocab }, TENSOR_NOT_REQUIRED | flags);
|
||||
layer.nextn.shared_head_head = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "weight", i), { n_embd, n_vocab }, TENSOR_NOT_REQUIRED | flags);
|
||||
layer.nextn.shared_head_norm = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "weight", i), { n_embd }, TENSOR_NOT_REQUIRED | flags);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
std::unique_ptr<llm_graph_context> llama_model_deepseek2::build_arch_graph(const llm_graph_params & params) const {
|
||||
if (params.gtype == LLM_GRAPH_TYPE_DECODER_MTP) {
|
||||
return std::make_unique<graph_mtp>(*this, params);
|
||||
}
|
||||
return std::make_unique<graph>(*this, params);
|
||||
}
|
||||
|
||||
llama_model_deepseek2::graph_mtp::graph_mtp(const llama_model & model, const llm_graph_params & params) :
|
||||
llm_graph_context(params) {
|
||||
GGML_ASSERT(hparams.n_layer_nextn > 0 && "GLM4 MTP requires n_layer_nextn > 0");
|
||||
GGML_ASSERT(hparams.n_layer_nextn == 1 && "GLM4 MTP currently only supports a single MTP block");
|
||||
GGML_ASSERT(hparams.is_mla() && "GLM4 MTP requires MLA");
|
||||
GGML_ASSERT(hparams.f_attn_temp_scale == 0.0f && "GLM4 MTP does not support attention temperature scaling");
|
||||
|
||||
// The appended MTP block is stored immediately after the main decoder layers.
|
||||
const int il = hparams.n_layer();
|
||||
const auto & layer = model.layers[il];
|
||||
|
||||
GGML_ASSERT(layer.nextn.eh_proj && "MTP block missing nextn.eh_proj");
|
||||
GGML_ASSERT(layer.nextn.enorm && "MTP block missing nextn.enorm");
|
||||
GGML_ASSERT(layer.nextn.hnorm && "MTP block missing nextn.hnorm");
|
||||
|
||||
GGML_ASSERT((uint32_t) il >= hparams.n_layer_dense_lead && "GLM4 MTP block expected to use MoE FFN");
|
||||
|
||||
const int64_t n_embd_head_k_mla = hparams.n_embd_head_k_mla();
|
||||
const int64_t n_embd_head_qk_rope = hparams.n_rot();
|
||||
const int64_t n_embd_head_qk_nope = n_embd_head_k_mla - n_embd_head_qk_rope;
|
||||
const int64_t kv_lora_rank = hparams.n_lora_kv;
|
||||
|
||||
GGML_ASSERT(n_embd_head_qk_nope >= 1);
|
||||
GGML_ASSERT(hparams.n_lora_q > 0);
|
||||
GGML_ASSERT(layer.wq_a);
|
||||
GGML_ASSERT(layer.attn_q_a_norm);
|
||||
GGML_ASSERT(layer.wq_b);
|
||||
GGML_ASSERT(layer.wkv_a_mqa);
|
||||
GGML_ASSERT(layer.attn_kv_a_norm);
|
||||
GGML_ASSERT(layer.wk_b);
|
||||
|
||||
const bool has_split_exps =
|
||||
layer.ffn_up_exps != nullptr &&
|
||||
layer.ffn_gate_exps != nullptr;
|
||||
|
||||
const bool has_fused_exps = layer.ffn_gate_up_exps != nullptr;
|
||||
|
||||
GGML_ASSERT(has_split_exps || has_fused_exps);
|
||||
GGML_ASSERT(layer.ffn_norm);
|
||||
GGML_ASSERT(layer.ffn_gate_inp);
|
||||
GGML_ASSERT(layer.ffn_down_exps);
|
||||
GGML_ASSERT(layer.ffn_gate_shexp);
|
||||
GGML_ASSERT(layer.ffn_down_shexp);
|
||||
GGML_ASSERT(layer.ffn_up_shexp);
|
||||
|
||||
auto inp = std::make_unique<llm_graph_input_embd_h>(hparams.n_embd);
|
||||
|
||||
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens);
|
||||
ggml_set_input(inp->tokens);
|
||||
|
||||
inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens);
|
||||
ggml_set_input(inp->embd);
|
||||
|
||||
ggml_tensor * tok_embd;
|
||||
if (ubatch.token) {
|
||||
ggml_tensor * tok_embd_w = layer.nextn.embed_tokens
|
||||
? layer.nextn.embed_tokens
|
||||
: model.tok_embd;
|
||||
|
||||
tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens);
|
||||
} else {
|
||||
tok_embd = inp->embd;
|
||||
}
|
||||
cb(tok_embd, "mtp_tok_embd", il);
|
||||
|
||||
inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens);
|
||||
ggml_set_input(inp->h);
|
||||
ggml_set_name(inp->h, "mtp_h_input");
|
||||
|
||||
ggml_tensor * h_embd = inp->h;
|
||||
|
||||
res->add_input(std::move(inp));
|
||||
|
||||
ggml_tensor * inp_pos = build_inp_pos();
|
||||
ggml_tensor * inp_out_ids = build_inp_out_ids();
|
||||
|
||||
auto * inp_attn_k = build_attn_inp_k();
|
||||
|
||||
ggml_tensor * h_norm = build_norm(h_embd, layer.nextn.hnorm, nullptr, LLM_NORM_RMS, il);
|
||||
cb(h_norm, "mtp_hnorm", il);
|
||||
|
||||
ggml_tensor * e_norm = build_norm(tok_embd, layer.nextn.enorm, nullptr, LLM_NORM_RMS, il);
|
||||
cb(e_norm, "mtp_enorm", il);
|
||||
|
||||
ggml_tensor * concat = ggml_concat(ctx0, e_norm, h_norm, 0);
|
||||
cb(concat, "mtp_concat", il);
|
||||
|
||||
ggml_tensor * cur = build_lora_mm(layer.nextn.eh_proj, concat, layer.nextn.eh_proj_s);
|
||||
cb(cur, "mtp_eh_proj", il);
|
||||
|
||||
ggml_tensor * inpSA = cur;
|
||||
|
||||
cur = build_norm(cur, layer.attn_norm, nullptr, LLM_NORM_RMS, il);
|
||||
cb(cur, "mtp_attn_norm", il);
|
||||
|
||||
ggml_tensor * q = ggml_mul_mat(ctx0, layer.wq_a, cur);
|
||||
cb(q, "mtp_q_a", il);
|
||||
|
||||
q = build_norm(q, layer.attn_q_a_norm, nullptr, LLM_NORM_RMS, il);
|
||||
cb(q, "mtp_q_a_norm", il);
|
||||
|
||||
q = ggml_mul_mat(ctx0, layer.wq_b, q);
|
||||
cb(q, "mtp_q_b", il);
|
||||
|
||||
ggml_tensor * q_nope =
|
||||
ggml_view_3d(ctx0, q, n_embd_head_qk_nope, n_head, n_tokens,
|
||||
ggml_row_size(q->type, n_embd_head_k_mla),
|
||||
ggml_row_size(q->type, n_embd_head_k_mla) * n_head, 0);
|
||||
cb(q_nope, "mtp_q_nope", il);
|
||||
|
||||
ggml_tensor * q_pe =
|
||||
ggml_view_3d(ctx0, q, n_embd_head_qk_rope, n_head, n_tokens,
|
||||
ggml_row_size(q->type, n_embd_head_k_mla),
|
||||
ggml_row_size(q->type, n_embd_head_k_mla) * n_head,
|
||||
ggml_row_size(q->type, n_embd_head_qk_nope));
|
||||
cb(q_pe, "mtp_q_pe", il);
|
||||
|
||||
ggml_tensor * kv_cmpr_pe = ggml_mul_mat(ctx0, layer.wkv_a_mqa, cur);
|
||||
cb(kv_cmpr_pe, "mtp_kv_cmpr_pe", il);
|
||||
|
||||
ggml_tensor * kv_cmpr =
|
||||
ggml_view_2d(ctx0, kv_cmpr_pe, kv_lora_rank, n_tokens,
|
||||
ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope), 0);
|
||||
cb(kv_cmpr, "mtp_kv_cmpr", il);
|
||||
|
||||
ggml_tensor * k_pe =
|
||||
ggml_view_3d(ctx0, kv_cmpr_pe, n_embd_head_qk_rope, 1, n_tokens,
|
||||
ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope),
|
||||
ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope),
|
||||
ggml_row_size(kv_cmpr_pe->type, kv_lora_rank));
|
||||
cb(k_pe, "mtp_k_pe", il);
|
||||
|
||||
kv_cmpr = build_norm(kv_cmpr, layer.attn_kv_a_norm, nullptr, LLM_NORM_RMS, il);
|
||||
cb(kv_cmpr, "mtp_kv_cmpr_norm", il);
|
||||
|
||||
GGML_ASSERT(ext_factor >= 0.0f);
|
||||
|
||||
const float attn_factor_org =
|
||||
attn_factor * (1.0f + 0.1f * logf(1.0f / freq_scale));
|
||||
|
||||
const float mscale =
|
||||
attn_factor_org * (1.0f + 0.1f * hparams.rope_yarn_log_mul * logf(1.0f / freq_scale));
|
||||
|
||||
const float kq_scale =
|
||||
1.0f * mscale * mscale / sqrtf(float(n_embd_head_k_mla));
|
||||
|
||||
q_pe = ggml_rope_ext(ctx0, q_pe, inp_pos, nullptr,
|
||||
n_rot, rope_type, n_ctx_orig, freq_base, freq_scale,
|
||||
ext_factor, attn_factor, beta_fast, beta_slow);
|
||||
cb(q_pe, "mtp_q_pe_rope", il);
|
||||
|
||||
k_pe = ggml_rope_ext(ctx0, k_pe, inp_pos, nullptr,
|
||||
n_rot, rope_type, n_ctx_orig, freq_base, freq_scale,
|
||||
ext_factor, attn_factor, beta_fast, beta_slow);
|
||||
cb(k_pe, "mtp_k_pe_rope", il);
|
||||
|
||||
q_nope = ggml_permute(ctx0, q_nope, 0, 2, 1, 3);
|
||||
cb(q_nope, "mtp_q_nope_perm", il);
|
||||
|
||||
ggml_tensor * q_nope_absorbed = ggml_mul_mat(ctx0, layer.wk_b, q_nope);
|
||||
cb(q_nope_absorbed, "mtp_q_nope_absorbed", il);
|
||||
|
||||
q_nope_absorbed = ggml_permute(ctx0, q_nope_absorbed, 0, 2, 1, 3);
|
||||
cb(q_nope_absorbed, "mtp_q_nope_absorbed_perm", il);
|
||||
|
||||
ggml_tensor * Qcur = ggml_concat(ctx0, q_nope_absorbed, q_pe, 0);
|
||||
cb(Qcur, "mtp_Qcur", il);
|
||||
|
||||
kv_cmpr = ggml_reshape_3d(ctx0, kv_cmpr, hparams.n_lora_kv, 1, n_tokens);
|
||||
cb(kv_cmpr, "mtp_kv_cmpr_reshape", il);
|
||||
|
||||
ggml_tensor * Kcur = ggml_concat(ctx0, kv_cmpr, k_pe, 0);
|
||||
cb(Kcur, "mtp_Kcur", il);
|
||||
|
||||
ggml_tensor * Vcur = kv_cmpr;
|
||||
cb(Vcur, "mtp_Vcur", il);
|
||||
|
||||
cur = build_attn(inp_attn_k,
|
||||
layer.wo, nullptr, layer.wo_s,
|
||||
Qcur, Kcur, Vcur, nullptr, nullptr, layer.wv_b, kq_scale, il);
|
||||
cb(cur, "mtp_attn_out", il);
|
||||
|
||||
ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpSA);
|
||||
cb(ffn_inp, "mtp_ffn_inp", il);
|
||||
|
||||
cur = build_norm(ffn_inp, layer.ffn_norm, nullptr, LLM_NORM_RMS, il);
|
||||
cb(cur, "mtp_ffn_norm", il);
|
||||
|
||||
ggml_tensor * moe_out = build_moe_ffn(cur,
|
||||
layer.ffn_gate_inp,
|
||||
layer.ffn_up_exps,
|
||||
layer.ffn_gate_exps,
|
||||
layer.ffn_down_exps,
|
||||
layer.ffn_exp_probs_b,
|
||||
n_expert, n_expert_used,
|
||||
LLM_FFN_SILU, hparams.expert_weights_norm,
|
||||
hparams.expert_weights_scale,
|
||||
(llama_expert_gating_func_type) hparams.expert_gating_func,
|
||||
il,
|
||||
nullptr,
|
||||
layer.ffn_gate_up_exps);
|
||||
cb(moe_out, "mtp_ffn_moe_out", il);
|
||||
|
||||
ggml_tensor * ffn_shexp = build_ffn(cur,
|
||||
layer.ffn_up_shexp, nullptr, nullptr,
|
||||
layer.ffn_gate_shexp, nullptr, nullptr,
|
||||
layer.ffn_down_shexp, nullptr, nullptr,
|
||||
nullptr, LLM_FFN_SILU, LLM_FFN_PAR, il);
|
||||
cb(ffn_shexp, "mtp_ffn_shexp", il);
|
||||
|
||||
cur = ggml_add(ctx0, moe_out, ffn_shexp);
|
||||
cb(cur, "mtp_ffn_out", il);
|
||||
|
||||
cur = ggml_add(ctx0, cur, ffn_inp);
|
||||
cb(cur, "mtp_post_ffn", il);
|
||||
|
||||
ggml_tensor * head_norm_w = layer.nextn.shared_head_norm
|
||||
? layer.nextn.shared_head_norm
|
||||
: model.output_norm;
|
||||
GGML_ASSERT(head_norm_w && "GLM4 MTP: missing both nextn.shared_head_norm and output_norm");
|
||||
|
||||
cur = build_norm(cur, head_norm_w, nullptr, LLM_NORM_RMS, -1);
|
||||
cb(cur, "h_nextn", -1);
|
||||
res->t_h_nextn = cur;
|
||||
|
||||
if (inp_out_ids) {
|
||||
cur = ggml_get_rows(ctx0, cur, inp_out_ids);
|
||||
}
|
||||
cb(cur, "mtp_shared_head_norm", -1);
|
||||
|
||||
ggml_tensor * head_w = layer.nextn.shared_head_head
|
||||
? layer.nextn.shared_head_head
|
||||
: model.output;
|
||||
|
||||
ggml_tensor * head_s = layer.nextn.shared_head_head
|
||||
? layer.nextn.shared_head_head_s
|
||||
: model.output_s;
|
||||
|
||||
GGML_ASSERT(head_w && "GLM4 MTP: missing LM head (nextn.shared_head_head or model.output)");
|
||||
|
||||
cur = build_lora_mm(head_w, cur, head_s);
|
||||
cb(cur, "result_output", -1);
|
||||
|
||||
res->t_logits = cur;
|
||||
ggml_build_forward_expand(gf, cur);
|
||||
}
|
||||
|
||||
llama_model_deepseek2::graph::graph(const llama_model & model, const llm_graph_params & params) :
|
||||
llm_graph_context(params) {
|
||||
// lite variants include DeepSeek-V2-Lite, GigaChat3-10B-A1.8B
|
||||
@@ -365,7 +641,7 @@ llama_model_deepseek2::graph::graph(const llama_model & model, const llm_graph_p
|
||||
Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, kq_scale, il);
|
||||
}
|
||||
}
|
||||
if (il == n_layer - 1 && inp_out_ids) {
|
||||
if (il == n_layer - 1 && inp_out_ids && (!cparams.embeddings_nextn || cparams.embeddings_nextn_masked)) {
|
||||
cur = ggml_get_rows(ctx0, cur, inp_out_ids);
|
||||
inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids);
|
||||
}
|
||||
@@ -425,6 +701,13 @@ llama_model_deepseek2::graph::graph(const llama_model & model, const llm_graph_p
|
||||
|
||||
cur = build_norm(cur, model.output_norm, NULL, LLM_NORM_RMS, -1);
|
||||
|
||||
cb(cur, "h_nextn", -1);
|
||||
res->t_h_nextn = cur;
|
||||
|
||||
if (cparams.embeddings_nextn && !cparams.embeddings_nextn_masked && inp_out_ids) {
|
||||
cur = ggml_get_rows(ctx0, cur, inp_out_ids);
|
||||
}
|
||||
|
||||
cb(cur, "result_norm", -1);
|
||||
res->t_embd = cur;
|
||||
|
||||
|
||||
+157
-75
@@ -1,5 +1,5 @@
|
||||
#include "models.h"
|
||||
#include "llama-kv-cache.h"
|
||||
#include "llama-kv-cache-msa.h"
|
||||
#include <cmath>
|
||||
#include <vector>
|
||||
#include <cstdint>
|
||||
@@ -7,7 +7,8 @@
|
||||
// MiniMax-M3: MiniMax-M2 style GQA (per-head QK-norm, partial rotary) with
|
||||
// DeepSeek-V3 leading-dense + routed/shared experts (sigmoid gating, routed scaling),
|
||||
// swigluoai activation, and MiniMax Sparse Attention (MSA). MTP is not in released model weights.
|
||||
// Notes: Blocks are anchored to absolute KV cache slots.
|
||||
// MSA blocks are defined over token positions. The graph translates between position space (block
|
||||
// selection) and cell space (K/V/indexer storage) via per-ubatch pos<->cell maps populated from llama_kv_cells
|
||||
|
||||
void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) {
|
||||
ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps);
|
||||
@@ -23,7 +24,6 @@ void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) {
|
||||
ml.get_key(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE, hparams.indexer_block_size);
|
||||
ml.get_key(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, hparams.indexer_local_blocks);
|
||||
msa_p = { (int) hparams.indexer_block_size, (int) hparams.indexer_top_k, (int) hparams.indexer_local_blocks };
|
||||
hparams.indexer_kv = true;
|
||||
|
||||
switch (hparams.n_layer()) {
|
||||
case 60: type = LLM_TYPE_428B_A23B; break;
|
||||
@@ -86,43 +86,83 @@ std::unique_ptr<llm_graph_context> llama_model_minimax_m3::build_arch_graph(cons
|
||||
return std::make_unique<graph>(*this, params);
|
||||
}
|
||||
|
||||
// per-query local-force bias for MSA selection
|
||||
// local window always wins a slot
|
||||
class llm_graph_input_msa_local : public llm_graph_input_i {
|
||||
class llm_graph_input_msa : public llm_graph_input_i {
|
||||
public:
|
||||
llm_graph_input_msa_local(int blk, int local, int64_t nblk) : blk(blk), local(local), nblk(nblk) {}
|
||||
llm_graph_input_msa(const llama_kv_cache_msa_context * mctx, int blk, int local) :
|
||||
mctx(mctx), blk(blk), local(local) {}
|
||||
|
||||
void set_input(const llama_ubatch * ubatch) override {
|
||||
if (!bias || !ubatch->pos) {
|
||||
return;
|
||||
}
|
||||
const int64_t n_tokens = ubatch->n_tokens;
|
||||
std::vector<float> data((size_t) nblk * n_tokens, 0.0f);
|
||||
for (int64_t i = 0; i < n_tokens; ++i) {
|
||||
const int64_t L = ubatch->pos[i] / blk;
|
||||
for (int l = 0; l < local && L - l >= 0; ++l) {
|
||||
if (L - l < nblk) {
|
||||
data[(size_t) i * nblk + (L - l)] = 1e30f;
|
||||
if (pos_slot_i) { mctx->set_input_pos_slot(pos_slot_i, ubatch); }
|
||||
if (pos_slot_f) { mctx->set_input_pos_slot(pos_slot_f, ubatch); }
|
||||
if (cell_blk) { mctx->set_input_cell_pos(cell_blk, ubatch, blk); }
|
||||
if (pos_mask) { mctx->set_input_pos_mask(pos_mask, ubatch); }
|
||||
|
||||
// local-force bias over position blocks
|
||||
if (bias && ubatch->pos) {
|
||||
const int64_t n_tokens = ubatch->n_tokens;
|
||||
const int64_t nblk = bias->ne[0];
|
||||
std::vector<float> data((size_t) nblk * n_tokens, 0.0f);
|
||||
for (int64_t i = 0; i < n_tokens; ++i) {
|
||||
const int64_t L = ubatch->pos[i] / blk;
|
||||
for (int l = 0; l < local && L - l >= 0; ++l) {
|
||||
if (L - l < nblk) {
|
||||
data[(size_t) i * nblk + (L - l)] = 1e30f;
|
||||
}
|
||||
}
|
||||
}
|
||||
ggml_backend_tensor_set(bias, data.data(), 0, data.size() * sizeof(float));
|
||||
}
|
||||
ggml_backend_tensor_set(bias, data.data(), 0, data.size() * sizeof(float));
|
||||
}
|
||||
|
||||
// valid as long as the bias tensor dims still match the new ubatch/cache window
|
||||
// valid as long as the tensor dims still match the new ubatch/cache window and the
|
||||
// ubatch is in the same regime (decode graphs have pos_slot_f, batch graphs cell_blk)
|
||||
bool can_reuse(const llm_graph_params & params) override {
|
||||
const auto * mctx = static_cast<const llama_kv_cache_context *>(params.mctx);
|
||||
const auto * mctx_new = static_cast<const llama_kv_cache_msa_context *>(params.mctx);
|
||||
|
||||
this->mctx = mctx_new;
|
||||
|
||||
const int64_t n_ps = GGML_PAD((int64_t) mctx_new->get_n_pos(), blk);
|
||||
const int64_t ns = params.cparams.kv_unified ? 1 : params.ubatch.n_seqs_unq;
|
||||
|
||||
const bool decode = params.ubatch.n_tokens == ns; // one token per stream
|
||||
|
||||
bool res = true;
|
||||
res &= bias->ne[1] == params.ubatch.n_tokens;
|
||||
res &= bias->ne[0] * blk == (int64_t) mctx->get_n_kv();
|
||||
|
||||
res &= bias->ne[0] * blk == n_ps;
|
||||
res &= bias->ne[1] == params.ubatch.n_tokens;
|
||||
|
||||
res &= pos_mask->ne[0] == n_ps;
|
||||
res &= pos_mask->ne[1] == params.ubatch.n_tokens;
|
||||
|
||||
res &= pos_slot_i->ne[0] == n_ps;
|
||||
res &= pos_slot_i->ne[1] == ns;
|
||||
|
||||
res &= decode == (pos_slot_f != nullptr);
|
||||
res &= decode == (cell_blk == nullptr);
|
||||
|
||||
if (pos_slot_f) {
|
||||
res &= pos_slot_f->ne[0] == n_ps;
|
||||
res &= pos_slot_f->ne[1] == ns;
|
||||
}
|
||||
|
||||
if (cell_blk) {
|
||||
res &= cell_blk->ne[0] == (int64_t) mctx_new->get_base()->get_n_kv();
|
||||
res &= cell_blk->ne[1] == ns;
|
||||
}
|
||||
|
||||
return res;
|
||||
}
|
||||
|
||||
ggml_tensor * bias = nullptr;
|
||||
int blk;
|
||||
int local;
|
||||
int64_t nblk;
|
||||
ggml_tensor * bias = nullptr; // F32 [nblk, n_tokens] local-force bias (position blocks)
|
||||
ggml_tensor * pos_mask = nullptr; // F32 [n_ps, n_tokens] 0/-inf visibility, by position
|
||||
ggml_tensor * pos_slot_i = nullptr; // I32 [n_ps, ns] pos -> cell (get_rows index)
|
||||
ggml_tensor * pos_slot_f = nullptr; // F32 [n_ps, ns] pos -> cell (gatherable values, decode)
|
||||
ggml_tensor * cell_blk = nullptr; // I32 [n_kv, ns] cell -> position block (batch)
|
||||
|
||||
const llama_kv_cache_msa_context * mctx;
|
||||
|
||||
int blk;
|
||||
int local;
|
||||
};
|
||||
|
||||
// One FA call for all GQA groups (and at multi-stream decode, all streams) by mapping them onto the FA sequence dim (ne[3])
|
||||
@@ -173,7 +213,9 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
||||
inpL = build_inp_embd(model.tok_embd);
|
||||
|
||||
ggml_tensor * inp_pos = build_inp_pos();
|
||||
auto inp_attn = build_attn_inp_kv();
|
||||
|
||||
// ==========================================
|
||||
// TODO: avoid such kind of complexity in the model graphs
|
||||
|
||||
// MSA calls ggml_flash_attn_ext directly and assumes the non-transposed V layout that
|
||||
// llama.cpp only provides when flash attention is enabled. Block selection is anchored
|
||||
@@ -185,6 +227,8 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
||||
const bool streams_ok = cparams.n_seq_max == 1 || !cparams.kv_unified;
|
||||
const bool msa_enabled = fa_on && streams_ok;
|
||||
|
||||
auto * inp_attn = build_attn_inp_kv_msa(msa_enabled);
|
||||
|
||||
static bool warned_no_fa = false;
|
||||
if (!fa_on && !warned_no_fa) {
|
||||
LLAMA_LOG_WARN("%s: flash attention disabled; MSA requires it -> running DENSE attention "
|
||||
@@ -197,36 +241,54 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
||||
"-> running DENSE attention. Output may be degraded. Drop --kv-unified to enable MSA.\n", __func__);
|
||||
warned_unified = true;
|
||||
}
|
||||
// ==========================================
|
||||
|
||||
// hoisted per-graph MSA state (shared by every sparse layer)
|
||||
llm_graph_input_msa_local * msa_loc = nullptr;
|
||||
llm_graph_input_msa * msa = nullptr;
|
||||
ggml_tensor * msa_kqm = nullptr;
|
||||
ggml_tensor * msa_mf = nullptr;
|
||||
int64_t n_kv = 0, nblk = 0, ns = 1, n_tps = 0;
|
||||
ggml_tensor * msa_mf = nullptr; // F32 copy of the FA mask for the final mask add
|
||||
int64_t n_kv = 0, n_ps = 0, nblk = 0, ns = 1, n_tps = 0;
|
||||
bool msa_decode = false; // gather (1 token per stream) vs mask
|
||||
const int blk = mm.msa_p.blk;
|
||||
const int64_t Hd = hparams.indexer_n_head; // one indexer head per GQA group
|
||||
|
||||
if (msa_enabled) {
|
||||
const auto * mctx_msa = static_cast<const llama_kv_cache_msa_context *>(mctx);
|
||||
|
||||
msa_kqm = inp_attn->get_kq_mask();
|
||||
n_kv = msa_kqm->ne[0];
|
||||
n_tps = msa_kqm->ne[1]; // tokens per stream
|
||||
ns = msa_kqm->ne[3]; // streams in this ubatch
|
||||
GGML_ASSERT(msa_kqm->type == GGML_TYPE_F16 && "MSA requires the FA (f16) mask");
|
||||
GGML_ASSERT(n_tps*ns == n_tokens);
|
||||
GGML_ASSERT(n_kv % blk == 0 &&
|
||||
"MSA: KV/mask n_kv must be a multiple of indexer.block_size (128); "
|
||||
"the flash-attention KV padding must be a multiple of the block size. "
|
||||
"A non-multiple would silently drop the partial tail block.");
|
||||
nblk = n_kv / blk;
|
||||
|
||||
// the position axis covers every position currently in the cache and is padded to whole blocks
|
||||
n_ps = GGML_PAD((int64_t) mctx_msa->get_n_pos(), blk);
|
||||
nblk = n_ps / blk;
|
||||
msa_decode = n_tps == 1;
|
||||
|
||||
msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32);
|
||||
auto inp = std::make_unique<llm_graph_input_msa>(mctx_msa, blk, mm.msa_p.local);
|
||||
|
||||
auto loc = std::make_unique<llm_graph_input_msa_local>(blk, mm.msa_p.local, nblk);
|
||||
loc->bias = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, nblk, n_tokens); // stream-grouped tokens
|
||||
ggml_set_input(loc->bias);
|
||||
msa_loc = (llm_graph_input_msa_local *) res->add_input(std::move(loc));
|
||||
inp->bias = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, nblk, n_tokens); // stream-grouped tokens
|
||||
ggml_set_input(inp->bias);
|
||||
|
||||
inp->pos_mask = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_ps, n_tokens);
|
||||
ggml_set_input(inp->pos_mask);
|
||||
|
||||
inp->pos_slot_i = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_ps, ns);
|
||||
ggml_set_input(inp->pos_slot_i);
|
||||
|
||||
if (msa_decode) {
|
||||
inp->pos_slot_f = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_ps, ns);
|
||||
ggml_set_input(inp->pos_slot_f);
|
||||
} else {
|
||||
inp->cell_blk = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, ns);
|
||||
ggml_set_input(inp->cell_blk);
|
||||
|
||||
msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32);
|
||||
}
|
||||
|
||||
msa = (llm_graph_input_msa *) res->add_input(std::move(inp));
|
||||
}
|
||||
|
||||
ggml_tensor * inp_out_ids = build_inp_out_ids();
|
||||
@@ -283,9 +345,11 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
||||
ik = ggml_rope_ext(ctx0, ik, inp_pos, nullptr, n_rot, rope_type, n_ctx_orig,
|
||||
freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow);
|
||||
|
||||
const auto * mctx_cur = inp_attn->mctx;
|
||||
ggml_build_forward_expand(gf, mctx_cur->cpy_k_idx(ctx0, ik, inp_attn->get_k_idxs(), il));
|
||||
ggml_tensor * ik_kv = mctx_cur->get_k_idx(ctx0, il);
|
||||
const auto * mctx_msa_l = static_cast<const llama_kv_cache_msa_context *>(mctx);
|
||||
const auto * mctx_cur = mctx_msa_l->get_base();
|
||||
const auto * mctx_idx = mctx_msa_l->get_idx();
|
||||
ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, ik, inp_attn->get_k_idxs_idx(), il));
|
||||
ggml_tensor * ik_kv = mctx_idx->get_k(ctx0, il);
|
||||
|
||||
if (inp_attn->self_k_rot) {
|
||||
Qcur = llama_mul_mat_hadamard(ctx0, Qcur, inp_attn->self_k_rot);
|
||||
@@ -316,42 +380,52 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
||||
|
||||
if (msa_decode) {
|
||||
// decode: batched over streams top-k + gather, one grouped FA
|
||||
// scores: per-stream batched matmul over the stream dim (ne[3]).
|
||||
// the cache views are not contiguous across streams (stride = kv_size, not n_kv)
|
||||
ggml_tensor * ikv4 = ggml_view_4d(ctx0, ik_kv, n_idx_dim, n_kv, 1, ns,
|
||||
ik_kv->nb[2], ik_kv->nb[3], ik_kv->nb[3], 0);
|
||||
// gather the indexer keys through the pos -> cell map
|
||||
ggml_tensor * ik3 = ggml_view_3d(ctx0, ik_kv, n_idx_dim, n_kv, ns,
|
||||
ik_kv->nb[2], ik_kv->nb[3], 0);
|
||||
ggml_tensor * ikp = ggml_get_rows(ctx0, ik3, msa->pos_slot_i); // [n_idx_dim, n_ps, ns]
|
||||
ggml_tensor * iq4 = ggml_reshape_4d(ctx0, iq, n_idx_dim, Hd, 1, ns);
|
||||
ggml_tensor * sc = ggml_mul_mat(ctx0, ikv4, iq4);
|
||||
ggml_tensor * sc = ggml_mul_mat(ctx0,
|
||||
ggml_reshape_4d(ctx0, ikp, n_idx_dim, n_ps, 1, ns), iq4);
|
||||
ggml_mul_mat_set_prec(sc, GGML_PREC_F32);
|
||||
sc = ggml_add_inplace(ctx0, sc, msa_mf);
|
||||
// unmapped positions come out -inf, so they can never rank into the top-k
|
||||
sc = ggml_add_inplace(ctx0, sc,
|
||||
ggml_reshape_4d(ctx0, msa->pos_mask, n_ps, 1, 1, ns));
|
||||
ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0);
|
||||
cb(bs, "msa_bs", il);
|
||||
|
||||
ggml_tensor * bsf = ggml_add(ctx0, bs,
|
||||
ggml_reshape_4d(ctx0, msa_loc->bias, nblk, 1, 1, ns));
|
||||
ggml_tensor * idx = ggml_top_k(ctx0, bsf, K);
|
||||
ggml_reshape_4d(ctx0, msa->bias, nblk, 1, 1, ns));
|
||||
ggml_tensor * idx = ggml_top_k(ctx0, bsf, K); // position blocks
|
||||
|
||||
// token idx: tj[t,k,h,s] = blk*idx[k,h,s] + t (for the mask gather)
|
||||
// row idx: tr[t,k,h,s] = tj*HKV + h (for the per-stream K/V gather)
|
||||
// pos idx: tj[t,k,h,s] = blk*idx[k,h,s] + t (positions - mask gather)
|
||||
// cell idx: cs[t,k,h,s] = pos_slot[tj] (pos -> cell translation)
|
||||
// row idx: tr[t,k,h,s] = cs*HKV + h (per-stream K/V gather)
|
||||
ggml_tensor * a = ggml_scale(ctx0, ggml_cast(ctx0, idx, GGML_TYPE_F32), (float) blk);
|
||||
a = ggml_reshape_4d(ctx0, a, 1, K, Hd, ns);
|
||||
ggml_tensor * tj = ggml_add(ctx0,
|
||||
ggml_repeat_4d(ctx0, a, blk, K, Hd, ns),
|
||||
ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) blk, 1.0f), blk, 1, 1));
|
||||
ggml_tensor * tr = ggml_add(ctx0,
|
||||
ggml_scale(ctx0, tj, (float) HKV),
|
||||
ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) HKV, 1.0f), 1, 1, Hd));
|
||||
|
||||
ggml_tensor * tokj = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tj, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32);
|
||||
|
||||
ggml_tensor * cs = ggml_get_rows(ctx0,
|
||||
ggml_reshape_3d(ctx0, msa->pos_slot_f, 1, n_ps, ns), tokj); // [1, blk*K*Hd, ns]
|
||||
cs = ggml_reshape_4d(ctx0, cs, blk, K, Hd, ns);
|
||||
|
||||
ggml_tensor * tr = ggml_add(ctx0,
|
||||
ggml_scale(ctx0, cs, (float) HKV),
|
||||
ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) HKV, 1.0f), 1, 1, Hd));
|
||||
|
||||
ggml_tensor * tokr = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tr, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32);
|
||||
|
||||
ggml_tensor * k3 = ggml_view_3d(ctx0, k, D, HKV*n_kv, ns, k->nb[1], k->nb[3], 0);
|
||||
ggml_tensor * v3 = ggml_view_3d(ctx0, v, D, HKV*n_kv, ns, v->nb[1], v->nb[3], 0);
|
||||
ggml_tensor * m3 = ggml_reshape_3d(ctx0, msa_kqm, 1, n_kv, ns);
|
||||
ggml_tensor * mp = ggml_reshape_3d(ctx0, msa->pos_mask, 1, n_ps, ns);
|
||||
|
||||
ggml_tensor * kg = ggml_get_rows(ctx0, k3, tokr);
|
||||
ggml_tensor * vg = ggml_get_rows(ctx0, v3, tokr);
|
||||
ggml_tensor * mg = ggml_get_rows(ctx0, m3, tokj);
|
||||
ggml_tensor * mg = ggml_get_rows(ctx0, mp, tokj);
|
||||
|
||||
// fold (group, stream) onto the FA channel dim
|
||||
const ggml_type kt = ggml_is_quantized(k->type) ? GGML_TYPE_F16 : k->type;
|
||||
@@ -372,12 +446,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
||||
iq->nb[1], iq->nb[2], st*n_tps*iq->nb[2]);
|
||||
ggml_tensor * ik_s = ggml_view_2d(ctx0, ik_kv, n_idx_dim, n_kv,
|
||||
ik_kv->nb[2], st*ik_kv->nb[3]);
|
||||
ggml_tensor * mf_s = ggml_view_3d(ctx0, msa_mf, n_kv, 1, n_tps,
|
||||
msa_mf->nb[1], msa_mf->nb[1], st*msa_mf->nb[3]);
|
||||
ggml_tensor * km_s = ggml_view_3d(ctx0, msa_kqm, n_kv, n_tps, 1,
|
||||
msa_kqm->nb[1], msa_kqm->nb[3], st*msa_kqm->nb[3]);
|
||||
ggml_tensor * bias_s = ggml_view_3d(ctx0, msa_loc->bias, nblk, 1, n_tps,
|
||||
msa_loc->bias->nb[1], msa_loc->bias->nb[1], st*n_tps*msa_loc->bias->nb[1]);
|
||||
ggml_tensor * psl_s = ggml_view_1d(ctx0, msa->pos_slot_i, n_ps,
|
||||
st*msa->pos_slot_i->nb[1]);
|
||||
ggml_tensor * pm_s = ggml_view_3d(ctx0, msa->pos_mask, n_ps, 1, n_tps,
|
||||
msa->pos_mask->nb[1], msa->pos_mask->nb[1], st*n_tps*msa->pos_mask->nb[1]);
|
||||
ggml_tensor * cb_s = ggml_view_1d(ctx0, msa->cell_blk, n_kv,
|
||||
st*msa->cell_blk->nb[1]);
|
||||
ggml_tensor * mf_s = ggml_view_3d(ctx0, msa_mf, n_kv, n_tps, 1,
|
||||
msa_mf->nb[1], msa_mf->nb[3], st*msa_mf->nb[3]);
|
||||
ggml_tensor * bias_s = ggml_view_3d(ctx0, msa->bias, nblk, 1, n_tps,
|
||||
msa->bias->nb[1], msa->bias->nb[1], st*n_tps*msa->bias->nb[1]);
|
||||
ggml_tensor * q_s = ggml_view_3d(ctx0, Qcur, D, n_head, n_tps,
|
||||
Qcur->nb[1], Qcur->nb[2], st*n_tps*Qcur->nb[2]);
|
||||
ggml_tensor * k_s = ggml_view_4d(ctx0, k, D, HKV, n_kv, 1,
|
||||
@@ -385,14 +463,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
||||
ggml_tensor * v_s = ggml_view_4d(ctx0, v, D, HKV, n_kv, 1,
|
||||
v->nb[1], v->nb[2], v->nb[3], st*v->nb[3]);
|
||||
|
||||
// block scores: bs = maxpool_blk(idx_q * idx_k^T + causal mask)
|
||||
// block scores: the indexer keys are gathered through the pos -> cell map first
|
||||
// scores are unscaled, only the top-k ordering matters
|
||||
ggml_tensor * sc = ggml_mul_mat(ctx0, ik_s,
|
||||
ggml_tensor * ikp = ggml_get_rows(ctx0, ik_s, psl_s); // [n_idx_dim, n_ps]
|
||||
ggml_tensor * sc = ggml_mul_mat(ctx0, ikp,
|
||||
ggml_reshape_2d(ctx0, iq_s, n_idx_dim, Hd*n_tps));
|
||||
// indexer scores run in F32
|
||||
ggml_mul_mat_set_prec(sc, GGML_PREC_F32);
|
||||
sc = ggml_reshape_3d(ctx0, sc, n_kv, Hd, n_tps);
|
||||
sc = ggml_add_inplace(ctx0, sc, mf_s);
|
||||
sc = ggml_reshape_3d(ctx0, sc, n_ps, Hd, n_tps);
|
||||
// unmapped positions (holes, padding, empty cells) come out -inf
|
||||
sc = ggml_add_inplace(ctx0, sc, pm_s);
|
||||
ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0);
|
||||
cb(bs, "msa_bs", il);
|
||||
|
||||
@@ -416,14 +496,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_
|
||||
bm = ggml_cont(ctx0, ggml_permute(ctx0, bm, 0, 2, 1, 3)); // [nblk, n_tps, Hd]
|
||||
cb(bm, "msa_block_mask", il);
|
||||
|
||||
// expand block -> token granularity (j = bk*blk + t),
|
||||
// then combine with the causal mask in place
|
||||
ggml_tensor * bmx = ggml_repeat_4d(ctx0,
|
||||
ggml_reshape_3d(ctx0, bm, 1, nblk, n_tps*Hd),
|
||||
blk, nblk, n_tps*Hd, 1);
|
||||
// expand block -> cell granularity through the cell -> position block
|
||||
// map, then combine with the causal mask. empty cells are masked by the causal mask.
|
||||
ggml_tensor * bm2 = ggml_cont(ctx0, ggml_transpose(ctx0,
|
||||
ggml_reshape_2d(ctx0, bm, nblk, n_tps*Hd))); // [n_tps*Hd, nblk]
|
||||
ggml_tensor * bmc = ggml_get_rows(ctx0, bm2, cb_s); // [n_tps*Hd, n_kv] F32
|
||||
ggml_tensor * bmx = ggml_cont(ctx0, ggml_transpose(ctx0, bmc));
|
||||
bmx = ggml_reshape_3d(ctx0, bmx, n_kv, n_tps, Hd);
|
||||
ggml_tensor * mask4 = ggml_add_inplace(ctx0, bmx, km_s);
|
||||
mask4 = ggml_reshape_4d(ctx0, mask4, n_kv, n_tps, 1, Hd);
|
||||
ggml_tensor * mask4 = ggml_add_inplace(ctx0, bmx, mf_s);
|
||||
mask4 = ggml_cast(ctx0,
|
||||
ggml_reshape_4d(ctx0, mask4, n_kv, n_tps, 1, Hd), GGML_TYPE_F16);
|
||||
cb(mask4, "msa_mask4", il);
|
||||
|
||||
// cache views with groups on ne[3];
|
||||
|
||||
@@ -1084,6 +1084,10 @@ struct llama_model_deepseek2 : public llama_model_base {
|
||||
graph(const llama_model & model, const llm_graph_params & params);
|
||||
};
|
||||
|
||||
struct graph_mtp : public llm_graph_context {
|
||||
graph_mtp(const llama_model & model, const llm_graph_params & params);
|
||||
};
|
||||
|
||||
std::unique_ptr<llm_graph_context> build_arch_graph(const llm_graph_params & params) const override;
|
||||
};
|
||||
|
||||
|
||||
@@ -258,6 +258,9 @@ llama_build_and_test(test-thread-safety.cpp ARGS -m "${MODEL_DEST}" -ngl 99 -p "
|
||||
set_tests_properties(test-thread-safety PROPERTIES FIXTURES_REQUIRED test-download-model)
|
||||
|
||||
llama_build_and_test(test-arg-parser.cpp)
|
||||
llama_build_and_test(test-model-resolution.cpp)
|
||||
# the test serves its repos from an httplib server, and the library links it privately
|
||||
target_link_libraries(test-model-resolution PRIVATE cpp-httplib)
|
||||
|
||||
if (NOT LLAMA_SANITIZE_ADDRESS AND NOT GGML_SCHED_NO_REALLOC)
|
||||
# TODO: repair known memory leaks
|
||||
|
||||
@@ -99,6 +99,34 @@ static void test(void) {
|
||||
argv = {"binary_name", "-sm", "hello"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON));
|
||||
|
||||
{
|
||||
common_params penalty_params;
|
||||
|
||||
argv = {"binary_name", "--repeat-penalty", "0"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
|
||||
|
||||
argv = {"binary_name", "--repeat-penalty", "-1"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
|
||||
|
||||
argv = {"binary_name", "--repeat-penalty", "nan"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
|
||||
|
||||
argv = {"binary_name", "--repeat-penalty", "inf"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
|
||||
|
||||
argv = {"binary_name", "--repeat-penalty", "-inf"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
|
||||
|
||||
const char * penalty_options[] = {"--frequency-penalty", "--presence-penalty"};
|
||||
const char * nonfinite_values[] = {"nan", "inf", "-inf"};
|
||||
for (const char * option : penalty_options) {
|
||||
for (const char * value : nonfinite_values) {
|
||||
argv = {"binary_name", option, value};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// non-existence arg in specific example (--draft cannot be used outside llama-speculative)
|
||||
argv = {"binary_name", "--draft", "123"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_EMBEDDING));
|
||||
|
||||
@@ -8,12 +8,15 @@
|
||||
#endif
|
||||
|
||||
#include <algorithm>
|
||||
#include <cmath>
|
||||
#include <cstdlib>
|
||||
#include <cstring>
|
||||
#include <fstream>
|
||||
#include <functional>
|
||||
#include <map>
|
||||
#include <string>
|
||||
#include <unordered_map>
|
||||
#include <unordered_set>
|
||||
#include <vector>
|
||||
|
||||
struct test_args {
|
||||
@@ -761,6 +764,563 @@ static void test_backend_logit_bias_sampling(const test_params & params) {
|
||||
printf("backend logit bias sampling test PASSED\n");
|
||||
}
|
||||
|
||||
static void accept_prompt(llama_sampler * smpl, const llama_vocab * vocab, const std::string & prompt) {
|
||||
const llama_token bos = llama_vocab_bos(vocab);
|
||||
if (bos != LLAMA_TOKEN_NULL) {
|
||||
llama_sampler_accept(smpl, bos);
|
||||
}
|
||||
|
||||
std::vector<llama_token> tokens(64);
|
||||
int32_t n_tokens = llama_tokenize(vocab, prompt.c_str(), (int32_t) prompt.size(),
|
||||
tokens.data(), (int32_t) tokens.size(), false, false);
|
||||
if (n_tokens < 0) {
|
||||
tokens.resize(-n_tokens);
|
||||
n_tokens = llama_tokenize(vocab, prompt.c_str(), (int32_t) prompt.size(),
|
||||
tokens.data(), (int32_t) tokens.size(), false, false);
|
||||
}
|
||||
|
||||
for (int32_t i = 0; i < n_tokens; ++i) {
|
||||
llama_sampler_accept(smpl, tokens[i]);
|
||||
}
|
||||
}
|
||||
|
||||
static std::vector<float> decode_raw_logits(const test_params & params, const std::string & prompt) {
|
||||
const int seq_id = 0;
|
||||
const int n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(params.model.get()));
|
||||
std::vector<llama_sampler_seq_config> empty_configs;
|
||||
test_context ctx(params, empty_configs);
|
||||
|
||||
GGML_ASSERT(ctx.decode({{ seq_id, prompt }}));
|
||||
|
||||
float * logits = llama_get_logits_ith(ctx.ctx.get(), ctx.idx_for_seq(seq_id));
|
||||
GGML_ASSERT(logits != nullptr);
|
||||
return std::vector<float>(logits, logits + n_vocab);
|
||||
}
|
||||
|
||||
static std::vector<llama_token_data> apply_cpu_sampler(
|
||||
const std::vector<float> & raw_logits,
|
||||
llama_sampler * sampler) {
|
||||
std::vector<llama_token_data> data;
|
||||
data.reserve(raw_logits.size());
|
||||
for (llama_token token = 0; token < (llama_token) raw_logits.size(); ++token) {
|
||||
data.push_back({ token, raw_logits[token], 0.0f });
|
||||
}
|
||||
|
||||
llama_token_data_array cur_p = { data.data(), data.size(), -1, false };
|
||||
llama_sampler_apply(sampler, &cur_p);
|
||||
data.resize(cur_p.size);
|
||||
return data;
|
||||
}
|
||||
|
||||
using sampler_setup_fn = std::function<void(llama_sampler *)>;
|
||||
using sampler_init_fn = std::function<llama_sampler *()>;
|
||||
|
||||
enum class penalties_position {
|
||||
before_filter,
|
||||
after_filter,
|
||||
};
|
||||
|
||||
static void add_filter_and_penalties(
|
||||
llama_sampler * chain,
|
||||
const sampler_init_fn & init_filter,
|
||||
int32_t penalty_last_n,
|
||||
float penalty_repeat,
|
||||
float penalty_freq,
|
||||
float penalty_present,
|
||||
penalties_position position) {
|
||||
const auto add_penalties = [&]() {
|
||||
llama_sampler_chain_add(chain, llama_sampler_init_penalties(
|
||||
penalty_last_n, penalty_repeat, penalty_freq, penalty_present));
|
||||
};
|
||||
|
||||
if (position == penalties_position::before_filter) {
|
||||
add_penalties();
|
||||
llama_sampler_chain_add(chain, init_filter());
|
||||
} else {
|
||||
llama_sampler_chain_add(chain, init_filter());
|
||||
add_penalties();
|
||||
}
|
||||
}
|
||||
|
||||
static llama_sampler_ptr make_sampler_chain(
|
||||
const sampler_setup_fn & add_samplers,
|
||||
const sampler_setup_fn & accept_history) {
|
||||
llama_sampler_ptr chain(llama_sampler_chain_init(llama_sampler_chain_default_params()));
|
||||
add_samplers(chain.get());
|
||||
accept_history(chain.get());
|
||||
return chain;
|
||||
}
|
||||
|
||||
struct backend_sampler_output {
|
||||
std::vector<float> logits;
|
||||
std::vector<llama_token> candidates;
|
||||
};
|
||||
|
||||
static backend_sampler_output run_backend_sampler(
|
||||
const test_params & params,
|
||||
const std::string & prompt,
|
||||
llama_sampler * sampler) {
|
||||
const int seq_id = 0;
|
||||
std::vector<llama_sampler_seq_config> configs = {{ seq_id, sampler }};
|
||||
test_context ctx(params, configs);
|
||||
|
||||
GGML_ASSERT(ctx.decode({{ seq_id, prompt }}));
|
||||
llama_synchronize(ctx.ctx.get());
|
||||
|
||||
const int32_t idx = ctx.idx_for_seq(seq_id);
|
||||
const uint32_t n_logits = llama_get_sampled_logits_count_ith(ctx.ctx.get(), idx);
|
||||
const uint32_t n_candidates = llama_get_sampled_candidates_count_ith(ctx.ctx.get(), idx);
|
||||
float * logits = llama_get_sampled_logits_ith(ctx.ctx.get(), idx);
|
||||
llama_token * candidates = llama_get_sampled_candidates_ith(ctx.ctx.get(), idx);
|
||||
GGML_ASSERT(logits != nullptr);
|
||||
|
||||
backend_sampler_output result;
|
||||
result.logits.assign(logits, logits + n_logits);
|
||||
result.candidates.resize(n_logits);
|
||||
|
||||
if (n_candidates == 0) {
|
||||
for (uint32_t i = 0; i < n_logits; ++i) {
|
||||
result.candidates[i] = (llama_token) i;
|
||||
}
|
||||
} else {
|
||||
GGML_ASSERT(candidates != nullptr);
|
||||
GGML_ASSERT(n_candidates == n_logits);
|
||||
std::memcpy(result.candidates.data(), candidates, n_candidates * sizeof(llama_token));
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
struct sampler_comparison_output {
|
||||
std::vector<llama_token_data> expected;
|
||||
backend_sampler_output actual;
|
||||
};
|
||||
|
||||
static sampler_comparison_output run_sampler_comparison(
|
||||
const test_params & params,
|
||||
const std::string & prompt,
|
||||
const std::vector<float> & raw_logits,
|
||||
const sampler_setup_fn & add_samplers,
|
||||
const sampler_setup_fn & accept_history) {
|
||||
llama_sampler_ptr cpu_chain = make_sampler_chain(add_samplers, accept_history);
|
||||
llama_sampler_ptr backend_chain = make_sampler_chain(add_samplers, accept_history);
|
||||
return {
|
||||
apply_cpu_sampler(raw_logits, cpu_chain.get()),
|
||||
run_backend_sampler(params, prompt, backend_chain.get()),
|
||||
};
|
||||
}
|
||||
|
||||
static std::unordered_map<llama_token, float> map_logits(const std::vector<llama_token_data> & data) {
|
||||
std::unordered_map<llama_token, float> result;
|
||||
result.reserve(data.size());
|
||||
for (const auto & item : data) {
|
||||
result[item.id] = item.logit;
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
struct sampler_comparison_stats {
|
||||
int n_mismatch = 0;
|
||||
int n_masked = 0;
|
||||
float max_diff = 0.0f;
|
||||
};
|
||||
|
||||
static sampler_comparison_stats compare_sampler_outputs(
|
||||
const char * name,
|
||||
const std::unordered_map<llama_token, float> & expected,
|
||||
const backend_sampler_output & actual,
|
||||
bool allow_extra_candidates = false) {
|
||||
GGML_ASSERT(actual.logits.size() == actual.candidates.size());
|
||||
|
||||
sampler_comparison_stats result;
|
||||
std::unordered_set<llama_token> seen;
|
||||
seen.reserve(actual.candidates.size());
|
||||
|
||||
for (size_t i = 0; i < actual.logits.size(); ++i) {
|
||||
const llama_token token = actual.candidates[i];
|
||||
const float logit = actual.logits[i];
|
||||
if (!seen.insert(token).second || std::isnan(logit)) {
|
||||
if (result.n_mismatch < 5) {
|
||||
printf("%s token %d has invalid backend output\n", name, token);
|
||||
}
|
||||
++result.n_mismatch;
|
||||
continue;
|
||||
}
|
||||
|
||||
const auto it = expected.find(token);
|
||||
if (it == expected.end()) {
|
||||
if (std::isinf(logit) && logit < 0.0f) {
|
||||
++result.n_masked;
|
||||
} else if (!allow_extra_candidates) {
|
||||
if (result.n_mismatch < 5) {
|
||||
printf("%s token %d was not masked\n", name, token);
|
||||
}
|
||||
++result.n_mismatch;
|
||||
}
|
||||
continue;
|
||||
}
|
||||
|
||||
const float diff = fabsf(it->second - logit);
|
||||
result.max_diff = std::max(result.max_diff, diff);
|
||||
if (!std::isfinite(logit) || diff > 1e-3f) {
|
||||
if (result.n_mismatch < 5) {
|
||||
printf("%s mismatch token %d: cpu=%.6f backend=%.6f diff=%.6f\n",
|
||||
name, token, it->second, logit, diff);
|
||||
}
|
||||
++result.n_mismatch;
|
||||
}
|
||||
}
|
||||
|
||||
for (const auto & item : expected) {
|
||||
if (seen.find(item.first) == seen.end()) {
|
||||
if (result.n_mismatch < 5) {
|
||||
printf("%s missing backend token %d\n", name, item.first);
|
||||
}
|
||||
++result.n_mismatch;
|
||||
}
|
||||
}
|
||||
|
||||
printf("%s logits: max_diff=%.6f n_masked=%d n_mismatch=%d\n",
|
||||
name, result.max_diff, result.n_masked, result.n_mismatch);
|
||||
return result;
|
||||
}
|
||||
|
||||
static float find_backend_logit(const backend_sampler_output & output, llama_token token) {
|
||||
for (size_t i = 0; i < output.candidates.size(); ++i) {
|
||||
if (output.candidates[i] == token) {
|
||||
return output.logits[i];
|
||||
}
|
||||
}
|
||||
GGML_ABORT("backend token not found");
|
||||
}
|
||||
|
||||
static sampler_comparison_output run_penalties_comparison(
|
||||
const test_params & params,
|
||||
int32_t penalty_last_n,
|
||||
float penalty_repeat,
|
||||
float penalty_freq,
|
||||
float penalty_present,
|
||||
const std::string & prompt,
|
||||
const std::function<void(llama_sampler *)> & extra_accept = {}) {
|
||||
const auto * vocab = llama_model_get_vocab(params.model.get());
|
||||
const std::vector<float> raw_logits = decode_raw_logits(params, prompt);
|
||||
const auto add_samplers = [&](llama_sampler * chain) {
|
||||
llama_sampler_chain_add(chain, llama_sampler_init_penalties(
|
||||
penalty_last_n, penalty_repeat, penalty_freq, penalty_present));
|
||||
};
|
||||
const auto accept_history = [&](llama_sampler * chain) {
|
||||
accept_prompt(chain, vocab, prompt);
|
||||
if (extra_accept) {
|
||||
extra_accept(chain);
|
||||
}
|
||||
};
|
||||
|
||||
return run_sampler_comparison(
|
||||
params, prompt, raw_logits, add_samplers, accept_history);
|
||||
}
|
||||
|
||||
static void compare_penalties_logits(
|
||||
const test_params & params,
|
||||
int32_t penalty_last_n,
|
||||
float penalty_repeat,
|
||||
float penalty_freq,
|
||||
float penalty_present,
|
||||
const std::string & prompt,
|
||||
const std::function<void(llama_sampler *)> & extra_accept = {}) {
|
||||
const sampler_comparison_output output = run_penalties_comparison(
|
||||
params, penalty_last_n, penalty_repeat, penalty_freq, penalty_present, prompt, extra_accept);
|
||||
|
||||
GGML_ASSERT(output.expected.size() == output.actual.logits.size());
|
||||
|
||||
const sampler_comparison_stats stats = compare_sampler_outputs(
|
||||
"penalties", map_logits(output.expected), output.actual);
|
||||
GGML_ASSERT(stats.n_masked == 0);
|
||||
GGML_ASSERT(stats.n_mismatch == 0);
|
||||
}
|
||||
|
||||
static void test_penalty_parameter_values(const test_params & params) {
|
||||
struct penalty_test_case {
|
||||
const char * name;
|
||||
float repeat;
|
||||
float frequency;
|
||||
float presence;
|
||||
};
|
||||
|
||||
const penalty_test_case cases[] = {
|
||||
{ "frequency -1", 1.0f, -1.0f, 0.0f },
|
||||
{ "frequency 0", 1.0f, 0.0f, 0.0f },
|
||||
{ "frequency 1", 1.0f, 1.0f, 0.0f },
|
||||
{ "presence -1", 1.0f, 0.0f, -1.0f },
|
||||
{ "presence 0", 1.0f, 0.0f, 0.0f },
|
||||
{ "presence 1", 1.0f, 0.0f, 1.0f },
|
||||
{ "repeat 1", 1.0f, 0.0f, 0.0f },
|
||||
};
|
||||
|
||||
int n_failed = 0;
|
||||
for (const auto & test : cases) {
|
||||
const sampler_comparison_output output = run_penalties_comparison(
|
||||
params, 64, test.repeat, test.frequency, test.presence, "Hello Hello world");
|
||||
GGML_ASSERT(output.expected.size() == output.actual.logits.size());
|
||||
const sampler_comparison_stats stats = compare_sampler_outputs(
|
||||
test.name, map_logits(output.expected), output.actual);
|
||||
n_failed += stats.n_mismatch != 0;
|
||||
}
|
||||
|
||||
GGML_ASSERT(n_failed == 0);
|
||||
}
|
||||
|
||||
static void compare_top_k_penalties_logits(
|
||||
const test_params & params,
|
||||
int32_t k,
|
||||
int32_t penalty_last_n,
|
||||
float penalty_repeat,
|
||||
float penalty_freq,
|
||||
float penalty_present,
|
||||
const std::string & prompt,
|
||||
penalties_position position) {
|
||||
const auto * vocab = llama_model_get_vocab(params.model.get());
|
||||
const std::vector<float> raw_logits = decode_raw_logits(params, prompt);
|
||||
const int n_vocab = (int) raw_logits.size();
|
||||
|
||||
GGML_ASSERT(n_vocab > k);
|
||||
|
||||
const sampler_init_fn init_top_k = [k]() {
|
||||
return llama_sampler_init_top_k(k);
|
||||
};
|
||||
llama_sampler_ptr top_k(init_top_k());
|
||||
const std::vector<llama_token_data> top_k_data = apply_cpu_sampler(raw_logits, top_k.get());
|
||||
GGML_ASSERT(top_k_data.size() == (size_t) k);
|
||||
const llama_token retained_history_token = top_k_data[0].id;
|
||||
|
||||
llama_token excluded_history_token = LLAMA_TOKEN_NULL;
|
||||
for (llama_token token = 0; token < n_vocab; ++token) {
|
||||
const auto it = std::find_if(top_k_data.begin(), top_k_data.end(), [token](const llama_token_data & data) {
|
||||
return data.id == token;
|
||||
});
|
||||
if (it == top_k_data.end()) {
|
||||
excluded_history_token = token;
|
||||
break;
|
||||
}
|
||||
}
|
||||
GGML_ASSERT(excluded_history_token != LLAMA_TOKEN_NULL);
|
||||
|
||||
const auto add_samplers = [&](llama_sampler * chain) {
|
||||
add_filter_and_penalties(chain, init_top_k,
|
||||
penalty_last_n, penalty_repeat, penalty_freq, penalty_present, position);
|
||||
};
|
||||
|
||||
auto accept_history = [&](llama_sampler * smpl) {
|
||||
accept_prompt(smpl, vocab, prompt);
|
||||
llama_sampler_accept(smpl, excluded_history_token);
|
||||
llama_sampler_accept(smpl, excluded_history_token);
|
||||
llama_sampler_accept(smpl, retained_history_token);
|
||||
llama_sampler_accept(smpl, retained_history_token);
|
||||
};
|
||||
|
||||
const sampler_comparison_output output = run_sampler_comparison(
|
||||
params, prompt, raw_logits, add_samplers, accept_history);
|
||||
|
||||
GGML_ASSERT(output.expected.size() == (size_t) k);
|
||||
GGML_ASSERT(output.actual.logits.size() == (size_t) k);
|
||||
|
||||
const std::unordered_map<llama_token, float> expected_logits = map_logits(output.expected);
|
||||
|
||||
if (position == penalties_position::after_filter) {
|
||||
GGML_ASSERT(expected_logits.find(retained_history_token) != expected_logits.end());
|
||||
GGML_ASSERT(fabsf(expected_logits.at(retained_history_token) - raw_logits[retained_history_token]) > 1e-6f);
|
||||
GGML_ASSERT(expected_logits.find(excluded_history_token) == expected_logits.end());
|
||||
GGML_ASSERT(std::find(output.actual.candidates.begin(), output.actual.candidates.end(),
|
||||
excluded_history_token) == output.actual.candidates.end());
|
||||
} else {
|
||||
const std::unordered_map<llama_token, float> unpenalized_logits = map_logits(top_k_data);
|
||||
bool changed = false;
|
||||
for (const auto & item : expected_logits) {
|
||||
const auto it = unpenalized_logits.find(item.first);
|
||||
if (it == unpenalized_logits.end() || fabsf(it->second - item.second) > 1e-6f) {
|
||||
changed = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
GGML_ASSERT(changed);
|
||||
}
|
||||
|
||||
const char * name = position == penalties_position::before_filter
|
||||
? "penalties top-k"
|
||||
: "top-k penalties";
|
||||
const sampler_comparison_stats stats = compare_sampler_outputs(
|
||||
name, expected_logits, output.actual);
|
||||
GGML_ASSERT(stats.n_masked == 0);
|
||||
GGML_ASSERT(stats.n_mismatch == 0);
|
||||
}
|
||||
|
||||
static void compare_masking_penalties_logits(
|
||||
const test_params & params,
|
||||
const char * filter_name,
|
||||
const sampler_init_fn & init_filter,
|
||||
int32_t penalty_last_n,
|
||||
float penalty_repeat,
|
||||
float penalty_freq,
|
||||
float penalty_present,
|
||||
const std::string & prompt,
|
||||
penalties_position position,
|
||||
bool allow_extra_candidates,
|
||||
bool add_history = true) {
|
||||
const auto * vocab = llama_model_get_vocab(params.model.get());
|
||||
const std::vector<float> raw_logits = decode_raw_logits(params, prompt);
|
||||
const int n_vocab = (int) raw_logits.size();
|
||||
llama_sampler_ptr filter(init_filter());
|
||||
const std::vector<llama_token_data> filtered_data = apply_cpu_sampler(raw_logits, filter.get());
|
||||
GGML_ASSERT(!filtered_data.empty());
|
||||
GGML_ASSERT(filtered_data.size() < (size_t) n_vocab);
|
||||
|
||||
const llama_token penalized_token = filtered_data[0].id;
|
||||
std::unordered_set<llama_token> retained_tokens;
|
||||
retained_tokens.reserve(filtered_data.size());
|
||||
for (const auto & data : filtered_data) {
|
||||
retained_tokens.insert(data.id);
|
||||
}
|
||||
|
||||
llama_token masked_token = LLAMA_TOKEN_NULL;
|
||||
for (llama_token token = 0; token < n_vocab; ++token) {
|
||||
if (retained_tokens.find(token) == retained_tokens.end()) {
|
||||
masked_token = token;
|
||||
break;
|
||||
}
|
||||
}
|
||||
GGML_ASSERT(masked_token != LLAMA_TOKEN_NULL);
|
||||
|
||||
const auto add_samplers = [&](llama_sampler * chain) {
|
||||
add_filter_and_penalties(chain, init_filter,
|
||||
penalty_last_n, penalty_repeat, penalty_freq, penalty_present, position);
|
||||
};
|
||||
auto accept_history = [&](llama_sampler * smpl) {
|
||||
if (!add_history) {
|
||||
return;
|
||||
}
|
||||
accept_prompt(smpl, vocab, prompt);
|
||||
llama_sampler_accept(smpl, penalized_token);
|
||||
llama_sampler_accept(smpl, penalized_token);
|
||||
llama_sampler_accept(smpl, masked_token);
|
||||
llama_sampler_accept(smpl, masked_token);
|
||||
};
|
||||
|
||||
const sampler_comparison_output output = run_sampler_comparison(
|
||||
params, prompt, raw_logits, add_samplers, accept_history);
|
||||
|
||||
GGML_ASSERT(output.actual.logits.size() == (size_t) n_vocab);
|
||||
|
||||
const std::unordered_map<llama_token, float> expected_logits = map_logits(output.expected);
|
||||
|
||||
GGML_ASSERT(expected_logits.find(masked_token) == expected_logits.end());
|
||||
if (add_history) {
|
||||
if (position == penalties_position::after_filter) {
|
||||
GGML_ASSERT(expected_logits.find(penalized_token) != expected_logits.end());
|
||||
GGML_ASSERT(fabsf(expected_logits.at(penalized_token) - raw_logits[penalized_token]) > 1e-6f);
|
||||
} else {
|
||||
llama_sampler_ptr penalties(llama_sampler_init_penalties(
|
||||
penalty_last_n, penalty_repeat, penalty_freq, penalty_present));
|
||||
accept_history(penalties.get());
|
||||
const std::unordered_map<llama_token, float> penalized_logits =
|
||||
map_logits(apply_cpu_sampler(raw_logits, penalties.get()));
|
||||
GGML_ASSERT(fabsf(penalized_logits.at(penalized_token) - raw_logits[penalized_token]) > 1e-6f);
|
||||
}
|
||||
}
|
||||
|
||||
const std::string name = position == penalties_position::before_filter
|
||||
? "penalties " + std::string(filter_name)
|
||||
: std::string(filter_name) + " penalties";
|
||||
const sampler_comparison_stats stats = compare_sampler_outputs(
|
||||
name.c_str(), expected_logits, output.actual, allow_extra_candidates);
|
||||
const float masked_logit = find_backend_logit(output.actual, masked_token);
|
||||
GGML_ASSERT(stats.n_masked > 0);
|
||||
GGML_ASSERT(std::isinf(masked_logit) && masked_logit < 0.0f);
|
||||
GGML_ASSERT(stats.n_mismatch == 0);
|
||||
}
|
||||
|
||||
static void test_backend_penalties_sampling(const test_params & params) {
|
||||
printf("Testing backend penalties (repeat + freq + presence)\n");
|
||||
compare_penalties_logits(params, 64, 1.1f, 0.5f, 0.25f, "Hello Hello world");
|
||||
|
||||
printf("Testing backend penalties with penalty_last_n > 64\n");
|
||||
const auto * vocab = llama_model_get_vocab(params.model.get());
|
||||
std::vector<llama_token> tokens(8);
|
||||
int32_t n_tok = llama_tokenize(vocab, "a", 1, tokens.data(), (int32_t) tokens.size(), false, false);
|
||||
if (n_tok < 0) {
|
||||
tokens.resize(-n_tok);
|
||||
n_tok = llama_tokenize(vocab, "a", 1, tokens.data(), (int32_t) tokens.size(), false, false);
|
||||
}
|
||||
GGML_ASSERT(n_tok > 0);
|
||||
const llama_token tok = tokens[0];
|
||||
|
||||
compare_penalties_logits(params, 80, 1.15f, 0.1f, 0.05f, "a", [tok](llama_sampler * smpl) {
|
||||
// accept_prompt already accepted BOS + one 'a'; fill the ring to n=80
|
||||
for (int i = 0; i < 78; ++i) {
|
||||
llama_sampler_accept(smpl, tok);
|
||||
}
|
||||
});
|
||||
|
||||
printf("Testing backend penalties without filler entries\n");
|
||||
compare_penalties_logits(params, 64, 1.1f, 0.5f, 0.25f, "Hello", [](llama_sampler * smpl) {
|
||||
for (llama_token token = 0; token < 64; ++token) {
|
||||
llama_sampler_accept(smpl, token);
|
||||
}
|
||||
});
|
||||
|
||||
printf("Testing backend top-k followed by penalties\n");
|
||||
compare_top_k_penalties_logits(params, 8, 64, 1.1f, 0.5f, 0.25f, "Hello",
|
||||
penalties_position::after_filter);
|
||||
|
||||
printf("Testing backend penalties followed by top-k\n");
|
||||
compare_top_k_penalties_logits(params, 8, 64, 1.1f, 0.5f, 0.25f, "Hello",
|
||||
penalties_position::before_filter);
|
||||
|
||||
printf("Testing backend top-p followed by penalties\n");
|
||||
compare_masking_penalties_logits(params, "top-p", []() {
|
||||
return llama_sampler_init_top_p(0.9f, 0);
|
||||
}, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true);
|
||||
|
||||
printf("Testing backend top-p followed by penalties with a large history window\n");
|
||||
compare_masking_penalties_logits(params, "top-p large-window", []() {
|
||||
return llama_sampler_init_top_p(0.9f, 0);
|
||||
}, 4096, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true);
|
||||
|
||||
printf("Testing backend penalties followed by top-p\n");
|
||||
compare_masking_penalties_logits(params, "top-p", []() {
|
||||
return llama_sampler_init_top_p(0.9f, 0);
|
||||
}, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::before_filter, true);
|
||||
|
||||
printf("Testing backend min-p followed by penalties\n");
|
||||
compare_masking_penalties_logits(params, "min-p", []() {
|
||||
return llama_sampler_init_min_p(0.1f, 0);
|
||||
}, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, false);
|
||||
|
||||
printf("Testing backend penalties followed by min-p\n");
|
||||
compare_masking_penalties_logits(params, "min-p", []() {
|
||||
return llama_sampler_init_min_p(0.1f, 0);
|
||||
}, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::before_filter, false);
|
||||
|
||||
printf("Testing backend top-p followed by penalties with empty history\n");
|
||||
compare_masking_penalties_logits(params, "top-p empty", []() {
|
||||
return llama_sampler_init_top_p(0.9f, 0);
|
||||
}, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true, false);
|
||||
|
||||
printf("Testing backend top-p followed by individual penalties\n");
|
||||
compare_masking_penalties_logits(params, "top-p repeat", []() {
|
||||
return llama_sampler_init_top_p(0.9f, 0);
|
||||
}, 64, 1.1f, 0.0f, 0.0f, "Hello", penalties_position::after_filter, true);
|
||||
compare_masking_penalties_logits(params, "top-p frequency", []() {
|
||||
return llama_sampler_init_top_p(0.9f, 0);
|
||||
}, 64, 1.0f, 0.5f, 0.0f, "Hello", penalties_position::after_filter, true);
|
||||
compare_masking_penalties_logits(params, "top-p presence", []() {
|
||||
return llama_sampler_init_top_p(0.9f, 0);
|
||||
}, 64, 1.0f, 0.0f, 0.25f, "Hello", penalties_position::after_filter, true);
|
||||
|
||||
printf("Testing backend penalty parameter values\n");
|
||||
test_penalty_parameter_values(params);
|
||||
|
||||
printf("backend penalties sampling test PASSED\n");
|
||||
}
|
||||
|
||||
// This test verifies that it is possible to have two different backend samplers,
|
||||
// one that uses the backend dist sampler, and another that uses CPU dist sampler.
|
||||
static void test_backend_mixed_sampling(const test_params & params) {
|
||||
@@ -1014,6 +1574,7 @@ struct backend_test_case {
|
||||
static const backend_test_case BACKEND_TESTS[] = {
|
||||
{ "greedy", test_backend_greedy_sampling, true },
|
||||
{ "logit_bias", test_backend_logit_bias_sampling, true },
|
||||
{ "penalties", test_backend_penalties_sampling, true },
|
||||
{ "temp", test_backend_temp_sampling, true },
|
||||
{ "temp_ext", test_backend_temp_ext_sampling, true },
|
||||
{ "top_k", test_backend_top_k_sampling, true },
|
||||
|
||||
@@ -3987,6 +3987,7 @@ static void test_template_output_peg_parsers(bool detailed_debug) {
|
||||
.expect_tool_calls({
|
||||
{ "special_function", R"({"arg1": 1})", {} },
|
||||
})
|
||||
.expect_reconstruction()
|
||||
.run();
|
||||
|
||||
// Tool call with negative number
|
||||
@@ -4212,6 +4213,7 @@ static void test_template_output_peg_parsers(bool detailed_debug) {
|
||||
.expect_tool_calls({
|
||||
{ "special_function", R"({"arg1": 1})", {} },
|
||||
})
|
||||
.expect_reconstruction()
|
||||
.run();
|
||||
|
||||
// Tool call with multiple params (mixed types)
|
||||
@@ -4268,6 +4270,24 @@ static void test_template_output_peg_parsers(bool detailed_debug) {
|
||||
.run();
|
||||
}
|
||||
|
||||
{
|
||||
// The DSML separator belongs to the tool call block, not assistant content.
|
||||
auto tst = peg_tester("models/templates/deepseek-ai-DeepSeek-V4-Flash-0731.jinja", detailed_debug);
|
||||
tst.test(
|
||||
"\n\n"
|
||||
"<|DSML|tool_calls>\n"
|
||||
"<|DSML|invoke name=\"special_function\">\n"
|
||||
"<|DSML|parameter name=\"arg1\" string=\"false\">1</|DSML|parameter>\n"
|
||||
"</|DSML|invoke>\n"
|
||||
"</|DSML|tool_calls>")
|
||||
.enable_thinking(false)
|
||||
.reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK)
|
||||
.tools({ special_function_tool })
|
||||
.expect(message_assist_call)
|
||||
.expect_reconstruction()
|
||||
.run();
|
||||
}
|
||||
|
||||
// GLM-4.6 tests - format: <tool_call>function_name\n<arg_key>...</arg_key>\n<arg_value>...</arg_value>\n</tool_call>
|
||||
{
|
||||
auto tst = peg_tester("models/templates/GLM-4.6.jinja", detailed_debug);
|
||||
@@ -6359,6 +6379,7 @@ static void test_template_generation_prompt() {
|
||||
std::vector<common_chat_msg> messages;
|
||||
bool add_generation_prompt = true;
|
||||
common_chat_continuation continue_final_message = COMMON_CHAT_CONTINUATION_NONE;
|
||||
bool enable_thinking = true;
|
||||
};
|
||||
|
||||
auto basic = [&]() {
|
||||
@@ -6390,6 +6411,7 @@ static void test_template_generation_prompt() {
|
||||
inputs.messages = opts.messages;
|
||||
inputs.add_generation_prompt = opts.add_generation_prompt;
|
||||
inputs.continue_final_message = opts.continue_final_message;
|
||||
inputs.enable_thinking = opts.enable_thinking;
|
||||
|
||||
auto params = common_chat_templates_apply(tmpls.get(), inputs);
|
||||
|
||||
@@ -6488,6 +6510,156 @@ static void test_template_generation_prompt() {
|
||||
check(tmpls, continuation_reasoning(), "<|Assistant|><think>I'm");
|
||||
}
|
||||
|
||||
const std::string deepseek_v4_reasoning_effort_max = "Reasoning Effort: Absolute maximum";
|
||||
const std::string deepseek_v4_flash_0731_reasoning_effort_max = "Reasoning Effort: Beyond maximum";
|
||||
|
||||
{
|
||||
auto tmpls = read_templates("models/templates/deepseek-ai-DeepSeek-V4.jinja");
|
||||
check(tmpls, basic(), "<|Assistant|><think>");
|
||||
check(tmpls, continuation_content(), "<|Assistant|><think>I'm thinking</think>Hello, ");
|
||||
check(tmpls, continuation_reasoning(), "<|Assistant|><think>I'm");
|
||||
|
||||
auto continuation_content_no_thinking = continuation_content();
|
||||
continuation_content_no_thinking.messages = { system_msg, message_user, simple_assist_msg("Hello, ") };
|
||||
continuation_content_no_thinking.enable_thinking = false;
|
||||
check(tmpls, continuation_content_no_thinking, "<|Assistant|></think>Hello, ");
|
||||
|
||||
common_chat_templates_inputs max_inputs;
|
||||
max_inputs.messages = { system_msg, message_user };
|
||||
max_inputs.chat_template_kwargs["reasoning_effort"] = R"("max")";
|
||||
auto max_params = common_chat_templates_apply(tmpls.get(), max_inputs);
|
||||
assert_contains(max_params.prompt, deepseek_v4_reasoning_effort_max);
|
||||
|
||||
auto high_inputs = max_inputs;
|
||||
high_inputs.chat_template_kwargs["reasoning_effort"] = R"("high")";
|
||||
auto high_params = common_chat_templates_apply(tmpls.get(), high_inputs);
|
||||
assert_not_contains(high_params.prompt, deepseek_v4_reasoning_effort_max);
|
||||
|
||||
auto low_inputs = max_inputs;
|
||||
low_inputs.chat_template_kwargs["reasoning_effort"] = R"("low")";
|
||||
auto low_params = common_chat_templates_apply(tmpls.get(), low_inputs);
|
||||
assert_not_contains(low_params.prompt, deepseek_v4_reasoning_effort_max);
|
||||
|
||||
common_chat_templates_inputs default_effort_inputs;
|
||||
default_effort_inputs.messages = { system_msg, message_user };
|
||||
auto default_effort_params = common_chat_templates_apply(tmpls.get(), default_effort_inputs);
|
||||
assert_not_contains(default_effort_params.prompt, deepseek_v4_reasoning_effort_max);
|
||||
|
||||
auto non_thinking_max_inputs = max_inputs;
|
||||
non_thinking_max_inputs.enable_thinking = false;
|
||||
auto non_thinking_max_params = common_chat_templates_apply(tmpls.get(), non_thinking_max_inputs);
|
||||
assert_not_contains(non_thinking_max_params.prompt, deepseek_v4_reasoning_effort_max);
|
||||
|
||||
common_chat_templates_inputs response_format_inputs;
|
||||
response_format_inputs.messages = { system_msg, message_user };
|
||||
response_format_inputs.tools = { get_time_tool };
|
||||
response_format_inputs.json_schema =
|
||||
R"({"type":"object","properties":{"answer":{"type":"string"}}})";
|
||||
auto response_format_params = common_chat_templates_apply(tmpls.get(), response_format_inputs);
|
||||
const auto tools_pos = response_format_params.prompt.find("## Tools");
|
||||
const auto response_format_pos = response_format_params.prompt.find(
|
||||
"## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n");
|
||||
if (tools_pos == std::string::npos || response_format_pos == std::string::npos || tools_pos > response_format_pos) {
|
||||
LOG_ERR("Expected response format after tools\nActual: %s\n", response_format_params.prompt.c_str());
|
||||
common_log_flush(common_log_main());
|
||||
throw std::runtime_error("Test failed");
|
||||
}
|
||||
assert_contains(response_format_params.prompt, R"("answer": {"type": "string"})");
|
||||
|
||||
response_format_inputs.json_schema = "{}";
|
||||
auto json_object_params = common_chat_templates_apply(tmpls.get(), response_format_inputs);
|
||||
assert_contains(json_object_params.prompt,
|
||||
"## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n{}");
|
||||
|
||||
common_chat_msg assistant_history;
|
||||
assistant_history.role = "assistant";
|
||||
assistant_history.content = "Previous answer";
|
||||
assistant_history.reasoning_content = "Previous reasoning";
|
||||
|
||||
common_chat_msg user_followup;
|
||||
user_followup.role = "user";
|
||||
user_followup.content = "Follow up";
|
||||
|
||||
common_chat_templates_inputs default_history_inputs;
|
||||
default_history_inputs.messages = { message_user, assistant_history, user_followup };
|
||||
auto default_history_params = common_chat_templates_apply(tmpls.get(), default_history_inputs);
|
||||
assert_contains(default_history_params.prompt, "<|Assistant|></think>Previous answer");
|
||||
|
||||
auto drop_thinking_inputs = default_history_inputs;
|
||||
drop_thinking_inputs.chat_template_kwargs["drop_thinking"] = "false";
|
||||
auto drop_thinking_params = common_chat_templates_apply(tmpls.get(), drop_thinking_inputs);
|
||||
assert_contains(drop_thinking_params.prompt, "<|Assistant|><think>Previous reasoning</think>Previous answer");
|
||||
|
||||
auto preserve_reasoning_inputs = default_history_inputs;
|
||||
preserve_reasoning_inputs.chat_template_kwargs["preserve_reasoning"] = "true";
|
||||
auto preserve_reasoning_params = common_chat_templates_apply(tmpls.get(), preserve_reasoning_inputs);
|
||||
assert_contains(preserve_reasoning_params.prompt, "<|Assistant|><think>Previous reasoning</think>Previous answer");
|
||||
assert_equals(true, common_chat_templates_get_caps(tmpls.get()).at("supports_preserve_reasoning"));
|
||||
|
||||
auto no_preserve_reasoning_inputs = default_history_inputs;
|
||||
no_preserve_reasoning_inputs.chat_template_kwargs["preserve_reasoning"] = "false";
|
||||
auto no_preserve_reasoning_params = common_chat_templates_apply(tmpls.get(), no_preserve_reasoning_inputs);
|
||||
assert_contains(no_preserve_reasoning_params.prompt, "<|Assistant|></think>Previous answer");
|
||||
|
||||
common_chat_msg empty_tool_call = simple_assist_msg("", "", "empty_args", "{}");
|
||||
common_chat_templates_inputs empty_tool_inputs;
|
||||
empty_tool_inputs.messages = { message_user, empty_tool_call };
|
||||
empty_tool_inputs.tools = { empty_args_tool };
|
||||
auto empty_tool_params = common_chat_templates_apply(tmpls.get(), empty_tool_inputs);
|
||||
assert_contains(empty_tool_params.prompt,
|
||||
"<|DSML|invoke name=\"empty_args\">\n\n</|DSML|invoke>");
|
||||
}
|
||||
|
||||
{
|
||||
auto tmpls = read_templates("models/templates/deepseek-ai-DeepSeek-V4-Flash-0731.jinja");
|
||||
check(tmpls, basic(), "<|Assistant|><think>");
|
||||
check(tmpls, continuation_content(), "<|Assistant|><think>I'm thinking</think>Hello, ");
|
||||
check(tmpls, continuation_reasoning(), "<|Assistant|><think>I'm");
|
||||
|
||||
auto continuation_content_no_thinking = continuation_content();
|
||||
continuation_content_no_thinking.messages = { system_msg, message_user, simple_assist_msg("Hello, ") };
|
||||
continuation_content_no_thinking.enable_thinking = false;
|
||||
check(tmpls, continuation_content_no_thinking, "<|Assistant|></think>Hello, ");
|
||||
|
||||
common_chat_templates_inputs high_inputs;
|
||||
high_inputs.messages = { system_msg, message_user };
|
||||
high_inputs.chat_template_kwargs["reasoning_effort"] = R"("high")";
|
||||
auto high_params = common_chat_templates_apply(tmpls.get(), high_inputs);
|
||||
assert_contains(high_params.prompt, deepseek_v4_reasoning_effort_max);
|
||||
|
||||
auto max_inputs = high_inputs;
|
||||
max_inputs.chat_template_kwargs["reasoning_effort"] = R"("max")";
|
||||
auto max_params = common_chat_templates_apply(tmpls.get(), max_inputs);
|
||||
assert_contains(max_params.prompt, deepseek_v4_flash_0731_reasoning_effort_max);
|
||||
|
||||
auto low_inputs = high_inputs;
|
||||
low_inputs.chat_template_kwargs["reasoning_effort"] = R"("low")";
|
||||
auto low_params = common_chat_templates_apply(tmpls.get(), low_inputs);
|
||||
assert_not_contains(low_params.prompt, deepseek_v4_reasoning_effort_max);
|
||||
assert_not_contains(low_params.prompt, deepseek_v4_flash_0731_reasoning_effort_max);
|
||||
|
||||
common_chat_templates_inputs default_effort_inputs;
|
||||
default_effort_inputs.messages = { system_msg, message_user };
|
||||
auto default_effort_params = common_chat_templates_apply(tmpls.get(), default_effort_inputs);
|
||||
assert_not_contains(default_effort_params.prompt, deepseek_v4_reasoning_effort_max);
|
||||
assert_not_contains(default_effort_params.prompt, deepseek_v4_flash_0731_reasoning_effort_max);
|
||||
|
||||
auto non_thinking_max_inputs = max_inputs;
|
||||
non_thinking_max_inputs.enable_thinking = false;
|
||||
auto non_thinking_max_params = common_chat_templates_apply(tmpls.get(), non_thinking_max_inputs);
|
||||
assert_not_contains(non_thinking_max_params.prompt, deepseek_v4_flash_0731_reasoning_effort_max);
|
||||
|
||||
common_chat_templates_inputs response_format_inputs;
|
||||
response_format_inputs.messages = { system_msg, message_user };
|
||||
response_format_inputs.tools = { get_time_tool };
|
||||
response_format_inputs.json_schema =
|
||||
R"({"type":"object","properties":{"answer":{"type":"string"}}})";
|
||||
auto response_format_params = common_chat_templates_apply(tmpls.get(), response_format_inputs);
|
||||
assert_contains(response_format_params.prompt,
|
||||
"## Response Format:\n\nYou MUST strictly adhere to the following schema to reply:\n");
|
||||
assert_contains(response_format_params.prompt, R"("answer": {"type": "string"})");
|
||||
}
|
||||
|
||||
{
|
||||
auto tmpls = read_templates("models/templates/openbmb-MiniCPM5-1B.jinja");
|
||||
check(tmpls, basic(), "<|im_start|>assistant\n<think>\n");
|
||||
|
||||
@@ -0,0 +1,506 @@
|
||||
// tests the HF model resolution and the model handler assembly end-to-end on
|
||||
// synthetic repo listings: a local httplib server bound to the loopback
|
||||
// serves hardcoded HF API responses, so the real client, hf_cache, resolution
|
||||
// and CLI parsing run against them without external network access
|
||||
|
||||
#include "arg.h"
|
||||
#include "common.h"
|
||||
#include "download.h"
|
||||
#include "http.h"
|
||||
#include "log.h"
|
||||
|
||||
#include <nlohmann/json.hpp>
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstdio>
|
||||
#include <cstdlib>
|
||||
#include <filesystem>
|
||||
#include <map>
|
||||
#include <thread>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
// the case and reordering being checked, printed with every failure
|
||||
static std::string g_context;
|
||||
|
||||
// independent of NDEBUG, so the checks stay alive in Release builds
|
||||
#define REQUIRE(x) do { \
|
||||
if (!(x)) { \
|
||||
fprintf(stderr, "%s:%d: [%s] REQUIRE(%s) failed\n", \
|
||||
__FILE__, __LINE__, g_context.c_str(), #x); \
|
||||
std::abort(); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
#define REQUIRE_EQ(actual, expected) do { \
|
||||
if (!((actual) == (expected))) { \
|
||||
fprintf(stderr, "%s:%d: [%s] REQUIRE_EQ(%s, %s) failed\n actual: '%s'\n expected: '%s'\n", \
|
||||
__FILE__, __LINE__, g_context.c_str(), #actual, #expected, \
|
||||
std::string(actual).c_str(), std::string(expected).c_str()); \
|
||||
std::abort(); \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
//
|
||||
// synthetic repos keyed by repo id, served over the loopback by a real
|
||||
// httplib server, so the tested code runs its own client and transport
|
||||
//
|
||||
|
||||
static std::map<std::string, std::vector<std::string>> g_repos;
|
||||
|
||||
static const char * COMMIT = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa";
|
||||
|
||||
// the server lives in main, so its destructor runs before the static teardown
|
||||
// tears down the winsock state httplib brings in
|
||||
static void serve_repos(httplib::Server & server) {
|
||||
server.Get(R"(/api/models/(.+)/refs)", [](const httplib::Request & req, httplib::Response & res) {
|
||||
if (g_repos.count(req.matches[1])) {
|
||||
res.set_content(nlohmann::json{{"branches", {{{"name", "main"}, {"targetCommit", COMMIT}}}}}.dump(),
|
||||
"application/json");
|
||||
} else {
|
||||
res.status = 404;
|
||||
}
|
||||
});
|
||||
server.Get(R"(/api/models/(.+)/tree/.+)", [](const httplib::Request & req, httplib::Response & res) {
|
||||
if (!g_repos.count(req.matches[1])) {
|
||||
res.status = 404;
|
||||
return;
|
||||
}
|
||||
auto files = nlohmann::json::array();
|
||||
size_t i = 0;
|
||||
for (const auto & p : g_repos[req.matches[1]]) {
|
||||
char oid[41];
|
||||
snprintf(oid, sizeof(oid), "%040lx", (unsigned long) ++i);
|
||||
files.push_back({{"type", "file"}, {"path", p}, {"size", 1}, {"oid", oid}});
|
||||
}
|
||||
res.set_content(files.dump(), "application/json");
|
||||
});
|
||||
}
|
||||
|
||||
static common_params_model model_ref(const std::string & hf_repo, const std::string & hf_file = "") {
|
||||
common_params_model m;
|
||||
m.hf_repo = hf_repo;
|
||||
m.hf_file = hf_file;
|
||||
return m;
|
||||
}
|
||||
|
||||
// the model cache is isolated under a temporary directory named after the
|
||||
// loopback port, so concurrent runs on a shared machine keep their own, and
|
||||
// the local path the handler wires for a file is snapshots/<commit>/<path>
|
||||
static std::filesystem::path cache_dir;
|
||||
|
||||
static std::string cached(std::string repo_id, const std::string & path) {
|
||||
string_replace_all(repo_id, "/", "--");
|
||||
return (cache_dir / ("models--" + repo_id) / "snapshots" / COMMIT / path).string();
|
||||
}
|
||||
|
||||
//
|
||||
// fixtures mimicking real repo layouts
|
||||
//
|
||||
|
||||
// flat layout in the style of ggml-org/gemma-4-31B-it-GGUF
|
||||
static const std::vector<std::string> flat = {
|
||||
"README.md",
|
||||
"model-BF16.gguf",
|
||||
"model-Q4_K_M.gguf",
|
||||
"model-Q8_0.gguf",
|
||||
"mmproj-model-BF16.gguf",
|
||||
"mmproj-model-Q8_0.gguf",
|
||||
"mtp-model-BF16.gguf",
|
||||
"mtp-model-Q4_0.gguf",
|
||||
"mtp-model-Q8_0.gguf",
|
||||
"dflash-model-BF16.gguf",
|
||||
"dflash-model-Q8_0.gguf",
|
||||
};
|
||||
|
||||
// quants in subdirectories with sharded files and root sidecars,
|
||||
// in the style of stepfun-ai/Step-3.7-Flash-GGUF
|
||||
static const std::vector<std::string> subdir = {
|
||||
"mmproj-model-f16.gguf",
|
||||
"model-mtp-BF16.gguf",
|
||||
"model-mtp-Q8_0.gguf",
|
||||
"Q3_K_M/model-Q3_K_M-00001-of-00003.gguf",
|
||||
"Q3_K_M/model-Q3_K_M-00002-of-00003.gguf",
|
||||
"Q3_K_M/model-Q3_K_M-00003-of-00003.gguf",
|
||||
"Q8_0/model-Q8_0-00001-of-00002.gguf",
|
||||
"Q8_0/model-Q8_0-00002-of-00002.gguf",
|
||||
};
|
||||
|
||||
// sidecar quants exist where the full model quant does not,
|
||||
// in the style of ggml-org/Qwen3.6-27B-GGUF
|
||||
static const std::vector<std::string> hole = {
|
||||
"model-BF16.gguf",
|
||||
"model-Q4_K_M.gguf",
|
||||
"model-Q8_0.gguf",
|
||||
"mtp-model-BF16.gguf",
|
||||
"mtp-model-Q4_0.gguf",
|
||||
"mtp-model-Q8_0.gguf",
|
||||
"dflash-model-BF16.gguf",
|
||||
"dflash-model-Q8_0.gguf",
|
||||
};
|
||||
|
||||
// unsloth-style naming with UD quants and a suffix MTP file
|
||||
static const std::vector<std::string> unsloth = {
|
||||
"model-UD-Q8_K_XL.gguf",
|
||||
"mmproj-BF16.gguf",
|
||||
"model-MTP-BF16.gguf",
|
||||
};
|
||||
|
||||
// bartowski-style vendor prefix and mradermacher-style dot quant
|
||||
static const std::vector<std::string> vendors = {
|
||||
"TheDrummer_Model-24B-v4.1-Q8_0.gguf",
|
||||
"BlackSheep-24B.Q8_0.gguf",
|
||||
};
|
||||
|
||||
// every speculative sidecar type at the same quant
|
||||
static const std::vector<std::string> quad = {
|
||||
"model-Q8_0.gguf",
|
||||
"mtp-model-Q8_0.gguf",
|
||||
"dflash-model-Q8_0.gguf",
|
||||
"eagle3-model-Q8_0.gguf",
|
||||
"dspark-model-Q8_0.gguf",
|
||||
};
|
||||
|
||||
static const std::vector<std::string> dflash_only = {
|
||||
"model-Q8_0.gguf",
|
||||
"dflash-model-Q8_0.gguf",
|
||||
};
|
||||
|
||||
static const std::vector<std::string> eagle3_only = {
|
||||
"model-Q8_0.gguf",
|
||||
"eagle3-model-Q8_0.gguf",
|
||||
};
|
||||
|
||||
// a single full quant with dspark sidecars at other quants,
|
||||
// in the style of ggml-org/DeepSeek-V4-Flash-0731-GGUF
|
||||
static const std::vector<std::string> spark = {
|
||||
"README.md",
|
||||
"model-MXFP4.gguf",
|
||||
"dspark-model-BF16.gguf",
|
||||
"dspark-model-MXFP4.gguf",
|
||||
};
|
||||
|
||||
// dspark outranks dflash in the type auto-selection
|
||||
static const std::vector<std::string> dspark_dflash = {
|
||||
"model-Q8_0.gguf",
|
||||
"dflash-model-Q8_0.gguf",
|
||||
"dspark-model-Q8_0.gguf",
|
||||
};
|
||||
|
||||
//
|
||||
// table-driven plan resolution through the real entry point,
|
||||
// each case replayed on multiple deterministic reorderings of the listing,
|
||||
// except the cases whose pick legitimately depends on the listing order
|
||||
//
|
||||
|
||||
struct plan_case {
|
||||
const char * name;
|
||||
const std::vector<std::string> & files;
|
||||
const char * hf_repo;
|
||||
const char * hf_file;
|
||||
bool sidecars; // request mmproj + mtp + dflash + eagle3 + dspark
|
||||
bool order_dependent; // the expected pick depends on the listing order
|
||||
const char * primary;
|
||||
std::vector<std::string> model_files;
|
||||
const char * mmproj;
|
||||
const char * mtp;
|
||||
const char * dflash;
|
||||
const char * eagle3;
|
||||
const char * dspark;
|
||||
};
|
||||
|
||||
static const plan_case plan_cases[] = {
|
||||
// exact tag picks the matching primary, sidecars follow the tag
|
||||
{"flat exact tag", flat, "test/repo:Q8_0", "", true, false,
|
||||
"model-Q8_0.gguf", {"model-Q8_0.gguf"},
|
||||
"mmproj-model-Q8_0.gguf", "mtp-model-Q8_0.gguf", "dflash-model-Q8_0.gguf", "", ""},
|
||||
|
||||
// no tag falls back to the default quant preference
|
||||
{"flat default", flat, "test/repo", "", false, false,
|
||||
"model-Q4_K_M.gguf", {"model-Q4_K_M.gguf"},
|
||||
"", "", "", "", ""},
|
||||
|
||||
// no tag and no default match falls back to the first model in the listing
|
||||
{"unsloth fallback", unsloth, "test/repo", "", true, true,
|
||||
"model-UD-Q8_K_XL.gguf", {"model-UD-Q8_K_XL.gguf"},
|
||||
"mmproj-BF16.gguf", "", "", "", ""},
|
||||
|
||||
// explicit hf_file picks that exact file
|
||||
{"flat hf_file", flat, "test/repo", "model-BF16.gguf", false, false,
|
||||
"model-BF16.gguf", {"model-BF16.gguf"},
|
||||
"", "", "", "", ""},
|
||||
|
||||
// missing hf_file resolves nothing
|
||||
{"flat missing hf_file", flat, "test/repo", "nope.gguf", false, false,
|
||||
"", {},
|
||||
"", "", "", "", ""},
|
||||
|
||||
// a sharded primary brings all its parts, a subdir primary finds the root sidecar
|
||||
{"subdir shards", subdir, "test/repo:Q3_K_M", "", true, false,
|
||||
"Q3_K_M/model-Q3_K_M-00001-of-00003.gguf",
|
||||
{"Q3_K_M/model-Q3_K_M-00001-of-00003.gguf",
|
||||
"Q3_K_M/model-Q3_K_M-00002-of-00003.gguf",
|
||||
"Q3_K_M/model-Q3_K_M-00003-of-00003.gguf"},
|
||||
"mmproj-model-f16.gguf", "model-mtp-Q8_0.gguf", "", "", ""},
|
||||
|
||||
// a tag with no matching full model still resolves the requested sidecars
|
||||
{"hole tag sidecar", hole, "test/repo:Q4_0", "", true, false,
|
||||
"", {},
|
||||
"", "mtp-model-Q4_0.gguf", "dflash-model-Q8_0.gguf", "", ""},
|
||||
|
||||
// the same tag without a requested sidecar resolves nothing
|
||||
{"hole tag alone", hole, "test/repo:Q4_0", "", false, false,
|
||||
"", {},
|
||||
"", "", "", "", ""},
|
||||
|
||||
// no tag anchors the sidecars on the primary quant
|
||||
{"hole default anchor", hole, "test/repo", "", true, false,
|
||||
"model-Q4_K_M.gguf", {"model-Q4_K_M.gguf"},
|
||||
"", "mtp-model-Q4_0.gguf", "dflash-model-Q8_0.gguf", "", ""},
|
||||
|
||||
// the mtp- keyword is case sensitive, a suffix -MTP file is not discovered
|
||||
{"unsloth suffix mtp", unsloth, "test/repo:Q8_K_XL", "", true, false,
|
||||
"model-UD-Q8_K_XL.gguf", {"model-UD-Q8_K_XL.gguf"},
|
||||
"mmproj-BF16.gguf", "", "", "", ""},
|
||||
|
||||
// vendor prefixes and the dot quant convention both match the tag,
|
||||
// first match wins between two files at the same quant
|
||||
{"vendor prefix", vendors, "test/repo:Q8_0", "", false, true,
|
||||
"TheDrummer_Model-24B-v4.1-Q8_0.gguf", {"TheDrummer_Model-24B-v4.1-Q8_0.gguf"},
|
||||
"", "", "", "", ""},
|
||||
|
||||
// every sidecar type resolves at the tag
|
||||
{"quad exact tag", quad, "test/repo:Q8_0", "", true, false,
|
||||
"model-Q8_0.gguf", {"model-Q8_0.gguf"},
|
||||
"", "mtp-model-Q8_0.gguf", "dflash-model-Q8_0.gguf", "eagle3-model-Q8_0.gguf", "dspark-model-Q8_0.gguf"},
|
||||
|
||||
// no tag anchors the dspark sidecar on the only full quant
|
||||
{"spark default anchor", spark, "test/repo", "", true, false,
|
||||
"model-MXFP4.gguf", {"model-MXFP4.gguf"},
|
||||
"", "", "", "", "dspark-model-MXFP4.gguf"},
|
||||
|
||||
// a tag with no matching full model still resolves the exact dspark sidecar
|
||||
{"spark tag sidecar", spark, "test/repo:BF16", "", true, false,
|
||||
"", {},
|
||||
"", "", "", "", "dspark-model-BF16.gguf"},
|
||||
};
|
||||
|
||||
static void check_plan(const plan_case & c) {
|
||||
common_download_opts opts;
|
||||
opts.download_mmproj = c.sidecars;
|
||||
opts.download_mtp = c.sidecars;
|
||||
opts.download_dflash = c.sidecars;
|
||||
opts.download_eagle3 = c.sidecars;
|
||||
opts.download_dspark = c.sidecars;
|
||||
|
||||
auto plan = common_download_get_hf_plan(model_ref(c.hf_repo, c.hf_file), opts);
|
||||
|
||||
REQUIRE_EQ(plan.primary.path, c.primary);
|
||||
REQUIRE_EQ(plan.mmproj.path, c.mmproj);
|
||||
REQUIRE_EQ(plan.mtp.path, c.mtp);
|
||||
REQUIRE_EQ(plan.dflash.path, c.dflash);
|
||||
REQUIRE_EQ(plan.eagle3.path, c.eagle3);
|
||||
REQUIRE_EQ(plan.dspark.path, c.dspark);
|
||||
|
||||
// exact shard set, order insensitive; the primary must be the first split
|
||||
std::vector<std::string> actual;
|
||||
for (const auto & f : plan.model_files) {
|
||||
actual.push_back(f.path);
|
||||
}
|
||||
std::sort(actual.begin(), actual.end());
|
||||
auto expected = c.model_files;
|
||||
std::sort(expected.begin(), expected.end());
|
||||
REQUIRE(actual == expected);
|
||||
if (!expected.empty()) {
|
||||
REQUIRE(plan.primary.path == expected.front());
|
||||
}
|
||||
}
|
||||
|
||||
static void test_plan_resolution() {
|
||||
printf("test-model-resolution: plan resolution on %zu cases\n", sizeof(plan_cases) / sizeof(plan_cases[0]));
|
||||
|
||||
for (const auto & c : plan_cases) {
|
||||
printf(" %s\n", c.name);
|
||||
// invariant: the resolution is insensitive to the listing order
|
||||
for (size_t rot = 0; rot < c.files.size(); ++rot) {
|
||||
if (c.order_dependent && rot > 0) {
|
||||
continue;
|
||||
}
|
||||
g_context = std::string(c.name) + ", reordering " + std::to_string(rot);
|
||||
auto files = c.files;
|
||||
std::rotate(files.begin(), files.begin() + rot, files.end());
|
||||
if (rot % 2 == 1) {
|
||||
std::reverse(files.begin(), files.end());
|
||||
}
|
||||
g_repos["test/repo"] = files;
|
||||
check_plan(c);
|
||||
}
|
||||
}
|
||||
g_repos.clear();
|
||||
}
|
||||
|
||||
//
|
||||
// end-to-end assembly: real CLI parsing, real handler init resolving over the
|
||||
// loopback, downloads skipped by flipping offline before apply
|
||||
//
|
||||
|
||||
static void assemble(std::vector<std::string> argv, common_params & params) {
|
||||
std::vector<char *> cargv;
|
||||
g_context.clear();
|
||||
for (auto & a : argv) {
|
||||
g_context += g_context.empty() ? a : " " + a;
|
||||
cargv.push_back(a.data());
|
||||
}
|
||||
bool ok = common_params_parse((int) cargv.size(), cargv.data(), params, LLAMA_EXAMPLE_SERVER);
|
||||
REQUIRE(ok);
|
||||
|
||||
auto handler = common_models_handler_init(params, LLAMA_EXAMPLE_SERVER);
|
||||
|
||||
// skip the network execution, on_done still wires the params
|
||||
params.offline = true;
|
||||
common_models_handler_apply(handler, params);
|
||||
}
|
||||
|
||||
static void test_task_assembly() {
|
||||
printf("test-model-resolution: end-to-end assembly\n");
|
||||
|
||||
g_repos["test/main"] = flat;
|
||||
g_repos["test/hole"] = hole;
|
||||
g_repos["test/quad"] = quad;
|
||||
g_repos["test/dflash"] = dflash_only;
|
||||
g_repos["test/eagle3"] = eagle3_only;
|
||||
g_repos["test/spark"] = spark;
|
||||
g_repos["test/pair"] = dspark_dflash;
|
||||
g_repos["test/small"] = {"draft-model-Q4_K_M.gguf"};
|
||||
g_repos["test/preset"] = {"preset.ini", "model-Q8_0.gguf"};
|
||||
|
||||
{
|
||||
// plain -hf wires the model and its mmproj, nothing speculative
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0"}, params);
|
||||
REQUIRE_EQ(params.model.path, cached("test/main", "model-Q8_0.gguf"));
|
||||
REQUIRE_EQ(params.mmproj.path, cached("test/main", "mmproj-model-Q8_0.gguf"));
|
||||
REQUIRE(params.speculative.draft.mparams.path.empty());
|
||||
}
|
||||
{
|
||||
// --no-mmproj disables the mmproj discovery
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "--no-mmproj"}, params);
|
||||
REQUIRE(params.mmproj.path.empty());
|
||||
}
|
||||
{
|
||||
// an explicit --mmproj wins over the discovery
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "--mmproj", "/local/mmproj.gguf"}, params);
|
||||
REQUIRE(params.mmproj.path == "/local/mmproj.gguf");
|
||||
}
|
||||
{
|
||||
// -hf with a spec type wires the sidecar of the main repo as fallback draft
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "--spec-type", "draft-mtp"}, params);
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/main", "mtp-model-Q8_0.gguf"));
|
||||
}
|
||||
{
|
||||
// -hfd with a spec type wires the draft repo sidecar at its tag,
|
||||
// not its full model, and suppresses the main repo fallback
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/hole:Q8_0", "-hfd", "test/hole:Q4_0", "--spec-type", "draft-mtp"}, params);
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/hole", "mtp-model-Q4_0.gguf"));
|
||||
}
|
||||
{
|
||||
// an explicit -md file wins over the sidecar resolution
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/main", "-md", "mtp-model-BF16.gguf", "--spec-type", "draft-mtp"}, params);
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/main", "mtp-model-BF16.gguf"));
|
||||
}
|
||||
{
|
||||
// -hfd without a spec type auto-selects the type, mtp first when all ship
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/quad:Q8_0"}, params);
|
||||
REQUIRE(params.speculative.types == std::vector<enum common_speculative_type>{COMMON_SPECULATIVE_TYPE_DRAFT_MTP});
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/quad", "mtp-model-Q8_0.gguf"));
|
||||
}
|
||||
{
|
||||
// auto-selection with only a dflash sidecar
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/dflash:Q8_0"}, params);
|
||||
REQUIRE(params.speculative.types == std::vector<enum common_speculative_type>{COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH});
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/dflash", "dflash-model-Q8_0.gguf"));
|
||||
}
|
||||
{
|
||||
// auto-selection with only an eagle3 sidecar
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/eagle3:Q8_0"}, params);
|
||||
REQUIRE(params.speculative.types == std::vector<enum common_speculative_type>{COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3});
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/eagle3", "eagle3-model-Q8_0.gguf"));
|
||||
}
|
||||
{
|
||||
// auto-selection prefers dspark over dflash when both ship
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/pair:Q8_0"}, params);
|
||||
REQUIRE(params.speculative.types == std::vector<enum common_speculative_type>{COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK});
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/pair", "dspark-model-Q8_0.gguf"));
|
||||
}
|
||||
{
|
||||
// -hf with the dspark spec type wires the sidecar of the main repo,
|
||||
// anchored on the only full quant
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/spark", "--spec-type", "draft-dspark"}, params);
|
||||
REQUIRE_EQ(params.model.path, cached("test/spark", "model-MXFP4.gguf"));
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/spark", "dspark-model-MXFP4.gguf"));
|
||||
}
|
||||
{
|
||||
// -hfd on a repo without sidecars keeps resolving a full model as draft
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/small"}, params);
|
||||
REQUIRE(params.speculative.types == std::vector<enum common_speculative_type>{COMMON_SPECULATIVE_TYPE_NONE});
|
||||
REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/small", "draft-model-Q4_K_M.gguf"));
|
||||
}
|
||||
{
|
||||
// a preset repo wires the preset and clears the model for router mode
|
||||
common_params params;
|
||||
assemble({"server", "-hf", "test/preset"}, params);
|
||||
REQUIRE_EQ(params.models_preset, cached("test/preset", "preset.ini"));
|
||||
REQUIRE(params.model.path.empty());
|
||||
REQUIRE(params.model.hf_repo.empty());
|
||||
}
|
||||
|
||||
g_repos.clear();
|
||||
}
|
||||
|
||||
int main(void) {
|
||||
// unbuffered, so a crash cannot swallow the reports already printed
|
||||
setvbuf(stdout, nullptr, _IONBF, 0);
|
||||
setvbuf(stderr, nullptr, _IONBF, 0);
|
||||
|
||||
// the negative cases legitimately log errors on every reordering,
|
||||
// keep the output down to the reports
|
||||
common_log_pause(common_log_main());
|
||||
|
||||
// the loopback endpoint also keeps the client init from rejecting
|
||||
// https on the builds without TLS support
|
||||
httplib::Server server;
|
||||
serve_repos(server);
|
||||
int port = server.bind_to_any_port("127.0.0.1");
|
||||
|
||||
// isolate the cache, its location is read once so it is set
|
||||
// before anything else
|
||||
cache_dir = std::filesystem::temp_directory_path() /
|
||||
("test-model-resolution-cache-" + std::to_string(port));
|
||||
std::filesystem::remove_all(cache_dir);
|
||||
common_set_env("LLAMA_CACHE", cache_dir.string());
|
||||
|
||||
std::thread server_thread([&server] { server.listen_after_bind(); });
|
||||
server.wait_until_ready();
|
||||
common_set_env("MODEL_ENDPOINT", "http://127.0.0.1:" + std::to_string(port) + "/");
|
||||
|
||||
test_plan_resolution();
|
||||
test_task_assembly();
|
||||
|
||||
server.stop();
|
||||
server_thread.join();
|
||||
|
||||
std::filesystem::remove_all(cache_dir);
|
||||
printf("test-model-resolution: all tests OK\n");
|
||||
return 0;
|
||||
}
|
||||
@@ -198,7 +198,9 @@ For the full list of features, please refer to [server's changelog](https://gith
|
||||
| `--ui-config, --webui-config JSON` | JSON that provides default UI settings (overrides UI defaults)<br/>(env: LLAMA_ARG_UI_CONFIG) |
|
||||
| `--ui-config-file, --webui-config-file PATH` | JSON file that provides default UI settings (overrides UI defaults)<br/>(env: LLAMA_ARG_UI_CONFIG_FILE) |
|
||||
| `--ui-mcp-proxy, --webui-mcp-proxy, --no-ui-mcp-proxy, --no-webui-mcp-proxy` | experimental: whether to enable MCP CORS proxy - do not enable in untrusted environments (default: disabled)<br/>(env: LLAMA_ARG_UI_MCP_PROXY) |
|
||||
| `--tools TOOL1,TOOL2,...` | experimental: whether to enable built-in tools for AI agents - do not enable in untrusted environments (default: no tools)<br/>specify "all" to enable all tools<br/>available tools: read_file, file_glob_search, grep_search, exec_shell_command, write_file, edit_file, get_datetime<br/>note: for security reasons, this will limit --cors-origins to localhost by default<br/>(env: LLAMA_ARG_TOOLS) |
|
||||
| `--tools TOOL1,TOOL2,...` | experimental: whether to enable built-in tools for AI agents - do not enable in untrusted environments (default: no tools)<br/>specify "all" to enable all tools<br/>available tools: read_file, file_glob_search, grep_search, exec_shell_command, write_file, edit_file, get_datetime, get_info<br/>note: for security reasons, this will limit --cors-origins to localhost by default<br/>(env: LLAMA_ARG_TOOLS) |
|
||||
| `--mcp-servers-config PATH` | experimental: path to JSON file with MCP server definitions (Cursor-compatible format) - do not enable in untrusted environments (default: none)<br/>note: for security reasons, this will limit --cors-origins to localhost by default<br/>(env: LLAMA_ARG_MCP_SERVERS_CONFIG) |
|
||||
| `--mcp-servers-json JSON` | experimental: inline JSON with MCP server definitions (Cursor-compatible format) - do not enable in untrusted environments (default: none)<br/>note: for security reasons, this will limit --cors-origins to localhost by default<br/>(env: LLAMA_ARG_MCP_SERVERS_JSON) |
|
||||
| `-ag, --agent, -no-ag, --no-agent` | whether to enable CORS proxy and all built-in tools - do not enable in untrusted environments (default: disabled)<br/>note: for security reasons, this will limit --cors-origins to localhost by default<br/>(env: LLAMA_ARG_AGENT) |
|
||||
| `--ui, --webui, --no-ui, --no-webui` | whether to enable the Web UI (default: enabled)<br/>(env: LLAMA_ARG_UI) |
|
||||
| `--embedding, --embeddings` | restrict to only support embedding use case; use only with dedicated embedding models (default: disabled)<br/>(env: LLAMA_ARG_EMBEDDINGS) |
|
||||
|
||||
@@ -1807,7 +1807,8 @@ private:
|
||||
// initialize samplers
|
||||
if (task.need_sampling()) {
|
||||
try {
|
||||
slot.smpl.reset(common_sampler_init(model_tgt, task.params.sampling));
|
||||
slot.smpl.reset(common_sampler_init(
|
||||
model_tgt, task.params.sampling, (int32_t) llama_n_ctx(ctx_tgt)));
|
||||
} catch (std::exception & e) {
|
||||
std::string err_msg = std::string("Failed to initialize samplers: ") + e.what();
|
||||
send_error(task, err_msg, ERROR_TYPE_INVALID_REQUEST);
|
||||
|
||||
@@ -1090,6 +1090,56 @@ struct server_tool_get_datetime : server_tool {
|
||||
}
|
||||
};
|
||||
|
||||
//
|
||||
// get_info: returns runtime info (OS name/version and cwd)
|
||||
//
|
||||
|
||||
struct server_tool_get_info : server_tool {
|
||||
server_tool_get_info() {
|
||||
name = "get_info";
|
||||
display_name = "Get Runtime Info";
|
||||
permission_write = false;
|
||||
}
|
||||
|
||||
json get_definition() const override {
|
||||
return {
|
||||
{"type", "function"},
|
||||
{"function", {
|
||||
{"name", name},
|
||||
{"description", "Returns runtime info: the OS name/version and the current working directory"},
|
||||
{"parameters", {
|
||||
{"type", "object"},
|
||||
{"properties", json::object()},
|
||||
}},
|
||||
}},
|
||||
};
|
||||
}
|
||||
|
||||
json invoke(json params, server_tool::stream *) const override {
|
||||
auto io = make_tools_io(params);
|
||||
|
||||
#ifdef _WIN32
|
||||
auto res = io->run({"cmd", "/c", "ver"}, 4096, 5);
|
||||
#else
|
||||
auto res = io->run({"uname", "-a"}, 4096, 5);
|
||||
#endif
|
||||
// "ver" prints a blank line before the version, so the output is stripped on both ends;
|
||||
// a failed spawn or a timeout leaves a diagnostic in res.output, which is not an OS name
|
||||
std::string os_info = res.exit_code == 0 && !res.timed_out ? string_strip(res.output) : "unknown";
|
||||
|
||||
std::string cwd = json_value(params, "cwd", std::string());
|
||||
if (cwd.empty()) {
|
||||
std::error_code ec;
|
||||
cwd = fs::current_path(ec).string();
|
||||
}
|
||||
|
||||
return {
|
||||
{"os", os_info},
|
||||
{"cwd", cwd},
|
||||
};
|
||||
}
|
||||
};
|
||||
|
||||
struct server_tool_stream_result : server_task_result {
|
||||
std::string chunk;
|
||||
bool done = false;
|
||||
@@ -1199,6 +1249,7 @@ static std::vector<std::unique_ptr<server_tool>> build_tools() {
|
||||
tools.push_back(std::make_unique<server_tool_write_file>());
|
||||
tools.push_back(std::make_unique<server_tool_edit_file>());
|
||||
tools.push_back(std::make_unique<server_tool_get_datetime>());
|
||||
tools.push_back(std::make_unique<server_tool_get_info>());
|
||||
return tools;
|
||||
}
|
||||
|
||||
|
||||
Vendored
+1
-1
@@ -41,7 +41,7 @@ if (LLAMA_BUILD_BORINGSSL)
|
||||
set(FIPS OFF CACHE BOOL "Enable FIPS (BoringSSL)")
|
||||
|
||||
set(BORINGSSL_GIT "https://boringssl.googlesource.com/boringssl" CACHE STRING "BoringSSL git repository")
|
||||
set(BORINGSSL_VERSION "0.20260730.0" CACHE STRING "BoringSSL version")
|
||||
set(BORINGSSL_VERSION "0.20260803.0" CACHE STRING "BoringSSL version")
|
||||
|
||||
message(STATUS "Fetching BoringSSL version ${BORINGSSL_VERSION}")
|
||||
|
||||
|
||||
Vendored
+411
-174
@@ -1412,6 +1412,46 @@ bool stream_line_reader::getline() {
|
||||
#endif
|
||||
|
||||
for (size_t i = 0;; i++) {
|
||||
// Fast path: whatever the stream has already buffered can be scanned for
|
||||
// the terminator in one pass. Asking for a byte at a time costs a virtual
|
||||
// call, a bounds check and a one-byte copy per character of the request.
|
||||
size_t buffered_size = 0;
|
||||
if (auto buffered = strm_.buffered_data(buffered_size)) {
|
||||
auto take = buffered_size;
|
||||
auto terminated = false;
|
||||
|
||||
for (size_t at = 0; at < buffered_size;) {
|
||||
auto nl = static_cast<const char *>(
|
||||
memchr(buffered + at, '\n', buffered_size - at));
|
||||
if (!nl) { break; }
|
||||
auto pos = static_cast<size_t>(nl - buffered);
|
||||
#ifdef CPPHTTPLIB_ALLOW_LF_AS_LINE_TERMINATOR
|
||||
take = pos + 1;
|
||||
terminated = true;
|
||||
break;
|
||||
#else
|
||||
// A bare LF does not end the line; keep looking for CRLF. The CR may
|
||||
// be the last byte of an earlier chunk, hence prev_byte.
|
||||
if ((pos > 0 ? buffered[pos - 1] : prev_byte) == '\r') {
|
||||
take = pos + 1;
|
||||
terminated = true;
|
||||
break;
|
||||
}
|
||||
at = pos + 1;
|
||||
#endif
|
||||
}
|
||||
|
||||
if (size() + take > CPPHTTPLIB_MAX_LINE_LENGTH) { return false; }
|
||||
#ifndef CPPHTTPLIB_ALLOW_LF_AS_LINE_TERMINATOR
|
||||
prev_byte = buffered[take - 1];
|
||||
#endif
|
||||
append(buffered, take);
|
||||
strm_.consume_buffered(take);
|
||||
i += take;
|
||||
if (terminated) { return true; }
|
||||
continue;
|
||||
}
|
||||
|
||||
if (size() >= CPPHTTPLIB_MAX_LINE_LENGTH) {
|
||||
// Treat exceptionally long lines as an error to
|
||||
// prevent infinite loops/memory exhaustion
|
||||
@@ -1443,16 +1483,26 @@ bool stream_line_reader::getline() {
|
||||
return true;
|
||||
}
|
||||
|
||||
void stream_line_reader::append(char c) {
|
||||
if (fixed_buffer_used_size_ < fixed_buffer_size_ - 1) {
|
||||
fixed_buffer_[fixed_buffer_used_size_++] = c;
|
||||
void stream_line_reader::append(char c) { append(&c, 1); }
|
||||
|
||||
void stream_line_reader::append(const char *data, size_t size) {
|
||||
// Once the line has outgrown the fixed buffer everything must keep going to
|
||||
// the growable one, even if a later chunk would have fit. Without the
|
||||
// emptiness check a short append after a long one would land in the fixed
|
||||
// buffer, which ptr() and size() no longer look at, and be lost.
|
||||
if (growable_buffer_.empty() &&
|
||||
fixed_buffer_used_size_ + size < fixed_buffer_size_) {
|
||||
memcpy(fixed_buffer_ + fixed_buffer_used_size_, data, size);
|
||||
fixed_buffer_used_size_ += size;
|
||||
fixed_buffer_[fixed_buffer_used_size_] = '\0';
|
||||
} else {
|
||||
// Unlike the per-character overload, this can be the very first append of
|
||||
// the line, so the fixed buffer may hold nothing and carry no terminator
|
||||
// yet. assign() takes an explicit length and does not need one.
|
||||
if (growable_buffer_.empty()) {
|
||||
assert(fixed_buffer_[fixed_buffer_used_size_] == '\0');
|
||||
growable_buffer_.assign(fixed_buffer_, fixed_buffer_used_size_);
|
||||
}
|
||||
growable_buffer_ += c;
|
||||
growable_buffer_.append(data, size);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1525,6 +1575,14 @@ bool mmap::open(const char *path) {
|
||||
is_open_empty_file = true;
|
||||
return false;
|
||||
}
|
||||
|
||||
if (addr_ == MAP_FAILED) {
|
||||
// Clear the sentinel before `close()`, since `is_open()` only checks
|
||||
// `addr_` against nullptr and `munmap()` must not be called with it.
|
||||
addr_ = nullptr;
|
||||
close();
|
||||
return false;
|
||||
}
|
||||
#endif
|
||||
|
||||
return true;
|
||||
@@ -1702,8 +1760,17 @@ public:
|
||||
socket_t socket() const override;
|
||||
time_t duration() const override;
|
||||
void set_read_timeout(time_t sec, time_t usec = 0) override;
|
||||
const char *buffered_data(size_t &size) const override;
|
||||
void consume_buffered(size_t size) override;
|
||||
|
||||
// The caller has just seen this socket become readable. Lets the next read
|
||||
// skip its own readiness wait, which would otherwise ask the kernel a
|
||||
// question that was answered a moment ago. Consumed by that read.
|
||||
void set_readable_hint() { readable_hint_ = true; }
|
||||
|
||||
private:
|
||||
bool ensure_readable();
|
||||
|
||||
socket_t sock_;
|
||||
time_t read_timeout_sec_;
|
||||
time_t read_timeout_usec_;
|
||||
@@ -1715,6 +1782,7 @@ private:
|
||||
std::vector<char> read_buff_;
|
||||
size_t read_buff_off_ = 0;
|
||||
size_t read_buff_content_size_ = 0;
|
||||
bool readable_hint_ = false;
|
||||
|
||||
static const size_t read_buff_size_ = 1024l * 4;
|
||||
};
|
||||
@@ -1782,6 +1850,9 @@ process_server_socket(const std::atomic<socket_t> &svr_sock, socket_t sock,
|
||||
[&](bool close_connection, bool &connection_closed) {
|
||||
SocketStream strm(sock, read_timeout_sec, read_timeout_usec,
|
||||
write_timeout_sec, write_timeout_usec);
|
||||
// process_server_socket_core() only gets here once keep_alive() has
|
||||
// seen the socket go readable.
|
||||
strm.set_readable_hint();
|
||||
return callback(strm, close_connection, connection_closed);
|
||||
});
|
||||
}
|
||||
@@ -3071,19 +3142,49 @@ bool zstd_decompressor::decompress(const char *data, size_t data_length,
|
||||
}
|
||||
#endif
|
||||
|
||||
bool contains_case_ignore(const std::string &s, const char *token) {
|
||||
auto token_end = token + std::strlen(token);
|
||||
return std::search(s.begin(), s.end(), token, token_end, [](char a, char b) {
|
||||
return case_ignore::to_lower(a) == case_ignore::to_lower(b);
|
||||
}) != s.end();
|
||||
}
|
||||
|
||||
// Content codings are case-insensitive (RFC 9110 8.4.1). Matching them
|
||||
// case-sensitively would make a response labeled e.g. "GZIP" look like an
|
||||
// unknown coding, and its payload would be handed back still compressed.
|
||||
bool is_zlib_encoding(const std::string &encoding) {
|
||||
return case_ignore::equal(encoding, "gzip") ||
|
||||
case_ignore::equal(encoding, "deflate");
|
||||
}
|
||||
|
||||
bool is_brotli_encoding(const std::string &encoding) {
|
||||
return contains_case_ignore(encoding, "br");
|
||||
}
|
||||
|
||||
bool is_zstd_encoding(const std::string &encoding) {
|
||||
return contains_case_ignore(encoding, "zstd");
|
||||
}
|
||||
|
||||
// Returns true if the content coding is one cpp-httplib is able to decompress
|
||||
// when the corresponding support is compiled in.
|
||||
bool is_known_content_encoding(const std::string &encoding) {
|
||||
return is_zlib_encoding(encoding) || is_brotli_encoding(encoding) ||
|
||||
is_zstd_encoding(encoding);
|
||||
}
|
||||
|
||||
std::unique_ptr<decompressor>
|
||||
create_decompressor(const std::string &encoding) {
|
||||
std::unique_ptr<decompressor> decompressor;
|
||||
|
||||
if (encoding == "gzip" || encoding == "deflate") {
|
||||
if (is_zlib_encoding(encoding)) {
|
||||
#ifdef CPPHTTPLIB_ZLIB_SUPPORT
|
||||
decompressor = detail::make_unique<gzip_decompressor>();
|
||||
#endif
|
||||
} else if (encoding.find("br") != std::string::npos) {
|
||||
} else if (is_brotli_encoding(encoding)) {
|
||||
#ifdef CPPHTTPLIB_BROTLI_SUPPORT
|
||||
decompressor = detail::make_unique<brotli_decompressor>();
|
||||
#endif
|
||||
} else if (encoding == "zstd" || encoding.find("zstd") != std::string::npos) {
|
||||
} else if (is_zstd_encoding(encoding)) {
|
||||
#ifdef CPPHTTPLIB_ZSTD_SUPPORT
|
||||
decompressor = detail::make_unique<zstd_decompressor>();
|
||||
#endif
|
||||
@@ -3145,8 +3246,7 @@ const char *get_header_value(const Headers &headers,
|
||||
|
||||
size_t get_header_value_count(const Headers &headers,
|
||||
const std::string &key) {
|
||||
auto r = headers.equal_range(key);
|
||||
return static_cast<size_t>(std::distance(r.first, r.second));
|
||||
return headers.count(key);
|
||||
}
|
||||
|
||||
template <typename Map>
|
||||
@@ -3370,44 +3470,33 @@ ReadContentResult read_content_chunked(Stream &strm, T &x,
|
||||
bool is_chunked_transfer_encoding(const Headers &headers) {
|
||||
// RFC 9112 6.1: a message is framed with the chunked coding when "chunked"
|
||||
// is the final transfer coding. A single field value may list several
|
||||
// codings ("gzip, chunked"), and the list may be split across multiple
|
||||
// Transfer-Encoding header lines (RFC 9110 5.3). Match the last coding token
|
||||
// case-insensitively rather than comparing the whole value against "chunked".
|
||||
// codings ("gzip, chunked"), and RFC 9110 5.3 lets that list be split across
|
||||
// several Transfer-Encoding lines, which combine into one comma-separated
|
||||
// list in the order the lines were received. Headers preserves that order,
|
||||
// so the final coding is the last token of the last line. Match it
|
||||
// case-insensitively rather than comparing the whole value against
|
||||
// "chunked".
|
||||
//
|
||||
// Security: reading a chunked message as unframed leaves its body in the
|
||||
// socket, where a keep-alive connection parses it as a smuggled request.
|
||||
// Headers is an unordered_multimap whose iteration order for duplicate keys
|
||||
// is not portable, so when there is more than one Transfer-Encoding line we
|
||||
// cannot tell which coding is truly final. In that ambiguous case we fail
|
||||
// safe by treating the message as chunked (a mis-parse just closes the
|
||||
// connection, whereas the opposite error enables smuggling).
|
||||
// Server::process_request() answers 400 and closes when the final coding is
|
||||
// not chunked, so a request whose framing cannot be determined never
|
||||
// reaches the "no body" path.
|
||||
auto rng = headers.equal_range("Transfer-Encoding");
|
||||
if (rng.first == rng.second) { return false; }
|
||||
|
||||
size_t line_count = 0;
|
||||
bool chunked_present = false;
|
||||
bool last_line_ends_with_chunked = false;
|
||||
// Cleared per line, so a trailing line carrying no coding at all leaves the
|
||||
// combined list ending in nothing rather than inheriting the line before it.
|
||||
std::string last_coding;
|
||||
|
||||
for (auto it = rng.first; it != rng.second; ++it) {
|
||||
line_count++;
|
||||
const auto &value = it->second;
|
||||
|
||||
std::string last_coding;
|
||||
bool line_has_chunked = false;
|
||||
last_coding.clear();
|
||||
split(value.data(), value.data() + value.size(), ',',
|
||||
[&](const char *b, const char *e) {
|
||||
last_coding.assign(b, e);
|
||||
if (case_ignore::equal(last_coding, "chunked")) {
|
||||
line_has_chunked = true;
|
||||
}
|
||||
});
|
||||
|
||||
if (line_has_chunked) { chunked_present = true; }
|
||||
last_line_ends_with_chunked = case_ignore::equal(last_coding, "chunked");
|
||||
[&](const char *b, const char *e) { last_coding.assign(b, e); });
|
||||
}
|
||||
|
||||
if (line_count == 0) { return false; }
|
||||
if (line_count == 1) { return last_line_ends_with_chunked; }
|
||||
return chunked_present;
|
||||
return case_ignore::equal(last_coding, "chunked");
|
||||
}
|
||||
|
||||
template <typename T, typename U>
|
||||
@@ -3420,9 +3509,12 @@ bool prepare_content_receiver(T &x, int &status,
|
||||
std::unique_ptr<decompressor> decompressor;
|
||||
|
||||
if (!encoding.empty()) {
|
||||
// A coding we know about but were not built with is an error. An
|
||||
// unrecognized coding (including "identity") is left alone and the
|
||||
// payload is passed through as-is, since some servers misuse the header,
|
||||
// e.g. by sending a character set such as "Content-Encoding: UTF-8".
|
||||
decompressor = detail::create_decompressor(encoding);
|
||||
if (!decompressor) {
|
||||
// Unsupported encoding or no support compiled in
|
||||
if (!decompressor && detail::is_known_content_encoding(encoding)) {
|
||||
status = StatusCode::UnsupportedMediaType_415;
|
||||
return false;
|
||||
}
|
||||
@@ -3845,6 +3937,19 @@ std::string params_to_query_str(const Params ¶ms) {
|
||||
return query;
|
||||
}
|
||||
|
||||
// Splits one "key=value" span of a query string at its first '='. A span with
|
||||
// no '=' at all lands entirely in key, leaving val empty, which is how a bare
|
||||
// "?flag" keeps its name.
|
||||
void divide_query_pair(const char *b, const char *e, std::string &key,
|
||||
std::string &val) {
|
||||
divide(b, static_cast<std::size_t>(e - b), '=',
|
||||
[&](const char *lhs_data, std::size_t lhs_size, const char *rhs_data,
|
||||
std::size_t rhs_size) {
|
||||
key.assign(lhs_data, lhs_size);
|
||||
val.assign(rhs_data, rhs_size);
|
||||
});
|
||||
}
|
||||
|
||||
void parse_query_text(const char *data, std::size_t size,
|
||||
Params ¶ms) {
|
||||
std::set<std::string> cache;
|
||||
@@ -3855,12 +3960,7 @@ void parse_query_text(const char *data, std::size_t size,
|
||||
|
||||
std::string key;
|
||||
std::string val;
|
||||
divide(b, static_cast<std::size_t>(e - b), '=',
|
||||
[&](const char *lhs_data, std::size_t lhs_size, const char *rhs_data,
|
||||
std::size_t rhs_size) {
|
||||
key.assign(lhs_data, lhs_size);
|
||||
val.assign(rhs_data, rhs_size);
|
||||
});
|
||||
divide_query_pair(b, e, key, val);
|
||||
|
||||
if (!key.empty()) {
|
||||
params.emplace(decode_query_component(key), decode_query_component(val));
|
||||
@@ -3874,20 +3974,18 @@ void parse_query_text(const std::string &s, Params ¶ms) {
|
||||
|
||||
// Normalize a query string by decoding and re-encoding each key/value pair
|
||||
// while preserving the original parameter order. This avoids double-encoding
|
||||
// and ensures consistent encoding without reordering (unlike Params which
|
||||
// uses std::multimap and sorts keys).
|
||||
// and ensures consistent encoding. It works on the raw string rather than
|
||||
// parsing into Params and re-serializing, because that round trip cannot
|
||||
// reproduce the input: params_to_query_str() always emits '=', so a bare
|
||||
// "flag" would come back as "flag=", and parse_query_text() drops exactly
|
||||
// duplicated pairs.
|
||||
std::string normalize_query_string(const std::string &query) {
|
||||
std::string result;
|
||||
split(query.data(), query.data() + query.size(), '&',
|
||||
[&](const char *b, const char *e) {
|
||||
std::string key;
|
||||
std::string val;
|
||||
divide(b, static_cast<std::size_t>(e - b), '=',
|
||||
[&](const char *lhs_data, std::size_t lhs_size,
|
||||
const char *rhs_data, std::size_t rhs_size) {
|
||||
key.assign(lhs_data, lhs_size);
|
||||
val.assign(rhs_data, rhs_size);
|
||||
});
|
||||
divide_query_pair(b, e, key, val);
|
||||
|
||||
if (!key.empty()) {
|
||||
auto dec_key = decode_query_component(key);
|
||||
@@ -3904,6 +4002,43 @@ std::string normalize_query_string(const std::string &query) {
|
||||
return result;
|
||||
}
|
||||
|
||||
// Build the request target that goes on the wire from a caller-supplied path.
|
||||
// Shared by the buffered send path and the streaming API so that both put the
|
||||
// same bytes in the request line for the same input.
|
||||
std::string encode_request_target(const std::string &target,
|
||||
bool path_encode) {
|
||||
// `substr(0, npos)` yields the whole string, which is what the no-query
|
||||
// case needs.
|
||||
auto query_pos = target.find('?');
|
||||
auto path_part = target.substr(0, query_pos);
|
||||
std::string query_part;
|
||||
if (query_pos != std::string::npos) {
|
||||
query_part = target.substr(query_pos + 1);
|
||||
}
|
||||
|
||||
auto result = path_encode ? encode_path(path_part) : std::move(path_part);
|
||||
|
||||
if (!query_part.empty()) {
|
||||
// When path encoding is disabled the caller has supplied an already-encoded
|
||||
// target and expects the exact bytes to be sent on the wire, so skip
|
||||
// normalization for the query too. Normalizing would decode-then-re-encode
|
||||
// it and corrupt pre-encoded binary payloads (e.g. turning `%20` into `+`,
|
||||
// which a strict RFC 3986 server decodes back as `+`, not a space).
|
||||
if (path_encode) {
|
||||
auto normalized = normalize_query_string(query_part);
|
||||
if (!normalized.empty()) {
|
||||
result += '?';
|
||||
result += normalized;
|
||||
}
|
||||
} else {
|
||||
result += '?';
|
||||
result += query_part;
|
||||
}
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
bool parse_multipart_boundary(const std::string &content_type,
|
||||
std::string &boundary) {
|
||||
std::map<std::string, std::string> params;
|
||||
@@ -4969,21 +5104,8 @@ bool is_field_valid(const std::string &name, const std::string &value) {
|
||||
|
||||
} // namespace fields
|
||||
|
||||
bool perform_websocket_handshake(Stream &strm, const std::string &host,
|
||||
int port, bool is_ssl,
|
||||
const std::string &path,
|
||||
const Headers &headers,
|
||||
bool perform_websocket_handshake(Stream &strm, Request &req,
|
||||
std::string &selected_subprotocol) {
|
||||
// Validate path and host
|
||||
if (!fields::is_field_value(path) || !fields::is_field_value(host)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Validate user-provided headers
|
||||
for (const auto &h : headers) {
|
||||
if (!fields::is_field_valid(h.first, h.second)) { return false; }
|
||||
}
|
||||
|
||||
// Generate random Sec-WebSocket-Key
|
||||
thread_local std::mt19937 rng(std::random_device{}());
|
||||
std::string key_bytes(16, '\0');
|
||||
@@ -4993,19 +5115,30 @@ bool perform_websocket_handshake(Stream &strm, const std::string &host,
|
||||
}
|
||||
auto client_key = base64_encode(key_bytes);
|
||||
|
||||
// Build upgrade request
|
||||
std::string req_str = "GET " + path + " HTTP/1.1\r\n";
|
||||
req_str += "Host: " + make_host_and_port_string(host, port, is_ssl) + "\r\n";
|
||||
req_str += "Upgrade: websocket\r\n";
|
||||
req_str += "Connection: Upgrade\r\n";
|
||||
req_str += "Sec-WebSocket-Key: " + client_key + "\r\n";
|
||||
req_str += "Sec-WebSocket-Version: 13\r\n";
|
||||
for (const auto &h : headers) {
|
||||
req_str += h.first + ": " + h.second + "\r\n";
|
||||
}
|
||||
req_str += "\r\n";
|
||||
req.headers.erase("Upgrade");
|
||||
req.headers.erase("Connection");
|
||||
req.headers.erase("Sec-WebSocket-Key");
|
||||
req.headers.erase("Sec-WebSocket-Version");
|
||||
req.headers.emplace("Upgrade", "websocket");
|
||||
req.headers.emplace("Connection", "Upgrade");
|
||||
req.headers.emplace("Sec-WebSocket-Key", client_key);
|
||||
req.headers.emplace("Sec-WebSocket-Version", "13");
|
||||
|
||||
if (strm.write(req_str.data(), req_str.size()) < 0) { return false; }
|
||||
// Build the request in memory first, like ClientImpl::write_request does.
|
||||
// Writing straight to the socket would leak a request line onto the wire
|
||||
// before check_and_write_headers gets a chance to reject an invalid header,
|
||||
// and would emit one small write per header.
|
||||
BufferStream bstrm;
|
||||
|
||||
if (write_request_line(bstrm, req.method, req.path) < 0) { return false; }
|
||||
|
||||
auto error = Error::Success;
|
||||
if (!check_and_write_headers(bstrm, req.headers, write_headers, error)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
const auto &data = bstrm.get_buffer();
|
||||
if (!write_data(strm, data.data(), data.size())) { return false; }
|
||||
|
||||
// Verify 101 response and Sec-WebSocket-Accept header
|
||||
auto expected_accept = websocket_accept_key(client_key);
|
||||
@@ -5013,6 +5146,39 @@ bool perform_websocket_handshake(Stream &strm, const std::string &host,
|
||||
selected_subprotocol);
|
||||
}
|
||||
|
||||
bool is_ip_address(const std::string &host) {
|
||||
struct in_addr addr4;
|
||||
struct in6_addr addr6;
|
||||
return inet_pton(AF_INET, host.c_str(), &addr4) == 1 ||
|
||||
inet_pton(AF_INET6, host.c_str(), &addr6) == 1;
|
||||
}
|
||||
|
||||
// Resolve where a client should connect for `host`, honoring a user-supplied
|
||||
// hostname-to-address map. `host` itself is never rewritten, so it keeps
|
||||
// supplying the Host header and SNI; only the connection target changes.
|
||||
//
|
||||
// A mapped IP literal goes to `ip`, which keeps create_socket's AI_NUMERICHOST
|
||||
// path. Anything else goes to `connect_host`, which create_socket resolves as
|
||||
// a name, or uses as the socket path when the address family is AF_UNIX. An
|
||||
// absent or empty mapping leaves `host` as the connection target; without the
|
||||
// empty check the value would reach getaddrinfo as a null node and silently
|
||||
// resolve to loopback.
|
||||
void apply_addr_map(const std::map<std::string, std::string> &addr_map,
|
||||
const std::string &host, std::string &connect_host,
|
||||
std::string &ip) {
|
||||
connect_host = host;
|
||||
ip.clear();
|
||||
|
||||
auto it = addr_map.find(host);
|
||||
if (it == addr_map.end() || it->second.empty()) { return; }
|
||||
|
||||
if (is_ip_address(it->second)) {
|
||||
ip = it->second;
|
||||
} else {
|
||||
connect_host = it->second;
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/*
|
||||
@@ -5044,7 +5210,12 @@ public:
|
||||
time_t duration() const override;
|
||||
void set_read_timeout(time_t sec, time_t usec = 0) override;
|
||||
|
||||
// See SocketStream::set_readable_hint().
|
||||
void set_readable_hint() { readable_hint_ = true; }
|
||||
|
||||
private:
|
||||
bool ensure_readable();
|
||||
|
||||
socket_t sock_;
|
||||
tls::session_t session_;
|
||||
time_t read_timeout_sec_;
|
||||
@@ -5053,6 +5224,7 @@ private:
|
||||
time_t write_timeout_usec_;
|
||||
time_t max_timeout_msec_;
|
||||
const std::chrono::time_point<std::chrono::steady_clock> start_time_;
|
||||
bool readable_hint_ = false;
|
||||
};
|
||||
|
||||
#ifdef CPPHTTPLIB_OPENSSL_SUPPORT
|
||||
@@ -5196,13 +5368,6 @@ std::string SHA_512(const std::string &s) {
|
||||
}
|
||||
#endif
|
||||
|
||||
bool is_ip_address(const std::string &host) {
|
||||
struct in_addr addr4;
|
||||
struct in6_addr addr6;
|
||||
return inet_pton(AF_INET, host.c_str(), &addr4) == 1 ||
|
||||
inet_pton(AF_INET6, host.c_str(), &addr6) == 1;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
bool process_server_socket_ssl(
|
||||
const std::atomic<socket_t> &svr_sock, tls::session_t session,
|
||||
@@ -5214,6 +5379,8 @@ bool process_server_socket_ssl(
|
||||
[&](bool close_connection, bool &connection_closed) {
|
||||
SSLSocketStream strm(sock, session, read_timeout_sec, read_timeout_usec,
|
||||
write_timeout_sec, write_timeout_usec);
|
||||
// See the non-TLS path in process_server_socket().
|
||||
strm.set_readable_hint();
|
||||
return callback(strm, close_connection, connection_closed);
|
||||
});
|
||||
}
|
||||
@@ -5665,6 +5832,7 @@ std::string to_string(const Error error) {
|
||||
case Error::UnsupportedAddressFamily: return "Unsupported address family";
|
||||
case Error::HTTPParsing: return "HTTP parsing failed";
|
||||
case Error::InvalidRangeHeader: return "Invalid Range header";
|
||||
case Error::UnsupportedContentEncoding: return "Unsupported Content-Encoding";
|
||||
default: break;
|
||||
}
|
||||
|
||||
@@ -6046,8 +6214,7 @@ std::string Request::get_trailer_value(const std::string &key,
|
||||
}
|
||||
|
||||
size_t Request::get_trailer_value_count(const std::string &key) const {
|
||||
auto r = trailers.equal_range(key);
|
||||
return static_cast<size_t>(std::distance(r.first, r.second));
|
||||
return trailers.count(key);
|
||||
}
|
||||
|
||||
bool Request::has_param(const std::string &key) const {
|
||||
@@ -6071,8 +6238,7 @@ Request::get_param_values(const std::string &key) const {
|
||||
}
|
||||
|
||||
size_t Request::get_param_value_count(const std::string &key) const {
|
||||
auto r = params.equal_range(key);
|
||||
return static_cast<size_t>(std::distance(r.first, r.second));
|
||||
return params.count(key);
|
||||
}
|
||||
|
||||
bool Request::is_multipart_form_data() const {
|
||||
@@ -6105,8 +6271,7 @@ bool MultipartFormData::has_field(const std::string &key) const {
|
||||
}
|
||||
|
||||
size_t MultipartFormData::get_field_count(const std::string &key) const {
|
||||
auto r = fields.equal_range(key);
|
||||
return static_cast<size_t>(std::distance(r.first, r.second));
|
||||
return fields.count(key);
|
||||
}
|
||||
|
||||
FormData MultipartFormData::get_file(const std::string &key,
|
||||
@@ -6129,8 +6294,7 @@ bool MultipartFormData::has_file(const std::string &key) const {
|
||||
}
|
||||
|
||||
size_t MultipartFormData::get_file_count(const std::string &key) const {
|
||||
auto r = files.equal_range(key);
|
||||
return static_cast<size_t>(std::distance(r.first, r.second));
|
||||
return files.count(key);
|
||||
}
|
||||
|
||||
// Multipart FormData writer implementation
|
||||
@@ -6209,8 +6373,7 @@ std::string Response::get_trailer_value(const std::string &key,
|
||||
}
|
||||
|
||||
size_t Response::get_trailer_value_count(const std::string &key) const {
|
||||
auto r = trailers.equal_range(key);
|
||||
return static_cast<size_t>(std::distance(r.first, r.second));
|
||||
return trailers.count(key);
|
||||
}
|
||||
|
||||
void Response::set_redirect(const std::string &url, int stat) {
|
||||
@@ -6306,8 +6469,7 @@ std::string Result::get_request_header_value(const std::string &key,
|
||||
|
||||
size_t
|
||||
Result::get_request_header_value_count(const std::string &key) const {
|
||||
auto r = request_headers_.equal_range(key);
|
||||
return static_cast<size_t>(std::distance(r.first, r.second));
|
||||
return request_headers_.count(key);
|
||||
}
|
||||
|
||||
// Stream implementation
|
||||
@@ -6595,6 +6757,24 @@ bool SocketStream::wait_writable() const {
|
||||
return select_write(sock_, write_timeout_sec_, write_timeout_usec_) > 0;
|
||||
}
|
||||
|
||||
bool SocketStream::ensure_readable() {
|
||||
if (readable_hint_) {
|
||||
readable_hint_ = false;
|
||||
return true;
|
||||
}
|
||||
return wait_readable();
|
||||
}
|
||||
|
||||
const char *SocketStream::buffered_data(size_t &size) const {
|
||||
size = read_buff_content_size_ - read_buff_off_;
|
||||
return size ? read_buff_.data() + read_buff_off_ : nullptr;
|
||||
}
|
||||
|
||||
void SocketStream::consume_buffered(size_t size) {
|
||||
assert(size <= read_buff_content_size_ - read_buff_off_);
|
||||
read_buff_off_ += size;
|
||||
}
|
||||
|
||||
bool SocketStream::is_peer_alive() const {
|
||||
return detail::is_socket_alive(sock_);
|
||||
}
|
||||
@@ -6621,7 +6801,7 @@ ssize_t SocketStream::read(char *ptr, size_t size) {
|
||||
}
|
||||
}
|
||||
|
||||
if (!wait_readable()) {
|
||||
if (!ensure_readable()) {
|
||||
error_ = Error::Timeout;
|
||||
return -1;
|
||||
}
|
||||
@@ -7099,6 +7279,14 @@ bool SSLSocketStream::wait_writable() const {
|
||||
!tls::is_peer_closed(session_, sock_);
|
||||
}
|
||||
|
||||
bool SSLSocketStream::ensure_readable() {
|
||||
if (readable_hint_) {
|
||||
readable_hint_ = false;
|
||||
return true;
|
||||
}
|
||||
return wait_readable();
|
||||
}
|
||||
|
||||
bool SSLSocketStream::is_peer_alive() const {
|
||||
return !tls::is_peer_closed(session_, sock_);
|
||||
}
|
||||
@@ -7111,7 +7299,7 @@ ssize_t SSLSocketStream::read(char *ptr, size_t size) {
|
||||
error_ = Error::ConnectionClosed;
|
||||
}
|
||||
return ret;
|
||||
} else if (wait_readable()) {
|
||||
} else if (ensure_readable()) {
|
||||
tls::TlsError err;
|
||||
auto ret = tls::read(session_, ptr, size, err);
|
||||
if (ret < 0) {
|
||||
@@ -7533,9 +7721,11 @@ void Server::wait_until_ready() const {
|
||||
}
|
||||
|
||||
void Server::stop() noexcept {
|
||||
if (is_running_) {
|
||||
assert(svr_sock_ != INVALID_SOCKET);
|
||||
std::atomic<socket_t> sock(svr_sock_.exchange(INVALID_SOCKET));
|
||||
// Release the listening socket whether or not the accept loop is running:
|
||||
// bind_to_port() without listen_after_bind() still owns the descriptor. The
|
||||
// exchange is what makes this safe to call concurrently with the accept loop.
|
||||
socket_t sock = svr_sock_.exchange(INVALID_SOCKET);
|
||||
if (sock != INVALID_SOCKET) {
|
||||
detail::shutdown_socket(sock);
|
||||
detail::close_socket(sock);
|
||||
}
|
||||
@@ -7697,7 +7887,15 @@ Server::write_content_with_provider(Stream &strm, const Request &req,
|
||||
};
|
||||
|
||||
if (res.content_length_ > 0) {
|
||||
if (req.ranges.empty()) {
|
||||
// Only a 206 response is served as a partial representation, matching the
|
||||
// condition `apply_ranges()` used to decide the Content-Length and the
|
||||
// multipart boundary. Since `detail::range_error()` validates `req.ranges`
|
||||
// only for a 2xx status, slicing under any other status would write a body
|
||||
// that disagrees with the header already sent, from an unchecked offset.
|
||||
auto is_partial =
|
||||
!req.ranges.empty() && res.status == StatusCode::PartialContent_206;
|
||||
|
||||
if (!is_partial) {
|
||||
return detail::write_content(strm, res.content_provider_, 0,
|
||||
res.content_length_, is_shutting_down);
|
||||
} else if (req.ranges.size() == 1) {
|
||||
@@ -8096,7 +8294,14 @@ int Server::bind_internal(const std::string &host, int port,
|
||||
}
|
||||
|
||||
bool Server::listen_internal() {
|
||||
if (is_decommissioned) { return false; }
|
||||
// A stop() between bind and listen leaves nothing to accept on. Report
|
||||
// failure instead of returning success without ever serving, and mark the
|
||||
// server decommissioned the way any failed listen does so that a concurrent
|
||||
// wait_until_ready() wakes up instead of spinning forever.
|
||||
if (is_decommissioned || svr_sock_ == INVALID_SOCKET) {
|
||||
is_decommissioned = true;
|
||||
return false;
|
||||
}
|
||||
|
||||
auto ret = true;
|
||||
is_running_ = true;
|
||||
@@ -8492,11 +8697,17 @@ Server::process_request(Stream &strm, const std::string &remote_addr,
|
||||
return write_response(strm, close_connection, req, res);
|
||||
}
|
||||
|
||||
// RFC 9112 §6.3: Reject requests with both a non-zero Content-Length and
|
||||
// any Transfer-Encoding to prevent request smuggling. Content-Length: 0 is
|
||||
// tolerated for compatibility with existing clients.
|
||||
if (req.get_header_value_u64("Content-Length") > 0 &&
|
||||
req.has_header("Transfer-Encoding")) {
|
||||
// RFC 9112 §6.3: Reject requests whose framing is ambiguous, which would
|
||||
// otherwise let an intermediary and this parser disagree on where the body
|
||||
// ends and enable request smuggling. Two cases: a non-zero Content-Length
|
||||
// alongside any Transfer-Encoding (Content-Length: 0 is tolerated for
|
||||
// compatibility with existing clients), and a Transfer-Encoding whose final
|
||||
// coding is not chunked, which leaves the body length undeterminable. The
|
||||
// latter must not fall through to the "no body" path, or the body bytes are
|
||||
// parsed as the next request on a persistent connection.
|
||||
if (req.has_header("Transfer-Encoding") &&
|
||||
(req.get_header_value_u64("Content-Length") > 0 ||
|
||||
!detail::is_chunked_transfer_encoding(req.headers))) {
|
||||
connection_closed = true;
|
||||
res.status = StatusCode::BadRequest_400;
|
||||
return write_response(strm, close_connection, req, res);
|
||||
@@ -8908,13 +9119,13 @@ socket_t ClientImpl::create_client_socket(Error &error) const {
|
||||
write_timeout_sec_, write_timeout_usec_, interface_, error);
|
||||
}
|
||||
|
||||
// Check is custom IP specified for host_
|
||||
// Check is custom IP or hostname specified for host_
|
||||
std::string connect_host;
|
||||
std::string ip;
|
||||
auto it = addr_map_.find(host_);
|
||||
if (it != addr_map_.end()) { ip = it->second; }
|
||||
detail::apply_addr_map(addr_map_, host_, connect_host, ip);
|
||||
|
||||
return detail::create_client_socket(
|
||||
host_, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_,
|
||||
connect_host, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_,
|
||||
socket_options_, connection_timeout_sec_, connection_timeout_usec_,
|
||||
read_timeout_sec_, read_timeout_usec_, write_timeout_sec_,
|
||||
write_timeout_usec_, interface_, error);
|
||||
@@ -9142,11 +9353,13 @@ void ClientImpl::prepare_default_headers(Request &r, bool for_stream,
|
||||
if (!r.has_header(header.first)) { r.headers.insert(header); }
|
||||
}
|
||||
|
||||
// RFC 9110 5.3 recommends sending control data such as Host first, so
|
||||
// prepend it rather than appending it after the caller's own fields.
|
||||
if (!r.has_header("Host")) {
|
||||
if (address_family_ == AF_UNIX) {
|
||||
r.headers.emplace("Host", "localhost");
|
||||
r.headers.emplace_front("Host", "localhost");
|
||||
} else {
|
||||
r.headers.emplace(
|
||||
r.headers.emplace_front(
|
||||
"Host", detail::make_host_and_port_string(host_, port_, is_ssl()));
|
||||
}
|
||||
}
|
||||
@@ -9197,7 +9410,12 @@ ClientImpl::open_stream(const std::string &method, const std::string &path,
|
||||
handle.response = detail::make_unique<Response>();
|
||||
handle.error = Error::Success;
|
||||
|
||||
auto query_path = params.empty() ? path : append_query_params(path, params);
|
||||
// Encode the target exactly like the buffered send path does, so that the
|
||||
// same `path` produces the same request line through either API.
|
||||
auto raw_query_path =
|
||||
params.empty() ? path : append_query_params(path, params);
|
||||
auto query_path = detail::encode_request_target(raw_query_path, path_encode_);
|
||||
|
||||
handle.connection_ = detail::make_unique<ClientConnection>();
|
||||
|
||||
{
|
||||
@@ -9311,7 +9529,20 @@ ClientImpl::open_stream(const std::string &method, const std::string &path,
|
||||
|
||||
auto content_encoding = handle.response->get_header_value("Content-Encoding");
|
||||
if (!content_encoding.empty()) {
|
||||
// Same policy as prepare_content_receiver(): reject a coding we know about
|
||||
// but were not built with, pass an unrecognized one through as-is.
|
||||
handle.decompressor_ = detail::create_decompressor(content_encoding);
|
||||
if (!handle.decompressor_) {
|
||||
if (detail::is_known_content_encoding(content_encoding)) {
|
||||
handle.error = Error::UnsupportedContentEncoding;
|
||||
handle.response.reset();
|
||||
return handle;
|
||||
}
|
||||
} else if (!handle.decompressor_->is_valid()) {
|
||||
handle.error = Error::Compression;
|
||||
handle.response.reset();
|
||||
return handle;
|
||||
}
|
||||
}
|
||||
|
||||
return handle;
|
||||
@@ -9842,52 +10073,26 @@ bool ClientImpl::write_request(Stream &strm, Request &req,
|
||||
{
|
||||
detail::BufferStream bstrm;
|
||||
|
||||
// Extract path and query from req.path
|
||||
std::string path_part, query_part;
|
||||
// Extract the query from req.path. The encoding itself is delegated to
|
||||
// `encode_request_target`; the raw query is still needed here to decide
|
||||
// between populating `req.params` from it and falling back to building a
|
||||
// query out of caller-supplied `req.params`.
|
||||
auto query_pos = req.path.find('?');
|
||||
if (query_pos != std::string::npos) {
|
||||
path_part = req.path.substr(0, query_pos);
|
||||
query_part = req.path.substr(query_pos + 1);
|
||||
} else {
|
||||
path_part = req.path;
|
||||
query_part = "";
|
||||
}
|
||||
auto query_part = query_pos == std::string::npos
|
||||
? std::string()
|
||||
: req.path.substr(query_pos + 1);
|
||||
|
||||
// Encode path part. If the original `req.path` already contained a
|
||||
// query component, preserve its raw query string (including parameter
|
||||
// order) instead of reparsing and reassembling it which may reorder
|
||||
// parameters due to container ordering (e.g. `Params` uses
|
||||
// `std::multimap`). When there is no query in `req.path`, fall back to
|
||||
// building a query from `req.params` so existing callers that pass
|
||||
// `Params` continue to work.
|
||||
auto path_with_query =
|
||||
path_encode_ ? detail::encode_path(path_part) : path_part;
|
||||
detail::encode_request_target(req.path, path_encode_);
|
||||
|
||||
if (!query_part.empty()) {
|
||||
// Normalize the query string (decode then re-encode) while preserving
|
||||
// the original parameter order. When path encoding is disabled the
|
||||
// caller has supplied an already-encoded target and expects the exact
|
||||
// bytes to be sent on the wire, so skip normalization for the query
|
||||
// too. Normalizing here would decode-then-re-encode the query and
|
||||
// corrupt pre-encoded binary payloads (e.g. turning `%20` into `+`,
|
||||
// which a strict RFC 3986 server decodes back as `+`, not a space).
|
||||
if (path_encode_) {
|
||||
auto normalized = detail::normalize_query_string(query_part);
|
||||
if (!normalized.empty()) { path_with_query += '?' + normalized; }
|
||||
} else {
|
||||
path_with_query += '?' + query_part;
|
||||
}
|
||||
|
||||
// Still populate req.params for handlers/users who read them.
|
||||
// The query already came in through `req.path`; still populate
|
||||
// `req.params` for handlers/users who read them.
|
||||
detail::parse_query_text(query_part, req.params);
|
||||
} else {
|
||||
// No query in path; parse any query_part (empty) and append params
|
||||
// from `req.params` when present (preserves prior behavior for
|
||||
// callers who provide Params separately).
|
||||
detail::parse_query_text(query_part, req.params);
|
||||
if (!req.params.empty()) {
|
||||
path_with_query = append_query_params(path_with_query, req.params);
|
||||
}
|
||||
} else if (!req.params.empty()) {
|
||||
// No query in `req.path`; build one from `req.params` so existing
|
||||
// callers that pass `Params` separately continue to work.
|
||||
path_with_query = append_query_params(path_with_query, req.params);
|
||||
}
|
||||
|
||||
// Write request line and headers
|
||||
@@ -10298,14 +10503,26 @@ bool ClientImpl::process_request(Stream &strm, Request &req,
|
||||
}
|
||||
|
||||
if (res.status != StatusCode::NotModified_304) {
|
||||
int dummy_status;
|
||||
auto content_status = 0;
|
||||
auto max_length = (!has_payload_max_length_ && req.content_receiver)
|
||||
? (std::numeric_limits<size_t>::max)()
|
||||
: payload_max_length_;
|
||||
if (!detail::read_content(strm, res, max_length, dummy_status,
|
||||
if (!detail::read_content(strm, res, max_length, content_status,
|
||||
std::move(progress), std::move(out),
|
||||
decompress_)) {
|
||||
if (error != Error::Canceled) { error = Error::Read; }
|
||||
if (error != Error::Canceled) {
|
||||
// Tell the caller apart from a plain read failure when the body could
|
||||
// not be decoded because of its Content-Encoding.
|
||||
switch (content_status) {
|
||||
case StatusCode::UnsupportedMediaType_415:
|
||||
error = Error::UnsupportedContentEncoding;
|
||||
break;
|
||||
case StatusCode::InternalServerError_500:
|
||||
error = Error::Compression;
|
||||
break;
|
||||
default: error = Error::Read; break;
|
||||
}
|
||||
}
|
||||
output_error_log(error, &req);
|
||||
return false;
|
||||
}
|
||||
@@ -16769,18 +16986,42 @@ bool WebSocketClient::create_stream(std::unique_ptr<Stream> &strm) {
|
||||
return true;
|
||||
}
|
||||
|
||||
void WebSocketClient::prepare_default_headers(Request &req) {
|
||||
#ifdef CPPHTTPLIB_SSL_ENABLED
|
||||
auto is_ssl = is_ssl_;
|
||||
#else
|
||||
auto is_ssl = false;
|
||||
#endif
|
||||
|
||||
if (!req.has_header("Host")) {
|
||||
if (address_family_ == AF_UNIX) {
|
||||
req.headers.emplace("Host", "localhost");
|
||||
} else {
|
||||
req.headers.emplace(
|
||||
"Host", detail::make_host_and_port_string(host_, port_, is_ssl));
|
||||
}
|
||||
}
|
||||
|
||||
#ifndef CPPHTTPLIB_NO_DEFAULT_USER_AGENT
|
||||
if (!req.has_header("User-Agent")) {
|
||||
auto agent = std::string("cpp-httplib/") + CPPHTTPLIB_VERSION;
|
||||
req.set_header("User-Agent", agent);
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
bool WebSocketClient::connect() {
|
||||
if (!is_valid_) { return false; }
|
||||
shutdown_and_close();
|
||||
|
||||
// Check is custom IP specified for host_
|
||||
// Check is custom IP or hostname specified for host_
|
||||
std::string connect_host;
|
||||
std::string ip;
|
||||
auto it = addr_map_.find(host_);
|
||||
if (it != addr_map_.end()) { ip = it->second; }
|
||||
detail::apply_addr_map(addr_map_, host_, connect_host, ip);
|
||||
|
||||
Error error;
|
||||
sock_ = detail::create_client_socket(
|
||||
host_, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_,
|
||||
connect_host, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_,
|
||||
socket_options_, connection_timeout_sec_, connection_timeout_usec_,
|
||||
read_timeout_sec_, read_timeout_usec_, write_timeout_sec_,
|
||||
write_timeout_usec_, interface_, error);
|
||||
@@ -16793,23 +17034,19 @@ bool WebSocketClient::connect() {
|
||||
return false;
|
||||
}
|
||||
|
||||
#ifdef CPPHTTPLIB_SSL_ENABLED
|
||||
auto is_ssl = is_ssl_;
|
||||
#else
|
||||
auto is_ssl = false;
|
||||
#endif
|
||||
Request req;
|
||||
req.method = "GET";
|
||||
req.path = path_;
|
||||
req.headers = headers_;
|
||||
prepare_default_headers(req);
|
||||
|
||||
std::string selected_subprotocol;
|
||||
if (!detail::perform_websocket_handshake(*strm, host_, port_, is_ssl, path_,
|
||||
headers_, selected_subprotocol)) {
|
||||
if (!detail::perform_websocket_handshake(*strm, req, selected_subprotocol)) {
|
||||
shutdown_and_close();
|
||||
return false;
|
||||
}
|
||||
subprotocol_ = std::move(selected_subprotocol);
|
||||
|
||||
Request req;
|
||||
req.method = "GET";
|
||||
req.path = path_;
|
||||
ws_ = std::unique_ptr<WebSocket>(new WebSocket(std::move(strm), req, false,
|
||||
websocket_ping_interval_sec_,
|
||||
websocket_max_missed_pongs_));
|
||||
|
||||
Vendored
+322
-11
@@ -8,8 +8,8 @@
|
||||
#ifndef CPPHTTPLIB_HTTPLIB_H
|
||||
#define CPPHTTPLIB_HTTPLIB_H
|
||||
|
||||
#define CPPHTTPLIB_VERSION "0.51.0"
|
||||
#define CPPHTTPLIB_VERSION_NUM "0x003300"
|
||||
#define CPPHTTPLIB_VERSION "0.52.0"
|
||||
#define CPPHTTPLIB_VERSION_NUM "0x003400"
|
||||
|
||||
#ifdef _WIN32
|
||||
#if defined(_WIN32_WINNT) && _WIN32_WINNT < 0x0A00
|
||||
@@ -182,7 +182,7 @@
|
||||
#endif
|
||||
|
||||
#ifndef CPPHTTPLIB_LISTEN_BACKLOG
|
||||
#define CPPHTTPLIB_LISTEN_BACKLOG 5
|
||||
#define CPPHTTPLIB_LISTEN_BACKLOG 128
|
||||
#endif
|
||||
|
||||
#ifndef CPPHTTPLIB_MAX_LINE_LENGTH
|
||||
@@ -321,6 +321,7 @@ using socket_t = int;
|
||||
#include <functional>
|
||||
#include <iomanip>
|
||||
#include <iostream>
|
||||
#include <iterator>
|
||||
#include <list>
|
||||
#include <map>
|
||||
#include <memory>
|
||||
@@ -333,9 +334,11 @@ using socket_t = int;
|
||||
#include <sys/stat.h>
|
||||
#include <system_error>
|
||||
#include <thread>
|
||||
#include <type_traits>
|
||||
#include <unordered_map>
|
||||
#include <unordered_set>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
// On macOS with a TLS backend, enable Keychain root certificates by default
|
||||
// unless the user explicitly opts out. Not enabled on iOS/tvOS/watchOS since
|
||||
@@ -968,11 +971,291 @@ enum StatusCode {
|
||||
NetworkAuthenticationRequired_511 = 511,
|
||||
};
|
||||
|
||||
using Headers =
|
||||
std::unordered_multimap<std::string, std::string, detail::case_ignore::hash,
|
||||
detail::case_ignore::equal_to>;
|
||||
namespace detail {
|
||||
|
||||
using Params = std::multimap<std::string, std::string>;
|
||||
// A multimap that keeps its entries in the order they were inserted.
|
||||
//
|
||||
// HTTP needs that order in two places. RFC 9110 5.3 makes the order of header
|
||||
// fields sharing a field name significant and forbids a proxy from reordering
|
||||
// them, and a query string's parameters are meaningful in the order the caller
|
||||
// wrote them. Neither standard container expresses it: std::unordered_multimap
|
||||
// gives no ordering guarantee at all for equivalent keys (libstdc++ yields
|
||||
// reverse insertion order, libc++ insertion order), and std::multimap sorts by
|
||||
// key, which would drop control data such as Host behind whatever else the
|
||||
// message carries and alphabetise a query string.
|
||||
//
|
||||
// Entries are therefore kept in a flat vector, in order. Lookup is a linear
|
||||
// scan, which beats hashing for the handful of entries a message carries
|
||||
// (headers are capped at CPPHTTPLIB_HEADER_MAX_COUNT).
|
||||
//
|
||||
// KeyEqual compares keys; it is what makes Headers case-insensitive and
|
||||
// Params, whose parameter names are case-sensitive, not.
|
||||
template <typename Mapped, typename KeyEqual> class insertion_ordered_multimap {
|
||||
public:
|
||||
using key_type = std::string;
|
||||
using mapped_type = Mapped;
|
||||
using value_type = std::pair<std::string, Mapped>;
|
||||
using size_type = std::size_t;
|
||||
using difference_type = std::ptrdiff_t;
|
||||
using reference = value_type &;
|
||||
using const_reference = const value_type &;
|
||||
|
||||
private:
|
||||
static size_type npos() { return static_cast<size_type>(-1); }
|
||||
|
||||
static bool keys_equal(const std::string &a, const std::string &b) {
|
||||
return KeyEqual()(a, b);
|
||||
}
|
||||
|
||||
// Iterating yields every entry in insertion order, but equal_range() and
|
||||
// find() have to walk only the entries sharing one key, which are not
|
||||
// adjacent. Both are the same iterator type: key_idx_ selects between the
|
||||
// two traversals, and since equality compares only the position, an iterator
|
||||
// restricted to one key still compares equal to end().
|
||||
template <typename V> class iterator_t {
|
||||
public:
|
||||
using iterator_category = std::bidirectional_iterator_tag;
|
||||
using value_type = insertion_ordered_multimap::value_type;
|
||||
using difference_type = insertion_ordered_multimap::difference_type;
|
||||
using pointer = V *;
|
||||
using reference = V &;
|
||||
|
||||
iterator_t() : data_(nullptr), idx_(0), size_(0), key_idx_(npos()) {}
|
||||
|
||||
template <typename U,
|
||||
typename std::enable_if<std::is_convertible<U *, V *>::value,
|
||||
int>::type = 0>
|
||||
iterator_t(const iterator_t<U> &rhs)
|
||||
: data_(rhs.data_), idx_(rhs.idx_), size_(rhs.size_),
|
||||
key_idx_(rhs.key_idx_) {}
|
||||
|
||||
reference operator*() const { return data_[idx_]; }
|
||||
pointer operator->() const { return data_ + idx_; }
|
||||
|
||||
iterator_t &operator++() {
|
||||
// Saturating, so that advancing past the last entry of a key (which
|
||||
// get_multimap_value() does when asked for an out-of-range id) stays at
|
||||
// end() instead of running off the container.
|
||||
if (idx_ >= size_) { return *this; }
|
||||
++idx_;
|
||||
if (key_idx_ != npos()) {
|
||||
while (idx_ < size_ && !matches(idx_)) {
|
||||
++idx_;
|
||||
}
|
||||
}
|
||||
return *this;
|
||||
}
|
||||
|
||||
iterator_t operator++(int) {
|
||||
auto tmp = *this;
|
||||
++*this;
|
||||
return tmp;
|
||||
}
|
||||
|
||||
iterator_t &operator--() {
|
||||
if (idx_ == 0) { return *this; }
|
||||
--idx_;
|
||||
if (key_idx_ != npos()) {
|
||||
while (idx_ > 0 && !matches(idx_)) {
|
||||
--idx_;
|
||||
}
|
||||
}
|
||||
return *this;
|
||||
}
|
||||
|
||||
iterator_t operator--(int) {
|
||||
auto tmp = *this;
|
||||
--*this;
|
||||
return tmp;
|
||||
}
|
||||
|
||||
template <typename U> bool operator==(const iterator_t<U> &rhs) const {
|
||||
return idx_ == rhs.idx_;
|
||||
}
|
||||
|
||||
template <typename U> bool operator!=(const iterator_t<U> &rhs) const {
|
||||
return idx_ != rhs.idx_;
|
||||
}
|
||||
|
||||
private:
|
||||
friend class insertion_ordered_multimap;
|
||||
template <typename> friend class iterator_t;
|
||||
|
||||
iterator_t(V *data, size_type idx, size_type size, size_type key_idx)
|
||||
: data_(data), idx_(idx), size_(size), key_idx_(key_idx) {}
|
||||
|
||||
bool matches(size_type i) const {
|
||||
return keys_equal(data_[i].first, data_[key_idx_].first);
|
||||
}
|
||||
|
||||
V *data_;
|
||||
size_type idx_;
|
||||
size_type size_;
|
||||
size_type key_idx_;
|
||||
};
|
||||
|
||||
public:
|
||||
using iterator = iterator_t<value_type>;
|
||||
using const_iterator = iterator_t<const value_type>;
|
||||
|
||||
insertion_ordered_multimap() = default;
|
||||
insertion_ordered_multimap(std::initializer_list<value_type> il)
|
||||
: entries_(il) {}
|
||||
template <typename InputIt>
|
||||
insertion_ordered_multimap(InputIt first, InputIt last)
|
||||
: entries_(first, last) {}
|
||||
|
||||
iterator begin() { return make_iter(0, npos()); }
|
||||
iterator end() { return make_iter(entries_.size(), npos()); }
|
||||
const_iterator begin() const { return make_citer(0, npos()); }
|
||||
const_iterator end() const { return make_citer(entries_.size(), npos()); }
|
||||
const_iterator cbegin() const { return begin(); }
|
||||
const_iterator cend() const { return end(); }
|
||||
|
||||
bool empty() const { return entries_.empty(); }
|
||||
size_type size() const { return entries_.size(); }
|
||||
void clear() { entries_.clear(); }
|
||||
void swap(insertion_ordered_multimap &rhs) { entries_.swap(rhs.entries_); }
|
||||
|
||||
iterator insert(const value_type &val) {
|
||||
entries_.push_back(val);
|
||||
return make_iter(entries_.size() - 1, npos());
|
||||
}
|
||||
|
||||
iterator insert(value_type &&val) {
|
||||
entries_.push_back(std::move(val));
|
||||
return make_iter(entries_.size() - 1, npos());
|
||||
}
|
||||
|
||||
template <typename... Args> iterator emplace(Args &&...args) {
|
||||
entries_.emplace_back(std::forward<Args>(args)...);
|
||||
return make_iter(entries_.size() - 1, npos());
|
||||
}
|
||||
|
||||
// For entries that have to lead the message, such as the Host header field
|
||||
// (RFC 9110 5.3 recommends sending control data first).
|
||||
template <typename... Args> iterator emplace_front(Args &&...args) {
|
||||
entries_.emplace(entries_.begin(), std::forward<Args>(args)...);
|
||||
return make_iter(0, npos());
|
||||
}
|
||||
|
||||
iterator find(const std::string &key) {
|
||||
auto i = index_of(key);
|
||||
return i == npos() ? end() : make_iter(i, i);
|
||||
}
|
||||
|
||||
const_iterator find(const std::string &key) const {
|
||||
auto i = index_of(key);
|
||||
return i == npos() ? end() : make_citer(i, i);
|
||||
}
|
||||
|
||||
size_type count(const std::string &key) const {
|
||||
size_type n = 0;
|
||||
for (const auto &entry : entries_) {
|
||||
if (keys_equal(entry.first, key)) { n++; }
|
||||
}
|
||||
return n;
|
||||
}
|
||||
|
||||
std::pair<iterator, iterator> equal_range(const std::string &key) {
|
||||
auto i = index_of(key);
|
||||
return i == npos() ? std::make_pair(end(), end())
|
||||
: std::make_pair(make_iter(i, i), end());
|
||||
}
|
||||
|
||||
std::pair<const_iterator, const_iterator>
|
||||
equal_range(const std::string &key) const {
|
||||
auto i = index_of(key);
|
||||
return i == npos() ? std::make_pair(end(), end())
|
||||
: std::make_pair(make_citer(i, i), end());
|
||||
}
|
||||
|
||||
size_type erase(const std::string &key) {
|
||||
auto before = entries_.size();
|
||||
entries_.erase(std::remove_if(entries_.begin(), entries_.end(),
|
||||
[&](const value_type &entry) {
|
||||
return keys_equal(entry.first, key);
|
||||
}),
|
||||
entries_.end());
|
||||
return before - entries_.size();
|
||||
}
|
||||
|
||||
iterator erase(const_iterator pos) {
|
||||
entries_.erase(entries_.begin() + static_cast<difference_type>(pos.idx_));
|
||||
return make_iter(pos.idx_, npos());
|
||||
}
|
||||
|
||||
// Erases what iterating [first, last) would actually visit, so erasing an
|
||||
// equal_range() removes only the entries with that key, not everything
|
||||
// positioned between them.
|
||||
iterator erase(const_iterator first, const_iterator last) {
|
||||
auto from = first.idx_;
|
||||
auto to = last.idx_;
|
||||
if (from >= to) { return make_iter(from, npos()); }
|
||||
|
||||
auto begin_it = entries_.begin();
|
||||
auto from_it = begin_it + static_cast<difference_type>(from);
|
||||
auto to_it = begin_it + static_cast<difference_type>(to);
|
||||
|
||||
if (first.key_idx_ == npos()) {
|
||||
entries_.erase(from_it, to_it);
|
||||
} else {
|
||||
auto key = entries_[first.key_idx_].first;
|
||||
auto keep = from_it;
|
||||
for (auto it = from_it; it != to_it; ++it) {
|
||||
if (!keys_equal(it->first, key)) {
|
||||
if (keep != it) { *keep = std::move(*it); }
|
||||
++keep;
|
||||
}
|
||||
}
|
||||
if (keep != to_it) {
|
||||
keep = std::move(to_it, entries_.end(), keep);
|
||||
} else {
|
||||
keep = entries_.end();
|
||||
}
|
||||
entries_.erase(keep, entries_.end());
|
||||
}
|
||||
return make_iter(from, npos());
|
||||
}
|
||||
|
||||
friend bool operator==(const insertion_ordered_multimap &lhs,
|
||||
const insertion_ordered_multimap &rhs) {
|
||||
return lhs.entries_ == rhs.entries_;
|
||||
}
|
||||
|
||||
friend bool operator!=(const insertion_ordered_multimap &lhs,
|
||||
const insertion_ordered_multimap &rhs) {
|
||||
return !(lhs == rhs);
|
||||
}
|
||||
|
||||
private:
|
||||
size_type index_of(const std::string &key) const {
|
||||
for (size_type i = 0; i < entries_.size(); i++) {
|
||||
if (keys_equal(entries_[i].first, key)) { return i; }
|
||||
}
|
||||
return npos();
|
||||
}
|
||||
|
||||
iterator make_iter(size_type idx, size_type key_idx) {
|
||||
return iterator(entries_.data(), idx, entries_.size(), key_idx);
|
||||
}
|
||||
|
||||
const_iterator make_citer(size_type idx, size_type key_idx) const {
|
||||
return const_iterator(entries_.data(), idx, entries_.size(), key_idx);
|
||||
}
|
||||
|
||||
std::vector<value_type> entries_;
|
||||
};
|
||||
|
||||
} // namespace detail
|
||||
|
||||
using Headers =
|
||||
detail::insertion_ordered_multimap<std::string,
|
||||
detail::case_ignore::equal_to>;
|
||||
|
||||
// Query parameter names are case-sensitive, unlike header field names.
|
||||
using Params =
|
||||
detail::insertion_ordered_multimap<std::string, std::equal_to<std::string>>;
|
||||
using Match = std::smatch;
|
||||
|
||||
using DownloadProgress = std::function<bool(size_t current, size_t total)>;
|
||||
@@ -1079,9 +1362,16 @@ struct FormField {
|
||||
std::string content;
|
||||
Headers headers;
|
||||
};
|
||||
using FormFields = std::multimap<std::string, FormField>;
|
||||
// RFC 7578 5.2: a form processor "SHOULD send back results in order" and
|
||||
// "Intermediaries MUST NOT reorder the results", so a handler walking these
|
||||
// should see the parts as they were sent. A std::multimap sorts by field name
|
||||
// and loses that. Field names are case-sensitive, hence std::equal_to rather
|
||||
// than the case-insensitive predicate Headers uses.
|
||||
using FormFields =
|
||||
detail::insertion_ordered_multimap<FormField, std::equal_to<std::string>>;
|
||||
|
||||
using FormFiles = std::multimap<std::string, FormData>;
|
||||
using FormFiles =
|
||||
detail::insertion_ordered_multimap<FormData, std::equal_to<std::string>>;
|
||||
|
||||
struct MultipartFormData {
|
||||
FormFields fields; // Text fields from multipart
|
||||
@@ -1514,6 +1804,7 @@ enum class Error {
|
||||
UnsupportedAddressFamily,
|
||||
HTTPParsing,
|
||||
InvalidRangeHeader,
|
||||
UnsupportedContentEncoding,
|
||||
|
||||
// For internal use only
|
||||
SSLPeerCouldBeClosed_,
|
||||
@@ -1545,6 +1836,18 @@ public:
|
||||
(void)usec;
|
||||
}
|
||||
|
||||
// Bytes already pulled off the socket and sitting in this stream's own
|
||||
// buffer. Exposing them lets a line reader scan for a terminator in one
|
||||
// pass instead of asking for a byte at a time. A stream that does no
|
||||
// buffering of its own reports none, and readers fall back to read().
|
||||
virtual const char *buffered_data(size_t &size) const {
|
||||
size = 0;
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
// Discards `size` bytes previously returned by buffered_data().
|
||||
virtual void consume_buffered(size_t size) { (void)size; }
|
||||
|
||||
ssize_t write(const char *ptr);
|
||||
ssize_t write(const std::string &s);
|
||||
|
||||
@@ -2452,7 +2755,8 @@ protected:
|
||||
std::thread::id socket_requests_are_from_thread_ = std::thread::id();
|
||||
bool socket_should_be_closed_when_request_is_done_ = false;
|
||||
|
||||
// Hostname-IP map
|
||||
// Hostname to connection target map. The value is an IP literal or another
|
||||
// hostname; only the connection target changes, never the identity.
|
||||
std::map<std::string, std::string> addr_map_;
|
||||
|
||||
// Default headers
|
||||
@@ -3154,6 +3458,10 @@ private:
|
||||
std::string make_host_and_port_string(const std::string &host, int port,
|
||||
bool is_ssl);
|
||||
|
||||
template <typename T>
|
||||
bool check_and_write_headers(Stream &strm, Headers &headers, T header_writer,
|
||||
Error &error);
|
||||
|
||||
std::string trim_copy(const std::string &s);
|
||||
|
||||
void divide(
|
||||
@@ -3364,6 +3672,7 @@ public:
|
||||
|
||||
private:
|
||||
void append(char c);
|
||||
void append(const char *data, size_t size);
|
||||
|
||||
Stream &strm_;
|
||||
char *fixed_buffer_;
|
||||
@@ -3992,6 +4301,7 @@ public:
|
||||
private:
|
||||
void shutdown_and_close();
|
||||
bool create_stream(std::unique_ptr<Stream> &strm);
|
||||
void prepare_default_headers(Request &req);
|
||||
|
||||
std::string host_;
|
||||
int port_;
|
||||
@@ -4016,7 +4326,8 @@ private:
|
||||
time_t connection_timeout_usec_ = CPPHTTPLIB_CONNECTION_TIMEOUT_USECOND;
|
||||
std::string interface_;
|
||||
|
||||
// Hostname-IP map
|
||||
// Hostname to connection target map. The value is an IP literal or another
|
||||
// hostname; only the connection target changes, never the identity.
|
||||
std::map<std::string, std::string> addr_map_;
|
||||
|
||||
#ifdef CPPHTTPLIB_SSL_ENABLED
|
||||
|
||||
Reference in New Issue
Block a user