diff --git a/common/arg.cpp b/common/arg.cpp index f10aea7a5..28334f4b8 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -2645,6 +2645,27 @@ common_params_context common_params_parser_init(common_params & params, llama_ex params.mtmd_batch_max_tokens = value; } ).set_examples({LLAMA_EXAMPLE_SERVER}).set_env("LLAMA_ARG_MTMD_BATCH_MAX_TOKENS")); + add_opt(common_arg( + {"--video-fps"}, "N", + string_format("target video frame rate (default: %.1f)", params.video_fps), + [](common_params & params, const std::string & value) { + params.video_fps = std::stof(value); + } + ).set_examples(mmproj_examples).set_env("LLAMA_ARG_VIDEO_FPS")); + add_opt(common_arg( + {"--video-timestamp-interval"}, "N", + string_format("interval in milliseconds between text timestamps (default: %" PRId64 ")", params.video_timestamp_interval_ms), + [](common_params & params, int value) { + params.video_timestamp_interval_ms = value; + } + ).set_examples(mmproj_examples).set_env("LLAMA_ARG_VIDEO_TIMESTAMP_INTERVAL")); + add_opt(common_arg( + {"--video-ffmpeg-dir"}, "DIR", + "path to the directory containing ffmpeg and ffprobe (default: search in PATH)", + [](common_params & params, const std::string & value) { + params.video_ffmpeg_bin_dir = value; + } + ).set_examples(mmproj_examples).set_env("LLAMA_ARG_VIDEO_FFMPEG_DIR")); if (params.is_gen_docs || llama_supports_rpc()) { add_opt(common_arg( {"--rpc"}, "SERVERS", @@ -2751,14 +2772,20 @@ common_params_context common_params_parser_init(common_params & params, llama_ex if (value < 0) { throw std::invalid_argument("invalid value"); } - for (int i = 0; i < value; ++i) { - // keep strings alive and avoid leaking memory by storing them in a static vector - static std::list buft_overrides; - buft_overrides.push_back(llm_ffn_exps_block_regex(i)); - params.tensor_buft_overrides.push_back({buft_overrides.back().c_str(), ggml_backend_cpu_buffer_type()}); - } + llm_add_n_cpu_ffn_overrides(value, LLM_FFN_EXPS_REGEX, params.tensor_buft_overrides); } ).set_env("LLAMA_ARG_N_CPU_MOE")); + add_opt(common_arg( + {"-ncffn", "--n-cpu-ffn"}, "N", + "keep the dense FFN weights of the first N layers in the CPU\n" + "(dense models; for MoE expert weights use --n-cpu-moe)", + [](common_params & params, int value) { + if (value < 0) { + throw std::invalid_argument("invalid value"); + } + llm_add_n_cpu_ffn_overrides(value, LLM_FFN_DENSE_REGEX, params.tensor_buft_overrides); + } + ).set_env("LLAMA_ARG_N_CPU_FFN")); GGML_ASSERT(params.n_gpu_layers < 0); // string_format would need to be extended for a default >= 0 add_opt(common_arg( {"-ngl", "--gpu-layers", "--n-gpu-layers"}, "N", @@ -4085,11 +4112,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex if (value < 0) { throw std::invalid_argument("invalid value"); } - for (int i = 0; i < value; ++i) { - static std::list buft_overrides_draft; - buft_overrides_draft.push_back(llm_ffn_exps_block_regex(i)); - params.speculative.draft.tensor_buft_overrides.push_back({buft_overrides_draft.back().c_str(), ggml_backend_cpu_buffer_type()}); - } + llm_add_n_cpu_ffn_overrides(value, LLM_FFN_EXPS_REGEX, params.speculative.draft.tensor_buft_overrides); } ).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_SPEC_DRAFT_N_CPU_MOE")); diff --git a/common/chat.cpp b/common/chat.cpp index bc133b55e..50778f483 100644 --- a/common/chat.cpp +++ b/common/chat.cpp @@ -1193,6 +1193,8 @@ static common_chat_params common_chat_params_init_qwen3_coder(const common_chat_ "", }; + auto is_qwen3_coder = !supports_reasoning; + if (supports_reasoning) { data.thinking_start_tag = ""; // Support both and as reasoning end sequences. @@ -1233,13 +1235,15 @@ static common_chat_params common_chat_params_init_qwen3_coder(const common_chat_ std::vector tool_call_starts = { "" }; - // Match complete opener for Qwen3-Coder models that occasionally omit the - // starting . The model may hallucinate a tool name, but it is preferable over - // constraining on - foreach_function(inputs.tools, [&](const json & tool) { - const std::string name = tool.at("function").at("name"); - tool_call_starts.push_back(""); - }); + if (is_qwen3_coder) { + // Match complete opener for Qwen3-Coder models that occasionally omit the + // starting . The model may hallucinate a tool name, but it is preferable over + // constraining on + foreach_function(inputs.tools, [&](const json & tool) { + const std::string name = tool.at("function").at("name"); + tool_call_starts.push_back(""); + }); + } auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) { auto generation_prompt = p.literal(GEN_PREFIX); @@ -1304,10 +1308,13 @@ static common_chat_params common_chat_params_init_qwen3_coder(const common_chat_ auto min_calls = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED ? 1 : 0; + auto tool_call_body = tool_choice + "" + p.space(); + auto tool_call = p.rule("tool-call", "\n" + tool_call_body); + // Qwen3-Coder models may occasionally omit the token. - auto tool_call_body = tool_choice + "" + p.space(); - auto tool_call_first = p.rule("tool-call-first", p.optional(p.literal("\n")) + tool_call_body); - auto tool_call = p.rule("tool-call", "\n" + tool_call_body); + auto tool_call_first = is_qwen3_coder ? + p.rule("tool-call-first", p.optional(p.literal("\n")) + tool_call_body) : + tool_call; auto calls = inputs.parallel_tool_calls ? tool_call_first + p.zero_or_more(tool_call) : tool_call_first; auto tool_calls = p.trigger_rule("tool-call-root", p.repeat(calls, min_calls, 1)); diff --git a/common/common.h b/common/common.h index 95d8e332a..74ef19c85 100644 --- a/common/common.h +++ b/common/common.h @@ -9,6 +9,7 @@ #include "ggml.h" #include "llama.h" +#include #include #include #include @@ -590,6 +591,11 @@ struct common_params { int image_max_tokens = -1; int mtmd_batch_max_tokens = 1024; + // for video input + float video_fps = 4.0f; + int64_t video_timestamp_interval_ms = 5000; + std::string video_ffmpeg_bin_dir = ""; + // finetune struct lr_opt lr; enum ggml_opt_optimizer_type optimizer = GGML_OPT_OPTIMIZER_TYPE_ADAMW; @@ -1109,19 +1115,30 @@ const char * const LLM_KV_SPLIT_TENSORS_COUNT = "split.tensors.count"; } // -// MoE utils +// FFN offload utils // const char * const LLM_FFN_EXPS_REGEX = "\\.ffn_(up|down|gate|gate_up)_(ch|)exps"; -inline std::string llm_ffn_exps_block_regex(int idx) { - return string_format("blk\\.%d%s", idx, LLM_FFN_EXPS_REGEX); +const char * const LLM_FFN_DENSE_REGEX = "\\.ffn_(up|down|gate)\\."; + +inline std::string llm_ffn_block_regex(int idx, const char * ffn_regex) { + return string_format("blk\\.%d%s", idx, ffn_regex); } inline llama_model_tensor_buft_override llm_ffn_exps_cpu_override() { return { LLM_FFN_EXPS_REGEX, ggml_backend_cpu_buffer_type() }; } +inline void llm_add_n_cpu_ffn_overrides(int n, const char * ffn_regex, std::vector & overrides) { + // keep strings alive and avoid leaking memory by storing them in a static list + static std::list buft_override_strings; + for (int i = 0; i < n; ++i) { + buft_override_strings.push_back(llm_ffn_block_regex(i, ffn_regex)); + overrides.push_back({buft_override_strings.back().c_str(), ggml_backend_cpu_buffer_type()}); + } +} + // // training utils // diff --git a/conversion/nemotron.py b/conversion/nemotron.py index e5d167185..cd0d48c8f 100644 --- a/conversion/nemotron.py +++ b/conversion/nemotron.py @@ -202,6 +202,10 @@ class NemotronHModel(GraniteHybridModel): is_moe: bool = False supports_mtp_export = True + _SSM_LAYER_TYPES = {"mamba", "linear_attention"} + _ATTN_LAYER_TYPES = {"attention", "full_attention"} + _MLP_LAYER_TYPES = {"moe"} + def __init__(self, *args, **kwargs): # We have to determine the correct model architecture (MoE vs non-MoE) before # calling the parent __init__. This is because the parent constructor @@ -242,8 +246,8 @@ class NemotronHModel(GraniteHybridModel): self._ssm_layers = [i for i, val in enumerate(pattern) if val == "M"] self._mlp_layers = [i for i, val in enumerate(pattern) if val == ("E" if self.is_moe else "-")] else: - self._ssm_layers = [i for i, val in enumerate(pattern) if val == "mamba"] - self._mlp_layers = [i for i, val in enumerate(pattern) if val == "moe"] + self._ssm_layers = [i for i, val in enumerate(pattern) if val in self._SSM_LAYER_TYPES] + self._mlp_layers = [i for i, val in enumerate(pattern) if val in self._MLP_LAYER_TYPES] # `--no-mtp` drops it entirely; `--mtp` exports only the MTP head self._mtp_bid: int | None = None @@ -272,7 +276,7 @@ class NemotronHModel(GraniteHybridModel): if isinstance(pattern, str): return [i for i, val in enumerate(pattern) if val == "*"] - return [i for i, val in enumerate(pattern) if val == "attention"] + return [i for i, val in enumerate(pattern) if val in self._ATTN_LAYER_TYPES] @classmethod def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None: diff --git a/ggml/include/ggml-rpc.h b/ggml/include/ggml-rpc.h index 059e44962..cbfe40013 100644 --- a/ggml/include/ggml-rpc.h +++ b/ggml/include/ggml-rpc.h @@ -6,8 +6,8 @@ extern "C" { #endif -#define RPC_PROTO_MAJOR_VERSION 5 -#define RPC_PROTO_MINOR_VERSION 1 +#define RPC_PROTO_MAJOR_VERSION 6 +#define RPC_PROTO_MINOR_VERSION 0 #define RPC_PROTO_PATCH_VERSION 0 #ifdef __cplusplus diff --git a/ggml/src/ggml-backend-impl.h b/ggml/src/ggml-backend-impl.h index 9c56ec30c..40cea024c 100644 --- a/ggml/src/ggml-backend-impl.h +++ b/ggml/src/ggml-backend-impl.h @@ -83,6 +83,7 @@ extern "C" { GGML_API ggml_backend_buffer_t ggml_backend_multi_buffer_alloc_buffer(ggml_backend_buffer_t * buffers, size_t n_buffers); GGML_API bool ggml_backend_buffer_is_multi_buffer(ggml_backend_buffer_t buffer); GGML_API void ggml_backend_multi_buffer_set_usage(ggml_backend_buffer_t buffer, enum ggml_backend_buffer_usage usage); + GGML_API void ggml_backend_meta_buffer_set_usage (ggml_backend_buffer_t buffer, enum ggml_backend_buffer_usage usage); // // Backend (meta) diff --git a/ggml/src/ggml-backend-meta.cpp b/ggml/src/ggml-backend-meta.cpp index fe58ea3bb..3ec40fb1a 100644 --- a/ggml/src/ggml-backend-meta.cpp +++ b/ggml/src/ggml-backend-meta.cpp @@ -1168,7 +1168,6 @@ static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state( } static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state(const struct ggml_tensor * tensor, bool assume_sync) { - GGML_ASSERT(ggml_backend_buffer_is_meta(tensor->buffer)); ggml_backend_meta_buffer_context * buf_ctx = (ggml_backend_meta_buffer_context *) tensor->buffer->context; return ggml_backend_meta_get_split_state(buf_ctx->get_simple_tensor_container(tensor), tensor, assume_sync); } @@ -1259,7 +1258,14 @@ static enum ggml_status ggml_backend_meta_buffer_init_tensor_impl(ggml_backend_m t_ij->data = (char *) ggml_backend_buffer_get_base(simple_buf) + size_t(tensor->data) - size_t(ggml_backend_buffer_get_base(tensor->buffer)); } - t_ij->extra = tensor->extra; + + if (simple_buf) { + // the backend that owns the buffer will set .extra + ggml_backend_buffer_init_tensor(simple_buf, t_ij); + } else { + t_ij->extra = tensor->extra; + } + for (int i = 0; i < GGML_MAX_SRC; i++) { t_ij->src[i] = tensor->src[i]; if (tensor->src[i] == tensor) { @@ -1668,6 +1674,16 @@ bool ggml_backend_buffer_is_meta(ggml_backend_buffer_t buf) { return buf != nullptr && buf->iface.free_buffer == ggml_backend_meta_buffer_iface.free_buffer; } +void ggml_backend_meta_buffer_set_usage(ggml_backend_buffer_t buffer, enum ggml_backend_buffer_usage usage) { + GGML_ASSERT(ggml_backend_buffer_is_meta(buffer)); + ggml_backend_meta_buffer_context * buf_ctx = (ggml_backend_meta_buffer_context *) buffer->context; + for (size_t i = 0; i < buf_ctx->bufs.size(); i++) { + if (buf_ctx->bufs[i]) { + ggml_backend_buffer_set_usage(buf_ctx->bufs[i].get(), usage); + } + } +} + static ggml_backend_buffer_t ggml_backend_meta_buffer_type_alloc_buffer(ggml_backend_buffer_type_t buft, size_t size) { const size_t n_simple_bufts = ggml_backend_meta_buft_n_bufts(buft); diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index 3b9316c9c..1da189de5 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -182,6 +182,8 @@ void ggml_backend_buffer_set_usage(ggml_backend_buffer_t buffer, enum ggml_backe // FIXME: add a generic callback to the buffer interface if (ggml_backend_buffer_is_multi_buffer(buffer)) { ggml_backend_multi_buffer_set_usage(buffer, usage); + } else if (ggml_backend_buffer_is_meta(buffer)) { + ggml_backend_meta_buffer_set_usage(buffer, usage); } } diff --git a/ggml/src/ggml-cuda/mmq-config-pascal.cuh b/ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh similarity index 99% rename from ggml/src/ggml-cuda/mmq-config-pascal.cuh rename to ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh index e7d4a9a3f..83eb7c146 100644 --- a/ggml/src/ggml-cuda/mmq-config-pascal.cuh +++ b/ggml/src/ggml-cuda/mmq-config-pascal-dp4a.cuh @@ -1,4 +1,4 @@ -static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_pascal(ggml_type type, int J, bool fallback) { +static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_pascal_dp4a(ggml_type type, int J, bool fallback) { CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); diff --git a/ggml/src/ggml-cuda/mmq-config-pascal-older.cuh b/ggml/src/ggml-cuda/mmq-config-pascal-older.cuh new file mode 100644 index 000000000..2a8dc9e1a --- /dev/null +++ b/ggml/src/ggml-cuda/mmq-config-pascal-older.cuh @@ -0,0 +1,273 @@ +static constexpr __host__ __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config_pascal_older(ggml_type type, int J, bool fallback) { + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q1_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_1, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q8_0, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + +// --------------------------------------------------------------------------------------------- + + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q2_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q2_K, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q3_K, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q4_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q5_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_Q6_K, 256, 1, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q6_K, MMQ_ITER_K, false, false); + +// --------------------------------------------------------------------------------------------- + + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ1_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XXS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_XS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ2_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q3_K, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_XXS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ3_S, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_XS, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_IQ4_NL, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, MMQ_ITER_K, false, false); + +// --------------------------------------------------------------------------------------------- + + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_MXFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_1, MMQ_ITER_K, false, false); + + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, true); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 8, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 16, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 24, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 32, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 40, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 48, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + CASE(GGML_TYPE_NVFP4, 256, 2, 64, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_NVFP4, MMQ_ITER_K, false, false); + + return ggml_cuda_mmq_config(GGML_TYPE_COUNT, 256, 2, 128, 64, GGML_CUDA_MMQ_SRAM_LAYOUT_Q8_0, 256, false, true); +} diff --git a/ggml/src/ggml-cuda/mmq.cu b/ggml/src/ggml-cuda/mmq.cu index 47a15cb63..09976eb5d 100644 --- a/ggml/src/ggml-cuda/mmq.cu +++ b/ggml/src/ggml-cuda/mmq.cu @@ -316,7 +316,9 @@ bool ggml_cuda_should_use_mmq(enum ggml_type type, int cc, int64_t ne11, int64_t } if (ggml_cuda_highest_compiled_arch(cc) < GGML_CUDA_CC_DP4A) { - return false; + // for MoE, mmq is faster even without native dp4a + // TODO: check if cards older than pascal might benefit from this as well + return cc >= GGML_CUDA_CC_PASCAL && n_experts > 0; } #ifdef GGML_CUDA_FORCE_MMQ diff --git a/ggml/src/ggml-cuda/mmq.cuh b/ggml/src/ggml-cuda/mmq.cuh index 0c2a82ee1..8f14fa947 100644 --- a/ggml/src/ggml-cuda/mmq.cuh +++ b/ggml/src/ggml-cuda/mmq.cuh @@ -214,7 +214,8 @@ struct ggml_cuda_mmq_config { return ggml_cuda_mmq_config((type_), (nthreads_), (occupancy_), (I_), (J_), (sram_layout_), (K_vram_), (stream_k_), (fallback_)); \ } \ -#include "mmq-config-pascal.cuh" +#include "mmq-config-pascal-older.cuh" +#include "mmq-config-pascal-dp4a.cuh" #include "mmq-config-ampere.cuh" #include "mmq-config-blackwell.cuh" @@ -248,7 +249,10 @@ static __host__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(const ggml_type ty if (ggml_cuda_highest_compiled_arch(cc) >= GGML_CUDA_CC_VOLTA) { return ggml_cuda_mmq_get_config_ampere(type, J, fallback); } - return ggml_cuda_mmq_get_config_pascal(type, J, fallback); + if (ggml_cuda_highest_compiled_arch(cc) >= GGML_CUDA_CC_DP4A) { + return ggml_cuda_mmq_get_config_pascal_dp4a(type, J, fallback); + } + return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback); } static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_type type, int J, bool fallback) { @@ -269,8 +273,10 @@ static constexpr __device__ ggml_cuda_mmq_config ggml_cuda_mmq_get_config(ggml_t return ggml_cuda_mmq_get_config_blackwell(type, J, fallback); #elif __CUDA_ARCH__ >= GGML_CUDA_CC_VOLTA return ggml_cuda_mmq_get_config_ampere(type, J, fallback); +#elif __CUDA_ARCH__ >= GGML_CUDA_CC_DP4A + return ggml_cuda_mmq_get_config_pascal_dp4a(type, J, fallback); #else - return ggml_cuda_mmq_get_config_pascal(type, J, fallback); + return ggml_cuda_mmq_get_config_pascal_older(type, J, fallback); #endif // BLACKWELL_MMA_AVAILABLE #endif // GGML_USE_HIP GGML_UNUSED_VARS(type, J, fallback); diff --git a/ggml/src/ggml-metal/ggml-metal-context.m b/ggml/src/ggml-metal/ggml-metal-context.m index 32d97cd5d..1227ed39a 100644 --- a/ggml/src/ggml-metal/ggml-metal-context.m +++ b/ggml/src/ggml-metal/ggml-metal-context.m @@ -84,106 +84,108 @@ struct ggml_metal { ggml_metal_t ggml_metal_init(ggml_metal_device_t dev) { GGML_LOG_INFO("%s: allocating\n", __func__); + @autoreleasepool { #if TARGET_OS_OSX && !GGML_METAL_NDEBUG - // Show all the Metal device instances in the system - NSArray * devices = MTLCopyAllDevices(); - for (id device in devices) { - GGML_LOG_INFO("%s: found device: %s\n", __func__, [[device name] UTF8String]); - } - [devices release]; // since it was created by a *Copy* C method + // Show all the Metal device instances in the system + NSArray * devices = MTLCopyAllDevices(); + for (id device in devices) { + GGML_LOG_INFO("%s: found device: %s\n", __func__, [[device name] UTF8String]); + } + [devices release]; // since it was created by a *Copy* C method #endif - // init context - ggml_metal_t res = calloc(1, sizeof(struct ggml_metal)); + // init context + ggml_metal_t res = calloc(1, sizeof(struct ggml_metal)); - id device = ggml_metal_device_get_obj(dev); + id device = ggml_metal_device_get_obj(dev); - GGML_LOG_INFO("%s: picking default device: %s\n", __func__, [[device name] UTF8String]); - - // TODO: would it be better to have one queue for the backend and one queue for the device? - // the graph encoders and async ops would use the backend queue while the sync ops would use the device queue? - //res->queue = [device newCommandQueue]; [TAG_QUEUE_PER_BACKEND] - id queue = ggml_metal_device_get_queue(dev); - if (queue == nil) { - GGML_LOG_ERROR("%s: error: failed to create command queue\n", __func__); - return NULL; - } - - res->dev = dev; - res->lib = ggml_metal_device_get_library(dev); - if (res->lib == NULL) { - GGML_LOG_WARN("%s: the device does not have a precompiled Metal library - this is unexpected\n", __func__); - GGML_LOG_WARN("%s: will try to compile it on the fly\n", __func__); - - res->lib = ggml_metal_library_init(dev); - if (res->lib == NULL) { - GGML_LOG_ERROR("%s: error: failed to initialize the Metal library\n", __func__); - - free(res); + GGML_LOG_INFO("%s: picking default device: %s\n", __func__, [[device name] UTF8String]); + // TODO: would it be better to have one queue for the backend and one queue for the device? + // the graph encoders and async ops would use the backend queue while the sync ops would use the device queue? + //res->queue = [device newCommandQueue]; [TAG_QUEUE_PER_BACKEND] + id queue = ggml_metal_device_get_queue(dev); + if (queue == nil) { + GGML_LOG_ERROR("%s: error: failed to create command queue\n", __func__); return NULL; } - } - res->ev_cpy = ggml_metal_device_event_init(dev); + res->dev = dev; + res->lib = ggml_metal_device_get_library(dev); + if (res->lib == NULL) { + GGML_LOG_WARN("%s: the device does not have a precompiled Metal library - this is unexpected\n", __func__); + GGML_LOG_WARN("%s: will try to compile it on the fly\n", __func__); - const struct ggml_metal_device_props * props_dev = ggml_metal_device_get_props(dev); + res->lib = ggml_metal_library_init(dev); + if (res->lib == NULL) { + GGML_LOG_ERROR("%s: error: failed to initialize the Metal library\n", __func__); - snprintf(res->name, sizeof(res->name), "%s", props_dev->name); + free(res); - res->d_queue = dispatch_queue_create("ggml-metal", DISPATCH_QUEUE_CONCURRENT); - - res->use_fusion = getenv("GGML_METAL_FUSION_DISABLE") == nil; - res->use_concurrency = getenv("GGML_METAL_CONCURRENCY_DISABLE") == nil; - - { - const char * val = getenv("GGML_METAL_GRAPH_DEBUG"); - res->debug_graph = val ? atoi(val) : 0; - } - - { - const char * val = getenv("GGML_METAL_FUSION_DEBUG"); - res->debug_fusion = val ? atoi(val) : 0; - } - - res->use_graph_optimize = true; - - if (getenv("GGML_METAL_GRAPH_OPTIMIZE_DISABLE") != NULL) { - res->use_graph_optimize = false; - } - - memset(res->fuse_cnt, 0, sizeof(res->fuse_cnt)); - - GGML_LOG_INFO("%s: use fusion = %s\n", __func__, res->use_fusion ? "true" : "false"); - GGML_LOG_INFO("%s: use concurrency = %s\n", __func__, res->use_concurrency ? "true" : "false"); - GGML_LOG_INFO("%s: use graph optimize = %s\n", __func__, res->use_graph_optimize ? "true" : "false"); - - res->capture_compute = 0; - res->capture_started = false; - res->capture_scope = nil; - - { - const char * val = getenv("GGML_METAL_CAPTURE_COMPUTE"); - if (val) { - res->capture_compute = atoi(val); + return NULL; + } } + + res->ev_cpy = ggml_metal_device_event_init(dev); + + const struct ggml_metal_device_props * props_dev = ggml_metal_device_get_props(dev); + + snprintf(res->name, sizeof(res->name), "%s", props_dev->name); + + res->d_queue = dispatch_queue_create("ggml-metal", DISPATCH_QUEUE_CONCURRENT); + + res->use_fusion = getenv("GGML_METAL_FUSION_DISABLE") == nil; + res->use_concurrency = getenv("GGML_METAL_CONCURRENCY_DISABLE") == nil; + + { + const char * val = getenv("GGML_METAL_GRAPH_DEBUG"); + res->debug_graph = val ? atoi(val) : 0; + } + + { + const char * val = getenv("GGML_METAL_FUSION_DEBUG"); + res->debug_fusion = val ? atoi(val) : 0; + } + + res->use_graph_optimize = true; + + if (getenv("GGML_METAL_GRAPH_OPTIMIZE_DISABLE") != NULL) { + res->use_graph_optimize = false; + } + + memset(res->fuse_cnt, 0, sizeof(res->fuse_cnt)); + + GGML_LOG_INFO("%s: use fusion = %s\n", __func__, res->use_fusion ? "true" : "false"); + GGML_LOG_INFO("%s: use concurrency = %s\n", __func__, res->use_concurrency ? "true" : "false"); + GGML_LOG_INFO("%s: use graph optimize = %s\n", __func__, res->use_graph_optimize ? "true" : "false"); + + res->capture_compute = 0; + res->capture_started = false; + res->capture_scope = nil; + + { + const char * val = getenv("GGML_METAL_CAPTURE_COMPUTE"); + if (val) { + res->capture_compute = atoi(val); + } + } + + res->has_error = false; + + res->gf = nil; + res->encode_async = nil; + for (int i = 0; i < GGML_METAL_MAX_COMMAND_BUFFERS; ++i) { + res->cmd_bufs[i].obj = nil; + } + + res->cmd_bufs_ext = [[NSMutableArray alloc] init]; + + res->cmd_buf_last = nil; + + res->pipelines_ext = ggml_metal_pipelines_init(); + + return res; } - - res->has_error = false; - - res->gf = nil; - res->encode_async = nil; - for (int i = 0; i < GGML_METAL_MAX_COMMAND_BUFFERS; ++i) { - res->cmd_bufs[i].obj = nil; - } - - res->cmd_bufs_ext = [[NSMutableArray alloc] init]; - - res->cmd_buf_last = nil; - - res->pipelines_ext = ggml_metal_pipelines_init(); - - return res; } void ggml_metal_free(ggml_metal_t ctx) { diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 4036cc21d..a82caa5e4 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -572,7 +572,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched return res; } -ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_metal_library_t lib, const ggml_tensor * op) { +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_metal_library_t lib, const ggml_tensor * op, bool tail) { GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); char base[256]; @@ -580,7 +580,7 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_me const int nsg = (ne00 + 31)/32; - snprintf(base, 256, "kernel_ssm_scan_%s", ggml_type_name(op->src[0]->type)); + snprintf(base, 256, "kernel_ssm_scan_%s%s", ggml_type_name(op->src[0]->type), tail ? "_tail" : ""); snprintf(name, 256, "%s_nsg=%d", base, nsg); ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); @@ -598,6 +598,27 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan(ggml_me return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan_ssd_mma(ggml_metal_library_t lib, const ggml_tensor * op) { + char base[256]; + char name[256]; + + snprintf(base, 256, "kernel_ssm_scan_ssd_mma_%s", ggml_type_name(op->src[0]->type)); + snprintf(name, 256, "%s", base); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr); + } + + // acs/exp(acs)/state-decay vectors + dtX + SAM rows + two 8x8 tiles per simdgroup + res.smem = (3*OP_SSM_SCAN_SSD_CS + + OP_SSM_SCAN_SSD_CS*OP_SSM_SCAN_SSD_HD + + OP_SSM_SCAN_SSD_NSG*8*OP_SSM_SCAN_SSD_CS + + OP_SSM_SCAN_SSD_NSG*2*8*8)*sizeof(float); + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rwkv(ggml_metal_library_t lib, const ggml_tensor * op) { char base[256]; char name[256]; diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index 6c39428c7..003b688db 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -129,7 +129,8 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc (ggml_metal_library_t lib, enum ggml_op op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs); -struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op, bool tail); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan_ssd_mma (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rwkv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_gated_delta_net (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_solve_tri (ggml_metal_library_t lib, const struct ggml_tensor * op); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index a6ef9cf6c..0cc6350e7 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -778,7 +778,9 @@ void ggml_metal_encoder_free(ggml_metal_encoder_t encoder) { } void ggml_metal_encoder_debug_group_push(ggml_metal_encoder_t encoder, const char * name) { - [encoder->obj pushDebugGroup:[NSString stringWithCString:name encoding:NSUTF8StringEncoding]]; + @autoreleasepool { + [encoder->obj pushDebugGroup:[NSString stringWithCString:name encoding:NSUTF8StringEncoding]]; + } } void ggml_metal_encoder_debug_group_pop (ggml_metal_encoder_t encoder) { @@ -1029,249 +1031,251 @@ ggml_metal_device_t ggml_metal_device_init(int device, int n_devices) { assert(dev != NULL); - if (dev->mtl_device == nil) { - dev->mtl_device = MTLCreateSystemDefaultDevice(); + @autoreleasepool { + if (dev->mtl_device == nil) { + dev->mtl_device = MTLCreateSystemDefaultDevice(); - if (dev->mtl_device) { - dev->mtl_queue = [dev->mtl_device newCommandQueue]; - if (dev->mtl_queue == nil) { - GGML_LOG_ERROR("%s: error: failed to create command queue\n", __func__); - } + if (dev->mtl_device) { + dev->mtl_queue = [dev->mtl_device newCommandQueue]; + if (dev->mtl_queue == nil) { + GGML_LOG_ERROR("%s: error: failed to create command queue\n", __func__); + } - dev->addr_virt = 0x000000400ULL; + dev->addr_virt = 0x000000400ULL; - dev->props.device = device; + dev->props.device = device; - // the Metal backend uses the system default device as the single physical device; - // additional (virtual) devices are emulated on top of it via GGML_METAL_DEVICES - dev->props.device_phys = 0; - dev->props.device_virt = device; + // the Metal backend uses the system default device as the single physical device; + // additional (virtual) devices are emulated on top of it via GGML_METAL_DEVICES + dev->props.device_phys = 0; + dev->props.device_virt = device; - dev->props.has_simdgroup_reduction = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; - dev->props.has_simdgroup_reduction |= [dev->mtl_device supportsFamily:MTLGPUFamilyMetal3_GGML]; + dev->props.has_simdgroup_reduction = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; + dev->props.has_simdgroup_reduction |= [dev->mtl_device supportsFamily:MTLGPUFamilyMetal3_GGML]; - dev->props.has_simdgroup_mm = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; - dev->props.has_unified_memory = dev->mtl_device.hasUnifiedMemory; + dev->props.has_simdgroup_mm = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; + dev->props.has_unified_memory = dev->mtl_device.hasUnifiedMemory; - dev->props.has_bfloat = [dev->mtl_device supportsFamily:MTLGPUFamilyMetal3_GGML]; - dev->props.has_bfloat |= [dev->mtl_device supportsFamily:MTLGPUFamilyApple6]; - if (getenv("GGML_METAL_BF16_DISABLE") != NULL) { - dev->props.has_bfloat = false; - } + dev->props.has_bfloat = [dev->mtl_device supportsFamily:MTLGPUFamilyMetal3_GGML]; + dev->props.has_bfloat |= [dev->mtl_device supportsFamily:MTLGPUFamilyApple6]; + if (getenv("GGML_METAL_BF16_DISABLE") != NULL) { + dev->props.has_bfloat = false; + } - dev->props.has_tensor = [dev->mtl_device supportsFamily:MTLGPUFamilyMetal4_GGML]; - if (getenv("GGML_METAL_TENSOR_DISABLE") != NULL) { - dev->props.has_tensor = false; - } - - // note: disable the tensor API by default for old chips because with the current implementation it is not useful - // - M2 Ultra: ~5% slower - // - M4, M4 Max: no significant difference - // - // TODO: try to update the tensor API kernels to at least match the simdgroup performance - if (getenv("GGML_METAL_TENSOR_ENABLE") == NULL && - ![[dev->mtl_device name] containsString:@"M5"] && - ![[dev->mtl_device name] containsString:@"M6"] && - ![[dev->mtl_device name] containsString:@"A19"] && - ![[dev->mtl_device name] containsString:@"A20"]) { - GGML_LOG_INFO("%s: tensor API disabled for pre-M5 and pre-A19 devices\n", __func__); - dev->props.has_tensor = false; - } - - // double-check that the tensor API compiles - if (dev->props.has_tensor) { - const char * src_tensor_f16 = "\n" - "#include \n" - "#include \n" - "#include \n" - " \n" - "using namespace metal; \n" - "using namespace mpp::tensor_ops; \n" - " \n" - "kernel void dummy_kernel( \n" - " tensor> A [[buffer(0)]], \n" - " tensor> B [[buffer(1)]], \n" - " device float * C [[buffer(2)]], \n" - " uint2 tgid [[threadgroup_position_in_grid]]) \n" - "{ \n" - " auto tA = A.slice(0, (int)tgid.y); \n" - " auto tB = B.slice((int)tgid.x, 0); \n" - " \n" - " matmul2d< \n" - " matmul2d_descriptor(16, 16, dynamic_extent), \n" - " execution_simdgroups<4>> mm; \n" - " \n" - " auto cT = mm.get_destination_cooperative_tensor(); \n" - " \n" - " auto sA = tA.slice(0, 0); \n" - " auto sB = tB.slice(0, 0); \n" - " mm.run(sB, sA, cT); \n" - " \n" - " auto tC = tensor, tensor_inline>(C, dextents(16, 16)); \n" - " \n" - " cT.store(tC); \n" - "}"; - - GGML_LOG_INFO("%s: testing tensor API for f16 support\n", __func__); - ggml_metal_library_t lib = ggml_metal_library_init_from_source(dev, src_tensor_f16, false); - if (lib == NULL) { - GGML_LOG_WARN("%s: - the tensor API is not supported in this environment - disabling\n", __func__); + dev->props.has_tensor = [dev->mtl_device supportsFamily:MTLGPUFamilyMetal4_GGML]; + if (getenv("GGML_METAL_TENSOR_DISABLE") != NULL) { dev->props.has_tensor = false; - } else { - struct ggml_metal_pipeline_with_params ppl = ggml_metal_library_compile_pipeline(lib, "dummy_kernel", "dummy_kernel", nil); - if (!ppl.pipeline) { + } + + // note: disable the tensor API by default for old chips because with the current implementation it is not useful + // - M2 Ultra: ~5% slower + // - M4, M4 Max: no significant difference + // + // TODO: try to update the tensor API kernels to at least match the simdgroup performance + if (getenv("GGML_METAL_TENSOR_ENABLE") == NULL && + ![[dev->mtl_device name] containsString:@"M5"] && + ![[dev->mtl_device name] containsString:@"M6"] && + ![[dev->mtl_device name] containsString:@"A19"] && + ![[dev->mtl_device name] containsString:@"A20"]) { + GGML_LOG_INFO("%s: tensor API disabled for pre-M5 and pre-A19 devices\n", __func__); + dev->props.has_tensor = false; + } + + // double-check that the tensor API compiles + if (dev->props.has_tensor) { + const char * src_tensor_f16 = "\n" + "#include \n" + "#include \n" + "#include \n" + " \n" + "using namespace metal; \n" + "using namespace mpp::tensor_ops; \n" + " \n" + "kernel void dummy_kernel( \n" + " tensor> A [[buffer(0)]], \n" + " tensor> B [[buffer(1)]], \n" + " device float * C [[buffer(2)]], \n" + " uint2 tgid [[threadgroup_position_in_grid]]) \n" + "{ \n" + " auto tA = A.slice(0, (int)tgid.y); \n" + " auto tB = B.slice((int)tgid.x, 0); \n" + " \n" + " matmul2d< \n" + " matmul2d_descriptor(16, 16, dynamic_extent), \n" + " execution_simdgroups<4>> mm; \n" + " \n" + " auto cT = mm.get_destination_cooperative_tensor(); \n" + " \n" + " auto sA = tA.slice(0, 0); \n" + " auto sB = tB.slice(0, 0); \n" + " mm.run(sB, sA, cT); \n" + " \n" + " auto tC = tensor, tensor_inline>(C, dextents(16, 16)); \n" + " \n" + " cT.store(tC); \n" + "}"; + + GGML_LOG_INFO("%s: testing tensor API for f16 support\n", __func__); + ggml_metal_library_t lib = ggml_metal_library_init_from_source(dev, src_tensor_f16, false); + if (lib == NULL) { GGML_LOG_WARN("%s: - the tensor API is not supported in this environment - disabling\n", __func__); dev->props.has_tensor = false; + } else { + struct ggml_metal_pipeline_with_params ppl = ggml_metal_library_compile_pipeline(lib, "dummy_kernel", "dummy_kernel", nil); + if (!ppl.pipeline) { + GGML_LOG_WARN("%s: - the tensor API is not supported in this environment - disabling\n", __func__); + dev->props.has_tensor = false; + } + + ggml_metal_library_free(lib); } - - ggml_metal_library_free(lib); } - } - // try to compile a dummy kernel to determine if the tensor API is supported for bfloat - if (dev->props.has_tensor && dev->props.has_bfloat) { - const char * src_tensor_bf16 = "\n" - "#include \n" - "#include \n" - "#include \n" - " \n" - "using namespace metal; \n" - "using namespace mpp::tensor_ops; \n" - " \n" - "kernel void dummy_kernel( \n" - " tensor> A [[buffer(0)]], \n" - " tensor> B [[buffer(1)]], \n" - " device float * C [[buffer(2)]], \n" - " uint2 tgid [[threadgroup_position_in_grid]]) \n" - "{ \n" - " auto tA = A.slice(0, (int)tgid.y); \n" - " auto tB = B.slice((int)tgid.x, 0); \n" - " \n" - " matmul2d< \n" - " matmul2d_descriptor(16, 16, dynamic_extent), \n" - " execution_simdgroups<4>> mm; \n" - " \n" - " auto cT = mm.get_destination_cooperative_tensor(); \n" - " \n" - " auto sA = tA.slice(0, 0); \n" - " auto sB = tB.slice(0, 0); \n" - " mm.run(sB, sA, cT); \n" - " \n" - " auto tC = tensor, tensor_inline>(C, dextents(16, 16)); \n" - " \n" - " cT.store(tC); \n" - "}"; + // try to compile a dummy kernel to determine if the tensor API is supported for bfloat + if (dev->props.has_tensor && dev->props.has_bfloat) { + const char * src_tensor_bf16 = "\n" + "#include \n" + "#include \n" + "#include \n" + " \n" + "using namespace metal; \n" + "using namespace mpp::tensor_ops; \n" + " \n" + "kernel void dummy_kernel( \n" + " tensor> A [[buffer(0)]], \n" + " tensor> B [[buffer(1)]], \n" + " device float * C [[buffer(2)]], \n" + " uint2 tgid [[threadgroup_position_in_grid]]) \n" + "{ \n" + " auto tA = A.slice(0, (int)tgid.y); \n" + " auto tB = B.slice((int)tgid.x, 0); \n" + " \n" + " matmul2d< \n" + " matmul2d_descriptor(16, 16, dynamic_extent), \n" + " execution_simdgroups<4>> mm; \n" + " \n" + " auto cT = mm.get_destination_cooperative_tensor(); \n" + " \n" + " auto sA = tA.slice(0, 0); \n" + " auto sB = tB.slice(0, 0); \n" + " mm.run(sB, sA, cT); \n" + " \n" + " auto tC = tensor, tensor_inline>(C, dextents(16, 16)); \n" + " \n" + " cT.store(tC); \n" + "}"; - GGML_LOG_INFO("%s: testing tensor API for bfloat support\n", __func__); - ggml_metal_library_t lib = ggml_metal_library_init_from_source(dev, src_tensor_bf16, false); - if (lib == NULL) { - GGML_LOG_WARN("%s: - the tensor API does not support bfloat - disabling bfloat support\n", __func__); - dev->props.has_bfloat = false; - } else { - struct ggml_metal_pipeline_with_params ppl = ggml_metal_library_compile_pipeline(lib, "dummy_kernel", "dummy_kernel", nil); - if (!ppl.pipeline) { + GGML_LOG_INFO("%s: testing tensor API for bfloat support\n", __func__); + ggml_metal_library_t lib = ggml_metal_library_init_from_source(dev, src_tensor_bf16, false); + if (lib == NULL) { GGML_LOG_WARN("%s: - the tensor API does not support bfloat - disabling bfloat support\n", __func__); dev->props.has_bfloat = false; + } else { + struct ggml_metal_pipeline_with_params ppl = ggml_metal_library_compile_pipeline(lib, "dummy_kernel", "dummy_kernel", nil); + if (!ppl.pipeline) { + GGML_LOG_WARN("%s: - the tensor API does not support bfloat - disabling bfloat support\n", __func__); + dev->props.has_bfloat = false; + } + + ggml_metal_library_free(lib); } - - ggml_metal_library_free(lib); } - } - dev->props.use_residency_sets = true; + dev->props.use_residency_sets = true; #if defined(GGML_METAL_HAS_RESIDENCY_SETS) - dev->props.use_residency_sets = getenv("GGML_METAL_NO_RESIDENCY") == nil; + dev->props.use_residency_sets = getenv("GGML_METAL_NO_RESIDENCY") == nil; #endif - dev->props.use_shared_buffers = dev->props.has_unified_memory; + dev->props.use_shared_buffers = dev->props.has_unified_memory; #if TARGET_OS_OSX - // In case of eGPU, shared memory may be preferable. - dev->props.use_shared_buffers |= [dev->mtl_device location] == MTLDeviceLocationExternal; + // In case of eGPU, shared memory may be preferable. + dev->props.use_shared_buffers |= [dev->mtl_device location] == MTLDeviceLocationExternal; #endif - if (getenv("GGML_METAL_SHARED_BUFFERS_DISABLE") != NULL) { - dev->props.use_shared_buffers = false; - } - if (getenv("GGML_METAL_SHARED_BUFFERS_ENABLE") != NULL) { - dev->props.use_shared_buffers = true; - } + if (getenv("GGML_METAL_SHARED_BUFFERS_DISABLE") != NULL) { + dev->props.use_shared_buffers = false; + } + if (getenv("GGML_METAL_SHARED_BUFFERS_ENABLE") != NULL) { + dev->props.use_shared_buffers = true; + } - dev->props.supports_gpu_family_apple7 = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; + dev->props.supports_gpu_family_apple7 = [dev->mtl_device supportsFamily:MTLGPUFamilyApple7]; - dev->props.device_id = ggml_metal_device_id_parse([[dev->mtl_device name] UTF8String]); + dev->props.device_id = ggml_metal_device_id_parse([[dev->mtl_device name] UTF8String]); - dev->props.op_offload_min_batch_size = getenv("GGML_OP_OFFLOAD_MIN_BATCH") ? atoi(getenv("GGML_OP_OFFLOAD_MIN_BATCH")) : 32; + dev->props.op_offload_min_batch_size = getenv("GGML_OP_OFFLOAD_MIN_BATCH") ? atoi(getenv("GGML_OP_OFFLOAD_MIN_BATCH")) : 32; - dev->props.max_buffer_size = dev->mtl_device.maxBufferLength; - dev->props.max_theadgroup_memory_size = dev->mtl_device.maxThreadgroupMemoryLength; - if (@available(macOS 10.12, iOS 16.0, *)) { - dev->props.max_working_set_size = dev->mtl_device.recommendedMaxWorkingSetSize; - } else { - dev->props.max_working_set_size = dev->mtl_device.maxBufferLength; - } + dev->props.max_buffer_size = dev->mtl_device.maxBufferLength; + dev->props.max_theadgroup_memory_size = dev->mtl_device.maxThreadgroupMemoryLength; + if (@available(macOS 10.12, iOS 16.0, *)) { + dev->props.max_working_set_size = dev->mtl_device.recommendedMaxWorkingSetSize; + } else { + dev->props.max_working_set_size = dev->mtl_device.maxBufferLength; + } - snprintf(dev->props.name, sizeof(dev->props.name), "%s%d", "MTL", device); - const char * gpu_name = [[dev->mtl_device name] UTF8String]; - if (n_devices > 1) { - snprintf(dev->props.desc, sizeof(dev->props.desc), "%s (dev p%d/v%d)", - gpu_name, dev->props.device_phys, dev->props.device_virt); - } else { - snprintf(dev->props.desc, sizeof(dev->props.desc), "%s", gpu_name); - } + snprintf(dev->props.name, sizeof(dev->props.name), "%s%d", "MTL", device); + const char * gpu_name = [[dev->mtl_device name] UTF8String]; + if (n_devices > 1) { + snprintf(dev->props.desc, sizeof(dev->props.desc), "%s (dev p%d/v%d)", + gpu_name, dev->props.device_phys, dev->props.device_virt); + } else { + snprintf(dev->props.desc, sizeof(dev->props.desc), "%s", gpu_name); + } - dev->library = ggml_metal_library_init(dev); - if (!dev->library) { - GGML_LOG_ERROR("%s: error: failed to create library\n", __func__); - } + dev->library = ggml_metal_library_init(dev); + if (!dev->library) { + GGML_LOG_ERROR("%s: error: failed to create library\n", __func__); + } - if (dev->props.use_residency_sets) { - dev->rsets = ggml_metal_rsets_init(dev); - } else { - dev->rsets = nil; - } + if (dev->props.use_residency_sets) { + dev->rsets = ggml_metal_rsets_init(dev); + } else { + dev->rsets = nil; + } - // print MTL GPU family: - GGML_LOG_INFO("%s: GPU name: %s (%s)\n", __func__, dev->props.name, dev->props.desc); + // print MTL GPU family: + GGML_LOG_INFO("%s: GPU name: %s (%s)\n", __func__, dev->props.name, dev->props.desc); - // determine max supported GPU family - // https://developer.apple.com/metal/Metal-Shading-Language-Specification.pdf - // https://developer.apple.com/metal/Metal-Feature-Set-Tables.pdf - { - for (int i = MTLGPUFamilyApple1 + 20; i >= MTLGPUFamilyApple1; --i) { - if ([dev->mtl_device supportsFamily:i]) { - dev->props.gpu_family = i - (int) MTLGPUFamilyApple1 + 1; - GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyApple%d (%d)\n", __func__, dev->props.gpu_family, i); - break; + // determine max supported GPU family + // https://developer.apple.com/metal/Metal-Shading-Language-Specification.pdf + // https://developer.apple.com/metal/Metal-Feature-Set-Tables.pdf + { + for (int i = MTLGPUFamilyApple1 + 20; i >= MTLGPUFamilyApple1; --i) { + if ([dev->mtl_device supportsFamily:i]) { + dev->props.gpu_family = i - (int) MTLGPUFamilyApple1 + 1; + GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyApple%d (%d)\n", __func__, dev->props.gpu_family, i); + break; + } + } + + for (int i = MTLGPUFamilyCommon1 + 5; i >= MTLGPUFamilyCommon1; --i) { + if ([dev->mtl_device supportsFamily:i]) { + GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyCommon%d (%d)\n", __func__, i - (int) MTLGPUFamilyCommon1 + 1, i); + break; + } + } + + for (int i = MTLGPUFamilyMetal3_GGML + 5; i >= MTLGPUFamilyMetal3_GGML; --i) { + if ([dev->mtl_device supportsFamily:i]) { + GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyMetal%d (%d)\n", __func__, i - (int) MTLGPUFamilyMetal3_GGML + 3, i); + break; + } } } - for (int i = MTLGPUFamilyCommon1 + 5; i >= MTLGPUFamilyCommon1; --i) { - if ([dev->mtl_device supportsFamily:i]) { - GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyCommon%d (%d)\n", __func__, i - (int) MTLGPUFamilyCommon1 + 1, i); - break; - } - } - - for (int i = MTLGPUFamilyMetal3_GGML + 5; i >= MTLGPUFamilyMetal3_GGML; --i) { - if ([dev->mtl_device supportsFamily:i]) { - GGML_LOG_INFO("%s: GPU family: MTLGPUFamilyMetal%d (%d)\n", __func__, i - (int) MTLGPUFamilyMetal3_GGML + 3, i); - break; - } - } - } - - GGML_LOG_INFO("%s: simdgroup reduction = %s\n", __func__, dev->props.has_simdgroup_reduction ? "true" : "false"); - GGML_LOG_INFO("%s: simdgroup matrix mul. = %s\n", __func__, dev->props.has_simdgroup_mm ? "true" : "false"); - GGML_LOG_INFO("%s: has unified memory = %s\n", __func__, dev->props.has_unified_memory ? "true" : "false"); - GGML_LOG_INFO("%s: has bfloat = %s\n", __func__, dev->props.has_bfloat ? "true" : "false"); - GGML_LOG_INFO("%s: has tensor = %s\n", __func__, dev->props.has_tensor ? "true" : "false"); - GGML_LOG_INFO("%s: use residency sets = %s\n", __func__, dev->props.use_residency_sets ? "true" : "false"); - GGML_LOG_INFO("%s: use shared buffers = %s\n", __func__, dev->props.use_shared_buffers ? "true" : "false"); + GGML_LOG_INFO("%s: simdgroup reduction = %s\n", __func__, dev->props.has_simdgroup_reduction ? "true" : "false"); + GGML_LOG_INFO("%s: simdgroup matrix mul. = %s\n", __func__, dev->props.has_simdgroup_mm ? "true" : "false"); + GGML_LOG_INFO("%s: has unified memory = %s\n", __func__, dev->props.has_unified_memory ? "true" : "false"); + GGML_LOG_INFO("%s: has bfloat = %s\n", __func__, dev->props.has_bfloat ? "true" : "false"); + GGML_LOG_INFO("%s: has tensor = %s\n", __func__, dev->props.has_tensor ? "true" : "false"); + GGML_LOG_INFO("%s: use residency sets = %s\n", __func__, dev->props.use_residency_sets ? "true" : "false"); + GGML_LOG_INFO("%s: use shared buffers = %s\n", __func__, dev->props.use_shared_buffers ? "true" : "false"); #if TARGET_OS_OSX || (TARGET_OS_IOS && __clang_major__ >= 15) - if (@available(macOS 10.12, iOS 16.0, *)) { - GGML_LOG_INFO("%s: recommendedMaxWorkingSetSize = %8.2f MB\n", __func__, dev->props.max_working_set_size / 1e6); - } + if (@available(macOS 10.12, iOS 16.0, *)) { + GGML_LOG_INFO("%s: recommendedMaxWorkingSetSize = %8.2f MB\n", __func__, dev->props.max_working_set_size / 1e6); + } #endif + } } } diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index f0b779979..9becf0479 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -158,6 +158,10 @@ #define OP_SUM_ROWS_NUM_SUM_ROWS 10 #define OP_SUM_ROWS_NUM_MEAN 11 +#define OP_SSM_SCAN_SSD_CS 64 // Metal-specific; Chunk Size; 64 is largest multiple of 8 (simdgroup tile) fitting into 32 KiB Metal threadgroup mem limit (~26.75 KiB shared mem; see smem layout comment in kernel_ssm_scan_ssd_mma_f32) +#define OP_SSM_SCAN_SSD_HD 64 // Metal-specific; Head Dim the MMA kernel is specialized for (Mamba-2); use_mma gates on d_inner == this +#define OP_SSM_SCAN_SSD_NSG 4 // Metal-specific; Number of SimdGroups per threadgroup; NSG*32 == threads dispatched per threadgroup + // kernel argument structs // // - element counters (e.g. ne00) typically use int32_t to reduce register usage @@ -893,6 +897,8 @@ typedef struct { int64_t n_head; int64_t n_group; int64_t n_seq_tokens; + int64_t n_seq_tokens_total; + int64_t token_offset; int64_t n_seqs; int64_t K; uint64_t s_off; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 1c3bb936b..75de0f6dd 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -1677,6 +1677,7 @@ int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) { ggml_metal_library_t lib = ctx->lib; ggml_metal_encoder_t enc = ctx->enc; + const ggml_metal_device_props * props_dev = ggml_metal_device_get_props(ctx->dev); GGML_TENSOR_LOCALS( int32_t, ne0, op->src[0], ne); GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb); @@ -1722,6 +1723,8 @@ int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) { /*.n_head =*/ n_head, /*.n_group =*/ n_group, /*.n_seq_tokens =*/ n_seq_tokens, + /*.n_seq_tokens_total =*/ n_seq_tokens, + /*.token_offset =*/ 0, /*.n_seqs =*/ n_seqs, /*.K =*/ K, /*.s_off =*/ ggml_nelements(op->src[1]) * sizeof(float), @@ -1751,26 +1754,53 @@ int ggml_metal_op_ssm_scan(ggml_metal_op_t ctx, int idx) { /*.nb0 =*/ nb0, }; - auto pipeline = ggml_metal_library_get_pipeline_ssm_scan(lib, op); + constexpr int64_t CHUNK = OP_SSM_SCAN_SSD_CS; - GGML_ASSERT(d_state <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)); + const int64_t snap_reserve = K > 1 ? K : 0; // tokens reserved for sequential kernel rollback snapshots + const int64_t mma_tokens = ((n_seq_tokens - snap_reserve) / CHUNK) * CHUNK; // largest multiple of CHUNK that leaves snap_reserve for the tail + const bool use_mma = + mma_tokens > 0 && + ne30 == 1 && // checks that A tensor is set to scalar decay per head (A shape {1, n_head}) + props_dev->has_simdgroup_mm && // hardware check for M1 or newer + d_state % 8 == 0 && // d_state must be multiple of 8 to align with simdgroup_float 8x8 tiles + d_inner == OP_SSM_SCAN_SSD_HD; // mma kernel is specialized for the Mamba-2 head dim; this checks it - const size_t smem = pipeline.smem; + const auto dispatch = [&](ggml_metal_pipeline_with_params pipeline, int64_t nth, int64_t n_tg_x) { + GGML_ASSERT(nth <= ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)); + GGML_ASSERT(pipeline.smem <= props_dev->max_theadgroup_memory_size); - ggml_metal_encoder_set_pipeline(enc, pipeline); - ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[0]), 1); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 2); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[2]), 3); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[3]), 4); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[4]), 5); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[5]), 6); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[6]), 7); - ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 8); + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[0]), 1); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[1]), 2); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[2]), 3); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[3]), 4); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[4]), 5); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[5]), 6); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op->src[6]), 7); + ggml_metal_encoder_set_buffer (enc, ggml_metal_get_buffer_id(op), 8); + ggml_metal_encoder_set_threadgroup_memory_size(enc, pipeline.smem, 0); + ggml_metal_encoder_dispatch_threadgroups(enc, n_tg_x, n_head, n_seqs, nth, 1, 1); + }; - ggml_metal_encoder_set_threadgroup_memory_size(enc, smem, 0); + if (!use_mma) { + dispatch(ggml_metal_library_get_pipeline_ssm_scan(lib, op, false), d_state, d_inner); + return 1; + } - ggml_metal_encoder_dispatch_threadgroups(enc, d_inner, n_head, n_seqs, d_state, 1, 1); + args.n_seq_tokens = mma_tokens; + dispatch( + ggml_metal_library_get_pipeline_ssm_scan_ssd_mma(lib, op), + OP_SSM_SCAN_SSD_NSG*32, + 1); + + if (mma_tokens < n_seq_tokens) { + ggml_metal_op_concurrency_reset(ctx); + + args.n_seq_tokens = n_seq_tokens - mma_tokens; + args.token_offset = mma_tokens; + dispatch(ggml_metal_library_get_pipeline_ssm_scan(lib, op, true), d_state, d_inner); + } return 1; } diff --git a/ggml/src/ggml-metal/kernels/ssm.metal b/ggml/src/ggml-metal/kernels/ssm.metal index be065c4ff..d3118a831 100644 --- a/ggml/src/ggml-metal/kernels/ssm.metal +++ b/ggml/src/ggml-metal/kernels/ssm.metal @@ -159,7 +159,9 @@ kernel void kernel_ssm_conv_f32_f32_batched_4( // ref: ggml.c:ggml_compute_forward_ssm_scan_f32, Mamba-2 part // Optimized version: reduces redundant memory loads by having one thread load shared values -kernel void kernel_ssm_scan_f32( +// TAIL == false is the whole-sequence / decode path: token_offset folds away at compile time. +template +kernel void kernel_ssm_scan_impl( constant ggml_metal_kargs_ssm_scan & args, device const void * src0, device const void * src1, @@ -200,13 +202,17 @@ kernel void kernel_ssm_scan_f32( const int32_t n_t = args.n_seq_tokens; const int32_t n_s = args.n_seqs; const int32_t K = args.K; + const int32_t n_t_total = TAIL ? args.n_seq_tokens_total : n_t; + const int32_t t_off = TAIL ? args.token_offset : 0; const int32_t s_off = args.s_off; device const int32_t * ids = (device const int32_t *) src6; - device const float * s0_buff = (device const float *) ((device const char *) src0 + ir*args.nb02 + ids[i3]*args.nb03); device float * s_buff = (device float *) ((device char *) dst + ir*args.nb02 + i3*args.nb03 + s_off); + device const float * s0_buff = t_off != 0 ? + s_buff : + (device const float *) ((device const char *) src0 + ir*args.nb02 + ids[i3]*args.nb03); const int32_t i = i0 + i1*nc; const int32_t g = ir / (nh / ng); // repeat_interleave @@ -218,12 +224,12 @@ kernel void kernel_ssm_scan_f32( const float A0 = A[i0%args.ne30]; - device const float * x = (device const float *)((device const char *) src1 + i1*args.nb10 + ir*args.nb11 + i3*args.nb13); // {dim, nh, nt, ns} - device const float * dt = (device const float *)((device const char *) src2 + ir*args.nb20 + i3*args.nb22); // {nh, nt, ns} - device const float * B = (device const float *)((device const char *) src4 + g*args.nb41 + i3*args.nb43); // {d_state, ng, nt, ns} - device const float * C = (device const float *)((device const char *) src5 + g*args.nb51 + i3*args.nb53); // {d_state, ng, nt, ns} + device const float * x = (device const float *)((device const char *) src1 + i1*args.nb10 + ir*args.nb11 + t_off*args.nb12 + i3*args.nb13); // {dim, nh, nt, ns} + device const float * dt = (device const float *)((device const char *) src2 + ir*args.nb20 + t_off*args.nb21 + i3*args.nb22); // {nh, nt, ns} + device const float * B = (device const float *)((device const char *) src4 + g*args.nb41 + t_off*args.nb42 + i3*args.nb43); // {d_state, ng, nt, ns} + device const float * C = (device const float *)((device const char *) src5 + g*args.nb51 + t_off*args.nb52 + i3*args.nb53); // {d_state, ng, nt, ns} - device float * y = dst + (i1 + ir*(nr) + i3*(n_t*nh*nr)); // {dim, nh, nt, ns} + device float * y = dst + (i1 + ir*nr + t_off*nh*nr + i3*(n_t_total*nh*nr)); // {dim, nh, nt, ns} for (int i2 = 0; i2 < n_t; i2 += sgptg) { threadgroup_barrier(mem_flags::mem_threadgroup); @@ -285,3 +291,183 @@ kernel void kernel_ssm_scan_f32( s_buff[i] = s; } + +typedef decltype(kernel_ssm_scan_impl) kernel_ssm_scan_t; + +template [[host_name("kernel_ssm_scan_f32")]] kernel kernel_ssm_scan_t kernel_ssm_scan_impl; +template [[host_name("kernel_ssm_scan_f32_tail")]] kernel kernel_ssm_scan_t kernel_ssm_scan_impl; + +// Chunked SSD SSM scan via Metal simdgroup MMatrix Multiply-Accumulate (simdgroup_float8x8) fast path. +// One threadgroup per (head, sequence) and tokens are processed in chunks. +// C*B^T computed in each chunk one time and reused across the head_dim channel tiles. +kernel void kernel_ssm_scan_ssd_mma_f32( + constant ggml_metal_kargs_ssm_scan & args, + device const void * src0, + device const void * src1, + device const void * src2, + device const void * src3, + device const void * src4, + device const void * src5, + device const void * src6, + device float * dst, + threadgroup float * shared [[threadgroup(0)]], + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiitg[[thread_index_in_threadgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]], + ushort tiisg[[thread_index_in_simdgroup]]) { + constexpr short CS = OP_SSM_SCAN_SSD_CS; + constexpr short TC = 8; // Tile Count of each edge in a simdgroup 8x8 tile + constexpr short HD = OP_SSM_SCAN_SSD_HD; + constexpr short NSG = OP_SSM_SCAN_SSD_NSG; + + // acs/exp(acs)/state-decay vectors, dtX[CS][HD], four private SAM row tiles [8][CS], + // and two 8x8 scratch tiles per simdgroup. Total: 26.75 KiB. + threadgroup float * shared_acs = shared; + threadgroup float * shared_exp_acs = shared + CS; + threadgroup float * shared_state_decay = shared + 2*CS; + threadgroup float * shared_dtx = shared + 3*CS; + threadgroup float * shared_sam = shared + 3*CS + CS*HD; + threadgroup float * sam_rows = shared_sam + sgitg*TC*CS; + threadgroup float * shared_tile = shared_sam + NSG*TC*CS; + threadgroup float * tile0 = shared_tile + sgitg*2*TC*TC; + threadgroup float * tile1 = tile0 + TC*TC; + + const int32_t ir = tgpig.y; // current head + const int32_t i3 = tgpig.z; // current seq + + const int32_t nc = args.d_state; + const int32_t nr = args.d_inner; + const int32_t nh = args.n_head; + const int32_t ng = args.n_group; + const int32_t n_t = args.n_seq_tokens; + const int32_t n_t_total = args.n_seq_tokens_total; + const int32_t g = ir / (nh / ng); + + device const int32_t * ids = (device const int32_t *) src6; + + device const float * s0_buff = (device const float *) ((device const char *) src0 + ir*args.nb02 + ids[i3]*args.nb03); + device float * s_buff = (device float *) ((device char *) dst + ir*args.nb02 + i3*args.nb03 + args.s_off); + + device const float * A = (device const float *) ((device const char *) src3 + ir*args.nb31); + device const float * x = (device const float *) ((device const char *) src1 + ir*args.nb11 + i3*args.nb13); + device const float * dt = (device const float *) ((device const char *) src2 + ir*args.nb20 + i3*args.nb22); + device const float * B = (device const float *) ((device const char *) src4 + g*args.nb41 + i3*args.nb43); + device const float * C = (device const float *) ((device const char *) src5 + g*args.nb51 + i3*args.nb53); + + device float * y = dst + (ir*nr + i3*(n_t_total*nh*nr)); + + for (int32_t t0 = 0; t0 < n_t; t0 += CS) { + for (int32_t idx = tiitg; idx < CS*HD; idx += NSG*N_SIMDWIDTH) { + const int32_t t = idx / HD; + const int32_t c = idx % HD; + const float dt0 = dt[(t0 + t) * (int32_t) args.ns21]; + const float dtsp = dt0 <= 20.0f ? log(1.0f + exp(dt0)) : dt0; + shared_dtx[idx] = x[(t0 + t) * (int32_t) args.ns12 + c] * dtsp; + } + if (tiitg < CS) { + const float dt0 = dt[(t0 + tiitg) * (int32_t) args.ns21]; + const float dtsp = dt0 <= 20.0f ? log(1.0f + exp(dt0)) : dt0; + shared_acs[tiitg] = dtsp * A[0]; + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + if (tiitg == 0) { + float acc = 0.0f; + for (short t = 0; t < CS; ++t) { + acc += shared_acs[t]; + shared_acs[t] = acc; + } + } + threadgroup_barrier(mem_flags::mem_threadgroup); + if (tiitg < CS) { + shared_exp_acs[tiitg] = exp(shared_acs[tiitg]); + shared_state_decay[tiitg] = exp(shared_acs[CS - 1] - shared_acs[tiitg]); + } + threadgroup_barrier(mem_flags::mem_threadgroup); + + device const float * state = t0 == 0 ? s0_buff : s_buff; + + // Build one 8x64 row tile of SAM per simdgroup, then reuse it across every channel tile. + for (short ib = sgitg; ib < CS/TC; ib += NSG) { + for (short jb = 0; jb <= ib; ++jb) { + simdgroup_float8x8 cb = make_filled_simdgroup_matrix(0.0f); + + for (int32_t k0 = 0; k0 < nc; k0 += TC) { + simdgroup_float8x8 mc; + simdgroup_float8x8 mb; + simdgroup_load(mc, C + (t0 + ib*TC)*(int32_t) args.ns52 + k0, args.ns52); + simdgroup_load(mb, B + (t0 + jb*TC)*(int32_t) args.ns42 + k0, args.ns42, 0, true); + simdgroup_multiply_accumulate(cb, mc, mb, cb); + } + + threadgroup float * sam = sam_rows + jb*TC; + simdgroup_store(cb, sam, CS); + simdgroup_barrier(mem_flags::mem_threadgroup); + for (short e = tiisg; e < TC*TC; e += N_SIMDWIDTH) { + const short ri = e / TC; + const short rj = e % TC; + const short i = ib*TC + ri; + const short j = jb*TC + rj; + sam[ri*CS + rj] = j <= i ? + sam[ri*CS + rj] * exp(shared_acs[i] - shared_acs[j]) : 0.0f; + } + simdgroup_barrier(mem_flags::mem_threadgroup); + } + + for (short ch = 0; ch < HD/TC; ++ch) { + simdgroup_float8x8 y_diag = make_filled_simdgroup_matrix(0.0f); + simdgroup_float8x8 y_inter = make_filled_simdgroup_matrix(0.0f); + + for (short jb = 0; jb <= ib; ++jb) { + simdgroup_float8x8 sam; + simdgroup_float8x8 mdtx; + simdgroup_load(sam, sam_rows + jb*TC, CS); + simdgroup_load(mdtx, shared_dtx + jb*TC*HD + ch*TC, HD); + simdgroup_multiply_accumulate(y_diag, sam, mdtx, y_diag); + } + + for (int32_t k0 = 0; k0 < nc; k0 += TC) { + simdgroup_float8x8 mc; + simdgroup_float8x8 ms; + simdgroup_load(mc, C + (t0 + ib*TC)*(int32_t) args.ns52 + k0, args.ns52); + simdgroup_load(ms, state + ch*TC*nc + k0, nc, 0, true); + simdgroup_multiply_accumulate(y_inter, mc, ms, y_inter); + } + + simdgroup_store(y_diag, tile0, TC); + simdgroup_store(y_inter, tile1, TC); + simdgroup_barrier(mem_flags::mem_threadgroup); + for (short e = tiisg; e < TC*TC; e += N_SIMDWIDTH) { + const short ri = e / TC; + const short ci = e % TC; + const int32_t token = t0 + ib*TC + ri; + y[token*nh*nr + ch*TC + ci] = + tile0[e] + shared_exp_acs[ib*TC + ri] * tile1[e]; + } + simdgroup_barrier(mem_flags::mem_threadgroup); + } + } + + // All simdgroups must finish reading s_buff before any thread overwrites it. + threadgroup_barrier(mem_flags::mem_device | mem_flags::mem_threadgroup); + + // Keep the carried-state reduction in token order. Reassociating this particular product + // with MMA compounds rounding differences at every chunk boundary; CB, y_diag, and C*S + // remain on the matrix unit. + const float chunk_decay = exp(shared_acs[CS - 1]); + for (int32_t idx = tiitg; idx < nc*HD; idx += NSG*N_SIMDWIDTH) { + const int32_t ci = idx / nc; + const int32_t si = idx % nc; + float state_c = 0.0f; + for (short t = 0; t < CS; ++t) { + state_c += shared_state_decay[t] * + B[(t0 + t)*(int32_t) args.ns42 + si] * + shared_dtx[t*HD + ci]; + } + s_buff[idx] = chunk_decay * state[idx] + state_c; + } + + // All state tiles must be visible before the next chunk consumes s_buff as S_prev. + threadgroup_barrier(mem_flags::mem_device | mem_flags::mem_threadgroup); + } +} diff --git a/ggml/src/ggml-rpc/ggml-rpc.cpp b/ggml/src/ggml-rpc/ggml-rpc.cpp index c4edd7190..aa2a93808 100644 --- a/ggml/src/ggml-rpc/ggml-rpc.cpp +++ b/ggml/src/ggml-rpc/ggml-rpc.cpp @@ -9,6 +9,9 @@ #include #include #include +#include +#include +#include #include #include #include @@ -17,6 +20,8 @@ #include #include #include +#include +#include static const char * RPC_DEBUG = std::getenv("GGML_RPC_DEBUG"); @@ -75,6 +80,7 @@ enum rpc_cmd { RPC_CMD_DEVICE_COUNT, RPC_CMD_GRAPH_RECOMPUTE, RPC_CMD_MEMSET_TENSOR, + RPC_CMD_NONE, RPC_CMD_COUNT, }; @@ -226,24 +232,24 @@ struct ggml_backend_rpc_buffer_type_context { size_t max_size; }; +class rpc_dispatcher; struct ggml_backend_rpc_context { - std::string endpoint; - uint32_t device; - std::string name; + std::shared_ptr dispatcher; + uint32_t device; + std::string name; }; struct ggml_backend_rpc_buffer_context { - std::shared_ptr sock; - void * base_ptr; - uint64_t remote_ptr; + std::shared_ptr dispatcher; + void * base_ptr; + uint64_t remote_ptr; }; // RPC helper functions // Computes FNV-1a hash of the data -static uint64_t fnv_hash(const uint8_t * data, size_t len) { +static uint64_t fnv_hash(const uint8_t * data, size_t len, uint64_t hash = 0xcbf29ce484222325ULL) { const uint64_t fnv_prime = 0x100000001b3ULL; - uint64_t hash = 0xcbf29ce484222325ULL; for (size_t i = 0; i < len; ++i) { hash ^= data[i]; @@ -256,7 +262,10 @@ static bool send_msg(socket_ptr sock, const void * msg, size_t msg_size) { if (!sock->send_data(&msg_size, sizeof(msg_size))) { return false; } - return sock->send_data(msg, msg_size); + if (!sock->send_data(msg, msg_size)) { + return false; + } + return sock->flush(); } static bool recv_msg(socket_ptr sock, void * msg, size_t msg_size) { @@ -311,7 +320,7 @@ static bool send_rpc_cmd(socket_ptr sock, enum rpc_cmd cmd, const void * input, if (!sock->send_data(input, input_size)) { return false; } - return true; + return sock->flush(); } // RPC request : | rpc_cmd (1 byte) | request_size (8 bytes) | request_data (request_size bytes) | @@ -357,44 +366,248 @@ static bool negotiate_hello(const std::shared_ptr & sock) { return true; } -static std::shared_ptr get_socket(const std::string & endpoint) { - static std::mutex mutex; - std::lock_guard lock(mutex); - static std::unordered_map> sockets; +template +class message_queue { +public: + message_queue() {} - auto it = sockets.find(endpoint); - if (it != sockets.end()) { - if (auto sock = it->second.lock()) { - return sock; + bool push(const T &value) { + std::unique_lock lock(mutex); + if (interrupted) { + return false; } + queue.push(value); + cvar.notify_all(); + return true; } + + bool pop(T* out) { + std::unique_lock lock(mutex); + cvar.wait(lock, [this] { return !queue.empty() || interrupted; }); + if (interrupted) { + return false; + } + *out = queue.front(); + queue.pop(); + return true; + } + + void interrupt() { + std::unique_lock lock(mutex); + interrupted = true; + lock.unlock(); + cvar.notify_all(); + } + +private: + bool interrupted = false; + std::queue queue; + std::mutex mutex; + std::condition_variable cvar; +}; + +class rpc_dispatcher { +public: + rpc_dispatcher() { + } + + void send(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size); + void send(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size, void * output, size_t output_size); + void send_async(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size); + void send_async(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size, void * output, size_t output_size); + + ggml_backend_event_t event_new(ggml_backend_dev_t dev); + void event_free(ggml_backend_event_t event); + void event_synchronize(ggml_backend_event_t event); + void event_record(ggml_backend_event_t event); + void synchronize(); + + void start(const std::string & endpoint); + void work(); + + ~rpc_dispatcher(); + +private: + struct rpc_msg { + rpc_cmd cmd; + std::shared_ptr input; + size_t input_size; + void * output; + size_t output_size; + std::promise completion; + }; + using rpc_msg_ptr = std::shared_ptr; + using rpc_msg_queue = message_queue; + struct rpc_event { + rpc_msg_ptr msg; + std::shared_future sf; + }; + rpc_msg_queue queue; + socket_ptr sock; + std::atomic_bool running; + std::thread thread; +}; + +static void rpc_dispatcher_trampoline(rpc_dispatcher * dispatcher) +{ + dispatcher->work(); +} + +void rpc_dispatcher::send(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size) { + auto msg = std::make_shared(); + msg->cmd = cmd; + msg->input = input; + msg->input_size = input_size; + msg->output = nullptr; + msg->output_size = 0; + GGML_ASSERT(queue.push(msg)); + auto future = msg->completion.get_future(); + future.wait(); +} + +void rpc_dispatcher::send_async(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size) { + auto msg = std::make_shared(); + msg->cmd = cmd; + msg->input = input; + msg->input_size = input_size; + msg->output = nullptr; + msg->output_size = 0; + GGML_ASSERT(queue.push(msg)); +} + +void rpc_dispatcher::send(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size, void * output, size_t output_size) { + auto msg = std::make_shared(); + msg->cmd = cmd; + msg->input = input; + msg->input_size = input_size; + msg->output = output; + msg->output_size = output_size; + GGML_ASSERT(queue.push(msg)); + auto future = msg->completion.get_future(); + future.wait(); +} + +void rpc_dispatcher::send_async(enum rpc_cmd cmd, std::shared_ptr input, size_t input_size, void * output, size_t output_size) { + auto msg = std::make_shared(); + msg->cmd = cmd; + msg->input = input; + msg->input_size = input_size; + msg->output = output; + msg->output_size = output_size; + GGML_ASSERT(queue.push(msg)); +} + +ggml_backend_event_t rpc_dispatcher::event_new(ggml_backend_dev_t dev) { + rpc_event * ev = new rpc_event; + ev->msg = std::make_shared(); + ev->msg->cmd = RPC_CMD_NONE; + ev->sf = ev->msg->completion.get_future().share(); + GGML_ASSERT(queue.push(ev->msg)); + return new ggml_backend_event { + /* .device = */ dev, + /* .context = */ ev, + }; +} + +void rpc_dispatcher::event_free(ggml_backend_event_t event) { + rpc_event * ev = (rpc_event *)event->context; + delete ev; +} + +void rpc_dispatcher::event_synchronize(ggml_backend_event_t event) { + rpc_event * ev = (rpc_event *)event->context; + ev->sf.wait(); +} + +void rpc_dispatcher::event_record(ggml_backend_event_t event) { + rpc_event * ev = (rpc_event *)event->context; + ev->msg = std::make_shared(); + ev->msg->cmd = RPC_CMD_NONE; + ev->sf = ev->msg->completion.get_future().share(); + GGML_ASSERT(queue.push(ev->msg)); +} + +void rpc_dispatcher::synchronize() { + // to ensure all messages are processed, submit dummy message and wait for it to complete + auto msg = std::make_shared(); + msg->cmd = RPC_CMD_NONE; + GGML_ASSERT(queue.push(msg)); + msg->completion.get_future().wait(); +} + +void rpc_dispatcher::start(const std::string & endpoint) { std::string host; int port; if (!parse_endpoint(endpoint, host, port)) { - GGML_LOG_ERROR("Failed to parse endpoint: %s\n", endpoint.c_str()); - return nullptr; + GGML_ABORT("Failed to parse endpoint: %s\n", endpoint.c_str()); + } + if (!rpc_transport_init()) { + GGML_ABORT("RPC transport initialization failed\n"); } - if (!rpc_transport_init()) { - return nullptr; - } - auto sock = socket_t::connect(host.c_str(), port); + sock = socket_t::connect(host.c_str(), port); if (sock == nullptr) { - return nullptr; + GGML_ABORT("Failed to connect to %s\n", endpoint.c_str()); } if (!negotiate_hello(sock)) { - return nullptr; + GGML_ABORT("RPC handshake failed for %s\n", endpoint.c_str()); } LOG_DBG("[%s] connected to %s\n", __func__, endpoint.c_str()); - sockets[endpoint] = sock; - return sock; + running = true; + thread = std::thread(rpc_dispatcher_trampoline, this); +} + +void rpc_dispatcher::work() { + while (running) { + rpc_msg_ptr msg_ptr; + if (!queue.pop(&msg_ptr)) { + break; + } + if (msg_ptr->cmd != RPC_CMD_NONE) { + if (msg_ptr->output) { + bool status = send_rpc_cmd(sock, msg_ptr->cmd, msg_ptr->input.get(), msg_ptr->input_size, msg_ptr->output, msg_ptr->output_size); + RPC_STATUS_ASSERT(status); + } else { + bool status = send_rpc_cmd(sock, msg_ptr->cmd, msg_ptr->input.get(), msg_ptr->input_size); + RPC_STATUS_ASSERT(status); + } + } + msg_ptr->completion.set_value(); + } +} + +rpc_dispatcher::~rpc_dispatcher() { + running = false; + queue.interrupt(); + sock = nullptr; + if (thread.joinable()) { + thread.join(); + } +} + +static std::shared_ptr get_dispatcher(const std::string & endpoint) { + static std::mutex mutex; + std::lock_guard lock(mutex); + static std::unordered_map> dispatchers; + + auto it = dispatchers.find(endpoint); + if (it != dispatchers.end()) { + if (auto dispatcher = it->second.lock()) { + return dispatcher; + } + } + + auto dispatcher = std::make_shared(); + dispatcher->start(endpoint); + dispatchers[endpoint] = dispatcher; + return dispatcher; } static void ggml_backend_rpc_buffer_free_buffer(ggml_backend_buffer_t buffer) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_free_buffer_req request = {ctx->remote_ptr}; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_FREE_BUFFER, &request, sizeof(request), nullptr, 0); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->remote_ptr = ctx->remote_ptr; + ctx->dispatcher->send(RPC_CMD_FREE_BUFFER, request, sizeof(*request)); delete ctx; } @@ -403,10 +616,10 @@ static void * ggml_backend_rpc_buffer_get_base(ggml_backend_buffer_t buffer) { if (ctx->base_ptr != nullptr) { return ctx->base_ptr; } - rpc_msg_buffer_get_base_req request = {ctx->remote_ptr}; + auto request = std::make_shared(); + request->remote_ptr = ctx->remote_ptr; rpc_msg_buffer_get_base_rsp response; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_BUFFER_GET_BASE, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + ctx->dispatcher->send(RPC_CMD_BUFFER_GET_BASE, request, sizeof(*request), &response, sizeof(response)); ctx->base_ptr = reinterpret_cast(response.base_ptr); return ctx->base_ptr; } @@ -463,12 +676,9 @@ static enum ggml_status ggml_backend_rpc_buffer_init_tensor(ggml_backend_buffer_ // Due to bandwidth constraints, we only call the server init tensor functions if necessary. // In particular, only quantized tensors need padding if (ggml_is_quantized(tensor->type) && (tensor->ne[0] % 512 != 0) && (tensor->view_src == nullptr)) { - rpc_msg_init_tensor_req request; - - request.tensor = serialize_tensor(tensor); - - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_INIT_TENSOR, &request, sizeof(request), nullptr, 0); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + ctx->dispatcher->send(RPC_CMD_INIT_TENSOR, request, sizeof(*request)); } return GGML_STATUS_SUCCESS; } @@ -476,27 +686,24 @@ static enum ggml_status ggml_backend_rpc_buffer_init_tensor(ggml_backend_buffer_ static void ggml_backend_rpc_buffer_memset_tensor( ggml_backend_buffer_t buffer, ggml_tensor * tensor, uint8_t value, size_t offset, size_t size) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_memset_tensor_req request = { - /* .tensor = */ serialize_tensor(tensor), - /* .offset = */ offset, - /* .size = */ size, - /* .value = */ value, - }; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_MEMSET_TENSOR, &request, sizeof(request), nullptr, 0); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + request->offset = offset; + request->size = size; + request->value = value; + ctx->dispatcher->send(RPC_CMD_MEMSET_TENSOR, request, sizeof(*request)); } static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggml_tensor * tensor, const void * data, size_t offset, size_t size) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; rpc_tensor rpc_tensor = serialize_tensor(tensor); if (size > HASH_THRESHOLD) { - rpc_msg_set_tensor_hash_req request; - request.tensor = rpc_tensor; - request.offset = offset; - request.hash = fnv_hash((const uint8_t*)data, size); + auto request = std::make_shared(); + request->tensor = rpc_tensor; + request->offset = offset; + request->hash = fnv_hash((const uint8_t*)data, size); rpc_msg_set_tensor_hash_rsp response; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_SET_TENSOR_HASH, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + ctx->dispatcher->send(RPC_CMD_SET_TENSOR_HASH, request, sizeof(*request), &response, sizeof(response)); if (response.result) { // the server has the same data, no need to send it return; @@ -504,22 +711,21 @@ static void ggml_backend_rpc_buffer_set_tensor(ggml_backend_buffer_t buffer, ggm } // input serialization format: | rpc_tensor | offset (8 bytes) | data (size bytes) size_t input_size = sizeof(rpc_tensor) + sizeof(uint64_t) + size; - std::vector input(input_size, 0); - memcpy(input.data(), &rpc_tensor, sizeof(rpc_tensor)); - memcpy(input.data() + sizeof(rpc_tensor), &offset, sizeof(offset)); - memcpy(input.data() + sizeof(rpc_tensor) + sizeof(offset), data, size); - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_SET_TENSOR, input.data(), input.size()); - RPC_STATUS_ASSERT(status); + uint8_t * input = new uint8_t[input_size](); + memcpy(input, &rpc_tensor, sizeof(rpc_tensor)); + memcpy(input + sizeof(rpc_tensor), &offset, sizeof(offset)); + memcpy(input + sizeof(rpc_tensor) + sizeof(offset), data, size); + std::shared_ptr input_ptr(input, std::default_delete()); + ctx->dispatcher->send(RPC_CMD_SET_TENSOR, input_ptr, input_size); } static void ggml_backend_rpc_buffer_get_tensor(ggml_backend_buffer_t buffer, const ggml_tensor * tensor, void * data, size_t offset, size_t size) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_get_tensor_req request; - request.tensor = serialize_tensor(tensor); - request.offset = offset; - request.size = size; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_GET_TENSOR, &request, sizeof(request), data, size); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + request->offset = offset; + request->size = size; + ctx->dispatcher->send(RPC_CMD_GET_TENSOR, request, sizeof(*request), data, size); } static bool ggml_backend_rpc_buffer_cpy_tensor(ggml_backend_buffer_t buffer, const ggml_tensor * src, ggml_tensor * dst) { @@ -529,16 +735,15 @@ static bool ggml_backend_rpc_buffer_cpy_tensor(ggml_backend_buffer_t buffer, con ggml_backend_rpc_buffer_context * src_ctx = (ggml_backend_rpc_buffer_context *)src_buffer->context; ggml_backend_buffer_t dst_buffer = dst->buffer; ggml_backend_rpc_buffer_context * dst_ctx = (ggml_backend_rpc_buffer_context *)dst_buffer->context; - if (src_ctx->sock != dst_ctx->sock) { + if (src_ctx->dispatcher != dst_ctx->dispatcher) { return false; } ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_copy_tensor_req request; - request.src = serialize_tensor(src); - request.dst = serialize_tensor(dst); + auto request = std::make_shared(); + request->src = serialize_tensor(src); + request->dst = serialize_tensor(dst); rpc_msg_copy_tensor_rsp response; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_COPY_TENSOR, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + ctx->dispatcher->send(RPC_CMD_COPY_TENSOR, request, sizeof(*request), &response, sizeof(response)); return response.result; } return false; @@ -546,9 +751,10 @@ static bool ggml_backend_rpc_buffer_cpy_tensor(ggml_backend_buffer_t buffer, con static void ggml_backend_rpc_buffer_clear(ggml_backend_buffer_t buffer, uint8_t value) { ggml_backend_rpc_buffer_context * ctx = (ggml_backend_rpc_buffer_context *)buffer->context; - rpc_msg_buffer_clear_req request = {ctx->remote_ptr, value}; - bool status = send_rpc_cmd(ctx->sock, RPC_CMD_BUFFER_CLEAR, &request, sizeof(request), nullptr, 0); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->remote_ptr = ctx->remote_ptr; + request->value = value; + ctx->dispatcher->send(RPC_CMD_BUFFER_CLEAR, request, sizeof(*request)); } static ggml_backend_buffer_i ggml_backend_rpc_buffer_interface = { @@ -572,15 +778,17 @@ static const char * ggml_backend_rpc_buffer_type_name(ggml_backend_buffer_type_t static ggml_backend_buffer_t ggml_backend_rpc_buffer_type_alloc_buffer(ggml_backend_buffer_type_t buft, size_t size) { ggml_backend_rpc_buffer_type_context * buft_ctx = (ggml_backend_rpc_buffer_type_context *)buft->context; - rpc_msg_alloc_buffer_req request = {buft_ctx->device, size}; + auto request = std::make_shared(); + request->device = buft_ctx->device; + request->size = size; rpc_msg_alloc_buffer_rsp response; - auto sock = get_socket(buft_ctx->endpoint); - bool status = send_rpc_cmd(sock, RPC_CMD_ALLOC_BUFFER, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + + auto dispatcher = get_dispatcher(buft_ctx->endpoint); + dispatcher->send(RPC_CMD_ALLOC_BUFFER, request, sizeof(*request), &response, sizeof(response)); if (response.remote_ptr != 0) { ggml_backend_buffer_t buffer = ggml_backend_buffer_init(buft, ggml_backend_rpc_buffer_interface, - new ggml_backend_rpc_buffer_context{sock, nullptr, response.remote_ptr}, + new ggml_backend_rpc_buffer_context{dispatcher, nullptr, response.remote_ptr}, response.remote_size); return buffer; } else { @@ -588,11 +796,11 @@ static ggml_backend_buffer_t ggml_backend_rpc_buffer_type_alloc_buffer(ggml_back } } -static size_t get_alignment(const std::shared_ptr & sock, uint32_t device) { - rpc_msg_get_alignment_req request = {device}; +static size_t get_alignment(const std::shared_ptr & dispatcher, uint32_t device) { + auto request = std::make_shared(); + request->device = device; rpc_msg_get_alignment_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_GET_ALIGNMENT, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + dispatcher->send(RPC_CMD_GET_ALIGNMENT, request, sizeof(*request), &response, sizeof(response)); return response.alignment; } @@ -601,11 +809,11 @@ static size_t ggml_backend_rpc_buffer_type_get_alignment(ggml_backend_buffer_typ return buft_ctx->alignment; } -static size_t get_max_size(const std::shared_ptr & sock, uint32_t device) { - rpc_msg_get_max_size_req request = {device}; +static size_t get_max_size(const std::shared_ptr & dispatcher, uint32_t device) { + auto request = std::make_shared(); + request->device = device; rpc_msg_get_max_size_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_GET_MAX_SIZE, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + dispatcher->send(RPC_CMD_GET_MAX_SIZE, request, sizeof(*request), &response, sizeof(response)); return response.max_size; } @@ -628,23 +836,63 @@ static size_t ggml_backend_rpc_buffer_type_get_alloc_size(ggml_backend_buffer_ty if (rpc_get) { ggml_backend_rpc_buffer_type_context * buft_ctx = (ggml_backend_rpc_buffer_type_context *)buft->context; - auto sock = get_socket(buft_ctx->endpoint); - rpc_msg_get_alloc_size_req request = { - /*.device =*/ buft_ctx->device, - /*.tensor =*/ serialize_tensor(tensor), - /*.srcs =*/ {}, + // Cache key for calls to read the alloc_size. + // We deliberately exclude src tensor dimensions from the key because: + // 1. For CPU backends, alloc_size = ggml_nbytes(output) regardless of src shapes + // 2. For GPU backends, the reservation graph uses max dimensions, so the + // cached value from reservation is always >= any subsequent request + // 3. Including src dims causes cache misses per-ubatch (e.g. growing KV cache) + // which blocks the main thread behind in-flight GRAPH_COMPUTE commands + struct alloc_size_cache_key { + uint32_t device; + uint32_t type; + uint32_t op; + int32_t op_params[GGML_MAX_OP_PARAMS / sizeof(int32_t)]; + uint32_t ne[GGML_MAX_DIMS]; }; + alloc_size_cache_key key = {}; + key.device = buft_ctx->device; + key.type = tensor->type; + key.op = tensor->op; + memcpy(key.op_params, tensor->op_params, sizeof(key.op_params)); + for (int i = 0; i < GGML_MAX_DIMS; i++) { + key.ne[i] = (uint32_t)tensor->ne[i]; + } + + uint64_t cache_hash = fnv_hash((const uint8_t *)&key, sizeof(key)); + cache_hash = fnv_hash((const uint8_t *)buft_ctx->endpoint.data(), buft_ctx->endpoint.size(), cache_hash); + + // alloc sizes are immutable for a given tensor configuration + static std::mutex cache_mutex; + static std::unordered_map cache; + + { + std::lock_guard lock(cache_mutex); + auto it = cache.find(cache_hash); + if (it != cache.end()) { + return it->second; + } + } + + auto request = std::make_shared(); + request->device = buft_ctx->device; + request->tensor = serialize_tensor(tensor); + // .get_alloc_size could be a function of the tensor's srcs, so we must serialize them as well for (int i = 0; i < GGML_MAX_SRC; i++) { - request.srcs[i] = serialize_tensor(tensor->src[i]); + request->srcs[i] = serialize_tensor(tensor->src[i]); } - // TODO: cache the alloc responses to avoid extra RPC calls? rpc_msg_get_alloc_size_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_GET_ALLOC_SIZE, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + auto dispatcher = get_dispatcher(buft_ctx->endpoint); + dispatcher->send(RPC_CMD_GET_ALLOC_SIZE, request, sizeof(*request), &response, sizeof(response)); + + { + std::lock_guard lock(cache_mutex); + cache[cache_hash] = response.alloc_size; + } return response.alloc_size; } @@ -673,9 +921,44 @@ static void ggml_backend_rpc_free(ggml_backend_t backend) { delete backend; } +static void ggml_backend_rpc_set_tensor_async(ggml_backend_t backend, ggml_tensor * tensor, const void * data, size_t offset, size_t size) { + ggml_backend_rpc_context * ctx = (ggml_backend_rpc_context *)backend->context; + rpc_tensor rpc_tensor = serialize_tensor(tensor); + if (size > HASH_THRESHOLD) { + auto request = std::make_shared(); + request->tensor = rpc_tensor; + request->offset = offset; + request->hash = fnv_hash((const uint8_t*)data, size); + rpc_msg_set_tensor_hash_rsp response; + // TODO: make this async + ctx->dispatcher->send(RPC_CMD_SET_TENSOR_HASH, request, sizeof(*request), &response, sizeof(response)); + if (response.result) { + // the server has the same data, no need to send it + return; + } + } + // input serialization format: | rpc_tensor | offset (8 bytes) | data (size bytes) + size_t input_size = sizeof(rpc_tensor) + sizeof(uint64_t) + size; + uint8_t * input = new uint8_t[input_size](); + memcpy(input, &rpc_tensor, sizeof(rpc_tensor)); + memcpy(input + sizeof(rpc_tensor), &offset, sizeof(offset)); + memcpy(input + sizeof(rpc_tensor) + sizeof(offset), data, size); + std::shared_ptr input_ptr(input, std::default_delete()); + ctx->dispatcher->send_async(RPC_CMD_SET_TENSOR, input_ptr, input_size); +} + +static void ggml_backend_rpc_get_tensor_async(ggml_backend_t backend, const ggml_tensor * tensor, void * data, size_t offset, size_t size) { + ggml_backend_rpc_context * ctx = (ggml_backend_rpc_context *)backend->context; + auto request = std::make_shared(); + request->tensor = serialize_tensor(tensor); + request->offset = offset; + request->size = size; + ctx->dispatcher->send_async(RPC_CMD_GET_TENSOR, request, sizeof(*request), data, size); +} + static void ggml_backend_rpc_synchronize(ggml_backend_t backend) { - GGML_UNUSED(backend); - // this is no-op because we don't have any async operations + ggml_backend_rpc_context * rpc_ctx = (ggml_backend_rpc_context *)backend->context; + rpc_ctx->dispatcher->synchronize(); } static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, std::vector & tensors, std::unordered_set & visited) { @@ -698,7 +981,7 @@ static void add_tensor(ggml_tensor * tensor, const ggml_cgraph * cgraph, std::ve tensors.push_back(result); } -static void serialize_graph(uint32_t device, const ggml_cgraph * cgraph, std::vector & output) { +static uint8_t * serialize_graph(uint32_t device, const ggml_cgraph * cgraph, size_t * output_size) { uint32_t n_nodes = cgraph->n_nodes; std::vector tensors; std::unordered_set visited; @@ -708,9 +991,9 @@ static void serialize_graph(uint32_t device, const ggml_cgraph * cgraph, std::ve // serialization format: // | device (4 bytes) | n_nodes (4 bytes) | nodes (n_nodes * sizeof(uint64_t) | n_tensors (4 bytes) | tensors (n_tensors * sizeof(rpc_tensor)) | uint32_t n_tensors = tensors.size(); - int output_size = 2*sizeof(uint32_t) + n_nodes * sizeof(uint64_t) + sizeof(uint32_t) + n_tensors * sizeof(rpc_tensor); - output.resize(output_size, 0); - uint8_t * dest = output.data(); + *output_size = 2*sizeof(uint32_t) + n_nodes * sizeof(uint64_t) + sizeof(uint32_t) + n_tensors * sizeof(rpc_tensor); + uint8_t * output = new uint8_t[*output_size](); + uint8_t * dest = output; memcpy(dest, &device, sizeof(device)); dest += sizeof(device); memcpy(dest, &n_nodes, sizeof(n_nodes)); @@ -723,6 +1006,7 @@ static void serialize_graph(uint32_t device, const ggml_cgraph * cgraph, std::ve dest += sizeof(n_tensors); rpc_tensor * out_tensors = (rpc_tensor *)dest; memcpy(out_tensors, tensors.data(), n_tensors * sizeof(rpc_tensor)); + return output; } static enum ggml_status ggml_backend_rpc_graph_compute(ggml_backend_t backend, ggml_cgraph * cgraph) { @@ -733,27 +1017,35 @@ static enum ggml_status ggml_backend_rpc_graph_compute(ggml_backend_t backend, g GGML_ASSERT(cgraph->n_nodes > 0); bool reuse = cgraph->uid != 0 && rpc_dev_ctx->last_graph_uid == cgraph->uid; if (reuse) { - rpc_msg_graph_recompute_req request; - request.device = rpc_ctx->device; - auto sock = get_socket(rpc_ctx->endpoint); - bool status = send_rpc_cmd(sock, RPC_CMD_GRAPH_RECOMPUTE, &request, sizeof(request)); - RPC_STATUS_ASSERT(status); + auto request = std::make_shared(); + request->device = rpc_ctx->device; + rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_RECOMPUTE, request, sizeof(*request)); } else { rpc_dev_ctx->last_graph_uid = cgraph->uid; - std::vector input; - serialize_graph(rpc_ctx->device, cgraph, input); - auto sock = get_socket(rpc_ctx->endpoint); - bool status = send_rpc_cmd(sock, RPC_CMD_GRAPH_COMPUTE, input.data(), input.size()); - RPC_STATUS_ASSERT(status); + size_t input_size = 0; + uint8_t * input = serialize_graph(rpc_ctx->device, cgraph, &input_size); + std::shared_ptr input_ptr(input, std::default_delete()); + rpc_ctx->dispatcher->send_async(RPC_CMD_GRAPH_COMPUTE, input_ptr, input_size); } return GGML_STATUS_SUCCESS; } +static void ggml_backend_rpc_event_record(ggml_backend_t backend, ggml_backend_event_t event) { + ggml_backend_rpc_context * rpc_ctx = (ggml_backend_rpc_context *)backend->context; + rpc_ctx->dispatcher->event_record(event); +} + +static void ggml_backend_rpc_event_wait(ggml_backend_t backend, ggml_backend_event_t event) { + // this is noop for RPC as we have a single stream + GGML_UNUSED(backend); + GGML_UNUSED(event); +} + static ggml_backend_i ggml_backend_rpc_interface = { /* .get_name = */ ggml_backend_rpc_name, /* .free = */ ggml_backend_rpc_free, - /* .set_tensor_async = */ NULL, - /* .get_tensor_async = */ NULL, + /* .set_tensor_async = */ ggml_backend_rpc_set_tensor_async, + /* .get_tensor_async = */ ggml_backend_rpc_get_tensor_async, /* .set_tensor_2d_async = */ NULL, /* .get_tensor_2d_async = */ NULL, /* .cpy_tensor_async = */ NULL, @@ -763,8 +1055,8 @@ static ggml_backend_i ggml_backend_rpc_interface = { /* .graph_plan_update = */ NULL, /* .graph_plan_compute = */ NULL, /* .graph_compute = */ ggml_backend_rpc_graph_compute, - /* .event_record = */ NULL, - /* .event_wait = */ NULL, + /* .event_record = */ ggml_backend_rpc_event_record, + /* .event_wait = */ ggml_backend_rpc_event_wait, /* .graph_optimize = */ NULL, }; @@ -778,13 +1070,9 @@ ggml_backend_buffer_type_t ggml_backend_rpc_buffer_type(const char * endpoint, u if (it != buft_map.end()) { return it->second; } - auto sock = get_socket(endpoint); - if (sock == nullptr) { - GGML_LOG_ERROR("Failed to connect to %s\n", endpoint); - return nullptr; - } - size_t alignment = get_alignment(sock, device); - size_t max_size = get_max_size(sock, device); + auto dispatcher = get_dispatcher(endpoint); + size_t alignment = get_alignment(dispatcher, device); + size_t max_size = get_max_size(dispatcher, device); ggml_backend_rpc_buffer_type_context * buft_ctx = new ggml_backend_rpc_buffer_type_context { /* .endpoint = */ endpoint, /* .device = */ device, @@ -804,10 +1092,11 @@ ggml_backend_buffer_type_t ggml_backend_rpc_buffer_type(const char * endpoint, u ggml_backend_t ggml_backend_rpc_init(const char * endpoint, uint32_t device) { std::string dev_name = "RPC" + std::to_string(device) + "[" + std::string(endpoint) + "]"; + auto dispatcher = get_dispatcher(endpoint); ggml_backend_rpc_context * ctx = new ggml_backend_rpc_context { - /* .endpoint = */ endpoint, - /* .device = */ device, - /* .name = */ dev_name, + /* .dispatcher = */ dispatcher, + /* .device = */ device, + /* .name = */ dev_name, }; auto reg = ggml_backend_rpc_add_server(endpoint); ggml_backend_t backend = new ggml_backend { @@ -823,26 +1112,16 @@ bool ggml_backend_is_rpc(ggml_backend_t backend) { return backend != NULL && ggml_guid_matches(backend->guid, ggml_backend_rpc_guid()); } -static void get_device_memory(const std::shared_ptr & sock, uint32_t device, size_t * free, size_t * total) { - rpc_msg_get_device_memory_req request; - request.device = device; +void ggml_backend_rpc_get_device_memory(const char * endpoint, uint32_t device, size_t * free, size_t * total) { + auto dispatcher = get_dispatcher(endpoint); + auto request = std::make_shared(); + request->device = device; rpc_msg_get_device_memory_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_GET_DEVICE_MEMORY, &request, sizeof(request), &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + dispatcher->send(RPC_CMD_GET_DEVICE_MEMORY, request, sizeof(*request), &response, sizeof(response)); *free = response.free_mem; *total = response.total_mem; } -void ggml_backend_rpc_get_device_memory(const char * endpoint, uint32_t device, size_t * free, size_t * total) { - auto sock = get_socket(endpoint); - if (sock == nullptr) { - *free = 0; - *total = 0; - return; - } - get_device_memory(sock, device, free, total); -} - // RPC server-side implementation class rpc_server { @@ -1647,9 +1926,6 @@ static void rpc_serve_client(const std::vector & backends, const if (!server.free_buffer(request)) { return; } - if (!send_msg(sock, nullptr, 0)) { - return; - } break; } case RPC_CMD_BUFFER_CLEAR: { @@ -1660,9 +1936,6 @@ static void rpc_serve_client(const std::vector & backends, const if (!server.buffer_clear(request)) { return; } - if (!send_msg(sock, nullptr, 0)) { - return; - } break; } case RPC_CMD_MEMSET_TENSOR: { @@ -1673,9 +1946,6 @@ static void rpc_serve_client(const std::vector & backends, const if (!server.memset_tensor(request)) { return; } - if (!send_msg(sock, nullptr, 0)) { - return; - } break; } case RPC_CMD_SET_TENSOR: { @@ -1710,9 +1980,6 @@ static void rpc_serve_client(const std::vector & backends, const if (!server.init_tensor(request)) { return; } - if (!send_msg(sock, nullptr, 0)) { - return; - } break; } case RPC_CMD_GET_TENSOR: { @@ -1889,10 +2156,10 @@ static void ggml_backend_rpc_device_get_props(ggml_backend_dev_t dev, struct ggm props->type = ggml_backend_rpc_device_get_type(dev); ggml_backend_rpc_device_get_memory(dev, &props->memory_free, &props->memory_total); props->caps = { - /* .async = */ false, + /* .async = */ true, /* .host_buffer = */ false, /* .buffer_from_host_ptr = */ false, - /* .events = */ false, + /* .events = */ true, /* .mmap_support = */ true, }; } @@ -1929,6 +2196,24 @@ static bool ggml_backend_rpc_device_supports_buft(ggml_backend_dev_t dev, ggml_b return buft_ctx->endpoint == dev_ctx->endpoint && buft_ctx->device == dev_ctx->device; } +static ggml_backend_event_t ggml_backend_rpc_device_event_new(ggml_backend_dev_t dev) { + ggml_backend_rpc_device_context * ctx = (ggml_backend_rpc_device_context *)dev->context; + auto dispatcher = get_dispatcher(ctx->endpoint); + return dispatcher->event_new(dev); +} + +static void ggml_backend_rpc_device_event_free(ggml_backend_dev_t dev, ggml_backend_event_t event) { + ggml_backend_rpc_device_context * ctx = (ggml_backend_rpc_device_context *)dev->context; + auto dispatcher = get_dispatcher(ctx->endpoint); + dispatcher->event_free(event); +} + +static void ggml_backend_rpc_device_event_synchronize(ggml_backend_dev_t dev, ggml_backend_event_t event) { + ggml_backend_rpc_device_context * ctx = (ggml_backend_rpc_device_context *)dev->context; + auto dispatcher = get_dispatcher(ctx->endpoint); + dispatcher->event_synchronize(event); +} + static const struct ggml_backend_device_i ggml_backend_rpc_device_i = { /* .get_name = */ ggml_backend_rpc_device_get_name, /* .get_description = */ ggml_backend_rpc_device_get_description, @@ -1942,9 +2227,9 @@ static const struct ggml_backend_device_i ggml_backend_rpc_device_i = { /* .supports_op = */ ggml_backend_rpc_device_supports_op, /* .supports_buft = */ ggml_backend_rpc_device_supports_buft, /* .offload_op = */ NULL, - /* .event_new = */ NULL, - /* .event_free = */ NULL, - /* .event_synchronize = */ NULL, + /* .event_new = */ ggml_backend_rpc_device_event_new, + /* .event_free = */ ggml_backend_rpc_device_event_free, + /* .event_synchronize = */ ggml_backend_rpc_device_event_synchronize, }; // backend reg interface @@ -2004,14 +2289,9 @@ ggml_backend_reg_t ggml_backend_rpc_reg(void) { } static uint32_t ggml_backend_rpc_get_device_count(const char * endpoint) { - auto sock = get_socket(endpoint); - if (sock == nullptr) { - GGML_LOG_ERROR("Failed to connect to %s\n", endpoint); - return 0; - } + auto dispatcher = get_dispatcher(endpoint); rpc_msg_device_count_rsp response; - bool status = send_rpc_cmd(sock, RPC_CMD_DEVICE_COUNT, nullptr, 0, &response, sizeof(response)); - RPC_STATUS_ASSERT(status); + dispatcher->send(RPC_CMD_DEVICE_COUNT, nullptr, 0, &response, sizeof(response)); return response.device_count; } diff --git a/ggml/src/ggml-rpc/transport-apple.cpp b/ggml/src/ggml-rpc/transport-apple.cpp new file mode 100644 index 000000000..c8be77a6d --- /dev/null +++ b/ggml/src/ggml-rpc/transport-apple.cpp @@ -0,0 +1,470 @@ +#include "transport-apple.h" +#include "transport.h" +#include "ggml-impl.h" + +#include + +#include +#include +#include +#include +#include +#include +#include + +// Apple RDMA-over-Thunderbolt (see Apple TN3205). +// +// Apple's RDMA is quite different from what's supported in Linux - deserving of its own transport implementation. +// see https://developer.apple.com/documentation/technotes/tn3205-low-latency-communication-with-rdma-over-thunderbolt for details +// at a high level the main differences are: +// UC(unreliable connection) on Apple vs RC(reliable connection) QP transport types on Linux (though in practice UC on Apple is still lossless) +// fixed 128KiB stride on Apple vs variable chunk size on Linux +// relying on Apple's hardware credit based flow control vs RNR NAKs + retries on Linux +// +// on Apple a SEND and its corresponding RECV must cover the same number of 4 KiB Thunderbolt frames, +// so every SEND posts a whole 128KiB stride over the wire, even when partially filled. +// (In testing 128KiB was the best performing among 32, 64, 128, 256) + +static constexpr uint32_t RDMA_SEG_MAGIC = 0x52534547u; // "RSEG" +static constexpr int RDMA_NBUF = 16; // ring depth (frames per direction) +static constexpr size_t RDMA_FRAME = 4096; // Thunderbolt frame (fixed on Apple) +static constexpr size_t RDMA_STRIDE = 128 * 1024; // 32 Thunderbolt frames; NBUF x this = 2 MiB pinned per direction +static constexpr uint32_t RDMA_PSN = 0; // any value works if both sides match: UC has no retransmit +static constexpr size_t RDMA_GID_SIZE = 16; + +static_assert(RDMA_STRIDE % RDMA_FRAME == 0, "RDMA_STRIDE must be a whole number of frames"); +// TN3205 counts queue depth in Thunderbolt frames, not work requests. +static constexpr uint32_t RDMA_QP_WR = (uint32_t)RDMA_NBUF * (RDMA_STRIDE / RDMA_FRAME); +static constexpr uint64_t RDMA_RECV_WR = 1ull << 20; // wr_id bit tagging recv completions +static constexpr uint64_t RDMA_WR_IDX_MASK = 0xffff; // buffer index in the low bits of wr_id +static constexpr uint8_t RDMA_SYNC_READY = 0x2A; // readiness-handshake byte (peer activated) + +struct rdma_seg_hdr { + uint32_t magic; // RDMA_SEG_MAGIC; a mismatch means the stream desynced + uint32_t len; // payload bytes in this frame; the rest of the stride is padding +}; +static constexpr size_t RDMA_PAYLOAD = RDMA_STRIDE - sizeof(rdma_seg_hdr); + +struct apple_rdma_caps { + uint32_t qpn; + uint16_t lid; + uint16_t reserved; + uint8_t gid[RDMA_GID_SIZE]; +}; + +static_assert(sizeof(apple_rdma_caps) == RPC_CONN_CAPS_SIZE, "apple_rdma_caps must match conn_caps size"); + +struct apple_rdma::impl { + int fd = -1; // bootstrap TCP socket, kept as the liveness anchor + + struct ibv_context * ctx = nullptr; + struct ibv_pd * pd = nullptr; + struct ibv_cq * cq = nullptr; // one CQ for both directions; RDMA_RECV_WR tags recv completions + struct ibv_qp * qp = nullptr; + + uint8_t * send_mem = nullptr; + struct ibv_mr * send_mr = nullptr; + uint8_t * recv_mem = nullptr; + struct ibv_mr * recv_mr = nullptr; + + int send_busy[RDMA_NBUF] = {}; // 1 while this buffer has a send in flight + // completed recv frames, oldest first: ring index, bytes already handed to + // the reader, and total payload length + struct { int buf; uint32_t off; uint32_t len; } inq[RDMA_NBUF] = {}; + int inq_head = 0; + int inq_count = 0; + int pend_buf = -1; + uint32_t pend_len = 0; + bool broken = false; + + uint32_t qpn = 0; + uint8_t port = 0; + int gid_idx = 0; + enum ibv_mtu path_mtu = IBV_MTU_1024; + + int progress(); + bool acquire_pending(); + bool post_pending(); + + bool post_recv(int i) { + struct ibv_sge sge = {}; + sge.addr = (uintptr_t)(recv_mem + (size_t)i * RDMA_STRIDE); + sge.length = (uint32_t)RDMA_STRIDE; + sge.lkey = recv_mr->lkey; + struct ibv_recv_wr wr = {}, * bad = nullptr; + wr.wr_id = RDMA_RECV_WR | (uint64_t)i; + wr.sg_list = &sge; + wr.num_sge = 1; + return ibv_post_recv(qp, &wr, &bad) == 0; + } + + bool post_send(int i, size_t len) { + struct ibv_sge sge = {}; + sge.addr = (uintptr_t)(send_mem + (size_t)i * RDMA_STRIDE); + sge.length = (uint32_t)len; + sge.lkey = send_mr->lkey; + struct ibv_send_wr wr = {}, * bad = nullptr; + wr.wr_id = (uint64_t)i; + wr.sg_list = &sge; + wr.num_sge = 1; + wr.opcode = IBV_WR_SEND; + wr.send_flags = IBV_SEND_SIGNALED; + return ibv_post_send(qp, &wr, &bad) == 0; + } + + ~impl() { + broken = true; + // the QP must be destroyed before the memory it can still write to is + // deregistered and freed: ERR only starts flushing the posted WQEs + if (qp) { + struct ibv_qp_attr a = {}; + a.qp_state = IBV_QPS_ERR; + ibv_modify_qp(qp, &a, IBV_QP_STATE); + struct ibv_wc wc[RDMA_NBUF * 2]; + while (ibv_poll_cq(cq, RDMA_NBUF * 2, wc) > 0) {} + ibv_destroy_qp(qp); + } + if (send_mr) ibv_dereg_mr(send_mr); + if (recv_mr) ibv_dereg_mr(recv_mr); + free(send_mem); + free(recv_mem); + if (cq) ibv_destroy_cq(cq); + if (pd) ibv_dealloc_pd(pd); + if (ctx) ibv_close_device(ctx); + } +}; + +apple_rdma::apple_rdma(std::unique_ptr p) : pimpl(std::move(p)) {} + +apple_rdma::~apple_rdma() = default; + +bool apple_rdma::broken() const { + return pimpl->broken; +} + +// The readiness handshake below still runs over the bootstrap socket, one byte +// each way, before the transport is declared live. +static bool tcp_send_byte(int fd, uint8_t b) { + ssize_t n; + do { n = ::send(fd, &b, sizeof(b), 0); } while (n < 0 && errno == EINTR); + return n == sizeof(b); +} + +static bool tcp_recv_byte(int fd, uint8_t * b) { + ssize_t n; + do { n = ::recv(fd, b, sizeof(*b), 0); } while (n < 0 && errno == EINTR); + return n == (ssize_t)sizeof(*b); +} + +// Index of the GID on this port equal to the target, or -1. Thunderbolt GIDs are +// RoCEv2 IPv4-mapped (::ffff:a.b.c.d), so this matches the local TCP address. +static int rdma_match_gid(struct ibv_context * ctx, uint8_t port, int gid_tbl_len, + const uint8_t * target, union ibv_gid * out) { + for (int i = 0; i < gid_tbl_len; i++) { + union ibv_gid g; + if (ibv_query_gid(ctx, port, i, &g) != 0) continue; + if (memcmp(g.raw, target, RDMA_GID_SIZE) != 0) continue; + if (out) *out = g; + return i; + } + return -1; +} + +// First ACTIVE port on the device. Only a cabled, up Thunderbolt link reports +// ACTIVE, and it is not always port 1, so the port cannot be hardcoded the way +// the Linux path does. Returns 0 if none. +static uint8_t rdma_first_active_port(struct ibv_context * ctx, struct ibv_port_attr * out) { + struct ibv_device_attr da; + if (ibv_query_device(ctx, &da) != 0) return 0; + for (uint8_t p = 1; p <= da.phys_port_cnt; p++) { + struct ibv_port_attr pa; + if (ibv_query_port(ctx, p, &pa) != 0) continue; + if (pa.state == IBV_PORT_ACTIVE) { if (out) *out = pa; return p; } + } + return 0; +} + +// Called before the endpoints are exchanged: pick the local device facing this +// peer, create a UC QP and register the frame rings. RDMA is point-to-point, so +// the device is the one whose GID equals the bootstrap connection's local +// address, i.e. the one cabled to the peer. +std::unique_ptr apple_rdma::probe(int fd, const uint8_t * target_gid, uint8_t * caps) { + int ndev = 0; + ibv_device ** devs = ibv_get_device_list(&ndev); + if (!devs) return nullptr; + + ibv_context * ctx = nullptr; + uint8_t port = 0; + struct ibv_port_attr pa = {}; + union ibv_gid gid = {}; + int gid_idx = -1; + std::string matched; + for (int d = 0; d < ndev; d++) { + ibv_context * c = ibv_open_device(devs[d]); + if (!c) continue; + struct ibv_port_attr p = {}; + uint8_t pt = rdma_first_active_port(c, &p); + int gi = pt ? rdma_match_gid(c, pt, p.gid_tbl_len, target_gid, &gid) : -1; + if (gi < 0) { ibv_close_device(c); continue; } + ctx = c; port = pt; pa = p; gid_idx = gi; + const char * name = ibv_get_device_name(devs[d]); + matched = name ? name : ""; + break; + } + ibv_free_device_list(devs); + if (!ctx) return nullptr; + + std::unique_ptr c(new impl()); + c->fd = fd; + c->ctx = ctx; + c->port = port; + c->gid_idx = gid_idx; + c->path_mtu = pa.active_mtu; + + c->pd = ibv_alloc_pd(ctx); + if (!c->pd) return nullptr; + + c->cq = ibv_create_cq(ctx, 2 * RDMA_QP_WR + 1, nullptr, nullptr, 0); + if (!c->cq) return nullptr; + + ibv_qp_init_attr qia = {}; + qia.send_cq = c->cq; + qia.recv_cq = c->cq; + qia.qp_type = IBV_QPT_UC; + qia.cap.max_send_wr = RDMA_QP_WR; + qia.cap.max_recv_wr = RDMA_QP_WR; + qia.cap.max_send_sge = 1; + qia.cap.max_recv_sge = 1; + c->qp = ibv_create_qp(c->pd, &qia); + if (!c->qp) return nullptr; + + { + ibv_qp_attr a = {}; + a.qp_state = IBV_QPS_INIT; + a.pkey_index = 0; + a.port_num = port; + a.qp_access_flags = IBV_ACCESS_LOCAL_WRITE | IBV_ACCESS_REMOTE_READ | IBV_ACCESS_REMOTE_WRITE; + if (ibv_modify_qp(c->qp, &a, + IBV_QP_STATE | IBV_QP_PKEY_INDEX | IBV_QP_PORT | IBV_QP_ACCESS_FLAGS) != 0) { + return nullptr; + } + } + + long page = sysconf(_SC_PAGESIZE); + if (page <= 0) page = 4096; + const size_t ring_bytes = (size_t)RDMA_NBUF * RDMA_STRIDE; + if (posix_memalign((void **)&c->send_mem, (size_t)page, ring_bytes) != 0) c->send_mem = nullptr; + if (posix_memalign((void **)&c->recv_mem, (size_t)page, ring_bytes) != 0) c->recv_mem = nullptr; + if (!c->send_mem || !c->recv_mem) return nullptr; + + // Apple's provider rejects LOCAL_WRITE-only MRs even for two-sided SEND/RECV. + const int mr_flags = IBV_ACCESS_LOCAL_WRITE | IBV_ACCESS_REMOTE_READ | IBV_ACCESS_REMOTE_WRITE; + c->send_mr = ibv_reg_mr(c->pd, c->send_mem, ring_bytes, mr_flags); + c->recv_mr = ibv_reg_mr(c->pd, c->recv_mem, ring_bytes, mr_flags); + if (!c->send_mr || !c->recv_mr) return nullptr; + + // Recvs are posted in activate() after the RTS transition, not here: Apple's + // provider rejects ibv_post_recv on a QP that has not reached RTS. + + c->qpn = c->qp->qp_num; + + apple_rdma_caps rc = {}; + rc.qpn = c->qpn; + rc.lid = pa.lid; + memcpy(rc.gid, gid.raw, RDMA_GID_SIZE); + memcpy(caps, &rc, sizeof(rc)); + + GGML_LOG_INFO("RDMA(Apple/UC) probed: dev=%s port=%u gid=%d qpn=%u lid=%u mtu=%d ring=%d x %zu KiB\n", + matched.c_str(), port, gid_idx, c->qpn, (unsigned)pa.lid, 128 << c->path_mtu, + RDMA_NBUF, RDMA_STRIDE / 1024); + return std::unique_ptr(new apple_rdma(std::move(c))); +} + +// Called once the peer's endpoint has arrived: INIT -> RTR -> RTS (UC: GID/GRH +// addressing, no timeout/retry/rnr/rd_atomic), then the readiness handshake. +bool apple_rdma::activate(const uint8_t * caps) { + impl * c = pimpl.get(); + + apple_rdma_caps rc = {}; + memcpy(&rc, caps, sizeof(rc)); + + bool ok = true; + { + ibv_qp_attr a = {}; + a.qp_state = IBV_QPS_RTR; + a.path_mtu = c->path_mtu; + a.rq_psn = RDMA_PSN; + a.dest_qp_num = rc.qpn; + a.ah_attr.is_global = 1; + a.ah_attr.port_num = c->port; + a.ah_attr.sl = 0; + a.ah_attr.src_path_bits = 0; + a.ah_attr.dlid = rc.lid; + a.ah_attr.grh.hop_limit = 1; + a.ah_attr.grh.sgid_index = (uint8_t)c->gid_idx; + memcpy(&a.ah_attr.grh.dgid, rc.gid, RDMA_GID_SIZE); + if (ibv_modify_qp(c->qp, &a, + IBV_QP_STATE | IBV_QP_AV | IBV_QP_PATH_MTU | IBV_QP_DEST_QPN | IBV_QP_RQ_PSN) != 0) { + GGML_LOG_ERROR("RDMA(Apple/UC) RTR failed: %s\n", strerror(errno)); + ok = false; + } + } + if (ok) { + ibv_qp_attr a = {}; + a.qp_state = IBV_QPS_RTS; + a.sq_psn = RDMA_PSN; + if (ibv_modify_qp(c->qp, &a, IBV_QP_STATE | IBV_QP_SQ_PSN) != 0) { + GGML_LOG_ERROR("RDMA(Apple/UC) RTS failed: %s\n", strerror(errno)); + ok = false; + } + } + + // Recvs are posted only now: the controller starts processing them at RTR. + for (int i = 0; ok && i < RDMA_NBUF; i++) { + if (!c->post_recv(i)) { + GGML_LOG_ERROR("RDMA(Apple/UC) post_recv %d/%d failed\n", i, RDMA_NBUF); + ok = false; + } + } + + // A queue pair processes receives only after RTR and the transitions above can + // fail on one side alone, so neither peer sends a frame until both report their + // recvs posted. + uint8_t peer_ready = 0; + if (!tcp_send_byte(c->fd, ok ? RDMA_SYNC_READY : 0) || !tcp_recv_byte(c->fd, &peer_ready)) { + return false; + } + if (!ok || peer_ready != RDMA_SYNC_READY) { + return false; + } + + GGML_LOG_INFO("RDMA(Apple/UC) activated: qpn=%u->%u mtu=%d rx_depth=%d\n", + c->qpn, rc.qpn, 128 << c->path_mtu, RDMA_NBUF); + return true; +} + +// Drain the CQ: release completed send buffers, queue completed recv frames for +// the reader. Returns the number of completions reaped, or -1 on error. +int apple_rdma::impl::progress() { + struct ibv_wc wc[RDMA_NBUF * 2]; + int n = ibv_poll_cq(cq, RDMA_NBUF * 2, wc); + if (n < 0) { GGML_LOG_ERROR("RDMA(Apple/UC) poll_cq failed\n"); broken = true; return -1; } + for (int j = 0; j < n; j++) { + uint64_t id = wc[j].wr_id; + bool is_recv = (id & RDMA_RECV_WR) != 0; + if (wc[j].status != IBV_WC_SUCCESS) { + GGML_LOG_ERROR("RDMA(Apple/UC) %s wc error: status=%d\n", is_recv ? "recv" : "send", wc[j].status); + broken = true; + return -1; + } + if (is_recv) { + int b = (int)(id & RDMA_WR_IDX_MASK); + const rdma_seg_hdr * h = (const rdma_seg_hdr *)(recv_mem + (size_t)b * RDMA_STRIDE); + if (h->magic != RDMA_SEG_MAGIC) { GGML_LOG_ERROR("RDMA(Apple/UC) bad frame magic\n"); broken = true; return -1; } + if (h->len > RDMA_PAYLOAD) { GGML_LOG_ERROR("RDMA(Apple/UC) frame len %u exceeds payload\n", h->len); broken = true; return -1; } + int slot = (inq_head + inq_count) % RDMA_NBUF; + inq[slot].buf = b; + inq[slot].off = 0; + inq[slot].len = h->len; + inq_count++; + } else { + send_busy[(int)(id & RDMA_WR_IDX_MASK)] = 0; + } + } + return n; +} + +// Reserve a free send buffer to coalesce into, waiting on progress if none free. +bool apple_rdma::impl::acquire_pending() { + if (pend_buf >= 0) return true; + for (;;) { + if (broken) return false; + for (int k = 0; k < RDMA_NBUF; k++) if (!send_busy[k]) { pend_buf = k; pend_len = 0; return true; } + if (progress() < 0) return false; + } +} + +// Post the pending frame. The whole STRIDE goes out even when only partly filled: +// TN3205 requires a SEND and its matching RECV to cover the same number of +// Thunderbolt frames, so a short send would fail the peer's receive. +bool apple_rdma::impl::post_pending() { + if (pend_buf < 0) return true; + int i = pend_buf; + rdma_seg_hdr * h = (rdma_seg_hdr *)(send_mem + (size_t)i * RDMA_STRIDE); + h->magic = RDMA_SEG_MAGIC; + h->len = pend_len; + if (!post_send(i, RDMA_STRIDE)) { broken = true; return false; } + send_busy[i] = 1; + pend_buf = -1; + pend_len = 0; + return true; +} + +// Coalescing write: append into the pending frame, posting a full frame when it +// fills. The trailing partial is posted by flush() at each message boundary. +bool apple_rdma::send(const void * data, size_t size) { + impl * c = pimpl.get(); + const uint8_t * p = (const uint8_t *)data; + while (size > 0) { + if (c->broken) return false; + if (!c->acquire_pending()) return false; + uint8_t * sb = c->send_mem + (size_t)c->pend_buf * RDMA_STRIDE; + size_t space = RDMA_PAYLOAD - c->pend_len; + size_t chunk = size < space ? size : space; + memcpy(sb + sizeof(rdma_seg_hdr) + c->pend_len, p, chunk); + c->pend_len += (uint32_t)chunk; + p += chunk; + size -= chunk; + if (c->pend_len == RDMA_PAYLOAD) { if (!c->post_pending()) return false; } + } + return true; +} + +bool apple_rdma::recv(void * data, size_t size) { + impl * c = pimpl.get(); + uint8_t * p = (uint8_t *)data; + if (!c->post_pending()) return false; // turnaround: flush the coalesced request + unsigned idle = 0; + while (size > 0) { + if (c->inq_count == 0) { + if (c->broken) return false; + int n = c->progress(); + if (n < 0) return false; + if (n == 0) { + // UC gives no disconnect notification, so the bootstrap TCP fd is + // the liveness anchor: nothing crosses it once RDMA is up, so any + // readability means the peer's FIN (macOS has no POLLRDHUP). + // Same idle interval as the Linux path. + if ((++idle & 0xFFFFF) == 0) { + struct pollfd pfd = { c->fd, POLLIN, 0 }; + if (poll(&pfd, 1, 0) > 0 && + (pfd.revents & (POLLIN | POLLHUP | POLLERR | POLLNVAL))) { + return false; + } + } + } else { + idle = 0; + } + continue; + } + idle = 0; + int slot = c->inq_head; + int b = c->inq[slot].buf; + uint32_t avail = c->inq[slot].len - c->inq[slot].off; + uint32_t take = (size < (size_t)avail) ? (uint32_t)size : avail; + memcpy(p, c->recv_mem + (size_t)b * RDMA_STRIDE + sizeof(rdma_seg_hdr) + c->inq[slot].off, take); + p += take; + size -= take; + c->inq[slot].off += take; + if (c->inq[slot].off == c->inq[slot].len) { + if (!c->post_recv(b)) { c->broken = true; return false; } + c->inq_head = (c->inq_head + 1) % RDMA_NBUF; + c->inq_count--; + } + } + return true; +} + +bool apple_rdma::flush() { + return pimpl->post_pending(); +} diff --git a/ggml/src/ggml-rpc/transport-apple.h b/ggml/src/ggml-rpc/transport-apple.h new file mode 100644 index 000000000..7968d38a1 --- /dev/null +++ b/ggml/src/ggml-rpc/transport-apple.h @@ -0,0 +1,27 @@ +#pragma once + +#include +#include +#include + +struct apple_rdma { + // target_gid is 16 bytes in, caps is RPC_CONN_CAPS_SIZE bytes out. + static std::unique_ptr probe(int fd, const uint8_t * target_gid, uint8_t * caps); + ~apple_rdma(); + + // Peer endpoint from its caps, which must be non-zero: this blocks on a + // readiness handshake over fd that the peer only joins if it also has RDMA. + bool activate(const uint8_t * caps); + + bool send(const void * data, size_t size); + bool recv(void * data, size_t size); + // Post the trailing partial frame; must be called at every message boundary. + bool flush(); + // True once the connection has failed; the caller should drop the socket. + bool broken() const; + +private: + struct impl; + explicit apple_rdma(std::unique_ptr p); + std::unique_ptr pimpl; +}; diff --git a/ggml/src/ggml-rpc/transport.cpp b/ggml/src/ggml-rpc/transport.cpp index a72815242..5ec15dc80 100644 --- a/ggml/src/ggml-rpc/transport.cpp +++ b/ggml/src/ggml-rpc/transport.cpp @@ -18,15 +18,20 @@ # include #endif #include +#include #include #include #ifdef GGML_RPC_RDMA # include +# include # include # ifndef _WIN32 # include # endif +# ifdef GGML_RPC_RDMA_APPLE +# include "transport-apple.h" +# endif #endif // GGML_RPC_RDMA #ifdef _WIN32 @@ -42,10 +47,13 @@ static const char * RPC_DEBUG = std::getenv("GGML_RPC_DEBUG"); do { if (RPC_DEBUG) GGML_LOG_DEBUG(__VA_ARGS__); } while (0) #ifdef GGML_RPC_RDMA -static constexpr size_t RDMA_CHUNK = 256 * 1024; // 256 KiB per send/recv (fits default 8 MiB memlock) -static constexpr int RDMA_RX_DEPTH = 24; // pre-posted recv ring: 24 × 256 KiB = 6 MiB static constexpr size_t RDMA_GID_SIZE = 16; // RoCE GID / IB GID is always 16 bytes using rdma_gid_t = std::array; +#endif // GGML_RPC_RDMA + +#if defined(GGML_RPC_RDMA) && !defined(GGML_RPC_RDMA_APPLE) +static constexpr size_t RDMA_CHUNK = 256 * 1024; // 256 KiB per send/recv (fits default 8 MiB memlock) +static constexpr int RDMA_RX_DEPTH = 24; // pre-posted recv ring: 24 × 256 KiB = 6 MiB struct rdma_conn { struct ibv_context * ctx = nullptr; @@ -111,27 +119,33 @@ struct rdma_caps { static_assert(sizeof(rdma_caps) == RPC_CONN_CAPS_SIZE, "rdma_caps must match conn_caps size"); -#endif // GGML_RPC_RDMA +#endif // GGML_RPC_RDMA && !GGML_RPC_RDMA_APPLE struct socket_t::impl { impl(sockfd_t fd) : use_rdma(false), fd(fd) {} ~impl(); bool send_data(const void * data, size_t size); bool recv_data(void * data, size_t size); + bool flush(); void get_caps(uint8_t * local_caps); void update_caps(const uint8_t * remote_caps); #ifdef GGML_RPC_RDMA - bool tcp_peer_closed(); std::optional rdma_build_target_gid(); + +# ifdef GGML_RPC_RDMA_APPLE + std::unique_ptr rdma; +# else bool rdma_probe(); - bool rdma_activate(uint32_t remote_qpn, uint32_t remote_psn, const uint8_t * remote_gid); - bool rdma_poll(struct ibv_cq * cq, struct ibv_wc * wc); bool rdma_send(const void * data, size_t size); bool rdma_recv(void * data, size_t size); + bool tcp_peer_closed(); + bool rdma_activate(uint32_t remote_qpn, uint32_t remote_psn, const uint8_t * remote_gid); + bool rdma_poll(struct ibv_cq * cq, struct ibv_wc * wc); std::unique_ptr rdma; rdma_local_info rdma_local = {}; +# endif #endif // GGML_RPC_RDMA bool use_rdma; sockfd_t fd; @@ -151,17 +165,6 @@ socket_t::impl::~impl() { #ifdef GGML_RPC_RDMA -bool socket_t::impl::tcp_peer_closed() { - if (fd < 0) return false; -#ifndef _WIN32 - struct pollfd pfd = { fd, POLLIN | POLLRDHUP, 0 }; - int r = poll(&pfd, 1, 0); - return r > 0 && (pfd.revents & (POLLHUP | POLLERR | POLLRDHUP)); -#else - return false; -#endif -} - // Build a RoCE GID-shaped 16-byte target from a TCP socket's local address. // Used to match the socket's local IP against the kernel's GID table so that // a single memcmp handles IPv4, IPv4-mapped IPv6, and native IPv6 uniformly: @@ -191,6 +194,19 @@ std::optional socket_t::impl::rdma_build_target_gid() { return std::nullopt; } +#ifndef GGML_RPC_RDMA_APPLE + +bool socket_t::impl::tcp_peer_closed() { + if (fd < 0) return false; +#ifndef _WIN32 + struct pollfd pfd = { fd, POLLIN | POLLRDHUP, 0 }; + int r = poll(&pfd, 1, 0); + return r > 0 && (pfd.revents & (POLLHUP | POLLERR | POLLRDHUP)); +#else + return false; +#endif +} + bool socket_t::impl::rdma_probe() { const char * dev_env = std::getenv("GGML_RDMA_DEV"); const char * gid_env = std::getenv("GGML_RDMA_GID"); @@ -457,10 +473,16 @@ bool socket_t::impl::rdma_recv(void * data, size_t size) { return true; } +#endif // !GGML_RPC_RDMA_APPLE (Linux RC transport) + #endif // GGML_RPC_RDMA bool socket_t::impl::send_data(const void * data, size_t size) { -#ifdef GGML_RPC_RDMA +#ifdef GGML_RPC_RDMA_APPLE + if (use_rdma) { + return rdma->send(data, size); + } +#elif defined(GGML_RPC_RDMA) if (use_rdma) { return rdma_send(data, size); } @@ -480,7 +502,11 @@ bool socket_t::impl::send_data(const void * data, size_t size) { } bool socket_t::impl::recv_data(void * data, size_t size) { -#ifdef GGML_RPC_RDMA +#ifdef GGML_RPC_RDMA_APPLE + if (use_rdma) { + return rdma->recv(data, size); + } +#elif defined(GGML_RPC_RDMA) if (use_rdma) { return rdma_recv(data, size); } @@ -506,6 +532,15 @@ bool socket_t::impl::recv_data(void * data, size_t size) { void socket_t::impl::get_caps(uint8_t * local_caps) { memset(local_caps, 0, RPC_CONN_CAPS_SIZE); #ifdef GGML_RPC_RDMA + if (std::getenv("GGML_RPC_NO_RDMA")) { + return; + } +# ifdef GGML_RPC_RDMA_APPLE + auto target_gid = rdma_build_target_gid(); + if (target_gid) { + rdma = apple_rdma::probe(fd, target_gid->data(), local_caps); + } +# else rdma_local = {}; if (rdma_probe()) { rdma_caps rc = {}; @@ -516,21 +551,30 @@ void socket_t::impl::get_caps(uint8_t * local_caps) { } else { rdma.reset(); } +# endif #endif // GGML_RPC_RDMA } void socket_t::impl::update_caps(const uint8_t * remote_caps) { #ifdef GGML_RPC_RDMA - if (!rdma) { - return; + // a peer that has no RDMA advertises all-zero caps and takes no further part + // in the negotiation, so drop to TCP without reporting a failure + bool remote_rdma = false; + for (size_t i = 0; i < RPC_CONN_CAPS_SIZE; i++) { + remote_rdma |= remote_caps[i] != 0; } - rdma_caps rc = {}; - memcpy(&rc, remote_caps, sizeof(rc)); - if (rc.qpn == 0) { + if (!rdma || !remote_rdma) { rdma.reset(); return; } - if (rdma_activate(rc.qpn, rc.psn, rc.gid)) { +# ifdef GGML_RPC_RDMA_APPLE + bool activated = rdma->activate(remote_caps); +# else + rdma_caps rc = {}; + memcpy(&rc, remote_caps, sizeof(rc)); + bool activated = rdma_activate(rc.qpn, rc.psn, rc.gid); +# endif + if (activated) { use_rdma = true; } else { GGML_LOG_ERROR("RDMA activate failed, staying on TCP\n"); @@ -541,6 +585,14 @@ void socket_t::impl::update_caps(const uint8_t * remote_caps) { #endif // GGML_RPC_RDMA } +bool socket_t::impl::flush() { +#ifdef GGML_RPC_RDMA_APPLE + if (use_rdma) { + return rdma->flush(); + } +#endif + return true; +} ///////////////////////////////////////////////////////////////////////////// @@ -556,6 +608,10 @@ bool socket_t::recv_data(void * data, size_t size) { return pimpl->recv_data(data, size); } +bool socket_t::flush() { + return pimpl->flush(); +} + void socket_t::get_caps(uint8_t * local_caps) { return pimpl->get_caps(local_caps); } diff --git a/ggml/src/ggml-rpc/transport.h b/ggml/src/ggml-rpc/transport.h index 73b85cc53..3f747ecff 100644 --- a/ggml/src/ggml-rpc/transport.h +++ b/ggml/src/ggml-rpc/transport.h @@ -15,6 +15,10 @@ struct socket_t { bool send_data(const void * data, size_t size); bool recv_data(void * data, size_t size); + // Must be called at every message boundary: the RDMA transport coalesces + // writes into fixed-size frames and posts the trailing partial frame only + // here. No-op on TCP. + bool flush(); socket_ptr accept(); diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index e6b364fa7..ffd8c3e47 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1048,6 +1048,8 @@ struct vk_device_struct { vk_pipeline pipeline_argsort_large_f32[num_argsort_pipelines]; vk_pipeline pipeline_topk_f32[num_topk_pipelines]; vk_pipeline pipeline_sum_rows_f32; + vk_pipeline pipeline_cross_entropy_loss_f32, pipeline_cross_entropy_loss_f32_wg512; + vk_pipeline pipeline_cross_entropy_loss_back_f32, pipeline_cross_entropy_loss_back_f32_wg512; vk_pipeline pipeline_fwht_f32[4]; vk_pipeline pipeline_cumsum_f32; vk_pipeline pipeline_cumsum_small_f32; @@ -4175,10 +4177,16 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t subgroup_size_16 = std::max(device->subgroup_size, 16u); const uint32_t subgroup_size_32 = std::max(device->subgroup_size, 32u); + // clamp WARP for l_/m_ warptiles so WM <= BM (breaks on subgroupSize > 64) + const uint32_t mm_warp_8 = std::min(subgroup_size_8, 64u); + const uint32_t mm_warp_16 = std::min(subgroup_size_16, 64u); + const uint32_t mul_mat_subgroup_size = (device->vendor_id == VK_VENDOR_ID_INTEL && device->subgroup_size_control) ? device->subgroup_min_size : device->subgroup_size; const uint32_t mul_mat_subgroup_size_8 = std::max(mul_mat_subgroup_size, 8u); const uint32_t mul_mat_subgroup_size_16 = std::max(mul_mat_subgroup_size, 16u); const uint32_t mul_mat_subgroup_size_32 = std::max(mul_mat_subgroup_size, 32u); + const uint32_t mul_mat_mm_warp_8 = std::min(mul_mat_subgroup_size_8, 64u); + const uint32_t mul_mat_mm_warp_16 = std::min(mul_mat_subgroup_size_16, 64u); const bool subgroup_min_size_16 = (!device->subgroup_size_control && device->subgroup_size >= 16) || (device->subgroup_size_control && device->subgroup_max_size >= 16); @@ -4259,39 +4267,39 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t s_warptile_wm = device->subgroup_size == 8 ? 8 : 32; - l_warptile = { 128, 128, 128, 16, subgroup_size_8 * 2, 64, 2, tm_l, tn_l, tk_l, subgroup_size_8 }; - m_warptile = { 128, 64, 64, 16, subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - s_warptile = { subgroup_size_32, 32, 32, 16, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; + l_warptile = { 128, 128, 128, 16, mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mm_warp_8 }; + m_warptile = { 128, 64, 64, 16, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + s_warptile = { subgroup_size_32, 32, 32, 16, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; - l_warptile_mmq = { 128, 128, 128, 32, subgroup_size_8 * 2, 64, 2, tm_l, tn_l, tk_l, subgroup_size_8 }; - m_warptile_mmq = { 128, 64, 64, 32, subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - s_warptile_mmq = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; + l_warptile_mmq = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mm_warp_8 }; + m_warptile_mmq = { 128, 64, 64, 32, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + s_warptile_mmq = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, subgroup_size_8 }; // Integer MMQ has a smaller shared memory profile, but heavier register use - l_warptile_mmq_int = { 128, 128, 128, 32, subgroup_size_8 * 2, 64, 2, 4, 4, 1, subgroup_size_8 }; - m_warptile_mmq_int = { 128, 64, 64, 32, subgroup_size_8, 32, 2, 2, 2, 1, subgroup_size_8 }; - s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, subgroup_size_8 }; + l_warptile_mmq_int = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 2, 4, 4, 1, mm_warp_8 }; + m_warptile_mmq_int = { 128, 64, 64, 32, mm_warp_8, 32, 2, 2, 2, 1, mm_warp_8 }; + s_warptile_mmq_int = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, subgroup_size_8 }; // K-quants use even more registers, mitigate by setting WMITER to 1 - l_warptile_mmq_int_k = { 128, 128, 128, 32, subgroup_size_8 * 2, 64, 1, 4, 4, 1, subgroup_size_8 }; - m_warptile_mmq_int_k = { 128, 64, 64, 32, subgroup_size_8, 32, 1, 2, 2, 1, subgroup_size_8 }; - s_warptile_mmq_int_k = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 1, 2, 1, 1, subgroup_size_8 }; + l_warptile_mmq_int_k = { 128, 128, 128, 32, mm_warp_8 * 2, 64, 1, 4, 4, 1, mm_warp_8 }; + m_warptile_mmq_int_k = { 128, 64, 64, 32, mm_warp_8, 32, 1, 2, 2, 1, mm_warp_8 }; + s_warptile_mmq_int_k = { subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 1, 2, 1, 1, subgroup_size_8 }; - l_warptile_id = { 128, 128, 128, 16, mul_mat_subgroup_size_16 * 2, 64, 2, tm_l, tn_l, tk_l, mul_mat_subgroup_size_16 }; - m_warptile_id = { 128, 64, 64, 16, mul_mat_subgroup_size_16, 32, 2, tm_m, tn_m, tk_m, mul_mat_subgroup_size_16 }; - s_warptile_id = { mul_mat_subgroup_size_16, 32, 32, 16, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, mul_mat_subgroup_size_16 }; + l_warptile_id = { 128, 128, 128, 16, mul_mat_mm_warp_16 * 2, 64, 2, tm_l, tn_l, tk_l, mul_mat_mm_warp_16 }; + m_warptile_id = { 128, 64, 64, 16, mul_mat_mm_warp_16, 32, 2, tm_m, tn_m, tk_m, mul_mat_mm_warp_16 }; + s_warptile_id = { mul_mat_subgroup_size_16, 32, 32, 16, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, mul_mat_subgroup_size_16 }; - l_warptile_mmqid = { 128, 128, 128, 32, mul_mat_subgroup_size_8 * 2, 64, 2, tm_l, tn_l, tk_l, mul_mat_subgroup_size_8 }; - m_warptile_mmqid = { 128, 64, 64, 32, mul_mat_subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, mul_mat_subgroup_size_8 }; - s_warptile_mmqid = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, mul_mat_subgroup_size_8 }; + l_warptile_mmqid = { 128, 128, 128, 32, mul_mat_mm_warp_8 * 2, 64, 2, tm_l, tn_l, tk_l, mul_mat_mm_warp_8 }; + m_warptile_mmqid = { 128, 64, 64, 32, mul_mat_mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mul_mat_mm_warp_8 }; + s_warptile_mmqid = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, tm_s, tn_s, tk_s, mul_mat_subgroup_size_8 }; - l_warptile_mmqid_int = { 128, 128, 128, 32, mul_mat_subgroup_size_8 * 2, 64, 2, 4, 4, 1, mul_mat_subgroup_size_8 }; - m_warptile_mmqid_int = { 128, 64, 64, 32, mul_mat_subgroup_size_8, 32, 2, 2, 2, 1, mul_mat_subgroup_size_8 }; - s_warptile_mmqid_int = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, mul_mat_subgroup_size_8 }; + l_warptile_mmqid_int = { 128, 128, 128, 32, mul_mat_mm_warp_8 * 2, 64, 2, 4, 4, 1, mul_mat_mm_warp_8 }; + m_warptile_mmqid_int = { 128, 64, 64, 32, mul_mat_mm_warp_8, 32, 2, 2, 2, 1, mul_mat_mm_warp_8 }; + s_warptile_mmqid_int = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 2, 2, 1, 1, mul_mat_subgroup_size_8 }; - l_warptile_mmqid_int_k = { 128, 128, 128, 32, mul_mat_subgroup_size_16 * 2, 64, 1, 4, 4, 1, mul_mat_subgroup_size_16 }; - m_warptile_mmqid_int_k = { 128, 64, 64, 32, mul_mat_subgroup_size_16, 32, 1, 2, 2, 1, mul_mat_subgroup_size_16 }; - s_warptile_mmqid_int_k = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 1, 2, 1, 1, mul_mat_subgroup_size_16 }; + l_warptile_mmqid_int_k = { 128, 128, 128, 32, mul_mat_mm_warp_16 * 2, 64, 1, 4, 4, 1, mul_mat_mm_warp_16 }; + m_warptile_mmqid_int_k = { 128, 64, 64, 32, mul_mat_mm_warp_16, 32, 1, 2, 2, 1, mul_mat_mm_warp_16 }; + s_warptile_mmqid_int_k = { mul_mat_subgroup_size_32, 32, 32, 32, s_warptile_wm, 32, 1, 2, 1, 1, mul_mat_subgroup_size_16 }; // chip specific tuning if ((device->architecture == AMD_GCN) && (device->driver_id != vk::DriverId::eAmdProprietary)) { @@ -4299,13 +4307,13 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { m_warptile_mmqid = m_warptile_mmqid_int = { 256, 64, 64, 32, 16, 16, 2, 2, 2, 1, 16 }; } else if (device->vendor_id == VK_VENDOR_ID_AMD && device->coopmat_support && device->driver_id != vk::DriverId::eAmdProprietary) { // This is intentionally using tx_m values, slight performance increase - l_warptile = { 256, 128, 128, 16, subgroup_size_8, 64, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - l_warptile_mmq = l_warptile_mmq_int = { 256, 128, 128, 32, subgroup_size_8, 64, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - l_warptile_mmq_int_k = { 256, 128, 128, 32, subgroup_size_16, 64, 1, 4, 2, 1, subgroup_size_16 }; + l_warptile = { 256, 128, 128, 16, mm_warp_8, 64, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + l_warptile_mmq = l_warptile_mmq_int = { 256, 128, 128, 32, mm_warp_8, 64, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + l_warptile_mmq_int_k = { 256, 128, 128, 32, mm_warp_16, 64, 1, 4, 2, 1, mm_warp_16 }; } else if (device->vendor_id == VK_VENDOR_ID_INTEL && device->coopmat_support) { // Xe2/Xe3 with coopmat enabled - warptile performance tuning - l_warptile = { 512, 128, 128, 16, subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; - l_warptile_mmq = { 512, 128, 128, 32, subgroup_size_8, 32, 2, tm_m, tn_m, tk_m, subgroup_size_8 }; + l_warptile = { 512, 128, 128, 16, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; + l_warptile_mmq = { 512, 128, 128, 32, mm_warp_8, 32, 2, tm_m, tn_m, tk_m, mm_warp_8 }; } l_mmq_wg_denoms = l_wg_denoms = {128, 128, 1 }; @@ -5178,8 +5186,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t s_warptile_wm = device->subgroup_size == 8 ? 8 : 32; // use scalar tile sizes - l_warptile = { 128, 128, 128, 16, subgroup_size_8 * 2, 64, 2, 4, 4, 1, subgroup_size_8 }; - m_warptile = { 128, 64, 64, 16, subgroup_size_8, 32, 2, 4, 2, 1, subgroup_size_8 }; + l_warptile = { 128, 128, 128, 16, mm_warp_8 * 2, 64, 2, 4, 4, 1, mm_warp_8 }; + m_warptile = { 128, 64, 64, 16, mm_warp_8, 32, 2, 4, 2, 1, mm_warp_8 }; s_warptile = { subgroup_size_32, 32, 32, 16, s_warptile_wm, 32, 2, 2, 2, 1, subgroup_size_8 }; l_wg_denoms = {128, 128, 1 }; @@ -5764,6 +5772,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_argmax_f32, "argmax_f32", argmax_f32_len, argmax_f32_data, "main", 2, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); ggml_vk_create_pipeline(device, device->pipeline_sum_rows_f32, "sum_rows_f32", sum_rows_f32_len, sum_rows_f32_data, "main", 2, sizeof(vk_op_sum_rows_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); + ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_f32, "cross_entropy_loss_f32", cross_entropy_loss_f32_len, cross_entropy_loss_f32_data, "main", 3, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); + ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_f32_wg512, "cross_entropy_loss_f32_wg512", cross_entropy_loss_f32_len, cross_entropy_loss_f32_data, "main", 3, sizeof(vk_op_push_constants), {1, 1, 1}, { 512 }, 1); + ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_back_f32, "cross_entropy_loss_back_f32", cross_entropy_loss_back_f32_len, cross_entropy_loss_back_f32_data, "main", 4, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1); + ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_back_f32_wg512, "cross_entropy_loss_back_f32_wg512", cross_entropy_loss_back_f32_len, cross_entropy_loss_back_f32_data, "main", 4, sizeof(vk_op_push_constants), {1, 1, 1}, { 512 }, 1); // Intel Windows driver in range [32.0.101.8509, 32.0.101.8860) will crash when using fwht kernels so we gate that here const bool can_use_fwht = device->driver_id != vk::DriverId::eIntelProprietaryWindows || !ggml_vk_intel_windows_driver_in_range(device->properties.driverVersion, 101, 8509, 101, 8860); @@ -11610,6 +11622,17 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const return ctx->device->pipeline_sum_rows_f32; } return nullptr; + case GGML_OP_CROSS_ENTROPY_LOSS: + if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { + return src0->ne[0] > 1024 ? ctx->device->pipeline_cross_entropy_loss_f32_wg512 : ctx->device->pipeline_cross_entropy_loss_f32; + } + return nullptr; + case GGML_OP_CROSS_ENTROPY_LOSS_BACK: + // src0 is the scalar grad; src1 is logits + if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && src2 && src2->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { + return src1->ne[0] > 1024 ? ctx->device->pipeline_cross_entropy_loss_back_f32_wg512 : ctx->device->pipeline_cross_entropy_loss_back_f32; + } + return nullptr; case GGML_OP_CUMSUM: if (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { if (src0->ne[0] <= 512) { @@ -13975,6 +13998,103 @@ static void ggml_vk_cumsum(ggml_backend_vk_context * ctx, vk_context& subctx, co ctx->prealloc_split_k_need_sync = true; } +static std::array ggml_vk_nrows_elements(uint32_t nr) { + if (nr > 262144) { + return { 512, 512, CEIL_DIV(nr, 262144) }; + } + if (nr > 512) { + return { 512, CEIL_DIV(nr, 512), 1 }; + } + return { nr, 1, 1 }; +} + +static void ggml_vk_cross_entropy_loss(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src1->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(src0)); + GGML_ASSERT(ggml_is_contiguous(src1)); + GGML_ASSERT(ggml_is_contiguous(dst)); + GGML_ASSERT(ggml_are_same_shape(src0, src1)); + GGML_ASSERT(ggml_is_scalar(dst)); + + const uint32_t nclasses = (uint32_t)src0->ne[0]; + const uint32_t nrows = (uint32_t)ggml_nrows(src0); + + vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, src0, src1, nullptr, dst, GGML_OP_CROSS_ENTROPY_LOSS); + GGML_ASSERT(pipeline != nullptr); + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_sum_rows_f32, 1); + + vk_subbuffer src0_buf = ggml_vk_tensor_subbuffer(ctx, src0); + vk_subbuffer src1_buf = ggml_vk_tensor_subbuffer(ctx, src1); + vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst, true); + + const vk_op_push_constants pc = { nclasses, nrows, 0.0f, 0.0f, 0.0f, 0.0f }; + + const size_t tmp_size = (size_t)nrows * sizeof(float); + if (ctx->prealloc_size_x < tmp_size) { + ctx->prealloc_size_x = tmp_size; + ggml_vk_preallocate_buffers(ctx, subctx); + } + if (ctx->prealloc_x_need_sync) { + ggml_vk_sync_buffers(ctx, subctx); + } + + vk_subbuffer tmp_buf = { ctx->prealloc_x, 0, tmp_size }; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { src0_buf, src1_buf, tmp_buf }, pc, ggml_vk_nrows_elements(nrows)); + ggml_vk_sync_buffers(ctx, subctx); + + vk_op_sum_rows_push_constants sp = {}; + sp.n_cols = nrows; + sp.ne01 = 1; + sp.ne02 = 1; + sp.weight = 1.0f; + init_pushconst_fastdiv(sp); + sp.misalign_offsets = get_misalign_bytes(ctx, dst) / ggml_type_size(dst->type); + + ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_sum_rows_f32, { tmp_buf, dst_buf }, sp, { 1, 1, 1 }); + ctx->prealloc_x_need_sync = true; +} + +static void ggml_vk_cross_entropy_loss_back(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const ggml_tensor * grad = dst->src[0]; + const ggml_tensor * logits = dst->src[1]; + const ggml_tensor * labels = dst->src[2]; + + GGML_ASSERT(grad->type == GGML_TYPE_F32); + GGML_ASSERT(logits->type == GGML_TYPE_F32); + GGML_ASSERT(labels->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_scalar(grad)); + GGML_ASSERT(ggml_is_contiguous(grad)); + GGML_ASSERT(ggml_is_contiguous(logits)); + GGML_ASSERT(ggml_is_contiguous(labels)); + GGML_ASSERT(ggml_is_contiguous(dst)); + GGML_ASSERT(ggml_are_same_shape(logits, labels)); + GGML_ASSERT(ggml_are_same_shape(logits, dst)); + + const uint32_t nclasses = (uint32_t)logits->ne[0]; + const uint32_t nrows = (uint32_t)ggml_nrows(logits); + + vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, grad, logits, labels, dst, GGML_OP_CROSS_ENTROPY_LOSS_BACK); + GGML_ASSERT(pipeline != nullptr); + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + vk_subbuffer grad_buf = ggml_vk_tensor_subbuffer(ctx, grad); + vk_subbuffer logits_buf = ggml_vk_tensor_subbuffer(ctx, logits); + vk_subbuffer labels_buf = ggml_vk_tensor_subbuffer(ctx, labels); + vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); + + const vk_op_push_constants pc = { nclasses, nrows, 0.0f, 0.0f, 0.0f, 0.0f }; + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { grad_buf, logits_buf, labels_buf, dst_buf }, pc, ggml_vk_nrows_elements(nrows)); +} + static void ggml_vk_argmax(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) { ggml_vk_op_f32(ctx, subctx, src0, nullptr, nullptr, nullptr, dst, GGML_OP_ARGMAX, { (uint32_t)src0->ne[0], (uint32_t)src0->ne[1], 0.0f, 0.0f, 0.0f, 0.0f }); } @@ -15720,6 +15840,14 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr case GGML_OP_ARGMAX: ggml_vk_argmax(ctx, compute_ctx, src0, node); + break; + case GGML_OP_CROSS_ENTROPY_LOSS: + ggml_vk_cross_entropy_loss(ctx, compute_ctx, node); + + break; + case GGML_OP_CROSS_ENTROPY_LOSS_BACK: + ggml_vk_cross_entropy_loss_back(ctx, compute_ctx, node); + break; case GGML_OP_COUNT_EQUAL: ggml_vk_count_equal(ctx, compute_ctx, src0, src1, node); @@ -18544,6 +18672,18 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm } case GGML_OP_ARGMAX: return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32; + case GGML_OP_CROSS_ENTROPY_LOSS: + return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32 + && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_F32 + && ggml_are_same_shape(op->src[0], op->src[1]) + && ggml_is_contiguous(op) && ggml_is_scalar(op) && op->type == GGML_TYPE_F32; + case GGML_OP_CROSS_ENTROPY_LOSS_BACK: + return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32 && ggml_is_scalar(op->src[0]) + && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_F32 + && ggml_is_contiguous(op->src[2]) && op->src[2]->type == GGML_TYPE_F32 + && ggml_are_same_shape(op->src[1], op->src[2]) + && ggml_are_same_shape(op->src[1], op) + && ggml_is_contiguous(op) && op->type == GGML_TYPE_F32; case GGML_OP_COUNT_EQUAL: return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_I32 && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_I32; @@ -19470,6 +19610,10 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph * tensor_clone = ggml_mean(ggml_ctx, src_clone[0]); } else if (tensor->op == GGML_OP_ARGMAX) { tensor_clone = ggml_argmax(ggml_ctx, src_clone[0]); + } else if (tensor->op == GGML_OP_CROSS_ENTROPY_LOSS) { + tensor_clone = ggml_cross_entropy_loss(ggml_ctx, src_clone[0], src_clone[1]); + } else if (tensor->op == GGML_OP_CROSS_ENTROPY_LOSS_BACK) { + tensor_clone = ggml_cross_entropy_loss_back(ggml_ctx, src_clone[0], src_clone[1], src_clone[2]); } else if (tensor->op == GGML_OP_COUNT_EQUAL) { tensor_clone = ggml_count_equal(ggml_ctx, src_clone[0], src_clone[1]); } else if (tensor->op == GGML_OP_SOLVE_TRI) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss.comp b/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss.comp new file mode 100644 index 000000000..0c135c6fd --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss.comp @@ -0,0 +1,78 @@ +#version 450 + +#include "generic_head.glsl" +#include "types.glsl" + +#extension GL_EXT_control_flow_attributes : enable + +layout(constant_id = 0) const uint BLOCK_SIZE = 32; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {A_TYPE data_a[];}; +layout (binding = 1) readonly buffer B {B_TYPE data_b[];}; +layout (binding = 2) writeonly buffer D {D_TYPE data_d[];}; + +shared FLOAT_TYPE tmp[BLOCK_SIZE]; + +FLOAT_TYPE wg_reduce_max(FLOAT_TYPE v) { + const uint tid = gl_LocalInvocationID.x; + tmp[tid] = v; + barrier(); + [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) { + if (tid < s) { + tmp[tid] = max(tmp[tid], tmp[tid + s]); + } + barrier(); + } + v = tmp[0]; + barrier(); + return v; +} + +FLOAT_TYPE wg_reduce_sum(FLOAT_TYPE v) { + const uint tid = gl_LocalInvocationID.x; + tmp[tid] = v; + barrier(); + [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) { + if (tid < s) { + tmp[tid] += tmp[tid + s]; + } + barrier(); + } + v = tmp[0]; + barrier(); + return v; +} + +void main() { + const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x; + const uint tid = gl_LocalInvocationID.x; + + if (row >= p.KY) { + return; + } + + const uint off = row * p.KX; + + FLOAT_TYPE max_logit = FLOAT_TYPE(uintBitsToFloat(0xFF800000)); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + max_logit = max(max_logit, FLOAT_TYPE(data_a[off + i])); + } + max_logit = wg_reduce_max(max_logit); + + FLOAT_TYPE sum_exp = FLOAT_TYPE(0.0f); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + sum_exp += exp(FLOAT_TYPE(data_a[off + i]) - max_logit); + } + const FLOAT_TYPE log_sum = log(wg_reduce_sum(sum_exp)); + + FLOAT_TYPE loss = FLOAT_TYPE(0.0f); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + loss += (FLOAT_TYPE(data_a[off + i]) - max_logit - log_sum) * FLOAT_TYPE(data_b[off + i]); + } + loss = -wg_reduce_sum(loss) / FLOAT_TYPE(p.KY); + + if (tid == 0) { + data_d[row] = D_TYPE(loss); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss_back.comp b/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss_back.comp new file mode 100644 index 000000000..3cdebe86e --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/cross_entropy_loss_back.comp @@ -0,0 +1,75 @@ +#version 450 + +#include "generic_head.glsl" +#include "types.glsl" + +#extension GL_EXT_control_flow_attributes : enable + +layout(constant_id = 0) const uint BLOCK_SIZE = 32; +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer G {A_TYPE data_g[];}; +layout (binding = 1) readonly buffer X {B_TYPE data_x[];}; +layout (binding = 2) readonly buffer Y {B_TYPE data_y[];}; +layout (binding = 3) writeonly buffer D {D_TYPE data_d[];}; + +shared FLOAT_TYPE tmp[BLOCK_SIZE]; + +FLOAT_TYPE wg_reduce_max(FLOAT_TYPE v) { + const uint tid = gl_LocalInvocationID.x; + tmp[tid] = v; + barrier(); + [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) { + if (tid < s) { + tmp[tid] = max(tmp[tid], tmp[tid + s]); + } + barrier(); + } + v = tmp[0]; + barrier(); + return v; +} + +FLOAT_TYPE wg_reduce_sum(FLOAT_TYPE v) { + const uint tid = gl_LocalInvocationID.x; + tmp[tid] = v; + barrier(); + [[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) { + if (tid < s) { + tmp[tid] += tmp[tid + s]; + } + barrier(); + } + v = tmp[0]; + barrier(); + return v; +} + +void main() { + const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x; + const uint tid = gl_LocalInvocationID.x; + + if (row >= p.KY) { + return; + } + + const uint off = row * p.KX; + const FLOAT_TYPE d_by_nrows = FLOAT_TYPE(data_g[0]) / FLOAT_TYPE(p.KY); + + FLOAT_TYPE max_logit = FLOAT_TYPE(uintBitsToFloat(0xFF800000)); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + max_logit = max(max_logit, FLOAT_TYPE(data_x[off + i])); + } + max_logit = wg_reduce_max(max_logit); + + FLOAT_TYPE sum_exp = FLOAT_TYPE(0.0f); + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + sum_exp += exp(FLOAT_TYPE(data_x[off + i]) - max_logit); + } + const FLOAT_TYPE inv_sum = FLOAT_TYPE(1.0f) / wg_reduce_sum(sum_exp); + + for (uint i = tid; i < p.KX; i += BLOCK_SIZE) { + const FLOAT_TYPE sm = exp(FLOAT_TYPE(data_x[off + i]) - max_logit) * inv_sum; + data_d[off + i] = D_TYPE((sm - FLOAT_TYPE(data_y[off + i])) * d_by_nrows); + } +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 59c5d72ce..da116976b 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -1055,6 +1055,8 @@ void process_shaders() { string_to_spv("argmax_f32", "argmax.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "int"}})); string_to_spv("sum_rows_f32", "sum_rows.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}})); + string_to_spv("cross_entropy_loss_f32", "cross_entropy_loss.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}})); + string_to_spv("cross_entropy_loss_back_f32", "cross_entropy_loss_back.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}})); string_to_spv("fwht_f32", "fwht.comp", {}); string_to_spv("fwht_shmem_f32", "fwht.comp", {{"FWHT_SHMEM", "1"}}); string_to_spv("count_equal_i32", "count_equal.comp", merge_maps(base_dict, {{"A_TYPE", "int"}, {"B_TYPE", "int"}, {"D_TYPE", "int"}})); diff --git a/gpttype_adapter.cpp b/gpttype_adapter.cpp index 83d25a1fc..b39735b07 100644 --- a/gpttype_adapter.cpp +++ b/gpttype_adapter.cpp @@ -5386,7 +5386,7 @@ static void PrepareMediaEmbds(const int nctx, const std::vector & media_int std::string media_obj = media_objects[i].b64data; const std::vector media_data_buffer = kcpp_base64_decode(media_obj); mtmd::bitmap bitmap(media_objects[i].is_audio - ? mtmd_helper_bitmap_init_from_buf(mtmd_ctx, media_data_buffer.data(), media_data_buffer.size(),false).bitmap + ? mtmd_helper_bitmap_init_from_buf(mtmd_ctx, media_data_buffer.data(), media_data_buffer.size(), false, mtmd_helper_init_opt_default()).bitmap : kcpp_mtmd_bitmap_init_image_from_buf(media_data_buffer.data(), media_data_buffer.size(), vision_max_res)); if(!bitmap.ptr) { diff --git a/include/llama.h b/include/llama.h index 8890edb22..38a0a306c 100644 --- a/include/llama.h +++ b/include/llama.h @@ -46,10 +46,10 @@ #define LLAMA_FILE_MAGIC_GGSQ 0x67677371u // 'ggsq' #define LLAMA_SESSION_MAGIC LLAMA_FILE_MAGIC_GGSN -#define LLAMA_SESSION_VERSION 9 +#define LLAMA_SESSION_VERSION 10 #define LLAMA_STATE_SEQ_MAGIC LLAMA_FILE_MAGIC_GGSQ -#define LLAMA_STATE_SEQ_VERSION 2 +#define LLAMA_STATE_SEQ_VERSION 3 #ifdef __cplusplus extern "C" { diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index ea1d720cb..ea54f9461 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -12,6 +12,7 @@ #include #include #include +#include static bool ggml_is_power_of_2(int n) { return (n & (n - 1)) == 0; @@ -1133,11 +1134,18 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & cells.pos_set(idx, ubatch.pos[i]); - if (ubatch.is_pos_2d()) { - llama_kv_cell_ext ext { - /*.x =*/ ubatch.pos[i + ubatch.n_tokens*2], - /*.y =*/ ubatch.pos[i + ubatch.n_tokens], - }; + if (ubatch.is_pos_2d() || ubatch.token) { + llama_kv_cell_ext ext; + + if (ubatch.is_pos_2d()) { + ext.x = ubatch.pos[i + ubatch.n_tokens*2]; + ext.y = ubatch.pos[i + ubatch.n_tokens]; + } + + if (ubatch.token) { + ext.tok = ubatch.token[i]; + } + cells.ext_set(idx, ext); } @@ -1810,6 +1818,69 @@ void llama_kv_cache::set_input_v_rot(ggml_tensor * dst) const { memcpy(dst->data, attn_rot_hadamard.at(n_rot).data(), ggml_nbytes(dst)); } +bool llama_kv_cache::has_cell_ext() const { + return hparams.n_pos_per_embd() > 1; +} + +void llama_kv_cache::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const { + const uint32_t n_tokens = ubatch.n_tokens; + + res.clear(); + res.resize(n_tokens*n, LLAMA_TOKEN_NULL); + + if (n == 0) { + return; + } + + // note: apply_ubatch() has already stored the current ubatch + // the window below thus covers tokens of this very ubatch as well, which is what we want + llama_pos p_min = std::numeric_limits::max(); + llama_pos p_max = std::numeric_limits::min(); + + std::bitset seqs; + + for (uint32_t i = 0; i < n_tokens; ++i) { + p_min = std::min(p_min, ubatch.pos[i]); + p_max = std::max(p_max, ubatch.pos[i]); + } + + for (uint32_t s = 0; s < ubatch.n_seqs_unq; ++s) { + seqs.set(ubatch.seq_id_unq[s]); + } + + // (seq_id, pos) -> token, for every cell that could be a predecessor of a ubatch token + std::unordered_map hist; + + const auto key = [](llama_seq_id seq_id, llama_pos pos) { + return ((uint64_t) seq_id << 32) | (uint32_t) pos; + }; + + for (uint32_t s = 0; s < n_stream; ++s) { + v_cells[s].for_each_token_in(seqs, p_min - (llama_pos) n, p_max, + [&](llama_seq_id seq_id, llama_pos pos, llama_token tok) { + hist[key(seq_id, pos)] = tok; + }); + } + + for (uint32_t i = 0; i < n_tokens; ++i) { + // TODO: a token that belongs to more than one sequence has an ambiguous history. + // the n-gram architectures have to reject such batches + const llama_seq_id seq_id = ubatch.seq_id[i][0]; + + for (uint32_t j = 0; j < n; ++j) { + const llama_pos p = ubatch.pos[i] - (llama_pos) (n - j); + if (p < 0) { + continue; + } + + const auto it = hist.find(key(seq_id, p)); + if (it != hist.end()) { + res[i*n + j] = it->second; + } + } + } +} + size_t llama_kv_cache::total_size() const { size_t size = 0; @@ -2111,7 +2182,7 @@ void llama_kv_cache::state_write_meta(llama_io_write_i & io, const cell_ranges_t io.write(&pos, sizeof(pos)); io.write(&n_seq_id, sizeof(n_seq_id)); - if (hparams.n_pos_per_embd() > 1) { + if (has_cell_ext()) { const llama_kv_cell_ext ext = cells.ext_get(i); io.write(&ext, sizeof(ext)); } @@ -2248,12 +2319,17 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 return false; } - if (hparams.n_pos_per_embd() > 1) { + if (has_cell_ext()) { llama_kv_cell_ext ext; io.read(&ext, sizeof(ext)); - ubatch.pos[i + ubatch.n_tokens] = ext.y; - ubatch.pos[i + ubatch.n_tokens*2] = ext.x; + if (hparams.n_pos_per_embd() > 1) { + ubatch.pos[i + ubatch.n_tokens] = ext.y; + ubatch.pos[i + ubatch.n_tokens*2] = ext.x; + } + + // apply_ubatch() below restores ext.tok from the ubatch tokens + ubatch.token[i] = ext.tok; } // read the sequence id, but directly discard it - we will use dest_seq_id instead @@ -2273,7 +2349,8 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 return false; } - // TODO: we cannot yet restore llama_kv_cell_ext as the apply_ubatch() does not support it yet + // note: apply_ubatch() rebuilds llama_kv_cell_ext from the ubatch + // only ext.tok and the M-RoPE 2D position round-trip through it // see: https://github.com/ggml-org/llama.cpp/pull/16825#issuecomment-3460868350 apply_ubatch(sinfo, ubatch); @@ -2306,7 +2383,7 @@ bool llama_kv_cache::state_read_meta(llama_io_read_i & io, uint32_t strm, uint32 cells.pos_set(i, pos); - if (hparams.n_pos_per_embd() > 1) { + if (has_cell_ext()) { llama_kv_cell_ext ext; io.read(&ext, sizeof(ext)); cells.ext_set(i, ext); @@ -2657,3 +2734,7 @@ void llama_kv_cache_context::set_input_k_rot(ggml_tensor * dst) const { void llama_kv_cache_context::set_input_v_rot(ggml_tensor * dst) const { kv->set_input_v_rot(dst); } + +void llama_kv_cache_context::get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const { + kv->get_prev_tokens(ubatch, n, res); +} diff --git a/src/llama-kv-cache.h b/src/llama-kv-cache.h index 6cb6dbd2f..9b225fae3 100644 --- a/src/llama-kv-cache.h +++ b/src/llama-kv-cache.h @@ -219,6 +219,14 @@ public: void set_input_k_rot(ggml_tensor * dst) const; void set_input_v_rot(ggml_tensor * dst) const; + // true if llama_kv_cell_ext holds information that has to survive a state save/restore + bool has_cell_ext() const; + + // for every token of the ubatch, the ids of the n tokens that precede it in its sequence + // entries with no matching cell are set to LLAMA_TOKEN_NULL + // note: used by n-gram input embeddings + void get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const; + private: const llama_model & model; const llama_hparams & hparams; @@ -401,6 +409,9 @@ public: void set_input_k_rot(ggml_tensor * dst) const; void set_input_v_rot(ggml_tensor * dst) const; + // see llama_kv_cache::get_prev_tokens() + void get_prev_tokens(const llama_ubatch & ubatch, uint32_t n, std::vector & res) const; + private: llama_memory_status status; diff --git a/src/llama-kv-cells.h b/src/llama-kv-cells.h index fddd31a0b..5167c037d 100644 --- a/src/llama-kv-cells.h +++ b/src/llama-kv-cells.h @@ -15,6 +15,10 @@ struct llama_kv_cell_ext { llama_pos x = 0; llama_pos y = 0; + // when tok = LLAMA_TOKEN_NULL when the cell is produced by embedding input (i.e. multimodal) + // use case: n-gram embeddings hash + llama_token tok = LLAMA_TOKEN_NULL; + // return true if the current 2D spatial position is greater than other bool is_2d_gt(llama_pos ox, llama_pos oy) const { return (y > oy) || (y == oy && x > ox); @@ -23,7 +27,7 @@ struct llama_kv_cell_ext { void reset() { static_assert(std::is_trivially_copyable_v); - memset(this, 0, sizeof(*this)); + *this = llama_kv_cell_ext{}; } }; @@ -305,6 +309,29 @@ public: return seq[i].test(seq_id); } + // gather the token ids of the cells in `seqs` with position in [p0, p1) + // the callback receives (seq_id, pos, token) for every such (cell, seq) pair + // note: used by n-gram input embeddings to recover the tokens preceding a ubatch + template + void for_each_token_in(const std::bitset & seqs, llama_pos p0, llama_pos p1, F && f) const { + for (const auto & i : used) { + if (pos[i] < p0 || pos[i] >= p1) { + continue; + } + + const auto m = seq[i] & seqs; + if (m.none()) { + continue; + } + + for (llama_seq_id s = 0; s < LLAMA_MAX_SEQ; ++s) { + if (m.test(s)) { + f(s, pos[i], ext[i].tok); + } + } + } + } + // note: call only if the cell is not empty and the seq_id is not in the cell void seq_add(uint32_t i, llama_seq_id seq_id) { assert(i < pos.size()); diff --git a/src/models/minimax-01.cpp b/src/models/minimax-01.cpp index a6ccee191..f14626b2c 100644 --- a/src/models/minimax-01.cpp +++ b/src/models/minimax-01.cpp @@ -174,11 +174,9 @@ public: bool can_reuse(const llm_graph_params & params) override { bool res = true; - if (params.ubatch.n_seq_tokens > 1) { - res &= ( inp_q_decay && inp_q_decay->ne[2] == params.ubatch.n_seq_tokens); - res &= ( inp_k_decay && inp_k_decay->ne[2] == params.ubatch.n_seq_tokens); - res &= (inp_diag_decay && inp_diag_decay->ne[1] == params.ubatch.n_seq_tokens); - } + res &= ( inp_q_decay && inp_q_decay->ne[2] == params.ubatch.n_seq_tokens); + res &= ( inp_k_decay && inp_k_decay->ne[2] == params.ubatch.n_seq_tokens); + res &= (inp_diag_decay && inp_diag_decay->ne[1] == params.ubatch.n_seq_tokens); return res; } @@ -223,19 +221,17 @@ llama_model_minimax_01::graph::graph(const llama_model & model, const llm_graph_ ggml_set_input(inp->inp_slopes); cb(inp->inp_slopes, "slopes", -1); - if (n_seq_tokens != 1) { - inp->inp_q_decay = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, n_head, n_seq_tokens, n_seqs); - ggml_set_input(inp->inp_q_decay); - cb(inp->inp_q_decay, "q_decay_exp", -1); + inp->inp_q_decay = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, n_head, n_seq_tokens, n_seqs); + ggml_set_input(inp->inp_q_decay); + cb(inp->inp_q_decay, "q_decay_exp", -1); - inp->inp_k_decay = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, n_head, n_seq_tokens, n_seqs); - ggml_set_input(inp->inp_k_decay); - cb(inp->inp_k_decay, "k_decay_exp", -1); + inp->inp_k_decay = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, n_head, n_seq_tokens, n_seqs); + ggml_set_input(inp->inp_k_decay); + cb(inp->inp_k_decay, "k_decay_exp", -1); - inp->inp_diag_decay = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, n_seq_tokens, n_seq_tokens, n_head, n_seqs); - ggml_set_input(inp->inp_diag_decay); - cb(inp->inp_diag_decay, "diag_decay_exp", -1); - } + inp->inp_diag_decay = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, n_seq_tokens, n_seq_tokens, n_head, n_seqs); + ggml_set_input(inp->inp_diag_decay); + cb(inp->inp_diag_decay, "diag_decay_exp", -1); la = (llm_graph_input_la *) res->add_input(std::move(inp)); @@ -319,41 +315,8 @@ llama_model_minimax_01::graph::graph(const llama_model & model, const llm_graph_ ggml_tensor * qkv = nullptr; ggml_tensor * kv_new = nullptr; - - if (n_seq_tokens == 1) { - // lightning attention - optimized single token case for TG - - ggml_tensor * slopes_neg = ggml_scale(ctx0, slope_rate, -1.0); - cb(slopes_neg, "slopes_neg", il); - - ggml_tensor * ratio = ggml_exp(ctx0, slopes_neg); - cb(ratio, "ratio", il); - - ggml_tensor * ratio_3d = ggml_reshape_3d(ctx0, ratio, 1, 1, n_head); - cb(ratio_3d, "ratio3d", il); - - ggml_tensor * v_trans = ggml_cont(ctx0, ggml_permute(ctx0, Vcur, 1, 2, 0, 3)); - cb(v_trans, "v_trans", il); - - ggml_tensor * k_trans = ggml_cont(ctx0, ggml_permute(ctx0, Kcur, 1, 2, 0, 3)); - cb(k_trans, "k_trans", il); - - ggml_tensor * kv_cur = ggml_mul_mat(ctx0, k_trans, v_trans); - cb(kv_cur, "kv_cur", il); - - ggml_tensor * kv_old_s = ggml_mul(ctx0, kv_old, ratio_3d); - cb(kv_old_s, "kv_old_s", il); - - kv_new = ggml_add(ctx0, kv_old_s, kv_cur); - cb(kv_new, "kv_new", il); - - ggml_tensor * q_trans = ggml_permute(ctx0, Qcur, 0, 2, 1, 3); - cb(q_trans, "q_trans", il); - - qkv = ggml_mul_mat(ctx0, kv_new, q_trans); - cb(qkv, "qkv", il); - } else if(n_seq_tokens > 1) { - // lightning attention - general multi token case for PP + { + // lightning attention ggml_tensor * q_decay_exp = la->inp_q_decay; ggml_tensor * k_decay_exp = la->inp_k_decay; diff --git a/src/models/nanbeige.cpp b/src/models/nanbeige.cpp index 3a546600f..7d5a6bbdb 100644 --- a/src/models/nanbeige.cpp +++ b/src/models/nanbeige.cpp @@ -103,6 +103,7 @@ llama_model_nanbeige::graph::graph(const llama_model & model, const llm_graph_pa 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; cur = build_norm(inpL, model.layers[il].attn_norm, NULL, LLM_NORM_RMS, il); diff --git a/tools/mtmd/mtmd-cli.cpp b/tools/mtmd/mtmd-cli.cpp index f6c787fdb..97678c6b2 100644 --- a/tools/mtmd/mtmd-cli.cpp +++ b/tools/mtmd/mtmd-cli.cpp @@ -87,6 +87,9 @@ struct mtmd_cli_context { mtmd::bitmaps bitmaps; std::vector videos; + mtmd_helper_init_opt init_opt = mtmd_helper_init_opt_default(); + std::string video_ffmpeg_bin_dir; + mtmd::batch_ptr mbatch; // chat template @@ -170,6 +173,12 @@ struct mtmd_cli_context { LOG_ERR("Failed to load vision model from %s\n", clip_path); exit(1); } + + video_ffmpeg_bin_dir = params.video_ffmpeg_bin_dir; + init_opt.video_params.fps_target = params.video_fps; + init_opt.video_params.timestamp_interval_ms = params.video_timestamp_interval_ms; + init_opt.video_params.ffmpeg_bin_dir = video_ffmpeg_bin_dir.empty() + ? nullptr : video_ffmpeg_bin_dir.c_str(); } bool check_antiprompt(const llama_tokens & generated_tokens) { @@ -184,7 +193,7 @@ struct mtmd_cli_context { } bool load_media(const std::string & fname) { - auto res = mtmd_helper_bitmap_init_from_file(ctx_vision.get(), fname.c_str(), false); + auto res = mtmd_helper_bitmap_init_from_file(ctx_vision.get(), fname.c_str(), false, init_opt); if (!res.bitmap) { return false; } diff --git a/tools/mtmd/mtmd-helper.cpp b/tools/mtmd/mtmd-helper.cpp index 42c52859f..30bd3c39b 100644 --- a/tools/mtmd/mtmd-helper.cpp +++ b/tools/mtmd/mtmd-helper.cpp @@ -370,14 +370,18 @@ static bool is_webp_file(const unsigned char * buf, size_t len) { } #ifdef MTMD_VIDEO -static mtmd_bitmap * decode_webp_with_ffmpeg(mtmd_context * mctx, const unsigned char * buf, size_t len, bool placeholder); +static mtmd_bitmap * decode_webp_with_ffmpeg(mtmd_context * mctx, const unsigned char * buf, size_t len, bool placeholder, + const mtmd_helper_video_init_params & params); #endif -mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_buf(mtmd_context * ctx, const unsigned char * buf, size_t len, bool placeholder) { +mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_buf(mtmd_context * ctx, const unsigned char * buf, size_t len, bool placeholder, + mtmd_helper_init_opt opt) { // calculate the hash if needed std::string id; mtmd_bitmap * result = nullptr; + GGML_UNUSED(opt); // only used by video code paths + if (!placeholder) { // use sha256 to prevent cache poisoning id = hash_sha256_hex(buf, len); @@ -415,7 +419,7 @@ mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_buf(mtmd_context * ctx, #ifdef MTMD_VIDEO // stb_image does not support webp; decode it with ffmpeg as a single frame if (!result && is_webp_file(buf, len)) { - result = decode_webp_with_ffmpeg(ctx, buf, len, placeholder); + result = decode_webp_with_ffmpeg(ctx, buf, len, placeholder, opt.video_params); if (!result) { LOG_ERR("%s: failed to decode webp buffer\n", __func__); return {nullptr, nullptr}; @@ -428,8 +432,7 @@ mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_buf(mtmd_context * ctx, // last try: load as video #ifdef MTMD_VIDEO if (!result) { - auto params = mtmd_helper_video_init_params_default(); - auto video_ctx = mtmd_helper_video_init_from_buf(ctx, buf, len, params); + auto video_ctx = mtmd_helper_video_init_from_buf(ctx, buf, len, opt.video_params); if (!video_ctx) { LOG_ERR("%s: failed to decode buffer as either image/audio/video\n", __func__); return {nullptr, nullptr}; @@ -457,7 +460,8 @@ mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_buf(mtmd_context * ctx, return {nullptr, nullptr}; } -mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_file(mtmd_context * ctx, const char * fname, bool placeholder) { +mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_file(mtmd_context * ctx, const char * fname, bool placeholder, + mtmd_helper_init_opt opt) { #ifdef _WIN32 int wlen = MultiByteToWideChar(CP_UTF8, 0, fname, -1, NULL, 0); if (!wlen) { @@ -498,7 +502,7 @@ mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_file(mtmd_context * ctx, return {nullptr, nullptr}; } - return mtmd_helper_bitmap_init_from_buf(ctx, buf.data(), buf.size(), placeholder); + return mtmd_helper_bitmap_init_from_buf(ctx, buf.data(), buf.size(), placeholder, opt); } bool mtmd_helper_support_video(mtmd_context * ctx) { @@ -856,6 +860,12 @@ mtmd_helper_video_init_params mtmd_helper_video_init_params_default() { }; } +mtmd_helper_init_opt mtmd_helper_init_opt_default() { + return { + /* video_params */ mtmd_helper_video_init_params_default(), + }; +} + static std::string video_resolve_bin(const char * bin_dir, const char * name) { if (!bin_dir || bin_dir[0] == '\0') { return name; // rely on PATH @@ -877,8 +887,8 @@ static std::string video_resolve_bin(const char * bin_dir, const char * name) { } #ifdef MTMD_VIDEO -static mtmd_bitmap * decode_webp_with_ffmpeg(mtmd_context * mctx, const unsigned char * buf, size_t len, bool placeholder) { - auto params = mtmd_helper_video_init_params_default(); +static mtmd_bitmap * decode_webp_with_ffmpeg(mtmd_context * mctx, const unsigned char * buf, size_t len, bool placeholder, + const mtmd_helper_video_init_params & params) { mtmd_helper_video vctx; vctx.mctx = mctx; vctx.input_buf.assign(buf, buf + len); diff --git a/tools/mtmd/mtmd-helper.h b/tools/mtmd/mtmd-helper.h index 58dfb1525..772e0f091 100644 --- a/tools/mtmd/mtmd-helper.h +++ b/tools/mtmd/mtmd-helper.h @@ -23,6 +23,23 @@ extern "C" { struct mtmd_helper_video; typedef struct mtmd_helper_video mtmd_helper_video; +struct mtmd_helper_video_init_params { + float fps_target; // desired output fps; <= 0 means use the video's native fps, defaulted to 4.0f + const char * ffmpeg_bin_dir; // directory containing ffmpeg/ffprobe binaries; NULL means search PATH + int64_t timestamp_interval_ms; // interval for adding timestamp as text chunk (example: "[10m50.5s]"); <= 0 means no timestamp, defaulted to 5000ms + // TODO @ngxson : allow "placeholder" bitmap output for counting tokens +}; + +MTMD_API struct mtmd_helper_video_init_params mtmd_helper_video_init_params_default(void); + +// opt for mtmd_helper_bitmap_init_from_*() +struct mtmd_helper_init_opt { + struct mtmd_helper_video_init_params video_params; +}; +typedef struct mtmd_helper_init_opt mtmd_helper_init_opt; + +MTMD_API struct mtmd_helper_init_opt mtmd_helper_init_opt_default(void); + // Set callback for all future logging events. // If this is not called, or NULL is supplied, everything is output on stderr. // Note: this also call mtmd_log_set() internally @@ -40,7 +57,11 @@ struct mtmd_helper_bitmap_wrapper { // it calls mtmd_helper_bitmap_init_from_buf() internally // returns nullptr on failure // this function is thread-safe -MTMD_API struct mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_file(mtmd_context * ctx, const char * fname, bool placeholder); +MTMD_API struct mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_file( + mtmd_context * ctx, + const char * fname, + bool placeholder, + struct mtmd_helper_init_opt opt); // helper function to construct a mtmd_bitmap from a buffer containing a file // supported formats: @@ -53,7 +74,11 @@ MTMD_API struct mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_file(mtm // - output bitmap will have SHA-256 hash (hex string) as the ID // returns nullptr on failure // this function is thread-safe -MTMD_API struct mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_buf(mtmd_context * ctx, const unsigned char * buf, size_t len, bool placeholder); +MTMD_API struct mtmd_helper_bitmap_wrapper mtmd_helper_bitmap_init_from_buf( + mtmd_context * ctx, + const unsigned char * buf, size_t len, + bool placeholder, + struct mtmd_helper_init_opt opt); // helper to count the total number of tokens from a list of chunks, useful to keep track of KV cache MTMD_API size_t mtmd_helper_get_n_tokens(const mtmd_input_chunks * chunks); @@ -124,14 +149,7 @@ struct mtmd_helper_video_info { int32_t n_frames; // estimated total frames at effective fps (-1 if unknown) }; -struct mtmd_helper_video_init_params { - float fps_target; // desired output fps; <= 0 means use the video's native fps, defaulted to 4.0f - const char * ffmpeg_bin_dir; // directory containing ffmpeg/ffprobe binaries; NULL means search PATH - int64_t timestamp_interval_ms; // interval for adding timestamp as text chunk (example: "[10m50.5s]"); <= 0 means no timestamp, defaulted to 5000ms - // TODO @ngxson : allow "placeholder" bitmap output for counting tokens -}; - -MTMD_API struct mtmd_helper_video_init_params mtmd_helper_video_init_params_default(void); +// note: mtmd_helper_video_init_params is defined at the top, as it is part of mtmd_helper_init_opt // returns NULL on failure (ffprobe not found, file unreadable, etc.) MTMD_API mtmd_helper_video * mtmd_helper_video_init( diff --git a/tools/server/server-common.cpp b/tools/server/server-common.cpp index 4f5b8202a..c30955e89 100644 --- a/tools/server/server-common.cpp +++ b/tools/server/server-common.cpp @@ -910,12 +910,17 @@ size_t validate_utf8(const std::string& text) { return len; } -server_tokens process_mtmd_prompt(mtmd_context * mctx, const std::string & prompt, const std::vector & files, bool is_placeholder) { +server_tokens process_mtmd_prompt( + mtmd_context * mctx, + const std::string & prompt, + const std::vector & files, + const mtmd_helper_init_opt & init_opt, + bool is_placeholder) { // these will be freed upon going out of scope mtmd::bitmaps bitmaps; std::vector videos; for (auto & file : files) { - auto out = mtmd_helper_bitmap_init_from_buf(mctx, file.data(), file.size(), is_placeholder); + auto out = mtmd_helper_bitmap_init_from_buf(mctx, file.data(), file.size(), is_placeholder, init_opt); if (!out.bitmap) { throw std::runtime_error("Failed to load image or audio file"); } @@ -956,7 +961,7 @@ server_tokens process_mtmd_prompt(mtmd_context * mctx, const std::string & promp * - "prompt": [12, 34, "string", 56, 78] * - "prompt": { "prompt_string": "string", "multimodal_data": [ "base64" ] } */ -static server_tokens tokenize_input_subprompt(const llama_vocab * vocab, mtmd_context * mctx, const json & json_prompt, bool add_special, bool parse_special) { +static server_tokens tokenize_input_subprompt(const llama_vocab * vocab, mtmd_context * mctx, const json & json_prompt, bool add_special, bool parse_special, const mtmd_helper_init_opt & init_opt) { constexpr char JSON_STRING_PROMPT_KEY[] = "prompt_string"; constexpr char JSON_MTMD_DATA_KEY[] = "multimodal_data"; const bool has_mtmd = mctx != nullptr; @@ -979,7 +984,7 @@ static server_tokens tokenize_input_subprompt(const llama_vocab * vocab, mtmd_co for (const auto & entry : json_prompt.at(JSON_MTMD_DATA_KEY)) { files.push_back(base64_decode(entry)); } - return process_mtmd_prompt(mctx, json_prompt.at(JSON_STRING_PROMPT_KEY), files); + return process_mtmd_prompt(mctx, json_prompt.at(JSON_STRING_PROMPT_KEY), files, init_opt); } else { // Not multimodal, but contains a subobject. llama_tokens tmp = tokenize_mixed(vocab, json_prompt.at(JSON_STRING_PROMPT_KEY), add_special, parse_special); @@ -990,15 +995,15 @@ static server_tokens tokenize_input_subprompt(const llama_vocab * vocab, mtmd_co } } -std::vector tokenize_input_prompts(const llama_vocab * vocab, mtmd_context * mctx, const json & json_prompt, bool add_special, bool parse_special) { +std::vector tokenize_input_prompts(const llama_vocab * vocab, mtmd_context * mctx, const json & json_prompt, bool add_special, bool parse_special, const mtmd_helper_init_opt & init_opt) { std::vector result; if (json_prompt.is_array() && !json_is_array_and_contains_numbers(json_prompt)) { result.reserve(json_prompt.size()); for (const auto & p : json_prompt) { - result.push_back(tokenize_input_subprompt(vocab, mctx, p,add_special, parse_special)); + result.push_back(tokenize_input_subprompt(vocab, mctx, p, add_special, parse_special, init_opt)); } } else { - result.push_back(tokenize_input_subprompt(vocab, mctx, json_prompt, add_special, parse_special)); + result.push_back(tokenize_input_subprompt(vocab, mctx, json_prompt, add_special, parse_special, init_opt)); } if (result.empty()) { throw std::runtime_error("\"prompt\" must not be empty"); @@ -1280,6 +1285,12 @@ json oaicompat_chat_params_parse( if (inputs.continue_final_message != COMMON_CHAT_CONTINUATION_NONE && inputs.add_generation_prompt) { throw std::invalid_argument("Cannot set both add_generation_prompt and continue_final_message to true."); } + if (inputs.continue_final_message != COMMON_CHAT_CONTINUATION_NONE + && !inputs.messages.empty() + && inputs.messages.back().role == "assistant" + && !inputs.messages.back().tool_calls.empty()) { + throw std::invalid_argument("Cannot continue an assistant message that contains tool calls."); + } inputs.reasoning_format = opt.reasoning_format; if (body.contains("reasoning_format")) { inputs.reasoning_format = common_reasoning_format_from_name(body.at("reasoning_format").get()); @@ -1781,7 +1792,8 @@ server_tokens format_prompt_rerank( const struct llama_vocab * vocab, mtmd_context * mctx, const std::string & query, - const std::string & doc) { + const std::string & doc, + const mtmd_helper_init_opt & init_opt) { server_tokens result = {}; const char * rerank_prompt = llama_model_chat_template(model, "rerank"); @@ -1790,12 +1802,12 @@ server_tokens format_prompt_rerank( std::string prompt = rerank_prompt; string_replace_all(prompt, "{query}" , query); string_replace_all(prompt, "{document}", doc ); - server_tokens tokens = tokenize_input_subprompt(vocab, mctx, prompt, false, true); + server_tokens tokens = tokenize_input_subprompt(vocab, mctx, prompt, false, true, init_opt); result.push_back(tokens); } else { // Get EOS token - use SEP token as fallback if EOS is not available - server_tokens query_tokens = tokenize_input_subprompt(vocab, mctx, query, false, false); - server_tokens doc_tokens = tokenize_input_subprompt(vocab, mctx, doc, false, false); + server_tokens query_tokens = tokenize_input_subprompt(vocab, mctx, query, false, false, init_opt); + server_tokens doc_tokens = tokenize_input_subprompt(vocab, mctx, doc, false, false, init_opt); llama_token eos_token = llama_vocab_eos(vocab); if (eos_token == LLAMA_TOKEN_NULL) { eos_token = llama_vocab_sep(vocab); diff --git a/tools/server/server-common.h b/tools/server/server-common.h index f8ea82ef4..6c681a2cf 100644 --- a/tools/server/server-common.h +++ b/tools/server/server-common.h @@ -5,6 +5,7 @@ #include "llama.h" #include "chat.h" #include "mtmd.h" +#include "mtmd-helper.h" #include "json.h" @@ -269,7 +270,12 @@ size_t validate_utf8(const std::string& text); // process mtmd prompt, return the server_tokens containing both text tokens and media chunks // if is_placeholder is true, the media chunk will be treated as placeholder for counting tokens; the output tokens are not usable for actual inference (e.g. for submitting a task to server_queue) -server_tokens process_mtmd_prompt(mtmd_context * mctx, const std::string & prompt, const std::vector & files, bool is_placeholder = false); +server_tokens process_mtmd_prompt( + mtmd_context * mctx, + const std::string & prompt, + const std::vector & files, + const mtmd_helper_init_opt & init_opt, + bool is_placeholder = false); /** * break the input "prompt" object into multiple prompt if needed, then tokenize them @@ -289,7 +295,8 @@ std::vector tokenize_input_prompts( mtmd_context * mctx, const json & json_prompt, bool add_special, - bool parse_special); + bool parse_special, + const mtmd_helper_init_opt & init_opt); // // OAI utils @@ -538,7 +545,8 @@ server_tokens format_prompt_rerank( const struct llama_vocab * vocab, mtmd_context * mctx, const std::string & query, - const std::string & doc); + const std::string & doc, + const mtmd_helper_init_opt & init_opt); // simple implementation of a pipe // used for streaming data between threads diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index a9edbd7be..9fdfcae56 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -794,6 +794,8 @@ public: llama_model * model_tgt = nullptr; mtmd_context * mctx = nullptr; + // note: video_params.ffmpeg_bin_dir points into params_base, which outlives this struct + mtmd_helper_init_opt init_opt = mtmd_helper_init_opt_default(); const llama_vocab * vocab = nullptr; server_queue queue_tasks; @@ -1118,6 +1120,11 @@ private: } SRV_INF("loaded multimodal model, '%s'\n", mmproj_path.c_str()); + init_opt.video_params.fps_target = params_base.video_fps; + init_opt.video_params.timestamp_interval_ms = params_base.video_timestamp_interval_ms; + init_opt.video_params.ffmpeg_bin_dir = params_base.video_ffmpeg_bin_dir.empty() + ? nullptr : params_base.video_ffmpeg_bin_dir.c_str(); + if (params_base.ctx_shift) { params_base.ctx_shift = false; SRV_WRN("%s\n", "ctx_shift is not supported by multimodal, it will be disabled"); @@ -2134,9 +2141,9 @@ private: try { auto & prompt = task.cli_prompt; if (mctx != nullptr) { - task.tokens = process_mtmd_prompt(mctx, prompt, task.cli_files); + task.tokens = process_mtmd_prompt(mctx, prompt, task.cli_files, init_opt); } else { - task.tokens = std::move(tokenize_input_prompts(vocab, mctx, prompt, true, true)[0]); + task.tokens = std::move(tokenize_input_prompts(vocab, mctx, prompt, true, true, init_opt)[0]); } task.cli_prompt.clear(); task.cli_files.clear(); @@ -4165,10 +4172,10 @@ std::unique_ptr server_routes::handle_completions_impl( if (res_type != TASK_RESPONSE_TYPE_NONE && ctx_server.mctx != nullptr) { // This is the case used by OAI compatible chat path with MTMD. TODO It can be moved to the path below. - inputs.push_back(process_mtmd_prompt(ctx_server.mctx, prompt.get(), files)); + inputs.push_back(process_mtmd_prompt(ctx_server.mctx, prompt.get(), files, ctx_server.init_opt)); } else { // Everything else, including multimodal completions. - inputs = tokenize_input_prompts(ctx_server.vocab, ctx_server.mctx, prompt, true, true); + inputs = tokenize_input_prompts(ctx_server.vocab, ctx_server.mctx, prompt, true, true, ctx_server.init_opt); } // tasks.reserve(inputs.size()); // TODO: this is inaccurate due to child tasks @@ -4752,7 +4759,7 @@ void server_routes::init_routes() { data["input_extra"] = input_extra; // default to empty array if it's not exist std::string prompt = json_value(data, "prompt", std::string()); - std::vector tokenized_prompts = tokenize_input_prompts(ctx_server.vocab, ctx_server.mctx, prompt, false, true); + std::vector tokenized_prompts = tokenize_input_prompts(ctx_server.vocab, ctx_server.mctx, prompt, false, true, ctx_server.init_opt); SRV_DBG("creating infill tasks, n_prompts = %d\n", (int) tokenized_prompts.size()); data["prompt"] = format_prompt_infill( ctx_server.vocab, @@ -4816,7 +4823,7 @@ void server_routes::init_routes() { }; this->post_chat_completions_tok = [this](const server_http_req & req) { - return handle_count_tokens(ctx_server.vocab, ctx_server.mctx, req, TASK_RESPONSE_TYPE_OAI_CHAT); + return handle_count_tokens(ctx_server.vocab, ctx_server.mctx, ctx_server.init_opt, req, TASK_RESPONSE_TYPE_OAI_CHAT); }; this->post_control = [this](const server_http_req & req) { @@ -4875,7 +4882,7 @@ void server_routes::init_routes() { }; this->post_responses_tok_oai = [this](const server_http_req & req) { - return handle_count_tokens(ctx_server.vocab, ctx_server.mctx, req, TASK_RESPONSE_TYPE_OAI_RESP); + return handle_count_tokens(ctx_server.vocab, ctx_server.mctx, ctx_server.init_opt, req, TASK_RESPONSE_TYPE_OAI_RESP); }; this->post_transcriptions_oai = [this](const server_http_req & req) { @@ -4925,7 +4932,7 @@ void server_routes::init_routes() { }; this->post_anthropic_count_tokens = [this](const server_http_req & req) { - return handle_count_tokens(ctx_server.vocab, ctx_server.mctx, req, TASK_RESPONSE_TYPE_ANTHROPIC); + return handle_count_tokens(ctx_server.vocab, ctx_server.mctx, ctx_server.init_opt, req, TASK_RESPONSE_TYPE_ANTHROPIC); }; // same with handle_chat_completions, but without inference part @@ -5058,7 +5065,7 @@ void server_routes::init_routes() { std::vector tasks; tasks.reserve(documents.size()); for (size_t i = 0; i < documents.size(); i++) { - auto tmp = format_prompt_rerank(ctx_server.model_tgt, ctx_server.vocab, ctx_server.mctx, query, documents[i]); + auto tmp = format_prompt_rerank(ctx_server.model_tgt, ctx_server.vocab, ctx_server.mctx, query, documents[i], ctx_server.init_opt); server_task task = server_task(SERVER_TASK_TYPE_RERANK); task.id = rd.get_new_id(); task.tokens = std::move(tmp); @@ -5296,7 +5303,7 @@ std::unique_ptr server_routes::handle_embeddings_impl(cons } } - auto tokenized_prompts = tokenize_input_prompts(ctx_server.vocab, ctx_server.mctx, prompt, true, true); + auto tokenized_prompts = tokenize_input_prompts(ctx_server.vocab, ctx_server.mctx, prompt, true, true, ctx_server.init_opt); for (const auto & tokens : tokenized_prompts) { // this check is necessary for models that do not add BOS token to the input if (tokens.empty()) { @@ -5357,7 +5364,7 @@ std::unique_ptr server_routes::handle_embeddings_impl(cons return res; } -std::unique_ptr server_routes::handle_count_tokens(const llama_vocab * vocab, mtmd_context * mctx, const server_http_req & req, task_response_type res_type) { +std::unique_ptr server_routes::handle_count_tokens(const llama_vocab * vocab, mtmd_context * mctx, const mtmd_helper_init_opt & init_opt, const server_http_req & req, task_response_type res_type) { auto res = create_response(); std::vector files; json body = json::parse(req.body); @@ -5395,7 +5402,7 @@ std::unique_ptr server_routes::handle_count_tokens(const l if (!prompt.is_string()) { throw std::runtime_error("for mtmd, input prompt must be a string."); } - n_tokens = process_mtmd_prompt(mctx, prompt.get(), files, true).size(); + n_tokens = process_mtmd_prompt(mctx, prompt.get(), files, init_opt, true).size(); } else { n_tokens = tokenize_mixed(vocab, prompt, true, true).size(); } diff --git a/tools/server/server-context.h b/tools/server/server-context.h index 5d464b8e8..0acbbffa9 100644 --- a/tools/server/server-context.h +++ b/tools/server/server-context.h @@ -169,7 +169,7 @@ private: std::unique_ptr handle_slots_restore(const server_http_req & req, int id_slot); std::unique_ptr handle_slots_erase(const server_http_req &, int id_slot); std::unique_ptr handle_embeddings_impl(const server_http_req & req, task_response_type res_type); - std::unique_ptr handle_count_tokens(const llama_vocab * vocab, mtmd_context * mctx, const server_http_req & req, task_response_type res_type); + std::unique_ptr handle_count_tokens(const llama_vocab * vocab, mtmd_context * mctx, const mtmd_helper_init_opt & init_opt, const server_http_req & req, task_response_type res_type); // using unique_ptr to allow late initialization of const std::unique_ptr meta; diff --git a/tools/tts/tts.cpp b/tools/tts/tts.cpp index 368123baf..6fd193632 100644 --- a/tools/tts/tts.cpp +++ b/tools/tts/tts.cpp @@ -103,7 +103,7 @@ int main(int argc, char ** argv) { mtmd::bitmap_ptr speaker_bitmap; if (!params.tts_speaker_file.empty()) { - auto wrapper = mtmd_helper_bitmap_init_from_file(mctx.get(), params.tts_speaker_file.c_str(), false); + auto wrapper = mtmd_helper_bitmap_init_from_file(mctx.get(), params.tts_speaker_file.c_str(), false, mtmd_helper_init_opt_default()); if (!wrapper.bitmap) { LOG_ERR("failed to load speaker file %s\n", params.tts_speaker_file.c_str()); return 1; diff --git a/tools/ui/eslint.config.js b/tools/ui/eslint.config.js index 6ad065f5a..9eab1734c 100644 --- a/tools/ui/eslint.config.js +++ b/tools/ui/eslint.config.js @@ -12,6 +12,107 @@ import { fileURLToPath } from 'node:url'; import ts from 'typescript-eslint'; const gitignorePath = fileURLToPath(new URL('./.gitignore', import.meta.url)); +// Require a blank line between sibling element-like nodes in a Svelte template +// (elements, components, and the {#if} / {#each} / {#await} / {#snippet} / +// {@render} blocks) that sit on separate lines at the same nesting level. +// Whitespace between siblings is a whitespace-only SvelteText node; when it +// holds a single newline (no blank line) the fix adds one, keeping the +// indentation of the second sibling. Real text content (e.g. `foo\n\nbar`) +// is left alone. +const ELEMENT_LIKE_TYPES = new Set([ + 'SvelteAwaitBlock', + 'SvelteComponent', + 'SvelteEachBlock', + 'SvelteElement', + 'SvelteIfBlock', + 'SvelteKeyBlock', + 'SvelteRenderTag', + 'SvelteSelf', + 'SvelteSnippetBlock' +]); +const paddingLineBetweenElements = { + create(context) { + // Check one list of template children. Each children array holds the + // element-like nodes plus the whitespace/comment text between them. + function checkChildren(children) { + if (!Array.isArray(children)) return; + + let lastElement = null; + let lastWhitespace = null; + + for (const child of children) { + if (child.type === 'SvelteText' && /^\s*$/.test(child.value)) { + lastWhitespace = child; + + continue; + } + + if (!ELEMENT_LIKE_TYPES.has(child.type)) continue; + + if ( + lastElement && + lastWhitespace && + child.loc.start.line - lastElement.loc.end.line === 1 + ) { + const textNode = lastWhitespace; + + context.report({ + fix(fixer) { + // Add a second newline so the two siblings are separated by a + // blank line, keeping the trailing indentation. + return fixer.replaceText(textNode, textNode.value.replace(/\n/, '\n\n')); + }, + message: 'Expected a blank line between sibling elements.', + node: child + }); + } + + lastElement = child; + lastWhitespace = null; + } + } + + return { + SvelteAwaitBlock(node) { + checkChildren(node.children); + checkChildren(node.then?.children); + checkChildren(node.else?.children); + }, + SvelteComponent(node) { + checkChildren(node.children); + }, + SvelteEachBlock(node) { + checkChildren(node.children); + checkChildren(node.else?.children); + }, + SvelteElement(node) { + checkChildren(node.children); + }, + SvelteFragment(node) { + checkChildren(node.children); + }, + SvelteIfBlock(node) { + checkChildren(node.children); + checkChildren(node.else?.children); + }, + SvelteKeyBlock(node) { + checkChildren(node.children); + }, + SvelteProgram(node) { + checkChildren(node.children); + }, + SvelteSnippetBlock(node) { + checkChildren(node.children); + } + }; + }, + meta: { + docs: { description: 'Require a blank line between sibling elements in a Svelte template.' }, + fixable: 'whitespace', + schema: [], + type: 'layout' + } +}; // Require a blank line between consecutive class accessors (get/set). The core // `padding-line-between-statements` rule only handles statements, not class // members, so this is enforced with a small custom rule. @@ -66,7 +167,12 @@ export default ts.config( { languageOptions: { globals: { ...globals.browser, ...globals.node } }, plugins: { - local: { rules: { 'blank-line-between-accessors': blankLineBetweenAccessors } }, + local: { + rules: { + 'blank-line-between-accessors': blankLineBetweenAccessors, + 'padding-line-between-elements': paddingLineBetweenElements + } + }, perfectionist, 'simple-import-sort': simpleImportSort }, @@ -82,6 +188,8 @@ export default ts.config( 'eol-last': 'error', // Enforce a blank line between consecutive get/set accessors 'local/blank-line-between-accessors': 'error', + // Require a blank line between sibling elements in a Svelte template + 'local/padding-line-between-elements': 'error', // typescript-eslint strongly recommend that you do not use the no-undef lint rule on TypeScript projects. // see: https://typescript-eslint.io/troubleshooting/faqs/eslint/#i-get-errors-from-the-no-undef-rule-about-global-variables-not-being-defined-even-though-there-are-no-typescript-errors 'no-undef': 'off', @@ -156,9 +264,49 @@ export default ts.config( // grouping); Prettier normalizes comma spacing afterwards. 'simple-import-sort/imports': ['error', { groups: [['.*']] }], 'svelte/no-at-html-tags': 'off', - // This app uses hash-based routing (#/) where resolve() from $app/paths does not apply - 'svelte/no-navigation-without-resolve': 'off' + 'svelte/no-navigation-without-resolve': 'off', + + // Sort HTML attributes alphabetically in the markup. The Svelte directives + // (bind:/use:/animate:/style:/in:/out:/transition:/class:) sort first, + // alphabetically among themselves, then all remaining attributes sort + // alphabetically. The rule keeps spread attributes in place and does not cross + // them. `this` stays first on because Prettier forces it there + // - reordering it alphabetically would fight the formatter. + 'svelte/sort-attributes': [ + 'error', + { + order: [ + 'this', + { + match: [ + '/^bind:/u', + '/^use:/u', + '/^animate:/u', + '/^style:/u', + '/^in:/u', + '/^out:/u', + '/^transition:/u', + '/^class:/u' + ], + sort: 'alphabetical' + }, + { + match: [ + '!/^bind:/u', + '!/^use:/u', + '!/^animate:/u', + '!/^style:/u', + '!/^in:/u', + '!/^out:/u', + '!/^transition:/u', + '!/^class:/u' + ], + sort: 'alphabetical' + } + ] + } + ] } }, { diff --git a/tools/ui/src/lib/components/app/actions/ActionIcon.svelte b/tools/ui/src/lib/components/app/actions/ActionIcon.svelte index e29b5ad67..0ed22d932 100644 --- a/tools/ui/src/lib/components/app/actions/ActionIcon.svelte +++ b/tools/ui/src/lib/components/app/actions/ActionIcon.svelte @@ -41,17 +41,17 @@ {#snippet button(props = {})} diff --git a/tools/ui/src/lib/components/app/chat/ChatAttachments/ChatAttachmentsPreview/ChatAttachmentsPreviewThumbnailStrip.svelte b/tools/ui/src/lib/components/app/chat/ChatAttachments/ChatAttachmentsPreview/ChatAttachmentsPreviewThumbnailStrip.svelte index f0ea9675f..e5ba09dba 100644 --- a/tools/ui/src/lib/components/app/chat/ChatAttachments/ChatAttachmentsPreview/ChatAttachmentsPreviewThumbnailStrip.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatAttachments/ChatAttachmentsPreview/ChatAttachmentsPreviewThumbnailStrip.svelte @@ -38,16 +38,16 @@ {#each items as item, index (item.id)} - {/each} - - - {#snippet footer()} - - - - Manage MCP Servers - - {/snippet} - - {:else} -
- No MCP servers configured -
- - - - - - - Add MCP Servers - - {/if} - - - diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddMcpSubmenu.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddMcpSubmenu.svelte new file mode 100644 index 000000000..07439afd6 --- /dev/null +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddMcpSubmenu.svelte @@ -0,0 +1,51 @@ + + + + + + + MCP + + + + + + + Servers + + + {#if chatFormActions.hasMcpPromptsSupport} + + + + Prompts + + {/if} + + {#if chatFormActions.hasMcpResourcesSupport} + + + + Resources + + {/if} + + diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddReasoningSubmenu.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddReasoningSubmenu.svelte index a3a0b3a20..1b6fc4b02 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddReasoningSubmenu.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddReasoningSubmenu.svelte @@ -64,6 +64,7 @@ +

Maximum reasoning effort with extended context usage

diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddSheet.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddSheet.svelte index 2e61bb07d..2f69dc96d 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddSheet.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddSheet.svelte @@ -78,7 +78,7 @@ {@render trigger({ disabled: chatFormActions.disabled, onclick: () => (sheetOpen = true) })} - + Add to chat @@ -90,8 +90,8 @@
{#if reasoning.modelSupportsThinking} (reasoningExpanded = open)} + open={reasoningExpanded} > {#if reasoningExpanded} @@ -120,10 +120,10 @@ {#each reasoning.levels as level (level.value)} {@const tokenLabel = reasoning.tokenLabel(level)} {/each} @@ -330,9 +330,9 @@ {/if} {/snippet} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActions.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActions.svelte index 118e54a0a..f1aa74369 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActions.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActions.svelte @@ -1,6 +1,5 @@
- +
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetailRow.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetailRow.svelte index 271997b39..572d4a42d 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetailRow.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetailRow.svelte @@ -11,6 +11,7 @@
{label} + {value}
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetails.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetails.svelte index 1153c70fd..0de508a61 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetails.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDetails.svelte @@ -57,12 +57,13 @@ {#if cumulativeRead > 0} 0 ? `${cumulativeCacheTotal.toLocaleString()} reused from KV cache` : undefined} + value={`${cumulativeRead.toLocaleString()} tok`} /> {/if} + {#if cumulativeOutput > 0} 0} 0 ? `${currentFresh.toLocaleString()} fresh + ${currentCache.toLocaleString()} cached` : undefined} + value={`${currentRead.toLocaleString()} tok`} /> {/if} @@ -100,6 +101,7 @@
KV cache total + {kvTotal.toLocaleString()} tok
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDial.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDial.svelte index 67d705ae4..32d08323d 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDial.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeDial.svelte @@ -18,7 +18,7 @@ const strokeWidth = $derived(size === 'md' ? 4 : 3); - + diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeLoadModel.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeLoadModel.svelte index 022e626ae..4edc72773 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeLoadModel.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugeLoadModel.svelte @@ -14,11 +14,13 @@ {#if modelId !== null && !isLoading}
Available context size is only visible once the model is loaded. - + +
{:else if isLoading}
+ Loading model...
{/if} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugePopup.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugePopup.svelte index 5d80ecc11..8fa09cf70 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugePopup.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormContextGauge/ContextGaugePopup.svelte @@ -54,17 +54,19 @@ {#if gaugePopup.open}
Context + · + {formatParameters(gauge.contextUsed)} / {gauge.contextTotal !== null ? formatParameters(gauge.contextTotal) : '-'} @@ -73,8 +75,8 @@ {#if gauge.activeModelId !== null && !gauge.isActiveModelLoaded} {:else if showProgressBar} @@ -91,6 +93,7 @@ {gauge.contextPercent}% used + {formatParameters(gauge.contextAvailable ?? 0)} remaining @@ -101,15 +104,15 @@ {#if gauge.hasAnyUsage} {/if} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectory.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectory.svelte index 99b7763e3..307d0e702 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectory.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectory.svelte @@ -323,80 +323,81 @@ - + event.preventDefault()} - onCloseAutoFocus={(event) => event.preventDefault()} class="w-[var(--bits-popover-anchor-width)] max-w-none rounded-xl border-border/50 p-0 shadow-xl" + {customAnchor} + onCloseAutoFocus={(event) => event.preventDefault()} + onOpenAutoFocus={(event) => event.preventDefault()} + onkeydown={handleKeydown} + preventScroll={false} + side="top" + sideOffset={12} >
{#if !fileSearchEnabled}
{searchUnavailableMessage}
{:else if query.trim() && (search.isSearching || queryResults.length > 0 || searchError)} nav.setHover(index)} + rawQuery={query} + results={queryResults} /> {/if} {#if pickerSupported && fileSearchEnabled} {/if} {#if homeBase && fileSearchEnabled} - + Searching in: diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectoryChip.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectoryChip.svelte index 23661d223..5a7b054ce 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectoryChip.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectoryChip.svelte @@ -29,8 +29,8 @@
@@ -42,6 +42,7 @@ {displayLabel} {/snippet} +

{displayLabelTitle}

@@ -56,14 +57,14 @@ class="w-0 overflow-hidden opacity-0 transition-[width,opacity] duration-200 ease-out group-hover:w-auto group-hover:opacity-100" >
{/if} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectoryResultsList.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectoryResultsList.svelte index e8087d967..db86a4ba4 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectoryResultsList.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormCurrentWorkingDirectory/ChatFormCurrentWorkingDirectoryResultsList.svelte @@ -34,8 +34,8 @@
{#if isSearching && results.length === 0}
Searching...
@@ -48,14 +48,15 @@
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormInput/ChatFormInputFileInputInvisible.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormInput/ChatFormInputFileInputInvisible.svelte index 395ecb201..dd9058690 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormInput/ChatFormInputFileInputInvisible.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormInput/ChatFormInputFileInputInvisible.svelte @@ -24,8 +24,8 @@ diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormInput/ChatFormInputRich.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormInput/ChatFormInputRich.svelte index 629001bb4..70251ea0c 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormInput/ChatFormInputRich.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormInput/ChatFormInputRich.svelte @@ -808,25 +808,25 @@
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormMcpResourcesList.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormMcpResourcesList.svelte index 18d86bcab..6513114a4 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormMcpResourcesList.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormMcpResourcesList.svelte @@ -27,8 +27,8 @@ {#each attachments as attachment, i (attachment.id)} handleResourceClick(attachment.resource.uri)} /> diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerItemHeader.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerItemHeader.svelte index d7c66f0d2..67e2790df 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerItemHeader.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerItemHeader.svelte @@ -21,12 +21,12 @@
{#if faviconUrl} { (e.currentTarget as HTMLImageElement).style.display = 'none'; }} + src={faviconUrl} /> {/if} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerList.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerList.svelte index 160c14ce8..2b3d6167a 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerList.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPicker/ChatFormPickerList.svelte @@ -1,4 +1,4 @@ - +
diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPickerMcpPrompts/ChatFormPromptPickerArgumentInput.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPickerMcpPrompts/ChatFormPromptPickerArgumentInput.svelte index 074c69b84..b20c13cdf 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPickerMcpPrompts/ChatFormPromptPickerArgumentInput.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormPickers/ChatFormPickerMcpPrompts/ChatFormPromptPickerArgumentInput.svelte @@ -36,7 +36,7 @@
-