mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-11 07:17:31 -05:00
Compare commits
15
Commits
b11524
...
xsn/json_patch
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f0440d9efc | ||
|
|
79e2e74eb1 | ||
|
|
8e2d31e0eb | ||
|
|
baef3ed9a1 | ||
|
|
f39148a953 | ||
|
|
50e3e3e480 | ||
|
|
64df9183f5 | ||
|
|
6184e92c57 | ||
|
|
8b54361025 | ||
|
|
a518119d30 | ||
|
|
8ae386707b | ||
|
|
e60eff95fd | ||
|
|
ba6439a6b5 | ||
|
|
5e4878e978 | ||
|
|
609290be6b |
@@ -45,6 +45,9 @@ insert_final_newline = unset
|
||||
trim_trailing_whitespace = unset
|
||||
insert_final_newline = unset
|
||||
|
||||
[vendor/**.patch]
|
||||
trim_trailing_whitespace = unset
|
||||
|
||||
[tools/ui/**]
|
||||
indent_style = unset
|
||||
indent_size = unset
|
||||
|
||||
@@ -61,8 +61,7 @@ common_chat_params peg_generator::generate_parser(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = autoparser.build_parser(inputs, parser_generation_prompt);
|
||||
data.parser = parser.save();
|
||||
data.parser = autoparser.build_parser(inputs, parser_generation_prompt);
|
||||
|
||||
// Build grammar if tools are present
|
||||
bool has_tools =
|
||||
@@ -78,7 +77,7 @@ common_chat_params peg_generator::generate_parser(const common_chat_template &
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !has_response_format && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
// Set grammar triggers based on tool section markers (fall back to per-call markers)
|
||||
|
||||
+178
-71
@@ -9,6 +9,7 @@
|
||||
#include "json.h"
|
||||
#include "log.h"
|
||||
#include "parsers/parsers.h"
|
||||
#include "sampling.h"
|
||||
|
||||
#include "jinja/value.h"
|
||||
#include "jinja/runtime.h"
|
||||
@@ -112,38 +113,6 @@ const char * common_chat_role_to_string(common_chat_role role) {
|
||||
return "";
|
||||
}
|
||||
|
||||
json common_chat_msg_delimiters::to_json() const {
|
||||
json result = json::array();
|
||||
for (const auto & d : delimiters) {
|
||||
result.push_back({
|
||||
{ "role", common_chat_role_to_string(d.role) },
|
||||
{ "delimiter", d.delimiter },
|
||||
});
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
common_chat_msg_delimiters common_chat_msg_delimiters_parse(const json & delimiters) {
|
||||
common_chat_msg_delimiters result;
|
||||
|
||||
if (!delimiters.is_array()) {
|
||||
return result;
|
||||
}
|
||||
|
||||
result.delimiters.reserve(delimiters.size());
|
||||
for (const auto & d : delimiters) {
|
||||
if (!d.is_object()) {
|
||||
continue;
|
||||
}
|
||||
result.delimiters.push_back({
|
||||
common_chat_role_from_string(d.value("role", std::string())),
|
||||
d.value("delimiter", std::string()),
|
||||
});
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
void common_chat_msg_delimiters::tokenize(const llama_vocab * vocab) {
|
||||
for (auto & d : delimiters) {
|
||||
d.tokens = common_tokenize(vocab, d.delimiter, false, true);
|
||||
@@ -620,8 +589,11 @@ std::vector<common_chat_tool> common_chat_tools_parse_oaicompat(const json & too
|
||||
}
|
||||
|
||||
common_chat_continuation common_chat_continuation_parse(const common_json & value) {
|
||||
if (value.is_boolean() && value.get<bool>()) {
|
||||
return COMMON_CHAT_CONTINUATION_AUTO;
|
||||
if (value.is_null()) {
|
||||
return COMMON_CHAT_CONTINUATION_NONE;
|
||||
}
|
||||
if (value.is_boolean()) {
|
||||
return value.get<bool>() ? COMMON_CHAT_CONTINUATION_AUTO : COMMON_CHAT_CONTINUATION_NONE;
|
||||
}
|
||||
if (value.is_string()) {
|
||||
auto value_str = value.get<std::string>();
|
||||
@@ -632,7 +604,7 @@ common_chat_continuation common_chat_continuation_parse(const common_json & valu
|
||||
return COMMON_CHAT_CONTINUATION_CONTENT;
|
||||
}
|
||||
}
|
||||
return COMMON_CHAT_CONTINUATION_NONE;
|
||||
throw std::invalid_argument("Invalid continue_final_message: expected a boolean, \"content\" or \"reasoning_content\"");
|
||||
}
|
||||
|
||||
bool common_chat_verify_template(const std::string & tmpl, bool use_jinja) {
|
||||
@@ -1087,41 +1059,55 @@ static json common_chat_extra_context() {
|
||||
return ctx;
|
||||
}
|
||||
|
||||
std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
const common_chat_template & tmpl,
|
||||
const std::string & src,
|
||||
autoparser::generation_params & params) {
|
||||
static common_chat_params common_chat_params_init_lfm2_tokens(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
|
||||
return common_chat_params_init_lfm2(tmpl, inputs, /* tool_list_tokens = */ true);
|
||||
}
|
||||
|
||||
static common_chat_params common_chat_params_init_lfm2_5(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
|
||||
return common_chat_params_init_lfm2(tmpl, inputs, /* tool_list_tokens = */ false);
|
||||
}
|
||||
|
||||
// Older gemma4 templates need their tool responses rewritten before rendering
|
||||
static common_chat_params common_chat_params_init_gemma4_legacy(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
|
||||
auto adjusted = inputs;
|
||||
workaround::convert_tool_responses_gemma4(adjusted.messages);
|
||||
return common_chat_params_init_gemma4(tmpl, adjusted);
|
||||
}
|
||||
|
||||
// Pick the dedicated handler for a template from its source, or null for the autoparser.
|
||||
// Order matters: the first match wins, and later checks assume the earlier ones did not match.
|
||||
static common_chat_params_init_fn common_chat_template_detect_params_init(const std::string & src) {
|
||||
// Ministral/Mistral Large 3 - uses special reasoning structure fixes, can't use autoparser
|
||||
// Note: Mistral Small 3.2 uses [CALL_ID] which Ministral doesn't have, so we can distinguish them
|
||||
if (src.find("[SYSTEM_PROMPT]") != std::string::npos && src.find("[TOOL_CALLS]") != std::string::npos &&
|
||||
src.find("[ARGS]") != std::string::npos && src.find("[CALL_ID]") == std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Ministral/Magistral Large 3\n");
|
||||
return common_chat_params_init_ministral_3(tmpl, params);
|
||||
return common_chat_params_init_ministral_3;
|
||||
}
|
||||
|
||||
// LLM-jp-4.1 - GPT-OSS dialect (spaces after special tokens, <|end|>-separated parallel calls)
|
||||
if (src.find("chat_format=llm-jp-harmony-v1") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: LLM-jp Harmony v1\n");
|
||||
return common_chat_params_init_llm_jp_harmony(tmpl, params);
|
||||
return common_chat_params_init_llm_jp_harmony;
|
||||
}
|
||||
|
||||
// GPT-OSS - has unique channel-based structure that needs dedicated handler
|
||||
if (src.find("<|channel|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: GPT-OSS\n");
|
||||
return common_chat_params_init_gpt_oss(tmpl, params);
|
||||
return common_chat_params_init_gpt_oss;
|
||||
}
|
||||
|
||||
// Muse Glimmer format using " to=<recipient>" recipients and <|eom|>/<|eot|> message terminators.
|
||||
if (src.find("<atem:function_calls>") != std::string::npos && src.find("<|eom|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Muse Glimmer\n");
|
||||
return common_chat_params_init_muse_glimmer(tmpl, params);
|
||||
return common_chat_params_init_muse_glimmer;
|
||||
}
|
||||
|
||||
// Functionary v3.2 - uses recipient-based format with >>>recipient\n{content}
|
||||
// Detection: template has ">>>all" for content and ">>>" prefix for tool calls
|
||||
if (src.find(">>>all") != std::string::npos && src.find(">>>${recipient}") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Functionary v3.2\n");
|
||||
return common_chat_params_init_functionary_v3_2(tmpl, params);
|
||||
return common_chat_params_init_functionary_v3_2;
|
||||
}
|
||||
|
||||
// Kimi K2 Thinking - uses unique tool call ID format: functions.<name>:<index>
|
||||
@@ -1129,14 +1115,14 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
if (src.find("<|tool_calls_section_begin|>") != std::string::npos &&
|
||||
src.find("<|tool_call_begin|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Kimi K2 Thinking\n");
|
||||
return common_chat_params_init_kimi_k2(tmpl, params);
|
||||
return common_chat_params_init_kimi_k2;
|
||||
}
|
||||
|
||||
// Kimi K3 - the <|open|>/<|close|>/<|end_of_msg|> markers are unique to it
|
||||
if (src.find("<|open|>") != std::string::npos && src.find("<|close|>") != std::string::npos &&
|
||||
src.find("<|end_of_msg|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Kimi K3\n");
|
||||
return common_chat_params_init_kimi_k3(tmpl, params);
|
||||
return common_chat_params_init_kimi_k3;
|
||||
}
|
||||
|
||||
// K2 Horizon - <|ifm|im_start|> turns, <ifm|think*> reasoning picked by reasoning_effort and
|
||||
@@ -1144,7 +1130,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
if (src.find("<|ifm|im_start|>") != std::string::npos &&
|
||||
src.find("<ifm|tool_calls>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: K2 Horizon\n");
|
||||
return common_chat_params_init_k2_horizon(tmpl, params);
|
||||
return common_chat_params_init_k2_horizon;
|
||||
}
|
||||
|
||||
// Ling 3.0 / Bailing V3 - <role>X</role> sections with <arg_key>/<arg_value> tagged
|
||||
@@ -1152,7 +1138,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
if (src.find("<role>ASSISTANT</role>") != std::string::npos &&
|
||||
src.find("<arg_key>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Ling 3.0 (Bailing V3)\n");
|
||||
return common_chat_params_init_ling3(tmpl, params);
|
||||
return common_chat_params_init_ling3;
|
||||
}
|
||||
|
||||
// Cohere2 MoE / North Code - marker-wrapped format with <|START_TEXT|> content and
|
||||
@@ -1161,19 +1147,19 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
if (src.find("<|START_TEXT|>") != std::string::npos &&
|
||||
src.find("<|START_ACTION|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Cohere2 MoE\n");
|
||||
return common_chat_params_init_cohere2moe(tmpl, params);
|
||||
return common_chat_params_init_cohere2moe;
|
||||
}
|
||||
|
||||
if (is_lfm2_template(src)) {
|
||||
LOG_DBG("Using specialized template: LFM2\n");
|
||||
return common_chat_params_init_lfm2(tmpl, params, /* tool_list_tokens = */ true);
|
||||
return common_chat_params_init_lfm2_tokens;
|
||||
}
|
||||
|
||||
// LFM2.5 format detection: template uses plain "List of tools: [...]" with no special tokens
|
||||
if (src.find("List of tools: [") != std::string::npos &&
|
||||
src.find("<|tool_list_start|>") == std::string::npos) {
|
||||
LOG_DBG("Using specialized template: LFM2.5\n");
|
||||
return common_chat_params_init_lfm2(tmpl, params, /* tool_list_tokens = */ false);
|
||||
return common_chat_params_init_lfm2_5;
|
||||
}
|
||||
|
||||
// GigaChatV3 format detection
|
||||
@@ -1181,7 +1167,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
src.find("<|message_sep|>") != std::string::npos &&
|
||||
src.find("<|function_call|>") == std::string::npos) {
|
||||
LOG_DBG("Using specialized template: GigaChatV3\n");
|
||||
return common_chat_params_init_gigachat_v3(tmpl, params);
|
||||
return common_chat_params_init_gigachat_v3;
|
||||
}
|
||||
|
||||
// MiniMax-M3: the namespace token "]<]minimax[>[" collides with the autoparser's
|
||||
@@ -1190,7 +1176,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
src.find("<tool_call>") != std::string::npos &&
|
||||
src.find("<invoke name=") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: MiniMax-M3\n");
|
||||
return common_chat_params_init_minimax_m3(tmpl, params);
|
||||
return common_chat_params_init_minimax_m3;
|
||||
}
|
||||
|
||||
// DeepSeek V3.2/V4 format detection: template defines dsml_token and uses it for tool calls.
|
||||
@@ -1201,18 +1187,18 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
(src.find("function_calls") != std::string::npos ||
|
||||
src.find("tool_calls") != std::string::npos)) {
|
||||
LOG_DBG("Using specialized template: DeepSeek V3.2/V4\n");
|
||||
return common_chat_params_init_deepseek_v3_2(tmpl, params);
|
||||
return common_chat_params_init_deepseek_v3_2;
|
||||
}
|
||||
|
||||
// Gemma4 format detection
|
||||
if (src.find("'<|tool_call>call:'") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Gemma4\n");
|
||||
if (src.find("{#- OpenAI Chat Completions:") == std::string::npos) {
|
||||
// apply workarounds if using the older gemma4 templates
|
||||
LOG_WRN("%s: detected an outdated gemma4 chat template, applying compatibility workarounds. "
|
||||
"Consider updating to the official template.\n", __func__);
|
||||
workaround::convert_tool_responses_gemma4(params.messages);
|
||||
return common_chat_params_init_gemma4_legacy;
|
||||
}
|
||||
return common_chat_params_init_gemma4(tmpl, params);
|
||||
return common_chat_params_init_gemma4;
|
||||
}
|
||||
|
||||
// MiniCPM5 - XML tool calls with <function name="..."><param name="...">...</param></function>
|
||||
@@ -1220,14 +1206,14 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
src.find("<function name=\"") != std::string::npos &&
|
||||
src.find("<param name=\"") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: MiniCPM5\n");
|
||||
return common_chat_params_init_minicpm5(tmpl, params);
|
||||
return common_chat_params_init_minicpm5;
|
||||
}
|
||||
|
||||
// TranslateGemma - user content must follow a custom schema with language codes
|
||||
if (src.find("[source_lang_code]") != std::string::npos &&
|
||||
src.find("[target_lang_code]") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: TranslateGemma\n");
|
||||
return common_chat_params_init_translate_gemma(tmpl, params);
|
||||
return common_chat_params_init_translate_gemma;
|
||||
}
|
||||
|
||||
// Qwen3-Coder XML tool calls, also used by Nemotron Nano 3, Qwen3.5 and StepFun-3.5-Flash
|
||||
@@ -1237,10 +1223,51 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
// Exclude models that don't use \n between tags
|
||||
src.find("'<tool_call><function=' ~ tool_call.name ~ '>'") == std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Qwen3-Coder\n");
|
||||
return common_chat_params_init_qwen3_coder(tmpl, params);
|
||||
return common_chat_params_init_qwen3_coder;
|
||||
}
|
||||
|
||||
return std::nullopt;
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
common_chat_template::common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token) {
|
||||
jinja::lexer lexer;
|
||||
auto lexer_res = lexer.tokenize(src);
|
||||
this->prog = jinja::parse_from_tokens(lexer_res);
|
||||
|
||||
this->src = lexer_res.source;
|
||||
this->bos_tok = bos_token;
|
||||
this->eos_tok = eos_token;
|
||||
|
||||
this->caps = jinja::caps_get(prog);
|
||||
// LOG_INF("%s: caps:\n%s\n", __func__, this->caps.to_string().c_str());
|
||||
|
||||
this->params_init = common_chat_template_detect_params_init(this->src);
|
||||
if (this->params_init) {
|
||||
return;
|
||||
}
|
||||
|
||||
// The analysis depends only on the template, so run it once here instead of on every apply.
|
||||
// A failure is kept for apply to report, so a bad template still loads like it did before.
|
||||
try {
|
||||
analysis = std::make_unique<autoparser::autoparser>();
|
||||
analysis->analyze_template(*this);
|
||||
} catch (const std::exception & e) {
|
||||
analysis.reset();
|
||||
analysis_error = e.what();
|
||||
}
|
||||
}
|
||||
|
||||
common_chat_template::~common_chat_template() = default;
|
||||
common_chat_template::common_chat_template(common_chat_template &&) = default;
|
||||
common_chat_template & common_chat_template::operator=(common_chat_template &&) = default;
|
||||
|
||||
std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
const common_chat_template & tmpl,
|
||||
const autoparser::generation_params & params) {
|
||||
if (!tmpl.params_init) {
|
||||
return std::nullopt;
|
||||
}
|
||||
return tmpl.params_init(tmpl, params);
|
||||
}
|
||||
|
||||
static common_chat_params common_chat_templates_apply_jinja(const struct common_chat_templates * tmpls,
|
||||
@@ -1342,21 +1369,23 @@ static common_chat_params common_chat_templates_apply_jinja(const struct common_
|
||||
data.prompt = common_chat_template_direct_apply_impl(tmpl, params_copy);
|
||||
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, params);
|
||||
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
|
||||
auto parser = build_chat_peg_parser([&data](common_chat_peg_builder &p) {
|
||||
data.parser = build_chat_peg_parser([&data](common_chat_peg_builder &p) {
|
||||
return p.literal(data.generation_prompt) << p.content(p.rest());
|
||||
});
|
||||
data.parser = parser.save();
|
||||
return data;
|
||||
}
|
||||
|
||||
if (auto result = common_chat_try_specialized_template(tmpl, src, params)) {
|
||||
if (auto result = common_chat_try_specialized_template(tmpl, params)) {
|
||||
return *result;
|
||||
}
|
||||
|
||||
if (!tmpl.analysis) {
|
||||
throw std::invalid_argument("Unable to generate parser for this template. Automatic parser generation failed: " + tmpl.analysis_error);
|
||||
}
|
||||
|
||||
try {
|
||||
LOG_DBG("%s: using differential autoparser\n", __func__);
|
||||
struct autoparser::autoparser autoparser;
|
||||
autoparser.analyze_template(tmpl);
|
||||
const auto & autoparser = *tmpl.analysis;
|
||||
auto auto_params = autoparser::peg_generator::generate_parser(tmpl, params, autoparser);
|
||||
|
||||
common_chat_msg_delimiters delimiters;
|
||||
@@ -1377,8 +1406,7 @@ static common_chat_params common_chat_templates_apply_jinja(const struct common_
|
||||
auto_params.thinking_end_tags = {std::move(end_tag)};
|
||||
}
|
||||
}
|
||||
common_peg_arena arena;
|
||||
arena.load(auto_params.parser);
|
||||
const auto & arena = auto_params.parser;
|
||||
LOG_DBG("%s: generated parser:\n%s\n\nparser generation prompt: %s\n", __func__, arena.dump(arena.root()).c_str(), auto_params.generation_prompt.c_str());
|
||||
return auto_params;
|
||||
} catch (const std::exception & e) {
|
||||
@@ -1525,9 +1553,10 @@ common_chat_msg common_chat_peg_parse(const common_peg_arena & src_pars
|
||||
const common_chat_input & input,
|
||||
bool is_partial,
|
||||
const common_chat_parser_params & params) {
|
||||
const common_peg_arena & parser = src_parser.empty() ?
|
||||
build_chat_peg_parser([](common_chat_peg_builder & p) { return p.content(p.rest()) + p.end(); }) :
|
||||
src_parser;
|
||||
// both branches must be lvalues, a temporary here would copy the arena on every call
|
||||
static const common_peg_arena content_only =
|
||||
build_chat_peg_parser([](common_chat_peg_builder & p) { return p.content(p.rest()) + p.end(); });
|
||||
const common_peg_arena & parser = src_parser.empty() ? content_only : src_parser;
|
||||
|
||||
if (src_parser.empty()) {
|
||||
LOG_DBG("No parser definition detected, assuming pure content parser.");
|
||||
@@ -1598,6 +1627,84 @@ common_chat_msg common_chat_peg_parse(const common_peg_arena & src_pars
|
||||
return msg;
|
||||
}
|
||||
|
||||
common_chat_session::common_chat_session(const common_chat_templates * tmpls,
|
||||
const llama_vocab * vocab,
|
||||
const common_chat_templates_inputs & inputs,
|
||||
const common_chat_session_params & params) {
|
||||
auto applied = common_chat_templates_apply(tmpls, inputs);
|
||||
|
||||
templated = true;
|
||||
prompt_text = std::move(applied.prompt);
|
||||
result.role = "assistant";
|
||||
|
||||
grammar_text = std::move(applied.grammar);
|
||||
grammar_lazy = applied.grammar_lazy;
|
||||
stops = std::move(applied.additional_stops);
|
||||
generation_prompt_text = applied.generation_prompt;
|
||||
thinking_start = std::move(applied.thinking_start_tag);
|
||||
thinking_ends = std::move(applied.thinking_end_tags);
|
||||
|
||||
parser_params.format = applied.format;
|
||||
parser_params.generation_prompt = vocab ? common_chat_input_tokenize(vocab, applied.generation_prompt)
|
||||
: common_chat_input(applied.generation_prompt);
|
||||
parser_params.debug = params.debug;
|
||||
parser_params.parser = std::move(applied.parser);
|
||||
|
||||
delimiters = std::move(applied.message_delimiters);
|
||||
|
||||
if (vocab) {
|
||||
common_params_sampling resolved;
|
||||
resolved.grammar_lazy = applied.grammar_lazy;
|
||||
common_sampling_add_preserved_tokens(resolved, vocab, applied.preserved_tokens);
|
||||
common_sampling_add_grammar_triggers(resolved, vocab, std::move(applied.grammar_triggers));
|
||||
preserved_tokens = std::move(resolved.preserved_tokens);
|
||||
grammar_triggers = std::move(resolved.grammar_triggers);
|
||||
|
||||
delimiters.tokenize(vocab);
|
||||
} else {
|
||||
grammar_triggers = std::move(applied.grammar_triggers);
|
||||
}
|
||||
|
||||
if (inputs.continue_final_message != COMMON_CHAT_CONTINUATION_NONE && !params.echo) {
|
||||
// start from the prefill so it is not emitted as part of the first delta
|
||||
result = common_chat_parse(input, true, parser_params);
|
||||
}
|
||||
}
|
||||
|
||||
void common_chat_session::apply_sampling(common_params_sampling & sampling) const {
|
||||
if (!templated) {
|
||||
return;
|
||||
}
|
||||
if (!grammar_text.empty()) {
|
||||
sampling.grammar = {COMMON_GRAMMAR_TYPE_TOOL_CALLS, grammar_text};
|
||||
}
|
||||
sampling.grammar_lazy = grammar_lazy;
|
||||
sampling.generation_prompt = generation_prompt_text;
|
||||
sampling.preserved_tokens.insert(preserved_tokens.begin(), preserved_tokens.end());
|
||||
sampling.grammar_triggers.insert(sampling.grammar_triggers.end(), grammar_triggers.begin(), grammar_triggers.end());
|
||||
}
|
||||
|
||||
const common_chat_msg & common_chat_session::feed(const common_chat_input & chunk) {
|
||||
GGML_ASSERT(!finished && "feed() after finish()");
|
||||
input.append(chunk);
|
||||
auto msg = common_chat_parse(input, true, parser_params);
|
||||
if (!msg.empty()) {
|
||||
result = std::move(msg);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
const common_chat_msg & common_chat_session::finish(const common_chat_input & chunk) {
|
||||
GGML_ASSERT(!finished && "finish() called twice");
|
||||
finished = true;
|
||||
input.append(chunk);
|
||||
auto msg = common_chat_parse(input, false, parser_params);
|
||||
if (!msg.empty()) {
|
||||
result = std::move(msg);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
std::map<std::string, bool> common_chat_templates_get_caps(const common_chat_templates * chat_templates) {
|
||||
GGML_ASSERT(chat_templates != nullptr);
|
||||
GGML_ASSERT(chat_templates->template_default != nullptr);
|
||||
|
||||
+79
-27
@@ -22,8 +22,16 @@ struct common_chat_templates;
|
||||
|
||||
namespace autoparser {
|
||||
struct generation_params;
|
||||
struct autoparser;
|
||||
} // namespace autoparser
|
||||
|
||||
struct common_chat_params;
|
||||
struct common_chat_template;
|
||||
|
||||
// Builds the prompt and parser for a template that has a dedicated handler (see common/parsers)
|
||||
using common_chat_params_init_fn = common_chat_params (*)(const common_chat_template & tmpl,
|
||||
const autoparser::generation_params & inputs);
|
||||
|
||||
struct common_chat_tool_call {
|
||||
std::string name;
|
||||
std::string arguments;
|
||||
@@ -54,19 +62,20 @@ struct common_chat_template {
|
||||
std::string eos_tok;
|
||||
std::string src;
|
||||
chat_template_caps caps;
|
||||
// Dedicated handler picked once from the source, null when the differential autoparser is used
|
||||
common_chat_params_init_fn params_init = nullptr;
|
||||
|
||||
common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token) {
|
||||
jinja::lexer lexer;
|
||||
auto lexer_res = lexer.tokenize(src);
|
||||
this->prog = jinja::parse_from_tokens(lexer_res);
|
||||
// Differential analysis, run once here when there is no dedicated handler. Null when there
|
||||
// is one, or when the analysis failed, in which case analysis_error says why.
|
||||
std::unique_ptr<autoparser::autoparser> analysis;
|
||||
std::string analysis_error;
|
||||
|
||||
this->src = lexer_res.source;
|
||||
this->bos_tok = bos_token;
|
||||
this->eos_tok = eos_token;
|
||||
common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token);
|
||||
|
||||
this->caps = jinja::caps_get(prog);
|
||||
// LOG_INF("%s: caps:\n%s\n", __func__, this->caps.to_string().c_str());
|
||||
}
|
||||
// autoparser is incomplete here, so these are defined where it is complete
|
||||
~common_chat_template();
|
||||
common_chat_template(common_chat_template &&);
|
||||
common_chat_template & operator=(common_chat_template &&);
|
||||
|
||||
const std::string & source() const { return src; }
|
||||
const std::string & bos_token() const { return bos_tok; }
|
||||
@@ -209,8 +218,6 @@ struct common_chat_msg_delimiters {
|
||||
|
||||
// split tokens into message spans. skips maps a start index to a length of a region to jump over without matching
|
||||
common_chat_msg_spans split(const llama_tokens & tokens, const std::map<size_t, size_t> & skips = {}) const;
|
||||
|
||||
common_json to_json() const;
|
||||
};
|
||||
|
||||
struct common_chat_tool {
|
||||
@@ -278,7 +285,7 @@ struct common_chat_params {
|
||||
std::vector<common_grammar_trigger> grammar_triggers;
|
||||
std::vector<std::string> preserved_tokens;
|
||||
std::vector<std::string> additional_stops;
|
||||
std::string parser;
|
||||
common_peg_arena parser;
|
||||
common_chat_msg_delimiters message_delimiters;
|
||||
};
|
||||
|
||||
@@ -310,16 +317,10 @@ common_chat_input common_chat_input_tokenize(const llama_vocab * vocab, const st
|
||||
// per-message parsing syntax
|
||||
// should be derived from common_chat_params
|
||||
struct common_chat_parser_params {
|
||||
common_chat_format format = COMMON_CHAT_FORMAT_CONTENT_ONLY;
|
||||
common_reasoning_format reasoning_format = COMMON_REASONING_FORMAT_NONE; // TODO: refactor this to "bool parse_reasoning"
|
||||
// Whether reasoning_content should be inlined in the content (e.g. for reasoning_format=deepseek in stream mode)
|
||||
bool reasoning_in_content = false;
|
||||
common_chat_input generation_prompt;
|
||||
bool parse_tool_calls = true;
|
||||
bool is_continuation = false;
|
||||
bool echo = false; // Include assistant prefilled msg in output
|
||||
bool debug = false; // Enable debug output for PEG parser
|
||||
common_peg_arena parser = {};
|
||||
common_chat_format format = COMMON_CHAT_FORMAT_CONTENT_ONLY;
|
||||
common_chat_input generation_prompt;
|
||||
bool debug = false; // Enable debug output for PEG parser
|
||||
common_peg_arena parser = {};
|
||||
common_chat_parser_params() = default;
|
||||
common_chat_parser_params(const common_chat_params & chat_params) {
|
||||
format = chat_params.format;
|
||||
@@ -365,6 +366,60 @@ const char * common_chat_format_name(common_chat_format format);
|
||||
common_chat_msg common_chat_parse(const common_chat_input & input, bool is_partial, const common_chat_parser_params & params);
|
||||
common_chat_msg common_chat_peg_parse(const common_peg_arena & src_parser, const common_chat_input & input, bool is_partial, const common_chat_parser_params & params);
|
||||
|
||||
struct common_chat_session_params {
|
||||
bool echo = false; // include the assistant prefill in the output when continuing a message
|
||||
bool debug = false; // enable debug output for the PEG parser
|
||||
};
|
||||
|
||||
class common_chat_session {
|
||||
public:
|
||||
common_chat_session() { result.role = "assistant"; }
|
||||
|
||||
common_chat_session(const common_chat_templates * tmpls,
|
||||
const llama_vocab * vocab,
|
||||
const common_chat_templates_inputs & inputs,
|
||||
const common_chat_session_params & params = {});
|
||||
|
||||
const std::string & prompt() const { return prompt_text; }
|
||||
common_chat_format format() const { return parser_params.format; }
|
||||
const common_chat_msg & msg() const { return result; }
|
||||
const common_peg_arena & parser() const { return parser_params.parser; }
|
||||
|
||||
const std::string & grammar() const { return grammar_text; }
|
||||
const std::string & generation_prompt() const { return generation_prompt_text; }
|
||||
const std::string & thinking_start_tag() const { return thinking_start; }
|
||||
const std::vector<std::string> & thinking_end_tags() const { return thinking_ends; }
|
||||
const std::vector<std::string> & additional_stops() const { return stops; }
|
||||
|
||||
const common_chat_msg_delimiters & message_delimiters() const { return delimiters; }
|
||||
|
||||
void apply_sampling(common_params_sampling & sampling) const;
|
||||
|
||||
bool has_template() const { return templated; }
|
||||
|
||||
const common_chat_msg & feed(const common_chat_input & chunk);
|
||||
|
||||
const common_chat_msg & finish(const common_chat_input & chunk = {});
|
||||
|
||||
private:
|
||||
std::string prompt_text;
|
||||
std::string grammar_text;
|
||||
bool grammar_lazy = false;
|
||||
std::vector<common_grammar_trigger> grammar_triggers;
|
||||
std::set<llama_token> preserved_tokens;
|
||||
std::vector<std::string> stops;
|
||||
std::string generation_prompt_text;
|
||||
std::string thinking_start;
|
||||
std::vector<std::string> thinking_ends;
|
||||
|
||||
common_chat_parser_params parser_params;
|
||||
common_chat_msg_delimiters delimiters;
|
||||
common_chat_input input;
|
||||
common_chat_msg result;
|
||||
bool templated = false;
|
||||
bool finished = false;
|
||||
};
|
||||
|
||||
// used by arg and server
|
||||
const char * common_reasoning_format_name(common_reasoning_format format);
|
||||
common_reasoning_format common_reasoning_format_from_name(const std::string & format);
|
||||
@@ -401,8 +456,7 @@ std::string common_chat_template_generation_prompt(
|
||||
|
||||
std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
const common_chat_template & tmpl,
|
||||
const std::string & src,
|
||||
autoparser::generation_params & params);
|
||||
const autoparser::generation_params & params);
|
||||
|
||||
|
||||
// specialized per-task preset
|
||||
@@ -412,5 +466,3 @@ struct common_chat_prompt_preset {
|
||||
};
|
||||
|
||||
common_chat_prompt_preset common_chat_get_asr_prompt(const common_chat_templates * chat_templates);
|
||||
|
||||
common_chat_msg_delimiters common_chat_msg_delimiters_parse(const common_json & delimiters);
|
||||
|
||||
+19
-16
@@ -1195,7 +1195,9 @@ common_decision_type common_get_decision_type(const struct llama_model * model)
|
||||
return common_decision_type_from_string(buf);
|
||||
}
|
||||
|
||||
common_decision_type common_get_decision_type(const std::string & fname) {
|
||||
common_gguf_info common_get_gguf_info(const std::string & fname) {
|
||||
common_gguf_info info;
|
||||
|
||||
struct gguf_init_params gguf_params = {
|
||||
/* .no_alloc = */ true,
|
||||
/* .ctx = */ nullptr,
|
||||
@@ -1203,31 +1205,32 @@ common_decision_type common_get_decision_type(const std::string & fname) {
|
||||
|
||||
gguf_context_ptr gguf_ctx(gguf_init_from_file(fname.c_str(), gguf_params));
|
||||
if (!gguf_ctx) {
|
||||
return COMMON_DECISION_TYPE_UNKNOWN; // missing or unreadable file
|
||||
return info; // missing or unreadable file
|
||||
}
|
||||
|
||||
std::string arch;
|
||||
const int64_t arch_id = gguf_find_key(gguf_ctx.get(), "general.architecture");
|
||||
if (arch_id < 0) {
|
||||
return COMMON_DECISION_TYPE_UNKNOWN; // no architecture in the metadata
|
||||
if (arch_id < 0 || gguf_get_kv_type(gguf_ctx.get(), arch_id) != GGUF_TYPE_STRING) {
|
||||
return info; // no architecture in the metadata
|
||||
}
|
||||
if (gguf_get_kv_type(gguf_ctx.get(), arch_id) != GGUF_TYPE_STRING) {
|
||||
return COMMON_DECISION_TYPE_UNKNOWN; // malformed metadata
|
||||
}
|
||||
arch = gguf_get_val_str(gguf_ctx.get(), arch_id);
|
||||
const std::string arch = gguf_get_val_str(gguf_ctx.get(), arch_id);
|
||||
if (arch.empty()) {
|
||||
return COMMON_DECISION_TYPE_UNKNOWN;
|
||||
return info;
|
||||
}
|
||||
|
||||
const std::string key = arch + ".decision.type";
|
||||
const int64_t type_id = gguf_find_key(gguf_ctx.get(), key.c_str());
|
||||
const int64_t type_id = gguf_find_key(gguf_ctx.get(), (arch + ".decision.type").c_str());
|
||||
if (type_id < 0) {
|
||||
return COMMON_DECISION_TYPE_NONE;
|
||||
info.decision_type = COMMON_DECISION_TYPE_NONE;
|
||||
} else if (gguf_get_kv_type(gguf_ctx.get(), type_id) == GGUF_TYPE_STRING) {
|
||||
info.decision_type = common_decision_type_from_string(gguf_get_val_str(gguf_ctx.get(), type_id));
|
||||
}
|
||||
if (gguf_get_kv_type(gguf_ctx.get(), type_id) != GGUF_TYPE_STRING) {
|
||||
return COMMON_DECISION_TYPE_UNKNOWN; // malformed metadata
|
||||
|
||||
// same key and type as the model loader
|
||||
const int64_t ctx_id = gguf_find_key(gguf_ctx.get(), (arch + ".context_length").c_str());
|
||||
if (ctx_id >= 0 && gguf_get_kv_type(gguf_ctx.get(), ctx_id) == GGUF_TYPE_UINT32) {
|
||||
info.n_ctx_train = gguf_get_val_u32(gguf_ctx.get(), ctx_id);
|
||||
}
|
||||
return common_decision_type_from_string(gguf_get_val_str(gguf_ctx.get(), type_id));
|
||||
|
||||
return info;
|
||||
}
|
||||
|
||||
common_init_result::common_init_result(common_params & params, bool model_only) :
|
||||
|
||||
+7
-3
@@ -970,9 +970,13 @@ enum common_decision_type {
|
||||
|
||||
common_decision_type common_get_decision_type(const struct llama_model * model);
|
||||
|
||||
// same as above, but reads a GGUF file; it does not load the model
|
||||
// returns COMMON_DECISION_TYPE_UNKNOWN if the file is missing, unreadable, or invalid
|
||||
common_decision_type common_get_decision_type(const std::string & fname);
|
||||
// metadata of a GGUF file, read without loading the model
|
||||
struct common_gguf_info {
|
||||
common_decision_type decision_type = COMMON_DECISION_TYPE_UNKNOWN; // UNKNOWN if the file is missing, unreadable, or invalid
|
||||
uint32_t n_ctx_train = 0; // 0 if unknown
|
||||
};
|
||||
|
||||
common_gguf_info common_get_gguf_info(const std::string & fname);
|
||||
|
||||
// note: defines the model, context, samplers, ets. lifetimes
|
||||
struct common_init_result {
|
||||
|
||||
@@ -77,7 +77,7 @@ common_chat_params common_chat_params_init_cohere2moe(const common_chat_template
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PREFIX);
|
||||
auto end = p.end();
|
||||
|
||||
@@ -124,12 +124,10 @@ common_chat_params common_chat_params_init_cohere2moe(const common_chat_template
|
||||
return generation_prompt + reasoning + body + p.optional(p.literal(TURN_END)) + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !has_response_format && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -145,7 +145,7 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
|
||||
bool require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
bool has_tool_calls = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PROMPT);
|
||||
auto end = p.end();
|
||||
|
||||
@@ -256,12 +256,10 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
|
||||
generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = has_tools && !require_tools;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -21,7 +21,7 @@ common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_te
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
// Functionary v3.2 format:
|
||||
// - Normal content: >>>all\n{content}
|
||||
// - Tool calls: >>>function_name\n{json_args}
|
||||
@@ -76,13 +76,11 @@ common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_te
|
||||
return generation_prompt + ret;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
// Grammar trigger for when the model starts outputting a tool call
|
||||
|
||||
@@ -198,7 +198,7 @@ common_chat_params common_chat_params_init_gemma4(const common_chat_template &
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto start = p.rule("start", p.optional(p.literal("<|turn>model\n")));
|
||||
|
||||
if (extract_reasoning) {
|
||||
@@ -290,12 +290,10 @@ common_chat_params common_chat_params_init_gemma4(const common_chat_template &
|
||||
return start + p.one_or_more(message);
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -25,7 +25,7 @@ common_chat_params common_chat_params_init_gigachat_v3(
|
||||
auto include_grammar = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
|
||||
const auto *tool_call_start_prefix = "<|message_sep|>\n\nfunction call<|role_sep|>\n";
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto ret = p.eps();
|
||||
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
|
||||
// Build a choice of all available tools
|
||||
@@ -60,13 +60,11 @@ common_chat_params common_chat_params_init_gigachat_v3(
|
||||
return p.literal("assistant<|role_sep|>\n") + ret;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -45,8 +45,7 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
|
||||
data.thinking_start_tag = "<|channel|>analysis<|message|>";
|
||||
data.thinking_end_tags = {"<|end|>"};
|
||||
|
||||
// These special tokens are required to parse properly, so we include them
|
||||
// even if parse_tool_calls is false.
|
||||
// These special tokens are required to parse properly
|
||||
data.preserved_tokens = {
|
||||
"<|channel|>", "<|constrain|>", "<|message|>", "<|start|>", "<|end|>",
|
||||
};
|
||||
@@ -68,7 +67,7 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto start = p.rule("start", p.literal("<|start|>assistant"));
|
||||
auto end = p.rule("end", p.literal("<|end|>"));
|
||||
auto content = p.rule("message-content", p.until("<|end|>"));
|
||||
@@ -138,12 +137,10 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
|
||||
return p.zero_or_more(start + any) + start + (final_msg | unsolicited);
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -77,7 +77,7 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PREFIX);
|
||||
|
||||
auto think_end = p.choice();
|
||||
@@ -174,12 +174,10 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template
|
||||
return generation_prompt + (reasoning << content << tool_calls);
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED);
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
if (data.grammar_lazy) {
|
||||
|
||||
@@ -48,7 +48,7 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
// Kimi K2 Thinking format:
|
||||
// - Reasoning: <think>{reasoning}</think>
|
||||
// - Content: text after reasoning
|
||||
@@ -111,12 +111,10 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
|
||||
return generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -66,7 +66,7 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto end = p.end();
|
||||
|
||||
auto start = p.optional(p.literal(MSG_START));
|
||||
@@ -151,12 +151,10 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
|
||||
return start + reasoning + response + tools + trailer + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -64,7 +64,7 @@ common_chat_params common_chat_params_init_lfm2(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PROMPT);
|
||||
auto end = p.end();
|
||||
|
||||
@@ -93,12 +93,10 @@ common_chat_params common_chat_params_init_lfm2(const common_chat_template &
|
||||
return generation_prompt + reasoning + content + tool_calls + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -80,7 +80,7 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
|
||||
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);
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto end = p.end();
|
||||
|
||||
// the effective parse input is generation_prompt + model output, so the
|
||||
@@ -185,12 +185,10 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
|
||||
return opener + reasoning + content + tools + tail + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !has_response_format && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -48,8 +48,7 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
|
||||
data.thinking_start_tag = "<|channel|>analysis<|message|>";
|
||||
data.thinking_end_tags = {"<|end|>"};
|
||||
|
||||
// These special tokens are required to parse properly, so we include them
|
||||
// even if parse_tool_calls is false.
|
||||
// These special tokens are required to parse properly
|
||||
data.preserved_tokens = {
|
||||
"<|channel|>", "<|constrain|>", "<|message|>", "<|start|>", "<|end|>",
|
||||
};
|
||||
@@ -71,7 +70,7 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
// tokenizer space after special tokens; not p.space() since GBNF `space` allows one space only
|
||||
auto sp = p.chars("[ ]", 0, -1);
|
||||
auto channel_tag = p.literal("<|channel|>") + sp;
|
||||
@@ -144,12 +143,10 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
|
||||
return p.zero_or_more(start + any) + start + final_msg;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -46,7 +46,7 @@ common_chat_params common_chat_params_init_minicpm5(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal("<|im_start|>assistant\n");
|
||||
|
||||
auto reasoning = p.eps();
|
||||
@@ -113,12 +113,10 @@ common_chat_params common_chat_params_init_minicpm5(const common_chat_template &
|
||||
return generation_prompt + reasoning + p.content(p.rest()) + p.end();
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -56,7 +56,7 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.prefix(GEN_PROMPT, THINK_START);
|
||||
auto end = p.end();
|
||||
|
||||
@@ -213,12 +213,10 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
|
||||
return generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -72,7 +72,7 @@ common_chat_params common_chat_params_init_ministral_3(const common_chat_templat
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.eps();
|
||||
auto reasoning =
|
||||
extract_reasoning ? p.optional("[THINK]" + p.reasoning(p.until("[/THINK]")) + "[/THINK]") : p.eps();
|
||||
@@ -108,13 +108,11 @@ common_chat_params common_chat_params_init_ministral_3(const common_chat_templat
|
||||
return generation_prompt + (reasoning << p.content(p.rest()));
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -48,7 +48,7 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
|
||||
// Constrained grammar whenever tools are offered or a response format is requested.
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto start = p.rule("start", p.literal("<|start|>assistant"));
|
||||
|
||||
if (!extract_reasoning && !include_grammar) {
|
||||
@@ -131,12 +131,10 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
|
||||
return p.zero_or_more(start + analysis) + start + final_msg;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
data.grammar_triggers = {
|
||||
{ COMMON_GRAMMAR_TRIGGER_TYPE_PATTERN,
|
||||
|
||||
@@ -71,7 +71,7 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
|
||||
});
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PREFIX);
|
||||
|
||||
auto reasoning = p.eps();
|
||||
@@ -174,13 +174,11 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
|
||||
return generation_prompt + (reasoning << p.content(p.rest()));
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
if (data.grammar_lazy) {
|
||||
|
||||
@@ -54,10 +54,9 @@ common_chat_params common_chat_params_init_translate_gemma(
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
return p.literal(data.generation_prompt) << p.content(p.rest());
|
||||
});
|
||||
data.parser = parser.save();
|
||||
|
||||
return data;
|
||||
}
|
||||
|
||||
@@ -1814,309 +1814,6 @@ void common_peg_arena::build_grammar(const common_grammar_builder & builder, boo
|
||||
}
|
||||
}
|
||||
|
||||
static common_json serialize_parser_variant(const common_peg_parser_variant & variant) {
|
||||
using json = common_json;
|
||||
|
||||
return std::visit([](const auto & p) -> json {
|
||||
using T = std::decay_t<decltype(p)>;
|
||||
|
||||
if constexpr (std::is_same_v<T, common_peg_epsilon_parser>) {
|
||||
return json{{"type", "epsilon"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_start_parser>) {
|
||||
return json{{"type", "start"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_end_parser>) {
|
||||
return json{{"type", "end"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_literal_parser>) {
|
||||
return json{{"type", "literal"}, {"literal", p.literal}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_sequence_parser>) {
|
||||
return json{{"type", "sequence"}, {"children", p.children}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_choice_parser>) {
|
||||
return json{{"type", "choice"}, {"children", p.children}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_repetition_parser>) {
|
||||
return json{
|
||||
{"type", "repetition"},
|
||||
{"child", p.child},
|
||||
{"min_count", p.min_count},
|
||||
{"max_count", p.max_count}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_and_parser>) {
|
||||
return json{{"type", "and"}, {"child", p.child}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_not_parser>) {
|
||||
return json{{"type", "not"}, {"child", p.child}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_any_parser>) {
|
||||
return json{{"type", "any"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_space_parser>) {
|
||||
return json{{"type", "space"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_chars_parser>) {
|
||||
json ranges = json::array();
|
||||
for (const auto & range : p.ranges) {
|
||||
ranges.push_back({{"start", range.start}, {"end", range.end}});
|
||||
}
|
||||
return json{
|
||||
{"type", "chars"},
|
||||
{"pattern", p.pattern},
|
||||
{"ranges", ranges},
|
||||
{"negated", p.negated},
|
||||
{"min_count", p.min_count},
|
||||
{"max_count", p.max_count}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_string_parser>) {
|
||||
return json{{"type", "string"}, {"delimiter", std::string(1, p.delimiter)}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_until_parser>) {
|
||||
return json{{"type", "until"}, {"delimiters", p.delimiters}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_schema_parser>) {
|
||||
return json{
|
||||
{"type", "schema"},
|
||||
{"child", p.child},
|
||||
{"name", p.name},
|
||||
{"raw", p.raw}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_rule_parser>) {
|
||||
return json{
|
||||
{"type", "rule"},
|
||||
{"name", p.name},
|
||||
{"child", p.child},
|
||||
{"trigger", p.trigger}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_ref_parser>) {
|
||||
return json{{"type", "ref"}, {"name", p.name}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_atomic_parser>) {
|
||||
return json{{"type", "atomic"}, {"child", p.child}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_tag_parser>) {
|
||||
return json{
|
||||
{"type", "tag"},
|
||||
{"child", p.child},
|
||||
{"tag", p.tag}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_gbnf_parser>) {
|
||||
return json{{"type", "gbnf"}, {"child", p.child}, {"grammar", p.grammar}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_ac_parser>) {
|
||||
return json{{"type", "ac"}, {"child", p.child}, {"delimiters", p.delimiters}};
|
||||
}
|
||||
}, variant);
|
||||
}
|
||||
|
||||
common_json common_peg_arena::to_json() const {
|
||||
auto parsers = common_json::array();
|
||||
for (const auto & parser : parsers_) {
|
||||
parsers.push_back(serialize_parser_variant(parser));
|
||||
}
|
||||
return common_json{
|
||||
{"parsers", parsers},
|
||||
{"rules", rules_},
|
||||
{"root", root_}
|
||||
};
|
||||
}
|
||||
|
||||
static common_peg_parser_variant deserialize_parser_variant(const common_json & j) {
|
||||
if (!j.contains("type") || !j["type"].is_string()) {
|
||||
throw std::runtime_error("Parser variant JSON missing or invalid 'type' field");
|
||||
}
|
||||
|
||||
std::string type = j["type"];
|
||||
|
||||
if (type == "epsilon") {
|
||||
return common_peg_epsilon_parser{};
|
||||
}
|
||||
if (type == "start") {
|
||||
return common_peg_start_parser{};
|
||||
}
|
||||
if (type == "end") {
|
||||
return common_peg_end_parser{};
|
||||
}
|
||||
if (type == "literal") {
|
||||
if (!j.contains("literal") || !j["literal"].is_string()) {
|
||||
throw std::runtime_error("literal parser missing or invalid 'literal' field");
|
||||
}
|
||||
return common_peg_literal_parser{j["literal"]};
|
||||
}
|
||||
if (type == "sequence") {
|
||||
if (!j.contains("children") || !j["children"].is_array()) {
|
||||
throw std::runtime_error("sequence parser missing or invalid 'children' field");
|
||||
}
|
||||
return common_peg_sequence_parser{j["children"].get<std::vector<common_peg_parser_id>>()};
|
||||
}
|
||||
if (type == "choice") {
|
||||
if (!j.contains("children") || !j["children"].is_array()) {
|
||||
throw std::runtime_error("choice parser missing or invalid 'children' field");
|
||||
}
|
||||
return common_peg_choice_parser{j["children"].get<std::vector<common_peg_parser_id>>()};
|
||||
}
|
||||
if (type == "repetition") {
|
||||
if (!j.contains("child") || !j.contains("min_count") || !j.contains("max_count")) {
|
||||
throw std::runtime_error("repetition parser missing required fields");
|
||||
}
|
||||
return common_peg_repetition_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["min_count"].get<int>(),
|
||||
j["max_count"].get<int>()
|
||||
};
|
||||
}
|
||||
if (type == "and") {
|
||||
if (!j.contains("child")) {
|
||||
throw std::runtime_error("and parser missing 'child' field");
|
||||
}
|
||||
return common_peg_and_parser{j["child"].get<common_peg_parser_id>()};
|
||||
}
|
||||
if (type == "not") {
|
||||
if (!j.contains("child")) {
|
||||
throw std::runtime_error("not parser missing 'child' field");
|
||||
}
|
||||
return common_peg_not_parser{j["child"].get<common_peg_parser_id>()};
|
||||
}
|
||||
if (type == "any") {
|
||||
return common_peg_any_parser{};
|
||||
}
|
||||
if (type == "space") {
|
||||
return common_peg_space_parser{};
|
||||
}
|
||||
if (type == "chars") {
|
||||
if (!j.contains("pattern") || !j.contains("ranges") || !j.contains("negated") ||
|
||||
!j.contains("min_count") || !j.contains("max_count")) {
|
||||
throw std::runtime_error("chars parser missing required fields");
|
||||
}
|
||||
common_peg_chars_parser parser;
|
||||
parser.pattern = j["pattern"];
|
||||
parser.negated = j["negated"].get<bool>();
|
||||
parser.min_count = j["min_count"].get<int>();
|
||||
parser.max_count = j["max_count"].get<int>();
|
||||
for (const auto & range_json : j["ranges"]) {
|
||||
if (!range_json.contains("start") || !range_json.contains("end")) {
|
||||
throw std::runtime_error("char_range missing 'start' or 'end' field");
|
||||
}
|
||||
parser.ranges.push_back({
|
||||
range_json["start"].get<uint32_t>(),
|
||||
range_json["end"].get<uint32_t>()
|
||||
});
|
||||
}
|
||||
return parser;
|
||||
}
|
||||
if (type == "string") {
|
||||
if (!j.contains("delimiter")) {
|
||||
throw std::runtime_error("string parser missing delimiter field.");
|
||||
}
|
||||
std::string delimiter = j["delimiter"];
|
||||
if (delimiter.empty()) {
|
||||
throw std::runtime_error("string parser delimiter is empty.");
|
||||
}
|
||||
return common_peg_string_parser{delimiter[0]};
|
||||
}
|
||||
if (type == "until") {
|
||||
if (!j.contains("delimiters") || !j["delimiters"].is_array()) {
|
||||
throw std::runtime_error("until parser missing or invalid 'delimiters' field");
|
||||
}
|
||||
return common_peg_until_parser{j["delimiters"].get<std::vector<std::string>>()};
|
||||
}
|
||||
if (type == "schema") {
|
||||
if (!j.contains("child") || !j.contains("name") || !j.contains("raw")) {
|
||||
throw std::runtime_error("schema parser missing required fields");
|
||||
}
|
||||
common_peg_schema_parser parser;
|
||||
parser.child = j["child"].get<common_peg_parser_id>();
|
||||
parser.name = j["name"];
|
||||
parser.raw = j["raw"].get<bool>();
|
||||
return parser;
|
||||
}
|
||||
if (type == "rule") {
|
||||
if (!j.contains("name") || !j.contains("child") || !j.contains("trigger")) {
|
||||
throw std::runtime_error("rule parser missing required fields");
|
||||
}
|
||||
return common_peg_rule_parser{
|
||||
j["name"].get<std::string>(),
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["trigger"].get<bool>()
|
||||
};
|
||||
}
|
||||
if (type == "ref") {
|
||||
if (!j.contains("name") || !j["name"].is_string()) {
|
||||
throw std::runtime_error("ref parser missing or invalid 'name' field");
|
||||
}
|
||||
return common_peg_ref_parser{j["name"]};
|
||||
}
|
||||
if (type == "atomic") {
|
||||
if (!j.contains("child")) {
|
||||
throw std::runtime_error("tag parser missing required fields");
|
||||
}
|
||||
return common_peg_atomic_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
};
|
||||
}
|
||||
if (type == "tag") {
|
||||
if (!j.contains("child") || !j.contains("tag")) {
|
||||
throw std::runtime_error("tag parser missing required fields");
|
||||
}
|
||||
return common_peg_tag_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["tag"].get<std::string>(),
|
||||
};
|
||||
}
|
||||
|
||||
if (type == "gbnf") {
|
||||
if (!j.contains("child") || !j.contains("grammar")) {
|
||||
throw std::runtime_error("gbnf parser missing required fields");
|
||||
}
|
||||
return common_peg_gbnf_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["grammar"].get<std::string>(),
|
||||
};
|
||||
}
|
||||
|
||||
if (type == "ac") {
|
||||
if (!j.contains("child") || !j.contains("delimiters") || !j["delimiters"].is_array() || j["delimiters"].empty()) {
|
||||
throw std::runtime_error("ac parser requires 'child' and a non-empty 'delimiters' array");
|
||||
}
|
||||
return common_peg_ac_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["delimiters"].get<std::vector<std::string>>(),
|
||||
};
|
||||
}
|
||||
|
||||
throw std::runtime_error("Unknown parser type: " + type);
|
||||
}
|
||||
|
||||
common_peg_arena common_peg_arena::from_json(const common_json & j) {
|
||||
if (!j.contains("parsers") || !j["parsers"].is_array()) {
|
||||
throw std::runtime_error("JSON missing or invalid 'parsers' array");
|
||||
}
|
||||
if (!j.contains("rules") || !j["rules"].is_object()) {
|
||||
throw std::runtime_error("JSON missing or invalid 'rules' object");
|
||||
}
|
||||
if (!j.contains("root")) {
|
||||
throw std::runtime_error("JSON missing 'root' field");
|
||||
}
|
||||
|
||||
common_peg_arena arena;
|
||||
|
||||
const auto & parsers_json = j["parsers"];
|
||||
arena.parsers_.reserve(parsers_json.size());
|
||||
for (const auto & parser_json : parsers_json) {
|
||||
arena.parsers_.push_back(deserialize_parser_variant(parser_json));
|
||||
}
|
||||
|
||||
arena.rules_ = j["rules"].get<std::unordered_map<std::string, common_peg_parser_id>>();
|
||||
|
||||
for (const auto & [name, id] : arena.rules_) {
|
||||
if (id >= arena.parsers_.size()) {
|
||||
throw std::runtime_error("Rule '" + name + "' references invalid parser ID: " + std::to_string(id));
|
||||
}
|
||||
}
|
||||
|
||||
arena.root_ = j["root"].get<common_peg_parser_id>();
|
||||
if (arena.root_ != COMMON_PEG_INVALID_PARSER_ID && arena.root_ >= arena.parsers_.size()) {
|
||||
throw std::runtime_error("Root references invalid parser ID: " + std::to_string(arena.root_));
|
||||
}
|
||||
|
||||
return arena;
|
||||
}
|
||||
|
||||
std::string common_peg_arena::save() const {
|
||||
return to_json().dump();
|
||||
}
|
||||
|
||||
void common_peg_arena::load(const std::string & data) {
|
||||
*this = from_json(common_json::parse(data));
|
||||
}
|
||||
|
||||
common_peg_arena build_peg_parser(const std::function<common_peg_parser(common_peg_parser_builder & builder)> & fn) {
|
||||
common_peg_parser_builder builder;
|
||||
builder.set_root(fn(builder));
|
||||
|
||||
@@ -357,12 +357,6 @@ class common_peg_arena {
|
||||
|
||||
std::string dump(common_peg_parser_id id) const;
|
||||
|
||||
common_json to_json() const;
|
||||
static common_peg_arena from_json(const common_json & j);
|
||||
|
||||
std::string save() const;
|
||||
void load(const std::string & data);
|
||||
|
||||
friend class common_peg_parser_builder;
|
||||
|
||||
private:
|
||||
|
||||
@@ -1050,3 +1050,41 @@ std::vector<common_sampler_type> common_sampler_types_from_chars(const std::stri
|
||||
|
||||
return samplers;
|
||||
}
|
||||
|
||||
void common_sampling_add_preserved_tokens(common_params_sampling & sampling, const llama_vocab * vocab, const std::vector<std::string> & tokens) {
|
||||
GGML_ASSERT(vocab != nullptr);
|
||||
for (const auto & t : tokens) {
|
||||
auto ids = common_tokenize(vocab, t, false, true);
|
||||
if (ids.size() == 1) {
|
||||
sampling.preserved_tokens.insert(ids[0]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void common_sampling_add_grammar_triggers(common_params_sampling & sampling, const llama_vocab * vocab, std::vector<common_grammar_trigger> triggers) {
|
||||
GGML_ASSERT(vocab != nullptr);
|
||||
for (auto & trigger : triggers) {
|
||||
if (trigger.type == COMMON_GRAMMAR_TRIGGER_TYPE_WORD) {
|
||||
const auto & word = trigger.value;
|
||||
auto ids = common_tokenize(vocab, word, false, true);
|
||||
if (ids.size() == 1) {
|
||||
auto token = ids[0];
|
||||
if (std::find(sampling.preserved_tokens.begin(), sampling.preserved_tokens.end(), (llama_token) token) == sampling.preserved_tokens.end()) {
|
||||
throw std::runtime_error("Grammar trigger word should be marked as preserved token: " + word);
|
||||
}
|
||||
common_grammar_trigger token_trigger;
|
||||
token_trigger.type = COMMON_GRAMMAR_TRIGGER_TYPE_TOKEN;
|
||||
token_trigger.value = word;
|
||||
token_trigger.token = token;
|
||||
sampling.grammar_triggers.push_back(std::move(token_trigger));
|
||||
} else {
|
||||
sampling.grammar_triggers.push_back({COMMON_GRAMMAR_TRIGGER_TYPE_WORD, word});
|
||||
}
|
||||
} else {
|
||||
sampling.grammar_triggers.push_back(std::move(trigger));
|
||||
}
|
||||
}
|
||||
if (sampling.grammar_lazy && sampling.grammar_triggers.empty()) {
|
||||
throw std::runtime_error("Error: no triggers set for lazy grammar!");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -118,6 +118,12 @@ std::string common_sampler_type_to_str(enum common_sampler_type cnstr);
|
||||
std::vector<enum common_sampler_type> common_sampler_types_from_names(const std::vector<std::string> & names);
|
||||
std::vector<enum common_sampler_type> common_sampler_types_from_chars(const std::string & chars);
|
||||
|
||||
// add the strings that are a single token in the vocab to the preserved tokens
|
||||
void common_sampling_add_preserved_tokens(common_params_sampling & sampling, const llama_vocab * vocab, const std::vector<std::string> & tokens);
|
||||
|
||||
// add grammar triggers, a trigger word that is a single token becomes a token trigger and must be a preserved token
|
||||
void common_sampling_add_grammar_triggers(common_params_sampling & sampling, const llama_vocab * vocab, std::vector<common_grammar_trigger> triggers);
|
||||
|
||||
llama_sampler * llama_sampler_init_llg(const llama_vocab * vocab,
|
||||
const char * grammar_kind, const char * grammar_data);
|
||||
|
||||
|
||||
+42
-43
@@ -439,6 +439,25 @@ class ModelBase:
|
||||
|
||||
return (unpacked * scale.unsqueeze(-1).float()).reshape(shape)
|
||||
|
||||
def dequant_fp8() -> None:
|
||||
for name in self.model_tensors.keys():
|
||||
if name.endswith(".weight_scale"):
|
||||
weight_name = name.removesuffix("_scale")
|
||||
if weight_name not in self.model_tensors:
|
||||
tensors_to_remove.append(name)
|
||||
continue
|
||||
w = self.model_tensors[weight_name]
|
||||
s = self.model_tensors[name]
|
||||
is_fp8_weight = False
|
||||
if self._fp8_as_q8:
|
||||
is_fp8_weight = w().dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
|
||||
self.model_tensors[weight_name] = lambda w=w, s=s: dequant_simple(w(), s(), None)
|
||||
tensors_to_remove.append(name)
|
||||
if is_fp8_weight:
|
||||
self._fp8_dequantized.add(weight_name)
|
||||
if name.endswith((".input_scale", ".k_scale", ".v_scale")):
|
||||
tensors_to_remove.append(name)
|
||||
|
||||
if quant_method == "bitnet":
|
||||
for name in self.model_tensors.keys():
|
||||
if name.endswith(".weight_scale"):
|
||||
@@ -498,18 +517,14 @@ class ModelBase:
|
||||
elif quant_method == "compressed-tensors":
|
||||
quant_format = quant_config["format"]
|
||||
groups = quant_config["config_groups"]
|
||||
nvfp4_compressed_tensors = (
|
||||
quant_format == "nvfp4-pack-quantized"
|
||||
or quant_format == "mixed-precision"
|
||||
and bool(groups)
|
||||
and all(g.get("format") == "nvfp4-pack-quantized" for g in groups.values() if isinstance(g, dict))
|
||||
)
|
||||
nvfp4_compressed_tensors = self._is_nvfp4_compressed_tensors(quant_method, quant_format, groups)
|
||||
|
||||
if len(groups) > 1 and not nvfp4_compressed_tensors:
|
||||
if nvfp4_compressed_tensors:
|
||||
dequant_fp8()
|
||||
elif len(groups) > 1:
|
||||
raise NotImplementedError("Can't handle multiple config groups for compressed-tensors yet")
|
||||
weight_config = tuple(groups.values())[0]["weights"]
|
||||
|
||||
if quant_format == "float-quantized" or quant_format == "int-quantized" or quant_format == "naive-quantized":
|
||||
elif quant_format == "float-quantized" or quant_format == "int-quantized" or quant_format == "naive-quantized":
|
||||
weight_config = tuple(groups.values())[0]["weights"]
|
||||
block_size = weight_config.get("block_structure", None)
|
||||
strategy = weight_config.get("strategy")
|
||||
assert strategy == "channel" or strategy == "block"
|
||||
@@ -529,6 +544,7 @@ class ModelBase:
|
||||
if self._fp8_as_q8 and is_fp8:
|
||||
self._fp8_dequantized.add(weight_name)
|
||||
elif quant_format == "pack-quantized":
|
||||
weight_config = tuple(groups.values())[0]["weights"]
|
||||
assert weight_config.get("strategy") == "group"
|
||||
assert weight_config.get("type", "int") == "int"
|
||||
num_bits = weight_config.get("num_bits")
|
||||
@@ -550,32 +566,10 @@ class ModelBase:
|
||||
tensors_to_remove += [base_name + n for n in ("_packed", "_shape", "_scale")]
|
||||
if (base_name + "_zero_point") in self.model_tensors:
|
||||
tensors_to_remove.append(base_name + "_zero_point")
|
||||
elif nvfp4_compressed_tensors:
|
||||
# Don't error from compressed-tensors, we'll handle them in _generate_nvfp4_tensors
|
||||
pass
|
||||
else:
|
||||
raise NotImplementedError(f"Quant format {quant_format!r} for method {quant_method!r} is not yet supported")
|
||||
elif quant_method == "modelopt":
|
||||
# Mixed-precision ModelOpt models: NVFP4 tensors are handled by
|
||||
# _generate_nvfp4_tensors; FP8 tensors have 1D weight_scale and
|
||||
# are dequantized here. k/v scale tensors are unused.
|
||||
for name in self.model_tensors.keys():
|
||||
if name.endswith(".weight_scale"):
|
||||
weight_name = name.removesuffix("_scale")
|
||||
if weight_name not in self.model_tensors:
|
||||
tensors_to_remove.append(name)
|
||||
continue
|
||||
w = self.model_tensors[weight_name]
|
||||
s = self.model_tensors[name]
|
||||
is_fp8_weight = False
|
||||
if self._fp8_as_q8:
|
||||
is_fp8_weight = w().dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
|
||||
self.model_tensors[weight_name] = lambda w=w, s=s: dequant_simple(w(), s(), None)
|
||||
tensors_to_remove.append(name)
|
||||
if is_fp8_weight:
|
||||
self._fp8_dequantized.add(weight_name)
|
||||
if name.endswith((".input_scale", ".k_scale", ".v_scale")):
|
||||
tensors_to_remove.append(name)
|
||||
dequant_fp8()
|
||||
elif quant_method is not None:
|
||||
raise NotImplementedError(f"Quant method is not yet supported: {quant_method!r}")
|
||||
|
||||
@@ -821,6 +815,18 @@ class ModelBase:
|
||||
func=load,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_nvfp4_compressed_tensors(quant_method, quant_format, groups) -> bool:
|
||||
# Some models use per-tensor quant_algo (e.g. "MIXED_PRECISION" with
|
||||
# per-layer NVFP4/FP8) instead of a single global "NVFP4" value.
|
||||
if quant_method != "compressed-tensors":
|
||||
return False
|
||||
if quant_format == "nvfp4-pack-quantized":
|
||||
return True
|
||||
if quant_format != "mixed-precision" or not groups:
|
||||
return False
|
||||
return any(g.get("format") == "nvfp4-pack-quantized" for g in groups.values() if isinstance(g, dict))
|
||||
|
||||
@staticmethod
|
||||
def _nvfp4_pack(weight: Tensor, scale: Tensor) -> tuple[np.ndarray, list[int]]:
|
||||
"""Repack NVFP4 ModelOpt tensors into ggml super-block layout.
|
||||
@@ -878,8 +884,8 @@ class ModelBase:
|
||||
weight = LazyTorchTensor.to_eager(self.model_tensors[name]())
|
||||
scale = LazyTorchTensor.to_eager(self.model_tensors[scale_name]())
|
||||
|
||||
# Skip non-NVFP4 tensors (e.g. FP8 with per-channel 1D scales)
|
||||
if scale.ndim < 2:
|
||||
# Skip non-NVFP4 tensors(e.g. 1D scale, or float8 weight)
|
||||
if scale.ndim < 2 or weight.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
continue
|
||||
|
||||
scale2 = LazyTorchTensor.to_eager(self.model_tensors.get(scale2_name, lambda: torch.tensor(1.0))())
|
||||
@@ -980,14 +986,7 @@ class ModelBase:
|
||||
quant_groups = quant_config.get("config_groups", quant_groups) or {}
|
||||
quant_layers = quant_config.get("quantized_layers", quant_layers) or {}
|
||||
|
||||
# Some models use per-tensor quant_algo (e.g. "MIXED_PRECISION" with
|
||||
# per-layer NVFP4/FP8) instead of a single global "NVFP4" value.
|
||||
nvfp4_compressed_tensors = quant_method == "compressed-tensors" and (
|
||||
quant_format == "nvfp4-pack-quantized"
|
||||
or quant_format == "mixed-precision"
|
||||
and bool(quant_groups)
|
||||
and all(g.get("format") == "nvfp4-pack-quantized" for g in quant_groups.values() if isinstance(g, dict))
|
||||
)
|
||||
nvfp4_compressed_tensors = self._is_nvfp4_compressed_tensors(quant_method, quant_format, quant_groups)
|
||||
|
||||
self._nvfp4_global_algo = quant_algo
|
||||
|
||||
|
||||
@@ -1176,7 +1176,22 @@ static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state(
|
||||
return ret;
|
||||
}
|
||||
|
||||
static bool ggml_backend_meta_is_host_view(const struct ggml_tensor * tensor) {
|
||||
return ggml_is_view(tensor) && ggml_backend_buffer_is_host(tensor->view_src->buffer);
|
||||
}
|
||||
|
||||
static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state(const struct ggml_tensor * tensor, bool assume_sync) {
|
||||
// [TAG_META_HOST_VIEWS]
|
||||
// TODO: technically, this check should not be needed if the backend scheduler correctly prevents assigning
|
||||
// such host-buffer views to the meta backend. figure out how to update the scheduler logic to achieve that
|
||||
// ref: https://github.com/ggml-org/llama.cpp/pull/30217
|
||||
if (!ggml_backend_buffer_is_meta(tensor->buffer)) {
|
||||
GGML_ASSERT(ggml_backend_meta_is_host_view(tensor));
|
||||
|
||||
// the view is not allocated in the meta buffer, it is not split across the sub-devices
|
||||
return { GGML_BACKEND_SPLIT_AXIS_NONE, {0}, {1}, 1 };
|
||||
}
|
||||
|
||||
ggml_backend_meta_buffer_context * buf_ctx = (ggml_backend_meta_buffer_context *) tensor->buffer->context;
|
||||
return ggml_backend_meta_get_split_state(buf_ctx->get_simple_tensor_container(tensor), tensor, assume_sync);
|
||||
}
|
||||
@@ -2026,9 +2041,11 @@ static enum ggml_status ggml_backend_meta_graph_compute(ggml_backend_t backend,
|
||||
|
||||
for (int i = 0; i < cgraph->n_nodes; i++) {
|
||||
ggml_tensor * node = cgraph->nodes[i];
|
||||
if (node->view_src != nullptr && node->view_src->op == GGML_OP_NONE && ggml_backend_buffer_is_host(node->view_src->buffer)) {
|
||||
// FIXME s_copy_main is on the CPU and its view seems to be incorrectly added to the graph nodes.
|
||||
// For regular usage this doesn't matter since it's a noop but trying to call ggml_backend_meta_buffer_simple_tensor results in a crash.
|
||||
if (!ggml_backend_buffer_is_meta(node->buffer)) {
|
||||
// [TAG_META_HOST_VIEWS]
|
||||
GGML_ASSERT(ggml_backend_meta_is_host_view(node));
|
||||
|
||||
// keep the node as is, mapping it to a simple tensor is not possible
|
||||
bcj.nodes[i] = node;
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -2899,6 +2899,79 @@ static int ggml_cuda_try_gdn_cache_fusion(
|
||||
return skip;
|
||||
}
|
||||
|
||||
// match ssm_scan + the strided cpy that scatters its state snapshots into the cache, so the kernel writes them and skips the cpy
|
||||
static int ggml_cuda_try_ssm_scan_cache_fusion(
|
||||
const ggml_cgraph * cgraph, int node_idx, ggml_cuda_ssm_scan_fused_cache & fused_state_cpy) {
|
||||
const ggml_tensor * ssm = cgraph->nodes[node_idx];
|
||||
// the kernel skips the snapshot tail, so the scan output must not be a graph output
|
||||
if (ssm->op != GGML_OP_SSM_SCAN || ssm->type != GGML_TYPE_F32 || (ssm->flags & GGML_TENSOR_FLAG_OUTPUT)) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
const int64_t K = ggml_get_op_params_i32(ssm, 0); // snapshot slot count
|
||||
|
||||
const ggml_tensor * s = ssm->src[0];
|
||||
const ggml_tensor * x = ssm->src[1];
|
||||
const ggml_tensor * A = ssm->src[3];
|
||||
|
||||
const int64_t d_state = s->ne[0];
|
||||
const int64_t D = d_state * s->ne[1] * x->ne[1]; // d_state * head_dim * n_head
|
||||
const int64_t n_tok = x->ne[2];
|
||||
const int64_t n_seqs = x->ne[3];
|
||||
|
||||
// only the mamba-2 kernels (group scan and SSD) write to the cache; mamba-1 still uses the cpy
|
||||
if (A->nb[1] != sizeof(float) || (d_state != 96 && d_state != 128 && d_state != 256)) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
// the scan reads its input rows from the cache (picked by ids), so with more than one seq a seq can read a row that another seq writes in the same launch
|
||||
if (n_seqs != 1) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
const int64_t n_written = std::min<int64_t>(n_tok, K);
|
||||
const size_t tail_off = ggml_row_size(GGML_TYPE_F32, ggml_nelements(x));
|
||||
|
||||
// snapshot cpy is the first real node after the scan (skip views/no-ops)
|
||||
const ggml_tensor * cpy = nullptr;
|
||||
int skip = 0;
|
||||
for (int j = node_idx + 1; j < cgraph->n_nodes && cpy == nullptr; ++j) {
|
||||
const ggml_tensor * n = cgraph->nodes[j];
|
||||
if (ggml_cuda_is_view_or_noop(n)) {
|
||||
continue;
|
||||
}
|
||||
if (n->op != GGML_OP_CPY || (n->flags & GGML_TENSOR_FLAG_OUTPUT)) {
|
||||
return 0;
|
||||
}
|
||||
cpy = n;
|
||||
skip = j - node_idx;
|
||||
}
|
||||
if (cpy == nullptr) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
const ggml_tensor * src = cpy->src[0]; // view of the scan snapshot tail
|
||||
const ggml_tensor * dst = cpy->src[1]; // cache view the kernel writes to
|
||||
|
||||
// src must be this scan's snapshot tail (contiguous, at the tail offset)
|
||||
if (src->op != GGML_OP_VIEW || src->view_src != ssm || src->view_offs != tail_off ||
|
||||
!ggml_is_contiguous(src)) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
// dst is the [D, n_seqs, n_written] cache view; require nb[1] == D, the per-seq stride the kernel takes from src0->nb[3]
|
||||
const std::array<int64_t, GGML_MAX_DIMS> expected_ne = { D, n_seqs, n_written, 1 };
|
||||
if (dst->op != GGML_OP_VIEW || dst->type != GGML_TYPE_F32 || dst->data == nullptr ||
|
||||
!std::equal(expected_ne.begin(), expected_ne.end(), dst->ne) ||
|
||||
dst->nb[0] != ggml_type_size(GGML_TYPE_F32) || dst->nb[1] != (size_t) ggml_row_size(GGML_TYPE_F32, D)) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
fused_state_cpy.data = (float *) dst->data; // rollback slot 0 (newest)
|
||||
fused_state_cpy.slot_stride = K > 1 ? (int64_t) (dst->nb[2] / sizeof(float)) : 0;
|
||||
return skip;
|
||||
}
|
||||
|
||||
static bool ggml_cuda_topk_moe_fusion(const struct ggml_cgraph * cgraph, int node_idx, ggml_cuda_topk_moe_args & args) {
|
||||
args.sigmoid = false;
|
||||
args.sqrt_softplus = false;
|
||||
@@ -3585,6 +3658,20 @@ static int ggml_cuda_try_fuse(ggml_backend_cuda_context * cuda_ctx, ggml_cgraph
|
||||
}
|
||||
}
|
||||
|
||||
// ssm_scan -> cpy: scatter recurrent-state snapshots into the cache
|
||||
if (node->op == GGML_OP_SSM_SCAN) {
|
||||
ggml_cuda_ssm_scan_fused_cache fused_state_cpy;
|
||||
const int nodes_to_skip = ggml_cuda_try_ssm_scan_cache_fusion(cgraph, i, fused_state_cpy);
|
||||
if (nodes_to_skip > 0) {
|
||||
#ifdef GGML_CUDA_DEBUG
|
||||
GGML_LOG_INFO("%s: fused ssm_scan snapshot copies for %s (skipped %d nodes)\n",
|
||||
__func__, node->name, nodes_to_skip);
|
||||
#endif
|
||||
ggml_cuda_op_ssm_scan_fused_cache(*cuda_ctx, node, fused_state_cpy);
|
||||
return nodes_to_skip;
|
||||
}
|
||||
}
|
||||
|
||||
//topk-moe
|
||||
if (cgraph->nodes[i]->op == GGML_OP_UNARY || cgraph->nodes[i]->op == GGML_OP_SOFT_MAX ||
|
||||
cgraph->nodes[i]->op == GGML_OP_ARGSORT) {
|
||||
|
||||
@@ -149,7 +149,7 @@ __global__ void __launch_bounds__(d_state, 1)
|
||||
const int src0_nb2, const int src0_nb3, const int src1_nb2, const int src1_nb3,
|
||||
const int src2_nb1, const int src2_nb2, const int src3_nb1,
|
||||
const int src4_nb2, const int src4_nb3, const int src5_nb2, const int src5_nb3,
|
||||
const int64_t s_off, const int64_t n_head, const int64_t d_head, const int64_t n_group, const int64_t n_tok, const int64_t K) {
|
||||
char * s_base, const int64_t s_slot_bytes, const int64_t n_head, const int64_t d_head, const int64_t n_group, const int64_t n_tok, const int64_t K) {
|
||||
const float * GGML_CUDA_RESTRICT src0 = src0_ptr;
|
||||
const float * GGML_CUDA_RESTRICT src1 = src1_ptr;
|
||||
const float * GGML_CUDA_RESTRICT src2 = src2_ptr;
|
||||
@@ -184,7 +184,7 @@ __global__ void __launch_bounds__(d_state, 1)
|
||||
const float * B_warp = (const float *) ((const char *) src4 + (seq_idx * src4_nb3) + (group_off));
|
||||
const float * C_warp = (const float *) ((const char *) src5 + (seq_idx * src5_nb3) + (group_off));
|
||||
float * y_warp = dst + (seq_idx * n_tok * n_head * d_head) + warp_idx;
|
||||
float * s_warp = (float *) ((char *) dst + s_off + seq_idx * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
|
||||
float * s_warp = (float *) (s_base + seq_idx * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
|
||||
|
||||
// strides across n_seq_tokens
|
||||
const int stride_x = src1_nb2 / sizeof(float);
|
||||
@@ -227,7 +227,7 @@ __global__ void __launch_bounds__(d_state, 1)
|
||||
// Slot 0 is the final state written below; slots 1..K-1 are rollback snapshots.
|
||||
const int64_t slot = n_tok - 1 - i;
|
||||
if (K > 1 && slot > 0 && slot < K) {
|
||||
float * s_snapshot_warp = (float *) ((char *) dst + s_off + (slot * gridDim.y + seq_idx) * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
|
||||
float * s_snapshot_warp = (float *) ((char *) s_warp + slot * s_slot_bytes);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < c_factor; j++) {
|
||||
s_snapshot_warp[WARP_SIZE * j + lane] = state[j];
|
||||
@@ -248,7 +248,11 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
const int src2_nb2, const int src3_nb1, const int src4_nb2, const int src4_nb3, const int src5_nb2,
|
||||
const int src5_nb3, const int64_t s_off, const int64_t d_state, const int64_t head_dim,
|
||||
const int64_t n_head, const int64_t n_group, const int64_t n_tok, const int64_t n_seq,
|
||||
const int64_t K, cudaStream_t stream) {
|
||||
const int64_t K, const ggml_cuda_ssm_scan_fused_cache * cache, cudaStream_t stream) {
|
||||
// when fused, the states go straight into the recurrent cache and the dst tail is left alone
|
||||
char * const s_base = cache ? (char *) cache->data : (char *) dst + s_off;
|
||||
const int64_t s_slot_bytes = cache ? cache->slot_stride * (int64_t) sizeof(float) : n_seq * (int64_t) src0_nb3;
|
||||
|
||||
// NOTE: if you change conditions here, be sure to update the corresponding supports_op condition!
|
||||
if (src3_nb1 == sizeof(float)) {
|
||||
// Mamba-2
|
||||
@@ -261,7 +265,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
ggml_cuda_kernel_launch(ssm_scan_f32_group<96/WARP_SIZE, 96>, launch_params,
|
||||
src0, src1, src2, src3, src4, src5, src6, dst,
|
||||
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
|
||||
} else if (d_state == 128) {
|
||||
constexpr int threads = 128;
|
||||
constexpr int num_warps = threads/WARP_SIZE;
|
||||
@@ -271,7 +275,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
ggml_cuda_kernel_launch(ssm_scan_f32_group<128/WARP_SIZE, 128>, launch_params,
|
||||
src0, src1, src2, src3, src4, src5, src6, dst,
|
||||
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
|
||||
} else if (d_state == 256) { // Falcon-H1
|
||||
constexpr int threads = 256;
|
||||
constexpr int num_warps = threads/WARP_SIZE;
|
||||
@@ -281,7 +285,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
ggml_cuda_kernel_launch(ssm_scan_f32_group<256/WARP_SIZE, 256>, launch_params,
|
||||
src0, src1, src2, src3, src4, src5, src6, dst,
|
||||
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
|
||||
} else {
|
||||
GGML_ABORT("doesn't support d_state!=(96, 128 or 256).");
|
||||
}
|
||||
@@ -570,12 +574,13 @@ __global__ void ssm_ssd_scale_state_kernel(
|
||||
}
|
||||
|
||||
// Copy initial state from src0[ids[s]] into s_cur for each sequence.
|
||||
// src0 and s_cur can alias when the state is written straight into the cache.
|
||||
// Grid: (ceil(d_state * head_dim * n_head / BLOCK), n_seqs)
|
||||
template <int BLOCK_SIZE>
|
||||
__global__ void ssm_ssd_init_state_kernel(
|
||||
const float * __restrict__ src0, // {d_state, head_dim, n_head, n_rs}
|
||||
const float * src0, // {d_state, head_dim, n_head, n_rs}
|
||||
const int32_t * __restrict__ ids, // {n_seqs}
|
||||
float * __restrict__ s_cur, // {d_state, head_dim, n_head, n_seqs}
|
||||
float * s_cur, // {d_state, head_dim, n_head, n_seqs}
|
||||
const int state_size, // d_state * head_dim * n_head
|
||||
const int64_t s0_stride_seq) { // elements between state rows
|
||||
const int s = blockIdx.y;
|
||||
@@ -599,7 +604,8 @@ static void ssm_scan_ssd_f32_cuda(
|
||||
const int A_stride, // A (src3) stride between heads
|
||||
const int B_stride_tok, const int B_stride_seq, // B (src4) strides
|
||||
const int C_stride_tok, const int C_stride_seq, // C (src5) strides
|
||||
const int64_t s_off, const int64_t d_state, const int64_t head_dim,
|
||||
float * s_cur, // state: dst state tail, or the cache when fused
|
||||
const int64_t d_state, const int64_t head_dim,
|
||||
const int64_t n_head, const int64_t n_group, const int64_t n_tok, const int64_t n_seq) {
|
||||
|
||||
cudaStream_t stream = ctx.stream();
|
||||
@@ -625,7 +631,6 @@ static void ssm_scan_ssd_f32_cuda(
|
||||
matmul_t * X_dt = X_dt_buf.get();
|
||||
matmul_t * B_weighted = B_w_buf.get();
|
||||
float * C_scaled = C_s_buf.get();
|
||||
float * s_cur = (float *)((char *)dst_d + s_off); // write state directly to dst
|
||||
|
||||
// Step 1: softplus(dt) and parallel prefix sum over full sequence
|
||||
{
|
||||
@@ -780,7 +785,8 @@ static void ssm_scan_ssd_f32_cuda(
|
||||
}
|
||||
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
|
||||
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
static void ggml_cuda_op_ssm_scan_impl(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
|
||||
const ggml_cuda_ssm_scan_fused_cache * cache) {
|
||||
const struct ggml_tensor * src0 = dst->src[0]; // s
|
||||
const struct ggml_tensor * src1 = dst->src[1]; // x
|
||||
const struct ggml_tensor * src2 = dst->src[2]; // dt
|
||||
@@ -864,12 +870,21 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
(int)(src3->nb[1] / sizeof(float)),
|
||||
(int)(src4->nb[2] / sizeof(float)), (int)(src4->nb[3] / sizeof(float)),
|
||||
(int)(src5->nb[2] / sizeof(float)), (int)(src5->nb[3] / sizeof(float)),
|
||||
s_off, nc, nr, nh, ng, n_t, n_s);
|
||||
cache ? cache->data : (float *) ((char *) dst_d + s_off), nc, nr, nh, ng, n_t, n_s);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
ssm_scan_f32_cuda(src0_d, src1_d, src2_d, src3_d, src4_d, src5_d, src6_d, dst_d,
|
||||
src0->nb[2], src0->nb[3], src1->nb[2], src1->nb[3], src2->nb[1], src2->nb[2],
|
||||
src3->nb[1], src4->nb[2], src4->nb[3], src5->nb[2], src5->nb[3],
|
||||
s_off, nc, nr, nh, ng, n_t, n_s, K, stream);
|
||||
s_off, nc, nr, nh, ng, n_t, n_s, K, cache, stream);
|
||||
}
|
||||
|
||||
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
ggml_cuda_op_ssm_scan_impl(ctx, dst, nullptr);
|
||||
}
|
||||
|
||||
void ggml_cuda_op_ssm_scan_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
|
||||
ggml_cuda_ssm_scan_fused_cache cache) {
|
||||
ggml_cuda_op_ssm_scan_impl(ctx, dst, &cache);
|
||||
}
|
||||
|
||||
@@ -1,3 +1,13 @@
|
||||
#include "common.cuh"
|
||||
|
||||
// fused-kernel recurrent-state output; strides in elements (per-seq stride is always the state row size, set in-kernel)
|
||||
struct ggml_cuda_ssm_scan_fused_cache {
|
||||
float * data; // rollback slot 0
|
||||
int64_t slot_stride; // between rollback slots
|
||||
};
|
||||
|
||||
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
|
||||
|
||||
// same op, but writes the state snapshot(s) into the cache instead of dst (see ggml_cuda_try_ssm_scan_cache_fusion)
|
||||
void ggml_cuda_op_ssm_scan_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
|
||||
ggml_cuda_ssm_scan_fused_cache cache);
|
||||
|
||||
@@ -107,7 +107,7 @@ static __device__ __forceinline__ float op_ceil(float x) {
|
||||
}
|
||||
|
||||
static __device__ __forceinline__ float op_round(float x) {
|
||||
return round(x);
|
||||
return roundf(x);
|
||||
}
|
||||
|
||||
static __device__ __forceinline__ float op_trunc(float x) {
|
||||
|
||||
@@ -280,7 +280,8 @@ static ADRENO_GPU_GEN get_adreno_gpu_gen(const char *device_name) {
|
||||
strstr(device_name, "613") || strstr(device_name, "615") ||
|
||||
strstr(device_name, "616") || strstr(device_name, "618") ||
|
||||
strstr(device_name, "619") || strstr(device_name, "620") ||
|
||||
strstr(device_name, "630") || strstr(device_name, "640") ||
|
||||
strstr(device_name, "623") || strstr(device_name, "630") ||
|
||||
strstr(device_name, "640") ||
|
||||
strstr(device_name, "642") || strstr(device_name, "643") ||
|
||||
strstr(device_name, "644") || strstr(device_name, "650") ||
|
||||
strstr(device_name, "660") || strstr(device_name, "663") ||
|
||||
@@ -863,7 +864,8 @@ struct ggml_backend_opencl_context {
|
||||
cl_kernel kernel_set_rows_q4_0_soa_i64, kernel_set_rows_q4_0_soa_i32;
|
||||
cl_kernel kernel_rope_norm_f32, kernel_rope_norm_f16, kernel_rope_neox_f32, kernel_rope_neox_f16;
|
||||
cl_kernel kernel_rope_multi_f32, kernel_rope_multi_f16, kernel_rope_vision_f32, kernel_rope_vision_f16;
|
||||
cl_kernel kernel_cpy_f16_f16, kernel_cpy_f16_f32, kernel_cpy_f32_f16, kernel_cpy_f32_f32, kernel_cpy_f32_f32_pack, kernel_cpy_i32_i32;
|
||||
cl_kernel kernel_cpy_f16_f16, kernel_cpy_f16_f32, kernel_cpy_f32_f16, kernel_cpy_f32_f32, kernel_cpy_i32_i32;
|
||||
cl_kernel kernel_cpy_f32_f32_pack = nullptr;
|
||||
cl_kernel kernel_cpy_f32_f32_flat = nullptr;
|
||||
cl_kernel kernel_mul_mat_f32_f32;
|
||||
cl_kernel kernel_mul_mat_f16_f16;
|
||||
@@ -1604,14 +1606,18 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
#else
|
||||
const std::string kernel_src = read_file("cpy.cl");
|
||||
#endif
|
||||
cl_program prog =
|
||||
build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts);
|
||||
const bool no_cpy_pack = backend_ctx->adreno_gen == ADRENO_GPU_GEN::A6X;
|
||||
|
||||
cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(),
|
||||
no_cpy_pack ? compile_opts + " -DGGML_CL_NO_CPY_PACK" : compile_opts);
|
||||
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f16_f16 = clCreateKernel(prog, "kernel_cpy_f16_f16", &err), err));
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f16_f32 = clCreateKernel(prog, "kernel_cpy_f16_f32", &err), err));
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f32_f16 = clCreateKernel(prog, "kernel_cpy_f32_f16", &err), err));
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f32_f32 = clCreateKernel(prog, "kernel_cpy_f32_f32", &err), err));
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f32_f32_pack = clCreateKernel(prog, "kernel_cpy_f32_f32_pack", &err), err));
|
||||
if (!no_cpy_pack) {
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f32_f32_pack = clCreateKernel(prog, "kernel_cpy_f32_f32_pack", &err), err));
|
||||
}
|
||||
{ // optional: without it ggml_cl_cpy keeps the row-mapped kernel
|
||||
cl_int err_flat = CL_SUCCESS;
|
||||
cl_kernel k = clCreateKernel(prog, "kernel_cpy_f32_f32_flat", &err_flat);
|
||||
@@ -3770,6 +3776,9 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
if (backend_ctx->has_vector_subgroup_broadcast) {
|
||||
CL_gemv_compile_opts += " -DVECTOR_SUB_GROUP_BROADCAST ";
|
||||
}
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::A6X) {
|
||||
CL_gemv_compile_opts += " -DGGML_CL_A6X_CONSTFOLD_FIX";
|
||||
}
|
||||
|
||||
#ifdef GGML_OPENCL_EMBED_KERNELS
|
||||
const std::string kernel_src_CL_gemv_general {
|
||||
@@ -4306,6 +4315,9 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
if (backend_ctx->has_vector_subgroup_broadcast) {
|
||||
CL_gemv_compile_opts += " -DVECTOR_SUB_GROUP_BROADCAST ";
|
||||
}
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::A6X) {
|
||||
CL_gemv_compile_opts += " -DGGML_CL_A6X_CONSTFOLD_FIX";
|
||||
}
|
||||
// Opt-in: dequant-once-per-block mc3 verify GEMV (factors q4_K dequant
|
||||
// out of the 3-column loop; byte-identical, lower spill). A/B vs the
|
||||
// shipped inline mc3 in the same binary.
|
||||
@@ -28335,7 +28347,8 @@ static void ggml_cl_cpy(ggml_backend_t backend, const ggml_tensor * src0, const
|
||||
kernel = backend_ctx->kernel_cpy_f32_f16;
|
||||
break;
|
||||
case GGML_TYPE_F32:
|
||||
kernel = ne00 < 32 ? backend_ctx->kernel_cpy_f32_f32_pack
|
||||
kernel = (ne00 < 32 && backend_ctx->kernel_cpy_f32_f32_pack)
|
||||
? backend_ctx->kernel_cpy_f32_f32_pack
|
||||
: backend_ctx->kernel_cpy_f32_f32;
|
||||
break;
|
||||
default:
|
||||
|
||||
@@ -183,6 +183,7 @@ kernel void kernel_cpy_f32_f32(
|
||||
}
|
||||
}
|
||||
|
||||
#ifndef GGML_CL_NO_CPY_PACK
|
||||
kernel void kernel_cpy_f32_f32_pack(
|
||||
global float * src0,
|
||||
ulong offset0,
|
||||
@@ -241,6 +242,7 @@ kernel void kernel_cpy_f32_f32_pack(
|
||||
dst_data[i00] = src[0];
|
||||
}
|
||||
}
|
||||
#endif // GGML_CL_NO_CPY_PACK
|
||||
|
||||
kernel void kernel_cpy_i32_i32(
|
||||
global int * src0,
|
||||
|
||||
@@ -7,6 +7,14 @@
|
||||
#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half")))
|
||||
#endif
|
||||
|
||||
// A6X compiler incorrectly constant-folds get_local_size() results;
|
||||
// force runtime materialization via a no-op ALU round-trip.
|
||||
#ifdef GGML_CL_A6X_CONSTFOLD_FIX
|
||||
#define MATERIALIZE_WG(x) do { (x) *= 2u; if ((x) > 1u) (x) /= 2u; } while(0)
|
||||
#else
|
||||
#define MATERIALIZE_WG(x)
|
||||
#endif
|
||||
|
||||
// assume
|
||||
#define QK4_0 32
|
||||
#define N_SIMDGROUP 4
|
||||
@@ -327,6 +335,7 @@ __kernel void kernel_gemv_noshuffle_q4_0_f32_mc3(
|
||||
uint BLOCK_STRIDE_A = N_SIMDGROUP * M; // = 4 * M (N_SIMDGROUP is the #define 4)
|
||||
uint COL_STRIDE = K / 4; // float4 pixels per activation column
|
||||
uint nsg = get_local_size(1); // runtime K-split (4 default, 8 small-M)
|
||||
MATERIALIZE_WG(nsg);
|
||||
|
||||
__private uint4 regA_hi, regA_lo;
|
||||
__private half2 regS;
|
||||
|
||||
@@ -11,6 +11,14 @@
|
||||
#define NSUBGROUPS 4
|
||||
#define SUBGROUP_SIZE 64
|
||||
|
||||
// A6X compiler incorrectly constant-folds get_local_size() results;
|
||||
// force runtime materialization via a no-op ALU round-trip.
|
||||
#ifdef GGML_CL_A6X_CONSTFOLD_FIX
|
||||
#define MATERIALIZE_WG(x) do { (x) *= 2u; if ((x) > 1u) (x) /= 2u; } while(0)
|
||||
#else
|
||||
#define MATERIALIZE_WG(x)
|
||||
#endif
|
||||
|
||||
// scales are transposed: consecutive codes of a row are `stride` apart
|
||||
inline void get_scale_min_k4(
|
||||
int j,
|
||||
@@ -233,6 +241,7 @@ kernel void kernel_gemv_noshuffle_q4_k_f32(
|
||||
// K-split (more waves/SP -> latency hiding) while large-M keeps 4. The
|
||||
// physical weight layout stride below is INDEPENDENT of this (see BLOCK_STRIDE_A).
|
||||
uint nsg = get_local_size(1);
|
||||
MATERIALIZE_WG(nsg);
|
||||
|
||||
uint K = ne00;
|
||||
uint M = ne01;
|
||||
@@ -400,6 +409,7 @@ kernel void kernel_gemv_noshuffle_q4_k_f32_glu(
|
||||
uint gid = get_global_id(0);
|
||||
ushort slid = get_sub_group_local_id();
|
||||
uint nsg = get_local_size(1);
|
||||
MATERIALIZE_WG(nsg);
|
||||
|
||||
uint K = ne00;
|
||||
uint M = ne01;
|
||||
@@ -514,6 +524,7 @@ kernel void kernel_gemv_noshuffle_q4_k_f32_splitk(
|
||||
uint gid = get_global_id(0);
|
||||
ushort slid = get_sub_group_local_id();
|
||||
uint nsg = get_local_size(1);
|
||||
MATERIALIZE_WG(nsg);
|
||||
uint ksplit = get_num_groups(1);
|
||||
uint kslice = get_group_id(1);
|
||||
|
||||
|
||||
@@ -9455,6 +9455,8 @@ static void ggml_vk_op_f32(ggml_backend_vk_context * ctx, vk_context& subctx, co
|
||||
elements = { (uint32_t)CEIL_DIV(ne00, 128), 1, 1 };
|
||||
} else {
|
||||
elements = { (uint32_t)ne01, (uint32_t)ne02, (uint32_t)ne03 };
|
||||
elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
|
||||
elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
|
||||
}
|
||||
break;
|
||||
|
||||
@@ -10777,7 +10779,11 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
||||
ggml_vk_tensor_subbuffer(ctx, src0, true),
|
||||
ggml_vk_tensor_subbuffer(ctx, set_rows, true),
|
||||
ggml_vk_tensor_subbuffer(ctx, indices),
|
||||
}, pc, { (uint32_t)src0->ne[1], (uint32_t)src0->ne[2], (uint32_t)src0->ne[3] });
|
||||
}, pc, {
|
||||
(uint32_t)src0->ne[1],
|
||||
std::min((uint32_t)src0->ne[2], ctx->device->properties.limits.maxComputeWorkGroupCount[1]),
|
||||
std::min((uint32_t)src0->ne[3], ctx->device->properties.limits.maxComputeWorkGroupCount[2]),
|
||||
});
|
||||
ggml_vk_rms_norm_finish(ctx, src0);
|
||||
return;
|
||||
}
|
||||
@@ -10824,7 +10830,11 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
||||
ggml_vk_tensor_subbuffer(ctx, dst, true),
|
||||
ggml_vk_tensor_subbuffer(ctx, residual),
|
||||
ggml_vk_tensor_subbuffer(ctx, post_scale),
|
||||
}, pc, { (uint32_t)src0->ne[1], (uint32_t)src0->ne[2], (uint32_t)src0->ne[3] });
|
||||
}, pc, {
|
||||
(uint32_t)src0->ne[1],
|
||||
std::min((uint32_t)src0->ne[2], ctx->device->properties.limits.maxComputeWorkGroupCount[1]),
|
||||
std::min((uint32_t)src0->ne[3], ctx->device->properties.limits.maxComputeWorkGroupCount[2]),
|
||||
});
|
||||
}
|
||||
ggml_vk_rms_norm_finish(ctx, src0);
|
||||
return;
|
||||
@@ -10911,6 +10921,8 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
||||
|
||||
std::array<uint32_t, 3> elements;
|
||||
elements = { (uint32_t)rms->src[0]->ne[1], (uint32_t)rms->src[0]->ne[2], (uint32_t)rms->src[0]->ne[3] };
|
||||
elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
|
||||
elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
|
||||
|
||||
static_assert(max_tensors == 7);
|
||||
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
|
||||
|
||||
@@ -53,101 +53,107 @@ shared FLOAT_TYPE sumsh[BLOCK_SIZE];
|
||||
void rms_norm(uint num_iters) {
|
||||
const uint ncols = p.ne00;
|
||||
const uint nrows = gl_NumWorkGroups.x;
|
||||
const uint nchannels = gl_NumWorkGroups.y;
|
||||
const uint nchannels = p.ne02;
|
||||
const uint nsamples = p.ne03;
|
||||
|
||||
const uint row = gl_WorkGroupID.x;
|
||||
const uint channel = gl_WorkGroupID.y;
|
||||
const uint samp = gl_WorkGroupID.z;
|
||||
const uint tid = gl_LocalInvocationID.x;
|
||||
|
||||
const uint stride_row = p.nb01;
|
||||
const uint stride_channel = p.nb02;
|
||||
const uint stride_sample = p.nb03;
|
||||
|
||||
uint32_t a_offset = samp*stride_sample + channel*stride_channel + row*stride_row + get_aoffset();
|
||||
uint32_t b_offset = src1_idx(0, row, channel, samp) + get_boffset();
|
||||
// grid.y/z are clamped to the device workgroup limit, iterate over the excess channels/samples
|
||||
for (uint samp = gl_WorkGroupID.z; samp < nsamples; samp += gl_NumWorkGroups.z) {
|
||||
for (uint channel = gl_WorkGroupID.y; channel < nchannels; channel += gl_NumWorkGroups.y) {
|
||||
barrier();
|
||||
|
||||
uint32_t a_offset = samp*stride_sample + channel*stride_channel + row*stride_row + get_aoffset();
|
||||
uint32_t b_offset = src1_idx(0, row, channel, samp) + get_boffset();
|
||||
#if RMS_NORM_ROPE_FUSION
|
||||
// Per-row offset in shared memory
|
||||
uint32_t d_offset = 0;
|
||||
// Per-row offset in shared memory
|
||||
uint32_t d_offset = 0;
|
||||
#elif RMS_NORM_SET_ROWS_FUSION
|
||||
uint32_t d_offset = data_i[channel].x*p.nb21 + row*ncols + get_doffset();
|
||||
uint32_t d_offset = data_i[channel].x*p.nb21 + row*ncols + get_doffset();
|
||||
#else
|
||||
uint32_t d_offset = ((samp*nchannels + channel)*nrows + row)*ncols + get_doffset();
|
||||
uint32_t d_offset = ((samp*nchannels + channel)*nrows + row)*ncols + get_doffset();
|
||||
#endif
|
||||
FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp
|
||||
FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp
|
||||
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
FLOAT_TYPE xi = FLOAT_TYPE(0);
|
||||
if (col < ncols) {
|
||||
xi = FLOAT_TYPE(data_a[a_offset + col]);
|
||||
}
|
||||
sum += xi * xi;
|
||||
}
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
FLOAT_TYPE xi = FLOAT_TYPE(0);
|
||||
if (col < ncols) {
|
||||
xi = FLOAT_TYPE(data_a[a_offset + col]);
|
||||
}
|
||||
sum += xi * xi;
|
||||
}
|
||||
|
||||
sumsh[tid] = sum;
|
||||
// sum up partial sums and write back result
|
||||
barrier();
|
||||
[[unroll]] for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
|
||||
if (tid < s) {
|
||||
sum += sumsh[tid + s];
|
||||
sumsh[tid] = sum;
|
||||
}
|
||||
barrier();
|
||||
}
|
||||
sum = sumsh[0];
|
||||
// sum up partial sums and write back result
|
||||
barrier();
|
||||
[[unroll]] for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
|
||||
if (tid < s) {
|
||||
sum += sumsh[tid + s];
|
||||
sumsh[tid] = sum;
|
||||
}
|
||||
barrier();
|
||||
}
|
||||
sum = sumsh[0];
|
||||
|
||||
const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
|
||||
const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
|
||||
const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
|
||||
const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
|
||||
|
||||
if (do_multiply) {
|
||||
if (ncols > p.ne10) {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
|
||||
if (do_multiply) {
|
||||
if (ncols > p.ne10) {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
|
||||
#if RMS_NORM_ADD_FUSION
|
||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||
if (do_post_multiply) {
|
||||
value *= FLOAT_TYPE(data_e[0]);
|
||||
}
|
||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||
if (do_post_multiply) {
|
||||
value *= FLOAT_TYPE(data_e[0]);
|
||||
}
|
||||
#endif
|
||||
data_d[d_offset + col] = D_TYPE(value);
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
|
||||
data_d[d_offset + col] = D_TYPE(value);
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
|
||||
#if RMS_NORM_ADD_FUSION
|
||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||
if (do_post_multiply) {
|
||||
value *= FLOAT_TYPE(data_e[0]);
|
||||
}
|
||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||
if (do_post_multiply) {
|
||||
value *= FLOAT_TYPE(data_e[0]);
|
||||
}
|
||||
#endif
|
||||
data_d[d_offset + col] = D_TYPE(value);
|
||||
data_d[d_offset + col] = D_TYPE(value);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
|
||||
}
|
||||
}
|
||||
#if RMS_NORM_ROPE_FUSION
|
||||
barrier();
|
||||
rope_params rp = p.rope;
|
||||
for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) {
|
||||
if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) {
|
||||
rope_neox(t, row, channel, samp, rp);
|
||||
} else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) {
|
||||
rope_norm(t, row, channel, samp, rp);
|
||||
barrier();
|
||||
rope_params rp = p.rope;
|
||||
for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) {
|
||||
if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) {
|
||||
rope_neox(t, row, channel, samp, rp);
|
||||
} else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) {
|
||||
rope_norm(t, row, channel, samp, rp);
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
void main() {
|
||||
|
||||
@@ -847,8 +847,11 @@ static webgpu_encoded_op ggml_webgpu_set(webgpu_context & ctx,
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_index, src1));
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_index + 1, dst));
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_pad(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -897,8 +900,11 @@ static webgpu_encoded_op ggml_webgpu_pad(webgpu_context & ctx, ggml_tensor * src
|
||||
ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst),
|
||||
};
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_solve_tri(webgpu_context & ctx,
|
||||
@@ -1495,8 +1501,11 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_set_rows(webgpu_context & ct
|
||||
} else {
|
||||
threads = src->ne[0] * src->ne[1] * src->ne[2] * src->ne[3];
|
||||
}
|
||||
uint32_t wg_x = CEIL_DIV(threads, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, 1);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(threads, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
// Workgroup size is a common constant
|
||||
@@ -1557,9 +1566,11 @@ static webgpu_encoded_op ggml_webgpu_get_rows(webgpu_context & ctx,
|
||||
uint32_t blocks_per_row = (uint32_t) (dst->ne[0] / (decisions->vectorized ? 4 : 1));
|
||||
uint32_t total_rows = (uint32_t) (dst->ne[1] * dst->ne[2] * dst->ne[3]);
|
||||
uint32_t total_threads = float_parallel ? blocks_per_row * total_rows : total_rows;
|
||||
uint32_t wg_x = CEIL_DIV(total_threads, decisions->wg_size);
|
||||
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(total_threads, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static void ggml_webgpu_quantize_q8_dispatch(webgpu_context & ctx,
|
||||
@@ -2347,7 +2358,8 @@ static webgpu_encoded_op ggml_webgpu_unary_op(webgpu_context & ctx, ggml_tensor
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
|
||||
}
|
||||
|
||||
uint32_t wg_x, wg_y;
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
@@ -2425,7 +2437,8 @@ static webgpu_encoded_op ggml_webgpu_binary_op(webgpu_context & ctx,
|
||||
}
|
||||
}
|
||||
|
||||
uint32_t wg_x, wg_y;
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
@@ -2556,8 +2569,11 @@ static webgpu_encoded_op ggml_webgpu_concat(webgpu_context & ctx,
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst));
|
||||
}
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_repeat(webgpu_context & ctx, ggml_tensor * src0, ggml_tensor * dst) {
|
||||
@@ -2591,8 +2607,12 @@ static webgpu_encoded_op ggml_webgpu_repeat(webgpu_context & ctx, ggml_tensor *
|
||||
|
||||
webgpu_pipeline pipeline = ctx->shader_lib->get_repeat_pipeline(shader_lib_ctx);
|
||||
auto * decisions = static_cast<ggml_webgpu_generic_shader_decisions *>(pipeline.context.get());
|
||||
uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context & ctx,
|
||||
@@ -2637,6 +2657,7 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context
|
||||
(uint32_t) dst->ne[0],
|
||||
(uint32_t) dst->ne[1],
|
||||
(uint32_t) dst->ne[2],
|
||||
(uint32_t) dst->ne[3],
|
||||
ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(rn_dst, 0)) // epsilon, treated as f32 in the shader
|
||||
};
|
||||
|
||||
@@ -2676,7 +2697,11 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst));
|
||||
}
|
||||
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(dst));
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
compute_2d_workgroups(ggml_nrows(dst), ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x,
|
||||
wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -2692,6 +2717,7 @@ static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor
|
||||
(uint32_t) src->ne[0],
|
||||
(uint32_t) src->ne[1],
|
||||
(uint32_t) src->ne[2],
|
||||
(uint32_t) src->ne[3],
|
||||
ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(dst, 0)) // epsilon, treated as f32 in the shader
|
||||
};
|
||||
|
||||
@@ -2707,7 +2733,12 @@ static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor
|
||||
if (!decisions->inplace) {
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
|
||||
}
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(src));
|
||||
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
compute_2d_workgroups(ggml_nrows(src), ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x,
|
||||
wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_rope(webgpu_context & ctx,
|
||||
@@ -2796,8 +2827,11 @@ static webgpu_encoded_op ggml_webgpu_rope(webgpu_context & ctx,
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, dst_binding, dst));
|
||||
}
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_glu(webgpu_context & ctx,
|
||||
@@ -2870,8 +2904,11 @@ static webgpu_encoded_op ggml_webgpu_glu(webgpu_context & ctx,
|
||||
}
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, dst_binding, dst));
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_scale(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -2909,7 +2946,8 @@ static webgpu_encoded_op ggml_webgpu_scale(webgpu_context & ctx, ggml_tensor * s
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
|
||||
}
|
||||
|
||||
uint32_t wg_x, wg_y;
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
@@ -2982,7 +3020,11 @@ static webgpu_encoded_op ggml_webgpu_soft_max(webgpu_context & ctx,
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num, dst));
|
||||
}
|
||||
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(dst));
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = ggml_nrows(dst);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_argmax(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -3165,8 +3207,11 @@ static webgpu_encoded_op ggml_webgpu_cumsum(webgpu_context & ctx, ggml_tensor *
|
||||
shader_lib_ctx.max_wg_size = ctx->global_ctx->capabilities.limits.maxComputeInvocationsPerWorkgroup;
|
||||
|
||||
webgpu_pipeline pipeline = ctx->shader_lib->get_cumsum_pipeline(shader_lib_ctx);
|
||||
uint32_t wg_x = ggml_nrows(dst);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = ggml_nrows(dst);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_sum_rows(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -3190,8 +3235,11 @@ static webgpu_encoded_op ggml_webgpu_sum_rows(webgpu_context & ctx, ggml_tensor
|
||||
|
||||
webgpu_pipeline pipeline = ctx->shader_lib->get_sum_rows_pipeline(shader_lib_ctx);
|
||||
|
||||
uint32_t wg_x = total_sum ? 1 : ggml_nrows(dst);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = total_sum ? 1 : ggml_nrows(dst);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static bool ggml_webgpu_can_fuse_rms_norm_mul(const struct ggml_cgraph * cgraph, int node_idx) {
|
||||
|
||||
@@ -53,10 +53,12 @@ var<storage, read_write> dst: array<DataType>;
|
||||
var<uniform> params: Params;
|
||||
#endif
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
if (gid.x < params.ne) {
|
||||
var i = gid.x;
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i < params.ne) {
|
||||
var i = gid_i;
|
||||
let i3 = i / (params.ne2 * params.ne1 * params.ne0);
|
||||
i = i % (params.ne2 * params.ne1 * params.ne0);
|
||||
let i2 = i / (params.ne1 * params.ne0);
|
||||
@@ -72,9 +74,9 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
ni[2] * params.stride_src0_2 +
|
||||
ni[3] * params.stride_src0_3;
|
||||
#ifdef SRC_OVERLAP
|
||||
dst[params.offset_dst + gid.x] = merged_src[params.offset_src0 + src_i];
|
||||
dst[params.offset_dst + gid_i] = merged_src[params.offset_src0 + src_i];
|
||||
#else
|
||||
dst[params.offset_dst + gid.x] = src0[params.offset_src0 + src_i];
|
||||
dst[params.offset_dst + gid_i] = src0[params.offset_src0 + src_i];
|
||||
#endif
|
||||
} else {
|
||||
ni[params.dim] -= params.src0_nedim;
|
||||
@@ -83,9 +85,9 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
ni[2] * params.stride_src1_2 +
|
||||
ni[3] * params.stride_src1_3;
|
||||
#ifdef SRC_OVERLAP
|
||||
dst[params.offset_dst + gid.x] = merged_src[params.offset_src1 + src_i];
|
||||
dst[params.offset_dst + gid_i] = merged_src[params.offset_src1 + src_i];
|
||||
#else
|
||||
dst[params.offset_dst + gid.x] = src1[params.offset_src1 + src_i];
|
||||
dst[params.offset_dst + gid_i] = src1[params.offset_src1 + src_i];
|
||||
#endif
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,8 +17,11 @@ var<workgroup> shared_sum: array<f32, WG_SIZE>;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(workgroup_id) wid: vec3<u32>,
|
||||
@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(local_invocation_id) lid: vec3<u32>) {
|
||||
let row_idx = params.offset_src + wid.x * params.ne0;
|
||||
|
||||
let wid_i = wid.x + wid.y * num_wg.x;
|
||||
let row_idx = params.offset_src + wid_i * params.ne0;
|
||||
let elems = (params.ne0 + WG_SIZE - 1) / WG_SIZE;
|
||||
var local_sum: f32 = 0.0;
|
||||
for (var col = lid.x * elems; col < (lid.x + 1) * elems && col < params.ne0; col ++) {
|
||||
|
||||
@@ -157,12 +157,15 @@ fn b_value(base: u32) -> DataType {
|
||||
#endif
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
if (gid.x >= params.ne) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i >= params.ne) {
|
||||
return;
|
||||
}
|
||||
|
||||
var i = gid.x;
|
||||
var i = gid_i;
|
||||
let i3 = i / (params.ne2 * params.ne1 * params.ne0);
|
||||
i = i % (params.ne2 * params.ne1 * params.ne0);
|
||||
let i2 = i / (params.ne1 * params.ne0);
|
||||
|
||||
@@ -45,12 +45,15 @@ fn wrap_around(idx: i32, n: u32) -> u32 {
|
||||
}
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
if (gid.x >= params.ne) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i >= params.ne) {
|
||||
return;
|
||||
}
|
||||
|
||||
var i = gid.x;
|
||||
var i = gid_i;
|
||||
let dst_plane = params.dst_ne2 * params.dst_ne1 * params.dst_ne0;
|
||||
let i3 = i / dst_plane;
|
||||
i = i % dst_plane;
|
||||
@@ -82,5 +85,5 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
}
|
||||
#endif
|
||||
|
||||
dst[params.offset_dst + gid.x] = value;
|
||||
dst[params.offset_dst + gid_i] = value;
|
||||
}
|
||||
|
||||
@@ -45,9 +45,12 @@ var<storage, read_write> dst: array<DataType>;
|
||||
var<uniform> params: Params;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
if (gid.x < params.ne) {
|
||||
var i = gid.x;
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i < params.ne) {
|
||||
var i = gid_i;
|
||||
let i3 = i / (params.ne2 * params.ne1 * params.ne0);
|
||||
i = i % (params.ne2 * params.ne1 * params.ne0);
|
||||
let i2 = i / (params.ne1 * params.ne0);
|
||||
@@ -65,6 +68,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
a_i2 * params.stride_src0_2 +
|
||||
a_i3 * params.stride_src0_3;
|
||||
|
||||
dst[params.offset_dst + gid.x] = src0[params.offset_src0 + a_index];
|
||||
dst[params.offset_dst + gid_i] = src0[params.offset_src0 + a_index];
|
||||
}
|
||||
}
|
||||
|
||||
@@ -88,6 +88,7 @@ struct Params {
|
||||
ne0: u32,
|
||||
ne1: u32,
|
||||
ne2: u32,
|
||||
ne3: u32,
|
||||
|
||||
eps: f32
|
||||
};
|
||||
@@ -96,10 +97,14 @@ var<workgroup> scratch: array<f32, WG_SIZE>;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(workgroup_id) wid: vec3<u32>,
|
||||
@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(local_invocation_id) lid: vec3<u32>) {
|
||||
|
||||
// one thread per row
|
||||
var i = wid.x;
|
||||
// one workgroup per row
|
||||
var i = wid.x + wid.y * num_wg.x;
|
||||
if (i >= params.ne1 * params.ne2 * params.ne3) {
|
||||
return;
|
||||
}
|
||||
let i3 = i / (params.ne2 * params.ne1);
|
||||
i = i % (params.ne2 * params.ne1);
|
||||
let i2 = i / params.ne1;
|
||||
|
||||
@@ -145,9 +145,12 @@ fn pair_offset(is_neox: bool, is_mrope: bool, is_vision: bool) -> u32 {
|
||||
}
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
// two elements per n_threads
|
||||
if (gid.x >= params.n_threads) {
|
||||
if (gid_i >= params.n_threads) {
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -156,7 +159,7 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
let is_imrope = params.mode == 40;
|
||||
let is_vision = params.mode == 24;
|
||||
|
||||
var i = gid.x * 2; // start index for this thread
|
||||
var i = gid_i * 2; // start index for this thread
|
||||
let i3 = i / (params.ne2 * params.ne1 * params.ne0);
|
||||
i = i % (params.ne2 * params.ne1 * params.ne0);
|
||||
let i2 = i / (params.ne1 * params.ne0);
|
||||
|
||||
@@ -31,6 +31,7 @@ struct Params {
|
||||
ne0: u32,
|
||||
ne1: u32,
|
||||
ne2: u32,
|
||||
ne3: u32,
|
||||
|
||||
eps: f32
|
||||
};
|
||||
@@ -53,10 +54,14 @@ var<workgroup> scratch: array<f32, WG_SIZE * 2u>;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(workgroup_id) wid: vec3<u32>,
|
||||
@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(local_invocation_id) lid: vec3<u32>) {
|
||||
|
||||
// one thread per row
|
||||
var i = wid.x;
|
||||
// one workgroup per row
|
||||
var i = wid.x + wid.y * num_wg.x;
|
||||
if (i >= params.ne1 * params.ne2 * params.ne3) {
|
||||
return;
|
||||
}
|
||||
let i3 = i / (params.ne2 * params.ne1);
|
||||
i = i % (params.ne2 * params.ne1);
|
||||
let i2 = i / params.ne1;
|
||||
|
||||
@@ -84,26 +84,29 @@ fn in_set_view(rel: u32, coords: vec4<u32>) -> bool {
|
||||
}
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
if (gid.x >= params.ne) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i >= params.ne) {
|
||||
return;
|
||||
}
|
||||
|
||||
#ifdef INPLACE
|
||||
let coords = decode_src1_coords(gid.x);
|
||||
let coords = decode_src1_coords(gid_i);
|
||||
|
||||
let src1_idx = params.offset_src1 + src1_idx_from_coords(coords);
|
||||
let dst_idx = params.offset_view + view_rel_from_coords(coords);
|
||||
|
||||
dst[dst_idx] = src1[src1_idx];
|
||||
#else
|
||||
let rel = select(params.ne, gid.x - params.offset_view, gid.x >= params.offset_view);
|
||||
let rel = select(params.ne, gid_i - params.offset_view, gid_i >= params.offset_view);
|
||||
let coords = decode_view_coords(rel);
|
||||
|
||||
if (rel < params.stride_dst13 * params.src1_ne3 && in_set_view(rel, coords)) {
|
||||
dst[gid.x] = src1[params.offset_src1 + src1_idx_from_coords(coords)];
|
||||
dst[gid_i] = src1[params.offset_src1 + src1_idx_from_coords(coords)];
|
||||
} else {
|
||||
dst[gid.x] = src0[params.offset_src0 + gid.x];
|
||||
dst[gid_i] = src0[params.offset_src0 + gid_i];
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
@@ -70,13 +70,16 @@ struct Params {
|
||||
var<uniform> params: Params;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
if (gid.x >= (params.ne3 * params.ne2 * params.n_rows * params.ne0) / VEC_SIZE) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i >= (params.ne3 * params.ne2 * params.n_rows * params.ne0) / VEC_SIZE) {
|
||||
return;
|
||||
}
|
||||
|
||||
let elems_per_row = params.ne0 / VEC_SIZE;
|
||||
var i = gid.x / elems_per_row;
|
||||
var i = gid_i / elems_per_row;
|
||||
|
||||
let i_src3 = i / (params.ne2 * params.n_rows);
|
||||
|
||||
@@ -107,6 +110,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
let i_dst_row = params.offset_dst + idx_val * params.stride_dst1 + i_src2 * params.stride_dst2 + i_src3 * params.stride_dst3;
|
||||
let i_src_row = params.offset_src + i_src1 * params.stride_src1 + i_src2 * params.stride_src2 + i_src3 * params.stride_src3;
|
||||
|
||||
let col_idx = gid.x % elems_per_row;
|
||||
let col_idx = gid_i % elems_per_row;
|
||||
dst[i_dst_row / VEC_SIZE + col_idx] = DST_TYPE(src[i_src_row / VEC_SIZE + col_idx]);
|
||||
}
|
||||
|
||||
@@ -124,9 +124,10 @@ var<workgroup> scratch: array<f32, WG_SIZE>;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(workgroup_id) wid: vec3<u32>,
|
||||
@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(local_invocation_id) lid: vec3<u32>) {
|
||||
|
||||
var i = wid.x;
|
||||
var i = wid.x + wid.y * num_wg.x;
|
||||
let i3 = i / (params.ne2 * params.ne1);
|
||||
i = i % (params.ne2 * params.ne1);
|
||||
let i2 = i / params.ne1;
|
||||
|
||||
@@ -25,9 +25,10 @@ var<workgroup> shared_sum: array<f32, WG_SIZE>;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(workgroup_id) wid: vec3<u32>,
|
||||
@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(local_invocation_id) lid: vec3<u32>) {
|
||||
|
||||
var i = wid.x;
|
||||
var i = wid.x + wid.y * num_wg.x;
|
||||
let i3 = i / (params.ne2 * params.ne1);
|
||||
i = i % (params.ne2 * params.ne1);
|
||||
let i2 = i / params.ne1;
|
||||
|
||||
@@ -2199,6 +2199,7 @@ static struct ggml_tensor * ggml_acc_impl(
|
||||
|
||||
struct ggml_tensor * result = inplace ? ggml_view_tensor(ctx, a) : ggml_dup_tensor(ctx, a);
|
||||
|
||||
GGML_ASSERT(offset < (size_t)(1 << 30));
|
||||
int32_t params[] = { nb1, nb2, nb3, offset, inplace ? 1 : 0 };
|
||||
ggml_set_op_params(result, params, sizeof(params));
|
||||
|
||||
|
||||
@@ -101,6 +101,13 @@ patches = {
|
||||
)],
|
||||
}
|
||||
|
||||
# local changes too large for the replacements above, kept as diffs and applied with git apply
|
||||
patch_files = [
|
||||
# backport of the fix for the stack overflow on deeply nested values (nlohmann/json#5387)
|
||||
# TODO: remove once nlohmann/json releases a version newer than 3.12.0
|
||||
"vendor/nlohmann/json-deep-nesting.patch",
|
||||
]
|
||||
|
||||
for url, filename in vendor.items():
|
||||
print(f"downloading {url} to {filename}") # noqa: NP100
|
||||
urllib.request.urlretrieve(url, filename)
|
||||
@@ -117,6 +124,14 @@ for filename, replacements in patches.items():
|
||||
with open(filename, "w", encoding="utf-8", newline="") as f:
|
||||
f.write(content)
|
||||
|
||||
for patch_file in patch_files:
|
||||
print(f"applying {patch_file}") # noqa: NP100
|
||||
try:
|
||||
subprocess.check_call(["git", "apply", patch_file])
|
||||
except subprocess.CalledProcessError:
|
||||
print(f"Error: cannot apply {patch_file}, upstream code has changed") # noqa: NP100
|
||||
sys.exit(1)
|
||||
|
||||
print("Splitting httplib.h...") # noqa: NP100
|
||||
try:
|
||||
subprocess.check_call([
|
||||
|
||||
@@ -1234,7 +1234,6 @@ bool llm_arch_supports_sm_tensor(const llm_arch & arch) {
|
||||
case LLM_ARCH_KIMI_K3:
|
||||
case LLM_ARCH_GLM5_NEXT:
|
||||
case LLM_ARCH_QWEN3TTS:
|
||||
case LLM_ARCH_K2_HORIZON:
|
||||
return false;
|
||||
default:
|
||||
return true;
|
||||
|
||||
+2
-16
@@ -2442,23 +2442,9 @@ uint32_t llama_context::graph_max_nodes(uint32_t n_tokens) const {
|
||||
}
|
||||
}
|
||||
|
||||
uint32_t n_sampling_nodes = 0;
|
||||
uint32_t n_sampling_nodes_max = 0;
|
||||
// every sampler builds n_outputs_max_per_seq chains, see llm_graph_context::build_sampling
|
||||
for (const auto & [seq_id, sampler] : sampling.samplers) {
|
||||
const uint32_t n_nodes = llama_sampler_backend_n_nodes(sampler);
|
||||
n_sampling_nodes += n_nodes;
|
||||
if (cparams.n_outputs_max_per_seq > 1) {
|
||||
n_sampling_nodes_max = std::max(n_sampling_nodes_max, n_nodes);
|
||||
}
|
||||
}
|
||||
|
||||
const uint32_t n_sampling_outputs_max = std::min<uint64_t>(
|
||||
std::min(n_tokens, cparams.n_outputs_max),
|
||||
(uint64_t) cparams.n_seq_max * cparams.n_outputs_max_per_seq);
|
||||
|
||||
res += n_sampling_nodes;
|
||||
if (n_sampling_outputs_max > 1) {
|
||||
res += (n_sampling_outputs_max - 1) * n_sampling_nodes_max;
|
||||
res += llama_sampler_backend_n_nodes(sampler) * cparams.n_outputs_max_per_seq;
|
||||
}
|
||||
|
||||
if (cparams.training) {
|
||||
|
||||
+77
-38
@@ -1963,6 +1963,11 @@ ggml_tensor * llm_graph_context::build_ffn(
|
||||
cur = ggml_geglu(ctx0, cur);
|
||||
cb(cur, "ffn_geglu", il);
|
||||
} break;
|
||||
case LLM_FFN_GEGLU_ERF:
|
||||
{
|
||||
cur = ggml_geglu_erf(ctx0, cur);
|
||||
cb(cur, "ffn_geglu_erf", il);
|
||||
} break;
|
||||
case LLM_FFN_REGLU:
|
||||
{
|
||||
cur = ggml_reglu(ctx0, cur);
|
||||
@@ -2466,19 +2471,57 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd, float to
|
||||
|
||||
auto inp = std::make_unique<llm_graph_input_embd>(n_embd_inp);
|
||||
|
||||
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_tokens);
|
||||
cb(inp->tokens, "inp_tokens", -1);
|
||||
ggml_set_input(inp->tokens);
|
||||
res->t_inp_tokens = inp->tokens;
|
||||
// mixed path (ubatch.is_mixed()): set_rows the token rows into a copy of the embd rows, with its own inputs as select branches must not share tensors
|
||||
// TODO: use inp->tokens and inp->embd once ggml_build_forward_select allows it
|
||||
const bool has_mixed = llm_arch_supports_mixed_batch(arch) && cparams.ctx_type == LLAMA_CONTEXT_TYPE_DEFAULT;
|
||||
|
||||
inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd_inp, ubatch.n_tokens);
|
||||
cb(inp->embd, "inp_embd", -1);
|
||||
ggml_set_input(inp->embd);
|
||||
const int64_t n_tok_rows = has_mixed ? llm_graph_n_tok_rows(ubatch) : 0;
|
||||
|
||||
// token embeddings with lora and padding
|
||||
auto build_tok = [&](ggml_tensor * ids) {
|
||||
ggml_tensor * cur = ggml_get_rows(ctx0, tok_embd, ids);
|
||||
// we have 3 standard paths to produce the input embeddings for the first layer:
|
||||
// - embd0: extract from the token embeddings weight (`tok_embd`) using the input token ids
|
||||
// - embd1: pass raw embeddings, skipping the `tok_embd`
|
||||
// - embd2: mixed path of both tokens ids + raw embeddings (if supported)
|
||||
ggml_tensor * embd0 = nullptr;
|
||||
ggml_tensor * embd1 = nullptr;
|
||||
ggml_tensor * embd2 = nullptr;
|
||||
|
||||
// construct the input tensors
|
||||
{
|
||||
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_tokens);
|
||||
cb(inp->tokens, "inp_tokens", -1);
|
||||
ggml_set_input(inp->tokens);
|
||||
res->t_inp_tokens = inp->tokens;
|
||||
|
||||
inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd_inp, ubatch.n_tokens);
|
||||
cb(inp->embd, "inp_embd", -1);
|
||||
ggml_set_input(inp->embd);
|
||||
|
||||
if (has_mixed) {
|
||||
inp->mixed_tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tok_rows);
|
||||
cb(inp->mixed_tokens, "inp_mixed_tokens", -1);
|
||||
ggml_set_input(inp->mixed_tokens);
|
||||
}
|
||||
}
|
||||
|
||||
// the embeddings placeholders for the 3 paths
|
||||
// we use ggml_build_forward_order to make the GET_ROWS ops stick at the beginning of the compute graph
|
||||
// this way the embeddings remain in the host buffer, and the GET_ROWS run before any other computations
|
||||
{
|
||||
embd0 = ggml_get_rows(ctx0, tok_embd, inp->tokens);
|
||||
ggml_build_forward_order(gf, embd0);
|
||||
|
||||
embd1 = inp->embd;
|
||||
|
||||
if (has_mixed) {
|
||||
embd2 = ggml_get_rows(ctx0, tok_embd, inp->mixed_tokens);
|
||||
ggml_build_forward_order(gf, embd2);
|
||||
}
|
||||
}
|
||||
|
||||
// helper for extracting token embeddings with lora and padding
|
||||
// TODO: when lora is active, this is likely going to cause issues similar to https://github.com/ggml-org/llama.cpp/pull/30160
|
||||
// need to add lora tests and refactor the logic to make the lora GET_ROWS go at the front of the graph
|
||||
auto build_tok = [&](ggml_tensor * cur, ggml_tensor * ids) {
|
||||
// apply lora for embedding tokens if needed
|
||||
for (const auto & lora : *loras) {
|
||||
llama_adapter_lora_weight * lw = lora.first->get_weight(tok_embd);
|
||||
@@ -2509,21 +2552,15 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd, float to
|
||||
std::array<ggml_tensor *, 3> inps = {};
|
||||
|
||||
// token embeddings path (ubatch.token != nullptr)
|
||||
inps[0] = build_tok(inp->tokens);
|
||||
inps[0] = build_tok(embd0, inp->tokens);
|
||||
|
||||
// vector embeddings path (ubatch.embd != nullptr)
|
||||
inps[1] = inp->embd;
|
||||
inps[1] = embd1;
|
||||
|
||||
assert(ggml_are_same_shape (inps[0], inps[1]));
|
||||
assert(ggml_are_same_stride(inps[0], inps[1]));
|
||||
|
||||
// mixed path (ubatch.is_mixed()): set_rows the token rows into a copy of the embd rows, with its own inputs as select branches must not share tensors
|
||||
// TODO: use inp->tokens and inp->embd once ggml_build_forward_select allows it
|
||||
const bool has_mixed = llm_arch_supports_mixed_batch(arch) && cparams.ctx_type == LLAMA_CONTEXT_TYPE_DEFAULT;
|
||||
if (has_mixed) {
|
||||
const int64_t n_tok_rows = llm_graph_n_tok_rows(ubatch);
|
||||
|
||||
inp->mixed_tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tok_rows);
|
||||
cb(inp->mixed_tokens, "inp_mixed_tokens", -1);
|
||||
ggml_set_input(inp->mixed_tokens);
|
||||
|
||||
inp->mixed_slots = ggml_new_tensor_1d(ctx0, GGML_TYPE_I64, n_tok_rows);
|
||||
cb(inp->mixed_slots, "inp_mixed_slots", -1);
|
||||
ggml_set_input(inp->mixed_slots);
|
||||
@@ -2533,11 +2570,12 @@ ggml_tensor * llm_graph_context::build_inp_embd(ggml_tensor * tok_embd, float to
|
||||
ggml_set_input(inp->mixed_embd);
|
||||
|
||||
// note: set_rows writes into its destination, so it gets a copy of the input
|
||||
inps[2] = ggml_set_rows(ctx0, ggml_dup(ctx0, inp->mixed_embd), build_tok(inp->mixed_tokens), inp->mixed_slots);
|
||||
}
|
||||
ggml_tensor * embd_mixed = build_tok(embd2, inp->mixed_tokens);
|
||||
inps[2] = ggml_set_rows(ctx0, ggml_dup(ctx0, inp->mixed_embd), embd_mixed, inp->mixed_slots);
|
||||
|
||||
assert(ggml_are_same_shape (inps[0], inps[1]));
|
||||
assert(ggml_are_same_stride(inps[0], inps[1]));
|
||||
assert(ggml_are_same_shape (inps[0], inps[2]));
|
||||
assert(ggml_are_same_stride(inps[0], inps[2]));
|
||||
}
|
||||
|
||||
const int idx = ubatch.is_mixed() ? 2 : ubatch.token ? 0 : 1;
|
||||
|
||||
@@ -3965,8 +4003,7 @@ void llm_graph_context::build_sampling() const {
|
||||
// res->t_logits will contain logits for all tokens that want the logits calculated (logits=1 or output=1)
|
||||
GGML_ASSERT(res->t_logits != nullptr && "missing t_logits tensor");
|
||||
|
||||
// add a dummy row to keep the single-output graph static regardless of active samplers
|
||||
// multi-output graphs can still vary with the number of output rows
|
||||
// the padding row gives the chains without a row of the ubatch a valid input, even with no output at all
|
||||
ggml_tensor * logits_t = ggml_pad(ctx0, res->t_logits, 0, 1, 0, 0);
|
||||
|
||||
for (const auto & entry : samplers) {
|
||||
@@ -3975,18 +4012,20 @@ void llm_graph_context::build_sampling() const {
|
||||
}
|
||||
}
|
||||
|
||||
static const std::vector<uint32_t> dummy_row = { 0 };
|
||||
static const std::vector<uint32_t> no_rows;
|
||||
|
||||
// every sampler builds n_outputs_max_per_seq chains, like the reserve, so the graph keeps its topology
|
||||
// whatever rows the ubatch outputs: a chain without a row works on the first row and is not selected
|
||||
for (const auto & [seq_id, sampler] : samplers) {
|
||||
const auto it = sampling_rows.find(seq_id);
|
||||
const auto & rows = it != sampling_rows.end() ? it->second : no_rows;
|
||||
|
||||
// inactive samplers always work on the first row
|
||||
const bool active = it != sampling_rows.end();
|
||||
const auto & rows = active ? it->second : dummy_row;
|
||||
const int i_out = active ? 1 : 0;
|
||||
for (uint32_t i = 0; i < cparams.n_outputs_max_per_seq; ++i) {
|
||||
const bool active = i < rows.size();
|
||||
const uint32_t row = active ? rows[i] : 0;
|
||||
const int i_out = active ? 1 : 0;
|
||||
|
||||
for (uint32_t i = 0; i < rows.size(); ++i) {
|
||||
ggml_tensor * logits_seq = ggml_view_1d(ctx0, logits_t, logits_t->ne[0], rows[i] * logits_t->nb[1]);
|
||||
ggml_tensor * logits_seq = ggml_view_1d(ctx0, logits_t, logits_t->ne[0], row * logits_t->nb[1]);
|
||||
ggml_format_name(logits_seq, "logits_seq_%d_%u", seq_id, i);
|
||||
|
||||
struct llama_sampler_data data = {
|
||||
@@ -4001,7 +4040,7 @@ void llm_graph_context::build_sampling() const {
|
||||
|
||||
if (data.sampled != nullptr) {
|
||||
if (active) {
|
||||
res->t_sampled[rows[i]] = data.sampled;
|
||||
res->t_sampled[row] = data.sampled;
|
||||
}
|
||||
outs[1] = data.sampled;
|
||||
ggml_build_forward_select(gf, outs.data(), outs.size(), i_out);
|
||||
@@ -4009,7 +4048,7 @@ void llm_graph_context::build_sampling() const {
|
||||
|
||||
if (data.probs != nullptr) {
|
||||
if (active) {
|
||||
res->t_sampled_probs[rows[i]] = data.probs;
|
||||
res->t_sampled_probs[row] = data.probs;
|
||||
}
|
||||
outs[1] = data.probs;
|
||||
ggml_build_forward_select(gf, outs.data(), outs.size(), i_out);
|
||||
@@ -4017,7 +4056,7 @@ void llm_graph_context::build_sampling() const {
|
||||
|
||||
if (data.logits != nullptr) {
|
||||
if (active) {
|
||||
res->t_sampled_logits[rows[i]] = data.logits;
|
||||
res->t_sampled_logits[row] = data.logits;
|
||||
}
|
||||
outs[1] = data.logits;
|
||||
ggml_build_forward_select(gf, outs.data(), outs.size(), i_out);
|
||||
@@ -4025,7 +4064,7 @@ void llm_graph_context::build_sampling() const {
|
||||
|
||||
if (data.candidates != nullptr) {
|
||||
if (active) {
|
||||
res->t_candidates[rows[i]] = data.candidates;
|
||||
res->t_candidates[row] = data.candidates;
|
||||
}
|
||||
outs[1] = data.candidates;
|
||||
ggml_build_forward_select(gf, outs.data(), outs.size(), i_out);
|
||||
|
||||
@@ -62,6 +62,7 @@ enum llm_ffn_op_type : int {
|
||||
LLM_FFN_RELU_SQR,
|
||||
LLM_FFN_SWIGLU,
|
||||
LLM_FFN_GEGLU,
|
||||
LLM_FFN_GEGLU_ERF,
|
||||
LLM_FFN_REGLU,
|
||||
LLM_FFN_SWIGLU_OAI_MOE,
|
||||
LLM_FFN_SITU, // kimi-k3
|
||||
|
||||
+14
-11
@@ -1064,18 +1064,21 @@ static llama_rope_scaling_type llama_rope_scaling_type_from_string(const std::st
|
||||
return LLAMA_ROPE_SCALING_TYPE_UNSPECIFIED;
|
||||
}
|
||||
|
||||
// Maps the GGUF `<arch>.hidden_activation` string to the FFN op type used by the
|
||||
// graph builders. Only gated activations that map cleanly to llm_ffn_op_type are
|
||||
// listed; unrecognized values fall back to GeGLU, which matches the historical
|
||||
// default for ModernBert-style architectures.
|
||||
// Maps GGUF activation names to the FFN op type used by the graph builders.
|
||||
static const std::map<std::string, llm_ffn_op_type> LLM_FFN_OP_TYPES_FROM_STRING = {
|
||||
{ "gelu", LLM_FFN_GEGLU },
|
||||
{ "geglu", LLM_FFN_GEGLU },
|
||||
{ "silu", LLM_FFN_SWIGLU },
|
||||
{ "swish", LLM_FFN_SWIGLU },
|
||||
{ "swiglu", LLM_FFN_SWIGLU },
|
||||
{ "relu", LLM_FFN_RELU },
|
||||
{ "reglu", LLM_FFN_REGLU },
|
||||
{ "gelu", LLM_FFN_GEGLU_ERF },
|
||||
{ "gelu_python", LLM_FFN_GEGLU_ERF },
|
||||
{ "gelu_pytorch_tanh", LLM_FFN_GEGLU },
|
||||
{ "gelu_new", LLM_FFN_GEGLU },
|
||||
{ "gelu_fast", LLM_FFN_GEGLU },
|
||||
{ "gelu_accurate", LLM_FFN_GEGLU },
|
||||
{ "gelu_python_tanh", LLM_FFN_GEGLU },
|
||||
{ "geglu", LLM_FFN_GEGLU },
|
||||
{ "silu", LLM_FFN_SWIGLU },
|
||||
{ "swish", LLM_FFN_SWIGLU },
|
||||
{ "swiglu", LLM_FFN_SWIGLU },
|
||||
{ "relu", LLM_FFN_RELU },
|
||||
{ "reglu", LLM_FFN_REGLU },
|
||||
};
|
||||
|
||||
// transformers names, "gelu" is the exact (erf) variant
|
||||
|
||||
+26
-18
@@ -159,6 +159,14 @@ llama_model_gemma4::graph::graph(const llama_model & model, const llm_graph_para
|
||||
ggml_tensor * cur;
|
||||
ggml_tensor * inpL;
|
||||
|
||||
// do the PLE first to guarantee it is done in the host buffer
|
||||
// ref: https://github.com/ggml-org/llama.cpp/pull/30160
|
||||
ggml_tensor * inp_per_layer = nullptr;
|
||||
if (model.per_layer_tok_embd) {
|
||||
inp_per_layer = build_inp_per_layer();
|
||||
ggml_build_forward_expand(gf, inp_per_layer);
|
||||
}
|
||||
|
||||
// important: do not normalize weights for raw embeddings input (i.e. encoded image emdeddings)
|
||||
inpL = build_inp_embd(model.tok_embd, sqrtf(n_embd));
|
||||
cb(inpL, "inp_scaled", -1);
|
||||
@@ -171,10 +179,11 @@ llama_model_gemma4::graph::graph(const llama_model & model, const llm_graph_para
|
||||
|
||||
ggml_tensor * inp_out_ids = build_inp_out_ids();
|
||||
|
||||
ggml_tensor * inp_per_layer = nullptr;
|
||||
if (model.per_layer_tok_embd) {
|
||||
inp_per_layer = build_inp_per_layer();
|
||||
ggml_build_forward_expand(gf, inp_per_layer);
|
||||
const float tok_embd_scale = sqrtf((float) n_embd_per_layer);
|
||||
|
||||
inp_per_layer = ggml_scale (ctx0, inp_per_layer, tok_embd_scale);
|
||||
inp_per_layer = ggml_reshape_3d(ctx0, inp_per_layer, n_embd_per_layer, n_layer, inp_per_layer->ne[1]);
|
||||
|
||||
// inp_per_layer shape: [n_embd_per_layer, n_tokens, n_layer]
|
||||
inp_per_layer = project_per_layer_inputs(inpL, inp_per_layer);
|
||||
@@ -448,10 +457,13 @@ public:
|
||||
llama_prefetch_rows(ple, ubatch->token, ubatch->n_tokens);
|
||||
}
|
||||
ggml_backend_tensor_set(tokens, ubatch->token, 0, ubatch->n_tokens * ggml_element_size(tokens));
|
||||
} else if (prefetch) {
|
||||
// [TAG_GEMMA4_IMG_PADDING]
|
||||
} else {
|
||||
const int32_t padding = 0;
|
||||
llama_prefetch_rows(ple, &padding, 1);
|
||||
if (prefetch) {
|
||||
// [TAG_GEMMA4_IMG_PADDING]
|
||||
llama_prefetch_rows(ple, &padding, 1);
|
||||
}
|
||||
ggml_backend_tensor_set(token0, &padding, 0, ggml_element_size(token0));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -460,6 +472,7 @@ public:
|
||||
}
|
||||
|
||||
ggml_tensor * tokens = nullptr;
|
||||
ggml_tensor * token0 = nullptr;
|
||||
|
||||
const llama_model & model;
|
||||
};
|
||||
@@ -470,30 +483,25 @@ ggml_tensor * llama_model_gemma4::graph::build_inp_per_layer() {
|
||||
auto inp = std::make_unique<llm_graph_input_gemma4_ple>(model);
|
||||
|
||||
ggml_tensor * inp_per_layer;
|
||||
float tok_embd_scale = sqrtf((float) n_embd_per_layer);
|
||||
// mixed ubatch: embd rows have token id 0, same padding row as below
|
||||
// TODO: use ggml_build_forward_select
|
||||
if (ubatch.token) {
|
||||
inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_tokens);
|
||||
ggml_set_input(inp->tokens);
|
||||
res->t_inp_tokens = inp->tokens;
|
||||
|
||||
inp_per_layer = ggml_get_rows (ctx0, model.per_layer_tok_embd, inp->tokens);
|
||||
inp_per_layer = ggml_reshape_3d(ctx0, inp_per_layer, n_embd_per_layer, n_layer, n_tokens);
|
||||
inp_per_layer = ggml_scale (ctx0, inp_per_layer, tok_embd_scale);
|
||||
inp_per_layer = ggml_get_rows(ctx0, model.per_layer_tok_embd, inp->tokens);
|
||||
cb(inp_per_layer, "inp_per_layer_selected", -1);
|
||||
} else {
|
||||
// [TAG_GEMMA4_IMG_PADDING]
|
||||
// Multimodal embedding path: use padding token (ID=0) embedding
|
||||
// TODO: verify if this is the correct behavior in transformers implementation
|
||||
const int64_t embd_size = model.per_layer_tok_embd->ne[0]; // n_embd_per_layer * n_layer
|
||||
|
||||
// Extract and dequantize padding token embedding (row 0)
|
||||
ggml_tensor * padding = ggml_view_1d(ctx0, model.per_layer_tok_embd, embd_size, 0);
|
||||
inp_per_layer = ggml_cast (ctx0, padding, GGML_TYPE_F32);
|
||||
inp_per_layer = ggml_scale(ctx0, inp_per_layer, tok_embd_scale);
|
||||
// [TAG_GEMMA4_IMG_PADDING]
|
||||
inp->token0 = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, 1);
|
||||
ggml_set_input(inp->token0);
|
||||
res->t_inp_tokens = inp->token0;
|
||||
|
||||
// Reshape to [n_embd_per_layer, n_layer, 1]
|
||||
inp_per_layer = ggml_reshape_3d(ctx0, inp_per_layer, n_embd_per_layer, n_layer, 1);
|
||||
inp_per_layer = ggml_get_rows(ctx0, model.per_layer_tok_embd, inp->token0);
|
||||
cb(inp_per_layer, "inp_per_layer_multimodal", -1);
|
||||
}
|
||||
res->add_input(std::move(inp));
|
||||
|
||||
@@ -17,10 +17,10 @@ void llama_model_modern_bert::load_arch_hparams(llama_model_loader & ml) {
|
||||
|
||||
// Some ModernBert derivatives (e.g. IBM Granite Embedding 97m R2) use
|
||||
// SiLU/SwiGLU in the FFN instead of the default GELU/GeGLU.
|
||||
hparams.llm_ffn_op = LLM_FFN_GEGLU;
|
||||
hparams.llm_ffn_op = LLM_FFN_GEGLU_ERF;
|
||||
std::string hidden_act;
|
||||
if (ml.get_key(LLM_KV_HIDDEN_ACT, hidden_act, false)) {
|
||||
hparams.llm_ffn_op = llm_ffn_op_type_from_string(hidden_act, LLM_FFN_GEGLU);
|
||||
hparams.llm_ffn_op = llm_ffn_op_type_from_string(hidden_act, LLM_FFN_GEGLU_ERF);
|
||||
}
|
||||
|
||||
// GGUFs without a classifier pooling type use mean (gte-reranker-modernbert-base)
|
||||
|
||||
+11
-11
@@ -405,10 +405,6 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa
|
||||
int sections[4];
|
||||
std::copy(std::begin(hparams.rope_sections), std::begin(hparams.rope_sections) + 4, sections);
|
||||
|
||||
ggml_tensor * inpL = build_inp_embd(model.tok_embd);
|
||||
cb(inpL, "model.input_embed", -1);
|
||||
ggml_build_forward_expand(gf, inpL);
|
||||
|
||||
auto * inp = build_inp_mem_hybrid();
|
||||
|
||||
// qwen4exp always builds llama_memory_hybrid_idx, so this downcast is safe
|
||||
@@ -421,6 +417,17 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa
|
||||
"the indexer cache must track the attention cache cell for cell");
|
||||
}
|
||||
|
||||
ggml_tensor * ple_emb = nullptr;
|
||||
if (hparams.ple_n_heads > 0) {
|
||||
ple_emb = build_inp_ple(mctx_hyb);
|
||||
// make sure ple_emb and build_inp_embd are in the same graph split
|
||||
ggml_build_forward_expand(gf, ple_emb);
|
||||
}
|
||||
|
||||
ggml_tensor * inpL = build_inp_embd(model.tok_embd);
|
||||
cb(inpL, "model.input_embed", -1);
|
||||
ggml_build_forward_expand(gf, inpL);
|
||||
|
||||
// the QSA layers share one set of k-pool inputs
|
||||
// the CUDA lightning indexer takes 32 or 64 heads, QSA has a few, so it scores with plain ops
|
||||
llm_graph_input_kpool * inp_kpool = nullptr;
|
||||
@@ -431,13 +438,6 @@ llama_model_qwen4exp::graph::graph(const llama_model & model, const llm_graph_pa
|
||||
ggml_tensor * inp_pos = build_inp_pos();
|
||||
ggml_tensor * inp_out_ids = build_inp_out_ids();
|
||||
|
||||
ggml_tensor * ple_emb = nullptr;
|
||||
if (hparams.ple_n_heads > 0) {
|
||||
ple_emb = build_inp_ple(mctx_hyb);
|
||||
// make sure ple_emb and build_inp_embd are in the same graph split
|
||||
ggml_build_forward_expand(gf, ple_emb);
|
||||
}
|
||||
|
||||
// the wide residual starts as hc identical copies of the embedding
|
||||
ggml_tensor * res_hc = ggml_repeat_4d(ctx0,
|
||||
ggml_reshape_3d(ctx0, inpL, n_embd, 1, n_tokens),
|
||||
|
||||
@@ -245,7 +245,6 @@ llama_build_and_test(
|
||||
peg-parser/test-basic.cpp
|
||||
peg-parser/test-gbnf-generation.cpp
|
||||
peg-parser/test-json-parser.cpp
|
||||
peg-parser/test-json-serialization.cpp
|
||||
peg-parser/test-python-dict-parser.cpp
|
||||
peg-parser/test-unicode.cpp
|
||||
peg-parser/tests.h
|
||||
|
||||
@@ -1,28 +0,0 @@
|
||||
#include "tests.h"
|
||||
|
||||
void test_json_serialization(testing &t) {
|
||||
auto original = build_peg_parser([](common_peg_parser_builder & p) {
|
||||
return "<tool_call>" + p.json() + "</tool_call>";
|
||||
});
|
||||
|
||||
auto json_serialized = original.to_json().dump();
|
||||
|
||||
t.test("compare before/after", [&](testing &t) {
|
||||
auto deserialized = common_peg_arena::from_json(common_json::parse(json_serialized));
|
||||
|
||||
// Test complex JSON
|
||||
std::string input = R"({"name": "test", "values": [1, 2, 3], "nested": {"a": true}})";
|
||||
common_peg_parse_context ctx1(input);
|
||||
common_peg_parse_context ctx2(input);
|
||||
|
||||
auto result1 = original.parse(ctx1);
|
||||
auto result2 = deserialized.parse(ctx2);
|
||||
|
||||
t.assert_equal("both_succeed", result1.success(), result2.success());
|
||||
t.assert_equal("same_end_pos", result1.end, result2.end);
|
||||
});
|
||||
|
||||
t.bench("deserialize", [&]() {
|
||||
auto deserialized = common_peg_arena::from_json(common_json::parse(json_serialized));
|
||||
}, 100);
|
||||
}
|
||||
@@ -21,5 +21,4 @@ void test_basic(testing &t);
|
||||
void test_json_parser(testing &t);
|
||||
void test_gbnf_generation(testing &t);
|
||||
void test_unicode(testing &t);
|
||||
void test_json_serialization(testing &t);
|
||||
void test_python_dict_parser(testing &t);
|
||||
|
||||
@@ -4768,6 +4768,104 @@ struct test_ssm_scan_rollback : public test_case {
|
||||
}
|
||||
};
|
||||
|
||||
// GGML_OP_SSM_SCAN + GGML_OP_CPY (recurrent cache fusion)
|
||||
struct test_ssm_scan_cache_fusion : public test_case {
|
||||
const ggml_type type;
|
||||
|
||||
const int64_t d_state;
|
||||
const int64_t head_dim;
|
||||
const int64_t n_head;
|
||||
const int64_t n_group;
|
||||
const int64_t n_seq_tokens;
|
||||
const int64_t n_seqs;
|
||||
const int64_t K; // snapshot slot count (1 = final state only)
|
||||
|
||||
ggml_tensor * cpy_node = nullptr;
|
||||
|
||||
std::string vars() override {
|
||||
return VARS_TO_STR8(type, d_state, head_dim, n_head, n_group, n_seq_tokens, n_seqs, K);
|
||||
}
|
||||
|
||||
test_ssm_scan_cache_fusion(ggml_type type = GGML_TYPE_F32,
|
||||
int64_t d_state = 128, int64_t head_dim = 64, int64_t n_head = 16, int64_t n_group = 2,
|
||||
int64_t n_seq_tokens = 4, int64_t n_seqs = 1, int64_t K = 4)
|
||||
: type(type), d_state(d_state), head_dim(head_dim), n_head(n_head), n_group(n_group),
|
||||
n_seq_tokens(n_seq_tokens), n_seqs(n_seqs), K(K) {}
|
||||
|
||||
ggml_tensor * build_graph(ggml_context * ctx) override {
|
||||
const int64_t D = d_state * head_dim * n_head;
|
||||
const int64_t n_written = std::min<int64_t>(n_seq_tokens, K);
|
||||
|
||||
// more cache rows per slot than seqs and a non-zero first row, so a wrong slot stride or offset shows up
|
||||
const int64_t mem_size = n_seqs + 2;
|
||||
const int64_t kv_head = 1;
|
||||
|
||||
ggml_tensor * s = ggml_new_tensor_4d(ctx, type, d_state, head_dim, n_head, n_seqs);
|
||||
ggml_tensor * x = ggml_new_tensor_4d(ctx, type, head_dim, n_head, n_seq_tokens, n_seqs);
|
||||
ggml_tensor * dt = ggml_new_tensor_3d(ctx, type, n_head, n_seq_tokens, n_seqs);
|
||||
ggml_tensor * A = ggml_new_tensor_2d(ctx, type, 1, n_head);
|
||||
ggml_tensor * B = ggml_new_tensor_4d(ctx, type, d_state, n_group, n_seq_tokens, n_seqs);
|
||||
ggml_tensor * C = ggml_new_tensor_4d(ctx, type, d_state, n_group, n_seq_tokens, n_seqs);
|
||||
ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, n_seqs);
|
||||
ggml_set_name(A, "A");
|
||||
ggml_set_name(ids, "ids");
|
||||
|
||||
ggml_tensor * out = ggml_ssm_scan(ctx, s, x, dt, A, B, C, ids, K);
|
||||
ggml_set_name(out, "ssm_out");
|
||||
|
||||
// snapshot tail view [D, n_seqs, n_written]
|
||||
ggml_tensor * src = ggml_view_3d(ctx, out,
|
||||
D, n_seqs, n_written,
|
||||
ggml_row_size(out->type, D),
|
||||
ggml_row_size(out->type, D * n_seqs),
|
||||
ggml_row_size(out->type, ggml_nelements(x)));
|
||||
|
||||
// recurrent cache view [D, n_seqs, n_written]
|
||||
ggml_tensor * cache = ggml_new_tensor_2d(ctx, type, D, mem_size * n_written);
|
||||
ggml_set_name(cache, "cache");
|
||||
ggml_tensor * dst = ggml_view_3d(ctx, cache,
|
||||
D, n_seqs, n_written,
|
||||
cache->nb[1],
|
||||
mem_size * cache->nb[1],
|
||||
kv_head * cache->nb[1]);
|
||||
|
||||
ggml_tensor * cpy = ggml_cpy(ctx, src, dst);
|
||||
ggml_set_name(cpy, "ssm_cache_cpy");
|
||||
cpy_node = cpy;
|
||||
|
||||
// read the cpy output so that neither the scan nor the cpy is the graph output (cont, since the cache view is strided)
|
||||
ggml_tensor * res = ggml_sum(ctx, ggml_cont(ctx, cpy));
|
||||
return res;
|
||||
}
|
||||
|
||||
std::string op_desc(ggml_tensor * t) override {
|
||||
GGML_UNUSED(t);
|
||||
return "SSM_SCAN_CACHE_FUSION";
|
||||
}
|
||||
|
||||
bool run_whole_graph() override { return true; }
|
||||
std::vector<ggml_tensor *> fusion_test_nodes() override { return { cpy_node }; }
|
||||
|
||||
void initialize_tensors(ggml_context * ctx) override {
|
||||
for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != nullptr; t = ggml_get_next_tensor(ctx, t)) {
|
||||
if (ggml_is_view_op(t->op)) { continue; }
|
||||
if (strcmp(t->name, "ids") == 0) {
|
||||
std::vector<int32_t> data(t->ne[0]);
|
||||
for (int i = 0; i < t->ne[0]; i++) {
|
||||
data[i] = i;
|
||||
}
|
||||
ggml_backend_tensor_set(t, data.data(), 0, t->ne[0] * sizeof(int32_t));
|
||||
} else if (strcmp(t->name, "A") == 0) {
|
||||
init_tensor_uniform(t, -1.0f, -0.5f);
|
||||
} else if (strcmp(t->name, "cache") == 0) {
|
||||
init_tensor_uniform(t, 0.0f, 0.0f);
|
||||
} else {
|
||||
init_tensor_uniform(t);
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
// GGML_OP_RWKV_WKV6
|
||||
struct test_rwkv_wkv6 : public test_case {
|
||||
const ggml_type type;
|
||||
@@ -10445,6 +10543,16 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
|
||||
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 128, 2)); // SSD multi-chunk, no tail (exercises the chunk-to-chunk state handoff)
|
||||
test_cases.emplace_back(new test_ssm_scan(GGML_TYPE_F32, 128, 64, 16, 2, 128, 2, false, /*K=*/1, /*weak_decay=*/true)); // SSD multi-chunk, carried state not numerically negligible
|
||||
|
||||
// ssm_scan + cache cpy fusion
|
||||
test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 4, 1, 4));
|
||||
test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 1, 1, 4)); // n_seq_tokens < K
|
||||
test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 8, 1, 3)); // n_seq_tokens > K
|
||||
test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 96, 64, 16, 2, 4, 1, 4));
|
||||
test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 256, 64, 8, 2, 4, 1, 4));
|
||||
test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 1, 1, 1)); // K == 1, final state only
|
||||
test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 4, 1, 1));
|
||||
test_cases.emplace_back(new test_ssm_scan_cache_fusion(GGML_TYPE_F32, 128, 64, 16, 2, 300, 1, 1)); // K == 1, SSD path over two chunks
|
||||
|
||||
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 1, 1));
|
||||
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 32, 1));
|
||||
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 32, 4));
|
||||
|
||||
@@ -450,7 +450,7 @@ static int debug_single_template(const debug_options & opts) {
|
||||
generation_params params = prepare_debug_params(opts, tools);
|
||||
common_chat_params parser_data;
|
||||
if (std::optional<common_chat_params> spec_tmpl =
|
||||
common_chat_try_specialized_template(chat_template, template_source, params)) {
|
||||
common_chat_try_specialized_template(chat_template, params)) {
|
||||
LOG_ERR("\n");
|
||||
LOG_ERR("This template uses a specialized parser, analysis results will not be available.\n");
|
||||
parser_data = *spec_tmpl;
|
||||
@@ -484,8 +484,7 @@ static int debug_single_template(const debug_options & opts) {
|
||||
|
||||
if (!std::empty(parser_data.parser)) {
|
||||
LOG_ERR("\n=== Generated Parser ===\n");
|
||||
common_peg_arena arena;
|
||||
arena.load(parser_data.parser);
|
||||
const common_peg_arena & arena = parser_data.parser;
|
||||
LOG_ERR("%s\n", arena.dump(arena.root()).c_str());
|
||||
|
||||
LOG_ERR("\n=== Generated Grammar ===\n");
|
||||
|
||||
+114
-51
@@ -1090,26 +1090,6 @@ struct peg_test_case {
|
||||
std::vector<std::string> expect_rules;
|
||||
};
|
||||
|
||||
struct make_peg_parser {
|
||||
common_chat_params params_;
|
||||
common_peg_arena arena_;
|
||||
bool detailed_debug_;
|
||||
|
||||
make_peg_parser(common_chat_templates * tmpls,
|
||||
const common_chat_templates_inputs & inputs,
|
||||
bool detailed_debug = false) {
|
||||
detailed_debug_ = detailed_debug;
|
||||
params_ = common_chat_templates_apply(tmpls, inputs);
|
||||
arena_.load(params_.parser);
|
||||
}
|
||||
|
||||
common_chat_msg parse(const std::string & msg, bool is_partial) const {
|
||||
common_chat_parser_params parser_params(params_);
|
||||
parser_params.debug = detailed_debug_;
|
||||
return common_chat_peg_parse(arena_, common_chat_input(msg), is_partial, parser_params);
|
||||
}
|
||||
};
|
||||
|
||||
// Global template filter for --template flag
|
||||
static std::string g_template_filter;
|
||||
|
||||
@@ -1161,14 +1141,19 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
tc.expect.role = "assistant";
|
||||
}
|
||||
|
||||
auto parser = make_peg_parser(tmpls, tc.params, detailed_debug);
|
||||
common_chat_session_params session_params;
|
||||
session_params.debug = detailed_debug;
|
||||
common_chat_session session(tmpls, nullptr, tc.params, session_params);
|
||||
const auto & parser = session.parser();
|
||||
common_params_sampling sampling;
|
||||
session.apply_sampling(sampling);
|
||||
if (detailed_debug) {
|
||||
LOG_DBG("Using parser: \n%s\n", parser.arena_.dump(parser.arena_.root()).c_str());
|
||||
LOG_DBG("Generation prompt: '%s'\n", parser.params_.generation_prompt.c_str());
|
||||
LOG_DBG("Using parser: \n%s\n", parser.dump(parser.root()).c_str());
|
||||
LOG_DBG("Generation prompt: '%s'\n", session.generation_prompt().c_str());
|
||||
}
|
||||
|
||||
for (const auto & rule : tc.expect_rules) {
|
||||
if (!parser.arena_.has_rule(rule)) {
|
||||
if (!parser.has_rule(rule)) {
|
||||
LOG_ERR("Missing rule: %s\n", rule.c_str());
|
||||
common_log_flush(common_log_main());
|
||||
throw std::runtime_error("Test failed");
|
||||
@@ -1179,12 +1164,14 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
common_chat_msg msg_prev;
|
||||
msg_accum.role = msg_prev.role = "assistant";
|
||||
|
||||
size_t fed = 0;
|
||||
for (size_t i = 1; i <= tc.input.size(); ++i) {
|
||||
auto is_partial = i < tc.input.size() || tc.is_partial;
|
||||
// Use UTF-8 safe truncation to avoid corrupting multi-byte characters
|
||||
size_t safe_len = utf8_truncate_safe_len(std::string_view(tc.input).substr(0, i));
|
||||
std::string prefix = tc.input.substr(0, safe_len);
|
||||
common_chat_msg msg_current = parser.parse(prefix, is_partial);
|
||||
common_chat_input chunk(tc.input.substr(fed, safe_len - fed));
|
||||
fed = safe_len;
|
||||
const common_chat_msg & msg_current = is_partial ? session.feed(chunk) : session.finish(chunk);
|
||||
|
||||
for (const auto & diff : common_chat_msg_diff::compute_diffs(msg_prev, msg_current)) {
|
||||
if (!diff.reasoning_content_delta.empty()) {
|
||||
@@ -1222,30 +1209,33 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
}
|
||||
|
||||
if (!tc.is_partial) {
|
||||
assert_msg_equals(tc.expect, parser.parse(tc.input, false), true);
|
||||
if (tc.input.empty()) {
|
||||
session.finish();
|
||||
}
|
||||
assert_msg_equals(tc.expect, session.msg(), true);
|
||||
}
|
||||
assert_msg_equals(tc.expect, msg_accum, true);
|
||||
|
||||
// A response format must be enforced by an eager grammar
|
||||
if (!tc.params.json_schema.empty()) {
|
||||
if (parser.params_.grammar.empty()) {
|
||||
if (session.grammar().empty()) {
|
||||
throw std::runtime_error("json_schema is set but no grammar was produced");
|
||||
}
|
||||
if (parser.params_.grammar_lazy) {
|
||||
if (sampling.grammar_lazy) {
|
||||
throw std::runtime_error("json_schema is set but the grammar is lazy");
|
||||
}
|
||||
}
|
||||
|
||||
// Test grammar if present in params
|
||||
if (!parser.params_.grammar.empty()) {
|
||||
auto grammar = build_grammar(parser.params_.grammar);
|
||||
if (!session.grammar().empty()) {
|
||||
auto grammar = build_grammar(session.grammar());
|
||||
if (!grammar) {
|
||||
throw std::runtime_error("Failed to build grammar: " + parser.params_.grammar);
|
||||
throw std::runtime_error("Failed to build grammar: " + session.grammar());
|
||||
}
|
||||
|
||||
// In production, grammar triggers match against the full generated text
|
||||
// including the generation prompt. All positions are in full_input coordinates.
|
||||
const auto & gen_prompt = parser.params_.generation_prompt;
|
||||
const auto & gen_prompt = sampling.generation_prompt;
|
||||
std::string full_input = gen_prompt + tc.input;
|
||||
|
||||
// Determine whether the reasoning-budget sampler path applies: tool-call grammar
|
||||
@@ -1253,9 +1243,9 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
// budget sampler inhibits grammar application while inside thinking blocks —
|
||||
// triggers inside <think>...</think> are suppressed.
|
||||
bool use_reasoning_budget_path = false;
|
||||
if (parser.params_.grammar_lazy && !parser.params_.thinking_end_tags.empty()) {
|
||||
if (sampling.grammar_lazy && !session.thinking_end_tags().empty()) {
|
||||
use_reasoning_budget_path = true;
|
||||
for (const auto & trigger : parser.params_.grammar_triggers) {
|
||||
for (const auto & trigger : sampling.grammar_triggers) {
|
||||
if (trigger.type != COMMON_GRAMMAR_TRIGGER_TYPE_WORD) {
|
||||
use_reasoning_budget_path = false;
|
||||
break;
|
||||
@@ -1270,8 +1260,8 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
// Reasoning-budget path: simulate thinking-aware trigger detection.
|
||||
// Walk through full_input tracking thinking state; only match triggers
|
||||
// when outside thinking blocks.
|
||||
const auto & think_start = parser.params_.thinking_start_tag;
|
||||
const auto & think_ends = parser.params_.thinking_end_tags;
|
||||
const auto & think_start = session.thinking_start_tag();
|
||||
const auto & think_ends = session.thinking_end_tags();
|
||||
|
||||
bool in_thinking = false;
|
||||
for (size_t i = 0; i < full_input.size(); ++i) {
|
||||
@@ -1292,7 +1282,7 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
continue;
|
||||
}
|
||||
// Outside thinking — check if any trigger word starts here
|
||||
for (const auto & trigger : parser.params_.grammar_triggers) {
|
||||
for (const auto & trigger : sampling.grammar_triggers) {
|
||||
if (full_input.compare(i, trigger.value.size(), trigger.value) == 0) {
|
||||
if (earliest_trigger_pos == std::string::npos || i < earliest_trigger_pos) {
|
||||
earliest_trigger_pos = i;
|
||||
@@ -1314,7 +1304,7 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
|
||||
if (!use_reasoning_budget_path) {
|
||||
// Legacy path: find triggers without thinking-awareness
|
||||
for (const auto & trigger : parser.params_.grammar_triggers) {
|
||||
for (const auto & trigger : sampling.grammar_triggers) {
|
||||
size_t pos = std::string::npos;
|
||||
std::smatch match;
|
||||
switch (trigger.type) {
|
||||
@@ -1370,16 +1360,16 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
// If the test expects tool calls and the grammar is lazy, the trigger must fire.
|
||||
// Otherwise the grammar would never activate in production and tool calls wouldn't
|
||||
// be constrained. A silent skip here would hide broken triggers.
|
||||
if (parser.params_.grammar_lazy && !tc.expect.tool_calls.empty() && !tc.is_partial
|
||||
if (sampling.grammar_lazy && !tc.expect.tool_calls.empty() && !tc.is_partial
|
||||
&& earliest_trigger_pos == std::string::npos) {
|
||||
std::string trigger_desc;
|
||||
for (const auto & trigger : parser.params_.grammar_triggers) {
|
||||
for (const auto & trigger : sampling.grammar_triggers) {
|
||||
trigger_desc += "\n [type=" + std::to_string(trigger.type) + "] " + trigger.value;
|
||||
}
|
||||
throw std::runtime_error(
|
||||
"Grammar trigger did not fire, but test expects tool calls (lazy grammar).\n"
|
||||
">>> Input: " + full_input + "\n"
|
||||
">>> Triggers (" + std::to_string(parser.params_.grammar_triggers.size()) + "):" + trigger_desc);
|
||||
">>> Triggers (" + std::to_string(sampling.grammar_triggers.size()) + "):" + trigger_desc);
|
||||
}
|
||||
|
||||
// Determine the constrained portion of input to test against grammar.
|
||||
@@ -1392,7 +1382,7 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
auto constrain_from = std::max(earliest_trigger_pos, gen_prompt.size());
|
||||
constrained = full_input.substr(constrain_from);
|
||||
grammar_triggered = true;
|
||||
} else if (!parser.params_.grammar_lazy) {
|
||||
} else if (!sampling.grammar_lazy) {
|
||||
// For non-lazy grammars, the entire input should match
|
||||
grammar_triggered = true;
|
||||
}
|
||||
@@ -1411,7 +1401,7 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
std::to_string(result.matched_codepoints) + " codepoints): " +
|
||||
(result.matched_prefix.size() > 100 ? result.matched_prefix.substr(0, 100) + "..." : result.matched_prefix) +
|
||||
"\n\n>>> Expected next: " + result.expected_description +
|
||||
"\n\n>>> Grammar: " + parser.params_.grammar;
|
||||
"\n\n>>> Grammar: " + session.grammar();
|
||||
} else {
|
||||
error_msg =
|
||||
"Grammar match failed:\n\n"
|
||||
@@ -1422,7 +1412,7 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
(result.matched_prefix.size() > 100 ? result.matched_prefix.substr(0, 100) + "..." : result.matched_prefix) +
|
||||
"\n\n>>> Failing character: " + result.failing_char +
|
||||
"\n\n>>> Expected: " + result.expected_description +
|
||||
"\n\n>>> Grammar: " + parser.params_.grammar;
|
||||
"\n\n>>> Grammar: " + session.grammar();
|
||||
}
|
||||
throw std::runtime_error(error_msg);
|
||||
}
|
||||
@@ -1437,7 +1427,7 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
// Start from tc.expect but copy tool call arguments from the actual parser
|
||||
// output, which preserves original JSON formatting (e.g. {"arg1":1} vs {"arg1": 1}).
|
||||
auto reconstruction_msg = tc.expect;
|
||||
auto parsed_msg = parser.parse(tc.input, false);
|
||||
const auto & parsed_msg = session.msg();
|
||||
for (size_t i = 0; i < reconstruction_msg.tool_calls.size() && i < parsed_msg.tool_calls.size(); i++) {
|
||||
reconstruction_msg.tool_calls[i].arguments = parsed_msg.tool_calls[i].arguments;
|
||||
}
|
||||
@@ -1446,7 +1436,7 @@ static void test_peg_parser(common_chat_templates * tmpls,
|
||||
reconstruction_inputs.add_generation_prompt = false;
|
||||
|
||||
auto reconstruction_params = common_chat_templates_apply(tmpls, reconstruction_inputs);
|
||||
std::string expected_text = parser.params_.prompt + tc.input;
|
||||
std::string expected_text = session.prompt() + tc.input;
|
||||
bool match = reconstruction_params.prompt == expected_text ||
|
||||
(reconstruction_params.prompt.size() > expected_text.size() &&
|
||||
reconstruction_params.prompt.compare(0, expected_text.size(), expected_text) == 0);
|
||||
@@ -4644,8 +4634,7 @@ static void test_template_output_peg_parsers(bool detailed_debug) {
|
||||
inputs.messages = { msg };
|
||||
|
||||
auto params = common_chat_templates_apply(tmpls.get(), inputs);
|
||||
common_peg_arena arena;
|
||||
arena.load(params.parser);
|
||||
const common_peg_arena & arena = params.parser;
|
||||
common_chat_parser_params pp(params);
|
||||
|
||||
// generation_prompt is non-empty for thinking models, so result.end
|
||||
@@ -7791,6 +7780,77 @@ static void test_deepseek_v4_tool_result_ordering() {
|
||||
}
|
||||
}
|
||||
|
||||
static void test_chat_session() {
|
||||
LOG_DBG("%s\n", __func__);
|
||||
auto tmpls = read_templates("models/templates/Qwen3.5-4B.jinja");
|
||||
|
||||
common_chat_templates_inputs inputs;
|
||||
inputs.messages = { message_user };
|
||||
inputs.tools = { special_function_tool };
|
||||
inputs.reasoning_format = COMMON_REASONING_FORMAT_AUTO;
|
||||
|
||||
const std::string output =
|
||||
"I'm\nthinking\n</think>\n\n"
|
||||
"<tool_call>\n"
|
||||
"<function=special_function>\n"
|
||||
"<parameter=arg1>\n1\n</parameter>\n"
|
||||
"</function>\n"
|
||||
"</tool_call>";
|
||||
|
||||
// the session renders the prompt and parses the output fed to it in small chunks
|
||||
{
|
||||
common_chat_session session(tmpls.get(), nullptr, inputs);
|
||||
|
||||
assert_contains(session.prompt(), "\"name\": \"special_function\"");
|
||||
assert_contains(session.prompt(), "<|im_start|>user\nHey there!<|im_end|>\n<|im_start|>assistant\n<think>\n");
|
||||
assert_equals(false, session.grammar().empty());
|
||||
assert_equals(std::string("<|im_start|>assistant\n<think>\n"), session.generation_prompt());
|
||||
|
||||
const std::string thinking = "I'm\nthinking\n</think>\n\n";
|
||||
for (size_t i = 0; i < thinking.size(); i += 3) {
|
||||
session.feed(common_chat_input(thinking.substr(i, 3)));
|
||||
}
|
||||
assert_msg_equals(simple_assist_msg("", "I'm\nthinking\n"), session.msg());
|
||||
|
||||
for (size_t i = thinking.size(); i < output.size(); i += 3) {
|
||||
session.feed(common_chat_input(output.substr(i, 3)));
|
||||
}
|
||||
assert_msg_equals(simple_assist_msg("", "I'm\nthinking\n", "special_function", "{\"arg1\":1}"), session.finish());
|
||||
}
|
||||
|
||||
// a copy does not see what is fed to the original
|
||||
{
|
||||
common_chat_session a(tmpls.get(), nullptr, inputs);
|
||||
common_chat_session b = a;
|
||||
a.feed(common_chat_input(output));
|
||||
assert_equals(true, b.msg().empty());
|
||||
b.feed(common_chat_input("I'm\nthinking\n</think>\n\nHello"));
|
||||
assert_equals(std::string("Hello"), b.msg().content);
|
||||
assert_equals(std::string("special_function"), a.finish().tool_calls.at(0).name);
|
||||
}
|
||||
|
||||
// a continued message starts from the prefill, unless it is echoed
|
||||
{
|
||||
common_chat_templates_inputs cont;
|
||||
cont.messages = { message_user, message_assist_prefill_content };
|
||||
cont.add_generation_prompt = false;
|
||||
cont.continue_final_message = COMMON_CHAT_CONTINUATION_CONTENT;
|
||||
cont.reasoning_format = COMMON_REASONING_FORMAT_AUTO;
|
||||
|
||||
common_chat_session session(tmpls.get(), nullptr, cont);
|
||||
assert_equals(std::string("<|im_start|>assistant\n<think>\nI'm thinking\n</think>\n\nHello, "),
|
||||
session.generation_prompt());
|
||||
assert_msg_equals(simple_assist_msg("Hello, ", "I'm thinking\n"), session.msg());
|
||||
session.feed(common_chat_input("world!"));
|
||||
assert_equals(std::string("Hello, world!"), session.msg().content);
|
||||
|
||||
common_chat_session_params echo;
|
||||
echo.echo = true;
|
||||
common_chat_session echoed(tmpls.get(), nullptr, cont, echo);
|
||||
assert_equals(true, echoed.msg().empty());
|
||||
}
|
||||
}
|
||||
|
||||
static void test_reasoning_budget_tokens_per_request() {
|
||||
LOG_DBG("%s\n", __func__);
|
||||
// Use Qwen3 template which has <think>...</think> reasoning markers.
|
||||
@@ -7811,7 +7871,8 @@ static void test_reasoning_budget_tokens_per_request() {
|
||||
{"reasoning_budget_tokens", 0},
|
||||
};
|
||||
std::vector<raw_buffer> out_files;
|
||||
auto llama_params = oaicompat_chat_params_parse(body, opt, out_files);
|
||||
common_chat_session out_session;
|
||||
auto llama_params = oaicompat_chat_params_parse(nullptr, body, opt, out_files, out_session);
|
||||
|
||||
// The per-request value must win over the server default (-1).
|
||||
if (!llama_params.contains("reasoning_budget_tokens")) {
|
||||
@@ -7844,7 +7905,8 @@ static void test_reasoning_budget_message_per_request() {
|
||||
{"reasoning_budget_message", per_request_message},
|
||||
};
|
||||
std::vector<raw_buffer> out_files;
|
||||
auto llama_params = oaicompat_chat_params_parse(body, opt, out_files);
|
||||
common_chat_session out_session;
|
||||
auto llama_params = oaicompat_chat_params_parse(nullptr, body, opt, out_files, out_session);
|
||||
|
||||
// The per-request value must win over the server default.
|
||||
if (!llama_params.contains("reasoning_budget_message")) {
|
||||
@@ -8035,6 +8097,7 @@ int main(int argc, char ** argv) {
|
||||
test_deepseek_v4_tool_result_ordering();
|
||||
test_template_generation_prompt();
|
||||
test_reasoning_effort_caps();
|
||||
test_chat_session();
|
||||
test_reasoning_budget_tokens_per_request();
|
||||
test_reasoning_budget_message_per_request();
|
||||
test_template_output_peg_parsers(detailed_debug);
|
||||
|
||||
@@ -19,7 +19,6 @@ int main(int argc, char *argv[]) {
|
||||
t.test("unicode", test_unicode);
|
||||
t.test("json", test_json_parser);
|
||||
t.test("gbnf", test_gbnf_generation);
|
||||
t.test("serialization", test_json_serialization);
|
||||
t.test("python-dict", test_python_dict_parser);
|
||||
|
||||
return t.summary();
|
||||
|
||||
@@ -35,7 +35,7 @@ options:
|
||||
--progress print test progress indicators
|
||||
--no-warmup skip warmup runs before benchmarking
|
||||
-fitt, --fit-target <MiB> fit model to device memory with this margin per device in MiB (default: off)
|
||||
-fitc, --fit-ctx <n> minimum ctx size for --fit-target (default: 4096)
|
||||
-fitc, --fit-ctx <n> minimum ctx size for --fit-target (default: 0)
|
||||
-rpc, --rpc <rpc_servers> register RPC devices (comma separated)
|
||||
|
||||
test parameters:
|
||||
|
||||
@@ -441,7 +441,7 @@ static void print_usage(int /* argc */, char ** argv) {
|
||||
printf(" --progress print test progress indicators\n");
|
||||
printf(" --no-warmup skip warmup runs before benchmarking\n");
|
||||
printf(" -fitt, --fit-target <MiB> fit model to device memory with this margin per device in MiB (default: off)\n");
|
||||
printf(" -fitc, --fit-ctx <n> minimum ctx size for --fit-target (default: 4096)\n");
|
||||
printf(" -fitc, --fit-ctx <n> minimum ctx size for --fit-target (default: 0)\n");
|
||||
if (llama_supports_rpc()) {
|
||||
printf(" -rpc, --rpc <rpc_servers> register RPC devices (comma separated)\n");
|
||||
}
|
||||
@@ -2348,7 +2348,8 @@ int llama_bench(int argc, char ** argv) {
|
||||
|
||||
std::vector<size_t> margins(llama_max_devices(), inst.fit_target * 1024 * 1024);
|
||||
|
||||
uint32_t n_ctx_needed = inst.n_prompt + inst.n_gen + inst.n_depth;
|
||||
// fit at least the requested minimum context size, not just the tokens the benchmark processes
|
||||
uint32_t n_ctx_needed = std::max<uint32_t>(inst.n_prompt + inst.n_gen + inst.n_depth, inst.fit_min_ctx);
|
||||
cparams.n_ctx = std::max(cparams.n_ctx, n_ctx_needed);
|
||||
|
||||
common_fit_params(inst.model.c_str(), &mparams, &cparams,
|
||||
|
||||
@@ -1353,10 +1353,6 @@ The `response_format` parameter supports both plain JSON output (e.g. `{"type":
|
||||
|
||||
`reasoning_control`: Arms realtime reasoning control for this completion so it can be ended early via `/v1/chat/completions/control`. Defaults to `false`.
|
||||
|
||||
`generation_prompt`: The generation prompt that was prefilled in by the template. Prepended to model output before parsing.
|
||||
|
||||
`parse_tool_calls`: Whether to parse the generated tool call.
|
||||
|
||||
`parallel_tool_calls` : Whether to enable parallel/multiple tool calls (only supported on some models, verification is based on jinja template).
|
||||
|
||||
For multimodal input (typed content, `messages[i].content[j]`):
|
||||
|
||||
@@ -1266,9 +1266,11 @@ server_tokens tokenize_oai_content_array(const llama_vocab * vocab, mtmd_context
|
||||
|
||||
// used by /chat/completions endpoint
|
||||
json oaicompat_chat_params_parse(
|
||||
const llama_vocab * vocab,
|
||||
json & body, /* openai api json semantics */
|
||||
const server_chat_params & opt,
|
||||
std::vector<raw_buffer> & out_files)
|
||||
std::vector<raw_buffer> & out_files,
|
||||
common_chat_session & out_session)
|
||||
{
|
||||
json llama_params;
|
||||
|
||||
@@ -1393,7 +1395,6 @@ json oaicompat_chat_params_parse(
|
||||
if (body.contains("grammar")) {
|
||||
throw std::invalid_argument("Cannot use custom grammar constraints with tools.");
|
||||
}
|
||||
llama_params["parse_tool_calls"] = true;
|
||||
}
|
||||
|
||||
// merge the template args provided from command line with the args provided in the user request
|
||||
@@ -1427,31 +1428,11 @@ json oaicompat_chat_params_parse(
|
||||
inputs.force_pure_content = opt.force_pure_content;
|
||||
|
||||
// Apply chat template to the list of messages
|
||||
auto chat_params = common_chat_templates_apply(opt.tmpls.get(), inputs);
|
||||
common_chat_session_params session_params;
|
||||
session_params.echo = json_value(body, "echo", false);
|
||||
out_session = common_chat_session(opt.tmpls.get(), vocab, inputs, session_params);
|
||||
|
||||
llama_params["chat_format"] = static_cast<int>(chat_params.format);
|
||||
llama_params["prompt"] = chat_params.prompt;
|
||||
if (!chat_params.grammar.empty()) {
|
||||
llama_params["grammar"] = chat_params.grammar;
|
||||
llama_params["grammar_type"] = std::string("tool_calls");
|
||||
}
|
||||
llama_params["grammar_lazy"] = chat_params.grammar_lazy;
|
||||
auto grammar_triggers = json::array();
|
||||
for (const auto & trigger : chat_params.grammar_triggers) {
|
||||
server_grammar_trigger ct(trigger);
|
||||
grammar_triggers.push_back(ct.to_json());
|
||||
}
|
||||
llama_params["grammar_triggers"] = grammar_triggers;
|
||||
llama_params["preserved_tokens"] = chat_params.preserved_tokens;
|
||||
llama_params["generation_prompt"] = chat_params.generation_prompt;
|
||||
for (const auto & stop : chat_params.additional_stops) {
|
||||
llama_params["stop"].push_back(stop);
|
||||
}
|
||||
if (!chat_params.parser.empty()) {
|
||||
llama_params["chat_parser"] = chat_params.parser;
|
||||
}
|
||||
|
||||
llama_params["message_delimiters"] = chat_params.message_delimiters.to_json();
|
||||
llama_params["prompt"] = out_session.prompt();
|
||||
|
||||
// Reasoning budget: pass parameters through to sampling layer
|
||||
{
|
||||
@@ -1461,10 +1442,10 @@ json oaicompat_chat_params_parse(
|
||||
reasoning_budget = opt.reasoning_budget;
|
||||
}
|
||||
|
||||
if (!chat_params.thinking_end_tags.empty()) {
|
||||
if (!out_session.thinking_end_tags().empty()) {
|
||||
llama_params["reasoning_budget_tokens"] = reasoning_budget;
|
||||
llama_params["reasoning_budget_start_tag"] = chat_params.thinking_start_tag;
|
||||
llama_params["reasoning_budget_end_tags"] = chat_params.thinking_end_tags;
|
||||
llama_params["reasoning_budget_start_tag"] = out_session.thinking_start_tag();
|
||||
llama_params["reasoning_budget_end_tags"] = out_session.thinking_end_tags();
|
||||
llama_params["reasoning_budget_message"] = json_value(body, "reasoning_budget_message", opt.reasoning_budget_message);
|
||||
llama_params["reasoning_control"] = json_value(body, "reasoning_control", false);
|
||||
}
|
||||
@@ -1491,6 +1472,15 @@ json oaicompat_chat_params_parse(
|
||||
}
|
||||
}
|
||||
|
||||
// the session owns these, the server applies them with server_task::apply_chat_session()
|
||||
for (const char * key : { "grammar_lazy", "grammar_triggers", "preserved_tokens" }) {
|
||||
llama_params.erase(key);
|
||||
}
|
||||
if (!out_session.grammar().empty()) {
|
||||
llama_params.erase("grammar");
|
||||
llama_params.erase("json_schema");
|
||||
}
|
||||
|
||||
return llama_params;
|
||||
}
|
||||
|
||||
|
||||
@@ -355,9 +355,11 @@ json oaicompat_completion_params_parse(const json & body);
|
||||
|
||||
// used by /chat/completions endpoint
|
||||
json oaicompat_chat_params_parse(
|
||||
const llama_vocab * vocab,
|
||||
json & body, /* openai api json semantics */
|
||||
const server_chat_params & opt,
|
||||
std::vector<raw_buffer> & out_files);
|
||||
std::vector<raw_buffer> & out_files,
|
||||
common_chat_session & out_session);
|
||||
|
||||
// used by /embeddings endpoint, content has the same format as a chat message content array
|
||||
server_tokens tokenize_oai_content_array(
|
||||
|
||||
@@ -4767,7 +4767,8 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl(
|
||||
server_task_type type,
|
||||
const json & data,
|
||||
const std::vector<raw_buffer> & files,
|
||||
task_response_type res_type) {
|
||||
task_response_type res_type,
|
||||
const common_chat_session & chat_session) {
|
||||
GGML_ASSERT(type == SERVER_TASK_TYPE_COMPLETION || type == SERVER_TASK_TYPE_INFILL);
|
||||
|
||||
auto res = create_response();
|
||||
@@ -4809,11 +4810,6 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl(
|
||||
|
||||
// tasks.reserve(inputs.size()); // TODO: this is inaccurate due to child tasks
|
||||
|
||||
// message delimiters for checkpointing
|
||||
json delims = json_value(data, "message_delimiters", json::array());
|
||||
auto delimiters = common_chat_msg_delimiters_parse(delims);
|
||||
delimiters.tokenize(ctx_server.vocab);
|
||||
|
||||
for (size_t i = 0; i < inputs.size(); i++) {
|
||||
server_task task = server_task(type);
|
||||
|
||||
@@ -4826,7 +4822,7 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl(
|
||||
meta->logit_bias_eog,
|
||||
data);
|
||||
|
||||
task.params.message_spans = task.tokens.find_message_spans(delimiters);
|
||||
task.apply_chat_session(chat_session);
|
||||
|
||||
task.id_slot = json_value(data, "id_slot", -1);
|
||||
sse_ping_interval = task.params.sse_ping_interval;
|
||||
@@ -4847,7 +4843,7 @@ std::unique_ptr<server_res_generator> server_routes::handle_completions_impl(
|
||||
tasks.push_back(std::move(task));
|
||||
}
|
||||
|
||||
rd.post_tasks(std::move(tasks));
|
||||
rd.post_tasks(std::move(tasks), chat_session);
|
||||
} catch (const std::exception & e) {
|
||||
res->error(format_error_response(e.what(), ERROR_TYPE_INVALID_REQUEST));
|
||||
return res;
|
||||
@@ -5444,16 +5440,20 @@ void server_routes::init_routes() {
|
||||
auto res = create_response();
|
||||
std::vector<raw_buffer> files;
|
||||
json body = json::parse(req.body);
|
||||
common_chat_session session;
|
||||
json body_parsed = oaicompat_chat_params_parse(
|
||||
ctx_server.vocab,
|
||||
body,
|
||||
meta->chat_params,
|
||||
files);
|
||||
files,
|
||||
session);
|
||||
return handle_completions_impl(
|
||||
req,
|
||||
SERVER_TASK_TYPE_COMPLETION,
|
||||
body_parsed,
|
||||
files,
|
||||
TASK_RESPONSE_TYPE_OAI_CHAT);
|
||||
TASK_RESPONSE_TYPE_OAI_CHAT,
|
||||
session);
|
||||
};
|
||||
|
||||
this->post_chat_completions_tok = [this](const server_http_req & req) {
|
||||
@@ -5503,16 +5503,20 @@ void server_routes::init_routes() {
|
||||
json body = server_chat_convert_responses_to_chatcmpl(json::parse(req.body));
|
||||
SRV_DBG("%s\n", "Request converted: OpenAI Responses -> OpenAI Chat Completions");
|
||||
SRV_DBG("converted request: %s\n", body.dump().c_str());
|
||||
common_chat_session session;
|
||||
json body_parsed = oaicompat_chat_params_parse(
|
||||
ctx_server.vocab,
|
||||
body,
|
||||
meta->chat_params,
|
||||
files);
|
||||
files,
|
||||
session);
|
||||
return handle_completions_impl(
|
||||
req,
|
||||
SERVER_TASK_TYPE_COMPLETION,
|
||||
body_parsed,
|
||||
files,
|
||||
TASK_RESPONSE_TYPE_OAI_RESP);
|
||||
TASK_RESPONSE_TYPE_OAI_RESP,
|
||||
session);
|
||||
};
|
||||
|
||||
this->post_responses_tok_oai = [this](const server_http_req & req) {
|
||||
@@ -5535,16 +5539,20 @@ void server_routes::init_routes() {
|
||||
files);
|
||||
SRV_DBG("%s\n", "Request converted: OpenAI Transcriptions -> OpenAI Chat Completions");
|
||||
SRV_DBG("converted request: %s\n", body.dump().c_str());
|
||||
common_chat_session session;
|
||||
json body_parsed = oaicompat_chat_params_parse(
|
||||
ctx_server.vocab,
|
||||
body,
|
||||
meta->chat_params,
|
||||
files);
|
||||
files,
|
||||
session);
|
||||
return handle_completions_impl(
|
||||
req,
|
||||
SERVER_TASK_TYPE_COMPLETION,
|
||||
body_parsed,
|
||||
files,
|
||||
TASK_RESPONSE_TYPE_OAI_ASR);
|
||||
TASK_RESPONSE_TYPE_OAI_ASR,
|
||||
session);
|
||||
};
|
||||
|
||||
this->post_anthropic_messages = [this](const server_http_req & req) {
|
||||
@@ -5553,16 +5561,20 @@ void server_routes::init_routes() {
|
||||
json body = server_chat_convert_anthropic_to_oai(json::parse(req.body));
|
||||
SRV_DBG("%s\n", "Request converted: Anthropic -> OpenAI Chat Completions");
|
||||
SRV_DBG("converted request: %s\n", body.dump().c_str());
|
||||
common_chat_session session;
|
||||
json body_parsed = oaicompat_chat_params_parse(
|
||||
ctx_server.vocab,
|
||||
body,
|
||||
meta->chat_params,
|
||||
files);
|
||||
files,
|
||||
session);
|
||||
return handle_completions_impl(
|
||||
req,
|
||||
SERVER_TASK_TYPE_COMPLETION,
|
||||
body_parsed,
|
||||
files,
|
||||
TASK_RESPONSE_TYPE_ANTHROPIC);
|
||||
TASK_RESPONSE_TYPE_ANTHROPIC,
|
||||
session);
|
||||
};
|
||||
|
||||
this->post_anthropic_count_tokens = [this](const server_http_req & req) {
|
||||
@@ -5573,11 +5585,14 @@ void server_routes::init_routes() {
|
||||
this->post_apply_template = [this](const server_http_req & req) {
|
||||
auto res = create_response();
|
||||
std::vector<raw_buffer> files; // dummy, unused
|
||||
common_chat_session session; // dummy, unused
|
||||
json body = json::parse(req.body);
|
||||
json data = oaicompat_chat_params_parse(
|
||||
ctx_server.vocab,
|
||||
body,
|
||||
meta->chat_params,
|
||||
files);
|
||||
files,
|
||||
session);
|
||||
res->ok({{ "prompt", std::move(data.at("prompt")) }});
|
||||
return res;
|
||||
};
|
||||
@@ -6129,10 +6144,13 @@ std::unique_ptr<server_res_generator> server_routes::handle_count_tokens(const s
|
||||
return res;
|
||||
}
|
||||
|
||||
common_chat_session session; // dummy, unused
|
||||
json body_parsed = oaicompat_chat_params_parse(
|
||||
ctx_server.vocab,
|
||||
body,
|
||||
meta->chat_params,
|
||||
files);
|
||||
files,
|
||||
session);
|
||||
json prompt = body_parsed.at("prompt");
|
||||
// SRV_DBG("prompt = %s\n", prompt.dump().c_str());
|
||||
|
||||
|
||||
@@ -166,7 +166,8 @@ private:
|
||||
server_task_type type,
|
||||
const json & data,
|
||||
const std::vector<raw_buffer> & files,
|
||||
task_response_type res_type);
|
||||
task_response_type res_type,
|
||||
const common_chat_session & chat_session = {});
|
||||
std::unique_ptr<server_res_generator> handle_slots_save(const server_http_req & req, int id_slot);
|
||||
std::unique_ptr<server_res_generator> handle_slots_restore(const server_http_req & req, int id_slot);
|
||||
std::unique_ptr<server_res_generator> handle_slots_erase(const server_http_req &, int id_slot);
|
||||
|
||||
@@ -540,6 +540,7 @@ void server_model_meta::update_args(common_preset_context & ctx_preset, std::str
|
||||
void server_model_meta::update_caps(const common_params & base) {
|
||||
// reset to the default so a failed refresh cannot keep old values
|
||||
architecture = server_model_architecture_json(false, false, false, {"text"});
|
||||
n_ctx_train = 0;
|
||||
|
||||
// resolve the model file offline; do not download
|
||||
common_params params;
|
||||
@@ -564,10 +565,12 @@ void server_model_meta::update_caps(const common_params & base) {
|
||||
return;
|
||||
}
|
||||
|
||||
// read the output modalities from the GGUF metadata
|
||||
// read the output modalities and the trained context from the GGUF metadata
|
||||
std::vector<std::string> output_modalities = {"text"};
|
||||
if (!params.model.path.empty()) {
|
||||
output_modalities = server_model_output_modalities(common_get_decision_type(params.model.path));
|
||||
const common_gguf_info info = common_get_gguf_info(params.model.path);
|
||||
output_modalities = server_model_output_modalities(info.decision_type);
|
||||
n_ctx_train = info.n_ctx_train;
|
||||
}
|
||||
|
||||
bool inp_image = false;
|
||||
@@ -2122,9 +2125,13 @@ void server_models_routes::init_routes() {
|
||||
{"source", server_model_source_to_string(meta.source)},
|
||||
{"can_remove", meta.source == SERVER_MODEL_SOURCE_CACHE},
|
||||
// {"need_download", meta.need_download},
|
||||
// TODO: add other fields, may require reading GGUF metadata
|
||||
// TODO: add other fields from the GGUF metadata
|
||||
};
|
||||
|
||||
if (meta.n_ctx_train > 0) {
|
||||
model_info["context_length"] = meta.n_ctx_train;
|
||||
}
|
||||
|
||||
// merge with loaded_info from the child process if available
|
||||
if (meta.is_running()) {
|
||||
for (auto it = meta.loaded_info.begin(); it != meta.loaded_info.end(); ++it) {
|
||||
|
||||
@@ -86,6 +86,7 @@ struct server_model_meta {
|
||||
int stop_timeout = 0; // seconds to wait before force-killing the model instance during shutdown
|
||||
bool hidden = false; // hidden from GET /models, but still accept if requested
|
||||
json architecture = server_model_architecture_json(false, false, false, {"text"});
|
||||
uint32_t n_ctx_train = 0; // trained context, read from the GGUF metadata; 0 when unknown
|
||||
|
||||
bool is_ready() const {
|
||||
return status == SERVER_MODEL_STATUS_LOADED;
|
||||
|
||||
@@ -517,23 +517,23 @@ void server_response_reader::post_task(server_task && task, bool front) {
|
||||
GGML_ASSERT(!task.is_parent() && "not supported, use post_tasks() instead");
|
||||
task.index = 0;
|
||||
id_tasks.insert(task.id);
|
||||
states.push_back(task.create_state());
|
||||
states.push_back(task_result_state());
|
||||
queue_results.add_waiting_task_id(task.id);
|
||||
queue_tasks.post(std::move(task), front);
|
||||
}
|
||||
|
||||
void server_response_reader::post_tasks(std::vector<server_task> && tasks, bool front) {
|
||||
void server_response_reader::post_tasks(std::vector<server_task> && tasks, const common_chat_session & session, bool front) {
|
||||
GGML_ASSERT(id_tasks.empty() && "post_tasks() can only be called once per reader");
|
||||
id_tasks = server_task::get_list_id(tasks);
|
||||
states.reserve(tasks.size());
|
||||
size_t index = 0;
|
||||
for (auto & task : tasks) {
|
||||
task.index = index++;
|
||||
states.push_back(task.create_state());
|
||||
states.push_back(task_result_state(session));
|
||||
// for child tasks
|
||||
for (auto & child_task : task.child_tasks) {
|
||||
child_task.index = index++;
|
||||
states.push_back(child_task.create_state());
|
||||
states.push_back(task_result_state(session));
|
||||
}
|
||||
}
|
||||
GGML_ASSERT(states.size() == id_tasks.size());
|
||||
|
||||
@@ -225,7 +225,7 @@ struct server_response_reader {
|
||||
|
||||
// if front = true, the task will be posted to the front of the queue (high priority)
|
||||
void post_task(server_task && task, bool front = false);
|
||||
void post_tasks(std::vector<server_task> && tasks, bool front = false);
|
||||
void post_tasks(std::vector<server_task> && tasks, const common_chat_session & session = {}, bool front = false);
|
||||
bool has_next() const;
|
||||
|
||||
// return nullptr if should_stop() is true before receiving a result
|
||||
|
||||
@@ -291,54 +291,23 @@ std::vector<std::unique_ptr<field>> make_llama_cmpl_schema(const common_params &
|
||||
// Chat parser params
|
||||
//
|
||||
|
||||
// TODO: change this to string field instead
|
||||
add((new field_json("chat_format"))
|
||||
->set_desc("Chat format used internally by the server")
|
||||
->set_handler([&](field_eval_context & ctx, const json & data) {
|
||||
ctx.params.chat_parser_params.format = static_cast<common_chat_format>(data.at("chat_format").get<int>());
|
||||
SRV_TRC("chat format: %s\n", common_chat_format_name(ctx.params.chat_parser_params.format));
|
||||
}));
|
||||
|
||||
add((new field_str("reasoning_format"))
|
||||
->set_desc("Reasoning format for chain-of-thought models")
|
||||
->set_handler([&](field_eval_context & ctx, const json & data) {
|
||||
auto reasoning_format = common_reasoning_format_from_name(data.at("reasoning_format").get<std::string>());
|
||||
ctx.params.chat_parser_params.reasoning_format = reasoning_format;
|
||||
ctx.params.chat_parser_params.reasoning_in_content = ctx.params.stream && (reasoning_format == COMMON_REASONING_FORMAT_DEEPSEEK_LEGACY);
|
||||
}));
|
||||
|
||||
add((new field_str("generation_prompt"))
|
||||
->set_desc("Generation prompt appended to the chat template output")
|
||||
->set_handler([&](field_eval_context & ctx, const json & data) {
|
||||
std::string s = data.at("generation_prompt").get<std::string>();
|
||||
ctx.params.sampling.generation_prompt = s;
|
||||
|
||||
if (ctx.vocab == nullptr) {
|
||||
ctx.params.chat_parser_params.generation_prompt = common_chat_input(s);
|
||||
return;
|
||||
}
|
||||
|
||||
ctx.params.chat_parser_params.generation_prompt = common_chat_input_tokenize(ctx.vocab, s);
|
||||
}));
|
||||
|
||||
add((new field_bool("parse_tool_calls", params.chat_parser_params.parse_tool_calls))
|
||||
->set_desc("Whether to parse tool calls from the generated output"));
|
||||
|
||||
add((new field_str("chat_parser"))
|
||||
->set_desc("Chat parser configuration string")
|
||||
->set_handler([&](field_eval_context & ctx, const json & data) {
|
||||
ctx.params.chat_parser_params.parser.load(data.at("chat_parser").get<std::string>());
|
||||
ctx.params.reasoning_format = common_reasoning_format_from_name(data.at("reasoning_format").get<std::string>());
|
||||
}));
|
||||
|
||||
add((new field_json("continue_final_message"))
|
||||
->set_desc("Whether to continue the final message of the chat template")
|
||||
->set_handler([&](field_eval_context & ctx, const json & data) {
|
||||
auto continuation = common_chat_continuation_parse(data.at("continue_final_message"));
|
||||
ctx.params.chat_parser_params.is_continuation = continuation != COMMON_CHAT_CONTINUATION_NONE;
|
||||
->set_handler([&](field_eval_context &, const json & data) {
|
||||
common_chat_continuation_parse(data.at("continue_final_message"));
|
||||
}));
|
||||
|
||||
add((new field_bool("echo", params.chat_parser_params.echo))
|
||||
->set_desc("Whether to echo the input tokens in the output"));
|
||||
add((new field_json("echo"))
|
||||
->set_desc("Whether to include the continued assistant message in the output")
|
||||
->set_handler([&](field_eval_context &, const json & data) {
|
||||
data.at("echo").get<bool>();
|
||||
}));
|
||||
|
||||
//
|
||||
// Token-level fields (require vocab)
|
||||
@@ -347,44 +316,17 @@ std::vector<std::unique_ptr<field>> make_llama_cmpl_schema(const common_params &
|
||||
add((new field_json("preserved_tokens"))
|
||||
->set_desc("List of token strings that must not be split during tokenization")
|
||||
->set_handler([&](field_eval_context & ctx, const json & data) {
|
||||
GGML_ASSERT(ctx.vocab != nullptr);
|
||||
for (const auto & t : data.at("preserved_tokens")) {
|
||||
auto ids = common_tokenize(ctx.vocab, t.get<std::string>(), false, true);
|
||||
if (ids.size() == 1) {
|
||||
ctx.params.sampling.preserved_tokens.insert(ids[0]);
|
||||
}
|
||||
}
|
||||
common_sampling_add_preserved_tokens(ctx.params.sampling, ctx.vocab, data.at("preserved_tokens").get<std::vector<std::string>>());
|
||||
}));
|
||||
|
||||
add((new field_json("grammar_triggers"))
|
||||
->set_desc("List of strings or patterns that trigger grammar-constrained generation")
|
||||
->set_handler([&](field_eval_context & ctx, const json & data) {
|
||||
GGML_ASSERT(ctx.vocab != nullptr);
|
||||
std::vector<common_grammar_trigger> triggers;
|
||||
for (const auto & t : data.at("grammar_triggers")) {
|
||||
server_grammar_trigger ct(t);
|
||||
if (ct.value.type == COMMON_GRAMMAR_TRIGGER_TYPE_WORD) {
|
||||
const auto & word = ct.value.value;
|
||||
auto ids = common_tokenize(ctx.vocab, word, false, true);
|
||||
if (ids.size() == 1) {
|
||||
auto token = ids[0];
|
||||
if (std::find(ctx.params.sampling.preserved_tokens.begin(), ctx.params.sampling.preserved_tokens.end(), (llama_token) token) == ctx.params.sampling.preserved_tokens.end()) {
|
||||
throw std::runtime_error("Grammar trigger word should be marked as preserved token: " + word);
|
||||
}
|
||||
common_grammar_trigger trigger;
|
||||
trigger.type = COMMON_GRAMMAR_TRIGGER_TYPE_TOKEN;
|
||||
trigger.value = word;
|
||||
trigger.token = token;
|
||||
ctx.params.sampling.grammar_triggers.push_back(std::move(trigger));
|
||||
} else {
|
||||
ctx.params.sampling.grammar_triggers.push_back({COMMON_GRAMMAR_TRIGGER_TYPE_WORD, word});
|
||||
}
|
||||
} else {
|
||||
ctx.params.sampling.grammar_triggers.emplace_back(std::move(ct.value));
|
||||
}
|
||||
}
|
||||
if (ctx.params.sampling.grammar_lazy && ctx.params.sampling.grammar_triggers.empty()) {
|
||||
throw std::runtime_error("Error: no triggers set for lazy grammar!");
|
||||
triggers.push_back(server_grammar_trigger(t).value);
|
||||
}
|
||||
common_sampling_add_grammar_triggers(ctx.params.sampling, ctx.vocab, std::move(triggers));
|
||||
}));
|
||||
|
||||
add((new field_bool("reasoning_control", params.sampling.reasoning_control))
|
||||
@@ -542,7 +484,7 @@ task_params eval_llama_cmpl_schema(
|
||||
// enabling this will output extra debug information in the HTTP responses from the server
|
||||
params.verbose = params_base.verbosity > 9;
|
||||
|
||||
params.chat_parser_params.reasoning_format = params_base.reasoning_format;
|
||||
params.reasoning_format = params_base.reasoning_format;
|
||||
|
||||
// create context and schema
|
||||
field_eval_context ctx(params);
|
||||
@@ -556,13 +498,6 @@ task_params eval_llama_cmpl_schema(
|
||||
f->eval(ctx, data);
|
||||
}
|
||||
|
||||
// post-processing
|
||||
{
|
||||
// if "reasoning_format" is not provided, its handler will not be called, we will need to handle it here
|
||||
auto reasoning_format = params.chat_parser_params.reasoning_format;
|
||||
params.chat_parser_params.reasoning_in_content = params.stream && (reasoning_format == COMMON_REASONING_FORMAT_DEEPSEEK_LEGACY);
|
||||
}
|
||||
|
||||
// debugging
|
||||
{
|
||||
auto budget = params.sampling.reasoning_budget_tokens;
|
||||
|
||||
@@ -73,10 +73,10 @@ json task_params::to_json(bool only_metrics) const {
|
||||
{"stream", stream},
|
||||
{"n_probs", sampling.n_probs},
|
||||
{"min_keep", sampling.min_keep},
|
||||
{"chat_format", common_chat_format_name(chat_parser_params.format)},
|
||||
{"reasoning_format", common_reasoning_format_name(chat_parser_params.reasoning_format)},
|
||||
{"reasoning_in_content", chat_parser_params.reasoning_in_content},
|
||||
{"generation_prompt", chat_parser_params.generation_prompt.text},
|
||||
{"chat_format", common_chat_format_name(chat_format)},
|
||||
{"reasoning_format", common_reasoning_format_name(reasoning_format)},
|
||||
{"reasoning_in_content", stream && reasoning_format == COMMON_REASONING_FORMAT_DEEPSEEK_LEGACY},
|
||||
{"generation_prompt", sampling.generation_prompt},
|
||||
{"samplers", samplers},
|
||||
{"speculative.types", common_speculative_type_name_str(speculative.types)},
|
||||
{"timings_per_token", timings_per_token},
|
||||
@@ -132,10 +132,10 @@ json task_params::to_json(bool only_metrics) const {
|
||||
{"grammar_lazy", sampling.grammar_lazy},
|
||||
{"grammar_triggers", grammar_triggers},
|
||||
{"preserved_tokens", sampling.preserved_tokens},
|
||||
{"chat_format", common_chat_format_name(chat_parser_params.format)},
|
||||
{"reasoning_format", common_reasoning_format_name(chat_parser_params.reasoning_format)},
|
||||
{"reasoning_in_content", chat_parser_params.reasoning_in_content},
|
||||
{"generation_prompt", chat_parser_params.generation_prompt.text},
|
||||
{"chat_format", common_chat_format_name(chat_format)},
|
||||
{"reasoning_format", common_reasoning_format_name(reasoning_format)},
|
||||
{"reasoning_in_content", stream && reasoning_format == COMMON_REASONING_FORMAT_DEEPSEEK_LEGACY},
|
||||
{"generation_prompt", sampling.generation_prompt},
|
||||
{"samplers", samplers},
|
||||
{"speculative.types", common_speculative_type_name_str(speculative.types)},
|
||||
{"timings_per_token", timings_per_token},
|
||||
@@ -148,15 +148,12 @@ json task_params::to_json(bool only_metrics) const {
|
||||
//
|
||||
// task_result_state
|
||||
//
|
||||
task_result_state::task_result_state(const common_chat_parser_params & chat_parser_params)
|
||||
: chat_parser_params(chat_parser_params)
|
||||
task_result_state::task_result_state(common_chat_session session)
|
||||
: chat_session(std::move(session))
|
||||
, chat_msg(chat_session.msg())
|
||||
, oai_resp_id("resp_" + random_string())
|
||||
, oai_resp_reasoning_id("rs_" + random_string())
|
||||
, oai_resp_message_id("msg_" + random_string()) {
|
||||
if (chat_parser_params.is_continuation && !chat_parser_params.echo) {
|
||||
// initialize chat_msg to avoid emitting a delta containing the assistant prefill
|
||||
chat_msg = common_chat_parse(generated_input, true, chat_parser_params);
|
||||
}
|
||||
}
|
||||
|
||||
common_chat_msg task_result_state::update_chat_msg(
|
||||
@@ -164,13 +161,8 @@ common_chat_msg task_result_state::update_chat_msg(
|
||||
bool is_partial,
|
||||
std::vector<common_chat_msg_diff> & diffs,
|
||||
bool filter_tool_calls) {
|
||||
generated_input.append(added);
|
||||
auto msg_prv_copy = chat_msg;
|
||||
//SRV_DBG("Parsing chat message: %s\n", generated_input.text.c_str());
|
||||
auto new_msg = common_chat_parse(
|
||||
generated_input,
|
||||
is_partial,
|
||||
chat_parser_params);
|
||||
auto new_msg = is_partial ? chat_session.feed(added) : chat_session.finish(added);
|
||||
if (!new_msg.empty()) {
|
||||
new_msg.set_tool_call_ids(generated_tool_call_ids, gen_tool_call_id);
|
||||
chat_msg = new_msg;
|
||||
|
||||
+16
-11
@@ -89,8 +89,9 @@ struct task_params {
|
||||
std::string control_action;
|
||||
std::string control_cmpl_id;
|
||||
|
||||
// per-request parameters for chat parsing
|
||||
common_chat_parser_params chat_parser_params;
|
||||
// reported in generation_settings, parsing itself is owned by the chat session
|
||||
common_chat_format chat_format = COMMON_CHAT_FORMAT_CONTENT_ONLY;
|
||||
common_reasoning_format reasoning_format = COMMON_REASONING_FORMAT_NONE;
|
||||
|
||||
// message spans for checkpointing
|
||||
common_chat_msg_spans message_spans;
|
||||
@@ -106,9 +107,8 @@ struct task_params {
|
||||
struct task_result_state {
|
||||
// tracking diffs for partial tool calls
|
||||
std::vector<common_chat_msg_diff> diffs;
|
||||
common_chat_parser_params chat_parser_params;
|
||||
common_chat_session chat_session; // owns all parsing for this generation
|
||||
common_chat_msg chat_msg;
|
||||
common_chat_input generated_input; // append new chunks of generated text here
|
||||
std::vector<std::string> generated_tool_call_ids;
|
||||
std::unordered_set<size_t> sent_tool_call_names;
|
||||
|
||||
@@ -124,7 +124,7 @@ struct task_result_state {
|
||||
const std::string oai_resp_message_id;
|
||||
std::string oai_resp_fc_id; // function call ID for current args delta
|
||||
|
||||
task_result_state(const common_chat_parser_params & chat_parser_params);
|
||||
task_result_state(common_chat_session session = {});
|
||||
|
||||
// parse partial tool calls and update the internal state
|
||||
common_chat_msg update_chat_msg(
|
||||
@@ -258,6 +258,17 @@ struct server_task {
|
||||
return ids;
|
||||
}
|
||||
|
||||
void apply_chat_session(const common_chat_session & session) {
|
||||
if (!session.has_template()) {
|
||||
return;
|
||||
}
|
||||
|
||||
session.apply_sampling(params.sampling);
|
||||
params.chat_format = session.format();
|
||||
params.antiprompt.insert(params.antiprompt.end(), session.additional_stops().begin(), session.additional_stops().end());
|
||||
params.message_spans = tokens.find_message_spans(session.message_delimiters());
|
||||
}
|
||||
|
||||
void add_child(int id_parent, int id_child) {
|
||||
server_task copy;
|
||||
|
||||
@@ -277,12 +288,6 @@ struct server_task {
|
||||
child_tasks.push_back(std::move(copy));
|
||||
}
|
||||
|
||||
// the task will be moved into queue, then onto slots
|
||||
// however, the state must be kept by caller (e.g., HTTP thread)
|
||||
task_result_state create_state() const {
|
||||
return task_result_state(params.chat_parser_params);
|
||||
}
|
||||
|
||||
bool is_parent() const {
|
||||
return child_tasks.size() > 0;
|
||||
}
|
||||
|
||||
@@ -12,37 +12,38 @@
|
||||
}
|
||||
|
||||
let { isFav, option, revealOnHover = true }: Props = $props();
|
||||
|
||||
// the favorite heart stays visible, so the reveal applies to the other icons
|
||||
const revealClass =
|
||||
'pointer-events-none opacity-0 group-hover:pointer-events-auto group-hover:opacity-100 [@media(pointer:coarse)]:pointer-events-auto [@media(pointer:coarse)]:opacity-100';
|
||||
</script>
|
||||
|
||||
<div
|
||||
class={[
|
||||
'flex items-center justify-center gap-1 max-md:gap-2.5',
|
||||
revealOnHover
|
||||
? 'pointer-events-none opacity-0 group-hover:pointer-events-auto group-hover:opacity-100 [@media(pointer:coarse)]:pointer-events-auto [@media(pointer:coarse)]:opacity-100'
|
||||
: ''
|
||||
]}
|
||||
class="flex items-center justify-center gap-1 max-md:gap-2.5"
|
||||
onclick={(event) => event.stopPropagation()}
|
||||
onkeydown={(event) => event.stopPropagation()}
|
||||
role="presentation"
|
||||
>
|
||||
<ActionIcon
|
||||
class="h-5 w-5 hover:text-foreground"
|
||||
icon={Info}
|
||||
iconSize="h-4 w-4"
|
||||
onclick={() =>
|
||||
// a phone has no manager: its information dialog takes the icon
|
||||
deviceStore.isMobile
|
||||
? uiStore.openModelInformation(option)
|
||||
: uiStore.openModelsManager(option.id)}
|
||||
tooltip="Manage model"
|
||||
tooltipAsTitle
|
||||
/>
|
||||
<span class={revealOnHover ? revealClass : ''}>
|
||||
<ActionIcon
|
||||
class="h-5 w-5 hover:text-foreground"
|
||||
icon={Info}
|
||||
iconSize="h-4 w-4"
|
||||
onclick={() =>
|
||||
// a phone has no manager: its information dialog takes the icon
|
||||
deviceStore.isMobile
|
||||
? uiStore.openModelInformation(option)
|
||||
: uiStore.openModelsManager(option.id)}
|
||||
tooltip="Manage model"
|
||||
tooltipAsTitle
|
||||
/>
|
||||
</span>
|
||||
|
||||
{#if isFav}
|
||||
<span class="flex h-5 w-5 items-center justify-center">
|
||||
<span class="flex group-hover:hidden [@media(pointer:coarse)]:hidden">
|
||||
<ActionIcon
|
||||
class="h-5 w-5 text-rose-500 hover:text-foreground"
|
||||
class="h-5 w-5 text-rose-500"
|
||||
icon={Heart}
|
||||
iconSize="h-4 w-4"
|
||||
onclick={() => modelsStore.toggleFavorite(option.model)}
|
||||
@@ -63,13 +64,15 @@
|
||||
</span>
|
||||
</span>
|
||||
{:else}
|
||||
<ActionIcon
|
||||
class="h-5 w-5 hover:text-foreground"
|
||||
icon={Heart}
|
||||
iconSize="h-4 w-4"
|
||||
onclick={() => modelsStore.toggleFavorite(option.model)}
|
||||
tooltip="Add to favorites"
|
||||
tooltipAsTitle
|
||||
/>
|
||||
<span class={revealOnHover ? revealClass : ''}>
|
||||
<ActionIcon
|
||||
class="h-5 w-5 hover:text-foreground"
|
||||
icon={Heart}
|
||||
iconSize="h-4 w-4"
|
||||
onclick={() => modelsStore.toggleFavorite(option.model)}
|
||||
tooltip="Add to favorites"
|
||||
tooltipAsTitle
|
||||
/>
|
||||
</span>
|
||||
{/if}
|
||||
</div>
|
||||
|
||||
@@ -16,7 +16,7 @@
|
||||
import type { ModelOption } from '$lib/types/models';
|
||||
import { filterModelOptions } from '$lib/utils';
|
||||
import { type Snippet, untrack } from 'svelte';
|
||||
import { SvelteMap, SvelteSet } from 'svelte/reactivity';
|
||||
import { SvelteMap } from 'svelte/reactivity';
|
||||
|
||||
interface Props {
|
||||
class?: string;
|
||||
@@ -170,16 +170,29 @@
|
||||
if (remaining.length > 0) rest.push({ ...entry, base: remaining[0], quants: remaining });
|
||||
}
|
||||
|
||||
const claimed = new SvelteSet<string>();
|
||||
const favorites = rest.filter((entry) =>
|
||||
entry.quants.some((q) => modelsStore.favoriteModelIds.has(q.model))
|
||||
);
|
||||
// a favorite quant stands on its own, listed flat like a loaded one, and the
|
||||
// quants left behind stay with their repo in the local block
|
||||
const favorites: ModelQuantGroup[] = [];
|
||||
const localRest: ModelQuantGroup[] = [];
|
||||
|
||||
for (const entry of favorites) claimed.add(entry.key);
|
||||
for (const entry of rest) {
|
||||
for (const quant of entry.quants) {
|
||||
if (modelsStore.favoriteModelIds.has(quant.model)) {
|
||||
favorites.push({ ...entry, base: quant, key: quant.id, quants: [quant] });
|
||||
}
|
||||
}
|
||||
|
||||
const { hidden, local } = splitHiddenQuants(
|
||||
rest.filter((entry) => !claimed.has(entry.key)),
|
||||
(option) => modelsStore.isHidden(option.id)
|
||||
const remaining = entry.quants.filter(
|
||||
(quant) => !modelsStore.favoriteModelIds.has(quant.model)
|
||||
);
|
||||
|
||||
if (remaining.length > 0) {
|
||||
localRest.push({ ...entry, base: remaining[0], quants: remaining });
|
||||
}
|
||||
}
|
||||
|
||||
const { hidden, local } = splitHiddenQuants(localRest, (option) =>
|
||||
modelsStore.isHidden(option.id)
|
||||
);
|
||||
const ordered: ModelsTableGroup[] = [];
|
||||
// loaded models lead the table, then favorites, then the local block
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
import ModelsManagerStatusCell from './ModelsManagerStatusCell.svelte';
|
||||
import { modelRowActions } from './row-actions';
|
||||
import { configuredContext, downloadProgressFor } from './utils';
|
||||
import { MoreHorizontal } from '@lucide/svelte';
|
||||
import { Heart, MoreHorizontal } from '@lucide/svelte';
|
||||
import { DropdownMenuActions } from '$lib/components/app';
|
||||
import { MODEL_ROW_GRID_CLASS, MODEL_ROW_TRAILING_CELL_CLASS } from '$lib/constants';
|
||||
import { ModelRowDownloadState } from '$lib/enums';
|
||||
@@ -74,6 +74,13 @@
|
||||
|
||||
<!-- a phone has no width for the modality icons, the id needs it more -->
|
||||
<ModelCapabilities hideModalities={deviceStore.isMobile} {option} />
|
||||
|
||||
{#if favorite}
|
||||
<!-- the heart is decorative, the button's text carries the state -->
|
||||
<Heart aria-hidden="true" class="h-3.5 w-3.5 shrink-0 text-rose-500" />
|
||||
|
||||
<span class="sr-only">favorited</span>
|
||||
{/if}
|
||||
</span>
|
||||
</button>
|
||||
|
||||
|
||||
+2
-9
@@ -7,8 +7,7 @@
|
||||
hasActiveFilters,
|
||||
modelContextLength,
|
||||
type ModelQuantGroup,
|
||||
type ModelsTableGroup,
|
||||
statusRank
|
||||
type ModelsTableGroup
|
||||
} from './utils';
|
||||
import {
|
||||
ArrowDown,
|
||||
@@ -155,10 +154,6 @@
|
||||
return (modelContextLength(a) ?? 0) - (modelContextLength(b) ?? 0);
|
||||
case ModelsTableSortKey.NAME:
|
||||
return a.model.localeCompare(b.model);
|
||||
case ModelsTableSortKey.STATUS:
|
||||
// a running model leads, then one that is being worked on (loading,
|
||||
// sleeping), then the rest; the reported status sorts the row's own cell
|
||||
return statusRank(b) - statusRank(a);
|
||||
default:
|
||||
return 0;
|
||||
}
|
||||
@@ -357,9 +352,7 @@
|
||||
{@render sortHeader(ModelsTableSortKey.CONTEXT, 'Context')}
|
||||
</span>
|
||||
|
||||
<span class="justify-self-center max-md:hidden">
|
||||
{@render sortHeader(ModelsTableSortKey.STATUS, 'Status')}
|
||||
</span>
|
||||
<span class="justify-self-center max-md:hidden">Status</span>
|
||||
|
||||
<span class="text-center max-md:hidden">Actions</span>
|
||||
</div>
|
||||
|
||||
@@ -1,10 +1,5 @@
|
||||
import { LOCAL_BACKEND_ID, type ModalityKey } from '$lib/constants';
|
||||
import {
|
||||
ModelCapability,
|
||||
ModelGroupKind,
|
||||
ModelsTableGroupKind,
|
||||
ServerModelStatus
|
||||
} from '$lib/enums';
|
||||
import { ModelCapability, ModelGroupKind, ModelsTableGroupKind } from '$lib/enums';
|
||||
import { HuggingFaceService, ModelsService } from '$lib/services';
|
||||
import { modelsStore } from '$lib/stores';
|
||||
import type { ModelDownloadEntry, ModelDownloadProgress, ModelOption } from '$lib/types/models';
|
||||
@@ -27,20 +22,6 @@ export function hasActiveFilters(
|
||||
return contextLimit > 0 || modalities.length > 0 || capabilities.length > 0;
|
||||
}
|
||||
|
||||
/**
|
||||
* Order a status sorts behind: loaded first, then a model being worked on
|
||||
* (loading, sleeping), then the rest.
|
||||
*/
|
||||
export function statusRank(option: ModelOption): number {
|
||||
const status = modelsStore.getModelStatus(option.model);
|
||||
|
||||
if (modelsStore.isModelRunning(option.model)) return 2;
|
||||
|
||||
if (status === ServerModelStatus.LOADING || status === ServerModelStatus.SLEEPING) return 1;
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
/** Byte counts of a tracked download: live while it runs, frozen while paused. */
|
||||
export function downloadProgressFor(repoWithTag: string): ModelDownloadProgress | null {
|
||||
return (
|
||||
|
||||
@@ -177,7 +177,7 @@
|
||||
|
||||
<DropdownMenu.Content
|
||||
align="end"
|
||||
class="w-full md:min-w-80 md:w-112 max-w-[calc(100vw-2rem)] p-0! max-h-[min(40rem,calc(var(--bits-dropdown-menu-content-available-height)-1rem))]"
|
||||
class="w-full md:min-w-80 md:w-md max-w-[calc(100vw-2rem)] p-0! max-h-[min(40rem,calc(var(--bits-dropdown-menu-content-available-height)-1rem))]"
|
||||
onOpenAutoFocus={(event) => event.preventDefault()}
|
||||
>
|
||||
<DropdownMenuSearchable
|
||||
|
||||
@@ -17,7 +17,7 @@
|
||||
<DropdownMenuPrimitive.Content
|
||||
bind:ref
|
||||
class={cn(
|
||||
'z-50 max-h-[calc(var(--bits-dropdown-menu-content-available-height)-1rem)] min-w-[8rem] origin-(--bits-dropdown-menu-content-transform-origin) overflow-x-hidden overflow-y-auto rounded-md border border-border bg-popover p-1.5 text-popover-foreground shadow-md outline-none data-[side=bottom]:slide-in-from-top-2 data-[side=left]:slide-in-from-right-2 data-[side=right]:slide-in-from-left-2 data-[side=top]:slide-in-from-bottom-2 data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=closed]:fill-mode-forwards data-[state=closed]:zoom-out-95 data-[state=open]:animate-in data-[state=open]:fade-in-0 data-[state=open]:zoom-in-95 dark:border-border/20',
|
||||
'z-50 max-h-[calc(var(--bits-dropdown-menu-content-available-height)-1rem)] min-w-[8rem] origin-(--bits-dropdown-menu-content-transform-origin) overflow-x-hidden overflow-y-auto rounded-xl border border-border bg-popover p-1.5 text-popover-foreground shadow-md outline-none data-[side=bottom]:slide-in-from-top-2 data-[side=left]:slide-in-from-right-2 data-[side=right]:slide-in-from-left-2 data-[side=top]:slide-in-from-bottom-2 data-[state=closed]:animate-out data-[state=closed]:fade-out-0 data-[state=closed]:fill-mode-forwards data-[state=closed]:zoom-out-95 data-[state=open]:animate-in data-[state=open]:fade-in-0 data-[state=open]:zoom-in-95 dark:border-border/20',
|
||||
className
|
||||
)}
|
||||
data-slot="dropdown-menu-content"
|
||||
|
||||
@@ -17,7 +17,7 @@
|
||||
<DropdownMenuPrimitive.Item
|
||||
bind:ref
|
||||
class={cn(
|
||||
"relative flex cursor-pointer items-center gap-2 rounded-sm px-2 py-1.5 text-sm outline-hidden select-none data-highlighted:bg-accent data-highlighted:text-accent-foreground data-[disabled]:pointer-events-none data-[disabled]:opacity-50 data-[inset]:pl-8 data-[variant=destructive]:text-destructive data-[variant=destructive]:data-highlighted:bg-destructive/10 data-[variant=destructive]:data-highlighted:text-destructive dark:data-[variant=destructive]:data-highlighted:bg-destructive/20 [&_svg]:pointer-events-none [&_svg]:shrink-0 [&_svg:not([class*='size-'])]:size-4 [&_svg:not([class*='text-'])]:text-muted-foreground data-[variant=destructive]:*:[svg]:!text-destructive",
|
||||
"relative flex cursor-pointer items-center gap-2 rounded-md px-2 py-1.5 text-sm outline-hidden select-none data-highlighted:bg-accent data-highlighted:text-accent-foreground data-[disabled]:pointer-events-none data-[disabled]:opacity-50 data-[inset]:pl-8 data-[variant=destructive]:text-destructive data-[variant=destructive]:data-highlighted:bg-destructive/10 data-[variant=destructive]:data-highlighted:text-destructive dark:data-[variant=destructive]:data-highlighted:bg-destructive/20 [&_svg]:pointer-events-none [&_svg]:shrink-0 [&_svg:not([class*='size-'])]:size-4 [&_svg:not([class*='text-'])]:text-muted-foreground data-[variant=destructive]:*:[svg]:!text-destructive",
|
||||
className
|
||||
)}
|
||||
data-inset={inset}
|
||||
|
||||
@@ -129,7 +129,7 @@ export const SETTINGS_REGISTRY: SettingsSectionEntry[] = [
|
||||
// Deliberately off for now: the natural place to turn it on is the first
|
||||
// run experience, once onboarding exists to ask the user about it.
|
||||
defaultValue: false,
|
||||
help: 'Fetch model metadata (avatars, context length, chat template, file sizes) from the Hugging Face Hub. When off, the UI only shows what the server reports for /v1/models and hides the org avatars.',
|
||||
help: 'Fetch model metadata (avatars, context length, chat template) from the Hugging Face Hub. When off, the UI only shows what the server reports for /v1/models and hides the org avatars.',
|
||||
key: SETTINGS_KEYS.USE_HUGGING_FACE_HUB,
|
||||
label: 'Use Hugging Face Hub API for models metadata',
|
||||
type: SettingsFieldType.CHECKBOX
|
||||
|
||||
@@ -80,6 +80,5 @@ export enum ModelRowDownloadState {
|
||||
/** Column the models manager table can be ordered by. */
|
||||
export enum ModelsTableSortKey {
|
||||
CONTEXT = 'context',
|
||||
NAME = 'name',
|
||||
STATUS = 'status'
|
||||
NAME = 'name'
|
||||
}
|
||||
|
||||
@@ -622,6 +622,13 @@ class ModelsStore implements ModelPropsHost, ModelStatusHost {
|
||||
capabilities: rawCapabilities.filter((value: unknown): value is string =>
|
||||
Boolean(value)
|
||||
),
|
||||
// 0 is the server's way of leaving the trained context unknown; a
|
||||
// model-mode listing reports it as meta.n_ctx_train instead
|
||||
contextLength:
|
||||
item.context_length ||
|
||||
(typeof item.meta?.n_ctx_train === 'number' && item.meta.n_ctx_train > 0
|
||||
? item.meta.n_ctx_train
|
||||
: undefined),
|
||||
description: details?.description,
|
||||
details: details?.details,
|
||||
draftSidecars: mergedDraftSidecars(
|
||||
|
||||
Vendored
+2
@@ -100,6 +100,8 @@ export interface ApiModelDataEntry {
|
||||
tags?: string[];
|
||||
/** Modality capabilities, reported by the router for every model regardless of load state */
|
||||
architecture?: ApiModelArchitecture;
|
||||
/** Trained context of the model, read from its GGUF metadata at registration */
|
||||
context_length?: number;
|
||||
/** Legacy meta field (may be present in older responses) */
|
||||
meta?: Record<string, unknown> | null;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,243 @@
|
||||
// Guards the manager table ordering and filtering: the context column sorts and
|
||||
// filters once a model's context is known, from the option or the Hub record.
|
||||
|
||||
import ModelsManagerWrapper from './components/ModelsManagerWrapper.svelte';
|
||||
import { SETTINGS_KEYS } from '$lib/constants';
|
||||
import { ServerModelStatus } from '$lib/enums';
|
||||
import { HuggingFaceService } from '$lib/services';
|
||||
import { modelsStore, settingsStore } from '$lib/stores';
|
||||
import type { ApiModelDataEntry } from '$lib/types';
|
||||
import type { ModelOption } from '$lib/types/models';
|
||||
import { SvelteMap } from 'svelte/reactivity';
|
||||
import { beforeEach, expect, it, vi } from 'vitest';
|
||||
import { render } from 'vitest-browser-svelte';
|
||||
|
||||
function option(model: string, contextLength?: number): ModelOption {
|
||||
return {
|
||||
capabilities: [],
|
||||
contextLength,
|
||||
id: model,
|
||||
model,
|
||||
name: model
|
||||
};
|
||||
}
|
||||
|
||||
const models = [
|
||||
option('org/alpha-8b:Q4_K_M', 8192),
|
||||
option('org/beta-8b:Q4_K_M', 131072),
|
||||
option('org/gamma-8b:Q4_K_M', 32768)
|
||||
];
|
||||
// the router listing carries no context, so a row starts without one
|
||||
const modelsWithoutContext = models.map((model) => ({ ...model, contextLength: undefined }));
|
||||
|
||||
beforeEach(() => {
|
||||
modelsStore.routerModels = [];
|
||||
modelsStore.models = modelsWithoutContext;
|
||||
modelsStore.favoriteModelIds = new Set();
|
||||
settingsStore.config[SETTINGS_KEYS.GROUP_MODELS_BY_FAMILY] = false;
|
||||
});
|
||||
|
||||
/** Renders the manager and waits for the test models to show. */
|
||||
async function renderWithModels(rows: ModelOption[] = modelsWithoutContext) {
|
||||
const screen = render(ModelsManagerWrapper);
|
||||
|
||||
modelsStore.models = rows;
|
||||
|
||||
await expect.element(screen.getByText(/gamma\s+8B/)).toBeVisible();
|
||||
|
||||
return screen;
|
||||
}
|
||||
|
||||
function rowNames(container: HTMLElement): string[] {
|
||||
return [...container.querySelectorAll('button[aria-pressed]')].map(
|
||||
(row) => row.textContent ?? ''
|
||||
);
|
||||
}
|
||||
|
||||
/** The Hub details cache the rows read their context from. */
|
||||
function detailsCache() {
|
||||
return (
|
||||
HuggingFaceService as unknown as {
|
||||
detailsCache: SvelteMap<string, { gguf?: { context_length?: number } } | null>;
|
||||
}
|
||||
).detailsCache;
|
||||
}
|
||||
|
||||
function warmCache() {
|
||||
const cache = detailsCache();
|
||||
|
||||
cache.set('org/alpha-8b', { gguf: { context_length: 8192 } });
|
||||
cache.set('org/beta-8b', { gguf: { context_length: 131072 } });
|
||||
cache.set('org/gamma-8b', { gguf: { context_length: 32768 } });
|
||||
}
|
||||
|
||||
it('sorts by context', async () => {
|
||||
const screen = await renderWithModels(models);
|
||||
|
||||
// lowest first
|
||||
await screen.getByTitle('Sort by context, lowest first').click();
|
||||
|
||||
const ascending = rowNames(screen.container);
|
||||
|
||||
await screen.getByTitle('Sort by context, highest first').click();
|
||||
|
||||
const descending = rowNames(screen.container);
|
||||
|
||||
expect(ascending).not.toEqual(descending);
|
||||
});
|
||||
|
||||
it('sorts by context with family grouping on', async () => {
|
||||
settingsStore.config[SETTINGS_KEYS.GROUP_MODELS_BY_FAMILY] = true;
|
||||
|
||||
const screen = await renderWithModels(models);
|
||||
|
||||
// lowest first
|
||||
await screen.getByTitle('Sort by context, lowest first').click();
|
||||
|
||||
const ascending = rowNames(screen.container);
|
||||
|
||||
await screen.getByTitle('Sort by context, highest first').click();
|
||||
|
||||
const descending = rowNames(screen.container);
|
||||
|
||||
expect(ascending).not.toEqual(descending);
|
||||
});
|
||||
|
||||
it('lists the favorited quants of a repo as flat rows', async () => {
|
||||
// one quant of the beta repo is a favorite, the other stays with the repo
|
||||
modelsStore.favoriteModelIds = new Set(['org/alpha-8b:Q4_K_M', 'org/beta-8b:Q4_K_M']);
|
||||
|
||||
const screen = await renderWithModels([
|
||||
modelsWithoutContext[0],
|
||||
modelsWithoutContext[1],
|
||||
option('org/beta-8b:Q8_0'),
|
||||
modelsWithoutContext[2]
|
||||
]);
|
||||
|
||||
// the favorite quant is a model row of its own, not a repo with subitems
|
||||
expect(screen.container.textContent).not.toContain('2 quants available');
|
||||
|
||||
// the repo appears once in favorites for its favorited quant and once in the
|
||||
// local block for the quant left behind
|
||||
expect(screen.getByText(/beta\s+8B/).elements().length).toBe(2);
|
||||
});
|
||||
|
||||
/** A router listing entry carrying only the meta context. */
|
||||
function entry(model: string, nCtxTrain: number): ApiModelDataEntry {
|
||||
return {
|
||||
created: 0,
|
||||
id: model,
|
||||
in_cache: false,
|
||||
meta: { n_ctx_train: nCtxTrain },
|
||||
object: 'model',
|
||||
owned_by: 'llamacpp',
|
||||
path: `/models/${model}`,
|
||||
status: { value: ServerModelStatus.UNLOADED }
|
||||
};
|
||||
}
|
||||
|
||||
it('sorts by the meta context of a listing without the router field', async () => {
|
||||
// a listing that skips the router's GGUF read reports the trained context
|
||||
// only as meta.n_ctx_train, so the option mapping falls back to it
|
||||
vi.spyOn(globalThis, 'fetch').mockImplementation(async (input: RequestInfo | URL) => {
|
||||
const url = typeof input === 'string' ? input : input instanceof URL ? input.href : input.url;
|
||||
|
||||
if (url.includes('/props')) {
|
||||
return new Response(
|
||||
JSON.stringify({
|
||||
default_generation_settings: { n_ctx: 0, params: {} },
|
||||
model_alias: 'llama-server',
|
||||
model_path: 'none',
|
||||
role: 'router'
|
||||
}),
|
||||
{ headers: { 'Content-Type': 'application/json' }, status: 200 }
|
||||
);
|
||||
}
|
||||
|
||||
if (url.includes('/server')) {
|
||||
return new Response(
|
||||
JSON.stringify({ git_branch: 'test', git_commit: 'test', mode: 'router', version: 'test' }),
|
||||
{ headers: { 'Content-Type': 'application/json' }, status: 200 }
|
||||
);
|
||||
}
|
||||
|
||||
if (/\/v1\/models|\/models\b/.test(url)) {
|
||||
return new Response(
|
||||
JSON.stringify({
|
||||
data: [
|
||||
entry('org/alpha-8b:Q4_K_M', 8192),
|
||||
entry('org/beta-8b:Q4_K_M', 131072),
|
||||
entry('org/gamma-8b:Q4_K_M', 32768)
|
||||
],
|
||||
object: 'list'
|
||||
}),
|
||||
{ headers: { 'Content-Type': 'application/json' }, status: 200 }
|
||||
);
|
||||
}
|
||||
|
||||
throw new Error(`unexpected fetch in the test: ${url}`);
|
||||
});
|
||||
|
||||
const screen = render(ModelsManagerWrapper);
|
||||
|
||||
// the fetch maps the listing into options, the meta context fills in
|
||||
await modelsStore.fetch(true);
|
||||
|
||||
await expect.element(screen.getByText(/gamma\s+8B/)).toBeVisible();
|
||||
|
||||
await screen.getByTitle('Sort by context, lowest first').click();
|
||||
|
||||
const names = rowNames(screen.container).join(' | ');
|
||||
|
||||
expect(names.indexOf('alpha')).toBeLessThan(names.indexOf('gamma'));
|
||||
expect(names.indexOf('gamma')).toBeLessThan(names.indexOf('beta'));
|
||||
});
|
||||
|
||||
it('re-sorts when the Hub details arrive after the sort was clicked', async () => {
|
||||
const screen = await renderWithModels();
|
||||
|
||||
// the user sorts while the contexts are still unknown
|
||||
await screen.getByTitle('Sort by context, lowest first').click();
|
||||
|
||||
// then the rows fetch their Hub records
|
||||
warmCache();
|
||||
|
||||
// the table re-sorts once the cache answers
|
||||
await vi.waitFor(() => {
|
||||
const names = rowNames(screen.container).join(' | ');
|
||||
|
||||
expect(names.indexOf('alpha')).toBeLessThan(names.indexOf('gamma'));
|
||||
expect(names.indexOf('gamma')).toBeLessThan(names.indexOf('beta'));
|
||||
|
||||
return names;
|
||||
});
|
||||
});
|
||||
|
||||
it('filters by search', async () => {
|
||||
const screen = await renderWithModels();
|
||||
|
||||
await screen.getByPlaceholder('Search your models').fill('beta');
|
||||
|
||||
await expect.element(screen.getByText(/alpha\s+8B/)).not.toBeVisible();
|
||||
await expect.element(screen.getByText(/beta\s+8B/)).toBeVisible();
|
||||
});
|
||||
|
||||
it('filters by context and sorts from the Hub details cache', async () => {
|
||||
warmCache();
|
||||
|
||||
const screen = await renderWithModels();
|
||||
|
||||
// open the context filter and ask for 32K or more
|
||||
await screen.getByText('Context:').click();
|
||||
await screen.getByText('32K or more').click();
|
||||
|
||||
await expect.element(screen.getByText(/alpha\s+8B/)).not.toBeVisible();
|
||||
await expect.element(screen.getByText(/gamma\s+8B/)).toBeVisible();
|
||||
|
||||
// sorting re-runs once the cache answers
|
||||
await screen.getByTitle('Sort by context, lowest first').click();
|
||||
|
||||
const names = rowNames(screen.container).join(' | ');
|
||||
|
||||
expect(names.indexOf('gamma')).toBeLessThan(names.indexOf('beta'));
|
||||
});
|
||||
+1184
File diff suppressed because it is too large
Load Diff
Vendored
+913
-72
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user