mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-10-03 11:35:46 +00:00
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:
commit
084d797f2d
52 changed files with 2420 additions and 631 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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()}},
|
||||
|
|
|
|||
|
|
@ -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 ||
|
||||
|
|
|
|||
|
|
@ -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\\|>)" },
|
||||
};
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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)) {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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 -----------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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];
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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"}});
|
||||
|
||||
|
|
|
|||
|
|
@ -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];
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]);
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -335,7 +335,7 @@
|
|||
|
||||
<ModeWatcher />
|
||||
|
||||
<Toaster richColors />
|
||||
<Toaster closeButton richColors />
|
||||
</Tooltip.Provider>
|
||||
|
||||
<!-- PWA update prompt + version -->
|
||||
|
|
|
|||
114
vendor/cpp-httplib/httplib.cpp
vendored
114
vendor/cpp-httplib/httplib.cpp
vendored
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
5
vendor/cpp-httplib/httplib.h
vendored
5
vendor/cpp-httplib/httplib.h
vendored
|
|
@ -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_;
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue