mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-10 23:07:21 -05:00
Compare commits
45
Commits
b11494
...
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 | ||
|
|
ff5888f999 | ||
|
|
dac3087394 | ||
|
|
03aa006acb | ||
|
|
08246a28f6 | ||
|
|
46baf1f1fe | ||
|
|
097f5b5332 | ||
|
|
d888016041 | ||
|
|
fda1866613 | ||
|
|
ff30363a0e | ||
|
|
000bee54a5 | ||
|
|
37ac634566 |
@@ -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);
|
||||
|
||||
|
||||
+46
-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
|
||||
|
||||
@@ -2352,6 +2351,10 @@ class TextModel(ModelBase):
|
||||
if classifier_pooling not in ("cls", "mean"):
|
||||
raise NotImplementedError(f"Unsupported classifier_pooling: {classifier_pooling}")
|
||||
self.gguf_writer.add_classifier_pooling_type(mode_mapping[classifier_pooling])
|
||||
if (classifier_activation := self.hparams.get("classifier_activation")) is not None:
|
||||
if classifier_activation not in ("gelu", "silu", "tanh"):
|
||||
raise NotImplementedError(f"Unsupported classifier_activation: {classifier_activation}")
|
||||
self.gguf_writer.add_classifier_activation(classifier_activation)
|
||||
|
||||
def _set_vocab_glmedge(self):
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
+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"))
|
||||
|
||||
@@ -803,6 +803,7 @@ User can use the device management in [docs/multi-gpu.md](https://github.com/ggm
|
||||
| GGML_SYCL_ENABLE_GRAPH | 0 (default) or 1 | Enable running computations through SYCL Graphs feature. Disabled by default because SYCL Graph is still on development, no better performance. |
|
||||
| GGML_SYCL_ENABLE_HOST_PINNED_MEM | 0 or 1 (default) | Enable host pinned memory to speed up copy data from host to device. When disable it, host memory will common malloc() on CPU. Disable it when use `--load-model mlock`.|
|
||||
| GGML_SYCL_HOST_PINNED_MEM_2G | 0 (default) or 1 | Limit the max memory allocation to be no more than 2GB when enable host pinned memory. USM allocations above 2 GiB take the relaxed/large-allocation path, which serializes H2D copies with compute and prevents copy/compute overlap. It will impact the startup time. Need more test. Depend on `GGML_SYCL_ENABLE_HOST_PINNED_MEM=1`.|
|
||||
| GGML_SYCL_UPLOAD_STAGING_SLOTS | 4 (default) or non-negative integer | Number of 8 MiB pinned host slots used to stage tensor uploads (model loading), so the host copy of one slot overlaps the transfer of the previous one. Set to 0 to use the old path: a malloc'd bounce buffer and a blocking copy per tensor. |
|
||||
| GGML_SYCL_GET_MEM_API | 0 (default) or 1 | Set to get memory info (free, total) by Level Zero or SYCL API:<br>0 - Level Zero API: support more GPUs, only run on Level Zero running time. When there is an error, fallback to call SYCL API. Depend on GGML_SYCL_SUPPORT_LEVEL_ZERO_API.<br>1 - SYCL API: legacy, support more running time, it can't get the free size of some GPUs (like Arc770). In such case, return the free size as value of total size.|
|
||||
| GGML_SYCL_USE_LEVEL_ZERO_API | 1 (default) or 0 | Use Level Zero API for device memory allocation instead of SYCL. Reduces system RAM usage on Intel dGPUs by avoiding DMA-buf/TTM host memory staging. Requires GGML_SYCL_SUPPORT_LEVEL_ZERO_API=ON at build time. SYCL backend always runs on Level Zero running time even if it's set as OFF (The SYCL api will be usage for memory allocation).|
|
||||
| GGML_SYCL_ENABLE_DNN | 0 or 1 (default)| Enable running computations through oneDNN and always use oneMKL. |
|
||||
|
||||
@@ -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) {
|
||||
@@ -5314,9 +5401,10 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
case GGML_UNARY_OP_CEIL:
|
||||
case GGML_UNARY_OP_ROUND:
|
||||
case GGML_UNARY_OP_TRUNC:
|
||||
// TODO: should become:
|
||||
//return ggml_is_contiguous_rows(op->src[0]);
|
||||
return ggml_is_contiguous(op->src[0]);
|
||||
if (op->src[0]->type == GGML_TYPE_BF16 && ggml_get_unary_op(op) == GGML_UNARY_OP_XIELU) {
|
||||
return false;
|
||||
}
|
||||
return op->src[0]->type == GGML_TYPE_F16 || op->src[0]->type == GGML_TYPE_F32 || op->src[0]->type == GGML_TYPE_BF16;
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
@@ -5658,7 +5746,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
|
||||
return max_bias == 0.0f;
|
||||
}
|
||||
case GGML_OP_ROLL:
|
||||
if(op->src[0]->type == GGML_TYPE_F32 && ggml_is_contiguous(op->src[0])) {
|
||||
if(op->src[0]->type == GGML_TYPE_F32) {
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
@@ -5688,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
|
||||
{
|
||||
@@ -5704,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) \
|
||||
|
||||
+138
-111
@@ -3,38 +3,46 @@
|
||||
|
||||
template <int block_size>
|
||||
static __global__ void norm_f32(
|
||||
const float * x, float * dst, const int ncols, const int64_t stride_row, const int64_t stride_channel,
|
||||
const int64_t stride_sample, const float eps) {
|
||||
const int nrows = gridDim.x;
|
||||
const int nchannels = gridDim.y;
|
||||
const float * x, float * dst, const int ncols, const int nchannels, const int nsamples,
|
||||
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps) {
|
||||
const int nrows = gridDim.x;
|
||||
const int row = blockIdx.x;
|
||||
const int tid = threadIdx.x;
|
||||
|
||||
const int row = blockIdx.x;
|
||||
const int channel = blockIdx.y;
|
||||
const int sample = blockIdx.z;
|
||||
const int tid = threadIdx.x;
|
||||
|
||||
x += sample*stride_sample + channel*stride_channel + row*stride_row;
|
||||
dst += ((sample*nchannels + channel)*nrows + row)*ncols;
|
||||
|
||||
float2 mean_var = make_float2(0.0f, 0.0f);
|
||||
extern __shared__ float2 s_sum2[];
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
const float xi = x[col];
|
||||
mean_var.x += xi;
|
||||
mean_var.y += xi * xi;
|
||||
}
|
||||
|
||||
// sum up partial sums
|
||||
extern __shared__ float2 s_sum2[];
|
||||
mean_var = block_reduce<block_reduce_method::SUM, block_size>(mean_var, s_sum2);
|
||||
// grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
|
||||
for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
|
||||
for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
|
||||
const float * xc = x + sample*stride_sample + channel*stride_channel + row*stride_row;
|
||||
float * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;
|
||||
|
||||
const float mean = mean_var.x / ncols;
|
||||
const float var = mean_var.y / ncols - mean * mean;
|
||||
const float inv_std = rsqrtf(var + eps);
|
||||
float2 mean_var = make_float2(0.0f, 0.0f);
|
||||
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
dst[col] = (x[col] - mean) * inv_std;
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
const float xi = xc[col];
|
||||
mean_var.x += xi;
|
||||
mean_var.y += xi * xi;
|
||||
}
|
||||
|
||||
// sum up partial sums
|
||||
mean_var = block_reduce<block_reduce_method::SUM, block_size>(mean_var, s_sum2);
|
||||
|
||||
const float mean = mean_var.x / ncols;
|
||||
const float var = mean_var.y / ncols - mean * mean;
|
||||
const float inv_std = rsqrtf(var + eps);
|
||||
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
dstc[col] = (xc[col] - mean) * inv_std;
|
||||
}
|
||||
|
||||
if constexpr (block_size > WARP_SIZE) {
|
||||
// sync is needed as we reuse s_sum2 across block_reduce invocations, see #26385
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -77,6 +85,8 @@ template <int block_size, bool do_multiply = false, bool do_add = false, bool do
|
||||
static __global__ void rms_norm_f32(const float * x,
|
||||
float * dst,
|
||||
const int ncols,
|
||||
const int nchannels,
|
||||
const int nsamples,
|
||||
const int64_t stride_row,
|
||||
const int64_t stride_channel,
|
||||
const int64_t stride_sample,
|
||||
@@ -99,61 +109,71 @@ static __global__ void rms_norm_f32(const float * x,
|
||||
const uint3 add_nsamples_packed = make_uint3(0, 0, 0),
|
||||
const float scale_out = 1.0f) {
|
||||
ggml_cuda_pdl_lc();
|
||||
const int nrows = gridDim.x;
|
||||
const int nchannels = gridDim.y;
|
||||
|
||||
const int row = blockIdx.x;
|
||||
const int channel = blockIdx.y;
|
||||
const int sample = blockIdx.z;
|
||||
const int tid = threadIdx.x;
|
||||
const int nrows = gridDim.x;
|
||||
const int row = blockIdx.x;
|
||||
const int tid = threadIdx.x;
|
||||
|
||||
static_assert(!do_add || do_multiply, "fusing add is not supported without multiplying");
|
||||
static_assert(!do_scale || !do_multiply, "fusing scale is not supported with multiplying");
|
||||
|
||||
x += sample*stride_sample + channel*stride_channel + row*stride_row;
|
||||
dst += ((sample*nchannels + channel)*nrows + row)*ncols;
|
||||
|
||||
if constexpr (do_multiply) {
|
||||
const uint32_t mul_row = fastmodulo(row, mul_nrows_packed);
|
||||
const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
|
||||
const uint32_t mul_sample = fastmodulo(sample, mul_nsamples_packed);
|
||||
mul += mul_sample * mul_stride_sample + mul_channel * mul_stride_channel + mul_row * mul_stride_row;
|
||||
}
|
||||
|
||||
if constexpr (do_add) {
|
||||
const int add_row = fastmodulo(row, add_nrows_packed);
|
||||
const int add_channel = fastmodulo(channel, add_nchannels_packed);
|
||||
const int add_sample = fastmodulo(sample, add_nsamples_packed);
|
||||
add += add_sample * add_stride_sample + add_channel * add_stride_channel + add_row * add_stride_row;
|
||||
}
|
||||
|
||||
float tmp = 0.0f; // partial sum for thread in warp
|
||||
extern __shared__ float s_sum[];
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
const float xi = x[col];
|
||||
tmp += xi * xi;
|
||||
}
|
||||
|
||||
// sum up partial sums
|
||||
extern __shared__ float s_sum[];
|
||||
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
|
||||
// grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
|
||||
for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
|
||||
for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
|
||||
const float * xc = x + sample*stride_sample + channel*stride_channel + row*stride_row;
|
||||
float * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;
|
||||
|
||||
const float mean = tmp / ncols;
|
||||
const float scale = rsqrtf(mean + eps);
|
||||
[[maybe_unused]] const float * mulc = nullptr;
|
||||
if constexpr (do_multiply) {
|
||||
const uint32_t mul_row = fastmodulo(row, mul_nrows_packed);
|
||||
const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
|
||||
const uint32_t mul_sample = fastmodulo(sample, mul_nsamples_packed);
|
||||
mulc = mul + mul_sample * mul_stride_sample + mul_channel * mul_stride_channel + mul_row * mul_stride_row;
|
||||
}
|
||||
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
if constexpr (do_multiply && do_add) {
|
||||
const int mul_col = fastmodulo(col, mul_ncols_packed);
|
||||
const int add_col = fastmodulo(col, add_ncols_packed);
|
||||
dst[col] = scale * x[col] * mul[mul_col] + add[add_col];
|
||||
} else if constexpr (do_multiply) {
|
||||
const int mul_col = fastmodulo(col, mul_ncols_packed);
|
||||
dst[col] = scale * x[col] * mul[mul_col];
|
||||
} else if constexpr (do_scale) {
|
||||
dst[col] = scale_out * (scale * x[col]);
|
||||
} else {
|
||||
dst[col] = scale * x[col];
|
||||
[[maybe_unused]] const float * addc = nullptr;
|
||||
if constexpr (do_add) {
|
||||
const int add_row = fastmodulo(row, add_nrows_packed);
|
||||
const int add_channel = fastmodulo(channel, add_nchannels_packed);
|
||||
const int add_sample = fastmodulo(sample, add_nsamples_packed);
|
||||
addc = add + add_sample * add_stride_sample + add_channel * add_stride_channel + add_row * add_stride_row;
|
||||
}
|
||||
|
||||
float tmp = 0.0f; // partial sum for thread in warp
|
||||
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
const float xi = xc[col];
|
||||
tmp += xi * xi;
|
||||
}
|
||||
|
||||
// sum up partial sums
|
||||
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
|
||||
|
||||
const float mean = tmp / ncols;
|
||||
const float scale = rsqrtf(mean + eps);
|
||||
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
if constexpr (do_multiply && do_add) {
|
||||
const int mul_col = fastmodulo(col, mul_ncols_packed);
|
||||
const int add_col = fastmodulo(col, add_ncols_packed);
|
||||
dstc[col] = scale * xc[col] * mulc[mul_col] + addc[add_col];
|
||||
} else if constexpr (do_multiply) {
|
||||
const int mul_col = fastmodulo(col, mul_ncols_packed);
|
||||
dstc[col] = scale * xc[col] * mulc[mul_col];
|
||||
} else if constexpr (do_scale) {
|
||||
dstc[col] = scale_out * (scale * xc[col]);
|
||||
} else {
|
||||
dstc[col] = scale * xc[col];
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (block_size > WARP_SIZE) {
|
||||
// sync is needed as we reuse s_sum across block_reduce invocations, see #26385
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -247,50 +267,57 @@ static __global__ void rms_norm_back_f32(
|
||||
|
||||
template <int block_size>
|
||||
static __global__ void l2_norm_f32(
|
||||
const float * x, float * dst, const int ncols, const int64_t stride_row, const int64_t stride_channel,
|
||||
const int64_t stride_sample, const float eps) {
|
||||
const int nrows = gridDim.x;
|
||||
const int nchannels = gridDim.y;
|
||||
const float * x, float * dst, const int ncols, const int nchannels, const int nsamples,
|
||||
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps) {
|
||||
const int nrows = gridDim.x;
|
||||
const int row = blockIdx.x;
|
||||
const int tid = threadIdx.x;
|
||||
|
||||
const int row = blockIdx.x;
|
||||
const int channel = blockIdx.y;
|
||||
const int sample = blockIdx.z;
|
||||
const int tid = threadIdx.x;
|
||||
|
||||
x += sample*stride_sample + channel*stride_channel + row*stride_row;
|
||||
dst += ((sample*nchannels + channel)*nrows + row)*ncols;
|
||||
|
||||
float tmp = 0.0f; // partial sum for thread in warp
|
||||
extern __shared__ float s_sum[];
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
const float xi = x[col];
|
||||
tmp += xi * xi;
|
||||
}
|
||||
|
||||
// sum up partial sums
|
||||
extern __shared__ float s_sum[];
|
||||
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
|
||||
ggml_cuda_pdl_lc();
|
||||
// grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
|
||||
for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
|
||||
for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
|
||||
const float * xc = x + sample*stride_sample + channel*stride_channel + row*stride_row;
|
||||
float * dstc = dst + ((sample*nchannels + channel)*nrows + row)*ncols;
|
||||
|
||||
// from https://pytorch.org/docs/stable/generated/torch.nn.functional.normalize.html
|
||||
const float scale = rsqrtf(fmaxf(tmp, eps * eps));
|
||||
float tmp = 0.0f; // partial sum for thread in warp
|
||||
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
dst[col] = scale * x[col];
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
const float xi = xc[col];
|
||||
tmp += xi * xi;
|
||||
}
|
||||
|
||||
// sum up partial sums
|
||||
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
|
||||
|
||||
// from https://pytorch.org/docs/stable/generated/torch.nn.functional.normalize.html
|
||||
const float scale = rsqrtf(fmaxf(tmp, eps * eps));
|
||||
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
dstc[col] = scale * xc[col];
|
||||
}
|
||||
|
||||
if constexpr (block_size > WARP_SIZE) {
|
||||
// sync is needed as we reuse s_sum across block_reduce invocations, see #26385
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
static void norm_f32_cuda(
|
||||
const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
|
||||
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream) {
|
||||
const dim3 blocks_num(nrows, nchannels, nsamples);
|
||||
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
|
||||
if (ncols < 1024) {
|
||||
const dim3 block_dims(WARP_SIZE, 1, 1);
|
||||
norm_f32<WARP_SIZE><<<blocks_num, block_dims, 0, stream>>>(x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
|
||||
norm_f32<WARP_SIZE><<<blocks_num, block_dims, 0, stream>>>(x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
|
||||
} else {
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
norm_f32<1024><<<blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float2): 0, stream>>>(x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
|
||||
norm_f32<1024><<<blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float2): 0, stream>>>(x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -310,19 +337,19 @@ static void rms_norm_f32_cuda(
|
||||
const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
|
||||
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream,
|
||||
const float scale_out = 1.0f) {
|
||||
const dim3 blocks_num(nrows, nchannels, nsamples);
|
||||
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
|
||||
if (ncols < 1024) {
|
||||
const dim3 block_dims(256, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<256, false, false, do_scale>, launch_params,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
|
||||
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps,
|
||||
// underlying cudaLaunchKernelEx does not support default params
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0),
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), scale_out);
|
||||
} else {
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<1024, false, false, do_scale>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps,
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<1024, false, false, do_scale>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps,
|
||||
// underlying cudaLaunchKernelEx does not support default params
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0),
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), scale_out);
|
||||
@@ -356,7 +383,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
const uint32_t add_nsamples,
|
||||
const float eps,
|
||||
cudaStream_t stream) {
|
||||
const dim3 blocks_num(nrows, nchannels, nsamples);
|
||||
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
|
||||
if (mul == nullptr) {
|
||||
rms_norm_f32_cuda(x, dst, ncols, nrows, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, stream);
|
||||
return;
|
||||
@@ -370,7 +397,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
const dim3 block_dims(256, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<256, true>, launch_params,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
|
||||
// underlying cudaLaunchKernelEx does not support default params
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), 1.0f);
|
||||
@@ -378,7 +405,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<1024, true>, launch_params,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
|
||||
// underlying cudaLaunchKernelEx does not support default params
|
||||
nullptr, 0, 0, 0, make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), make_uint3(0, 0, 0), 1.0f);
|
||||
@@ -397,7 +424,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
const dim3 block_dims(256, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims,block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<256, true, true>, launch_params,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed, add,
|
||||
add_stride_row, add_stride_channel, add_stride_sample, add_ncols_packed, add_nrows_packed,
|
||||
add_nchannels_packed, add_nsamples_packed, 1.0f);
|
||||
@@ -405,7 +432,7 @@ static void rms_norm_mul_f32_cuda(const float * x,
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(rms_norm_f32<1024, true, true>, launch_params,
|
||||
x, dst, ncols, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps, mul, mul_stride_row, mul_stride_channel,
|
||||
mul_stride_sample, mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed, add,
|
||||
add_stride_row, add_stride_channel, add_stride_sample, add_ncols_packed, add_nrows_packed,
|
||||
add_nchannels_packed, add_nsamples_packed, 1.0f);
|
||||
@@ -426,15 +453,15 @@ static void rms_norm_back_f32_cuda(const float * grad, const float * xf, float *
|
||||
static void l2_norm_f32_cuda(
|
||||
const float * x, float * dst, const int ncols, const int nrows, const int nchannels, const int nsamples,
|
||||
const int64_t stride_row, const int64_t stride_channel, const int64_t stride_sample, const float eps, cudaStream_t stream) {
|
||||
const dim3 blocks_num(nrows, nchannels, nsamples);
|
||||
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
|
||||
if (ncols < 1024) {
|
||||
const dim3 block_dims(WARP_SIZE, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, 0, stream};
|
||||
ggml_cuda_kernel_launch(l2_norm_f32<WARP_SIZE>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
|
||||
ggml_cuda_kernel_launch(l2_norm_f32<WARP_SIZE>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
|
||||
} else {
|
||||
const dim3 block_dims(1024, 1, 1);
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params{blocks_num, block_dims, block_dims.x > WARP_SIZE ? 32 * sizeof(float): 0, stream};
|
||||
ggml_cuda_kernel_launch(l2_norm_f32<1024>, launch_params, x, dst, ncols, stride_row, stride_channel, stride_sample, eps);
|
||||
ggml_cuda_kernel_launch(l2_norm_f32<1024>, launch_params, x, dst, ncols, nchannels, nsamples, stride_row, stride_channel, stride_sample, eps);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
+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);
|
||||
|
||||
@@ -17,6 +17,10 @@ static __global__ void roll_f32_cuda(const float * __restrict__ src,
|
||||
const int64_t ne01,
|
||||
const int64_t ne02,
|
||||
const int64_t ne03,
|
||||
const int64_t nb00,
|
||||
const int64_t nb01,
|
||||
const int64_t nb02,
|
||||
const int64_t nb03,
|
||||
const int s0,
|
||||
const int s1,
|
||||
const int s2,
|
||||
@@ -39,7 +43,7 @@ static __global__ void roll_f32_cuda(const float * __restrict__ src,
|
||||
const int64_t d3 = wrap_index(i3 - s3, ne03);
|
||||
|
||||
dst[i3 * (ne00 * ne01 * ne02) + i2 * (ne01 * ne00) + i1 * ne00 + i0] =
|
||||
src[d3 * (ne00 * ne01 * ne02) + d2 * (ne01 * ne00) + d1 * ne00 + d0];
|
||||
src[(d3 * nb03 + d2 * nb02 + d1 * nb01 + d0 * nb00) / sizeof(float)];
|
||||
}
|
||||
|
||||
void ggml_cuda_op_roll(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
@@ -63,5 +67,5 @@ void ggml_cuda_op_roll(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
int64_t num_blocks = (sz + CUDA_ROLL_BLOCK_SIZE - 1) / CUDA_ROLL_BLOCK_SIZE;
|
||||
|
||||
roll_f32_cuda<<<num_blocks, CUDA_ROLL_BLOCK_SIZE, 0, stream>>>(
|
||||
src0_d, dst_d, ne00, ne01, ne02, ne03, s0, s1, s2, s3);
|
||||
src0_d, dst_d, ne00, ne01, ne02, ne03, nb00, nb01, nb02, nb03, s0, s1, s2, s3);
|
||||
}
|
||||
|
||||
+70
-60
@@ -709,7 +709,7 @@ void ggml_cuda_op_rope_fused(ggml_backend_cuda_context & ctx, ggml_tensor * rope
|
||||
// one block per row: block_reduce gives the norm scale, then each thread applies mul and rope to the elements it owns
|
||||
template <int block_size, bool has_ff, typename D>
|
||||
static __global__ void rms_norm_mul_rope_f32(
|
||||
const float * x, D * dst, const int ncols,
|
||||
const float * x, D * dst, const int ncols, const int nchannels, const int nsamples,
|
||||
const int64_t s01, const int64_t s02, const int64_t s03,
|
||||
const int64_t s1, const int64_t s2, const int64_t s3,
|
||||
const float eps,
|
||||
@@ -724,66 +724,76 @@ static __global__ void rms_norm_mul_rope_f32(
|
||||
const int64_t * row_indices, const int set_rows_stride,
|
||||
const bool is_neox) {
|
||||
ggml_cuda_pdl_lc();
|
||||
const int row = blockIdx.x;
|
||||
const int channel = blockIdx.y;
|
||||
const int sample = blockIdx.z;
|
||||
const int tid = threadIdx.x;
|
||||
|
||||
x += sample*s03 + channel*s02 + row*s01;
|
||||
|
||||
const uint32_t mul_row = fastmodulo(row, mul_nrows_packed);
|
||||
const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
|
||||
const uint32_t mul_sample = fastmodulo(sample, mul_nsamples_packed);
|
||||
mul += mul_sample*mul_s03 + mul_channel*mul_s02 + mul_row*mul_s01;
|
||||
|
||||
float tmp = 0.0f;
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
const float xi = x[col];
|
||||
tmp += xi * xi;
|
||||
}
|
||||
const int row = blockIdx.x;
|
||||
const int tid = threadIdx.x;
|
||||
|
||||
extern __shared__ float s_sum[];
|
||||
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
|
||||
|
||||
const float scale = rsqrtf(tmp/ncols + eps);
|
||||
ggml_cuda_pdl_sync();
|
||||
|
||||
int64_t idst = sample*s3 + channel*s2 + row*s1;
|
||||
if (set_rows_stride != 0) {
|
||||
idst = row*s1 + row_indices[channel]*set_rows_stride;
|
||||
}
|
||||
dst += idst;
|
||||
// grid.y and grid.z are clamped to the CUDA limit, iterate over the excess channels/samples
|
||||
for (int sample = blockIdx.z; sample < nsamples; sample += gridDim.z) {
|
||||
for (int channel = blockIdx.y; channel < nchannels; channel += gridDim.y) {
|
||||
const float * xc = x + sample*s03 + channel*s02 + row*s01;
|
||||
|
||||
for (int i0 = 2*tid; i0 < ncols; i0 += 2*block_size) {
|
||||
int ix0;
|
||||
int ix1;
|
||||
if (is_neox && i0 < n_dims) {
|
||||
ix0 = i0/2;
|
||||
ix1 = i0/2 + n_dims/2;
|
||||
} else {
|
||||
ix0 = i0 + 0;
|
||||
ix1 = i0 + 1;
|
||||
const uint32_t mul_row = fastmodulo(row, mul_nrows_packed);
|
||||
const uint32_t mul_channel = fastmodulo(channel, mul_nchannels_packed);
|
||||
const uint32_t mul_sample = fastmodulo(sample, mul_nsamples_packed);
|
||||
const float * mulc = mul + mul_sample*mul_s03 + mul_channel*mul_s02 + mul_row*mul_s01;
|
||||
|
||||
float tmp = 0.0f;
|
||||
|
||||
for (int col = tid; col < ncols; col += block_size) {
|
||||
const float xi = xc[col];
|
||||
tmp += xi * xi;
|
||||
}
|
||||
|
||||
tmp = block_reduce<block_reduce_method::SUM, block_size>(tmp, s_sum);
|
||||
|
||||
const float scale = rsqrtf(tmp/ncols + eps);
|
||||
|
||||
int64_t idst = sample*s3 + channel*s2 + row*s1;
|
||||
if (set_rows_stride != 0) {
|
||||
idst = row*s1 + row_indices[channel]*set_rows_stride;
|
||||
}
|
||||
D * dstc = dst + idst;
|
||||
|
||||
for (int i0 = 2*tid; i0 < ncols; i0 += 2*block_size) {
|
||||
int ix0;
|
||||
int ix1;
|
||||
if (is_neox && i0 < n_dims) {
|
||||
ix0 = i0/2;
|
||||
ix1 = i0/2 + n_dims/2;
|
||||
} else {
|
||||
ix0 = i0 + 0;
|
||||
ix1 = i0 + 1;
|
||||
}
|
||||
|
||||
const float x0 = scale * xc[ix0] * mulc[fastmodulo(ix0, mul_ncols_packed)];
|
||||
const float x1 = scale * xc[ix1] * mulc[fastmodulo(ix1, mul_ncols_packed)];
|
||||
|
||||
if (i0 >= n_dims) {
|
||||
dstc[ix0] = ggml_cuda_cast<D>(x0);
|
||||
dstc[ix1] = ggml_cuda_cast<D>(x1);
|
||||
continue;
|
||||
}
|
||||
|
||||
const float theta_base = pos[channel]*powf(theta_scale, i0/2.0f);
|
||||
const float freq_factor = has_ff ? freq_factors[i0/2] : 1.0f;
|
||||
|
||||
float cos_theta;
|
||||
float sin_theta;
|
||||
rope_yarn<true>(theta_base/freq_factor, freq_scale, corr_dims, i0, ext_factor, attn_factor, cos_theta, sin_theta);
|
||||
|
||||
dstc[ix0] = ggml_cuda_cast<D>(x0*cos_theta - x1*sin_theta);
|
||||
dstc[ix1] = ggml_cuda_cast<D>(x0*sin_theta + x1*cos_theta);
|
||||
}
|
||||
|
||||
if constexpr (block_size > WARP_SIZE) {
|
||||
// sync is needed as we reuse s_sum across block_reduce invocations, see #26385
|
||||
__syncthreads();
|
||||
}
|
||||
}
|
||||
|
||||
const float x0 = scale * x[ix0] * mul[fastmodulo(ix0, mul_ncols_packed)];
|
||||
const float x1 = scale * x[ix1] * mul[fastmodulo(ix1, mul_ncols_packed)];
|
||||
|
||||
if (i0 >= n_dims) {
|
||||
dst[ix0] = ggml_cuda_cast<D>(x0);
|
||||
dst[ix1] = ggml_cuda_cast<D>(x1);
|
||||
continue;
|
||||
}
|
||||
|
||||
const float theta_base = pos[channel]*powf(theta_scale, i0/2.0f);
|
||||
const float freq_factor = has_ff ? freq_factors[i0/2] : 1.0f;
|
||||
|
||||
float cos_theta;
|
||||
float sin_theta;
|
||||
rope_yarn<true>(theta_base/freq_factor, freq_scale, corr_dims, i0, ext_factor, attn_factor, cos_theta, sin_theta);
|
||||
|
||||
dst[ix0] = ggml_cuda_cast<D>(x0*cos_theta - x1*sin_theta);
|
||||
dst[ix1] = ggml_cuda_cast<D>(x0*sin_theta + x1*cos_theta);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -806,7 +816,7 @@ static void rms_norm_mul_rope_cuda(
|
||||
const bool is_neox, cudaStream_t stream) {
|
||||
GGML_ASSERT(ncols % 2 == 0);
|
||||
|
||||
const dim3 blocks_num(nrows, nchannels, nsamples);
|
||||
const dim3 blocks_num(nrows, MIN(nchannels, UINT16_MAX), MIN(nsamples, UINT16_MAX));
|
||||
|
||||
const float theta_scale = powf(freq_base, -2.0f/n_dims);
|
||||
|
||||
@@ -820,13 +830,13 @@ static void rms_norm_mul_rope_cuda(
|
||||
const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, 32*sizeof(float), stream};
|
||||
if (freq_factors == nullptr) {
|
||||
ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<256, false, D>, launch_params,
|
||||
x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
|
||||
x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
|
||||
mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
|
||||
n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
|
||||
freq_factors, row_indices, set_rows_stride, is_neox);
|
||||
} else {
|
||||
ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<256, true, D>, launch_params,
|
||||
x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
|
||||
x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
|
||||
mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
|
||||
n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
|
||||
freq_factors, row_indices, set_rows_stride, is_neox);
|
||||
@@ -836,13 +846,13 @@ static void rms_norm_mul_rope_cuda(
|
||||
const ggml_cuda_kernel_launch_params launch_params = {blocks_num, block_dims, 32*sizeof(float), stream};
|
||||
if (freq_factors == nullptr) {
|
||||
ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<1024, false, D>, launch_params,
|
||||
x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
|
||||
x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
|
||||
mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
|
||||
n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
|
||||
freq_factors, row_indices, set_rows_stride, is_neox);
|
||||
} else {
|
||||
ggml_cuda_kernel_launch(rms_norm_mul_rope_f32<1024, true, D>, launch_params,
|
||||
x, dst, ncols, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
|
||||
x, dst, ncols, nchannels, nsamples, s01, s02, s03, s1, s2, s3, eps, mul, mul_s01, mul_s02, mul_s03,
|
||||
mul_ncols_packed, mul_nrows_packed, mul_nchannels_packed, mul_nsamples_packed,
|
||||
n_dims, pos, freq_scale, ext_factor, attn_factor, corr_dims, theta_scale,
|
||||
freq_factors, row_indices, set_rows_stride, is_neox);
|
||||
|
||||
@@ -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) {
|
||||
@@ -134,24 +134,67 @@ static void unary_cuda(const T * x, T * dst, const int k, cudaStream_t stream) {
|
||||
ggml_cuda_kernel_launch(unary_op_kernel<op, T>, launch_params, x, dst, k);
|
||||
}
|
||||
|
||||
template <float (*op)(float), typename T>
|
||||
static __global__ void unary_op_kernel_strided(const T * x, T * dst, const int k,const int64_t ne00,const int64_t ne01,const int64_t ne02,const size_t nb00,const size_t nb01,const size_t nb02,const size_t nb03) {
|
||||
ggml_cuda_pdl_lc();
|
||||
const int i = blockDim.x*blockIdx.x + threadIdx.x;
|
||||
|
||||
if (i >= k) {
|
||||
return;
|
||||
}
|
||||
|
||||
int64_t rem = i;
|
||||
const int64_t i0 = rem % ne00; rem /= ne00;
|
||||
const int64_t i1 = rem % ne01; rem /= ne01;
|
||||
const int64_t i2 = rem % ne02;
|
||||
const int64_t i3 = rem / ne02;
|
||||
const size_t src_byte_offset = i0 * nb00 + i1 * nb01 + i2 * nb02 + i3 * nb03;
|
||||
const T * src_ptr = (const T *)((const char *)x + src_byte_offset);
|
||||
|
||||
ggml_cuda_pdl_sync();
|
||||
dst[i] = ggml_cuda_cast<T>(op(ggml_cuda_cast<float>(*src_ptr)));
|
||||
}
|
||||
|
||||
template <float (*op)(float), typename T>
|
||||
static void unary_cuda_strided(const T * x, T * dst, const int k,const int64_t ne00,const int64_t ne01,const int64_t ne02,const size_t nb00,const size_t nb01,const size_t nb02,const size_t nb03, cudaStream_t stream) {
|
||||
const int num_blocks = (k + CUDA_NEG_BLOCK_SIZE - 1) / CUDA_NEG_BLOCK_SIZE;
|
||||
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params((dim3)num_blocks, CUDA_NEG_BLOCK_SIZE, 0, stream);
|
||||
ggml_cuda_kernel_launch(unary_op_kernel_strided<op, T>, launch_params, x, dst, k, ne00,ne01,ne02,nb00,nb01,nb02,nb03);
|
||||
}
|
||||
|
||||
template <float (*op)(float)>
|
||||
void ggml_cuda_op_unary(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
|
||||
const ggml_tensor * src0 = dst->src[0];
|
||||
const void * src0_d = src0->data;
|
||||
void * dst_d = dst->data;
|
||||
cudaStream_t stream = ctx.stream();
|
||||
|
||||
GGML_ASSERT(ggml_is_contiguous(src0));
|
||||
cudaStream_t stream = ctx.stream();
|
||||
|
||||
GGML_ASSERT(src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_BF16);
|
||||
GGML_ASSERT(src0->type == dst->type);
|
||||
|
||||
if (src0->type == GGML_TYPE_F16) {
|
||||
unary_cuda<op>((const half *)src0_d, (half *)dst_d, ggml_nelements(src0), stream);
|
||||
} else if (src0->type == GGML_TYPE_BF16) {
|
||||
unary_cuda<op>((const nv_bfloat16 *)src0_d, (nv_bfloat16 *)dst_d, ggml_nelements(src0), stream);
|
||||
if (ggml_is_contiguous(src0)) {
|
||||
if (src0->type == GGML_TYPE_F16) {
|
||||
unary_cuda<op>((const half *)src0_d, (half *)dst_d, ggml_nelements(src0), stream);
|
||||
} else if (src0->type == GGML_TYPE_BF16) {
|
||||
unary_cuda<op>((const nv_bfloat16 *)src0_d, (nv_bfloat16 *)dst_d, ggml_nelements(src0), stream);
|
||||
} else {
|
||||
unary_cuda<op>((const float *)src0_d, (float *)dst_d, ggml_nelements(src0), stream);
|
||||
}
|
||||
} else {
|
||||
unary_cuda<op>((const float *)src0_d, (float *)dst_d, ggml_nelements(src0), stream);
|
||||
if (src0->type == GGML_TYPE_F16) {
|
||||
unary_cuda_strided<op>((const half *)src0_d, (half *)dst_d, ggml_nelements(src0),
|
||||
src0->ne[0], src0->ne[1], src0->ne[2],
|
||||
src0->nb[0], src0->nb[1], src0->nb[2], src0->nb[3], stream);
|
||||
} else if (src0->type == GGML_TYPE_BF16) {
|
||||
unary_cuda_strided<op>((const nv_bfloat16 *)src0_d, (nv_bfloat16 *)dst_d, ggml_nelements(src0),
|
||||
src0->ne[0], src0->ne[1], src0->ne[2],
|
||||
src0->nb[0], src0->nb[1], src0->nb[2], src0->nb[3], stream);
|
||||
} else {
|
||||
unary_cuda_strided<op>((const float *)src0_d, (float *)dst_d, ggml_nelements(src0),
|
||||
src0->ne[0], src0->ne[1], src0->ne[2],
|
||||
src0->nb[0], src0->nb[1], src0->nb[2], src0->nb[3], stream);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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; \
|
||||
|
||||
@@ -425,7 +425,7 @@ static void hvx_mv_2d_repacked_##SUFFIX(unsigned int nth, unsigned int ith, void
|
||||
if (copy_cnt > 0) { \
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, ct_end); \
|
||||
if (src2) { \
|
||||
hvx_add_f32_uaa((uint8_t *) &dst_col[src0_start_row], \
|
||||
hvx_add_f32_uuu((uint8_t *) &dst_col[src0_start_row], \
|
||||
(const uint8_t *) tmp, \
|
||||
(const uint8_t *) ((const float *) mmctx->vtcm_bias + src0_start_row), \
|
||||
copy_cnt); \
|
||||
@@ -1108,7 +1108,7 @@ static void hvx_mv_2d(unsigned int nth, unsigned int ith, void * data) {
|
||||
if (copy_cnt > 0) {
|
||||
htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, src0_end_row);
|
||||
if (src2) {
|
||||
hvx_add_f32_uaa((uint8_t *) &dst_col[src0_start_row],
|
||||
hvx_add_f32_uuu((uint8_t *) &dst_col[src0_start_row],
|
||||
(const uint8_t *) tmp,
|
||||
(const uint8_t *) ((const float *) mmctx->vtcm_bias + src0_start_row),
|
||||
copy_cnt);
|
||||
|
||||
@@ -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);
|
||||
|
||||
|
||||
@@ -452,6 +452,46 @@ static void unary_mul_sycl(const T * x, const T * g, T * dst, const int64_t k, c
|
||||
});
|
||||
}
|
||||
|
||||
// ADD(bias) + UNARY + MUL(scale) with both broadcast over dim 0, the delta-net alpha gate:
|
||||
// dst[i] = op(a[i] + bias[i % ne0]) * scale[i % ne0]. k == ne0 makes that the flat index.
|
||||
template<typename F>
|
||||
static void add_unary_mul_flat_kernel(const float * a, const float * bias, const float * scale, float * dst,
|
||||
const int64_t k, const sycl::nd_item<1> &item_ct1, F op) {
|
||||
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
|
||||
dst[i] = op(a[i] + bias[i]) * scale[i];
|
||||
}
|
||||
}
|
||||
|
||||
template<typename F>
|
||||
static void add_unary_mul_bcast_kernel(const float * a, const float * bias, const float * scale, float * dst,
|
||||
const int64_t k, const sycl::uint3 ne0_fd, const sycl::nd_item<1> &item_ct1, F op) {
|
||||
SYCL_GLOBAL_ID_LOOP(k, item_ct1) {
|
||||
const uint32_t h = fastmodulo((uint32_t) i, ne0_fd);
|
||||
dst[i] = op(a[i] + bias[h]) * scale[h];
|
||||
}
|
||||
}
|
||||
|
||||
template<typename F>
|
||||
static void add_unary_mul_sycl(const float * a, const float * bias, const float * scale, float * dst,
|
||||
const int64_t k, const int64_t ne0, queue_ptr main_stream, F op) {
|
||||
const size_t num_blocks = ceil_div((size_t) k, (size_t) SYCL_GLU_BLOCK_SIZE);
|
||||
const sycl::nd_range<1> range(num_blocks * sycl::range<1>(SYCL_GLU_BLOCK_SIZE), sycl::range<1>(SYCL_GLU_BLOCK_SIZE));
|
||||
|
||||
if (k == ne0) {
|
||||
main_stream->parallel_for(range, [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
||||
add_unary_mul_flat_kernel(a, bias, scale, dst, k, item_ct1, op);
|
||||
});
|
||||
return;
|
||||
}
|
||||
|
||||
// 32-bit fastdiv, exact only below 2^31; ggml_sycl_can_fuse() already declined past that
|
||||
GGML_ASSERT(k < ((int64_t) 1 << 31));
|
||||
const sycl::uint3 ne0_fd = init_fastdiv_values((uint32_t) ne0);
|
||||
main_stream->parallel_for(range, [=](sycl::nd_item<1> item_ct1) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
||||
add_unary_mul_bcast_kernel(a, bias, scale, dst, k, ne0_fd, item_ct1, op);
|
||||
});
|
||||
}
|
||||
|
||||
namespace ggml_sycl_detail {
|
||||
static void acc_f32_sycl(const char *x, const char *y, float *dst,
|
||||
const int64_t n_elements,
|
||||
@@ -995,6 +1035,19 @@ static inline void ggml_sycl_op_swiglu(ggml_backend_sycl_context & ctx, ggml_ten
|
||||
});
|
||||
}
|
||||
|
||||
// Hands `launch` the functor for the unary op of a fused unary chain. Anything else
|
||||
// ggml_sycl_can_fuse() has already declined, so the default is a dispatcher bug.
|
||||
template<typename F>
|
||||
static void dispatch_fused_unary_op(ggml_unary_op uop, F && launch) {
|
||||
switch (uop) {
|
||||
case GGML_UNARY_OP_SILU: launch([](float v) { return op_silu(v); }); break;
|
||||
case GGML_UNARY_OP_SIGMOID: launch([](float v) { return op_sigmoid(v); }); break;
|
||||
case GGML_UNARY_OP_SOFTPLUS: launch([](float v) { return op_softplus(v); }); break;
|
||||
default:
|
||||
GGML_ABORT("fused unary chain: unsupported unary op %s", ggml_unary_op_name(uop));
|
||||
}
|
||||
}
|
||||
|
||||
// dst = op(unary_node->src[0]) * other, written straight to the MUL output, saving the
|
||||
// standalone unary launch. Preconditions come from ggml_sycl_can_fuse(); re-asserted here.
|
||||
void ggml_sycl_op_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor * unary_node, ggml_tensor * mul_node) {
|
||||
@@ -1032,13 +1085,41 @@ void ggml_sycl_op_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor *
|
||||
}
|
||||
};
|
||||
|
||||
switch (ggml_get_unary_op(unary_node)) {
|
||||
case GGML_UNARY_OP_SILU: dispatch_type([](float v) { return op_silu(v); }); break;
|
||||
case GGML_UNARY_OP_SIGMOID: dispatch_type([](float v) { return op_sigmoid(v); }); break;
|
||||
case GGML_UNARY_OP_SOFTPLUS: dispatch_type([](float v) { return op_softplus(v); }); break;
|
||||
default:
|
||||
GGML_ABORT("fused unary+mul: unsupported unary op %s", ggml_unary_op_name(ggml_get_unary_op(unary_node)));
|
||||
}
|
||||
dispatch_fused_unary_op(ggml_get_unary_op(unary_node), dispatch_type);
|
||||
}
|
||||
|
||||
// dst = op(a + bias) * scale for an ADD + UNARY + MUL chain whose bias and scale broadcast
|
||||
// over dim 0. Preconditions come from ggml_sycl_can_fuse(); re-asserted here.
|
||||
void ggml_sycl_op_add_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor * add_node,
|
||||
ggml_tensor * unary_node, ggml_tensor * mul_node) {
|
||||
// the dst-arity convention the other fusions follow; a and bias live on add_node
|
||||
scope_op_debug_print scope_dbg_print(__func__, mul_node, /*num_src=*/2);
|
||||
|
||||
const ggml_tensor * a = add_node->src[0];
|
||||
const ggml_tensor * bias = add_node->src[1];
|
||||
const ggml_tensor * scale = (mul_node->src[0] == unary_node) ? mul_node->src[1] : mul_node->src[0];
|
||||
|
||||
// scale is picked by elimination; ggml_can_fuse()'s single-use rule rules out MUL(unary, unary)
|
||||
GGML_ASSERT(scale != unary_node);
|
||||
GGML_ASSERT(a->type == GGML_TYPE_F32 && bias->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(scale->type == GGML_TYPE_F32 && mul_node->type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(ggml_are_same_shape(a, mul_node));
|
||||
// a and dst are indexed flat
|
||||
GGML_ASSERT(ggml_is_contiguous(a) && ggml_is_contiguous(mul_node));
|
||||
// bias and scale are one contiguous ne0-length row each, broadcast over the outer dims
|
||||
GGML_ASSERT(bias->ne[0] == a->ne[0] && scale->ne[0] == a->ne[0]);
|
||||
GGML_ASSERT(ggml_nrows(bias) == 1 && ggml_nrows(scale) == 1);
|
||||
GGML_ASSERT(ggml_is_contiguous(bias) && ggml_is_contiguous(scale));
|
||||
|
||||
queue_ptr main_stream = ctx.stream();
|
||||
SYCL_CHECK(ggml_sycl_set_device(ctx.device));
|
||||
|
||||
const auto dispatch_op = [&](auto op) {
|
||||
add_unary_mul_sycl((const float *) a->data, (const float *) bias->data, (const float *) scale->data,
|
||||
(float *) mul_node->data, ggml_nelements(mul_node), mul_node->ne[0], main_stream, op);
|
||||
};
|
||||
|
||||
dispatch_fused_unary_op(ggml_get_unary_op(unary_node), dispatch_op);
|
||||
}
|
||||
|
||||
__dpct_inline__ float ggml_sycl_op_swiglu_oai_single(float x, float g, float alpha = 1.702f, float limit = 7.0f) {
|
||||
|
||||
@@ -132,4 +132,9 @@ void ggml_sycl_arange(ggml_backend_sycl_context & ctx, ggml_tensor * dst);
|
||||
// fused UNARY(silu|sigmoid|softplus) + MUL; see ggml_sycl_can_fuse() for the accepted shapes
|
||||
void ggml_sycl_op_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor * unary_node, ggml_tensor * mul_node);
|
||||
|
||||
// fused f32 ADD + UNARY(silu|sigmoid|softplus) + MUL with the bias and the scale broadcast
|
||||
// over dim 0; see ggml_sycl_can_fuse() for the accepted shapes
|
||||
void ggml_sycl_op_add_unary_mul_fused(ggml_backend_sycl_context & ctx, ggml_tensor * add_node,
|
||||
ggml_tensor * unary_node, ggml_tensor * mul_node);
|
||||
|
||||
#endif // GGML_SYCL_ELEMENTWISE_HPP
|
||||
|
||||
@@ -64,6 +64,12 @@ static bool ggml_sycl_should_fuse_mul_mat_glu(const ggml_tensor * gate, const gg
|
||||
return true;
|
||||
}
|
||||
|
||||
// the unary ops the fused unary chains in element_wise.cpp have a functor for
|
||||
static bool ggml_sycl_fused_unary_has_kernel(ggml_unary_op unary_op) {
|
||||
return unary_op == GGML_UNARY_OP_SILU || unary_op == GGML_UNARY_OP_SIGMOID ||
|
||||
unary_op == GGML_UNARY_OP_SOFTPLUS;
|
||||
}
|
||||
|
||||
bool ggml_sycl_can_fuse(const ggml_cgraph * cgraph, int node_idx, std::initializer_list<enum ggml_op> ops,
|
||||
std::initializer_list<enum ggml_unary_op> unary_ops) {
|
||||
#ifndef NDEBUG
|
||||
@@ -184,9 +190,7 @@ bool ggml_sycl_can_fuse(const ggml_cgraph * cgraph, int node_idx, std::initializ
|
||||
return false;
|
||||
}
|
||||
|
||||
// the ops ggml_sycl_op_unary_mul_fused() has a kernel for
|
||||
if (unary_op != GGML_UNARY_OP_SILU && unary_op != GGML_UNARY_OP_SIGMOID &&
|
||||
unary_op != GGML_UNARY_OP_SOFTPLUS) {
|
||||
if (!ggml_sycl_fused_unary_has_kernel(unary_op)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
@@ -233,6 +237,55 @@ bool ggml_sycl_can_fuse(const ggml_cgraph * cgraph, int node_idx, std::initializ
|
||||
return true;
|
||||
}
|
||||
|
||||
// ADD(bias) + UNARY + MUL(scale): the delta-net alpha gate, softplus(alpha + dt) * a.
|
||||
// The broadcast is what stops the same-shape UNARY + MUL branch above firing past one token.
|
||||
if (ops.size() == 3 && ops.begin()[0] == GGML_OP_ADD && ops.begin()[1] == GGML_OP_UNARY &&
|
||||
ops.begin()[2] == GGML_OP_MUL && unary_ops.size() == 1) {
|
||||
const ggml_tensor * add = cgraph->nodes[node_idx];
|
||||
const ggml_tensor * unary = cgraph->nodes[node_idx + 1];
|
||||
const ggml_tensor * mul = cgraph->nodes[node_idx + 2];
|
||||
|
||||
const ggml_unary_op unary_op = ggml_get_unary_op(unary);
|
||||
if (unary_op != unary_ops.begin()[0]) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!ggml_sycl_fused_unary_has_kernel(unary_op)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// ggml_can_fuse() has already pinned the chain: unary consumes add, mul consumes
|
||||
// unary, add and unary have one use each, and all three have the same shape
|
||||
const ggml_tensor * a = add->src[0];
|
||||
const ggml_tensor * bias = add->src[1];
|
||||
const ggml_tensor * scale = (mul->src[0] == unary) ? mul->src[1] : mul->src[0];
|
||||
|
||||
if (a->type != GGML_TYPE_F32 || bias->type != GGML_TYPE_F32 ||
|
||||
scale->type != GGML_TYPE_F32 || mul->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// the activation and the destination are indexed flat
|
||||
if (!ggml_is_contiguous(a) || !ggml_is_contiguous(mul) || !ggml_are_same_shape(a, mul)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// the kernel reads the bias and the scale as v[col], so each must be a single
|
||||
// contiguous row spanning ne0
|
||||
if (bias->ne[0] != a->ne[0] || scale->ne[0] != a->ne[0] ||
|
||||
ggml_nrows(bias) != 1 || ggml_nrows(scale) != 1 ||
|
||||
!ggml_is_contiguous(bias) || !ggml_is_contiguous(scale)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// the 32-bit fastdiv is inexact past 2^31; decline, the unfused path handles it
|
||||
if (ggml_nelements(mul) >= ((int64_t) 1 << 31)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
if (ops.size() == 3 && ops.begin()[0] == GGML_OP_SSM_CONV && ops.begin()[1] == GGML_OP_ADD &&
|
||||
ops.begin()[2] == GGML_OP_UNARY && unary_ops.size() == 1 && unary_ops.begin()[0] == GGML_UNARY_OP_SILU) {
|
||||
const ggml_tensor * ssm_conv = cgraph->nodes[node_idx];
|
||||
|
||||
+42
-30
@@ -46,8 +46,8 @@ static constexpr float H20[20][20] = {
|
||||
#undef P
|
||||
#undef N
|
||||
|
||||
template <int N>
|
||||
static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst, const int64_t n_rows,
|
||||
template <int N, typename T>
|
||||
static void fwht_kernel(const T * __restrict__ src, float * __restrict__ dst, const int64_t n_rows,
|
||||
const float scale, const sycl::nd_item<2> & item) {
|
||||
const sycl::sub_group sg = item.get_sub_group();
|
||||
|
||||
@@ -67,7 +67,7 @@ static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst
|
||||
|
||||
#pragma unroll
|
||||
for (int i = 0; i < el_w; ++i) {
|
||||
reg[i] = src[i * WARP_SIZE + lane] * scale;
|
||||
reg[i] = static_cast<float>(src[i * WARP_SIZE + lane]) * scale;
|
||||
}
|
||||
|
||||
// Butterflies inside the sub-group. The partner of a lane with bit h clear is the
|
||||
@@ -107,8 +107,8 @@ static void fwht_kernel(const float * __restrict__ src, float * __restrict__ dst
|
||||
}
|
||||
}
|
||||
|
||||
template <int N>
|
||||
static void launch_fwht(const float * src, float * dst, const int64_t n_rows, const float scale,
|
||||
template <int N, typename T>
|
||||
static void launch_fwht(const T * src, float * dst, const int64_t n_rows, const float scale,
|
||||
dpct::queue_ptr stream) {
|
||||
constexpr int rows_per_block = 4;
|
||||
|
||||
@@ -120,7 +120,7 @@ static void launch_fwht(const float * src, float * dst, const int64_t n_rows, co
|
||||
|
||||
stream->parallel_for(sycl::nd_range<2>(global, local),
|
||||
[=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
||||
fwht_kernel<N>(src, dst, n_rows, scale, item);
|
||||
fwht_kernel<N, T>(src, dst, n_rows, scale, item);
|
||||
});
|
||||
}
|
||||
|
||||
@@ -128,8 +128,8 @@ static void launch_fwht(const float * src, float * dst, const int64_t n_rows, co
|
||||
// keeps N/NT values rather than N/WARP_SIZE. Butterflies below the sub-group width
|
||||
// still shuffle; those up to NT go through work-group local memory; the rest stay
|
||||
// in registers.
|
||||
template <int N, int NT>
|
||||
static void fwht_kernel_wide(const float * __restrict__ src,
|
||||
template <int N, int NT, typename T>
|
||||
static void fwht_kernel_wide(const T * __restrict__ src,
|
||||
float * __restrict__ dst,
|
||||
const int64_t n_rows,
|
||||
const float scale,
|
||||
@@ -151,7 +151,7 @@ static void fwht_kernel_wide(const float * __restrict__ src,
|
||||
float reg[el_w];
|
||||
#pragma unroll
|
||||
for (int i = 0; i < el_w; ++i) {
|
||||
reg[i] = src[i * NT + tid] * scale;
|
||||
reg[i] = static_cast<float>(src[i * NT + tid]) * scale;
|
||||
}
|
||||
|
||||
const sycl::sub_group sg = item.get_sub_group();
|
||||
@@ -207,8 +207,8 @@ static void fwht_kernel_wide(const float * __restrict__ src,
|
||||
}
|
||||
}
|
||||
|
||||
template <int N, int NT>
|
||||
static void launch_fwht_wide(const float * src,
|
||||
template <int N, int NT, typename T>
|
||||
static void launch_fwht_wide(const T * src,
|
||||
float * dst,
|
||||
const int64_t n_rows,
|
||||
const float scale,
|
||||
@@ -220,13 +220,13 @@ static void launch_fwht_wide(const float * src,
|
||||
sycl::local_accessor<float, 1> smem(sycl::range<1>(N), cgh);
|
||||
cgh.parallel_for(sycl::nd_range<2>(global, local),
|
||||
[=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
||||
fwht_kernel_wide<N, NT>(src, dst, n_rows, scale, item, get_pointer(smem));
|
||||
fwht_kernel_wide<N, NT, T>(src, dst, n_rows, scale, item, get_pointer(smem));
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
template <int N, int m>
|
||||
static void kronecker_kernel(const float * __restrict__ src,
|
||||
template <int N, int m, typename T>
|
||||
static void kronecker_kernel(const T * __restrict__ src,
|
||||
float * __restrict__ dst,
|
||||
const int64_t n_rows,
|
||||
const float scale,
|
||||
@@ -255,7 +255,7 @@ static void kronecker_kernel(const float * __restrict__ src,
|
||||
|
||||
#pragma unroll
|
||||
for (int j = 0; j < m; ++j) {
|
||||
reg[i * m + j] = src[b_idx * m + j] * scale;
|
||||
reg[i * m + j] = static_cast<float>(src[b_idx * m + j]) * scale;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -321,8 +321,8 @@ static void kronecker_kernel(const float * __restrict__ src,
|
||||
}
|
||||
}
|
||||
|
||||
template <int N, int m>
|
||||
static void launch_kronecker(const float * src,
|
||||
template <int N, int m, typename T>
|
||||
static void launch_kronecker(const T * src,
|
||||
float * dst,
|
||||
const int64_t n_rows,
|
||||
const float scale,
|
||||
@@ -337,25 +337,16 @@ static void launch_kronecker(const float * src,
|
||||
|
||||
stream->parallel_for(sycl::nd_range<2>(global, local),
|
||||
[=](sycl::nd_item<2> item) [[sycl::reqd_sub_group_size(WARP_SIZE)]] {
|
||||
kronecker_kernel<N, m>(src, dst, n_rows, scale, item);
|
||||
kronecker_kernel<N, m, T>(src, dst, n_rows, scale, item);
|
||||
});
|
||||
}
|
||||
|
||||
bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
|
||||
if (src->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
if (!ggml_are_same_shape(src, dst)) {
|
||||
return false;
|
||||
}
|
||||
if (!ggml_is_contiguous(src) || !ggml_is_contiguous(dst)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static bool ggml_sycl_op_fwht_impl(ggml_backend_sycl_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
|
||||
const int n = (int) src->ne[0];
|
||||
const int64_t rows = ggml_nrows(src);
|
||||
|
||||
const float * src_d = (const float *) src->data;
|
||||
const T * src_d = (const T *) src->data;
|
||||
float * dst_d = (float *) dst->data;
|
||||
dpct::queue_ptr stream = ctx.stream();
|
||||
|
||||
@@ -402,3 +393,24 @@ bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src,
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
bool ggml_sycl_op_fwht(ggml_backend_sycl_context & ctx, const ggml_tensor * src, ggml_tensor * dst) {
|
||||
if (dst->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
}
|
||||
if (!ggml_are_same_shape(src, dst)) {
|
||||
return false;
|
||||
}
|
||||
if (!ggml_is_contiguous(src) || !ggml_is_contiguous(dst)) {
|
||||
return false;
|
||||
}
|
||||
|
||||
switch (src->type) {
|
||||
case GGML_TYPE_F32:
|
||||
return ggml_sycl_op_fwht_impl<float>(ctx, src, dst);
|
||||
case GGML_TYPE_F16:
|
||||
return ggml_sycl_op_fwht_impl<sycl::half>(ctx, src, dst);
|
||||
default:
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -139,6 +139,7 @@ int g_ggml_sycl_dev2dev_memcpy = DEV2DEV_MEMCPY_SYCL;
|
||||
int g_ggml_sycl_usm_system = 0;
|
||||
int g_ggml_sycl_enable_host_pinned_mem = 1;
|
||||
int g_ggml_sycl_host_pinned_mem_2g = 0;
|
||||
int g_ggml_sycl_upload_staging_slots = 4;
|
||||
int g_ggml_sycl_get_mem_api = MEMORY_API_TYPE_LEVEL_ZERO;
|
||||
int g_ggml_sycl_enable_sparse_fa = 0;
|
||||
int g_ggml_sycl_debug_sparse_fa = 0;
|
||||
@@ -458,6 +459,7 @@ static void ggml_check_sycl() try {
|
||||
|
||||
g_ggml_sycl_host_pinned_mem_2g =
|
||||
ggml_sycl_get_env("GGML_SYCL_HOST_PINNED_MEM_2G", 0) & g_ggml_sycl_enable_host_pinned_mem;
|
||||
g_ggml_sycl_upload_staging_slots = std::max(0, ggml_sycl_get_env("GGML_SYCL_UPLOAD_STAGING_SLOTS", 4));
|
||||
|
||||
g_ggml_sycl_enable_sparse_fa = ggml_sycl_get_env("GGML_SYCL_SPARSE_FA", 0);
|
||||
g_ggml_sycl_debug_sparse_fa = ggml_sycl_get_env("GGML_SYCL_SPARSE_FA_DEBUG", 0);
|
||||
@@ -555,6 +557,7 @@ static void ggml_check_sycl() try {
|
||||
#endif
|
||||
|
||||
GGML_LOG_INFO(" GGML_SYCL_ENABLE_FUSION: %d\n", g_ggml_sycl_enable_fusion);
|
||||
GGML_LOG_INFO(" GGML_SYCL_UPLOAD_STAGING_SLOTS: %d\n", g_ggml_sycl_upload_staging_slots);
|
||||
|
||||
#if defined(__INTEL_LLVM_COMPILER)
|
||||
GGML_LOG_INFO(" GGML_SYCL_ENABLE_ESIMD: %d\n", g_ggml_sycl_enable_esimd);
|
||||
@@ -667,12 +670,23 @@ inline void free_aligned_mem_host(void * memblock) {
|
||||
// sycl buffer
|
||||
|
||||
struct ggml_backend_sycl_buffer_context {
|
||||
// pinned staging for uploads; the host fills one slot while the previous one transfers
|
||||
static constexpr size_t staging_slot_size = 8*1024*1024;
|
||||
|
||||
struct host_staging {
|
||||
void * data = nullptr;
|
||||
std::vector<sycl::event> events;
|
||||
std::vector<bool> submitted;
|
||||
int next = 0;
|
||||
};
|
||||
|
||||
int device;
|
||||
void * dev_ptr = nullptr;
|
||||
queue_ptr stream;
|
||||
std::string name;
|
||||
optimize_feature opt_feature;
|
||||
std::vector<ggml_tensor_extra_gpu *> tensor_extras;
|
||||
host_staging staging;
|
||||
bool is_usm_system;
|
||||
|
||||
ggml_backend_sycl_buffer_context(int device, void * dev_ptr, queue_ptr stream, bool is_usm_system) :
|
||||
@@ -682,7 +696,22 @@ struct ggml_backend_sycl_buffer_context {
|
||||
opt_feature = ggml_sycl_info().devices[device].opt_feature;
|
||||
}
|
||||
|
||||
// waits for every queued upload, then releases the pinned block
|
||||
void drop_host_staging() {
|
||||
for (size_t i = 0; i < staging.submitted.size(); ++i) {
|
||||
if (staging.submitted[i]) {
|
||||
staging.events[i].wait_and_throw();
|
||||
staging.submitted[i] = false;
|
||||
}
|
||||
}
|
||||
if (staging.data != nullptr) {
|
||||
sycl::free(staging.data, *stream);
|
||||
staging.data = nullptr;
|
||||
}
|
||||
}
|
||||
|
||||
~ggml_backend_sycl_buffer_context() {
|
||||
drop_host_staging();
|
||||
if (dev_ptr != nullptr) {
|
||||
ggml_sycl_set_device(device);
|
||||
if (is_usm_system)
|
||||
@@ -783,6 +812,40 @@ static void ggml_backend_sycl_buffer_set_tensor(ggml_backend_buffer_t buffer,
|
||||
GGML_SYCL_DEBUG(" size=%zu offset=%zu\n", size, offset);
|
||||
ggml_backend_sycl_buffer_context * ctx = ( ggml_backend_sycl_buffer_context *)buffer->context;
|
||||
ggml_sycl_set_device(ctx->device);
|
||||
|
||||
// copy through pinned memory so the device never reads mmap()ed pages directly
|
||||
// chunks pipeline on the in-order compute queue, so no drain per tensor is needed
|
||||
const int n_slots = g_ggml_sycl_upload_staging_slots;
|
||||
if (n_slots > 0 && ctx->staging.data == nullptr) {
|
||||
ctx->staging.data = sycl::malloc_host(n_slots * ctx->staging_slot_size, *ctx->stream);
|
||||
if (ctx->staging.data != nullptr) {
|
||||
ctx->staging.events.resize(n_slots);
|
||||
ctx->staging.submitted.assign(n_slots, false);
|
||||
}
|
||||
}
|
||||
if (ctx->staging.data != nullptr) {
|
||||
queue_ptr stream = ctx->stream;
|
||||
char * dst = (char *) tensor->data + offset;
|
||||
const char * src = (const char *) data;
|
||||
size_t remaining = size;
|
||||
while (remaining > 0) {
|
||||
const size_t chunk = std::min(remaining, ctx->staging_slot_size);
|
||||
const int slot = ctx->staging.next;
|
||||
ctx->staging.next = (ctx->staging.next + 1) % (int) ctx->staging.submitted.size();
|
||||
if (ctx->staging.submitted[slot]) {
|
||||
ctx->staging.events[slot].wait_and_throw();
|
||||
}
|
||||
void * stage = (char *) ctx->staging.data + slot * ctx->staging_slot_size;
|
||||
memcpy(stage, src, chunk);
|
||||
ctx->staging.events[slot] = stream->memcpy(dst, stage, chunk);
|
||||
ctx->staging.submitted[slot] = true;
|
||||
src += chunk;
|
||||
dst += chunk;
|
||||
remaining -= chunk;
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
auto stream = &(dpct::dev_mgr::instance().get_device(ctx->device).default_queue());
|
||||
SYCL_CHECK(CHECK_TRY_ERROR(dpct::dev_mgr::instance().get_device(ctx->device).queues_wait_and_throw()));
|
||||
#ifndef _WIN32
|
||||
@@ -5039,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);
|
||||
@@ -6249,6 +6319,16 @@ static void ggml_backend_sycl_graph_compute_impl(ggml_backend_sycl_context * syc
|
||||
i++;
|
||||
continue;
|
||||
}
|
||||
// ADD(bias) + UNARY + MUL(scale) with both broadcast over dim 0, the form the branch
|
||||
// above cannot take; ggml_get_unary_op() asserts, so check the op first.
|
||||
if (node->op == GGML_OP_ADD && i + 2 < cgraph->n_nodes &&
|
||||
cgraph->nodes[i + 1]->op == GGML_OP_UNARY &&
|
||||
ggml_sycl_can_fuse(cgraph, i, { GGML_OP_ADD, GGML_OP_UNARY, GGML_OP_MUL },
|
||||
{ ggml_get_unary_op(cgraph->nodes[i + 1]) })) {
|
||||
ggml_sycl_op_add_unary_mul_fused(*sycl_ctx, node, cgraph->nodes[i + 1], cgraph->nodes[i + 2]);
|
||||
i += 2;
|
||||
continue;
|
||||
}
|
||||
|
||||
// Batch consecutive independent same-shape F32 L2_NORM siblings (the GDN q/k
|
||||
// norms) into one launch; sources are strided views of the fused qkv buffer, so
|
||||
|
||||
+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,
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user