diff --git a/common/arg.cpp b/common/arg.cpp index bcd20591cd..86af0ba10a 100644 --- a/common/arg.cpp +++ b/common/arg.cpp @@ -27,6 +27,7 @@ #include #include #include +#include #include #include #include @@ -2021,7 +2022,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()); @@ -2029,14 +2036,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( @@ -2552,7 +2567,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)", @@ -3287,7 +3302,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/common.cpp b/common/common.cpp index ff27d392fb..d9ce575516 100644 --- a/common/common.cpp +++ b/common/common.cpp @@ -998,6 +998,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) { @@ -1299,8 +1316,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()) }; } @@ -1462,18 +1480,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 4bc924927f..3444aa157e 100644 --- a/common/common.h +++ b/common/common.h @@ -734,6 +734,8 @@ struct common_params { // TTS params std::string tts_lang = ""; std::string tts_speaker_file = ""; + + 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 @@ -858,6 +860,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/sampling.cpp b/common/sampling.cpp index 256ac161e2..5698c0263b 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; diff --git a/common/sampling.h b/common/sampling.h index 4191988bb8..91e2cea787 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/conversion/deepseek.py b/conversion/deepseek.py index 5b69e23437..1c9b325d5e 100644 --- a/conversion/deepseek.py +++ b/conversion/deepseek.py @@ -447,12 +447,43 @@ class DeepseekV2Model(TextModel): class DeepseekV32Model(DeepseekV2Model): model_arch = gguf.MODEL_ARCH.DEEPSEEK32 skip_mtp = False + supports_mtp_export = True + _n_main_layers: int | None = None def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.block_count = self.hparams["num_hidden_layers"] + self.hparams.get("num_nextn_predict_layers", 0) + self.block_count = self.hparams["num_hidden_layers"] + if not self.no_mtp: + self.block_count += self.hparams.get("num_nextn_predict_layers", 0) self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count) + 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: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None: + if (titem := super().filter_tensors(item)) is None: + return None + name, gen = titem + + # DeepSeek V3.2 appends the NextN/MTP block past num_hidden_layers + # (model.layers.61 -> blk.61 in the 62-block file). + assert cls._n_main_layers is not None + is_mtp = (m := re.match(r"model\.layers\.(\d+)\.", name)) is not None and int(m.group(1)) >= cls._n_main_layers + + # --no-mtp: drop the appended NextN block entirely. + if is_mtp and cls.no_mtp: + return None + # --mtp: keep ONLY NextN-block tensors plus the shared embeddings/ + # norm/lm_head (so the resulting GGUF carries just the draft head). + 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 set_vocab(self): from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained(self.dir_model) @@ -463,7 +494,7 @@ class DeepseekV32Model(DeepseekV2Model): super().set_gguf_parameters() # NextN/MTP prediction layers - if (num_nextn_predict_layers := self.hparams.get("num_nextn_predict_layers")) is not None: + if not self.no_mtp and (num_nextn_predict_layers := self.hparams.get("num_nextn_predict_layers")) is not None: self.gguf_writer.add_nextn_predict_layers(num_nextn_predict_layers) # DSA indexer parameters diff --git a/conversion/glm.py b/conversion/glm.py index cc34cddbf8..e28f54574e 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/qwen.py b/conversion/qwen.py index d1127f7431..7e3d8c0d12 100644 --- a/conversion/qwen.py +++ b/conversion/qwen.py @@ -268,8 +268,101 @@ class Qwen3MoeModel(Qwen2MoeModel): super().set_vocab() +class _QwenMtpMixin: + """Shared MTP wiring for Qwen3-Next and Qwen3.5/3.6 text variants. The HF + config carries the MTP block under `mtp_num_hidden_layers` (computed from + the checkpoint when absent, e.g. Qwen3-Next) and the tensors under + `mtp.*`; we extend block_count, emit the nextn metadata key, and remap + `mtp.*` to the standard layer-indexed nextn naming so the existing + tensor_map handles them.""" + + supports_mtp_export = True + hparams: dict[str, Any] + model_arch: gguf.MODEL_ARCH + gguf_writer: gguf.GGUFWriter + block_count: int + tensor_map: gguf.TensorNameMap + no_mtp: bool + mtp_only: bool + _original_block_count: int | None = None + opt_num_mtp_layers: int = 0 + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.block_count = self.hparams["num_hidden_layers"] + if not self.no_mtp: + n_mtp = self.hparams.get("mtp_num_hidden_layers", 0) + # Qwen-3-Next doesn't include `mtp_num_hidden_layers` in config. + if n_mtp == 0: + assert self.opt_num_mtp_layers != 0 + n_mtp = self.opt_num_mtp_layers + self.block_count += n_mtp + self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count) + + def index_tensors(self, remote_hf_model_id: str | None = None) -> dict[str, Callable[[], Tensor]]: + hparams = {**self.hparams, **self.hparams.get("text_config", {})} + key = next((k for k in ["n_layers", "num_hidden_layers", "n_layer", "num_layers"] if k in hparams), None) + type(self)._original_block_count = hparams.get(key) + type(self).opt_num_mtp_layers = 0 + return super().index_tensors(remote_hf_model_id=remote_hf_model_id) # ty: ignore[unresolved-attribute] + + @classmethod + def filter_tensors(cls, item): + assert cls._original_block_count is not None + # TODO: change TextModel to super() + if (titem := TextModel.filter_tensors(item)) is None: + return None + name, gen = titem + if name.startswith("model.mtp."): + name = name.replace("model.", "", 1) + if name.startswith("mtp."): + if cls.no_mtp: + return None + remapper = { + "fc": "eh_proj", + "pre_fc_norm_embedding": "enorm", + "pre_fc_norm_hidden": "hnorm", + "norm": "shared_head.norm", + } + parts = name.split(".", 3) + if len(parts) == 4 and parts[1] == "layers" and parts[2].isdecimal(): + mtp_idx = int(parts[2]) + name = f"model.layers.{cls._original_block_count + mtp_idx}.{parts[3]}" + cls.opt_num_mtp_layers = max(cls.opt_num_mtp_layers, mtp_idx + 1) + elif len(parts) == 3 and parts[1] in remapper: + name = f"model.layers.{cls._original_block_count}.{remapper[parts[1]]}.{parts[2]}" + elif cls.mtp_only: + keep = name in ( + "model.embed_tokens.weight", "model.norm.weight", "lm_head.weight", + "embed_tokens.weight", "norm.weight", + ) + if not keep: + return None + return name, gen + + def set_gguf_parameters(self): + super().set_gguf_parameters() # ty: ignore[unresolved-attribute] + if self.no_mtp: + return + if (n := self.block_count - self.hparams["num_hidden_layers"]) > 0: + self.gguf_writer.add_nextn_predict_layers(n) + + def prepare_metadata(self, vocab_only: bool): + from_dir = self.fname_out.is_dir() + super().prepare_metadata(vocab_only=vocab_only) # ty: ignore[unresolved-attribute] + + if not self.mtp_only or not from_dir: + return + + output_type: str = self.ftype.name.partition("_")[2] # pyright: ignore[reportAttributeAccessIssue] # ty: ignore[unresolved-attribute] + fname_default: str = gguf.naming_convention( + self.metadata.name, self.metadata.basename, self.metadata.finetune, # pyright: ignore[reportAttributeAccessIssue] # ty: ignore[unresolved-attribute] + self.metadata.version, size_label=None, output_type=output_type, model_type=None) # pyright: ignore[reportAttributeAccessIssue] # ty: ignore[unresolved-attribute] + self.fname_out = self.fname_out.parent / f"mtp-{fname_default}.gguf" + + @ModelBase.register("Qwen3NextForCausalLM") -class Qwen3NextModel(Qwen2MoeModel): +class Qwen3NextModel(_QwenMtpMixin, Qwen2MoeModel): model_arch = gguf.MODEL_ARCH.QWEN3NEXT def set_gguf_parameters(self): @@ -284,16 +377,6 @@ class Qwen3NextModel(Qwen2MoeModel): rope_dim = self.hparams["hidden_size"] // self.hparams["num_attention_heads"] self.gguf_writer.add_rope_dimension_count(int(rope_dim * self.rope_parameters.get("partial_rotary_factor", 0.25))) - @classmethod - def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None: - name, gen = item - - if name.startswith("mtp"): - # ignore MTP layers for now - return None - - return super().filter_tensors(item) - def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: if name.endswith(".A_log"): data_torch = -torch.exp(data_torch) @@ -536,97 +619,13 @@ class _Qwen35MRopeMixin: self.gguf_writer.add_rope_dimension_sections(self._QWEN35_DEFAULT_MROPE_SECTION) -class _Qwen35MtpMixin: - """Shared MTP wiring for Qwen3.5/3.6 text variants. The HF config carries - the MTP block under `mtp_num_hidden_layers` and the tensors under - `mtp.*`; we extend block_count, emit the nextn metadata key, and remap - `mtp.*` to the standard layer-indexed nextn naming so the existing - tensor_map handles them.""" - - supports_mtp_export = True - hparams: dict[str, Any] - model_arch: gguf.MODEL_ARCH - gguf_writer: gguf.GGUFWriter - block_count: int - tensor_map: gguf.TensorNameMap - no_mtp: bool - mtp_only: bool - _original_block_count: int | None = None - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - self.block_count = self.hparams["num_hidden_layers"] - if not self.no_mtp: - self.block_count += self.hparams.get("mtp_num_hidden_layers", 0) - self.tensor_map = gguf.get_tensor_name_map(self.model_arch, self.block_count) - - def index_tensors(self, remote_hf_model_id: str | None = None) -> dict[str, Callable[[], Tensor]]: - hparams = {**self.hparams, **self.hparams.get("text_config", {})} - key = next((k for k in ["n_layers", "num_hidden_layers", "n_layer", "num_layers"] if k in hparams), None) - type(self)._original_block_count = hparams.get(key) - return super().index_tensors(remote_hf_model_id=remote_hf_model_id) # ty: ignore[unresolved-attribute] - - @classmethod - def filter_tensors(cls, item): - assert cls._original_block_count is not None - # TODO: change TextModel to super() - if (titem := TextModel.filter_tensors(item)) is None: - return None - name, gen = titem - if name.startswith("model.mtp."): - name = name.replace("model.", "", 1) - if name.startswith("mtp."): - if cls.no_mtp: - return None - remapper = { - "fc": "eh_proj", - "pre_fc_norm_embedding": "enorm", - "pre_fc_norm_hidden": "hnorm", - "norm": "shared_head.norm", - } - parts = name.split(".", 3) - if len(parts) == 4 and parts[1] == "layers" and parts[2].isdecimal(): - mtp_idx = int(parts[2]) - name = f"model.layers.{cls._original_block_count + mtp_idx}.{parts[3]}" - elif len(parts) == 3 and parts[1] in remapper: - name = f"model.layers.{cls._original_block_count}.{remapper[parts[1]]}.{parts[2]}" - elif cls.mtp_only: - keep = name in ( - "model.embed_tokens.weight", "model.norm.weight", "lm_head.weight", - "embed_tokens.weight", "norm.weight", - ) - if not keep: - return None - return name, gen - - def set_gguf_parameters(self): - super().set_gguf_parameters() # ty: ignore[unresolved-attribute] - if self.no_mtp: - return - if (n := self.hparams.get("mtp_num_hidden_layers", 0)) > 0: - self.gguf_writer.add_nextn_predict_layers(n) - - def prepare_metadata(self, vocab_only: bool): - from_dir = self.fname_out.is_dir() - super().prepare_metadata(vocab_only=vocab_only) # ty: ignore[unresolved-attribute] - - if not self.mtp_only or not from_dir: - return - - output_type: str = self.ftype.name.partition("_")[2] # pyright: ignore[reportAttributeAccessIssue] # ty: ignore[unresolved-attribute] - fname_default: str = gguf.naming_convention( - self.metadata.name, self.metadata.basename, self.metadata.finetune, # pyright: ignore[reportAttributeAccessIssue] # ty: ignore[unresolved-attribute] - self.metadata.version, size_label=None, output_type=output_type, model_type=None) # pyright: ignore[reportAttributeAccessIssue] # ty: ignore[unresolved-attribute] - self.fname_out = self.fname_out.parent / f"mtp-{fname_default}.gguf" - - @ModelBase.register("Qwen3_5ForConditionalGeneration", "Qwen3_5ForCausalLM") -class Qwen3_5TextModel(_Qwen35MtpMixin, _Qwen35MRopeMixin, _LinearAttentionVReorderBase): +class Qwen3_5TextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase): model_arch = gguf.MODEL_ARCH.QWEN35 @ModelBase.register("Qwen3_5MoeForConditionalGeneration", "Qwen3_5MoeForCausalLM") -class Qwen3_5MoeTextModel(_Qwen35MtpMixin, _Qwen35MRopeMixin, _LinearAttentionVReorderBase): +class Qwen3_5MoeTextModel(_Qwen35MRopeMixin, _LinearAttentionVReorderBase): model_arch = gguf.MODEL_ARCH.QWEN35MOE diff --git a/examples/gen-docs/gen-docs.cpp b/examples/gen-docs/gen-docs.cpp index baf61bf27b..114416719c 100644 --- a/examples/gen-docs/gen-docs.cpp +++ b/examples/gen-docs/gen-docs.cpp @@ -70,6 +70,8 @@ static void write_table(std::ostringstream & ss, std::vector & opt static void write_help(std::ostringstream & ss, const md_file & md) { common_params params; + params.is_gen_docs = true; + auto ctx_arg = common_params_parser_init(params, md.ex); std::vector common_options; diff --git a/ggml/src/ggml-backend.cpp b/ggml/src/ggml-backend.cpp index 7f4e252dca..f6fb91798c 100644 --- a/ggml/src/ggml-backend.cpp +++ b/ggml/src/ggml-backend.cpp @@ -765,8 +765,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; }; @@ -805,8 +806,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; @@ -832,6 +834,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++) { @@ -1297,7 +1329,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); @@ -1313,10 +1345,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; @@ -1353,7 +1389,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; } } @@ -1373,7 +1411,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); @@ -1399,7 +1439,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; @@ -1782,6 +1826,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]); @@ -1814,7 +1861,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 33be16dc5c..d27d8acb1d 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -627,7 +627,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 09d9f3a7d6..c3758cd50c 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 285c0e9543..f320c6f004 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-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index e5ef797555..c153bd8217 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -477,6 +477,23 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max(ggml_me return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer( + ggml_metal_library_t lib, + const ggml_tensor * op) { + GGML_ASSERT(op->op == GGML_OP_LIGHTNING_INDEXER); + + char name[256]; + + snprintf(name, 256, "kernel_lightning_indexer_%s", ggml_type_name(op->src[1]->type)); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, name, name, nullptr); + } + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_metal_library_t lib, ggml_op op) { const char * name = nullptr; diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index 63e6d0c8aa..7e1deeaa21 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -125,6 +125,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_bl struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_add (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_tri (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc (ggml_metal_library_t lib, enum ggml_op op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 827d7ca320..2dc6eb8fdb 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -2,6 +2,7 @@ #import "ggml-impl.h" #import "ggml-backend-impl.h" +#import "ggml-metal-impl.h" #include @@ -1308,6 +1309,36 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te return false; } return has_simdgroup_mm; // TODO: over-restricted for vec-kernels + case GGML_OP_LIGHTNING_INDEXER: + if (op->src[0]->ne[0] != OP_LIGHTNING_INDEXER_DK || + op->src[0]->ne[1] != OP_LIGHTNING_INDEXER_NH) { + return false; + } + if (!has_simdgroup_mm || + op->src[0]->type != GGML_TYPE_F32 || + op->src[2]->type != GGML_TYPE_F32 || + op->src[3]->type != GGML_TYPE_F16 || + op->type != GGML_TYPE_F32 || + !ggml_is_contiguous_rows(op->src[0]) || + !ggml_is_contiguous_rows(op->src[1]) || + !ggml_is_contiguous_rows(op->src[2]) || + !ggml_is_contiguous_rows(op->src[3])) { + return false; + } + switch (op->src[1]->type) { + case GGML_TYPE_F32: + case GGML_TYPE_F16: + case GGML_TYPE_Q4_0: + case GGML_TYPE_Q4_1: + case GGML_TYPE_Q5_0: + case GGML_TYPE_Q5_1: + case GGML_TYPE_Q8_0: + return true; + case GGML_TYPE_BF16: + return has_bfloat; + default: + return false; + } case GGML_OP_DSV4_HC_COMB: return has_simdgroup_reduction && op->src[0]->type == GGML_TYPE_F32 && diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index 8e8926cda4..e173b91c0c 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -112,6 +112,13 @@ #define OP_FLASH_ATTN_EXT_VEC_NQPSG 1 #define OP_FLASH_ATTN_EXT_VEC_NCPSG 32 +#define OP_LIGHTNING_INDEXER_DK 128 +#define OP_LIGHTNING_INDEXER_NH 64 +#define OP_LIGHTNING_INDEXER_NHPTG 8 +#define OP_LIGHTNING_INDEXER_NKPSG 8 +#define OP_LIGHTNING_INDEXER_NSG 8 +#define OP_LIGHTNING_INDEXER_NBPTG 8 + #define OP_UNARY_NUM_SCALE 10 #define OP_UNARY_NUM_FILL 11 #define OP_UNARY_NUM_CLAMP 12 @@ -1171,6 +1178,23 @@ typedef struct { int64_t val; } ggml_metal_kargs_memset; +typedef struct { + int32_t n_kv; + int32_t n_batch; + int32_t mask_ne3; + uint64_t nb1; + uint64_t nb3; + uint64_t nbq1; + uint64_t nbq2; + uint64_t nbq3; + uint64_t nbk2; + uint64_t nbk3; + uint64_t nbw1; + uint64_t nbw3; + uint64_t nbm1; + uint64_t nbm3; +} ggml_metal_kargs_lightning_indexer; + typedef struct { int32_t n_tokens; int32_t n_iter; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 2d84804b92..c5d7619c12 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -320,6 +320,10 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) { { n_fuse = ggml_metal_op_cumsum(ctx, idx); } break; + case GGML_OP_LIGHTNING_INDEXER: + { + n_fuse = ggml_metal_op_lightning_indexer(ctx, idx); + } break; case GGML_OP_DSV4_HC_COMB: case GGML_OP_DSV4_HC_PRE: case GGML_OP_DSV4_HC_POST: @@ -1307,6 +1311,72 @@ int ggml_metal_op_diag(ggml_metal_op_t ctx, int idx) { return 1; } +int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { + ggml_tensor * op = ctx->node(idx); + + ggml_metal_encoder_t enc = ctx->enc; + + GGML_ASSERT(op->op == GGML_OP_LIGHTNING_INDEXER); + + const ggml_tensor * q = op->src[0]; + const ggml_tensor * k = op->src[1]; + const ggml_tensor * w = op->src[2]; + const ggml_tensor * m = op->src[3]; + + GGML_ASSERT(q->type == GGML_TYPE_F32); + GGML_ASSERT(k->type == GGML_TYPE_F32 || + k->type == GGML_TYPE_F16 || + k->type == GGML_TYPE_BF16 || + k->type == GGML_TYPE_Q4_0 || + k->type == GGML_TYPE_Q4_1 || + k->type == GGML_TYPE_Q5_0 || + k->type == GGML_TYPE_Q5_1 || + k->type == GGML_TYPE_Q8_0); + GGML_ASSERT(w->type == GGML_TYPE_F32); + GGML_ASSERT(m->type == GGML_TYPE_F16); + GGML_ASSERT(op->type == GGML_TYPE_F32); + + GGML_ASSERT(q->ne[0] == OP_LIGHTNING_INDEXER_DK); + GGML_ASSERT(q->ne[1] == OP_LIGHTNING_INDEXER_NH); + + ggml_metal_kargs_lightning_indexer args = { + /*.n_kv =*/ (int32_t) k->ne[2], + /*.n_batch =*/ (int32_t) q->ne[2], + /*.mask_ne3 =*/ (int32_t) m->ne[3], + /*.nb1 =*/ op->nb[1], + /*.nb3 =*/ op->nb[3], + /*.nbq1 =*/ q->nb[1], + /*.nbq2 =*/ q->nb[2], + /*.nbq3 =*/ q->nb[3], + /*.nbk2 =*/ k->nb[2], + /*.nbk3 =*/ k->nb[3], + /*.nbw1 =*/ w->nb[1], + /*.nbw3 =*/ w->nb[3], + /*.nbm1 =*/ m->nb[1], + /*.nbm3 =*/ m->nb[3], + }; + + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(q), 1); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(k), 2); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(w), 3); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(m), 4); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5); + + const int nsg = OP_LIGHTNING_INDEXER_NSG; + const int nkptg = OP_LIGHTNING_INDEXER_NKPSG*nsg; + const int nbptg = OP_LIGHTNING_INDEXER_NBPTG; + + auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, op); + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_dispatch_threadgroups(enc, + (k->ne[2] + nkptg - 1)/nkptg, + (q->ne[2] + nbptg - 1)/nbptg, + q->ne[3], 32, nsg, 1); + + return 1; +} + int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) { ggml_tensor * op = ctx->node(idx); diff --git a/ggml/src/ggml-metal/ggml-metal-ops.h b/ggml/src/ggml-metal/ggml-metal-ops.h index 4fdad1dd59..b03b59e0bd 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.h +++ b/ggml/src/ggml-metal/ggml-metal-ops.h @@ -54,6 +54,7 @@ int ggml_metal_op_cumsum (ggml_metal_op_t ctx, int idx); int ggml_metal_op_get_rows (ggml_metal_op_t ctx, int idx); int ggml_metal_op_set_rows (ggml_metal_op_t ctx, int idx); int ggml_metal_op_diag (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_lightning_indexer (ggml_metal_op_t ctx, int idx); int ggml_metal_op_dsv4_hc (ggml_metal_op_t ctx, int idx); int ggml_metal_op_soft_max (ggml_metal_op_t ctx, int idx); int ggml_metal_op_ssm_conv (ggml_metal_op_t ctx, int idx); diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal index 47d4f685a0..7d12cb0fe3 100644 --- a/ggml/src/ggml-metal/ggml-metal.metal +++ b/ggml/src/ggml-metal/ggml-metal.metal @@ -11295,6 +11295,155 @@ typedef decltype(kernel_count_equal) kernel_count_equal_t; template [[host_name("kernel_count_equal_i32")]] kernel kernel_count_equal_t kernel_count_equal; +template< + typename kd4x4_t, + short nl_k, + void (*deq_k)(device const kd4x4_t *, short, thread half4x4 &)> +kernel void kernel_lightning_indexer( + constant ggml_metal_kargs_lightning_indexer & args, + device const char * q, + device const char * k, + device const char * w, + device const char * m, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiitg[[thread_index_in_threadgroup]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]]) { + constexpr short DK = OP_LIGHTNING_INDEXER_DK; + constexpr short NH = OP_LIGHTNING_INDEXER_NH; + constexpr short NHPTG = OP_LIGHTNING_INDEXER_NHPTG; + constexpr short NKPSG = OP_LIGHTNING_INDEXER_NKPSG; + constexpr short NSG = OP_LIGHTNING_INDEXER_NSG; + constexpr short NBPTG = OP_LIGHTNING_INDEXER_NBPTG; + + constexpr short DK4 = DK/4; + constexpr short DK8 = DK/8; + constexpr short DK16 = DK/16; + + constexpr short NK = NKPSG*NSG; // keys per threadgroup + constexpr short NTG = 32*NSG; // threads per threadgroup + + const int i_stream = tgpig.z; + const int i_kv_0 = tgpig.x*NK; // first key of this threadgroup + const int i_kv = i_kv_0 + sgitg*NKPSG; // first key of this simdgroup + + threadgroup half4x4 sk4x4[NK*DK16]; + threadgroup half * sk = (threadgroup half *) sk4x4; + + for (short i = tiitg; i < NK*DK16; i += NTG) { + const short ik = i/DK16; + const short i16 = i%DK16; + + half4x4 tmp; + + if (i_kv_0 + ik < args.n_kv) { + device const kd4x4_t * kr = (device const kd4x4_t *) (k + (i_kv_0 + ik)*args.nbk2 + i_stream*args.nbk3); + + deq_k(kr + i16/nl_k, i16%nl_k, tmp); + } else { + FOR_UNROLL (short j = 0; j < 4; ++j) { + tmp[j] = half4(0.0h); + } + } + + sk4x4[i] = tmp; + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + // K tile of this simdgroup, transposed to [DK, NKPSG] + simdgroup_half8x8 mk[DK8]; + + FOR_UNROLL (short i = 0; i < DK8; ++i) { + simdgroup_load(mk[i], sk + sgitg*NKPSG*DK + 8*i, DK, 0, true); + } + + threadgroup half4 sq4[NHPTG*DK4]; + threadgroup half * sq = (threadgroup half *) sq4; + + threadgroup float sw [NHPTG]; + threadgroup float sqk[NSG*NHPTG*NKPSG]; + + const int i_batch_0 = tgpig.y*NBPTG; + const int n_batch = min((int) NBPTG, args.n_batch - i_batch_0); + + for (short ib = 0; ib < n_batch; ++ib) { + const int i_batch = i_batch_0 + ib; + + device const char * pq = q + i_batch*args.nbq2 + i_stream*args.nbq3; + device const char * pw = w + i_batch*args.nbw1 + i_stream*args.nbw3; + + float score = 0.0f; + + FOR_UNROLL (short i_head = 0; i_head < NH; i_head += NHPTG) { + // stage the Q tile [DK, NHPTG] and the (prescaled) head weights + for (short i = tiitg; i < NHPTG*DK4; i += NTG) { + const short ih = i/DK4; + const short i4 = i%DK4; + + device const float4 * q4 = (device const float4 *) (pq + (i_head + ih)*args.nbq1); + + sq4[ih*DK4 + i4] = half4(q4[i4]); + } + + if (tiitg < NHPTG) { + sw[tiitg] = ((device const float *) pw)[i_head + tiitg]; + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + simdgroup_float8x8 mqk = make_filled_simdgroup_matrix(0.0f); + + FOR_UNROLL (short i = 0; i < DK8; ++i) { + simdgroup_half8x8 mq; + + simdgroup_load(mq, sq + 8*i, DK, 0, false); + simdgroup_multiply_accumulate(mqk, mq, mk[i], mqk); + } + + threadgroup float * pqk = sqk + sgitg*NHPTG*NKPSG; + + simdgroup_store(mqk, pqk, NKPSG, 0, false); + simdgroup_barrier(mem_flags::mem_threadgroup); + + // one lane per key: ReLU, apply the head weight and accumulate over the head tile + if (tiisg < NKPSG) { + FOR_UNROLL (short ih = 0; ih < NHPTG; ++ih) { + score += max(pqk[ih*NKPSG + tiisg], 0.0f)*sw[ih]; + } + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + if (tiisg < NKPSG) { + const int ik = i_kv + tiisg; + if (ik < args.n_kv) { + device const half * pm = (device const half *) (m + i_batch*args.nbm1 + (i_stream % args.mask_ne3)*args.nbm3); + device float * pd = (device float *) (dst + i_batch*args.nb1 + i_stream*args.nb3); + + pd[ik] = score + (float) pm[ik]; + } + } + } +} + +typedef decltype(kernel_lightning_indexer) kernel_lightning_indexer_t; + +template [[host_name("kernel_lightning_indexer_f32")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_f16")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; + +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_lightning_indexer_bf16")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +#endif + +template [[host_name("kernel_lightning_indexer_q4_0")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_q4_1")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_q5_0")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_q5_1")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_q8_0")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; + kernel void kernel_dsv4_hc_comb_f32( constant ggml_metal_kargs_dsv4_hc_comb & args, device const char * mixes, @@ -11452,4 +11601,3 @@ kernel void kernel_dsv4_hc_post_f32( *(device float *) (dst + i0*args.nb_d0 + idst*args.nb_d1 + it*args.nb_d2) = result[idst]; } } - diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 915c8e7b90..fc0fce0d78 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -7065,7 +7065,7 @@ static inline bool use_flat_gemv_for_large_m_q4_K(const ggml_tensor *tensor) { return tensor->ne[1] >= 32768 && tensor->ne[2] == 1 && tensor->ne[3] == 1; } -static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_tensor *tensor) { +static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) { // gemv_noshuffle variant perf drops for large M, use flat variant for large M. // threshold is well above typical hidden/FFN dims, but below typical vocab sizes. // q6_K flat gemv is worse for smaller K; 2048 seems to be a reasonable threshold. @@ -7083,7 +7083,15 @@ static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_tensor *tensor) { if ((tensor->ne[1] % 128 != 0) && tensor->ne[2] == 1 && tensor->ne[3] == 1) { return true; } - return tensor->ne[1] >= 32768 && tensor->ne[0] >= 2048 && tensor->ne[2] == 1 && tensor->ne[3] == 1; + + // The gemv_noshuffle slowdown tracks TOTAL weight size, not ne0 alone; ne0 >= 2048 is a + // proxy for "large weight" that misses a narrow-hidden vocab-scale lm_head. + // Add a direct size escape so such weights also take the flat path, without changing + // which weights ne0 >= 2048 already routes there. + // The size escape is not taken on the A7X since its compiler miscompiles the flat K-quant GEMV + return tensor->ne[1] >= 32768 + && (tensor->ne[0] >= 2048 || (backend_ctx->adreno_gen != ADRENO_GPU_GEN::A7X && ggml_nbytes(tensor) >= (256ull << 20))) + && tensor->ne[2] == 1 && tensor->ne[3] == 1; } static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_tensor * op) { @@ -9403,7 +9411,7 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer, cl_kernel kernel; #ifdef GGML_OPENCL_USE_ADRENO_KERNELS kernel = backend_ctx->kernel_convert_block_q6_K; - if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(tensor)) { + if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(backend_ctx, tensor)) { kernel = backend_ctx->kernel_convert_block_q6_K_noshuffle; } #else @@ -9436,7 +9444,7 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer, tensor->extra = extra; #ifdef GGML_OPENCL_USE_ADRENO_KERNELS - if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(tensor)) { + if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(backend_ctx, tensor)) { cl_int M = tensor->ne[1]; // ne01 cl_int K = tensor->ne[0]; // ne00 @@ -10473,7 +10481,7 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer, CL_CHECK(clReleaseMemObject(data_device)); return; } - if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(tensor)) { + if (use_adreno_kernels(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(backend_ctx, tensor)) { static ggml_cl_buffer buf_trans_ql; static ggml_cl_buffer buf_trans_qh; static ggml_cl_buffer buf_trans_s; @@ -18895,7 +18903,7 @@ static void ggml_cl_mul_mat(ggml_backend_t backend, const ggml_tensor * src0, co } // q6_K x fp32 - if (src0t == GGML_TYPE_Q6_K && src1t == GGML_TYPE_F32 && !use_flat_gemv_for_large_m_q6_K(src0)) { + if (src0t == GGML_TYPE_Q6_K && src1t == GGML_TYPE_F32 && !use_flat_gemv_for_large_m_q6_K(backend_ctx, src0)) { ggml_cl_mul_mat_q6_K_f32_adreno(backend, src0, src1, dst); return; } diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index 0e23fce2b7..20f50733df 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -2518,7 +2518,13 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = { MODEL_TENSOR.SSM_NORM, MODEL_TENSOR.SSM_IN, MODEL_TENSOR.SSM_BETA_ALPHA, - MODEL_TENSOR.SSM_OUT + MODEL_TENSOR.SSM_OUT, + 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.QWEN3VL: [ MODEL_TENSOR.TOKEN_EMBD, @@ -3404,6 +3410,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/include/llama.h b/include/llama.h index 6e53e22972..f2d7e38858 100644 --- a/include/llama.h +++ b/include/llama.h @@ -1256,6 +1256,7 @@ extern "C" { struct ggml_tensor * probs; struct ggml_tensor * sampled; struct ggml_tensor * candidates; + int64_t n_vocab; }; // user code can implement the interface below in order to create custom llama_sampler @@ -1425,9 +1426,9 @@ 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 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/CMakeLists.txt b/src/CMakeLists.txt index 320784c3a8..24f05cc916 100644 --- a/src/CMakeLists.txt +++ b/src/CMakeLists.txt @@ -25,6 +25,7 @@ add_library(llama llama-kv-cache.cpp llama-kv-cache-iswa.cpp llama-kv-cache-dsa.cpp + llama-kv-cache-msa.cpp llama-kv-cache-dsv4.cpp llama-memory.cpp llama-memory-hybrid.cpp diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index e12a8cdc2a..9dde345df4 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); @@ -3187,6 +3222,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. @@ -3620,6 +3683,7 @@ void llm_graph_context::build_sampling() const { /*.probs =*/ nullptr, /*.sampled =*/ nullptr, /*.candidates =*/ nullptr, + /*.n_vocab =*/ logits_seq->ne[0], }; assert(sampler->iface->backend_apply); diff --git a/src/llama-graph.h b/src/llama-graph.h index 160e294135..32d8d395aa 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 50af97f358..846d4c69a6 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 fc770bf003..6e8336c987 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 241c50365a..96cb045d2e 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 e2b330993b..e74fc4d910 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 0000000000..55ef286caf --- /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 0000000000..f09b6d32b0 --- /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 44cb1668da..8678a326d9 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]]; @@ -906,10 +846,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]); - } } } } @@ -1058,44 +994,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 -> @@ -1104,6 +1002,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 @@ -1210,15 +1113,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); @@ -1262,8 +1156,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); } } @@ -1283,12 +1176,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; } @@ -1337,6 +1224,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; @@ -1405,23 +1298,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); @@ -1523,28 +1399,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; @@ -1979,18 +1833,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, @@ -2303,36 +2145,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; @@ -2581,68 +2393,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; @@ -2844,10 +2594,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]); } @@ -2856,10 +2602,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 d5a92f4405..6cb6dbd2f9 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.cpp b/src/llama-model.cpp index ae219347f8..0d1a9b4bf4 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" @@ -2073,9 +2074,13 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, { res = nullptr; } break; - case LLM_ARCH_DEEPSEEK32: + case LLM_ARCH_MINIMAX_M3: { - res = new llama_kv_cache_dsa( + // 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, @@ -2088,9 +2093,11 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, hparams.n_swa, hparams.swa_type, nullptr, + filter_idx, nullptr); } break; case LLM_ARCH_GLM_DSA: + case LLM_ARCH_DEEPSEEK32: { if (params.ctx_type == LLAMA_CONTEXT_TYPE_MTP && hparams.n_layer_nextn > 0) { // The NextN/MTP draft head runs dense MLA (no DSA indexer), so the @@ -2119,10 +2126,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, @@ -2136,7 +2144,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; @@ -2213,11 +2222,11 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, // checks default: { - // The MTP head is dense-attention only on hybrid Qwen3.5/3.6, so use a plain + // The MTP head is dense-attention only on hybrid Qwen3-Next/3.5/3.6, so use a plain // attention KV cache for the MTP context instead of the hybrid wrapper. - const bool mtp_on_hybrid_qwen35 = + const bool mtp_on_hybrid_qwen = params.ctx_type == LLAMA_CONTEXT_TYPE_MTP && - (arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE); + (arch == LLM_ARCH_QWEN3NEXT || arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE); if (llm_arch_is_recurrent(arch)) { res = new llama_memory_recurrent( @@ -2229,7 +2238,7 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, cparams.n_seq_max, cparams.n_rs_seq, nullptr); - } else if (llm_arch_is_hybrid(arch) && !mtp_on_hybrid_qwen35) { + } else if (llm_arch_is_hybrid(arch) && !mtp_on_hybrid_qwen) { // The main difference between hybrid architectures is the // layer filters, so pick the right one here llama_memory_hybrid::layer_filter_cb filter_attn = nullptr; @@ -2244,7 +2253,7 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, filter_recr = [&](uint32_t il) { return hparams.is_recr(il) && hparams.n_ff(il) == 0; }; - } else if (arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE) { + } else if (arch == LLM_ARCH_QWEN3NEXT || arch == LLM_ARCH_QWEN35 || arch == LLM_ARCH_QWEN35MOE) { filter_attn = [&](uint32_t il) { return il < hparams.n_layer() && !hparams.is_recr(il); }; @@ -2310,12 +2319,12 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, }; } - if (mtp_on_hybrid_qwen35) { + if (mtp_on_hybrid_qwen) { filter = [&](uint32_t il) { return il >= hparams.n_layer(); }; } if ((arch == LLM_ARCH_STEP35 || arch == LLM_ARCH_HY_V3 || arch == LLM_ARCH_GLM_DSA || - arch == LLM_ARCH_MIMO2) && + arch == LLM_ARCH_MIMO2 || arch == LLM_ARCH_DEEPSEEK32) && hparams.n_layer_nextn > 0) { if (params.ctx_type == LLAMA_CONTEXT_TYPE_MTP) { filter = [&](uint32_t il) { return il >= hparams.n_layer(); }; diff --git a/src/llama-sampler.cpp b/src/llama-sampler.cpp index a9cb6bee5f..b2f1abe737 100644 --- a/src/llama-sampler.cpp +++ b/src/llama-sampler.cpp @@ -589,6 +589,7 @@ static bool llama_sampler_backend_support( /*.probs = */ nullptr, /*.sampled = */ nullptr, /*.candidates = */ ggml_new_tensor_1d(ctx, GGML_TYPE_I32, n), + /*.n_vocab = */ n, }; ggml_cgraph * gf = ggml_new_graph(ctx); @@ -2638,7 +2639,7 @@ 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 penalty_last_n; const float penalty_repeat; const float penalty_freq; @@ -2648,10 +2649,49 @@ 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_vocab = 0; + 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 penalty_last_n, + float penalty_repeat, + float penalty_freq, + float penalty_present) + : llama_sampler_backend("penalties") + , 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 +2728,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; } @@ -2736,7 +2775,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 +2786,171 @@ 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(data->n_vocab > 0 && data->n_vocab <= INT32_MAX); + + sctx->has_candidates = data->candidates != nullptr; + sctx->n_vocab = (int32_t) data->n_vocab; + 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,10 +2958,10 @@ 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( @@ -2766,22 +2971,18 @@ struct llama_sampler * llama_sampler_init_penalties( 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( + penalty_last_n, + penalty_repeat, + penalty_freq, + penalty_present) ); } diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index 443cd46408..10032a8c64 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -2532,6 +2532,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; diff --git a/src/models/deepseek2.cpp b/src/models/deepseek2.cpp index a9e8bc5140..ba90c0d077 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/deepseek32.cpp b/src/models/deepseek32.cpp index 32262e6840..8a07a0b71c 100644 --- a/src/models/deepseek32.cpp +++ b/src/models/deepseek32.cpp @@ -44,13 +44,24 @@ void llama_model_deepseek32::load_arch_hparams(llama_model_loader & ml) { GGML_ASSERT(hparams.n_layer_nextn < hparams.n_layer_all && "n_layer_nextn must be < n_layer"); switch (hparams.n_layer()) { - case 62: type = LLM_TYPE_685B_A37B; break; + case 61: type = LLM_TYPE_685B_A37B; break; default: type = LLM_TYPE_UNKNOWN; } } -void llama_model_deepseek32::load_arch_tensors(llama_model_loader &) { +void llama_model_deepseek32::load_arch_tensors(llama_model_loader & ml) { LLAMA_LOAD_LOCALS; + + 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(); if (!is_mla) { throw std::runtime_error("DEEPSEEK32 architecture requires MLA"); @@ -80,12 +91,7 @@ void llama_model_deepseek32::load_arch_tensors(llama_model_loader &) { } for (int i = 0; i < n_layer_all; ++i) { - int flags = 0; - if (i >= n_layer) { - // skip all tensors in the NextN layers - // TODO @ngxson : TENSOR_NOT_REQUIRED was a hack, need to remove it later - flags |= TENSOR_SKIP | TENSOR_NOT_REQUIRED; - } + const int flags = (i >= n_layer) ? mtp_flags : trunk_flags; auto & layer = layers[i]; @@ -138,7 +144,7 @@ void llama_model_deepseek32::load_arch_tensors(llama_model_loader &) { 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 (preserved but unused) - conditionally load for last nextn_predict_layers + // NextN/MTP tensors - conditionally load for last nextn_predict_layers if (i >= n_layer) { layer.nextn.eh_proj = create_tensor(tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", i), { 2 * n_embd, n_embd }, flags); layer.nextn.enorm = create_tensor(tn(LLM_TENSOR_NEXTN_ENORM, "weight", i), { n_embd }, flags); @@ -153,6 +159,9 @@ void llama_model_deepseek32::load_arch_tensors(llama_model_loader &) { } std::unique_ptr llama_model_deepseek32::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); } @@ -430,7 +439,9 @@ llama_model_deepseek32::graph::graph(const llama_model & model, const llm_graph_ Qcur, Kcur, Vcur, nullptr, nullptr, model.layers[il].wv_b, top_k, kq_scale, il); } } - if (il == n_layer - 1 && inp_out_ids) { + // when unmasked nextn embeddings are requested, t_h_nextn must keep all rows, + // so the early output masking has to be skipped (it is applied after the final norm instead) + 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); } @@ -493,6 +504,14 @@ llama_model_deepseek32::graph::graph(const llama_model & model, const llm_graph_ cur = build_norm(cur, model.output_norm, NULL, LLM_NORM_RMS, -1); + // post-norm hidden state feeds the NextN/MTP draft head + 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; @@ -504,3 +523,243 @@ llama_model_deepseek32::graph::graph(const llama_model & model, const llm_graph_ ggml_build_forward_expand(gf, cur); } + +// LLM_GRAPH_TYPE_DECODER_MTP draft head for DeepSeek V3.2 (DEEPSEEK32). +// Semantics mirror the deepseek-family NextN/MTP layer: +// enorm(embed) + hnorm(prev_hidden) -> concat(e, h) -> eh_proj -> +// full deepseek32 decoder block (dense MLA attention + sigmoid-gated MoE FFN +// with shared expert, exactly as the trunk deepseek2 graph builds it) -> +// shared_head_norm (fallback output_norm) -> shared LM head. +// The DSA indexer is not used at runtime. +llama_model_deepseek32::graph_mtp::graph_mtp(const llama_model & model, const llm_graph_params & params) + : llm_graph_context(params) { + GGML_ASSERT(hparams.n_layer_nextn > 0 && "DEEPSEEK32 MTP requires n_layer_nextn > 0"); + GGML_ASSERT(hparams.n_layer_nextn == 1 && "DEEPSEEK32 MTP currently only supports a single MTP block"); + GGML_ASSERT(hparams.is_mla() && "DEEPSEEK32 MTP requires MLA"); + + const int il = hparams.n_layer() + cparams.nextn_layer_offset; + GGML_ASSERT(cparams.nextn_layer_offset >= 0 && + cparams.nextn_layer_offset < (int) hparams.n_layer_nextn && + "nextn_layer_offset out of range [0, n_layer_nextn)"); + 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(layer.ffn_gate_inp && "MTP block missing ffn_gate_inp"); + + // note: these are the actual head sizes you get when treating as MHA or after "decompression" using wv_b for MLA + const int64_t n_embd_head_k = 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 - n_embd_head_qk_rope; + + const uint32_t kv_lora_rank = hparams.n_lora_kv; + + // We have to pre-scale kq_scale and attn_factor to make the YaRN RoPE work correctly. + // See the deepseek2 trunk graph for the detailed explanation - this must match it EXACTLY. + 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)); + + // TODO: extract in a common llm_graph_context::build_inp_embd_h() + 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(); + + // MLA with the absorption optimization uses a K-only cache (V is a view of K) + auto * inp_attn = 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, /*dim=*/ 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); + + // self-attention: dense MLA, same construction as the deepseek2 trunk graph + { + ggml_tensor * q = ggml_mul_mat(ctx0, layer.wq_a, cur); + cb(q, "mtp_q", il); + + q = build_norm(q, layer.attn_q_a_norm, nullptr, LLM_NORM_RMS, il); + cb(q, "mtp_q", il); + + q = ggml_mul_mat(ctx0, layer.wq_b, q); + cb(q, "mtp_q", il); + + // split into {n_embd_head_qk_nope, n_head, n_tokens} + 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), + ggml_row_size(q->type, n_embd_head_k) * n_head, 0); + cb(q_nope, "mtp_q_nope", il); + + // and {n_embd_head_qk_rope, n_head, n_tokens} + 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), + ggml_row_size(q->type, n_embd_head_k) * 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); + + // split into {kv_lora_rank, n_tokens} + 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); + + // and {n_embd_head_qk_rope, 1, n_tokens} + 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); + + 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", 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", il); + + kv_cmpr = build_norm(kv_cmpr, layer.attn_kv_a_norm, nullptr, LLM_NORM_RMS, il); + cb(kv_cmpr, "mtp_kv_cmpr", il); + + // {n_embd_head_qk_nope, n_tokens, n_head} + q_nope = ggml_permute(ctx0, q_nope, 0, 2, 1, 3); + cb(q_nope, "mtp_q_nope_perm", il); + + // {n_embd_head_qk_nope, kv_lora_rank, n_head} x {n_embd_head_qk_nope, n_tokens, n_head} + ggml_tensor * q_nope_absorbed = ggml_mul_mat(ctx0, layer.wk_b, q_nope); + cb(q_nope_absorbed, "mtp_q_nope_absorbed", il); + + // {kv_lora_rank, n_head, n_tokens} + q_nope_absorbed = ggml_permute(ctx0, q_nope_absorbed, 0, 2, 1, 3); + cb(q_nope_absorbed, "mtp_q_nope_absorbed_perm", il); + + // {n_embd_head_qk_rope + kv_lora_rank, n_head, n_tokens} + // note: rope must go first for in-place context shifting in build_rope_shift() + 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, kv_lora_rank, 1, n_tokens); + cb(kv_cmpr, "mtp_kv_cmpr_reshape", il); + + // {n_embd_head_qk_rope + kv_lora_rank, 1, n_tokens} + ggml_tensor * Kcur = ggml_concat(ctx0, kv_cmpr, k_pe, 0); + cb(Kcur, "mtp_Kcur", il); + + // {kv_lora_rank, 1, n_tokens} + ggml_tensor * Vcur = kv_cmpr; + cb(Vcur, "mtp_Vcur", il); + + // note: MLA with the absorption optimization converts into MQA (ie: GQA with 1 group) + cur = build_attn(inp_attn, + layer.wo, NULL, 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, NULL, LLM_NORM_RMS, il); + cb(cur, "mtp_ffn_norm", il); + + // MoE FFN with shared expert - same construction as the deepseek2 trunk graph + 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, + layer.ffn_up_exps_s, + layer.ffn_gate_exps_s, + layer.ffn_down_exps_s); + cb(moe_out, "mtp_ffn_moe_out", il); + + // FFN shared expert + ggml_tensor * ffn_shexp = + build_ffn(cur, + layer.ffn_up_shexp, NULL, layer.ffn_up_shexp_s, + layer.ffn_gate_shexp, NULL, layer.ffn_gate_shexp_s, + layer.ffn_down_shexp, NULL, layer.ffn_down_shexp_s, + NULL, 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); + + // shared_head_norm applied after the decoder block, before the shared LM head. + // The post-norm hidden state seeds the next MTP step. + ggml_tensor * head_norm_w = layer.nextn.shared_head_norm + ? layer.nextn.shared_head_norm + : model.output_norm; + GGML_ASSERT(head_norm_w && "DEEPSEEK32 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; + + 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 && "DEEPSEEK32 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); +} + diff --git a/src/models/minimax-m3.cpp b/src/models/minimax-m3.cpp index 0773ad5435..854d5aed0f 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 37f842b5a0..ad3dadaf39 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -1089,6 +1089,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; }; @@ -1102,6 +1106,10 @@ struct llama_model_deepseek32 : 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; }; @@ -2038,6 +2046,10 @@ struct llama_model_qwen3next : public llama_model_base { const llama_model & model; }; + 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/src/models/qwen3next.cpp b/src/models/qwen3next.cpp index 09b66423d5..0808fd87aa 100644 --- a/src/models/qwen3next.cpp +++ b/src/models/qwen3next.cpp @@ -13,7 +13,11 @@ void llama_model_qwen3next::load_arch_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_SSM_TIME_STEP_RANK, hparams.ssm_dt_rank); ml.get_key(LLM_KV_SSM_GROUP_COUNT, hparams.ssm_n_group); - // Mark recurrent layers (linear attention layers) + // NextN/MTP: extra decoder block appended beyond the main stack + ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.n_layer_nextn, false); + GGML_ASSERT(hparams.n_layer_nextn < hparams.n_layer_all && "n_layer_nextn must be < n_layer_all"); + + // Mark recurrent layers (linear attention layers). if (!ml.get_key_or_arr(LLM_KV_ATTENTION_RECURRENT_LAYERS, hparams.is_recr_impl, hparams.n_layer_all, false)) { uint32_t full_attn_interval = 4; ml.get_key(LLM_KV_FULL_ATTENTION_INTERVAL, full_attn_interval, false); @@ -28,13 +32,17 @@ void llama_model_qwen3next::load_arch_hparams(llama_model_loader & ml) { } } -void llama_model_qwen3next::load_arch_tensors(llama_model_loader &) { +void llama_model_qwen3next::load_arch_tensors(llama_model_loader & ml) { LLAMA_LOAD_LOCALS; if (n_expert == 0) { throw std::runtime_error(arch_name() + " model cannot have zero experts"); } + const bool mtp_only = (hparams.n_layer_nextn > 0) && (ml.get_weight("blk.0.attn_norm.weight") == nullptr); + const int trunk_flags = mtp_only ? TENSOR_NOT_REQUIRED : 0; + int mtp_flags = !ml.load_mtp ? TENSOR_SKIP : 0; + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), { n_embd, n_vocab }, 0); // output @@ -61,49 +69,73 @@ void llama_model_qwen3next::load_arch_tensors(llama_model_loader &) { const int64_t qkvz_dim = key_dim * 2 + value_dim * 2; const int64_t ba_dim = n_v_heads * 2; - for (int i = 0; i < n_layer; ++i) { - auto & layer = layers[i]; - const uint32_t n_ff_shexp = hparams.n_ff_shexp > 0 ? hparams.n_ff_shexp : hparams.n_ff(i); + auto load_block_trunk = [&](int il, int flags) { + auto & layer = layers[il]; + const uint32_t n_ff_shexp = hparams.n_ff_shexp > 0 ? hparams.n_ff_shexp : hparams.n_ff(il); - layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), { n_embd }, 0); - layer.attn_post_norm = create_tensor(tn(LLM_TENSOR_ATTN_POST_NORM, "weight", i), { n_embd }, 0); + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", il), { n_embd }, flags); + layer.attn_post_norm = create_tensor(tn(LLM_TENSOR_ATTN_POST_NORM, "weight", il), { n_embd }, flags); - if (!hparams.is_recr(i)) { + if (!hparams.is_recr(il)) { // Attention layers - create_tensor_qkv(layer, i, n_embd, n_embd_head_k * n_head * 2, n_embd_k_gqa, n_embd_v_gqa, 0); - layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), { n_embd_head_k * n_head, n_embd }, 0); - + create_tensor_qkv(layer, il, n_embd, n_embd_head_k * n_head * 2, n_embd_k_gqa, n_embd_v_gqa, flags); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", il), { n_embd_head_k * n_head, n_embd }, flags); // Q/K normalization for attention layers - layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", i), { n_embd_head_k }, 0); - layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", i), { n_embd_head_k }, 0); + layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", il), { n_embd_head_k }, flags); + layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", il), { n_embd_head_k }, flags); } else { // Linear attention (gated delta net) specific tensors // Create tensors with calculated dimensions // note: ssm_in is used by legacy GGUF - layer.ssm_in = create_tensor(tn(LLM_TENSOR_SSM_IN, "weight", i), { n_embd, qkvz_dim }, TENSOR_NOT_REQUIRED); - layer.wqkv = create_tensor(tn(LLM_TENSOR_ATTN_QKV, "weight", i), { n_embd, key_dim * 2 + value_dim }, TENSOR_NOT_REQUIRED); - layer.wqkv_gate = create_tensor(tn(LLM_TENSOR_ATTN_GATE, "weight", i), { n_embd, value_dim }, TENSOR_NOT_REQUIRED); - layer.ssm_conv1d = create_tensor(tn(LLM_TENSOR_SSM_CONV1D, "weight", i), { hparams.ssm_d_conv, conv_dim }, 0); - layer.ssm_dt = create_tensor(tn(LLM_TENSOR_SSM_DT, "bias", i), { hparams.ssm_dt_rank }, 0); - layer.ssm_a = create_tensor(tn(LLM_TENSOR_SSM_A_NOSCAN, i), { hparams.ssm_dt_rank }, 0); - layer.ssm_beta_alpha = create_tensor(tn(LLM_TENSOR_SSM_BETA_ALPHA, "weight", i), { n_embd, ba_dim }, 0); - layer.ssm_norm = create_tensor(tn(LLM_TENSOR_SSM_NORM, "weight", i), { head_v_dim }, 0); - layer.ssm_out = create_tensor(tn(LLM_TENSOR_SSM_OUT, "weight", i), { value_dim, n_embd }, 0); + layer.ssm_in = create_tensor(tn(LLM_TENSOR_SSM_IN, "weight", il), { n_embd, qkvz_dim }, TENSOR_NOT_REQUIRED | flags); + layer.wqkv = create_tensor(tn(LLM_TENSOR_ATTN_QKV, "weight", il), { n_embd, key_dim * 2 + value_dim }, TENSOR_NOT_REQUIRED | flags); + layer.wqkv_gate = create_tensor(tn(LLM_TENSOR_ATTN_GATE, "weight", il), { n_embd, value_dim }, TENSOR_NOT_REQUIRED | flags); + layer.ssm_conv1d = create_tensor(tn(LLM_TENSOR_SSM_CONV1D, "weight", il), { hparams.ssm_d_conv, conv_dim }, flags); + layer.ssm_dt = create_tensor(tn(LLM_TENSOR_SSM_DT, "bias", il), { hparams.ssm_dt_rank }, flags); + layer.ssm_a = create_tensor(tn(LLM_TENSOR_SSM_A_NOSCAN, il), { hparams.ssm_dt_rank }, flags); + layer.ssm_beta_alpha = create_tensor(tn(LLM_TENSOR_SSM_BETA_ALPHA, "weight", il), { n_embd, ba_dim }, flags); + layer.ssm_norm = create_tensor(tn(LLM_TENSOR_SSM_NORM, "weight", il), { head_v_dim }, flags); + layer.ssm_out = create_tensor(tn(LLM_TENSOR_SSM_OUT, "weight", il), { value_dim, n_embd }, flags); } - layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), { n_embd, n_expert }, 0); - 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_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", il), { n_embd, n_expert }, flags); + layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", il), { n_ff_exp, n_embd, n_expert }, flags); + create_tensor_gate_up_exps(layer, il, n_embd, n_ff_exp, n_expert, flags); // Shared experts - layer.ffn_gate_inp_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP_SHEXP, "weight", i), { n_embd }, 0); - layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", i), { n_embd, n_ff_shexp }, 0); - layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", i), { n_embd, n_ff_shexp }, 0); - layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", i), { n_ff_shexp, n_embd }, 0); + layer.ffn_gate_inp_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP_SHEXP, "weight", il), { n_embd }, flags); + layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_SHEXP, "weight", il), { n_embd, n_ff_shexp }, flags); + layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", il), { n_embd, n_ff_shexp }, flags); + layer.ffn_down_shexp = create_tensor(tn(LLM_TENSOR_FFN_DOWN_SHEXP, "weight", il), { n_ff_shexp, n_embd }, flags); + }; + + auto load_block_mtp = [&](int il) { + // MTP head is identical to the trunk block (full attention + FFN) + load_block_trunk(il, mtp_flags); + + auto & layer = layers[il]; + + // NextN-specific tensors that define the MTP block. + layer.nextn.eh_proj = create_tensor(tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", il), { 2 * n_embd, n_embd }, mtp_flags); + layer.nextn.enorm = create_tensor(tn(LLM_TENSOR_NEXTN_ENORM, "weight", il), { n_embd }, mtp_flags); + layer.nextn.hnorm = create_tensor(tn(LLM_TENSOR_NEXTN_HNORM, "weight", il), { n_embd }, mtp_flags); + layer.nextn.embed_tokens = create_tensor(tn(LLM_TENSOR_NEXTN_EMBED_TOKENS, "weight", il), { n_embd, n_vocab }, mtp_flags | TENSOR_NOT_REQUIRED); + layer.nextn.shared_head_head = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "weight", il), { n_embd, n_vocab }, mtp_flags | TENSOR_NOT_REQUIRED); + layer.nextn.shared_head_norm = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "weight", il), { n_embd }, mtp_flags | TENSOR_NOT_REQUIRED); + }; + + for (int i = 0; i < n_layer; i++) { + load_block_trunk(i, trunk_flags); + } + for (int i = n_layer; i < n_layer_all; i++) { + load_block_mtp(i); } } std::unique_ptr llama_model_qwen3next::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); } @@ -120,6 +152,7 @@ llama_model_qwen3next::graph::graph(const llama_model & model, const llm_graph_p ggml_tensor * inp_pos = build_inp_pos(); ggml_tensor * inp_out_ids = build_inp_out_ids(); + // MTP/NextN layers are loaded as extra decoder blocks but not executed in the main pass. for (int il = 0; il < n_layer; ++il) { res->t_layer_inp[il] = inpL; @@ -139,7 +172,7 @@ llama_model_qwen3next::graph::graph(const llama_model & model, const llm_graph_p cur = build_layer_attn(inp->get_attn(), cur, inp_pos, il); } - if (il == n_layer - 1 && inp_out_ids) { + if (il == n_layer - 1 && inp_out_ids && cparams.embeddings_nextn_masked) { cur = ggml_get_rows(ctx0, cur, inp_out_ids); inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids); } @@ -171,9 +204,16 @@ llama_model_qwen3next::graph::graph(const llama_model & model, const llm_graph_p } cur = inpL; - // Final norm + // post-norm hidden state is input to both the LM head and the MTP head cur = build_norm(cur, model.output_norm, nullptr, LLM_NORM_RMS, -1); + cb(cur, "h_nextn", -1); + res->t_h_nextn = cur; + + if (!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; @@ -186,14 +226,6 @@ llama_model_qwen3next::graph::graph(const llama_model & model, const llm_graph_p ggml_build_forward_expand(gf, cur); } -// utility to get one slice from the third dimension -// input dim: [x, y, c, b] -// output dim: [x, y, 1, b] -static ggml_tensor * get_slice_2d(ggml_context * ctx0, ggml_tensor * t, int64_t c) { - return ggml_view_4d(ctx0, t, t->ne[0], t->ne[1], 1, t->ne[3], - t->nb[1], t->nb[2], t->nb[3], t->nb[2] * c); -} - ggml_tensor * llama_model_qwen3next::graph::build_norm_gated( ggml_tensor * input, ggml_tensor * weights, @@ -216,7 +248,7 @@ ggml_tensor * llama_model_qwen3next::graph::build_layer_attn( // Order: joint QG projection, QG split, Q norm, KV projection, K norm, RoPE, attention // Qwen3Next uses a single Q projection that outputs query + gate - ggml_tensor * Qcur_full = build_lora_mm(model.layers[il].wq, cur); + ggml_tensor * Qcur_full = build_lora_mm(model.layers[il].wq, cur, model.layers[il].wq_s); cb(Qcur_full, "Qcur_full", il); Qcur_full = ggml_reshape_4d(ctx0, Qcur_full, n_embd_head * 2, n_head, n_tokens, 1); @@ -232,10 +264,10 @@ ggml_tensor * llama_model_qwen3next::graph::build_layer_attn( Qcur_full->nb[1], Qcur_full->nb[2], Qcur_full->nb[3], n_embd_head * ggml_element_size(Qcur_full)); cb(gate, "gate", il); - ggml_tensor * Kcur = build_lora_mm(model.layers[il].wk, cur); + ggml_tensor * Kcur = build_lora_mm(model.layers[il].wk, cur, model.layers[il].wk_s); cb(Kcur, "Kcur", il); - ggml_tensor * Vcur = build_lora_mm(model.layers[il].wv, cur); + ggml_tensor * Vcur = build_lora_mm(model.layers[il].wv, cur, model.layers[il].wv_s); cb(Vcur, "Vcur", il); Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens); @@ -274,8 +306,6 @@ ggml_tensor * llama_model_qwen3next::graph::build_layer_attn( gate = ggml_sigmoid(ctx0, gate); cb(gate, "gate_sigmoid", il); - gate = ggml_reshape_2d(ctx0, gate, n_embd_head * n_head, n_tokens); - cur = ggml_mul(ctx0, cur, gate); cb(cur, "attn_gated", il); @@ -550,16 +580,19 @@ ggml_tensor * llama_model_qwen3next::graph::build_layer_ffn(ggml_tensor * cur, c LLM_FFN_SILU, true, hparams.expert_weights_scale, LLAMA_EXPERT_GATING_FUNC_TYPE_SOFTMAX, il, - nullptr, model.layers[il].ffn_gate_up_exps); + nullptr, model.layers[il].ffn_gate_up_exps, + model.layers[il].ffn_up_exps_s, + model.layers[il].ffn_gate_exps_s, + model.layers[il].ffn_down_exps_s); cb(moe_out, "ffn_moe_out", il); // Add shared experts if present - following Qwen3Next reference implementation if (model.layers[il].ffn_up_shexp != nullptr) { ggml_tensor * ffn_shexp = build_ffn(cur, - model.layers[il].ffn_up_shexp, NULL, NULL, - model.layers[il].ffn_gate_shexp, NULL, NULL, - model.layers[il].ffn_down_shexp, NULL, NULL, + model.layers[il].ffn_up_shexp, NULL, model.layers[il].ffn_up_shexp_s, + model.layers[il].ffn_gate_shexp, NULL, model.layers[il].ffn_gate_shexp_s, + model.layers[il].ffn_down_shexp, NULL, model.layers[il].ffn_down_shexp_s, NULL, LLM_FFN_SILU, LLM_FFN_PAR, il); cb(ffn_shexp, "ffn_shexp", il); @@ -593,3 +626,198 @@ ggml_tensor * llama_model_qwen3next::graph::build_layer_ffn(ggml_tensor * cur, c } return cur; } + +// LLM_GRAPH_TYPE_DECODER_MTP draft head for Qwen3-Next +llama_model_qwen3next::graph_mtp::graph_mtp(const llama_model & model, const llm_graph_params & params) + : llm_graph_context(params) { + GGML_ASSERT(hparams.n_layer_nextn > 0 && "QWEN3NEXT MTP requires n_layer_nextn > 0"); + GGML_ASSERT(hparams.n_layer_nextn == 1 && "QWEN3NEXT MTP currently only supports a single MTP block"); + + const int64_t n_embd_head = hparams.n_embd_head_v(); + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + 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(layer.ffn_gate_inp && "MTP block missing ffn_gate_inp"); + + // TODO: extract in a common llm_graph_context::build_inp_embd_h() + 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); + + // TODO: make static using `ggml_build_forward_select()` + // see llm_graph_context::build_inp_embd() for reference + 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 = build_attn_inp_kv(); + + 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, /*dim=*/ 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 * Qcur_full = build_lora_mm(layer.wq, cur, layer.wq_s); + cb(Qcur_full, "mtp_Qcur_full", il); + + ggml_tensor * Qcur = ggml_view_3d(ctx0, Qcur_full, + n_embd_head, n_head, n_tokens, + ggml_element_size(Qcur_full) * n_embd_head * 2, + ggml_element_size(Qcur_full) * n_embd_head * 2 * n_head, + 0); + Qcur = build_norm(Qcur, layer.attn_q_norm, nullptr, LLM_NORM_RMS, il); + cb(Qcur, "mtp_Qcur_normed", il); + + ggml_tensor * Kcur = build_lora_mm(layer.wk, cur, layer.wk_s); + Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens); + Kcur = build_norm(Kcur, layer.attn_k_norm, nullptr, LLM_NORM_RMS, il); + cb(Kcur, "mtp_Kcur_normed", il); + + ggml_tensor * Vcur = build_lora_mm(layer.wv, cur, layer.wv_s); + Vcur = ggml_reshape_3d(ctx0, Vcur, n_embd_head, n_head_kv, n_tokens); + + Qcur = ggml_rope_ext(ctx0, Qcur, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + Kcur = ggml_rope_ext(ctx0, Kcur, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + + cb(Qcur, "mtp_Qcur", il); + cb(Kcur, "mtp_Kcur", il); + cb(Vcur, "mtp_Vcur", il); + + const float kq_scale = hparams.f_attention_scale == 0.0f + ? 1.0f / sqrtf(float(n_embd_head)) : hparams.f_attention_scale; + + cur = build_attn(inp_attn, + nullptr, nullptr, nullptr, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, kq_scale, il); + cb(cur, "mtp_attn_pregate", il); + + ggml_tensor * gate = ggml_view_3d(ctx0, Qcur_full, + n_embd_head, n_head, n_tokens, + ggml_element_size(Qcur_full) * n_embd_head * 2, + ggml_element_size(Qcur_full) * n_embd_head * 2 * n_head, + ggml_element_size(Qcur_full) * n_embd_head); + + // TODO: CUDA is missing non-contiguous unary ops. when implemented: remove this cont + gate = ggml_cont_2d(ctx0, gate, n_embd_head * n_head, n_tokens); + cb(gate, "mtp_gate", il); + + cur = ggml_mul(ctx0, cur, ggml_sigmoid(ctx0, gate)); + cur = build_lora_mm(layer.wo, cur, layer.wo_s); + cb(cur, "mtp_attn_out", il); + + if (inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids); + } + + cur = ggml_add(ctx0, cur, inpSA); + cb(cur, "mtp_attn_residual", il); + + ggml_tensor * ffn_residual = cur; + cur = build_norm(cur, layer.attn_post_norm, nullptr, LLM_NORM_RMS, il); + cb(cur, "mtp_attn_post_norm", il); + + // MoE FFN — routed experts plus gated shared expert (mirrors the trunk). + ggml_tensor * moe_out = + build_moe_ffn(cur, + layer.ffn_gate_inp, + layer.ffn_up_exps, + layer.ffn_gate_exps, + layer.ffn_down_exps, + nullptr, + n_expert, n_expert_used, + LLM_FFN_SILU, true, + hparams.expert_weights_scale, + LLAMA_EXPERT_GATING_FUNC_TYPE_SOFTMAX, il, + nullptr, layer.ffn_gate_up_exps, + layer.ffn_up_exps_s, + layer.ffn_gate_exps_s, + layer.ffn_down_exps_s); + cb(moe_out, "mtp_ffn_moe_out", il); + + if (layer.ffn_up_shexp != nullptr) { + ggml_tensor * ffn_shexp = + build_ffn(cur, + layer.ffn_up_shexp, nullptr, layer.ffn_up_shexp_s, + layer.ffn_gate_shexp, nullptr, layer.ffn_gate_shexp_s, + layer.ffn_down_shexp, nullptr, layer.ffn_down_shexp_s, + nullptr, + LLM_FFN_SILU, LLM_FFN_PAR, il); + cb(ffn_shexp, "mtp_ffn_shexp", il); + + ggml_tensor * shared_gate = build_lora_mm(layer.ffn_gate_inp_shexp, cur); + shared_gate = ggml_sigmoid(ctx0, shared_gate); + cb(shared_gate, "mtp_shared_expert_gate_sigmoid", il); + + ffn_shexp = ggml_mul(ctx0, ffn_shexp, shared_gate); + cb(ffn_shexp, "mtp_ffn_shexp_gated", il); + + cur = ggml_add(ctx0, moe_out, ffn_shexp); + } else { + cur = moe_out; + } + cb(cur, "mtp_ffn_out", il); + + cur = ggml_add(ctx0, cur, ffn_residual); + 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 && "QWEN3NEXT 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; + + 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 && "QWEN3NEXT 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); +} diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 881e55c75a..419e1eba4c 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -258,6 +258,9 @@ llama_build_and_test(test-thread-safety.cpp ARGS -m "${MODEL_DEST}" -ngl 99 -p " set_tests_properties(test-thread-safety PROPERTIES FIXTURES_REQUIRED test-download-model) llama_build_and_test(test-arg-parser.cpp) +llama_build_and_test(test-model-resolution.cpp) +# the test serves its repos from an httplib server, and the library links it privately +target_link_libraries(test-model-resolution PRIVATE cpp-httplib) if (NOT LLAMA_SANITIZE_ADDRESS AND NOT GGML_SCHED_NO_REALLOC) # TODO: repair known memory leaks diff --git a/tests/test-arg-parser.cpp b/tests/test-arg-parser.cpp index 1d3584f903..fd5adb740e 100644 --- a/tests/test-arg-parser.cpp +++ b/tests/test-arg-parser.cpp @@ -99,6 +99,34 @@ static void test(void) { argv = {"binary_name", "-sm", "hello"}; assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON)); + { + common_params penalty_params; + + argv = {"binary_name", "--repeat-penalty", "0"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON)); + + argv = {"binary_name", "--repeat-penalty", "-1"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON)); + + argv = {"binary_name", "--repeat-penalty", "nan"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON)); + + argv = {"binary_name", "--repeat-penalty", "inf"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON)); + + argv = {"binary_name", "--repeat-penalty", "-inf"}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON)); + + const char * penalty_options[] = {"--frequency-penalty", "--presence-penalty"}; + const char * nonfinite_values[] = {"nan", "inf", "-inf"}; + for (const char * option : penalty_options) { + for (const char * value : nonfinite_values) { + argv = {"binary_name", option, value}; + assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), penalty_params, LLAMA_EXAMPLE_COMMON)); + } + } + } + // non-existence arg in specific example (--draft cannot be used outside llama-speculative) argv = {"binary_name", "--draft", "123"}; assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_EMBEDDING)); diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 32468bf22d..8cb5989358 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9731,6 +9731,12 @@ static std::vector> make_test_cases_eval() { } } + for (int kv : { 1, 7, 8, 63, 64, 65 }) { + for (ggml_type type_K : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q8_0, GGML_TYPE_Q5_1, GGML_TYPE_Q5_0, GGML_TYPE_Q4_1, GGML_TYPE_Q4_0}) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 32, 4, 1, type_K)); + } + } + return test_cases; } #ifdef _MSC_VER diff --git a/tests/test-backend-sampler.cpp b/tests/test-backend-sampler.cpp index c24076e313..1a46468ba2 100644 --- a/tests/test-backend-sampler.cpp +++ b/tests/test-backend-sampler.cpp @@ -8,12 +8,15 @@ #endif #include +#include #include #include #include +#include #include #include #include +#include #include struct test_args { @@ -761,6 +764,563 @@ static void test_backend_logit_bias_sampling(const test_params & params) { printf("backend logit bias sampling test PASSED\n"); } +static void accept_prompt(llama_sampler * smpl, const llama_vocab * vocab, const std::string & prompt) { + const llama_token bos = llama_vocab_bos(vocab); + if (bos != LLAMA_TOKEN_NULL) { + llama_sampler_accept(smpl, bos); + } + + std::vector tokens(64); + int32_t n_tokens = llama_tokenize(vocab, prompt.c_str(), (int32_t) prompt.size(), + tokens.data(), (int32_t) tokens.size(), false, false); + if (n_tokens < 0) { + tokens.resize(-n_tokens); + n_tokens = llama_tokenize(vocab, prompt.c_str(), (int32_t) prompt.size(), + tokens.data(), (int32_t) tokens.size(), false, false); + } + + for (int32_t i = 0; i < n_tokens; ++i) { + llama_sampler_accept(smpl, tokens[i]); + } +} + +static std::vector decode_raw_logits(const test_params & params, const std::string & prompt) { + const int seq_id = 0; + const int n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(params.model.get())); + std::vector empty_configs; + test_context ctx(params, empty_configs); + + GGML_ASSERT(ctx.decode({{ seq_id, prompt }})); + + float * logits = llama_get_logits_ith(ctx.ctx.get(), ctx.idx_for_seq(seq_id)); + GGML_ASSERT(logits != nullptr); + return std::vector(logits, logits + n_vocab); +} + +static std::vector apply_cpu_sampler( + const std::vector & raw_logits, + llama_sampler * sampler) { + std::vector data; + data.reserve(raw_logits.size()); + for (llama_token token = 0; token < (llama_token) raw_logits.size(); ++token) { + data.push_back({ token, raw_logits[token], 0.0f }); + } + + llama_token_data_array cur_p = { data.data(), data.size(), -1, false }; + llama_sampler_apply(sampler, &cur_p); + data.resize(cur_p.size); + return data; +} + +using sampler_setup_fn = std::function; +using sampler_init_fn = std::function; + +enum class penalties_position { + before_filter, + after_filter, +}; + +static void add_filter_and_penalties( + llama_sampler * chain, + const sampler_init_fn & init_filter, + int32_t penalty_last_n, + float penalty_repeat, + float penalty_freq, + float penalty_present, + penalties_position position) { + const auto add_penalties = [&]() { + llama_sampler_chain_add(chain, llama_sampler_init_penalties( + penalty_last_n, penalty_repeat, penalty_freq, penalty_present)); + }; + + if (position == penalties_position::before_filter) { + add_penalties(); + llama_sampler_chain_add(chain, init_filter()); + } else { + llama_sampler_chain_add(chain, init_filter()); + add_penalties(); + } +} + +static llama_sampler_ptr make_sampler_chain( + const sampler_setup_fn & add_samplers, + const sampler_setup_fn & accept_history) { + llama_sampler_ptr chain(llama_sampler_chain_init(llama_sampler_chain_default_params())); + add_samplers(chain.get()); + accept_history(chain.get()); + return chain; +} + +struct backend_sampler_output { + std::vector logits; + std::vector candidates; +}; + +static backend_sampler_output run_backend_sampler( + const test_params & params, + const std::string & prompt, + llama_sampler * sampler) { + const int seq_id = 0; + std::vector configs = {{ seq_id, sampler }}; + test_context ctx(params, configs); + + GGML_ASSERT(ctx.decode({{ seq_id, prompt }})); + llama_synchronize(ctx.ctx.get()); + + const int32_t idx = ctx.idx_for_seq(seq_id); + const uint32_t n_logits = llama_get_sampled_logits_count_ith(ctx.ctx.get(), idx); + const uint32_t n_candidates = llama_get_sampled_candidates_count_ith(ctx.ctx.get(), idx); + float * logits = llama_get_sampled_logits_ith(ctx.ctx.get(), idx); + llama_token * candidates = llama_get_sampled_candidates_ith(ctx.ctx.get(), idx); + GGML_ASSERT(logits != nullptr); + + backend_sampler_output result; + result.logits.assign(logits, logits + n_logits); + result.candidates.resize(n_logits); + + if (n_candidates == 0) { + for (uint32_t i = 0; i < n_logits; ++i) { + result.candidates[i] = (llama_token) i; + } + } else { + GGML_ASSERT(candidates != nullptr); + GGML_ASSERT(n_candidates == n_logits); + std::memcpy(result.candidates.data(), candidates, n_candidates * sizeof(llama_token)); + } + + return result; +} + +struct sampler_comparison_output { + std::vector expected; + backend_sampler_output actual; +}; + +static sampler_comparison_output run_sampler_comparison( + const test_params & params, + const std::string & prompt, + const std::vector & raw_logits, + const sampler_setup_fn & add_samplers, + const sampler_setup_fn & accept_history) { + llama_sampler_ptr cpu_chain = make_sampler_chain(add_samplers, accept_history); + llama_sampler_ptr backend_chain = make_sampler_chain(add_samplers, accept_history); + return { + apply_cpu_sampler(raw_logits, cpu_chain.get()), + run_backend_sampler(params, prompt, backend_chain.get()), + }; +} + +static std::unordered_map map_logits(const std::vector & data) { + std::unordered_map result; + result.reserve(data.size()); + for (const auto & item : data) { + result[item.id] = item.logit; + } + return result; +} + +struct sampler_comparison_stats { + int n_mismatch = 0; + int n_masked = 0; + float max_diff = 0.0f; +}; + +static sampler_comparison_stats compare_sampler_outputs( + const char * name, + const std::unordered_map & expected, + const backend_sampler_output & actual, + bool allow_extra_candidates = false) { + GGML_ASSERT(actual.logits.size() == actual.candidates.size()); + + sampler_comparison_stats result; + std::unordered_set seen; + seen.reserve(actual.candidates.size()); + + for (size_t i = 0; i < actual.logits.size(); ++i) { + const llama_token token = actual.candidates[i]; + const float logit = actual.logits[i]; + if (!seen.insert(token).second || std::isnan(logit)) { + if (result.n_mismatch < 5) { + printf("%s token %d has invalid backend output\n", name, token); + } + ++result.n_mismatch; + continue; + } + + const auto it = expected.find(token); + if (it == expected.end()) { + if (std::isinf(logit) && logit < 0.0f) { + ++result.n_masked; + } else if (!allow_extra_candidates) { + if (result.n_mismatch < 5) { + printf("%s token %d was not masked\n", name, token); + } + ++result.n_mismatch; + } + continue; + } + + const float diff = fabsf(it->second - logit); + result.max_diff = std::max(result.max_diff, diff); + if (!std::isfinite(logit) || diff > 1e-3f) { + if (result.n_mismatch < 5) { + printf("%s mismatch token %d: cpu=%.6f backend=%.6f diff=%.6f\n", + name, token, it->second, logit, diff); + } + ++result.n_mismatch; + } + } + + for (const auto & item : expected) { + if (seen.find(item.first) == seen.end()) { + if (result.n_mismatch < 5) { + printf("%s missing backend token %d\n", name, item.first); + } + ++result.n_mismatch; + } + } + + printf("%s logits: max_diff=%.6f n_masked=%d n_mismatch=%d\n", + name, result.max_diff, result.n_masked, result.n_mismatch); + return result; +} + +static float find_backend_logit(const backend_sampler_output & output, llama_token token) { + for (size_t i = 0; i < output.candidates.size(); ++i) { + if (output.candidates[i] == token) { + return output.logits[i]; + } + } + GGML_ABORT("backend token not found"); +} + +static sampler_comparison_output run_penalties_comparison( + const test_params & params, + int32_t penalty_last_n, + float penalty_repeat, + float penalty_freq, + float penalty_present, + const std::string & prompt, + const std::function & extra_accept = {}) { + const auto * vocab = llama_model_get_vocab(params.model.get()); + const std::vector raw_logits = decode_raw_logits(params, prompt); + const auto add_samplers = [&](llama_sampler * chain) { + llama_sampler_chain_add(chain, llama_sampler_init_penalties( + penalty_last_n, penalty_repeat, penalty_freq, penalty_present)); + }; + const auto accept_history = [&](llama_sampler * chain) { + accept_prompt(chain, vocab, prompt); + if (extra_accept) { + extra_accept(chain); + } + }; + + return run_sampler_comparison( + params, prompt, raw_logits, add_samplers, accept_history); +} + +static void compare_penalties_logits( + const test_params & params, + int32_t penalty_last_n, + float penalty_repeat, + float penalty_freq, + float penalty_present, + const std::string & prompt, + const std::function & extra_accept = {}) { + const sampler_comparison_output output = run_penalties_comparison( + params, penalty_last_n, penalty_repeat, penalty_freq, penalty_present, prompt, extra_accept); + + GGML_ASSERT(output.expected.size() == output.actual.logits.size()); + + const sampler_comparison_stats stats = compare_sampler_outputs( + "penalties", map_logits(output.expected), output.actual); + GGML_ASSERT(stats.n_masked == 0); + GGML_ASSERT(stats.n_mismatch == 0); +} + +static void test_penalty_parameter_values(const test_params & params) { + struct penalty_test_case { + const char * name; + float repeat; + float frequency; + float presence; + }; + + const penalty_test_case cases[] = { + { "frequency -1", 1.0f, -1.0f, 0.0f }, + { "frequency 0", 1.0f, 0.0f, 0.0f }, + { "frequency 1", 1.0f, 1.0f, 0.0f }, + { "presence -1", 1.0f, 0.0f, -1.0f }, + { "presence 0", 1.0f, 0.0f, 0.0f }, + { "presence 1", 1.0f, 0.0f, 1.0f }, + { "repeat 1", 1.0f, 0.0f, 0.0f }, + }; + + int n_failed = 0; + for (const auto & test : cases) { + const sampler_comparison_output output = run_penalties_comparison( + params, 64, test.repeat, test.frequency, test.presence, "Hello Hello world"); + GGML_ASSERT(output.expected.size() == output.actual.logits.size()); + const sampler_comparison_stats stats = compare_sampler_outputs( + test.name, map_logits(output.expected), output.actual); + n_failed += stats.n_mismatch != 0; + } + + GGML_ASSERT(n_failed == 0); +} + +static void compare_top_k_penalties_logits( + const test_params & params, + int32_t k, + int32_t penalty_last_n, + float penalty_repeat, + float penalty_freq, + float penalty_present, + const std::string & prompt, + penalties_position position) { + const auto * vocab = llama_model_get_vocab(params.model.get()); + const std::vector raw_logits = decode_raw_logits(params, prompt); + const int n_vocab = (int) raw_logits.size(); + + GGML_ASSERT(n_vocab > k); + + const sampler_init_fn init_top_k = [k]() { + return llama_sampler_init_top_k(k); + }; + llama_sampler_ptr top_k(init_top_k()); + const std::vector top_k_data = apply_cpu_sampler(raw_logits, top_k.get()); + GGML_ASSERT(top_k_data.size() == (size_t) k); + const llama_token retained_history_token = top_k_data[0].id; + + llama_token excluded_history_token = LLAMA_TOKEN_NULL; + for (llama_token token = 0; token < n_vocab; ++token) { + const auto it = std::find_if(top_k_data.begin(), top_k_data.end(), [token](const llama_token_data & data) { + return data.id == token; + }); + if (it == top_k_data.end()) { + excluded_history_token = token; + break; + } + } + GGML_ASSERT(excluded_history_token != LLAMA_TOKEN_NULL); + + const auto add_samplers = [&](llama_sampler * chain) { + add_filter_and_penalties(chain, init_top_k, + penalty_last_n, penalty_repeat, penalty_freq, penalty_present, position); + }; + + auto accept_history = [&](llama_sampler * smpl) { + accept_prompt(smpl, vocab, prompt); + llama_sampler_accept(smpl, excluded_history_token); + llama_sampler_accept(smpl, excluded_history_token); + llama_sampler_accept(smpl, retained_history_token); + llama_sampler_accept(smpl, retained_history_token); + }; + + const sampler_comparison_output output = run_sampler_comparison( + params, prompt, raw_logits, add_samplers, accept_history); + + GGML_ASSERT(output.expected.size() == (size_t) k); + GGML_ASSERT(output.actual.logits.size() == (size_t) k); + + const std::unordered_map expected_logits = map_logits(output.expected); + + if (position == penalties_position::after_filter) { + GGML_ASSERT(expected_logits.find(retained_history_token) != expected_logits.end()); + GGML_ASSERT(fabsf(expected_logits.at(retained_history_token) - raw_logits[retained_history_token]) > 1e-6f); + GGML_ASSERT(expected_logits.find(excluded_history_token) == expected_logits.end()); + GGML_ASSERT(std::find(output.actual.candidates.begin(), output.actual.candidates.end(), + excluded_history_token) == output.actual.candidates.end()); + } else { + const std::unordered_map unpenalized_logits = map_logits(top_k_data); + bool changed = false; + for (const auto & item : expected_logits) { + const auto it = unpenalized_logits.find(item.first); + if (it == unpenalized_logits.end() || fabsf(it->second - item.second) > 1e-6f) { + changed = true; + break; + } + } + GGML_ASSERT(changed); + } + + const char * name = position == penalties_position::before_filter + ? "penalties top-k" + : "top-k penalties"; + const sampler_comparison_stats stats = compare_sampler_outputs( + name, expected_logits, output.actual); + GGML_ASSERT(stats.n_masked == 0); + GGML_ASSERT(stats.n_mismatch == 0); +} + +static void compare_masking_penalties_logits( + const test_params & params, + const char * filter_name, + const sampler_init_fn & init_filter, + int32_t penalty_last_n, + float penalty_repeat, + float penalty_freq, + float penalty_present, + const std::string & prompt, + penalties_position position, + bool allow_extra_candidates, + bool add_history = true) { + const auto * vocab = llama_model_get_vocab(params.model.get()); + const std::vector raw_logits = decode_raw_logits(params, prompt); + const int n_vocab = (int) raw_logits.size(); + llama_sampler_ptr filter(init_filter()); + const std::vector filtered_data = apply_cpu_sampler(raw_logits, filter.get()); + GGML_ASSERT(!filtered_data.empty()); + GGML_ASSERT(filtered_data.size() < (size_t) n_vocab); + + const llama_token penalized_token = filtered_data[0].id; + std::unordered_set retained_tokens; + retained_tokens.reserve(filtered_data.size()); + for (const auto & data : filtered_data) { + retained_tokens.insert(data.id); + } + + llama_token masked_token = LLAMA_TOKEN_NULL; + for (llama_token token = 0; token < n_vocab; ++token) { + if (retained_tokens.find(token) == retained_tokens.end()) { + masked_token = token; + break; + } + } + GGML_ASSERT(masked_token != LLAMA_TOKEN_NULL); + + const auto add_samplers = [&](llama_sampler * chain) { + add_filter_and_penalties(chain, init_filter, + penalty_last_n, penalty_repeat, penalty_freq, penalty_present, position); + }; + auto accept_history = [&](llama_sampler * smpl) { + if (!add_history) { + return; + } + accept_prompt(smpl, vocab, prompt); + llama_sampler_accept(smpl, penalized_token); + llama_sampler_accept(smpl, penalized_token); + llama_sampler_accept(smpl, masked_token); + llama_sampler_accept(smpl, masked_token); + }; + + const sampler_comparison_output output = run_sampler_comparison( + params, prompt, raw_logits, add_samplers, accept_history); + + GGML_ASSERT(output.actual.logits.size() == (size_t) n_vocab); + + const std::unordered_map expected_logits = map_logits(output.expected); + + GGML_ASSERT(expected_logits.find(masked_token) == expected_logits.end()); + if (add_history) { + if (position == penalties_position::after_filter) { + GGML_ASSERT(expected_logits.find(penalized_token) != expected_logits.end()); + GGML_ASSERT(fabsf(expected_logits.at(penalized_token) - raw_logits[penalized_token]) > 1e-6f); + } else { + llama_sampler_ptr penalties(llama_sampler_init_penalties( + penalty_last_n, penalty_repeat, penalty_freq, penalty_present)); + accept_history(penalties.get()); + const std::unordered_map penalized_logits = + map_logits(apply_cpu_sampler(raw_logits, penalties.get())); + GGML_ASSERT(fabsf(penalized_logits.at(penalized_token) - raw_logits[penalized_token]) > 1e-6f); + } + } + + const std::string name = position == penalties_position::before_filter + ? "penalties " + std::string(filter_name) + : std::string(filter_name) + " penalties"; + const sampler_comparison_stats stats = compare_sampler_outputs( + name.c_str(), expected_logits, output.actual, allow_extra_candidates); + const float masked_logit = find_backend_logit(output.actual, masked_token); + GGML_ASSERT(stats.n_masked > 0); + GGML_ASSERT(std::isinf(masked_logit) && masked_logit < 0.0f); + GGML_ASSERT(stats.n_mismatch == 0); +} + +static void test_backend_penalties_sampling(const test_params & params) { + printf("Testing backend penalties (repeat + freq + presence)\n"); + compare_penalties_logits(params, 64, 1.1f, 0.5f, 0.25f, "Hello Hello world"); + + printf("Testing backend penalties with penalty_last_n > 64\n"); + const auto * vocab = llama_model_get_vocab(params.model.get()); + std::vector tokens(8); + int32_t n_tok = llama_tokenize(vocab, "a", 1, tokens.data(), (int32_t) tokens.size(), false, false); + if (n_tok < 0) { + tokens.resize(-n_tok); + n_tok = llama_tokenize(vocab, "a", 1, tokens.data(), (int32_t) tokens.size(), false, false); + } + GGML_ASSERT(n_tok > 0); + const llama_token tok = tokens[0]; + + compare_penalties_logits(params, 80, 1.15f, 0.1f, 0.05f, "a", [tok](llama_sampler * smpl) { + // accept_prompt already accepted BOS + one 'a'; fill the ring to n=80 + for (int i = 0; i < 78; ++i) { + llama_sampler_accept(smpl, tok); + } + }); + + printf("Testing backend penalties without filler entries\n"); + compare_penalties_logits(params, 64, 1.1f, 0.5f, 0.25f, "Hello", [](llama_sampler * smpl) { + for (llama_token token = 0; token < 64; ++token) { + llama_sampler_accept(smpl, token); + } + }); + + printf("Testing backend top-k followed by penalties\n"); + compare_top_k_penalties_logits(params, 8, 64, 1.1f, 0.5f, 0.25f, "Hello", + penalties_position::after_filter); + + printf("Testing backend penalties followed by top-k\n"); + compare_top_k_penalties_logits(params, 8, 64, 1.1f, 0.5f, 0.25f, "Hello", + penalties_position::before_filter); + + printf("Testing backend top-p followed by penalties\n"); + compare_masking_penalties_logits(params, "top-p", []() { + return llama_sampler_init_top_p(0.9f, 0); + }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true); + + printf("Testing backend top-p followed by penalties with a large history window\n"); + compare_masking_penalties_logits(params, "top-p large-window", []() { + return llama_sampler_init_top_p(0.9f, 0); + }, 4096, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true); + + printf("Testing backend penalties followed by top-p\n"); + compare_masking_penalties_logits(params, "top-p", []() { + return llama_sampler_init_top_p(0.9f, 0); + }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::before_filter, true); + + printf("Testing backend min-p followed by penalties\n"); + compare_masking_penalties_logits(params, "min-p", []() { + return llama_sampler_init_min_p(0.1f, 0); + }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, false); + + printf("Testing backend penalties followed by min-p\n"); + compare_masking_penalties_logits(params, "min-p", []() { + return llama_sampler_init_min_p(0.1f, 0); + }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::before_filter, false); + + printf("Testing backend top-p followed by penalties with empty history\n"); + compare_masking_penalties_logits(params, "top-p empty", []() { + return llama_sampler_init_top_p(0.9f, 0); + }, 64, 1.1f, 0.5f, 0.25f, "Hello", penalties_position::after_filter, true, false); + + printf("Testing backend top-p followed by individual penalties\n"); + compare_masking_penalties_logits(params, "top-p repeat", []() { + return llama_sampler_init_top_p(0.9f, 0); + }, 64, 1.1f, 0.0f, 0.0f, "Hello", penalties_position::after_filter, true); + compare_masking_penalties_logits(params, "top-p frequency", []() { + return llama_sampler_init_top_p(0.9f, 0); + }, 64, 1.0f, 0.5f, 0.0f, "Hello", penalties_position::after_filter, true); + compare_masking_penalties_logits(params, "top-p presence", []() { + return llama_sampler_init_top_p(0.9f, 0); + }, 64, 1.0f, 0.0f, 0.25f, "Hello", penalties_position::after_filter, true); + + printf("Testing backend penalty parameter values\n"); + test_penalty_parameter_values(params); + + printf("backend penalties sampling test PASSED\n"); +} + // This test verifies that it is possible to have two different backend samplers, // one that uses the backend dist sampler, and another that uses CPU dist sampler. static void test_backend_mixed_sampling(const test_params & params) { @@ -1014,6 +1574,7 @@ struct backend_test_case { static const backend_test_case BACKEND_TESTS[] = { { "greedy", test_backend_greedy_sampling, true }, { "logit_bias", test_backend_logit_bias_sampling, true }, + { "penalties", test_backend_penalties_sampling, true }, { "temp", test_backend_temp_sampling, true }, { "temp_ext", test_backend_temp_ext_sampling, true }, { "top_k", test_backend_top_k_sampling, true }, diff --git a/tests/test-model-resolution.cpp b/tests/test-model-resolution.cpp new file mode 100644 index 0000000000..a96e40bb2b --- /dev/null +++ b/tests/test-model-resolution.cpp @@ -0,0 +1,506 @@ +// tests the HF model resolution and the model handler assembly end-to-end on +// synthetic repo listings: a local httplib server bound to the loopback +// serves hardcoded HF API responses, so the real client, hf_cache, resolution +// and CLI parsing run against them without external network access + +#include "arg.h" +#include "common.h" +#include "download.h" +#include "http.h" +#include "log.h" + +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +// the case and reordering being checked, printed with every failure +static std::string g_context; + +// independent of NDEBUG, so the checks stay alive in Release builds +#define REQUIRE(x) do { \ + if (!(x)) { \ + fprintf(stderr, "%s:%d: [%s] REQUIRE(%s) failed\n", \ + __FILE__, __LINE__, g_context.c_str(), #x); \ + std::abort(); \ + } \ +} while (0) + +#define REQUIRE_EQ(actual, expected) do { \ + if (!((actual) == (expected))) { \ + fprintf(stderr, "%s:%d: [%s] REQUIRE_EQ(%s, %s) failed\n actual: '%s'\n expected: '%s'\n", \ + __FILE__, __LINE__, g_context.c_str(), #actual, #expected, \ + std::string(actual).c_str(), std::string(expected).c_str()); \ + std::abort(); \ + } \ +} while (0) + +// +// synthetic repos keyed by repo id, served over the loopback by a real +// httplib server, so the tested code runs its own client and transport +// + +static std::map> g_repos; + +static const char * COMMIT = "aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"; + +// the server lives in main, so its destructor runs before the static teardown +// tears down the winsock state httplib brings in +static void serve_repos(httplib::Server & server) { + server.Get(R"(/api/models/(.+)/refs)", [](const httplib::Request & req, httplib::Response & res) { + if (g_repos.count(req.matches[1])) { + res.set_content(nlohmann::json{{"branches", {{{"name", "main"}, {"targetCommit", COMMIT}}}}}.dump(), + "application/json"); + } else { + res.status = 404; + } + }); + server.Get(R"(/api/models/(.+)/tree/.+)", [](const httplib::Request & req, httplib::Response & res) { + if (!g_repos.count(req.matches[1])) { + res.status = 404; + return; + } + auto files = nlohmann::json::array(); + size_t i = 0; + for (const auto & p : g_repos[req.matches[1]]) { + char oid[41]; + snprintf(oid, sizeof(oid), "%040lx", (unsigned long) ++i); + files.push_back({{"type", "file"}, {"path", p}, {"size", 1}, {"oid", oid}}); + } + res.set_content(files.dump(), "application/json"); + }); +} + +static common_params_model model_ref(const std::string & hf_repo, const std::string & hf_file = "") { + common_params_model m; + m.hf_repo = hf_repo; + m.hf_file = hf_file; + return m; +} + +// the model cache is isolated under a temporary directory named after the +// loopback port, so concurrent runs on a shared machine keep their own, and +// the local path the handler wires for a file is snapshots// +static std::filesystem::path cache_dir; + +static std::string cached(std::string repo_id, const std::string & path) { + string_replace_all(repo_id, "/", "--"); + return (cache_dir / ("models--" + repo_id) / "snapshots" / COMMIT / path).string(); +} + +// +// fixtures mimicking real repo layouts +// + +// flat layout in the style of ggml-org/gemma-4-31B-it-GGUF +static const std::vector flat = { + "README.md", + "model-BF16.gguf", + "model-Q4_K_M.gguf", + "model-Q8_0.gguf", + "mmproj-model-BF16.gguf", + "mmproj-model-Q8_0.gguf", + "mtp-model-BF16.gguf", + "mtp-model-Q4_0.gguf", + "mtp-model-Q8_0.gguf", + "dflash-model-BF16.gguf", + "dflash-model-Q8_0.gguf", +}; + +// quants in subdirectories with sharded files and root sidecars, +// in the style of stepfun-ai/Step-3.7-Flash-GGUF +static const std::vector subdir = { + "mmproj-model-f16.gguf", + "model-mtp-BF16.gguf", + "model-mtp-Q8_0.gguf", + "Q3_K_M/model-Q3_K_M-00001-of-00003.gguf", + "Q3_K_M/model-Q3_K_M-00002-of-00003.gguf", + "Q3_K_M/model-Q3_K_M-00003-of-00003.gguf", + "Q8_0/model-Q8_0-00001-of-00002.gguf", + "Q8_0/model-Q8_0-00002-of-00002.gguf", +}; + +// sidecar quants exist where the full model quant does not, +// in the style of ggml-org/Qwen3.6-27B-GGUF +static const std::vector hole = { + "model-BF16.gguf", + "model-Q4_K_M.gguf", + "model-Q8_0.gguf", + "mtp-model-BF16.gguf", + "mtp-model-Q4_0.gguf", + "mtp-model-Q8_0.gguf", + "dflash-model-BF16.gguf", + "dflash-model-Q8_0.gguf", +}; + +// unsloth-style naming with UD quants and a suffix MTP file +static const std::vector unsloth = { + "model-UD-Q8_K_XL.gguf", + "mmproj-BF16.gguf", + "model-MTP-BF16.gguf", +}; + +// bartowski-style vendor prefix and mradermacher-style dot quant +static const std::vector vendors = { + "TheDrummer_Model-24B-v4.1-Q8_0.gguf", + "BlackSheep-24B.Q8_0.gguf", +}; + +// every speculative sidecar type at the same quant +static const std::vector quad = { + "model-Q8_0.gguf", + "mtp-model-Q8_0.gguf", + "dflash-model-Q8_0.gguf", + "eagle3-model-Q8_0.gguf", + "dspark-model-Q8_0.gguf", +}; + +static const std::vector dflash_only = { + "model-Q8_0.gguf", + "dflash-model-Q8_0.gguf", +}; + +static const std::vector eagle3_only = { + "model-Q8_0.gguf", + "eagle3-model-Q8_0.gguf", +}; + +// a single full quant with dspark sidecars at other quants, +// in the style of ggml-org/DeepSeek-V4-Flash-0731-GGUF +static const std::vector spark = { + "README.md", + "model-MXFP4.gguf", + "dspark-model-BF16.gguf", + "dspark-model-MXFP4.gguf", +}; + +// dspark outranks dflash in the type auto-selection +static const std::vector dspark_dflash = { + "model-Q8_0.gguf", + "dflash-model-Q8_0.gguf", + "dspark-model-Q8_0.gguf", +}; + +// +// table-driven plan resolution through the real entry point, +// each case replayed on multiple deterministic reorderings of the listing, +// except the cases whose pick legitimately depends on the listing order +// + +struct plan_case { + const char * name; + const std::vector & files; + const char * hf_repo; + const char * hf_file; + bool sidecars; // request mmproj + mtp + dflash + eagle3 + dspark + bool order_dependent; // the expected pick depends on the listing order + const char * primary; + std::vector model_files; + const char * mmproj; + const char * mtp; + const char * dflash; + const char * eagle3; + const char * dspark; +}; + +static const plan_case plan_cases[] = { + // exact tag picks the matching primary, sidecars follow the tag + {"flat exact tag", flat, "test/repo:Q8_0", "", true, false, + "model-Q8_0.gguf", {"model-Q8_0.gguf"}, + "mmproj-model-Q8_0.gguf", "mtp-model-Q8_0.gguf", "dflash-model-Q8_0.gguf", "", ""}, + + // no tag falls back to the default quant preference + {"flat default", flat, "test/repo", "", false, false, + "model-Q4_K_M.gguf", {"model-Q4_K_M.gguf"}, + "", "", "", "", ""}, + + // no tag and no default match falls back to the first model in the listing + {"unsloth fallback", unsloth, "test/repo", "", true, true, + "model-UD-Q8_K_XL.gguf", {"model-UD-Q8_K_XL.gguf"}, + "mmproj-BF16.gguf", "", "", "", ""}, + + // explicit hf_file picks that exact file + {"flat hf_file", flat, "test/repo", "model-BF16.gguf", false, false, + "model-BF16.gguf", {"model-BF16.gguf"}, + "", "", "", "", ""}, + + // missing hf_file resolves nothing + {"flat missing hf_file", flat, "test/repo", "nope.gguf", false, false, + "", {}, + "", "", "", "", ""}, + + // a sharded primary brings all its parts, a subdir primary finds the root sidecar + {"subdir shards", subdir, "test/repo:Q3_K_M", "", true, false, + "Q3_K_M/model-Q3_K_M-00001-of-00003.gguf", + {"Q3_K_M/model-Q3_K_M-00001-of-00003.gguf", + "Q3_K_M/model-Q3_K_M-00002-of-00003.gguf", + "Q3_K_M/model-Q3_K_M-00003-of-00003.gguf"}, + "mmproj-model-f16.gguf", "model-mtp-Q8_0.gguf", "", "", ""}, + + // a tag with no matching full model still resolves the requested sidecars + {"hole tag sidecar", hole, "test/repo:Q4_0", "", true, false, + "", {}, + "", "mtp-model-Q4_0.gguf", "dflash-model-Q8_0.gguf", "", ""}, + + // the same tag without a requested sidecar resolves nothing + {"hole tag alone", hole, "test/repo:Q4_0", "", false, false, + "", {}, + "", "", "", "", ""}, + + // no tag anchors the sidecars on the primary quant + {"hole default anchor", hole, "test/repo", "", true, false, + "model-Q4_K_M.gguf", {"model-Q4_K_M.gguf"}, + "", "mtp-model-Q4_0.gguf", "dflash-model-Q8_0.gguf", "", ""}, + + // the mtp- keyword is case sensitive, a suffix -MTP file is not discovered + {"unsloth suffix mtp", unsloth, "test/repo:Q8_K_XL", "", true, false, + "model-UD-Q8_K_XL.gguf", {"model-UD-Q8_K_XL.gguf"}, + "mmproj-BF16.gguf", "", "", "", ""}, + + // vendor prefixes and the dot quant convention both match the tag, + // first match wins between two files at the same quant + {"vendor prefix", vendors, "test/repo:Q8_0", "", false, true, + "TheDrummer_Model-24B-v4.1-Q8_0.gguf", {"TheDrummer_Model-24B-v4.1-Q8_0.gguf"}, + "", "", "", "", ""}, + + // every sidecar type resolves at the tag + {"quad exact tag", quad, "test/repo:Q8_0", "", true, false, + "model-Q8_0.gguf", {"model-Q8_0.gguf"}, + "", "mtp-model-Q8_0.gguf", "dflash-model-Q8_0.gguf", "eagle3-model-Q8_0.gguf", "dspark-model-Q8_0.gguf"}, + + // no tag anchors the dspark sidecar on the only full quant + {"spark default anchor", spark, "test/repo", "", true, false, + "model-MXFP4.gguf", {"model-MXFP4.gguf"}, + "", "", "", "", "dspark-model-MXFP4.gguf"}, + + // a tag with no matching full model still resolves the exact dspark sidecar + {"spark tag sidecar", spark, "test/repo:BF16", "", true, false, + "", {}, + "", "", "", "", "dspark-model-BF16.gguf"}, +}; + +static void check_plan(const plan_case & c) { + common_download_opts opts; + opts.download_mmproj = c.sidecars; + opts.download_mtp = c.sidecars; + opts.download_dflash = c.sidecars; + opts.download_eagle3 = c.sidecars; + opts.download_dspark = c.sidecars; + + auto plan = common_download_get_hf_plan(model_ref(c.hf_repo, c.hf_file), opts); + + REQUIRE_EQ(plan.primary.path, c.primary); + REQUIRE_EQ(plan.mmproj.path, c.mmproj); + REQUIRE_EQ(plan.mtp.path, c.mtp); + REQUIRE_EQ(plan.dflash.path, c.dflash); + REQUIRE_EQ(plan.eagle3.path, c.eagle3); + REQUIRE_EQ(plan.dspark.path, c.dspark); + + // exact shard set, order insensitive; the primary must be the first split + std::vector actual; + for (const auto & f : plan.model_files) { + actual.push_back(f.path); + } + std::sort(actual.begin(), actual.end()); + auto expected = c.model_files; + std::sort(expected.begin(), expected.end()); + REQUIRE(actual == expected); + if (!expected.empty()) { + REQUIRE(plan.primary.path == expected.front()); + } +} + +static void test_plan_resolution() { + printf("test-model-resolution: plan resolution on %zu cases\n", sizeof(plan_cases) / sizeof(plan_cases[0])); + + for (const auto & c : plan_cases) { + printf(" %s\n", c.name); + // invariant: the resolution is insensitive to the listing order + for (size_t rot = 0; rot < c.files.size(); ++rot) { + if (c.order_dependent && rot > 0) { + continue; + } + g_context = std::string(c.name) + ", reordering " + std::to_string(rot); + auto files = c.files; + std::rotate(files.begin(), files.begin() + rot, files.end()); + if (rot % 2 == 1) { + std::reverse(files.begin(), files.end()); + } + g_repos["test/repo"] = files; + check_plan(c); + } + } + g_repos.clear(); +} + +// +// end-to-end assembly: real CLI parsing, real handler init resolving over the +// loopback, downloads skipped by flipping offline before apply +// + +static void assemble(std::vector argv, common_params & params) { + std::vector cargv; + g_context.clear(); + for (auto & a : argv) { + g_context += g_context.empty() ? a : " " + a; + cargv.push_back(a.data()); + } + bool ok = common_params_parse((int) cargv.size(), cargv.data(), params, LLAMA_EXAMPLE_SERVER); + REQUIRE(ok); + + auto handler = common_models_handler_init(params, LLAMA_EXAMPLE_SERVER); + + // skip the network execution, on_done still wires the params + params.offline = true; + common_models_handler_apply(handler, params); +} + +static void test_task_assembly() { + printf("test-model-resolution: end-to-end assembly\n"); + + g_repos["test/main"] = flat; + g_repos["test/hole"] = hole; + g_repos["test/quad"] = quad; + g_repos["test/dflash"] = dflash_only; + g_repos["test/eagle3"] = eagle3_only; + g_repos["test/spark"] = spark; + g_repos["test/pair"] = dspark_dflash; + g_repos["test/small"] = {"draft-model-Q4_K_M.gguf"}; + g_repos["test/preset"] = {"preset.ini", "model-Q8_0.gguf"}; + + { + // plain -hf wires the model and its mmproj, nothing speculative + common_params params; + assemble({"server", "-hf", "test/main:Q8_0"}, params); + REQUIRE_EQ(params.model.path, cached("test/main", "model-Q8_0.gguf")); + REQUIRE_EQ(params.mmproj.path, cached("test/main", "mmproj-model-Q8_0.gguf")); + REQUIRE(params.speculative.draft.mparams.path.empty()); + } + { + // --no-mmproj disables the mmproj discovery + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "--no-mmproj"}, params); + REQUIRE(params.mmproj.path.empty()); + } + { + // an explicit --mmproj wins over the discovery + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "--mmproj", "/local/mmproj.gguf"}, params); + REQUIRE(params.mmproj.path == "/local/mmproj.gguf"); + } + { + // -hf with a spec type wires the sidecar of the main repo as fallback draft + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "--spec-type", "draft-mtp"}, params); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/main", "mtp-model-Q8_0.gguf")); + } + { + // -hfd with a spec type wires the draft repo sidecar at its tag, + // not its full model, and suppresses the main repo fallback + common_params params; + assemble({"server", "-hf", "test/hole:Q8_0", "-hfd", "test/hole:Q4_0", "--spec-type", "draft-mtp"}, params); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/hole", "mtp-model-Q4_0.gguf")); + } + { + // an explicit -md file wins over the sidecar resolution + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/main", "-md", "mtp-model-BF16.gguf", "--spec-type", "draft-mtp"}, params); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/main", "mtp-model-BF16.gguf")); + } + { + // -hfd without a spec type auto-selects the type, mtp first when all ship + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/quad:Q8_0"}, params); + REQUIRE(params.speculative.types == std::vector{COMMON_SPECULATIVE_TYPE_DRAFT_MTP}); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/quad", "mtp-model-Q8_0.gguf")); + } + { + // auto-selection with only a dflash sidecar + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/dflash:Q8_0"}, params); + REQUIRE(params.speculative.types == std::vector{COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH}); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/dflash", "dflash-model-Q8_0.gguf")); + } + { + // auto-selection with only an eagle3 sidecar + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/eagle3:Q8_0"}, params); + REQUIRE(params.speculative.types == std::vector{COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3}); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/eagle3", "eagle3-model-Q8_0.gguf")); + } + { + // auto-selection prefers dspark over dflash when both ship + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/pair:Q8_0"}, params); + REQUIRE(params.speculative.types == std::vector{COMMON_SPECULATIVE_TYPE_DRAFT_DSPARK}); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/pair", "dspark-model-Q8_0.gguf")); + } + { + // -hf with the dspark spec type wires the sidecar of the main repo, + // anchored on the only full quant + common_params params; + assemble({"server", "-hf", "test/spark", "--spec-type", "draft-dspark"}, params); + REQUIRE_EQ(params.model.path, cached("test/spark", "model-MXFP4.gguf")); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/spark", "dspark-model-MXFP4.gguf")); + } + { + // -hfd on a repo without sidecars keeps resolving a full model as draft + common_params params; + assemble({"server", "-hf", "test/main:Q8_0", "-hfd", "test/small"}, params); + REQUIRE(params.speculative.types == std::vector{COMMON_SPECULATIVE_TYPE_NONE}); + REQUIRE_EQ(params.speculative.draft.mparams.path, cached("test/small", "draft-model-Q4_K_M.gguf")); + } + { + // a preset repo wires the preset and clears the model for router mode + common_params params; + assemble({"server", "-hf", "test/preset"}, params); + REQUIRE_EQ(params.models_preset, cached("test/preset", "preset.ini")); + REQUIRE(params.model.path.empty()); + REQUIRE(params.model.hf_repo.empty()); + } + + g_repos.clear(); +} + +int main(void) { + // unbuffered, so a crash cannot swallow the reports already printed + setvbuf(stdout, nullptr, _IONBF, 0); + setvbuf(stderr, nullptr, _IONBF, 0); + + // the negative cases legitimately log errors on every reordering, + // keep the output down to the reports + common_log_pause(common_log_main()); + + // the loopback endpoint also keeps the client init from rejecting + // https on the builds without TLS support + httplib::Server server; + serve_repos(server); + int port = server.bind_to_any_port("127.0.0.1"); + + // isolate the cache, its location is read once so it is set + // before anything else + cache_dir = std::filesystem::temp_directory_path() / + ("test-model-resolution-cache-" + std::to_string(port)); + std::filesystem::remove_all(cache_dir); + common_set_env("LLAMA_CACHE", cache_dir.string()); + + std::thread server_thread([&server] { server.listen_after_bind(); }); + server.wait_until_ready(); + common_set_env("MODEL_ENDPOINT", "http://127.0.0.1:" + std::to_string(port) + "/"); + + test_plan_resolution(); + test_task_assembly(); + + server.stop(); + server_thread.join(); + + std::filesystem::remove_all(cache_dir); + printf("test-model-resolution: all tests OK\n"); + return 0; +} diff --git a/tools/server/README-dev.md b/tools/server/README-dev.md index b41d70c63a..45bcdcca76 100644 --- a/tools/server/README-dev.md +++ b/tools/server/README-dev.md @@ -199,6 +199,9 @@ Invoke a tool call, request body is a JSON object with: - `tool` (string): the name of the tool - `params` (object): a mapping from argument name (string) to argument value +Headers: +- `x-tool-cwd`: optional; if set, use as the CWD for tool; this is not part of tool's params because it's meant to be set by the runtime, not the LLM itself + Returns JSON object. There are two response formats (MCP tools use the same two formats: their result content is concatenated into `plain_text_response`, and RPC or tool errors are surfaced as the `error` string): Format 1: Plain text. The text will be placed into a field called `plain_text_response`, example: diff --git a/tools/server/README.md b/tools/server/README.md index f45c018972..a0956f9e65 100644 --- a/tools/server/README.md +++ b/tools/server/README.md @@ -198,7 +198,9 @@ For the full list of features, please refer to [server's changelog](https://gith | `--ui-config, --webui-config JSON` | JSON that provides default UI settings (overrides UI defaults)
(env: LLAMA_ARG_UI_CONFIG) | | `--ui-config-file, --webui-config-file PATH` | JSON file that provides default UI settings (overrides UI defaults)
(env: LLAMA_ARG_UI_CONFIG_FILE) | | `--ui-mcp-proxy, --webui-mcp-proxy, --no-ui-mcp-proxy, --no-webui-mcp-proxy` | experimental: whether to enable MCP CORS proxy - do not enable in untrusted environments (default: disabled)
(env: LLAMA_ARG_UI_MCP_PROXY) | -| `--tools TOOL1,TOOL2,...` | experimental: whether to enable built-in tools for AI agents - do not enable in untrusted environments (default: no tools)
specify "all" to enable all tools
available tools: read_file, file_glob_search, grep_search, exec_shell_command, write_file, edit_file, get_datetime
note: for security reasons, this will limit --cors-origins to localhost by default
(env: LLAMA_ARG_TOOLS) | +| `--tools TOOL1,TOOL2,...` | experimental: whether to enable built-in tools for AI agents - do not enable in untrusted environments (default: no tools)
specify "all" to enable all tools
available tools: read_file, file_glob_search, grep_search, exec_shell_command, write_file, edit_file, get_datetime, get_info
note: for security reasons, this will limit --cors-origins to localhost by default
(env: LLAMA_ARG_TOOLS) | +| `--mcp-servers-config PATH` | experimental: path to JSON file with MCP server definitions (Cursor-compatible format) - do not enable in untrusted environments (default: none)
note: for security reasons, this will limit --cors-origins to localhost by default
(env: LLAMA_ARG_MCP_SERVERS_CONFIG) | +| `--mcp-servers-json JSON` | experimental: inline JSON with MCP server definitions (Cursor-compatible format) - do not enable in untrusted environments (default: none)
note: for security reasons, this will limit --cors-origins to localhost by default
(env: LLAMA_ARG_MCP_SERVERS_JSON) | | `-ag, --agent, -no-ag, --no-agent` | whether to enable CORS proxy and all built-in tools - do not enable in untrusted environments (default: disabled)
note: for security reasons, this will limit --cors-origins to localhost by default
(env: LLAMA_ARG_AGENT) | | `--ui, --webui, --no-ui, --no-webui` | whether to enable the Web UI (default: enabled)
(env: LLAMA_ARG_UI) | | `--embedding, --embeddings` | restrict to only support embedding use case; use only with dedicated embedding models (default: disabled)
(env: LLAMA_ARG_EMBEDDINGS) | diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 4655b518e2..5d2798cc14 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 90b7e2a9f0..984bb478ea 100644 --- a/tools/server/server-tools.cpp +++ b/tools/server/server-tools.cpp @@ -64,24 +64,27 @@ public: class tools_io_basic : public tools_io { public: + // cwd, if non-empty, is used to resolve relative paths and as the working directory for run() + explicit tools_io_basic(std::string cwd = "") : cwd(std::move(cwd)) {} + bool is_directory(const std::string & path) const override { std::error_code ec; - return fs::is_directory(path, ec) && !ec; + return fs::is_directory(resolve(path), ec) && !ec; } bool is_regular_file(const std::string & path) const override { std::error_code ec; - return fs::is_regular_file(path, ec) && !ec; + return fs::is_regular_file(resolve(path), ec) && !ec; } bool file_size(const std::string & path, uintmax_t & out_size) const override { std::error_code ec; - out_size = fs::file_size(path, ec); + out_size = fs::file_size(resolve(path), ec); return !ec; } bool read_file(const std::string & path, std::string & out) const override { - std::ifstream f(path, std::ios::binary); + std::ifstream f(resolve(path), std::ios::binary); if (!f) return false; std::ostringstream ss; ss << f.rdbuf(); @@ -91,12 +94,12 @@ public: bool write_file(const std::string & path, const std::string & content) const override { std::error_code ec; - fs::path fpath(path); + fs::path fpath(resolve(path)); if (fpath.has_parent_path()) { fs::create_directories(fpath.parent_path(), ec); if (ec) return false; } - std::ofstream f(path, std::ios::binary); + std::ofstream f(fpath, std::ios::binary); if (!f) return false; f << content; return (bool) f; @@ -104,13 +107,14 @@ public: std::vector list_files(const std::string & base, std::string & err) const override { err.clear(); + std::string abs_base = resolve(base); if (!is_directory(base)) { err = "path does not exist or is not a directory: " + base; return {}; } auto res = run( - {"git", "-C", base, "ls-files", "--cached", "--others", "--exclude-standard"}, + {"git", "-C", abs_base, "ls-files", "--cached", "--others", "--exclude-standard"}, SERVER_TOOL_GIT_LS_FILES_MAX_OUTPUT, SERVER_TOOL_GIT_LS_FILES_TIMEOUT); if (res.exit_code == 0 && !res.timed_out) { @@ -128,7 +132,7 @@ public: return result; } - return list_files_fallback(base); + return list_files_fallback(abs_base); } exec_result run( @@ -145,7 +149,7 @@ public: | subprocess_option_inherit_environment | subprocess_option_search_user_path; - if (!proc.create(args, options)) { + if (!proc.create(args, options, {}, cwd.empty() ? nullptr : cwd.c_str())) { res.output = "failed to spawn process"; return res; } @@ -205,6 +209,16 @@ public: } private: + std::string cwd; + + // resolves `path` against `cwd` if `path` is relative and `cwd` is set; otherwise returns `path` unchanged + std::string resolve(const std::string & path) const { + if (cwd.empty() || fs::path(path).is_absolute()) { + return path; + } + return (fs::path(cwd) / path).string(); + } + static const std::unordered_set & junk_dir_names() { static const std::unordered_set names = { ".git", ".svn", ".hg", "node_modules", "__pycache__", @@ -244,8 +258,8 @@ private: }; static std::unique_ptr make_tools_io(const json & params) { - GGML_UNUSED(params); // TODO in follow-up PR - return std::make_unique(); + std::string cwd = json_value(params, "cwd", std::string()); + return std::make_unique(cwd); } // no '/' in pattern -> match basename at any depth; else match full relative path @@ -1076,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; @@ -1185,9 +1249,26 @@ 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; } +static std::string str_to_lower(const std::string & value) { + std::string lowered(value.size(), '\0'); + std::transform(value.begin(), value.end(), lowered.begin(), [](unsigned char c) { return std::tolower(c); }); + return lowered; +} + +static std::string get_header(const std::map & headers, const std::string & key, std::string default_value = "") { + const auto lowered_key = str_to_lower(key); + for (const auto & h : headers) { + if (str_to_lower(h.first) == lowered_key) { + return h.second; + } + } + return default_value; +} + void server_tools::setup(const std::vector & enabled_tools, server_mcp & mcp_mgr) { if (!enabled_tools.empty()) { @@ -1271,6 +1352,12 @@ void server_tools::setup(const std::vector & enabled_tools, json params = body.value("params", json::object()); bool stream = body.value("stream", false); + // accept x-tool-cwd header to override of the process + auto cwd = get_header(req.headers, "x-tool-cwd"); + if (!cwd.empty()) { + params["cwd"] = cwd; + } + server_tool & tool = find_tool(tools, tool_name, stream); if (stream) { diff --git a/tools/server/server.cpp b/tools/server/server.cpp index a3b2a8b0fe..aafb1f3079 100644 --- a/tools/server/server.cpp +++ b/tools/server/server.cpp @@ -486,6 +486,13 @@ int llama_server(common_params & params, int argc, char ** argv) { SRV_INF("listening on %s\n", ctx_http.listening_address.c_str()); + // TODO: remove this in the future + // check the string to also handle the .sock case + if (string_ends_with(ctx_http.listening_address, ":8080")) { + SRV_WRN("%s", "NOTICE: server default port will be changed to :9931 in a future release\n"); + SRV_WRN("%s", " ref: https://github.com/ggml-org/llama.cpp/pull/26508\n"); + } + if (is_router_server) { if (!params.models_preset_hf.empty()) { SRV_WRN( "NOTE: using preset.ini from HF repo '%s'\n", params.models_preset_hf.c_str()); diff --git a/tools/server/tests/unit/test_tools_builtin.py b/tools/server/tests/unit/test_tools_builtin.py index 1b2d0db432..fb194cac66 100755 --- a/tools/server/tests/unit/test_tools_builtin.py +++ b/tools/server/tests/unit/test_tools_builtin.py @@ -19,8 +19,8 @@ def create_server(): server.server_tools = "all" -def call_tool(name: str, params: dict) -> dict: - res = server.make_request("POST", "/tools", data={"tool": name, "params": params}) +def call_tool(name: str, params: dict, headers: dict | None = None) -> dict: + res = server.make_request("POST", "/tools", data={"tool": name, "params": params}, headers=headers) assert res.status_code == 200, res.body assert "error" not in res.body, res.body return res.body @@ -123,6 +123,29 @@ def test_tools_builtin_exec_shell_command_stream(): assert "[exit code: 0]" in chunks +def test_tools_builtin_cwd_header(): + global server + server.start() + + cwd_dir = os.path.join(PROJECT_ROOT, "tools", "server", "tests", "unit") + headers = {"x-tool-cwd": cwd_dir} + + res = call_tool("read_file", {"path": "test_tools_builtin.py"}, headers=headers) + assert GREP_MARKER in res["plain_text_response"] + + # exec_shell_command should also run with that directory as its working directory: + # writing to a relative filename must land inside cwd_dir + marker_name = "llama_cpp_test_tools_builtin_cwd_marker.txt" + marker_path = os.path.join(cwd_dir, marker_name) + try: + command = f"echo hello > {marker_name}" + call_tool("exec_shell_command", {"command": command}, headers=headers) + assert os.path.exists(marker_path) + finally: + if os.path.exists(marker_path): + os.remove(marker_path) + + def test_tools_builtin_edit_file_rejects_overlapping_edits(): global server server.start() diff --git a/vendor/cpp-httplib/CMakeLists.txt b/vendor/cpp-httplib/CMakeLists.txt index bf674583ff..49c419094a 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}")