mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-10 23:07:21 -05:00
Compare commits
34
Commits
b11505
...
xsn/json_patch
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f0440d9efc | ||
|
|
79e2e74eb1 | ||
|
|
8e2d31e0eb | ||
|
|
baef3ed9a1 | ||
|
|
f39148a953 | ||
|
|
50e3e3e480 | ||
|
|
64df9183f5 | ||
|
|
6184e92c57 | ||
|
|
8b54361025 | ||
|
|
a518119d30 | ||
|
|
8ae386707b | ||
|
|
e60eff95fd | ||
|
|
ba6439a6b5 | ||
|
|
5e4878e978 | ||
|
|
609290be6b | ||
|
|
86a2835320 | ||
|
|
e94acad853 | ||
|
|
5e1d74043e | ||
|
|
b42b7e6d30 | ||
|
|
b013e56a71 | ||
|
|
2c564f43df | ||
|
|
1f8fa52318 | ||
|
|
8a1a9b5126 | ||
|
|
d4d82d67f4 | ||
|
|
3d65c90d04 | ||
|
|
de7fa0a3c6 | ||
|
|
71ad0590f4 | ||
|
|
a11f57ba93 | ||
|
|
fc9ce6b9d5 | ||
|
|
c35b66744f | ||
|
|
1167d3f42c | ||
|
|
4f92965a7b | ||
|
|
c811cb8f0a | ||
|
|
033df86b69 |
@@ -155,6 +155,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app
|
||||
|
||||
|
||||
@@ -114,6 +114,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app
|
||||
|
||||
|
||||
@@ -123,6 +123,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app
|
||||
|
||||
|
||||
@@ -151,6 +151,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/lib/ /app
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app
|
||||
|
||||
@@ -130,6 +130,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app
|
||||
|
||||
|
||||
@@ -227,6 +227,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app/
|
||||
|
||||
|
||||
@@ -136,6 +136,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app
|
||||
|
||||
|
||||
@@ -133,6 +133,7 @@ ENTRYPOINT [ "/llama.cpp/bin/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
WORKDIR /llama.cpp/bin
|
||||
|
||||
|
||||
@@ -117,6 +117,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app
|
||||
|
||||
|
||||
@@ -107,6 +107,7 @@ ENTRYPOINT [ "/app/llama-cli" ]
|
||||
FROM base AS server
|
||||
|
||||
ENV LLAMA_ARG_HOST=0.0.0.0
|
||||
ENV LLAMA_ARG_PORT=8080
|
||||
|
||||
COPY --from=build /app/full/llama /app/full/llama-server /app
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -9,6 +9,8 @@ on:
|
||||
branches:
|
||||
- master
|
||||
|
||||
run-name: "Publish ${{ github.event.workflow_run.display_title }}"
|
||||
|
||||
cache-mode: none
|
||||
permissions:
|
||||
actions: read
|
||||
|
||||
+2
-2
@@ -1479,7 +1479,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
|
||||
));
|
||||
add_opt(common_arg(
|
||||
{"--server-base"}, "URL",
|
||||
string_format("connect to this server instead of starting a new one, example: 'http://localhost:8080' (default: none)"),
|
||||
string_format("connect to this server instead of starting a new one, example: 'http://localhost:9931' (default: none)"),
|
||||
[](common_params & params, const std::string & value) {
|
||||
params.server_base = value;
|
||||
}
|
||||
@@ -2778,7 +2778,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
|
||||
).set_env("LLAMA_ARG_N_CPU_MOE"));
|
||||
add_opt(common_arg(
|
||||
{"--moe-cache-mib"}, "N",
|
||||
"GPU cache size in MiB for the MoE experts kept in the CPU (default: 0, disabled)",
|
||||
"GPU cache size in MiB for the MoE experts kept in the CPU. with multiple GPUs, it is split among them like the layers (--tensor-split) (default: 0, disabled)",
|
||||
[](common_params & params, int value) {
|
||||
if (value < 0) {
|
||||
throw std::invalid_argument("invalid value");
|
||||
|
||||
@@ -61,8 +61,7 @@ common_chat_params peg_generator::generate_parser(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = autoparser.build_parser(inputs, parser_generation_prompt);
|
||||
data.parser = parser.save();
|
||||
data.parser = autoparser.build_parser(inputs, parser_generation_prompt);
|
||||
|
||||
// Build grammar if tools are present
|
||||
bool has_tools =
|
||||
@@ -78,7 +77,7 @@ common_chat_params peg_generator::generate_parser(const common_chat_template &
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !has_response_format && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
// Set grammar triggers based on tool section markers (fall back to per-call markers)
|
||||
|
||||
+178
-71
@@ -9,6 +9,7 @@
|
||||
#include "json.h"
|
||||
#include "log.h"
|
||||
#include "parsers/parsers.h"
|
||||
#include "sampling.h"
|
||||
|
||||
#include "jinja/value.h"
|
||||
#include "jinja/runtime.h"
|
||||
@@ -112,38 +113,6 @@ const char * common_chat_role_to_string(common_chat_role role) {
|
||||
return "";
|
||||
}
|
||||
|
||||
json common_chat_msg_delimiters::to_json() const {
|
||||
json result = json::array();
|
||||
for (const auto & d : delimiters) {
|
||||
result.push_back({
|
||||
{ "role", common_chat_role_to_string(d.role) },
|
||||
{ "delimiter", d.delimiter },
|
||||
});
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
common_chat_msg_delimiters common_chat_msg_delimiters_parse(const json & delimiters) {
|
||||
common_chat_msg_delimiters result;
|
||||
|
||||
if (!delimiters.is_array()) {
|
||||
return result;
|
||||
}
|
||||
|
||||
result.delimiters.reserve(delimiters.size());
|
||||
for (const auto & d : delimiters) {
|
||||
if (!d.is_object()) {
|
||||
continue;
|
||||
}
|
||||
result.delimiters.push_back({
|
||||
common_chat_role_from_string(d.value("role", std::string())),
|
||||
d.value("delimiter", std::string()),
|
||||
});
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
void common_chat_msg_delimiters::tokenize(const llama_vocab * vocab) {
|
||||
for (auto & d : delimiters) {
|
||||
d.tokens = common_tokenize(vocab, d.delimiter, false, true);
|
||||
@@ -620,8 +589,11 @@ std::vector<common_chat_tool> common_chat_tools_parse_oaicompat(const json & too
|
||||
}
|
||||
|
||||
common_chat_continuation common_chat_continuation_parse(const common_json & value) {
|
||||
if (value.is_boolean() && value.get<bool>()) {
|
||||
return COMMON_CHAT_CONTINUATION_AUTO;
|
||||
if (value.is_null()) {
|
||||
return COMMON_CHAT_CONTINUATION_NONE;
|
||||
}
|
||||
if (value.is_boolean()) {
|
||||
return value.get<bool>() ? COMMON_CHAT_CONTINUATION_AUTO : COMMON_CHAT_CONTINUATION_NONE;
|
||||
}
|
||||
if (value.is_string()) {
|
||||
auto value_str = value.get<std::string>();
|
||||
@@ -632,7 +604,7 @@ common_chat_continuation common_chat_continuation_parse(const common_json & valu
|
||||
return COMMON_CHAT_CONTINUATION_CONTENT;
|
||||
}
|
||||
}
|
||||
return COMMON_CHAT_CONTINUATION_NONE;
|
||||
throw std::invalid_argument("Invalid continue_final_message: expected a boolean, \"content\" or \"reasoning_content\"");
|
||||
}
|
||||
|
||||
bool common_chat_verify_template(const std::string & tmpl, bool use_jinja) {
|
||||
@@ -1087,41 +1059,55 @@ static json common_chat_extra_context() {
|
||||
return ctx;
|
||||
}
|
||||
|
||||
std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
const common_chat_template & tmpl,
|
||||
const std::string & src,
|
||||
autoparser::generation_params & params) {
|
||||
static common_chat_params common_chat_params_init_lfm2_tokens(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
|
||||
return common_chat_params_init_lfm2(tmpl, inputs, /* tool_list_tokens = */ true);
|
||||
}
|
||||
|
||||
static common_chat_params common_chat_params_init_lfm2_5(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
|
||||
return common_chat_params_init_lfm2(tmpl, inputs, /* tool_list_tokens = */ false);
|
||||
}
|
||||
|
||||
// Older gemma4 templates need their tool responses rewritten before rendering
|
||||
static common_chat_params common_chat_params_init_gemma4_legacy(const common_chat_template & tmpl, const autoparser::generation_params & inputs) {
|
||||
auto adjusted = inputs;
|
||||
workaround::convert_tool_responses_gemma4(adjusted.messages);
|
||||
return common_chat_params_init_gemma4(tmpl, adjusted);
|
||||
}
|
||||
|
||||
// Pick the dedicated handler for a template from its source, or null for the autoparser.
|
||||
// Order matters: the first match wins, and later checks assume the earlier ones did not match.
|
||||
static common_chat_params_init_fn common_chat_template_detect_params_init(const std::string & src) {
|
||||
// Ministral/Mistral Large 3 - uses special reasoning structure fixes, can't use autoparser
|
||||
// Note: Mistral Small 3.2 uses [CALL_ID] which Ministral doesn't have, so we can distinguish them
|
||||
if (src.find("[SYSTEM_PROMPT]") != std::string::npos && src.find("[TOOL_CALLS]") != std::string::npos &&
|
||||
src.find("[ARGS]") != std::string::npos && src.find("[CALL_ID]") == std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Ministral/Magistral Large 3\n");
|
||||
return common_chat_params_init_ministral_3(tmpl, params);
|
||||
return common_chat_params_init_ministral_3;
|
||||
}
|
||||
|
||||
// LLM-jp-4.1 - GPT-OSS dialect (spaces after special tokens, <|end|>-separated parallel calls)
|
||||
if (src.find("chat_format=llm-jp-harmony-v1") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: LLM-jp Harmony v1\n");
|
||||
return common_chat_params_init_llm_jp_harmony(tmpl, params);
|
||||
return common_chat_params_init_llm_jp_harmony;
|
||||
}
|
||||
|
||||
// GPT-OSS - has unique channel-based structure that needs dedicated handler
|
||||
if (src.find("<|channel|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: GPT-OSS\n");
|
||||
return common_chat_params_init_gpt_oss(tmpl, params);
|
||||
return common_chat_params_init_gpt_oss;
|
||||
}
|
||||
|
||||
// Muse Glimmer format using " to=<recipient>" recipients and <|eom|>/<|eot|> message terminators.
|
||||
if (src.find("<atem:function_calls>") != std::string::npos && src.find("<|eom|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Muse Glimmer\n");
|
||||
return common_chat_params_init_muse_glimmer(tmpl, params);
|
||||
return common_chat_params_init_muse_glimmer;
|
||||
}
|
||||
|
||||
// Functionary v3.2 - uses recipient-based format with >>>recipient\n{content}
|
||||
// Detection: template has ">>>all" for content and ">>>" prefix for tool calls
|
||||
if (src.find(">>>all") != std::string::npos && src.find(">>>${recipient}") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Functionary v3.2\n");
|
||||
return common_chat_params_init_functionary_v3_2(tmpl, params);
|
||||
return common_chat_params_init_functionary_v3_2;
|
||||
}
|
||||
|
||||
// Kimi K2 Thinking - uses unique tool call ID format: functions.<name>:<index>
|
||||
@@ -1129,14 +1115,14 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
if (src.find("<|tool_calls_section_begin|>") != std::string::npos &&
|
||||
src.find("<|tool_call_begin|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Kimi K2 Thinking\n");
|
||||
return common_chat_params_init_kimi_k2(tmpl, params);
|
||||
return common_chat_params_init_kimi_k2;
|
||||
}
|
||||
|
||||
// Kimi K3 - the <|open|>/<|close|>/<|end_of_msg|> markers are unique to it
|
||||
if (src.find("<|open|>") != std::string::npos && src.find("<|close|>") != std::string::npos &&
|
||||
src.find("<|end_of_msg|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Kimi K3\n");
|
||||
return common_chat_params_init_kimi_k3(tmpl, params);
|
||||
return common_chat_params_init_kimi_k3;
|
||||
}
|
||||
|
||||
// K2 Horizon - <|ifm|im_start|> turns, <ifm|think*> reasoning picked by reasoning_effort and
|
||||
@@ -1144,7 +1130,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
if (src.find("<|ifm|im_start|>") != std::string::npos &&
|
||||
src.find("<ifm|tool_calls>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: K2 Horizon\n");
|
||||
return common_chat_params_init_k2_horizon(tmpl, params);
|
||||
return common_chat_params_init_k2_horizon;
|
||||
}
|
||||
|
||||
// Ling 3.0 / Bailing V3 - <role>X</role> sections with <arg_key>/<arg_value> tagged
|
||||
@@ -1152,7 +1138,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
if (src.find("<role>ASSISTANT</role>") != std::string::npos &&
|
||||
src.find("<arg_key>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Ling 3.0 (Bailing V3)\n");
|
||||
return common_chat_params_init_ling3(tmpl, params);
|
||||
return common_chat_params_init_ling3;
|
||||
}
|
||||
|
||||
// Cohere2 MoE / North Code - marker-wrapped format with <|START_TEXT|> content and
|
||||
@@ -1161,19 +1147,19 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
if (src.find("<|START_TEXT|>") != std::string::npos &&
|
||||
src.find("<|START_ACTION|>") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Cohere2 MoE\n");
|
||||
return common_chat_params_init_cohere2moe(tmpl, params);
|
||||
return common_chat_params_init_cohere2moe;
|
||||
}
|
||||
|
||||
if (is_lfm2_template(src)) {
|
||||
LOG_DBG("Using specialized template: LFM2\n");
|
||||
return common_chat_params_init_lfm2(tmpl, params, /* tool_list_tokens = */ true);
|
||||
return common_chat_params_init_lfm2_tokens;
|
||||
}
|
||||
|
||||
// LFM2.5 format detection: template uses plain "List of tools: [...]" with no special tokens
|
||||
if (src.find("List of tools: [") != std::string::npos &&
|
||||
src.find("<|tool_list_start|>") == std::string::npos) {
|
||||
LOG_DBG("Using specialized template: LFM2.5\n");
|
||||
return common_chat_params_init_lfm2(tmpl, params, /* tool_list_tokens = */ false);
|
||||
return common_chat_params_init_lfm2_5;
|
||||
}
|
||||
|
||||
// GigaChatV3 format detection
|
||||
@@ -1181,7 +1167,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
src.find("<|message_sep|>") != std::string::npos &&
|
||||
src.find("<|function_call|>") == std::string::npos) {
|
||||
LOG_DBG("Using specialized template: GigaChatV3\n");
|
||||
return common_chat_params_init_gigachat_v3(tmpl, params);
|
||||
return common_chat_params_init_gigachat_v3;
|
||||
}
|
||||
|
||||
// MiniMax-M3: the namespace token "]<]minimax[>[" collides with the autoparser's
|
||||
@@ -1190,7 +1176,7 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
src.find("<tool_call>") != std::string::npos &&
|
||||
src.find("<invoke name=") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: MiniMax-M3\n");
|
||||
return common_chat_params_init_minimax_m3(tmpl, params);
|
||||
return common_chat_params_init_minimax_m3;
|
||||
}
|
||||
|
||||
// DeepSeek V3.2/V4 format detection: template defines dsml_token and uses it for tool calls.
|
||||
@@ -1201,18 +1187,18 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
(src.find("function_calls") != std::string::npos ||
|
||||
src.find("tool_calls") != std::string::npos)) {
|
||||
LOG_DBG("Using specialized template: DeepSeek V3.2/V4\n");
|
||||
return common_chat_params_init_deepseek_v3_2(tmpl, params);
|
||||
return common_chat_params_init_deepseek_v3_2;
|
||||
}
|
||||
|
||||
// Gemma4 format detection
|
||||
if (src.find("'<|tool_call>call:'") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Gemma4\n");
|
||||
if (src.find("{#- OpenAI Chat Completions:") == std::string::npos) {
|
||||
// apply workarounds if using the older gemma4 templates
|
||||
LOG_WRN("%s: detected an outdated gemma4 chat template, applying compatibility workarounds. "
|
||||
"Consider updating to the official template.\n", __func__);
|
||||
workaround::convert_tool_responses_gemma4(params.messages);
|
||||
return common_chat_params_init_gemma4_legacy;
|
||||
}
|
||||
return common_chat_params_init_gemma4(tmpl, params);
|
||||
return common_chat_params_init_gemma4;
|
||||
}
|
||||
|
||||
// MiniCPM5 - XML tool calls with <function name="..."><param name="...">...</param></function>
|
||||
@@ -1220,14 +1206,14 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
src.find("<function name=\"") != std::string::npos &&
|
||||
src.find("<param name=\"") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: MiniCPM5\n");
|
||||
return common_chat_params_init_minicpm5(tmpl, params);
|
||||
return common_chat_params_init_minicpm5;
|
||||
}
|
||||
|
||||
// TranslateGemma - user content must follow a custom schema with language codes
|
||||
if (src.find("[source_lang_code]") != std::string::npos &&
|
||||
src.find("[target_lang_code]") != std::string::npos) {
|
||||
LOG_DBG("Using specialized template: TranslateGemma\n");
|
||||
return common_chat_params_init_translate_gemma(tmpl, params);
|
||||
return common_chat_params_init_translate_gemma;
|
||||
}
|
||||
|
||||
// Qwen3-Coder XML tool calls, also used by Nemotron Nano 3, Qwen3.5 and StepFun-3.5-Flash
|
||||
@@ -1237,10 +1223,51 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
// Exclude models that don't use \n between tags
|
||||
src.find("'<tool_call><function=' ~ tool_call.name ~ '>'") == std::string::npos) {
|
||||
LOG_DBG("Using specialized template: Qwen3-Coder\n");
|
||||
return common_chat_params_init_qwen3_coder(tmpl, params);
|
||||
return common_chat_params_init_qwen3_coder;
|
||||
}
|
||||
|
||||
return std::nullopt;
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
common_chat_template::common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token) {
|
||||
jinja::lexer lexer;
|
||||
auto lexer_res = lexer.tokenize(src);
|
||||
this->prog = jinja::parse_from_tokens(lexer_res);
|
||||
|
||||
this->src = lexer_res.source;
|
||||
this->bos_tok = bos_token;
|
||||
this->eos_tok = eos_token;
|
||||
|
||||
this->caps = jinja::caps_get(prog);
|
||||
// LOG_INF("%s: caps:\n%s\n", __func__, this->caps.to_string().c_str());
|
||||
|
||||
this->params_init = common_chat_template_detect_params_init(this->src);
|
||||
if (this->params_init) {
|
||||
return;
|
||||
}
|
||||
|
||||
// The analysis depends only on the template, so run it once here instead of on every apply.
|
||||
// A failure is kept for apply to report, so a bad template still loads like it did before.
|
||||
try {
|
||||
analysis = std::make_unique<autoparser::autoparser>();
|
||||
analysis->analyze_template(*this);
|
||||
} catch (const std::exception & e) {
|
||||
analysis.reset();
|
||||
analysis_error = e.what();
|
||||
}
|
||||
}
|
||||
|
||||
common_chat_template::~common_chat_template() = default;
|
||||
common_chat_template::common_chat_template(common_chat_template &&) = default;
|
||||
common_chat_template & common_chat_template::operator=(common_chat_template &&) = default;
|
||||
|
||||
std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
const common_chat_template & tmpl,
|
||||
const autoparser::generation_params & params) {
|
||||
if (!tmpl.params_init) {
|
||||
return std::nullopt;
|
||||
}
|
||||
return tmpl.params_init(tmpl, params);
|
||||
}
|
||||
|
||||
static common_chat_params common_chat_templates_apply_jinja(const struct common_chat_templates * tmpls,
|
||||
@@ -1342,21 +1369,23 @@ static common_chat_params common_chat_templates_apply_jinja(const struct common_
|
||||
data.prompt = common_chat_template_direct_apply_impl(tmpl, params_copy);
|
||||
data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, params);
|
||||
data.format = COMMON_CHAT_FORMAT_PEG_NATIVE;
|
||||
auto parser = build_chat_peg_parser([&data](common_chat_peg_builder &p) {
|
||||
data.parser = build_chat_peg_parser([&data](common_chat_peg_builder &p) {
|
||||
return p.literal(data.generation_prompt) << p.content(p.rest());
|
||||
});
|
||||
data.parser = parser.save();
|
||||
return data;
|
||||
}
|
||||
|
||||
if (auto result = common_chat_try_specialized_template(tmpl, src, params)) {
|
||||
if (auto result = common_chat_try_specialized_template(tmpl, params)) {
|
||||
return *result;
|
||||
}
|
||||
|
||||
if (!tmpl.analysis) {
|
||||
throw std::invalid_argument("Unable to generate parser for this template. Automatic parser generation failed: " + tmpl.analysis_error);
|
||||
}
|
||||
|
||||
try {
|
||||
LOG_DBG("%s: using differential autoparser\n", __func__);
|
||||
struct autoparser::autoparser autoparser;
|
||||
autoparser.analyze_template(tmpl);
|
||||
const auto & autoparser = *tmpl.analysis;
|
||||
auto auto_params = autoparser::peg_generator::generate_parser(tmpl, params, autoparser);
|
||||
|
||||
common_chat_msg_delimiters delimiters;
|
||||
@@ -1377,8 +1406,7 @@ static common_chat_params common_chat_templates_apply_jinja(const struct common_
|
||||
auto_params.thinking_end_tags = {std::move(end_tag)};
|
||||
}
|
||||
}
|
||||
common_peg_arena arena;
|
||||
arena.load(auto_params.parser);
|
||||
const auto & arena = auto_params.parser;
|
||||
LOG_DBG("%s: generated parser:\n%s\n\nparser generation prompt: %s\n", __func__, arena.dump(arena.root()).c_str(), auto_params.generation_prompt.c_str());
|
||||
return auto_params;
|
||||
} catch (const std::exception & e) {
|
||||
@@ -1525,9 +1553,10 @@ common_chat_msg common_chat_peg_parse(const common_peg_arena & src_pars
|
||||
const common_chat_input & input,
|
||||
bool is_partial,
|
||||
const common_chat_parser_params & params) {
|
||||
const common_peg_arena & parser = src_parser.empty() ?
|
||||
build_chat_peg_parser([](common_chat_peg_builder & p) { return p.content(p.rest()) + p.end(); }) :
|
||||
src_parser;
|
||||
// both branches must be lvalues, a temporary here would copy the arena on every call
|
||||
static const common_peg_arena content_only =
|
||||
build_chat_peg_parser([](common_chat_peg_builder & p) { return p.content(p.rest()) + p.end(); });
|
||||
const common_peg_arena & parser = src_parser.empty() ? content_only : src_parser;
|
||||
|
||||
if (src_parser.empty()) {
|
||||
LOG_DBG("No parser definition detected, assuming pure content parser.");
|
||||
@@ -1598,6 +1627,84 @@ common_chat_msg common_chat_peg_parse(const common_peg_arena & src_pars
|
||||
return msg;
|
||||
}
|
||||
|
||||
common_chat_session::common_chat_session(const common_chat_templates * tmpls,
|
||||
const llama_vocab * vocab,
|
||||
const common_chat_templates_inputs & inputs,
|
||||
const common_chat_session_params & params) {
|
||||
auto applied = common_chat_templates_apply(tmpls, inputs);
|
||||
|
||||
templated = true;
|
||||
prompt_text = std::move(applied.prompt);
|
||||
result.role = "assistant";
|
||||
|
||||
grammar_text = std::move(applied.grammar);
|
||||
grammar_lazy = applied.grammar_lazy;
|
||||
stops = std::move(applied.additional_stops);
|
||||
generation_prompt_text = applied.generation_prompt;
|
||||
thinking_start = std::move(applied.thinking_start_tag);
|
||||
thinking_ends = std::move(applied.thinking_end_tags);
|
||||
|
||||
parser_params.format = applied.format;
|
||||
parser_params.generation_prompt = vocab ? common_chat_input_tokenize(vocab, applied.generation_prompt)
|
||||
: common_chat_input(applied.generation_prompt);
|
||||
parser_params.debug = params.debug;
|
||||
parser_params.parser = std::move(applied.parser);
|
||||
|
||||
delimiters = std::move(applied.message_delimiters);
|
||||
|
||||
if (vocab) {
|
||||
common_params_sampling resolved;
|
||||
resolved.grammar_lazy = applied.grammar_lazy;
|
||||
common_sampling_add_preserved_tokens(resolved, vocab, applied.preserved_tokens);
|
||||
common_sampling_add_grammar_triggers(resolved, vocab, std::move(applied.grammar_triggers));
|
||||
preserved_tokens = std::move(resolved.preserved_tokens);
|
||||
grammar_triggers = std::move(resolved.grammar_triggers);
|
||||
|
||||
delimiters.tokenize(vocab);
|
||||
} else {
|
||||
grammar_triggers = std::move(applied.grammar_triggers);
|
||||
}
|
||||
|
||||
if (inputs.continue_final_message != COMMON_CHAT_CONTINUATION_NONE && !params.echo) {
|
||||
// start from the prefill so it is not emitted as part of the first delta
|
||||
result = common_chat_parse(input, true, parser_params);
|
||||
}
|
||||
}
|
||||
|
||||
void common_chat_session::apply_sampling(common_params_sampling & sampling) const {
|
||||
if (!templated) {
|
||||
return;
|
||||
}
|
||||
if (!grammar_text.empty()) {
|
||||
sampling.grammar = {COMMON_GRAMMAR_TYPE_TOOL_CALLS, grammar_text};
|
||||
}
|
||||
sampling.grammar_lazy = grammar_lazy;
|
||||
sampling.generation_prompt = generation_prompt_text;
|
||||
sampling.preserved_tokens.insert(preserved_tokens.begin(), preserved_tokens.end());
|
||||
sampling.grammar_triggers.insert(sampling.grammar_triggers.end(), grammar_triggers.begin(), grammar_triggers.end());
|
||||
}
|
||||
|
||||
const common_chat_msg & common_chat_session::feed(const common_chat_input & chunk) {
|
||||
GGML_ASSERT(!finished && "feed() after finish()");
|
||||
input.append(chunk);
|
||||
auto msg = common_chat_parse(input, true, parser_params);
|
||||
if (!msg.empty()) {
|
||||
result = std::move(msg);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
const common_chat_msg & common_chat_session::finish(const common_chat_input & chunk) {
|
||||
GGML_ASSERT(!finished && "finish() called twice");
|
||||
finished = true;
|
||||
input.append(chunk);
|
||||
auto msg = common_chat_parse(input, false, parser_params);
|
||||
if (!msg.empty()) {
|
||||
result = std::move(msg);
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
std::map<std::string, bool> common_chat_templates_get_caps(const common_chat_templates * chat_templates) {
|
||||
GGML_ASSERT(chat_templates != nullptr);
|
||||
GGML_ASSERT(chat_templates->template_default != nullptr);
|
||||
|
||||
+79
-27
@@ -22,8 +22,16 @@ struct common_chat_templates;
|
||||
|
||||
namespace autoparser {
|
||||
struct generation_params;
|
||||
struct autoparser;
|
||||
} // namespace autoparser
|
||||
|
||||
struct common_chat_params;
|
||||
struct common_chat_template;
|
||||
|
||||
// Builds the prompt and parser for a template that has a dedicated handler (see common/parsers)
|
||||
using common_chat_params_init_fn = common_chat_params (*)(const common_chat_template & tmpl,
|
||||
const autoparser::generation_params & inputs);
|
||||
|
||||
struct common_chat_tool_call {
|
||||
std::string name;
|
||||
std::string arguments;
|
||||
@@ -54,19 +62,20 @@ struct common_chat_template {
|
||||
std::string eos_tok;
|
||||
std::string src;
|
||||
chat_template_caps caps;
|
||||
// Dedicated handler picked once from the source, null when the differential autoparser is used
|
||||
common_chat_params_init_fn params_init = nullptr;
|
||||
|
||||
common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token) {
|
||||
jinja::lexer lexer;
|
||||
auto lexer_res = lexer.tokenize(src);
|
||||
this->prog = jinja::parse_from_tokens(lexer_res);
|
||||
// Differential analysis, run once here when there is no dedicated handler. Null when there
|
||||
// is one, or when the analysis failed, in which case analysis_error says why.
|
||||
std::unique_ptr<autoparser::autoparser> analysis;
|
||||
std::string analysis_error;
|
||||
|
||||
this->src = lexer_res.source;
|
||||
this->bos_tok = bos_token;
|
||||
this->eos_tok = eos_token;
|
||||
common_chat_template(const std::string & src, const std::string & bos_token, const std::string & eos_token);
|
||||
|
||||
this->caps = jinja::caps_get(prog);
|
||||
// LOG_INF("%s: caps:\n%s\n", __func__, this->caps.to_string().c_str());
|
||||
}
|
||||
// autoparser is incomplete here, so these are defined where it is complete
|
||||
~common_chat_template();
|
||||
common_chat_template(common_chat_template &&);
|
||||
common_chat_template & operator=(common_chat_template &&);
|
||||
|
||||
const std::string & source() const { return src; }
|
||||
const std::string & bos_token() const { return bos_tok; }
|
||||
@@ -209,8 +218,6 @@ struct common_chat_msg_delimiters {
|
||||
|
||||
// split tokens into message spans. skips maps a start index to a length of a region to jump over without matching
|
||||
common_chat_msg_spans split(const llama_tokens & tokens, const std::map<size_t, size_t> & skips = {}) const;
|
||||
|
||||
common_json to_json() const;
|
||||
};
|
||||
|
||||
struct common_chat_tool {
|
||||
@@ -278,7 +285,7 @@ struct common_chat_params {
|
||||
std::vector<common_grammar_trigger> grammar_triggers;
|
||||
std::vector<std::string> preserved_tokens;
|
||||
std::vector<std::string> additional_stops;
|
||||
std::string parser;
|
||||
common_peg_arena parser;
|
||||
common_chat_msg_delimiters message_delimiters;
|
||||
};
|
||||
|
||||
@@ -310,16 +317,10 @@ common_chat_input common_chat_input_tokenize(const llama_vocab * vocab, const st
|
||||
// per-message parsing syntax
|
||||
// should be derived from common_chat_params
|
||||
struct common_chat_parser_params {
|
||||
common_chat_format format = COMMON_CHAT_FORMAT_CONTENT_ONLY;
|
||||
common_reasoning_format reasoning_format = COMMON_REASONING_FORMAT_NONE; // TODO: refactor this to "bool parse_reasoning"
|
||||
// Whether reasoning_content should be inlined in the content (e.g. for reasoning_format=deepseek in stream mode)
|
||||
bool reasoning_in_content = false;
|
||||
common_chat_input generation_prompt;
|
||||
bool parse_tool_calls = true;
|
||||
bool is_continuation = false;
|
||||
bool echo = false; // Include assistant prefilled msg in output
|
||||
bool debug = false; // Enable debug output for PEG parser
|
||||
common_peg_arena parser = {};
|
||||
common_chat_format format = COMMON_CHAT_FORMAT_CONTENT_ONLY;
|
||||
common_chat_input generation_prompt;
|
||||
bool debug = false; // Enable debug output for PEG parser
|
||||
common_peg_arena parser = {};
|
||||
common_chat_parser_params() = default;
|
||||
common_chat_parser_params(const common_chat_params & chat_params) {
|
||||
format = chat_params.format;
|
||||
@@ -365,6 +366,60 @@ const char * common_chat_format_name(common_chat_format format);
|
||||
common_chat_msg common_chat_parse(const common_chat_input & input, bool is_partial, const common_chat_parser_params & params);
|
||||
common_chat_msg common_chat_peg_parse(const common_peg_arena & src_parser, const common_chat_input & input, bool is_partial, const common_chat_parser_params & params);
|
||||
|
||||
struct common_chat_session_params {
|
||||
bool echo = false; // include the assistant prefill in the output when continuing a message
|
||||
bool debug = false; // enable debug output for the PEG parser
|
||||
};
|
||||
|
||||
class common_chat_session {
|
||||
public:
|
||||
common_chat_session() { result.role = "assistant"; }
|
||||
|
||||
common_chat_session(const common_chat_templates * tmpls,
|
||||
const llama_vocab * vocab,
|
||||
const common_chat_templates_inputs & inputs,
|
||||
const common_chat_session_params & params = {});
|
||||
|
||||
const std::string & prompt() const { return prompt_text; }
|
||||
common_chat_format format() const { return parser_params.format; }
|
||||
const common_chat_msg & msg() const { return result; }
|
||||
const common_peg_arena & parser() const { return parser_params.parser; }
|
||||
|
||||
const std::string & grammar() const { return grammar_text; }
|
||||
const std::string & generation_prompt() const { return generation_prompt_text; }
|
||||
const std::string & thinking_start_tag() const { return thinking_start; }
|
||||
const std::vector<std::string> & thinking_end_tags() const { return thinking_ends; }
|
||||
const std::vector<std::string> & additional_stops() const { return stops; }
|
||||
|
||||
const common_chat_msg_delimiters & message_delimiters() const { return delimiters; }
|
||||
|
||||
void apply_sampling(common_params_sampling & sampling) const;
|
||||
|
||||
bool has_template() const { return templated; }
|
||||
|
||||
const common_chat_msg & feed(const common_chat_input & chunk);
|
||||
|
||||
const common_chat_msg & finish(const common_chat_input & chunk = {});
|
||||
|
||||
private:
|
||||
std::string prompt_text;
|
||||
std::string grammar_text;
|
||||
bool grammar_lazy = false;
|
||||
std::vector<common_grammar_trigger> grammar_triggers;
|
||||
std::set<llama_token> preserved_tokens;
|
||||
std::vector<std::string> stops;
|
||||
std::string generation_prompt_text;
|
||||
std::string thinking_start;
|
||||
std::vector<std::string> thinking_ends;
|
||||
|
||||
common_chat_parser_params parser_params;
|
||||
common_chat_msg_delimiters delimiters;
|
||||
common_chat_input input;
|
||||
common_chat_msg result;
|
||||
bool templated = false;
|
||||
bool finished = false;
|
||||
};
|
||||
|
||||
// used by arg and server
|
||||
const char * common_reasoning_format_name(common_reasoning_format format);
|
||||
common_reasoning_format common_reasoning_format_from_name(const std::string & format);
|
||||
@@ -401,8 +456,7 @@ std::string common_chat_template_generation_prompt(
|
||||
|
||||
std::optional<common_chat_params> common_chat_try_specialized_template(
|
||||
const common_chat_template & tmpl,
|
||||
const std::string & src,
|
||||
autoparser::generation_params & params);
|
||||
const autoparser::generation_params & params);
|
||||
|
||||
|
||||
// specialized per-task preset
|
||||
@@ -412,5 +466,3 @@ struct common_chat_prompt_preset {
|
||||
};
|
||||
|
||||
common_chat_prompt_preset common_chat_get_asr_prompt(const common_chat_templates * chat_templates);
|
||||
|
||||
common_chat_msg_delimiters common_chat_msg_delimiters_parse(const common_json & delimiters);
|
||||
|
||||
+27
-28
@@ -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) :
|
||||
@@ -2388,40 +2391,36 @@ void common_prompt_checkpoint::update_dft(
|
||||
}
|
||||
}
|
||||
|
||||
void common_prompt_checkpoint::load_tgt(
|
||||
bool common_prompt_checkpoint::load_tgt(
|
||||
llama_context * ctx,
|
||||
llama_seq_id seq_id,
|
||||
llama_state_seq_flags flags) const {
|
||||
if (ctx == nullptr) {
|
||||
return;
|
||||
return true;
|
||||
}
|
||||
|
||||
if (data_tgt.empty()) {
|
||||
return;
|
||||
return true;
|
||||
}
|
||||
|
||||
const size_t n = llama_state_seq_set_data_ext(ctx, data_tgt.data(), data_tgt.size(), seq_id, flags);
|
||||
if (n != data_tgt.size()) {
|
||||
GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", data_tgt.size(), n);
|
||||
}
|
||||
return n == data_tgt.size();
|
||||
}
|
||||
|
||||
void common_prompt_checkpoint::load_dft(
|
||||
bool common_prompt_checkpoint::load_dft(
|
||||
llama_context * ctx,
|
||||
llama_seq_id seq_id,
|
||||
llama_state_seq_flags flags) const {
|
||||
if (ctx == nullptr) {
|
||||
return;
|
||||
return true;
|
||||
}
|
||||
|
||||
if (data_dft.empty()) {
|
||||
return;
|
||||
return true;
|
||||
}
|
||||
|
||||
const size_t n = llama_state_seq_set_data_ext(ctx, data_dft.data(), data_dft.size(), seq_id, flags);
|
||||
if (n != data_dft.size()) {
|
||||
GGML_ABORT("checkpoint size mismatch: expected %zu, got %zu\n", data_dft.size(), n);
|
||||
}
|
||||
return n == data_dft.size();
|
||||
}
|
||||
|
||||
void common_prompt_checkpoint::clear_tgt() {
|
||||
|
||||
+12
-7
@@ -593,7 +593,7 @@ struct common_params {
|
||||
ggml_type cache_type_k = GGML_TYPE_F16; // KV cache data type for the K
|
||||
ggml_type cache_type_v = GGML_TYPE_F16; // KV cache data type for the V
|
||||
|
||||
size_t moe_cache_size = 0; // GPU cache size in bytes for the MoE experts kept in the CPU
|
||||
size_t moe_cache_size = 0; // GPU cache size in bytes for the MoE experts kept in the CPU, split among the GPUs like the layers
|
||||
|
||||
common_conversation_mode conversation_mode = COMMON_CONVERSATION_MODE_AUTO;
|
||||
|
||||
@@ -625,7 +625,7 @@ struct common_params {
|
||||
std::string cls_sep = "\t"; // separator of classification sequences
|
||||
|
||||
// server params
|
||||
int32_t port = 8080; // server listens on this network port
|
||||
int32_t port = 9931; // server listens on this network port
|
||||
bool reuse_port = false; // allow multiple sockets to bind to the same port
|
||||
int32_t timeout_read = 3600; // http read timeout in seconds
|
||||
int32_t timeout_write = timeout_read; // http write timeout in seconds
|
||||
@@ -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 {
|
||||
@@ -1296,12 +1300,13 @@ struct common_prompt_checkpoint {
|
||||
llama_seq_id seq_id,
|
||||
llama_state_seq_flags flags);
|
||||
|
||||
void load_tgt(
|
||||
// return false if the state could not be restored
|
||||
bool load_tgt(
|
||||
llama_context * ctx,
|
||||
llama_seq_id seq_id,
|
||||
llama_state_seq_flags flags) const;
|
||||
|
||||
void load_dft(
|
||||
bool load_dft(
|
||||
llama_context * ctx,
|
||||
llama_seq_id seq_id,
|
||||
llama_state_seq_flags flags) const;
|
||||
|
||||
@@ -77,7 +77,7 @@ common_chat_params common_chat_params_init_cohere2moe(const common_chat_template
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PREFIX);
|
||||
auto end = p.end();
|
||||
|
||||
@@ -124,12 +124,10 @@ common_chat_params common_chat_params_init_cohere2moe(const common_chat_template
|
||||
return generation_prompt + reasoning + body + p.optional(p.literal(TURN_END)) + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !has_response_format && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -145,7 +145,7 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
|
||||
bool require_tools = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
bool has_tool_calls = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PROMPT);
|
||||
auto end = p.end();
|
||||
|
||||
@@ -256,12 +256,10 @@ common_chat_params common_chat_params_init_deepseek_v3_2(const common_chat_templ
|
||||
generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = has_tools && !require_tools;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -21,7 +21,7 @@ common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_te
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
// Functionary v3.2 format:
|
||||
// - Normal content: >>>all\n{content}
|
||||
// - Tool calls: >>>function_name\n{json_args}
|
||||
@@ -76,13 +76,11 @@ common_chat_params common_chat_params_init_functionary_v3_2(const common_chat_te
|
||||
return generation_prompt + ret;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
// Grammar trigger for when the model starts outputting a tool call
|
||||
|
||||
@@ -198,7 +198,7 @@ common_chat_params common_chat_params_init_gemma4(const common_chat_template &
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto start = p.rule("start", p.optional(p.literal("<|turn>model\n")));
|
||||
|
||||
if (extract_reasoning) {
|
||||
@@ -290,12 +290,10 @@ common_chat_params common_chat_params_init_gemma4(const common_chat_template &
|
||||
return start + p.one_or_more(message);
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -25,7 +25,7 @@ common_chat_params common_chat_params_init_gigachat_v3(
|
||||
auto include_grammar = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE;
|
||||
const auto *tool_call_start_prefix = "<|message_sep|>\n\nfunction call<|role_sep|>\n";
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto ret = p.eps();
|
||||
if (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE) {
|
||||
// Build a choice of all available tools
|
||||
@@ -60,13 +60,11 @@ common_chat_params common_chat_params_init_gigachat_v3(
|
||||
return p.literal("assistant<|role_sep|>\n") + ret;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -45,8 +45,7 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
|
||||
data.thinking_start_tag = "<|channel|>analysis<|message|>";
|
||||
data.thinking_end_tags = {"<|end|>"};
|
||||
|
||||
// These special tokens are required to parse properly, so we include them
|
||||
// even if parse_tool_calls is false.
|
||||
// These special tokens are required to parse properly
|
||||
data.preserved_tokens = {
|
||||
"<|channel|>", "<|constrain|>", "<|message|>", "<|start|>", "<|end|>",
|
||||
};
|
||||
@@ -68,7 +67,7 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto start = p.rule("start", p.literal("<|start|>assistant"));
|
||||
auto end = p.rule("end", p.literal("<|end|>"));
|
||||
auto content = p.rule("message-content", p.until("<|end|>"));
|
||||
@@ -138,12 +137,10 @@ common_chat_params common_chat_params_init_gpt_oss(const common_chat_template &
|
||||
return p.zero_or_more(start + any) + start + (final_msg | unsolicited);
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -77,7 +77,7 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PREFIX);
|
||||
|
||||
auto think_end = p.choice();
|
||||
@@ -174,12 +174,10 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template
|
||||
return generation_prompt + (reasoning << content << tool_calls);
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED);
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
if (data.grammar_lazy) {
|
||||
|
||||
@@ -48,7 +48,7 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
// Kimi K2 Thinking format:
|
||||
// - Reasoning: <think>{reasoning}</think>
|
||||
// - Content: text after reasoning
|
||||
@@ -111,12 +111,10 @@ common_chat_params common_chat_params_init_kimi_k2(const common_chat_template &
|
||||
return generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -66,7 +66,7 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto end = p.end();
|
||||
|
||||
auto start = p.optional(p.literal(MSG_START));
|
||||
@@ -151,12 +151,10 @@ common_chat_params common_chat_params_init_kimi_k3(const common_chat_template &
|
||||
return start + reasoning + response + tools + trailer + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -64,7 +64,7 @@ common_chat_params common_chat_params_init_lfm2(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PROMPT);
|
||||
auto end = p.end();
|
||||
|
||||
@@ -93,12 +93,10 @@ common_chat_params common_chat_params_init_lfm2(const common_chat_template &
|
||||
return generation_prompt + reasoning + content + tool_calls + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -80,7 +80,7 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto end = p.end();
|
||||
|
||||
// the effective parse input is generation_prompt + model output, so the
|
||||
@@ -185,12 +185,10 @@ common_chat_params common_chat_params_init_ling3(const common_chat_template &
|
||||
return opener + reasoning + content + tools + tail + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !has_response_format && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED;
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -48,8 +48,7 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
|
||||
data.thinking_start_tag = "<|channel|>analysis<|message|>";
|
||||
data.thinking_end_tags = {"<|end|>"};
|
||||
|
||||
// These special tokens are required to parse properly, so we include them
|
||||
// even if parse_tool_calls is false.
|
||||
// These special tokens are required to parse properly
|
||||
data.preserved_tokens = {
|
||||
"<|channel|>", "<|constrain|>", "<|message|>", "<|start|>", "<|end|>",
|
||||
};
|
||||
@@ -71,7 +70,7 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE;
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
// tokenizer space after special tokens; not p.space() since GBNF `space` allows one space only
|
||||
auto sp = p.chars("[ ]", 0, -1);
|
||||
auto channel_tag = p.literal("<|channel|>") + sp;
|
||||
@@ -144,12 +143,10 @@ common_chat_params common_chat_params_init_llm_jp_harmony(const common_chat_temp
|
||||
return p.zero_or_more(start + any) + start + final_msg;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -46,7 +46,7 @@ common_chat_params common_chat_params_init_minicpm5(const common_chat_template &
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal("<|im_start|>assistant\n");
|
||||
|
||||
auto reasoning = p.eps();
|
||||
@@ -113,12 +113,10 @@ common_chat_params common_chat_params_init_minicpm5(const common_chat_template &
|
||||
return generation_prompt + reasoning + p.content(p.rest()) + p.end();
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -56,7 +56,7 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.prefix(GEN_PROMPT, THINK_START);
|
||||
auto end = p.end();
|
||||
|
||||
@@ -213,12 +213,10 @@ common_chat_params common_chat_params_init_minimax_m3(const common_chat_template
|
||||
return generation_prompt + reasoning + content_before_tools + tool_calls + end;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -72,7 +72,7 @@ common_chat_params common_chat_params_init_ministral_3(const common_chat_templat
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.eps();
|
||||
auto reasoning =
|
||||
extract_reasoning ? p.optional("[THINK]" + p.reasoning(p.until("[/THINK]")) + "[/THINK]") : p.eps();
|
||||
@@ -108,13 +108,11 @@ common_chat_params common_chat_params_init_ministral_3(const common_chat_templat
|
||||
return generation_prompt + (reasoning << p.content(p.rest()));
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
data.grammar_triggers = {
|
||||
|
||||
@@ -48,7 +48,7 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
|
||||
// Constrained grammar whenever tools are offered or a response format is requested.
|
||||
auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE);
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto start = p.rule("start", p.literal("<|start|>assistant"));
|
||||
|
||||
if (!extract_reasoning && !include_grammar) {
|
||||
@@ -131,12 +131,10 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
|
||||
return p.zero_or_more(start + analysis) + start + final_msg;
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = !(has_response_format || (has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED));
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
data.grammar_triggers = {
|
||||
{ COMMON_GRAMMAR_TRIGGER_TYPE_PATTERN,
|
||||
|
||||
@@ -71,7 +71,7 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
|
||||
});
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
auto generation_prompt = p.literal(GEN_PREFIX);
|
||||
|
||||
auto reasoning = p.eps();
|
||||
@@ -174,13 +174,11 @@ common_chat_params common_chat_params_init_qwen3_coder(const common_chat_templat
|
||||
return generation_prompt + (reasoning << p.content(p.rest()));
|
||||
});
|
||||
|
||||
data.parser = parser.save();
|
||||
|
||||
if (include_grammar) {
|
||||
data.grammar_lazy = has_tools && inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_AUTO;
|
||||
|
||||
data.grammar = build_grammar([&](const common_grammar_builder & builder) {
|
||||
parser.build_grammar(builder, data.grammar_lazy);
|
||||
data.parser.build_grammar(builder, data.grammar_lazy);
|
||||
});
|
||||
|
||||
if (data.grammar_lazy) {
|
||||
|
||||
@@ -54,10 +54,9 @@ common_chat_params common_chat_params_init_translate_gemma(
|
||||
data.prompt += data.generation_prompt;
|
||||
}
|
||||
|
||||
auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
data.parser = build_chat_peg_parser([&](common_chat_peg_builder & p) {
|
||||
return p.literal(data.generation_prompt) << p.content(p.rest());
|
||||
});
|
||||
data.parser = parser.save();
|
||||
|
||||
return data;
|
||||
}
|
||||
|
||||
@@ -1814,309 +1814,6 @@ void common_peg_arena::build_grammar(const common_grammar_builder & builder, boo
|
||||
}
|
||||
}
|
||||
|
||||
static common_json serialize_parser_variant(const common_peg_parser_variant & variant) {
|
||||
using json = common_json;
|
||||
|
||||
return std::visit([](const auto & p) -> json {
|
||||
using T = std::decay_t<decltype(p)>;
|
||||
|
||||
if constexpr (std::is_same_v<T, common_peg_epsilon_parser>) {
|
||||
return json{{"type", "epsilon"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_start_parser>) {
|
||||
return json{{"type", "start"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_end_parser>) {
|
||||
return json{{"type", "end"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_literal_parser>) {
|
||||
return json{{"type", "literal"}, {"literal", p.literal}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_sequence_parser>) {
|
||||
return json{{"type", "sequence"}, {"children", p.children}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_choice_parser>) {
|
||||
return json{{"type", "choice"}, {"children", p.children}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_repetition_parser>) {
|
||||
return json{
|
||||
{"type", "repetition"},
|
||||
{"child", p.child},
|
||||
{"min_count", p.min_count},
|
||||
{"max_count", p.max_count}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_and_parser>) {
|
||||
return json{{"type", "and"}, {"child", p.child}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_not_parser>) {
|
||||
return json{{"type", "not"}, {"child", p.child}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_any_parser>) {
|
||||
return json{{"type", "any"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_space_parser>) {
|
||||
return json{{"type", "space"}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_chars_parser>) {
|
||||
json ranges = json::array();
|
||||
for (const auto & range : p.ranges) {
|
||||
ranges.push_back({{"start", range.start}, {"end", range.end}});
|
||||
}
|
||||
return json{
|
||||
{"type", "chars"},
|
||||
{"pattern", p.pattern},
|
||||
{"ranges", ranges},
|
||||
{"negated", p.negated},
|
||||
{"min_count", p.min_count},
|
||||
{"max_count", p.max_count}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_string_parser>) {
|
||||
return json{{"type", "string"}, {"delimiter", std::string(1, p.delimiter)}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_until_parser>) {
|
||||
return json{{"type", "until"}, {"delimiters", p.delimiters}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_schema_parser>) {
|
||||
return json{
|
||||
{"type", "schema"},
|
||||
{"child", p.child},
|
||||
{"name", p.name},
|
||||
{"raw", p.raw}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_rule_parser>) {
|
||||
return json{
|
||||
{"type", "rule"},
|
||||
{"name", p.name},
|
||||
{"child", p.child},
|
||||
{"trigger", p.trigger}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_ref_parser>) {
|
||||
return json{{"type", "ref"}, {"name", p.name}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_atomic_parser>) {
|
||||
return json{{"type", "atomic"}, {"child", p.child}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_tag_parser>) {
|
||||
return json{
|
||||
{"type", "tag"},
|
||||
{"child", p.child},
|
||||
{"tag", p.tag}
|
||||
};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_gbnf_parser>) {
|
||||
return json{{"type", "gbnf"}, {"child", p.child}, {"grammar", p.grammar}};
|
||||
} else if constexpr (std::is_same_v<T, common_peg_ac_parser>) {
|
||||
return json{{"type", "ac"}, {"child", p.child}, {"delimiters", p.delimiters}};
|
||||
}
|
||||
}, variant);
|
||||
}
|
||||
|
||||
common_json common_peg_arena::to_json() const {
|
||||
auto parsers = common_json::array();
|
||||
for (const auto & parser : parsers_) {
|
||||
parsers.push_back(serialize_parser_variant(parser));
|
||||
}
|
||||
return common_json{
|
||||
{"parsers", parsers},
|
||||
{"rules", rules_},
|
||||
{"root", root_}
|
||||
};
|
||||
}
|
||||
|
||||
static common_peg_parser_variant deserialize_parser_variant(const common_json & j) {
|
||||
if (!j.contains("type") || !j["type"].is_string()) {
|
||||
throw std::runtime_error("Parser variant JSON missing or invalid 'type' field");
|
||||
}
|
||||
|
||||
std::string type = j["type"];
|
||||
|
||||
if (type == "epsilon") {
|
||||
return common_peg_epsilon_parser{};
|
||||
}
|
||||
if (type == "start") {
|
||||
return common_peg_start_parser{};
|
||||
}
|
||||
if (type == "end") {
|
||||
return common_peg_end_parser{};
|
||||
}
|
||||
if (type == "literal") {
|
||||
if (!j.contains("literal") || !j["literal"].is_string()) {
|
||||
throw std::runtime_error("literal parser missing or invalid 'literal' field");
|
||||
}
|
||||
return common_peg_literal_parser{j["literal"]};
|
||||
}
|
||||
if (type == "sequence") {
|
||||
if (!j.contains("children") || !j["children"].is_array()) {
|
||||
throw std::runtime_error("sequence parser missing or invalid 'children' field");
|
||||
}
|
||||
return common_peg_sequence_parser{j["children"].get<std::vector<common_peg_parser_id>>()};
|
||||
}
|
||||
if (type == "choice") {
|
||||
if (!j.contains("children") || !j["children"].is_array()) {
|
||||
throw std::runtime_error("choice parser missing or invalid 'children' field");
|
||||
}
|
||||
return common_peg_choice_parser{j["children"].get<std::vector<common_peg_parser_id>>()};
|
||||
}
|
||||
if (type == "repetition") {
|
||||
if (!j.contains("child") || !j.contains("min_count") || !j.contains("max_count")) {
|
||||
throw std::runtime_error("repetition parser missing required fields");
|
||||
}
|
||||
return common_peg_repetition_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["min_count"].get<int>(),
|
||||
j["max_count"].get<int>()
|
||||
};
|
||||
}
|
||||
if (type == "and") {
|
||||
if (!j.contains("child")) {
|
||||
throw std::runtime_error("and parser missing 'child' field");
|
||||
}
|
||||
return common_peg_and_parser{j["child"].get<common_peg_parser_id>()};
|
||||
}
|
||||
if (type == "not") {
|
||||
if (!j.contains("child")) {
|
||||
throw std::runtime_error("not parser missing 'child' field");
|
||||
}
|
||||
return common_peg_not_parser{j["child"].get<common_peg_parser_id>()};
|
||||
}
|
||||
if (type == "any") {
|
||||
return common_peg_any_parser{};
|
||||
}
|
||||
if (type == "space") {
|
||||
return common_peg_space_parser{};
|
||||
}
|
||||
if (type == "chars") {
|
||||
if (!j.contains("pattern") || !j.contains("ranges") || !j.contains("negated") ||
|
||||
!j.contains("min_count") || !j.contains("max_count")) {
|
||||
throw std::runtime_error("chars parser missing required fields");
|
||||
}
|
||||
common_peg_chars_parser parser;
|
||||
parser.pattern = j["pattern"];
|
||||
parser.negated = j["negated"].get<bool>();
|
||||
parser.min_count = j["min_count"].get<int>();
|
||||
parser.max_count = j["max_count"].get<int>();
|
||||
for (const auto & range_json : j["ranges"]) {
|
||||
if (!range_json.contains("start") || !range_json.contains("end")) {
|
||||
throw std::runtime_error("char_range missing 'start' or 'end' field");
|
||||
}
|
||||
parser.ranges.push_back({
|
||||
range_json["start"].get<uint32_t>(),
|
||||
range_json["end"].get<uint32_t>()
|
||||
});
|
||||
}
|
||||
return parser;
|
||||
}
|
||||
if (type == "string") {
|
||||
if (!j.contains("delimiter")) {
|
||||
throw std::runtime_error("string parser missing delimiter field.");
|
||||
}
|
||||
std::string delimiter = j["delimiter"];
|
||||
if (delimiter.empty()) {
|
||||
throw std::runtime_error("string parser delimiter is empty.");
|
||||
}
|
||||
return common_peg_string_parser{delimiter[0]};
|
||||
}
|
||||
if (type == "until") {
|
||||
if (!j.contains("delimiters") || !j["delimiters"].is_array()) {
|
||||
throw std::runtime_error("until parser missing or invalid 'delimiters' field");
|
||||
}
|
||||
return common_peg_until_parser{j["delimiters"].get<std::vector<std::string>>()};
|
||||
}
|
||||
if (type == "schema") {
|
||||
if (!j.contains("child") || !j.contains("name") || !j.contains("raw")) {
|
||||
throw std::runtime_error("schema parser missing required fields");
|
||||
}
|
||||
common_peg_schema_parser parser;
|
||||
parser.child = j["child"].get<common_peg_parser_id>();
|
||||
parser.name = j["name"];
|
||||
parser.raw = j["raw"].get<bool>();
|
||||
return parser;
|
||||
}
|
||||
if (type == "rule") {
|
||||
if (!j.contains("name") || !j.contains("child") || !j.contains("trigger")) {
|
||||
throw std::runtime_error("rule parser missing required fields");
|
||||
}
|
||||
return common_peg_rule_parser{
|
||||
j["name"].get<std::string>(),
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["trigger"].get<bool>()
|
||||
};
|
||||
}
|
||||
if (type == "ref") {
|
||||
if (!j.contains("name") || !j["name"].is_string()) {
|
||||
throw std::runtime_error("ref parser missing or invalid 'name' field");
|
||||
}
|
||||
return common_peg_ref_parser{j["name"]};
|
||||
}
|
||||
if (type == "atomic") {
|
||||
if (!j.contains("child")) {
|
||||
throw std::runtime_error("tag parser missing required fields");
|
||||
}
|
||||
return common_peg_atomic_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
};
|
||||
}
|
||||
if (type == "tag") {
|
||||
if (!j.contains("child") || !j.contains("tag")) {
|
||||
throw std::runtime_error("tag parser missing required fields");
|
||||
}
|
||||
return common_peg_tag_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["tag"].get<std::string>(),
|
||||
};
|
||||
}
|
||||
|
||||
if (type == "gbnf") {
|
||||
if (!j.contains("child") || !j.contains("grammar")) {
|
||||
throw std::runtime_error("gbnf parser missing required fields");
|
||||
}
|
||||
return common_peg_gbnf_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["grammar"].get<std::string>(),
|
||||
};
|
||||
}
|
||||
|
||||
if (type == "ac") {
|
||||
if (!j.contains("child") || !j.contains("delimiters") || !j["delimiters"].is_array() || j["delimiters"].empty()) {
|
||||
throw std::runtime_error("ac parser requires 'child' and a non-empty 'delimiters' array");
|
||||
}
|
||||
return common_peg_ac_parser{
|
||||
j["child"].get<common_peg_parser_id>(),
|
||||
j["delimiters"].get<std::vector<std::string>>(),
|
||||
};
|
||||
}
|
||||
|
||||
throw std::runtime_error("Unknown parser type: " + type);
|
||||
}
|
||||
|
||||
common_peg_arena common_peg_arena::from_json(const common_json & j) {
|
||||
if (!j.contains("parsers") || !j["parsers"].is_array()) {
|
||||
throw std::runtime_error("JSON missing or invalid 'parsers' array");
|
||||
}
|
||||
if (!j.contains("rules") || !j["rules"].is_object()) {
|
||||
throw std::runtime_error("JSON missing or invalid 'rules' object");
|
||||
}
|
||||
if (!j.contains("root")) {
|
||||
throw std::runtime_error("JSON missing 'root' field");
|
||||
}
|
||||
|
||||
common_peg_arena arena;
|
||||
|
||||
const auto & parsers_json = j["parsers"];
|
||||
arena.parsers_.reserve(parsers_json.size());
|
||||
for (const auto & parser_json : parsers_json) {
|
||||
arena.parsers_.push_back(deserialize_parser_variant(parser_json));
|
||||
}
|
||||
|
||||
arena.rules_ = j["rules"].get<std::unordered_map<std::string, common_peg_parser_id>>();
|
||||
|
||||
for (const auto & [name, id] : arena.rules_) {
|
||||
if (id >= arena.parsers_.size()) {
|
||||
throw std::runtime_error("Rule '" + name + "' references invalid parser ID: " + std::to_string(id));
|
||||
}
|
||||
}
|
||||
|
||||
arena.root_ = j["root"].get<common_peg_parser_id>();
|
||||
if (arena.root_ != COMMON_PEG_INVALID_PARSER_ID && arena.root_ >= arena.parsers_.size()) {
|
||||
throw std::runtime_error("Root references invalid parser ID: " + std::to_string(arena.root_));
|
||||
}
|
||||
|
||||
return arena;
|
||||
}
|
||||
|
||||
std::string common_peg_arena::save() const {
|
||||
return to_json().dump();
|
||||
}
|
||||
|
||||
void common_peg_arena::load(const std::string & data) {
|
||||
*this = from_json(common_json::parse(data));
|
||||
}
|
||||
|
||||
common_peg_arena build_peg_parser(const std::function<common_peg_parser(common_peg_parser_builder & builder)> & fn) {
|
||||
common_peg_parser_builder builder;
|
||||
builder.set_root(fn(builder));
|
||||
|
||||
@@ -357,12 +357,6 @@ class common_peg_arena {
|
||||
|
||||
std::string dump(common_peg_parser_id id) const;
|
||||
|
||||
common_json to_json() const;
|
||||
static common_peg_arena from_json(const common_json & j);
|
||||
|
||||
std::string save() const;
|
||||
void load(const std::string & data);
|
||||
|
||||
friend class common_peg_parser_builder;
|
||||
|
||||
private:
|
||||
|
||||
@@ -1050,3 +1050,41 @@ std::vector<common_sampler_type> common_sampler_types_from_chars(const std::stri
|
||||
|
||||
return samplers;
|
||||
}
|
||||
|
||||
void common_sampling_add_preserved_tokens(common_params_sampling & sampling, const llama_vocab * vocab, const std::vector<std::string> & tokens) {
|
||||
GGML_ASSERT(vocab != nullptr);
|
||||
for (const auto & t : tokens) {
|
||||
auto ids = common_tokenize(vocab, t, false, true);
|
||||
if (ids.size() == 1) {
|
||||
sampling.preserved_tokens.insert(ids[0]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void common_sampling_add_grammar_triggers(common_params_sampling & sampling, const llama_vocab * vocab, std::vector<common_grammar_trigger> triggers) {
|
||||
GGML_ASSERT(vocab != nullptr);
|
||||
for (auto & trigger : triggers) {
|
||||
if (trigger.type == COMMON_GRAMMAR_TRIGGER_TYPE_WORD) {
|
||||
const auto & word = trigger.value;
|
||||
auto ids = common_tokenize(vocab, word, false, true);
|
||||
if (ids.size() == 1) {
|
||||
auto token = ids[0];
|
||||
if (std::find(sampling.preserved_tokens.begin(), sampling.preserved_tokens.end(), (llama_token) token) == sampling.preserved_tokens.end()) {
|
||||
throw std::runtime_error("Grammar trigger word should be marked as preserved token: " + word);
|
||||
}
|
||||
common_grammar_trigger token_trigger;
|
||||
token_trigger.type = COMMON_GRAMMAR_TRIGGER_TYPE_TOKEN;
|
||||
token_trigger.value = word;
|
||||
token_trigger.token = token;
|
||||
sampling.grammar_triggers.push_back(std::move(token_trigger));
|
||||
} else {
|
||||
sampling.grammar_triggers.push_back({COMMON_GRAMMAR_TRIGGER_TYPE_WORD, word});
|
||||
}
|
||||
} else {
|
||||
sampling.grammar_triggers.push_back(std::move(trigger));
|
||||
}
|
||||
}
|
||||
if (sampling.grammar_lazy && sampling.grammar_triggers.empty()) {
|
||||
throw std::runtime_error("Error: no triggers set for lazy grammar!");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -118,6 +118,12 @@ std::string common_sampler_type_to_str(enum common_sampler_type cnstr);
|
||||
std::vector<enum common_sampler_type> common_sampler_types_from_names(const std::vector<std::string> & names);
|
||||
std::vector<enum common_sampler_type> common_sampler_types_from_chars(const std::string & chars);
|
||||
|
||||
// add the strings that are a single token in the vocab to the preserved tokens
|
||||
void common_sampling_add_preserved_tokens(common_params_sampling & sampling, const llama_vocab * vocab, const std::vector<std::string> & tokens);
|
||||
|
||||
// add grammar triggers, a trigger word that is a single token becomes a token trigger and must be a preserved token
|
||||
void common_sampling_add_grammar_triggers(common_params_sampling & sampling, const llama_vocab * vocab, std::vector<common_grammar_trigger> triggers);
|
||||
|
||||
llama_sampler * llama_sampler_init_llg(const llama_vocab * vocab,
|
||||
const char * grammar_kind, const char * grammar_data);
|
||||
|
||||
|
||||
+42
-43
@@ -439,6 +439,25 @@ class ModelBase:
|
||||
|
||||
return (unpacked * scale.unsqueeze(-1).float()).reshape(shape)
|
||||
|
||||
def dequant_fp8() -> None:
|
||||
for name in self.model_tensors.keys():
|
||||
if name.endswith(".weight_scale"):
|
||||
weight_name = name.removesuffix("_scale")
|
||||
if weight_name not in self.model_tensors:
|
||||
tensors_to_remove.append(name)
|
||||
continue
|
||||
w = self.model_tensors[weight_name]
|
||||
s = self.model_tensors[name]
|
||||
is_fp8_weight = False
|
||||
if self._fp8_as_q8:
|
||||
is_fp8_weight = w().dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
|
||||
self.model_tensors[weight_name] = lambda w=w, s=s: dequant_simple(w(), s(), None)
|
||||
tensors_to_remove.append(name)
|
||||
if is_fp8_weight:
|
||||
self._fp8_dequantized.add(weight_name)
|
||||
if name.endswith((".input_scale", ".k_scale", ".v_scale")):
|
||||
tensors_to_remove.append(name)
|
||||
|
||||
if quant_method == "bitnet":
|
||||
for name in self.model_tensors.keys():
|
||||
if name.endswith(".weight_scale"):
|
||||
@@ -498,18 +517,14 @@ class ModelBase:
|
||||
elif quant_method == "compressed-tensors":
|
||||
quant_format = quant_config["format"]
|
||||
groups = quant_config["config_groups"]
|
||||
nvfp4_compressed_tensors = (
|
||||
quant_format == "nvfp4-pack-quantized"
|
||||
or quant_format == "mixed-precision"
|
||||
and bool(groups)
|
||||
and all(g.get("format") == "nvfp4-pack-quantized" for g in groups.values() if isinstance(g, dict))
|
||||
)
|
||||
nvfp4_compressed_tensors = self._is_nvfp4_compressed_tensors(quant_method, quant_format, groups)
|
||||
|
||||
if len(groups) > 1 and not nvfp4_compressed_tensors:
|
||||
if nvfp4_compressed_tensors:
|
||||
dequant_fp8()
|
||||
elif len(groups) > 1:
|
||||
raise NotImplementedError("Can't handle multiple config groups for compressed-tensors yet")
|
||||
weight_config = tuple(groups.values())[0]["weights"]
|
||||
|
||||
if quant_format == "float-quantized" or quant_format == "int-quantized" or quant_format == "naive-quantized":
|
||||
elif quant_format == "float-quantized" or quant_format == "int-quantized" or quant_format == "naive-quantized":
|
||||
weight_config = tuple(groups.values())[0]["weights"]
|
||||
block_size = weight_config.get("block_structure", None)
|
||||
strategy = weight_config.get("strategy")
|
||||
assert strategy == "channel" or strategy == "block"
|
||||
@@ -529,6 +544,7 @@ class ModelBase:
|
||||
if self._fp8_as_q8 and is_fp8:
|
||||
self._fp8_dequantized.add(weight_name)
|
||||
elif quant_format == "pack-quantized":
|
||||
weight_config = tuple(groups.values())[0]["weights"]
|
||||
assert weight_config.get("strategy") == "group"
|
||||
assert weight_config.get("type", "int") == "int"
|
||||
num_bits = weight_config.get("num_bits")
|
||||
@@ -550,32 +566,10 @@ class ModelBase:
|
||||
tensors_to_remove += [base_name + n for n in ("_packed", "_shape", "_scale")]
|
||||
if (base_name + "_zero_point") in self.model_tensors:
|
||||
tensors_to_remove.append(base_name + "_zero_point")
|
||||
elif nvfp4_compressed_tensors:
|
||||
# Don't error from compressed-tensors, we'll handle them in _generate_nvfp4_tensors
|
||||
pass
|
||||
else:
|
||||
raise NotImplementedError(f"Quant format {quant_format!r} for method {quant_method!r} is not yet supported")
|
||||
elif quant_method == "modelopt":
|
||||
# Mixed-precision ModelOpt models: NVFP4 tensors are handled by
|
||||
# _generate_nvfp4_tensors; FP8 tensors have 1D weight_scale and
|
||||
# are dequantized here. k/v scale tensors are unused.
|
||||
for name in self.model_tensors.keys():
|
||||
if name.endswith(".weight_scale"):
|
||||
weight_name = name.removesuffix("_scale")
|
||||
if weight_name not in self.model_tensors:
|
||||
tensors_to_remove.append(name)
|
||||
continue
|
||||
w = self.model_tensors[weight_name]
|
||||
s = self.model_tensors[name]
|
||||
is_fp8_weight = False
|
||||
if self._fp8_as_q8:
|
||||
is_fp8_weight = w().dtype in (torch.float8_e4m3fn, torch.float8_e5m2)
|
||||
self.model_tensors[weight_name] = lambda w=w, s=s: dequant_simple(w(), s(), None)
|
||||
tensors_to_remove.append(name)
|
||||
if is_fp8_weight:
|
||||
self._fp8_dequantized.add(weight_name)
|
||||
if name.endswith((".input_scale", ".k_scale", ".v_scale")):
|
||||
tensors_to_remove.append(name)
|
||||
dequant_fp8()
|
||||
elif quant_method is not None:
|
||||
raise NotImplementedError(f"Quant method is not yet supported: {quant_method!r}")
|
||||
|
||||
@@ -821,6 +815,18 @@ class ModelBase:
|
||||
func=load,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _is_nvfp4_compressed_tensors(quant_method, quant_format, groups) -> bool:
|
||||
# Some models use per-tensor quant_algo (e.g. "MIXED_PRECISION" with
|
||||
# per-layer NVFP4/FP8) instead of a single global "NVFP4" value.
|
||||
if quant_method != "compressed-tensors":
|
||||
return False
|
||||
if quant_format == "nvfp4-pack-quantized":
|
||||
return True
|
||||
if quant_format != "mixed-precision" or not groups:
|
||||
return False
|
||||
return any(g.get("format") == "nvfp4-pack-quantized" for g in groups.values() if isinstance(g, dict))
|
||||
|
||||
@staticmethod
|
||||
def _nvfp4_pack(weight: Tensor, scale: Tensor) -> tuple[np.ndarray, list[int]]:
|
||||
"""Repack NVFP4 ModelOpt tensors into ggml super-block layout.
|
||||
@@ -878,8 +884,8 @@ class ModelBase:
|
||||
weight = LazyTorchTensor.to_eager(self.model_tensors[name]())
|
||||
scale = LazyTorchTensor.to_eager(self.model_tensors[scale_name]())
|
||||
|
||||
# Skip non-NVFP4 tensors (e.g. FP8 with per-channel 1D scales)
|
||||
if scale.ndim < 2:
|
||||
# Skip non-NVFP4 tensors(e.g. 1D scale, or float8 weight)
|
||||
if scale.ndim < 2 or weight.dtype in (torch.float8_e4m3fn, torch.float8_e5m2):
|
||||
continue
|
||||
|
||||
scale2 = LazyTorchTensor.to_eager(self.model_tensors.get(scale2_name, lambda: torch.tensor(1.0))())
|
||||
@@ -980,14 +986,7 @@ class ModelBase:
|
||||
quant_groups = quant_config.get("config_groups", quant_groups) or {}
|
||||
quant_layers = quant_config.get("quantized_layers", quant_layers) or {}
|
||||
|
||||
# Some models use per-tensor quant_algo (e.g. "MIXED_PRECISION" with
|
||||
# per-layer NVFP4/FP8) instead of a single global "NVFP4" value.
|
||||
nvfp4_compressed_tensors = quant_method == "compressed-tensors" and (
|
||||
quant_format == "nvfp4-pack-quantized"
|
||||
or quant_format == "mixed-precision"
|
||||
and bool(quant_groups)
|
||||
and all(g.get("format") == "nvfp4-pack-quantized" for g in quant_groups.values() if isinstance(g, dict))
|
||||
)
|
||||
nvfp4_compressed_tensors = self._is_nvfp4_compressed_tensors(quant_method, quant_format, quant_groups)
|
||||
|
||||
self._nvfp4_global_algo = quant_algo
|
||||
|
||||
|
||||
+2
-2
@@ -849,8 +849,8 @@ class Gemma4DSparkModel(DFlashModel):
|
||||
raise ValueError("Gemma4 DSpark attention bias and MoE are not supported")
|
||||
if (self.hparams.get("draft_vocab_size") or self.hparams["vocab_size"]) != self.hparams["vocab_size"]:
|
||||
raise ValueError("Gemma4 DSpark currently requires a full draft vocabulary")
|
||||
if "model.lm_head.weight" not in self.model_tensors and self.hparams.get("tie_word_embeddings") is not True:
|
||||
raise ValueError("Gemma4 DSpark requires lm_head.weight unless tie_word_embeddings is true")
|
||||
if "model.lm_head.weight" not in self.model_tensors:
|
||||
raise ValueError("Gemma4 DSpark requires lm_head.weight")
|
||||
|
||||
self.dflash_config = self.hparams.get("dflash_config", {})
|
||||
markov_type = self.dflash_config.get("markov_head_type", self.hparams.get("markov_head_type", "vanilla"))
|
||||
|
||||
@@ -164,11 +164,11 @@ export ZENDNNL_MATMUL_ALGO=1 # Blocked AOCL DLP algo for best performance
|
||||
./build/bin/llama-server \
|
||||
-m models/Llama-3.1-8B-Instruct.BF16.gguf \
|
||||
--host 0.0.0.0 \
|
||||
--port 8080 \
|
||||
--port 9931 \
|
||||
-t 64
|
||||
```
|
||||
|
||||
Access the server at `http://localhost:8080`.
|
||||
Access the server at `http://localhost:9931`.
|
||||
|
||||
**Performance tips**:
|
||||
- Use `ZENDNNL_MATMUL_ALGO=1` for optimal performance
|
||||
|
||||
+1
-1
@@ -351,7 +351,7 @@ cmake --build build --config Release
|
||||
|
||||
#### Override Compute Capability Specifications
|
||||
|
||||
By default, all supported compute capabilities are enabled. To customize this behavior, you can specify the `MUSA_ARCHITECTURES` option in the CMake command:
|
||||
By default, compute capabilities `2.2` (MTT S4000) and `3.1` (MTT S5000) are enabled, compute capability `2.1` (MTT S70, MTT S80, MTT S3000) is deprecated and has to be enabled explicitly. To customize this behavior, you can specify the `MUSA_ARCHITECTURES` option in the CMake command:
|
||||
|
||||
```bash
|
||||
cmake -B build -DGGML_MUSA=ON -DMUSA_ARCHITECTURES="31"
|
||||
|
||||
@@ -282,7 +282,7 @@ This table can be generated with:
|
||||
|
||||
# Usage - need tool-aware Jinja template
|
||||
|
||||
First, start a server with any model, but make sure it has a tools-enabled template: you can verify this by inspecting the `chat_template` or `chat_template_tool_use` properties in `http://localhost:8080/props`).
|
||||
First, start a server with any model, but make sure it has a tools-enabled template: you can verify this by inspecting the `chat_template` or `chat_template_tool_use` properties in `http://localhost:9931/props`).
|
||||
|
||||
Here are some models known to work (w/ chat template override when needed):
|
||||
|
||||
@@ -336,7 +336,7 @@ To get the official template from original HuggingFace repos, you can use [scrip
|
||||
Test in CLI (or with any library / software that can use OpenAI-compatible API backends):
|
||||
|
||||
```bash
|
||||
curl http://localhost:8080/v1/chat/completions -d '{
|
||||
curl http://localhost:9931/v1/chat/completions -d '{
|
||||
"model": "gpt-3.5-turbo",
|
||||
"tools": [
|
||||
{
|
||||
@@ -366,7 +366,7 @@ curl http://localhost:8080/v1/chat/completions -d '{
|
||||
}'
|
||||
|
||||
|
||||
curl http://localhost:8080/v1/chat/completions -d '{
|
||||
curl http://localhost:9931/v1/chat/completions -d '{
|
||||
"model": "gpt-3.5-turbo",
|
||||
"messages": [
|
||||
{"role": "system", "content": "You are a chatbot that uses tools/functions. Dont overthink things."},
|
||||
|
||||
@@ -10,7 +10,7 @@ import json, requests
|
||||
|
||||
if True:
|
||||
|
||||
def create_completion(*, response_model=None, endpoint="http://localhost:8080/v1/chat/completions", messages, **kwargs):
|
||||
def create_completion(*, response_model=None, endpoint="http://localhost:9931/v1/chat/completions", messages, **kwargs):
|
||||
'''
|
||||
Creates a chat completion using an OpenAI-compatible endpoint w/ JSON schema support
|
||||
(llama.cpp server, llama-cpp-python, Anyscale / Together...)
|
||||
@@ -45,7 +45,7 @@ else:
|
||||
#! pip install instructor openai
|
||||
import instructor, openai
|
||||
client = instructor.patch(
|
||||
openai.OpenAI(api_key="123", base_url="http://localhost:8080"),
|
||||
openai.OpenAI(api_key="123", base_url="http://localhost:9931"),
|
||||
mode=instructor.Mode.JSON_SCHEMA)
|
||||
create_completion = client.chat.completions.create
|
||||
|
||||
|
||||
@@ -10,4 +10,4 @@ Recommended way to run this model:
|
||||
llama-server -hf {namespace}/{model_name}-GGUF
|
||||
```
|
||||
|
||||
Then, access http://localhost:8080
|
||||
Then, access http://localhost:9931
|
||||
|
||||
@@ -10,11 +10,11 @@ Recommended way to run this model:
|
||||
llama-server -hf {namespace}/{model_name}-GGUF --embeddings
|
||||
```
|
||||
|
||||
Then the endpoint can be accessed at http://localhost:8080/embedding, for
|
||||
Then the endpoint can be accessed at http://localhost:9931/embedding, for
|
||||
example using `curl`:
|
||||
```console
|
||||
curl --request POST \
|
||||
--url http://localhost:8080/embedding \
|
||||
--url http://localhost:9931/embedding \
|
||||
--header "Content-Type: application/json" \
|
||||
--data '{{"input": "Hello embeddings"}}' \
|
||||
--silent
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
#!/usr/bin/env bash
|
||||
curl --request POST \
|
||||
--url http://localhost:8080/embedding \
|
||||
--url http://localhost:9931/embedding \
|
||||
--header "Content-Type: application/json" \
|
||||
--data '{"input": "Hello world today"}' \
|
||||
--silent
|
||||
|
||||
@@ -295,7 +295,7 @@ def example_concurrent(host):
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description=sys.modules[__name__].__doc__)
|
||||
parser.add_argument("--host", default="localhost:8080", help="llama.cpp server")
|
||||
parser.add_argument("--host", default="localhost:9931", help="llama.cpp server")
|
||||
parser.add_argument("-v", "--verbose", action="store_true", help="enables logging")
|
||||
args = parser.parse_args()
|
||||
logging.basicConfig(level=logging.INFO if args.verbose else logging.ERROR)
|
||||
|
||||
@@ -206,7 +206,7 @@ int main(int argc, char ** argv) {
|
||||
// reset the draft context to the checkpoint before verification
|
||||
if (ctx_dft) {
|
||||
if (use_ckpt_dft) {
|
||||
ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
|
||||
GGML_ASSERT(ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
|
||||
}
|
||||
|
||||
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, ckpt.pos_max + 1, -1);
|
||||
@@ -269,13 +269,13 @@ int main(int argc, char ** argv) {
|
||||
draft = std::move(ids);
|
||||
|
||||
{
|
||||
ckpt.load_tgt(ctx_tgt, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
|
||||
GGML_ASSERT(ckpt.load_tgt(ctx_tgt, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
|
||||
|
||||
llama_memory_seq_rm(llama_get_memory(ctx_tgt), seq_id, ckpt.pos_max + 1, -1);
|
||||
}
|
||||
|
||||
if (ctx_dft) {
|
||||
ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY);
|
||||
GGML_ASSERT(ckpt.load_dft(ctx_dft, seq_id, LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY));
|
||||
|
||||
llama_memory_seq_rm(llama_get_memory(ctx_dft), seq_id, ckpt.pos_max + 1, -1);
|
||||
}
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -2,7 +2,8 @@
|
||||
|
||||
#ifdef GGML_CUDA_USE_CUB
|
||||
# include <cub/cub.cuh>
|
||||
# if (CCCL_MAJOR_VERSION >= 3 && CCCL_MINOR_VERSION >= 1)
|
||||
// strided_iterator was added in CCCL 3.1
|
||||
# if (CCCL_MAJOR_VERSION > 3 || (CCCL_MAJOR_VERSION == 3 && CCCL_MINOR_VERSION >= 1))
|
||||
# define STRIDED_ITERATOR_AVAILABLE
|
||||
# include <cuda/iterator>
|
||||
# endif
|
||||
@@ -27,21 +28,21 @@ static __global__ void init_offsets(int * offsets, const int ncols, const int nr
|
||||
}
|
||||
#endif // STRIDED_ITERATOR_AVAILABLE
|
||||
|
||||
#ifdef GGML_CUDA_USE_CUB
|
||||
|
||||
// returns the suggested maximum number of rows to process during one argsort_f32_i32_cuda_cub() call
|
||||
int argsort_f32_i32_cuda_cub_chunk_nrows(const size_t nb01, const int64_t nrows) {
|
||||
// perform argsort in chunks up to approximately this size (currently 64MB)
|
||||
// returns the suggested maximum number of rows to process at once, given the temporary buffer bytes per row
|
||||
int ggml_cuda_chunk_nrows(const size_t row_bytes, const int64_t nrows) {
|
||||
// process rows in chunks up to approximately this size (currently 64MB)
|
||||
// to avoid excessive temporary buffers memory usage
|
||||
const int chunk_bytes = 1 << 26;
|
||||
|
||||
// calculate how many rows will fit in one chunk (must be at least one)
|
||||
const int chunk_nrows = std::max((int) (chunk_bytes / nb01), 1);
|
||||
const int chunk_nrows = std::max((int) (chunk_bytes / row_bytes), 1);
|
||||
|
||||
// limit the resulting amount to total nrows
|
||||
return std::min((int64_t) chunk_nrows, nrows);
|
||||
}
|
||||
|
||||
#ifdef GGML_CUDA_USE_CUB
|
||||
|
||||
void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
|
||||
const float * x,
|
||||
int * dst,
|
||||
@@ -289,7 +290,7 @@ void ggml_cuda_op_argsort(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int chunk_nrows = argsort_f32_i32_cuda_cub_chunk_nrows(src0->nb[1], nrows);
|
||||
const int chunk_nrows = ggml_cuda_chunk_nrows(src0->nb[1], nrows);
|
||||
|
||||
ggml_cuda_pool & pool = ctx.pool();
|
||||
|
||||
|
||||
@@ -4,8 +4,9 @@
|
||||
|
||||
void ggml_cuda_op_argsort(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
|
||||
|
||||
int ggml_cuda_chunk_nrows(const size_t row_bytes, const int64_t nrows);
|
||||
|
||||
#ifdef GGML_CUDA_USE_CUB
|
||||
int argsort_f32_i32_cuda_cub_chunk_nrows(const size_t nb01, const int64_t nrows);
|
||||
void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool,
|
||||
const float * x,
|
||||
int * dst,
|
||||
|
||||
@@ -200,9 +200,12 @@ static bool ggml_cuda_op_fwht_impl(ggml_backend_cuda_context & ctx, const ggml_t
|
||||
case 4096:
|
||||
ggml_cuda_kernel_launch(fwht_cuda_block<4096, nt, T>, launch_params_w, src_d, dst_d, rows, scale);
|
||||
return true;
|
||||
#if !defined(GGML_USE_MUSA)
|
||||
// 32 KB of shared memory, above the MUSA limit; falls back there
|
||||
case 8192:
|
||||
ggml_cuda_kernel_launch(fwht_cuda_block<8192, nt, T>, launch_params_w, src_d, dst_d, rows, scale);
|
||||
return true;
|
||||
#endif // !defined(GGML_USE_MUSA)
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
@@ -5689,11 +5776,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
case GGML_OP_SUM:
|
||||
return ggml_is_contiguous_rows(op->src[0]);
|
||||
case GGML_OP_TOP_K:
|
||||
#if defined(GGML_USE_HIP) || defined(GGML_CUDA_USE_CUB)
|
||||
return true;
|
||||
#else
|
||||
return op->src[0]->ne[0] <= 1024;
|
||||
#endif // defined(GGML_USE_HIP) || defined(GGML_CUDA_USE_CUB)
|
||||
return op->src[0]->ne[0] <= INT_MAX;
|
||||
case GGML_OP_ARGSORT:
|
||||
#ifndef GGML_CUDA_USE_CUB
|
||||
{
|
||||
@@ -5705,7 +5788,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
return ncols_pad * sizeof(int) <= ggml_cuda_info().devices[dev_ctx->device].smpb;
|
||||
}
|
||||
#else
|
||||
return true;
|
||||
return op->src[0]->ne[0] <= INT_MAX;
|
||||
#endif
|
||||
case GGML_OP_SUM_ROWS:
|
||||
return op->src[0]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && ggml_is_contiguous_rows(op->src[0]);
|
||||
|
||||
@@ -7,9 +7,9 @@
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q1_0(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -98,9 +98,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q2_0(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -187,9 +187,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_0(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -250,9 +250,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_1(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -313,9 +313,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_0(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -393,9 +393,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_1(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -471,9 +471,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q8_0(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -537,9 +537,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q2_K(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -598,9 +598,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q3_K(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -711,9 +711,9 @@ static __device__ __forceinline__ int unpack_scales_q45_K(const int * scales, co
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q4_K(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -822,9 +822,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q5_K(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -946,9 +946,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_q6_K(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1036,9 +1036,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq1_s(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1098,9 +1098,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_xxs(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1162,9 +1162,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_xs(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1227,9 +1227,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq2_s(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1295,9 +1295,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq3_xxs(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1359,9 +1359,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq3_s(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1428,9 +1428,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq4_xs(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1495,9 +1495,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_iq4_nl(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1564,9 +1564,9 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_mxfp4(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kbx0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
int * x_qs = (int *) x_tile;
|
||||
@@ -1670,7 +1670,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
}
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4(
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_load_tiles_nvfp4(
|
||||
const char * __restrict__ x, int * __restrict__ x_tile, const int kb0, const int i_max, const int stride) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
|
||||
|
||||
@@ -10,8 +10,8 @@ using namespace ggml_cuda_mma;
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_0_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_0, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -60,8 +60,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_1_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_1, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -110,8 +110,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q8_0, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -148,8 +148,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
|
||||
typedef tile<16, 8, int, input_layout> tile_B;
|
||||
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -203,8 +203,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
|
||||
typedef tile< 8, 8, int> tile_B;
|
||||
typedef tile<16, 8, int> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -281,8 +281,8 @@ static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma(
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q5_1, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -318,8 +318,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile<16, 8, int, input_layout> tile_B;
|
||||
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -368,8 +368,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile< 8, 8, int> tile_B;
|
||||
typedef tile<16, 8, int> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -442,8 +442,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(type, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -474,7 +474,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
}
|
||||
|
||||
// Used for Q3_K, IQ2_S, and IQ2_XS:
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma(
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
#if defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
constexpr data_layout input_layout = get_input_data_layout();
|
||||
@@ -483,7 +483,7 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
|
||||
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -533,7 +533,7 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
|
||||
typedef tile<16, 8, int> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, prec_src1);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -610,8 +610,8 @@ template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q2_K_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q2_K, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -680,8 +680,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile<16, 4, int, input_layout> tile_B;
|
||||
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -749,8 +749,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile< 8, 4, int> tile_B;
|
||||
typedef tile<16, 8, int> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -870,8 +870,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q3_K_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q3_K, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -905,8 +905,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q4_K_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q4_K, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -940,8 +940,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q5_K_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q5_K, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -975,8 +975,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_vec_dot_q6_K_q8_1_dp4a(
|
||||
const int * __restrict__ x, const int * __restrict__ y, float * __restrict__ sum, const int k00) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, GGML_PREC_Q8) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, GGML_PREC_Q8);
|
||||
|
||||
constexpr tile_x_sizes txs = mmq_get_dp4a_tile_x_sizes(GGML_TYPE_Q6_K, I);
|
||||
const int * x_qs = (const int *) x;
|
||||
@@ -1015,8 +1015,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile<16, 4, int, input_layout> tile_B;
|
||||
typedef tile<16, 16, int, DATA_LAYOUT_J_MAJOR> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -1066,8 +1066,8 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile< 8, 4, int> tile_B;
|
||||
typedef tile<16, 8, int> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q8);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
y += (threadIdx.y % ntx) * (tile_C::J*MMQ_TILE_Y_K);
|
||||
@@ -1181,7 +1181,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
typedef tile<16, 8, float> tile_C;
|
||||
|
||||
constexpr int sram_stride = ggml_cuda_mmq_get_sram_stride(type, J, fallback, GGML_PREC_Q4);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, GGML_PREC_Q4);
|
||||
constexpr int ntx = rows_per_warp / tile_C::I;
|
||||
constexpr int nfrags = MMQ_TILE_NE_K / tile_A::J;
|
||||
|
||||
|
||||
+76
-38
@@ -8,66 +8,66 @@
|
||||
static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream, const ggml_prec prec_src1) {
|
||||
switch (args.type_x) {
|
||||
case GGML_TYPE_Q1_0:
|
||||
mul_mat_q_case<GGML_TYPE_Q1_0>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q1_0, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q2_0:
|
||||
mul_mat_q_case<GGML_TYPE_Q2_0>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q2_0, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q4_0:
|
||||
mul_mat_q_case<GGML_TYPE_Q4_0>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q4_0, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q4_1:
|
||||
mul_mat_q_case<GGML_TYPE_Q4_1>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q4_1, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q5_0:
|
||||
mul_mat_q_case<GGML_TYPE_Q5_0>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q5_0, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q5_1:
|
||||
mul_mat_q_case<GGML_TYPE_Q5_1>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q5_1, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q8_0:
|
||||
mul_mat_q_case<GGML_TYPE_Q8_0>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q8_0, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
// -----------------------------------------------------------------------
|
||||
case GGML_TYPE_Q2_K:
|
||||
mul_mat_q_case<GGML_TYPE_Q2_K>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q2_K, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q3_K:
|
||||
mul_mat_q_case<GGML_TYPE_Q3_K>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q3_K, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q4_K:
|
||||
mul_mat_q_case<GGML_TYPE_Q4_K>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q4_K, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q5_K:
|
||||
mul_mat_q_case<GGML_TYPE_Q5_K>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q5_K, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_Q6_K:
|
||||
mul_mat_q_case<GGML_TYPE_Q6_K>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_Q6_K, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
// -----------------------------------------------------------------------
|
||||
case GGML_TYPE_IQ1_S:
|
||||
mul_mat_q_case<GGML_TYPE_IQ1_S>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_IQ1_S, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_IQ2_XXS:
|
||||
mul_mat_q_case<GGML_TYPE_IQ2_XXS>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_IQ2_XXS, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_IQ2_XS:
|
||||
mul_mat_q_case<GGML_TYPE_IQ2_XS>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_IQ2_XS, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_IQ2_S:
|
||||
mul_mat_q_case<GGML_TYPE_IQ2_S>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_IQ2_S, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_IQ3_XXS:
|
||||
mul_mat_q_case<GGML_TYPE_IQ3_XXS>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_IQ3_XXS, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_IQ3_S:
|
||||
mul_mat_q_case<GGML_TYPE_IQ3_S>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_IQ3_S, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_IQ4_XS:
|
||||
mul_mat_q_case<GGML_TYPE_IQ4_XS>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_IQ4_XS, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_IQ4_NL:
|
||||
mul_mat_q_case<GGML_TYPE_IQ4_NL>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_IQ4_NL, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
// -----------------------------------------------------------------------
|
||||
case GGML_TYPE_MXFP4:
|
||||
@@ -76,14 +76,14 @@ static void ggml_cuda_mul_mat_q_switch_type(ggml_backend_cuda_context & ctx, con
|
||||
mul_mat_q_case<GGML_TYPE_MXFP4, GGML_PREC_Q4>(ctx, args, stream);
|
||||
break;
|
||||
}
|
||||
mul_mat_q_case<GGML_TYPE_MXFP4>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_MXFP4, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
case GGML_TYPE_NVFP4:
|
||||
if (prec_src1 == GGML_PREC_Q4) {
|
||||
mul_mat_q_case<GGML_TYPE_NVFP4, GGML_PREC_Q4>(ctx, args, stream);
|
||||
break;
|
||||
}
|
||||
mul_mat_q_case<GGML_TYPE_NVFP4>(ctx, args, stream);
|
||||
mul_mat_q_case<GGML_TYPE_NVFP4, GGML_PREC_Q8>(ctx, args, stream);
|
||||
break;
|
||||
default:
|
||||
GGML_ABORT("fatal error");
|
||||
@@ -141,7 +141,10 @@ void ggml_cuda_mul_mat_q(
|
||||
GGML_TENSOR_BINARY_OP_LOCALS;
|
||||
|
||||
cudaStream_t stream = ctx.stream();
|
||||
const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc;
|
||||
|
||||
const int id = ggml_cuda_get_device();
|
||||
const int cc = ggml_cuda_info().devices[id].cc;
|
||||
const size_t smpbo = ggml_cuda_info().devices[id].smpbo;
|
||||
|
||||
const size_t ts_src0 = ggml_type_size(src0->type);
|
||||
const size_t ts_src1 = ggml_type_size(src1->type);
|
||||
@@ -176,7 +179,7 @@ void ggml_cuda_mul_mat_q(
|
||||
const int64_t s03 = src0->nb[3] / ts_src0;
|
||||
const int64_t s3 = dst->nb[3] / ts_dst;
|
||||
|
||||
const bool fallback = ne01 % 128 != 0;
|
||||
const bool fallback = ggml_cuda_mmq_needs_fallback(ne01);
|
||||
|
||||
const ggml_prec prec_src1 = ggml_cuda_mmq_get_prec_src1(src0, dst, cc);
|
||||
|
||||
@@ -184,9 +187,52 @@ void ggml_cuda_mul_mat_q(
|
||||
const size_t y_block_size = use_native_fp4 ? sizeof(block_fp4_mmq) : sizeof(block_q8_1_mmq);
|
||||
const size_t y_values_per_block = use_native_fp4 ? QK_FP4_MMQ : QK8_1_MMQ;
|
||||
|
||||
int J_best = 0;
|
||||
int nthreads_best = 0;
|
||||
{
|
||||
int64_t ncols_opt = ne11;
|
||||
if (ids) {
|
||||
const int64_t n_expert_used = ids->ne[0];
|
||||
ncols_opt = ne12;
|
||||
|
||||
// Each expert only sees ne12*n_expert_used/ne02 tokens on average.
|
||||
// On RDNA3 and RDNA4 it is faster to pick the tile size against this value instead of ne12.
|
||||
if (GGML_CUDA_CC_IS_RDNA3(cc) || GGML_CUDA_CC_IS_RDNA4(cc)) {
|
||||
ncols_opt = (ne12*n_expert_used + ne02 - 1) / ne02;
|
||||
}
|
||||
}
|
||||
|
||||
int ntiles_J_best = INT_MAX;
|
||||
|
||||
for (int J = 8; J <= 128 && ntiles_J_best > 1; J += 8) {
|
||||
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(src0->type, J, fallback, cc, prec_src1);
|
||||
if (config.type == GGML_TYPE_COUNT) {
|
||||
continue;
|
||||
}
|
||||
|
||||
if (mmq_get_nbytes_shared(config, cc) > smpbo) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int ntiles_x = (ncols_opt + config.J - 1) / config.J;
|
||||
|
||||
if (ntiles_x < ntiles_J_best) {
|
||||
J_best = J;
|
||||
nthreads_best = config.nthreads;
|
||||
ntiles_J_best = ntiles_x;
|
||||
}
|
||||
}
|
||||
}
|
||||
GGML_ASSERT(J_best > 0);
|
||||
|
||||
// A tile of size J can read in at most J - 1 extra columns.
|
||||
// For simplicity, round up the padding of a full tile to a multiple of the number of bytes that nthreads can load in parallel.
|
||||
const size_t src1_load_chunk_size = nthreads_best * sizeof(int);
|
||||
const size_t src1_q8_1_padding = ((J_best * sizeof(block_q8_1_mmq) + src1_load_chunk_size - 1) / src1_load_chunk_size)
|
||||
* src1_load_chunk_size;
|
||||
|
||||
if (!ids) {
|
||||
const size_t nbytes_src1_q8_1 = ne13*ne12 * ne11*ne10_padded * y_block_size/y_values_per_block +
|
||||
ggml_cuda_mmq_get_J_max(src0->type, fallback, cc, ne11) * sizeof(block_q8_1_mmq);
|
||||
const size_t nbytes_src1_q8_1 = ne13*ne12 * ne11*ne10_padded * y_block_size/y_values_per_block + src1_q8_1_padding;
|
||||
ggml_cuda_pool_alloc<char> src1_q8_1(ctx.pool(), nbytes_src1_q8_1);
|
||||
ggml_cuda_pool_alloc<float> src1_scale(ctx.pool());
|
||||
if (src0->type == GGML_TYPE_NVFP4 && use_native_fp4) {
|
||||
@@ -223,7 +269,7 @@ void ggml_cuda_mul_mat_q(
|
||||
ne00, ne01, ne1, s01, ne11, s1,
|
||||
ne02, ne12, s02, s12, s2,
|
||||
ne03, ne13, s03, s13, s3,
|
||||
ne1, ne1};
|
||||
ne1, J_best};
|
||||
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream, prec_src1);
|
||||
return;
|
||||
}
|
||||
@@ -237,7 +283,7 @@ void ggml_cuda_mul_mat_q(
|
||||
GGML_ASSERT(ne1 == n_expert_used);
|
||||
|
||||
ggml_cuda_pool_alloc<int32_t> ids_src1(ctx.pool(), ne_get_rows);
|
||||
ggml_cuda_pool_alloc<int32_t> ids_dst(ctx.pool(), ne_get_rows);
|
||||
ggml_cuda_pool_alloc<int32_t> ids_dst(ctx.pool(), ne_get_rows + J_best-1); // Needs to be padded for unconditional memory access.
|
||||
ggml_cuda_pool_alloc<int32_t> expert_bounds(ctx.pool(), ne02 + 1);
|
||||
|
||||
// gate/up activations are broadcast across experts (ne11 == 1): quantize each token once and
|
||||
@@ -254,8 +300,7 @@ void ggml_cuda_mul_mat_q(
|
||||
CUDA_CHECK(cudaGetLastError());
|
||||
}
|
||||
|
||||
const size_t nbytes_src1_q8_1 = ne12*n_expert_used*ne10_padded * y_block_size/y_values_per_block +
|
||||
ggml_cuda_mmq_get_J_max(src0->type, fallback, cc, ne12) * sizeof(block_q8_1_mmq);
|
||||
const size_t nbytes_src1_q8_1 = ne12*n_expert_used*ne10_padded * y_block_size/y_values_per_block + src1_q8_1_padding;
|
||||
ggml_cuda_pool_alloc<char> src1_q8_1(ctx.pool(), nbytes_src1_q8_1);
|
||||
ggml_cuda_pool_alloc<float> src1_scale(ctx.pool());
|
||||
if (src0->type == GGML_TYPE_NVFP4 && use_native_fp4) {
|
||||
@@ -296,13 +341,6 @@ void ggml_cuda_mul_mat_q(
|
||||
ne11 * ne10_padded * sizeof(block_q8_1) / (QK8_1 * sizeof(int));
|
||||
const int64_t s13 = ne12*s12;
|
||||
|
||||
// Each expert only sees ne12*n_expert_used/ne02 tokens on average.
|
||||
// On RDNA3 and RDNA4 it is faster to pick the tile size against this value instead of ne12.
|
||||
int64_t ncols_opt = ne12;
|
||||
if (GGML_CUDA_CC_IS_RDNA3(cc) || GGML_CUDA_CC_IS_RDNA4(cc)) {
|
||||
ncols_opt = (ne12*n_expert_used + ne02 - 1) / ne02;
|
||||
}
|
||||
|
||||
// Note that ne02 is used instead of ne12 because the number of y channels determines the z dimension of the CUDA grid.
|
||||
const mmq_args args = {
|
||||
src0_d, src0->type, (const int *) src1_q8_1.get(), ids_dst.get(), expert_bounds.get(), dst_d,
|
||||
@@ -310,7 +348,7 @@ void ggml_cuda_mul_mat_q(
|
||||
ne00, ne01, ne_get_rows, s01, ne_get_rows, s1,
|
||||
ne02, ne02, s02, s12, s2,
|
||||
ne03, ne13, s03, s13, s3,
|
||||
ne12, ncols_opt};
|
||||
ne12, J_best};
|
||||
|
||||
ggml_cuda_mul_mat_q_switch_type(ctx, args, stream, prec_src1);
|
||||
}
|
||||
|
||||
+106
-138
@@ -208,7 +208,7 @@ struct ggml_cuda_mmq_config {
|
||||
static_assert((nthreads_) % 32 == 0 && (nthreads_) <= 512, "bad nthreads"); \
|
||||
static_assert( (occupancy_) <= 8, "bad occupancy"); \
|
||||
static_assert((I_) % 32 == 0, "bad I"); \
|
||||
static_assert((J_) % 8 == 0, "bad J"); \
|
||||
static_assert((J_) % 8 == 0 && (J_) <= 128, "bad J"); \
|
||||
static_assert((K_vram_) % 256 == 0, "bad K_vram"); \
|
||||
return ggml_cuda_mmq_config((type_), (nthreads_), (occupancy_), (I_), (J_), (sram_layout_), (K_vram_), (stream_k_), (fallback_)); \
|
||||
} \
|
||||
@@ -227,7 +227,7 @@ struct ggml_cuda_mmq_config {
|
||||
|
||||
#undef CASE
|
||||
|
||||
static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
if (GGML_CUDA_CC_IS_AMD(cc)) {
|
||||
if (GGML_CUDA_CC_IS_GCN(cc)) {
|
||||
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
|
||||
@@ -262,7 +262,7 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty
|
||||
return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback);
|
||||
}
|
||||
|
||||
static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
#ifdef GGML_USE_HIP
|
||||
#ifdef GCN
|
||||
return ggml_cuda_mmq_get_config_gcn(type, J, fallback);
|
||||
@@ -295,93 +295,86 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t
|
||||
GGML_UNUSED_VARS(type, J, fallback, prec_src1);
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_type(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).type;
|
||||
static __host__ int ggml_cuda_mmq_get_type(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).type;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_type(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).type;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_nthreads(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).nthreads;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_occupancy(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).occupancy;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_I(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).I;
|
||||
static __host__ int ggml_cuda_mmq_get_I(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).I;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_I(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).I;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_J(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).J;
|
||||
static __host__ int ggml_cuda_mmq_get_J(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).J;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_J(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).J;
|
||||
}
|
||||
|
||||
static __host__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).sram_layout;
|
||||
static __host__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).sram_layout;
|
||||
}
|
||||
|
||||
static constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ ggml_cuda_mmq_sram_layout ggml_cuda_mmq_get_sram_layout(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).sram_layout;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_K_vram(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).K_vram;
|
||||
static __host__ int ggml_cuda_mmq_get_K_vram(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).K_vram;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_K_vram(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).K_vram;
|
||||
}
|
||||
|
||||
static __host__ bool ggml_cuda_mmq_get_stream_k(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).stream_k;
|
||||
static __host__ bool ggml_cuda_mmq_get_stream_k(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).stream_k;
|
||||
}
|
||||
|
||||
static constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ bool ggml_cuda_mmq_get_stream_k(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).stream_k;
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_fallback(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc).fallback;
|
||||
static __host__ int ggml_cuda_mmq_get_fallback(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1).fallback;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_fallback(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).fallback;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_sram_stride(const ggml_type type, const int J, const bool fallback, const int cc) {
|
||||
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, cc));
|
||||
static __host__ int ggml_cuda_mmq_get_sram_stride(const ggml_type type, const int J, const bool fallback, const int cc, const ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, cc, prec_src1));
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8) {
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_sram_stride(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_sram_stride(ggml_cuda_mmq_get_sram_layout(type, J, fallback, prec_src1));
|
||||
}
|
||||
|
||||
static __host__ int ggml_cuda_mmq_get_J_max(const ggml_type type, const bool fallback, const int cc, const int64_t ne11) {
|
||||
int ret = std::min(ne11, int64_t(512));
|
||||
ret -= ret % 8;
|
||||
for (;ret > 0; ret -= 8) {
|
||||
if (ggml_cuda_mmq_get_config(type, ret, fallback, cc).type != GGML_TYPE_COUNT) {
|
||||
return ret;
|
||||
}
|
||||
}
|
||||
return ret;
|
||||
static __host__ bool ggml_cuda_mmq_needs_fallback(const int64_t nrows_x) {
|
||||
return nrows_x % 128 != 0;
|
||||
}
|
||||
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_rows_per_warp(ggml_type type, int J, bool fallback) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback).rows_per_warp();
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_rows_per_warp(ggml_type type, int J, bool fallback, ggml_prec prec_src1) {
|
||||
return ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).rows_per_warp();
|
||||
}
|
||||
|
||||
#define MMQ_DP4A_TXS_Q4_0 tile_x_sizes{I*MMQ_TILE_NE_K + I, I*MMQ_TILE_NE_K/QI4_0 + I/QI4_0, 0}
|
||||
@@ -437,12 +430,12 @@ static __host__ int ggml_cuda_mmq_get_nbytes_shared_x(const ggml_cuda_mmq_config
|
||||
#include "mmq-load-tiles.cuh"
|
||||
#include "mmq-vec-dot.cuh"
|
||||
|
||||
template <ggml_type type, int J, bool fallback> static __device__ __forceinline__ void ggml_cuda_mmq_write_back_dp4a(
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1> static __device__ __forceinline__ void ggml_cuda_mmq_write_back_dp4a(
|
||||
const float * __restrict__ sum, const int32_t * __restrict__ ids_dst, float * __restrict__ dst,
|
||||
const float * __restrict__ y_scale, const int stride, const int i_max, const int j_max) {
|
||||
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback);
|
||||
constexpr int nwarps = ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1) / warp_size;
|
||||
constexpr int I = ggml_cuda_mmq_get_I(type, J, fallback, prec_src1);
|
||||
|
||||
const bool y_scale_used = y_scale != nullptr;
|
||||
|
||||
@@ -476,7 +469,7 @@ template <ggml_type type, int J, bool fallback> static __device__ __forceinline_
|
||||
}
|
||||
}
|
||||
|
||||
template<ggml_type type, int J, bool fallback>
|
||||
template<ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
static __device__ __forceinline__ void ggml_cuda_mmq_write_back_mma(
|
||||
const float * __restrict__ sum, const int * __restrict__ ids_dst, float * __restrict__ dst,
|
||||
const float * __restrict__ y_scale, const int stride, const int i_max, const int j_max) {
|
||||
@@ -487,7 +480,7 @@ static __device__ __forceinline__ void ggml_cuda_mmq_write_back_mma(
|
||||
typedef tile<16, 8, int> tile_C;
|
||||
#endif // defined(AMD_MFMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE)
|
||||
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback);
|
||||
constexpr int rows_per_warp = ggml_cuda_mmq_get_rows_per_warp(type, J, fallback, prec_src1);
|
||||
constexpr int ntx = rows_per_warp/tile_C::I; // Number of x minitiles per warp.
|
||||
|
||||
const int i0 = (threadIdx.y / ntx) * (ntx*tile_C::I);
|
||||
@@ -541,7 +534,7 @@ struct ggml_cuda_mmq_util_funcs {
|
||||
vdr(vdr), load_tiles(load_tiles), vec_dot(vec_dot), write_back(write_back) {}
|
||||
};
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_funcs() {
|
||||
if (!ggml_cuda_mmq_get_config(type, J, fallback, prec_src1).use_mma_data_layout()) {
|
||||
switch (type) {
|
||||
@@ -550,136 +543,136 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
|
||||
VDR_Q1_0_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q1_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q2_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q2_0_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q2_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q4_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q4_0_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q4_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q4_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q4_1:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q4_1_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q4_1<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q4_1_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q5_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q5_0_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q5_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q5_1:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q5_1_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q5_1<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q8_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q8_0_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q8_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
case GGML_TYPE_Q2_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q2_K_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q2_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q2_K_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q3_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q3_K_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q3_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q3_K_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q4_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q4_K_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q4_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q4_K_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q5_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q5_K_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q5_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q5_K_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q6_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_Q6_K_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_q6_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q6_K_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
case GGML_TYPE_IQ1_S:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_IQ1_S_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_iq1_s<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_1_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ2_XXS:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_IQ2_XXS_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_iq2_xxs<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ2_XS:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_IQ2_XS_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_iq2_xs<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ2_S:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_IQ2_S_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_iq2_s<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ3_XXS:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_IQ3_XXS_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_iq3_xxs<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ3_S:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_IQ3_S_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_iq3_s<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ4_XS:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_IQ4_XS_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_iq4_xs<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ4_NL:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_IQ4_NL_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_iq4_nl<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
case GGML_TYPE_MXFP4:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_MXFP4_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_mxfp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_NVFP4:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
VDR_NVFP4_Q8_1_MMQ,
|
||||
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback, prec_src1>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_dp4a<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_dp4a<type, J, fallback, prec_src1>);
|
||||
default:
|
||||
return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr);
|
||||
}
|
||||
@@ -695,7 +688,7 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_mxfp4_fp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
}
|
||||
break;
|
||||
case GGML_TYPE_NVFP4:
|
||||
@@ -704,7 +697,7 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_nvfp4_nvfp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_fp4_fp4_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
}
|
||||
break;
|
||||
default:
|
||||
@@ -720,164 +713,164 @@ static constexpr __device__ ggml_cuda_mmq_util_funcs ggml_cuda_mmq_get_util_func
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q1_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q2_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q2_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q4_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q4_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_DS4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q4_1:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q4_1<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q5_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q5_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q5_1:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q5_1<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q8_0:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q8_0<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
case GGML_TYPE_Q2_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q2_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q2_K_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q3_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q3_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q4_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q4_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q5_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q5_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_Q6_K:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_q6_K<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q6_K_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
case GGML_TYPE_IQ1_S:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_iq1_s<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_1_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ2_XXS:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_iq2_xxs<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ2_XS:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_iq2_xs<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ2_S:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_iq2_s<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ3_XXS:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_iq3_xxs<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ3_S:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_iq3_s<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ4_XS:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_iq4_xs<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_IQ4_NL:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_iq4_nl<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
case GGML_TYPE_MXFP4:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_mxfp4<type, J, fallback>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_q8_1_mma<type, J, fallback, MMQ_Q8_1_DS_LAYOUT_D4>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
case GGML_TYPE_NVFP4:
|
||||
return ggml_cuda_mmq_util_funcs(
|
||||
-1,
|
||||
ggml_cuda_mmq_load_tiles_nvfp4<type, J, fallback, prec_src1>,
|
||||
ggml_cuda_mmq_vec_dot_q8_0_16_q8_1_mma<type, J, fallback, prec_src1>,
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback>);
|
||||
ggml_cuda_mmq_write_back_mma<type, J, fallback, prec_src1>);
|
||||
default:
|
||||
return ggml_cuda_mmq_util_funcs(1, nullptr, nullptr, nullptr);
|
||||
}
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
static constexpr __device__ int ggml_cuda_mmq_get_vdr() {
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vdr;
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
static constexpr __device__ ggml_cuda_mmq_load_tiles_t ggml_cuda_mmq_get_load_tiles() {
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().load_tiles;
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
static constexpr __device__ ggml_cuda_mmq_vec_dot_t ggml_cuda_mmq_get_vec_dot() {
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().vec_dot;
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
static constexpr __device__ ggml_cuda_mmq_write_back_t ggml_cuda_mmq_get_write_back() {
|
||||
return ggml_cuda_mmq_get_util_funcs<type, J, fallback, prec_src1>().write_back;
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------------------------
|
||||
|
||||
template <ggml_type type, int J, bool fallback, bool fixup, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, bool fixup, ggml_prec prec_src1>
|
||||
static __device__ __forceinline__ void mul_mat_q_process_tile(
|
||||
const char * __restrict__ x, const int offset_x, const int * __restrict__ y,
|
||||
const int * __restrict__ ids_dst, float * __restrict__ dst, float * __restrict__ tmp_fixup,
|
||||
@@ -958,7 +951,7 @@ static __device__ __forceinline__ void mul_mat_q_process_tile(
|
||||
|
||||
// The mul_mat_q kernel implements "stream-k" work partitioning as described in https://arxiv.org/abs/2301.03598
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1), ggml_cuda_mmq_get_occupancy(type, J, fallback, prec_src1))
|
||||
static __global__ void mul_mat_q(
|
||||
const char * __restrict__ x, const int * __restrict__ y, const int32_t * __restrict__ ids_dst,
|
||||
@@ -1245,7 +1238,7 @@ static __global__ void mul_mat_q(
|
||||
tile_x_max_i, tile_y_max_j, kb0_start, kb0_stop);
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
__launch_bounds__(ggml_cuda_mmq_get_nthreads(type, J, fallback, prec_src1)/2, 1)
|
||||
static __global__ void mul_mat_q_stream_k_fixup(
|
||||
const int32_t * __restrict__ ids_dst, const int32_t * __restrict__ expert_bounds, float * __restrict__ dst,
|
||||
@@ -1390,7 +1383,7 @@ struct mmq_args {
|
||||
int64_t nchannels_x; int64_t nchannels_y; int64_t stride_channel_x; int64_t stride_channel_y; int64_t stride_channel_dst;
|
||||
int64_t nsamples_x; int64_t nsamples_y; int64_t stride_sample_x; int64_t stride_sample_y; int64_t stride_sample_dst;
|
||||
int64_t ncols_max;
|
||||
int64_t ncols_opt; // value to optimize the tile size against, launch grid still uses ncols_max
|
||||
int J_best; // Tile width in ne11(dense)/ne12(MoE) direction to use for optimal performance.
|
||||
};
|
||||
|
||||
static size_t mmq_get_nbytes_shared(const ggml_cuda_mmq_config & config, const int cc) {
|
||||
@@ -1400,7 +1393,7 @@ static size_t mmq_get_nbytes_shared(const ggml_cuda_mmq_config & config, const i
|
||||
return nbs_ids + nbs_x + GGML_PAD(nbs_y, config.nthreads*sizeof(int));
|
||||
}
|
||||
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, int J, bool fallback, ggml_prec prec_src1>
|
||||
static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
|
||||
const int id = ggml_cuda_get_device();
|
||||
const int cc = ggml_cuda_info().devices[id].cc;
|
||||
@@ -1482,34 +1475,9 @@ static void launch_mul_mat_q(ggml_backend_cuda_context & ctx, const mmq_args & a
|
||||
ntx_fd);
|
||||
}
|
||||
|
||||
template <ggml_type type, bool fallback, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, bool fallback, ggml_prec prec_src1>
|
||||
void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
|
||||
const int id = ggml_cuda_get_device();
|
||||
const int cc = ggml_cuda_info().devices[id].cc;
|
||||
const size_t smpbo = ggml_cuda_info().devices[id].smpbo;
|
||||
|
||||
int J_best = 0;
|
||||
int ntiles_J_best = INT_MAX;
|
||||
|
||||
for (int J = 8; J <= 128 && ntiles_J_best > 1; J += 8) {
|
||||
const ggml_cuda_mmq_config config = ggml_cuda_mmq_get_config(type, J, fallback, cc, prec_src1);
|
||||
if (config.type == GGML_TYPE_COUNT) {
|
||||
continue;
|
||||
}
|
||||
|
||||
if (mmq_get_nbytes_shared(config, cc) > smpbo) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int ntiles_x = (args.ncols_opt + config.J - 1) / config.J;
|
||||
|
||||
if (ntiles_x < ntiles_J_best) {
|
||||
J_best = J;
|
||||
ntiles_J_best = ntiles_x;
|
||||
}
|
||||
}
|
||||
|
||||
switch (J_best) {
|
||||
switch (args.J_best) {
|
||||
case 8:
|
||||
launch_mul_mat_q<type, 8, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
@@ -1559,25 +1527,25 @@ void mul_mat_q_switch_J(ggml_backend_cuda_context & ctx, const mmq_args & args,
|
||||
launch_mul_mat_q<type, 128, fallback, prec_src1>(ctx, args, stream);
|
||||
break;
|
||||
default:
|
||||
fprintf(stderr, "J_best=%d\n", J_best);
|
||||
fprintf(stderr, "J_best=%d\n", args.J_best);
|
||||
GGML_ABORT("fatal error");
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
template <ggml_type type, ggml_prec prec_src1 = GGML_PREC_Q8>
|
||||
template <ggml_type type, ggml_prec prec_src1>
|
||||
void mul_mat_q_case(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) {
|
||||
if (args.nrows_x % 128 == 0) {
|
||||
constexpr bool fallback = false;
|
||||
if (ggml_cuda_mmq_needs_fallback(args.nrows_x)) {
|
||||
constexpr bool fallback = true;
|
||||
mul_mat_q_switch_J<type, fallback, prec_src1>(ctx, args, stream);
|
||||
} else {
|
||||
constexpr bool fallback = true;
|
||||
constexpr bool fallback = false;
|
||||
mul_mat_q_switch_J<type, fallback, prec_src1>(ctx, args, stream);
|
||||
}
|
||||
}
|
||||
|
||||
#define DECL_MMQ_CASE(type) \
|
||||
template void mul_mat_q_case<type>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
|
||||
template void mul_mat_q_case<type, GGML_PREC_Q8>(ggml_backend_cuda_context & ctx, const mmq_args & args, cudaStream_t stream) \
|
||||
|
||||
// FP4 variant: uses native FP4 MMA instead of keeping src1 at Q8_1.
|
||||
#define DECL_MMQ_CASE_W4A4(type) \
|
||||
|
||||
+38
-35
@@ -15,49 +15,52 @@ static __global__ void pad_f32(const float * src, size_t s00, size_t s01, size_t
|
||||
// blockIdx.z: i3*ne2+i2
|
||||
// blockIdx.y: i1
|
||||
// blockIDx.x: i0 / CUDA_PAD_BLOCK_SIZE
|
||||
// gridDim.y: ne1
|
||||
// gridDim.y and gridDim.z are capped at 65535, blocks stride over larger ne1 and ne2*ne3
|
||||
int i0 = threadIdx.x + blockIdx.x * blockDim.x;
|
||||
int i1 = blockIdx.y;
|
||||
int i2 = blockIdx.z % ne2;
|
||||
int i3 = blockIdx.z / ne2;
|
||||
|
||||
if (i0 >= ne0 || i1 >= ne1 || i2 >= ne2 || i3 >= ne3) {
|
||||
if (i0 >= ne0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int64_t dst_idx = i3 * (ne0 * ne1 * ne2) + i2 * (ne0 * ne1) + i1 * ne0 + i0;
|
||||
for (int i1 = blockIdx.y; i1 < ne1; i1 += gridDim.y) {
|
||||
for (int i23 = blockIdx.z; i23 < ne2 * ne3; i23 += gridDim.z) {
|
||||
int i2 = i23 % ne2;
|
||||
int i3 = i23 / ne2;
|
||||
|
||||
if (!circular) {
|
||||
if ((i0 >= lp0 && i0 < ne0 - rp0) && (i1 >= lp1 && i1 < ne1 - rp1) && (i2 >= lp2 && i2 < ne2 - rp2) &&
|
||||
(i3 >= lp3 && i3 < ne3 - rp3)) {
|
||||
const int64_t i00 = i0 - lp0;
|
||||
const int64_t i01 = i1 - lp1;
|
||||
const int64_t i02 = i2 - lp2;
|
||||
const int64_t i03 = i3 - lp3;
|
||||
const int64_t dst_idx = i3 * (ne0 * ne1 * ne2) + i2 * (ne0 * ne1) + i1 * ne0 + i0;
|
||||
|
||||
const int64_t src_idx = i03 * s03 + i02 * s02 + i01 * s01 + i00 * s00;
|
||||
if (!circular) {
|
||||
if ((i0 >= lp0 && i0 < ne0 - rp0) && (i1 >= lp1 && i1 < ne1 - rp1) && (i2 >= lp2 && i2 < ne2 - rp2) &&
|
||||
(i3 >= lp3 && i3 < ne3 - rp3)) {
|
||||
const int64_t i00 = i0 - lp0;
|
||||
const int64_t i01 = i1 - lp1;
|
||||
const int64_t i02 = i2 - lp2;
|
||||
const int64_t i03 = i3 - lp3;
|
||||
|
||||
dst[dst_idx] = src[src_idx];
|
||||
} else {
|
||||
dst[dst_idx] = 0.0f;
|
||||
const int64_t src_idx = i03 * s03 + i02 * s02 + i01 * s01 + i00 * s00;
|
||||
|
||||
dst[dst_idx] = src[src_idx];
|
||||
} else {
|
||||
dst[dst_idx] = 0.0f;
|
||||
}
|
||||
}
|
||||
// circular means on a torus, so x and y wrap around
|
||||
else {
|
||||
const int64_t ne00 = ne0 - lp0 - rp0;
|
||||
const int64_t ne01 = ne1 - lp1 - rp1;
|
||||
const int64_t ne02 = ne2 - lp2 - rp2;
|
||||
const int64_t ne03 = ne3 - lp3 - rp3;
|
||||
|
||||
const int64_t i00 = wrap_around(i0 - lp0, ne00);
|
||||
const int64_t i01 = wrap_around(i1 - lp1, ne01);
|
||||
const int64_t i02 = wrap_around(i2 - lp2, ne02);
|
||||
const int64_t i03 = wrap_around(i3 - lp3, ne03);
|
||||
|
||||
const int64_t src_idx = i03 * s03 + i02 * s02 + i01 * s01 + i00 * s00;
|
||||
|
||||
dst[dst_idx] = src[src_idx];
|
||||
}
|
||||
}
|
||||
}
|
||||
// circular means on a torus, so x and y wrap around
|
||||
else {
|
||||
const int64_t ne00 = ne0 - lp0 - rp0;
|
||||
const int64_t ne01 = ne1 - lp1 - rp1;
|
||||
const int64_t ne02 = ne2 - lp2 - rp2;
|
||||
const int64_t ne03 = ne3 - lp3 - rp3;
|
||||
|
||||
const int64_t i00 = wrap_around(i0 - lp0, ne00);
|
||||
const int64_t i01 = wrap_around(i1 - lp1, ne01);
|
||||
const int64_t i02 = wrap_around(i2 - lp2, ne02);
|
||||
const int64_t i03 = wrap_around(i3 - lp3, ne03);
|
||||
|
||||
const int64_t src_idx = i03 * s03 + i02 * s02 + i01 * s01 + i00 * s00;
|
||||
|
||||
dst[dst_idx] = src[src_idx];
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -67,7 +70,7 @@ static void pad_f32_cuda(const float * src, size_t s00, size_t s01, size_t s02,
|
||||
const int ne0, const int ne1, const int ne2, const int ne3,
|
||||
const bool circular, cudaStream_t stream) {
|
||||
int num_blocks = (ne0 + CUDA_PAD_BLOCK_SIZE - 1) / CUDA_PAD_BLOCK_SIZE;
|
||||
dim3 gridDim(num_blocks, ne1, ne2 * ne3);
|
||||
dim3 gridDim(num_blocks, std::min(ne1, 65535), std::min(ne2 * ne3, 65535));
|
||||
pad_f32<<<gridDim, CUDA_PAD_BLOCK_SIZE, 0, stream>>>(src, s00, s01, s02, s03, dst,
|
||||
lp0, rp0, lp1, rp1, lp2, rp2, lp3, rp3,
|
||||
ne0, ne1, ne2, ne3, circular);
|
||||
|
||||
@@ -149,7 +149,7 @@ __global__ void __launch_bounds__(d_state, 1)
|
||||
const int src0_nb2, const int src0_nb3, const int src1_nb2, const int src1_nb3,
|
||||
const int src2_nb1, const int src2_nb2, const int src3_nb1,
|
||||
const int src4_nb2, const int src4_nb3, const int src5_nb2, const int src5_nb3,
|
||||
const int64_t s_off, const int64_t n_head, const int64_t d_head, const int64_t n_group, const int64_t n_tok, const int64_t K) {
|
||||
char * s_base, const int64_t s_slot_bytes, const int64_t n_head, const int64_t d_head, const int64_t n_group, const int64_t n_tok, const int64_t K) {
|
||||
const float * GGML_CUDA_RESTRICT src0 = src0_ptr;
|
||||
const float * GGML_CUDA_RESTRICT src1 = src1_ptr;
|
||||
const float * GGML_CUDA_RESTRICT src2 = src2_ptr;
|
||||
@@ -184,7 +184,7 @@ __global__ void __launch_bounds__(d_state, 1)
|
||||
const float * B_warp = (const float *) ((const char *) src4 + (seq_idx * src4_nb3) + (group_off));
|
||||
const float * C_warp = (const float *) ((const char *) src5 + (seq_idx * src5_nb3) + (group_off));
|
||||
float * y_warp = dst + (seq_idx * n_tok * n_head * d_head) + warp_idx;
|
||||
float * s_warp = (float *) ((char *) dst + s_off + seq_idx * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
|
||||
float * s_warp = (float *) (s_base + seq_idx * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
|
||||
|
||||
// strides across n_seq_tokens
|
||||
const int stride_x = src1_nb2 / sizeof(float);
|
||||
@@ -227,7 +227,7 @@ __global__ void __launch_bounds__(d_state, 1)
|
||||
// Slot 0 is the final state written below; slots 1..K-1 are rollback snapshots.
|
||||
const int64_t slot = n_tok - 1 - i;
|
||||
if (K > 1 && slot > 0 && slot < K) {
|
||||
float * s_snapshot_warp = (float *) ((char *) dst + s_off + (slot * gridDim.y + seq_idx) * src0_nb3 + head_idx * src0_nb2 + head_off * d_state);
|
||||
float * s_snapshot_warp = (float *) ((char *) s_warp + slot * s_slot_bytes);
|
||||
#pragma unroll
|
||||
for (int j = 0; j < c_factor; j++) {
|
||||
s_snapshot_warp[WARP_SIZE * j + lane] = state[j];
|
||||
@@ -248,7 +248,11 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
const int src2_nb2, const int src3_nb1, const int src4_nb2, const int src4_nb3, const int src5_nb2,
|
||||
const int src5_nb3, const int64_t s_off, const int64_t d_state, const int64_t head_dim,
|
||||
const int64_t n_head, const int64_t n_group, const int64_t n_tok, const int64_t n_seq,
|
||||
const int64_t K, cudaStream_t stream) {
|
||||
const int64_t K, const ggml_cuda_ssm_scan_fused_cache * cache, cudaStream_t stream) {
|
||||
// when fused, the states go straight into the recurrent cache and the dst tail is left alone
|
||||
char * const s_base = cache ? (char *) cache->data : (char *) dst + s_off;
|
||||
const int64_t s_slot_bytes = cache ? cache->slot_stride * (int64_t) sizeof(float) : n_seq * (int64_t) src0_nb3;
|
||||
|
||||
// NOTE: if you change conditions here, be sure to update the corresponding supports_op condition!
|
||||
if (src3_nb1 == sizeof(float)) {
|
||||
// Mamba-2
|
||||
@@ -261,7 +265,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
ggml_cuda_kernel_launch(ssm_scan_f32_group<96/WARP_SIZE, 96>, launch_params,
|
||||
src0, src1, src2, src3, src4, src5, src6, dst,
|
||||
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
|
||||
} else if (d_state == 128) {
|
||||
constexpr int threads = 128;
|
||||
constexpr int num_warps = threads/WARP_SIZE;
|
||||
@@ -271,7 +275,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
ggml_cuda_kernel_launch(ssm_scan_f32_group<128/WARP_SIZE, 128>, launch_params,
|
||||
src0, src1, src2, src3, src4, src5, src6, dst,
|
||||
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
|
||||
} else if (d_state == 256) { // Falcon-H1
|
||||
constexpr int threads = 256;
|
||||
constexpr int num_warps = threads/WARP_SIZE;
|
||||
@@ -281,7 +285,7 @@ static void ssm_scan_f32_cuda(const float * src0, const float * src1, const floa
|
||||
ggml_cuda_kernel_launch(ssm_scan_f32_group<256/WARP_SIZE, 256>, launch_params,
|
||||
src0, src1, src2, src3, src4, src5, src6, dst,
|
||||
src0_nb2, src0_nb3, src1_nb2, src1_nb3, src2_nb1, src2_nb2, src3_nb1,
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_off, n_head, head_dim, n_group, n_tok, K);
|
||||
src4_nb2, src4_nb3, src5_nb2, src5_nb3, s_base, s_slot_bytes, n_head, head_dim, n_group, n_tok, K);
|
||||
} else {
|
||||
GGML_ABORT("doesn't support d_state!=(96, 128 or 256).");
|
||||
}
|
||||
@@ -570,12 +574,13 @@ __global__ void ssm_ssd_scale_state_kernel(
|
||||
}
|
||||
|
||||
// Copy initial state from src0[ids[s]] into s_cur for each sequence.
|
||||
// src0 and s_cur can alias when the state is written straight into the cache.
|
||||
// Grid: (ceil(d_state * head_dim * n_head / BLOCK), n_seqs)
|
||||
template <int BLOCK_SIZE>
|
||||
__global__ void ssm_ssd_init_state_kernel(
|
||||
const float * __restrict__ src0, // {d_state, head_dim, n_head, n_rs}
|
||||
const float * src0, // {d_state, head_dim, n_head, n_rs}
|
||||
const int32_t * __restrict__ ids, // {n_seqs}
|
||||
float * __restrict__ s_cur, // {d_state, head_dim, n_head, n_seqs}
|
||||
float * s_cur, // {d_state, head_dim, n_head, n_seqs}
|
||||
const int state_size, // d_state * head_dim * n_head
|
||||
const int64_t s0_stride_seq) { // elements between state rows
|
||||
const int s = blockIdx.y;
|
||||
@@ -599,7 +604,8 @@ static void ssm_scan_ssd_f32_cuda(
|
||||
const int A_stride, // A (src3) stride between heads
|
||||
const int B_stride_tok, const int B_stride_seq, // B (src4) strides
|
||||
const int C_stride_tok, const int C_stride_seq, // C (src5) strides
|
||||
const int64_t s_off, const int64_t d_state, const int64_t head_dim,
|
||||
float * s_cur, // state: dst state tail, or the cache when fused
|
||||
const int64_t d_state, const int64_t head_dim,
|
||||
const int64_t n_head, const int64_t n_group, const int64_t n_tok, const int64_t n_seq) {
|
||||
|
||||
cudaStream_t stream = ctx.stream();
|
||||
@@ -625,7 +631,6 @@ static void ssm_scan_ssd_f32_cuda(
|
||||
matmul_t * X_dt = X_dt_buf.get();
|
||||
matmul_t * B_weighted = B_w_buf.get();
|
||||
float * C_scaled = C_s_buf.get();
|
||||
float * s_cur = (float *)((char *)dst_d + s_off); // write state directly to dst
|
||||
|
||||
// Step 1: softplus(dt) and parallel prefix sum over full sequence
|
||||
{
|
||||
@@ -780,7 +785,8 @@ static void ssm_scan_ssd_f32_cuda(
|
||||
}
|
||||
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
|
||||
|
||||
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
static void ggml_cuda_op_ssm_scan_impl(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
|
||||
const ggml_cuda_ssm_scan_fused_cache * cache) {
|
||||
const struct ggml_tensor * src0 = dst->src[0]; // s
|
||||
const struct ggml_tensor * src1 = dst->src[1]; // x
|
||||
const struct ggml_tensor * src2 = dst->src[2]; // dt
|
||||
@@ -864,12 +870,21 @@ void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
(int)(src3->nb[1] / sizeof(float)),
|
||||
(int)(src4->nb[2] / sizeof(float)), (int)(src4->nb[3] / sizeof(float)),
|
||||
(int)(src5->nb[2] / sizeof(float)), (int)(src5->nb[3] / sizeof(float)),
|
||||
s_off, nc, nr, nh, ng, n_t, n_s);
|
||||
cache ? cache->data : (float *) ((char *) dst_d + s_off), nc, nr, nh, ng, n_t, n_s);
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
ssm_scan_f32_cuda(src0_d, src1_d, src2_d, src3_d, src4_d, src5_d, src6_d, dst_d,
|
||||
src0->nb[2], src0->nb[3], src1->nb[2], src1->nb[3], src2->nb[1], src2->nb[2],
|
||||
src3->nb[1], src4->nb[2], src4->nb[3], src5->nb[2], src5->nb[3],
|
||||
s_off, nc, nr, nh, ng, n_t, n_s, K, stream);
|
||||
s_off, nc, nr, nh, ng, n_t, n_s, K, cache, stream);
|
||||
}
|
||||
|
||||
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
ggml_cuda_op_ssm_scan_impl(ctx, dst, nullptr);
|
||||
}
|
||||
|
||||
void ggml_cuda_op_ssm_scan_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
|
||||
ggml_cuda_ssm_scan_fused_cache cache) {
|
||||
ggml_cuda_op_ssm_scan_impl(ctx, dst, &cache);
|
||||
}
|
||||
|
||||
@@ -1,3 +1,13 @@
|
||||
#include "common.cuh"
|
||||
|
||||
// fused-kernel recurrent-state output; strides in elements (per-seq stride is always the state row size, set in-kernel)
|
||||
struct ggml_cuda_ssm_scan_fused_cache {
|
||||
float * data; // rollback slot 0
|
||||
int64_t slot_stride; // between rollback slots
|
||||
};
|
||||
|
||||
void ggml_cuda_op_ssm_scan(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
|
||||
|
||||
// same op, but writes the state snapshot(s) into the cache instead of dst (see ggml_cuda_try_ssm_scan_cache_fusion)
|
||||
void ggml_cuda_op_ssm_scan_fused_cache(ggml_backend_cuda_context & ctx, ggml_tensor * dst,
|
||||
ggml_cuda_ssm_scan_fused_cache cache);
|
||||
|
||||
+124
-66
@@ -1,6 +1,29 @@
|
||||
#include "argsort.cuh"
|
||||
#include "top-k.cuh"
|
||||
|
||||
// Adjusted implementation thresholds from #28547, can be overridden at build time
|
||||
#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC
|
||||
# if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
|
||||
// not measured on HIP/MUSA, keep the old split
|
||||
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC 1024
|
||||
# else
|
||||
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC 512
|
||||
# endif
|
||||
#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC
|
||||
|
||||
#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT
|
||||
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT 4096
|
||||
#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT
|
||||
|
||||
// bitonic up to this width while nrows fits in one wave of SMs, 0 disables
|
||||
#ifndef GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS
|
||||
# if defined(GGML_USE_HIP) || defined(GGML_USE_MUSA)
|
||||
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS 0
|
||||
# else
|
||||
# define GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS 1024
|
||||
# endif
|
||||
#endif // GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS
|
||||
|
||||
#ifdef GGML_CUDA_USE_CUB
|
||||
# include <cub/cub.cuh>
|
||||
// DeviceTopK has a race condition before CCCL 3.4.3.
|
||||
@@ -14,6 +37,15 @@ using namespace cub;
|
||||
# endif // CCCL >= 3.4.3
|
||||
#endif // GGML_CUDA_USE_CUB
|
||||
|
||||
// max rows for the per-row DeviceTopK / CUB argsort path before switching to radix / bitonic
|
||||
#ifndef GGML_CUDA_TOP_K_NROWS_THRESHOLD
|
||||
# ifdef CUB_TOP_K_AVAILABLE
|
||||
# define GGML_CUDA_TOP_K_NROWS_THRESHOLD 2
|
||||
# else
|
||||
# define GGML_CUDA_TOP_K_NROWS_THRESHOLD 1
|
||||
# endif
|
||||
#endif // GGML_CUDA_TOP_K_NROWS_THRESHOLD
|
||||
|
||||
#ifdef CUB_TOP_K_AVAILABLE
|
||||
|
||||
static void top_k_cub(ggml_cuda_pool & pool,
|
||||
@@ -40,7 +72,7 @@ static void top_k_cub(ggml_cuda_pool & pool,
|
||||
ncols, k, env));
|
||||
}
|
||||
|
||||
#elif defined(GGML_CUDA_USE_CUB) // CUB_TOP_K_AVAILABLE
|
||||
#endif // CUB_TOP_K_AVAILABLE
|
||||
|
||||
static int next_power_of_2(int x) {
|
||||
int n = 1;
|
||||
@@ -50,10 +82,6 @@ static int next_power_of_2(int x) {
|
||||
return n;
|
||||
}
|
||||
|
||||
#endif // CUB_TOP_K_AVAILABLE
|
||||
|
||||
#if !defined(GGML_CUDA_USE_CUB) && defined(GGML_USE_HIP)
|
||||
|
||||
static __device__ __forceinline__ uint32_t top_k_float_to_ordered(float value) {
|
||||
const uint32_t bits = __float_as_uint(value);
|
||||
const uint32_t mask = (uint32_t) (-(int32_t) (bits >> 31)) | 0x80000000U;
|
||||
@@ -95,7 +123,7 @@ static __global__ void top_k_radix_histogram(
|
||||
__syncthreads();
|
||||
|
||||
const top_k_radix_state state = states[row];
|
||||
for (int col = row_block * BLOCK_SIZE + tid;
|
||||
for (int64_t col = row_block * BLOCK_SIZE + tid;
|
||||
col < ncols;
|
||||
col += blocks_per_row * BLOCK_SIZE) {
|
||||
const uint32_t key = top_k_float_to_ordered(row_src[col]);
|
||||
@@ -165,7 +193,7 @@ static __global__ void top_k_radix_gather(
|
||||
int * row_dst = dst + (size_t) row * k;
|
||||
top_k_radix_state * state = &states[row];
|
||||
|
||||
for (int col = row_block * BLOCK_SIZE + tid;
|
||||
for (int64_t col = row_block * BLOCK_SIZE + tid;
|
||||
col < ncols;
|
||||
col += blocks_per_row * BLOCK_SIZE) {
|
||||
const uint32_t key = top_k_float_to_ordered(row_src[col]);
|
||||
@@ -183,36 +211,72 @@ static __global__ void top_k_radix_gather(
|
||||
|
||||
static void top_k_radix_cuda(
|
||||
ggml_cuda_pool & pool,
|
||||
const float * src, int * dst, int ncols, int nrows, int k, cudaStream_t stream) {
|
||||
const float * src, int * dst, int ncols, int64_t nrows, int k, cudaStream_t stream) {
|
||||
constexpr int BLOCK_SIZE = 256;
|
||||
constexpr int RADIX_BITS = 8;
|
||||
constexpr int NBINS = 1 << RADIX_BITS;
|
||||
const int blocks_per_row = std::min((ncols + 1023) / 1024, 64);
|
||||
const int blocks_per_row = (int) std::min<int64_t>(((int64_t) ncols + 1023) / 1024, 64);
|
||||
|
||||
ggml_cuda_pool_alloc<top_k_radix_state> states_alloc(pool, nrows);
|
||||
ggml_cuda_pool_alloc<int> histograms_alloc(pool, (size_t) nrows * blocks_per_row * NBINS);
|
||||
// chunk the rows to bound the histogram memory to 64 MB
|
||||
const int64_t chunk_nrows = ggml_cuda_chunk_nrows((size_t) blocks_per_row * NBINS * sizeof(int), nrows);
|
||||
|
||||
ggml_cuda_pool_alloc<top_k_radix_state> states_alloc(pool, chunk_nrows);
|
||||
ggml_cuda_pool_alloc<int> histograms_alloc(pool, (size_t) chunk_nrows * blocks_per_row * NBINS);
|
||||
top_k_radix_state * states = states_alloc.get();
|
||||
int * histograms = histograms_alloc.get();
|
||||
|
||||
top_k_radix_init<<<(nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, nrows, k);
|
||||
for (int64_t i = 0; i < nrows; i += chunk_nrows) {
|
||||
const int iter_nrows = std::min(chunk_nrows, nrows - i);
|
||||
|
||||
const dim3 row_grid(blocks_per_row * nrows);
|
||||
for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) {
|
||||
top_k_radix_histogram<BLOCK_SIZE, RADIX_BITS>
|
||||
top_k_radix_init<<<(iter_nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, iter_nrows, k);
|
||||
|
||||
const dim3 row_grid(blocks_per_row * iter_nrows);
|
||||
for (int shift = 32 - RADIX_BITS; shift >= 0; shift -= RADIX_BITS) {
|
||||
top_k_radix_histogram<BLOCK_SIZE, RADIX_BITS>
|
||||
<<<row_grid, BLOCK_SIZE, 0, stream>>>(
|
||||
src, states, histograms, ncols, blocks_per_row, shift);
|
||||
top_k_radix_select<BLOCK_SIZE, RADIX_BITS>
|
||||
<<<iter_nrows, BLOCK_SIZE, 0, stream>>>(histograms, states, blocks_per_row, shift);
|
||||
}
|
||||
|
||||
top_k_radix_reset_counters
|
||||
<<<(iter_nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, iter_nrows);
|
||||
top_k_radix_gather<BLOCK_SIZE>
|
||||
<<<row_grid, BLOCK_SIZE, 0, stream>>>(
|
||||
src, states, histograms, ncols, blocks_per_row, shift);
|
||||
top_k_radix_select<BLOCK_SIZE, RADIX_BITS>
|
||||
<<<nrows, BLOCK_SIZE, 0, stream>>>(histograms, states, blocks_per_row, shift);
|
||||
}
|
||||
src, dst, states, ncols, k, blocks_per_row);
|
||||
|
||||
top_k_radix_reset_counters
|
||||
<<<(nrows + BLOCK_SIZE - 1) / BLOCK_SIZE, BLOCK_SIZE, 0, stream>>>(states, nrows);
|
||||
top_k_radix_gather<BLOCK_SIZE>
|
||||
<<<row_grid, BLOCK_SIZE, 0, stream>>>(
|
||||
src, dst, states, ncols, k, blocks_per_row);
|
||||
src += (size_t) ncols * iter_nrows;
|
||||
dst += (size_t) k * iter_nrows;
|
||||
}
|
||||
}
|
||||
|
||||
#endif // !defined(GGML_CUDA_USE_CUB) && defined(GGML_USE_HIP)
|
||||
static void top_k_argsort_cuda(
|
||||
ggml_cuda_pool & pool,
|
||||
const float * src, int * dst, int ncols, int64_t nrows, int k, bool use_cub, cudaStream_t stream) {
|
||||
const int64_t chunk_nrows = ggml_cuda_chunk_nrows((size_t) ncols * sizeof(int), nrows);
|
||||
|
||||
ggml_cuda_pool_alloc<int> tmp_alloc(pool, (size_t) ncols * chunk_nrows);
|
||||
int * tmp = tmp_alloc.get();
|
||||
|
||||
for (int64_t i = 0; i < nrows; i += chunk_nrows) {
|
||||
const int iter_nrows = std::min(chunk_nrows, nrows - i);
|
||||
|
||||
if (use_cub) {
|
||||
#ifdef GGML_CUDA_USE_CUB
|
||||
argsort_f32_i32_cuda_cub(pool, src, tmp, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
|
||||
#else
|
||||
GGML_ABORT("CUB is not available");
|
||||
#endif // GGML_CUDA_USE_CUB
|
||||
} else {
|
||||
argsort_f32_i32_cuda_bitonic(src, tmp, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
|
||||
}
|
||||
CUDA_CHECK(cudaMemcpy2DAsync(dst, k * sizeof(int), tmp, ncols * sizeof(int), k * sizeof(int), iter_nrows,
|
||||
cudaMemcpyDeviceToDevice, stream));
|
||||
|
||||
src += (size_t) ncols * iter_nrows;
|
||||
dst += (size_t) k * iter_nrows;
|
||||
}
|
||||
}
|
||||
|
||||
void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
const ggml_tensor * src0 = dst->src[0];
|
||||
@@ -229,51 +293,45 @@ void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
const int64_t nrows = ggml_nrows(src0);
|
||||
const int64_t k = dst->ne[0];
|
||||
ggml_cuda_pool & pool = ctx.pool();
|
||||
|
||||
const int device = ggml_cuda_get_device();
|
||||
|
||||
#ifdef CUB_TOP_K_AVAILABLE
|
||||
// TODO: Switch to `DeviceSegmentedTopK` for multi-row TopK once implemented
|
||||
// https://github.com/NVIDIA/cccl/issues/6391
|
||||
// TODO: investigate if there exists a point where parallelized argsort is faster than sequential top-k
|
||||
for (int i = 0; i < nrows; i++) {
|
||||
// a single row always uses DeviceTopK if available
|
||||
const bool bitonic_short = nrows > 1 && ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC;
|
||||
#else
|
||||
const bool bitonic_short = ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC;
|
||||
#endif // CUB_TOP_K_AVAILABLE
|
||||
const bool bitonic_few_rows = nrows > GGML_CUDA_TOP_K_NROWS_THRESHOLD &&
|
||||
ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_BITONIC_FEW_ROWS &&
|
||||
nrows <= ggml_cuda_info().devices[device].nsm;
|
||||
|
||||
if (bitonic_short || bitonic_few_rows) {
|
||||
// the padded row must fit in shared memory
|
||||
const int ncols_pad = next_power_of_2(ncols);
|
||||
if (ncols_pad * sizeof(int) <= ggml_cuda_info().devices[device].smpb) {
|
||||
top_k_argsort_cuda(pool, src0_d, dst_d, ncols, nrows, k, false, stream);
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
if (nrows > GGML_CUDA_TOP_K_NROWS_THRESHOLD) {
|
||||
top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
|
||||
return;
|
||||
}
|
||||
|
||||
#ifdef CUB_TOP_K_AVAILABLE
|
||||
// TODO: Assess perf of `DeviceBatchedTopK` for multi-row TopK & CCCL >= 3.5.0, re-running perf sweep of https://github.com/ggml-org/llama.cpp/pull/28713
|
||||
for (int64_t i = 0; i < nrows; i++) {
|
||||
top_k_cub(pool, src0_d + i * ncols, dst_d + i * k, ncols, k, stream);
|
||||
}
|
||||
#elif defined(GGML_CUDA_USE_CUB) // CUB_TOP_K_AVAILABLE
|
||||
// Fall back to argsort + copy
|
||||
const int ncols_pad = next_power_of_2(ncols);
|
||||
const size_t shared_mem = ncols_pad * sizeof(int);
|
||||
const size_t max_shared_mem = ggml_cuda_info().devices[ggml_cuda_get_device()].smpb;
|
||||
const bool use_bitonic = shared_mem <= max_shared_mem && ncols <= 1024;
|
||||
const int chunk_nrows = argsort_f32_i32_cuda_cub_chunk_nrows(src0->nb[1], nrows);
|
||||
|
||||
ggml_cuda_pool_alloc<int> temp_dst_alloc(pool, ncols * chunk_nrows);
|
||||
int * tmp_dst = temp_dst_alloc.get();
|
||||
|
||||
for (int64_t i = 0; i < nrows; i += chunk_nrows) {
|
||||
int iter_nrows = std::min((int64_t) chunk_nrows, nrows - i);
|
||||
|
||||
if (use_bitonic) {
|
||||
argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
|
||||
} else {
|
||||
argsort_f32_i32_cuda_cub(pool, src0_d, tmp_dst, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream);
|
||||
}
|
||||
CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), iter_nrows,
|
||||
cudaMemcpyDeviceToDevice, stream));
|
||||
|
||||
src0_d += ncols * iter_nrows;
|
||||
dst_d += k * iter_nrows;
|
||||
if (ncols <= GGML_CUDA_TOP_K_NCOLS_THRESHOLD_ARGSORT) {
|
||||
top_k_argsort_cuda(pool, src0_d, dst_d, ncols, nrows, k, true, stream);
|
||||
} else {
|
||||
top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
|
||||
}
|
||||
#else // GGML_CUDA_USE_CUB
|
||||
#if defined(GGML_USE_HIP)
|
||||
if (ncols > 1024) {
|
||||
top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
|
||||
} else {
|
||||
#endif // defined(GGML_USE_HIP)
|
||||
ggml_cuda_pool_alloc<int> temp_dst_alloc(pool, ncols * nrows);
|
||||
int * tmp_dst = temp_dst_alloc.get();
|
||||
argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, nrows, GGML_SORT_ORDER_DESC, stream);
|
||||
CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), nrows,
|
||||
cudaMemcpyDeviceToDevice, stream));
|
||||
#if defined(GGML_USE_HIP)
|
||||
}
|
||||
#endif // defined(GGML_USE_HIP)
|
||||
#endif
|
||||
top_k_radix_cuda(pool, src0_d, dst_d, ncols, nrows, k, stream);
|
||||
#endif // CUB_TOP_K_AVAILABLE
|
||||
}
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -351,12 +351,15 @@ IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_f32_thread, float, hvx_copy_f32_uu,
|
||||
const dma_addr_t vsrc = ok \
|
||||
? (src_data + (size_t) ((in * IC + iic) * IH + iih) * IW * sizeof(float)) \
|
||||
: src_data; \
|
||||
dma_queue_push(dma_q, dma_make_data(vdst, vsrc), \
|
||||
IW * sizeof(float), IW * sizeof(float), IW * sizeof(float), ok ? 1 : 0); \
|
||||
/* IC*KH descriptors per row can exceed the ring capacity: retire the oldest when full */ \
|
||||
while (!dma_queue_push(dma_q, dma_make_data(vdst, vsrc), \
|
||||
IW * sizeof(float), IW * sizeof(float), IW * sizeof(float), \
|
||||
ok ? 1 : 0)) { \
|
||||
dma_queue_pop(dma_q); \
|
||||
} \
|
||||
} \
|
||||
} \
|
||||
for (uint32_t i = 0; i < IC * KH; i++) \
|
||||
dma_queue_pop(dma_q); \
|
||||
dma_queue_flush(dma_q); \
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r); \
|
||||
for (uint32_t iow = 0; iow < OW; iow++) { \
|
||||
DST_CTYPE * dst_patch = dstb + (uint64_t) iow * patch_stride; \
|
||||
|
||||
@@ -1741,10 +1741,13 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
|
||||
op->src[0]->ne[0] != 576) {
|
||||
return false;
|
||||
}
|
||||
if (op->src[1]->ne[0] == 72 && op->src[1]->ne[0] != op->src[2]->ne[0]) {
|
||||
return false;
|
||||
}
|
||||
if (op->src[1]->ne[0] < op->src[2]->ne[0]) {
|
||||
// the kernels exist for K == V and for these K > V pairs only
|
||||
if (op->src[1]->ne[0] != op->src[2]->ne[0] &&
|
||||
!(op->src[1]->ne[0] == 96 && op->src[2]->ne[0] == 64) &&
|
||||
!(op->src[1]->ne[0] == 128 && op->src[2]->ne[0] == 96) &&
|
||||
!(op->src[1]->ne[0] == 192 && op->src[2]->ne[0] == 128) &&
|
||||
!(op->src[1]->ne[0] == 320 && op->src[2]->ne[0] == 256) &&
|
||||
!(op->src[1]->ne[0] == 576 && op->src[2]->ne[0] == 512)) {
|
||||
return false;
|
||||
}
|
||||
if (op->src[1]->type != op->src[2]->type) {
|
||||
|
||||
@@ -3180,6 +3180,7 @@ static int ggml_metal_op_flash_attn_ext_n_kv_max_sparse(const ggml_tensor * op)
|
||||
(dk == 96 && dv == 96) ||
|
||||
(dk == 96 && dv == 64) ||
|
||||
(dk == 128 && dv == 128) ||
|
||||
(dk == 128 && dv == 96) ||
|
||||
(dk == 192 && dv == 128) ||
|
||||
(dk == 192 && dv == 192) ||
|
||||
(dk == 256 && dv == 256) ||
|
||||
|
||||
@@ -40,6 +40,9 @@ int fa_vec_baseline_ne(int dk, int dv) {
|
||||
if (dk == 128 && dv == 128) {
|
||||
return 1;
|
||||
}
|
||||
if (dk == 128 && dv == 96) {
|
||||
return 4;
|
||||
}
|
||||
if (dk == 192 && dv == 192) {
|
||||
return 2;
|
||||
}
|
||||
|
||||
@@ -44,6 +44,7 @@ template [[host_name("kernel_flash_attn_ext_f16_dk96_dv96" )]] kernel flash_at
|
||||
template [[host_name("kernel_flash_attn_ext_f16_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 96, 64>;
|
||||
template [[host_name("kernel_flash_attn_ext_f16_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 112, 112>;
|
||||
template [[host_name("kernel_flash_attn_ext_f16_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 128, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_f16_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 128, 96>;
|
||||
template [[host_name("kernel_flash_attn_ext_f16_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 192, 192>;
|
||||
template [[host_name("kernel_flash_attn_ext_f16_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 192, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_f16_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, half4x4, 1, dequantize_f16, half4x4, 1, dequantize_f16, 256, 256>;
|
||||
@@ -62,6 +63,7 @@ template [[host_name("kernel_flash_attn_ext_bf16_dk96_dv96" )]] kernel flash_at
|
||||
template [[host_name("kernel_flash_attn_ext_bf16_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 96, 64>;
|
||||
template [[host_name("kernel_flash_attn_ext_bf16_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 112, 112>;
|
||||
template [[host_name("kernel_flash_attn_ext_bf16_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 128, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_bf16_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 128, 96>;
|
||||
template [[host_name("kernel_flash_attn_ext_bf16_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 192, 192>;
|
||||
template [[host_name("kernel_flash_attn_ext_bf16_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 192, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_bf16_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_BF, bfloat4x4, 1, dequantize_bf16, bfloat4x4, 1, dequantize_bf16, 256, 256>;
|
||||
|
||||
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_f32_dk96_dv96" )]] kernel flash_at
|
||||
template [[host_name("kernel_flash_attn_ext_f32_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 96, 64>;
|
||||
template [[host_name("kernel_flash_attn_ext_f32_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 112, 112>;
|
||||
template [[host_name("kernel_flash_attn_ext_f32_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 128, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_f32_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 128, 96>;
|
||||
template [[host_name("kernel_flash_attn_ext_f32_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 192, 192>;
|
||||
template [[host_name("kernel_flash_attn_ext_f32_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 192, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_f32_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES_F32, float4x4, 1, dequantize_f32, float4x4, 1, dequantize_f32, 256, 256>;
|
||||
|
||||
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q4_0_dk96_dv96" )]] kernel flash_at
|
||||
template [[host_name("kernel_flash_attn_ext_q4_0_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 96, 64>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_0_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 112, 112>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_0_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 128, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_0_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 128, 96>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_0_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 192, 192>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_0_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 192, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_0_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_0, 2, dequantize_q4_0, block_q4_0, 2, dequantize_q4_0, 256, 256>;
|
||||
|
||||
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q4_1_dk96_dv96" )]] kernel flash_at
|
||||
template [[host_name("kernel_flash_attn_ext_q4_1_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 96, 64>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_1_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 112, 112>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_1_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 128, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_1_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 128, 96>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_1_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 192, 192>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_1_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 192, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q4_1_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q4_1, 2, dequantize_q4_1, block_q4_1, 2, dequantize_q4_1, 256, 256>;
|
||||
|
||||
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q5_0_dk96_dv96" )]] kernel flash_at
|
||||
template [[host_name("kernel_flash_attn_ext_q5_0_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 96, 64>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_0_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 112, 112>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_0_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 128, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_0_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 128, 96>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_0_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 192, 192>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_0_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 192, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_0_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_0, 2, dequantize_q5_0, block_q5_0, 2, dequantize_q5_0, 256, 256>;
|
||||
|
||||
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q5_1_dk96_dv96" )]] kernel flash_at
|
||||
template [[host_name("kernel_flash_attn_ext_q5_1_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 96, 64>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_1_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 112, 112>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_1_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 128, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_1_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 128, 96>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_1_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 192, 192>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_1_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 192, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q5_1_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q5_1, 2, dequantize_q5_1, block_q5_1, 2, dequantize_q5_1, 256, 256>;
|
||||
|
||||
@@ -41,6 +41,7 @@ template [[host_name("kernel_flash_attn_ext_q8_0_dk96_dv96" )]] kernel flash_at
|
||||
template [[host_name("kernel_flash_attn_ext_q8_0_dk96_dv64" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 96, 64>;
|
||||
template [[host_name("kernel_flash_attn_ext_q8_0_dk112_dv112")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 112, 112>;
|
||||
template [[host_name("kernel_flash_attn_ext_q8_0_dk128_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 128, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q8_0_dk128_dv96" )]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 128, 96>;
|
||||
template [[host_name("kernel_flash_attn_ext_q8_0_dk192_dv192")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 192, 192>;
|
||||
template [[host_name("kernel_flash_attn_ext_q8_0_dk192_dv128")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 192, 128>;
|
||||
template [[host_name("kernel_flash_attn_ext_q8_0_dk256_dv256")]] kernel flash_attn_ext_t kernel_flash_attn_ext<FA_TYPES, block_q8_0, 2, dequantize_q8_0, block_q8_0, 2, dequantize_q8_0, 256, 256>;
|
||||
|
||||
@@ -56,8 +56,12 @@ template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv128_q2_ne4")]] kerne
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 128, 1, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 128, 2, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 128, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 96, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 96, 4, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f16_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 128, 96, 4, 4>;
|
||||
#if defined(GGML_METAL_HAS_BF16)
|
||||
template [[host_name("kernel_flash_attn_ext_vec_bf16_dk128_dv128")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, bfloat4, 1, dequantize_bf16_t4, bfloat4, 1, dequantize_bf16_t4, 128, 128, 1>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_bf16_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, bfloat4, 1, dequantize_bf16_t4, bfloat4, 1, dequantize_bf16_t4, 128, 96, 4>;
|
||||
#endif
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f16_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 192, 192, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f16_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, half4, 1, dequantize_f16_t4, half4, 1, dequantize_f16_t4, 192, 192, 4, 1>;
|
||||
|
||||
@@ -25,6 +25,7 @@ template [[host_name("kernel_flash_attn_ext_vec_f32_dk64_dv64")]] kernel flas
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f32_dk96_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 96, 96, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f32_dk96_dv64")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 96, 64, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f32_dk128_dv128")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 128, 128, 1>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f32_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 128, 96, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f32_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 192, 192, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f32_dk192_dv128")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 192, 128, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_f32_dk256_dv256")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES_F32, float4, 1, dequantize_f32_t4, float4, 1, dequantize_f32_t4, 256, 256, 1>;
|
||||
|
||||
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv128_q2_ne4")]] kern
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 128, 1, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 128, 2, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 128, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 96, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 96, 4, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 128, 96, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 192, 192, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 192, 192, 4, 1>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_0_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_0, 8, dequantize_q4_0_t4, block_q4_0, 8, dequantize_q4_0_t4, 192, 192, 2, 2>;
|
||||
|
||||
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv128_q2_ne4")]] kern
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 128, 1, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 128, 2, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 128, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 96, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 96, 4, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 128, 96, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 192, 192, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 192, 192, 4, 1>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q4_1_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q4_1, 8, dequantize_q4_1_t4, block_q4_1, 8, dequantize_q4_1_t4, 192, 192, 2, 2>;
|
||||
|
||||
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv128_q2_ne4")]] kern
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 128, 1, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 128, 2, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 128, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 96, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 96, 4, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 128, 96, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 192, 192, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 192, 192, 4, 1>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_0_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_0, 8, dequantize_q5_0_t4, block_q5_0, 8, dequantize_q5_0_t4, 192, 192, 2, 2>;
|
||||
|
||||
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv128_q2_ne4")]] kern
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 128, 1, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 128, 2, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 128, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 96, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 96, 4, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 128, 96, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 192, 192, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 192, 192, 4, 1>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q5_1_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q5_1, 8, dequantize_q5_1_t4, block_q5_1, 8, dequantize_q5_1_t4, 192, 192, 2, 2>;
|
||||
|
||||
@@ -44,6 +44,9 @@ template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv128_q2_ne4")]] kern
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv128_q4_ne1")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 128, 1, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv128_q4_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 128, 2, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv128_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 128, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv96")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 96, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv96_q2_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 96, 4, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk128_dv96_q4_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 128, 96, 4, 4>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk192_dv192")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 192, 192, 2>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk192_dv192_q1_ne4")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 192, 192, 4, 1>;
|
||||
template [[host_name("kernel_flash_attn_ext_vec_q8_0_dk192_dv192_q2_ne2")]] kernel flash_attn_ext_vec_t kernel_flash_attn_ext_vec<FA_TYPES, block_q8_0, 8, dequantize_q8_0_t4, block_q8_0, 8, dequantize_q8_0_t4, 192, 192, 2, 2>;
|
||||
|
||||
@@ -21,7 +21,10 @@ if (MUSAToolkit_FOUND)
|
||||
message(STATUS "MUSA Toolkit found")
|
||||
|
||||
if (NOT DEFINED MUSA_ARCHITECTURES)
|
||||
set(MUSA_ARCHITECTURES "21;22;31")
|
||||
set(MUSA_ARCHITECTURES "22;31")
|
||||
endif()
|
||||
if ("21" IN_LIST MUSA_ARCHITECTURES)
|
||||
message(WARNING "MUSA architecture 21 (MTT S70, MTT S80, MTT S3000) is deprecated and no longer tested")
|
||||
endif()
|
||||
message(STATUS "Using MUSA architectures: ${MUSA_ARCHITECTURES}")
|
||||
|
||||
|
||||
@@ -280,7 +280,8 @@ static ADRENO_GPU_GEN get_adreno_gpu_gen(const char *device_name) {
|
||||
strstr(device_name, "613") || strstr(device_name, "615") ||
|
||||
strstr(device_name, "616") || strstr(device_name, "618") ||
|
||||
strstr(device_name, "619") || strstr(device_name, "620") ||
|
||||
strstr(device_name, "630") || strstr(device_name, "640") ||
|
||||
strstr(device_name, "623") || strstr(device_name, "630") ||
|
||||
strstr(device_name, "640") ||
|
||||
strstr(device_name, "642") || strstr(device_name, "643") ||
|
||||
strstr(device_name, "644") || strstr(device_name, "650") ||
|
||||
strstr(device_name, "660") || strstr(device_name, "663") ||
|
||||
@@ -863,7 +864,8 @@ struct ggml_backend_opencl_context {
|
||||
cl_kernel kernel_set_rows_q4_0_soa_i64, kernel_set_rows_q4_0_soa_i32;
|
||||
cl_kernel kernel_rope_norm_f32, kernel_rope_norm_f16, kernel_rope_neox_f32, kernel_rope_neox_f16;
|
||||
cl_kernel kernel_rope_multi_f32, kernel_rope_multi_f16, kernel_rope_vision_f32, kernel_rope_vision_f16;
|
||||
cl_kernel kernel_cpy_f16_f16, kernel_cpy_f16_f32, kernel_cpy_f32_f16, kernel_cpy_f32_f32, kernel_cpy_f32_f32_pack, kernel_cpy_i32_i32;
|
||||
cl_kernel kernel_cpy_f16_f16, kernel_cpy_f16_f32, kernel_cpy_f32_f16, kernel_cpy_f32_f32, kernel_cpy_i32_i32;
|
||||
cl_kernel kernel_cpy_f32_f32_pack = nullptr;
|
||||
cl_kernel kernel_cpy_f32_f32_flat = nullptr;
|
||||
cl_kernel kernel_mul_mat_f32_f32;
|
||||
cl_kernel kernel_mul_mat_f16_f16;
|
||||
@@ -1604,14 +1606,18 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
#else
|
||||
const std::string kernel_src = read_file("cpy.cl");
|
||||
#endif
|
||||
cl_program prog =
|
||||
build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts);
|
||||
const bool no_cpy_pack = backend_ctx->adreno_gen == ADRENO_GPU_GEN::A6X;
|
||||
|
||||
cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(),
|
||||
no_cpy_pack ? compile_opts + " -DGGML_CL_NO_CPY_PACK" : compile_opts);
|
||||
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f16_f16 = clCreateKernel(prog, "kernel_cpy_f16_f16", &err), err));
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f16_f32 = clCreateKernel(prog, "kernel_cpy_f16_f32", &err), err));
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f32_f16 = clCreateKernel(prog, "kernel_cpy_f32_f16", &err), err));
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f32_f32 = clCreateKernel(prog, "kernel_cpy_f32_f32", &err), err));
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f32_f32_pack = clCreateKernel(prog, "kernel_cpy_f32_f32_pack", &err), err));
|
||||
if (!no_cpy_pack) {
|
||||
CL_CHECK((backend_ctx->kernel_cpy_f32_f32_pack = clCreateKernel(prog, "kernel_cpy_f32_f32_pack", &err), err));
|
||||
}
|
||||
{ // optional: without it ggml_cl_cpy keeps the row-mapped kernel
|
||||
cl_int err_flat = CL_SUCCESS;
|
||||
cl_kernel k = clCreateKernel(prog, "kernel_cpy_f32_f32_flat", &err_flat);
|
||||
@@ -3770,6 +3776,9 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
if (backend_ctx->has_vector_subgroup_broadcast) {
|
||||
CL_gemv_compile_opts += " -DVECTOR_SUB_GROUP_BROADCAST ";
|
||||
}
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::A6X) {
|
||||
CL_gemv_compile_opts += " -DGGML_CL_A6X_CONSTFOLD_FIX";
|
||||
}
|
||||
|
||||
#ifdef GGML_OPENCL_EMBED_KERNELS
|
||||
const std::string kernel_src_CL_gemv_general {
|
||||
@@ -4306,6 +4315,9 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
|
||||
if (backend_ctx->has_vector_subgroup_broadcast) {
|
||||
CL_gemv_compile_opts += " -DVECTOR_SUB_GROUP_BROADCAST ";
|
||||
}
|
||||
if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::A6X) {
|
||||
CL_gemv_compile_opts += " -DGGML_CL_A6X_CONSTFOLD_FIX";
|
||||
}
|
||||
// Opt-in: dequant-once-per-block mc3 verify GEMV (factors q4_K dequant
|
||||
// out of the 3-column loop; byte-identical, lower spill). A/B vs the
|
||||
// shipped inline mc3 in the same binary.
|
||||
@@ -28335,7 +28347,8 @@ static void ggml_cl_cpy(ggml_backend_t backend, const ggml_tensor * src0, const
|
||||
kernel = backend_ctx->kernel_cpy_f32_f16;
|
||||
break;
|
||||
case GGML_TYPE_F32:
|
||||
kernel = ne00 < 32 ? backend_ctx->kernel_cpy_f32_f32_pack
|
||||
kernel = (ne00 < 32 && backend_ctx->kernel_cpy_f32_f32_pack)
|
||||
? backend_ctx->kernel_cpy_f32_f32_pack
|
||||
: backend_ctx->kernel_cpy_f32_f32;
|
||||
break;
|
||||
default:
|
||||
|
||||
@@ -183,6 +183,7 @@ kernel void kernel_cpy_f32_f32(
|
||||
}
|
||||
}
|
||||
|
||||
#ifndef GGML_CL_NO_CPY_PACK
|
||||
kernel void kernel_cpy_f32_f32_pack(
|
||||
global float * src0,
|
||||
ulong offset0,
|
||||
@@ -241,6 +242,7 @@ kernel void kernel_cpy_f32_f32_pack(
|
||||
dst_data[i00] = src[0];
|
||||
}
|
||||
}
|
||||
#endif // GGML_CL_NO_CPY_PACK
|
||||
|
||||
kernel void kernel_cpy_i32_i32(
|
||||
global int * src0,
|
||||
|
||||
@@ -7,6 +7,14 @@
|
||||
#define REQD_SUBGROUP_SIZE_64 __attribute__((qcom_reqd_sub_group_size("half")))
|
||||
#endif
|
||||
|
||||
// A6X compiler incorrectly constant-folds get_local_size() results;
|
||||
// force runtime materialization via a no-op ALU round-trip.
|
||||
#ifdef GGML_CL_A6X_CONSTFOLD_FIX
|
||||
#define MATERIALIZE_WG(x) do { (x) *= 2u; if ((x) > 1u) (x) /= 2u; } while(0)
|
||||
#else
|
||||
#define MATERIALIZE_WG(x)
|
||||
#endif
|
||||
|
||||
// assume
|
||||
#define QK4_0 32
|
||||
#define N_SIMDGROUP 4
|
||||
@@ -327,6 +335,7 @@ __kernel void kernel_gemv_noshuffle_q4_0_f32_mc3(
|
||||
uint BLOCK_STRIDE_A = N_SIMDGROUP * M; // = 4 * M (N_SIMDGROUP is the #define 4)
|
||||
uint COL_STRIDE = K / 4; // float4 pixels per activation column
|
||||
uint nsg = get_local_size(1); // runtime K-split (4 default, 8 small-M)
|
||||
MATERIALIZE_WG(nsg);
|
||||
|
||||
__private uint4 regA_hi, regA_lo;
|
||||
__private half2 regS;
|
||||
|
||||
@@ -11,6 +11,14 @@
|
||||
#define NSUBGROUPS 4
|
||||
#define SUBGROUP_SIZE 64
|
||||
|
||||
// A6X compiler incorrectly constant-folds get_local_size() results;
|
||||
// force runtime materialization via a no-op ALU round-trip.
|
||||
#ifdef GGML_CL_A6X_CONSTFOLD_FIX
|
||||
#define MATERIALIZE_WG(x) do { (x) *= 2u; if ((x) > 1u) (x) /= 2u; } while(0)
|
||||
#else
|
||||
#define MATERIALIZE_WG(x)
|
||||
#endif
|
||||
|
||||
// scales are transposed: consecutive codes of a row are `stride` apart
|
||||
inline void get_scale_min_k4(
|
||||
int j,
|
||||
@@ -233,6 +241,7 @@ kernel void kernel_gemv_noshuffle_q4_k_f32(
|
||||
// K-split (more waves/SP -> latency hiding) while large-M keeps 4. The
|
||||
// physical weight layout stride below is INDEPENDENT of this (see BLOCK_STRIDE_A).
|
||||
uint nsg = get_local_size(1);
|
||||
MATERIALIZE_WG(nsg);
|
||||
|
||||
uint K = ne00;
|
||||
uint M = ne01;
|
||||
@@ -400,6 +409,7 @@ kernel void kernel_gemv_noshuffle_q4_k_f32_glu(
|
||||
uint gid = get_global_id(0);
|
||||
ushort slid = get_sub_group_local_id();
|
||||
uint nsg = get_local_size(1);
|
||||
MATERIALIZE_WG(nsg);
|
||||
|
||||
uint K = ne00;
|
||||
uint M = ne01;
|
||||
@@ -514,6 +524,7 @@ kernel void kernel_gemv_noshuffle_q4_k_f32_splitk(
|
||||
uint gid = get_global_id(0);
|
||||
ushort slid = get_sub_group_local_id();
|
||||
uint nsg = get_local_size(1);
|
||||
MATERIALIZE_WG(nsg);
|
||||
uint ksplit = get_num_groups(1);
|
||||
uint kslice = get_group_id(1);
|
||||
|
||||
|
||||
@@ -5102,12 +5102,19 @@ static bool ggml_sycl_mul_mat_glu_mmvq_fused(ggml_backend_sycl_context & ctx, gg
|
||||
return false;
|
||||
}
|
||||
|
||||
// quant pairs the reorder kernel cannot serve (mixed gate/up types) take the
|
||||
// standard-layout fused path instead; q4_K keeps the reorder path below
|
||||
if (wg->type != GGML_TYPE_Q4_K || wu->type != GGML_TYPE_Q4_K) {
|
||||
// quant pairs the reorder kernel does not serve (mixed gate/up types, q5_K off BMG) take the
|
||||
// standard-layout fused path instead; same-type q4_K / q5_K keep the reorder path below
|
||||
const bool reorder_pair = wg->type == wu->type &&
|
||||
(wu->type == GGML_TYPE_Q4_K || (wu->type == GGML_TYPE_Q5_K && ggml_sycl_q5_k_mmvq_reuse(ctx.device)));
|
||||
if (!reorder_pair) {
|
||||
return ggml_sycl_mul_mat_glu_mmvq_plain(ctx, glu, gate, up, wu, wg, act);
|
||||
}
|
||||
|
||||
// past 5 columns the two unfused q5_K GEMVs are faster than the fused kernel
|
||||
if (wu->type == GGML_TYPE_Q5_K && act->ne[1] > 5) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// install the reorder (SoA) layout the fused kernel needs, as the unfused mmvq path would;
|
||||
// a no-op once done. after the bail checks so a declined op does not pay for it.
|
||||
opt_for_reorder(&ctx, wu, act, up, mul_mat_algo::MMVQ);
|
||||
|
||||
+60
-48
@@ -110,7 +110,8 @@ static void mul_mat_vec_q_reorder(const void * __restrict__ vx, const void * __r
|
||||
|
||||
// With has_fusion, `vgate` is a second weight matrix sharing vx's shape, stride and reorder
|
||||
// layout: one pass computes both row dot products and the epilogue writes glu(gate, up).
|
||||
template <typename reorder_vec_dot_q_sycl, int ncols_dst, bool has_fusion = false, int rows_per_sg = 1>
|
||||
template <typename reorder_vec_dot_q_sycl, int ncols_dst, bool has_fusion = false, int rows_per_sg = 1,
|
||||
bool shared_weights = reorder_vec_dot_shared_weights<reorder_vec_dot_q_sycl::gtype>::value>
|
||||
static void mul_mat_vec_q_reorder_ncols(const void * __restrict__ vx, const void * __restrict__ vgate,
|
||||
const void * __restrict__ vy, float * __restrict__ dst, const int ncols,
|
||||
const int nrows, const int stride_col_y_bytes, const int stride_col_dst,
|
||||
@@ -181,7 +182,7 @@ static void mul_mat_vec_q_reorder_ncols(const void * __restrict__ vx, const void
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if constexpr (reorder_vec_dot_shared_weights<reorder_vec_dot_q_sycl::gtype>::value) {
|
||||
} else if constexpr (shared_weights) {
|
||||
const int ibx = row0 * blocks_per_row + i;
|
||||
const auto bx_offset = block_type::get_block_offset(ibx, nblocks);
|
||||
const auto d_offset = block_type::get_d_offset(nrows, ncols, ibx);
|
||||
@@ -1945,8 +1946,8 @@ static void reorder_mul_mat_vec_q5_k_q8_1_sycl(const void * vx, const void * vy,
|
||||
});
|
||||
}
|
||||
|
||||
template <int ncols_dst>
|
||||
static void reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols(
|
||||
template <int ncols_dst, int rows_per_sg, bool shared_weights>
|
||||
static void reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols_impl(
|
||||
const void * vx, const void * vy, float * dst,
|
||||
const int ncols, const int nrows,
|
||||
const int stride_col_y_bytes, const int stride_col_dst,
|
||||
@@ -1954,20 +1955,35 @@ static void reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols(
|
||||
GGML_ASSERT(ncols % QK_K == 0);
|
||||
|
||||
constexpr size_t num_subgroups = WARP_SIZE;
|
||||
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups);
|
||||
const int block_num_y = ceil_div(nrows, GGML_SYCL_MMV_Y * (int) num_subgroups * rows_per_sg);
|
||||
const sycl::range<3> block_nums(1, 1, block_num_y);
|
||||
const sycl::range<3> block_dims(1, GGML_SYCL_MMV_Y, num_subgroups * WARP_SIZE);
|
||||
|
||||
stream->submit([&](sycl::handler & cgh) {
|
||||
cgh.parallel_for(sycl::nd_range<3>(block_nums * block_dims, block_dims),
|
||||
[=](sycl::nd_item<3> nd_item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
||||
mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K>, ncols_dst>(
|
||||
mul_mat_vec_q_reorder_ncols<reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K>, ncols_dst,
|
||||
/*has_fusion=*/ false, rows_per_sg, shared_weights>(
|
||||
vx, /*vgate=*/ nullptr, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst,
|
||||
/*glu_op=*/ GGML_GLU_OP_SWIGLU, nd_item);
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
template <int ncols_dst>
|
||||
static void reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols(
|
||||
const void * vx, const void * vy, float * dst,
|
||||
const int ncols, const int nrows,
|
||||
const int stride_col_y_bytes, const int stride_col_dst,
|
||||
dpct::queue_ptr stream) {
|
||||
if (ggml_sycl_q5_k_mmvq_reuse(ggml_sycl_get_device())) {
|
||||
constexpr int rows_per_sg = ncols_dst >= 3 ? 2 : 1;
|
||||
reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols_impl<ncols_dst, rows_per_sg, true>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream);
|
||||
} else {
|
||||
reorder_mul_mat_vec_q5_k_q8_1_sycl_ncols_impl<ncols_dst, 1, false>(vx, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, stream);
|
||||
}
|
||||
}
|
||||
|
||||
static void reorder_mul_mat_vec_q5_k_q8_1_sycl_switch_ncols(
|
||||
const void * vx, const void * vy, float * dst,
|
||||
const int ncols, const int nrows, const int ncols_dst,
|
||||
@@ -3129,8 +3145,11 @@ static void launch_mul_mat_vec_q_reorder_glu(const void * vx, const void * vgate
|
||||
const int ncols, const int nrows, const int stride_col_y_bytes,
|
||||
const int stride_col_dst, const ggml_glu_op glu_op,
|
||||
dpct::queue_ptr stream) {
|
||||
// q4_K pairs rows for 3..4 columns, q5_K for 3..5
|
||||
constexpr int row_pair_max = reorder_vec_dot_q_sycl::gtype == GGML_TYPE_Q5_K ? 5 : 4;
|
||||
constexpr int rows_per_sg =
|
||||
reorder_vec_dot_shared_activations<reorder_vec_dot_q_sycl::gtype>::value && ncols_dst >= 3 && ncols_dst <= 4
|
||||
reorder_vec_dot_shared_activations<reorder_vec_dot_q_sycl::gtype>::value && ncols_dst >= 3 &&
|
||||
ncols_dst <= row_pair_max
|
||||
? 2
|
||||
: 1;
|
||||
launch_mul_mat_vec_q_reorder_glu_impl<reorder_vec_dot_q_sycl, ncols_dst, rows_per_sg>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, glu_op, stream);
|
||||
@@ -3321,55 +3340,48 @@ bool ggml_sycl_mul_mat_vec_q_glu_plain(enum ggml_type gate_type, enum ggml_type
|
||||
return false;
|
||||
}
|
||||
|
||||
template <ggml_type type, int... Ns>
|
||||
static bool mul_mat_vec_q_glu_reorder_ncols(enum ggml_glu_op glu_op, const void * vx, const void * vgate,
|
||||
const void * vy, float * dst, int ncols, int nrows, int ncols_dst,
|
||||
int stride_col_y_bytes, int stride_col_dst, dpct::queue_ptr stream) {
|
||||
using vec_dot = reorder_vec_dot_q_sycl<type>;
|
||||
|
||||
auto launch = [&](auto I) -> bool {
|
||||
constexpr int n = decltype(I)::value;
|
||||
if (ncols_dst != n) {
|
||||
return false;
|
||||
}
|
||||
if constexpr (type == GGML_TYPE_Q4_K && n == 2) {
|
||||
if (nrows >= Q4_K_MMVQ_ROW_PAIR_MIN_NROWS) {
|
||||
launch_mul_mat_vec_q_reorder_glu_impl<vec_dot, 2, 2>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
}
|
||||
}
|
||||
launch_mul_mat_vec_q_reorder_glu<vec_dot, n>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes,
|
||||
stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
};
|
||||
|
||||
// unary fold over launch
|
||||
return (launch(std::integral_constant<int, Ns>{}) || ...);
|
||||
}
|
||||
|
||||
bool ggml_sycl_mul_mat_vec_q_glu_reorder(enum ggml_type src0_type, enum ggml_glu_op glu_op, const void * vx,
|
||||
const void * vgate, const void * vy, float * dst, int ncols, int nrows,
|
||||
int ncols_dst, int stride_col_y_bytes, int stride_col_dst,
|
||||
dpct::queue_ptr stream) {
|
||||
if (src0_type != GGML_TYPE_Q4_K) {
|
||||
return false;
|
||||
}
|
||||
if (glu_op != GGML_GLU_OP_SWIGLU && glu_op != GGML_GLU_OP_GEGLU) {
|
||||
return false;
|
||||
}
|
||||
|
||||
using vec_dot = reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>;
|
||||
|
||||
switch (ncols_dst) {
|
||||
case 1:
|
||||
launch_mul_mat_vec_q_reorder_glu<vec_dot, 1>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes,
|
||||
stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
case 2:
|
||||
if (nrows >= Q4_K_MMVQ_ROW_PAIR_MIN_NROWS) {
|
||||
launch_mul_mat_vec_q_reorder_glu_impl<vec_dot, 2, 2>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, glu_op, stream);
|
||||
} else {
|
||||
launch_mul_mat_vec_q_reorder_glu_impl<vec_dot, 2, 1>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes, stride_col_dst, glu_op, stream);
|
||||
}
|
||||
return true;
|
||||
case 3:
|
||||
launch_mul_mat_vec_q_reorder_glu<vec_dot, 3>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes,
|
||||
stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
case 4:
|
||||
launch_mul_mat_vec_q_reorder_glu<vec_dot, 4>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes,
|
||||
stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
case 5:
|
||||
launch_mul_mat_vec_q_reorder_glu<vec_dot, 5>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes,
|
||||
stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
case 6:
|
||||
launch_mul_mat_vec_q_reorder_glu<vec_dot, 6>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes,
|
||||
stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
case 7:
|
||||
launch_mul_mat_vec_q_reorder_glu<vec_dot, 7>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes,
|
||||
stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
case 8:
|
||||
launch_mul_mat_vec_q_reorder_glu<vec_dot, 8>(vx, vgate, vy, dst, ncols, nrows, stride_col_y_bytes,
|
||||
stride_col_dst, glu_op, stream);
|
||||
return true;
|
||||
switch (src0_type) {
|
||||
case GGML_TYPE_Q4_K:
|
||||
return mul_mat_vec_q_glu_reorder_ncols<GGML_TYPE_Q4_K, 1, 2, 3, 4, 5, 6, 7, 8>(
|
||||
glu_op, vx, vgate, vy, dst, ncols, nrows, ncols_dst, stride_col_y_bytes, stride_col_dst, stream);
|
||||
case GGML_TYPE_Q5_K:
|
||||
// fusion declines q5_K past 5 columns
|
||||
return mul_mat_vec_q_glu_reorder_ncols<GGML_TYPE_Q5_K, 1, 2, 3, 4, 5>(
|
||||
glu_op, vx, vgate, vy, dst, ncols, nrows, ncols_dst, stride_col_y_bytes, stride_col_dst, stream);
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -15,6 +15,12 @@
|
||||
|
||||
#include "common.hpp"
|
||||
|
||||
// q5_K multi-column MMVQ shares weights across columns, pairs rows and fuses gate/up in the reorder
|
||||
// layout: faster on Xe2 (BMG), so untested archs keep the per-column kernel
|
||||
inline bool ggml_sycl_q5_k_mmvq_reuse(int device) {
|
||||
const gpu_arch arch = ggml_sycl_info().devices[device].hw_info.arch;
|
||||
return arch == gpu_arch::intel_gpu_bmg_g21 || arch == gpu_arch::intel_gpu_bmg_g31;
|
||||
}
|
||||
|
||||
void ggml_sycl_op_mul_mat_vec_q(
|
||||
ggml_backend_sycl_context & ctx,
|
||||
|
||||
@@ -362,6 +362,10 @@ template <> struct reorder_vec_dot_shared_weights<GGML_TYPE_Q4_K> {
|
||||
static constexpr bool value = true;
|
||||
};
|
||||
|
||||
template <> struct reorder_vec_dot_shared_weights<GGML_TYPE_Q5_K> {
|
||||
static constexpr bool value = true;
|
||||
};
|
||||
|
||||
template <ggml_type T> struct reorder_vec_dot_shared_activations {
|
||||
static constexpr bool value = false;
|
||||
};
|
||||
@@ -370,6 +374,10 @@ template <> struct reorder_vec_dot_shared_activations<GGML_TYPE_Q4_K> {
|
||||
static constexpr bool value = true;
|
||||
};
|
||||
|
||||
template <> struct reorder_vec_dot_shared_activations<GGML_TYPE_Q5_K> {
|
||||
static constexpr bool value = true;
|
||||
};
|
||||
|
||||
template <> struct reorder_vec_dot_q_sycl<GGML_TYPE_Q4_0> {
|
||||
static constexpr ggml_type gtype = GGML_TYPE_Q4_0;
|
||||
|
||||
@@ -676,56 +684,72 @@ template <> struct reorder_vec_dot_q_sycl<GGML_TYPE_Q5_K> {
|
||||
using q5_k_block = ggml_sycl_reordered::block_q_t<GGML_TYPE_Q5_K>;
|
||||
using q5_k_traits = typename q5_k_block::traits;
|
||||
|
||||
__dpct_inline__ float operator()(const void * __restrict__ vbq, const std::pair<int, int> ibx_offset,
|
||||
const std::pair<int, int> d_offset, const int8_t * q8_1_quant_ptr,
|
||||
const sycl::half2 * q8_1_ds, const int & iqs) {
|
||||
const uint8_t * base = static_cast<const uint8_t *>(vbq);
|
||||
const uint8_t * qs = base + ibx_offset.first; // low 4 bits
|
||||
const uint8_t * qh_base = base + ibx_offset.second; // high bit
|
||||
const uint8_t * scs = base + d_offset.first;
|
||||
const ggml_half2 * dms = reinterpret_cast<const ggml_half2 *>(base + d_offset.second);
|
||||
struct weights {
|
||||
int vl[2];
|
||||
int vh[2];
|
||||
uint16_t aux[2];
|
||||
ggml_half2 dm;
|
||||
};
|
||||
|
||||
// same activation layout as Q4_K
|
||||
static_assert(QR5_K == QR4_K);
|
||||
using activations = reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>::activations;
|
||||
|
||||
__dpct_inline__ static weights load(const void * __restrict__ vbq, const std::pair<int, int> ibx_offset,
|
||||
const std::pair<int, int> d_offset, const int & iqs) {
|
||||
const uint8_t * base = static_cast<const uint8_t *>(vbq);
|
||||
const uint8_t * qs = base + ibx_offset.first; // low 4 bits
|
||||
const uint8_t * qh_base = base + ibx_offset.second; // high bit
|
||||
const uint8_t * scs = base + d_offset.first;
|
||||
const ggml_half2 * dms = reinterpret_cast<const ggml_half2 *>(base + d_offset.second);
|
||||
|
||||
const int bq8_offset = QR5_K * ((iqs / 2) / (QI8_1 / 2));
|
||||
const int * ql_ptr = (const int *) (qs + 16 * bq8_offset + 4 * ((iqs / 2) % 4));
|
||||
const int * qh_ptr = (const int *) (qh_base + 4 * ((iqs / 2) % 4));
|
||||
const uint16_t * scales = (const uint16_t *) scs;
|
||||
|
||||
int vl[2];
|
||||
int vh[2];
|
||||
int u[2 * QR5_K];
|
||||
float d8[QR5_K];
|
||||
weights w;
|
||||
w.vl[0] = ql_ptr[0];
|
||||
w.vl[1] = ql_ptr[4];
|
||||
|
||||
vl[0] = ql_ptr[0];
|
||||
vl[1] = ql_ptr[4];
|
||||
w.vh[0] = qh_ptr[0] >> bq8_offset;
|
||||
w.vh[1] = qh_ptr[4] >> bq8_offset;
|
||||
|
||||
vh[0] = qh_ptr[0] >> bq8_offset;
|
||||
vh[1] = qh_ptr[4] >> bq8_offset;
|
||||
|
||||
uint16_t aux[2];
|
||||
const int j = (QR5_K * ((iqs / 2) / (QI8_1 / 2))) / 2;
|
||||
if (j < 2) {
|
||||
aux[0] = scales[j + 0] & 0x3f3f;
|
||||
aux[1] = scales[j + 2] & 0x3f3f;
|
||||
w.aux[0] = scales[j + 0] & 0x3f3f;
|
||||
w.aux[1] = scales[j + 2] & 0x3f3f;
|
||||
} else {
|
||||
aux[0] = ((scales[j + 2] >> 0) & 0x0f0f) | ((scales[j - 2] & 0xc0c0) >> 2);
|
||||
aux[1] = ((scales[j + 2] >> 4) & 0x0f0f) | ((scales[j - 0] & 0xc0c0) >> 2);
|
||||
w.aux[0] = ((scales[j + 2] >> 0) & 0x0f0f) | ((scales[j - 2] & 0xc0c0) >> 2);
|
||||
w.aux[1] = ((scales[j + 2] >> 4) & 0x0f0f) | ((scales[j - 0] & 0xc0c0) >> 2);
|
||||
}
|
||||
|
||||
const uint8_t * sc = (const uint8_t *) aux;
|
||||
w.dm = *dms;
|
||||
|
||||
return w;
|
||||
}
|
||||
|
||||
__dpct_inline__ static activations load_activations(const int8_t * q8_1_quant_ptr,
|
||||
const sycl::half2 * q8_1_ds, const int & iqs) {
|
||||
return reorder_vec_dot_q_sycl<GGML_TYPE_Q4_K>::load_activations(q8_1_quant_ptr, q8_1_ds, iqs);
|
||||
}
|
||||
|
||||
__dpct_inline__ static float apply(const weights & w, const activations & a) {
|
||||
const uint8_t * sc = (const uint8_t *) w.aux;
|
||||
const uint8_t * m = sc + 2;
|
||||
|
||||
for (int i = 0; i < QR5_K; ++i) {
|
||||
const int8_t* quant_base_ptr = q8_1_quant_ptr + (bq8_offset + i) * QK8_1;
|
||||
sycl::half2 ds_values = *(q8_1_ds + bq8_offset + i);
|
||||
return vec_dot_q5_K_q8_1_impl_vmmq(w.vl, w.vh, a.u, sc, m, w.dm, a.d8);
|
||||
}
|
||||
|
||||
d8[i] = ds_values[0];
|
||||
__dpct_inline__ static float dot(const weights & w, const int8_t * q8_1_quant_ptr,
|
||||
const sycl::half2 * q8_1_ds, const int & iqs) {
|
||||
return apply(w, load_activations(q8_1_quant_ptr, q8_1_ds, iqs));
|
||||
}
|
||||
|
||||
const int * q8 = (const int *) quant_base_ptr + ((iqs / 2) % 4);
|
||||
u[2 * i + 0] = q8[0];
|
||||
u[2 * i + 1] = q8[4];
|
||||
}
|
||||
|
||||
return vec_dot_q5_K_q8_1_impl_vmmq(vl, vh, u, sc, m, *dms, d8);
|
||||
__dpct_inline__ float operator()(const void * __restrict__ vbq, const std::pair<int, int> ibx_offset,
|
||||
const std::pair<int, int> d_offset, const int8_t * q8_1_quant_ptr,
|
||||
const sycl::half2 * q8_1_ds, const int & iqs) {
|
||||
return dot(load(vbq, ibx_offset, d_offset, iqs), q8_1_quant_ptr, q8_1_ds, iqs);
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -9455,6 +9455,8 @@ static void ggml_vk_op_f32(ggml_backend_vk_context * ctx, vk_context& subctx, co
|
||||
elements = { (uint32_t)CEIL_DIV(ne00, 128), 1, 1 };
|
||||
} else {
|
||||
elements = { (uint32_t)ne01, (uint32_t)ne02, (uint32_t)ne03 };
|
||||
elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
|
||||
elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
|
||||
}
|
||||
break;
|
||||
|
||||
@@ -10777,7 +10779,11 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
||||
ggml_vk_tensor_subbuffer(ctx, src0, true),
|
||||
ggml_vk_tensor_subbuffer(ctx, set_rows, true),
|
||||
ggml_vk_tensor_subbuffer(ctx, indices),
|
||||
}, pc, { (uint32_t)src0->ne[1], (uint32_t)src0->ne[2], (uint32_t)src0->ne[3] });
|
||||
}, pc, {
|
||||
(uint32_t)src0->ne[1],
|
||||
std::min((uint32_t)src0->ne[2], ctx->device->properties.limits.maxComputeWorkGroupCount[1]),
|
||||
std::min((uint32_t)src0->ne[3], ctx->device->properties.limits.maxComputeWorkGroupCount[2]),
|
||||
});
|
||||
ggml_vk_rms_norm_finish(ctx, src0);
|
||||
return;
|
||||
}
|
||||
@@ -10824,7 +10830,11 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
||||
ggml_vk_tensor_subbuffer(ctx, dst, true),
|
||||
ggml_vk_tensor_subbuffer(ctx, residual),
|
||||
ggml_vk_tensor_subbuffer(ctx, post_scale),
|
||||
}, pc, { (uint32_t)src0->ne[1], (uint32_t)src0->ne[2], (uint32_t)src0->ne[3] });
|
||||
}, pc, {
|
||||
(uint32_t)src0->ne[1],
|
||||
std::min((uint32_t)src0->ne[2], ctx->device->properties.limits.maxComputeWorkGroupCount[1]),
|
||||
std::min((uint32_t)src0->ne[3], ctx->device->properties.limits.maxComputeWorkGroupCount[2]),
|
||||
});
|
||||
}
|
||||
ggml_vk_rms_norm_finish(ctx, src0);
|
||||
return;
|
||||
@@ -10911,6 +10921,8 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
||||
|
||||
std::array<uint32_t, 3> elements;
|
||||
elements = { (uint32_t)rms->src[0]->ne[1], (uint32_t)rms->src[0]->ne[2], (uint32_t)rms->src[0]->ne[3] };
|
||||
elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
|
||||
elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
|
||||
|
||||
static_assert(max_tensors == 7);
|
||||
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
|
||||
|
||||
@@ -53,101 +53,107 @@ shared FLOAT_TYPE sumsh[BLOCK_SIZE];
|
||||
void rms_norm(uint num_iters) {
|
||||
const uint ncols = p.ne00;
|
||||
const uint nrows = gl_NumWorkGroups.x;
|
||||
const uint nchannels = gl_NumWorkGroups.y;
|
||||
const uint nchannels = p.ne02;
|
||||
const uint nsamples = p.ne03;
|
||||
|
||||
const uint row = gl_WorkGroupID.x;
|
||||
const uint channel = gl_WorkGroupID.y;
|
||||
const uint samp = gl_WorkGroupID.z;
|
||||
const uint tid = gl_LocalInvocationID.x;
|
||||
|
||||
const uint stride_row = p.nb01;
|
||||
const uint stride_channel = p.nb02;
|
||||
const uint stride_sample = p.nb03;
|
||||
|
||||
uint32_t a_offset = samp*stride_sample + channel*stride_channel + row*stride_row + get_aoffset();
|
||||
uint32_t b_offset = src1_idx(0, row, channel, samp) + get_boffset();
|
||||
// grid.y/z are clamped to the device workgroup limit, iterate over the excess channels/samples
|
||||
for (uint samp = gl_WorkGroupID.z; samp < nsamples; samp += gl_NumWorkGroups.z) {
|
||||
for (uint channel = gl_WorkGroupID.y; channel < nchannels; channel += gl_NumWorkGroups.y) {
|
||||
barrier();
|
||||
|
||||
uint32_t a_offset = samp*stride_sample + channel*stride_channel + row*stride_row + get_aoffset();
|
||||
uint32_t b_offset = src1_idx(0, row, channel, samp) + get_boffset();
|
||||
#if RMS_NORM_ROPE_FUSION
|
||||
// Per-row offset in shared memory
|
||||
uint32_t d_offset = 0;
|
||||
// Per-row offset in shared memory
|
||||
uint32_t d_offset = 0;
|
||||
#elif RMS_NORM_SET_ROWS_FUSION
|
||||
uint32_t d_offset = data_i[channel].x*p.nb21 + row*ncols + get_doffset();
|
||||
uint32_t d_offset = data_i[channel].x*p.nb21 + row*ncols + get_doffset();
|
||||
#else
|
||||
uint32_t d_offset = ((samp*nchannels + channel)*nrows + row)*ncols + get_doffset();
|
||||
uint32_t d_offset = ((samp*nchannels + channel)*nrows + row)*ncols + get_doffset();
|
||||
#endif
|
||||
FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp
|
||||
FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp
|
||||
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
FLOAT_TYPE xi = FLOAT_TYPE(0);
|
||||
if (col < ncols) {
|
||||
xi = FLOAT_TYPE(data_a[a_offset + col]);
|
||||
}
|
||||
sum += xi * xi;
|
||||
}
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
FLOAT_TYPE xi = FLOAT_TYPE(0);
|
||||
if (col < ncols) {
|
||||
xi = FLOAT_TYPE(data_a[a_offset + col]);
|
||||
}
|
||||
sum += xi * xi;
|
||||
}
|
||||
|
||||
sumsh[tid] = sum;
|
||||
// sum up partial sums and write back result
|
||||
barrier();
|
||||
[[unroll]] for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
|
||||
if (tid < s) {
|
||||
sum += sumsh[tid + s];
|
||||
sumsh[tid] = sum;
|
||||
}
|
||||
barrier();
|
||||
}
|
||||
sum = sumsh[0];
|
||||
// sum up partial sums and write back result
|
||||
barrier();
|
||||
[[unroll]] for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
|
||||
if (tid < s) {
|
||||
sum += sumsh[tid + s];
|
||||
sumsh[tid] = sum;
|
||||
}
|
||||
barrier();
|
||||
}
|
||||
sum = sumsh[0];
|
||||
|
||||
const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
|
||||
const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
|
||||
const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
|
||||
const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
|
||||
|
||||
if (do_multiply) {
|
||||
if (ncols > p.ne10) {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
|
||||
if (do_multiply) {
|
||||
if (ncols > p.ne10) {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
|
||||
#if RMS_NORM_ADD_FUSION
|
||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||
if (do_post_multiply) {
|
||||
value *= FLOAT_TYPE(data_e[0]);
|
||||
}
|
||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||
if (do_post_multiply) {
|
||||
value *= FLOAT_TYPE(data_e[0]);
|
||||
}
|
||||
#endif
|
||||
data_d[d_offset + col] = D_TYPE(value);
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
|
||||
data_d[d_offset + col] = D_TYPE(value);
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
|
||||
#if RMS_NORM_ADD_FUSION
|
||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||
if (do_post_multiply) {
|
||||
value *= FLOAT_TYPE(data_e[0]);
|
||||
}
|
||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||
if (do_post_multiply) {
|
||||
value *= FLOAT_TYPE(data_e[0]);
|
||||
}
|
||||
#endif
|
||||
data_d[d_offset + col] = D_TYPE(value);
|
||||
data_d[d_offset + col] = D_TYPE(value);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||
if (col >= ncols) {
|
||||
continue;
|
||||
}
|
||||
data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
|
||||
}
|
||||
}
|
||||
#if RMS_NORM_ROPE_FUSION
|
||||
barrier();
|
||||
rope_params rp = p.rope;
|
||||
for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) {
|
||||
if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) {
|
||||
rope_neox(t, row, channel, samp, rp);
|
||||
} else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) {
|
||||
rope_norm(t, row, channel, samp, rp);
|
||||
barrier();
|
||||
rope_params rp = p.rope;
|
||||
for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) {
|
||||
if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) {
|
||||
rope_neox(t, row, channel, samp, rp);
|
||||
} else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) {
|
||||
rope_norm(t, row, channel, samp, rp);
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
void main() {
|
||||
|
||||
@@ -847,8 +847,11 @@ static webgpu_encoded_op ggml_webgpu_set(webgpu_context & ctx,
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_index, src1));
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_index + 1, dst));
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_pad(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -897,8 +900,11 @@ static webgpu_encoded_op ggml_webgpu_pad(webgpu_context & ctx, ggml_tensor * src
|
||||
ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst),
|
||||
};
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_solve_tri(webgpu_context & ctx,
|
||||
@@ -1495,8 +1501,11 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_set_rows(webgpu_context & ct
|
||||
} else {
|
||||
threads = src->ne[0] * src->ne[1] * src->ne[2] * src->ne[3];
|
||||
}
|
||||
uint32_t wg_x = CEIL_DIV(threads, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, 1);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(threads, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
// Workgroup size is a common constant
|
||||
@@ -1557,9 +1566,11 @@ static webgpu_encoded_op ggml_webgpu_get_rows(webgpu_context & ctx,
|
||||
uint32_t blocks_per_row = (uint32_t) (dst->ne[0] / (decisions->vectorized ? 4 : 1));
|
||||
uint32_t total_rows = (uint32_t) (dst->ne[1] * dst->ne[2] * dst->ne[3]);
|
||||
uint32_t total_threads = float_parallel ? blocks_per_row * total_rows : total_rows;
|
||||
uint32_t wg_x = CEIL_DIV(total_threads, decisions->wg_size);
|
||||
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(total_threads, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static void ggml_webgpu_quantize_q8_dispatch(webgpu_context & ctx,
|
||||
@@ -2347,7 +2358,8 @@ static webgpu_encoded_op ggml_webgpu_unary_op(webgpu_context & ctx, ggml_tensor
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
|
||||
}
|
||||
|
||||
uint32_t wg_x, wg_y;
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
@@ -2425,7 +2437,8 @@ static webgpu_encoded_op ggml_webgpu_binary_op(webgpu_context & ctx,
|
||||
}
|
||||
}
|
||||
|
||||
uint32_t wg_x, wg_y;
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
@@ -2556,8 +2569,11 @@ static webgpu_encoded_op ggml_webgpu_concat(webgpu_context & ctx,
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst));
|
||||
}
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_repeat(webgpu_context & ctx, ggml_tensor * src0, ggml_tensor * dst) {
|
||||
@@ -2591,8 +2607,12 @@ static webgpu_encoded_op ggml_webgpu_repeat(webgpu_context & ctx, ggml_tensor *
|
||||
|
||||
webgpu_pipeline pipeline = ctx->shader_lib->get_repeat_pipeline(shader_lib_ctx);
|
||||
auto * decisions = static_cast<ggml_webgpu_generic_shader_decisions *>(pipeline.context.get());
|
||||
uint32_t wg_x = CEIL_DIV(ne, decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ne, decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context & ctx,
|
||||
@@ -2637,6 +2657,7 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context
|
||||
(uint32_t) dst->ne[0],
|
||||
(uint32_t) dst->ne[1],
|
||||
(uint32_t) dst->ne[2],
|
||||
(uint32_t) dst->ne[3],
|
||||
ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(rn_dst, 0)) // epsilon, treated as f32 in the shader
|
||||
};
|
||||
|
||||
@@ -2676,7 +2697,11 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_rms_norm_mul(webgpu_context
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 2, dst));
|
||||
}
|
||||
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(dst));
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
compute_2d_workgroups(ggml_nrows(dst), ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x,
|
||||
wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -2692,6 +2717,7 @@ static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor
|
||||
(uint32_t) src->ne[0],
|
||||
(uint32_t) src->ne[1],
|
||||
(uint32_t) src->ne[2],
|
||||
(uint32_t) src->ne[3],
|
||||
ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(dst, 0)) // epsilon, treated as f32 in the shader
|
||||
};
|
||||
|
||||
@@ -2707,7 +2733,12 @@ static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor
|
||||
if (!decisions->inplace) {
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
|
||||
}
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(src));
|
||||
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
compute_2d_workgroups(ggml_nrows(src), ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x,
|
||||
wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_rope(webgpu_context & ctx,
|
||||
@@ -2796,8 +2827,11 @@ static webgpu_encoded_op ggml_webgpu_rope(webgpu_context & ctx,
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, dst_binding, dst));
|
||||
}
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_glu(webgpu_context & ctx,
|
||||
@@ -2870,8 +2904,11 @@ static webgpu_encoded_op ggml_webgpu_glu(webgpu_context & ctx,
|
||||
}
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, dst_binding, dst));
|
||||
|
||||
uint32_t wg_x = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_scale(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -2909,7 +2946,8 @@ static webgpu_encoded_op ggml_webgpu_scale(webgpu_context & ctx, ggml_tensor * s
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, 1, dst));
|
||||
}
|
||||
|
||||
uint32_t wg_x, wg_y;
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = CEIL_DIV(ggml_nelements(dst), decisions->wg_size);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
@@ -2982,7 +3020,11 @@ static webgpu_encoded_op ggml_webgpu_soft_max(webgpu_context & ctx,
|
||||
entries.push_back(ggml_webgpu_make_tensor_bind_group_entry(ctx, binding_num, dst));
|
||||
}
|
||||
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, ggml_nrows(dst));
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = ggml_nrows(dst);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_argmax(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -3165,8 +3207,11 @@ static webgpu_encoded_op ggml_webgpu_cumsum(webgpu_context & ctx, ggml_tensor *
|
||||
shader_lib_ctx.max_wg_size = ctx->global_ctx->capabilities.limits.maxComputeInvocationsPerWorkgroup;
|
||||
|
||||
webgpu_pipeline pipeline = ctx->shader_lib->get_cumsum_pipeline(shader_lib_ctx);
|
||||
uint32_t wg_x = ggml_nrows(dst);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = ggml_nrows(dst);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static webgpu_encoded_op ggml_webgpu_sum_rows(webgpu_context & ctx, ggml_tensor * src, ggml_tensor * dst) {
|
||||
@@ -3190,8 +3235,11 @@ static webgpu_encoded_op ggml_webgpu_sum_rows(webgpu_context & ctx, ggml_tensor
|
||||
|
||||
webgpu_pipeline pipeline = ctx->shader_lib->get_sum_rows_pipeline(shader_lib_ctx);
|
||||
|
||||
uint32_t wg_x = total_sum ? 1 : ggml_nrows(dst);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x);
|
||||
uint32_t wg_x;
|
||||
uint32_t wg_y;
|
||||
uint32_t total_wg = total_sum ? 1 : ggml_nrows(dst);
|
||||
compute_2d_workgroups(total_wg, ctx->global_ctx->capabilities.limits.maxComputeWorkgroupsPerDimension, wg_x, wg_y);
|
||||
return ggml_backend_webgpu_build(ctx, pipeline, params, entries, wg_x, wg_y);
|
||||
}
|
||||
|
||||
static bool ggml_webgpu_can_fuse_rms_norm_mul(const struct ggml_cgraph * cgraph, int node_idx) {
|
||||
@@ -3356,6 +3404,7 @@ static std::optional<webgpu_encoded_op> ggml_webgpu_encode(webgpu_context ctx,
|
||||
case GGML_OP_TRANSPOSE:
|
||||
case GGML_OP_RESHAPE:
|
||||
return std::nullopt;
|
||||
case GGML_OP_DUP:
|
||||
case GGML_OP_CPY:
|
||||
case GGML_OP_CONT:
|
||||
return ggml_webgpu_cpy(ctx, src0, node);
|
||||
@@ -4419,6 +4468,7 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
|
||||
supports_op = (src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_I32 ||
|
||||
src0->type == GGML_TYPE_I16);
|
||||
break;
|
||||
case GGML_OP_DUP:
|
||||
case GGML_OP_CPY:
|
||||
case GGML_OP_CONT:
|
||||
supports_op = (op->type == GGML_TYPE_F16 || op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_I32) &&
|
||||
|
||||
@@ -53,10 +53,12 @@ var<storage, read_write> dst: array<DataType>;
|
||||
var<uniform> params: Params;
|
||||
#endif
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
if (gid.x < params.ne) {
|
||||
var i = gid.x;
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i < params.ne) {
|
||||
var i = gid_i;
|
||||
let i3 = i / (params.ne2 * params.ne1 * params.ne0);
|
||||
i = i % (params.ne2 * params.ne1 * params.ne0);
|
||||
let i2 = i / (params.ne1 * params.ne0);
|
||||
@@ -72,9 +74,9 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
ni[2] * params.stride_src0_2 +
|
||||
ni[3] * params.stride_src0_3;
|
||||
#ifdef SRC_OVERLAP
|
||||
dst[params.offset_dst + gid.x] = merged_src[params.offset_src0 + src_i];
|
||||
dst[params.offset_dst + gid_i] = merged_src[params.offset_src0 + src_i];
|
||||
#else
|
||||
dst[params.offset_dst + gid.x] = src0[params.offset_src0 + src_i];
|
||||
dst[params.offset_dst + gid_i] = src0[params.offset_src0 + src_i];
|
||||
#endif
|
||||
} else {
|
||||
ni[params.dim] -= params.src0_nedim;
|
||||
@@ -83,9 +85,9 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
ni[2] * params.stride_src1_2 +
|
||||
ni[3] * params.stride_src1_3;
|
||||
#ifdef SRC_OVERLAP
|
||||
dst[params.offset_dst + gid.x] = merged_src[params.offset_src1 + src_i];
|
||||
dst[params.offset_dst + gid_i] = merged_src[params.offset_src1 + src_i];
|
||||
#else
|
||||
dst[params.offset_dst + gid.x] = src1[params.offset_src1 + src_i];
|
||||
dst[params.offset_dst + gid_i] = src1[params.offset_src1 + src_i];
|
||||
#endif
|
||||
}
|
||||
}
|
||||
|
||||
@@ -17,8 +17,11 @@ var<workgroup> shared_sum: array<f32, WG_SIZE>;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(workgroup_id) wid: vec3<u32>,
|
||||
@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(local_invocation_id) lid: vec3<u32>) {
|
||||
let row_idx = params.offset_src + wid.x * params.ne0;
|
||||
|
||||
let wid_i = wid.x + wid.y * num_wg.x;
|
||||
let row_idx = params.offset_src + wid_i * params.ne0;
|
||||
let elems = (params.ne0 + WG_SIZE - 1) / WG_SIZE;
|
||||
var local_sum: f32 = 0.0;
|
||||
for (var col = lid.x * elems; col < (lid.x + 1) * elems && col < params.ne0; col ++) {
|
||||
|
||||
@@ -157,12 +157,15 @@ fn b_value(base: u32) -> DataType {
|
||||
#endif
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
if (gid.x >= params.ne) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i >= params.ne) {
|
||||
return;
|
||||
}
|
||||
|
||||
var i = gid.x;
|
||||
var i = gid_i;
|
||||
let i3 = i / (params.ne2 * params.ne1 * params.ne0);
|
||||
i = i % (params.ne2 * params.ne1 * params.ne0);
|
||||
let i2 = i / (params.ne1 * params.ne0);
|
||||
|
||||
@@ -45,12 +45,15 @@ fn wrap_around(idx: i32, n: u32) -> u32 {
|
||||
}
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
if (gid.x >= params.ne) {
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i >= params.ne) {
|
||||
return;
|
||||
}
|
||||
|
||||
var i = gid.x;
|
||||
var i = gid_i;
|
||||
let dst_plane = params.dst_ne2 * params.dst_ne1 * params.dst_ne0;
|
||||
let i3 = i / dst_plane;
|
||||
i = i % dst_plane;
|
||||
@@ -82,5 +85,5 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
}
|
||||
#endif
|
||||
|
||||
dst[params.offset_dst + gid.x] = value;
|
||||
dst[params.offset_dst + gid_i] = value;
|
||||
}
|
||||
|
||||
@@ -45,9 +45,12 @@ var<storage, read_write> dst: array<DataType>;
|
||||
var<uniform> params: Params;
|
||||
|
||||
@compute @workgroup_size(WG_SIZE)
|
||||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
if (gid.x < params.ne) {
|
||||
var i = gid.x;
|
||||
fn main(@builtin(num_workgroups) num_wg: vec3<u32>,
|
||||
@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
|
||||
let gid_i = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y;
|
||||
if (gid_i < params.ne) {
|
||||
var i = gid_i;
|
||||
let i3 = i / (params.ne2 * params.ne1 * params.ne0);
|
||||
i = i % (params.ne2 * params.ne1 * params.ne0);
|
||||
let i2 = i / (params.ne1 * params.ne0);
|
||||
@@ -65,6 +68,6 @@ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||||
a_i2 * params.stride_src0_2 +
|
||||
a_i3 * params.stride_src0_3;
|
||||
|
||||
dst[params.offset_dst + gid.x] = src0[params.offset_src0 + a_index];
|
||||
dst[params.offset_dst + gid_i] = src0[params.offset_src0 + a_index];
|
||||
}
|
||||
}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user