Compare commits

...
15 Commits
Author SHA1 Message Date
Xuan Son Nguyen f0440d9efc vendor: apply deep nested json patch from upstream 2026-10-09 23:35:12 +02:00
Dante 79e2e74eb1 CUDA: fix round issue, under MSVC the CPU and GPU agree (#30229) 2026-10-09 19:59:50 +02:00
Georgi Gerganov 8e2d31e0eb graph : reorder get_rows for embeddings (#30160)
* graph : reorder get_rows for embeddings

* cont : fix gemma4 and improve input embedding construction logic

* cont : add TODO for lora

* cont : fix raw embeddings path

* gemma4 : avoid ple cast in embeddings path
2026-10-09 20:43:05 +03:00
Aleksander GrygierandPascal baef3ed9a1 ui: Models Manager Follow-up Improvements (#30228)
* common : read a GGUF's trained context from its metadata

common_get_gguf_n_ctx_train opens only the file's metadata (no_alloc,
like common_get_decision_type) and reads <arch>.context_length, so a
caller can learn the trained context without loading the model. It
accepts both u32 and u64 values and returns 0 when the file is missing,
unreadable, invalid, or reports no context length.

Assisted-by: pi:zai-org/GLM-5.3-Flash

* server : report the trained context in the models listing

update_caps already resolves the model file offline to read its
modalities, so it now reads the trained context from the same GGUF
metadata, and GET /models reports it as context_length when it is
known. A router listing then carries the context without any Hub
request, which lets the UI sort and filter by it offline.

Assisted-by: pi:zai-org/GLM-5.3-Flash

* ui : take the trained context from the models listing

The router now reports context_length per model, so the option mapping
fills contextLength from it and the manager reads it before the Hub
record. The Context column, the context sort and the context filter
then work with the Hugging Face Hub API turned off. A browser suite
guards the sort and the search, the Hub-cache driven context filter and
the re-sort when details arrive after the sort was clicked.

Assisted-by: pi:zai-org/GLM-5.3-Flash

* ui : mark favorite models with a heart

A favorited model shows a rose heart in the selector even before its
row is hovered, and the crossed heart takes its place on hover, so
unfavoriting stays one hover away. The manager table marks its
favorited rows with the same heart after the badges and capabilities.

Assisted-by: pi:zai-org/GLM-5.3-Flash

* fix: UI text nit

* fix: UI nits

* fix: Favorite models grouping in models table

* feat: Remove sorting from Status column in Models Table

* server: read the GGUF metadata once per model

Read the decision type and the trained context in a single GGUF open,
accept only a UINT32 context length like the model loader, and reset
n_ctx_train with the other caps so a failed refresh drops it.

* fix: Post-review fixes

---------

Co-authored-by: Pascal <admin@serveurperso.com>
2026-10-09 19:28:54 +02:00
Martin Emrich f39148a953 llama-bench: respect -fitc if bigger than required benchmark size (#28331)
Assisted-By: opencode,llama.cpp,Qwen3.6-35B-A3B,Qwen3.8-27B
2026-10-09 19:14:30 +02:00
Anav Prasad 50e3e3e480 CUDA: Remove redundant CUDA copies after SSM_SCAN (#29807)
* CUDA: fuse copy of updated state snapshots into recurrent cache with ssm_scan

* CUDA: remove redundant cuda copies with K==1 (non spec-dec) scenario as well
2026-10-09 18:23:28 +02:00
jingzhou 64df9183f5 opencl: fix kernel compilation for a6x GPUs (#30176)
* opencl: skip kernel_cpy_f32_f32_pack on A6X to avoid shader compiler crash

* The A6x compiler backend found in iot device with a623 (E031.50.31.01)
  cannot handle kernels with a large number of arguments. Skip this
  kernel for A6x to avoid compiler crash

* opencl: A6X constant-fold workaround for get_local_size in GEMV kernels

* opencl: add Adreno 623 to A6X GPU detection list
2026-10-09 09:19:20 -07:00
bosh 6184e92c57 model : use exact GELU for ModernBERT encoders (#30108)
* model : use exact GELU for ModernBERT encoders

Assisted-by: Codex

* model : keep tanh GELU aliases on ggml_geglu

Assisted-by: Claude Opus 5.5

* model : map gelu_python to ggml_geglu_erf

Assisted-by: Claude Opus 5.5
2026-10-09 23:26:56 +08:00
Aldehir Rojas 8b54361025 chat : refactor API (#30210) 2026-10-09 10:06:22 -05:00
Pascal a518119d30 llama: keep the backend sampling graph static across ubatches (#30223)
The reserve builds n_outputs_max_per_seq sampling chains per sampler,
while a decode built one per output row, so the graph changed its
topology after the reserve and GGML_SCHED_NO_REALLOC builds aborted on
the next same sized graph. Every sampler now builds
n_outputs_max_per_seq chains, the ones without a row of the ubatch on
the padding row and not selected, and graph_max_nodes counts them.
2026-10-09 17:20:25 +03:00
Jasmine-tim 8ae386707b ggml: fix OOB write in ggml_acc with negative offset (#30135)
ggml_acc_impl narrowed a size_t offset to int32_t without checking that it
fits, so a large offset could truncate to a negative int32_t. The forward
then sign-extended it to a huge size_t and the bounds assertion wrapped,
allowing an OOB write below the dst buffer. Check the offset before the
narrowing, matching the existing check in ggml_set_impl.
2026-10-09 16:57:33 +03:00
Georgi Gerganov e60eff95fd meta : handle host views (#30217)
* meta : handle views of tensors allocated on the host

A view shares the memory of its view_src, so ggml-alloc never allocates a view in
the buffer of the split it lands in - the scheduler copies the source into the
split and the ops that use the view read that copy. The view node itself is a noop
and does not need a split of its own, but the meta backend asserted when one was
left inside a meta split:

  - ggml_backend_meta_get_split_state() dereferenced tensor->buffer->context
  - the graph rebuild mapped every node with ggml_backend_meta_buffer_simple_tensor()

Accept such nodes when they are views of host tensors, which also generalizes the
previous s_copy_main workaround. This fixes the assert hit by KV cache views when
using --split-mode tensor with partial offload.

Assisted-by: pi:llama.cpp/Qwen3.8-Flash-Next

* archs : re-enable sm tensor for K2 Horizon

* cont : add TODO and reference
2026-10-09 16:01:04 +03:00
Masashi Yoshimura ba6439a6b5 webgpu: use 2D workgroup dispatch for all the ops which use 1D dispatch (e.g., rms_norm) (#30219)
* remove 1D workgroups dispatching

* formatting
2026-10-09 16:00:37 +03:00
Ruben Ortlam 5e4878e978 vulkan: fix rms_norm workgroup count overflow (#30145) 2026-10-09 15:00:14 +02:00
ynankani 609290be6b convert : support compressed-tensor mixed-precision NVFP4 checkpoint (#28636)
Signed-off-by: ynankani <ynankani@nvidia.com>
2026-10-09 13:54:58 +02:00
99 changed files with 3748 additions and 1206 deletions
+3
View File
@@ -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
+2 -3
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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 {
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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 = {
+3 -6
View File
@@ -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 = {
+2 -4
View File
@@ -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) {
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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 = {
+3 -6
View File
@@ -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 = {
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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 = {
+2 -4
View File
@@ -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,
+2 -4
View File
@@ -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) {
+1 -2
View File
@@ -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;
}
-303
View File
@@ -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));
-6
View File
@@ -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:
+38
View File
@@ -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!");
}
}
+6
View File
@@ -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
View File
@@ -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
+20 -3
View File
@@ -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;
}
+87
View File
@@ -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) {
+29 -14
View File
@@ -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);
}
+10
View File
@@ -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);
+1 -1
View File
@@ -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) {
+19 -6
View File
@@ -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:
+2
View File
@@ -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);
+14 -2
View File
@@ -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() {
+75 -27
View File
@@ -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 ++) {
+6 -3
View File
@@ -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);
+7 -4
View File
@@ -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;
+6 -3
View File
@@ -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;
+9 -6
View File
@@ -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;
+1
View File
@@ -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));
+15
View File
@@ -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([
-1
View File
@@ -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
View File
@@ -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
View File
@@ -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);
+1
View File
@@ -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
View File
@@ -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
View File
@@ -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));
+2 -2
View File
@@ -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
View File
@@ -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),
-1
View File
@@ -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);
}
-1
View File
@@ -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);
+108
View File
@@ -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));
+2 -3
View File
@@ -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
View File
@@ -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);
-1
View File
@@ -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();
+1 -1
View File
@@ -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:
+3 -2
View File
@@ -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,
-4
View File
@@ -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]`):
+19 -29
View File
@@ -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;
}
+3 -1
View File
@@ -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(
+36 -18
View File
@@ -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());
+2 -1
View File
@@ -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);
+10 -3
View File
@@ -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) {
+1
View File
@@ -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;
+4 -4
View File
@@ -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());
+1 -1
View File
@@ -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
+13 -78
View File
@@ -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;
+12 -20
View File
@@ -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
View File
@@ -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>
@@ -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
+1 -2
View File
@@ -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(
+2
View File
@@ -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'));
});
File diff suppressed because it is too large Load Diff
+913 -72
View File
File diff suppressed because it is too large Load Diff