Merge commit '4ceb171910' into concedo_experimental

# Conflicts:
#	.github/workflows/build-sycl.yml
#	.github/workflows/docker.yml
#	.github/workflows/release.yml
#	cmake/llama-config.cmake.in
#	docs/backend/snapdragon/README.md
#	docs/backend/snapdragon/developer.md
#	examples/simple-cmake-pkg/CMakeLists.txt
#	ggml/include/ggml-sycl.h
#	ggml/src/ggml-cpu/repack.cpp
#	ggml/src/ggml-cpu/repack.h
#	ggml/src/ggml-hexagon/ggml-hexagon.cpp
#	ggml/src/ggml-hexagon/htp-opnode.h
#	ggml/src/ggml-hexagon/htp/act-ops.c
#	ggml/src/ggml-hexagon/htp/allreduce-ops.c
#	ggml/src/ggml-hexagon/htp/argsort-ops.c
#	ggml/src/ggml-hexagon/htp/binary-ops.c
#	ggml/src/ggml-hexagon/htp/concat-ops.c
#	ggml/src/ggml-hexagon/htp/cpy-ops.c
#	ggml/src/ggml-hexagon/htp/cumsum-ops.c
#	ggml/src/ggml-hexagon/htp/diag-ops.c
#	ggml/src/ggml-hexagon/htp/dma-queue.c
#	ggml/src/ggml-hexagon/htp/dma-queue.h
#	ggml/src/ggml-hexagon/htp/fill-ops.c
#	ggml/src/ggml-hexagon/htp/flash-attn-ops.c
#	ggml/src/ggml-hexagon/htp/flash-attn-ops.h
#	ggml/src/ggml-hexagon/htp/gated-delta-net-ops.c
#	ggml/src/ggml-hexagon/htp/get-rows-ops.c
#	ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h
#	ggml/src/ggml-hexagon/htp/htp-ctx.h
#	ggml/src/ggml-hexagon/htp/htp-ops.h
#	ggml/src/ggml-hexagon/htp/htp-tensor.c
#	ggml/src/ggml-hexagon/htp/htp-tensor.h
#	ggml/src/ggml-hexagon/htp/htp_iface.idl
#	ggml/src/ggml-hexagon/htp/hvx-exp.h
#	ggml/src/ggml-hexagon/htp/hvx-mm-kernels-tiled.h
#	ggml/src/ggml-hexagon/htp/im2col-ops.c
#	ggml/src/ggml-hexagon/htp/main.c
#	ggml/src/ggml-hexagon/htp/matmul-ops.c
#	ggml/src/ggml-hexagon/htp/matmul-ops.h
#	ggml/src/ggml-hexagon/htp/pad-ops.c
#	ggml/src/ggml-hexagon/htp/repeat-ops.c
#	ggml/src/ggml-hexagon/htp/roll-ops.c
#	ggml/src/ggml-hexagon/htp/rope-ops.c
#	ggml/src/ggml-hexagon/htp/rope-ops.h
#	ggml/src/ggml-hexagon/htp/set-rows-ops.c
#	ggml/src/ggml-hexagon/htp/softmax-ops.c
#	ggml/src/ggml-hexagon/htp/solve-tri-ops.c
#	ggml/src/ggml-hexagon/htp/ssm-conv.c
#	ggml/src/ggml-hexagon/htp/sum-rows-ops.c
#	ggml/src/ggml-hexagon/htp/unary-ops.c
#	ggml/src/ggml-opencl/ggml-opencl.cpp
#	ggml/src/ggml-sycl/dsv4-hc.cpp
#	ggml/src/ggml-sycl/fattn-mkl.cpp
#	ggml/src/ggml-sycl/fattn-tile.hpp
#	ggml/src/ggml-sycl/ggml-sycl.cpp
#	ggml/src/ggml-webgpu/ggml-webgpu.cpp
#	scripts/snapdragon/ggml-hexagon-profile.py
#	scripts/snapdragon/ggml-hexagon-trace.py
#	scripts/snapdragon/run.py
#	scripts/sync_vendor.py
#	tests/test-backend-ops.cpp
#	tests/test-chat.cpp
#	tests/test-llama-archs.cpp
#	tests/test-recurrent-state-rollback.cpp
#	tests/test-save-load-state.cpp
#	tools/cli/README.md
#	tools/completion/README.md
#	tools/server/README.md
This commit is contained in:
Concedo 2026-09-24 16:37:57 +08:00
commit 084d797f2d
52 changed files with 2420 additions and 631 deletions

View file

@ -2017,7 +2017,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
params.sampling.temp = std::max(params.sampling.temp, 0.0f);
params.sampling.user_sampling_config |= common_params_sampling_config::COMMON_PARAMS_SAMPLING_CONFIG_TEMP;
}
).set_sampling());
).set_sampling().set_env("LLAMA_ARG_TEMPERATURE"));
add_opt(common_arg(
{"--top-k"}, "N",
string_format("top-k sampling (default: %d, 0 = disabled)", params.sampling.top_k),
@ -2033,7 +2033,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
params.sampling.top_p = std::stof(value);
params.sampling.user_sampling_config |= common_params_sampling_config::COMMON_PARAMS_SAMPLING_CONFIG_TOP_P;
}
).set_sampling());
).set_sampling().set_env("LLAMA_ARG_TOP_P"));
add_opt(common_arg(
{"--min-p"}, "N",
string_format("min-p sampling (default: %.2f, 0.0 = disabled)", (double)params.sampling.min_p),
@ -2041,7 +2041,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
params.sampling.min_p = std::stof(value);
params.sampling.user_sampling_config |= common_params_sampling_config::COMMON_PARAMS_SAMPLING_CONFIG_MIN_P;
}
).set_sampling());
).set_sampling().set_env("LLAMA_ARG_MIN_P"));
add_opt(common_arg(
{"--top-nsigma", "--top-n-sigma"}, "N",
string_format("top-n-sigma sampling (default: %.2f, -1.0 = disabled)", params.sampling.top_n_sigma),
@ -2097,7 +2097,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
params.sampling.penalty_repeat = penalty_repeat;
params.sampling.user_sampling_config |= common_params_sampling_config::COMMON_PARAMS_SAMPLING_CONFIG_PENALTY_REPEAT;
}
).set_sampling());
).set_sampling().set_env("LLAMA_ARG_REPEAT_PENALTY"));
add_opt(common_arg(
{"--presence-penalty"}, "N",
string_format("repeat alpha presence penalty (default: %.2f, 0.0 = disabled)", (double)params.sampling.penalty_present),
@ -2108,7 +2108,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
}
params.sampling.penalty_present = penalty_present;
}
).set_sampling());
).set_sampling().set_env("LLAMA_ARG_PRESENCE_PENALTY"));
add_opt(common_arg(
{"--frequency-penalty"}, "N",
string_format("repeat alpha frequency penalty (default: %.2f, 0.0 = disabled)", (double)params.sampling.penalty_freq),
@ -2119,7 +2119,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
}
params.sampling.penalty_freq = penalty_freq;
}
).set_sampling());
).set_sampling().set_env("LLAMA_ARG_FREQUENCY_PENALTY"));
add_opt(common_arg(
{"--dry-multiplier"}, "N",
string_format("set DRY sampling multiplier (default: %.2f, 0.0 = disabled)", (double)params.sampling.dry_multiplier),
@ -3309,9 +3309,18 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
).set_examples({LLAMA_EXAMPLE_EMBEDDING}));
add_opt(common_arg(
{"--host"}, "HOST",
string_format("ip address to listen, or bind to an UNIX socket if the address ends with .sock (default: %s)", params.hostname.c_str()),
string_format("IP addresses to listen on, comma-separated, or UNIX socket paths ending in .sock; with multiple TCP addresses, :: binds IPv6 only; overlapping addresses result in undefined behavior (default: %s)", params.hostnames[0].c_str()),
[](common_params & params, const std::string & value) {
params.hostname = value;
params.hostnames.clear();
for (auto & host : parse_csv_row(value)) {
host = string_strip(host);
if (!host.empty()) {
params.hostnames.push_back(host);
}
}
if (params.hostnames.empty()) {
throw std::invalid_argument("--host requires at least one address");
}
}
).set_examples({LLAMA_EXAMPLE_SERVER}).set_env("LLAMA_ARG_HOST"));
add_opt(common_arg(

View file

@ -1244,7 +1244,9 @@ std::optional<common_chat_params> common_chat_try_specialized_template(
// Qwen3-Coder XML tool calls, also used by Nemotron Nano 3, Qwen3.5 and StepFun-3.5-Flash
if (src.find("<tool_call>") != std::string::npos &&
src.find("<function=") != std::string::npos &&
src.find("<parameter=") != std::string::npos) {
src.find("<parameter=") != std::string::npos &&
// 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);
}

View file

@ -632,10 +632,10 @@ struct common_params {
int32_t checkpoint_min_step = 8192; // minimum spacing between context checkpoints
int32_t cache_ram_mib = 8192; // -1 = no limit, 0 - disable, 1 = 1 MiB, etc.
std::string hostname = "127.0.0.1";
std::string public_path = ""; // NOLINT
std::string api_prefix = ""; // NOLINT
std::string chat_template = ""; // NOLINT
std::vector<std::string> hostnames = {"127.0.0.1"};
bool use_jinja = true; // NOLINT
// server CORS params

View file

@ -54,7 +54,7 @@ static void ensure_key_type_allowed(const value & val) {
}
// execute with error handling
value statement::execute(context & ctx) {
value statement::execute(context & ctx) const {
try {
return execute_impl(ctx);
} catch (const continue_statement::signal & /* ex */) {
@ -83,7 +83,7 @@ value statement::execute(context & ctx) {
}
}
value identifier::execute_impl(context & ctx) {
value identifier::execute_impl(context & ctx) const {
auto it = ctx.get_val(val);
auto builtins = global_builtins();
if (!it->is_undefined()) {
@ -101,7 +101,7 @@ value identifier::execute_impl(context & ctx) {
}
}
value object_literal::execute_impl(context & ctx) {
value object_literal::execute_impl(context & ctx) const {
auto obj = mk_val<value_object>();
for (const auto & pair : val) {
value key = pair.first->execute(ctx);
@ -112,7 +112,7 @@ value object_literal::execute_impl(context & ctx) {
return obj;
}
value binary_expression::execute_impl(context & ctx) {
value binary_expression::execute_impl(context & ctx) const {
value left_val = left->execute(ctx);
// Logical operators
@ -320,9 +320,7 @@ static value try_builtin_func(context & ctx, const std::string & name, value & i
throw std::runtime_error("Unknown (built-in) filter '" + name + "' for type " + input->type());
}
value filter_expression::execute_impl(context & ctx) {
value input = operand ? operand->execute(ctx) : val;
static value apply_filter(context & ctx, const statement_ptr & filter, value input) {
JJ_DEBUG("Applying filter to %s", input->type().c_str());
auto set_filter_alias = [](auto & filter_id) {
@ -378,22 +376,21 @@ value filter_expression::execute_impl(context & ctx) {
}
}
value filter_statement::execute_impl(context & ctx) {
value filter_expression::execute_impl(context & ctx) const {
return apply_filter(ctx, filter, operand->execute(ctx));
}
value filter_statement::execute_impl(context & ctx) const {
// eval body as string, then apply filter
auto body_val = exec_statements(body, ctx);
value_string parts = mk_val<value_string>();
gather_string_parts_recursive(body_val, parts);
JJ_DEBUG("FilterStatement: applying filter to body string of length %zu", parts->val_str.length());
filter_expression filter_expr(std::move(parts), std::move(filter));
value out = filter_expr.execute(ctx);
// this node can be reused later, make sure filter is preserved
this->filter = std::move(filter_expr.filter);
return out;
return apply_filter(ctx, filter, parts);
}
value test_expression::execute_impl(context & ctx) {
value test_expression::execute_impl(context & ctx) const {
// NOTE: "value is something" translates to function call "test_is_something(value)"
const auto & builtins = global_builtins();
@ -442,7 +439,7 @@ value test_expression::execute_impl(context & ctx) {
}
}
value unary_expression::execute_impl(context & ctx) {
value unary_expression::execute_impl(context & ctx) const {
value operand_val = argument->execute(ctx);
JJ_DEBUG("Executing unary expression with operator '%s'", op.value.c_str());
@ -461,7 +458,7 @@ value unary_expression::execute_impl(context & ctx) {
throw std::runtime_error("Unknown unary operator '" + op.value + "'");
}
value if_statement::execute_impl(context & ctx) {
value if_statement::execute_impl(context & ctx) const {
value test_val = test->execute(ctx);
auto out = mk_val<value_array>();
@ -482,17 +479,17 @@ value if_statement::execute_impl(context & ctx) {
return str;
}
value for_statement::execute_impl(context & ctx) {
value for_statement::execute_impl(context & ctx) const {
context scope(ctx); // new scope for loop variables
jinja::select_expression * select_expr = cast_stmt<select_expression>(iterable);
const jinja::select_expression * select_expr = cast_stmt<select_expression>(iterable);
statement_ptr test_expr_nullptr;
statement_ptr & iter_expr = [&]() -> statement_ptr & {
const statement_ptr & iter_expr = [&]() -> const statement_ptr & {
auto tmp = cast_stmt<select_expression>(iterable);
return tmp ? tmp->lhs : iterable;
}();
statement_ptr & test_expr = [&]() -> statement_ptr & {
const statement_ptr & test_expr = [&]() -> const statement_ptr & {
auto tmp = cast_stmt<select_expression>(iterable);
return tmp ? tmp->test : test_expr_nullptr;
}();
@ -648,7 +645,7 @@ value for_statement::execute_impl(context & ctx) {
return str;
}
value set_statement::execute_impl(context & ctx) {
value set_statement::execute_impl(context & ctx) const {
auto rhs = val ? val->execute(ctx) : exec_statements(body, ctx);
if (is_stmt<identifier>(assignee)) {
@ -747,7 +744,7 @@ static inline void bind_parameters(const std::string & name, const statements &
}
}
value macro_statement::execute_impl(context & ctx) {
value macro_statement::execute_impl(context & ctx) const {
if (!is_stmt<identifier>(this->name)) {
throw std::runtime_error("Macro name must be an identifier");
}
@ -770,7 +767,7 @@ value macro_statement::execute_impl(context & ctx) {
return mk_val<value_undefined>();
}
value call_statement::execute_impl(context & ctx) {
value call_statement::execute_impl(context & ctx) const {
auto call_expr = cast_stmt<call_expression>(this->call);
if (!call_expr) {
throw std::runtime_error("Call statement requires a valid call expression");
@ -810,7 +807,7 @@ value call_statement::execute_impl(context & ctx) {
return callee_func->invoke(args);
}
value member_expression::execute_impl(context & ctx) {
value member_expression::execute_impl(context & ctx) const {
value object = this->object->execute(ctx);
value property;
@ -943,7 +940,7 @@ value member_expression::execute_impl(context & ctx) {
return val;
}
value call_expression::execute_impl(context & ctx) {
value call_expression::execute_impl(context & ctx) const {
// gather arguments
func_args args(ctx);
for (auto & arg_stmt : this->args) {
@ -961,7 +958,7 @@ value call_expression::execute_impl(context & ctx) {
return callee_func->invoke(args);
}
value keyword_argument_expression::execute_impl(context & ctx) {
value keyword_argument_expression::execute_impl(context & ctx) const {
if (!is_stmt<identifier>(key)) {
throw std::runtime_error("Keyword argument key must be identifiers");
}
@ -985,7 +982,7 @@ std::string runtime::debug_dump_program(const program & prog, const std::string
return std::string(lvl * 2, ' ');
};
ctx.visitor = [&](bool is_leaf, statement * node, std::vector<visitor_pair> children) {
ctx.visitor = [&](bool is_leaf, const statement * node, std::vector<visitor_pair> children) {
oss << indent(lvl) << node->type() << ":\n";
lvl++;
if (is_leaf) {

View file

@ -48,9 +48,9 @@ const T * cast_stmt(const statement_ptr & ptr) {
void enable_debug(bool enable);
// for visiting AST nodes
// function signature: void(bool is_leaf, statement * node, pair of <label, children>)
using visitor_pair = std::pair<std::string, std::vector<statement *>>;
using visitor_fn = std::function<void(bool, statement *, std::vector<visitor_pair>)>;
// function signature: void(bool is_leaf, const statement * node, pair of <label, children>)
using visitor_pair = std::pair<std::string, std::vector<const statement *>>;
using visitor_fn = std::function<void(bool, const statement *, std::vector<visitor_pair>)>;
struct context {
std::shared_ptr<std::string> src; // for debugging; use shared_ptr to avoid copying on scope creation
@ -107,8 +107,8 @@ private:
};
// utils for visiting AST nodes
static std::vector<statement *> stmts_to_ptr(const statements & stmts) {
std::vector<statement *> children;
static std::vector<const statement *> stmts_to_ptr(const statements & stmts) {
std::vector<const statement *> children;
for (const auto & stmt : stmts) {
children.push_back(stmt.get());
}
@ -117,17 +117,18 @@ static std::vector<statement *> stmts_to_ptr(const statements & stmts) {
/**
* Base class for all nodes in the AST.
* The AST is shared between threads, so visit and execute must be const.
*/
struct statement {
size_t pos; // position in source, for debugging
virtual ~statement() = default;
virtual std::string type() const { return "Statement"; }
virtual void visit(context & ctx) { ctx.visitor(true, this, {}); }
virtual void visit(context & ctx) const { ctx.visitor(true, this, {}); }
// execute_impl must be overridden by derived classes
virtual value execute_impl(context &) { throw_exec_error(); }
virtual value execute_impl(context &) const { throw_exec_error(); }
// execute is the public method to execute a statement with error handling
value execute(context &);
value execute(context &) const;
private:
[[noreturn]] void throw_exec_error() const {
@ -166,7 +167,7 @@ struct program : public statement {
program() = default;
explicit program(statements && body) : body(std::move(body)) {}
std::string type() const override { return "Program"; }
[[noreturn]] value execute_impl(context &) override {
[[noreturn]] value execute_impl(context &) const override {
throw std::runtime_error("Cannot execute program directly, use jinja::runtime instead");
}
};
@ -182,8 +183,8 @@ struct if_statement : public statement {
}
std::string type() const override { return "If"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"test", {test.get()}},
{"body", stmts_to_ptr(body)},
@ -213,8 +214,8 @@ struct for_statement : public statement {
}
std::string type() const override { return "For"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"loopvar", {loopvar.get()}},
{"iterable", {iterable.get()}},
@ -233,7 +234,7 @@ struct break_statement : public statement {
}
};
[[noreturn]] value execute_impl(context &) override {
[[noreturn]] value execute_impl(context &) const override {
throw break_statement::signal();
}
};
@ -247,7 +248,7 @@ struct continue_statement : public statement {
}
};
[[noreturn]] value execute_impl(context &) override {
[[noreturn]] value execute_impl(context &) const override {
throw continue_statement::signal();
}
};
@ -255,7 +256,7 @@ struct continue_statement : public statement {
// do nothing
struct noop_statement : public statement {
std::string type() const override { return "Noop"; }
value execute_impl(context &) override {
value execute_impl(context &) const override {
return mk_val<value_undefined>();
}
};
@ -272,8 +273,8 @@ struct set_statement : public statement {
}
std::string type() const override { return "Set"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"assignee", {assignee.get()}},
{"value", {val.get()}},
@ -294,8 +295,8 @@ struct macro_statement : public statement {
}
std::string type() const override { return "Macro"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"name", {name.get()}},
{"args", stmts_to_ptr(args)},
@ -308,7 +309,7 @@ struct comment_statement : public statement {
std::string val;
explicit comment_statement(const std::string & v) : val(v) {}
std::string type() const override { return "Comment"; }
value execute_impl(context &) override {
value execute_impl(context &) const override {
return mk_val<value_undefined>();
}
};
@ -318,7 +319,7 @@ struct comment_statement : public statement {
// Represents an omitted expression in a computed member, e.g. `a[]`.
struct blank_expression : public expression {
std::string type() const override { return "BlankExpression"; }
value execute_impl(context &) override {
value execute_impl(context &) const override {
return mk_val<value_undefined>();
}
};
@ -334,8 +335,8 @@ struct member_expression : public expression {
chk_type<expression>(this->property);
}
std::string type() const override { return "MemberExpression"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"object", {object.get()}},
{"property", {property.get()}}
@ -353,8 +354,8 @@ struct call_expression : public expression {
for (const auto& arg : this->args) chk_type<expression>(arg);
}
std::string type() const override { return "CallExpression"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"callee", {callee.get()}},
{"args", stmts_to_ptr(args)}
@ -369,7 +370,7 @@ struct identifier : public expression {
std::string val;
explicit identifier(const std::string & val) : val(val) {}
std::string type() const override { return "Identifier"; }
value execute_impl(context & ctx) override;
value execute_impl(context & ctx) const override;
};
// Literals
@ -378,7 +379,7 @@ struct integer_literal : public expression {
int64_t val;
explicit integer_literal(int64_t val) : val(val) {}
std::string type() const override { return "IntegerLiteral"; }
value execute_impl(context &) override {
value execute_impl(context &) const override {
return mk_val<value_int>(val);
}
};
@ -387,7 +388,7 @@ struct float_literal : public expression {
double val;
explicit float_literal(double val) : val(val) {}
std::string type() const override { return "FloatLiteral"; }
value execute_impl(context &) override {
value execute_impl(context &) const override {
return mk_val<value_float>(val);
}
};
@ -396,7 +397,7 @@ struct string_literal : public expression {
std::string val;
explicit string_literal(const std::string & val) : val(val) {}
std::string type() const override { return "StringLiteral"; }
value execute_impl(context &) override {
value execute_impl(context &) const override {
return mk_val<value_string>(val);
}
};
@ -407,7 +408,7 @@ struct array_literal : public expression {
for (const auto& item : this->val) chk_type<expression>(item);
}
std::string type() const override { return "ArrayLiteral"; }
value execute_impl(context & ctx) override {
value execute_impl(context & ctx) const override {
auto arr = mk_val<value_array>();
for (const auto & item_stmt : val) {
arr->push_back(item_stmt->execute(ctx));
@ -422,7 +423,7 @@ struct tuple_literal : public expression {
for (const auto& item : this->val) chk_type<expression>(item);
}
std::string type() const override { return "TupleLiteral"; }
value execute_impl(context & ctx) override {
value execute_impl(context & ctx) const override {
auto arr = mk_val<value_array>();
for (const auto & item_stmt : val) {
arr->push_back(item_stmt->execute(ctx));
@ -441,7 +442,7 @@ struct object_literal : public expression {
}
}
std::string type() const override { return "ObjectLiteral"; }
value execute_impl(context & ctx) override;
value execute_impl(context & ctx) const override;
};
// Complex Expressions
@ -462,8 +463,8 @@ struct binary_expression : public expression {
chk_type<expression>(this->right);
}
std::string type() const override { return "BinaryExpression"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"left", {left.get()}},
{"right", {right.get()}}
@ -476,10 +477,7 @@ struct binary_expression : public expression {
* Operator precedence: https://github.com/pallets/jinja/issues/379#issuecomment-168076202
*/
struct filter_expression : public expression {
// either an expression or a value is allowed
statement_ptr operand;
value_string val; // will be set by filter_statement
statement_ptr filter;
filter_expression(statement_ptr && operand, statement_ptr && filter)
@ -488,14 +486,9 @@ struct filter_expression : public expression {
chk_type<identifier, call_expression>(this->filter);
}
filter_expression(value_string && val, statement_ptr && filter)
: val(std::move(val)), filter(std::move(filter)) {
chk_type<identifier, call_expression>(this->filter);
}
std::string type() const override { return "FilterExpression"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"operand", {operand.get()}},
{"filter", {filter.get()}}
@ -512,8 +505,8 @@ struct filter_statement : public statement {
chk_type<identifier, call_expression>(this->filter);
}
std::string type() const override { return "FilterStatement"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"filter", {filter.get()}},
{"body", stmts_to_ptr(body)}
@ -537,14 +530,14 @@ struct select_expression : public expression {
chk_type<expression>(this->test);
}
std::string type() const override { return "SelectExpression"; }
value execute_impl(context & ctx) override {
value execute_impl(context & ctx) const override {
auto predicate = test->execute_impl(ctx);
if (!predicate->as_bool()) {
return mk_val<value_undefined>();
}
return lhs->execute_impl(ctx);
}
void visit(context & ctx) override {
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"lhs", {lhs.get()}},
{"test", {test.get()}}
@ -567,8 +560,8 @@ struct test_expression : public expression {
chk_type<identifier, call_expression>(this->test);
}
std::string type() const override { return "TestExpression"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"operand", {operand.get()}},
{"test", {test.get()}}
@ -588,8 +581,8 @@ struct unary_expression : public expression {
chk_type<expression>(this->argument);
}
std::string type() const override { return "UnaryExpression"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"argument", {argument.get()}}
});
@ -608,10 +601,10 @@ struct slice_expression : public expression {
chk_type<expression>(this->step_expr);
}
std::string type() const override { return "SliceExpression"; }
[[noreturn]] value execute_impl(context &) override {
[[noreturn]] value execute_impl(context &) const override {
throw std::runtime_error("must be handled by MemberExpression");
}
void visit(context & ctx) override {
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"start_expr", {start_expr.get()}},
{"stop_expr", {stop_expr.get()}},
@ -630,8 +623,8 @@ struct keyword_argument_expression : public expression {
chk_type<expression>(this->val);
}
std::string type() const override { return "KeywordArgumentExpression"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"key", {key.get()}},
{"val", {val.get()}}
@ -645,7 +638,7 @@ struct spread_expression : public expression {
chk_type<expression>(this->argument);
}
std::string type() const override { return "SpreadExpression"; }
void visit(context & ctx) override {
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"argument", {argument.get()}}
});
@ -663,8 +656,8 @@ struct call_statement : public statement {
for (const auto & arg : this->caller_args) chk_type<expression>(arg);
}
std::string type() const override { return "CallStatement"; }
value execute_impl(context & ctx) override;
void visit(context & ctx) override {
value execute_impl(context & ctx) const override;
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"call", {call.get()}},
{"caller_args", stmts_to_ptr(caller_args)},
@ -685,7 +678,7 @@ struct ternary_expression : public expression {
chk_type<expression>(this->false_expr);
}
std::string type() const override { return "Ternary"; }
value execute_impl(context & ctx) override {
value execute_impl(context & ctx) const override {
value cond_val = condition->execute(ctx);
if (cond_val->as_bool()) {
return true_expr->execute(ctx);
@ -693,7 +686,7 @@ struct ternary_expression : public expression {
return false_expr->execute(ctx);
}
}
void visit(context & ctx) override {
void visit(context & ctx) const override {
ctx.visitor(false, this, {
{"condition", {condition.get()}},
{"true_expr", {true_expr.get()}},

View file

@ -82,6 +82,9 @@ struct common_json_value {
// note: a nested pair {"a", "b"} does not build, use common_json::array({"a", "b"}) for an array
common_json_value(std::initializer_list<common_json_item> items);
template <typename T, typename std::enable_if<std::is_enum<T>::value, int>::type = 0>
common_json_value(T val) : common_json_value((typename std::underlying_type<T>::type) val) {}
template <typename T, typename std::enable_if<std::is_integral<T>::value && !std::is_same<T, bool>::value, int>::type = 0>
common_json_value(T val) : type(std::is_signed<T>::value ? VAL_INT : VAL_UINT) {
if (std::is_signed<T>::value) {
@ -111,6 +114,7 @@ struct common_json_item {
// the types common_json_value holds on its own
// anything else reaches its common_json ctor and recurses forever
template <typename T> struct common_json_is_value : std::integral_constant<bool,
std::is_enum<T>::value ||
std::is_arithmetic<T>::value ||
std::is_same<T, std::nullptr_t>::value ||
std::is_same<T, std::string>::value ||

View file

@ -130,7 +130,7 @@ common_chat_params common_chat_params_init_muse_glimmer(const common_chat_templa
});
data.grammar_triggers = {
{ COMMON_GRAMMAR_TRIGGER_TYPE_PATTERN,
"<\\|start\\|>assistant( to=(?!self<\\|message\\|>)(?!user<\\|message\\|>)[^<]*?<\\|message\\|>)" },
"(?:^|<\\|start\\|>assistant)( to=(?!self<\\|message\\|>)(?!user<\\|message\\|>)[^<]*?<\\|message\\|>)" },
};
}

View file

@ -776,6 +776,36 @@ class ModelBase:
raw = torch.cat((s.unsqueeze(-1), qs.to(torch.uint8)), dim=-1)
return raw.reshape(rows, n_blocks * 17).cpu().numpy()
def _mxfp4_expert_tensor(self, loaders: list[tuple[Callable[[], Tensor], Callable[[], Tensor]]]):
"""
One stacked [n_expert, rows, cols] MXFP4 tensor, built lazily.
gguf_writer holds every added tensor until the final write, so building
this eagerly (like the DeepSeek-V4 path does) keeps every expert in
memory at once. lazy means only the tensor being written is resident.
"""
# meta shapes, so this does not read any weights
rows, packed_cols = loaders[0][0]().shape
n_blocks = (packed_cols * 2) // 32
byte_shape = (len(loaders), rows, n_blocks * 17)
def load(fns: list[tuple[Callable[[], Tensor], Callable[[], Tensor]]]) -> np.ndarray:
out = np.empty(byte_shape, dtype=np.uint8)
for eid, (packed_fn, scale_fn) in enumerate(fns):
out[eid] = self.repack_mxfp4_blocks(
LazyTorchTensor.to_eager(packed_fn()),
LazyTorchTensor.to_eager(scale_fn()),
)
return out
# loaders goes through args, not the closure, so that `func` matches
# LazyBase's single-argument shape
return gguf.LazyNumpyTensor(
meta=gguf.LazyNumpyTensor.meta_with_dtype_and_shape(np.uint8, byte_shape),
args=(loaders,),
func=load,
)
@staticmethod
def _nvfp4_pack(weight: Tensor, scale: Tensor) -> tuple[np.ndarray, list[int]]:
"""Repack NVFP4 ModelOpt tensors into ggml super-block layout.

View file

@ -159,32 +159,14 @@ class HunYuanMoEModel(TextModel):
class HunYuanModel(TextModel):
model_arch = gguf.MODEL_ARCH.HUNYUAN_DENSE
def _get_eod_token_id(self) -> int | None:
"""Get the actual end-of-generation token from config (eod_token_id)."""
return self.hparams.get("eod_token_id")
def _get_eot_token_id(self) -> int | None:
"""Get the end-of-turn token from generation_config.json.
This is the first entry in eos_token_id when it's a list."""
gen_cfg_path = self.dir_model / "generation_config.json"
if gen_cfg_path.is_file():
with open(gen_cfg_path, encoding="utf-8") as f:
gen_cfg = json.load(f)
eos = gen_cfg.get("eos_token_id")
if isinstance(eos, list) and len(eos) >= 2:
return eos[0]
return None
def _fix_special_tokens(self):
"""Fix EOS/EOT tokens that are incorrect in upstream configs."""
eod_id = self._get_eod_token_id()
if eod_id is not None:
self.gguf_writer.add_eos_token_id(eod_id)
eot_id = self._get_eot_token_id()
if eot_id is not None:
self.gguf_writer.add_eot_token_id(eot_id)
def set_vocab(self):
# Also called by draft models (e.g. DFlash), with dir_model pointing at
# the target model.
config = ModelBase.load_hparams(self.dir_model, self.is_mistral_format)
config = {**config, **config.get("text_config", {})}
self.hparams["pad_token_id"] = config.get("pad_token_id")
self.hparams["eod_token_id"] = config.get("eod_token_id")
if (self.dir_model / "tokenizer.json").is_file():
tokens, toktypes, tokpre = self.get_vocab_base()
self.gguf_writer.add_tokenizer_model("gpt2")
@ -199,7 +181,6 @@ class HunYuanModel(TextModel):
token_types = ('bos', 'eos', 'unk', 'sep', 'cls', 'mask')
special_vocab = gguf.SpecialVocab(self.dir_model, load_merges=True, special_token_types=token_types)
special_vocab.add_to_gguf(self.gguf_writer)
self._fix_special_tokens()
else:
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(self.dir_model, trust_remote_code=True)
@ -251,7 +232,18 @@ class HunYuanModel(TextModel):
# FIX for BOS token: Overwrite incorrect id read from config.json
if self.hparams['hidden_size'] == 4096:
self.gguf_writer.add_bos_token_id(127958) # only for 7b dense, fix <|bos|> token
self._fix_special_tokens()
# Fix EOS/EOT tokens that are incorrect in upstream configs.
eod_id = self.hparams.get("eod_token_id")
if eod_id is not None:
self.gguf_writer.add_eos_token_id(eod_id)
gen_cfg = self.dir_model / "generation_config.json"
if gen_cfg.is_file():
with open(gen_cfg, encoding="utf-8") as f:
eos = json.load(f).get("eos_token_id")
if isinstance(eos, list) and len(eos) >= 2:
self.gguf_writer.add_eot_token_id(eos[0])
def set_gguf_parameters(self):
# Some HunYuanVL variants set num_experts=1 (not real MoE);

View file

@ -2,15 +2,14 @@ from __future__ import annotations
import re
from pathlib import Path
from typing import Callable, Iterable, Iterator, TYPE_CHECKING
from typing import Iterable, Iterator, TYPE_CHECKING
import numpy as np
import torch
if TYPE_CHECKING:
from torch import Tensor
from .base import LazyTorchTensor, ModelBase, TextModel, gguf, logger
from .base import ModelBase, TextModel, gguf, logger
from .kimi_linear import KimiLinearModel
@ -104,36 +103,6 @@ class KimiK3Model(TextModel):
"only the routed experts have a repack path"
)
def _mxfp4_expert_tensor(self, loaders: list[tuple[Callable[[], Tensor], Callable[[], Tensor]]]):
"""
One stacked [n_expert, rows, cols] MXFP4 tensor, built lazily.
gguf_writer holds every added tensor until the final write, so building
this eagerly (like the DeepSeek-V4 path does) keeps all ~1.38 TB of
experts in memory. lazy means only the tensor being written is resident.
"""
# meta shapes, so this does not read any weights
rows, packed_cols = loaders[0][0]().shape
n_blocks = (packed_cols * 2) // 32
byte_shape = (len(loaders), rows, n_blocks * 17)
def load(fns: list[tuple[Callable[[], Tensor], Callable[[], Tensor]]]) -> np.ndarray:
out = np.empty(byte_shape, dtype=np.uint8)
for eid, (packed_fn, scale_fn) in enumerate(fns):
out[eid] = self.repack_mxfp4_blocks(
LazyTorchTensor.to_eager(packed_fn()),
LazyTorchTensor.to_eager(scale_fn()),
)
return out
# loaders goes through args, not the closure, so that `func` matches
# LazyBase's single-argument shape
return gguf.LazyNumpyTensor(
meta=gguf.LazyNumpyTensor.meta_with_dtype_and_shape(np.uint8, byte_shape),
args=(loaders,),
func=load,
)
def _write_mxfp4_experts(self) -> None:
n_experts = self.hparams["num_experts"]

View file

@ -10,7 +10,7 @@ import torch
if TYPE_CHECKING:
from torch import Tensor
from .base import MmprojModel, ModelBase, TextModel, gguf
from .base import MmprojModel, ModelBase, TextModel, gguf, logger
@ModelBase.register("MiMoV2FlashForCausalLM", "MiMoV2ForCausalLM")
@ -167,6 +167,84 @@ class MimoV2Model(TextModel):
self.gguf_writer.add_nextn_predict_layers(self._n_nextn)
_MXFP4_EXPERT_RE = re.compile(
r"^model\.layers\.(\d+)\.mlp\.experts\.(\d+)\.(gate|up|down)_proj\.weight$"
)
_MXFP4_PROJ = {
"gate": gguf.MODEL_TENSOR.FFN_GATE_EXP,
"up": gguf.MODEL_TENSOR.FFN_UP_EXP,
"down": gguf.MODEL_TENSOR.FFN_DOWN_EXP,
}
def _is_mxfp4_packed(self) -> bool:
quant_config = self.hparams.get("quantization_config") or {}
if quant_config.get("store_dtype") != "mxfp4":
return False
# repack_mxfp4_blocks assumes ggml's 32-element group
block_size = quant_config.get("mxfp4_block_size", 32)
if block_size != 32:
raise NotImplementedError(
f"MXFP4 block size {block_size} is not ggml's QK_MXFP4 (32)")
return True
def _write_mxfp4_experts(self) -> None:
n_experts = self.hparams["n_routed_experts"]
# the FP8 half uses `weight_scale_inv` and is left to dequant_model
stray = [n for n in self.model_tensors
if n.endswith(".weight_scale") and not self._MXFP4_EXPERT_RE.match(n.removesuffix("_scale"))]
if stray:
raise NotImplementedError(
f"{len(stray)} MXFP4 tensor(s) outside the routed experts, e.g. {stray[0]!r}; "
"only the routed experts have a repack path"
)
# (bid, proj) -> {expert id: (weight name, scale name)}
groups: dict[tuple[int, str], dict[int, tuple[str, str]]] = {}
for name in self.model_tensors:
m = self._MXFP4_EXPERT_RE.match(name)
if m is None:
continue
bid, eid, proj = int(m.group(1)), int(m.group(2)), m.group(3)
scale_name = name + "_scale"
if scale_name not in self.model_tensors:
raise KeyError(f"missing {scale_name} for {name}")
groups.setdefault((bid, proj), {})[eid] = (name, scale_name)
consumed: list[str] = []
for (bid, proj), experts in sorted(groups.items()):
missing = [e for e in range(n_experts) if e not in experts]
if missing or len(experts) != n_experts:
raise KeyError(
f"layer {bid} {proj}_proj: {len(experts)} of {n_experts} experts present"
+ (f", first missing is {missing[0]}" if missing else "")
)
loaders = []
for eid in range(n_experts):
weight_name, scale_name = experts[eid]
loaders.append((self.model_tensors[weight_name], self.model_tensors[scale_name]))
consumed += [weight_name, scale_name]
data = self._mxfp4_expert_tensor(loaders)
new_name = self.format_tensor_name(self._MXFP4_PROJ[proj], bid)
shape = gguf.quant_shape_from_byte_shape(data.shape, gguf.GGMLQuantizationType.MXFP4)
logger.info(
f"{new_name}: repacked {n_experts} experts to MXFP4, "
f"shape = {{{', '.join(str(n) for n in reversed(shape))}}}"
)
self.gguf_writer.add_tensor(new_name, data, raw_dtype=gguf.GGMLQuantizationType.MXFP4)
for name in consumed:
del self.model_tensors[name]
def generate_extra_tensors(self) -> Iterable[tuple[str, Tensor]]:
# not a generator on purpose: base.py chains this with get_tensors(), so the
# tensors used here must be removed from model_tensors before that starts
if self._is_mxfp4_packed():
self._write_mxfp4_experts()
return ()
_experts: list[dict[str, Tensor]] | None = None
@classmethod
@ -192,7 +270,7 @@ class MimoV2Model(TextModel):
bid = new_bid
# process the experts separately
if name.find("mlp.experts") != -1:
if ".mlp.experts." in name and name.endswith(".weight"):
n_experts = self.hparams["n_routed_experts"]
assert bid is not None
@ -229,6 +307,10 @@ class MimoV2Model(TextModel):
if len(experts) > 0:
raise ValueError(f"Unprocessed experts: {experts}")
if self._is_mxfp4_packed():
self._is_mxfp4 = True
self.ftype = gguf.LlamaFileType.MOSTLY_MXFP4_MOE
@ModelBase.register("MiMoV2ForCausalLM")
@ModelBase.example("XiaomiMiMo/MiMo-V2.5")
@ -382,6 +464,8 @@ class MiMoV2VisionAudioModel(MmprojModel):
"_codebook.inited",
)
for name, tensor in state_dict.items():
if name.startswith("decoder."):
continue
if name.endswith(skip_suffixes):
continue
if m := codebook_re.match(name):

View file

@ -39,6 +39,8 @@
#define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8
#define ggml_gemv_q1_0_4x4_q8_0_generic ggml_gemv_q1_0_4x4_q8_0
#define ggml_gemv_q1_0_4x8_q8_0_generic ggml_gemv_q1_0_4x8_q8_0
#define ggml_gemv_q4_0_4x4_q8_0_generic ggml_gemv_q4_0_4x4_q8_0
#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0
#define ggml_gemv_q4_0_8x8_q8_0_generic ggml_gemv_q4_0_8x8_q8_0
@ -52,6 +54,8 @@
#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0
#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0
#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0
#define ggml_gemm_q1_0_4x4_q8_0_generic ggml_gemm_q1_0_4x4_q8_0
#define ggml_gemm_q1_0_4x8_q8_0_generic ggml_gemm_q1_0_4x8_q8_0
#define ggml_gemm_q4_0_4x4_q8_0_generic ggml_gemm_q4_0_4x4_q8_0
#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0
#define ggml_gemm_q4_0_8x8_q8_0_generic ggml_gemm_q4_0_8x8_q8_0
@ -77,6 +81,8 @@
// repack.cpp
#define ggml_quantize_mat_q8_0_4x4_generic ggml_quantize_mat_q8_0_4x4
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
#define ggml_gemv_q1_0_4x4_q8_0_generic ggml_gemv_q1_0_4x4_q8_0
#define ggml_gemv_q1_0_4x8_q8_0_generic ggml_gemv_q1_0_4x8_q8_0
#define ggml_gemv_q4_0_4x4_q8_0_generic ggml_gemv_q4_0_4x4_q8_0
#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0
#define ggml_gemv_q4_K_8x4_q8_K_generic ggml_gemv_q4_K_8x4_q8_K
@ -87,6 +93,8 @@
#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0
#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0
#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0
#define ggml_gemm_q1_0_4x4_q8_0_generic ggml_gemm_q1_0_4x4_q8_0
#define ggml_gemm_q1_0_4x8_q8_0_generic ggml_gemm_q1_0_4x8_q8_0
#define ggml_gemm_q4_0_4x4_q8_0_generic ggml_gemm_q4_0_4x4_q8_0
#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0
#define ggml_gemm_q4_K_8x4_q8_K_generic ggml_gemm_q4_K_8x4_q8_K
@ -112,6 +120,8 @@
#define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8
#define ggml_gemv_q1_0_4x4_q8_0_generic ggml_gemv_q1_0_4x4_q8_0
#define ggml_gemv_q1_0_4x8_q8_0_generic ggml_gemv_q1_0_4x8_q8_0
#define ggml_gemv_q4_0_4x4_q8_0_generic ggml_gemv_q4_0_4x4_q8_0
#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0
#define ggml_gemv_q4_0_8x8_q8_0_generic ggml_gemv_q4_0_8x8_q8_0
@ -125,6 +135,8 @@
#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0
#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0
#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0
#define ggml_gemm_q1_0_4x4_q8_0_generic ggml_gemm_q1_0_4x4_q8_0
#define ggml_gemm_q1_0_4x8_q8_0_generic ggml_gemm_q1_0_4x8_q8_0
#define ggml_gemm_q4_0_4x4_q8_0_generic ggml_gemm_q4_0_4x4_q8_0
#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0
#define ggml_gemm_q4_0_8x8_q8_0_generic ggml_gemm_q4_0_8x8_q8_0
@ -153,6 +165,8 @@
#define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8
#define ggml_gemv_q1_0_4x4_q8_0_generic ggml_gemv_q1_0_4x4_q8_0
#define ggml_gemv_q1_0_4x8_q8_0_generic ggml_gemv_q1_0_4x8_q8_0
#define ggml_gemv_q4_0_4x4_q8_0_generic ggml_gemv_q4_0_4x4_q8_0
#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0
#define ggml_gemv_q4_0_8x8_q8_0_generic ggml_gemv_q4_0_8x8_q8_0
@ -166,6 +180,8 @@
#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0
#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0
#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0
#define ggml_gemm_q1_0_4x4_q8_0_generic ggml_gemm_q1_0_4x4_q8_0
#define ggml_gemm_q1_0_4x8_q8_0_generic ggml_gemm_q1_0_4x8_q8_0
#define ggml_gemm_q4_0_4x4_q8_0_generic ggml_gemm_q4_0_4x4_q8_0
#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0
#define ggml_gemm_q4_0_8x8_q8_0_generic ggml_gemm_q4_0_8x8_q8_0
@ -188,6 +204,8 @@
#define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8
#define ggml_gemv_q1_0_4x4_q8_0_generic ggml_gemv_q1_0_4x4_q8_0
#define ggml_gemv_q1_0_4x8_q8_0_generic ggml_gemv_q1_0_4x8_q8_0
#define ggml_gemv_q4_0_4x4_q8_0_generic ggml_gemv_q4_0_4x4_q8_0
#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0
#define ggml_gemv_q2_K_8x8_q8_K_generic ggml_gemv_q2_K_8x8_q8_K
@ -200,6 +218,8 @@
#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0
#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0
#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0
#define ggml_gemm_q1_0_4x4_q8_0_generic ggml_gemm_q1_0_4x4_q8_0
#define ggml_gemm_q1_0_4x8_q8_0_generic ggml_gemm_q1_0_4x8_q8_0
#define ggml_gemm_q4_0_4x4_q8_0_generic ggml_gemm_q4_0_4x4_q8_0
#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0
#define ggml_gemm_q2_K_8x8_q8_K_generic ggml_gemm_q2_K_8x8_q8_K
@ -231,6 +251,8 @@
#define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8
#define ggml_gemv_q1_0_4x4_q8_0_generic ggml_gemv_q1_0_4x4_q8_0
#define ggml_gemv_q1_0_4x8_q8_0_generic ggml_gemv_q1_0_4x8_q8_0
#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0
#define ggml_gemv_q4_0_8x8_q8_0_generic ggml_gemv_q4_0_8x8_q8_0
#define ggml_gemv_q2_K_8x8_q8_K_generic ggml_gemv_q2_K_8x8_q8_K
@ -243,6 +265,8 @@
#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0
#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0
#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0
#define ggml_gemm_q1_0_4x4_q8_0_generic ggml_gemm_q1_0_4x4_q8_0
#define ggml_gemm_q1_0_4x8_q8_0_generic ggml_gemm_q1_0_4x8_q8_0
#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0
#define ggml_gemm_q4_0_8x8_q8_0_generic ggml_gemm_q4_0_8x8_q8_0
#define ggml_gemm_q2_K_8x8_q8_K_generic ggml_gemm_q2_K_8x8_q8_K
@ -277,6 +301,8 @@
#define ggml_quantize_mat_q8_0_4x8_generic ggml_quantize_mat_q8_0_4x8
#define ggml_quantize_mat_q8_K_4x4_generic ggml_quantize_mat_q8_K_4x4
#define ggml_quantize_mat_q8_K_4x8_generic ggml_quantize_mat_q8_K_4x8
#define ggml_gemv_q1_0_4x4_q8_0_generic ggml_gemv_q1_0_4x4_q8_0
#define ggml_gemv_q1_0_4x8_q8_0_generic ggml_gemv_q1_0_4x8_q8_0
#define ggml_gemv_q4_0_4x4_q8_0_generic ggml_gemv_q4_0_4x4_q8_0
#define ggml_gemv_q4_0_4x8_q8_0_generic ggml_gemv_q4_0_4x8_q8_0
#define ggml_gemv_q4_0_8x8_q8_0_generic ggml_gemv_q4_0_8x8_q8_0
@ -290,6 +316,8 @@
#define ggml_gemv_iq4_nl_4x4_q8_0_generic ggml_gemv_iq4_nl_4x4_q8_0
#define ggml_gemv_q8_0_4x4_q8_0_generic ggml_gemv_q8_0_4x4_q8_0
#define ggml_gemv_q8_0_4x8_q8_0_generic ggml_gemv_q8_0_4x8_q8_0
#define ggml_gemm_q1_0_4x4_q8_0_generic ggml_gemm_q1_0_4x4_q8_0
#define ggml_gemm_q1_0_4x8_q8_0_generic ggml_gemm_q1_0_4x8_q8_0
#define ggml_gemm_q4_0_4x4_q8_0_generic ggml_gemm_q4_0_4x4_q8_0
#define ggml_gemm_q4_0_4x8_q8_0_generic ggml_gemm_q4_0_4x8_q8_0
#define ggml_gemm_q4_0_8x8_q8_0_generic ggml_gemm_q4_0_8x8_q8_0

View file

@ -48,6 +48,24 @@ static inline void decode_q_Kx8_6bit_scales(const uint8_t * scales_in, int16x8_t
}
#endif
#if defined(__aarch64__) && defined(__ARM_NEON) && (defined(__ARM_FEATURE_DOTPROD) || defined(__ARM_FEATURE_MATMUL_INT8))
#define B1(c,s,n) 0x ## n ## c , 0x ## n ## s
#define B2(c,s,n) B1(c,s,n ## c), B1(c,s,n ## s)
#define B3(c,s,n) B2(c,s,n ## c), B2(c,s,n ## s)
#define B4(c,s,n) B3(c,s,n ## c), B3(c,s,n ## s)
#define B5(c,s,n) B4(c,s,n ## c), B4(c,s,n ## s)
#define B6(c,s,n) B5(c,s,n ## c), B5(c,s,n ## s)
#define B7(c,s,n) B6(c,s,n ## c), B6(c,s,n ## s)
#define B8(c,s ) B7(c,s, c), B7(c,s, s)
static const uint64_t table_q1_signs[256] = { B8(ff, 01) };
static inline int8x16_t ggml_q1_0_unpack_pair(uint8_t bits0, uint8_t bits1) {
return vreinterpretq_s8_u8(vcombine_u8(vcreate_u8(table_q1_signs[bits0]),
vcreate_u8(table_q1_signs[bits1])));
}
#endif
void ggml_quantize_mat_q8_0_4x4(const float * GGML_RESTRICT x, void * GGML_RESTRICT vy, int64_t k) {
assert(QK8_0 == 32);
assert(k % QK8_0 == 0);
@ -1748,6 +1766,132 @@ void ggml_gemv_q8_0_4x8_q8_0(int n,
ggml_gemv_q8_0_4x8_q8_0_generic(n, s, bs, vx, vy, nr, nc);
}
void ggml_gemv_q1_0_4x4_q8_0(int n,
float * GGML_RESTRICT s,
size_t bs,
const void * GGML_RESTRICT vx,
const void * GGML_RESTRICT vy,
int nr,
int nc) {
const int qk = QK1_0;
const int nb = n / qk;
const int ncols_interleaved = 4;
assert(n % qk == 0);
assert(nc % ncols_interleaved == 0);
UNUSED(nb);
UNUSED(ncols_interleaved);
#if defined(__aarch64__) && defined(__ARM_NEON) && defined(__ARM_FEATURE_DOTPROD)
for (int c = 0; c < nc; c += ncols_interleaved) {
const block_q1_0x4 * b_ptr = (const block_q1_0x4 *) vx + (c / ncols_interleaved) * nb;
const block_q8_0 * a_ptr = (const block_q8_0 *) vy;
float32x4_t acc = vdupq_n_f32(0);
for (int l = 0; l < nb; l++) {
const float32x4_t b_d = vcvt_f32_f16(vld1_f16((const float16_t *) b_ptr[l].d));
float32x4_t accb = vdupq_n_f32(0);
for (int k = 0; k < 4; k++) {
const block_q8_0 * GGML_RESTRICT a_blk = a_ptr + l * 4 + k;
const float ad = GGML_CPU_FP16_TO_FP32(a_blk->d);
int32x4_t ret = vdupq_n_s32(0);
for (int tile = 0; tile < 8; tile += 4) {
const int8x16_t signs0 = ggml_q1_0_unpack_pair(b_ptr[l].qs[k * 16 + 2 * (tile + 0) + 0],
b_ptr[l].qs[k * 16 + 2 * (tile + 0) + 1]);
const int8x16_t signs1 = ggml_q1_0_unpack_pair(b_ptr[l].qs[k * 16 + 2 * (tile + 1) + 0],
b_ptr[l].qs[k * 16 + 2 * (tile + 1) + 1]);
const int8x16_t signs2 = ggml_q1_0_unpack_pair(b_ptr[l].qs[k * 16 + 2 * (tile + 2) + 0],
b_ptr[l].qs[k * 16 + 2 * (tile + 2) + 1]);
const int8x16_t signs3 = ggml_q1_0_unpack_pair(b_ptr[l].qs[k * 16 + 2 * (tile + 3) + 0],
b_ptr[l].qs[k * 16 + 2 * (tile + 3) + 1]);
const int8x16_t q_tiles = vld1q_s8(a_blk->qs + tile * 4);
ret = vdotq_laneq_s32(ret, signs0, q_tiles, 0);
ret = vdotq_laneq_s32(ret, signs1, q_tiles, 1);
ret = vdotq_laneq_s32(ret, signs2, q_tiles, 2);
ret = vdotq_laneq_s32(ret, signs3, q_tiles, 3);
}
accb = vfmaq_n_f32(accb, vcvtq_f32_s32(ret), ad);
}
acc = vfmaq_f32(acc, accb, b_d);
}
vst1q_f32(s, acc);
s += ncols_interleaved;
}
return;
#endif
ggml_gemv_q1_0_4x4_q8_0_generic(n, s, bs, vx, vy, nr, nc);
}
void ggml_gemv_q1_0_4x8_q8_0(int n,
float * GGML_RESTRICT s,
size_t bs,
const void * GGML_RESTRICT vx,
const void * GGML_RESTRICT vy,
int nr,
int nc) {
const int qk = QK1_0;
const int nb = n / qk;
const int ncols_interleaved = 4;
assert(n % qk == 0);
assert(nc % ncols_interleaved == 0);
UNUSED(nb);
UNUSED(ncols_interleaved);
#if defined(__aarch64__) && defined(__ARM_NEON) && defined(__ARM_FEATURE_DOTPROD)
for (int c = 0; c < nc; c += ncols_interleaved) {
const block_q1_0x4 * b_ptr = (const block_q1_0x4 *) vx + (c / ncols_interleaved) * nb;
const block_q8_0 * a_ptr = (const block_q8_0 *) vy;
float32x4_t acc = vdupq_n_f32(0);
for (int l = 0; l < nb; l++) {
const float32x4_t b_d = vcvt_f32_f16(vld1_f16((const float16_t *) b_ptr[l].d));
float32x4_t accb = vdupq_n_f32(0);
for (int k = 0; k < 4; ++k) {
const block_q8_0 * GGML_RESTRICT a_blk = a_ptr + l * 4 + k;
const uint8_t * GGML_RESTRICT b_qs = (const uint8_t *) b_ptr[l].qs + k * 16;
const float ad = GGML_CPU_FP16_TO_FP32(a_blk->d);
int8x8x4_t a_chunks = vld1_s8_x4(a_blk->qs);
int8x16_t a0 = vcombine_s8(a_chunks.val[0], a_chunks.val[0]);
int8x16_t a1 = vcombine_s8(a_chunks.val[1], a_chunks.val[1]);
int8x16_t a2 = vcombine_s8(a_chunks.val[2], a_chunks.val[2]);
int8x16_t a3 = vcombine_s8(a_chunks.val[3], a_chunks.val[3]);
int32x4_t ret0 = vdupq_n_s32(0);
int32x4_t ret1 = vdupq_n_s32(0);
ret0 = vdotq_s32(ret0, ggml_q1_0_unpack_pair(b_qs[0], b_qs[1]), a0);
ret1 = vdotq_s32(ret1, ggml_q1_0_unpack_pair(b_qs[2], b_qs[3]), a0);
ret0 = vdotq_s32(ret0, ggml_q1_0_unpack_pair(b_qs[4], b_qs[5]), a1);
ret1 = vdotq_s32(ret1, ggml_q1_0_unpack_pair(b_qs[6], b_qs[7]), a1);
ret0 = vdotq_s32(ret0, ggml_q1_0_unpack_pair(b_qs[8], b_qs[9]), a2);
ret1 = vdotq_s32(ret1, ggml_q1_0_unpack_pair(b_qs[10], b_qs[11]), a2);
ret0 = vdotq_s32(ret0, ggml_q1_0_unpack_pair(b_qs[12], b_qs[13]), a3);
ret1 = vdotq_s32(ret1, ggml_q1_0_unpack_pair(b_qs[14], b_qs[15]), a3);
accb = vfmaq_n_f32(accb, vcvtq_f32_s32(vpaddq_s32(ret0, ret1)), ad);
}
acc = vfmaq_f32(acc, accb, b_d);
}
vst1q_f32(s, acc);
s += ncols_interleaved;
}
return;
#endif
ggml_gemv_q1_0_4x8_q8_0_generic(n, s, bs, vx, vy, nr, nc);
}
void ggml_gemm_q4_0_4x4_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, const void * GGML_RESTRICT vy, int nr, int nc) {
const int qk = QK8_0;
const int nb = n / qk;
@ -4998,3 +5142,168 @@ void ggml_gemm_q8_0_4x8_q8_0(int n,
#endif // defined(__aarch64__) && defined(__ARM_NEON) && defined(__ARM_FEATURE_MATMUL_INT8)
ggml_gemm_q8_0_4x8_q8_0_generic(n, s, bs, vx, vy, nr, nc);
}
void ggml_gemm_q1_0_4x4_q8_0(int n,
float * GGML_RESTRICT s,
size_t bs,
const void * GGML_RESTRICT vx,
const void * GGML_RESTRICT vy,
int nr,
int nc) {
const int qk = QK1_0;
const int nb = n / qk;
const int ncols_interleaved = 4;
assert(n % qk == 0);
assert(nr % 4 == 0);
assert(nc % ncols_interleaved == 0);
UNUSED(nb);
UNUSED(ncols_interleaved);
#if defined(__aarch64__) && defined(__ARM_NEON) && defined(__ARM_FEATURE_DOTPROD)
for (int y = 0; y < nr / 4; y++) {
const block_q8_0x4 * a_ptr = (const block_q8_0x4 *) vy + (4 * y * nb);
for (int x = 0; x < nc / ncols_interleaved; x++) {
const block_q1_0x4 * b_ptr = (const block_q1_0x4 *) vx + (x * nb);
float32x4_t sumf[4];
for (int m = 0; m < 4; m++) {
sumf[m] = vdupq_n_f32(0);
}
for (int l = 0; l < nb; l++) {
float32x4_t b_d = vcvt_f32_f16(vld1_f16((const float16_t *) b_ptr[l].d));
float32x4_t blockf_0 = vdupq_n_f32(0);
float32x4_t blockf_1 = vdupq_n_f32(0);
float32x4_t blockf_2 = vdupq_n_f32(0);
float32x4_t blockf_3 = vdupq_n_f32(0);
for (int k = 0; k < 4; ++k) {
const block_q8_0x4 * GGML_RESTRICT a_blk = a_ptr + 4 * l + k;
float32x4_t a_d = vcvt_f32_f16(vld1_f16((const float16_t *) a_blk->d));
int32x4_t sumi_0 = vdupq_n_s32(0);
int32x4_t sumi_1 = vdupq_n_s32(0);
int32x4_t sumi_2 = vdupq_n_s32(0);
int32x4_t sumi_3 = vdupq_n_s32(0);
for (int tile = 0; tile < 8; ++tile) {
const int8x16_t signs = ggml_q1_0_unpack_pair(b_ptr[l].qs[k * 16 + 2 * tile + 0],
b_ptr[l].qs[k * 16 + 2 * tile + 1]);
const int8x16_t a_tile = vld1q_s8(a_blk->qs + tile * 16);
sumi_0 = vdotq_laneq_s32(sumi_0, signs, a_tile, 0);
sumi_1 = vdotq_laneq_s32(sumi_1, signs, a_tile, 1);
sumi_2 = vdotq_laneq_s32(sumi_2, signs, a_tile, 2);
sumi_3 = vdotq_laneq_s32(sumi_3, signs, a_tile, 3);
}
blockf_0 = vfmaq_laneq_f32(blockf_0, vcvtq_f32_s32(sumi_0), a_d, 0);
blockf_1 = vfmaq_laneq_f32(blockf_1, vcvtq_f32_s32(sumi_1), a_d, 1);
blockf_2 = vfmaq_laneq_f32(blockf_2, vcvtq_f32_s32(sumi_2), a_d, 2);
blockf_3 = vfmaq_laneq_f32(blockf_3, vcvtq_f32_s32(sumi_3), a_d, 3);
}
sumf[0] = vfmaq_f32(sumf[0], blockf_0, b_d);
sumf[1] = vfmaq_f32(sumf[1], blockf_1, b_d);
sumf[2] = vfmaq_f32(sumf[2], blockf_2, b_d);
sumf[3] = vfmaq_f32(sumf[3], blockf_3, b_d);
}
for (int m = 0; m < 4; m++) {
vst1q_f32(s + (y * 4 + m) * bs + x * 4, sumf[m]);
}
}
}
return;
#endif
ggml_gemm_q1_0_4x4_q8_0_generic(n, s, bs, vx, vy, nr, nc);
}
void ggml_gemm_q1_0_4x8_q8_0(int n,
float * GGML_RESTRICT s,
size_t bs,
const void * GGML_RESTRICT vx,
const void * GGML_RESTRICT vy,
int nr,
int nc) {
const int qk = QK1_0;
const int nb = n / qk;
const int ncols_interleaved = 4;
assert(n % qk == 0);
assert(nr % 4 == 0);
assert(nc % ncols_interleaved == 0);
UNUSED(nb);
UNUSED(ncols_interleaved);
#if defined(__aarch64__) && defined(__ARM_NEON) && defined(__ARM_FEATURE_MATMUL_INT8)
for (int y = 0; y < nr / 4; y++) {
const block_q8_0x4 * a_ptr = (const block_q8_0x4 *) vy + (4 * y * nb);
for (int x = 0; x < nc / ncols_interleaved; x++) {
const block_q1_0x4 * b_ptr = (const block_q1_0x4 *) vx + (x * nb);
float32x4_t sumf[4];
for (int m = 0; m < 4; ++m) {
sumf[m] = vdupq_n_f32(0);
}
for (int l = 0; l < nb; l++) {
const float32x4_t b_d = vcvt_f32_f16(vld1_f16((const float16_t *) b_ptr[l].d));
float32x4_t blockf[4];
for (int m = 0; m < 4; ++m) {
blockf[m] = vdupq_n_f32(0);
}
for (int k = 0; k < 4; ++k) {
const block_q8_0x4 * GGML_RESTRICT a_blk = a_ptr + 4 * l + k;
const uint8_t * GGML_RESTRICT b_qs = (const uint8_t *) b_ptr[l].qs + k * 16;
int32x4_t acc[4];
for (int i = 0; i < 4; ++i) {
acc[i] = vdupq_n_s32(0);
}
for (int chunk = 0; chunk < 4; ++chunk) {
const int8x16_t a01 = vld1q_s8(a_blk->qs + chunk * 32);
const int8x16_t a23 = vld1q_s8(a_blk->qs + chunk * 32 + 16);
const int8x16_t b01 = ggml_q1_0_unpack_pair(b_qs[chunk * 4 + 0], b_qs[chunk * 4 + 1]);
const int8x16_t b23 = ggml_q1_0_unpack_pair(b_qs[chunk * 4 + 2], b_qs[chunk * 4 + 3]);
acc[0] = vmmlaq_s32(acc[0], a01, b01);
acc[1] = vmmlaq_s32(acc[1], a01, b23);
acc[2] = vmmlaq_s32(acc[2], a23, b01);
acc[3] = vmmlaq_s32(acc[3], a23, b23);
}
const int32x4_t row0 = vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1]));
const int32x4_t row1 = vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1]));
const int32x4_t row2 = vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3]));
const int32x4_t row3 = vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3]));
const float32x4_t a_d = vcvt_f32_f16(vld1_f16((const float16_t *) a_blk->d));
blockf[0] = vfmaq_laneq_f32(blockf[0], vcvtq_f32_s32(row0), a_d, 0);
blockf[1] = vfmaq_laneq_f32(blockf[1], vcvtq_f32_s32(row1), a_d, 1);
blockf[2] = vfmaq_laneq_f32(blockf[2], vcvtq_f32_s32(row2), a_d, 2);
blockf[3] = vfmaq_laneq_f32(blockf[3], vcvtq_f32_s32(row3), a_d, 3);
}
sumf[0] = vfmaq_f32(sumf[0], blockf[0], b_d);
sumf[1] = vfmaq_f32(sumf[1], blockf[1], b_d);
sumf[2] = vfmaq_f32(sumf[2], blockf[2], b_d);
sumf[3] = vfmaq_f32(sumf[3], blockf[3], b_d);
}
for (int m = 0; m < 4; ++m) {
vst1q_f32(s + (y * 4 + m) * bs + x * 4, sumf[m]);
}
}
}
return;
#endif
ggml_gemm_q1_0_4x8_q8_0_generic(n, s, bs, vx, vy, nr, nc);
}

View file

@ -1,5 +1,6 @@
#include "conv2d.cuh"
#include "convert.cuh"
#include "mma.cuh"
struct conv_params {
const int64_t IW, IH;
@ -111,6 +112,220 @@ static void conv2d_cuda(const float * X_D, const T * K_D, float * Y_D, const con
conv2d_kernel<T, whcn_layout><<<blocks, CUDA_CONV2D_BLOCK_SIZE, 0, st>>>(X_D, K_D, Y_D, P);
}
static __global__ void
conv2d_pad_f16(const float * input, half * output, int iw, int ih, int pw, int ph, int px, int py, int total) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= total) {
return;
}
const int x = i % pw - px, y = i / pw % ph - py, nc = i / (pw * ph);
output[i] = __float2half(
(unsigned) x < (unsigned) iw && (unsigned) y < (unsigned) ih ? input[(nc * ih + y) * iw + x] : 0.0f);
}
template <int KW, int KH, bool use_mma>
static __global__ void conv2d_implicit_gemm_f16(const half * __restrict__ input,
const half * __restrict__ weight,
float * __restrict__ output,
const conv_params P,
const int split_k) {
using namespace ggml_cuda_mma;
constexpr int warp_size = ggml_cuda_get_physical_warp_size();
constexpr int nthreads = 4 * warp_size;
constexpr int BM = 64, BN = 64, BK = 64;
constexpr int AS = BK / 2 + 4;
constexpr int BS = BN / 2 + 4;
__shared__ __align__(16) half2 a_s[BM][AS];
__shared__ __align__(16) half2 b_s[BK][BS];
const int tid = threadIdx.y * warp_size + threadIdx.x;
const int iw = int(P.IW), ih = int(P.IH), ow = int(P.OW), oh = int(P.OH);
const int kw = KW ? KW : int(P.KW), kh = KH ? KH : int(P.KH);
const int ic = int(P.IC), oc = int(P.OC);
const int sx = int(P.ST_X), sy = int(P.ST_Y);
const int dx = int(P.DL_X), dy = int(P.DL_Y);
const int n = blockIdx.z / split_k, split = blockIdx.z % split_k;
const int m0 = blockIdx.y * BM, n0 = blockIdx.x * BN;
const int k_total = ic * kw * kh;
const int load_lane = warp_size == 32 ? threadIdx.x : threadIdx.x % (BN / 2);
const int load_row = threadIdx.y * (warp_size / (BN / 2)) + (warp_size == 32 ? 0 : threadIdx.x / (BN / 2));
const int spatial = n0 + 2 * load_lane;
const int spatial0 = min(spatial, ow * oh - 1), spatial1 = min(spatial + 1, ow * oh - 1);
const int y0 = spatial0 / ow, x0 = spatial0 % ow;
const int y1 = spatial1 / ow, x1 = spatial1 % ow;
const int pos0 = y0 * sy * iw + x0 * sx, pos1 = y1 * sy * iw + x1 * sx;
[[maybe_unused]] const int wm = threadIdx.y / 2 * 32, wn = threadIdx.y % 2 * 32;
#if defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)
using tile_ab = tile<16, 8, half2, get_input_data_layout()>;
# if defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)
// AMD accumulator fragments transpose the input fragment's row/column mapping.
using tile_c = tile<16, 16, float, DATA_LAYOUT_J_MAJOR>;
# else
using tile_c = tile<16, 16, float>;
# endif
[[maybe_unused]] tile_c c[2][2];
#else
if constexpr (use_mma) {
NO_DEVICE_CODE;
return;
}
#endif
constexpr int RM = 4, RN = BM * BN / (nthreads * RM);
[[maybe_unused]] const int simt_m = tid / (BN / RN) * RM, simt_n = tid % (BN / RN) * RN;
[[maybe_unused]] float c_simt[RM][RN] = {};
const int tiles = (k_total + BK - 1) / BK;
const int begin = int(int64_t(tiles) * split / split_k) * BK;
const int end = int(int64_t(tiles) * (split + 1) / split_k) * BK;
for (int k0 = begin; k0 < end; k0 += BK) {
if (k_total % 8 == 0 && uintptr_t(weight) % 16 == 0) {
#pragma unroll
for (int i = tid; i < BM * BK / 8; i += nthreads) {
const int row = i / (BK / 8), col = 8 * (i % (BK / 8));
const int4 v = m0 + row < oc && k0 + col < k_total ?
((const int4 *) weight)[((m0 + row) * k_total + k0 + col) / 8] :
make_int4(0, 0, 0, 0);
*(int4 *) &a_s[row][col / 2] = v;
}
} else {
#pragma unroll
for (int i = tid; i < BM * BK / 2; i += nthreads) {
const int row = i / (BK / 2), col = 2 * (i % (BK / 2));
half lo = __float2half(0.0f), hi = lo;
if (m0 + row < oc && k0 + col < k_total) {
lo = weight[(m0 + row) * k_total + k0 + col];
if (k0 + col + 1 < k_total) {
hi = weight[(m0 + row) * k_total + k0 + col + 1];
}
}
a_s[row][col / 2] = __halves2half2(lo, hi);
}
}
#pragma unroll
for (int k = load_row; k < BK; k += nthreads / (BN / 2)) {
const int ki = k0 + k;
const int ci = ki / (kw * kh), ky = ki / kw % kh, kx = ki % kw;
const int offset = ki < k_total ? (n * ic + ci) * ih * iw + ky * dy * iw + kx * dx : 0;
half lo = __float2half(0.0f), hi = lo;
if (ki < k_total && spatial < ow * oh) {
lo = input[offset + pos0];
}
if (ki < k_total && spatial + 1 < ow * oh) {
hi = input[offset + pos1];
}
b_s[k][load_lane] = __halves2half2(lo, hi);
}
__syncthreads();
if constexpr (use_mma) {
#if defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)
# pragma unroll
for (int k = 0; k < BK; k += 16) {
tile_ab a[2], b[2];
# pragma unroll
for (int i = 0; i < 2; ++i) {
load_ldmatrix(a[i], &a_s[wm + 16 * i][k / 2], AS);
load_ldmatrix_trans(b[i], &b_s[k][(wn + 16 * i) / 2], BS);
}
# pragma unroll
for (int i = 0; i < 2; ++i) {
# pragma unroll
for (int j = 0; j < 2; ++j) {
mma(c[i][j], a[i], b[j]);
}
}
}
#endif
} else {
#pragma unroll 4
for (int k = 0; k < BK; ++k) {
float a[RM], b[RN];
#pragma unroll
for (int i = 0; i < RM; ++i) {
a[i] = __half2float(((const half *) a_s[simt_m + i])[k]);
}
#pragma unroll
for (int j = 0; j < RN; ++j) {
b[j] = __half2float(((const half *) b_s[k])[simt_n + j]);
}
#pragma unroll
for (int i = 0; i < RM; ++i) {
#pragma unroll
for (int j = 0; j < RN; ++j) {
c_simt[i][j] += a[i] * b[j];
}
}
}
}
__syncthreads();
}
if constexpr (use_mma) {
#if defined(TURING_MMA_AVAILABLE) || defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)
# pragma unroll
for (int i = 0; i < 2; ++i) {
# pragma unroll
for (int j = 0; j < 2; ++j) {
# pragma unroll
for (int l = 0; l < c[i][j].ne; ++l) {
const int co = m0 + wm + 16 * i + c[i][j].get_i(l);
const int pos = n0 + wn + 16 * j + c[i][j].get_j(l);
if (co < oc && pos < ow * oh) {
output[(int64_t(blockIdx.z) * oc + co) * ow * oh + pos] = c[i][j].x[l];
}
}
}
}
#endif
} else {
#pragma unroll
for (int i = 0; i < RM; ++i) {
#pragma unroll
for (int j = 0; j < RN; ++j) {
const int co = m0 + simt_m + i, pos = n0 + simt_n + j;
if (co < oc && pos < ow * oh) {
output[(int64_t(blockIdx.z) * oc + co) * ow * oh + pos] = c_simt[i][j];
}
}
}
}
}
static __global__ void conv2d_reduce_split_k(const float * __restrict__ partial,
float * __restrict__ output,
const int total,
const int per_batch,
const int split_k) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= total) {
return;
}
const int n = i / per_batch;
const float * src = partial + int64_t(n) * (split_k - 1) * per_batch + i;
float sum = 0.0f;
for (int k = 0; k < split_k; ++k) {
sum += src[int64_t(k) * per_batch];
}
output[i] = sum;
}
template <bool use_mma>
static void conv2d_launch_implicit_gemm(const half * input,
const half * weight,
float * output,
const conv_params & params,
int split_k,
dim3 grid,
dim3 block,
cudaStream_t stream) {
if (params.KW == 3 && params.KH == 3) {
conv2d_implicit_gemm_f16<3, 3, use_mma><<<grid, block, 0, stream>>>(input, weight, output, params, split_k);
} else if (params.KW == 1 && params.KH == 1) {
conv2d_implicit_gemm_f16<1, 1, use_mma><<<grid, block, 0, stream>>>(input, weight, output, params, split_k);
} else {
conv2d_implicit_gemm_f16<0, 0, use_mma><<<grid, block, 0, stream>>>(input, weight, output, params, split_k);
}
}
static void conv2d_cuda_f16(const float * X_D, const half * K_D, float * Y_D, const conv_params P, cudaStream_t st) {
conv2d_cuda<half>(X_D, K_D, Y_D, P, st);
}
@ -126,6 +341,7 @@ void ggml_cuda_op_conv2d(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
const float * X_D = (const float *) input->data;
float * Y_D = (float *) dst->data;
GGML_ASSERT(input->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32);
GGML_ASSERT(ggml_is_contiguous(input));
GGML_ASSERT(ggml_is_contiguous(kernel));
GGML_ASSERT(kernel->type == GGML_TYPE_F16 || kernel->type == GGML_TYPE_F32);
@ -146,19 +362,86 @@ void ggml_cuda_op_conv2d(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
// No cwhn
GGML_ASSERT(p[6] == false);
const int IW = input->ne[0]; // input_w
const int IH = input->ne[1]; // input_h
const int OW = dst->ne[0]; // output_w
const int OH = dst->ne[1]; // output_h
const int KW = kernel->ne[0]; // kernel_w
const int KH = kernel->ne[1]; // kernel_h
const int IC = input->ne[2]; // input_channels
const int OC = kernel->ne[3]; // ouptut_chanles
const int B = input->ne[3]; // n_batches
const int64_t IW = input->ne[0]; // input_w
const int64_t IH = input->ne[1]; // input_h
const int64_t OW = dst->ne[0]; // output_w
const int64_t OH = dst->ne[1]; // output_h
const int64_t KW = kernel->ne[0]; // kernel_w
const int64_t KH = kernel->ne[1]; // kernel_h
const int64_t IC = input->ne[2]; // input_channels
const int64_t OC = kernel->ne[3]; // ouptut_chanles
const int64_t B = input->ne[3]; // n_batches
const int64_t total = B * OC * OH * OW;
conv_params params = { IW, IH, OW, OH, KW, KH, ST_X, ST_Y, PD_X, PD_Y, DL_X, DL_Y, IC, OC, B, total };
const auto & device = ggml_cuda_info().devices[ctx.device];
const bool use_mma =
turing_mma_available(device.cc) || amd_wmma_available(device.cc) || amd_mfma_available(device.cc);
// MUSA can share the tiling without a native fragment implementation in mma.cuh.
const bool use_simt = GGML_CUDA_CC_IS_MTHREADS(device.cc);
const bool pointwise = KW == 1 && KH == 1 && ST_X == 1 && ST_Y == 1 && PD_X == 0 && PD_Y == 0;
const bool use_blas = pointwise && fast_fp16_hardware_available(device.cc);
// Short reductions on small maps do not amortize conversion and launch costs.
const bool small_conv = IC * KW * KH < 64 && OW * OH < 512;
const int64_t limit = INT_MAX - 256;
const int64_t padded_w = IW + 2 * int64_t(PD_X), padded_h = IH + 2 * int64_t(PD_Y);
const bool padded_fits = padded_w > 0 && padded_w <= limit && padded_h > 0 && padded_h <= limit &&
padded_w * padded_h <= limit && IC * B <= limit / (padded_w * padded_h);
if (kernel->type == GGML_TYPE_F16 && (use_mma || use_blas || use_simt) && (use_blas || !small_conv) &&
ggml_nelements(input) <= limit && ggml_nelements(kernel) <= limit && total <= limit && padded_fits &&
PD_X >= 0 && PD_Y >= 0 && ST_X > 0 && ST_Y > 0 && DL_X > 0 && DL_Y > 0 &&
(OW - 1) * ST_X + (KW - 1) * DL_X < padded_w && (OH - 1) * ST_Y + (KH - 1) * DL_Y < padded_h &&
(OC + 63) / 64 <= 65535 && B <= 65535) {
const int pw = int(padded_w), ph = int(padded_h);
const int padded_total = int(padded_w * padded_h * IC * B);
ggml_cuda_pool_alloc<half> x_half(ctx.pool(), padded_total);
// Match im2col's F16 input precision, but expand patches only in shared memory and accumulate in F32.
if (PD_X == 0 && PD_Y == 0) {
ggml_get_to_fp16_cuda(input->type)(X_D, x_half.get(), padded_total, st);
} else {
conv2d_pad_f16<<<(padded_total + 255) / 256, 256, 0, st>>>(X_D, x_half.get(), int(IW), int(IH), pw, ph,
PD_X, PD_Y, padded_total);
}
const conv_params padded_params = { pw, ph, OW, OH, KW, KH, ST_X, ST_Y, 0, 0, DL_X, DL_Y, IC, OC, B, total };
if (use_blas) {
const float alpha = 1.0f, beta = 0.0f;
const int positions = int(OW * OH);
cublasHandle_t cublas_h = ctx.cublas_handle();
for (int n = 0; n < B; ++n) {
CUBLAS_CHECK(cublasGemmEx(cublas_h, CUBLAS_OP_N, CUBLAS_OP_N, positions, int(OC), int(IC), &alpha,
x_half.get() + int64_t(n) * IC * positions, CUDA_R_16F, positions, K_D,
CUDA_R_16F, int(IC), &beta, Y_D + int64_t(n) * OC * positions, CUDA_R_32F,
positions, CUBLAS_COMPUTE_32F, CUBLAS_GEMM_DEFAULT_TENSOR_OP));
}
return;
}
const int64_t blocks = ((OW * OH + 63) / 64) * ((OC + 63) / 64) * B;
const int target = 8 * ggml_cuda_info().devices[ctx.device].nsm;
// Split long reductions so small spatial maps still occupy the GPU.
const int split_k = int(std::min({ int64_t(32), int64_t(65535) / B, (IC * KW * KH + 63) / 64,
std::max(int64_t(1), (target + blocks - 1) / blocks) }));
ggml_cuda_pool_alloc<float> partial(ctx.pool());
float * result = split_k == 1 ? Y_D : partial.alloc(total * split_k);
const dim3 block(device.warp_size, 4);
const dim3 grid(unsigned((OW * OH + 63) / 64), unsigned((OC + 63) / 64), unsigned(B * split_k));
if (use_mma) {
conv2d_launch_implicit_gemm<true>(x_half.get(), (const half *) K_D, result, padded_params, split_k, grid,
block, st);
} else {
conv2d_launch_implicit_gemm<false>(x_half.get(), (const half *) K_D, result, padded_params, split_k, grid,
block, st);
}
if (split_k > 1) {
conv2d_reduce_split_k<<<(total + 255) / 256, 256, 0, st>>>(result, Y_D, int(total), int(OC * OW * OH),
split_k);
}
return;
}
if (kernel->type == GGML_TYPE_F16) {
conv2d_cuda_f16(X_D, (half *) K_D, Y_D, params, st);
} else {

View file

@ -439,6 +439,29 @@ static __global__ void convert_unary(
}
}
template <typename T> struct alignas(sizeof(T)*4) cvt_vec4 { T v[4]; };
// four elements per thread, so a warp moves 512B (RDNA) / 1k (CDNA) per load
template <typename src_t, typename dst_t>
static __global__ void convert_unary_cont_vec4(
const void * __restrict__ vx, dst_t * __restrict__ y, const int64_t k4) {
const int64_t i = (int64_t)blockDim.x*blockIdx.x + threadIdx.x;
if (i >= k4) {
return;
}
const cvt_vec4<src_t> xv = ((const cvt_vec4<src_t> *) vx)[i];
cvt_vec4<dst_t> yv;
#pragma unroll
for (int j = 0; j < 4; ++j) {
yv.v[j] = ggml_cuda_cast<dst_t>(xv.v[j]);
}
((cvt_vec4<dst_t> *) y)[i] = yv;
}
template <typename src_t, typename dst_t>
static void convert_unary_cuda(const void * vx, dst_t * y,
const int64_t ne00, const int64_t ne01, const int64_t ne02, const int64_t ne03,
@ -452,6 +475,15 @@ static void convert_unary_cuda(const void * vx, dst_t * y,
template <typename src_t, typename dst_t>
static void convert_unary_cont_cuda(const void * vx, dst_t * y, const int64_t k, cudaStream_t stream) {
if (k % 4 == 0 &&
(uintptr_t) vx % alignof(cvt_vec4<src_t>) == 0 &&
(uintptr_t) y % alignof(cvt_vec4<dst_t>) == 0) {
const int64_t k4 = k/4;
const int64_t num_blocks = (k4 + CUDA_DEQUANTIZE_BLOCK_SIZE - 1) / CUDA_DEQUANTIZE_BLOCK_SIZE;
convert_unary_cont_vec4<src_t, dst_t><<<num_blocks, CUDA_DEQUANTIZE_BLOCK_SIZE, 0, stream>>>(vx, y, k4);
return;
}
convert_unary_cuda<src_t>(vx, y, k, 1, 1, 1, k, k, k, stream);
}

View file

@ -2,7 +2,6 @@
#include "cp-async.cuh"
#include "mma.cuh"
#include "fattn-common.cuh"
#include "fattn-swizzle.cuh"
using namespace ggml_cuda_mma;
@ -327,6 +326,32 @@ static constexpr __device__ bool ggml_cuda_fattn_mma_get_Q_in_reg(const int DKQ,
return ggml_cuda_fattn_mma_get_config(DKQ, DV, ncols).Q_in_reg;
}
// Swizzling needs a tile stride that is a multiple of 32 half2 columns.
static constexpr __host__ __device__ bool ggml_cuda_fattn_mma_bank_aligned(const int nbatch_2) {
return nbatch_2 >= 32 && nbatch_2 % 32 == 0;
}
// Swizzling needs ldmatrix, on other hardware the tiles keep the row padding.
static __host__ bool ggml_cuda_fattn_mma_get_swizzled(const int DKQ, const int DV, const int ncols1, const int ncols2, const int cc) {
const fattn_mma_config cfg = ggml_cuda_fattn_mma_get_config(DKQ, DV, ncols1*ncols2, cc);
return turing_mma_available(cc) && ggml_cuda_fattn_mma_bank_aligned(cfg.nbatch_K2) && ggml_cuda_fattn_mma_bank_aligned(cfg.nbatch_V2);
}
static constexpr __device__ bool ggml_cuda_fattn_mma_get_swizzled(const int DKQ, const int DV, const int ncols1, const int ncols2) {
#if defined(TURING_MMA_AVAILABLE)
const fattn_mma_config cfg = ggml_cuda_fattn_mma_get_config(DKQ, DV, ncols1*ncols2);
return ggml_cuda_fattn_mma_bank_aligned(cfg.nbatch_K2) && ggml_cuda_fattn_mma_bank_aligned(cfg.nbatch_V2);
#else
GGML_UNUSED_VARS(DKQ, DV, ncols1, ncols2);
return false;
#endif // defined(TURING_MMA_AVAILABLE)
}
// Row padding is only needed if the tile is not swizzled.
static constexpr __host__ __device__ int ggml_cuda_fattn_mma_get_stride_tile(const int nbatch_2, const bool swizzled) {
return swizzled ? nbatch_2 : nbatch_2 + 4;
}
static constexpr __device__ int get_cols_per_thread() {
#if defined(AMD_WMMA_AVAILABLE) || defined(AMD_MFMA_AVAILABLE)
return 1; // AMD has a single column per thread.
@ -411,12 +436,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile(
for (int k0 = k0_start; k0 < k0_stop; k0 += stride_k) {
const int k = k0 + (stride_k == warp_size ? threadIdx.x : threadIdx.x % stride_k);
if constexpr (swz) {
const int smem_offs_b = ggml_cuda_fattn_smem_swizzle::bytes_rc<stride_tile>(i, k*h2_per_chunk);
cp_async_cg_16<preload>(tile_KV_32 + smem_offs_b, KV + i_KV*stride_KV + k*h2_per_chunk);
} else {
cp_async_cg_16<preload>(tile_KV_32 + i*(stride_tile*sizeof(half2)) + k*16, KV + i_KV*stride_KV + k*h2_per_chunk);
}
cp_async_cg_16<preload>(tile_KV_32 + swizzle_bytes<swz, half2>(i, k*h2_per_chunk, stride_tile), KV + i_KV*stride_KV + k*h2_per_chunk);
}
}
};
@ -458,11 +478,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_load_tile(
} else {
src = !oob_check || i < i_sup ? KV + int64_t(k_VKQ_0 + i)*stride_KV + k*h2_per_chunk : zero;
}
if constexpr (swz) {
ggml_cuda_memcpy_1<16>((char *) tile_KV + ggml_cuda_fattn_smem_swizzle::bytes_rc<stride_tile>(i, k*h2_per_chunk), src);
} else {
ggml_cuda_memcpy_1<16>(tile_KV + i*stride_tile + k*4, src);
}
ggml_cuda_memcpy_1<16>((char *) tile_KV + swizzle_bytes<swz, half2>(i, k*h2_per_chunk, stride_tile), src);
}
}
};
@ -605,11 +621,9 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
constexpr bool Q_in_reg = ggml_cuda_fattn_mma_get_Q_in_reg (DKQ, DV, ncols);
constexpr int nstages = ggml_cuda_fattn_mma_get_nstages (DKQ, DV, ncols1, ncols2, use_sparse);
// swizzle the tile stride for K and V based on the batch size.
constexpr int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2);
constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2);
constexpr bool swz_K = ggml_cuda_fattn_smem_swizzle::enabled(nbatch_K2);
constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_smem_swizzle::enabled(nbatch_V2);
constexpr bool swz = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols1, ncols2);
constexpr int stride_tile_K = ggml_cuda_fattn_mma_get_stride_tile(nbatch_K2, swz);
constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_mma_get_stride_tile(nbatch_V2, swz);
const int k_VKQ_0 = kb0 * nbatch_fa;
#if defined(TURING_MMA_AVAILABLE)
@ -627,7 +641,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
constexpr bool use_cp_async = true;
cp_async_wait_all();
__syncthreads();
flash_attn_ext_f16_load_tile<stride_tile_V, swz_V, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
flash_attn_ext_f16_load_tile<stride_tile_V, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(V_h2, tile_V, nbatch_V2, stride_V, k_VKQ_0, k_VKQ_sup, nullptr);
} else {
// the sparse mask values are gathered per element, always load them synchronously
@ -647,7 +661,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
if constexpr (nstages <= 1) {
const int k0_diff = k0_stop - k0_start;
constexpr bool use_cp_async = nstages == 1;
flash_attn_ext_f16_load_tile<stride_tile_K, swz_K, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
flash_attn_ext_f16_load_tile<stride_tile_K, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(K_h2 + k0_start, tile_K, k0_diff, stride_K, k_VKQ_0, k_VKQ_sup, indices);
if (use_cp_async) {
cp_async_wait_all();
@ -663,7 +677,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
#pragma unroll
for (int k_KQ_0 = k0_start; k_KQ_0 < k0_stop; k_KQ_0 += T_A_KQ::J) {
T_A_KQ K_A;
ggml_cuda_fattn_smem_swizzle::load_ldmatrix<stride_tile_K, swz_K>(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start);
load_ldmatrix<swz>(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start, stride_tile_K);
if constexpr (cols_per_warp == 8) {
mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[k_KQ_0/T_A_KQ::J]);
} else {
@ -689,7 +703,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
const int i_KQ_0 = i_KQ_00 + (threadIdx.y % np)*T_A_KQ::I;
T_A_KQ K_A;
ggml_cuda_fattn_smem_swizzle::load_ldmatrix<stride_tile_K, swz_K>(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start);
load_ldmatrix<swz>(K_A, tile_K, i_KQ_0, k_KQ_0 - k0_start, stride_tile_K);
if constexpr (cols_per_warp == 8) {
mma(KQ_C[i_KQ_00/(np*T_A_KQ::I)], K_A, Q_B[0]);
@ -984,7 +998,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
flash_attn_ext_f16_load_mask<ncols1, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(mask_h, tile_mask, stride_mask, k_VKQ_0 + nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, nullptr);
}
flash_attn_ext_f16_load_tile<stride_tile_K, swz_K, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
flash_attn_ext_f16_load_tile<stride_tile_K, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(K_h2, tile_K, nbatch_K2, stride_K, k_VKQ_0 + nbatch_fa, k_VKQ_sup, nullptr);
}
}
@ -1000,7 +1014,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
const int i0_diff = i0_stop - i0_start;
if (!V_is_K_view || i0_stop > 2*nbatch_K2) {
constexpr bool use_cp_async = nstages == 1;
flash_attn_ext_f16_load_tile<stride_tile_V, swz_V, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
flash_attn_ext_f16_load_tile<stride_tile_V, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(V_h2 + i0_start/2, tile_V, i0_diff/2, stride_V, k_VKQ_0, k_VKQ_sup, indices);
if (use_cp_async) {
cp_async_wait_all();
@ -1019,7 +1033,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::J;
T_A_VKQ A; // Transposed in SRAM but not in registers, gets transposed on load.
ggml_cuda_fattn_smem_swizzle::load_ldmatrix_trans<stride_tile_V, swz_V>(A, tile_V, (int)(tile_V_i - tile_V) + 2*k0*stride_tile_V + (i_VKQ_0 - i0_start)/2);
load_ldmatrix_trans<swz>(A, tile_V, 2*k0, (int)(tile_V_i - tile_V) + (i_VKQ_0 - i0_start)/2, stride_tile_V);
if constexpr (T_B_KQ::I == 8) {
mma(VKQ_C[i_VKQ_0/T_A_VKQ::I], A, B[k00/(np*T_A_VKQ::J)]);
} else {
@ -1045,7 +1059,8 @@ static __device__ __forceinline__ void flash_attn_ext_f16_iter(
const int k0 = k00 + (threadIdx.y % np)*T_A_VKQ::I;
T_A_VKQ A; // Transposed in both SRAM and registers, load normally.
ggml_cuda_fattn_smem_swizzle::load_ldmatrix<stride_tile_V, swz_V>(A, tile_V, (int)(tile_V_i - tile_V) + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2);
static_assert(!swz, "Volta has no ldmatrix");
load_ldmatrix(A, tile_V_i + k0*stride_tile_V + (i_VKQ_0 - i0_start)/2, stride_tile_V);
mma(VKQ_C[i_VKQ_0/i0_stride], B[k00/(np*T_A_VKQ::I)], A);
}
}
@ -1236,12 +1251,10 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(
static_assert(nwarps * (cols_per_warp/ncols2) % ncols1 == 0, "bad nwarps");
constexpr int stride_tile_Q = DKQ/2 + 4;
// swizzle the tile stride for K and V based on the batch size.
constexpr int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2);
constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2);
constexpr bool swz = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols1, ncols2);
constexpr int stride_tile_K = ggml_cuda_fattn_mma_get_stride_tile(nbatch_K2, swz);
constexpr int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_mma_get_stride_tile(nbatch_V2, swz);
constexpr int stride_tile_KV_max = stride_tile_K > stride_tile_V ? stride_tile_K : stride_tile_V;
constexpr bool swz_K = ggml_cuda_fattn_smem_swizzle::enabled(nbatch_K2);
constexpr bool swz_V = V_is_K_view ? swz_K : ggml_cuda_fattn_smem_swizzle::enabled(nbatch_V2);
extern __shared__ half2 tile_Q[];
half2 * tile_K = Q_in_reg ? tile_Q : tile_Q + ncols * stride_tile_Q;
@ -1338,7 +1351,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(
flash_attn_ext_f16_load_mask<ncols1, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(mask_h, tile_mask, stride_mask, kb0*nbatch_fa, k_VKQ_sup, jt*ncols1, ne01, nullptr);
}
flash_attn_ext_f16_load_tile<stride_tile_K, swz_K, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
flash_attn_ext_f16_load_tile<stride_tile_K, swz, nwarps, nbatch_fa, use_cp_async, oob_check, use_sparse>
(K_h2, tile_K, nbatch_K2, stride_K, kb0*nbatch_fa, k_VKQ_sup, nullptr);
}
@ -1503,14 +1516,12 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(
constexpr int tile_stride = nbatch_combine + 4;
static_assert((DV/2) % nbatch_combine == 0, "bad nbatch_combine");
constexpr bool combine_needs_sync = swz_K || swz_V;
if constexpr (cols_per_warp == 8) {
const int jc_cwmo = (threadIdx.x % (2*T_C_VKQ::J)) / T_C_VKQ::J; // jc combine write meta offset
const int jc_cwm = threadIdx.y*(2*T_C_VKQ::J) + 2*T_C_VKQ::get_j(-1) + jc_cwmo; // jc combine write meta
const float2 KQ_cmr = make_float2(KQ_max[jc_cwmo], KQ_rowsum[jc_cwmo]); // KQ combine max rowsum
if constexpr (combine_needs_sync) {
if constexpr (swz) {
__syncthreads();
}
@ -1550,7 +1561,7 @@ static __device__ __forceinline__ void flash_attn_ext_f16_process_tile(
const bool thread_should_write = T_C_KQ::J == 8 || T_C_KQ::get_j(threadIdx.x & 2) < 8;
#endif // defined(TURING_MMA_AVAILABLE)
if constexpr (combine_needs_sync) {
if constexpr (swz) {
__syncthreads();
}
@ -2025,8 +2036,9 @@ void ggml_cuda_flash_attn_ext_mma_f16_case(ggml_backend_cuda_context & ctx, ggml
constexpr bool V_is_K_view = DKQ == 576; // Guaranteed by the kernel selection logic in fattn.cu
// KV tile strides must match flash_attn_ext_f16_iter / _process_tile.
const int stride_tile_K = ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_K2, cc);
const int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_smem_swizzle::tile_stride(nbatch_V2, cc);
const bool swizzled = ggml_cuda_fattn_mma_get_swizzled(DKQ, DV, ncols1, ncols2, cc);
const int stride_tile_K = ggml_cuda_fattn_mma_get_stride_tile(nbatch_K2, swizzled);
const int stride_tile_V = V_is_K_view ? stride_tile_K : ggml_cuda_fattn_mma_get_stride_tile(nbatch_V2, swizzled);
const size_t nbytes_shared_KV_1stage = nbatch_fa * std::max(stride_tile_K, stride_tile_V) * sizeof(half2);
const size_t nbytes_shared_KV_2stage = nbatch_fa * (stride_tile_K + stride_tile_V) * sizeof(half2);
const size_t nbytes_shared_Q = ncols * (DKQ/2 + 4) * sizeof(half2);

View file

@ -1,126 +0,0 @@
#pragma once
#include "common.cuh"
#include "mma.cuh"
// XOR swizzle for K/V SMEM tiles to avoid bank conflicts without row padding (Turing+ only).
// Stride must be a multiple of 32 half2 columns, otherwise we keep +4 row padding.
namespace ggml_cuda_fattn_smem_swizzle {
static __host__ __device__ constexpr bool bank_aligned(const int nbatch_2) {
return nbatch_2 >= 32 && nbatch_2 % 32 == 0;
}
static __device__ constexpr bool enabled(const int nbatch_2) {
#if defined(TURING_MMA_AVAILABLE)
return bank_aligned(nbatch_2);
#else
GGML_UNUSED(nbatch_2);
return false;
#endif // defined(TURING_MMA_AVAILABLE)
}
static __host__ bool enabled(const int nbatch_2, const int cc) {
#ifdef GGML_USE_HIP
GGML_UNUSED(nbatch_2);
GGML_UNUSED(cc);
return false;
#else
return turing_mma_available(cc) && bank_aligned(nbatch_2);
#endif // GGML_USE_HIP
}
static __device__ constexpr int tile_stride(const int nbatch_2) {
return enabled(nbatch_2) ? nbatch_2 : nbatch_2 + 4;
}
static __host__ int tile_stride(const int nbatch_2, const int cc) {
return enabled(nbatch_2, cc) ? nbatch_2 : nbatch_2 + 4;
}
// Swizzled byte offset for tile element (row, col_h2), same map used for writes and reads.
template<int stride_h2>
static __device__ __forceinline__ int bytes_rc(const int row, const int col_h2) {
static_assert(bank_aligned(stride_h2), "swizzled tile needs a stride that is a multiple of 32");
return ((row * stride_h2 + col_h2) * (int) sizeof(half2)) ^ ((row & 7) << 4);
}
// ldmatrix.x4 via 64-bit generic pointer.
static __device__ __forceinline__ void ldmatrix_x4(int * xi, const half2 * addr) {
#if defined(TURING_MMA_AVAILABLE)
asm volatile("ldmatrix.sync.aligned.m8n8.x4.b16 {%0, %1, %2, %3}, [%4];"
: "=r"(xi[0]), "=r"(xi[1]), "=r"(xi[2]), "=r"(xi[3])
: "l"(addr));
#else
GGML_UNUSED_VARS(xi, addr);
NO_DEVICE_CODE;
#endif // defined(TURING_MMA_AVAILABLE)
}
static __device__ __forceinline__ void ldmatrix_x4_trans(int * xi, const half2 * addr) {
#if defined(TURING_MMA_AVAILABLE)
asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.b16 {%0, %1, %2, %3}, [%4];"
: "=r"(xi[0]), "=r"(xi[2]), "=r"(xi[1]), "=r"(xi[3])
: "l"(addr));
#else
GGML_UNUSED_VARS(xi, addr);
NO_DEVICE_CODE;
#endif // defined(TURING_MMA_AVAILABLE)
}
// Per-lane swizzled address for one tile<16, 8, half2> ldmatrix: 16 rows, 4 half2 columns per lane.
template<int stride_h2>
static __device__ __forceinline__ const half2 * lane_addr(
const half2 * tile_base, const int base_row, const int base_col_h2, const int I, const int J) {
static_assert(bank_aligned(stride_h2), "swizzled tile needs a stride that is a multiple of 32");
const int lane_row = threadIdx.x % I;
const int lane_col = (threadIdx.x / I) * (J / 2);
uint32_t byte_off = (uint32_t) ((base_row + lane_row)*stride_h2 + base_col_h2 + lane_col) * (uint32_t) sizeof(half2);
byte_off ^= (uint32_t) (((base_row + lane_row) & 7) << 4);
return (const half2 *) ((const char *) tile_base + byte_off);
}
template<int stride_h2, bool swz, typename TileT>
static __device__ __forceinline__ void load_ldmatrix(
TileT & t, const half2 * tile_base, const int base_row, const int base_col_h2) {
if constexpr (swz) {
static_assert(std::is_same_v<TileT, ggml_cuda_mma::tile<16, 8, half2>>,
"the swizzled layout is only supported for tile<16, 8, half2>");
ldmatrix_x4((int *) t.x, lane_addr<stride_h2>(tile_base, base_row, base_col_h2, TileT::I, TileT::J));
} else {
ggml_cuda_mma::load_ldmatrix(t, tile_base + base_row*stride_h2 + base_col_h2, stride_h2);
}
}
template<int stride_h2, bool swz, typename TileT>
static __device__ __forceinline__ void load_ldmatrix(TileT & t, const half2 * tile_base, const int off_h2) {
if constexpr (swz) {
load_ldmatrix<stride_h2, swz>(t, tile_base, off_h2 / stride_h2, off_h2 % stride_h2);
} else {
ggml_cuda_mma::load_ldmatrix(t, tile_base + off_h2, stride_h2);
}
}
template<int stride_h2, bool swz, typename TileT>
static __device__ __forceinline__ void load_ldmatrix_trans(
TileT & t, const half2 * tile_base, const int base_row, const int base_col_h2) {
if constexpr (swz) {
static_assert(std::is_same_v<TileT, ggml_cuda_mma::tile<16, 8, half2>>,
"the swizzled layout is only supported for tile<16, 8, half2>");
ldmatrix_x4_trans((int *) t.x, lane_addr<stride_h2>(tile_base, base_row, base_col_h2, TileT::I, TileT::J));
} else {
ggml_cuda_mma::load_ldmatrix_trans(t, tile_base + base_row*stride_h2 + base_col_h2, stride_h2);
}
}
template<int stride_h2, bool swz, typename TileT>
static __device__ __forceinline__ void load_ldmatrix_trans(TileT & t, const half2 * tile_base, const int off_h2) {
if constexpr (swz) {
load_ldmatrix_trans<stride_h2, swz>(t, tile_base, off_h2 / stride_h2, off_h2 % stride_h2);
} else {
ggml_cuda_mma::load_ldmatrix_trans(t, tile_base + off_h2, stride_h2);
}
}
} // namespace ggml_cuda_fattn_smem_swizzle

View file

@ -5145,6 +5145,9 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
if (b->type == GGML_TYPE_F16 && a->type != GGML_TYPE_F16) {
return false;
}
if (op->op == GGML_OP_MUL_MAT_ID && ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
return false;
}
#ifdef GGML_USE_MUSA
const int cc = ggml_cuda_info().devices[dev_ctx->device].cc;
if (b->ne[2]*b->ne[3] > 1 && !ggml_is_transposed(a) && !ggml_is_transposed(b)) {

View file

@ -782,6 +782,20 @@ namespace ggml_cuda_mma {
}
}
// Byte offset of tile element (i, j). If swz, XOR swizzle it to avoid bank conflicts without row padding.
template <bool swz, typename T>
static __device__ __forceinline__ int swizzle_bytes(const int i, const int j, const int stride) {
static_assert(!swz || sizeof(T) == 4, "swizzled tiles need 32 bit elements");
const int off = (i*stride + j) * (int) sizeof(T);
return swz ? off ^ ((i & 7) << 4) : off;
}
template <bool swz, typename T>
static __device__ __forceinline__ const T * swizzle(
const T * __restrict__ tile_base, const int i, const int j, const int stride) {
return (const T *) ((const char *) tile_base + swizzle_bytes<swz, T>(i, j, stride));
}
template <typename T>
static __device__ __forceinline__ void load_ldmatrix(
tile<8, 8, T> & t, const T * __restrict__ xs0, const int stride) {
@ -858,6 +872,29 @@ namespace ggml_cuda_mma {
#endif // TURING_MMA_AVAILABLE
}
// Load from tile element (i0, j0), swz tells if the tile is stored swizzled.
template <bool swz, int I, int J, typename T, data_layout dl>
static __device__ __forceinline__ void load_ldmatrix(
tile<I, J, T, dl> & t, const T * __restrict__ tile_base, const int i0, const int j0, const int stride) {
if constexpr (!swz) {
load_ldmatrix(t, tile_base + i0*stride + j0, stride);
return;
}
#if defined(TURING_MMA_AVAILABLE)
static_assert(I == 16, "bad tile width");
static_assert(J == 8, "bad tile height");
const int i = i0 + threadIdx.x % t.I;
const int j = j0 + (threadIdx.x / t.I) * (t.J / 2);
int * xi = (int *) t.x;
asm volatile("ldmatrix.sync.aligned.m8n8.x4.b16 {%0, %1, %2, %3}, [%4];"
: "=r"(xi[0]), "=r"(xi[1]), "=r"(xi[2]), "=r"(xi[3])
: "l"(swizzle<true>(tile_base, i, j, stride)));
#else
GGML_UNUSED_VARS(t, tile_base, i0, j0, stride);
NO_DEVICE_CODE;
#endif // defined(TURING_MMA_AVAILABLE)
}
static __device__ __forceinline__ void load_ldmatrix(
tile<8, 4, half2, DATA_LAYOUT_I_MAJOR_MIRRORED> & t, const half2 * __restrict__ xs0, const int stride) {
ggml_cuda_memcpy_1<4*sizeof(half2)>(t.x, xs0 + t.get_i(0)*stride);
@ -917,6 +954,29 @@ namespace ggml_cuda_mma {
#endif // TURING_MMA_AVAILABLE
}
// Load from tile element (i0, j0), swz tells if the tile is stored swizzled.
template <bool swz, int I, typename T, data_layout dl>
static __device__ __forceinline__ void load_ldmatrix_trans(
tile<I, 8, T, dl> & t, const T * __restrict__ tile_base, const int i0, const int j0, const int stride) {
if constexpr (!swz) {
load_ldmatrix_trans(t, tile_base + i0*stride + j0, stride);
return;
}
#if defined(TURING_MMA_AVAILABLE)
static_assert(I == 16, "bad tile width");
static_assert(dl == DATA_LAYOUT_I_MAJOR, "bad data layout");
const int i = i0 + threadIdx.x % t.I;
const int j = j0 + (threadIdx.x / t.I) * (t.J / 2);
int * xi = (int *) t.x;
asm volatile("ldmatrix.sync.aligned.m8n8.x4.trans.b16 {%0, %1, %2, %3}, [%4];"
: "=r"(xi[0]), "=r"(xi[2]), "=r"(xi[1]), "=r"(xi[3])
: "l"(swizzle<true>(tile_base, i, j, stride)));
#else
GGML_UNUSED_VARS(t, tile_base, i0, j0, stride);
NO_DEVICE_CODE;
#endif // defined(TURING_MMA_AVAILABLE)
}
static __device__ __forceinline__ void mma(
tile<16, 8, int> & D, const tile<16, 4, int> & A, const tile<8, 4, int> & B) {
#ifdef TURING_MMA_AVAILABLE

View file

@ -1156,14 +1156,18 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm_id(ggml_m
const bool bc_inp = op->src[0]->ne[0] % 32 != 0;
// src1 prec [TAG_GGML_PREC]
const bool amax = ggml_get_op_params_i32(op, 3) == GGML_PREC_F32;
snprintf(base, 256, "kernel_mul_mm_id_%s_%s", ggml_type_name(tsrc0), ggml_type_name(tsrc1));
snprintf(name, 256, "%s_bci=%d", base, bc_inp);
snprintf(name, 256, "%s_bci=%d_amax=%d", base, bc_inp, amax);
ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
if (!res.pipeline) {
ggml_metal_cv_t cv = ggml_metal_cv_init();
ggml_metal_cv_set_bool(cv, bc_inp, FC_MUL_MM + 0);
ggml_metal_cv_set_bool(cv, amax, FC_MUL_MM + 6);
res = ggml_metal_library_compile_pipeline(lib, base, name, cv);

View file

@ -10,10 +10,23 @@
#include <string>
#include <vector>
// derive the non-empty op sequence from the raw `ops_all` sequence
static std::vector<ggml_op> ggml_metal_fusion_filter_ops(const std::vector<ggml_op> & ops_all) {
std::vector<ggml_op> ops;
for (ggml_op op : ops_all) {
if (!ggml_op_is_empty(op)) {
ops.push_back(op);
}
}
return ops;
}
struct ggml_metal_fusion {
ggml_metal_fusion_id id;
std::vector<ggml_op> ops; // op sequence (fixed length, non-empty nodes)
std::vector<ggml_op> ops; // non-empty op sequence, derived from ops_all
std::vector<ggml_op> ops_all; // full raw op sequence (may include empty RESHAPE/VIEW nodes)
std::vector<int> outs; // additional fused output nodes, relative to ops
@ -30,6 +43,25 @@ struct ggml_metal_fusion {
const int * node_idxs,
int idx,
ggml_metal_fusion_mode mode);
ggml_metal_fusion(
ggml_metal_fusion_id id,
const std::vector<ggml_op> & ops_all,
const std::vector<int> & outs,
bool unsafe,
bool (*check)(const struct ggml_metal_fusion * fusion,
const struct ggml_tensor * const * nodes,
const struct ggml_cgraph * gf,
const int * node_idxs,
int idx,
ggml_metal_fusion_mode mode))
: id(id),
ops(ggml_metal_fusion_filter_ops(ops_all)),
ops_all(ops_all),
outs(outs),
unsafe(unsafe),
check(check) {
}
};
ggml_metal_fusion_id ggml_metal_fusion_get_id(const ggml_metal_fusion * fusion) {
@ -297,17 +329,17 @@ static bool ggml_metal_fusion_check_snake(
// SOFT_MAX + ARGSORT + GET_ROWS (plus optional norm/scale) for MoE routing.
// This is a multi-output elision chain: the fused kernel writes both the selected
// expert ids and the gathered/normalized routing weights.
static const std::vector<ggml_op> ops_topk_moe_all = {
static const std::vector<ggml_op> ops_topk_moe = {
GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS
};
static const std::vector<ggml_op> ops_topk_moe_scale_all = {
static const std::vector<ggml_op> ops_topk_moe_scale = {
GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS, GGML_OP_SCALE
};
static const std::vector<ggml_op> ops_topk_moe_norm_all = {
static const std::vector<ggml_op> ops_topk_moe_norm = {
GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS,
GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE
};
static const std::vector<ggml_op> ops_topk_moe_norm_scale_all = {
static const std::vector<ggml_op> ops_topk_moe_norm_scale = {
GGML_OP_SOFT_MAX, GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS,
GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE, GGML_OP_SCALE
};
@ -607,86 +639,63 @@ static const std::vector<ggml_op> ops_snake = { GGML_OP_MUL, GGML_OP_SIN, GGML_O
static const std::vector<ggml_op> ops_gdn_cache = { GGML_OP_GATED_DELTA_NET, GGML_OP_CPY };
static const std::vector<ggml_op> ops_topk_moe = {
GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS
};
static const std::vector<ggml_op> ops_topk_moe_scale = {
GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS, GGML_OP_SCALE
};
static const std::vector<ggml_op> ops_topk_moe_norm = {
GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS,
GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV
};
static const std::vector<ggml_op> ops_topk_moe_norm_scale = {
GGML_OP_SOFT_MAX, GGML_OP_ARGSORT, GGML_OP_GET_ROWS,
GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_SCALE
};
static const std::vector<ggml_op> ops_ssm_conv_silu = { GGML_OP_SSM_CONV, GGML_OP_UNARY };
static const std::vector<ggml_op> ops_moe_reduce_2 = { GGML_OP_MUL, GGML_OP_ADD };
static const std::vector<ggml_op> ops_moe_reduce_3 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD };
static const std::vector<ggml_op> ops_moe_reduce_4 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const std::vector<ggml_op> ops_moe_reduce_5 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const std::vector<ggml_op> ops_moe_reduce_6 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const std::vector<ggml_op> ops_moe_reduce_7 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const std::vector<ggml_op> ops_moe_reduce_8 = { GGML_OP_MUL, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD };
static const std::vector<ggml_op> ops_moe_reduce_all_2 = {
static const std::vector<ggml_op> ops_moe_reduce_2 = {
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_ADD
};
static const std::vector<ggml_op> ops_moe_reduce_all_3 = {
static const std::vector<ggml_op> ops_moe_reduce_3 = {
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_ADD, GGML_OP_ADD
};
static const std::vector<ggml_op> ops_moe_reduce_all_4 = {
static const std::vector<ggml_op> ops_moe_reduce_4 = {
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
};
static const std::vector<ggml_op> ops_moe_reduce_all_5 = {
static const std::vector<ggml_op> ops_moe_reduce_5 = {
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
};
static const std::vector<ggml_op> ops_moe_reduce_all_6 = {
static const std::vector<ggml_op> ops_moe_reduce_6 = {
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
};
static const std::vector<ggml_op> ops_moe_reduce_all_7 = {
static const std::vector<ggml_op> ops_moe_reduce_7 = {
GGML_OP_MUL, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
};
static const std::vector<ggml_op> ops_moe_reduce_all_8 = {
static const std::vector<ggml_op> ops_moe_reduce_8 = {
GGML_OP_MUL,
GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW, GGML_OP_VIEW,
GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD, GGML_OP_ADD
};
static const std::vector<ggml_metal_fusion> ggml_metal_fusions = {
{ GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, ops_norm_mul, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, ops_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_SCALE, ops_norm_scale, ops_norm_scale, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL, ops_rms_norm_mul, ops_rms_norm_mul, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_rms_norm_mul_add, ops_rms_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_SCALE, ops_rms_norm_scale, ops_rms_norm_scale, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_2, ops_add_2, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_3, ops_add_3, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_4, ops_add_4, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_5, ops_add_5, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_6, ops_add_6, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_7, ops_add_7, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_SNAKE, ops_snake, ops_snake, {}, false, ggml_metal_fusion_check_snake },
{ GGML_METAL_FUSION_GDN_CACHE, ops_gdn_cache, ops_gdn_cache, {}, true, ggml_metal_fusion_check_gdn_cache },
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe, ops_topk_moe_all, {1}, true, ggml_metal_fusion_check_topk_moe },
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_scale, ops_topk_moe_scale_all, {1}, true, ggml_metal_fusion_check_topk_moe },
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm, ops_topk_moe_norm_all, {1}, true, ggml_metal_fusion_check_topk_moe },
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm_scale, ops_topk_moe_norm_scale_all, {1}, true, ggml_metal_fusion_check_topk_moe },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_2, ops_moe_reduce_all_2, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_3, ops_moe_reduce_all_3, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_4, ops_moe_reduce_all_4, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_5, ops_moe_reduce_all_5, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_6, ops_moe_reduce_all_6, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_7, ops_moe_reduce_all_7, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_8, ops_moe_reduce_all_8, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_SSM_CONV_SILU, ops_ssm_conv_silu, ops_ssm_conv_silu, {}, false, ggml_metal_fusion_check_ssm_conv_silu },
{ GGML_METAL_FUSION_NORM_MUL, ops_norm_mul, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_SCALE, ops_norm_scale, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL, ops_rms_norm_mul, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_MUL_ADD, ops_rms_norm_mul_add, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_NORM_SCALE, ops_rms_norm_scale, {}, false, ggml_metal_fusion_check_norm },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_2, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_3, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_4, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_5, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_6, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_ADD_CHAIN, ops_add_7, {}, false, ggml_metal_fusion_check_add_chain },
{ GGML_METAL_FUSION_SNAKE, ops_snake, {}, false, ggml_metal_fusion_check_snake },
{ GGML_METAL_FUSION_GDN_CACHE, ops_gdn_cache, {}, true, ggml_metal_fusion_check_gdn_cache },
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe, {1}, true, ggml_metal_fusion_check_topk_moe },
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_scale, {1}, true, ggml_metal_fusion_check_topk_moe },
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm, {1}, true, ggml_metal_fusion_check_topk_moe },
{ GGML_METAL_FUSION_TOPK_MOE, ops_topk_moe_norm_scale, {1}, true, ggml_metal_fusion_check_topk_moe },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_2, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_3, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_4, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_5, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_6, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_7, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_MOE_REDUCE, ops_moe_reduce_8, {}, true, ggml_metal_fusion_check_moe_reduce },
{ GGML_METAL_FUSION_SSM_CONV_SILU, ops_ssm_conv_silu, {}, false, ggml_metal_fusion_check_ssm_conv_silu },
};
// ---- alloc deps -----------------------------------------------------------

View file

@ -2719,9 +2719,12 @@ int ggml_metal_op_mul_mat_id(ggml_metal_op_t ctx, int idx) {
ggml_metal_buffer_id bid_amax = bid_ids;
bid_amax.offs += ggml_metal_op_mul_mat_id_extra_ids(op);
// src1 prec [TAG_GGML_PREC]
const bool use_amax = ggml_get_op_params_i32(op, 3) == GGML_PREC_F32;
// src1 rescale factors, computed before the matmul
// ref: https://github.com/ggml-org/llama.cpp/pull/26223
{
if (use_amax) {
ggml_metal_kargs_mul_mm_id_amax args = {
/*.ne00 =*/ ne10,
/*.ne01 =*/ ne11,
@ -2779,17 +2782,17 @@ int ggml_metal_op_mul_mat_id(ggml_metal_op_t ctx, int idx) {
ggml_metal_op_concurrency_reset(ctx);
{
if (use_amax) {
auto pipeline = ggml_metal_library_get_pipeline_mul_mm_id_amax(lib);
ggml_metal_encoder_set_pipeline(enc, pipeline);
ggml_metal_encoder_set_buffer (enc, bid_amax, 0);
ggml_metal_encoder_dispatch_threadgroups(enc, 1, 1, 1, 32, 1, 1);
}
// the next kernel has to wait for the amax data
ggml_metal_op_concurrency_reset(ctx);
// the next kernel has to wait for the amax data
ggml_metal_op_concurrency_reset(ctx);
}
{
auto pipeline = ggml_metal_library_get_pipeline_mul_mm_id(lib, op);

View file

@ -142,7 +142,7 @@ kernel void kernel_flash_attn_ext_blk(
const int32_t i1 = tgpig[1];
const int32_t i0 = tgpig[0];
char res = i0*C + C > args.ne30 ? 1 : 0;
char res = i0*C + C > args.ne30 || i1*Q + Q > args.ne31 ? 1 : 0;
device const half * mask_src = (device const half *) (mask + (i1*Q)*args.nb31 + i2*args.nb32 + i3*args.nb33) + i0*C + tiisg;

View file

@ -7,6 +7,7 @@ constant short FC_mul_mm_ne12 [[function_constant(FC_MUL_MM + 2)]];
constant short FC_mul_mm_ne13 [[function_constant(FC_MUL_MM + 3)]];
constant short FC_mul_mm_r2 [[function_constant(FC_MUL_MM + 4)]];
constant short FC_mul_mm_r3 [[function_constant(FC_MUL_MM + 5)]];
constant bool FC_mul_mm_id_amax [[function_constant(FC_MUL_MM + 6)]];
// each block_q contains 16*nl weights
#ifdef GGML_METAL_HAS_TENSOR
@ -584,8 +585,8 @@ kernel void kernel_mul_mm_id(
const short lb1 = (short) tiitg/NL1; // 0 .. NR1-1, this thread's row of the B tile
// power-of-two rescaling
const float s1_inv = ((device const float *) amax)[0];
const float s1_scale = ((device const float *) amax)[1];
const float s1_inv = FC_mul_mm_id_amax ? ((device const float *) amax)[0] : 1.0f;
const float s1_scale = FC_mul_mm_id_amax ? ((device const float *) amax)[1] : 1.0f;
#ifndef GGML_METAL_HAS_TENSOR
S0_8x8 ma[4];

View file

@ -4771,80 +4771,51 @@ static void quantize_row_iq1_m_impl(const float * GGML_RESTRICT x, void * GGML_R
// 1: +, -
// 2: -, +
// 3: -, -
for (int i1 = 0; i1 <= block_size; ++i1) {
for (int i2 = i1; i2 <= block_size; ++i2) {
memset(sumqx, 0, 4*sizeof(float));
memset(sumq2, 0, 4*sizeof(float));
for (int j = 0; j < i1; ++j) {
int i = idx[2*j];
if (i < block_size/2) {
sumqx[0] += weight[i]*x_p[0]*xb[i];
sumqx[1] += weight[i]*x_p[0]*xb[i];
sumqx[2] += weight[i]*x_m[0]*xb[i];
sumqx[3] += weight[i]*x_m[0]*xb[i];
sumq2[0] += weight[i]*x_p[0]*x_p[0];
sumq2[1] += weight[i]*x_p[0]*x_p[0];
sumq2[2] += weight[i]*x_m[0]*x_m[0];
sumq2[3] += weight[i]*x_m[0]*x_m[0];
} else {
sumqx[0] += weight[i]*x_p[0]*xb[i];
sumqx[2] += weight[i]*x_p[0]*xb[i];
sumqx[1] += weight[i]*x_m[0]*xb[i];
sumqx[3] += weight[i]*x_m[0]*xb[i];
sumq2[0] += weight[i]*x_p[0]*x_p[0];
sumq2[2] += weight[i]*x_p[0]*x_p[0];
sumq2[1] += weight[i]*x_m[0]*x_m[0];
sumq2[3] += weight[i]*x_m[0]*x_m[0];
// prefix sums are kept per half of the block because each half can use a different sign (x_p or x_m)
// since v[0]-v[1] = v[1]-v[2] = -1 for both x_p and x_m, the 3-group sum for a split collapses to T*v[2] - px[i1] - px[i2]
{
float px[2][IQ1M_BLOCK_SIZE+1];
float pw[2][IQ1M_BLOCK_SIZE+1];
px[0][0] = px[1][0] = 0;
pw[0][0] = pw[1][0] = 0;
for (int j = 0; j < block_size; ++j) {
const int i = idx[2*j];
const int h = i < block_size/2 ? 0 : 1;
px[h][j+1] = px[h][j] + weight[i]*xb[i];
px[1-h][j+1] = px[1-h][j];
pw[h][j+1] = pw[h][j] + weight[i];
pw[1-h][j+1] = pw[1-h][j];
}
const float txs[2] = {px[0][block_size], px[1][block_size]}; // total weight*x per half
const float tws[2] = {pw[0][block_size], pw[1][block_size]}; // total weight per half
const float p2 = x_p[2], m2 = x_m[2];
const float cp1 = x_p[0]*x_p[0] - x_p[1]*x_p[1];
const float cp2 = x_p[1]*x_p[1] - x_p[2]*x_p[2];
const float cm1 = x_m[0]*x_m[0] - x_m[1]*x_m[1];
const float cm2 = x_m[1]*x_m[1] - x_m[2]*x_m[2];
for (int i1 = 0; i1 <= block_size; ++i1) {
for (int i2 = i1; i2 <= block_size; ++i2) {
float qx_p[2], qx_m[2], q2_p[2], q2_m[2];
for (int h = 0; h < 2; ++h) {
const float sx = px[h][i1] + px[h][i2];
qx_p[h] = txs[h]*p2 - sx;
qx_m[h] = txs[h]*m2 - sx;
q2_p[h] = tws[h]*p2*p2 + pw[h][i1]*cp1 + pw[h][i2]*cp2;
q2_m[h] = tws[h]*m2*m2 + pw[h][i1]*cm1 + pw[h][i2]*cm2;
}
}
for (int j = i1; j < i2; ++j) {
int i = idx[2*j];
if (i < block_size/2) {
sumqx[0] += weight[i]*x_p[1]*xb[i];
sumqx[1] += weight[i]*x_p[1]*xb[i];
sumqx[2] += weight[i]*x_m[1]*xb[i];
sumqx[3] += weight[i]*x_m[1]*xb[i];
sumq2[0] += weight[i]*x_p[1]*x_p[1];
sumq2[1] += weight[i]*x_p[1]*x_p[1];
sumq2[2] += weight[i]*x_m[1]*x_m[1];
sumq2[3] += weight[i]*x_m[1]*x_m[1];
} else {
sumqx[0] += weight[i]*x_p[1]*xb[i];
sumqx[2] += weight[i]*x_p[1]*xb[i];
sumqx[1] += weight[i]*x_m[1]*xb[i];
sumqx[3] += weight[i]*x_m[1]*xb[i];
sumq2[0] += weight[i]*x_p[1]*x_p[1];
sumq2[2] += weight[i]*x_p[1]*x_p[1];
sumq2[1] += weight[i]*x_m[1]*x_m[1];
sumq2[3] += weight[i]*x_m[1]*x_m[1];
}
}
for (int j = i2; j < block_size; ++j) {
int i = idx[2*j];
if (i < block_size/2) {
sumqx[0] += weight[i]*x_p[2]*xb[i];
sumqx[1] += weight[i]*x_p[2]*xb[i];
sumqx[2] += weight[i]*x_m[2]*xb[i];
sumqx[3] += weight[i]*x_m[2]*xb[i];
sumq2[0] += weight[i]*x_p[2]*x_p[2];
sumq2[1] += weight[i]*x_p[2]*x_p[2];
sumq2[2] += weight[i]*x_m[2]*x_m[2];
sumq2[3] += weight[i]*x_m[2]*x_m[2];
} else {
sumqx[0] += weight[i]*x_p[2]*xb[i];
sumqx[2] += weight[i]*x_p[2]*xb[i];
sumqx[1] += weight[i]*x_m[2]*xb[i];
sumqx[3] += weight[i]*x_m[2]*xb[i];
sumq2[0] += weight[i]*x_p[2]*x_p[2];
sumq2[2] += weight[i]*x_p[2]*x_p[2];
sumq2[1] += weight[i]*x_m[2]*x_m[2];
sumq2[3] += weight[i]*x_m[2]*x_m[2];
}
}
for (int k = 0; k < 4; ++k) {
if (sumq2[k] > 0 && sumqx[k]*sumqx[k] > best_score*sumq2[k]) {
scale = sumqx[k]/sumq2[k]; best_score = scale*sumqx[k];
besti1 = i1; besti2 = i2; best_k = k;
sumqx[0] = qx_p[0] + qx_p[1];
sumqx[1] = qx_p[0] + qx_m[1];
sumqx[2] = qx_m[0] + qx_p[1];
sumqx[3] = qx_m[0] + qx_m[1];
sumq2[0] = q2_p[0] + q2_p[1];
sumq2[1] = q2_p[0] + q2_m[1];
sumq2[2] = q2_m[0] + q2_p[1];
sumq2[3] = q2_m[0] + q2_m[1];
for (int k = 0; k < 4; ++k) {
if (sumq2[k] > 0 && sumqx[k]*sumqx[k] > best_score*sumq2[k]) {
scale = sumqx[k]/sumq2[k]; best_score = scale*sumqx[k];
besti1 = i1; besti2 = i2; best_k = k;
}
}
}
}

View file

@ -124,6 +124,24 @@ struct vk_flash_attn_push_constants {
static_assert(sizeof(vk_flash_attn_push_constants) <= 128, "sizeof(vk_flash_attn_push_constants) must be <= 128");
struct vk_fa_xe_opt_push_constants {
uint32_t kv_seq_len;
uint32_t activation_length;
uint32_t q_head;
uint32_t kv_head;
uint32_t qk_ratio;
uint32_t qk_sub_groups;
uint32_t flag;
uint32_t nbkv_tok;
uint32_t nbkv_head;
uint32_t batch_stride_q;
uint32_t batch_stride_k;
uint32_t batch_stride_v;
uint32_t batch_stride_m;
uint32_t batch_stride_o;
float softmax_scale;
};
struct vk_op_push_constants {
uint32_t KX;
uint32_t KY;

View file

@ -993,6 +993,7 @@ struct vk_device_struct {
bool fa_sparse_compact_use_subgroups;
vk_pipeline pipeline_flash_attn_split_k_reduce;
std::map<std::tuple<uint32_t, uint32_t, uint32_t, uint32_t>, std::pair<vk_pipeline, vk_pipeline>> pipeline_xe_fa_decode_dual_phases;
vk_pipeline pipeline_count_experts;
// [2] is for whether to take n_experts from spec constant (0) or push constant (1)

View file

@ -2987,6 +2987,46 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_matmul_split_k_reduce, "split_k_reduce", split_k_reduce_len, split_k_reduce_data, "main", 2, 2 * sizeof(uint32_t), {256 * 4, 1, 1}, {}, 1);
ggml_vk_create_pipeline(device, device->pipeline_flash_attn_split_k_reduce, "fa_split_k_reduce", fa_split_k_reduce_len, fa_split_k_reduce_data, "main", 3, sizeof(vk_op_flash_attn_split_k_reduce_push_constants), {1, device->subgroup_size, 1}, {device->subgroup_size}, 1, true);
if (device->vendor_id == VK_VENDOR_ID_INTEL && (device->architecture == INTEL_XE2 || (device->architecture == INTEL_XE1 && device->coopmat_support && device->uma))) {
auto upper_power_of_2 = [&](uint32_t in) {
GGML_ASSERT(in != 0);
if (in <= 1) return 1u;
uint32_t ret = in - 1;
ret |= ret >> 1;
ret |= ret >> 2;
ret |= ret >> 4;
ret |= ret >> 8;
ret |= ret >> 16;
return ret + 1;
};
uint32_t xe_native_sub_group_size = 16;
if (device->architecture == INTEL_XE1) {
xe_native_sub_group_size = 8;
}
for (auto& it : device->pipeline_xe_fa_decode_dual_phases) {
const uint32_t split_p_chunk = 32;
auto HdQk = it.first;
auto& pipelines = it.second;
uint32_t head_dim_qk = std::get<0>(HdQk);
uint32_t head_dim_pv = std::get<1>(HdQk);
uint32_t gqa_ratio = std::get<2>(HdQk);
uint32_t q_len = std::get<3>(HdQk);
const uint32_t out_dim_per_wg = gqa_ratio > 16 ? 8 : 16;
uint32_t aligned_q_len = upper_power_of_2(q_len);
uint32_t group_sz_ph1 = std::min(std::max(aligned_q_len * xe_native_sub_group_size, 64u), 256u);
uint32_t out_per_wg_ph1 = std::min(q_len, 256u / xe_native_sub_group_size);
uint32_t aligned_gqa_ratio = upper_power_of_2(gqa_ratio);
uint32_t split_p_per_iter_ph2 = 256;
uint32_t split_p_per_warp = 16;
uint32_t group_sz_ph2 = (split_p_per_iter_ph2 / split_p_per_warp) * xe_native_sub_group_size;
uint32_t out_per_wg_ph2 = std::min(std::max(16u / aligned_gqa_ratio, 1u), q_len);
ggml_vk_create_pipeline(device, pipelines.first, "xe_fa_decode_ph1", fa_decode_ph1_cm1_len, fa_decode_ph1_cm1_data, "main", 5, sizeof(vk_fa_xe_opt_push_constants), { 1, 32, 1 }, { group_sz_ph1, gqa_ratio, head_dim_qk, xe_native_sub_group_size, split_p_chunk, out_per_wg_ph1 }, 1, false, true, xe_native_sub_group_size);
ggml_vk_create_pipeline(device, pipelines.second, "xe_fa_decode_ph2", fa_decode_ph2_cm1_len, fa_decode_ph2_cm1_data, "main", 5, sizeof(vk_fa_xe_opt_push_constants), { 1, 1, 1 }, { group_sz_ph2, gqa_ratio, head_dim_pv, out_per_wg_ph2, xe_native_sub_group_size, split_p_per_iter_ph2, split_p_chunk, out_dim_per_wg }, 1, false, true, xe_native_sub_group_size);
}
}
for (auto &it : device->pipeline_fa_mask_opt) {
auto BrBc = it.first;
ggml_vk_create_pipeline(device, it.second, "fa_mask_opt", fa_mask_opt_len, fa_mask_opt_data, "main", 2, sizeof(vk_op_flash_attn_mask_opt_push_constants), {1, 1, 1}, {128, 128 / device->subgroup_size, BrBc.first, BrBc.second}, 1, true, true, device->subgroup_size);
@ -7931,6 +7971,18 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
vk_pipeline pipeline = nullptr;
bool xe_fa_opt = false;
bool fa_copy_qstate = false;
bool xe_fa_supported_platform =
(ctx->device.get()->architecture == INTEL_XE2 && ctx->device.get()->properties.deviceID != 0xFD80 && ctx->device.get()->properties.deviceID != 0xFD81) ||
(ctx->device.get()->architecture == INTEL_XE1 && ctx->device.get()->coopmat_support && ctx->device.get()->uma);
bool xe_fa_supported_usage = neq0 % 32 == 0 && nev0 % 16 == 0 && q->nb[1] > q->nb[2] && k->nb[1] > k->nb[2] && v->nb[1] > v->nb[2] && mask != nullptr;
bool xe_fa_supported_dtype = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_F16 && v->type == GGML_TYPE_F16 && (mask != nullptr && mask->type == GGML_TYPE_F16);
std::pair<vk_pipeline, vk_pipeline> xe_fa_pipeline_dual_phases = { nullptr , nullptr };
vk_pipeline xe_fa_pipeline = nullptr;
size_t size_p = 0;
size_t size_group_max = 0;
{
std::lock_guard<std::mutex> guard(ctx->device->compile_mutex);
auto &pipelines = ctx->device->pipeline_flash_attn_f32_f16;
@ -7988,6 +8040,37 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
// of "align", so recompute split_k based on that.
split_kv = ROUNDUP_POW2(std::max(1u, KV / split_k), alignment);
split_k = CEIL_DIV(KV, split_kv);
xe_fa_opt = xe_fa_supported_platform && xe_fa_supported_usage && xe_fa_supported_dtype;
if (xe_fa_opt) {
std::lock_guard<std::mutex> guard(ctx->device->compile_mutex);
const uint32_t split_p_size = 32;
const size_t max_dim = (nek1 + split_p_size - 1) / split_p_size;
const size_t p_dim = max_dim * split_p_size;
auto& pipelines = ctx->device->pipeline_xe_fa_decode_dual_phases;
auto it = pipelines.find({ (uint32_t)neq0, (uint32_t)nev0, qk_ratio, (uint32_t)neq1 });
if (it != pipelines.end()) {
xe_fa_pipeline_dual_phases = it->second;
} else {
pipelines[{(uint32_t)neq0, (uint32_t)nev0, qk_ratio, (uint32_t)neq1}] = xe_fa_pipeline_dual_phases = std::make_pair(std::make_shared<vk_pipeline_struct>(), std::make_shared<vk_pipeline_struct>());
}
size_p = neq1 * neq2 * p_dim * neq3 * sizeof(ggml_fp16_t);
size_group_max = neq1 * neq2 * max_dim * neq3 * sizeof(float);
size_t temp_size = ggml_nelements(q) * sizeof(ggml_fp16_t) + size_p + size_group_max;
fa_copy_qstate = true;
if (ctx->prealloc_size_x < temp_size) {
ctx->prealloc_size_x = temp_size;
ggml_vk_preallocate_buffers(ctx, subctx);
}
if (ctx->prealloc_x_need_sync) {
ggml_vk_sync_buffers(ctx, subctx);
}
}
}
if (xe_fa_opt == true) {
use_mask_opt = false;
}
// Reserve space for split_k temporaries. For each split x batch, we need to store the O matrix (D x ne1)
@ -8143,7 +8226,68 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
mask_n_head_log2, m0, m1,
gqa_ratio, split_kv, split_k };
if (split_k > 1) {
if (xe_fa_opt && split_k > 1) {
auto upper_power_of_2 = [&](uint32_t in) {
GGML_ASSERT(in != 0);
if (in <= 1) return 1u;
uint32_t ret = in - 1;
ret |= ret >> 1;
ret |= ret >> 2;
ret |= ret >> 4;
ret |= ret >> 8;
ret |= ret >> 16;
return ret + 1;
};
auto to_fp16_vk_0 = ggml_vk_get_to_fp16(ctx, q->type);
const uint32_t out_dim_per_wg = qk_ratio > 16 ? 8 : 16;
size_t x_ne = ggml_nelements(q);
size_t temp_buf_offset = 0;
uint32_t head_stride_k = uint32_t(nbk2 / ggml_type_size(k->type));
uint32_t head_stride_v = uint32_t(nbv2 / ggml_type_size(v->type));
uint32_t batch_stride_q = uint32_t(nbq3 / ggml_type_size(q->type));
uint32_t batch_stride_k = uint32_t(nbk3 / ggml_type_size(k->type));
uint32_t batch_stride_v = uint32_t(nbv3 / ggml_type_size(v->type));
uint32_t batch_stride_m = mask ? uint32_t(mask->nb[3] / ggml_type_size(mask->type)) : 0u;
uint32_t batch_stride_o = uint32_t(nb3 / ggml_type_size(dst->type));
vk_fa_xe_opt_push_constants pc_ph1 = { (uint32_t)nek1, (uint32_t)neq1, (uint32_t)neq2, (uint32_t)nek2, qk_ratio, 1, (sinks != nullptr) ? 1u : 0u, (uint32_t)k_stride, head_stride_k,
batch_stride_q, batch_stride_k, batch_stride_v, batch_stride_m, batch_stride_o, scale };
vk_fa_xe_opt_push_constants pc_ph2 = pc_ph1;
pc_ph2.nbkv_tok = v_stride;
pc_ph2.nbkv_head = head_stride_v;
vk_subbuffer q_temp_buf = fa_copy_qstate ? ggml_vk_subbuffer(ctx, ctx->prealloc_x, temp_buf_offset) : q_buf;
temp_buf_offset += fa_copy_qstate ? x_ne * sizeof(ggml_fp16_t) : 0;
vk_subbuffer p_temp_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_x, temp_buf_offset);
temp_buf_offset += size_p;
vk_subbuffer max_temp_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_x, temp_buf_offset);
temp_buf_offset += size_group_max;
uint32_t xe_native_sub_group_size = ctx->device.get()->architecture == INTEL_XE1 ? 8 : 16;
uint32_t aligned_gqa_ratio = upper_power_of_2(qk_ratio);
uint32_t out_per_wg_ph1 = std::min(256u / xe_native_sub_group_size, (uint32_t)neq1);
uint32_t out_per_wg_ph2 = std::min(std::max(16u / aligned_gqa_ratio, 1u), (uint32_t)neq1);
uint32_t ph1_wg = ((neq1 + out_per_wg_ph1 - 1) / out_per_wg_ph1) * nek2;
uint32_t ph2_wg = ((neq1 + out_per_wg_ph2 - 1) / out_per_wg_ph2) * ne0 / out_dim_per_wg;
if (fa_copy_qstate) {
const std::vector<uint32_t> pc_cpy_fp16 =
{ (uint32_t)q->ne[0], (uint32_t)q->ne[1], (uint32_t)q->ne[2], (uint32_t)q->ne[3], (uint32_t)(x_ne) };
ggml_vk_sync_buffers(ctx, subctx);
ggml_pipeline_request_descriptor_sets(ctx, to_fp16_vk_0, 1);
ggml_vk_dispatch_pipeline(ctx, subctx, to_fp16_vk_0, { q_buf, q_temp_buf }, pc_cpy_fp16, { (uint32_t)(x_ne), 1, 1 });
}
ggml_vk_sync_buffers(ctx, subctx);
ggml_pipeline_request_descriptor_sets(ctx, xe_fa_pipeline_dual_phases.first, 1);
ggml_vk_dispatch_pipeline(ctx, subctx, xe_fa_pipeline_dual_phases.first,
{ q_temp_buf, k_buf, mask_buf, p_temp_buf, max_temp_buf },
pc_ph1, { (uint32_t)ph1_wg, (uint32_t)nek1, (uint32_t)neq3 });
ggml_vk_sync_buffers(ctx, subctx);
ggml_pipeline_request_descriptor_sets(ctx, xe_fa_pipeline_dual_phases.second, 1);
ggml_vk_dispatch_pipeline(ctx, subctx, xe_fa_pipeline_dual_phases.second,
{ p_temp_buf, v_buf, max_temp_buf, sinks_buf, dst_buf },
pc_ph2, { (uint32_t)ph2_wg, (uint32_t)nev2, (uint32_t)neq3 });
ctx->prealloc_x_need_sync = true;
} else if (split_k > 1) {
ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_flash_attn_split_k_reduce, 1);
if (ctx->prealloc_split_k_need_sync) {
@ -14985,6 +15129,9 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
// If there's not enough shared memory for row_ids and the result tile, fallback to CPU
return false;
}
if (ggml_get_op_params_i32(op, 3) == GGML_PREC_F32) {
return false;
}
}
switch (src0_type) {
case GGML_TYPE_F32:

View file

@ -0,0 +1,263 @@
#version 450
#extension GL_EXT_control_flow_attributes : enable
#extension GL_EXT_shader_16bit_storage : require
#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require
#extension GL_KHR_memory_scope_semantics : enable
#extension GL_KHR_shader_subgroup_basic : enable
#extension GL_KHR_shader_subgroup_ballot : enable
#extension GL_KHR_shader_subgroup_arithmetic : enable
#extension GL_KHR_cooperative_matrix : enable
#extension GL_EXT_shared_memory_block : enable
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
layout (binding = 0) readonly buffer Q {float16_t qState[];};
layout (binding = 1) readonly buffer K_VEC4 {f16vec4 kStateVec4[];};
layout (binding = 2) buffer MASK_F16 {float16_t mState_f16[];};
layout (binding = 3) buffer P_FP16 {float16_t matP_f16[];};
layout (binding = 4) buffer OUT_MAX {float out_max_f32[];};
layout (push_constant) uniform parameter
{
uint kvSeqLen;
uint activationLength;
uint qHead;
uint kvHead;
uint qkRatio;
uint qkSubGroups;
uint flag;
uint kvStride1;
uint kvStride2;
uint batchStrideQ;
uint batchStrideK;
uint batchStrideV;
uint batchStrideM;
uint batchStrideO;
float softMaxScale;
} p;
layout (constant_id = 0) const uint GROUPSIZE = 128;
layout (constant_id = 1) const uint GQA_RATIO = 8;
layout (constant_id = 2) const uint HEAD_DIM = 128;
layout (constant_id = 3) const uint WARPSIZE = 16;
layout (constant_id = 4) const uint MATP_REDUCE = 32;
layout (constant_id = 5) const uint N_TOK = 1;
layout (constant_id = 6) const uint COOP_MAT_P_PER_LOOP = 4;
#define MAX_HEADS 8
#define TN WARPSIZE
#define TM 8
#define TK 16
#define SUBGROUP_COUNT (GROUPSIZE / WARPSIZE)
#define MATP_PER_LOOP (COOP_MAT_P_PER_LOOP * TM)
#define P_LOOP_COUNT (MATP_REDUCE / MATP_PER_LOOP)
#define COOP_MAT_Q_PER_TOKEN ((GQA_RATIO + TN - 1) / TN)
#define COOP_MAT_P_M COOP_MAT_Q_PER_TOKEN
#define COOP_MAT_P_N (MATP_REDUCE / TM)
#define SLM_PV_SIZE (MATP_REDUCE * COOP_MAT_P_M * TN)
#define SLM_MASK_SIZE (N_TOK * MATP_REDUCE)
#define SLM_POOL_SIZE_K (MATP_PER_LOOP * HEAD_DIM)
#define K_LOAD_PER_LOOP (GROUPSIZE * 4)
#define HEAD_DIM_VEC4 (HEAD_DIM / 4)
#define SLM_CHUNK_SIZE (TK / 4)
#define K_LOAD_LOOPS ((SLM_POOL_SIZE_K + K_LOAD_PER_LOOP - 1) / K_LOAD_PER_LOOP)
#define O_COUNT ((GQA_RATIO + SUBGROUP_COUNT - 1) / SUBGROUP_COUNT)
shared slm_pool_block {
float slm_pool_pv[SLM_PV_SIZE + SLM_MASK_SIZE];
} slm_pool_f32;
shared slm_pool_alias_block {
float16_t slm_pool_k[SLM_POOL_SIZE_K];
} slm_pool_f16;
void main() {
const uint lane = gl_SubgroupInvocationID;
const uint kHeadIdx = gl_WorkGroupID.x % p.kvHead;
const uint outGroupIdx = gl_WorkGroupID.x / p.kvHead;
const uint v = gl_WorkGroupID.y;
const uint d = gl_WorkGroupID.z;
const uint localLinearId = gl_SubgroupID;
const uint wgLane = localLinearId * WARPSIZE + lane;
const uint qDim = p.qHead * HEAD_DIM;
const uint kvDim = p.kvStride1;
const uint maskDim = p.kvSeqLen;
const uint maxDim = (p.kvSeqLen + MATP_REDUCE - 1) / MATP_REDUCE;
const uint pDim = maxDim * MATP_REDUCE;
const uint tokFlatIdx = localLinearId + outGroupIdx * N_TOK;
uint offsetBaseQ = min(tokFlatIdx, p.activationLength - 1) * qDim;
offsetBaseQ = offsetBaseQ + d * p.batchStrideQ + kHeadIdx * HEAD_DIM * GQA_RATIO;
const uint offsetBaseK = (d * p.batchStrideK + (v * MATP_REDUCE) * kvDim + kHeadIdx * p.kvStride2) / 4;
uint offsetOut = d * p.qHead * p.activationLength * pDim + v * MATP_REDUCE + kHeadIdx * GQA_RATIO * pDim + (localLinearId * O_COUNT + outGroupIdx * N_TOK * p.qHead) * pDim + lane;
uint offsetMax = d * p.qHead * p.activationLength * maxDim + v + kHeadIdx * GQA_RATIO * maxDim + (localLinearId * O_COUNT + outGroupIdx * N_TOK * p.qHead) * maxDim;
const uint offsetSlmLoadPv = (localLinearId * O_COUNT * MATP_REDUCE + lane);
const uint offsetBaseM = v * MATP_REDUCE + lane;
const float fp32Min = uintBitsToFloat(0xFEFFFFFF);
const uint loopCount = HEAD_DIM / TK;
float maskFp32[MATP_REDUCE / WARPSIZE];
if (tokFlatIdx < p.activationLength) {
[[unroll]] for (uint mk = 0; mk < MATP_REDUCE / WARPSIZE; mk++) {
const uint maskOffset = mk * WARPSIZE + offsetBaseM;
if (maskOffset < maskDim) {
maskFp32[mk] = float(mState_f16[d * p.batchStrideM + tokFlatIdx * maskDim + maskOffset]);
} else {
maskFp32[mk] = fp32Min;
}
}
}
coopmat<float, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> matP[COOP_MAT_P_M][COOP_MAT_P_N];
[[unroll]] for (uint mp = 0; mp < COOP_MAT_P_M; mp++) {
[[unroll]] for (uint np = 0; np < COOP_MAT_P_N; np++) {
matP[mp][np] = coopmat<float, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(0.0f);
}
}
[[unroll]] for (uint kLoad = 0; kLoad < K_LOAD_LOOPS; kLoad++) {
const uint flatOffset = kLoad * GROUPSIZE + wgLane;
const uint kRowIdx = flatOffset / HEAD_DIM_VEC4;
const uint kColIdx = flatOffset % HEAD_DIM_VEC4;
const uint slmChunkCol = kColIdx % SLM_CHUNK_SIZE;
const uint slmChunkRow = kColIdx / SLM_CHUNK_SIZE;
const uint offsetK = offsetBaseK + kRowIdx * kvDim / 4 + kColIdx;
const uint offsetSlmK = kRowIdx * TK + slmChunkRow * TK * MATP_PER_LOOP + slmChunkCol * 4;
slm_pool_f16.slm_pool_k[offsetSlmK + 0] = kStateVec4[offsetK].x;
slm_pool_f16.slm_pool_k[offsetSlmK + 1] = kStateVec4[offsetK].y;
slm_pool_f16.slm_pool_k[offsetSlmK + 2] = kStateVec4[offsetK].z;
slm_pool_f16.slm_pool_k[offsetSlmK + 3] = kStateVec4[offsetK].w;
}
[[unroll]] for (uint pLoop = 0; pLoop < P_LOOP_COUNT; pLoop++) {
f16vec4 kTemp[K_LOAD_LOOPS];
if (pLoop + 1 < P_LOOP_COUNT) {
[[unroll]] for (uint kLoad = 0; kLoad < K_LOAD_LOOPS; kLoad++) {
const uint flatOffset = kLoad * GROUPSIZE + wgLane;
const uint kRowIdx = flatOffset / HEAD_DIM_VEC4 + (pLoop + 1) * MATP_PER_LOOP;
const uint kColIdx = flatOffset % HEAD_DIM_VEC4;
const uint offsetK = offsetBaseK + kRowIdx * kvDim / 4 + kColIdx;
kTemp[kLoad] = kStateVec4[offsetK];
}
}
barrier();
if (localLinearId < N_TOK) {
[[unroll]] for (uint loop = 0; loop < loopCount; loop++) {
coopmat<float16_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> matQ[COOP_MAT_P_M];
coopmat<float16_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> matK[COOP_MAT_P_PER_LOOP];
[[unroll]] for (uint mq = 0; mq < COOP_MAT_P_M; mq++) {
coopMatLoad(
matQ[mq],
qState,
offsetBaseQ + mq * TN * HEAD_DIM + loop * TK,
HEAD_DIM,
gl_CooperativeMatrixLayoutColumnMajor);
}
[[unroll]] for (uint np = 0; np < COOP_MAT_P_PER_LOOP; np++) {
coopMatLoad(
matK[np],
slm_pool_f16.slm_pool_k,
loop * TK * MATP_PER_LOOP + np * TM * TK,
TK,
gl_CooperativeMatrixLayoutRowMajor);
}
[[unroll]] for (uint mp = 0; mp < COOP_MAT_P_M; mp++) {
[[unroll]] for (uint np = 0; np < COOP_MAT_P_PER_LOOP; np++) {
matP[mp][pLoop * COOP_MAT_P_PER_LOOP + np] = coopMatMulAdd(matK[np], matQ[mp], matP[mp][pLoop * COOP_MAT_P_PER_LOOP + np]);
}
}
}
}
barrier();
if (pLoop + 1 < P_LOOP_COUNT) {
[[unroll]] for (uint kLoad = 0; kLoad < K_LOAD_LOOPS; kLoad++) {
const uint flatOffset = kLoad * GROUPSIZE + wgLane;
const uint kRowIdx = flatOffset / HEAD_DIM_VEC4;
const uint kColIdx = flatOffset % HEAD_DIM_VEC4;
const uint slmChunkCol = kColIdx % SLM_CHUNK_SIZE;
const uint slmChunkRow = kColIdx / SLM_CHUNK_SIZE;
const uint offsetSlmK = kRowIdx * TK + slmChunkRow * TK * MATP_PER_LOOP + slmChunkCol * 4;
slm_pool_f16.slm_pool_k[offsetSlmK + 0] = kTemp[kLoad].x;
slm_pool_f16.slm_pool_k[offsetSlmK + 1] = kTemp[kLoad].y;
slm_pool_f16.slm_pool_k[offsetSlmK + 2] = kTemp[kLoad].z;
slm_pool_f16.slm_pool_k[offsetSlmK + 3] = kTemp[kLoad].w;
}
}
}
barrier();
if (tokFlatIdx < p.activationLength) {
[[unroll]] for (uint mk = 0; mk < MATP_REDUCE / WARPSIZE; mk++) {
slm_pool_f32.slm_pool_pv[SLM_PV_SIZE + localLinearId * MATP_REDUCE + mk * WARPSIZE + lane] = maskFp32[mk];
}
}
[[unroll]] for (uint oLoop = 0; oLoop < N_TOK; oLoop++) {
if (oLoop + outGroupIdx * N_TOK < p.activationLength) {
if (localLinearId == oLoop) {
[[unroll]] for (uint mp = 0; mp < COOP_MAT_P_M; mp++) {
[[unroll]] for (uint np = 0; np < COOP_MAT_P_N; np++) {
coopMatStore(matP[mp][np], slm_pool_f32.slm_pool_pv, mp * MATP_REDUCE * TN + np * TM, MATP_REDUCE, gl_CooperativeMatrixLayoutColumnMajor);
}
}
}
barrier();
[[unroll]] for (uint maskIdx = 0; maskIdx < MATP_REDUCE / WARPSIZE; maskIdx++) {
maskFp32[maskIdx] = slm_pool_f32.slm_pool_pv[SLM_PV_SIZE + oLoop * MATP_REDUCE + maskIdx * WARPSIZE + lane];
}
float fp32O[O_COUNT][MATP_REDUCE / WARPSIZE];
float maxOut[O_COUNT];
[[unroll]] for (uint oc = 0; oc < O_COUNT; oc++) {
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
fp32O[oc][os] = slm_pool_f32.slm_pool_pv[offsetSlmLoadPv + os * WARPSIZE + oc * MATP_REDUCE] * p.softMaxScale;
}
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
fp32O[oc][os] = fp32O[oc][os] + maskFp32[os];
}
float maxTemp = fp32Min;
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
maxTemp = max(maxTemp, fp32O[oc][os]);
}
maxOut[oc] = subgroupMax(maxTemp);
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
fp32O[oc][os] = exp(fp32O[oc][os] - maxOut[oc]);
}
}
[[unroll]] for (uint oc = 0; oc < O_COUNT; oc++) {
if (localLinearId * O_COUNT + oc < GQA_RATIO) {
[[unroll]] for (uint os = 0; os < MATP_REDUCE / WARPSIZE; os++) {
matP_f16[offsetOut + oc * pDim + os * WARPSIZE] = float16_t(fp32O[oc][os]);
}
if (lane == 0) {
out_max_f32[offsetMax + oc * maxDim] = maxOut[oc];
}
}
}
offsetOut = offsetOut + p.qHead * pDim;
offsetMax = offsetMax + p.qHead * maxDim;
barrier();
}
}
}

View file

@ -0,0 +1,408 @@
#version 450
#extension GL_EXT_control_flow_attributes : enable
#extension GL_EXT_shader_16bit_storage : require
#extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require
#extension GL_KHR_memory_scope_semantics : enable
#extension GL_KHR_shader_subgroup_basic : enable
#extension GL_KHR_shader_subgroup_ballot : enable
#extension GL_KHR_shader_subgroup_arithmetic : enable
#extension GL_KHR_cooperative_matrix : enable
#extension GL_EXT_shared_memory_block : enable
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
layout (binding = 0) readonly buffer P {f16vec4 pStateVec4[];};
layout (binding = 1) readonly buffer V {float16_t vState[];};
layout (binding = 1) readonly buffer V_VEC4 {f16vec4 vStateVec4[];};
layout (binding = 2) buffer MAX_FP32 {float max_f32[];};
layout (binding = 3) buffer SINK_FP32 {float sink_f32[];};
layout (binding = 4) buffer OUT_FP32 {float out_f32[];};
layout (binding = 4) buffer OUT_VEC4 {vec4 out_f32_vec4[];};
layout (binding = 4) buffer OUT_F16 {float16_t out_f16[];};
layout (push_constant) uniform parameter
{
uint kvSeqLen;
uint activationLength;
uint qHead;
uint kvHead;
uint qkRatio;
uint qkSubGroups;
uint flag;
uint kvStride1;
uint kvStride2;
uint batchStrideQ;
uint batchStrideK;
uint batchStrideV;
uint batchStrideM;
uint batchStrideO;
float softMaxScale;
} p;
layout (constant_id = 0) const uint GROUPSIZE = 256;
layout (constant_id = 1) const uint GQA_RATIO = 8;
layout (constant_id = 2) const uint HEAD_DIM = 128;
layout (constant_id = 3) const uint N_TOKS_PER_GROUP = 1;
layout (constant_id = 4) const uint WARPSIZE = 16;
layout (constant_id = 5) const uint MATP_PER_LOOP = 64;
layout (constant_id = 6) const uint MATP_REDUCE = 32;
layout (constant_id = 7) const uint WARP_V_DIM = 16;
#define TN WARPSIZE
#define TM 8
#define TK 16
#define MAT_O_N (WARP_V_DIM / TM)
#define MAT_P_M (GQA_RATIO * N_TOKS_PER_GROUP)
#define ALIGNED_P_M ((MAT_P_M + WARPSIZE - 1) / WARPSIZE)
#define V_HEAD_GROUPS (HEAD_DIM / WARP_V_DIM)
#define SUBGROUP_COUNT (GROUPSIZE / WARPSIZE)
#define SPLIT_P_GROUPS (MATP_PER_LOOP / TK)
#define SLM_POOL_SIZE_O (SUBGROUP_COUNT * ALIGNED_P_M * TN * MAT_O_N * TM)
#define P_LOAD_PER_LOOP (GROUPSIZE * 4)
#define P_LOAD_LOOPS ((MAT_P_M * MATP_PER_LOOP + P_LOAD_PER_LOOP - 1) / P_LOAD_PER_LOOP)
#define SLM_POOL_SIZE_P (P_LOAD_LOOPS * P_LOAD_PER_LOOP)
#define SIZE_LOCAL_MAX (MAT_P_M * MATP_PER_LOOP / MATP_REDUCE)
#define MAX_LOAD_LOOPS ((SIZE_LOCAL_MAX + GROUPSIZE - 1) / GROUPSIZE)
#define SLM_POOL_SIZE_LOCAL_MAX (MAX_LOAD_LOOPS * GROUPSIZE)
#define MAX_REDUCE_COUNT ((MAT_P_M + SUBGROUP_COUNT - 1) / SUBGROUP_COUNT)
#define GLOBAL_MAX_SIZE (MAX_REDUCE_COUNT * SUBGROUP_COUNT)
#define SLM_POOL_SIZE_SOFTMAX_SUM (SUBGROUP_COUNT * P_LOAD_LOOPS)
#define SLM_OFFSET_P (GLOBAL_MAX_SIZE * 2 + SLM_POOL_SIZE_SOFTMAX_SUM * 2 + SLM_POOL_SIZE_LOCAL_MAX * 2 * 2)
#define SLM_OFFSET_GLOBAL_MAX 0
#define SLM_OFFSET_SOFTMAX_SUM (SLM_OFFSET_GLOBAL_MAX + GLOBAL_MAX_SIZE)
#define SLM_OFFSET_O (SLM_OFFSET_SOFTMAX_SUM + SLM_POOL_SIZE_SOFTMAX_SUM)
#define SLM_OFFSET_LOCAL_MAX (GLOBAL_MAX_SIZE + SLM_POOL_SIZE_SOFTMAX_SUM)
#define P_REDUCE_VEC4 (MATP_PER_LOOP / 4)
#define MAX_PER_LOOP (MATP_PER_LOOP / MATP_REDUCE)
#define SLM_MAX_STRIDE (MATP_REDUCE / 4)
#define SUB_GROUPS_PER_LINE (MATP_PER_LOOP / WARPSIZE / 4)
shared slm_pool_block {
float slm_pool_o[GLOBAL_MAX_SIZE + SLM_POOL_SIZE_SOFTMAX_SUM + SLM_POOL_SIZE_O];
} slm_pool_f32;
shared slm_pool_alias_block {
float16_t slm_pool_pv[GLOBAL_MAX_SIZE * 2 + SLM_POOL_SIZE_SOFTMAX_SUM * 2 + SLM_POOL_SIZE_LOCAL_MAX * 2 * 2 + SLM_POOL_SIZE_P * 2];
} slm_pool_alias_f16;
void main() {
const uint lane = gl_SubgroupInvocationID;
const uint v = gl_WorkGroupID.y;
const uint d = gl_WorkGroupID.z;
const uint vWarpIdx = gl_WorkGroupID.x % V_HEAD_GROUPS;
const uint outTokIdx = gl_WorkGroupID.x / V_HEAD_GROUPS;
const uint localLinearId = gl_SubgroupID;
const uint wgLane = localLinearId * WARPSIZE + lane;
const uint splitIdx = localLinearId;
const uint maxDim = (p.kvSeqLen + MATP_REDUCE - 1) / MATP_REDUCE;
const uint pDim = maxDim * MATP_REDUCE;
const uint kvDim = p.kvStride1;
const uint oDim = p.qHead * HEAD_DIM;
const uint offsetBaseP = (d * p.activationLength * p.qHead + v * GQA_RATIO + outTokIdx * N_TOKS_PER_GROUP * p.qHead) * pDim / 4;
const uint offsetBaseMax = (d * p.activationLength * p.qHead + v * GQA_RATIO + outTokIdx * N_TOKS_PER_GROUP * p.qHead) * maxDim;
const uint offsetBaseV = (d * p.batchStrideV + v * p.kvStride2 + vWarpIdx * WARP_V_DIM + splitIdx * TK * kvDim);
const uint offsetSlmP = (SLM_OFFSET_P + wgLane * 4);
const float fp32Min = uintBitsToFloat(0xFEFFFFFF);
const float fp32Max = uintBitsToFloat(0x7EFFFFFF);
uint offsetV = offsetBaseV;
coopmat<float, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> sums[ALIGNED_P_M][MAT_O_N];
f16vec4 pStateTemp[P_LOAD_LOOPS];
float fp32CompensationP[P_LOAD_LOOPS];
uint loadRowBase[P_LOAD_LOOPS];
uint loadColBase[P_LOAD_LOOPS];
float fp32SoftMaxSum[P_LOAD_LOOPS];
float fp32GlobalMaxP[P_LOAD_LOOPS];
uint maxRowBase[MAX_LOAD_LOOPS];
uint maxColBase[MAX_LOAD_LOOPS];
uint outOffsets[ALIGNED_P_M];
bool outputMask[ALIGNED_P_M];
float fp32SinkCoeff[ALIGNED_P_M];
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
const uint flatOffset = pm * WARPSIZE + lane;
const uint inGroupTokIdx = flatOffset / GQA_RATIO;
const uint inGroupHeadIdx = flatOffset % GQA_RATIO;
outputMask[pm] = (N_TOKS_PER_GROUP * outTokIdx + inGroupTokIdx < p.activationLength) && (inGroupHeadIdx < GQA_RATIO) && (inGroupTokIdx < N_TOKS_PER_GROUP);
outOffsets[pm] = (inGroupTokIdx * oDim + inGroupHeadIdx * HEAD_DIM) / 4;
if ((0x1 & p.flag) != 0) {
fp32SinkCoeff[pm] = sink_f32[inGroupHeadIdx + v * GQA_RATIO];
}
}
[[unroll]] for (uint maxCount = 0; maxCount < MAX_REDUCE_COUNT; maxCount++) {
const uint flatIdx = maxCount * SUBGROUP_COUNT + localLinearId;
const uint rowIdx = flatIdx % GQA_RATIO;
const uint tokIdx = flatIdx / GQA_RATIO;
if (tokIdx < N_TOKS_PER_GROUP) {
float fp32MaxReduce = fp32Min;
const uint maxOffset = offsetBaseMax + (tokIdx * p.qHead + rowIdx) * maxDim;
[[unroll]] for (uint maxReduce = 0; maxReduce < (maxDim + WARPSIZE - 1) / WARPSIZE; maxReduce++) {
if (maxReduce * WARPSIZE + lane < maxDim) {
fp32MaxReduce = max(fp32MaxReduce, max_f32[maxOffset + maxReduce * WARPSIZE + lane]);
}
}
fp32MaxReduce = subgroupMax(fp32MaxReduce);
if (lane == 0) {
slm_pool_f32.slm_pool_o[SLM_OFFSET_GLOBAL_MAX + maxCount * SUBGROUP_COUNT + localLinearId] = fp32MaxReduce;
}
} else {
if (lane == 0) {
slm_pool_f32.slm_pool_o[SLM_OFFSET_GLOBAL_MAX + maxCount * SUBGROUP_COUNT + localLinearId] = fp32Max;
}
}
}
barrier();
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
const uint flatOffset = (pLoad * GROUPSIZE + wgLane) / P_REDUCE_VEC4;
const uint rowIdxFlat = flatOffset % GQA_RATIO;
const uint tokenIdxFlat = min(flatOffset / GQA_RATIO, N_TOKS_PER_GROUP - 1);
loadColBase[pLoad] = (pLoad * GROUPSIZE + wgLane) % P_REDUCE_VEC4;
loadRowBase[pLoad] = (tokenIdxFlat * p.qHead + rowIdxFlat);
fp32SoftMaxSum[pLoad] = 0.0f;
fp32GlobalMaxP[pLoad] = slm_pool_f32.slm_pool_o[SLM_OFFSET_GLOBAL_MAX + flatOffset];
}
[[unroll]] for (uint maxLoad = 0; maxLoad < MAX_LOAD_LOOPS; maxLoad++) {
const uint flatOffset = (maxLoad * GROUPSIZE + wgLane) / MAX_PER_LOOP;
const uint rowIdxFlat = flatOffset % GQA_RATIO;
const uint tokenIdxFlat = min(flatOffset / GQA_RATIO, N_TOKS_PER_GROUP - 1);
maxColBase[maxLoad] = (maxLoad * GROUPSIZE + wgLane) % MAX_PER_LOOP;
maxRowBase[maxLoad] = (tokenIdxFlat * p.qHead + rowIdxFlat);
}
[[unroll]] for (uint maxLoad = 0; maxLoad < MAX_LOAD_LOOPS; maxLoad++) {
const uint flatMaxOffset = maxRowBase[maxLoad] * maxDim + maxColBase[maxLoad];
slm_pool_f32.slm_pool_o[SLM_OFFSET_LOCAL_MAX + maxLoad * GROUPSIZE + wgLane] = max_f32[offsetBaseMax + flatMaxOffset];
maxColBase[maxLoad] = maxColBase[maxLoad] + MATP_PER_LOOP / MATP_REDUCE;
}
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
const uint flatOffset = loadRowBase[pLoad] * pDim / 4 + loadColBase[pLoad];
pStateTemp[pLoad] = pStateVec4[offsetBaseP + flatOffset];
}
[[unroll]] for (uint n = 0; n < ALIGNED_P_M; n++) {
[[unroll]] for (uint i = 0; i < MAT_O_N; i++) {
sums[n][i] = coopmat<float, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(0.0f);
}
}
barrier();
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
const uint maxOffset = (pLoad * GROUPSIZE + wgLane) / SLM_MAX_STRIDE;
if (loadColBase[pLoad] < pDim / 4) {
fp32CompensationP[pLoad] = slm_pool_f32.slm_pool_o[SLM_OFFSET_LOCAL_MAX + maxOffset];
float pTemp[4] = float[4](pStateTemp[pLoad].x, pStateTemp[pLoad].y, pStateTemp[pLoad].z, pStateTemp[pLoad].w);
float compTemp = exp(fp32CompensationP[pLoad] - fp32GlobalMaxP[pLoad]);
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
pTemp[kk] = pTemp[kk] * compTemp;
fp32SoftMaxSum[pLoad] = fp32SoftMaxSum[pLoad] + pTemp[kk];
slm_pool_alias_f16.slm_pool_pv[offsetSlmP + pLoad * GROUPSIZE * 4 + kk] = float16_t(pTemp[kk]);
}
} else {
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
slm_pool_alias_f16.slm_pool_pv[offsetSlmP + pLoad * GROUPSIZE * 4 + kk] = float16_t(0.0f);
}
}
loadColBase[pLoad] = loadColBase[pLoad] + P_REDUCE_VEC4;
}
const uint loopCount = (p.kvSeqLen + MATP_PER_LOOP - 1) / MATP_PER_LOOP;
for (uint loop = 0; loop < loopCount; loop++) {
const uint slmPingPongLoad = (loop & 0x1);
const uint slmPingPongStore = ((loop + 1) & 0x1);
if (loop + 1 < loopCount) {
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
const uint flatOffset = loadRowBase[pLoad] * pDim / 4 + loadColBase[pLoad];
pStateTemp[pLoad] = pStateVec4[offsetBaseP + flatOffset];
}
[[unroll]] for (uint maxLoad = 0; maxLoad < MAX_LOAD_LOOPS; maxLoad++) {
const uint flatMaxOffset = maxRowBase[maxLoad] * maxDim + maxColBase[maxLoad];
slm_pool_f32.slm_pool_o[SLM_OFFSET_LOCAL_MAX + slmPingPongStore * SLM_POOL_SIZE_LOCAL_MAX + maxLoad * GROUPSIZE + wgLane] = max_f32[offsetBaseMax + flatMaxOffset];
maxColBase[maxLoad] = maxColBase[maxLoad] + MATP_PER_LOOP / MATP_REDUCE;
}
}
barrier();
{
const uint coopMatOffsetP = SLM_OFFSET_P + slmPingPongLoad * SLM_POOL_SIZE_P + splitIdx * TK;
coopmat<float16_t, gl_ScopeSubgroup, TM, TK, gl_MatrixUseA> matV[MAT_O_N];
[[unroll]] for (uint cc = 0; cc < MAT_O_N; cc++) {
coopMatLoad(
matV[cc],
vState,
offsetV + TM * cc,
kvDim,
gl_CooperativeMatrixLayoutColumnMajor);
}
[[unroll]] for (uint mo = 0; mo < ALIGNED_P_M; mo++) {
coopmat<float16_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> matP;
coopMatLoad(
matP,
slm_pool_alias_f16.slm_pool_pv,
coopMatOffsetP + mo * TN * MATP_PER_LOOP,
MATP_PER_LOOP,
gl_CooperativeMatrixLayoutColumnMajor);
[[unroll]] for (uint no = 0; no < MAT_O_N; no++) {
sums[mo][no] = coopMatMulAdd(matV[no], matP, sums[mo][no]);
}
}
}
offsetV += MATP_PER_LOOP * kvDim;
if (loop * MATP_PER_LOOP + splitIdx * TK >= p.kvSeqLen) {
offsetV = 0;
}
if (loop + 1 < loopCount) {
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
const uint maxOffset = (pLoad * GROUPSIZE + wgLane) / SLM_MAX_STRIDE;
if (loadColBase[pLoad] < pDim / 4) {
fp32CompensationP[pLoad] = slm_pool_f32.slm_pool_o[SLM_OFFSET_LOCAL_MAX + slmPingPongStore * SLM_POOL_SIZE_LOCAL_MAX + maxOffset];
float pTemp[4] = float[4](pStateTemp[pLoad].x, pStateTemp[pLoad].y, pStateTemp[pLoad].z, pStateTemp[pLoad].w);
float compTemp = exp(fp32CompensationP[pLoad] - fp32GlobalMaxP[pLoad]);
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
pTemp[kk] = pTemp[kk] * compTemp;
fp32SoftMaxSum[pLoad] = fp32SoftMaxSum[pLoad] + pTemp[kk];
slm_pool_alias_f16.slm_pool_pv[offsetSlmP + slmPingPongStore * SLM_POOL_SIZE_P + pLoad * GROUPSIZE * 4 + kk] = float16_t(pTemp[kk]);
}
} else {
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
slm_pool_alias_f16.slm_pool_pv[offsetSlmP + slmPingPongStore * SLM_POOL_SIZE_P + pLoad * GROUPSIZE * 4 + kk] = float16_t(0.0f);
}
}
loadColBase[pLoad] = loadColBase[pLoad] + P_REDUCE_VEC4;
}
}
}
barrier();
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
fp32SoftMaxSum[pLoad] = subgroupAdd(fp32SoftMaxSum[pLoad]);
}
[[unroll]] for (uint mo = 0; mo < ALIGNED_P_M; mo++) {
[[unroll]] for (uint no = 0; no < MAT_O_N; no++) {
coopMatStore(
sums[mo][no],
slm_pool_f32.slm_pool_o,
SLM_OFFSET_O + mo * TN * WARP_V_DIM + TM * no + localLinearId * ALIGNED_P_M * TN * WARP_V_DIM,
WARP_V_DIM,
gl_CooperativeMatrixLayoutColumnMajor);
}
}
[[unroll]] for (uint pLoad = 0; pLoad < P_LOAD_LOOPS; pLoad++) {
slm_pool_f32.slm_pool_o[SLM_OFFSET_SOFTMAX_SUM + pLoad * SUBGROUP_COUNT + localLinearId] = fp32SoftMaxSum[pLoad];
}
barrier();
if (localLinearId == 1) {
const uint sumBase = SLM_OFFSET_SOFTMAX_SUM + lane * SUB_GROUPS_PER_LINE;
float sumTemp[ALIGNED_P_M][SUB_GROUPS_PER_LINE];
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
[[unroll]] for (uint reduce = 0; reduce < SUB_GROUPS_PER_LINE; reduce++) {
sumTemp[pm][reduce] = slm_pool_f32.slm_pool_o[sumBase + pm * WARPSIZE * SUB_GROUPS_PER_LINE + reduce];
}
}
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
[[unroll]] for (uint reduce = 1; reduce < SUB_GROUPS_PER_LINE; reduce++) {
sumTemp[pm][0] = sumTemp[pm][0] + sumTemp[pm][reduce];
}
}
if ((0x1 & p.flag) != 0) {
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
float fp32GlobalMax = slm_pool_f32.slm_pool_o[SLM_OFFSET_GLOBAL_MAX + pm * WARPSIZE + lane];
float sinkCompensation = fp32GlobalMax - fp32SinkCoeff[pm];
sinkCompensation = exp(sinkCompensation);
float softmaxSumTemp = sumTemp[pm][0] * sinkCompensation;
sumTemp[pm][0] = sumTemp[pm][0] + 1.0f / sinkCompensation;
sumTemp[pm][0] = 1.0f / sumTemp[pm][0];
sinkCompensation = sinkCompensation / (1.0f + softmaxSumTemp);
sumTemp[pm][0] = fp32GlobalMax < fp32SinkCoeff[pm] ? sinkCompensation : sumTemp[pm][0];
slm_pool_f32.slm_pool_o[SLM_OFFSET_SOFTMAX_SUM + pm * WARPSIZE + lane] = sumTemp[pm][0];
}
} else {
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
slm_pool_f32.slm_pool_o[SLM_OFFSET_SOFTMAX_SUM + pm * WARPSIZE + lane] = 1.0f / sumTemp[pm][0];
}
}
}
[[unroll]] for (uint reduce = 2; reduce < SPLIT_P_GROUPS; reduce = reduce << 1 ) {
const uint stride = (reduce >> 1) * ALIGNED_P_M * TN * MAT_O_N * TM;
if ((localLinearId % reduce) == 0) {
const uint reduceBase = localLinearId * ALIGNED_P_M * TN * MAT_O_N * TM + SLM_OFFSET_O;
float sumTemp0[4];
float sumTemp1[4];
const uint reduceVec4Count = ALIGNED_P_M * TN * MAT_O_N * TM / 4 / WARPSIZE;
[[unroll]] for (uint totalLoads = 0; totalLoads < reduceVec4Count; totalLoads++) {
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
sumTemp0[kk] = slm_pool_f32.slm_pool_o[reduceBase + totalLoads * 4 * WARPSIZE + 4 * lane + kk];
sumTemp1[kk] = slm_pool_f32.slm_pool_o[reduceBase + stride + totalLoads * 4 * WARPSIZE + 4 * lane + kk];
}
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
sumTemp0[kk] = sumTemp0[kk] + sumTemp1[kk];
}
[[unroll]] for (uint kk = 0; kk < 4; kk++) {
slm_pool_f32.slm_pool_o[reduceBase + totalLoads * 4 * WARPSIZE + 4 * lane + kk] = sumTemp0[kk];
}
}
}
barrier();
}
if (localLinearId == 0) {
const uint slmBase0 = SLM_OFFSET_O + lane * WARP_V_DIM;
const uint slmBase1 = slmBase0 + SPLIT_P_GROUPS / 2 * ALIGNED_P_M * TN * MAT_O_N * TM;
const uint offsetOutBase = (d * p.batchStrideO + vWarpIdx * WARP_V_DIM + v * GQA_RATIO * HEAD_DIM + outTokIdx * oDim * N_TOKS_PER_GROUP) / 4;
float fp32SoftMaxMul[ALIGNED_P_M];
float fp32Output[ALIGNED_P_M][4];
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
fp32SoftMaxMul[pm] = slm_pool_f32.slm_pool_o[SLM_OFFSET_SOFTMAX_SUM + pm * WARPSIZE + lane];
}
[[unroll]] for (uint vg = 0; vg < WARP_V_DIM / 4; vg++) {
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
[[unroll]] for (uint vc = 0; vc < 4; vc++) {
fp32Output[pm][vc] = slm_pool_f32.slm_pool_o[slmBase0 + pm * WARPSIZE * WARP_V_DIM + vg * 4 + vc] * fp32SoftMaxMul[pm];
fp32Output[pm][vc] = fp32Output[pm][vc] + slm_pool_f32.slm_pool_o[slmBase1 + pm * WARPSIZE * WARP_V_DIM + vg * 4 + vc] * fp32SoftMaxMul[pm];
}
}
[[unroll]] for (uint pm = 0; pm < ALIGNED_P_M; pm++) {
if (outputMask[pm] == true) {
out_f32_vec4[offsetOutBase + outOffsets[pm] + vg] = vec4(fp32Output[pm][0], fp32Output[pm][1], fp32Output[pm][2], fp32Output[pm][3]);
}
}
}
}
}

View file

@ -948,6 +948,10 @@ void process_shaders() {
string_to_spv("fa_split_k_reduce", "flash_attn_split_k_reduce.comp", {});
string_to_spv("fa_mask_opt", "flash_attn_mask_opt.comp", {});
string_to_spv("fa_decode_ph1", "flash_attn_decode_phase_1.comp", {}, true, true, false, false);
string_to_spv("fa_decode_ph2", "flash_attn_decode_phase_2.comp", {}, true, true, false, false);
string_to_spv("fa_sparse_compact", "flash_attn_sparse_compact.comp", {});
string_to_spv("fa_sparse_compact_subgroup", "flash_attn_sparse_compact.comp", {{"USE_SUBGROUPS", "1"}});

View file

@ -3913,8 +3913,8 @@ struct ggml_tensor * ggml_permute(
struct ggml_tensor * result = ggml_view_tensor(ctx, a);
ggml_format_name(result, "%s (permuted)", a->name);
int ne[GGML_MAX_DIMS];
int nb[GGML_MAX_DIMS];
int64_t ne[GGML_MAX_DIMS];
size_t nb[GGML_MAX_DIMS];
ne[axis0] = a->ne[0];
ne[axis1] = a->ne[1];

View file

@ -19,6 +19,7 @@
#include <limits>
#include <stdexcept>
#include <string>
#include <unordered_map>
//
// llama_context
@ -587,6 +588,40 @@ void llama_context::resolve_fused_ops(const llama_memory_context_i * mctx, uint3
}
}
static int llama_graph_n_input_tensors(ggml_cgraph * gf) {
std::unordered_map<const ggml_tensor *, std::vector<ggml_tensor *>> users;
for (int i = 0; i < ggml_graph_n_nodes(gf); ++i) {
ggml_tensor * node = ggml_graph_node(gf, i);
if (node->flags & GGML_TENSOR_FLAG_INPUT) {
users[node].push_back(node);
}
for (int j = 0; j < GGML_MAX_SRC; ++j) {
ggml_tensor * src = node->src[j];
if (!src) {
break;
}
if (src->flags & GGML_TENSOR_FLAG_INPUT) {
users[src].push_back(node);
}
}
}
for (const auto & [tensor, nodes] : users) {
if (tensor->op != GGML_OP_NONE) {
LLAMA_LOG_WARN("%s: input tensor '%32s' has op %s, expected GGML_OP_NONE\n",
__func__, tensor->name, ggml_op_name(tensor->op));
}
for (const ggml_tensor * node : nodes) {
LLAMA_LOG_DEBUG("%s: input tensor '%32s' [%s, ne = { %5" PRId64 ", %5" PRId64 ", %5" PRId64 ", %5" PRId64 " }] is used by node '%s' (%s)\n",
__func__, tensor->name, ggml_type_name(tensor->type),
tensor->ne[0], tensor->ne[1], tensor->ne[2], tensor->ne[3],
node->name, ggml_op_name(node->op));
}
}
return (int) users.size();
}
void llama_context::sched_reserve() {
if (!sched_need_reserve) {
return;
@ -633,11 +668,15 @@ void llama_context::sched_reserve() {
resolve_fused_ops(mctx.get(), n_seqs);
// reserve worst-case graph
int n_splits_pp = -1;
int n_nodes_pp = -1;
int n_splits_pp = -1;
int n_nodes_pp = -1;
int n_inputs_pp = -1;
int n_input_tensors_pp = -1;
int n_splits_tg = -1;
int n_nodes_tg = -1;
int n_splits_tg = -1;
int n_nodes_tg = -1;
int n_inputs_tg = -1;
int n_input_tensors_tg = -1;
const uint32_t n_outputs_pp = std::min(n_tokens, cparams.n_outputs_max);
@ -657,8 +696,10 @@ void llama_context::sched_reserve() {
}
}
n_splits_pp = ggml_backend_sched_get_n_splits(sched.get());
n_nodes_pp = ggml_graph_n_nodes(gf);
n_splits_pp = ggml_backend_sched_get_n_splits(sched.get());
n_nodes_pp = ggml_graph_n_nodes(gf);
n_inputs_pp = get_gf_res_reserve()->inputs.size();
n_input_tensors_pp = this->n_input_tensors;
}
// reserve with tg (token generation) graph to get the number of splits and nodes
@ -668,8 +709,10 @@ void llama_context::sched_reserve() {
throw std::runtime_error("failed to allocate compute tg buffers");
}
n_splits_tg = ggml_backend_sched_get_n_splits(sched.get());
n_nodes_tg = ggml_graph_n_nodes(gf);
n_splits_tg = ggml_backend_sched_get_n_splits(sched.get());
n_nodes_tg = ggml_graph_n_nodes(gf);
n_inputs_tg = get_gf_res_reserve()->inputs.size();
n_input_tensors_tg = this->n_input_tensors;
}
// reserve again with pp graph to avoid ggml-alloc reallocations during inference
@ -707,16 +750,21 @@ void llama_context::sched_reserve() {
}
}
if (n_nodes_pp == n_nodes_tg) {
LLAMA_LOG_INFO("%s: graph nodes = %d\n", __func__, n_nodes_pp);
} else {
LLAMA_LOG_INFO("%s: graph nodes = %d (with bs=%d), %d (with bs=1)\n", __func__, n_nodes_pp, n_tokens, n_nodes_tg);
}
{
const bool diff = n_nodes_pp != n_nodes_tg || n_splits_pp != n_splits_tg ||
n_inputs_pp != n_inputs_tg || n_input_tensors_pp != n_input_tensors_tg;
if (n_splits_pp == n_splits_tg) {
LLAMA_LOG_INFO("%s: graph splits = %d\n", __func__, n_splits_pp);
} else {
LLAMA_LOG_INFO("%s: graph splits = %d (with bs=%d), %d (with bs=1)\n", __func__, n_splits_pp, n_tokens, n_splits_tg);
const auto val = [diff](int v_pp, int v_tg) -> std::string {
return diff ? format("%d / %d", v_pp, v_tg) : format("%d", v_pp);
};
LLAMA_LOG_INFO("%s: graph%s: nodes = %s, splits = %s, input objects = %s, input tensors = %s\n",
__func__,
diff ? format(" (pp bs=%d, tg bs=%d)", n_tokens, n_seqs).c_str() : "",
val(n_nodes_pp, n_nodes_tg).c_str(),
val(n_splits_pp, n_splits_tg).c_str(),
val(n_inputs_pp, n_inputs_tg).c_str(),
val(n_input_tensors_pp, n_input_tensors_tg).c_str());
}
const int64_t t_end_us = ggml_time_us();
@ -2485,6 +2533,7 @@ ggml_cgraph * llama_context::graph_reserve(
auto * gf = model.build_graph(gparams);
this->n_input_tensors = llama_graph_n_input_tensors(gf);
this->n_outputs = save_n_outputs;
// initialize scheduler with the specified graph

View file

@ -333,6 +333,7 @@ private:
// reuse the batch_allocr to avoid unnecessary memory allocations
std::unique_ptr<llama_batch_allocr> balloc;
uint32_t n_input_tensors = 0; // number of tensors marked as input during the last graph reserve
uint32_t n_outputs = 0; // number of actually-used outputs in the current ubatch or last logical batch
std::vector<int32_t> output_ids; // map batch token positions to ids of the logits and embd buffers

View file

@ -2303,6 +2303,10 @@ ggml_tensor * llm_graph_context::build_moe_ffn(
}
experts = build_lora_mm_id(down_exps, cur, selected_experts, down_exps_s); // [n_embd, n_expert_used, n_tokens]
if (arch == LLM_ARCH_MISTRAL4) {
// src1 can exceed F16 range
ggml_prec_set_src(experts, GGML_PREC_F32, 1);
}
cb(experts, "ffn_moe_down", il);
if (down_exps_s) {
@ -2452,6 +2456,7 @@ ggml_tensor * llm_graph_context::build_inp_pos() const {
cur = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, (int64_t)n_tokens*hparams.n_pos_per_embd());
ggml_set_input(cur);
cb(cur, "inp_pos", -1);
res->add_input(std::move(inp));
@ -2466,7 +2471,7 @@ ggml_tensor * llm_graph_context::build_inp_attn_scale() const {
// this need to be 1x1xN for broadcasting
cur = ggml_new_tensor_3d(ctx0, GGML_TYPE_F32, 1, 1, n_tokens);
ggml_set_input(cur);
ggml_set_name(cur, "attn_scale");
cb(cur, "inp_attn_scale", -1);
res->add_input(std::move(inp));
@ -2488,6 +2493,7 @@ ggml_tensor * llm_graph_context::build_inp_out_ids() const {
cur = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_outputs);
ggml_set_input(cur);
ggml_set_name(cur, "out_ids");
res->add_input(std::move(inp));
@ -2501,6 +2507,7 @@ ggml_tensor * llm_graph_context::build_inp_mean() const {
cur = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_tokens, ubatch.n_seqs_unq);
ggml_set_input(cur);
ggml_set_name(cur, "mean");
res->add_input(std::move(inp));
@ -2514,6 +2521,7 @@ ggml_tensor * llm_graph_context::build_inp_cls() const {
cur = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ubatch.n_seqs_unq);
ggml_set_input(cur);
ggml_set_name(cur, "cls");
res->add_input(std::move(inp));
@ -2538,6 +2546,7 @@ ggml_tensor * llm_graph_context::build_inp_cross_embd() const {
cur = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_embd, n_enc);
ggml_set_input(cur);
ggml_set_name(cur, "cross_embd");
res->add_input(std::move(inp));
@ -2551,6 +2560,7 @@ ggml_tensor * llm_graph_context::build_inp_pos_bucket_enc() const {
cur = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_tokens, n_tokens);
ggml_set_input(cur);
ggml_set_name(cur, "pos_bucket_enc");
res->add_input(std::move(inp));
@ -2568,6 +2578,7 @@ ggml_tensor * llm_graph_context::build_inp_pos_bucket_dec() const {
cur = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, n_tokens);
ggml_set_input(cur);
ggml_set_name(cur, "pos_bucket_dec");
res->add_input(std::move(inp));
@ -2736,6 +2747,7 @@ llm_graph_input_attn_no_cache * llm_graph_context::build_attn_inp_no_cache() con
// note: there is no KV cache, so the number of KV values is equal to the number of tokens in the batch
inp->self_kq_mask = ggml_new_tensor_4d(ctx0, type_mask, n_tokens, n_tokens, 1, 1);
ggml_set_input(inp->self_kq_mask);
cb(inp->self_kq_mask, "self_kq_mask", -1);
inp->self_kq_mask_cnv = inp->self_kq_mask;
@ -3512,6 +3524,7 @@ static std::unique_ptr<llm_graph_input_rs> build_rs_inp_impl(
inp->s_copy = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_rs);
ggml_set_input(inp->s_copy);
ggml_set_name(inp->s_copy, "rs_s_copy");
inp->s_copy_main = ggml_view_1d(ctx0, inp->s_copy, n_seqs, 0);
inp->s_copy_extra = ggml_view_1d(ctx0, inp->s_copy, n_rs - n_seqs, n_seqs * inp->s_copy->nb[0]);

View file

@ -1417,6 +1417,7 @@ ggml_tensor * llama_kv_cache::build_input_k_idxs(ggml_context * ctx, const llama
ggml_tensor * k_idxs = ggml_new_tensor_1d(ctx, GGML_TYPE_I64, n_tokens);
ggml_set_input(k_idxs);
ggml_set_name(k_idxs, "attn_inp_k_idxs");
return k_idxs;
}
@ -1433,6 +1434,7 @@ ggml_tensor * llama_kv_cache::build_input_v_idxs(ggml_context * ctx, const llama
}
ggml_set_input(v_idxs);
ggml_set_name(v_idxs, "attn_inp_v_idxs");
return v_idxs;
}

View file

@ -83,6 +83,8 @@ llama_model_hunyuan_vl::graph::graph(const llama_model & model, const llm_graph_
ggml_tensor * inp_out_ids = build_inp_out_ids();
for (int il = 0; il < n_layer; ++il) {
res->t_layer_inp[il] = inpL;
ggml_tensor * inpSA = inpL;
// norm

View file

@ -23,7 +23,12 @@ void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) {
ml.get_key(LLM_KV_ATTENTION_INDEXER_TOP_K, hparams.indexer_top_k);
ml.get_key(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE, hparams.indexer_block_size);
ml.get_key(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, hparams.indexer_local_blocks);
msa_p = { (int) hparams.indexer_block_size, (int) hparams.indexer_top_k, (int) hparams.indexer_local_blocks };
msa_p = {
/*.blk =*/ (int) hparams.indexer_block_size,
/*.topk_blocks =*/ (int) hparams.indexer_top_k,
/*.local =*/ (int) hparams.indexer_local_blocks,
};
GGML_ASSERT(hparams.indexer_block_size > 0); // avoid div by zero

View file

@ -1490,9 +1490,14 @@ struct clip_model_loader {
}
// Load the vision/audio feature layer indices if they are explicitly provided
// NOTE: gguf conversions should standardize the values of the vision feature layer to
// be non-negative, since we use -1 to mark values as unset here.
// NOTE: gguf conversions should standardize the values of the vision feature layer to be non-negative, since we use -1 to mark values as unset here.
get_arr_int(string_format(KEY_FEATURE_LAYERS, prefix), hparams.feature_layers, false);
for (const auto & v : hparams.feature_layers) {
if (v > (int) hparams.n_layer) {
throw std::runtime_error(string_format("%s: feature layer index %d is out of range (n_layer: %d)",
__func__, v, hparams.n_layer));
}
}
// model-specific params
switch (model.proj_type) {
@ -1538,7 +1543,12 @@ struct clip_model_loader {
std::vector<int> wa_layer_indexes_vec;
get_arr_int(KEY_WIN_ATTN_LAYER_INDEXES, wa_layer_indexes_vec, false);
if (!wa_layer_indexes_vec.empty()) {
hparams.insert_layer_id = wa_layer_indexes_vec[0];
const int insert_lid = wa_layer_indexes_vec[0];
if (insert_lid < 0 || insert_lid >= (int) hparams.n_layer) {
throw std::runtime_error(string_format("%s: layer index %d is out of range (n_layer: %d)",
__func__, insert_lid, hparams.n_layer));
}
hparams.insert_layer_id = insert_lid;
}
} break;
case PROJECTOR_TYPE_INTERNVL:
@ -3318,6 +3328,7 @@ struct clip_model_loader {
model.pos_embed = get_tensor(string_format(TN_SAM_POS_EMBD, "weight"));
model.patch_embed_proj_w = get_tensor(string_format(TN_SAM_PATCH_EMBD, "weight"));
model.patch_embed_proj_b = get_tensor(string_format(TN_SAM_PATCH_EMBD, "bias"));
model.n_sam_layers = hparams.sam_n_layer;
model.sam_layers.resize(model.n_sam_layers);
for (int il = 0; il < model.n_sam_layers; ++il) {
auto & layer = model.sam_layers[il];
@ -4744,8 +4755,10 @@ bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params) {
// -> https://huggingface.co/HuggingFaceM4/siglip-so400m-14-980-flash-attn2-navit
// -> https://huggingface.co/HuggingFaceM4/siglip-so400m-14-980-flash-attn2-navit/blob/d66538faeba44480d0bfaa42145eef26f9423199/modeling_siglip.py#L316
std::vector<int32_t> positions(pos_h * pos_w);
int bucket_coords_h[1024];
int bucket_coords_w[1024];
// note: sized by the actual patch counts; a tall/wide image produces more
// than 1024 patches per side and a fixed [1024] array would be overrun
std::vector<int> bucket_coords_h(pos_h);
std::vector<int> bucket_coords_w(pos_w);
for (int i = 0; i < pos_h; i++){
bucket_coords_h[i] = std::floor(70.0*i/pos_h);
}
@ -4788,8 +4801,8 @@ bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params) {
// SigLIP position buckets (same as resampler path)
std::vector<int32_t> positions(pos_h * pos_w);
int bucket_coords_h[1024];
int bucket_coords_w[1024];
std::vector<int> bucket_coords_h(pos_h);
std::vector<int> bucket_coords_w(pos_w);
for (int i = 0; i < pos_h; i++){
bucket_coords_h[i] = std::floor(70.0*i/pos_h);
}

View file

@ -294,8 +294,8 @@ private:
support = filter_support * filterscale; // Widen filter when downsampling
ksize = static_cast<int>(std::ceil(support)) * 2 + 1; // Total pixels in kernel
std::vector<double> pre_weights(outSize * ksize); // Temporary weights
bounds.resize(outSize * 2);
std::vector<double> pre_weights((size_t) outSize * ksize); // Temporary weights
bounds.resize((size_t) outSize * 2);
// For each output pixel, compute its filter coefficients
@ -322,20 +322,20 @@ private:
for (x = 0; x < xmax; x++) {
// Distance from input pixel center to output pixel center in input space
double w = resample_filter((x + xmin - center + 0.5) * ss);
pre_weights[xx * ksize + x] = w;
pre_weights[(size_t) xx * ksize + x] = w;
ww += w; // Accumulate for normalization
}
// Normalize weights to sum to 1.0 (preserves brightness)
for (x = 0; x < xmax; x++) {
if (ww != 0.0) {
pre_weights[xx * ksize + x] /= ww;
pre_weights[(size_t) xx * ksize + x] /= ww;
}
}
// Zero-pad remaining kernel positions
for (; x < ksize; x++) {
pre_weights[xx * ksize + x] = 0;
pre_weights[(size_t) xx * ksize + x] = 0;
}
// Store input pixel range for this output pixel
@ -345,11 +345,11 @@ private:
// Convert floating-point coefficients to fixed-point integers
// Formula: int32 = round(float * 2^PRECISION_BITS)
weights.resize(outSize * ksize);
weights.resize((size_t) outSize * ksize);
const double fxp_scale = std::ldexp(1.0, PRECISION_BITS); // 1.0 * 2^PRECISION_BITS
for (int i = 0; i < outSize * ksize; i++) {
for (size_t i = 0; i < (size_t) outSize * ksize; i++) {
// Pillow adds +/- 0.5 then truncates toward zero; std::round would round twice
const double rounded = pre_weights[i] * fxp_scale + (pre_weights[i] < 0 ? -0.5 : 0.5);
weights[i] = static_cast<int32_t>(rounded);
@ -442,6 +442,12 @@ private:
const int src_width = img.get_size().width;
const int src_height = img.get_size().height;
// sanity check on the target size
if (target_width <= 0 || target_width > 65536 || target_height <= 0 || target_height > 65536) {
throw std::runtime_error("resize target " + std::to_string(target_width) + "x" +
std::to_string(target_height) + " is out of range (max 65536)");
}
bool need_horizontal = (target_width != src_width);
bool need_vertical = (target_height != src_height);

View file

@ -18,14 +18,37 @@
class server_http_context::Impl {
public:
std::unique_ptr<httplib::Server> srv;
std::vector<std::unique_ptr<httplib::Server>> servers;
std::vector<std::string> hosts;
std::vector<std::thread> threads; // one thread per listener
std::unique_ptr<httplib::ThreadPool> pool; // single pool shared among all listeners
int n_threads_http = 0;
};
class server_http_task_queue : public httplib::TaskQueue {
httplib::ThreadPool & pool;
public:
explicit server_http_task_queue(httplib::ThreadPool & pool) : pool(pool) {}
bool enqueue(std::function<void()> fn) override { return pool.enqueue(std::move(fn)); }
// note: must call join() to drain the pool
void shutdown() override { /* no-op */ }
};
server_http_context::server_http_context()
: pimpl(std::make_unique<Impl>())
{}
server_http_context::~server_http_context() = default;
server_http_context::~server_http_context() {
// just in case any exit paths that forget to call join()
try {
stop();
join();
} catch (const std::exception & e) {
SRV_ERR("failed to stop HTTP server: %s\n", e.what());
} catch (...) {
SRV_ERR("%s", "failed to stop HTTP server\n");
}
}
static void log_server_request(const httplib::Request & req, const httplib::Response & res) {
// skip logging requests that are regularly sent, to avoid log spam
@ -90,7 +113,6 @@ bool server_http_context::init(const common_params & params) {
path_prefix = params.api_prefix;
port = params.port;
hostname = params.hostname;
if (gcp.enabled) {
SRV_TRC("Google Cloud Platform compat: health route = %s, predict route = %s, port = %d\n", gcp.path_health.c_str(), gcp.path_predict.c_str(), gcp.port);
@ -102,7 +124,39 @@ bool server_http_context::init(const common_params & params) {
port = gcp.port;
}
auto & srv = pimpl->srv;
pimpl->hosts = params.hostnames;
size_t n_tcp_hosts = 0;
for (const auto & host : pimpl->hosts) {
if (!string_ends_with(host, ".sock")) {
n_tcp_hosts++;
}
}
if (port == 0 && n_tcp_hosts > 1) {
SRV_ERR("%s", "--port 0 is not supported with multiple TCP addresses\n");
return false;
}
for (size_t i = 0; i < pimpl->hosts.size(); ++i) {
pimpl->servers.emplace_back();
if (!init_listener(params)) {
return false;
}
// with multiple TCP addresses, [::] must not also claim 0.0.0.0
if (n_tcp_hosts > 1) {
pimpl->servers.back()->set_ipv6_v6only(true);
}
}
pimpl->n_threads_http = params.n_threads_http;
if (pimpl->n_threads_http < 1) {
// +4 threads for monitoring, health and MCP.
pimpl->n_threads_http = std::max(params.n_parallel + 4, static_cast<int32_t>(std::thread::hardware_concurrency() - 1));
}
SRV_TRC("using %d threads for HTTP server\n", pimpl->n_threads_http);
return true;
}
bool server_http_context::init_listener(const common_params & params) {
auto & srv = pimpl->servers.back();
#ifdef CPPHTTPLIB_OPENSSL_SUPPORT
if (!params.ssl_file_key.empty() && !params.ssl_file_cert.empty()) {
@ -306,18 +360,8 @@ bool server_http_context::init(const common_params & params) {
return httplib::Server::HandlerResponse::Unhandled;
});
auto n_threads_http = params.n_threads_http;
if (n_threads_http < 1) {
// +4 threads for monitoring, health and some threads reserved for MCP and other tasks in the future
n_threads_http = std::max(params.n_parallel + 4, static_cast<int32_t>(std::thread::hardware_concurrency() - 1));
}
SRV_TRC("using %d threads for HTTP server\n", n_threads_http);
srv->new_task_queue = [n_threads_http] {
// spawn n_threads_http fixed thread (always alive), while allow up to 1024 max possible additional threads
// when n_threads_http is used, server will create new "dynamic" threads that will be destroyed after processing each request
// ref: https://github.com/yhirose/cpp-httplib/pull/2368
const auto max_threads = static_cast<size_t>(n_threads_http + 1024);
return new httplib::ThreadPool(n_threads_http, max_threads);
srv->new_task_queue = [this] {
return new server_http_task_queue(*pimpl->pool);
};
//
@ -432,47 +476,76 @@ bool server_http_context::init(const common_params & params) {
bool server_http_context::start() {
// Bind and listen
const auto & srv = pimpl->srv;
auto was_bound = false;
auto is_sock = false;
if (string_ends_with(std::string(hostname), ".sock")) {
is_sock = true;
SRV_TRC("%s", "setting address family to AF_UNIX\n");
srv->set_address_family(AF_UNIX);
// bind_to_port requires a second arg, any value other than 0 should
// simply get ignored
was_bound = srv->bind_to_port(hostname, 8080);
} else {
SRV_TRC("%s", "binding port with default address family\n");
// bind HTTP listen port
if (port == 0) {
const auto bound_port = srv->bind_to_any_port(hostname);
was_bound = (bound_port >= 0);
listening_addresses.clear();
for (size_t i = 0; i < pimpl->servers.size(); ++i) {
const auto & srv = pimpl->servers[i];
const auto & host = pimpl->hosts[i];
const bool is_sock = string_ends_with(host, ".sock");
bool was_bound;
if (is_sock) {
SRV_TRC("%s", "setting address family to AF_UNIX\n");
srv->set_address_family(AF_UNIX);
// AF_UNIX ignores the port, but bind_to_port requires a nonzero value.
was_bound = srv->bind_to_port(host, 8080);
} else if (port == 0) {
const auto bound_port = srv->bind_to_any_port(host);
was_bound = bound_port >= 0;
if (was_bound) {
port = bound_port;
}
} else {
was_bound = srv->bind_to_port(hostname, port);
was_bound = srv->bind_to_port(host, port);
}
if (!was_bound) {
SRV_ERR("couldn't bind HTTP server socket, hostname: %s, port: %d\n", host.c_str(), port);
stop();
listening_addresses.clear();
return false;
}
listening_addresses.push_back(is_sock ? string_format("unix://%s", host.c_str())
: string_format("%s://%s:%d", is_ssl ? "https" : "http", common_http_format_host(host).c_str(), port));
}
// n_threads_http fixed threads (always alive), plus up to 1024 dynamic threads destroyed after each request
// ref: https://github.com/yhirose/cpp-httplib/pull/2368
pimpl->pool = std::make_unique<httplib::ThreadPool>(pimpl->n_threads_http, pimpl->n_threads_http + 1024);
for (size_t i = 0; i < pimpl->servers.size(); ++i) {
const auto & srv = pimpl->servers[i];
pimpl->threads.emplace_back([srv = srv.get(), addr = listening_addresses[i]] {
if (!srv->listen_after_bind()) {
SRV_ERR("listener on %s stopped unexpectedly\n", addr.c_str());
}
});
srv->wait_until_ready();
if (!srv->is_running()) {
SRV_ERR("couldn't start HTTP listener on %s\n", listening_addresses[i].c_str());
stop();
join();
listening_addresses.clear();
return false;
}
}
if (!was_bound) {
SRV_ERR("couldn't bind HTTP server socket, hostname: %s, port: %d\n", hostname.c_str(), port);
return false;
}
// run the HTTP server in a thread
thread = std::thread([this] { pimpl->srv->listen_after_bind(); });
srv->wait_until_ready();
listening_address = is_sock ? string_format("unix://%s", hostname.c_str())
: string_format("%s://%s:%d", is_ssl ? "https" : "http", common_http_format_host(hostname).c_str(), port);
return true;
}
void server_http_context::stop() const {
if (pimpl->srv) {
pimpl->srv->stop();
for (const auto & srv : pimpl->servers) {
if (srv) {
srv->stop();
}
}
}
void server_http_context::join() {
for (auto & thread : pimpl->threads) {
if (thread.joinable()) {
thread.join();
}
}
// Queued requests still refer to their servers until the workers finish.
if (pimpl->pool) {
pimpl->pool->shutdown();
pimpl->pool.reset();
}
}
@ -584,7 +657,7 @@ static void process_handler_response(server_http_req_ptr && request, server_http
void server_http_context::get(const std::string & path, const server_http_context::handler_t & handler) const {
handlers.emplace(path, handler);
pimpl->srv->Get(path_prefix + path, [handler](const httplib::Request & req, httplib::Response & res) {
auto callback = [handler](const httplib::Request & req, httplib::Response & res) {
server_http_req_ptr request = std::make_unique<server_http_req>(server_http_req{
get_params(req),
get_headers(req),
@ -596,12 +669,16 @@ void server_http_context::get(const std::string & path, const server_http_contex
});
server_http_res_ptr response = handler(*request);
process_handler_response(std::move(request), response, res);
});
};
const std::string full_path = path_prefix + path;
for (const auto & srv : pimpl->servers) {
srv->Get(full_path, callback);
}
}
void server_http_context::post(const std::string & path, const server_http_context::handler_t & handler) const {
handlers.emplace(path, handler);
pimpl->srv->Post(path_prefix + path, [handler](const httplib::Request & req, httplib::Response & res) {
auto callback = [handler](const httplib::Request & req, httplib::Response & res) {
std::string body = req.body;
std::map<std::string, uploaded_file> files;
@ -643,12 +720,16 @@ void server_http_context::post(const std::string & path, const server_http_conte
});
server_http_res_ptr response = handler(*request);
process_handler_response(std::move(request), response, res);
});
};
const std::string full_path = path_prefix + path;
for (const auto & srv : pimpl->servers) {
srv->Post(full_path, callback);
}
}
void server_http_context::del(const std::string & path, const server_http_context::handler_t & handler) const {
handlers.emplace(path, handler);
pimpl->srv->Delete(path_prefix + path, [handler](const httplib::Request & req, httplib::Response & res) {
auto callback = [handler](const httplib::Request & req, httplib::Response & res) {
server_http_req_ptr request = std::make_unique<server_http_req>(server_http_req{
get_params(req),
get_headers(req),
@ -660,7 +741,11 @@ void server_http_context::del(const std::string & path, const server_http_contex
});
server_http_res_ptr response = handler(*request);
process_handler_response(std::move(request), response, res);
});
};
const std::string full_path = path_prefix + path;
for (const auto & srv : pimpl->servers) {
srv->Delete(full_path, callback);
}
}
//

View file

@ -68,7 +68,6 @@ struct server_http_context {
class Impl;
std::unique_ptr<Impl> pimpl;
std::thread thread; // server thread
std::atomic<bool> is_ready = false;
// note: the handler should never throw exceptions
@ -76,7 +75,6 @@ struct server_http_context {
mutable std::unordered_map<std::string, handler_t> handlers;
std::string path_prefix;
std::string hostname;
int port = 8080;
bool is_ssl = false;
@ -86,6 +84,7 @@ struct server_http_context {
bool init(const common_params & params);
bool start();
void stop() const;
void join();
void get(const std::string & path, const handler_t & handler) const;
void post(const std::string & path, const handler_t & handler) const;
@ -96,5 +95,8 @@ struct server_http_context {
void register_gcp_compat() const;
// for debugging
std::string listening_address;
std::vector<std::string> listening_addresses;
private:
bool init_listener(const common_params & params);
};

View file

@ -470,6 +470,7 @@ static void unset_reserved_args(common_preset & preset, bool unset_model_args) {
preset.unset_option("LLAMA_ARG_SSL_KEY_FILE");
preset.unset_option("LLAMA_ARG_SSL_CERT_FILE");
preset.unset_option("LLAMA_API_KEY");
preset.unset_option("LLAMA_ARG_API_KEY_FILE");
preset.unset_option("LLAMA_ARG_MODELS_DIR");
preset.unset_option("LLAMA_ARG_MODELS_MAX");
preset.unset_option("LLAMA_ARG_MODELS_PRESET");

View file

@ -111,7 +111,9 @@ int llama_server(int argc, char ** argv) {
llama_backend_init();
llama_numa_init(params.numa);
return llama_server(params, argc, argv);
const int result = llama_server(params, argc, argv);
common_log_flush(common_log_main());
return result;
}
int llama_server(common_params & params, int argc, char ** argv) {
@ -183,12 +185,6 @@ int llama_server(common_params & params, int argc, char ** argv) {
// struct that contains llama context and inference
server_context ctx_server;
server_http_context ctx_http;
if (!ctx_http.init(params)) {
SRV_ERR("%s", "failed to initialize HTTP server\n");
return 1;
}
//
// Router
//
@ -199,6 +195,13 @@ int llama_server(common_params & params, int argc, char ** argv) {
server_tools tools;
std::optional<server_models_routes> models_routes{};
server_http_context ctx_http;
if (!ctx_http.init(params)) {
SRV_ERR("%s", "failed to initialize HTTP server\n");
return 1;
}
if (is_router_server) {
// setup server instances manager
try {
@ -438,9 +441,7 @@ int llama_server(common_params & params, int argc, char ** argv) {
} catch (const std::exception & e) {
SRV_ERR("failed to load models on startup: %s\n", e.what());
ctx_http.stop();
if (ctx_http.thread.joinable()) {
ctx_http.thread.join();
}
ctx_http.join();
clean_up();
return 1;
}
@ -473,9 +474,7 @@ int llama_server(common_params & params, int argc, char ** argv) {
if (!ctx_server.load_model(params)) {
clean_up();
if (ctx_http.thread.joinable()) {
ctx_http.thread.join();
}
ctx_http.join();
SRV_ERR("%s", "exiting due to model loading error\n");
return 1;
}
@ -509,11 +508,15 @@ int llama_server(common_params & params, int argc, char ** argv) {
#endif
}
SRV_INF("listening on %s\n", ctx_http.listening_address.c_str());
bool uses_default_port = false;
for (const auto & address : ctx_http.listening_addresses) {
SRV_INF("listening on %s\n", address.c_str());
uses_default_port |= string_ends_with(address, ":8080");
}
// TODO: remove this in the future
// check the string to also handle the .sock case
if (string_ends_with(ctx_http.listening_address, ":8080")) {
if (uses_default_port) {
SRV_WRN("%s", "notice: server default port will be changed to :9931 in a future release (ref: https://github.com/ggml-org/llama.cpp/pull/26508)\n");
}
@ -523,9 +526,7 @@ int llama_server(common_params & params, int argc, char ** argv) {
SRV_WRN("%s", " please only use presets that you can trust! Unknown presets may be unsafe\n");
}
if (ctx_http.thread.joinable()) {
ctx_http.thread.join(); // keep the main thread alive
}
ctx_http.join(); // keep the main thread alive
// when the HTTP server stops, clean up and exit
clean_up();
@ -541,9 +542,7 @@ int llama_server(common_params & params, int argc, char ** argv) {
ctx_server.start_loop();
clean_up();
if (ctx_http.thread.joinable()) {
ctx_http.thread.join();
}
ctx_http.join();
if (monitor_thread.joinable()) {
monitor_thread.join();
}

View file

@ -1,5 +1,6 @@
import pytest
import requests
import socket
from utils import *
server = ServerPreset.tinyllama2()
@ -18,6 +19,37 @@ def test_server_start_simple():
assert res.status_code == 200
def test_server_multiple_addresses(monkeypatch):
# The CLI value replaces the environment value, including an unavailable address.
monkeypatch.setenv("LLAMA_ARG_HOST", "192.0.2.1")
try:
with socket.socket(socket.AF_INET6, socket.SOCK_STREAM) as probe:
probe.bind(("::1", 0))
except OSError:
pytest.skip("IPv6 loopback is unavailable") # ty: ignore[too-many-positional-arguments]
server.server_host = "127.0.0.1,::1"
server.api_key = "test-multiple-addresses"
server.start()
def check_address(host):
res = server.make_request("GET", "/health", host=host)
assert res.status_code == 200
res = server.make_request("POST", "/v1/completions", data={}, host=host)
assert res.status_code == 401
events = list(server.make_stream_request("POST", "/v1/completions", data={
"prompt": "Once upon a time",
"max_tokens": 8,
"stream": True,
}, headers={"Authorization": f"Bearer {server.api_key}"}, host=host))
assert len(events) > 1
return True
# parallel_function_calls swallows exceptions, a failed check leaves None in the results
results = parallel_function_calls([(check_address, (host,)) for host in ["127.0.0.1", "[::1]"]])
assert all(results)
def test_server_props():
global server
server.start()

View file

@ -155,8 +155,6 @@ class ServerProcess:
else:
server_path = "../../../build/bin/llama-server"
server_args = [
"--host",
self.server_host,
"--port",
self.server_port,
"--temp",
@ -164,6 +162,7 @@ class ServerProcess:
"--seed",
self.seed,
]
server_args.extend(["--host", self.server_host])
if self.offline:
server_args.append("--offline")
if self.model_file:
@ -365,6 +364,11 @@ class ServerProcess:
if hasattr(self, '_log') and self._log != sys.stdout:
self._log.close()
def make_url(self, path: str, host: str | None = None) -> str:
if host is None:
host = self.server_host.split(",")[0].strip()
return f"http://{host}:{self.server_port}{path}"
def make_request(
self,
method: str,
@ -372,8 +376,9 @@ class ServerProcess:
data: dict | Any | None = None,
headers: dict | None = None,
timeout: float | None = DEFAULT_REQUEST_TIMEOUT,
host: str | None = None,
) -> ServerResponse:
url = f"http://{self.server_host}:{self.server_port}{path}"
url = self.make_url(path, host)
parse_body = False
if method == "GET":
response = requests.get(url, headers=headers, timeout=timeout)
@ -407,8 +412,9 @@ class ServerProcess:
path: str,
data: dict | None = None,
headers: dict | None = None,
host: str | None = None,
) -> Iterator[dict]:
url = f"http://{self.server_host}:{self.server_port}{path}"
url = self.make_url(path, host)
if method == "POST":
response = requests.post(url, headers=headers, json=data, stream=True)
else:

View file

@ -40,6 +40,10 @@ export const VIDEO_FILE_TYPES = {
[FileTypeVideo.OGG]: {
extensions: [FileExtensionVideo.OGG],
mimeTypes: [MimeTypeVideo.OGG]
},
[FileTypeVideo.WEBM]: {
extensions: [FileExtensionVideo.WEBM],
mimeTypes: [MimeTypeVideo.WEBM]
}
} as const;

View file

@ -38,7 +38,8 @@ export enum FileTypeAudio {
export enum FileTypeVideo {
MP4 = 'mp4',
OGG = 'ogg'
OGG = 'ogg',
WEBM = 'webm'
}
export enum FileTypePdf {
@ -104,7 +105,8 @@ export enum FileExtensionAudio {
export enum FileExtensionVideo {
MP4 = '.mp4',
OGG = '.ogg'
OGG = '.ogg',
WEBM = '.webm'
}
export enum FileExtensionPdf {
@ -203,7 +205,8 @@ export enum MimeTypeAudio {
export enum MimeTypeVideo {
MP4 = 'video/mp4',
OGG = 'video/ogg'
OGG = 'video/ogg',
WEBM = 'video/webm'
}
export enum MimeTypeImage {

View file

@ -51,6 +51,7 @@ export function getFileTypeCategory(mimeType: string): FileTypeCategory | null {
// Video
case MimeTypeVideo.MP4:
case MimeTypeVideo.OGG:
case MimeTypeVideo.WEBM:
return FileTypeCategory.VIDEO;
// PDF

View file

@ -335,7 +335,7 @@
<ModeWatcher />
<Toaster richColors />
<Toaster closeButton richColors />
</Tooltip.Provider>
<!-- PWA update prompt + version -->

View file

@ -1248,6 +1248,13 @@ bool parse_trailers(stream_line_reader &line_reader, Headers &dest,
// to look up.
split(trailer_header.data(), trailer_header.data() + trailer_header.size(),
',', [&](const char *b, const char *e) {
// A legitimate message declares only a handful of trailers. Cap the
// set so a peer cannot grow it without bound: an oversized set only
// arises from an attempt to force many colliding names into
// quadratic lookups (case_ignore::hash is unkeyed).
if (declared_trailers.size() >= CPPHTTPLIB_HEADER_MAX_COUNT) {
return;
}
std::string key(b, e);
if (prohibited_trailers.find(key) == prohibited_trailers.end()) {
declared_trailers.insert(key);
@ -1258,6 +1265,8 @@ bool parse_trailers(stream_line_reader &line_reader, Headers &dest,
size_t trailer_header_count = 0;
while (strcmp(line_reader.ptr(), "\r\n") != 0) {
if (line_reader.size() > CPPHTTPLIB_HEADER_MAX_LENGTH) { return false; }
// Count every received trailer field, not only the declared ones stored in
// dest, so undeclared fields cannot keep this loop running past the limit.
if (trailer_header_count >= CPPHTTPLIB_HEADER_MAX_COUNT) { return false; }
constexpr auto line_terminator_len = 2;
@ -1270,12 +1279,13 @@ bool parse_trailers(stream_line_reader &line_reader, Headers &dest,
if (declared_trailers.find(key) !=
declared_trailers.end()) {
dest.emplace(key, val);
trailer_header_count++;
}
})) {
return false;
}
trailer_header_count++;
if (!line_reader.getline()) { return false; }
}
@ -3844,11 +3854,13 @@ bool read_content(Stream &strm, T &x, size_t payload_max_length, int &status,
ssize_t write_request_line(Stream &strm, const std::string &method,
const std::string &path) {
// A request target must not carry CR/LF (or other control octets); otherwise
// a value smuggled into it splits the request line and injects headers or a
// whole request. The same field-value check already guards header values in
// check_and_write_headers and the request target in
// perform_websocket_handshake; apply it here too.
// Neither the method nor the request target may carry CR/LF (or other
// control octets); otherwise a value smuggled into either splits the request
// line and injects headers or a whole request. The method must be a token
// (RFC 9110 Section 9.1), which also rejects an empty method and embedded
// spaces. The target gets the same field-value check that already guards
// header values in check_and_write_headers.
if (!fields::is_token(method)) { return -1; }
if (!fields::is_field_value(path)) { return -1; }
std::string s = method;
@ -8626,7 +8638,10 @@ bool Server::write_response_core(Stream &strm, bool close_connection,
// Prepare additional headers
if (close_connection ||
detail::has_header_token(req.headers, "Connection", "close") ||
400 <= res.status) { // Don't leave connections open after errors
400 <= res.status || // Don't leave connections open after errors
// The client withholds the body until `100 Continue`, which was never
// sent, so whether and when the body follows is unknown.
(req.expect_100_continue_pending_ && detail::has_framed_body(req))) {
res.set_header("Connection", "close");
} else {
std::string s = "timeout=";
@ -8877,6 +8892,13 @@ bool Server::read_content_core(
}
#endif
// The client is waiting for this before it sends the body.
if (req.expect_100_continue_pending_) {
req.expect_100_continue_pending_ = false;
detail::write_response_line(strm, StatusCode::Continue_100);
strm.write("\r\n");
}
if (!detail::read_content(strm, req, payload_max_length_, res.status, nullptr,
out, true)) {
return false;
@ -9714,19 +9736,20 @@ Server::process_request(Stream &strm, const std::string &remote_addr,
// case-insensitive, and a 100-continue expectation in an HTTP/1.0 request
// must be ignored. An expectation we do not recognize is left alone; the
// 417 the section allows for one is a MAY, not a requirement.
//
// `100 Continue` itself is deferred until the body is actually read (see
// read_content_core), so a request rejected by a later handler never
// invites the client to send a body nobody will read.
if (req.version != "HTTP/1.0" &&
detail::has_header_token(req.headers, "Expect", "100-continue")) {
int status = StatusCode::Continue_100;
if (expect_100_continue_handler_) {
status = expect_100_continue_handler_(req, res);
}
switch (status) {
case StatusCode::Continue_100:
case StatusCode::ExpectationFailed_417:
detail::write_response_line(strm, status);
strm.write("\r\n");
break;
default:
if (status == StatusCode::Continue_100) {
req.expect_100_continue_pending_ = true;
} else {
if (res.status == -1) { res.status = status; }
connection_closed = true;
return write_response(strm, true, req, res);
}
@ -9739,18 +9762,25 @@ Server::process_request(Stream &strm, const std::string &remote_addr,
};
// WebSocket upgrade
// Check pre_routing_handler_ before upgrading so that authentication
// and other middleware can reject the request with an HTTP response
// (e.g., 401) before the protocol switches.
// Run pre_routing_handler_ and pre_request_handler_ before upgrading so
// that authentication and other middleware can reject the request with an
// HTTP response (e.g., 401) before the protocol switches.
if (detail::is_websocket_upgrade(req)) {
if (pre_routing_handler_ &&
pre_routing_handler_(req, res) == HandlerResponse::Handled) {
if (res.status == -1) { res.status = StatusCode::OK_200; }
return write_response(strm, close_connection, req, res);
return write_response_with_content(strm, close_connection, req, res);
}
// Find matching WebSocket handler
for (const auto &entry : websocket_handlers_) {
if (entry.matcher->match(req)) {
req.matched_route = entry.matcher->pattern();
if (pre_request_handler_ &&
pre_request_handler_(req, res) == HandlerResponse::Handled) {
if (res.status == -1) { res.status = StatusCode::OK_200; }
return write_response_with_content(strm, close_connection, req, res);
}
// Compute accept key
auto client_key = req.get_header_value("Sec-WebSocket-Key");
auto accept_key = detail::websocket_accept_key(client_key);
@ -10610,22 +10640,45 @@ ssize_t ChunkedDecoder::read_payload(char *buf, size_t len,
stream_line_reader lr(strm, line_buf, sizeof(line_buf));
if (!lr.getline()) { return -1; }
// Everything below is bounded by eol rather than by the buffer's NUL, so
// the line terminator is never mistaken for line content.
const char *eol = lr.ptr() + lr.size();
if (lr.end_with_crlf()) {
eol -= 2;
} else if (eol != lr.ptr() && eol[-1] == '\n') {
// Only reachable under CPPHTTPLIB_ALLOW_LF_AS_LINE_TERMINATOR, where
// getline() ends the line on a bare LF. That LF is the terminator, so it
// has to come off here or the check below would reject the line.
eol -= 1;
}
// RFC 9112 §7.1: chunk-size = 1*HEXDIG
const char *p = lr.ptr();
int v = 0;
if (!is_hex(*p, v)) { return -1; }
if (p == eol || !is_hex(*p, v)) { return -1; }
size_t chunk_len = 0;
constexpr size_t chunk_len_max = (std::numeric_limits<size_t>::max)();
for (; is_hex(*p, v); ++p) {
for (; p < eol && is_hex(*p, v); ++p) {
if (chunk_len > (chunk_len_max >> 4)) { return -1; }
chunk_len = (chunk_len << 4) | static_cast<size_t>(v);
}
while (is_space_or_tab(*p)) {
while (p < eol && is_space_or_tab(*p)) {
++p;
}
if (*p != '\0' && *p != ';' && *p != '\r' && *p != '\n') { return -1; }
// RFC 9112 §7.1.1: only a chunk-ext may sit between the size and the line
// terminator, and it is built from tokens and quoted-strings, so it never
// holds a CR, LF or any other control character. getline() reads up to the
// CRLF, so a bare LF left in here would be swallowed as extension text
// while an intermediary that ends the line on it delimits the chunks
// differently, and the two disagree on where the body ends (request
// smuggling).
if (p < eol && *p != ';') { return -1; }
for (; p < eol; ++p) {
if (!is_space_or_tab(*p) && !fields::is_field_vchar(*p)) { return -1; }
}
if (chunk_len == 0) {
chunk_remaining = 0;
@ -11054,9 +11107,10 @@ bool ClientImpl::write_request(Stream &strm, Request &req,
// Write request line and headers
if (detail::write_request_line(bstrm, req.method, path_with_query) < 0) {
// A rejected target (e.g. CR/LF smuggled in via a decoded redirect
// Location under set_path_encode(false)) must fail the request cleanly
// instead of emitting a request-line-less, header-injecting request.
// A rejected method (not a token, e.g. carrying CR/LF) or target (e.g.
// CR/LF smuggled in via a decoded redirect Location under
// set_path_encode(false)) must fail the request cleanly instead of
// emitting a request-line-less, header-injecting request.
error = Error::Write;
output_error_log(error, &req);
return false;
@ -14650,11 +14704,11 @@ void shutdown(session_t session, bool graceful) {
auto ssl = static_cast<SSL *>(session);
if (graceful) {
// First call sends close_notify
if (SSL_shutdown(ssl) == 0) {
// Second call waits for peer's close_notify
SSL_shutdown(ssl);
}
// Send close_notify without waiting for the peer's. The connection is
// closed right after this, so a unidirectional shutdown is enough, and an
// idle peer that never answers would otherwise hold this thread until the
// read timeout. The other backends do not wait either.
SSL_shutdown(ssl);
}
}

View file

@ -8,8 +8,8 @@
#ifndef CPPHTTPLIB_HTTPLIB_H
#define CPPHTTPLIB_HTTPLIB_H
#define CPPHTTPLIB_VERSION "0.56.0"
#define CPPHTTPLIB_VERSION_NUM "0x003800"
#define CPPHTTPLIB_VERSION "0.57.1"
#define CPPHTTPLIB_VERSION_NUM "0x003901"
#ifdef _WIN32
#if defined(_WIN32_WINNT) && _WIN32_WINNT < 0x0A00
@ -1756,6 +1756,7 @@ struct Request {
// private members...
bool body_consumed_ = false;
bool expect_100_continue_pending_ = false;
size_t redirect_count_ = CPPHTTPLIB_REDIRECT_MAX_COUNT;
size_t content_length_ = 0;
ContentProvider content_provider_;