diff --git a/Makefile b/Makefile index 351ea6b39..5539eabf3 100644 --- a/Makefile +++ b/Makefile @@ -667,7 +667,7 @@ ggml-vulkan-shaders-noext.o: ggml/src/ggml-vulkan-shaders-noext.cpp ggml/include $(CXX) $(CXXFLAGS) $(VKGEN_NOEXT_FORCE) $(VULKAN_FLAGS) -c $< -o $@ # intermediate objects -llama.o: src/llama.cpp ggml/include/ggml.h ggml/include/ggml-alloc.h ggml/include/ggml-backend.h ggml/include/ggml-cuda.h ggml/include/ggml-metal.h include/llama.h otherarch/llama-util.h src/llama-chat.cpp src/llama-mmap.cpp src/llama-context.cpp src/llama-adapter.cpp src/llama-arch.cpp src/llama-batch.cpp src/llama-vocab.cpp src/llama-grammar.cpp src/llama-sampler.cpp src/llama-kv-cache.cpp src/llama-kv-cache-dsa.cpp src/llama-kv-cache-dsv4.cpp src/llama-kv-cache-iswa.cpp src/llama-memory-hybrid.cpp src/llama-memory-hybrid-iswa.cpp src/llama-memory-recurrent.cpp src/llama-model-loader.cpp src/llama-model-saver.cpp src/llama-quant.cpp src/llama-hparams.cpp src/llama-graph.cpp src/llama-io.cpp src/llama-memory.cpp common/fit.cpp ggml/include/ggml.h ggml/include/ggml-cpu.h ggml/include/ggml-cuda.h include/llama.h otherarch/llama-util.h +llama.o: src/llama.cpp ggml/include/ggml.h ggml/include/ggml-alloc.h ggml/include/ggml-backend.h ggml/include/ggml-cuda.h ggml/include/ggml-metal.h include/llama.h otherarch/llama-util.h src/llama-chat.cpp src/llama-mmap.cpp src/llama-context.cpp src/llama-adapter.cpp src/llama-arch.cpp src/llama-batch.cpp src/llama-vocab.cpp src/llama-grammar.cpp src/llama-sampler.cpp src/llama-kv-cache.cpp src/llama-kv-cache-dsa.cpp src/llama-kv-cache-dsv4.cpp src/llama-kv-cache-iswa.cpp src/llama-kv-cache-msa.cpp src/llama-memory-hybrid.cpp src/llama-memory-hybrid-iswa.cpp src/llama-memory-recurrent.cpp src/llama-model-loader.cpp src/llama-model-saver.cpp src/llama-quant.cpp src/llama-hparams.cpp src/llama-graph.cpp src/llama-io.cpp src/llama-memory.cpp common/fit.cpp ggml/include/ggml.h ggml/include/ggml-cpu.h ggml/include/ggml-cuda.h include/llama.h otherarch/llama-util.h $(CXX) $(CXXFLAGS) -c $< -o $@ llama-model.o: src/llama-model.cpp src/llama-model.h src/models/models.h ggml/include/ggml.h include/llama.h $(CXX) $(CXXFLAGS) -c $< -o $@ diff --git a/common/arg.cpp b/common/arg.cpp index 0833307fa..1cb45118a 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -28,6 +28,7 @@ #include #include #include +#include #include #include #include @@ -2037,7 +2038,13 @@ common_params_context common_params_parser_init(common_params & params, llama_ex {"--repeat-penalty"}, "N", string_format("penalize repeat sequence of tokens (default: %.2f, 1.0 = disabled)", (double)params.sampling.penalty_repeat), [](common_params & params, const std::string & value) { - params.sampling.penalty_repeat = std::stof(value); + const float penalty_repeat = std::stof(value); + if (!std::isfinite(penalty_repeat) || + penalty_repeat <= 0.0f || + !std::isfinite(1.0f/penalty_repeat)) { + throw std::runtime_error("error: repeat-penalty must be finite and greater than 0\n"); + } + params.sampling.penalty_repeat = penalty_repeat; params.sampling.user_sampling_config |= common_params_sampling_config::COMMON_PARAMS_SAMPLING_CONFIG_PENALTY_REPEAT; } ).set_sampling()); @@ -2045,14 +2052,22 @@ common_params_context common_params_parser_init(common_params & params, llama_ex {"--presence-penalty"}, "N", string_format("repeat alpha presence penalty (default: %.2f, 0.0 = disabled)", (double)params.sampling.penalty_present), [](common_params & params, const std::string & value) { - params.sampling.penalty_present = std::stof(value); + const float penalty_present = std::stof(value); + if (!std::isfinite(penalty_present)) { + throw std::runtime_error("error: presence-penalty must be finite\n"); + } + params.sampling.penalty_present = penalty_present; } ).set_sampling()); add_opt(common_arg( {"--frequency-penalty"}, "N", string_format("repeat alpha frequency penalty (default: %.2f, 0.0 = disabled)", (double)params.sampling.penalty_freq), [](common_params & params, const std::string & value) { - params.sampling.penalty_freq = std::stof(value); + const float penalty_freq = std::stof(value); + if (!std::isfinite(penalty_freq)) { + throw std::runtime_error("error: frequency-penalty must be finite\n"); + } + params.sampling.penalty_freq = penalty_freq; } ).set_sampling()); add_opt(common_arg( @@ -2568,7 +2583,7 @@ 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")); - if (llama_supports_rpc()) { + if (params.is_gen_docs || llama_supports_rpc()) { add_opt(common_arg( {"--rpc"}, "SERVERS", "comma-separated list of RPC servers (host:port)", @@ -3317,7 +3332,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex {"--tools"}, "TOOL1,TOOL2,...", "experimental: whether to enable built-in tools for AI agents - do not enable in untrusted environments (default: no tools)\n" "specify \"all\" to enable all tools\n" - "available tools: read_file, file_glob_search, grep_search, exec_shell_command, write_file, edit_file, get_datetime\n" + "available tools: read_file, file_glob_search, grep_search, exec_shell_command, write_file, edit_file, get_datetime, get_info\n" "note: for security reasons, this will limit --cors-origins to localhost by default", [](common_params & params, const std::string & value) { params.server_tools = parse_csv_row(value); diff --git a/common/chat.cpp b/common/chat.cpp index ba68e0d73..aabc8e659 100644 --- a/common/chat.cpp +++ b/common/chat.cpp @@ -2128,6 +2128,11 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE; auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE); + std::optional additional_context; + if (is_v4 && has_response_format) { + additional_context = json{ { "response_format", inputs.json_schema } }; + } + const std::string DSML = "|DSML|"; const std::string THINK_START = ""; const std::string THINK_END = ""; @@ -2139,9 +2144,12 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha const std::string PARAM_START = "<" + DSML + "parameter"; const std::string PARAM_END = ""; const std::string GEN_PROMPT = "<|Assistant|>"; + const std::string TC_SEPARATOR = "\n\n"; - data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs, adjusted_messages); - data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs, adjusted_messages); + data.prompt = common_chat_template_direct_apply_impl( + tmpl, inputs, adjusted_messages, std::nullopt, additional_context); + data.generation_prompt = common_chat_template_generation_prompt_impl( + tmpl, inputs, adjusted_messages, std::nullopt, additional_context); data.format = COMMON_CHAT_FORMAT_PEG_NATIVE; data.supports_thinking = true; data.thinking_start_tag = THINK_START; @@ -2155,9 +2163,16 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha if (inputs.has_continuation()) { const auto & msg = inputs.continue_msg; - data.generation_prompt = GEN_PROMPT + THINK_START + msg.reasoning_content; - if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) { - data.generation_prompt += THINK_END + msg.render_content(); + if (is_v4 && msg.reasoning_content.empty()) { + data.generation_prompt = GEN_PROMPT + THINK_END; + if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) { + data.generation_prompt += msg.render_content(); + } + } else { + data.generation_prompt = GEN_PROMPT + THINK_START + msg.reasoning_content; + if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) { + data.generation_prompt += THINK_END + msg.render_content(); + } } data.prompt += data.generation_prompt; @@ -2256,7 +2271,9 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha if (extract_reasoning && inputs.enable_thinking) { reasoning = p.optional(THINK_START + p.reasoning(p.until(THINK_END)) + THINK_END); - reasoning_with_tc = THINK_START + p.reasoning(p.until_one_of({ FC_START, THINK_END })) + obligatory_tool_calls; + reasoning_with_tc = THINK_START + + p.reasoning(p.until_one_of({ TC_SEPARATOR + FC_START, FC_START, THINK_END })) + + p.space() + obligatory_tool_calls; allow_reasoning_with_tc = true; } else if (extract_reasoning) { // Thinking disabled but reasoning extraction requested: the generation prompt @@ -2279,7 +2296,9 @@ static common_chat_params common_chat_params_init_deepseek_v3_2(const common_cha return generation_prompt + reasoning + p.content(p.rest()) + end; } - auto content_before_tools = p.negate(p.literal(THINK_START)) + p.content(p.until(FC_START)); + auto content_before_tools = p.negate(p.literal(THINK_START)) + + p.content(p.until_one_of({ TC_SEPARATOR + FC_START, FC_START })) + + p.space(); return allow_reasoning_with_tc ? generation_prompt + (reasoning_with_tc | (reasoning + content_before_tools + tool_calls)) + end : generation_prompt + reasoning + content_before_tools + tool_calls + end; }); diff --git a/common/common.cpp b/common/common.cpp index 5a79855df..949cdca2a 100644 --- a/common/common.cpp +++ b/common/common.cpp @@ -1004,6 +1004,23 @@ bool fs_is_directory(const std::string & path) { return std::filesystem::exists(dir) && std::filesystem::is_directory(dir); } +std::string common_get_env(const std::string & name) { + const char * value = std::getenv(name.c_str()); + return value == nullptr ? "" : value; +} + +void common_set_env(const std::string & name, const std::string & value) { +#if defined(_WIN32) + _putenv_s(name.c_str(), value.c_str()); +#else + if (value.empty()) { + unsetenv(name.c_str()); + } else { + setenv(name.c_str(), value.c_str(), 1); + } +#endif +} + std::string fs_get_cache_directory() { std::string cache_directory = ""; auto ensure_trailing_slash = [](std::string p) { @@ -1305,8 +1322,9 @@ common_init_result::common_init_result(common_params & params, bool model_only) pimpl->samplers.resize(cparams.n_seq_max); pimpl->samplers_seq_config.resize(cparams.n_seq_max); + const int32_t n_ctx = cparams.n_ctx > 0 ? (int32_t) cparams.n_ctx : llama_model_n_ctx_train(model); for (int i = 0; i < (int) cparams.n_seq_max; ++i) { - pimpl->samplers[i].reset(common_sampler_init(model, params.sampling)); + pimpl->samplers[i].reset(common_sampler_init(model, params.sampling, n_ctx)); pimpl->samplers_seq_config[i] = { i, common_sampler_get(pimpl->samplers[i].get()) }; } @@ -1468,18 +1486,18 @@ common_init_result_ptr common_init_from_params(common_params & params, bool mode common_init_result::~common_init_result() = default; std::string common_get_model_endpoint() { - const char * model_endpoint_env = getenv("MODEL_ENDPOINT"); - // We still respect the use of environment-variable "HF_ENDPOINT" for backward-compatibility. - const char * hf_endpoint_env = getenv("HF_ENDPOINT"); - const char * endpoint_env = model_endpoint_env ? model_endpoint_env : hf_endpoint_env; - std::string model_endpoint = "https://huggingface.co/"; - if (endpoint_env) { - model_endpoint = endpoint_env; - if (model_endpoint.back() != '/') { - model_endpoint += '/'; - } + std::string endpoint = common_get_env("MODEL_ENDPOINT"); + if (endpoint.empty()) { + // the HF_ENDPOINT variable is respected for backward compatibility + endpoint = common_get_env("HF_ENDPOINT"); } - return model_endpoint; + if (endpoint.empty()) { + return "https://huggingface.co/"; + } + if (endpoint.back() != '/') { + endpoint += '/'; + } + return endpoint; } char * common_get_model_or_exit(int argc, char * argv[]) { diff --git a/common/common.h b/common/common.h index e6d5d892e..62bc8029d 100644 --- a/common/common.h +++ b/common/common.h @@ -740,6 +740,8 @@ struct common_params { llama_progress_callback load_progress_callback = NULL; void * load_progress_callback_user_data = NULL; bool no_alloc = false; // Don't allocate model buffers + + bool is_gen_docs = false; // whether we are running inside llama-gen-docs }; // call once at the start of a program if it uses libcommon @@ -864,6 +866,15 @@ std::string string_from(const struct llama_context * ctx, const struct llama_bat bool glob_match(const std::string & pattern, const std::string & str); +// +// Environment utils +// + +// portable environment access, an unset variable reads as an empty string +// and setting an empty value unsets the variable +std::string common_get_env(const std::string & name); +void common_set_env(const std::string & name, const std::string & value); + // // Filesystem utils // diff --git a/common/jinja/caps.cpp b/common/jinja/caps.cpp index 26306bd91..6f4364072 100644 --- a/common/jinja/caps.cpp +++ b/common/jinja/caps.cpp @@ -485,6 +485,7 @@ caps caps_get(jinja::program & prog) { }); }, [&](context & ctx) { + ctx.set_val("enable_thinking", mk_val(true)); caps_apply_preserve_reasoning(ctx, true); }, nullptr, // tools_fn diff --git a/common/sampling.cpp b/common/sampling.cpp index 256ac161e..ba5504ed0 100644 --- a/common/sampling.cpp +++ b/common/sampling.cpp @@ -184,9 +184,26 @@ std::string common_params_sampling::print() const { return std::string(result); } -struct common_sampler * common_sampler_init(const struct llama_model * model, struct common_params_sampling & params) { - const llama_vocab * vocab = llama_model_get_vocab(model); +struct common_sampler * common_sampler_init( + const struct llama_model * model, + struct common_params_sampling & params, + int32_t n_ctx) { + if (!std::isfinite(params.penalty_repeat) || + params.penalty_repeat <= 0.0f || + !std::isfinite(1.0f/params.penalty_repeat)) { + throw std::invalid_argument("penalty_repeat must be finite and greater than 0"); + } + if (!std::isfinite(params.penalty_freq)) { + throw std::invalid_argument("penalty_freq must be finite"); + } + if (!std::isfinite(params.penalty_present)) { + throw std::invalid_argument("penalty_present must be finite"); + } + if (params.penalty_last_n == -1) { + params.penalty_last_n = n_ctx > 0 ? n_ctx : llama_model_n_ctx_train(model); + } + const llama_vocab * vocab = llama_model_get_vocab(model); llama_sampler_chain_params lparams = llama_sampler_chain_default_params(); lparams.no_perf = params.no_perf; @@ -366,7 +383,7 @@ struct common_sampler * common_sampler_init(const struct llama_model * model, st samplers.push_back(llama_sampler_init_infill(vocab)); break; case COMMON_SAMPLER_TYPE_PENALTIES: - samplers.push_back(llama_sampler_init_penalties(params.penalty_last_n, params.penalty_repeat, params.penalty_freq, params.penalty_present)); + samplers.push_back(llama_sampler_init_penalties(llama_vocab_n_tokens(vocab), params.penalty_last_n, params.penalty_repeat, params.penalty_freq, params.penalty_present)); break; case COMMON_SAMPLER_TYPE_ADAPTIVE_P: // the `adaptive-p` sampler is like `dist` and `mirostat` in that it selects diff --git a/common/sampling.h b/common/sampling.h index 4191988bb..91e2cea78 100644 --- a/common/sampling.h +++ b/common/sampling.h @@ -37,7 +37,10 @@ struct common_sampler; // llama_sampler API overloads // note: can mutate params in some cases -struct common_sampler * common_sampler_init(const struct llama_model * model, struct common_params_sampling & params); +struct common_sampler * common_sampler_init( + const struct llama_model * model, + struct common_params_sampling & params, + int32_t n_ctx = 0); void common_sampler_free(struct common_sampler * gsmpl); diff --git a/common/speculative.cpp b/common/speculative.cpp index e9bf31d68..4ce1ea325 100644 --- a/common/speculative.cpp +++ b/common/speculative.cpp @@ -2385,57 +2385,28 @@ common_speculative * common_speculative_init(common_params_speculative & params, { uint32_t enabled_configs = common_get_enabled_speculative_configs(params.types); - bool has_draft_simple = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_DRAFT_SIMPLE)); - bool has_draft_eagle3 = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3)) && params.draft.ctx_dft != nullptr; - bool has_draft_mtp = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_DRAFT_MTP)) && params.draft.ctx_dft != nullptr; - bool has_draft_dflash = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH)) && params.draft.ctx_dft != nullptr; - bool has_draft_dspark = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK)) && params.draft.ctx_dft != nullptr; - - - - bool has_ngram_cache = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_NGRAM_CACHE)); - bool has_ngram_simple = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE)); - bool has_ngram_map_k = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K)); - bool has_ngram_map_k4v = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V)); - bool has_ngram_mod = (enabled_configs & (1u << COMMON_SPECULATIVE_TYPE_NGRAM_MOD)); + auto add_config_if_enabled = [&](common_speculative_type type, bool available = true) { + if (available && (enabled_configs & (1u << type))) { + configs.emplace_back(type, params); + } + }; // when adding a new type - update here the logic above static_assert(COMMON_SPECULATIVE_TYPE_COUNT == 11); // this list here defines the priority of the speculators // the one with highest priority are listed first - if (has_ngram_simple) { - // This implementation can guess a lot of tokens without any draft model. - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE, params)); - } - if (has_ngram_map_k) { - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K, params)); - } - if (has_ngram_map_k4v) { - // This implementation can guess tokens with high acceptance rate but is more expensive. - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V, params)); - } - if (has_ngram_mod) { - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_MOD, params)); - } - if (has_ngram_cache) { - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_NGRAM_CACHE, params)); - } - if (has_draft_simple) { - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_DRAFT_SIMPLE, params)); - } - if (has_draft_eagle3) { - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3, params)); - } - if (has_draft_mtp) { - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_DRAFT_MTP, params)); - } - if (has_draft_dflash) { - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH, params)); - } - if (has_draft_dspark) { - configs.push_back(common_speculative_config(COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK, params)); - } + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE); + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K); + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K4V); + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_NGRAM_MOD); + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_NGRAM_CACHE); + + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_DRAFT_SIMPLE); + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3, params.draft.ctx_dft != nullptr); + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_DRAFT_MTP, params.draft.ctx_dft != nullptr); + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH, params.draft.ctx_dft != nullptr); + add_config_if_enabled(COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK, params.draft.ctx_dft != nullptr); } std::vector> impls = {}; diff --git a/conversion/chatglm.py b/conversion/chatglm.py index 801913075..d63855038 100644 --- a/conversion/chatglm.py +++ b/conversion/chatglm.py @@ -81,7 +81,7 @@ class ChatGLMModel(TextModel): @staticmethod def token_bytes_to_string(b): - from transformers.models.gpt2.tokenization_gpt2 import bytes_to_unicode # ty: ignore[unresolved-import] + from transformers.convert_slow_tokenizer import bytes_to_unicode byte_encoder = bytes_to_unicode() return ''.join([byte_encoder[ord(char)] for char in b.decode('latin-1')]) diff --git a/conversion/deepseek.py b/conversion/deepseek.py index 1c9b325d5..0518fcecc 100644 --- a/conversion/deepseek.py +++ b/conversion/deepseek.py @@ -535,7 +535,10 @@ class DeepseekV4Model(TextModel): logger.info("Skipping %d DeepSeek-V4 MTP tensor(s) for conversion v0", type(self)._skipped_mtp_tensors) # add a default chat template; if the model has a built-in template, it will be overridden later - template_path = Path(__file__).parent.parent / "models" / "templates" / "deepseek-ai-DeepSeek-V4.jinja" + model_id_hint = self.remote_hf_model_id or self.dir_model.name + is_0731 = "0731" in model_id_hint + template_name = "deepseek-ai-DeepSeek-V4-Flash-0731.jinja" if is_0731 else "deepseek-ai-DeepSeek-V4.jinja" + template_path = Path(__file__).parent.parent / "models" / "templates" / template_name if template_path.is_file(): with open(template_path, "r", encoding="utf-8") as f: self.gguf_writer.add_chat_template(f.read()) diff --git a/conversion/glm.py b/conversion/glm.py index cc34cddbf..e28f54574 100644 --- a/conversion/glm.py +++ b/conversion/glm.py @@ -206,10 +206,70 @@ class Glm4MoeModel(TextModel): @ModelBase.register("Glm4MoeLiteForCausalLM") class Glm4MoeLiteModel(DeepseekV2Model): model_arch = gguf.MODEL_ARCH.DEEPSEEK2 + skip_mtp = False + supports_mtp_export = True + _n_main_layers: int | None = None def set_vocab(self): return self._set_vocab_glm() + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + num_hidden_layers = self.hparams["num_hidden_layers"] + self.num_nextn_predict_layers = self.hparams.get("num_nextn_predict_layers", 0) + self.skip_mtp = self.no_mtp or self.num_nextn_predict_layers == 0 + + if self.skip_mtp: + self.block_count = num_hidden_layers + else: + self.block_count = num_hidden_layers + self.num_nextn_predict_layers + + self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count) + + def set_gguf_parameters(self): + super().set_gguf_parameters() + + if self.skip_mtp: + return + + self.gguf_writer.add_nextn_predict_layers(self.num_nextn_predict_layers) + + def index_tensors(self, remote_hf_model_id: str | None = None): + type(self)._n_main_layers = self.hparams["num_hidden_layers"] + return super().index_tensors(remote_hf_model_id=remote_hf_model_id) + + @classmethod + def filter_tensors(cls, item): + if (titem := super().filter_tensors(item)) is None: + return None + name, gen = titem + + if cls._n_main_layers is not None: + match = re.match(r"model\.layers\.(\d+)\.", name) + is_mtp = match is not None and int(match.group(1)) >= cls._n_main_layers + if is_mtp and cls.no_mtp: + return None + if cls.mtp_only and not is_mtp and name not in ( + "model.embed_tokens.weight", "model.norm.weight", "lm_head.weight", + ): + return None + + return name, gen + + def prepare_metadata(self, vocab_only: bool): + from_dir = self.fname_out.is_dir() + super().prepare_metadata(vocab_only=vocab_only) + + if not self.mtp_only or not from_dir: + return + + output_type: str = self.ftype.name.partition("_")[2] + fname_default: str = gguf.naming_convention( + self.metadata.name, self.metadata.basename, self.metadata.finetune, + self.metadata.version, size_label=None, output_type=output_type, model_type=None) + self.fname_out = self.fname_out.parent / f"mtp-{fname_default}.gguf" + @ModelBase.register("GlmMoeDsaForCausalLM") class GlmMoeDsaModel(DeepseekV2Model): diff --git a/conversion/llama.py b/conversion/llama.py index 9b3373f91..1aced49c5 100644 --- a/conversion/llama.py +++ b/conversion/llama.py @@ -119,7 +119,7 @@ class LlamaModel(TextModel): path_tekken_json = self.dir_model / "tekken.json" path_tokenizer_json = self.dir_model / "tokenizer.json" if path_tekken_json.is_file() and not path_tokenizer_json.is_file(): - self._set_vocab_mistral() + return self._set_vocab_mistral() tokenizer_config_file = self.dir_model / 'tokenizer_config.json' if tokenizer_config_file.is_file(): diff --git a/conversion/qwen.py b/conversion/qwen.py index 7e3d8c0d1..b4ae528bf 100644 --- a/conversion/qwen.py +++ b/conversion/qwen.py @@ -18,7 +18,7 @@ class QwenModel(TextModel): @staticmethod def token_bytes_to_string(b): - from transformers.models.gpt2.tokenization_gpt2 import bytes_to_unicode # ty: ignore[unresolved-import] + from transformers.convert_slow_tokenizer import bytes_to_unicode byte_encoder = bytes_to_unicode() return ''.join([byte_encoder[ord(char)] for char in b.decode('latin-1')]) diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index 9a5a28e6b..8e65c448e 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -766,8 +766,9 @@ struct ggml_backend_sched_split { int backend_id; int i_start; int i_end; - struct ggml_tensor * inputs[GGML_SCHED_MAX_SPLIT_INPUTS]; + struct ggml_tensor ** inputs; int n_inputs; + int inputs_capacity; // graph view of this split struct ggml_cgraph graph; }; @@ -806,8 +807,9 @@ struct ggml_backend_sched { int cur_copy; int next_copy; ggml_backend_event_t events[GGML_SCHED_MAX_BACKENDS][GGML_SCHED_MAX_COPIES]; - struct ggml_tensor * graph_inputs[GGML_SCHED_MAX_SPLIT_INPUTS]; + struct ggml_tensor ** graph_inputs; int n_graph_inputs; + int graph_inputs_capacity; struct ggml_context * ctx; @@ -833,6 +835,36 @@ struct ggml_backend_sched { #define tensor_id_copy(id, backend_id, copy_id) sched->hv_tensor_copies[(id) * sched->n_backends * sched->n_copies + (backend_id) * sched->n_copies + (copy_id)] #define tensor_copy(tensor, backend_id, copy_id) tensor_id_copy(hash_id(tensor), backend_id, copy_id) +static void ggml_backend_sched_split_inputs_grow(struct ggml_backend_sched_split * split) { + int new_cap = GGML_SCHED_MAX_SPLIT_INPUTS; + if (split->inputs_capacity > 0) { + new_cap = 2*split->inputs_capacity; + GGML_LOG_WARN("%s: increasing split inputs capacity from %d to %d\n", __func__, split->inputs_capacity, new_cap); + } + auto * pnew = (struct ggml_tensor **) realloc((void *) split->inputs, new_cap * sizeof(struct ggml_tensor *)); + if (pnew == NULL) { + GGML_LOG_ERROR("%s: failed to allocate %zu bytes\n", __func__, new_cap * sizeof(struct ggml_tensor *)); + GGML_ABORT("failed to grow split inputs container"); + } + split->inputs = pnew; + split->inputs_capacity = new_cap; +} + +static void ggml_backend_sched_graph_inputs_grow(ggml_backend_sched_t sched) { + int new_cap = GGML_SCHED_MAX_SPLIT_INPUTS; + if (sched->graph_inputs_capacity > 0) { + new_cap = 2*sched->graph_inputs_capacity; + GGML_LOG_WARN("%s: increasing graph inputs capacity from %d to %d\n", __func__, sched->graph_inputs_capacity, new_cap); + } + auto * pnew = (struct ggml_tensor **) realloc((void *) sched->graph_inputs, new_cap * sizeof(struct ggml_tensor *)); + if (pnew == NULL) { + GGML_LOG_ERROR("%s: failed to allocate %zu bytes\n", __func__, new_cap * sizeof(struct ggml_tensor *)); + GGML_ABORT("failed to grow graph inputs container"); + } + sched->graph_inputs = pnew; + sched->graph_inputs_capacity = new_cap; +} + // returns the priority of the backend, lower id is higher priority static int ggml_backend_sched_backend_id(ggml_backend_sched_t sched, ggml_backend_t backend) { for (int i = 0; i < sched->n_backends; i++) { @@ -1304,7 +1336,7 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra } // check if the split has too many inputs // FIXME: count the number of inputs instead of only checking when full - if (split->n_inputs == GGML_SCHED_MAX_SPLIT_INPUTS) { + if (split->n_inputs >= split->inputs_capacity) { const size_t id = hash_id(src); int src_backend_id = sched->hv_tensor_backend_ids[id]; bool supported = ggml_backend_sched_buffer_supported(sched, src, cur_backend_id); @@ -1320,10 +1352,14 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra split->i_end = i; i_split++; if (i_split >= sched->splits_capacity) { + int old_cap = sched->splits_capacity; sched->splits_capacity *= 2; sched->splits = (ggml_backend_sched_split *) realloc(sched->splits, sched->splits_capacity * sizeof(struct ggml_backend_sched_split)); GGML_ASSERT(sched->splits != NULL); + for (int k = old_cap; k < sched->splits_capacity; k++) { + memset(&sched->splits[k], 0, sizeof(struct ggml_backend_sched_split)); + } } split = &sched->splits[i_split]; split->backend_id = node_backend_id; @@ -1360,7 +1396,9 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra SET_CAUSE(tensor_copy, "4.cpy"); } int n_graph_inputs = sched->n_graph_inputs++; - GGML_ASSERT(n_graph_inputs < GGML_SCHED_MAX_SPLIT_INPUTS); + if (n_graph_inputs >= sched->graph_inputs_capacity) { + ggml_backend_sched_graph_inputs_grow(sched); + } sched->graph_inputs[n_graph_inputs] = src; } } @@ -1380,7 +1418,9 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra SET_CAUSE(tensor_copy, "4.cpy"); } int n_inputs = split->n_inputs++; - GGML_ASSERT(n_inputs < GGML_SCHED_MAX_SPLIT_INPUTS); + if (n_inputs >= split->inputs_capacity) { + ggml_backend_sched_split_inputs_grow(split); + } split->inputs[n_inputs] = src; } node->src[j] = tensor_id_copy(src_id, cur_backend_id, sched->cur_copy); @@ -1406,7 +1446,11 @@ void ggml_backend_sched_split_graph(ggml_backend_sched_t sched, struct ggml_cgra sched->prev_leaf_backend_ids = tmp; } - int graph_size = std::max(graph->n_nodes, graph->n_leafs) + sched->n_splits*GGML_SCHED_MAX_SPLIT_INPUTS*2*sched->n_copies; + int total_inputs = sched->n_graph_inputs; + for (int i = 0; i < sched->n_splits; i++) { + total_inputs += sched->splits[i].n_inputs; + } + int graph_size = std::max(graph->n_nodes, graph->n_leafs) + total_inputs * 2 * sched->n_copies; // remember the actual graph_size for performing reallocation checks later [GGML_SCHED_DEBUG_REALLOC] sched->debug_prev_graph_size = sched->debug_graph_size; @@ -1793,6 +1837,9 @@ ggml_backend_sched_t ggml_backend_sched_new( sched->splits = (ggml_backend_sched_split *) calloc(initial_splits_capacity, sizeof(sched->splits[0])); sched->splits_capacity = initial_splits_capacity; + sched->graph_inputs_capacity = GGML_SCHED_MAX_SPLIT_INPUTS; + sched->graph_inputs = (struct ggml_tensor **) calloc(sched->graph_inputs_capacity, sizeof(struct ggml_tensor *)); + for (int b = 0; b < n_backends; b++) { sched->backends[b] = backends[b]; sched->bufts[b] = bufts ? bufts[b] : ggml_backend_get_default_buffer_type(backends[b]); @@ -1825,7 +1872,11 @@ void ggml_backend_sched_free(ggml_backend_sched_t sched) { ggml_gallocr_free(sched->galloc); ggml_free(sched->ctx); ggml_hash_set_free(&sched->hash_set); + for (int i = 0; i < sched->splits_capacity; i++) { + free(sched->splits[i].inputs); + } free(sched->splits); + free(sched->graph_inputs); free(sched->hv_tensor_backend_ids); free(sched->hv_tensor_copies); free(sched->node_backend_ids); diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index 14fbc4a88..f1ce30e6c 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -634,7 +634,8 @@ template struct block_reduce_policy { }; template -static __device__ T block_reduce(T val, T * shared_vals) { +static __device__ T block_reduce(T val, [[maybe_unused]] T * shared_vals) { + // for multi-warp reductions, callers must not reuse shared_vals until all reads from this invocation have completed val = block_reduce_policy::reduce(val); const unsigned int block_size = block_size_template == 0 ? blockDim.x : block_size_template; if (block_size > WARP_SIZE) { diff --git a/ggml/src/ggml-cuda/norm.cu b/ggml/src/ggml-cuda/norm.cu index 09d9f3a7d..c3758cd50 100644 --- a/ggml/src/ggml-cuda/norm.cu +++ b/ggml/src/ggml-cuda/norm.cu @@ -64,7 +64,7 @@ static __global__ void group_norm_f32(const float * x, float * dst, const int gr tmp += xi * xi; } - tmp = block_reduce(tmp, s_sum); + tmp = block_reduce(tmp, s_sum + 32); const float variance = tmp / group_size; const float scale = rsqrtf(variance + eps); @@ -297,7 +297,7 @@ static void group_norm_f32_cuda( group_norm_f32<<>>(x, dst, group_size, ne_elements, eps); } else { const dim3 block_dims(1024, 1, 1); - group_norm_f32<1024><< WARP_SIZE ? 32 * sizeof(float): 0, stream>>>(x, dst, group_size, ne_elements, eps); + group_norm_f32<1024><< WARP_SIZE ? 2 * 32 * sizeof(float): 0, stream>>>(x, dst, group_size, ne_elements, eps); } } diff --git a/ggml/src/ggml-cuda/softmax.cu b/ggml/src/ggml-cuda/softmax.cu index 285c0e954..f320c6f00 100644 --- a/ggml/src/ggml-cuda/softmax.cu +++ b/ggml/src/ggml-cuda/softmax.cu @@ -116,6 +116,11 @@ static __global__ void soft_max_f32( vals[col] = val; } + if (block_size > WARP_SIZE) { + // sync is needed as we reuse buf_iw across block_reduce invocations, see #26385 + // for block_size <= WARP_SIZE, block_reduce does not access buf_iw + __syncthreads(); + } // find the sum of exps in the block tmp = block_reduce(tmp, buf_iw); @@ -142,6 +147,8 @@ static __device__ void soft_max_f32_parallelize_cols_single_row(const float * __ float * __restrict__ dst, float * __restrict__ tmp_maxs, float * __restrict__ tmp_sums, + float * shared_vals_max, + float * shared_vals_sum, const soft_max_params p) { namespace cg = cooperative_groups; @@ -154,7 +161,6 @@ static __device__ void soft_max_f32_parallelize_cols_single_row(const float * __ float local_vals[n_elem_per_thread] = { -INFINITY, -INFINITY, -INFINITY, -INFINITY }; float local_max = -INFINITY; const int step_size = gridDim.x * blockDim.x; - __shared__ float shared_vals[32]; // Compute thread-local max for (int col = col_start; col < p.ncols;) { @@ -171,7 +177,7 @@ static __device__ void soft_max_f32_parallelize_cols_single_row(const float * __ } // Compute CTA-level max - local_max = block_reduce(local_max, shared_vals); + local_max = block_reduce(local_max, shared_vals_max); // Store CTA-level max to GMEM if (tid == 0) { @@ -186,7 +192,7 @@ static __device__ void soft_max_f32_parallelize_cols_single_row(const float * __ } else { local_max = -INFINITY; } - local_max = block_reduce(local_max, shared_vals); + local_max = block_reduce(local_max, shared_vals_max); // Compute softmax dividends, accumulate divisor float tmp_expf = 0.0f; @@ -209,7 +215,7 @@ static __device__ void soft_max_f32_parallelize_cols_single_row(const float * __ } // Reduce divisor within CTA - tmp_expf = block_reduce(tmp_expf, shared_vals); + tmp_expf = block_reduce(tmp_expf, shared_vals_sum); // Store CTA-level sum to GMEM if (tid == 0) { @@ -223,7 +229,7 @@ static __device__ void soft_max_f32_parallelize_cols_single_row(const float * __ } else { tmp_expf = 0.0f; } - tmp_expf = block_reduce(tmp_expf, shared_vals); + tmp_expf = block_reduce(tmp_expf, shared_vals_sum); // Divide dividend by global sum + store data for (int col = col_start; col < p.ncols;) { @@ -310,9 +316,11 @@ __launch_bounds__(8*WARP_SIZE, 1) static __global__ void soft_max_f32_paralleliz // https://docs.nvidia.com/cuda/cuda-programming-guide/05-appendices/device-callable-apis.html#grid-synchronization // https://docs.nvidia.com/cuda/cuda-programming-guide/05-appendices/device-callable-apis.html#class-cluster-group { + __shared__ float shared_vals[2][32]; + for (int rowx = 0; rowx < p.ne01 * p.ne02 * p.ne03; rowx++) { soft_max_f32_parallelize_cols_single_row(x + int64_t(rowx) * p.ncols, dst + int64_t(rowx) * p.ncols, tmp_maxs, - tmp_sums, p); + tmp_sums, shared_vals[0], shared_vals[1], p); } } diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 11c86614d..479b9d473 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -1032,6 +1032,7 @@ struct vk_device_struct { vk_pipeline pipeline_pool2d_f32; vk_pipeline pipeline_rwkv_wkv6_f32; vk_pipeline pipeline_rwkv_wkv7_f32; + vk_pipeline pipeline_gated_linear_attn_f32; // [size_idx][kda] where size_idx: 0=d16, 1=d32, 2=d64, 3=d128 vk_pipeline pipeline_gated_delta_net[4][2]; vk_pipeline pipeline_ssm_scan_f32_d128; @@ -1753,6 +1754,13 @@ struct vk_op_rwkv_wkv7_push_constants { uint32_t C; uint32_t H; }; +struct vk_op_gated_linear_attn_push_constants { + uint32_t B; + uint32_t T; + uint32_t C; + uint32_t H; + float scale; +}; struct vk_op_gated_delta_net_push_constants { uint32_t H; uint32_t n_tokens; @@ -5671,6 +5679,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_rwkv_wkv7_f32, "rwkv_wkv7_f32", rwkv_wkv7_f32_len, rwkv_wkv7_f32_data, "main", 8, sizeof(vk_op_rwkv_wkv7_push_constants), {1, 1, 1}, {device->subgroup_size}, 1); + ggml_vk_create_pipeline(device, device->pipeline_gated_linear_attn_f32, "gated_linear_attn_f32", gated_linear_attn_f32_len, gated_linear_attn_f32_data, "main", 6, sizeof(vk_op_gated_linear_attn_push_constants), {1, 1, 1}, {}, 1); + { const uint32_t gdn_sizes[] = {16, 32, 64, 128}; const char * gdn_names[][2] = { @@ -11425,6 +11435,11 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const return ctx->device->pipeline_rwkv_wkv7_f32; } return nullptr; + case GGML_OP_GATED_LINEAR_ATTN: + if (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { + return ctx->device->pipeline_gated_linear_attn_f32; + } + return nullptr; case GGML_OP_GATED_DELTA_NET: if (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { const uint32_t S_v = dst->src[2]->ne[0]; @@ -12455,6 +12470,41 @@ static void ggml_vk_rwkv_wkv7(ggml_backend_vk_context * ctx, vk_context& subctx, ); } +static void ggml_vk_gated_linear_attn(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { + const size_t seq_length = dst->src[0]->ne[2]; + const size_t n_embed = dst->ne[0]; + const size_t n_heads = dst->src[0]->ne[1]; + const size_t n_seqs = dst->src[4]->ne[1]; + + float scale; + memcpy(&scale, dst->op_params, sizeof(float)); + + GGML_ASSERT(dst->buffer != nullptr); + + vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, dst->src[0], dst->src[1], dst->src[2], dst, dst->op); + GGML_ASSERT(pipeline != nullptr); + + ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1); + + vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst); + vk_subbuffer src_buf[5] = {}; + for (int i = 0; i < 5; i++) { + src_buf[i] = ggml_vk_tensor_subbuffer(ctx, dst->src[i]); + } + + const vk_op_gated_linear_attn_push_constants pc = { + (uint32_t)n_seqs, + (uint32_t)seq_length, + (uint32_t)n_embed, + (uint32_t)n_heads, + scale, + }; + + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {src_buf[0], src_buf[1], src_buf[2], src_buf[3], src_buf[4], dst_buf}, + pc, { (uint32_t)(n_seqs * n_heads), 1, 1 }); +} + static void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { const ggml_tensor * src_q = dst->src[0]; const ggml_tensor * src_v = dst->src[2]; @@ -15454,6 +15504,11 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr break; + case GGML_OP_GATED_LINEAR_ATTN: + ggml_vk_gated_linear_attn(ctx, compute_ctx, node); + + break; + case GGML_OP_GATED_DELTA_NET: ggml_vk_gated_delta_net(ctx, compute_ctx, node); @@ -18161,6 +18216,9 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm case GGML_OP_RWKV_WKV6: case GGML_OP_RWKV_WKV7: return true; // all inputs are contiguous, see ggml.c + case GGML_OP_GATED_LINEAR_ATTN: + // the shader block size is hardcoded to head_size 64 + return op->src[0]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && op->src[0]->ne[0] == 64; case GGML_OP_GATED_DELTA_NET: { const uint32_t S_v = op->src[2]->ne[0]; @@ -19150,6 +19208,10 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph * } else if (tensor->op == GGML_OP_RWKV_WKV7) { tensor_clone = ggml_rwkv_wkv7(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], src_clone[3], src_clone[4], src_clone[5], src_clone[6]); + } else if (tensor->op == GGML_OP_GATED_LINEAR_ATTN) { + const float * op_params = (const float *)tensor->op_params; + tensor_clone = ggml_gated_linear_attn(ggml_ctx, src_clone[0], src_clone[1], + src_clone[2], src_clone[3], src_clone[4], op_params[0]); } else if (tensor->op == GGML_OP_GATED_DELTA_NET) { tensor_clone = ggml_gated_delta_net(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], src_clone[3], src_clone[4], src_clone[5], diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/gla.comp b/ggml/src/ggml-vulkan/vulkan-shaders/gla.comp new file mode 100644 index 000000000..b3387616b --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/gla.comp @@ -0,0 +1,82 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : require + +#define BLOCK_SIZE 64 +layout(local_size_x = BLOCK_SIZE, local_size_y = 1, local_size_z = 1) in; + +layout(push_constant) uniform Parameters { + uint B; + uint T; + uint C; + uint H; + float scale; +}; + +layout(binding = 0) readonly buffer KBuf { A_TYPE k[]; }; +layout(binding = 1) readonly buffer VBuf { A_TYPE v[]; }; +layout(binding = 2) readonly buffer QBuf { A_TYPE q[]; }; +layout(binding = 3) readonly buffer GBuf { A_TYPE g[]; }; +layout(binding = 4) readonly buffer StateBuf { A_TYPE state_in[]; }; +layout(binding = 5) buffer DstBuf { A_TYPE dst[]; }; + +shared A_TYPE _k[BLOCK_SIZE], _q[BLOCK_SIZE], _g[BLOCK_SIZE]; + +void main() { + const uint head_size = BLOCK_SIZE; + const uint batch_id = gl_WorkGroupID.x / H; + const uint head_id = gl_WorkGroupID.x % H; + const uint tid = gl_LocalInvocationID.x; + + const uint state_size = C * head_size; + const uint n_seq_tokens = T / B; + + if (batch_id >= B || head_id >= H) { + return; + } + + // state[i] holds column tid of this head's state matrix: S[i][tid] + A_TYPE state[BLOCK_SIZE]; + [[unroll]] for (uint i = 0; i < head_size; i++) { + state[i] = state_in[batch_id * state_size + head_id * head_size * head_size + + i * head_size + tid]; + } + + const uint start_t = batch_id * n_seq_tokens * C + head_id * head_size + tid; + const uint end_t = (batch_id + 1) * n_seq_tokens * C + head_id * head_size + tid; + + for (uint t = start_t; t < end_t; t += C) { + barrier(); + _k[tid] = k[t]; + _q[tid] = q[t]; + _g[tid] = g[t]; + barrier(); + + const A_TYPE v_val = v[t]; + A_TYPE y = 0.0; + + [[unroll]] for (uint i = 0; i < head_size; i += 4) { + vec4 k_vec = vec4(_k[i], _k[i+1], _k[i+2], _k[i+3]); + vec4 q_vec = vec4(_q[i], _q[i+1], _q[i+2], _q[i+3]); + vec4 g_vec = vec4(_g[i], _g[i+1], _g[i+2], _g[i+3]); + vec4 s_vec = vec4(state[i], state[i+1], state[i+2], state[i+3]); + + vec4 kv = k_vec * v_val; + + s_vec = s_vec * g_vec + kv; + y += dot(q_vec, s_vec); + + state[i] = s_vec.x; + state[i+1] = s_vec.y; + state[i+2] = s_vec.z; + state[i+3] = s_vec.w; + } + + dst[t] = y * scale; + } + + [[unroll]] for (uint i = 0; i < head_size; i++) { + dst[T * C + batch_id * state_size + head_id * head_size * head_size + + i * head_size + tid] = state[i]; + } +} 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 87fc3cf19..099fb15cc 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -1083,6 +1083,8 @@ void process_shaders() { string_to_spv("rwkv_wkv6_f32", "wkv6.comp", merge_maps(base_dict, {{"A_TYPE", "float"}})); + string_to_spv("gated_linear_attn_f32", "gla.comp", merge_maps(base_dict, {{"A_TYPE", "float"}})); + string_to_spv("rwkv_wkv7_f32", "wkv7.comp", merge_maps(base_dict, {{"A_TYPE", "float"}})); string_to_spv("gated_delta_net_f32", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}, {"USE_SUBGROUP_CLUSTERED", "1"}})); diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index 7df984432..6b0a26b63 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -11,6 +11,7 @@ GGUF_MAGIC = 0x46554747 # "GGUF" GGUF_VERSION = 3 GGUF_DEFAULT_ALIGNMENT = 32 GGML_QUANT_VERSION = 2 # GGML_QNT_VERSION from ggml.h +GGML_MAX_DIMS = 4 # GGML_MAX_DIMS from ggml.h # # metadata keys @@ -3221,6 +3222,13 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = { MODEL_TENSOR.FFN_DOWN_SHEXP, MODEL_TENSOR.FFN_UP_SHEXP, MODEL_TENSOR.FFN_EXP_PROBS_B, + # NextN/MTP tensors + MODEL_TENSOR.NEXTN_EH_PROJ, + MODEL_TENSOR.NEXTN_EMBED_TOKENS, + MODEL_TENSOR.NEXTN_ENORM, + MODEL_TENSOR.NEXTN_HNORM, + MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD, + MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM, ], MODEL_ARCH.DEEPSEEK2OCR: [ MODEL_TENSOR.TOKEN_EMBD, diff --git a/gguf-py/gguf/gguf_reader.py b/gguf-py/gguf/gguf_reader.py index 0a1b85f50..ea241ada2 100644 --- a/gguf-py/gguf/gguf_reader.py +++ b/gguf-py/gguf/gguf_reader.py @@ -22,6 +22,7 @@ if __name__ == "__main__": sys.path.insert(0, str(Path(__file__).parent.parent)) from gguf.constants import ( + GGML_MAX_DIMS, GGML_QUANT_SIZES, GGUF_DEFAULT_ALIGNMENT, GGUF_MAGIC, @@ -266,6 +267,8 @@ class GGUFReader: # Get Tensor Dimensions Count n_dims = self._get(offs, np.uint32) offs += int(n_dims.nbytes) + if n_dims[0] > GGML_MAX_DIMS: + raise ValueError(f'Tensor dimensions count {n_dims[0]} exceeds GGML_MAX_DIMS ({GGML_MAX_DIMS})') # Get Tensor Dimension Array dims = self._get(offs, np.uint64, n_dims[0]) @@ -326,7 +329,10 @@ class GGUFReader: raise ValueError(f'Found duplicated tensor with name {tensor_name}') tensor_names.add(tensor_name) ggml_type = GGMLQuantizationType(raw_dtype[0]) - n_elems = int(np.prod(dims)) + # use Python ints: np.prod on uint64 wraps silently on overflow + n_elems = 1 + for dim in dims.tolist(): + n_elems *= int(dim) np_dims = tuple(reversed(dims.tolist())) block_size, type_size = GGML_QUANT_SIZES[ggml_type] n_bytes = n_elems * type_size // block_size diff --git a/include/llama.h b/include/llama.h index 1534bfc1c..22455b369 100644 --- a/include/llama.h +++ b/include/llama.h @@ -1427,10 +1427,11 @@ extern "C" { /// NOTE: Avoid using on the full vocabulary as searching for repeated tokens can become slow. For example, apply top-k or top-p sampling first. LLAMA_API struct llama_sampler * llama_sampler_init_penalties( + int32_t n_vocab, int32_t penalty_last_n, // last n tokens to penalize (0 = disable penalty, -1 = context size) - float penalty_repeat, // 1.0 = disabled - float penalty_freq, // 0.0 = disabled - float penalty_present); // 0.0 = disabled + float penalty_repeat, // must be > 0.0, 1.0 = disabled + float penalty_freq, // must be finite, 0.0 = disabled + float penalty_present); // must be finite, 0.0 = disabled /// @details DRY sampler, designed by p-e-w, as described in: https://github.com/oobabooga/text-generation-webui/pull/5677, porting Koboldcpp implementation authored by pi6am: https://github.com/LostRuins/koboldcpp/pull/982 LLAMA_API struct llama_sampler * llama_sampler_init_dry( diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index 1f12b1251..7382aa2ad 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -8,6 +8,7 @@ #include "llama-kv-cache.h" #include "llama-kv-cache-iswa.h" #include "llama-kv-cache-dsa.h" +#include "llama-kv-cache-msa.h" #include "llama-kv-cache-dsv4.h" #include "llama-memory-hybrid.h" #include "llama-memory-hybrid-iswa.h" @@ -518,6 +519,40 @@ bool llm_graph_input_attn_k::can_reuse(const llm_graph_params & params) { return res; } +llm_graph_input_attn_kv_msa::llm_graph_input_attn_kv_msa( + const llama_hparams & hparams, + const llama_cparams & cparams, + const llama_kv_cache_msa_context * mctx) : + llm_graph_input_attn_kv(hparams, cparams, mctx->get_base()), + mctx_msa(mctx) { +} + +void llm_graph_input_attn_kv_msa::set_input(const llama_ubatch * ubatch) { + llm_graph_input_attn_kv::set_input(ubatch); + + if (self_k_idxs_idx) { + mctx_msa->get_idx()->set_input_k_idxs(self_k_idxs_idx, ubatch); + } +} + +bool llm_graph_input_attn_kv_msa::can_reuse(const llm_graph_params & params) { + mctx_msa = static_cast(params.mctx); + + // the parent class operates on the base cache context + this->mctx = mctx_msa->get_base(); + + bool res = true; + + res &= self_k_idxs->ne[0] == params.ubatch.n_tokens; + if (self_k_idxs_idx) { + res &= self_k_idxs_idx->ne[0] == params.ubatch.n_tokens; + } + + res &= can_reuse_kq_mask(self_kq_mask, this->mctx, params.ubatch, params.cparams); + + return res; +} + void llm_graph_input_attn_k_dsa::set_input(const llama_ubatch * ubatch) { mctx->get_mla()->set_input_k_idxs(self_k_idxs_mla, ubatch); @@ -3188,6 +3223,34 @@ llm_graph_input_attn_k_dsa * llm_graph_context::build_attn_inp_k_dsa() const { return (llm_graph_input_attn_k_dsa *) res->add_input(std::move(inp)); } +llm_graph_input_attn_kv_msa * llm_graph_context::build_attn_inp_kv_msa(bool msa_enabled) const { + const auto * mctx_cur = static_cast(mctx); + + auto inp = std::make_unique(hparams, cparams, mctx_cur); + + const auto * mctx_base = mctx_cur->get_base(); + const auto * mctx_idx = mctx_cur->get_idx(); + + { + GGML_ASSERT(hparams.swa_type == LLAMA_SWA_TYPE_NONE && "Use llama_kv_cache_iswa for SWA"); + + inp->self_k_idxs = mctx_base->build_input_k_idxs(ctx0, ubatch); + inp->self_v_idxs = mctx_base->build_input_v_idxs(ctx0, ubatch); + + inp->self_kq_mask = build_attn_inp_kq_mask(ctx0, mctx_base, ubatch, cparams); + inp->self_kq_mask_cnv = inp->self_kq_mask; + } + + inp->self_k_rot = mctx_base->build_input_k_rot(ctx0); + inp->self_v_rot = mctx_base->build_input_v_rot(ctx0); + + if (msa_enabled) { + inp->self_k_idxs_idx = mctx_idx->build_input_k_idxs(ctx0, ubatch); + } + + return (llm_graph_input_attn_kv_msa *) res->add_input(std::move(inp)); +} + // TODO: maybe separate the inner implementation into a separate function // like with the non-sliding window equivalent // once sliding-window hybrid caches are a thing. diff --git a/src/llama-graph.h b/src/llama-graph.h index 160e29413..32d8d395a 100644 --- a/src/llama-graph.h +++ b/src/llama-graph.h @@ -23,6 +23,7 @@ struct llama_memory_context_i; class llama_kv_cache_context; class llama_kv_cache_dsa_context; +class llama_kv_cache_msa_context; class llama_kv_cache_dsv4_raw_context; class llama_kv_cache_dsv4_context; class llama_kv_cache_iswa_context; @@ -425,6 +426,26 @@ public: const llama_kv_cache_dsa_context * mctx; }; +// standard K/V attention input against the base cache, plus destination indices for the indexer key cache +class llm_graph_input_attn_kv_msa : public llm_graph_input_attn_kv { +public: + llm_graph_input_attn_kv_msa( + const llama_hparams & hparams, + const llama_cparams & cparams, + const llama_kv_cache_msa_context * mctx); + ~llm_graph_input_attn_kv_msa() = default; + + void set_input(const llama_ubatch * ubatch) override; + + bool can_reuse(const llm_graph_params & params) override; + + ggml_tensor * get_k_idxs_idx() const { return self_k_idxs_idx; } + + ggml_tensor * self_k_idxs_idx = nullptr; // I64 [n_batch] + + const llama_kv_cache_msa_context * mctx_msa; +}; + class llm_graph_input_attn_kv_iswa : public llm_graph_input_i { public: llm_graph_input_attn_kv_iswa( @@ -1169,6 +1190,8 @@ struct llm_graph_context { llm_graph_input_attn_k_dsa * build_attn_inp_k_dsa() const; + llm_graph_input_attn_kv_msa * build_attn_inp_kv_msa(bool msa_enabled) const; + ggml_tensor * build_attn( llm_graph_input_attn_k_dsa * inp, ggml_tensor * wo, diff --git a/src/llama-hparams.cpp b/src/llama-hparams.cpp index 50af97f35..846d4c69a 100644 --- a/src/llama-hparams.cpp +++ b/src/llama-hparams.cpp @@ -180,16 +180,6 @@ uint32_t llama_hparams::n_embd_v_gqa_max() const { return val; } -uint32_t llama_hparams::n_embd_k_idx(uint32_t il) const { - if (!indexer_kv || indexer_head_size == 0) { - return 0; // arch without a MSA indexer - } - if (il < n_layer_dense_lead) { - return 0; // leading dense layers carry no indexer - } - return indexer_head_size; // 128 -} - uint32_t llama_hparams::n_embd_r() const { if (wkv_head_size != 0) { // for RWKV models diff --git a/src/llama-hparams.h b/src/llama-hparams.h index fc770bf00..6e8336c98 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -230,8 +230,6 @@ struct llama_hparams { // MSA uint32_t indexer_block_size = 0; uint32_t indexer_local_blocks = 0; - // MSA stores its indexer keys in the main KV cache (k_idx tensors); - bool indexer_kv = false; // Indexer is "full" (1) or "shared" (0) // Shared indexers reuse top-k from previous full layer @@ -356,9 +354,6 @@ struct llama_hparams { uint32_t n_embd_k_gqa_max() const; uint32_t n_embd_v_gqa_max() const; - // dimension of the single-head MSA indexer key stream - uint32_t n_embd_k_idx(uint32_t il = 0) const; - // dimension of the rolling state embeddings // corresponds to Mamba's conv_states size or RWKV's token_shift states size uint32_t n_embd_r() const; diff --git a/src/llama-kv-cache-dsa.cpp b/src/llama-kv-cache-dsa.cpp index 241c50365..96cb045d2 100644 --- a/src/llama-kv-cache-dsa.cpp +++ b/src/llama-kv-cache-dsa.cpp @@ -23,7 +23,8 @@ llama_kv_cache_dsa::llama_kv_cache_dsa( uint32_t n_pad, uint32_t n_swa, llama_swa_type swa_type, - const layer_filter_cb & filter, + const layer_filter_cb & filter_mla, + const layer_filter_cb & filter_lid, const layer_reuse_cb & reuse) : hparams_lid(model.hparams), n_stream(unified ? 1 : n_seq_max) { @@ -32,7 +33,7 @@ llama_kv_cache_dsa::llama_kv_cache_dsa( kv_mla = std::make_unique( model, model.hparams, type_k, type_v, v_trans, offload, unified, kv_size, n_seq_max, n_pad, - n_swa, swa_type, nullptr, filter, reuse, nullptr); + n_swa, swa_type, nullptr, filter_mla, reuse, nullptr); // we use llama_kv_cache for caching indexer keys // by hand-tweaking some hparams we fool it to create @@ -49,7 +50,7 @@ llama_kv_cache_dsa::llama_kv_cache_dsa( kv_lid = std::make_unique( model, hparams_lid, type_k, type_v, v_trans, offload, unified, kv_size, n_seq_max, n_pad, - n_swa, swa_type, nullptr, filter, reuse, nullptr); + n_swa, swa_type, nullptr, filter_lid, reuse, nullptr); } void llama_kv_cache_dsa::clear(bool data) { diff --git a/src/llama-kv-cache-dsa.h b/src/llama-kv-cache-dsa.h index e2b330993..e74fc4d91 100644 --- a/src/llama-kv-cache-dsa.h +++ b/src/llama-kv-cache-dsa.h @@ -26,7 +26,8 @@ public: uint32_t n_pad, uint32_t n_swa, llama_swa_type swa_type, - const layer_filter_cb & filter, + const layer_filter_cb & filter_mla, + const layer_filter_cb & filter_lid, const layer_reuse_cb & reuse); ~llama_kv_cache_dsa() = default; diff --git a/src/llama-kv-cache-msa.cpp b/src/llama-kv-cache-msa.cpp new file mode 100644 index 000000000..55ef286ca --- /dev/null +++ b/src/llama-kv-cache-msa.cpp @@ -0,0 +1,395 @@ +#include "llama-kv-cache-msa.h" + +#include "llama-impl.h" +#include "llama-batch.h" +#include "llama-model.h" + +#include +#include +#include + +// llama_kv_cache_msa + +llama_kv_cache_msa::llama_kv_cache_msa( + const llama_model & model, + ggml_type type_k, + ggml_type type_v, + bool v_trans, + bool offload, + bool unified, + uint32_t kv_size, + uint32_t n_seq_max, + uint32_t n_pad, + uint32_t n_swa, + llama_swa_type swa_type, + const layer_filter_cb & filter, + const layer_filter_cb & filter_idx, + const layer_reuse_cb & reuse) : + hparams_idx(model.hparams), + n_stream(unified ? 1 : n_seq_max), n_seq_max(n_seq_max), n_pad(n_pad), + n_swa(n_swa), swa_type(swa_type) { + + LLAMA_LOG_INFO("%s: creating main KV cache, size = %u cells\n", __func__, kv_size); + + kv_base = std::make_unique( + model, model.hparams, type_k, type_v, + v_trans, offload, unified, kv_size, n_seq_max, n_pad, + n_swa, swa_type, nullptr, filter, reuse, nullptr); + + // the MSA indexer uses a single key head per layer + std::fill(hparams_idx.n_head_kv_arr.begin(), hparams_idx.n_head_kv_arr.end(), 1); + hparams_idx.n_embd_head_k_full = model.hparams.indexer_head_size; + // the rope parameters are kept identical to the main cache + + LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size); + + kv_idx = std::make_unique( + model, hparams_idx, type_k, type_v, + v_trans, offload, unified, kv_size, n_seq_max, n_pad, + n_swa, swa_type, nullptr, filter_idx, reuse, nullptr); +} + +void llama_kv_cache_msa::clear(bool data) { + kv_base->clear(data); + kv_idx ->clear(data); +} + +bool llama_kv_cache_msa::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { + bool res = true; + + res = res & kv_base->seq_rm(seq_id, p0, p1); + res = res & kv_idx ->seq_rm(seq_id, p0, p1); + + return res; +} + +void llama_kv_cache_msa::seq_cp(llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) { + kv_base->seq_cp(seq_id_src, seq_id_dst, p0, p1); + kv_idx ->seq_cp(seq_id_src, seq_id_dst, p0, p1); +} + +void llama_kv_cache_msa::seq_keep(llama_seq_id seq_id) { + kv_base->seq_keep(seq_id); + kv_idx ->seq_keep(seq_id); +} + +void llama_kv_cache_msa::seq_add(llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) { + kv_base->seq_add(seq_id, p0, p1, shift); + kv_idx ->seq_add(seq_id, p0, p1, shift); +} + +void llama_kv_cache_msa::seq_div(llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) { + kv_base->seq_div(seq_id, p0, p1, d); + kv_idx ->seq_div(seq_id, p0, p1, d); +} + +llama_pos llama_kv_cache_msa::seq_pos_min(llama_seq_id seq_id) const { + return kv_base->seq_pos_min(seq_id); +} + +llama_pos llama_kv_cache_msa::seq_pos_max(llama_seq_id seq_id) const { + return kv_base->seq_pos_max(seq_id); +} + +std::map llama_kv_cache_msa::memory_breakdown() const { + std::map mb = kv_base->memory_breakdown(); + for (const auto & buft_size : kv_idx->memory_breakdown()) { + mb[buft_size.first] += buft_size.second; + } + return mb; +} + +llama_memory_context_ptr llama_kv_cache_msa::init_batch( + llama_batch_allocr & balloc, + uint32_t n_ubatch, + bool embd_all) { + GGML_UNUSED(embd_all); + + do { + balloc.split_reset(); + + std::vector ubatches; + while (true) { + auto ubatch = n_stream == 1 ? balloc.split_simple(n_ubatch) : balloc.split_equal(n_ubatch, true, 0); + + if (ubatch.n_tokens == 0) { + break; + } + + ubatches.push_back(std::move(ubatch)); + } + + if (balloc.get_n_used() < balloc.get_n_tokens()) { + // failed to find a suitable split + break; + } + + auto sinfos_base = kv_base->prepare(ubatches); + if (sinfos_base.empty()) { + break; + } + + auto sinfos_idx = kv_idx->prepare(ubatches); + if (sinfos_idx.empty()) { + break; + } + + assert(sinfos_base.size() == sinfos_idx.size()); + + return std::make_unique( + this, std::move(sinfos_base), std::move(sinfos_idx), std::move(ubatches)); + } while (false); + + return std::make_unique(LLAMA_MEMORY_STATUS_FAILED_PREPARE); +} + +llama_memory_context_ptr llama_kv_cache_msa::init_full() { + return std::make_unique(this); +} + +llama_memory_context_ptr llama_kv_cache_msa::init_update(llama_context * lctx, bool optimize) { + return std::make_unique(this, lctx, optimize); +} + +bool llama_kv_cache_msa::get_can_shift() const { + return kv_base->get_can_shift() && + kv_idx ->get_can_shift() && + kv_base->get_size() == kv_idx->get_size(); +} + +void llama_kv_cache_msa::state_write(llama_io_write_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) const { + kv_base->state_write(io, seq_id, flags); + kv_idx ->state_write(io, seq_id, flags); +} + +void llama_kv_cache_msa::state_read(llama_io_read_i & io, llama_seq_id seq_id, llama_state_seq_flags flags) { + kv_base->state_read(io, seq_id, flags); + kv_idx ->state_read(io, seq_id, flags); +} + +llama_kv_cache * llama_kv_cache_msa::get_base() const { + return kv_base.get(); +} + +llama_kv_cache * llama_kv_cache_msa::get_idx() const { + return kv_idx.get(); +} + +// llama_kv_cache_msa_context + +llama_kv_cache_msa_context::llama_kv_cache_msa_context(llama_memory_status status) : + kv(nullptr), status(status) {} + +llama_kv_cache_msa_context::llama_kv_cache_msa_context( + llama_kv_cache_msa * kv) : + kv(kv), + ctx_base(kv->get_base()->init_full()), + ctx_idx (kv->get_idx ()->init_full()), + status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) { +} + +llama_kv_cache_msa_context::llama_kv_cache_msa_context( + llama_kv_cache_msa * kv, + llama_context * lctx, + bool optimize) : + kv(kv), + ctx_base(kv->get_base()->init_update(lctx, optimize)), + ctx_idx (kv->get_idx ()->init_update(lctx, optimize)), + status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) { +} + +llama_kv_cache_msa_context::llama_kv_cache_msa_context( + llama_kv_cache_msa * kv, + slot_info_vec_t sinfos_base, + slot_info_vec_t sinfos_idx, + std::vector ubatches) : + kv(kv), + ubatches(std::move(ubatches)), + // here we copy the ubatches. not sure if this is ideal + ctx_base(new llama_kv_cache_context(kv->get_base(), std::move(sinfos_base), this->ubatches)), + ctx_idx (new llama_kv_cache_context(kv->get_idx (), std::move(sinfos_idx), this->ubatches)), + status(llama_memory_status_combine(ctx_base->get_status(), ctx_idx->get_status())) { +} + +llama_kv_cache_msa_context::~llama_kv_cache_msa_context() = default; + +bool llama_kv_cache_msa_context::next() { + assert(status == LLAMA_MEMORY_STATUS_SUCCESS); + + ctx_base->next(); + ctx_idx ->next(); + + if (++i_next >= ubatches.size()) { + return false; + } + + return true; +} + +bool llama_kv_cache_msa_context::apply() { + assert(!llama_memory_status_is_fail(status)); + + bool res = true; + + res = res & ctx_base->apply(); + res = res & ctx_idx ->apply(); + + return res; +} + +llama_memory_status llama_kv_cache_msa_context::get_status() const { + return status; +} + +const llama_ubatch & llama_kv_cache_msa_context::get_ubatch() const { + assert(status == LLAMA_MEMORY_STATUS_SUCCESS); + + return ubatches[i_next]; +} + +const llama_kv_cache_context * llama_kv_cache_msa_context::get_base() const { + assert(status == LLAMA_MEMORY_STATUS_SUCCESS); + + return static_cast(ctx_base.get()); +} + +const llama_kv_cache_context * llama_kv_cache_msa_context::get_idx() const { + assert(status == LLAMA_MEMORY_STATUS_SUCCESS); + + return static_cast(ctx_idx.get()); +} + +uint32_t llama_kv_cache_msa_context::get_n_pos() const { + // pad the value so that the graph remains constant across batches and can be reused + const uint32_t n_pad_cur = std::max(kv->get_n_pad(), 256u); + + llama_pos pos_max = -1; + + for (llama_seq_id seq_id = 0; seq_id < (llama_seq_id) kv->get_n_seq_max(); ++seq_id) { + pos_max = std::max(pos_max, kv->seq_pos_max(seq_id)); + } + + return std::max(n_pad_cur, GGML_PAD((uint32_t) (pos_max + 1), n_pad_cur)); +} + +void llama_kv_cache_msa_context::set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const { + GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer)); + GGML_ASSERT(dst->type == GGML_TYPE_I32); + GGML_ASSERT(div > 0); + + const int64_t n_tokens = ubatch->n_tokens; + const int64_t n_kv = dst->ne[0]; + const int64_t n_stream_ub = dst->ne[1]; + + GGML_ASSERT(n_tokens % n_stream_ub == 0); + const int64_t n_tps = n_tokens/n_stream_ub; + + int32_t * data = (int32_t *) dst->data; + + for (int64_t s = 0; s < n_stream_ub; ++s) { + const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0]; + + const auto & cells = kv->get_base()->get_cells(seq_id); + + for (int64_t j = 0; j < n_kv; ++j) { + // the value for empty or other-sequence cells is irrelevant as consumers mask them + data[s*n_kv + j] = + cells.is_empty(j) || !cells.seq_has(j, seq_id) + ? 0 + : (int32_t) (cells.pos_get(j)/div); + } + } +} + +void llama_kv_cache_msa_context::set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const { + GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer)); + GGML_ASSERT(dst->type == GGML_TYPE_I32 || dst->type == GGML_TYPE_F32); + + const int64_t n_tokens = ubatch->n_tokens; + const int64_t n_pos = dst->ne[0]; + const int64_t n_stream_ub = dst->ne[1]; + + GGML_ASSERT(n_tokens % n_stream_ub == 0); + const int64_t n_tps = n_tokens/n_stream_ub; + + for (int64_t s = 0; s < n_stream_ub; ++s) { + const llama_seq_id seq_id = ubatch->seq_id[s*n_tps][0]; + + const auto & cells = kv->get_base()->get_cells(seq_id); + + std::vector map(n_pos, 0); + + for (uint32_t j = 0; j < cells.size(); ++j) { + if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) { + continue; + } + + const llama_pos p0 = cells.pos_get(j); + + if (p0 < 0 || p0 >= n_pos) { + continue; + } + + map[p0] = (int32_t) j; + } + + if (dst->type == GGML_TYPE_I32) { + int32_t * data = (int32_t *) dst->data + s*n_pos; + std::copy(map.begin(), map.end(), data); + } else { + float * data = (float *) dst->data + s*n_pos; + for (int64_t p = 0; p < n_pos; ++p) { + data[p] = (float) map[p]; + } + } + } +} + +void llama_kv_cache_msa_context::set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const { + GGML_ASSERT(ggml_backend_buffer_is_host(dst->buffer)); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + const int64_t n_tokens = ubatch->n_tokens; + const int64_t n_pos = dst->ne[0]; + + GGML_ASSERT(dst->ne[1] == n_tokens); + + const uint32_t n_swa = kv->get_n_swa(); + const llama_swa_type swa_type = kv->get_swa_type(); + + float * data = (float *) dst->data; + + std::fill(data, data + n_pos*n_tokens, -INFINITY); + + for (int64_t i = 0; i < n_tokens; ++i) { + const llama_seq_id seq_id = ubatch->seq_id[i][0]; + + const auto & cells = kv->get_base()->get_cells(seq_id); + + const llama_pos p1 = ubatch->pos[i]; + + for (uint32_t j = 0; j < cells.size(); ++j) { + if (cells.is_empty(j) || !cells.seq_has(j, seq_id)) { + continue; + } + + const llama_pos p0 = cells.pos_get(j); + + if (p0 < 0 || p0 >= n_pos) { + continue; + } + + // causal mask + if (p0 > p1) { + continue; + } + + // apply SWA if any + if (llama_hparams::is_masked_swa(n_swa, swa_type, p0, p1)) { + continue; + } + + data[i*n_pos + p0] = 0.0f; + } + } +} diff --git a/src/llama-kv-cache-msa.h b/src/llama-kv-cache-msa.h new file mode 100644 index 000000000..f09b6d32b --- /dev/null +++ b/src/llama-kv-cache-msa.h @@ -0,0 +1,153 @@ +#pragma once + +#include "llama-kv-cache.h" + +#include + +// llama_kv_cache_msa + +// uses two instances of llama_kv_cache, one for K/V tensors, and one for the MSA indexer tensors +// both receive identical sequence operations and identical ubatches, so their cell layouts stay in synced. +// the context also exposes per-ubatch pos - cell translation maps populated from llama_kv_cells via +// llama_kv_cache::get_cells(), which the model graph uses to run MSA block selection in position space + +class llama_kv_cache_msa : public llama_memory_i { +public: + llama_kv_cache_msa( + const llama_model & model, + ggml_type type_k, + ggml_type type_v, + bool v_trans, + bool offload, + bool unified, + uint32_t kv_size, + uint32_t n_seq_max, + uint32_t n_pad, + uint32_t n_swa, + llama_swa_type swa_type, + const layer_filter_cb & filter, + const layer_filter_cb & filter_idx, + const layer_reuse_cb & reuse); + + ~llama_kv_cache_msa() = default; + + // llama_memory_i + + llama_memory_context_ptr init_batch( + llama_batch_allocr & balloc, + uint32_t n_ubatch, + bool embd_all) override; + + llama_memory_context_ptr init_full() override; + + llama_memory_context_ptr init_update(llama_context * lctx, bool optimize) override; + + bool get_can_shift() const override; + + void clear(bool data) override; + + bool seq_rm (llama_seq_id seq_id, llama_pos p0, llama_pos p1) override; + void seq_cp (llama_seq_id seq_id_src, llama_seq_id seq_id_dst, llama_pos p0, llama_pos p1) override; + void seq_keep(llama_seq_id seq_id) override; + void seq_add (llama_seq_id seq_id, llama_pos p0, llama_pos p1, llama_pos shift) override; + void seq_div (llama_seq_id seq_id, llama_pos p0, llama_pos p1, int d) override; + + llama_pos seq_pos_min(llama_seq_id seq_id) const override; + llama_pos seq_pos_max(llama_seq_id seq_id) const override; + + std::map memory_breakdown() const override; + + // state write/load + + void state_write(llama_io_write_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) const override; + void state_read (llama_io_read_i & io, llama_seq_id seq_id = -1, llama_state_seq_flags flags = 0) override; + + // llama_kv_cache_msa specific API + + llama_kv_cache * get_base() const; + llama_kv_cache * get_idx () const; + + uint32_t get_n_pad() const { return n_pad; } + uint32_t get_n_seq_max() const { return n_seq_max; } + uint32_t get_n_swa() const { return n_swa; } + llama_swa_type get_swa_type() const { return swa_type; } + +private: + // keep the indexer KV cache hparams instance here as llama_kv_cache stores only a reference + llama_hparams hparams_idx; + + const uint32_t n_stream = 1; + const uint32_t n_seq_max = 1; + const uint32_t n_pad = 1; + + const uint32_t n_swa = 0; + const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE; + + std::unique_ptr kv_base; + std::unique_ptr kv_idx; +}; + +class llama_kv_cache_msa_context : public llama_memory_context_i { +public: + using slot_info_vec_t = llama_kv_cache::slot_info_vec_t; + + // used for errors + llama_kv_cache_msa_context(llama_memory_status status); + + // used to create a full-cache context + llama_kv_cache_msa_context( + llama_kv_cache_msa * kv); + + // used to create an update context + llama_kv_cache_msa_context( + llama_kv_cache_msa * kv, + llama_context * lctx, + bool optimize); + + // used to create a batch processing context from a batch + llama_kv_cache_msa_context( + llama_kv_cache_msa * kv, + slot_info_vec_t sinfos_base, + slot_info_vec_t sinfos_idx, + std::vector ubatches); + + virtual ~llama_kv_cache_msa_context(); + + // llama_memory_context_i + + bool next() override; + bool apply() override; + + llama_memory_status get_status() const override; + const llama_ubatch & get_ubatch() const override; + + // llama_kv_cache_msa_context specific API + + const llama_kv_cache_context * get_base() const; + const llama_kv_cache_context * get_idx () const; + + // max position currently present in the cache plus one, padded MSA blocks are defined over token positions + // so the block-selection tensors are sized by this value rather than by the number of cells + uint32_t get_n_pos() const; + + // position <-> cell translation maps, populated from the base cache cells + // the model graph relates cache contents to token positions only through these per ubatch inputs + // value for empty or other-sequence cells is 0 so consumers must mask them + void set_input_cell_pos(ggml_tensor * dst, const llama_ubatch * ubatch, int32_t div) const; + // positions without a cell map to cell 0, consumers must mask them assumes one sequence per stream + void set_input_pos_slot(ggml_tensor * dst, const llama_ubatch * ubatch) const; + void set_input_pos_mask(ggml_tensor * dst, const llama_ubatch * ubatch) const; + +private: + llama_kv_cache_msa * kv; + + // the index of the next ubatch to process + size_t i_next = 0; + + std::vector ubatches; + + const llama_memory_context_ptr ctx_base; + const llama_memory_context_ptr ctx_idx; + + const llama_memory_status status; +}; diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index 09ee65978..46bbc992a 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -112,7 +112,7 @@ llama_kv_cache::llama_kv_cache( auto it = ctx_map.find(buft); if (it == ctx_map.end()) { ggml_init_params params = { - /*.mem_size =*/ size_t(3u*(1 + n_stream)*n_layer*ggml_tensor_overhead()), //Reserve tensor metadata for up to 3 tensors per layer (K, V, and optional K_idx), plus one view per tensor per stream. + /*.mem_size =*/ size_t(2u*(1 + n_stream)*n_layer*ggml_tensor_overhead()), /*.mem_buffer =*/ NULL, /*.no_alloc =*/ true, }; @@ -242,25 +242,9 @@ llama_kv_cache::llama_kv_cache( v_stream.push_back(has_v ? ggml_view_2d(ctx, v, n_embd_v_gqa, kv_size, v->nb[1], s*v->nb[2]) : nullptr); } - const uint32_t n_embd_k_idx = hparams.n_embd_k_idx(il); - ggml_tensor * k_idx = n_embd_k_idx > 0 - ? ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd_k_idx, kv_size, n_stream) - : nullptr; - if (k_idx) { - ggml_format_name(k_idx, "cache_k_idx_l%d", il); - msa_strict_slots = (n_stream == n_seq_max); - } - - std::vector k_idx_stream; - for (uint32_t s = 0; s < n_stream; ++s) { - k_idx_stream.push_back(k_idx - ? ggml_view_2d(ctx, k_idx, n_embd_k_idx, kv_size, k_idx->nb[1], s*k_idx->nb[2]) - : nullptr); - } - map_layer_ids[il] = layers.size(); - layers.push_back({ il, k, v, k_idx, k_stream, v_stream, k_idx_stream }); + layers.push_back({ il, k, v, k_stream, v_stream, }); } if (reuse) { @@ -309,24 +293,13 @@ llama_kv_cache::llama_kv_cache( } { - const size_t memory_size_k = size_k_bytes(); - const size_t memory_size_v = size_v_bytes(); - const size_t memory_size_k_idx = size_k_idx_bytes(); - const size_t memory_size_total = memory_size_k + memory_size_v + memory_size_k_idx; + const size_t memory_size_k = size_k_bytes(); + const size_t memory_size_v = size_v_bytes(); - constexpr float mib = 1024.0f * 1024.0f; - - const std::string k_log = format(", K (%s): %7.2f MiB", ggml_type_name(type_k), (float) memory_size_k / mib); - const std::string v_log = format(", V (%s): %7.2f MiB", ggml_type_name(type_v), (float) memory_size_v / mib); - - std::string k_idx_log; - if (memory_size_k_idx > 0) { - k_idx_log = format(", K_idx (%s): %7.2f MiB", ggml_type_name(GGML_TYPE_F32), (float) memory_size_k_idx / mib); - } - - LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs)%s%s%s\n", __func__, - (float) memory_size_total / mib, kv_size, (int) layers.size(), n_seq_max, n_stream, - k_log.c_str(), v_log.c_str(), k_idx_log.c_str()); + LLAMA_LOG_INFO("%s: size = %7.2f MiB (%6u cells, %3d layers, %2u/%u seqs), K (%s): %7.2f MiB, V (%s): %7.2f MiB\n", __func__, + (float)(memory_size_k + memory_size_v) / (1024.0f * 1024.0f), kv_size, (int) layers.size(), n_seq_max, n_stream, + ggml_type_name(type_k), (float)memory_size_k / (1024.0f * 1024.0f), + ggml_type_name(type_v), (float)memory_size_v / (1024.0f * 1024.0f)); } // TODO: refactor [TAG_KV_CACHE_SHARE_CELLS] @@ -419,39 +392,6 @@ bool llama_kv_cache::seq_rm(llama_seq_id seq_id, llama_pos p0, llama_pos p1) { p1 = std::numeric_limits::max(); } - // empty range - nothing to remove - if (p0 >= p1) { - return true; - } - - // MSA anchors block selection to absolute cache slots (slot == position). Tail trim and full removal preserve this invariant, but removing a prefix - // or middle range would free slots while later cells survive, desynchronizing the indexer cache. Reject such removals before modifying the cache. - if (msa_strict_slots) { - for (llama_seq_id sid = 0; sid < (llama_seq_id) seq_to_stream.size(); ++sid) { - if (seq_id >= 0 && sid != seq_id) { - continue; - } - - const auto & cells = v_cells[seq_to_stream[sid]]; - - const llama_pos pmin = cells.seq_pos_min(sid); - const llama_pos pmax = cells.seq_pos_max(sid); - - if (pmin < 0) { - continue; // empty sequence - } - - const bool overlaps = p0 <= pmax && p1 > pmin; // the range removes something - const bool leaves_tail = p1 <= pmax; // cells beyond the range survive - - if (overlaps && leaves_tail) { - LLAMA_LOG_WARN("%s: MSA: partial (non-suffix) removal [%d, %d) for seq %d is not supported " - "(block selection is anchored to cache slots) - rejected\n", __func__, p0, p1, sid); - return false; - } - } - } - if (seq_id >= 0) { auto & cells = v_cells[seq_to_stream[seq_id]]; auto & head = v_heads[seq_to_stream[seq_id]]; @@ -907,10 +847,6 @@ bool llama_kv_cache::update(llama_context * lctx, bool do_shift, const stream_co if (layer.v_stream[ssrc]) { ggml_backend_tensor_copy(layer.v_stream[ssrc], layer.v_stream[sdst]); } - if (layer.k_idx_stream[ssrc]) { - GGML_ASSERT(layer.k_idx_stream[sdst]); - ggml_backend_tensor_copy(layer.k_idx_stream[ssrc], layer.k_idx_stream[sdst]); - } } } } @@ -1063,44 +999,6 @@ llama_kv_cache::slot_info llama_kv_cache::find_slot(const llama_ubatch & ubatch, const auto & cells = v_cells[seq_to_stream[seq_id]]; - if (n_tokens > cells.size()) { - LLAMA_LOG_ERROR("%s: n_tokens = %d > size = %u\n", __func__, n_tokens, cells.size()); - return { }; - } - - // MSA block selection assumes slot == logical position (append-only streams). - if (msa_strict_slots) { - for (uint32_t ii = 0; ii < n_tokens; ++ii) { - const llama_pos pos = ubatch.pos[s*n_tokens + ii]; - - if (pos < 0 || (uint64_t) pos >= cells.size()) { - LLAMA_LOG_WARN("%s: MSA: position %d is outside the cache range [0, %u)\n", - __func__, pos, cells.size()); - return { }; - } - - const uint32_t idx = (uint32_t) pos; - - if (!cells.is_empty(idx)) { - LLAMA_LOG_WARN("%s: MSA: required slot %u is already occupied (stream %u)\n", - __func__, idx, seq_to_stream[seq_id]); - return { }; - } - - // strictly increasing positions, rules out duplicates and, for contiguous requests, is tightened to exact adjacency - if (!res.idxs[s].empty() && (cont ? idx != res.idxs[s].back() + 1 - : idx <= res.idxs[s].back())) { - LLAMA_LOG_WARN("%s: MSA: token positions are not %s within the ubatch\n", - __func__, cont ? "contiguous" : "strictly increasing"); - return { }; - } - - res.idxs[s].push_back(idx); - } - - continue; - } - uint32_t head_cur = v_heads[seq_to_stream[seq_id]]; // if we have enough unused cells before the current head -> @@ -1109,6 +1007,11 @@ llama_kv_cache::slot_info llama_kv_cache::find_slot(const llama_ubatch & ubatch, head_cur = 0; } + if (n_tokens > cells.size()) { + LLAMA_LOG_ERROR("%s: n_tokens = %d > size = %u\n", __func__, n_tokens, cells.size()); + return { }; + } + uint32_t n_tested = 0; // for continuous slots, we test that all tokens in the ubatch fit, starting from the current head @@ -1215,15 +1118,6 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & const auto idx = sinfo.idxs[s][ii]; - if (msa_strict_slots && (llama_pos) idx != ubatch.pos[i]) { - LLAMA_LOG_ERROR("%s: MSA slot/position invariant violated: " - "writing pos %d into cell %u (stream %u). The indexer cache " - "would desync and block selection would silently corrupt. " - "This is a bug, please report it with reproduction steps.\n", - __func__, ubatch.pos[i], idx, sinfo.strm[s]); - GGML_ABORT("MSA: slot != pos"); - } - if (!cells.is_empty(idx)) { assert(cells.seq_count(idx) == 1); @@ -1267,8 +1161,7 @@ void llama_kv_cache::apply_ubatch(const slot_info & sinfo, const llama_ubatch & LLAMA_LOG_DEBUG("%s: purging positions [%d, %d] of sequence %d from KV cache\n", __func__, cells.seq_pos_min(s), seq_pos_max_rm[s], s); - // under MSA strict slots this path should be unreachable, since strict MSA placement never selects occupied cells - GGML_ASSERT(seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1)); + seq_rm(s, cells.seq_pos_min(s), seq_pos_max_rm[s] + 1); } } @@ -1288,12 +1181,6 @@ bool llama_kv_cache::get_can_shift() const { if (hparams.n_pos_per_embd() > 1) { return false; } - // shifting would leave k_idx stale - for (const auto & layer : layers) { - if (layer.k_idx) { - return false; - } - } return true; } @@ -1342,6 +1229,12 @@ ggml_tensor * llama_kv_cache::get_k_storage(int32_t il) const { return layers[ikv].k; } +const llama_kv_cells & llama_kv_cache::get_cells(llama_seq_id seq_id) const { + GGML_ASSERT(seq_id >= 0 && (size_t) seq_id < seq_to_stream.size()); + + return v_cells[seq_to_stream[seq_id]]; +} + uint32_t llama_kv_cache::get_n_kv(const slot_info & sinfo) const { uint32_t result = 0; @@ -1410,23 +1303,6 @@ ggml_tensor * llama_kv_cache::get_v(ggml_context * ctx, int32_t il, uint32_t n_k ggml_row_size(v->type, kv_size*n_embd_v_gqa)*sinfo.s0); } -ggml_tensor * llama_kv_cache::get_k_idx(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const { - const int32_t ikv = map_layer_ids.at(il); - auto * k_idx = layers[ikv].k_idx; - GGML_ASSERT(k_idx); - - const uint64_t kv_size = get_size(); - const int64_t n_idx = k_idx->ne[0]; // 128 - const uint32_t ns = sinfo.s1 - sinfo.s0 + 1; - - return ggml_view_4d(ctx, k_idx, - n_idx, 1, n_kv, ns, - ggml_row_size(k_idx->type, n_idx), // nb1 (single head) - ggml_row_size(k_idx->type, n_idx), // nb2 (per cell) - ggml_row_size(k_idx->type, n_idx*kv_size), // nb3 (per stream) - ggml_row_size(k_idx->type, n_idx*kv_size)*sinfo.s0); -} - ggml_tensor * llama_kv_cache::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const { GGML_UNUSED(sinfo); @@ -1528,28 +1404,6 @@ ggml_tensor * llama_kv_cache::build_input_k_idxs(ggml_context * ctx, const llama return k_idxs; } -ggml_tensor * llama_kv_cache::cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const { - GGML_UNUSED(sinfo); - const int32_t ikv = map_layer_ids.at(il); - ggml_tensor * k_idx = layers[ikv].k_idx; - GGML_ASSERT(k_idx && "cpy_k_idx on a layer with no indexer cache"); - - const int64_t n_embd_head = k_idx_cur->ne[0]; // 128 - const int64_t n_head = k_idx_cur->ne[1]; // 1 - const int64_t n_tokens = k_idx_cur->ne[2]; - const int64_t n_embd_gqa = n_embd_head*n_head; // 128 - - GGML_ASSERT(ggml_row_size(k_idx_cur->type, n_embd_head) == k_idx_cur->nb[1]); - k_idx_cur = ggml_view_2d(ctx, k_idx_cur, n_embd_gqa, n_tokens, k_idx_cur->nb[2], 0); - - const int64_t n_stream = k_idx->ne[2]; - if (n_stream > 1) { - const int64_t kv_size = get_size(); - k_idx = ggml_reshape_2d(ctx, k_idx, n_embd_gqa, kv_size*n_stream); - } - return ggml_set_rows(ctx, k_idx, k_idx_cur, k_idxs); // same k_idxs as the K store -} - ggml_tensor * llama_kv_cache::build_input_v_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const { const uint32_t n_tokens = ubatch.n_tokens; @@ -1984,18 +1838,6 @@ size_t llama_kv_cache::size_v_bytes() const { return size_v_bytes; } -size_t llama_kv_cache::size_k_idx_bytes() const { - size_t size_k_idx_bytes = 0; - - for (const auto & layer : layers) { - if (layer.k_idx) { - size_k_idx_bytes += ggml_nbytes(layer.k_idx); - } - } - - return size_k_idx_bytes; -} - ggml_tensor * llama_kv_cache::build_rope_shift( const llama_cparams & cparams, ggml_context * ctx, @@ -2308,36 +2150,6 @@ void llama_kv_cache::state_write_data(llama_io_write_i & io, const cell_ranges_t } } - if (size_k_idx_bytes() > 0) { - const uint32_t has_k_idx_u32 = 1; - io.write(&has_k_idx_u32, sizeof(has_k_idx_u32)); - - for (const auto & layer : layers) { - const uint32_t layer_has_k_idx = layer.k_idx ? 1 : 0; - io.write(&layer_has_k_idx, sizeof(layer_has_k_idx)); - - if (!layer_has_k_idx) { - continue; - } - - GGML_ASSERT(layer.k_idx_stream[cr.strm]); - - const int32_t k_idx_type_i = (int32_t) layer.k_idx->type; - io.write(&k_idx_type_i, sizeof(k_idx_type_i)); - - const uint64_t k_idx_size_row = ggml_row_size(layer.k_idx->type, layer.k_idx->ne[0]); - io.write(&k_idx_size_row, sizeof(k_idx_size_row)); - - for (const auto & range : cr.data) { - const size_t range_size = range.second - range.first; - const size_t buf_size = range_size * k_idx_size_row; - const size_t offset = range.first * k_idx_size_row; - - io.write_tensor(layer.k_idx_stream[cr.strm], offset, buf_size); - } - } - } - if (!v_trans) { for (const auto & layer : layers) { const uint32_t il = layer.il; @@ -2586,68 +2398,6 @@ bool llama_kv_cache::state_read_data(llama_io_read_i & io, uint32_t strm, uint32 } } - if (size_k_idx_bytes() > 0) { - uint32_t has_k_idx_u32 = 0; - io.read(&has_k_idx_u32, sizeof(has_k_idx_u32)); - - if (has_k_idx_u32 != 1) { - LLAMA_LOG_ERROR("%s: missing k_idx data in KV cache state\n", __func__); - return false; - } - - for (const auto & layer : layers) { - uint32_t layer_has_k_idx = 0; - io.read(&layer_has_k_idx, sizeof(layer_has_k_idx)); - - const uint32_t expected_layer_has_k_idx = layer.k_idx ? 1 : 0; - - if (layer_has_k_idx != expected_layer_has_k_idx) { - LLAMA_LOG_ERROR( - "%s: mismatched k_idx state for layer: got %u, expected %u\n", - __func__, layer_has_k_idx, expected_layer_has_k_idx); - return false; - } - - if (!layer_has_k_idx) { - continue; - } - - GGML_ASSERT(layer.k_idx_stream[strm]); - - int32_t k_idx_type_i = -1; - io.read(&k_idx_type_i, sizeof(k_idx_type_i)); - - if (k_idx_type_i != (int32_t) layer.k_idx->type) { - LLAMA_LOG_ERROR( - "%s: mismatched k_idx type: got %d, expected %d\n", - __func__, k_idx_type_i, (int32_t) layer.k_idx->type); - return false; - } - - uint64_t k_idx_size_row = 0; - io.read(&k_idx_size_row, sizeof(k_idx_size_row)); - - const uint64_t expected_k_idx_size_row = ggml_row_size(layer.k_idx->type, layer.k_idx->ne[0]); - - if (k_idx_size_row != expected_k_idx_size_row) { - LLAMA_LOG_ERROR( - "%s: mismatched k_idx row size: got %zu, expected %zu\n", - __func__, (size_t) k_idx_size_row, (size_t) expected_k_idx_size_row); - return false; - } - - if (cell_count) { - if (sinfo.is_contiguous()) { - io.read_tensor(layer.k_idx_stream[strm], sinfo.head() * k_idx_size_row, cell_count * k_idx_size_row); - } else { - for (uint32_t i = 0; i < cell_count; ++i) { - io.read_tensor(layer.k_idx_stream[strm], sinfo.idxs[0][i] * k_idx_size_row, k_idx_size_row); - } - } - } - } - } - if (!this->v_trans) { for (const auto & layer : layers) { const uint32_t il = layer.il; @@ -2849,10 +2599,6 @@ ggml_tensor * llama_kv_cache_context::get_v(ggml_context * ctx, int32_t il) cons return kv->get_v(ctx, il, n_kv, sinfos[i_cur]); } -ggml_tensor * llama_kv_cache_context::get_k_idx(ggml_context * ctx, int32_t il) const { - return kv->get_k_idx(ctx, il, n_kv, sinfos[i_cur]); -} - ggml_tensor * llama_kv_cache_context::cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const { return kv->cpy_k(ctx, k_cur, k_idxs, il, sinfos[i_cur]); } @@ -2861,10 +2607,6 @@ ggml_tensor * llama_kv_cache_context::cpy_v(ggml_context * ctx, ggml_tensor * v_ return kv->cpy_v(ctx, v_cur, v_idxs, il, sinfos[i_cur]); } -ggml_tensor * llama_kv_cache_context::cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il) const { - return kv->cpy_k_idx(ctx, k_idx_cur, k_idxs, il, sinfos[i_cur]); -} - ggml_tensor * llama_kv_cache_context::build_input_k_idxs(ggml_context * ctx, const llama_ubatch & ubatch) const { return kv->build_input_k_idxs(ctx, ubatch); } diff --git a/src/llama-kv-cache.h b/src/llama-kv-cache.h index d5a92f440..6cb6dbd2f 100644 --- a/src/llama-kv-cache.h +++ b/src/llama-kv-cache.h @@ -164,6 +164,8 @@ public: std::vector get_layer_ids() const; ggml_tensor * get_k_storage(int32_t il) const; + const llama_kv_cells & get_cells(llama_seq_id seq_id) const; + // // graph_build API // @@ -173,12 +175,10 @@ public: // get views of the current state of the cache ggml_tensor * get_k(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; ggml_tensor * get_v(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; - ggml_tensor * get_k_idx(ggml_context * ctx, int32_t il, uint32_t n_kv, const slot_info & sinfo) const; // store k_cur and v_cur in the cache based on the provided head location ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const; ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il, const slot_info & sinfo) const; - ggml_tensor * cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il, const slot_info & sinfo) const; // // preparation API @@ -230,11 +230,9 @@ private: ggml_tensor * k; ggml_tensor * v; - ggml_tensor * k_idx; // MSA single-head indexer keys, F32 std::vector k_stream; std::vector v_stream; - std::vector k_idx_stream; }; bool v_trans = true; // the value tensor is transposed @@ -263,9 +261,6 @@ private: // env: LLAMA_KV_CACHE_DEBUG int debug = 0; - // set when a k_idx (indexer) cache exists and the stream layout supports MSA (single seq, or one stream per seq) - bool msa_strict_slots = false; - // this is the SWA type of the cache - not to be confused with the model SWA type const llama_swa_type swa_type = LLAMA_SWA_TYPE_NONE; @@ -298,7 +293,6 @@ private: size_t size_k_bytes() const; size_t size_v_bytes() const; - size_t size_k_idx_bytes() const; ggml_tensor * build_rope_shift( const llama_cparams & cparams, @@ -378,7 +372,6 @@ public: // get views of the current state of the cache ggml_tensor * get_k(ggml_context * ctx, int32_t il) const; ggml_tensor * get_v(ggml_context * ctx, int32_t il) const; - ggml_tensor * get_k_idx(ggml_context * ctx, int32_t il) const; // store k_cur and v_cur in the cache based on the provided head location // note: the heads in k_cur and v_cur should be laid out contiguously in memory @@ -388,7 +381,6 @@ public: // - v_idxs [n_tokens] or [n_tokens*n_embd_v_gqa] depending if V cache is transposed ggml_tensor * cpy_k(ggml_context * ctx, ggml_tensor * k_cur, ggml_tensor * k_idxs, int32_t il) const; ggml_tensor * cpy_v(ggml_context * ctx, ggml_tensor * v_cur, ggml_tensor * v_idxs, int32_t il) const; - ggml_tensor * cpy_k_idx(ggml_context * ctx, ggml_tensor * k_idx_cur, ggml_tensor * k_idxs, int32_t il) const; // create destination indices for each head of the current batch for where it would be written in the KV cache // the indices address the global KV cache (not per stream) - this is not relevant for the user of this API, but diff --git a/src/llama-model-loader.cpp b/src/llama-model-loader.cpp index 72185c8da..3177bf692 100644 --- a/src/llama-model-loader.cpp +++ b/src/llama-model-loader.cpp @@ -858,7 +858,11 @@ struct ggml_tensor * llama_model_loader::require_tensor_meta(const std::string & return tensor; } -const struct ggml_tensor * llama_model_loader::check_tensor_dims(const std::string & name, const std::vector & ne, bool required) const { +const struct ggml_tensor * llama_model_loader::check_tensor_dims( + const std::string & name, + const std::vector & ne, + bool required, + bool allow_reshape) const { const struct ggml_tensor * cur = get_tensor_meta(name.c_str()); if (cur == NULL) { @@ -868,21 +872,33 @@ const struct ggml_tensor * llama_model_loader::check_tensor_dims(const std::stri throw std::runtime_error(format("%s: tensor '%s' not found", __func__, name.c_str())); } - { - bool is_ok = true; + bool is_ok = true; + + if (allow_reshape) { + // check total number of elements only + const int64_t ncur = ggml_nelements(cur); + int64_t nexp = 1; + for (size_t i = 0; i < ne.size(); ++i) { + nexp *= ne[i]; + } + if (ncur != nexp) { + is_ok = false; + } + } else { for (size_t i = 0; i < GGML_MAX_DIMS; ++i) { if ((i < ne.size() && ne[i] != cur->ne[i]) || (i >= ne.size() && cur->ne[i] != 1)) { is_ok = false; break; } } - if (!is_ok) { - throw std::runtime_error( - format("%s: tensor '%s' has wrong shape; expected %s, got %s", - __func__, name.c_str(), - llama_format_tensor_shape(ne).c_str(), - llama_format_tensor_shape(cur).c_str())); - } + } + + if (!is_ok) { + throw std::runtime_error( + format("%s: tensor '%s' has wrong shape; expected %s, got %s", + __func__, name.c_str(), + llama_format_tensor_shape(ne).c_str(), + llama_format_tensor_shape(cur).c_str())); } return cur; @@ -1247,11 +1263,25 @@ struct ggml_tensor * llama_model_loader::create_tensor( return ret; } - ggml_tensor * t_meta = get_tensor_meta(tn.str().c_str()); - ggml_backend_buffer_type_t buft = buft_for_tensor(t_meta); - if (buft == nullptr) { - return nullptr; // return type is ggml_tensor * + LLAMA_LOG_DEBUG("%s: loading tensor %s\n", __func__, tn.str().c_str()); + const struct ggml_tensor * cur = check_tensor_dims(tn.str(), ne, !(flags & TENSOR_NOT_REQUIRED), flags & TENSOR_ALLOW_RESHAPE); + if (cur == NULL) { + return NULL; } + + ggml_tensor t_meta = *cur; + if (flags & TENSOR_ALLOW_RESHAPE) { + for (size_t dim = 0; dim < GGML_MAX_DIMS; dim++) { + t_meta.ne[dim] = dim < ne.size() ? ne.begin()[dim] : 1; + t_meta.nb[dim] = dim == 0 ? ggml_type_size(t_meta.type) : t_meta.ne[dim-1]*t_meta.nb[dim-1]; + } + } + + ggml_backend_buffer_type_t buft = buft_for_tensor(&t_meta); + if (buft == nullptr) { + return nullptr; + } + ggml_context * ctx = ctx_for_buft(buft); // if duplicated, check if the original tensor was allocated in the same buffer type context and avoid creating a new one @@ -1262,20 +1292,13 @@ struct ggml_tensor * llama_model_loader::create_tensor( } } - // LLAMA_LOG_DEBUG("%s: loading tensor %s\n", __func__, tn.str().c_str()); - const struct ggml_tensor * cur = check_tensor_dims(tn.str(), ne, !(flags & TENSOR_NOT_REQUIRED)); - - if (cur == NULL) { - return NULL; - } - const bool duplicated = flags & TENSOR_DUPLICATED; - struct ggml_tensor * tensor = ggml_dup_tensor(ctx, cur); - ggml_set_name(tensor, ggml_get_name(cur)); + struct ggml_tensor * tensor = ggml_dup_tensor(ctx, &t_meta); + ggml_set_name(tensor, ggml_get_name(&t_meta)); if (duplicated) { - size_data += ggml_nbytes(cur); + size_data += ggml_nbytes(&t_meta); } else { n_created++; } @@ -1283,34 +1306,6 @@ struct ggml_tensor * llama_model_loader::create_tensor( return tensor; } -struct ggml_tensor * llama_model_loader::create_tensor_as_view(struct ggml_context * ctx, struct ggml_tensor * base, const std::string & name, const std::initializer_list & ne, size_t offset, bool required) { - const struct ggml_tensor * cur = check_tensor_dims(name, ne, required); - - if (cur == NULL) { - return NULL; - } - - if (cur->type != base->type) { - throw std::runtime_error(format("%s: tensor '%s' has wrong type; expected %s, got %s", __func__, name.c_str(), ggml_type_name(base->type), ggml_type_name(cur->type))); - } - - std::array dims; - for (size_t i = 0; i < GGML_MAX_DIMS; ++i) { - dims[i] = i < ne.size() ? ne.begin()[i] : 1; - } - - struct ggml_tensor * tensor = ggml_view_4d(ctx, base, - dims[0], dims[1], dims[2], dims[3], - cur->nb[1], cur->nb[2], cur->nb[3], - offset); - - ggml_set_name(tensor, name.c_str()); - - n_created++; - - return tensor; -} - void llama_model_loader::done_getting_tensors(bool partial) const { if (n_created > n_tensors) { throw std::runtime_error(format("%s: too many tensors created; expected %d, got %d", __func__, n_tensors, n_created)); diff --git a/src/llama-model-loader.h b/src/llama-model-loader.h index 7ad380782..d6b31c231 100644 --- a/src/llama-model-loader.h +++ b/src/llama-model-loader.h @@ -67,6 +67,7 @@ struct llama_model_loader { static const int TENSOR_DUPLICATED = 1 << 1; static const int TENSOR_SKIP = 1 << 2; static const int TENSOR_SKIP_IF_VIRTUAL = 1 << 3; + static const int TENSOR_ALLOW_RESHAPE = 1 << 4; int n_kv = 0; int n_tensors = 0; @@ -177,14 +178,16 @@ struct llama_model_loader { struct ggml_tensor * require_tensor_meta(const std::string & name) const; - const struct ggml_tensor * check_tensor_dims(const std::string & name, const std::vector & ne, bool required) const; + const struct ggml_tensor * check_tensor_dims( + const std::string & name, + const std::vector & ne, + bool required, + bool allow_reshape) const; struct ggml_tensor * create_tensor( const llama_hparams & hparams, const buft_list_t * buft_list_cpu, const buft_list_t * buft_list_input, const buft_list_t * buft_list_output, const buft_list_t * buft_list_layer, const LLM_TN_IMPL & tn, const std::initializer_list & ne, int flags); - struct ggml_tensor * create_tensor_as_view(struct ggml_context * ctx, struct ggml_tensor * base, const std::string & name, const std::initializer_list & ne, size_t offset, bool required = true); - void done_getting_tensors(bool partial = false) const; void init_mappings(bool prefetch = true, llama_mlocks * mlock_mmaps = nullptr); diff --git a/src/llama-model.cpp b/src/llama-model.cpp index 9d5e751fb..3b1fd2c2d 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -11,6 +11,7 @@ #include "llama-kv-cache.h" #include "llama-kv-cache-iswa.h" #include "llama-kv-cache-dsa.h" +#include "llama-kv-cache-msa.h" #include "llama-kv-cache-dsv4.h" #include "llama-memory-hybrid.h" #include "llama-memory-hybrid-iswa.h" @@ -2213,6 +2214,28 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, { res = nullptr; } break; + case LLM_ARCH_MINIMAX_M3: + { + // sparse (MSA) layers carry an indexer key cache, but leading dense layers do not + llama_kv_cache::layer_filter_cb filter_idx = + [&](int32_t il) { return (uint32_t) il >= hparams.n_layer_dense_lead; }; + + res = new llama_kv_cache_msa( + *this, + params.type_k, + params.type_v, + !cparams.flash_attn, + cparams.offload_kqv, + cparams.kv_unified, + cparams.n_ctx_seq, + cparams.n_seq_max, + 1, + hparams.n_swa, + hparams.swa_type, + nullptr, + filter_idx, + nullptr); + } break; case LLM_ARCH_GLM_DSA: case LLM_ARCH_DEEPSEEK32: { @@ -2243,10 +2266,11 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, } else { // Main context: DSA cache for the trunk layers only - the nextn // layer(s) are never attended by the trunk graph. - llama_kv_cache::layer_filter_cb filter = nullptr; + llama_kv_cache::layer_filter_cb filter_mla = nullptr; if (hparams.n_layer_nextn > 0) { - filter = [&](uint32_t il) { return il < hparams.n_layer(); }; + filter_mla = [&](uint32_t il) { return il < hparams.n_layer(); }; } + llama_kv_cache::layer_filter_cb filter_lid = [&](uint32_t il) { return il < hparams.n_layer() && (arch != LLM_ARCH_GLM_DSA || hparams.is_indexer_full(il)); }; res = new llama_kv_cache_dsa( *this, @@ -2260,7 +2284,8 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, 1, hparams.n_swa, hparams.swa_type, - filter, + filter_mla, + filter_lid, nullptr); } } break; @@ -2984,7 +3009,8 @@ llama_model_base::llama_model_base(const struct llama_model_params & params) : l TENSOR_DUPLICATED (llama_model_loader::TENSOR_DUPLICATED), TENSOR_NOT_REQUIRED (llama_model_loader::TENSOR_NOT_REQUIRED), TENSOR_SKIP (llama_model_loader::TENSOR_SKIP), - TENSOR_SKIP_IF_VIRTUAL(llama_model_loader::TENSOR_SKIP_IF_VIRTUAL) {} + TENSOR_SKIP_IF_VIRTUAL(llama_model_loader::TENSOR_SKIP_IF_VIRTUAL), + TENSOR_ALLOW_RESHAPE (llama_model_loader::TENSOR_ALLOW_RESHAPE) {} ggml_tensor * llama_model_base::create_tensor(const LLM_TN_IMPL & tn, const std::initializer_list & ne, int flags) { GGML_ASSERT(ml != nullptr); diff --git a/src/llama-model.h b/src/llama-model.h index 056a6efa5..6b9e94a0a 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -719,6 +719,7 @@ struct llama_model_base : public llama_model { const int TENSOR_NOT_REQUIRED; const int TENSOR_SKIP; const int TENSOR_SKIP_IF_VIRTUAL; + const int TENSOR_ALLOW_RESHAPE; explicit llama_model_base(const llama_model_params & params); virtual ~llama_model_base() = default; diff --git a/src/llama-sampler.cpp b/src/llama-sampler.cpp index a9cb6bee5..6cf2d27cf 100644 --- a/src/llama-sampler.cpp +++ b/src/llama-sampler.cpp @@ -2638,7 +2638,8 @@ struct llama_sampler * llama_sampler_init_grammar_lazy_patterns( // penalties -struct llama_sampler_penalties { +struct llama_sampler_penalties : public llama_sampler_backend { + const int32_t n_vocab; const int32_t penalty_last_n; const float penalty_repeat; const float penalty_freq; @@ -2648,10 +2649,50 @@ struct llama_sampler_penalties { // a frequency map to count token occurrences std::unordered_map token_count; + + // backend graph inputs + ggml_tensor * inp_token_ids = nullptr; + ggml_tensor * inp_counts = nullptr; + + // backend helpers + int32_t n_max = 0; + bool has_candidates = false; + + std::vector host_token_ids; + std::vector host_counts; + + static bool is_disabled( + int32_t penalty_last_n, + float penalty_repeat, + float penalty_freq, + float penalty_present) { + return penalty_last_n == 0 || + (penalty_repeat == 1.0f && penalty_freq == 0.0f && penalty_present == 0.0f); + } + + bool is_disabled() const { + return is_disabled(penalty_last_n, penalty_repeat, penalty_freq, penalty_present); + } + + llama_sampler_penalties( + int32_t n_vocab, + int32_t penalty_last_n, + float penalty_repeat, + float penalty_freq, + float penalty_present) + : llama_sampler_backend("penalties") + , n_vocab (n_vocab) + , penalty_last_n (penalty_last_n) + , penalty_repeat (penalty_repeat) + , penalty_freq (penalty_freq) + , penalty_present (penalty_present) + , prev (penalty_last_n) { + } }; -static const char * llama_sampler_penalties_name(const struct llama_sampler * /*smpl*/) { - return "penalties"; +static const char * llama_sampler_penalties_name(const struct llama_sampler * smpl) { + auto * ctx = (llama_sampler_penalties *) smpl->ctx; + return ctx->get_name(); } static void llama_sampler_penalties_accept(struct llama_sampler * smpl, llama_token token) { @@ -2688,8 +2729,7 @@ static void llama_sampler_penalties_accept(struct llama_sampler * smpl, llama_to static void llama_sampler_penalties_apply(struct llama_sampler * smpl, llama_token_data_array * cur_p) { auto * ctx = (llama_sampler_penalties *) smpl->ctx; - if ((ctx->penalty_last_n == 0) || - (ctx->penalty_repeat == 1.0f && ctx->penalty_freq == 0.0f && ctx->penalty_present == 0.0f)) { + if (ctx->is_disabled()) { return; } @@ -2727,6 +2767,7 @@ static void llama_sampler_penalties_reset(struct llama_sampler * smpl) { static struct llama_sampler * llama_sampler_penalties_clone(const struct llama_sampler * smpl) { const auto * ctx = (const llama_sampler_penalties *) smpl->ctx; auto * result = llama_sampler_init_penalties( + ctx->n_vocab, ctx->penalty_last_n, ctx->penalty_repeat, ctx->penalty_freq, @@ -2736,7 +2777,8 @@ static struct llama_sampler * llama_sampler_penalties_clone(const struct llama_s { auto * result_ctx = (llama_sampler_penalties *) result->ctx; - result_ctx->prev = ctx->prev; + result_ctx->prev = ctx->prev; + result_ctx->token_count = ctx->token_count; } return result; @@ -2746,6 +2788,170 @@ static void llama_sampler_penalties_free(struct llama_sampler * smpl) { delete (llama_sampler_penalties *) smpl->ctx; } +static bool llama_sampler_penalties_backend_init( + struct llama_sampler * smpl, + ggml_backend_buffer_type_t buft) { + auto * sctx = (llama_sampler_penalties *) smpl->ctx; + + const bool res = llama_sampler_backend_support(smpl, buft); + + sctx->init(res); + + return res; +} + +static void llama_sampler_penalties_backend_apply( + struct llama_sampler * smpl, + struct ggml_context * ctx, + struct ggml_cgraph * gf, + struct llama_sampler_data * data) { + GGML_UNUSED(gf); + + auto * sctx = (llama_sampler_penalties *) smpl->ctx; + + if (sctx->is_disabled()) { + return; + } + + GGML_ASSERT(sctx->n_vocab > 0); + + sctx->has_candidates = data->candidates != nullptr; + sctx->n_max = std::min(sctx->penalty_last_n, sctx->n_vocab); + + sctx->inp_token_ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, sctx->n_max); + ggml_set_name(sctx->inp_token_ids, "penalties_token_ids"); + ggml_set_input(sctx->inp_token_ids); + + sctx->inp_counts = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, sctx->n_max); + ggml_set_name(sctx->inp_counts, "penalties_counts"); + ggml_set_input(sctx->inp_counts); + + if ((int32_t) sctx->host_token_ids.size() != sctx->n_max) { + sctx->host_token_ids.assign(sctx->n_max, 0); + sctx->host_counts.assign(sctx->n_max, 0); + } + + // flatten + ggml_tensor * logits = ggml_reshape_1d(ctx, data->logits, ggml_nelements(data->logits)); + ggml_tensor * gathered = logits; + ggml_tensor * counts_f32 = ggml_cast(ctx, sctx->inp_counts, GGML_TYPE_F32); + + if (sctx->has_candidates) { + ggml_tensor * candidates = ggml_reshape_1d( + ctx, data->candidates, ggml_nelements(data->candidates)); + const int64_t n_candidates = candidates->ne[0]; + GGML_ASSERT(n_candidates == ggml_nelements(logits)); + + ggml_tensor * counts_rows = ggml_fill( + ctx, ggml_new_tensor_2d(ctx, GGML_TYPE_F32, 1, sctx->n_vocab), 0.0f); + ggml_tensor * scatter_rows = ggml_reshape_2d(ctx, counts_f32, 1, sctx->n_max); + counts_rows = ggml_set_rows(ctx, counts_rows, scatter_rows, sctx->inp_token_ids); + counts_f32 = ggml_get_rows(ctx, counts_rows, candidates); + counts_f32 = ggml_reshape_1d(ctx, counts_f32, n_candidates); + } else { + ggml_tensor * logits_rows = ggml_reshape_2d(ctx, logits, 1, ggml_nelements(logits)); + gathered = ggml_get_rows(ctx, logits_rows, sctx->inp_token_ids); + gathered = ggml_reshape_1d(ctx, gathered, sctx->n_max); + } + + ggml_tensor * active_mask = ggml_step(ctx, counts_f32); + ggml_tensor * inactive_mask = ggml_sub(ctx, ggml_fill(ctx, active_mask, 1.0f), active_mask); + + ggml_tensor * penalized = gathered; + + if (sctx->penalty_repeat != 1.0f) { + ggml_tensor * pos_mask = ggml_step(ctx, penalized); + ggml_tensor * neg_mask = ggml_sub(ctx, ggml_fill(ctx, pos_mask, 1.0f), pos_mask); + + ggml_tensor * pos_scale = ggml_scale(ctx, pos_mask, 1.0f/sctx->penalty_repeat); + ggml_tensor * neg_scale = ggml_scale(ctx, neg_mask, sctx->penalty_repeat); + ggml_tensor * repeat_scale = ggml_add(ctx, pos_scale, neg_scale); + + // scale inactive entries with 1 to avoid -INF * 0 = NaN for values masked by top-p + repeat_scale = ggml_mul(ctx, repeat_scale, active_mask); + repeat_scale = ggml_add(ctx, repeat_scale, inactive_mask); + penalized = ggml_mul(ctx, gathered, repeat_scale); + } + + if (sctx->penalty_freq != 0.0f) { + ggml_tensor * penalty_freq = ggml_scale(ctx, counts_f32, sctx->penalty_freq); + penalized = ggml_sub(ctx, penalized, penalty_freq); + } + + if (sctx->penalty_present != 0.0f) { + ggml_tensor * penalty_present = ggml_scale(ctx, active_mask, sctx->penalty_present); + penalized = ggml_sub(ctx, penalized, penalty_present); + } + + if (sctx->has_candidates) { + data->logits = penalized; + } else { + ggml_tensor * logits_rows = ggml_reshape_2d(ctx, logits, 1, ggml_nelements(logits)); + ggml_tensor * scatter_rows = ggml_reshape_2d(ctx, penalized, 1, sctx->n_max); + logits_rows = ggml_set_rows(ctx, logits_rows, scatter_rows, sctx->inp_token_ids); + data->logits = ggml_reshape_1d(ctx, logits_rows, ggml_nelements(logits)); + } +} + +static void llama_sampler_penalties_backend_set_input(struct llama_sampler * smpl) { + auto * sctx = (llama_sampler_penalties *) smpl->ctx; + + if (!sctx->inp_token_ids || !sctx->inp_counts || sctx->n_max <= 0 || sctx->n_vocab <= 0) { + return; + } + + if (sctx->is_disabled()) { + return; + } + + // fill active entries from the map + int32_t n_active = 0; + + for (const auto & it : sctx->token_count) { + GGML_ASSERT(n_active < sctx->n_max); + sctx->host_token_ids[n_active] = it.first; + sctx->host_counts [n_active] = it.second; + ++n_active; + } + + // Sorting is required because backend_apply uses ggml_set_rows (a scatter-back operation) + std::vector> entries; + entries.reserve(n_active); + for (int32_t i = 0; i < n_active; ++i) { + entries.emplace_back(sctx->host_token_ids[i], sctx->host_counts[i]); + } + std::sort(entries.begin(), entries.end(), [](const auto & a, const auto & b) { + return a.first < b.first; + }); + for (int32_t i = 0; i < n_active; ++i) { + sctx->host_token_ids[i] = entries[i].first; + sctx->host_counts [i] = entries[i].second; + } + + // Padding: Finds a filler token id that is not present in token_count. + // Use it to do padding for the arrays, it avoids resizing every time. + // The arrays must always have exactly n_max entries (the GPU tensor is a fixed size). + int32_t filler = 0; + if (n_active < sctx->n_max) { + while (sctx->token_count.find(filler) != sctx->token_count.end()) { + ++filler; + } + GGML_ASSERT(filler < sctx->n_vocab); + } + + // Fill the rest of the arrays with the filler token id and count 0. + // Inactive slots are padded with a unique dummy token ID (count = 0). + // The uniqueness matters because ggml_set_rows with duplicate indices can produce non-deterministic or incorrect results. + // Using a filler token with count 0 that isn't in the active set is safe, because the active_mask step in backend_apply filters them out via ggml_step(counts_f32) + for (int32_t i = n_active; i < sctx->n_max; ++i) { + sctx->host_token_ids[i] = filler; + sctx->host_counts [i] = 0; + } + + ggml_backend_tensor_set(sctx->inp_token_ids, sctx->host_token_ids.data(), 0, sctx->n_max * sizeof(int32_t)); + ggml_backend_tensor_set(sctx->inp_counts, sctx->host_counts.data(), 0, sctx->n_max * sizeof(int32_t)); +} + static struct llama_sampler_i llama_sampler_penalties_i = { /* .name = */ llama_sampler_penalties_name, /* .accept = */ llama_sampler_penalties_accept, @@ -2753,35 +2959,33 @@ static struct llama_sampler_i llama_sampler_penalties_i = { /* .reset = */ llama_sampler_penalties_reset, /* .clone = */ llama_sampler_penalties_clone, /* .free = */ llama_sampler_penalties_free, - /* .backend_init = */ nullptr, + /* .backend_init = */ llama_sampler_penalties_backend_init, /* .backend_accept = */ nullptr, - /* .backend_apply = */ nullptr, - /* .backend_set_input = */ nullptr, + /* .backend_apply = */ llama_sampler_penalties_backend_apply, + /* .backend_set_input = */ llama_sampler_penalties_backend_set_input, }; struct llama_sampler * llama_sampler_init_penalties( + int32_t n_vocab, int32_t penalty_last_n, float penalty_repeat, float penalty_freq, float penalty_present) { penalty_last_n = std::max(penalty_last_n, 0); - const bool is_empty = (penalty_last_n == 0 || (penalty_repeat == 1.0f && penalty_freq == 0.0f && penalty_present == 0.0f)); - - if (is_empty) { + if (llama_sampler_penalties::is_disabled( + penalty_last_n, penalty_repeat, penalty_freq, penalty_present)) { return llama_sampler_init_empty("?penalties"); } return llama_sampler_init( /* .iface = */ &llama_sampler_penalties_i, - /* .ctx = */ new llama_sampler_penalties { - /* .penalty_last_n = */ penalty_last_n, - /* .penalty_repeat = */ penalty_repeat, - /* .penalty_freq = */ penalty_freq, - /* .penalty_present = */ penalty_present, - /* .prev = */ ring_buffer(penalty_last_n), - /* .token_count = */ {}, - } + /* .ctx = */ new llama_sampler_penalties( + n_vocab, + penalty_last_n, + penalty_repeat, + penalty_freq, + penalty_present) ); } diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index bbcaf9ee7..45dad6579 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -1598,8 +1598,10 @@ struct llm_tokenizer_plamo2 : llm_tokenizer { if (vocab.is_byte(token_id)) { if (entry.text.length() == 6 && entry.text.substr(0, 3) == "<0x" && entry.text.back() == '>') { std::string hex_str = entry.text.substr(3, 2); - int byte_val = std::stoi(hex_str, nullptr, 16); - bytes_[byte_val] = static_cast(token_id); + if (std::isxdigit(static_cast(hex_str[0])) && std::isxdigit(static_cast(hex_str[1]))) { + int byte_val = std::stoi(hex_str, nullptr, 16); + bytes_[byte_val] = static_cast(token_id); + } } continue; } @@ -2771,6 +2773,12 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) { const std::string & key = kv(std::get<0>(it)); int32_t & id = std::get<1>(it); + if (id >= 0 && static_cast(id) >= id_to_token.size()) { + LLAMA_LOG_WARN("%s: default special token '%s' = %d out of vocab range, disabling\n", + __func__, key.c_str(), id); + id = LLAMA_TOKEN_NULL; + } + uint32_t new_id; if (!ml.get_key(std::get<0>(it), new_id, false)) { continue; @@ -3906,12 +3914,15 @@ int32_t llama_vocab::impl::token_to_piece(llama_token token, char * buf, int32_t if (vocab.is_byte(token)) { // Handle byte tokens like <0xXX> if (token_text.length() == 6 && token_text.substr(0, 3) == "<0x" && token_text.back() == '>') { - int hex_val = std::stoi(token_text.substr(3, 2), nullptr, 16); - if (length < 1) { - return -1; + std::string hex_str = token_text.substr(3, 2); + if (std::isxdigit(static_cast(hex_str[0])) && std::isxdigit(static_cast(hex_str[1]))) { + int hex_val = std::stoi(hex_str, nullptr, 16); + if (length < 1) { + return -1; + } + buf[0] = static_cast(hex_val); + return 1; } - buf[0] = static_cast(hex_val); - return 1; } } diff --git a/src/llama.cpp b/src/llama.cpp index 386a48c58..f60971292 100644 --- a/src/llama.cpp +++ b/src/llama.cpp @@ -14,6 +14,7 @@ #include "llama-kv-cache-dsa.cpp" #include "llama-kv-cache-dsv4.cpp" #include "llama-kv-cache-iswa.cpp" +#include "llama-kv-cache-msa.cpp" #include "llama-memory-hybrid.cpp" #include "llama-memory-hybrid-iswa.cpp" #include "llama-memory-recurrent.cpp" diff --git a/src/models/deepseek2.cpp b/src/models/deepseek2.cpp index a9e8bc514..ba90c0d07 100644 --- a/src/models/deepseek2.cpp +++ b/src/models/deepseek2.cpp @@ -37,6 +37,11 @@ void llama_model_deepseek2::load_arch_hparams(llama_model_loader & ml) { hparams.rope_yarn_log_mul /= 0.1f; } + // NextN/MTP + ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.n_layer_nextn, false); + GGML_ASSERT(hparams.n_layer_nextn == 0 || + hparams.n_layer() + hparams.n_layer_nextn == hparams.n_layer_all); + // (optional) temperature tuning - used by mistral-large ml.get_key(LLM_KV_ATTENTION_TEMPERATURE_SCALE, hparams.f_attn_temp_scale, false); ml.get_key(LLM_KV_ATTENTION_TEMPERATURE_LENGTH, hparams.n_attn_temp_floor_scale, false); // FIXME why not use temperature_length? @@ -52,10 +57,20 @@ void llama_model_deepseek2::load_arch_hparams(llama_model_loader & ml) { } } -void llama_model_deepseek2::load_arch_tensors(llama_model_loader &) { +void llama_model_deepseek2::load_arch_tensors(llama_model_loader & ml) { LLAMA_LOAD_LOCALS; const int64_t n_expert_shared = hparams.n_expert_shared; + const bool mtp_only = (hparams.n_layer_nextn > 0) && (ml.get_weight("blk.0.attn_norm.weight") == nullptr); + const std::string mtp_probe = "blk." + std::to_string(n_layer) + ".nextn.eh_proj.weight"; + const bool trunk_only = (hparams.n_layer_nextn > 0) && (ml.get_weight(mtp_probe.c_str()) == nullptr); + const int trunk_flags = mtp_only ? TENSOR_NOT_REQUIRED : 0; + int mtp_flags = trunk_only ? TENSOR_NOT_REQUIRED : 0; + + if (!ml.load_mtp) { + mtp_flags |= TENSOR_SKIP; + } + const bool is_mla = hparams.is_mla(); // note: these are the actual head sizes you get when treating as MHA or after "decompression" using wv_b for MLA @@ -81,44 +96,45 @@ void llama_model_deepseek2::load_arch_tensors(llama_model_loader &) { output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED); } - for (int i = 0; i < n_layer; ++i) { + for (int i = 0; i < n_layer_all; ++i) { auto & layer = layers[i]; + const int flags = i < n_layer ? trunk_flags : mtp_flags; - layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0); + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, flags); if (q_lora_rank > 0) { - layer.attn_q_a_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_A_NORM, "weight", i), {q_lora_rank}, 0); + layer.attn_q_a_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_A_NORM, "weight", i), {q_lora_rank}, flags); } - layer.attn_kv_a_norm = create_tensor(tn(LLM_TENSOR_ATTN_KV_A_NORM, "weight", i), {kv_lora_rank}, 0); + layer.attn_kv_a_norm = create_tensor(tn(LLM_TENSOR_ATTN_KV_A_NORM, "weight", i), {kv_lora_rank}, flags); if (q_lora_rank > 0) { - layer.wq_a = create_tensor(tn(LLM_TENSOR_ATTN_Q_A, "weight", i), {n_embd, q_lora_rank}, 0); - layer.wq_b = create_tensor(tn(LLM_TENSOR_ATTN_Q_B, "weight", i), {q_lora_rank, n_head * n_embd_head_k_mla}, 0); + layer.wq_a = create_tensor(tn(LLM_TENSOR_ATTN_Q_A, "weight", i), {n_embd, q_lora_rank}, flags); + layer.wq_b = create_tensor(tn(LLM_TENSOR_ATTN_Q_B, "weight", i), {q_lora_rank, n_head * n_embd_head_k_mla}, flags); } else { - layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_head * n_embd_head_k_mla}, 0); + layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_head * n_embd_head_k_mla}, flags); } - layer.wkv_a_mqa = create_tensor(tn(LLM_TENSOR_ATTN_KV_A_MQA, "weight", i), {n_embd, kv_lora_rank + n_embd_head_qk_rope}, 0); + layer.wkv_a_mqa = create_tensor(tn(LLM_TENSOR_ATTN_KV_A_MQA, "weight", i), {n_embd, kv_lora_rank + n_embd_head_qk_rope}, flags); // note: only old legacy GGUF files will have the unsplit wkv_b tensor in if (is_mla) { - layer.wk_b = create_tensor(tn(LLM_TENSOR_ATTN_K_B, "weight", i), {n_embd_head_qk_nope, kv_lora_rank, n_head}, 0); - layer.wv_b = create_tensor(tn(LLM_TENSOR_ATTN_V_B, "weight", i), {kv_lora_rank, n_embd_head_v_mla, n_head}, 0); + layer.wk_b = create_tensor(tn(LLM_TENSOR_ATTN_K_B, "weight", i), {n_embd_head_qk_nope, kv_lora_rank, n_head}, flags); + layer.wv_b = create_tensor(tn(LLM_TENSOR_ATTN_V_B, "weight", i), {kv_lora_rank, n_embd_head_v_mla, n_head}, flags); } else { - layer.wkv_b = create_tensor(tn(LLM_TENSOR_ATTN_KV_B, "weight", i), {kv_lora_rank, n_head * (n_embd_head_qk_nope + n_embd_head_v_mla)}, 0); + layer.wkv_b = create_tensor(tn(LLM_TENSOR_ATTN_KV_B, "weight", i), {kv_lora_rank, n_head * (n_embd_head_qk_nope + n_embd_head_v_mla)}, flags); } - layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_head * n_embd_head_v_mla, n_embd}, 0); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_head * n_embd_head_v_mla, n_embd}, flags); - layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, 0); + layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, flags); if (i < (int) hparams.n_layer_dense_lead) { - layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd, n_ff}, 0); - layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), { n_ff, n_embd}, 0); - layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, 0); + layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd, n_ff}, flags); + layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), { n_ff, n_embd}, flags); + layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, flags); } else { - layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), {n_embd, n_expert}, 0); - layer.ffn_exp_probs_b = create_tensor(tn(LLM_TENSOR_FFN_EXP_PROBS_B, "bias", i), {n_expert}, TENSOR_NOT_REQUIRED); + layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), {n_embd, n_expert}, flags); + layer.ffn_exp_probs_b = create_tensor(tn(LLM_TENSOR_FFN_EXP_PROBS_B, "bias", i), {n_expert}, TENSOR_NOT_REQUIRED | flags); if (n_expert == 0) { throw std::runtime_error("n_expert must be > 0"); @@ -128,21 +144,281 @@ void llama_model_deepseek2::load_arch_tensors(llama_model_loader &) { } // MoE branch - layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), {n_ff_exp, n_embd, n_expert}, 0); - create_tensor_gate_up_exps(layer, i, n_embd, n_ff_exp, n_expert, 0); + layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), {n_ff_exp, n_embd, n_expert}, flags); + create_tensor_gate_up_exps(layer, i, n_embd, n_ff_exp, n_expert, flags); // Shared expert branch - layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, 0); - layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", i), { n_ff_exp * n_expert_shared, n_embd}, 0); - layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, 0); + layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, flags); + layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", i), { n_ff_exp * n_expert_shared, n_embd}, flags); + layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", i), {n_embd, n_ff_exp * n_expert_shared}, flags); + } + + // NextN/MTP tensors + if (i >= n_layer) { + layer.nextn.eh_proj = create_tensor(tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", i), { 2 * n_embd, n_embd }, mtp_flags); + layer.nextn.enorm = create_tensor(tn(LLM_TENSOR_NEXTN_ENORM, "weight", i), { n_embd }, mtp_flags); + layer.nextn.hnorm = create_tensor(tn(LLM_TENSOR_NEXTN_HNORM, "weight", i), { n_embd }, mtp_flags); + layer.nextn.embed_tokens = create_tensor(tn(LLM_TENSOR_NEXTN_EMBED_TOKENS, "weight", i), { n_embd, n_vocab }, TENSOR_NOT_REQUIRED | flags); + layer.nextn.shared_head_head = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "weight", i), { n_embd, n_vocab }, TENSOR_NOT_REQUIRED | flags); + layer.nextn.shared_head_norm = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "weight", i), { n_embd }, TENSOR_NOT_REQUIRED | flags); } } } std::unique_ptr llama_model_deepseek2::build_arch_graph(const llm_graph_params & params) const { + if (params.gtype == LLM_GRAPH_TYPE_DECODER_MTP) { + return std::make_unique(*this, params); + } return std::make_unique(*this, params); } +llama_model_deepseek2::graph_mtp::graph_mtp(const llama_model & model, const llm_graph_params & params) : + llm_graph_context(params) { + GGML_ASSERT(hparams.n_layer_nextn > 0 && "GLM4 MTP requires n_layer_nextn > 0"); + GGML_ASSERT(hparams.n_layer_nextn == 1 && "GLM4 MTP currently only supports a single MTP block"); + GGML_ASSERT(hparams.is_mla() && "GLM4 MTP requires MLA"); + GGML_ASSERT(hparams.f_attn_temp_scale == 0.0f && "GLM4 MTP does not support attention temperature scaling"); + + // The appended MTP block is stored immediately after the main decoder layers. + const int il = hparams.n_layer(); + const auto & layer = model.layers[il]; + + GGML_ASSERT(layer.nextn.eh_proj && "MTP block missing nextn.eh_proj"); + GGML_ASSERT(layer.nextn.enorm && "MTP block missing nextn.enorm"); + GGML_ASSERT(layer.nextn.hnorm && "MTP block missing nextn.hnorm"); + + GGML_ASSERT((uint32_t) il >= hparams.n_layer_dense_lead && "GLM4 MTP block expected to use MoE FFN"); + + const int64_t n_embd_head_k_mla = hparams.n_embd_head_k_mla(); + const int64_t n_embd_head_qk_rope = hparams.n_rot(); + const int64_t n_embd_head_qk_nope = n_embd_head_k_mla - n_embd_head_qk_rope; + const int64_t kv_lora_rank = hparams.n_lora_kv; + + GGML_ASSERT(n_embd_head_qk_nope >= 1); + GGML_ASSERT(hparams.n_lora_q > 0); + GGML_ASSERT(layer.wq_a); + GGML_ASSERT(layer.attn_q_a_norm); + GGML_ASSERT(layer.wq_b); + GGML_ASSERT(layer.wkv_a_mqa); + GGML_ASSERT(layer.attn_kv_a_norm); + GGML_ASSERT(layer.wk_b); + + const bool has_split_exps = + layer.ffn_up_exps != nullptr && + layer.ffn_gate_exps != nullptr; + + const bool has_fused_exps = layer.ffn_gate_up_exps != nullptr; + + GGML_ASSERT(has_split_exps || has_fused_exps); + GGML_ASSERT(layer.ffn_norm); + GGML_ASSERT(layer.ffn_gate_inp); + GGML_ASSERT(layer.ffn_down_exps); + GGML_ASSERT(layer.ffn_gate_shexp); + GGML_ASSERT(layer.ffn_down_shexp); + GGML_ASSERT(layer.ffn_up_shexp); + + auto inp = std::make_unique(hparams.n_embd); + + inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens); + ggml_set_input(inp->tokens); + + inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd_inp(), n_tokens); + ggml_set_input(inp->embd); + + ggml_tensor * tok_embd; + if (ubatch.token) { + ggml_tensor * tok_embd_w = layer.nextn.embed_tokens + ? layer.nextn.embed_tokens + : model.tok_embd; + + tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens); + } else { + tok_embd = inp->embd; + } + cb(tok_embd, "mtp_tok_embd", il); + + inp->h = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens); + ggml_set_input(inp->h); + ggml_set_name(inp->h, "mtp_h_input"); + + ggml_tensor * h_embd = inp->h; + + res->add_input(std::move(inp)); + + ggml_tensor * inp_pos = build_inp_pos(); + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + auto * inp_attn_k = build_attn_inp_k(); + + ggml_tensor * h_norm = build_norm(h_embd, layer.nextn.hnorm, nullptr, LLM_NORM_RMS, il); + cb(h_norm, "mtp_hnorm", il); + + ggml_tensor * e_norm = build_norm(tok_embd, layer.nextn.enorm, nullptr, LLM_NORM_RMS, il); + cb(e_norm, "mtp_enorm", il); + + ggml_tensor * concat = ggml_concat(ctx0, e_norm, h_norm, 0); + cb(concat, "mtp_concat", il); + + ggml_tensor * cur = build_lora_mm(layer.nextn.eh_proj, concat, layer.nextn.eh_proj_s); + cb(cur, "mtp_eh_proj", il); + + ggml_tensor * inpSA = cur; + + cur = build_norm(cur, layer.attn_norm, nullptr, LLM_NORM_RMS, il); + cb(cur, "mtp_attn_norm", il); + + ggml_tensor * q = ggml_mul_mat(ctx0, layer.wq_a, cur); + cb(q, "mtp_q_a", il); + + q = build_norm(q, layer.attn_q_a_norm, nullptr, LLM_NORM_RMS, il); + cb(q, "mtp_q_a_norm", il); + + q = ggml_mul_mat(ctx0, layer.wq_b, q); + cb(q, "mtp_q_b", il); + + ggml_tensor * q_nope = + ggml_view_3d(ctx0, q, n_embd_head_qk_nope, n_head, n_tokens, + ggml_row_size(q->type, n_embd_head_k_mla), + ggml_row_size(q->type, n_embd_head_k_mla) * n_head, 0); + cb(q_nope, "mtp_q_nope", il); + + ggml_tensor * q_pe = + ggml_view_3d(ctx0, q, n_embd_head_qk_rope, n_head, n_tokens, + ggml_row_size(q->type, n_embd_head_k_mla), + ggml_row_size(q->type, n_embd_head_k_mla) * n_head, + ggml_row_size(q->type, n_embd_head_qk_nope)); + cb(q_pe, "mtp_q_pe", il); + + ggml_tensor * kv_cmpr_pe = ggml_mul_mat(ctx0, layer.wkv_a_mqa, cur); + cb(kv_cmpr_pe, "mtp_kv_cmpr_pe", il); + + ggml_tensor * kv_cmpr = + ggml_view_2d(ctx0, kv_cmpr_pe, kv_lora_rank, n_tokens, + ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope), 0); + cb(kv_cmpr, "mtp_kv_cmpr", il); + + ggml_tensor * k_pe = + ggml_view_3d(ctx0, kv_cmpr_pe, n_embd_head_qk_rope, 1, n_tokens, + ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope), + ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope), + ggml_row_size(kv_cmpr_pe->type, kv_lora_rank)); + cb(k_pe, "mtp_k_pe", il); + + kv_cmpr = build_norm(kv_cmpr, layer.attn_kv_a_norm, nullptr, LLM_NORM_RMS, il); + cb(kv_cmpr, "mtp_kv_cmpr_norm", il); + + GGML_ASSERT(ext_factor >= 0.0f); + + const float attn_factor_org = + attn_factor * (1.0f + 0.1f * logf(1.0f / freq_scale)); + + const float mscale = + attn_factor_org * (1.0f + 0.1f * hparams.rope_yarn_log_mul * logf(1.0f / freq_scale)); + + const float kq_scale = + 1.0f * mscale * mscale / sqrtf(float(n_embd_head_k_mla)); + + q_pe = ggml_rope_ext(ctx0, q_pe, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + cb(q_pe, "mtp_q_pe_rope", il); + + k_pe = ggml_rope_ext(ctx0, k_pe, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + cb(k_pe, "mtp_k_pe_rope", il); + + q_nope = ggml_permute(ctx0, q_nope, 0, 2, 1, 3); + cb(q_nope, "mtp_q_nope_perm", il); + + ggml_tensor * q_nope_absorbed = ggml_mul_mat(ctx0, layer.wk_b, q_nope); + cb(q_nope_absorbed, "mtp_q_nope_absorbed", il); + + q_nope_absorbed = ggml_permute(ctx0, q_nope_absorbed, 0, 2, 1, 3); + cb(q_nope_absorbed, "mtp_q_nope_absorbed_perm", il); + + ggml_tensor * Qcur = ggml_concat(ctx0, q_nope_absorbed, q_pe, 0); + cb(Qcur, "mtp_Qcur", il); + + kv_cmpr = ggml_reshape_3d(ctx0, kv_cmpr, hparams.n_lora_kv, 1, n_tokens); + cb(kv_cmpr, "mtp_kv_cmpr_reshape", il); + + ggml_tensor * Kcur = ggml_concat(ctx0, kv_cmpr, k_pe, 0); + cb(Kcur, "mtp_Kcur", il); + + ggml_tensor * Vcur = kv_cmpr; + cb(Vcur, "mtp_Vcur", il); + + cur = build_attn(inp_attn_k, + layer.wo, nullptr, layer.wo_s, + Qcur, Kcur, Vcur, nullptr, nullptr, layer.wv_b, kq_scale, il); + cb(cur, "mtp_attn_out", il); + + ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpSA); + cb(ffn_inp, "mtp_ffn_inp", il); + + cur = build_norm(ffn_inp, layer.ffn_norm, nullptr, LLM_NORM_RMS, il); + cb(cur, "mtp_ffn_norm", il); + + ggml_tensor * moe_out = build_moe_ffn(cur, + layer.ffn_gate_inp, + layer.ffn_up_exps, + layer.ffn_gate_exps, + layer.ffn_down_exps, + layer.ffn_exp_probs_b, + n_expert, n_expert_used, + LLM_FFN_SILU, hparams.expert_weights_norm, + hparams.expert_weights_scale, + (llama_expert_gating_func_type) hparams.expert_gating_func, + il, + nullptr, + layer.ffn_gate_up_exps); + cb(moe_out, "mtp_ffn_moe_out", il); + + ggml_tensor * ffn_shexp = build_ffn(cur, + layer.ffn_up_shexp, nullptr, nullptr, + layer.ffn_gate_shexp, nullptr, nullptr, + layer.ffn_down_shexp, nullptr, nullptr, + nullptr, LLM_FFN_SILU, LLM_FFN_PAR, il); + cb(ffn_shexp, "mtp_ffn_shexp", il); + + cur = ggml_add(ctx0, moe_out, ffn_shexp); + cb(cur, "mtp_ffn_out", il); + + cur = ggml_add(ctx0, cur, ffn_inp); + cb(cur, "mtp_post_ffn", il); + + ggml_tensor * head_norm_w = layer.nextn.shared_head_norm + ? layer.nextn.shared_head_norm + : model.output_norm; + GGML_ASSERT(head_norm_w && "GLM4 MTP: missing both nextn.shared_head_norm and output_norm"); + + cur = build_norm(cur, head_norm_w, nullptr, LLM_NORM_RMS, -1); + cb(cur, "h_nextn", -1); + res->t_h_nextn = cur; + + if (inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + } + cb(cur, "mtp_shared_head_norm", -1); + + ggml_tensor * head_w = layer.nextn.shared_head_head + ? layer.nextn.shared_head_head + : model.output; + + ggml_tensor * head_s = layer.nextn.shared_head_head + ? layer.nextn.shared_head_head_s + : model.output_s; + + GGML_ASSERT(head_w && "GLM4 MTP: missing LM head (nextn.shared_head_head or model.output)"); + + cur = build_lora_mm(head_w, cur, head_s); + cb(cur, "result_output", -1); + + res->t_logits = cur; + ggml_build_forward_expand(gf, cur); +} + llama_model_deepseek2::graph::graph(const llama_model & model, const llm_graph_params & params) : llm_graph_context(params) { // lite variants include DeepSeek-V2-Lite, GigaChat3-10B-A1.8B @@ -365,7 +641,7 @@ llama_model_deepseek2::graph::graph(const llama_model & model, const llm_graph_p Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, kq_scale, il); } } - if (il == n_layer - 1 && inp_out_ids) { + if (il == n_layer - 1 && inp_out_ids && (!cparams.embeddings_nextn || cparams.embeddings_nextn_masked)) { cur = ggml_get_rows(ctx0, cur, inp_out_ids); inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids); } @@ -425,6 +701,13 @@ llama_model_deepseek2::graph::graph(const llama_model & model, const llm_graph_p cur = build_norm(cur, model.output_norm, NULL, LLM_NORM_RMS, -1); + cb(cur, "h_nextn", -1); + res->t_h_nextn = cur; + + if (cparams.embeddings_nextn && !cparams.embeddings_nextn_masked && inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + } + cb(cur, "result_norm", -1); res->t_embd = cur; diff --git a/src/models/deepseek4.cpp b/src/models/deepseek4.cpp index e68dc49b6..89cd46176 100644 --- a/src/models/deepseek4.cpp +++ b/src/models/deepseek4.cpp @@ -114,7 +114,9 @@ void llama_model_deepseek4::load_arch_tensors(llama_model_loader & ml) { layer.wq_b = create_tensor(tn(LLM_TENSOR_ATTN_Q_B, "weight", i), {q_lora_rank, n_head * n_embd_head}, flags); layer.wkv = create_tensor(tn(LLM_TENSOR_ATTN_KV, "weight", i), {n_embd, n_embd_head}, flags); layer.attn_kv_norm = create_tensor(tn(LLM_TENSOR_ATTN_KV_NORM, "weight", i), {n_embd_head}, flags); - layer.wo_a = create_tensor(tn(LLM_TENSOR_ATTN_OUT_A, "weight", i), {n_head * n_embd_head / o_groups, o_lora_rank * o_groups}, flags); + // for wo_a, the shape in the file is (n_head * n_embd_head / o_groups, o_lora_rank*o_groups) + // so we reshape here, to avoid reshaping the tensor in the graph + layer.wo_a = create_tensor(tn(LLM_TENSOR_ATTN_OUT_A, "weight", i), {n_head * n_embd_head / o_groups, o_lora_rank, o_groups}, flags | TENSOR_ALLOW_RESHAPE); layer.wo_b = create_tensor(tn(LLM_TENSOR_ATTN_OUT_B, "weight", i), {o_groups * o_lora_rank, n_embd}, flags); layer.hc_attn_fn = create_tensor(tn(LLM_TENSOR_HC_ATTN_FN, "weight", i), {hc_dim, hc_mix_dim}, flags); @@ -1258,7 +1260,7 @@ ggml_tensor * llama_model_deepseek4::graph::build_attention_impl( out = ggml_reshape_3d(ctx0, out, o_group_dim, n_groups, nt); out = ggml_permute(ctx0, out, 0, 2, 1, 3); - ggml_tensor * oa = ggml_mul_mat(ctx0, ggml_reshape_3d(ctx0, layer.wo_a, layer.wo_a->ne[0], o_lora_rank, n_groups), out); + ggml_tensor * oa = ggml_mul_mat(ctx0, layer.wo_a, out); cb(oa, "attn_wo_a", il); oa = ggml_permute(ctx0, oa, 0, 2, 1, 3); oa = ggml_cont_2d(ctx0, oa, o_lora_rank*n_groups, nt); diff --git a/src/models/dflash.cpp b/src/models/dflash.cpp index 6c82ab3da..daff6e78f 100644 --- a/src/models/dflash.cpp +++ b/src/models/dflash.cpp @@ -125,7 +125,7 @@ void llama_model_dflash::load_arch_tensors(llama_model_loader &) { layer.wq_b = create_tensor(tn(LLM_TENSOR_ATTN_Q_B, "weight", i), {q_lora_rank, n_head * n_embd_head}, 0); layer.wkv = create_tensor(tn(LLM_TENSOR_ATTN_KV, "weight", i), {n_embd, n_embd_head}, 0); layer.attn_kv_norm = create_tensor(tn(LLM_TENSOR_ATTN_KV_NORM, "weight", i), {n_embd_head}, 0); - layer.wo_a = create_tensor(tn(LLM_TENSOR_ATTN_OUT_A, "weight", i), {n_head * n_embd_head / o_groups, o_lora_rank * o_groups}, 0); + layer.wo_a = create_tensor(tn(LLM_TENSOR_ATTN_OUT_A, "weight", i), {n_head * n_embd_head / o_groups, o_lora_rank, o_groups}, TENSOR_ALLOW_RESHAPE); layer.wo_b = create_tensor(tn(LLM_TENSOR_ATTN_OUT_B, "weight", i), {o_groups * o_lora_rank, n_embd}, 0); layer.hc_attn_fn = create_tensor(tn(LLM_TENSOR_HC_ATTN_FN, "weight", i), {hc_dim, hc_mix_dim}, 0); diff --git a/src/models/minimax-m3.cpp b/src/models/minimax-m3.cpp index 0773ad543..854d5aed0 100644 --- a/src/models/minimax-m3.cpp +++ b/src/models/minimax-m3.cpp @@ -1,5 +1,5 @@ #include "models.h" -#include "llama-kv-cache.h" +#include "llama-kv-cache-msa.h" #include #include #include @@ -7,7 +7,8 @@ // MiniMax-M3: MiniMax-M2 style GQA (per-head QK-norm, partial rotary) with // DeepSeek-V3 leading-dense + routed/shared experts (sigmoid gating, routed scaling), // swigluoai activation, and MiniMax Sparse Attention (MSA). MTP is not in released model weights. -// Notes: Blocks are anchored to absolute KV cache slots. +// MSA blocks are defined over token positions. The graph translates between position space (block +// selection) and cell space (K/V/indexer storage) via per-ubatch pos<->cell maps populated from llama_kv_cells void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); @@ -23,7 +24,6 @@ void llama_model_minimax_m3::load_arch_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_ATTENTION_INDEXER_BLOCK_SIZE, hparams.indexer_block_size); ml.get_key(LLM_KV_ATTENTION_INDEXER_LOCAL_BLOCKS, hparams.indexer_local_blocks); msa_p = { (int) hparams.indexer_block_size, (int) hparams.indexer_top_k, (int) hparams.indexer_local_blocks }; - hparams.indexer_kv = true; switch (hparams.n_layer()) { case 60: type = LLM_TYPE_428B_A23B; break; @@ -86,43 +86,83 @@ std::unique_ptr llama_model_minimax_m3::build_arch_graph(cons return std::make_unique(*this, params); } -// per-query local-force bias for MSA selection -// local window always wins a slot -class llm_graph_input_msa_local : public llm_graph_input_i { +class llm_graph_input_msa : public llm_graph_input_i { public: - llm_graph_input_msa_local(int blk, int local, int64_t nblk) : blk(blk), local(local), nblk(nblk) {} + llm_graph_input_msa(const llama_kv_cache_msa_context * mctx, int blk, int local) : + mctx(mctx), blk(blk), local(local) {} void set_input(const llama_ubatch * ubatch) override { - if (!bias || !ubatch->pos) { - return; - } - const int64_t n_tokens = ubatch->n_tokens; - std::vector data((size_t) nblk * n_tokens, 0.0f); - for (int64_t i = 0; i < n_tokens; ++i) { - const int64_t L = ubatch->pos[i] / blk; - for (int l = 0; l < local && L - l >= 0; ++l) { - if (L - l < nblk) { - data[(size_t) i * nblk + (L - l)] = 1e30f; + if (pos_slot_i) { mctx->set_input_pos_slot(pos_slot_i, ubatch); } + if (pos_slot_f) { mctx->set_input_pos_slot(pos_slot_f, ubatch); } + if (cell_blk) { mctx->set_input_cell_pos(cell_blk, ubatch, blk); } + if (pos_mask) { mctx->set_input_pos_mask(pos_mask, ubatch); } + + // local-force bias over position blocks + if (bias && ubatch->pos) { + const int64_t n_tokens = ubatch->n_tokens; + const int64_t nblk = bias->ne[0]; + std::vector data((size_t) nblk * n_tokens, 0.0f); + for (int64_t i = 0; i < n_tokens; ++i) { + const int64_t L = ubatch->pos[i] / blk; + for (int l = 0; l < local && L - l >= 0; ++l) { + if (L - l < nblk) { + data[(size_t) i * nblk + (L - l)] = 1e30f; + } } } + ggml_backend_tensor_set(bias, data.data(), 0, data.size() * sizeof(float)); } - ggml_backend_tensor_set(bias, data.data(), 0, data.size() * sizeof(float)); } - // valid as long as the bias tensor dims still match the new ubatch/cache window + // valid as long as the tensor dims still match the new ubatch/cache window and the + // ubatch is in the same regime (decode graphs have pos_slot_f, batch graphs cell_blk) bool can_reuse(const llm_graph_params & params) override { - const auto * mctx = static_cast(params.mctx); + const auto * mctx_new = static_cast(params.mctx); + + this->mctx = mctx_new; + + const int64_t n_ps = GGML_PAD((int64_t) mctx_new->get_n_pos(), blk); + const int64_t ns = params.cparams.kv_unified ? 1 : params.ubatch.n_seqs_unq; + + const bool decode = params.ubatch.n_tokens == ns; // one token per stream bool res = true; - res &= bias->ne[1] == params.ubatch.n_tokens; - res &= bias->ne[0] * blk == (int64_t) mctx->get_n_kv(); + + res &= bias->ne[0] * blk == n_ps; + res &= bias->ne[1] == params.ubatch.n_tokens; + + res &= pos_mask->ne[0] == n_ps; + res &= pos_mask->ne[1] == params.ubatch.n_tokens; + + res &= pos_slot_i->ne[0] == n_ps; + res &= pos_slot_i->ne[1] == ns; + + res &= decode == (pos_slot_f != nullptr); + res &= decode == (cell_blk == nullptr); + + if (pos_slot_f) { + res &= pos_slot_f->ne[0] == n_ps; + res &= pos_slot_f->ne[1] == ns; + } + + if (cell_blk) { + res &= cell_blk->ne[0] == (int64_t) mctx_new->get_base()->get_n_kv(); + res &= cell_blk->ne[1] == ns; + } + return res; } - ggml_tensor * bias = nullptr; - int blk; - int local; - int64_t nblk; + ggml_tensor * bias = nullptr; // F32 [nblk, n_tokens] local-force bias (position blocks) + ggml_tensor * pos_mask = nullptr; // F32 [n_ps, n_tokens] 0/-inf visibility, by position + ggml_tensor * pos_slot_i = nullptr; // I32 [n_ps, ns] pos -> cell (get_rows index) + ggml_tensor * pos_slot_f = nullptr; // F32 [n_ps, ns] pos -> cell (gatherable values, decode) + ggml_tensor * cell_blk = nullptr; // I32 [n_kv, ns] cell -> position block (batch) + + const llama_kv_cache_msa_context * mctx; + + int blk; + int local; }; // One FA call for all GQA groups (and at multi-stream decode, all streams) by mapping them onto the FA sequence dim (ne[3]) @@ -173,7 +213,9 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ inpL = build_inp_embd(model.tok_embd); ggml_tensor * inp_pos = build_inp_pos(); - auto inp_attn = build_attn_inp_kv(); + + // ========================================== + // TODO: avoid such kind of complexity in the model graphs // MSA calls ggml_flash_attn_ext directly and assumes the non-transposed V layout that // llama.cpp only provides when flash attention is enabled. Block selection is anchored @@ -185,6 +227,8 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ const bool streams_ok = cparams.n_seq_max == 1 || !cparams.kv_unified; const bool msa_enabled = fa_on && streams_ok; + auto * inp_attn = build_attn_inp_kv_msa(msa_enabled); + static bool warned_no_fa = false; if (!fa_on && !warned_no_fa) { LLAMA_LOG_WARN("%s: flash attention disabled; MSA requires it -> running DENSE attention " @@ -197,36 +241,54 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ "-> running DENSE attention. Output may be degraded. Drop --kv-unified to enable MSA.\n", __func__); warned_unified = true; } + // ========================================== // hoisted per-graph MSA state (shared by every sparse layer) - llm_graph_input_msa_local * msa_loc = nullptr; + llm_graph_input_msa * msa = nullptr; ggml_tensor * msa_kqm = nullptr; - ggml_tensor * msa_mf = nullptr; - int64_t n_kv = 0, nblk = 0, ns = 1, n_tps = 0; + ggml_tensor * msa_mf = nullptr; // F32 copy of the FA mask for the final mask add + int64_t n_kv = 0, n_ps = 0, nblk = 0, ns = 1, n_tps = 0; bool msa_decode = false; // gather (1 token per stream) vs mask const int blk = mm.msa_p.blk; const int64_t Hd = hparams.indexer_n_head; // one indexer head per GQA group if (msa_enabled) { + const auto * mctx_msa = static_cast(mctx); + msa_kqm = inp_attn->get_kq_mask(); n_kv = msa_kqm->ne[0]; n_tps = msa_kqm->ne[1]; // tokens per stream ns = msa_kqm->ne[3]; // streams in this ubatch GGML_ASSERT(msa_kqm->type == GGML_TYPE_F16 && "MSA requires the FA (f16) mask"); GGML_ASSERT(n_tps*ns == n_tokens); - GGML_ASSERT(n_kv % blk == 0 && - "MSA: KV/mask n_kv must be a multiple of indexer.block_size (128); " - "the flash-attention KV padding must be a multiple of the block size. " - "A non-multiple would silently drop the partial tail block."); - nblk = n_kv / blk; + + // the position axis covers every position currently in the cache and is padded to whole blocks + n_ps = GGML_PAD((int64_t) mctx_msa->get_n_pos(), blk); + nblk = n_ps / blk; msa_decode = n_tps == 1; - msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32); + auto inp = std::make_unique(mctx_msa, blk, mm.msa_p.local); - auto loc = std::make_unique(blk, mm.msa_p.local, nblk); - loc->bias = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, nblk, n_tokens); // stream-grouped tokens - ggml_set_input(loc->bias); - msa_loc = (llm_graph_input_msa_local *) res->add_input(std::move(loc)); + inp->bias = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, nblk, n_tokens); // stream-grouped tokens + ggml_set_input(inp->bias); + + inp->pos_mask = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_ps, n_tokens); + ggml_set_input(inp->pos_mask); + + inp->pos_slot_i = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_ps, ns); + ggml_set_input(inp->pos_slot_i); + + if (msa_decode) { + inp->pos_slot_f = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_ps, ns); + ggml_set_input(inp->pos_slot_f); + } else { + inp->cell_blk = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, n_kv, ns); + ggml_set_input(inp->cell_blk); + + msa_mf = ggml_cast(ctx0, msa_kqm, GGML_TYPE_F32); + } + + msa = (llm_graph_input_msa *) res->add_input(std::move(inp)); } ggml_tensor * inp_out_ids = build_inp_out_ids(); @@ -283,9 +345,11 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ ik = ggml_rope_ext(ctx0, ik, inp_pos, nullptr, n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow); - const auto * mctx_cur = inp_attn->mctx; - ggml_build_forward_expand(gf, mctx_cur->cpy_k_idx(ctx0, ik, inp_attn->get_k_idxs(), il)); - ggml_tensor * ik_kv = mctx_cur->get_k_idx(ctx0, il); + const auto * mctx_msa_l = static_cast(mctx); + const auto * mctx_cur = mctx_msa_l->get_base(); + const auto * mctx_idx = mctx_msa_l->get_idx(); + ggml_build_forward_expand(gf, mctx_idx->cpy_k(ctx0, ik, inp_attn->get_k_idxs_idx(), il)); + ggml_tensor * ik_kv = mctx_idx->get_k(ctx0, il); if (inp_attn->self_k_rot) { Qcur = llama_mul_mat_hadamard(ctx0, Qcur, inp_attn->self_k_rot); @@ -316,42 +380,52 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ if (msa_decode) { // decode: batched over streams top-k + gather, one grouped FA - // scores: per-stream batched matmul over the stream dim (ne[3]). - // the cache views are not contiguous across streams (stride = kv_size, not n_kv) - ggml_tensor * ikv4 = ggml_view_4d(ctx0, ik_kv, n_idx_dim, n_kv, 1, ns, - ik_kv->nb[2], ik_kv->nb[3], ik_kv->nb[3], 0); + // gather the indexer keys through the pos -> cell map + ggml_tensor * ik3 = ggml_view_3d(ctx0, ik_kv, n_idx_dim, n_kv, ns, + ik_kv->nb[2], ik_kv->nb[3], 0); + ggml_tensor * ikp = ggml_get_rows(ctx0, ik3, msa->pos_slot_i); // [n_idx_dim, n_ps, ns] ggml_tensor * iq4 = ggml_reshape_4d(ctx0, iq, n_idx_dim, Hd, 1, ns); - ggml_tensor * sc = ggml_mul_mat(ctx0, ikv4, iq4); + ggml_tensor * sc = ggml_mul_mat(ctx0, + ggml_reshape_4d(ctx0, ikp, n_idx_dim, n_ps, 1, ns), iq4); ggml_mul_mat_set_prec(sc, GGML_PREC_F32); - sc = ggml_add_inplace(ctx0, sc, msa_mf); + // unmapped positions come out -inf, so they can never rank into the top-k + sc = ggml_add_inplace(ctx0, sc, + ggml_reshape_4d(ctx0, msa->pos_mask, n_ps, 1, 1, ns)); ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0); cb(bs, "msa_bs", il); ggml_tensor * bsf = ggml_add(ctx0, bs, - ggml_reshape_4d(ctx0, msa_loc->bias, nblk, 1, 1, ns)); - ggml_tensor * idx = ggml_top_k(ctx0, bsf, K); + ggml_reshape_4d(ctx0, msa->bias, nblk, 1, 1, ns)); + ggml_tensor * idx = ggml_top_k(ctx0, bsf, K); // position blocks - // token idx: tj[t,k,h,s] = blk*idx[k,h,s] + t (for the mask gather) - // row idx: tr[t,k,h,s] = tj*HKV + h (for the per-stream K/V gather) + // pos idx: tj[t,k,h,s] = blk*idx[k,h,s] + t (positions - mask gather) + // cell idx: cs[t,k,h,s] = pos_slot[tj] (pos -> cell translation) + // row idx: tr[t,k,h,s] = cs*HKV + h (per-stream K/V gather) ggml_tensor * a = ggml_scale(ctx0, ggml_cast(ctx0, idx, GGML_TYPE_F32), (float) blk); a = ggml_reshape_4d(ctx0, a, 1, K, Hd, ns); ggml_tensor * tj = ggml_add(ctx0, ggml_repeat_4d(ctx0, a, blk, K, Hd, ns), ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) blk, 1.0f), blk, 1, 1)); - ggml_tensor * tr = ggml_add(ctx0, - ggml_scale(ctx0, tj, (float) HKV), - ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) HKV, 1.0f), 1, 1, Hd)); ggml_tensor * tokj = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tj, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32); + + ggml_tensor * cs = ggml_get_rows(ctx0, + ggml_reshape_3d(ctx0, msa->pos_slot_f, 1, n_ps, ns), tokj); // [1, blk*K*Hd, ns] + cs = ggml_reshape_4d(ctx0, cs, blk, K, Hd, ns); + + ggml_tensor * tr = ggml_add(ctx0, + ggml_scale(ctx0, cs, (float) HKV), + ggml_reshape_3d(ctx0, ggml_arange(ctx0, 0.0f, (float) HKV, 1.0f), 1, 1, Hd)); + ggml_tensor * tokr = ggml_cast(ctx0, ggml_reshape_2d(ctx0, tr, (int64_t) blk*K*Hd, ns), GGML_TYPE_I32); ggml_tensor * k3 = ggml_view_3d(ctx0, k, D, HKV*n_kv, ns, k->nb[1], k->nb[3], 0); ggml_tensor * v3 = ggml_view_3d(ctx0, v, D, HKV*n_kv, ns, v->nb[1], v->nb[3], 0); - ggml_tensor * m3 = ggml_reshape_3d(ctx0, msa_kqm, 1, n_kv, ns); + ggml_tensor * mp = ggml_reshape_3d(ctx0, msa->pos_mask, 1, n_ps, ns); ggml_tensor * kg = ggml_get_rows(ctx0, k3, tokr); ggml_tensor * vg = ggml_get_rows(ctx0, v3, tokr); - ggml_tensor * mg = ggml_get_rows(ctx0, m3, tokj); + ggml_tensor * mg = ggml_get_rows(ctx0, mp, tokj); // fold (group, stream) onto the FA channel dim const ggml_type kt = ggml_is_quantized(k->type) ? GGML_TYPE_F16 : k->type; @@ -372,12 +446,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ iq->nb[1], iq->nb[2], st*n_tps*iq->nb[2]); ggml_tensor * ik_s = ggml_view_2d(ctx0, ik_kv, n_idx_dim, n_kv, ik_kv->nb[2], st*ik_kv->nb[3]); - ggml_tensor * mf_s = ggml_view_3d(ctx0, msa_mf, n_kv, 1, n_tps, - msa_mf->nb[1], msa_mf->nb[1], st*msa_mf->nb[3]); - ggml_tensor * km_s = ggml_view_3d(ctx0, msa_kqm, n_kv, n_tps, 1, - msa_kqm->nb[1], msa_kqm->nb[3], st*msa_kqm->nb[3]); - ggml_tensor * bias_s = ggml_view_3d(ctx0, msa_loc->bias, nblk, 1, n_tps, - msa_loc->bias->nb[1], msa_loc->bias->nb[1], st*n_tps*msa_loc->bias->nb[1]); + ggml_tensor * psl_s = ggml_view_1d(ctx0, msa->pos_slot_i, n_ps, + st*msa->pos_slot_i->nb[1]); + ggml_tensor * pm_s = ggml_view_3d(ctx0, msa->pos_mask, n_ps, 1, n_tps, + msa->pos_mask->nb[1], msa->pos_mask->nb[1], st*n_tps*msa->pos_mask->nb[1]); + ggml_tensor * cb_s = ggml_view_1d(ctx0, msa->cell_blk, n_kv, + st*msa->cell_blk->nb[1]); + ggml_tensor * mf_s = ggml_view_3d(ctx0, msa_mf, n_kv, n_tps, 1, + msa_mf->nb[1], msa_mf->nb[3], st*msa_mf->nb[3]); + ggml_tensor * bias_s = ggml_view_3d(ctx0, msa->bias, nblk, 1, n_tps, + msa->bias->nb[1], msa->bias->nb[1], st*n_tps*msa->bias->nb[1]); ggml_tensor * q_s = ggml_view_3d(ctx0, Qcur, D, n_head, n_tps, Qcur->nb[1], Qcur->nb[2], st*n_tps*Qcur->nb[2]); ggml_tensor * k_s = ggml_view_4d(ctx0, k, D, HKV, n_kv, 1, @@ -385,14 +463,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ ggml_tensor * v_s = ggml_view_4d(ctx0, v, D, HKV, n_kv, 1, v->nb[1], v->nb[2], v->nb[3], st*v->nb[3]); - // block scores: bs = maxpool_blk(idx_q * idx_k^T + causal mask) + // block scores: the indexer keys are gathered through the pos -> cell map first // scores are unscaled, only the top-k ordering matters - ggml_tensor * sc = ggml_mul_mat(ctx0, ik_s, + ggml_tensor * ikp = ggml_get_rows(ctx0, ik_s, psl_s); // [n_idx_dim, n_ps] + ggml_tensor * sc = ggml_mul_mat(ctx0, ikp, ggml_reshape_2d(ctx0, iq_s, n_idx_dim, Hd*n_tps)); // indexer scores run in F32 ggml_mul_mat_set_prec(sc, GGML_PREC_F32); - sc = ggml_reshape_3d(ctx0, sc, n_kv, Hd, n_tps); - sc = ggml_add_inplace(ctx0, sc, mf_s); + sc = ggml_reshape_3d(ctx0, sc, n_ps, Hd, n_tps); + // unmapped positions (holes, padding, empty cells) come out -inf + sc = ggml_add_inplace(ctx0, sc, pm_s); ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0); cb(bs, "msa_bs", il); @@ -416,14 +496,16 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ bm = ggml_cont(ctx0, ggml_permute(ctx0, bm, 0, 2, 1, 3)); // [nblk, n_tps, Hd] cb(bm, "msa_block_mask", il); - // expand block -> token granularity (j = bk*blk + t), - // then combine with the causal mask in place - ggml_tensor * bmx = ggml_repeat_4d(ctx0, - ggml_reshape_3d(ctx0, bm, 1, nblk, n_tps*Hd), - blk, nblk, n_tps*Hd, 1); + // expand block -> cell granularity through the cell -> position block + // map, then combine with the causal mask. empty cells are masked by the causal mask. + ggml_tensor * bm2 = ggml_cont(ctx0, ggml_transpose(ctx0, + ggml_reshape_2d(ctx0, bm, nblk, n_tps*Hd))); // [n_tps*Hd, nblk] + ggml_tensor * bmc = ggml_get_rows(ctx0, bm2, cb_s); // [n_tps*Hd, n_kv] F32 + ggml_tensor * bmx = ggml_cont(ctx0, ggml_transpose(ctx0, bmc)); bmx = ggml_reshape_3d(ctx0, bmx, n_kv, n_tps, Hd); - ggml_tensor * mask4 = ggml_add_inplace(ctx0, bmx, km_s); - mask4 = ggml_reshape_4d(ctx0, mask4, n_kv, n_tps, 1, Hd); + ggml_tensor * mask4 = ggml_add_inplace(ctx0, bmx, mf_s); + mask4 = ggml_cast(ctx0, + ggml_reshape_4d(ctx0, mask4, n_kv, n_tps, 1, Hd), GGML_TYPE_F16); cb(mask4, "msa_mask4", il); // cache views with groups on ne[3]; diff --git a/src/models/models.h b/src/models/models.h index 930cc3184..5f206621d 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -1084,6 +1084,10 @@ struct llama_model_deepseek2 : public llama_model_base { graph(const llama_model & model, const llm_graph_params & params); }; + struct graph_mtp : public llm_graph_context { + graph_mtp(const llama_model & model, const llm_graph_params & params); + }; + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; }; diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 4655b518e..5d2798cc1 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -1807,7 +1807,8 @@ private: // initialize samplers if (task.need_sampling()) { try { - slot.smpl.reset(common_sampler_init(model_tgt, task.params.sampling)); + slot.smpl.reset(common_sampler_init( + model_tgt, task.params.sampling, (int32_t) llama_n_ctx(ctx_tgt))); } catch (std::exception & e) { std::string err_msg = std::string("Failed to initialize samplers: ") + e.what(); send_error(task, err_msg, ERROR_TYPE_INVALID_REQUEST); diff --git a/tools/server/server-tools.cpp b/tools/server/server-tools.cpp index 4a6c5ed44..984bb478e 100644 --- a/tools/server/server-tools.cpp +++ b/tools/server/server-tools.cpp @@ -1090,6 +1090,56 @@ struct server_tool_get_datetime : server_tool { } }; +// +// get_info: returns runtime info (OS name/version and cwd) +// + +struct server_tool_get_info : server_tool { + server_tool_get_info() { + name = "get_info"; + display_name = "Get Runtime Info"; + permission_write = false; + } + + json get_definition() const override { + return { + {"type", "function"}, + {"function", { + {"name", name}, + {"description", "Returns runtime info: the OS name/version and the current working directory"}, + {"parameters", { + {"type", "object"}, + {"properties", json::object()}, + }}, + }}, + }; + } + + json invoke(json params, server_tool::stream *) const override { + auto io = make_tools_io(params); + +#ifdef _WIN32 + auto res = io->run({"cmd", "/c", "ver"}, 4096, 5); +#else + auto res = io->run({"uname", "-a"}, 4096, 5); +#endif + // "ver" prints a blank line before the version, so the output is stripped on both ends; + // a failed spawn or a timeout leaves a diagnostic in res.output, which is not an OS name + std::string os_info = res.exit_code == 0 && !res.timed_out ? string_strip(res.output) : "unknown"; + + std::string cwd = json_value(params, "cwd", std::string()); + if (cwd.empty()) { + std::error_code ec; + cwd = fs::current_path(ec).string(); + } + + return { + {"os", os_info}, + {"cwd", cwd}, + }; + } +}; + struct server_tool_stream_result : server_task_result { std::string chunk; bool done = false; @@ -1199,6 +1249,7 @@ static std::vector> build_tools() { tools.push_back(std::make_unique()); tools.push_back(std::make_unique()); tools.push_back(std::make_unique()); + tools.push_back(std::make_unique()); return tools; } diff --git a/vendor/cpp-httplib/CMakeLists.txt b/vendor/cpp-httplib/CMakeLists.txt index bf674583f..49c419094 100644 --- a/vendor/cpp-httplib/CMakeLists.txt +++ b/vendor/cpp-httplib/CMakeLists.txt @@ -41,7 +41,7 @@ if (LLAMA_BUILD_BORINGSSL) set(FIPS OFF CACHE BOOL "Enable FIPS (BoringSSL)") set(BORINGSSL_GIT "https://boringssl.googlesource.com/boringssl" CACHE STRING "BoringSSL git repository") - set(BORINGSSL_VERSION "0.20260730.0" CACHE STRING "BoringSSL version") + set(BORINGSSL_VERSION "0.20260803.0" CACHE STRING "BoringSSL version") message(STATUS "Fetching BoringSSL version ${BORINGSSL_VERSION}") diff --git a/vendor/cpp-httplib/httplib.cpp b/vendor/cpp-httplib/httplib.cpp index 155179638..30e2de896 100644 --- a/vendor/cpp-httplib/httplib.cpp +++ b/vendor/cpp-httplib/httplib.cpp @@ -1412,6 +1412,46 @@ bool stream_line_reader::getline() { #endif for (size_t i = 0;; i++) { + // Fast path: whatever the stream has already buffered can be scanned for + // the terminator in one pass. Asking for a byte at a time costs a virtual + // call, a bounds check and a one-byte copy per character of the request. + size_t buffered_size = 0; + if (auto buffered = strm_.buffered_data(buffered_size)) { + auto take = buffered_size; + auto terminated = false; + + for (size_t at = 0; at < buffered_size;) { + auto nl = static_cast( + memchr(buffered + at, '\n', buffered_size - at)); + if (!nl) { break; } + auto pos = static_cast(nl - buffered); +#ifdef CPPHTTPLIB_ALLOW_LF_AS_LINE_TERMINATOR + take = pos + 1; + terminated = true; + break; +#else + // A bare LF does not end the line; keep looking for CRLF. The CR may + // be the last byte of an earlier chunk, hence prev_byte. + if ((pos > 0 ? buffered[pos - 1] : prev_byte) == '\r') { + take = pos + 1; + terminated = true; + break; + } + at = pos + 1; +#endif + } + + if (size() + take > CPPHTTPLIB_MAX_LINE_LENGTH) { return false; } +#ifndef CPPHTTPLIB_ALLOW_LF_AS_LINE_TERMINATOR + prev_byte = buffered[take - 1]; +#endif + append(buffered, take); + strm_.consume_buffered(take); + i += take; + if (terminated) { return true; } + continue; + } + if (size() >= CPPHTTPLIB_MAX_LINE_LENGTH) { // Treat exceptionally long lines as an error to // prevent infinite loops/memory exhaustion @@ -1443,16 +1483,26 @@ bool stream_line_reader::getline() { return true; } -void stream_line_reader::append(char c) { - if (fixed_buffer_used_size_ < fixed_buffer_size_ - 1) { - fixed_buffer_[fixed_buffer_used_size_++] = c; +void stream_line_reader::append(char c) { append(&c, 1); } + +void stream_line_reader::append(const char *data, size_t size) { + // Once the line has outgrown the fixed buffer everything must keep going to + // the growable one, even if a later chunk would have fit. Without the + // emptiness check a short append after a long one would land in the fixed + // buffer, which ptr() and size() no longer look at, and be lost. + if (growable_buffer_.empty() && + fixed_buffer_used_size_ + size < fixed_buffer_size_) { + memcpy(fixed_buffer_ + fixed_buffer_used_size_, data, size); + fixed_buffer_used_size_ += size; fixed_buffer_[fixed_buffer_used_size_] = '\0'; } else { + // Unlike the per-character overload, this can be the very first append of + // the line, so the fixed buffer may hold nothing and carry no terminator + // yet. assign() takes an explicit length and does not need one. if (growable_buffer_.empty()) { - assert(fixed_buffer_[fixed_buffer_used_size_] == '\0'); growable_buffer_.assign(fixed_buffer_, fixed_buffer_used_size_); } - growable_buffer_ += c; + growable_buffer_.append(data, size); } } @@ -1525,6 +1575,14 @@ bool mmap::open(const char *path) { is_open_empty_file = true; return false; } + + if (addr_ == MAP_FAILED) { + // Clear the sentinel before `close()`, since `is_open()` only checks + // `addr_` against nullptr and `munmap()` must not be called with it. + addr_ = nullptr; + close(); + return false; + } #endif return true; @@ -1702,8 +1760,17 @@ public: socket_t socket() const override; time_t duration() const override; void set_read_timeout(time_t sec, time_t usec = 0) override; + const char *buffered_data(size_t &size) const override; + void consume_buffered(size_t size) override; + + // The caller has just seen this socket become readable. Lets the next read + // skip its own readiness wait, which would otherwise ask the kernel a + // question that was answered a moment ago. Consumed by that read. + void set_readable_hint() { readable_hint_ = true; } private: + bool ensure_readable(); + socket_t sock_; time_t read_timeout_sec_; time_t read_timeout_usec_; @@ -1715,6 +1782,7 @@ private: std::vector read_buff_; size_t read_buff_off_ = 0; size_t read_buff_content_size_ = 0; + bool readable_hint_ = false; static const size_t read_buff_size_ = 1024l * 4; }; @@ -1782,6 +1850,9 @@ process_server_socket(const std::atomic &svr_sock, socket_t sock, [&](bool close_connection, bool &connection_closed) { SocketStream strm(sock, read_timeout_sec, read_timeout_usec, write_timeout_sec, write_timeout_usec); + // process_server_socket_core() only gets here once keep_alive() has + // seen the socket go readable. + strm.set_readable_hint(); return callback(strm, close_connection, connection_closed); }); } @@ -3071,19 +3142,49 @@ bool zstd_decompressor::decompress(const char *data, size_t data_length, } #endif +bool contains_case_ignore(const std::string &s, const char *token) { + auto token_end = token + std::strlen(token); + return std::search(s.begin(), s.end(), token, token_end, [](char a, char b) { + return case_ignore::to_lower(a) == case_ignore::to_lower(b); + }) != s.end(); +} + +// Content codings are case-insensitive (RFC 9110 8.4.1). Matching them +// case-sensitively would make a response labeled e.g. "GZIP" look like an +// unknown coding, and its payload would be handed back still compressed. +bool is_zlib_encoding(const std::string &encoding) { + return case_ignore::equal(encoding, "gzip") || + case_ignore::equal(encoding, "deflate"); +} + +bool is_brotli_encoding(const std::string &encoding) { + return contains_case_ignore(encoding, "br"); +} + +bool is_zstd_encoding(const std::string &encoding) { + return contains_case_ignore(encoding, "zstd"); +} + +// Returns true if the content coding is one cpp-httplib is able to decompress +// when the corresponding support is compiled in. +bool is_known_content_encoding(const std::string &encoding) { + return is_zlib_encoding(encoding) || is_brotli_encoding(encoding) || + is_zstd_encoding(encoding); +} + std::unique_ptr create_decompressor(const std::string &encoding) { std::unique_ptr decompressor; - if (encoding == "gzip" || encoding == "deflate") { + if (is_zlib_encoding(encoding)) { #ifdef CPPHTTPLIB_ZLIB_SUPPORT decompressor = detail::make_unique(); #endif - } else if (encoding.find("br") != std::string::npos) { + } else if (is_brotli_encoding(encoding)) { #ifdef CPPHTTPLIB_BROTLI_SUPPORT decompressor = detail::make_unique(); #endif - } else if (encoding == "zstd" || encoding.find("zstd") != std::string::npos) { + } else if (is_zstd_encoding(encoding)) { #ifdef CPPHTTPLIB_ZSTD_SUPPORT decompressor = detail::make_unique(); #endif @@ -3145,8 +3246,7 @@ const char *get_header_value(const Headers &headers, size_t get_header_value_count(const Headers &headers, const std::string &key) { - auto r = headers.equal_range(key); - return static_cast(std::distance(r.first, r.second)); + return headers.count(key); } template @@ -3370,44 +3470,33 @@ ReadContentResult read_content_chunked(Stream &strm, T &x, bool is_chunked_transfer_encoding(const Headers &headers) { // RFC 9112 6.1: a message is framed with the chunked coding when "chunked" // is the final transfer coding. A single field value may list several - // codings ("gzip, chunked"), and the list may be split across multiple - // Transfer-Encoding header lines (RFC 9110 5.3). Match the last coding token - // case-insensitively rather than comparing the whole value against "chunked". + // codings ("gzip, chunked"), and RFC 9110 5.3 lets that list be split across + // several Transfer-Encoding lines, which combine into one comma-separated + // list in the order the lines were received. Headers preserves that order, + // so the final coding is the last token of the last line. Match it + // case-insensitively rather than comparing the whole value against + // "chunked". // // Security: reading a chunked message as unframed leaves its body in the // socket, where a keep-alive connection parses it as a smuggled request. - // Headers is an unordered_multimap whose iteration order for duplicate keys - // is not portable, so when there is more than one Transfer-Encoding line we - // cannot tell which coding is truly final. In that ambiguous case we fail - // safe by treating the message as chunked (a mis-parse just closes the - // connection, whereas the opposite error enables smuggling). + // Server::process_request() answers 400 and closes when the final coding is + // not chunked, so a request whose framing cannot be determined never + // reaches the "no body" path. auto rng = headers.equal_range("Transfer-Encoding"); + if (rng.first == rng.second) { return false; } - size_t line_count = 0; - bool chunked_present = false; - bool last_line_ends_with_chunked = false; + // Cleared per line, so a trailing line carrying no coding at all leaves the + // combined list ending in nothing rather than inheriting the line before it. + std::string last_coding; for (auto it = rng.first; it != rng.second; ++it) { - line_count++; const auto &value = it->second; - - std::string last_coding; - bool line_has_chunked = false; + last_coding.clear(); split(value.data(), value.data() + value.size(), ',', - [&](const char *b, const char *e) { - last_coding.assign(b, e); - if (case_ignore::equal(last_coding, "chunked")) { - line_has_chunked = true; - } - }); - - if (line_has_chunked) { chunked_present = true; } - last_line_ends_with_chunked = case_ignore::equal(last_coding, "chunked"); + [&](const char *b, const char *e) { last_coding.assign(b, e); }); } - if (line_count == 0) { return false; } - if (line_count == 1) { return last_line_ends_with_chunked; } - return chunked_present; + return case_ignore::equal(last_coding, "chunked"); } template @@ -3420,9 +3509,12 @@ bool prepare_content_receiver(T &x, int &status, std::unique_ptr decompressor; if (!encoding.empty()) { + // A coding we know about but were not built with is an error. An + // unrecognized coding (including "identity") is left alone and the + // payload is passed through as-is, since some servers misuse the header, + // e.g. by sending a character set such as "Content-Encoding: UTF-8". decompressor = detail::create_decompressor(encoding); - if (!decompressor) { - // Unsupported encoding or no support compiled in + if (!decompressor && detail::is_known_content_encoding(encoding)) { status = StatusCode::UnsupportedMediaType_415; return false; } @@ -3845,6 +3937,19 @@ std::string params_to_query_str(const Params ¶ms) { return query; } +// Splits one "key=value" span of a query string at its first '='. A span with +// no '=' at all lands entirely in key, leaving val empty, which is how a bare +// "?flag" keeps its name. +void divide_query_pair(const char *b, const char *e, std::string &key, + std::string &val) { + divide(b, static_cast(e - b), '=', + [&](const char *lhs_data, std::size_t lhs_size, const char *rhs_data, + std::size_t rhs_size) { + key.assign(lhs_data, lhs_size); + val.assign(rhs_data, rhs_size); + }); +} + void parse_query_text(const char *data, std::size_t size, Params ¶ms) { std::set cache; @@ -3855,12 +3960,7 @@ void parse_query_text(const char *data, std::size_t size, std::string key; std::string val; - divide(b, static_cast(e - b), '=', - [&](const char *lhs_data, std::size_t lhs_size, const char *rhs_data, - std::size_t rhs_size) { - key.assign(lhs_data, lhs_size); - val.assign(rhs_data, rhs_size); - }); + divide_query_pair(b, e, key, val); if (!key.empty()) { params.emplace(decode_query_component(key), decode_query_component(val)); @@ -3874,20 +3974,18 @@ void parse_query_text(const std::string &s, Params ¶ms) { // Normalize a query string by decoding and re-encoding each key/value pair // while preserving the original parameter order. This avoids double-encoding -// and ensures consistent encoding without reordering (unlike Params which -// uses std::multimap and sorts keys). +// and ensures consistent encoding. It works on the raw string rather than +// parsing into Params and re-serializing, because that round trip cannot +// reproduce the input: params_to_query_str() always emits '=', so a bare +// "flag" would come back as "flag=", and parse_query_text() drops exactly +// duplicated pairs. std::string normalize_query_string(const std::string &query) { std::string result; split(query.data(), query.data() + query.size(), '&', [&](const char *b, const char *e) { std::string key; std::string val; - divide(b, static_cast(e - b), '=', - [&](const char *lhs_data, std::size_t lhs_size, - const char *rhs_data, std::size_t rhs_size) { - key.assign(lhs_data, lhs_size); - val.assign(rhs_data, rhs_size); - }); + divide_query_pair(b, e, key, val); if (!key.empty()) { auto dec_key = decode_query_component(key); @@ -3904,6 +4002,43 @@ std::string normalize_query_string(const std::string &query) { return result; } +// Build the request target that goes on the wire from a caller-supplied path. +// Shared by the buffered send path and the streaming API so that both put the +// same bytes in the request line for the same input. +std::string encode_request_target(const std::string &target, + bool path_encode) { + // `substr(0, npos)` yields the whole string, which is what the no-query + // case needs. + auto query_pos = target.find('?'); + auto path_part = target.substr(0, query_pos); + std::string query_part; + if (query_pos != std::string::npos) { + query_part = target.substr(query_pos + 1); + } + + auto result = path_encode ? encode_path(path_part) : std::move(path_part); + + if (!query_part.empty()) { + // When path encoding is disabled the caller has supplied an already-encoded + // target and expects the exact bytes to be sent on the wire, so skip + // normalization for the query too. Normalizing would decode-then-re-encode + // it and corrupt pre-encoded binary payloads (e.g. turning `%20` into `+`, + // which a strict RFC 3986 server decodes back as `+`, not a space). + if (path_encode) { + auto normalized = normalize_query_string(query_part); + if (!normalized.empty()) { + result += '?'; + result += normalized; + } + } else { + result += '?'; + result += query_part; + } + } + + return result; +} + bool parse_multipart_boundary(const std::string &content_type, std::string &boundary) { std::map params; @@ -4969,21 +5104,8 @@ bool is_field_valid(const std::string &name, const std::string &value) { } // namespace fields -bool perform_websocket_handshake(Stream &strm, const std::string &host, - int port, bool is_ssl, - const std::string &path, - const Headers &headers, +bool perform_websocket_handshake(Stream &strm, Request &req, std::string &selected_subprotocol) { - // Validate path and host - if (!fields::is_field_value(path) || !fields::is_field_value(host)) { - return false; - } - - // Validate user-provided headers - for (const auto &h : headers) { - if (!fields::is_field_valid(h.first, h.second)) { return false; } - } - // Generate random Sec-WebSocket-Key thread_local std::mt19937 rng(std::random_device{}()); std::string key_bytes(16, '\0'); @@ -4993,19 +5115,30 @@ bool perform_websocket_handshake(Stream &strm, const std::string &host, } auto client_key = base64_encode(key_bytes); - // Build upgrade request - std::string req_str = "GET " + path + " HTTP/1.1\r\n"; - req_str += "Host: " + make_host_and_port_string(host, port, is_ssl) + "\r\n"; - req_str += "Upgrade: websocket\r\n"; - req_str += "Connection: Upgrade\r\n"; - req_str += "Sec-WebSocket-Key: " + client_key + "\r\n"; - req_str += "Sec-WebSocket-Version: 13\r\n"; - for (const auto &h : headers) { - req_str += h.first + ": " + h.second + "\r\n"; - } - req_str += "\r\n"; + req.headers.erase("Upgrade"); + req.headers.erase("Connection"); + req.headers.erase("Sec-WebSocket-Key"); + req.headers.erase("Sec-WebSocket-Version"); + req.headers.emplace("Upgrade", "websocket"); + req.headers.emplace("Connection", "Upgrade"); + req.headers.emplace("Sec-WebSocket-Key", client_key); + req.headers.emplace("Sec-WebSocket-Version", "13"); - if (strm.write(req_str.data(), req_str.size()) < 0) { return false; } + // Build the request in memory first, like ClientImpl::write_request does. + // Writing straight to the socket would leak a request line onto the wire + // before check_and_write_headers gets a chance to reject an invalid header, + // and would emit one small write per header. + BufferStream bstrm; + + if (write_request_line(bstrm, req.method, req.path) < 0) { return false; } + + auto error = Error::Success; + if (!check_and_write_headers(bstrm, req.headers, write_headers, error)) { + return false; + } + + const auto &data = bstrm.get_buffer(); + if (!write_data(strm, data.data(), data.size())) { return false; } // Verify 101 response and Sec-WebSocket-Accept header auto expected_accept = websocket_accept_key(client_key); @@ -5013,6 +5146,39 @@ bool perform_websocket_handshake(Stream &strm, const std::string &host, selected_subprotocol); } +bool is_ip_address(const std::string &host) { + struct in_addr addr4; + struct in6_addr addr6; + return inet_pton(AF_INET, host.c_str(), &addr4) == 1 || + inet_pton(AF_INET6, host.c_str(), &addr6) == 1; +} + +// Resolve where a client should connect for `host`, honoring a user-supplied +// hostname-to-address map. `host` itself is never rewritten, so it keeps +// supplying the Host header and SNI; only the connection target changes. +// +// A mapped IP literal goes to `ip`, which keeps create_socket's AI_NUMERICHOST +// path. Anything else goes to `connect_host`, which create_socket resolves as +// a name, or uses as the socket path when the address family is AF_UNIX. An +// absent or empty mapping leaves `host` as the connection target; without the +// empty check the value would reach getaddrinfo as a null node and silently +// resolve to loopback. +void apply_addr_map(const std::map &addr_map, + const std::string &host, std::string &connect_host, + std::string &ip) { + connect_host = host; + ip.clear(); + + auto it = addr_map.find(host); + if (it == addr_map.end() || it->second.empty()) { return; } + + if (is_ip_address(it->second)) { + ip = it->second; + } else { + connect_host = it->second; + } +} + } // namespace detail /* @@ -5044,7 +5210,12 @@ public: time_t duration() const override; void set_read_timeout(time_t sec, time_t usec = 0) override; + // See SocketStream::set_readable_hint(). + void set_readable_hint() { readable_hint_ = true; } + private: + bool ensure_readable(); + socket_t sock_; tls::session_t session_; time_t read_timeout_sec_; @@ -5053,6 +5224,7 @@ private: time_t write_timeout_usec_; time_t max_timeout_msec_; const std::chrono::time_point start_time_; + bool readable_hint_ = false; }; #ifdef CPPHTTPLIB_OPENSSL_SUPPORT @@ -5196,13 +5368,6 @@ std::string SHA_512(const std::string &s) { } #endif -bool is_ip_address(const std::string &host) { - struct in_addr addr4; - struct in6_addr addr6; - return inet_pton(AF_INET, host.c_str(), &addr4) == 1 || - inet_pton(AF_INET6, host.c_str(), &addr6) == 1; -} - template bool process_server_socket_ssl( const std::atomic &svr_sock, tls::session_t session, @@ -5214,6 +5379,8 @@ bool process_server_socket_ssl( [&](bool close_connection, bool &connection_closed) { SSLSocketStream strm(sock, session, read_timeout_sec, read_timeout_usec, write_timeout_sec, write_timeout_usec); + // See the non-TLS path in process_server_socket(). + strm.set_readable_hint(); return callback(strm, close_connection, connection_closed); }); } @@ -5665,6 +5832,7 @@ std::string to_string(const Error error) { case Error::UnsupportedAddressFamily: return "Unsupported address family"; case Error::HTTPParsing: return "HTTP parsing failed"; case Error::InvalidRangeHeader: return "Invalid Range header"; + case Error::UnsupportedContentEncoding: return "Unsupported Content-Encoding"; default: break; } @@ -6046,8 +6214,7 @@ std::string Request::get_trailer_value(const std::string &key, } size_t Request::get_trailer_value_count(const std::string &key) const { - auto r = trailers.equal_range(key); - return static_cast(std::distance(r.first, r.second)); + return trailers.count(key); } bool Request::has_param(const std::string &key) const { @@ -6071,8 +6238,7 @@ Request::get_param_values(const std::string &key) const { } size_t Request::get_param_value_count(const std::string &key) const { - auto r = params.equal_range(key); - return static_cast(std::distance(r.first, r.second)); + return params.count(key); } bool Request::is_multipart_form_data() const { @@ -6105,8 +6271,7 @@ bool MultipartFormData::has_field(const std::string &key) const { } size_t MultipartFormData::get_field_count(const std::string &key) const { - auto r = fields.equal_range(key); - return static_cast(std::distance(r.first, r.second)); + return fields.count(key); } FormData MultipartFormData::get_file(const std::string &key, @@ -6129,8 +6294,7 @@ bool MultipartFormData::has_file(const std::string &key) const { } size_t MultipartFormData::get_file_count(const std::string &key) const { - auto r = files.equal_range(key); - return static_cast(std::distance(r.first, r.second)); + return files.count(key); } // Multipart FormData writer implementation @@ -6209,8 +6373,7 @@ std::string Response::get_trailer_value(const std::string &key, } size_t Response::get_trailer_value_count(const std::string &key) const { - auto r = trailers.equal_range(key); - return static_cast(std::distance(r.first, r.second)); + return trailers.count(key); } void Response::set_redirect(const std::string &url, int stat) { @@ -6306,8 +6469,7 @@ std::string Result::get_request_header_value(const std::string &key, size_t Result::get_request_header_value_count(const std::string &key) const { - auto r = request_headers_.equal_range(key); - return static_cast(std::distance(r.first, r.second)); + return request_headers_.count(key); } // Stream implementation @@ -6595,6 +6757,24 @@ bool SocketStream::wait_writable() const { return select_write(sock_, write_timeout_sec_, write_timeout_usec_) > 0; } +bool SocketStream::ensure_readable() { + if (readable_hint_) { + readable_hint_ = false; + return true; + } + return wait_readable(); +} + +const char *SocketStream::buffered_data(size_t &size) const { + size = read_buff_content_size_ - read_buff_off_; + return size ? read_buff_.data() + read_buff_off_ : nullptr; +} + +void SocketStream::consume_buffered(size_t size) { + assert(size <= read_buff_content_size_ - read_buff_off_); + read_buff_off_ += size; +} + bool SocketStream::is_peer_alive() const { return detail::is_socket_alive(sock_); } @@ -6621,7 +6801,7 @@ ssize_t SocketStream::read(char *ptr, size_t size) { } } - if (!wait_readable()) { + if (!ensure_readable()) { error_ = Error::Timeout; return -1; } @@ -7099,6 +7279,14 @@ bool SSLSocketStream::wait_writable() const { !tls::is_peer_closed(session_, sock_); } +bool SSLSocketStream::ensure_readable() { + if (readable_hint_) { + readable_hint_ = false; + return true; + } + return wait_readable(); +} + bool SSLSocketStream::is_peer_alive() const { return !tls::is_peer_closed(session_, sock_); } @@ -7111,7 +7299,7 @@ ssize_t SSLSocketStream::read(char *ptr, size_t size) { error_ = Error::ConnectionClosed; } return ret; - } else if (wait_readable()) { + } else if (ensure_readable()) { tls::TlsError err; auto ret = tls::read(session_, ptr, size, err); if (ret < 0) { @@ -7533,9 +7721,11 @@ void Server::wait_until_ready() const { } void Server::stop() noexcept { - if (is_running_) { - assert(svr_sock_ != INVALID_SOCKET); - std::atomic sock(svr_sock_.exchange(INVALID_SOCKET)); + // Release the listening socket whether or not the accept loop is running: + // bind_to_port() without listen_after_bind() still owns the descriptor. The + // exchange is what makes this safe to call concurrently with the accept loop. + socket_t sock = svr_sock_.exchange(INVALID_SOCKET); + if (sock != INVALID_SOCKET) { detail::shutdown_socket(sock); detail::close_socket(sock); } @@ -7697,7 +7887,15 @@ Server::write_content_with_provider(Stream &strm, const Request &req, }; if (res.content_length_ > 0) { - if (req.ranges.empty()) { + // Only a 206 response is served as a partial representation, matching the + // condition `apply_ranges()` used to decide the Content-Length and the + // multipart boundary. Since `detail::range_error()` validates `req.ranges` + // only for a 2xx status, slicing under any other status would write a body + // that disagrees with the header already sent, from an unchecked offset. + auto is_partial = + !req.ranges.empty() && res.status == StatusCode::PartialContent_206; + + if (!is_partial) { return detail::write_content(strm, res.content_provider_, 0, res.content_length_, is_shutting_down); } else if (req.ranges.size() == 1) { @@ -8096,7 +8294,14 @@ int Server::bind_internal(const std::string &host, int port, } bool Server::listen_internal() { - if (is_decommissioned) { return false; } + // A stop() between bind and listen leaves nothing to accept on. Report + // failure instead of returning success without ever serving, and mark the + // server decommissioned the way any failed listen does so that a concurrent + // wait_until_ready() wakes up instead of spinning forever. + if (is_decommissioned || svr_sock_ == INVALID_SOCKET) { + is_decommissioned = true; + return false; + } auto ret = true; is_running_ = true; @@ -8492,11 +8697,17 @@ Server::process_request(Stream &strm, const std::string &remote_addr, return write_response(strm, close_connection, req, res); } - // RFC 9112 §6.3: Reject requests with both a non-zero Content-Length and - // any Transfer-Encoding to prevent request smuggling. Content-Length: 0 is - // tolerated for compatibility with existing clients. - if (req.get_header_value_u64("Content-Length") > 0 && - req.has_header("Transfer-Encoding")) { + // RFC 9112 §6.3: Reject requests whose framing is ambiguous, which would + // otherwise let an intermediary and this parser disagree on where the body + // ends and enable request smuggling. Two cases: a non-zero Content-Length + // alongside any Transfer-Encoding (Content-Length: 0 is tolerated for + // compatibility with existing clients), and a Transfer-Encoding whose final + // coding is not chunked, which leaves the body length undeterminable. The + // latter must not fall through to the "no body" path, or the body bytes are + // parsed as the next request on a persistent connection. + if (req.has_header("Transfer-Encoding") && + (req.get_header_value_u64("Content-Length") > 0 || + !detail::is_chunked_transfer_encoding(req.headers))) { connection_closed = true; res.status = StatusCode::BadRequest_400; return write_response(strm, close_connection, req, res); @@ -8908,13 +9119,13 @@ socket_t ClientImpl::create_client_socket(Error &error) const { write_timeout_sec_, write_timeout_usec_, interface_, error); } - // Check is custom IP specified for host_ + // Check is custom IP or hostname specified for host_ + std::string connect_host; std::string ip; - auto it = addr_map_.find(host_); - if (it != addr_map_.end()) { ip = it->second; } + detail::apply_addr_map(addr_map_, host_, connect_host, ip); return detail::create_client_socket( - host_, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_, + connect_host, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_, socket_options_, connection_timeout_sec_, connection_timeout_usec_, read_timeout_sec_, read_timeout_usec_, write_timeout_sec_, write_timeout_usec_, interface_, error); @@ -9142,11 +9353,13 @@ void ClientImpl::prepare_default_headers(Request &r, bool for_stream, if (!r.has_header(header.first)) { r.headers.insert(header); } } + // RFC 9110 5.3 recommends sending control data such as Host first, so + // prepend it rather than appending it after the caller's own fields. if (!r.has_header("Host")) { if (address_family_ == AF_UNIX) { - r.headers.emplace("Host", "localhost"); + r.headers.emplace_front("Host", "localhost"); } else { - r.headers.emplace( + r.headers.emplace_front( "Host", detail::make_host_and_port_string(host_, port_, is_ssl())); } } @@ -9197,7 +9410,12 @@ ClientImpl::open_stream(const std::string &method, const std::string &path, handle.response = detail::make_unique(); handle.error = Error::Success; - auto query_path = params.empty() ? path : append_query_params(path, params); + // Encode the target exactly like the buffered send path does, so that the + // same `path` produces the same request line through either API. + auto raw_query_path = + params.empty() ? path : append_query_params(path, params); + auto query_path = detail::encode_request_target(raw_query_path, path_encode_); + handle.connection_ = detail::make_unique(); { @@ -9311,7 +9529,20 @@ ClientImpl::open_stream(const std::string &method, const std::string &path, auto content_encoding = handle.response->get_header_value("Content-Encoding"); if (!content_encoding.empty()) { + // Same policy as prepare_content_receiver(): reject a coding we know about + // but were not built with, pass an unrecognized one through as-is. handle.decompressor_ = detail::create_decompressor(content_encoding); + if (!handle.decompressor_) { + if (detail::is_known_content_encoding(content_encoding)) { + handle.error = Error::UnsupportedContentEncoding; + handle.response.reset(); + return handle; + } + } else if (!handle.decompressor_->is_valid()) { + handle.error = Error::Compression; + handle.response.reset(); + return handle; + } } return handle; @@ -9842,52 +10073,26 @@ bool ClientImpl::write_request(Stream &strm, Request &req, { detail::BufferStream bstrm; - // Extract path and query from req.path - std::string path_part, query_part; + // Extract the query from req.path. The encoding itself is delegated to + // `encode_request_target`; the raw query is still needed here to decide + // between populating `req.params` from it and falling back to building a + // query out of caller-supplied `req.params`. auto query_pos = req.path.find('?'); - if (query_pos != std::string::npos) { - path_part = req.path.substr(0, query_pos); - query_part = req.path.substr(query_pos + 1); - } else { - path_part = req.path; - query_part = ""; - } + auto query_part = query_pos == std::string::npos + ? std::string() + : req.path.substr(query_pos + 1); - // Encode path part. If the original `req.path` already contained a - // query component, preserve its raw query string (including parameter - // order) instead of reparsing and reassembling it which may reorder - // parameters due to container ordering (e.g. `Params` uses - // `std::multimap`). When there is no query in `req.path`, fall back to - // building a query from `req.params` so existing callers that pass - // `Params` continue to work. auto path_with_query = - path_encode_ ? detail::encode_path(path_part) : path_part; + detail::encode_request_target(req.path, path_encode_); if (!query_part.empty()) { - // Normalize the query string (decode then re-encode) while preserving - // the original parameter order. When path encoding is disabled the - // caller has supplied an already-encoded target and expects the exact - // bytes to be sent on the wire, so skip normalization for the query - // too. Normalizing here would decode-then-re-encode the query and - // corrupt pre-encoded binary payloads (e.g. turning `%20` into `+`, - // which a strict RFC 3986 server decodes back as `+`, not a space). - if (path_encode_) { - auto normalized = detail::normalize_query_string(query_part); - if (!normalized.empty()) { path_with_query += '?' + normalized; } - } else { - path_with_query += '?' + query_part; - } - - // Still populate req.params for handlers/users who read them. + // The query already came in through `req.path`; still populate + // `req.params` for handlers/users who read them. detail::parse_query_text(query_part, req.params); - } else { - // No query in path; parse any query_part (empty) and append params - // from `req.params` when present (preserves prior behavior for - // callers who provide Params separately). - detail::parse_query_text(query_part, req.params); - if (!req.params.empty()) { - path_with_query = append_query_params(path_with_query, req.params); - } + } else if (!req.params.empty()) { + // No query in `req.path`; build one from `req.params` so existing + // callers that pass `Params` separately continue to work. + path_with_query = append_query_params(path_with_query, req.params); } // Write request line and headers @@ -10298,14 +10503,26 @@ bool ClientImpl::process_request(Stream &strm, Request &req, } if (res.status != StatusCode::NotModified_304) { - int dummy_status; + auto content_status = 0; auto max_length = (!has_payload_max_length_ && req.content_receiver) ? (std::numeric_limits::max)() : payload_max_length_; - if (!detail::read_content(strm, res, max_length, dummy_status, + if (!detail::read_content(strm, res, max_length, content_status, std::move(progress), std::move(out), decompress_)) { - if (error != Error::Canceled) { error = Error::Read; } + if (error != Error::Canceled) { + // Tell the caller apart from a plain read failure when the body could + // not be decoded because of its Content-Encoding. + switch (content_status) { + case StatusCode::UnsupportedMediaType_415: + error = Error::UnsupportedContentEncoding; + break; + case StatusCode::InternalServerError_500: + error = Error::Compression; + break; + default: error = Error::Read; break; + } + } output_error_log(error, &req); return false; } @@ -16769,18 +16986,42 @@ bool WebSocketClient::create_stream(std::unique_ptr &strm) { return true; } +void WebSocketClient::prepare_default_headers(Request &req) { +#ifdef CPPHTTPLIB_SSL_ENABLED + auto is_ssl = is_ssl_; +#else + auto is_ssl = false; +#endif + + if (!req.has_header("Host")) { + if (address_family_ == AF_UNIX) { + req.headers.emplace("Host", "localhost"); + } else { + req.headers.emplace( + "Host", detail::make_host_and_port_string(host_, port_, is_ssl)); + } + } + +#ifndef CPPHTTPLIB_NO_DEFAULT_USER_AGENT + if (!req.has_header("User-Agent")) { + auto agent = std::string("cpp-httplib/") + CPPHTTPLIB_VERSION; + req.set_header("User-Agent", agent); + } +#endif +} + bool WebSocketClient::connect() { if (!is_valid_) { return false; } shutdown_and_close(); - // Check is custom IP specified for host_ + // Check is custom IP or hostname specified for host_ + std::string connect_host; std::string ip; - auto it = addr_map_.find(host_); - if (it != addr_map_.end()) { ip = it->second; } + detail::apply_addr_map(addr_map_, host_, connect_host, ip); Error error; sock_ = detail::create_client_socket( - host_, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_, + connect_host, ip, port_, address_family_, tcp_nodelay_, ipv6_v6only_, socket_options_, connection_timeout_sec_, connection_timeout_usec_, read_timeout_sec_, read_timeout_usec_, write_timeout_sec_, write_timeout_usec_, interface_, error); @@ -16793,23 +17034,19 @@ bool WebSocketClient::connect() { return false; } -#ifdef CPPHTTPLIB_SSL_ENABLED - auto is_ssl = is_ssl_; -#else - auto is_ssl = false; -#endif + Request req; + req.method = "GET"; + req.path = path_; + req.headers = headers_; + prepare_default_headers(req); std::string selected_subprotocol; - if (!detail::perform_websocket_handshake(*strm, host_, port_, is_ssl, path_, - headers_, selected_subprotocol)) { + if (!detail::perform_websocket_handshake(*strm, req, selected_subprotocol)) { shutdown_and_close(); return false; } subprotocol_ = std::move(selected_subprotocol); - Request req; - req.method = "GET"; - req.path = path_; ws_ = std::unique_ptr(new WebSocket(std::move(strm), req, false, websocket_ping_interval_sec_, websocket_max_missed_pongs_)); diff --git a/vendor/cpp-httplib/httplib.h b/vendor/cpp-httplib/httplib.h index 6df357d70..94605defa 100644 --- a/vendor/cpp-httplib/httplib.h +++ b/vendor/cpp-httplib/httplib.h @@ -8,8 +8,8 @@ #ifndef CPPHTTPLIB_HTTPLIB_H #define CPPHTTPLIB_HTTPLIB_H -#define CPPHTTPLIB_VERSION "0.51.0" -#define CPPHTTPLIB_VERSION_NUM "0x003300" +#define CPPHTTPLIB_VERSION "0.52.0" +#define CPPHTTPLIB_VERSION_NUM "0x003400" #ifdef _WIN32 #if defined(_WIN32_WINNT) && _WIN32_WINNT < 0x0A00 @@ -182,7 +182,7 @@ #endif #ifndef CPPHTTPLIB_LISTEN_BACKLOG -#define CPPHTTPLIB_LISTEN_BACKLOG 5 +#define CPPHTTPLIB_LISTEN_BACKLOG 128 #endif #ifndef CPPHTTPLIB_MAX_LINE_LENGTH @@ -321,6 +321,7 @@ using socket_t = int; #include #include #include +#include #include #include #include @@ -333,9 +334,11 @@ using socket_t = int; #include #include #include +#include #include #include #include +#include // On macOS with a TLS backend, enable Keychain root certificates by default // unless the user explicitly opts out. Not enabled on iOS/tvOS/watchOS since @@ -968,11 +971,291 @@ enum StatusCode { NetworkAuthenticationRequired_511 = 511, }; -using Headers = - std::unordered_multimap; +namespace detail { -using Params = std::multimap; +// A multimap that keeps its entries in the order they were inserted. +// +// HTTP needs that order in two places. RFC 9110 5.3 makes the order of header +// fields sharing a field name significant and forbids a proxy from reordering +// them, and a query string's parameters are meaningful in the order the caller +// wrote them. Neither standard container expresses it: std::unordered_multimap +// gives no ordering guarantee at all for equivalent keys (libstdc++ yields +// reverse insertion order, libc++ insertion order), and std::multimap sorts by +// key, which would drop control data such as Host behind whatever else the +// message carries and alphabetise a query string. +// +// Entries are therefore kept in a flat vector, in order. Lookup is a linear +// scan, which beats hashing for the handful of entries a message carries +// (headers are capped at CPPHTTPLIB_HEADER_MAX_COUNT). +// +// KeyEqual compares keys; it is what makes Headers case-insensitive and +// Params, whose parameter names are case-sensitive, not. +template class insertion_ordered_multimap { +public: + using key_type = std::string; + using mapped_type = Mapped; + using value_type = std::pair; + using size_type = std::size_t; + using difference_type = std::ptrdiff_t; + using reference = value_type &; + using const_reference = const value_type &; + +private: + static size_type npos() { return static_cast(-1); } + + static bool keys_equal(const std::string &a, const std::string &b) { + return KeyEqual()(a, b); + } + + // Iterating yields every entry in insertion order, but equal_range() and + // find() have to walk only the entries sharing one key, which are not + // adjacent. Both are the same iterator type: key_idx_ selects between the + // two traversals, and since equality compares only the position, an iterator + // restricted to one key still compares equal to end(). + template class iterator_t { + public: + using iterator_category = std::bidirectional_iterator_tag; + using value_type = insertion_ordered_multimap::value_type; + using difference_type = insertion_ordered_multimap::difference_type; + using pointer = V *; + using reference = V &; + + iterator_t() : data_(nullptr), idx_(0), size_(0), key_idx_(npos()) {} + + template ::value, + int>::type = 0> + iterator_t(const iterator_t &rhs) + : data_(rhs.data_), idx_(rhs.idx_), size_(rhs.size_), + key_idx_(rhs.key_idx_) {} + + reference operator*() const { return data_[idx_]; } + pointer operator->() const { return data_ + idx_; } + + iterator_t &operator++() { + // Saturating, so that advancing past the last entry of a key (which + // get_multimap_value() does when asked for an out-of-range id) stays at + // end() instead of running off the container. + if (idx_ >= size_) { return *this; } + ++idx_; + if (key_idx_ != npos()) { + while (idx_ < size_ && !matches(idx_)) { + ++idx_; + } + } + return *this; + } + + iterator_t operator++(int) { + auto tmp = *this; + ++*this; + return tmp; + } + + iterator_t &operator--() { + if (idx_ == 0) { return *this; } + --idx_; + if (key_idx_ != npos()) { + while (idx_ > 0 && !matches(idx_)) { + --idx_; + } + } + return *this; + } + + iterator_t operator--(int) { + auto tmp = *this; + --*this; + return tmp; + } + + template bool operator==(const iterator_t &rhs) const { + return idx_ == rhs.idx_; + } + + template bool operator!=(const iterator_t &rhs) const { + return idx_ != rhs.idx_; + } + + private: + friend class insertion_ordered_multimap; + template friend class iterator_t; + + iterator_t(V *data, size_type idx, size_type size, size_type key_idx) + : data_(data), idx_(idx), size_(size), key_idx_(key_idx) {} + + bool matches(size_type i) const { + return keys_equal(data_[i].first, data_[key_idx_].first); + } + + V *data_; + size_type idx_; + size_type size_; + size_type key_idx_; + }; + +public: + using iterator = iterator_t; + using const_iterator = iterator_t; + + insertion_ordered_multimap() = default; + insertion_ordered_multimap(std::initializer_list il) + : entries_(il) {} + template + insertion_ordered_multimap(InputIt first, InputIt last) + : entries_(first, last) {} + + iterator begin() { return make_iter(0, npos()); } + iterator end() { return make_iter(entries_.size(), npos()); } + const_iterator begin() const { return make_citer(0, npos()); } + const_iterator end() const { return make_citer(entries_.size(), npos()); } + const_iterator cbegin() const { return begin(); } + const_iterator cend() const { return end(); } + + bool empty() const { return entries_.empty(); } + size_type size() const { return entries_.size(); } + void clear() { entries_.clear(); } + void swap(insertion_ordered_multimap &rhs) { entries_.swap(rhs.entries_); } + + iterator insert(const value_type &val) { + entries_.push_back(val); + return make_iter(entries_.size() - 1, npos()); + } + + iterator insert(value_type &&val) { + entries_.push_back(std::move(val)); + return make_iter(entries_.size() - 1, npos()); + } + + template iterator emplace(Args &&...args) { + entries_.emplace_back(std::forward(args)...); + return make_iter(entries_.size() - 1, npos()); + } + + // For entries that have to lead the message, such as the Host header field + // (RFC 9110 5.3 recommends sending control data first). + template iterator emplace_front(Args &&...args) { + entries_.emplace(entries_.begin(), std::forward(args)...); + return make_iter(0, npos()); + } + + iterator find(const std::string &key) { + auto i = index_of(key); + return i == npos() ? end() : make_iter(i, i); + } + + const_iterator find(const std::string &key) const { + auto i = index_of(key); + return i == npos() ? end() : make_citer(i, i); + } + + size_type count(const std::string &key) const { + size_type n = 0; + for (const auto &entry : entries_) { + if (keys_equal(entry.first, key)) { n++; } + } + return n; + } + + std::pair equal_range(const std::string &key) { + auto i = index_of(key); + return i == npos() ? std::make_pair(end(), end()) + : std::make_pair(make_iter(i, i), end()); + } + + std::pair + equal_range(const std::string &key) const { + auto i = index_of(key); + return i == npos() ? std::make_pair(end(), end()) + : std::make_pair(make_citer(i, i), end()); + } + + size_type erase(const std::string &key) { + auto before = entries_.size(); + entries_.erase(std::remove_if(entries_.begin(), entries_.end(), + [&](const value_type &entry) { + return keys_equal(entry.first, key); + }), + entries_.end()); + return before - entries_.size(); + } + + iterator erase(const_iterator pos) { + entries_.erase(entries_.begin() + static_cast(pos.idx_)); + return make_iter(pos.idx_, npos()); + } + + // Erases what iterating [first, last) would actually visit, so erasing an + // equal_range() removes only the entries with that key, not everything + // positioned between them. + iterator erase(const_iterator first, const_iterator last) { + auto from = first.idx_; + auto to = last.idx_; + if (from >= to) { return make_iter(from, npos()); } + + auto begin_it = entries_.begin(); + auto from_it = begin_it + static_cast(from); + auto to_it = begin_it + static_cast(to); + + if (first.key_idx_ == npos()) { + entries_.erase(from_it, to_it); + } else { + auto key = entries_[first.key_idx_].first; + auto keep = from_it; + for (auto it = from_it; it != to_it; ++it) { + if (!keys_equal(it->first, key)) { + if (keep != it) { *keep = std::move(*it); } + ++keep; + } + } + if (keep != to_it) { + keep = std::move(to_it, entries_.end(), keep); + } else { + keep = entries_.end(); + } + entries_.erase(keep, entries_.end()); + } + return make_iter(from, npos()); + } + + friend bool operator==(const insertion_ordered_multimap &lhs, + const insertion_ordered_multimap &rhs) { + return lhs.entries_ == rhs.entries_; + } + + friend bool operator!=(const insertion_ordered_multimap &lhs, + const insertion_ordered_multimap &rhs) { + return !(lhs == rhs); + } + +private: + size_type index_of(const std::string &key) const { + for (size_type i = 0; i < entries_.size(); i++) { + if (keys_equal(entries_[i].first, key)) { return i; } + } + return npos(); + } + + iterator make_iter(size_type idx, size_type key_idx) { + return iterator(entries_.data(), idx, entries_.size(), key_idx); + } + + const_iterator make_citer(size_type idx, size_type key_idx) const { + return const_iterator(entries_.data(), idx, entries_.size(), key_idx); + } + + std::vector entries_; +}; + +} // namespace detail + +using Headers = + detail::insertion_ordered_multimap; + +// Query parameter names are case-sensitive, unlike header field names. +using Params = + detail::insertion_ordered_multimap>; using Match = std::smatch; using DownloadProgress = std::function; @@ -1079,9 +1362,16 @@ struct FormField { std::string content; Headers headers; }; -using FormFields = std::multimap; +// RFC 7578 5.2: a form processor "SHOULD send back results in order" and +// "Intermediaries MUST NOT reorder the results", so a handler walking these +// should see the parts as they were sent. A std::multimap sorts by field name +// and loses that. Field names are case-sensitive, hence std::equal_to rather +// than the case-insensitive predicate Headers uses. +using FormFields = + detail::insertion_ordered_multimap>; -using FormFiles = std::multimap; +using FormFiles = + detail::insertion_ordered_multimap>; struct MultipartFormData { FormFields fields; // Text fields from multipart @@ -1514,6 +1804,7 @@ enum class Error { UnsupportedAddressFamily, HTTPParsing, InvalidRangeHeader, + UnsupportedContentEncoding, // For internal use only SSLPeerCouldBeClosed_, @@ -1545,6 +1836,18 @@ public: (void)usec; } + // Bytes already pulled off the socket and sitting in this stream's own + // buffer. Exposing them lets a line reader scan for a terminator in one + // pass instead of asking for a byte at a time. A stream that does no + // buffering of its own reports none, and readers fall back to read(). + virtual const char *buffered_data(size_t &size) const { + size = 0; + return nullptr; + } + + // Discards `size` bytes previously returned by buffered_data(). + virtual void consume_buffered(size_t size) { (void)size; } + ssize_t write(const char *ptr); ssize_t write(const std::string &s); @@ -2452,7 +2755,8 @@ protected: std::thread::id socket_requests_are_from_thread_ = std::thread::id(); bool socket_should_be_closed_when_request_is_done_ = false; - // Hostname-IP map + // Hostname to connection target map. The value is an IP literal or another + // hostname; only the connection target changes, never the identity. std::map addr_map_; // Default headers @@ -3154,6 +3458,10 @@ private: std::string make_host_and_port_string(const std::string &host, int port, bool is_ssl); +template +bool check_and_write_headers(Stream &strm, Headers &headers, T header_writer, + Error &error); + std::string trim_copy(const std::string &s); void divide( @@ -3364,6 +3672,7 @@ public: private: void append(char c); + void append(const char *data, size_t size); Stream &strm_; char *fixed_buffer_; @@ -3992,6 +4301,7 @@ public: private: void shutdown_and_close(); bool create_stream(std::unique_ptr &strm); + void prepare_default_headers(Request &req); std::string host_; int port_; @@ -4016,7 +4326,8 @@ private: time_t connection_timeout_usec_ = CPPHTTPLIB_CONNECTION_TIMEOUT_USECOND; std::string interface_; - // Hostname-IP map + // Hostname to connection target map. The value is an IP literal or another + // hostname; only the connection target changes, never the identity. std::map addr_map_; #ifdef CPPHTTPLIB_SSL_ENABLED