From 490c4195045cbd768d954646d3cb1d5658322654 Mon Sep 17 00:00:00 2001 From: Ryandito Diandaru Date: Sat, 29 Aug 2026 14:34:19 +0000 Subject: [PATCH 01/22] model: K2 Horizon gguf conversion code --- conversion/__init__.py | 1 + conversion/k2_horizon.py | 196 +++++++++++++++++++++++++++++++++ gguf-py/gguf/constants.py | 44 ++++++++ gguf-py/gguf/gguf_writer.py | 6 + gguf-py/gguf/tensor_mapping.py | 9 ++ 5 files changed, 256 insertions(+) create mode 100644 conversion/k2_horizon.py diff --git a/conversion/__init__.py b/conversion/__init__.py index a5632fcc4bb9..a8af0831a878 100644 --- a/conversion/__init__.py +++ b/conversion/__init__.py @@ -134,6 +134,7 @@ "JinaBertForMaskedLM": "bert", "JinaBertModel": "bert", "JinaEmbeddingsV5Model": "bert", + "K2HorizonForCausalLM": "k2_horizon", "KORMoForCausalLM": "qwen", "KimiK25ForConditionalGeneration": "deepseek", "KimiK3ForConditionalGeneration": "kimi_k3", diff --git a/conversion/k2_horizon.py b/conversion/k2_horizon.py new file mode 100644 index 000000000000..d3f5af9236ba --- /dev/null +++ b/conversion/k2_horizon.py @@ -0,0 +1,196 @@ +from __future__ import annotations + +import re +from typing import Iterable + +import torch +from torch import Tensor + +from .base import ModelBase, TextModel, gguf + +@ModelBase.register("K2HorizonForCausalLM") +@ModelBase.example( + "IFM/K2-Horizon-0.9B", + "IFM/K2-Horizon-36B", +) +class K2HorizonModel(TextModel): + model_arch = gguf.MODEL_ARCH.K2HORIZON + + def set_gguf_parameters(self): + super().set_gguf_parameters() + # generic + rope_head_dim = self.hparams.get("rope_head_dim") + norm_groups = int(self.hparams.get("layernorm_num_groups", 1)) + + self.gguf_writer.add_group_norm_groups(norm_groups) + if rope_head_dim is not None: + self.gguf_writer.add_rope_dimension_count(int(rope_head_dim)) + + # moe + num_experts = int(self.hparams.get("num_experts", 0)) + if num_experts > 0: + moe_ff = int(self.hparams["moe_intermediate_size"]) + dense_layers = self.hparams.get("num_dense_layers") + mlp_only_layers = {int(layer) for layer in self.hparams.get("mlp_only_layers", [])} + sparse_step = int(self.hparams.get("decoder_sparse_step", 1)) + shared_experts = int(self.hparams.get("num_shared_experts", 0)) + router_scale = self.hparams.get("router_scaling_factor") + normalize_topk = bool(self.hparams.get("norm_topk_prob", False)) + router_func = self.hparams.get("router_score_func") + + if dense_layers is None: + dense_layers = 0 + while dense_layers in mlp_only_layers: + dense_layers += 1 + + self.gguf_writer.add_expert_feed_forward_length(moe_ff) + self.gguf_writer.add_leading_dense_block_count(dense_layers) + self.gguf_writer.add_moe_every_n_layers(sparse_step) + self.gguf_writer.add_expert_shared_count(shared_experts) + self.gguf_writer.add_expert_weights_norm(normalize_topk) + if shared_experts > 0: + self.gguf_writer.add_expert_shared_feed_forward_length(moe_ff * shared_experts) + if router_scale is not None: + self.gguf_writer.add_expert_weights_scale(float(router_scale)) + match router_func: + case "sigmoid": + gating_func = gguf.ExpertGatingFuncType.SIGMOID + case "softmax": + gating_func = gguf.ExpertGatingFuncType.SOFTMAX + case _: + raise ValueError(f"Unsupported router_score_func: {router_func!r}") + self.gguf_writer.add_expert_gating_func(gating_func) + + # mova + value_experts = int(self.hparams.get("mova_num_experts", 0)) + value_experts_used = int(self.hparams.get("mova_num_experts_per_tok", 0)) + + if value_experts > 0 and value_experts_used > 0: + assert value_experts_used <= value_experts + self.gguf_writer.add_attention_value_expert_count(value_experts) + self.gguf_writer.add_attention_value_expert_used_count(value_experts_used) + + # gate func, only making sure it exists and is softplus + gate_func = self.hparams.get("attention_gate_func") + if gate_func not in (None, "softplus"): + raise ValueError(f"Unsupported attention_gate_func: {gate_func!r}") + + _experts: list[dict[str, Tensor]] | None = None + _value_experts: list[dict[str, Tensor]] | None = None + def modify_tensors( + self, + data_torch: Tensor, + name: str, + bid: int | None + ) -> Iterable[tuple[str, Tensor]]: + # MoE: router + if name.endswith(".mlp.gate.bias"): + assert bid is not None + yield ( + self.format_tensor_name( + gguf.MODEL_TENSOR.FFN_EXP_PROBS_B, + bid, + ".bias" + ), + data_torch + ) + return + + # MoE: actual up down or gate + is_moe_tensor = re.fullmatch(r"model\.layers\.\d+\.mlp\.experts\.\d+\.(down_proj|gate_proj|up_proj)\.weight", name) + if is_moe_tensor: + assert bid is not None + num_experts = int(self.hparams["num_experts"]) + + # allocate on first layer that has experts + if self._experts is None: + self._experts = [{} for _ in range(self.block_count)] + + # atp, this_blocks_experts contains all experts + this_blocks_experts = self._experts[bid] + this_blocks_experts[name] = data_torch + + # filling up self._experts until up down gate are all inside, then continue + if len(this_blocks_experts) < num_experts * 3: + return + + for projection in ("down_proj", "gate_proj", "up_proj"): + tensors = [] + for expert_id in range(num_experts): + expert_name = f"model.layers.{bid}.mlp.experts.{expert_id}.{projection}.weight" + tensors.append(this_blocks_experts.pop(expert_name)) + merged = torch.stack(tensors, dim=0) + merged_name = f"model.layers.{bid}.mlp.experts.{projection}.weight" + yield from super().modify_tensors( + merged, + merged_name, + bid + ) + return + + # MoVA + is_mova_weights = re.fullmatch(r"model\.layers\.\d+\.self_attn\.v_experts\.\d+\.weight", name) + if is_mova_weights: + assert bid is not None + num_value_experts = int(self.hparams["mova_num_experts"]) + if self._value_experts is None: + self._value_experts = [{} for _ in range(self.block_count)] + + this_blocks_value_expert = self._value_experts[bid] + this_blocks_value_expert[name] = data_torch + + # no need to * 3 because no up down gate like normal moe + if len(this_blocks_value_expert) < num_value_experts: + return + + tensors = [] + for value_exp_id in range(num_value_experts): + value_exp_name = f"model.layers.{bid}.self_attn.v_experts.{value_exp_id}.weight" + tensors.append(this_blocks_value_expert.pop(value_exp_name)) + + merged = torch.stack(tensors, dim = 0) + merged_name = f"model.layers.{bid}.self_attn.v_experts.weight" + yield from super().modify_tensors( + merged, + merged_name, + bid + ) + return + + # fallback, the default way basically + yield from super().modify_tensors( + data_torch, + name, + bid + ) + + def prepare_tensors(self): + super().prepare_tensors() + + # this is just checks basically + if self._experts is not None: + remaining_experts = [ + name + for block in self._experts + for name in block + ] + + if remaining_experts: + raise ValueError( + f"Unprocessed MoE experts: {remaining_experts}" + ) + + if self._value_experts is not None: + remaining_value_experts = [ + name + for block in self._value_experts + for name in block + ] + + if remaining_value_experts: + raise ValueError( + "Unprocessed MoVA value experts: " + f"{remaining_value_experts}" + ) + + diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index c99feb3c795c..13d5567384b0 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -217,6 +217,8 @@ class Attention: SLIDING_WINDOW_PATTERN = "{arch}.attention.sliding_window_pattern" TEMPERATURE_SCALE = "{arch}.attention.temperature_scale" ROPE_PATTERN = "{arch}.attention.rope_pattern" + VALUE_EXPERT_COUNT = "{arch}.attention.value_expert_count" + VALUE_EXPERT_USED_COUNT = "{arch}.attention.value_expert_used_count" class Indexer: HEAD_COUNT = "{arch}.attention.indexer.head_count" @@ -625,6 +627,7 @@ class MODEL_ARCH(IntEnum): NANBEIGE = auto() QWEN3TTS = auto() POCKETTTS = auto() + K2HORIZON = auto() class VISION_PROJECTOR_TYPE(IntEnum): @@ -888,6 +891,9 @@ class MODEL_TENSOR(IntEnum): INDEXER_COMPRESSOR_WGATE = auto() INDEXER_COMPRESSOR_APE = auto() INDEXER_COMPRESSOR_NORM = auto() + ATTN_V_GATE = auto() # K2Horizon + ATTN_V_EXP = auto() # K2Horizon + # vision V_MMPROJ = auto() V_MMPROJ_FC = auto() @@ -1374,6 +1380,7 @@ class MODEL_TENSOR(IntEnum): MODEL_ARCH.NANBEIGE: "nanbeige", MODEL_ARCH.QWEN3TTS: "qwen3tts", MODEL_ARCH.POCKETTTS: "pockettts", + MODEL_ARCH.K2HORIZON: "k2-horizon", } VISION_PROJECTOR_TYPE_NAMES: dict[VISION_PROJECTOR_TYPE, str] = { @@ -1964,6 +1971,10 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.DFLASH_SELECTOR_NEXT: "selector_successor", MODEL_TENSOR.DFLASH_SELECTOR_HIDDEN: "selector_hidden", MODEL_TENSOR.D2T: "d2t", + # K2 Horizon + MODEL_TENSOR.ATTN_V_GATE: "blk.{bid}.attn_v_gate", + MODEL_TENSOR.ATTN_V_EXP: "blk.{bid}.attn_v_exps", + } MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = { @@ -5374,6 +5385,39 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.FFN_DOWN, MODEL_TENSOR.FFN_UP, ], + MODEL_ARCH.K2HORIZON: [ + MODEL_TENSOR.TOKEN_EMBD, + MODEL_TENSOR.OUTPUT_NORM, + MODEL_TENSOR.OUTPUT, + MODEL_TENSOR.ATTN_NORM, + MODEL_TENSOR.ATTN_Q, + MODEL_TENSOR.ATTN_Q_NORM, + MODEL_TENSOR.ATTN_K, + MODEL_TENSOR.ATTN_K_NORM, + MODEL_TENSOR.ATTN_V, + MODEL_TENSOR.ATTN_V_GATE, # MoVA + MODEL_TENSOR.ATTN_V_EXP, # MoVA + MODEL_TENSOR.ATTN_OUT, + MODEL_TENSOR.ATTN_GATE, + MODEL_TENSOR.FFN_NORM, + + # Dense MLP + MODEL_TENSOR.FFN_GATE, + MODEL_TENSOR.FFN_UP, + MODEL_TENSOR.FFN_DOWN, + + # MoE + MODEL_TENSOR.FFN_GATE_INP, + MODEL_TENSOR.FFN_EXP_PROBS_B, + MODEL_TENSOR.FFN_GATE_EXP, + MODEL_TENSOR.FFN_UP_EXP, + MODEL_TENSOR.FFN_DOWN_EXP, + + # Shared Expert + MODEL_TENSOR.FFN_GATE_SHEXP, + MODEL_TENSOR.FFN_UP_SHEXP, + MODEL_TENSOR.FFN_DOWN_SHEXP, + ] } # tensors that will not be serialized diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index 1f309ad2eafd..68ee81e4b60b 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -1556,6 +1556,12 @@ def add_xielu_beta(self, values: Sequence[float]): def add_xielu_eps(self, values: Sequence[float]): self.add_array(Keys.xIELU.EPS, values) + def add_attention_value_expert_count(self, count: int): + self.add_uint32(Keys.Attention.VALUE_EXPERT_COUNT.format(arch=self.arch), count) + + def add_attention_value_expert_used_count(self, count: int): + self.add_uint32(Keys.Attention.VALUE_EXPERT_USED_COUNT.format(arch=self.arch), count) + # diffusion models def add_diffusion_shift_logits(self, value: bool) -> None: diff --git a/gguf-py/gguf/tensor_mapping.py b/gguf-py/gguf/tensor_mapping.py index 861acfe181fe..6fdcfb418f5a 100644 --- a/gguf-py/gguf/tensor_mapping.py +++ b/gguf-py/gguf/tensor_mapping.py @@ -392,6 +392,7 @@ class TensorNameMap: "model.layers.{bid}.linear_attn.in_proj_z", # qwen3.5 "model.layers.{bid}.self_attn.g_proj", # step3.5 head-wise attention gate "model.layers.{bid}.self_attn.output_gate", # minimax-01 + "model.layers.{bid}.self_attn.attn_gate_proj", # K2Horizon ), # Feed-forward norm @@ -2696,6 +2697,14 @@ class TensorNameMap: MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM: ( "model.layers.{bid}.shared_head.norm", ), + + MODEL_TENSOR.ATTN_V_GATE: ( + "model.layers.{bid}.self_attn.v_router", + ), + + MODEL_TENSOR.ATTN_V_EXP: ( + "model.layers.{bid}.self_attn.v_experts", + ), } # architecture-specific block mappings From 992bdde80308f453821c2833bd11778e2959f22a Mon Sep 17 00:00:00 2001 From: Ryandito Diandaru Date: Sun, 30 Aug 2026 17:14:05 +0000 Subject: [PATCH 02/22] model: loading hparams and tensors in k2-horizon.cpp --- src/llama-arch.cpp | 10 ++ src/llama-arch.h | 8 ++ src/llama-hparams.h | 4 + src/llama-model.h | 4 + src/models/k2-horizon.cpp | 249 ++++++++++++++++++++++++++++++++++++++ src/models/models.h | 22 ++++ 6 files changed, 297 insertions(+) create mode 100644 src/models/k2-horizon.cpp diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index 5e61f61f7f0d..daea7c10314a 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -154,6 +154,7 @@ static const std::map LLM_ARCH_NAMES = { { LLM_ARCH_NANBEIGE, "nanbeige" }, { LLM_ARCH_QWEN3TTS, "qwen3tts" }, { LLM_ARCH_POCKETTTS, "pockettts" }, + { LLM_ARCH_K2_HORIZON, "k2-horizon" }, { LLM_ARCH_UNKNOWN, "(unknown)" }, }; @@ -415,6 +416,10 @@ static const std::map LLM_KV_NAMES = { { LLM_KV_XIELU_BETA, "xielu.beta" }, { LLM_KV_XIELU_EPS, "xielu.eps" }, + // K2 Horizon MoVA + { LLM_KV_ATTENTION_VALUE_EXPERT_COUNT, "%s.attention.value_expert_count"}, + { LLM_KV_ATTENTION_VALUE_EXPERT_USED_COUNT, "%s.attention.value_expert_used_count"}, + // deprecated { LLM_KV_TOKENIZER_PREFIX_ID, "tokenizer.ggml.prefix_token_id" }, { LLM_KV_TOKENIZER_SUFFIX_ID, "tokenizer.ggml.suffix_token_id" }, @@ -693,6 +698,8 @@ static const std::map LLM_TENSOR_NAMES = { { LLM_TENSOR_DFLASH_SELECTOR_PREV, "selector_predecessor" }, { LLM_TENSOR_DFLASH_SELECTOR_NEXT, "selector_successor" }, { LLM_TENSOR_DFLASH_SELECTOR_HIDDEN, "selector_hidden" }, + { LLM_TENSOR_ATTN_V_GATE, "blk.%d.attn_v_gate"}, + { LLM_TENSOR_ATTN_V_EXPS, "blk.%d.attn_v_exps"}, }; // declare information about the model weight tensors: @@ -982,6 +989,9 @@ static const std::map LLM_TENSOR_INFOS = { {LLM_TENSOR_DFLASH_SELECTOR_PREV, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_GET_ROWS}}, {LLM_TENSOR_DFLASH_SELECTOR_NEXT, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_GET_ROWS}}, {LLM_TENSOR_DFLASH_SELECTOR_HIDDEN, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + // K2 Horizon MoVA + {LLM_TENSOR_ATTN_V_GATE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ATTN_V_EXPS, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT_ID}}, }; LLM_KV::LLM_KV(llm_arch arch, const char * suffix) : arch(arch), suffix(suffix) {} diff --git a/src/llama-arch.h b/src/llama-arch.h index ca7d55a5fd78..b3915ac89ca1 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -159,6 +159,7 @@ enum llm_arch { LLM_ARCH_QWEN3TTS, LLM_ARCH_POCKETTTS, LLM_ARCH_MINIMAX_01, + LLM_ARCH_K2_HORIZON, LLM_ARCH_UNKNOWN, }; @@ -424,6 +425,10 @@ enum llm_kv { LLM_KV_DENSE_2_FEAT_OUT, LLM_KV_DENSE_3_FEAT_IN, LLM_KV_DENSE_3_FEAT_OUT, + + // K2 Horizon MoVA + LLM_KV_ATTENTION_VALUE_EXPERT_COUNT, + LLM_KV_ATTENTION_VALUE_EXPERT_USED_COUNT, }; enum llm_tensor { @@ -700,6 +705,9 @@ enum llm_tensor { LLM_TENSOR_DFLASH_SELECTOR_PREV, LLM_TENSOR_DFLASH_SELECTOR_NEXT, LLM_TENSOR_DFLASH_SELECTOR_HIDDEN, + // K2 Horizon MoVA + LLM_TENSOR_ATTN_V_GATE, + LLM_TENSOR_ATTN_V_EXPS, }; diff --git a/src/llama-hparams.h b/src/llama-hparams.h index 1411692a8909..1d877b0cc5c9 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -65,6 +65,10 @@ struct llama_hparams { uint32_t n_expert_used = 0; uint32_t n_rel_attn_bkts = 0; + // K2 Horizon MoVA + uint32_t n_value_expert = 0; + uint32_t n_value_expert_used = 0; + // TODO: this needs to be reworked int32_t n_layer_kv_from_start = -1; // if non-negative, the first n_layer_kv_from_start layers have KV cache diff --git a/src/llama-model.h b/src/llama-model.h index 38066538ed10..a727bd812f20 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -297,6 +297,10 @@ struct llama_layer { struct ggml_tensor * wv_enc = nullptr; struct ggml_tensor * wo_enc = nullptr; struct ggml_tensor * wqkv_gate = nullptr; + // K2 Horizon MoVA + struct ggml_tensor * attn_v_gate = nullptr; + struct ggml_tensor * attn_v_gate_b = nullptr; + struct ggml_tensor * attn_v_exps = nullptr; // relative position bias struct ggml_tensor * attn_rel_b = nullptr; diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp new file mode 100644 index 000000000000..986ca7c81aa3 --- /dev/null +++ b/src/models/k2-horizon.cpp @@ -0,0 +1,249 @@ +#include "models.h" + +void llama_model_k2_horizon::load_arch_hparams(llama_model_loader & ml) { + // generic + ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); + ml.get_key(LLM_KV_ATTENTION_GROUPNORM_GROUPS, hparams.n_norm_groups, false); + + hparams.f_norm_group_eps = hparams.f_norm_rms_eps; + if (hparams.n_norm_groups == 0) hparams.n_norm_groups = 1; + + // moe + if (hparams.n_expert > 0) { + ml.get_key(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, hparams.n_ff_exp); + ml.get_key(LLM_KV_LEADING_DENSE_BLOCK_COUNT, hparams.n_layer_dense_lead, false); + ml.get_key(LLM_KV_MOE_EVERY_N_LAYERS, hparams.moe_every_n_layers, false); + ml.get_key(LLM_KV_EXPERT_SHARED_COUNT, hparams.n_expert_shared, false); + ml.get_key(LLM_KV_EXPERT_SHARED_FEED_FORWARD_LENGTH, hparams.n_ff_shexp, false); + ml.get_key(LLM_KV_EXPERT_WEIGHTS_SCALE, hparams.expert_weights_scale, false); + ml.get_key(LLM_KV_EXPERT_WEIGHTS_NORM, hparams.expert_weights_norm, false); + ml.get_key(LLM_KV_EXPERT_GATING_FUNC, hparams.expert_gating_func, false); + if (hparams.expert_gating_func == LLAMA_EXPERT_GATING_FUNC_TYPE_NONE) { + hparams.expert_gating_func = LLAMA_EXPERT_GATING_FUNC_TYPE_SIGMOID; + } + } + + // mova + ml.get_key(LLM_KV_ATTENTION_VALUE_EXPERT_COUNT, hparams.n_value_expert, false); + ml.get_key(LLM_KV_ATTENTION_VALUE_EXPERT_USED_COUNT, hparams.n_value_expert_used, false); + if (hparams.n_value_expert > 0) { + GGML_ASSERT(hparams.n_value_expert <= LLAMA_MAX_EXPERTS); + GGML_ASSERT(hparams.n_value_expert_used > 0); + GGML_ASSERT(hparams.n_value_expert_used <= hparams.n_value_expert); + } + else { + GGML_ASSERT(hparams.n_value_expert_used == 0); + } + + // model size info + if (hparams.n_layer() == 28 && hparams.n_embd == 1536) { + type = LLM_TYPE_1B; + } + else if (hparams.n_layer() == 48 && hparams.n_embd == 2560) { + type = LLM_TYPE_36B; + } + else { + type = LLM_TYPE_UNKNOWN; + } +} + +void llama_model_k2_horizon::load_arch_tensors(llama_model_loader & ml) { + GGML_UNUSED(ml); + LLAMA_LOAD_LOCALS; // initializing variables basically + + // embeddings + tok_embd = create_tensor( + tn(LLM_TENSOR_TOKEN_EMBD, "weight"), + {n_embd, n_vocab}, + 0 + ); + + // final norm and output projection + output_norm = create_tensor( + tn(LLM_TENSOR_OUTPUT_NORM, "weight"), + {n_embd}, + 0 + ); + + // output + output = create_tensor( + tn(LLM_TENSOR_OUTPUT, "weight"), + {n_embd, n_vocab}, + TENSOR_NOT_REQUIRED // can be tied with embedding (indicated by tensor not found in .gguf). see next conditional + ); + if (output == nullptr) { + output = create_tensor( + tn(LLM_TENSOR_TOKEN_EMBD, "weight"), + {n_embd, n_vocab}, + TENSOR_DUPLICATED + ); + } + + for (int i = 0; i < n_layer; i++){ + auto & layer = layers[i]; + const bool is_moe_layer = n_expert > 0 && static_cast(i) >= hparams.n_layer_dense_lead; + const bool is_mova_layer = is_moe_layer && hparams.n_value_expert > 0; // in the architecture, if mova is moe as well + + // attn normalization + layer.attn_norm = create_tensor( + tn(LLM_TENSOR_ATTN_NORM, "weight", i), + {n_embd}, + 0 + ); + + // query and key tensors, always dense. and their optional normalization + // query + layer.wq = create_tensor( + tn(LLM_TENSOR_ATTN_Q, "weight", i), + {n_embd, n_embd_head_k * n_head}, + 0 + ); + layer.attn_q_norm = create_tensor( + tn(LLM_TENSOR_ATTN_Q_NORM, "weight", i), + {n_embd_head_k * n_head}, + TENSOR_NOT_REQUIRED + ); + + // key + layer.wk = create_tensor( + tn(LLM_TENSOR_ATTN_K, "weight", i), + {n_embd, n_embd_k_gqa}, + 0 + ); + layer.attn_k_norm = create_tensor( + tn(LLM_TENSOR_ATTN_K_NORM, "weight", i), + {n_embd_k_gqa}, + TENSOR_NOT_REQUIRED + ); + + // value tensors, possible MoVA + if (is_mova_layer) { + layer.attn_v_gate = create_tensor( + tn(LLM_TENSOR_ATTN_V_GATE, "weight", i), + {n_embd, hparams.n_value_expert}, + 0 + ); + layer.attn_v_gate_b = create_tensor( + tn(LLM_TENSOR_ATTN_V_GATE, "bias", i), + {hparams.n_value_expert}, + TENSOR_NOT_REQUIRED + ); + layer.attn_v_exps = create_tensor( + tn(LLM_TENSOR_ATTN_V_EXPS, "weight", i), + {n_embd, n_embd_v_gqa, hparams.n_value_expert}, + 0 + ); + } + else { + layer.wv = create_tensor( + tn(LLM_TENSOR_ATTN_V, "weight", i), + {n_embd, n_embd_v_gqa}, + 0 + ); + } + + // attn output projection + layer.wo = create_tensor( + tn(LLM_TENSOR_ATTN_OUT, "weight", i), + {n_embd_head_v * n_head, n_embd}, + 0 + ); + + // optional softplus gate + layer.wqkv_gate = create_tensor( + tn(LLM_TENSOR_ATTN_GATE, "weight", i), + {n_embd, n_embd_head_v * n_head}, + TENSOR_NOT_REQUIRED + ); + + // FFN normalization + layer.ffn_norm = create_tensor( + tn(LLM_TENSOR_FFN_NORM, "weight", i), + {n_embd}, + 0 + ); + + // MoE stuff + if (is_moe_layer) { + if (hparams.n_ff_exp == 0){ + throw std::runtime_error("K2 MoE layer requires expert_feed_forward_length"); + } + + // moe router and it's optional bias + 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 + ); + + // routed experts (up, gate, and down) + layer.ffn_up_exps = create_tensor( + tn(LLM_TENSOR_FFN_UP_EXPS, "weight", i), + {n_embd, hparams.n_ff_exp, n_expert}, + 0 + ); + layer.ffn_gate_exps = create_tensor( + tn(LLM_TENSOR_FFN_GATE_EXPS, "weight", i), + {n_embd, hparams.n_ff_exp, n_expert}, + 0 + ); + layer.ffn_down_exps = create_tensor( + tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), + {hparams.n_ff_exp, n_embd, n_expert}, + 0 + ); + + // shared experts (always evaluated) + if (hparams.n_expert_shared > 0) { + int64_t n_ff_shexp; + if (hparams.n_ff_shexp > 0) { + n_ff_shexp = hparams.n_ff_shexp; + } else { + n_ff_shexp = hparams.n_ff_exp * hparams.n_expert_shared; + } + + // up gate down + layer.ffn_up_shexp = create_tensor( + tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", i), + {n_embd, n_ff_shexp}, + 0 + ); + layer.ffn_gate_shexp = create_tensor( + tn(LLM_TENSOR_FFN_GATE_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 + ); + } + } + else { + // ordinary up gate down + 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}, + 0 + ); + layer.ffn_down = create_tensor( + tn(LLM_TENSOR_FFN_DOWN, "weight", i), + {n_ff, n_embd}, + 0 + ); + } + + } + +} \ No newline at end of file diff --git a/src/models/models.h b/src/models/models.h index af60764c2f7f..a5781786506c 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -2540,3 +2540,25 @@ struct llama_model_step35 : public llama_model_base { std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; }; + + +struct llama_model_k2_horizon : public llama_model_base { + llama_model_k2_horizon( + const llama_model_params & params + ) : llama_model_base(params) {} + + void load_arch_hparams(llama_model_loader & ml) override; + + void load_arch_tensors(llama_model_loader & ml) override; + + struct graph: public llm_graph_context { + graph( + const llama_model & model, + const llm_graph_params & params + ); + }; + + std::unique_ptr build_arch_graph( + const llm_graph_params & params + ) const override; +}; From d63ad5e64f4b5668ef85283e00f14f31545a07ea Mon Sep 17 00:00:00 2001 From: Ryandito Diandaru Date: Mon, 31 Aug 2026 11:41:29 +0000 Subject: [PATCH 03/22] model: K2 Horizon compute graph --- src/models/k2-horizon.cpp | 418 ++++++++++++++++++++++++++++++++++++++ src/models/models.h | 6 + 2 files changed, 424 insertions(+) diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp index 986ca7c81aa3..4ee66cadee97 100644 --- a/src/models/k2-horizon.cpp +++ b/src/models/k2-horizon.cpp @@ -246,4 +246,422 @@ void llama_model_k2_horizon::load_arch_tensors(llama_model_loader & ml) { } +} + +// helper for grouped RMS norm +static ggml_tensor * k2_horizon_group_rms_norm( + ggml_context * ctx, + ggml_tensor * cur, + ggml_tensor * weight, + int64_t n_groups, + float eps +) { + GGML_ASSERT(n_groups > 0); + GGML_ASSERT(cur->ne[0] % n_groups == 0); + + const int64_t n_embd = cur->ne[0]; + const int64_t n_tokens = cur->ne[1]; + + // separate embeddings into groups + cur = ggml_reshape_3d( + ctx, + cur, + n_embd / n_groups, + n_groups, + n_tokens + ); + + // norm it + cur = ggml_rms_norm(ctx, cur, eps); + + // bring back shape + cur = ggml_reshape_2d(ctx, cur, n_embd, n_tokens); + + // additional normalized * (1 + weights) + if (weight != nullptr) { + cur = ggml_add( + ctx, + ggml_mul(ctx, cur, weight), + cur + ); + } + + return cur; +} + +ggml_tensor * llama_model_k2_horizon::graph::build_routed_value( + const llama_layer & layer, + ggml_tensor * cur, + int il +) const { + const int64_t n_embd = cur->ne[0]; + const int64_t n_tokens = cur->ne[1]; + const int64_t n_embd_gqa = hparams.n_embd_v_gqa(il); + const int64_t n_values = hparams.n_value_expert; + const int64_t n_used = hparams.n_value_expert_used; + + GGML_ASSERT(layer.attn_v_gate != nullptr); + GGML_ASSERT(layer.attn_v_exps != nullptr); + GGML_ASSERT(n_values > 0); + GGML_ASSERT(n_used > 0); + + // router. logits and probs + ggml_tensor * logits = build_lora_mm(layer.attn_v_gate, cur); + ggml_tensor * probs = nullptr; + + // probs + llama_expert_gating_func_type gating_func = static_cast(hparams.expert_gating_func); + switch(gating_func){ + case LLAMA_EXPERT_GATING_FUNC_TYPE_SOFTMAX: + probs = ggml_soft_max(ctx0, logits); + break; + case LLAMA_EXPERT_GATING_FUNC_TYPE_SIGMOID: + probs = ggml_sigmoid(ctx0, logits); + break; + default: + GGML_ABORT("Unsupported K2 Horizon value-router gating function"); + } + + // selection probs + ggml_tensor * selection_probs = probs; + if (layer.attn_v_gate_b != nullptr){ + selection_probs = ggml_add(ctx0, probs, layer.attn_v_gate_b); + cb(selection_probs, "v_moe_probs_biased", il); + } + + // select expert values + ggml_tensor * selected_value_experts = ggml_argsort_top_k(ctx0, selection_probs, n_used); + + // reshaping and selecting the weights (probs) of the selected experts + probs = ggml_reshape_3d(ctx0, probs, 1, n_values, n_tokens); + ggml_tensor * selected_weights = ggml_get_rows(ctx0, probs, selected_value_experts); + + // if weights of value experts are to be normalized + if (hparams.expert_weights_norm) { + selected_weights = ggml_reshape_2d(ctx0, selected_weights, n_used, n_tokens); + ggml_tensor * selected_weights_sum = ggml_sum_rows(ctx0, selected_weights); + selected_weights_sum = ggml_clamp(ctx0, selected_weights_sum, 6.103515625e-5f, INFINITY); + selected_weights = ggml_div(ctx0, selected_weights, selected_weights_sum); + selected_weights = ggml_reshape_3d(ctx0, selected_weights, 1, n_used, n_tokens); + cb(selected_weights, "v_moe_weights_norm", il); + } + + // scaling + if (hparams.expert_weights_scale != 0.0f && hparams.expert_weights_scale != 1.0f) { + selected_weights = ggml_scale(ctx0, selected_weights, hparams.expert_weights_scale); + cb(selected_weights, "v_moe_weights_scaled", il); + } + + // labeling + cb(logits, "v_moe_logits", il); + cb(probs, "v_moe_probs", il); + cb(selected_value_experts->src[0], "v_moe_argsort", il); + cb(selected_value_experts, "v_moe_topk", il); + cb(selected_weights, "v_moe_weights", il); + + ggml_tensor * value_inp = ggml_reshape_3d(ctx0, cur, n_embd, 1, n_tokens); + // computing only on selected experts (the _id in the api) + ggml_tensor * values = build_lora_mm_id(layer.attn_v_exps, value_inp, selected_value_experts); + values = ggml_silu(ctx0, values); + values = ggml_mul(ctx0, values, selected_weights); + cb(values, "v_moe_weighted", il); + + // sum the multiple value outputs + ggml_tensor * value_parts[LLAMA_MAX_EXPERTS] = {}; + for(int64_t i = 0; i < n_used; i++) { + value_parts[i] = ggml_view_2d(ctx0, values, n_embd_gqa, n_tokens, values->nb[2], i * values->nb[1]); + } + ggml_tensor * value_out = value_parts[0]; + for (int64_t i = 1; i < n_used; ++i) { + value_out = ggml_add(ctx0, value_out, value_parts[i]); + } + + // making it contiguous in case it isn't (for one expert only) + if (n_used == 1) value_out = ggml_cont(ctx0, value_out); + + cb(value_out, "Vcur_routed", il); + return value_out; +} + +llama_model_k2_horizon::graph::graph( + const llama_model & model, + const llm_graph_params & params +) : llm_graph_context(params) { + const int64_t n_embd_head = hparams.n_embd_head_v(); + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + // initialization or placeholders for computational artifacts + ggml_tensor * cur; + ggml_tensor * inpL = build_inp_embd(model.tok_embd); + ggml_tensor * inp_pos = build_inp_pos(); + auto * inp_attn = build_attn_inp_kv(); + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + + for (int il = 0; il < n_layer; ++il) { + res->t_layer_inp[il] = inpL; + ggml_tensor * inpSA = inpL; // for residuals + + const bool is_moe_layer = n_expert > 0 && static_cast(il) >= hparams.n_layer_dense_lead; + const bool is_mova_layer = is_moe_layer && hparams.n_value_expert > 0; + + // ============ grouped rms norm + cur = k2_horizon_group_rms_norm( + ctx0, + inpL, + model.layers[il].attn_norm, + hparams.n_norm_groups, + hparams.f_norm_rms_eps + ); + cb(cur, "attn_norm", il); + + // ============ setup attention tensors + ggml_tensor * attn_inp = cur; + + // query + ggml_tensor * Qcur = build_lora_mm(model.layers[il].wq, cur, model.layers[il].wq_s); + if (model.layers[il].attn_q_norm != nullptr) { + Qcur = k2_horizon_group_rms_norm( + ctx0, + Qcur, + model.layers[il].attn_q_norm, + n_head, + hparams.f_norm_rms_eps + ); + } + + // key + ggml_tensor * Kcur = build_lora_mm(model.layers[il].wk, cur, model.layers[il].wk_s); + if (model.layers[il].attn_k_norm != nullptr) { + Kcur = k2_horizon_group_rms_norm( + ctx0, + Kcur, + model.layers[il].attn_k_norm, + n_head_kv, + hparams.f_norm_rms_eps + ); + } + + // value + ggml_tensor * Vcur; + if (is_mova_layer) { + Vcur = build_routed_value(model.layers[il], cur, il); // handle MoVA + } + else { + Vcur = build_lora_mm(model.layers[il].wv, cur, model.layers[il].wv_s); + } + + // reshaping + Qcur = ggml_reshape_3d(ctx0, Qcur, n_embd_head, n_head, n_tokens); + Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens); + Vcur = ggml_reshape_3d(ctx0, Vcur, n_embd_head, n_head_kv, n_tokens); + + // applying RoPE + 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, "Qcur", il); + cb(Kcur, "Kcur", il); + cb(Vcur, "Vcur", il); + + // ============ attention (with and without gating) + const float kq_scale = 1.0f / sqrtf(static_cast(n_embd_head)); + if(model.layers[il].wqkv_gate == nullptr){ // without gating + cur = build_attn( + inp_attn, + model.layers[il].wo, + model.layers[il].wo_b, + model.layers[il].wo_s, + Qcur, + Kcur, + Vcur, + nullptr, // attention score bias + nullptr, // attn sink + nullptr, // MLA value transformation + kq_scale, + il + ); + } + else { // with gating + // no output yet + cur = build_attn( + inp_attn, + nullptr, + nullptr, + nullptr, + Qcur, + Kcur, + Vcur, + nullptr, + nullptr, + nullptr, + kq_scale, + il + ); + + // building the gate + constexpr float LN2 = 0.6931471805599453f; + constexpr float ONE_OVER_LN2 = 1.4426950408889634f; + + ggml_tensor * gate = build_lora_mm(model.layers[il].wqkv_gate, attn_inp, model.layers[il].wqkv_gate_s); + gate = ggml_scale(ctx0, gate, LN2); + gate = ggml_softplus(ctx0, gate); + gate = ggml_scale(ctx0, gate, ONE_OVER_LN2); + + // applying the gate + cur = ggml_mul(ctx0, cur, gate); + + // projection + cur = build_lora_mm(model.layers[il].wo, cur, model.layers[il].wo_s); + + // bias + if (model.layers[il].wo_b != nullptr) { + cur = ggml_add(ctx0, cur, model.layers[il].wo_b); + } + } + + // ============ output layer, and take (usually) last token for generation + if (il == n_layer - 1 && inp_out_ids != nullptr) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids); // pull the same positions for inpSA + } + + // ============ add residuals + ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpSA); + cb(ffn_inp, "ffn_inp", il); + + // ============ group RMSNorm before FFN + cur = k2_horizon_group_rms_norm( + ctx0, + ffn_inp, + model.layers[il].ffn_norm, + hparams.n_norm_groups, + hparams.f_norm_rms_eps + ); + cb(cur, "ffn_norm", il); + + // ============ Mixture of Experts + if (is_moe_layer) { + ggml_tensor * moe_out = build_moe_ffn( + cur, + model.layers[il].ffn_gate_inp, + model.layers[il].ffn_up_exps, + model.layers[il].ffn_gate_exps, + model.layers[il].ffn_down_exps, + model.layers[il].ffn_exp_probs_b, + n_expert, + n_expert_used, + LLM_FFN_SILU, + hparams.expert_weights_norm, + hparams.expert_weights_scale, + static_cast(hparams.expert_gating_func), + il + ); + + // shared experts + if (model.layers[il].ffn_gate_shexp != nullptr){ + ggml_tensor * shared_moe_out = build_ffn( + cur, + model.layers[il].ffn_up_shexp, + nullptr, + nullptr, + model.layers[il].ffn_gate_shexp, + nullptr, + nullptr, + model.layers[il].ffn_down_shexp, + nullptr, + nullptr, + nullptr, + LLM_FFN_SILU, + LLM_FFN_PAR, + il + ); + cur = ggml_add(ctx0, moe_out, shared_moe_out); + } + else{ + cur = moe_out; + } + } + else { // normal non moe FFN + cur = build_ffn( + cur, + model.layers[il].ffn_up, + nullptr, + nullptr, + model.layers[il].ffn_gate, + nullptr, + nullptr, + model.layers[il].ffn_down, + nullptr, + nullptr, + nullptr, + LLM_FFN_SILU, + LLM_FFN_PAR, + il + ); + } + cb(cur, "ffn_out", il); + + // ============ FFN residual + cur = ggml_add(ctx0, cur, ffn_inp); + cur = build_cvec(cur, il); + cb(cur, "l_out", il); + + // for next layer + inpL = cur; + } + + // final group rms norm. also becomes last layer embedding + cur = k2_horizon_group_rms_norm( + ctx0, + inpL, + model.output_norm, + hparams.n_norm_groups, + hparams.f_norm_rms_eps + ); + cb(cur, "result_norm", -1); + res->t_embd = cur; + + // ============ vocab projection. also becomes logits + cur = build_lora_mm(model.output, cur,model.output_s); + cb(cur, "result_output", -1); + res->t_logits = cur; + + // build everything + ggml_build_forward_expand(gf, cur); +} + + +std::unique_ptr llama_model_k2_horizon::build_arch_graph ( + const llm_graph_params & params +) const { + return std::make_unique(*this, params); } \ No newline at end of file diff --git a/src/models/models.h b/src/models/models.h index a5781786506c..5327182a2540 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -2556,6 +2556,12 @@ struct llama_model_k2_horizon : public llama_model_base { const llama_model & model, const llm_graph_params & params ); + + ggml_tensor * build_routed_value ( + const llama_layer & layer, + ggml_tensor * cur, + int il // layer index + ) const; }; std::unique_ptr build_arch_graph( From d6b148a7954728a6569898f1692adeed36dca4ec Mon Sep 17 00:00:00 2001 From: Ryandito Diandaru Date: Tue, 1 Sep 2026 15:55:03 +0000 Subject: [PATCH 04/22] model: K2 Horizon compute graph adjustment and registering tokenizers --- conversion/__init__.py | 1 + conversion/base.py | 6 ++++++ conversion/k2_horizon.py | 5 ++++- convert_hf_to_gguf_update.py | 3 +++ src/llama-model.cpp | 5 ++++- src/llama-vocab.cpp | 9 +++++++++ src/llama-vocab.h | 1 + src/models/k2-horizon.cpp | 10 +++------- 8 files changed, 31 insertions(+), 9 deletions(-) diff --git a/conversion/__init__.py b/conversion/__init__.py index a8af0831a878..648dd73fb39c 100644 --- a/conversion/__init__.py +++ b/conversion/__init__.py @@ -135,6 +135,7 @@ "JinaBertModel": "bert", "JinaEmbeddingsV5Model": "bert", "K2HorizonForCausalLM": "k2_horizon", + "K2AuroraForCausalLM": "k2_horizon", # TODO: DELETE "KORMoForCausalLM": "qwen", "KimiK25ForConditionalGeneration": "deepseek", "KimiK3ForConditionalGeneration": "kimi_k3", diff --git a/conversion/base.py b/conversion/base.py index daae28e92adc..4b8594701cc7 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -1540,6 +1540,12 @@ def get_vocab_base_pre(self, tokenizer) -> str: if chkhsh == "9e454714343b69b99b71795c1d27a68c2a1d15dab111f4d353109f966af29da7": # ref: https://huggingface.co/LiquidAI/LFM2.5-8B-A1B res = "lfm2" + if chkhsh == "1f9825a388f700a6b591722f17d470cbbcf10973ece35d2fd14239a14110ae1a": + # ref: https://huggingface.co/IFM/K2-Horizon-0.9B + res = "k2-horizon" + if chkhsh == "a9af07a84191f55098b248ae6f3dfe9e32d3190bebe8eafd91c1ddec9bc3449f": + # ref: https://huggingface.co/IFM/K2-Horizon-36B + res = "k2-horizon" if chkhsh == "0ef9807a4087ebef797fc749390439009c3b9eda9ad1a097abbe738f486c01e5": # ref: https://huggingface.co/meta-llama/Meta-Llama-3-8B res = "llama-bpe" diff --git a/conversion/k2_horizon.py b/conversion/k2_horizon.py index d3f5af9236ba..6ceb30b29f0f 100644 --- a/conversion/k2_horizon.py +++ b/conversion/k2_horizon.py @@ -8,7 +8,10 @@ from .base import ModelBase, TextModel, gguf -@ModelBase.register("K2HorizonForCausalLM") +@ModelBase.register( + "K2HorizonForCausalLM", + "K2AuroraForCausalLM", # TODO: DELETE +) @ModelBase.example( "IFM/K2-Horizon-0.9B", "IFM/K2-Horizon-36B", diff --git a/convert_hf_to_gguf_update.py b/convert_hf_to_gguf_update.py index e5d3196efe41..17f9488af96c 100755 --- a/convert_hf_to_gguf_update.py +++ b/convert_hf_to_gguf_update.py @@ -190,6 +190,9 @@ class TOKENIZER_TYPE(IntEnum): {"name": "gpt-2", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/evilfreelancer/ruGPT3XL", "chkhsh": "0fe1cf6eda062318a1af7270f3331a85c539a01778ff948e24388e949c5282f4"}, # lfm2 variants {"name": "lfm2", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/LiquidAI/LFM2.5-8B-A1B", "chkhsh": "9e454714343b69b99b71795c1d27a68c2a1d15dab111f4d353109f966af29da7"}, + # K2 Horizon. 2 hashes because various sets of tokens depending on size + {"name": "k2-horizon", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/IFM/K2-Horizon-0.9B", "chkhsh": "1f9825a388f700a6b591722f17d470cbbcf10973ece35d2fd14239a14110ae1a"}, + {"name": "k2-horizon", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/IFM/K2-Horizon-36B", "chkhsh": "a9af07a84191f55098b248ae6f3dfe9e32d3190bebe8eafd91c1ddec9bc3449f"}, ] diff --git a/src/llama-model.cpp b/src/llama-model.cpp index fc83658dd7ff..682fee0d814b 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -336,7 +336,9 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params return new llama_model_kimi_k3(params); case LLM_ARCH_STEP35: return new llama_model_step35(params); - default: + case LLM_ARCH_K2_HORIZON: + return new llama_model_k2_horizon(params); + default: throw std::runtime_error(std::string("unsupported model architecture: '") + llm_arch_name(arch) + "'"); } @@ -2932,6 +2934,7 @@ llama_rope_type llama_model_rope_type(const llama_model * model) { case LLM_ARCH_MIMO2: case LLM_ARCH_STEP35: case LLM_ARCH_TALKIE: + case LLM_ARCH_K2_HORIZON: case LLM_ARCH_MELLUM: return LLAMA_ROPE_TYPE_NEOX; diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index ff926ceecd17..4d054db403fc 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -535,6 +535,11 @@ struct llm_tokenizer_bpe : llm_tokenizer { "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}+| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", }; break; + case LLAMA_VOCAB_PRE_TYPE_K2_HORIZON: + regex_exprs = { + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + break; case LLAMA_VOCAB_PRE_TYPE_WHITESPACE: // whitespace pre-tokenizer (jinaai/jina-embeddings-v2-base-zh) regex_exprs = { @@ -2381,6 +2386,10 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) { } else if ( tokenizer_pre == "mellum2") { pre_type = LLAMA_VOCAB_PRE_TYPE_MELLUM2; + } else if ( + tokenizer_pre == "k2-horizon") { + pre_type = LLAMA_VOCAB_PRE_TYPE_K2_HORIZON; + clean_spaces = false; } else { throw std::runtime_error(format("unknown pre-tokenizer type: '%s'", tokenizer_pre.c_str())); } diff --git a/src/llama-vocab.h b/src/llama-vocab.h index b7c28926338b..2f3f7d56bc57 100644 --- a/src/llama-vocab.h +++ b/src/llama-vocab.h @@ -65,6 +65,7 @@ enum llama_vocab_pre_type { LLAMA_VOCAB_PRE_TYPE_GRANITE_EMB_MULTI = 54, LLAMA_VOCAB_PRE_TYPE_MELLUM2 = 55, LLAMA_VOCAB_PRE_TYPE_LAGUNA = 56, + LLAMA_VOCAB_PRE_TYPE_K2_HORIZON = 57, }; struct LLM_KV; diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp index 4ee66cadee97..ac901da50a51 100644 --- a/src/models/k2-horizon.cpp +++ b/src/models/k2-horizon.cpp @@ -277,13 +277,9 @@ static ggml_tensor * k2_horizon_group_rms_norm( // bring back shape cur = ggml_reshape_2d(ctx, cur, n_embd, n_tokens); - // additional normalized * (1 + weights) + // apply the learned normalization weights if (weight != nullptr) { - cur = ggml_add( - ctx, - ggml_mul(ctx, cur, weight), - cur - ); + cur = ggml_mul(ctx, cur, weight); } return cur; @@ -664,4 +660,4 @@ std::unique_ptr llama_model_k2_horizon::build_arch_graph ( const llm_graph_params & params ) const { return std::make_unique(*this, params); -} \ No newline at end of file +} From 35999d101cf2233fc54f09c3c8d599da7303ce02 Mon Sep 17 00:00:00 2001 From: Ryandito Diandaru Date: Tue, 1 Sep 2026 18:04:20 +0000 Subject: [PATCH 05/22] model: K2 Horizon chat template and accomodate safetensors naming --- conversion/base.py | 2 + conversion/k2_horizon.py | 15 +- models/templates/k2-horizon.jinja | 883 ++++++++++++++++++++++++++++++ 3 files changed, 899 insertions(+), 1 deletion(-) create mode 100644 models/templates/k2-horizon.jinja diff --git a/conversion/base.py b/conversion/base.py index 4b8594701cc7..6f0c3194debf 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -221,6 +221,8 @@ def index_tensors(self, remote_hf_model_id: str | None = None) -> dict[str, Call prefix = "model" if not self.is_mistral_format else "consolidated" part_names: list[str] = ModelBase.get_model_part_names(self.dir_model, prefix, ".safetensors") + if not part_names and not self.is_mistral_format: + part_names = ModelBase.get_model_part_names(self.dir_model, "pytorch_model", ".safetensors") is_safetensors: bool = len(part_names) > 0 if not is_safetensors: part_names = ModelBase.get_model_part_names(self.dir_model, "pytorch_model", ".bin") diff --git a/conversion/k2_horizon.py b/conversion/k2_horizon.py index 6ceb30b29f0f..2d7722209517 100644 --- a/conversion/k2_horizon.py +++ b/conversion/k2_horizon.py @@ -1,6 +1,7 @@ from __future__ import annotations import re +from pathlib import Path from typing import Iterable import torch @@ -19,6 +20,19 @@ class K2HorizonModel(TextModel): model_arch = gguf.MODEL_ARCH.K2HORIZON + def set_vocab(self): + super().set_vocab() + + template_path = ( + Path(__file__).parent.parent + / "models" + / "templates" + / "k2-horizon.jinja" + ) + template = template_path.read_text(encoding="utf-8") + self.gguf_writer.remove_key(gguf.Keys.Tokenizer.CHAT_TEMPLATE) + self.gguf_writer.add_chat_template(template) + def set_gguf_parameters(self): super().set_gguf_parameters() # generic @@ -196,4 +210,3 @@ def prepare_tensors(self): f"{remaining_value_experts}" ) - diff --git a/models/templates/k2-horizon.jinja b/models/templates/k2-horizon.jinja new file mode 100644 index 000000000000..9dc680f7a7e8 --- /dev/null +++ b/models/templates/k2-horizon.jinja @@ -0,0 +1,883 @@ +{{- bos_token }} +{%- if tool_presentation is defined -%} + {{- raise_exception("Unsupported argument: tool_presentation. Use tool_presentation_format with one of: json, xml, markdown.") -}} +{%- endif -%} +{%- if tool_calling_format is defined -%} + {{- raise_exception("Unsupported argument: tool_calling_format. Use tool_call_format with one of: json, xml, xml_typed.") -}} +{%- endif -%} +{%- if tool_format is defined -%} + {{- raise_exception("Unsupported argument: tool_format. Use tool_call_format with one of: json, xml, xml_typed.") -}} +{%- endif -%} +{%- set tool_presentation_fmt = tool_presentation_format | default('markdown') -%} +{%- set tool_call_fmt = tool_call_format | default('xml') -%} +{%- if tool_presentation_fmt != 'json' and tool_presentation_fmt != 'xml' and tool_presentation_fmt != 'markdown' -%} + {{- raise_exception("Unsupported tool_presentation_format: '" ~ tool_presentation_fmt ~ "'. Supported formats: json, xml, markdown.") -}} +{%- endif -%} +{%- if tool_call_fmt != 'json' and tool_call_fmt != 'xml' and tool_call_fmt != 'xml_typed' -%} + {{- raise_exception("Unsupported tool_call_format: '" ~ tool_call_fmt ~ "'. Supported formats: json, xml, xml_typed.") -}} +{%- endif -%} + +{#- Renderability state, computed during validate_tools (single walk, no extra -#} +{#- traversal at render time): ok = working flag for the tool being validated; -#} +{#- bad = pipe-delimited indices of tools that must render as verbatim JSON. -#} +{%- set RB = namespace(ok=true, bad='|') -%} + +{%- macro value_contains_mapping(v) -%} +{%- if v is mapping -%} +true +{%- elif v is sequence and v is not string -%} +{%- set f = namespace(x='false') -%} +{%- for c in v -%}{%- if value_contains_mapping(c) == 'true' -%}{%- set f.x = 'true' -%}{%- endif -%}{%- endfor -%} +{{- f.x -}} +{%- else -%} +false +{%- endif -%} +{%- endmacro -%} + +{%- macro render_compact_type_name(type_name, spec) -%} +{%- if type_name == "array" -%} +array[{%- if 'items' in spec -%}{{ render_compact_type(spec['items']) }}{%- else -%}any{%- endif -%}] +{%- elif type_name -%} +{{- type_name -}} +{%- else -%} +any +{%- endif -%} +{%- endmacro -%} + +{%- macro render_compact_type(spec) -%} +{%- if spec is not mapping -%} +any +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string and spec.type | length > 0 -%} +{%- for type_name in spec.type -%}{{ render_compact_type_name(type_name, spec) }}{%- if not loop.last -%}|{%- endif -%}{%- endfor -%} +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string -%} +any +{%- elif spec.type -%} +{{- render_compact_type_name(spec.type, spec) -}} +{%- elif spec['$ref'] is string -%} +{{- spec['$ref'].split('/') | last -}} +{%- elif spec.oneOf -%} +oneOf[{%- for variant in spec.oneOf -%}{{ render_compact_type(variant) }}{%- if not loop.last -%}|{%- endif -%}{%- endfor -%}] +{%- elif spec.anyOf -%} +anyOf[{%- for variant in spec.anyOf -%}{{ render_compact_type(variant) }}{%- if not loop.last -%}|{%- endif -%}{%- endfor -%}] +{%- elif spec.properties -%} +object +{%- elif 'items' in spec -%} +array[{{ render_compact_type(spec['items']) }}] +{%- else -%} +any +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_type_name(type_name, spec) -%} +{%- if type_name == "array" -%} +array of {% if 'items' in spec %}{{ render_markdown_type(spec['items']) }}{% else %}any{% endif %} +{%- elif type_name -%} +{{- type_name -}} +{%- else -%} +any +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_type(spec) -%} +{%- if spec is true -%} +True +{%- elif spec is false -%} +False +{%- elif spec is not mapping -%} +any +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string and spec.type | length > 0 -%} +{%- for type_name in spec.type -%}{{ render_markdown_type_name(type_name, spec) }}{% if not loop.last %} or {% endif %}{%- endfor -%} +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string -%} +any +{%- elif spec.type -%} +{{- render_markdown_type_name(spec.type, spec) -}} +{%- elif spec['$ref'] is string -%} +{{- spec['$ref'].split('/') | last -}} +{%- elif spec.oneOf -%} +oneOf[{%- for variant in spec.oneOf -%}{{ render_markdown_type(variant) }}{% if not loop.last %} or {% endif %}{%- endfor -%}] +{%- elif spec.anyOf -%} +anyOf[{%- for variant in spec.anyOf -%}{{ render_markdown_type(variant) }}{% if not loop.last %} or {% endif %}{%- endfor -%}] +{%- elif spec.properties -%} +object +{%- elif 'items' in spec -%} +array of {{ render_markdown_type(spec['items']) }} +{%- else -%} +any +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_text(value) -%} +{{- value.split() | join(" ") -}} +{%- endmacro -%} + +{%- macro render_python_string(value) -%} +'{{- value.split() | join(" ") | replace("\\", "\\\\") | replace("'", "\\'") -}}' +{%- endmacro -%} + +{%- macro render_python_repr(value) -%} +{%- if value is string -%} +{{ render_python_string(value) }} +{%- elif value is true -%} +True +{%- elif value is false -%} +False +{%- elif value is none -%} +None +{%- elif value is mapping -%} +{{- "{" -}} +{%- for key, child in value | items -%} +{{ render_python_repr(key) }}: {{ render_python_repr(child) }}{%- if not loop.last -%}, {% endif -%} +{%- endfor -%} +{{- "}" -}} +{%- elif value is sequence -%} +{{- "[" -}} +{%- for child in value -%} +{{ render_python_repr(child) }}{%- if not loop.last -%}, {% endif -%} +{%- endfor -%} +{{- "]" -}} +{%- else -%} +{{- value -}} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_value(value) -%} +{%- if value is string -%}{{ render_xml_text(value) }}{%- else -%}{{ render_python_repr(value) }}{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_enum_value(value) -%} +{%- if value is string -%}"{{- value | replace("\\", "\\\\") | replace("\"", "\\\"") -}}"{%- else -%}"{{- render_python_repr(value) | replace("\\", "\\\\") | replace("\"", "\\\"") -}}"{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_enum(values) -%} +{%- for value in values -%}{{ render_xml_enum_value(value) }}{%- if not loop.last -%}|{%- endif -%}{%- endfor -%} +{%- endmacro -%} + +{%- macro render_xml_default_attr(value) -%} +{{- " default=" }}{%- if value is string -%}"{{- value | replace("\\", "\\\\") | replace("\"", "\\\"") -}}"{%- else -%}{{ render_xml_value(value) }}{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_attr(name, value) -%} +{{- " " + name + "=" }}{%- if value == "" -%}""{%- else -%}{{ render_xml_value(value) }}{%- endif -%} +{%- endmacro -%} + +{%- macro validate_schema(spec, path, lenient=false, classify=true, in_variant=false) -%} +{%- if spec is mapping -%} + {%- if not lenient -%} + {%- if spec.required is defined -%} + {%- if spec.required is string or spec.required is not sequence -%} + {{- raise_exception("Schema '" + path + "' has 'required' but it is not a list.") -}} + {%- endif -%} + {%- if spec.required | length > 0 and not spec.properties and not in_variant -%} + {{- raise_exception("Schema '" + path + "' has required fields but no properties object to define them.") -}} + {%- endif -%} + {%- if spec.properties -%} + {%- for required_name in spec.required -%} + {%- if required_name not in spec.properties -%} + {{- raise_exception("Schema '" + path + "' marks '" + required_name + "' as required, but that property is not defined in properties.") -}} + {%- endif -%} + {%- endfor -%} + {%- endif -%} + {%- endif -%} + {%- endif -%} + {#- renderability classification, piggybacking on this walk (no raises here): -#} + {#- constructs the pretty renderer does not fully handle flip RB.ok so the -#} + {#- tool falls back to verbatim JSON. Skipped entirely for json presentation. -#} + {%- if classify -%} + {%- for key, value in spec | items -%} + {%- if key == '$ref' -%} + {#- llama.cpp's Jinja has no dictionary constructor, so $ref inlining stays -#} + {#- template-local by falling back to the exact JSON presentation. -#} + {%- set RB.ok = false -%} + {%- elif key == '$defs' or key == 'definitions' -%} + {%- if value is mapping -%} + {%- for dk, dv in value | items -%} + {{- validate_schema(dv, path + ".$defs." + dk, true) -}} + {%- endfor -%} + {%- else -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == 'type' -%} + {%- if value is mapping -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == 'enum' -%} + {%- if value is string or value is mapping or value is not sequence -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == 'items' -%} + {#- any items shape renders: mapping structurally, others via repr detail -#} + {%- elif key == 'oneOf' or key == 'anyOf' -%} + {%- if value is mapping or value is string or value is not sequence -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == 'required' -%} + {%- if value and not spec.properties -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif ('|' ~ key ~ '|') in '|description|default|title|examples|properties|patternProperties|additionalProperties|returns|' -%} + {%- elif value is mapping -%} + {%- for uk, uv in value | items -%} + {%- if value_contains_mapping(uv) == 'true' -%}{%- set RB.ok = false -%}{%- endif -%} + {%- endfor -%} + {%- elif value is sequence and value is not string -%} + {%- if value_contains_mapping(value) == 'true' -%}{%- set RB.ok = false -%}{%- endif -%} + {%- endif -%} + {%- endfor -%} + {%- endif -%} + {%- if spec.properties -%} + {%- for child_name, child_spec in spec.properties | items -%} + {{- validate_schema(child_spec, path + "." + child_name, lenient, classify) -}} + {%- endfor -%} + {%- endif -%} + {%- if 'items' in spec -%}{{- validate_schema(spec['items'], path + "[]", lenient, classify) -}}{%- endif -%} + {%- if spec.oneOf -%} + {%- for variant in spec.oneOf -%}{{- validate_schema(variant, path + ".oneOf[" + (loop.index0 | string) + "]", lenient, classify, true) -}}{%- endfor -%} + {%- endif -%} + {%- if spec.anyOf -%} + {%- for variant in spec.anyOf -%}{{- validate_schema(variant, path + ".anyOf[" + (loop.index0 | string) + "]", lenient, classify, true) -}}{%- endfor -%} + {%- endif -%} + {%- if spec.additionalProperties is mapping -%}{{- validate_schema(spec.additionalProperties, path + ".additionalProperties", lenient, classify) -}}{%- endif -%} + {%- if spec.patternProperties is mapping -%} + {%- for pattern, pattern_spec in spec.patternProperties | items -%} + {{- validate_schema(pattern_spec, path + ".patternProperties[" + pattern + "]", lenient, classify) -}} + {%- endfor -%} + {%- endif -%} + {%- if spec.returns is mapping -%}{{- validate_schema(spec.returns, path + ".returns", lenient, classify) -}}{%- endif -%} +{%- endif -%} +{%- endmacro -%} + +{%- macro validate_tools(tools_list, classify=true) -%} +{%- set RB.bad = '|' -%} +{%- for tool in tools_list -%} + {%- set fn = tool.function if tool.function is defined else tool -%} + {%- set RB.ok = true -%} + {%- if fn.parameters is defined and fn.parameters is string -%} + {{- raise_exception("tool.function.parameters must be a dict, not a JSON string. Parse it before passing to the template.") -}} + {%- endif -%} + {%- if fn.parameters is not defined or fn.parameters is none -%} + {%- if fn.arguments is defined -%} + {{- raise_exception("Tool '" + fn.name + "' has 'arguments' instead of 'parameters'. Rename 'arguments' to 'parameters'.") -}} + {%- else -%} + {{- raise_exception("Tool '" + fn.name + "' is missing required 'parameters' field. Each tool must have a 'parameters' dict with 'type', 'properties', and 'required' keys.") -}} + {%- endif -%} + {%- endif -%} + {{- validate_schema(fn.parameters, "tool." + fn.name + ".parameters", false, classify) -}} + {%- if classify -%} + {%- if fn.parameters is mapping -%} + {#- unknown container-valued keys at the parameters ROOT are never rendered -#} + {#- by the pretty path (root extras are dropped) -> verbatim fallback. -#} + {%- for rk, rv in fn.parameters | items -%} + {%- if rk not in ['type', 'description', 'enum', 'default', 'properties', 'required', 'optional', 'title', 'items', 'oneOf', 'anyOf', 'additionalProperties', 'patternProperties', 'returns', 'examples', '$defs', 'definitions', '$ref'] -%} + {%- if rv is mapping or (rv is sequence and rv is not string) -%}{%- set RB.ok = false -%}{%- endif -%} + {%- endif -%} + {%- endfor -%} + {%- else -%} + {%- set RB.ok = false -%} + {%- endif -%} + {%- endif -%} + {%- if fn.returns is mapping -%}{{- validate_schema(fn.returns, "tool." + fn.name + ".returns", false, classify) -}}{%- endif -%} + {%- if classify and fn.returns is not defined and fn.response is mapping -%}{{- validate_schema(fn.response, "tool." + fn.name + ".response", true) -}}{%- endif -%} + {#- unknown container-valued keys at the FUNCTION level are never rendered -> fallback. -#} + {%- if classify -%} + {%- for fk, fv in fn | items -%} + {%- if fk not in ['name', 'description', 'parameters', 'returns', 'response', 'type', 'function'] -%} + {%- if fv is mapping or (fv is sequence and fv is not string) -%}{%- set RB.ok = false -%}{%- endif -%} + {%- endif -%} + {%- endfor -%} + {%- endif -%} + {%- if not RB.ok -%}{%- set RB.bad = RB.bad ~ loop.index0 ~ '|' -%}{%- endif -%} +{%- endfor -%} +{%- endmacro -%} + +{%- macro render_tools_json(tools_list) -%} +{{- "" }} +{%- for tool in tools_list %} +{{- "\n" }} +{{- tool | tojson }} +{%- endfor %} +{{- "\n" }} +{%- endmacro -%} + +{%- macro render_xml_schema_attrs(spec, include_value_attrs) -%} +{%- if spec is mapping -%} +{%- set structural_keys = ["type", "description", "enum", "default", "properties", "required", "items", "oneOf", "anyOf", "additionalProperties", "patternProperties", "returns"] -%} +{%- if include_value_attrs and spec.enum -%}{{- " enum=" }}{{ render_xml_enum(spec.enum) }}{%- endif -%} +{%- if include_value_attrs and spec.default is defined -%}{{ render_xml_default_attr(spec.default) }}{%- endif -%} +{%- if spec.additionalProperties is defined and spec.additionalProperties is not mapping -%}{{ render_xml_attr("additionalProperties", spec.additionalProperties) }}{%- endif -%} +{%- if spec.patternProperties is defined and spec.patternProperties is not mapping -%}{{ render_xml_attr("patternProperties", spec.patternProperties) }}{%- endif -%} +{%- for key, value in spec | items -%} + {%- if key not in structural_keys -%} +{{ render_xml_attr(key, value) }} + {%- endif -%} +{%- endfor -%} +{%- endif -%} +{%- endmacro -%} + +{%- macro xml_schema_has_children(spec, include_properties, include_description) -%} +{%- if spec is not mapping -%} +false +{%- elif (include_description and spec.description is defined) or (include_properties and spec.properties) or 'items' in spec or spec.oneOf or spec.anyOf or spec.additionalProperties is mapping or spec.patternProperties is mapping or spec.returns is defined -%} +true +{%- else -%} +false +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_schema_node(tag, spec, include_properties) -%} +{%- if spec is mapping -%} +{{- "<" + tag + " type=" + render_compact_type(spec) }}{{ render_xml_schema_attrs(spec, true) }} +{%- if xml_schema_has_children(spec, include_properties, true) == 'true' -%} +{{- ">" }}{{ render_xml_schema_children(spec, include_properties, true) }}{{- "" }} +{%- else -%} +{{- "/>" }} +{%- endif -%} +{%- else -%} +{{- "<" + tag + ">" }}{{ render_xml_value(spec) }}{{- "" }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_pattern_property(pattern, spec) -%} +{%- if spec is mapping -%} +{{- "" }}{{ render_xml_schema_children(spec, true, true) }}{{- "" }} +{%- else -%} +{{- "/>" }} +{%- endif -%} +{%- else -%} +{{- "" }}{{ render_xml_value(spec) }}{{- "" }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_schema_children(spec, include_properties, include_description) -%} +{%- if include_description and spec.description is defined -%}{{- "" }}{{ spec.description }}{{- "" }}{%- endif -%} +{%- if include_properties and spec.properties -%} +{%- for child_name, child_spec in spec.properties | items -%} +{{- render_xml_param(child_name, child_spec, spec.required or []) }} +{%- endfor -%} +{%- endif -%} +{%- if 'items' in spec -%}{{ render_xml_schema_node("items", spec['items'], true) }}{%- endif -%} +{%- if spec.oneOf -%} +{{- "" }} +{%- for variant in spec.oneOf -%}{{ render_xml_schema_node("variant", variant, true) }}{%- endfor -%} +{{- "" }} +{%- endif -%} +{%- if spec.anyOf -%} +{{- "" }} +{%- for variant in spec.anyOf -%}{{ render_xml_schema_node("variant", variant, true) }}{%- endfor -%} +{{- "" }} +{%- endif -%} +{%- if spec.additionalProperties is mapping -%}{{ render_xml_schema_node("additionalProperties", spec.additionalProperties, true) }}{%- endif -%} +{%- if spec.patternProperties is mapping -%} +{{- "" }} +{%- for pattern, pattern_spec in spec.patternProperties | items -%}{{ render_xml_pattern_property(pattern, pattern_spec) }}{%- endfor -%} +{{- "" }} +{%- elif spec.patternProperties is defined -%}{{ render_xml_value(spec.patternProperties) }}{%- endif -%} +{%- if spec.returns is mapping -%}{{ render_xml_schema_node("returns", spec.returns, true) }}{%- elif spec.returns is defined -%}{{ render_xml_value(spec.returns) }}{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_param(name, spec, required_list) -%} +{{- "" }} +{%- if spec.description -%}{{ spec.description }}{%- endif -%} +{{- render_xml_schema_children(spec, true, false) }} +{{- "" }} +{%- else -%} +{{- "/>" }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_tools_xml(tools_list) -%} +{{- "" }} +{%- for tool in tools_list -%} + {%- set fn = tool.function if tool.function is defined else tool -%} + {%- set fnp = namespace(p=fn.parameters) -%} +{{- "\n" }} +{%- if fn.description -%} +{{- "" }}{{ fn.description }}{{- "" }} +{%- endif -%} +{{- "" }} +{%- if fnp.p and fnp.p.properties -%} + {%- for pname, pspec in fnp.p.properties | items -%} +{{- render_xml_param(pname, pspec, fnp.p.required or []) }} + {%- endfor -%} +{%- elif fnp.p is mapping and (fnp.p.oneOf or fnp.p.anyOf or 'items' in fnp.p) -%} +{{- render_xml_schema_children(fnp.p, true, false) }} +{%- endif -%} +{{- "" }} +{%- set fn_ret = fn.returns if fn.returns is defined else fn.response -%} +{%- if fn_ret is mapping -%}{{ render_xml_schema_node("returns", fn_ret, true) }}{%- elif fn_ret is defined -%}{{ render_xml_value(fn_ret) }}{%- endif -%} +{{- "" }} +{%- endfor -%} +{{- "\n" }} +{%- endmacro -%} + +{%- macro render_markdown_literal(value) -%} +{%- if value is string and value == "" -%}"" +{%- elif value is string -%}`{{ value | replace("\n", "\\n") }}` +{%- else -%}`{{ render_python_repr(value) }}` +{%- endif -%} +{%- endmacro -%} + +{%- macro render_allowed_values(values) -%} +{%- for value in values -%}{{ render_markdown_literal(value) }}{% if not loop.last %}, {% endif %}{%- endfor -%} +{%- endmacro -%} + +{%- macro render_markdown_value(value) -%} +{%- if value is string and value == "" -%}""{%- elif value is string -%}{{ value }}{%- else -%}{{ render_python_repr(value) }}{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_detail(indent, label, value) -%} +{{- "\n" + indent + " - " + label + ": " }}{{ render_markdown_value(value) }} +{%- endmacro -%} + +{%- macro render_markdown_metadata_detail(label, value) -%} +{{- "\n- " + label + ": " }}{{ render_markdown_value(value) }} +{%- endmacro -%} + +{%- macro render_markdown_schema_annotations(spec, indent, include_value_details) -%} +{%- if include_value_details and spec.description is defined -%}{{ render_markdown_detail(indent, "Description", spec.description | replace("\n", "\n" + indent + " ")) }}{%- endif -%} +{%- if include_value_details and spec.enum is defined -%}{{- "\n" + indent + " - Allowed values: " }}{{ render_allowed_values(spec.enum) }}{%- endif -%} +{%- if include_value_details and spec.default is defined -%}{{- "\n" + indent + " - Default: " }}{{ render_markdown_literal(spec.default) }}{%- endif -%} +{%- if spec.additionalProperties is defined -%} + {%- if spec.additionalProperties is mapping -%} +{{- "\n" + indent + " - Additional properties *(" + render_markdown_type(spec.additionalProperties) + ")*" }} +{{- render_markdown_schema_details(spec.additionalProperties, indent + " ", true) }} + {%- else -%} +{{ render_markdown_detail(indent, "Additional properties", spec.additionalProperties) }} + {%- endif -%} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_metadata_annotations(spec) -%} +{%- if spec.description is defined -%}{{ render_markdown_metadata_detail("Description", spec.description | replace("\n", "\n ")) }}{%- endif -%} +{%- if spec.enum is defined -%}{{- "\n- Allowed values: " }}{{ render_allowed_values(spec.enum) }}{%- endif -%} +{%- if spec.default is defined -%}{{- "\n- Default: " }}{{ render_markdown_literal(spec.default) }}{%- endif -%} +{%- if spec.additionalProperties is defined -%} + {%- if spec.additionalProperties is mapping -%} +{{- "\n- Additional properties *(" + render_markdown_type(spec.additionalProperties) + ")*" }} +{{- render_markdown_schema_details(spec.additionalProperties, "", true) }} + {%- else -%} +{{ render_markdown_metadata_detail("Additional properties", spec.additionalProperties) }} + {%- endif -%} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_schema_extras(spec, indent) -%} +{%- set rendered_keys = ["type", "description", "enum", "default", "properties", "required", "items", "oneOf", "anyOf", "additionalProperties", "patternProperties", "returns"] -%} +{%- for key, value in spec | items -%} + {%- if key not in rendered_keys -%} +{{- "\n" + indent + " - " + key + ": " }}{{ render_markdown_value(value) }} + {%- endif -%} +{%- endfor -%} +{%- endmacro -%} + +{%- macro render_markdown_metadata_extras(spec) -%} +{%- set rendered_keys = ["type", "description", "enum", "default", "properties", "required", "items", "oneOf", "anyOf", "additionalProperties", "patternProperties", "returns"] -%} +{%- for key, value in spec | items -%} + {%- if key not in rendered_keys -%} +{{- "\n- " + key + ": " }}{{ render_markdown_value(value) }} + {%- endif -%} +{%- endfor -%} +{%- endmacro -%} + +{%- macro markdown_schema_has_extra(spec) -%} +{%- set rendered_keys = ["type", "description", "enum", "default", "properties", "required", "items", "oneOf", "anyOf", "additionalProperties", "patternProperties", "returns"] -%} +{%- set found = namespace(value='false') -%} +{%- for key, value in spec | items -%} + {%- if key not in rendered_keys -%}{%- set found.value = 'true' -%}{%- endif -%} +{%- endfor -%} +{{- found.value -}} +{%- endmacro -%} + +{%- macro markdown_parameter_schema_has_details(spec) -%} +{%- if spec.description is defined or spec.enum is defined or spec.default is defined or spec.additionalProperties is defined or spec.patternProperties is defined or 'items' in spec or spec.oneOf or spec.anyOf or spec.returns is defined or markdown_schema_has_extra(spec) == 'true' -%} +true +{%- else -%} +false +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_schema_structure(spec, indent, include_properties) -%} +{%- if include_properties and spec.properties -%} + {%- for child_name, child_spec in spec.properties | items -%} +{{- render_markdown_param(child_name, child_spec, spec.required or [], indent + " ") }} + {%- endfor -%} +{%- endif -%} +{%- if 'items' in spec and spec['items'] is mapping -%} +{{- "\n" + indent + " - Items *(" + render_markdown_type(spec['items']) + ")*" }} +{{- render_markdown_schema_details(spec['items'], indent + " ", true) }} +{%- elif 'items' in spec -%} +{{ render_markdown_detail(indent, "Items", spec['items']) }} +{%- endif -%} +{%- if spec.oneOf -%} +{{- "\n" + indent + " - oneOf:" }} + {%- for variant in spec.oneOf -%} +{{- "\n" + indent + " - Variant " }}{{ loop.index }}{{- " *(" + render_markdown_type(variant) + ")*" }} +{{- render_markdown_schema_details(variant, indent + " ", true) }} + {%- endfor -%} +{%- endif -%} +{%- if spec.anyOf -%} +{{- "\n" + indent + " - anyOf:" }} + {%- for variant in spec.anyOf -%} +{{- "\n" + indent + " - Variant " }}{{ loop.index }}{{- " *(" + render_markdown_type(variant) + ")*" }} +{{- render_markdown_schema_details(variant, indent + " ", true) }} + {%- endfor -%} +{%- endif -%} +{%- if spec.patternProperties is mapping -%} +{{- "\n" + indent + " - Pattern properties:" }} + {%- for pattern, pattern_spec in spec.patternProperties | items -%} + {%- if pattern_spec is mapping -%} +{{- "\n" + indent + " - `" + pattern + "` *(" + render_markdown_type(pattern_spec) + ")*" }} +{{- render_markdown_schema_details(pattern_spec, indent + " ", true) }} + {%- else -%} +{{- "\n" + indent + " - `" + pattern + "`: " }}{{ render_markdown_value(pattern_spec) }} + {%- endif -%} + {%- endfor -%} +{%- elif spec.patternProperties is defined -%} +{{ render_markdown_detail(indent, "Pattern properties", spec.patternProperties) }} +{%- endif -%} +{%- if spec.returns is mapping -%} +{{- "\n" + indent + " - Returns *(" + render_markdown_type(spec.returns) + ")*" }} +{{- render_markdown_schema_details(spec.returns, indent + " ", true) }} +{%- elif spec.returns is defined -%} +{{ render_markdown_detail(indent, "Returns", spec.returns) }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_schema_details(spec, indent, include_value_details) -%} +{%- if spec is mapping -%} +{{- render_markdown_schema_annotations(spec, indent, include_value_details) }} +{{- render_markdown_schema_structure(spec, indent, true) }} +{{- render_markdown_schema_extras(spec, indent) }} +{%- elif spec is not boolean -%} +{{- "\n" + indent + " - Value: " }}{{ render_markdown_literal(spec) }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_parameter_schema(spec) -%} +{%- if spec is mapping -%} +{{- render_markdown_metadata_annotations(spec) }} +{%- if 'items' in spec and spec['items'] is mapping -%} +{{- "\n- Items *(" + render_markdown_type(spec['items']) + ")*" }} +{{- render_markdown_schema_details(spec['items'], "", true) }} +{%- elif 'items' in spec -%} +{{ render_markdown_metadata_detail("Items", spec['items']) }} +{%- endif -%} +{%- if spec.oneOf -%} +{{- "\n- oneOf:" }} + {%- for variant in spec.oneOf -%} +{{- "\n - Variant " }}{{ loop.index }}{{- " *(" + render_markdown_type(variant) + ")*" }} +{{- render_markdown_schema_details(variant, " ", true) }} + {%- endfor -%} +{%- endif -%} +{%- if spec.anyOf -%} +{{- "\n- anyOf:" }} + {%- for variant in spec.anyOf -%} +{{- "\n - Variant " }}{{ loop.index }}{{- " *(" + render_markdown_type(variant) + ")*" }} +{{- render_markdown_schema_details(variant, " ", true) }} + {%- endfor -%} +{%- endif -%} +{%- if spec.patternProperties is mapping -%} +{{- "\n- Pattern properties:" }} + {%- for pattern, pattern_spec in spec.patternProperties | items -%} + {%- if pattern_spec is mapping -%} +{{- "\n - `" + pattern + "` *(" + render_markdown_type(pattern_spec) + ")*" }} +{{- render_markdown_schema_details(pattern_spec, " ", true) }} + {%- else -%} +{{- "\n - `" + pattern + "`: " }}{{ render_markdown_value(pattern_spec) }} + {%- endif -%} + {%- endfor -%} +{%- elif spec.patternProperties is defined -%} +{{ render_markdown_metadata_detail("Pattern properties", spec.patternProperties) }} +{%- endif -%} +{%- if spec.returns is mapping -%} +{{- "\n- Returns *(" + render_markdown_type(spec.returns) + ")*" }} +{{- render_markdown_schema_details(spec.returns, "", true) }} +{%- elif spec.returns is defined -%} +{{ render_markdown_metadata_detail("Returns", spec.returns) }} +{%- endif -%} +{{- render_markdown_metadata_extras(spec) }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_param(name, spec, required_list, indent) -%} +{{- "\n" + indent + "- `" + name + "` *(" + render_markdown_type(spec) }} +{%- if name in (required_list or []) -%}{{- ", required" }}{%- endif -%} +{{- ")*" }} +{%- if spec.description -%}{{- " - " + spec.description | replace("\n", "\n" + indent + " ") }}{%- endif -%} +{%- if spec.enum -%} +{{- "\n" + indent + " - Allowed values: " }}{{ render_allowed_values(spec.enum) }} +{%- endif -%} +{%- if spec.default is defined -%} +{{- "\n" + indent + " - Default: " }}{{ render_markdown_literal(spec.default) }} +{%- endif -%} +{{- render_markdown_schema_details(spec, indent, false) }} +{%- endmacro -%} + +{%- macro render_tools_markdown(tools_list) -%} +{{- "" }} +{%- for tool in tools_list -%} + {%- set fn = tool.function if tool.function is defined else tool -%} + {%- set fnp = namespace(p=fn.parameters) -%} +{{- "\n## " + fn.name }} +{%- if fn.description -%} +{{- "\n" + fn.description }} +{%- endif -%} +{{- "\n\n**Parameters**" }} +{%- if fnp.p and fnp.p.properties -%} + {%- for pname, pspec in fnp.p.properties | items -%} +{{- render_markdown_param(pname, pspec, fnp.p.required or [], "") }} + {%- endfor -%} +{%- elif fnp.p is mapping and (fnp.p.oneOf or fnp.p.anyOf or 'items' in fnp.p) -%} +{{- render_markdown_parameter_schema(fnp.p) }} +{%- else -%} +{{- "\n- None" }} +{%- endif -%} +{%- set fn_ret = fn.returns if fn.returns is defined else fn.response -%} +{%- if fn_ret is mapping -%} +{{- "\n\n**Returns**" }} +{{- "\n- Return *(" + render_markdown_type(fn_ret) + ")*" }} +{{- render_markdown_schema_details(fn_ret, "", true) }} +{%- elif fn_ret is defined -%} +{{- "\n\n**Returns**\n- " }}{{ render_markdown_value(fn_ret) }} +{%- endif -%} +{%- if not loop.last -%}{{- "\n" }}{%- endif -%} +{%- endfor -%} +{{- "\n" }} +{%- endmacro -%} + +{%- macro render_tool_presentation(tools_list, fmt) -%} +{%- if fmt == 'json' -%} +{{- render_tools_json(tools_list) }} +{%- elif RB.bad != '|' -%} +{#- some tool uses constructs the pretty renderers cannot represent (verdicts -#} +{#- computed during validate_tools): render the WHOLE toolset exactly as the -#} +{#- json presentation would, so the block stays uniform and model-familiar. -#} +{{- render_tools_json(tools_list) }} +{%- elif fmt == 'xml' -%} +{{- render_tools_xml(tools_list) }} +{%- elif fmt == 'markdown' -%} +{{- render_tools_markdown(tools_list) }} +{%- else -%} +{{- raise_exception("Unsupported tool_presentation_format: '" + fmt + "'. Supported formats: json, xml, markdown.") }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_call_instructions(fmt) -%} +{%- if fmt == 'json' -%} +{{- "Wrap all tool calls in a single block. For each call, emit one JSON object with the function name and arguments on the same line inside tags:\n\n\n{\"name\": , \"arguments\": }\n" }} +{%- elif fmt == 'xml' -%} +{{- "Wrap all tool calls in a single block. For each call, write the function name at the start of , followed by paired and tags for each argument:\n\n\n$FUNCTION_NAME\n$PARAMETER_NAME\n$PARAMETER_VALUE\n...\n\n\n\nString and scalar parameters should be written as plain text. Array and object parameters should be written as JSON literals." }} +{%- elif fmt == 'xml_typed' -%} +{{- "Wrap all tool calls in a single block. For each call, write the function name at the start of , followed by , , and tags for each argument:\n\n\n$FUNCTION_NAME\n$PARAMETER_NAME\n$ARGUMENT_TYPE\n$PARAMETER_VALUE\n...\n\n\n\nUse the parameter type shown in the tool definition. If that type contains anyOf or oneOf, use the actual argument value type instead. String and scalar parameters should be written as plain text. Array and object parameters should be written as JSON literals." }} +{%- else -%} +{{- raise_exception("Unsupported tool_call_format: '" + fmt + "'. Supported formats: json, xml, xml_typed.") }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_system_with_tools(tools_list, system_content, presentation_fmt, call_fmt) -%} +{{- "<|ifm|im_start|>system\n# Tools\nYou may call one or more tools to assist with the user query.\n\nAvailable tools are:\n\n" }} +{{- render_tool_presentation(tools_list, presentation_fmt) }} +{{- "\n\nWhen calling tools, you MUST follow the tool-call format below:\n\n" }} +{{- render_call_instructions(call_fmt) }} +{%- if system_content -%} +{{- "\n\n" + system_content }} +{%- endif -%} +{{- "<|ifm|im_end|>" }} +{%- endmacro -%} + +{%- macro render_argument_value(value) -%} +{%- if value is string -%}{{- value -}}{%- else -%}{{- value | tojson -}}{%- endif -%} +{%- endmacro -%} + +{%- macro render_value_type(value) -%} +{%- if value is none -%}null +{%- elif value is boolean -%}boolean +{%- elif value is integer -%}integer +{%- elif value is number -%}number +{%- elif value is string -%}string +{%- elif value is mapping -%}object +{%- elif value is sequence -%}array +{%- else -%}any +{%- endif -%} +{%- endmacro -%} + +{%- macro schema_has_combinator(spec) -%} +{%- if spec.oneOf or spec.anyOf -%} +true +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string and spec.type | length > 1 -%} +true +{%- elif spec.type == "array" and 'items' in spec -%} +{{- schema_has_combinator(spec['items']) -}} +{%- elif spec.properties -%} + {%- set found = namespace(value='false') -%} + {%- for child_name, child_spec in spec.properties | items -%} + {%- if schema_has_combinator(child_spec) == 'true' -%} + {%- set found.value = 'true' -%} + {%- endif -%} + {%- endfor -%} +{{- found.value -}} +{%- else -%} +false +{%- endif -%} +{%- endmacro -%} + +{%- macro render_arg_type(tools_list, tool_name, arg_name, value) -%} +{%- set found = namespace(type='any') -%} +{%- for tool in tools_list -%} + {%- set fn = tool.function if tool.function is defined else tool -%} + {%- if fn.name == tool_name and fn.parameters and fn.parameters.properties and arg_name in fn.parameters.properties -%} + {%- set spec = fn.parameters.properties[arg_name] -%} + {%- if spec is mapping and spec['$ref'] is string -%} + {%- set found.type = render_value_type(value) -%} + {%- elif schema_has_combinator(spec) == 'true' -%} + {%- set found.type = render_value_type(value) -%} + {%- else -%} + {%- set found.type = render_compact_type(spec) -%} + {%- endif -%} + {%- endif -%} +{%- endfor -%} +{{- found.type -}} +{%- endmacro -%} + +{%- macro render_tool_calls_block(tool_calls, fmt, tools_list) -%} +{{- "" }} +{%- for raw_tool_call in tool_calls -%} + {%- set tool_call = raw_tool_call.function if raw_tool_call.function else raw_tool_call -%} + {%- if tool_call.arguments is string -%} + {{- raise_exception("tool_call.arguments must be a dict, not a JSON string. Parse it before passing to the template.") -}} + {%- endif -%} + {%- if fmt == 'json' -%} +{{- "\n{\"name\": \"" + tool_call.name + "\", \"arguments\": " }}{{ tool_call.arguments | tojson }}{{- "}" }} + {%- elif fmt == 'xml' or fmt == 'xml_typed' -%} +{{- "\n" + tool_call.name + "\n" }} + {%- for key, value in tool_call.arguments | items -%} +{{- "" + key + "\n" }} +{%- if fmt == 'xml_typed' -%} +{{- "" + render_arg_type(tools_list, tool_call.name, key, value) + "\n" }} +{%- endif -%} +{{- "" }}{{ render_argument_value(value) }}{{- "\n" }} + {%- endfor -%} +{{- "" }} + {%- else -%} + {{- raise_exception("Unsupported tool_call_format: '" + fmt + "'. Supported formats: json, xml, xml_typed.") -}} + {%- endif -%} +{%- endfor -%} +{{- "\n" }} +{%- endmacro -%} + +{%- macro render_tool_response_messages(raw_content) -%} +{%- if raw_content is string -%} +{{- '<|ifm|im_start|>tool\n' + raw_content + '<|ifm|im_end|>' }} +{%- elif raw_content is sequence and raw_content is not string and raw_content is not mapping -%} + {%- if raw_content | length == 0 -%} + {{- raise_exception("tool message content list must not be empty.") -}} + {%- endif -%} +{{- '<|ifm|im_start|>tool\n' -}} + {%- for item in raw_content -%} + {%- if not loop.first -%}{{- '\n' -}}{%- endif -%} + {%- if item is string -%} +{{- item -}} + {%- elif item is mapping and item.text is string -%} +{{- item.text -}} + {%- else -%} +{{- (item | tojson) -}} + {%- endif -%} + {%- endfor -%} +{{- '<|ifm|im_end|>' -}} +{%- else -%} +{{- '<|ifm|im_start|>tool\n' }}{{ raw_content | tojson }}{{- '<|ifm|im_end|>' }} +{%- endif -%} +{%- endmacro -%} + +{%- set available_tools = tools if tools else [] -%} +{%- if (not available_tools) and messages[0].role == 'system' and messages[0].get('tools') -%} + {%- set available_tools = messages[0]['tools'] -%} +{%- endif -%} +{%- if available_tools -%} + {{- validate_tools(available_tools, tool_presentation_fmt != 'json') }} + {%- set system_content = '' -%} + {%- if messages[0].role == 'system' and messages[0].content -%} + {%- set system_content = messages[0].content -%} + {%- endif -%} + {{- render_system_with_tools(available_tools, system_content, tool_presentation_fmt, tool_call_fmt) }} +{%- else -%} + {%- if messages[0].role == 'system' -%} + {{- '<|ifm|im_start|>system\n' + messages[0].content + '<|ifm|im_end|>' }} + {%- endif -%} +{%- endif -%} + +{%- for message in messages -%} + {%- if message.content is string -%} + {%- set content = message.content -%} + {%- else -%} + {%- set content = '' -%} + {%- endif -%} + {%- if (message.role == "user") or (message.role == "system" and not loop.first) -%} + {{- '<|ifm|im_start|>' + message.role + '\n' + content + '<|ifm|im_end|>' }} + {%- elif message.role == "assistant" -%} + {%- set thinking_content = '' -%} + {%- set think_tag = 'ifm|think' -%} + {%- if message.think is defined and message.think is string -%} + {%- set thinking_content = message.think -%} + {%- set think_tag = 'ifm|think' -%} + {%- elif message.think_fast is defined and message.think_fast is string -%} + {%- set thinking_content = message.think_fast -%} + {%- set think_tag = 'ifm|think_fast' -%} + {%- elif message.think_faster is defined and message.think_faster is string -%} + {%- set thinking_content = message.think_faster -%} + {%- set think_tag = 'ifm|think_faster' -%} + {%- elif message.reasoning_content is defined and message.reasoning_content is string -%} + {%- set thinking_content = message.reasoning_content -%} + {%- set think_tag = 'ifm|think' -%} + {%- elif message.reasoning is defined and message.reasoning is string -%} + {%- set thinking_content = message.reasoning -%} + {%- set think_tag = 'ifm|think' -%} + {%- else -%} + {%- if '' in content -%} + {%- set thinking_content = content.split('')[0].rstrip('\n').split('')[-1].lstrip('\n') -%} + {%- set content = content.split('')[-1].lstrip('\n') -%} + {%- set think_tag = 'ifm|think' -%} + {%- elif '' in content -%} + {%- set thinking_content = content.split('')[0].rstrip('\n').split('')[-1].lstrip('\n') -%} + {%- set content = content.split('')[-1].lstrip('\n') -%} + {%- set think_tag = 'ifm|think_fast' -%} + {%- elif '' in content -%} + {%- set thinking_content = content.split('')[0].rstrip('\n').split('')[-1].lstrip('\n') -%} + {%- set content = content.split('')[-1].lstrip('\n') -%} + {%- set think_tag = 'ifm|think_faster' -%} + {%- endif -%} + {%- endif -%} + {{- '<|ifm|im_start|>' + message.role }} + {% generation %} + {%- if think_tag -%} + {%- if thinking_content -%} + {{- '<' + think_tag + '>\n' + thinking_content + '\n\n' + content.lstrip('\n') }} + {%- else -%} + {{- '<' + think_tag + '>\n\n' + content.lstrip('\n') }} + {%- endif -%} + {%- else -%} + {{- content }} + {%- endif -%} + {%- if message.tool_calls -%} + {%- if content -%} + {{- '\n' }} + {%- endif -%} + {{- render_tool_calls_block(message.tool_calls, tool_call_fmt, available_tools) }} + {%- endif -%} + {{- '<|ifm|im_end|>' -}} + {%- endgeneration -%} + {%- elif message.role == "tool" -%} + {{- render_tool_response_messages(message.content) }} + {%- endif -%} +{%- endfor -%} +{%- if add_generation_prompt -%} + {%- set effort = reasoning_effort | default('high') -%} + {%- if enable_thinking is defined and enable_thinking is false -%} + {{- '<|ifm|im_start|>assistant\n\n\n' }} + {%- elif effort == 'high' -%} + {{- '<|ifm|im_start|>assistant\n\n' }} + {%- elif effort == 'medium' -%} + {{- '<|ifm|im_start|>assistant\n\n' }} + {%- elif effort == 'low' -%} + {{- '<|ifm|im_start|>assistant\n\n' }} + {%- else -%} + {{- raise_exception("Unsupported reasoning_effort: '" + effort + "'. Supported values: high, medium, low.") -}} + {%- endif -%} +{%- endif -%} From 69d3a4e825c242f1378cd7d12ed9c68b38120725 Mon Sep 17 00:00:00 2001 From: WestWaters Date: Fri, 11 Sep 2026 01:08:48 -0700 Subject: [PATCH 06/22] unicode : add the K2-Horizon pre-tokenizer splitter MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The K2-Horizon regex had no arm in unicode_regex_split_custom and fell through to the general std::regex fallback, which fails two ways. On MSVC std::regex rejects \p{...}, so no K2-Horizon GGUF loads on Windows at all: llama-quantize, llama-imatrix and llama-perplexity all abort with regex_error(error_escape) before a token is produced. Where the fallback does compile it is still wrong. unicode_regex_split collapses each codepoint to a single byte naming its Unicode category before matching, and U+200C/U+200D are category Control, which has no entry in k_ucat_cpt, so both become the 0xD0 fallback byte. The literal ‌ and ‍ alternatives in K2's regex can then never match and every ZWNJ or ZWJ ends a letter run. The splitter is the existing llama3 one with a single rule widened, since K2's regex differs from llama3's only in that a letter run also takes marks, ZWNJ and ZWJ. tests/test-unicode.cpp gains a case for this: it fails before the change with [Amy] [ZWNJ khaham] and passes after with the run intact. --- src/unicode.cpp | 148 +++++++++++++++++++++++++++++++++++++++++ tests/test-unicode.cpp | 30 +++++++-- 2 files changed, 171 insertions(+), 7 deletions(-) diff --git a/src/unicode.cpp b/src/unicode.cpp index 93996f9dd542..d2f166208f4a 100644 --- a/src/unicode.cpp +++ b/src/unicode.cpp @@ -470,6 +470,150 @@ static std::vector unicode_regex_split_custom_llama3(const std::string & return bpe_offsets; } +static std::vector unicode_regex_split_custom_k2_horizon(const std::string & text, const std::vector & offsets) { + std::vector bpe_offsets; // store the offset of each word + bpe_offsets.reserve(offsets.size()); // Reserve memory for the approximate size + + const auto cpts = unicode_cpts_from_utf8(text); + + size_t start = 0; + for (auto offset : offsets) { + const size_t offset_ini = start; + const size_t offset_end = start + offset; + assert(offset_end <= cpts.size()); + start = offset_end; + + static const uint32_t OUT_OF_RANGE = 0xFFFFFFFF; + auto _get_cpt = [&] (const size_t pos) -> uint32_t { + return (offset_ini <= pos && pos < offset_end) ? cpts[pos] : OUT_OF_RANGE; + }; + + auto _get_flags = [&] (const size_t pos) -> unicode_cpt_flags { + return (offset_ini <= pos && pos < offset_end) ? unicode_cpt_flags_from_cpt(cpts[pos]) : unicode_cpt_flags{}; + }; + + // K2-Horizon: letter runs are (?:\p{L}|\p{M}|\u200C|\u200D)+ + auto _is_k2_letter = [&] (const size_t pos) -> bool { + const uint32_t c = _get_cpt(pos); + if (c == 0x200C || c == 0x200D) { + return true; + } + const auto f = _get_flags(pos); + return f.is_letter || f.is_accent_mark; + }; + + size_t _prev_end = offset_ini; + auto _add_token = [&] (const size_t end) -> size_t { + assert(_prev_end <= end && end <= offset_end); + size_t len = end - _prev_end; + if (len > 0) { + bpe_offsets.push_back(len); + } + _prev_end = end; + return len; + }; + + for (size_t pos = offset_ini; pos < offset_end; /*pos++*/ ) { + const uint32_t cpt = _get_cpt(pos); + const auto flags = _get_flags(pos); + + // regex: (?i:'s|'t|'re|'ve|'m|'ll|'d) // case insensitive + if (cpt == '\'' && pos+1 < offset_end) { + uint32_t cpt_next = unicode_tolower(_get_cpt(pos+1)); + if (cpt_next == 's' || cpt_next == 't' || cpt_next == 'm' || cpt_next == 'd') { + pos += _add_token(pos+2); + continue; + } + if (pos+2 < offset_end) { + uint32_t cpt_next_next = unicode_tolower(_get_cpt(pos+2)); + if ((cpt_next == 'r' && cpt_next_next == 'e') || + (cpt_next == 'v' && cpt_next_next == 'e') || + (cpt_next == 'l' && cpt_next_next == 'l')) { + pos += _add_token(pos+3); + continue; + } + } + } + + // regex: [^\r\n\p{L}\p{N}]?(?:\p{L}|\p{M}|\u200C|\u200D)+ + if (!(cpt == '\r' || cpt == '\n' || flags.is_number)) { + if (_is_k2_letter(pos) || _is_k2_letter(pos+1)) { // one or more letters/marks/ZWNJ/ZWJ + pos++; + while (_is_k2_letter(pos)) { + pos++; + } + _add_token(pos); + continue; + } + } + + // regex: \p{N}{1,3} + if (flags.is_number) { + size_t ini = pos; + while (_get_flags(pos).is_number) { + if (++pos - ini >= 3 ) { + _add_token(pos); + ini = pos; + } + } + _add_token(pos); + continue; + } + + // regex: ?[^\s\p{L}\p{N}]+[\r\n]* + auto flags2 = (cpt == ' ' ? _get_flags(pos+1) : flags); + if (!(flags2.is_whitespace | flags2.is_letter | flags2.is_number) && flags.as_uint()) { + pos += (cpt == ' '); + while (!(flags2.is_whitespace | flags2.is_letter | flags2.is_number) && flags2.as_uint()) { + flags2 = _get_flags(++pos); + } + uint32_t cpt2 = _get_cpt(pos); + while (cpt2 == '\r' || cpt2 == '\n') { + cpt2 = _get_cpt(++pos); + } + _add_token(pos); + continue; + } + + size_t num_whitespaces = 0; + size_t last_end_r_or_n = 0; + while (_get_flags(pos+num_whitespaces).is_whitespace) { + uint32_t cpt2 = _get_cpt(pos+num_whitespaces); + if (cpt2 == '\r' || cpt2 == '\n') { + last_end_r_or_n = pos + num_whitespaces + 1; + } + num_whitespaces++; + } + + // regex: \s*[\r\n]+ + if (last_end_r_or_n > 0) { + pos = last_end_r_or_n; + _add_token(pos); + continue; + } + + // regex: \s+(?!\S) + if (num_whitespaces > 1 && _get_cpt(pos+num_whitespaces) != OUT_OF_RANGE) { + pos += num_whitespaces - 1; + _add_token(pos); + continue; + } + + // regex: \s+ + if (num_whitespaces > 0) { + pos += num_whitespaces; + _add_token(pos); + continue; + } + + // no matches + _add_token(++pos); + } + } + + return bpe_offsets; +} + // Qwen2 system regex: "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+" static std::vector unicode_regex_split_custom_qwen2(const std::string & text, const std::vector & offsets) { std::vector bpe_offsets; // store the offset of each word @@ -1062,6 +1206,10 @@ static std::vector unicode_regex_split_custom(const std::string & text, } else if ( regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { bpe_offsets = unicode_regex_split_custom_qwen35(text, offsets); + } else if (regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { + // K2-Horizon: llama3 splitter with marks + ZWNJ/ZWJ inside letter runs + // (the generic std::regex fallback cannot parse \p{..} on MSVC) + bpe_offsets = unicode_regex_split_custom_k2_horizon(text, offsets); } else if (regex_expr == "\\p{Han}+") { // K2's first pattern - handle all K2 patterns together bpe_offsets = unicode_regex_split_custom_kimi_k2(text, offsets); diff --git a/tests/test-unicode.cpp b/tests/test-unicode.cpp index 2347d9000a8e..52edf759bbdf 100644 --- a/tests/test-unicode.cpp +++ b/tests/test-unicode.cpp @@ -4,15 +4,13 @@ #include #include -int main() { - const std::vector regex_exprs = { - "[~][A-Za-z]+| ?[\\p{S}]+|\\s+", - }; - const std::vector expected = { " ~", "foo" }; - const auto actual = unicode_regex_split(" ~foo", regex_exprs, false); +static int check(const char * name, const std::string & text, + const std::vector & regex_exprs, + const std::vector & expected) { + const auto actual = unicode_regex_split(text, regex_exprs, false); if (actual != expected) { - fprintf(stderr, "unexpected split:"); + fprintf(stderr, "%s: unexpected split:", name); for (const auto & piece : actual) { fprintf(stderr, " [%s]", piece.c_str()); } @@ -22,3 +20,21 @@ int main() { return 0; } + +int main() { + int n_fail = 0; + + n_fail += check("simple", " ~foo", + { "[~][A-Za-z]+| ?[\\p{S}]+|\\s+" }, + { " ~", "foo" }); + + // K2-Horizon letter runs take marks and ZWNJ/ZWJ, so a ZWNJ must not end a run. + // "A" + "mi" + U+200C + "khaham" + " " + "1" + const std::string k2_text = "Aمی‌خواهم 1"; + + n_fail += check("k2-horizon zwnj", k2_text, + { "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+" }, + { "Aمی‌خواهم", " ", "1" }); + + return n_fail; +} From a8104b5532af306b72e2e811d46e23fce9531514 Mon Sep 17 00:00:00 2001 From: "Natani L. Mayday" <71436458+TaskPuppyNatani@users.noreply.github.com> Date: Mon, 14 Sep 2026 14:18:00 -0700 Subject: [PATCH 07/22] tests: expand K2 Horizon unicode splitter coverage --- tests/test-unicode.cpp | 216 ++++++++++++++++++++++++++++++++++++----- 1 file changed, 192 insertions(+), 24 deletions(-) diff --git a/tests/test-unicode.cpp b/tests/test-unicode.cpp index 52edf759bbdf..3ea2991fb46a 100644 --- a/tests/test-unicode.cpp +++ b/tests/test-unicode.cpp @@ -1,40 +1,208 @@ #include "../src/unicode.h" #include +#include #include #include -static int check(const char * name, const std::string & text, - const std::vector & regex_exprs, - const std::vector & expected) { - const auto actual = unicode_regex_split(text, regex_exprs, false); +int main() { + { + const std::vector regex_exprs = { + "[~][A-Za-z]+| ?[\\p{S}]+|\\s+", + }; + const std::vector expected = { " ~", "foo" }; + const auto actual = unicode_regex_split(" ~foo", regex_exprs, false); - if (actual != expected) { - fprintf(stderr, "%s: unexpected split:", name); - for (const auto & piece : actual) { - fprintf(stderr, " [%s]", piece.c_str()); + if (actual != expected) { + fprintf(stderr, "unexpected split:"); + for (const auto & piece : actual) { + fprintf(stderr, " [%s]", piece.c_str()); + } + fprintf(stderr, "\n"); + return 1; } - fprintf(stderr, "\n"); - return 1; } - return 0; -} + { + const std::vector regex_exprs = { + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" + "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" + "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" + "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; -int main() { - int n_fail = 0; + const std::string input = "ab\u200ccd ef\u200dgh cafe\u0301"; + const std::vector expected = { + "ab\u200ccd", + " ef\u200dgh", + " cafe\u0301", + }; + + try { + const auto actual = unicode_regex_split(input, regex_exprs, false); + + if (actual != expected) { + fprintf(stderr, "unexpected K2-Horizon split:"); + for (const auto & piece : actual) { + fprintf(stderr, " [%s]", piece.c_str()); + } + fprintf(stderr, "\n"); + return 1; + } + } catch (const std::exception & e) { + fprintf(stderr, "K2-Horizon regex split threw exception: %s\n", e.what()); + return 1; + } + } + + + { + const std::vector llama3_regex = { + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" + "[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|" + "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" + "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + + const std::vector k2_horizon_regex = { + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" + "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" + "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" + "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + + const std::vector inputs = { + "Hello, world!", + "can't won't we're they'll", + "123 1234 1234567", + "Hello!!!\\nNext line", + "alpha beta gamma", + " café résumé", + }; + + for (const auto & input : inputs) { + const auto llama3 = unicode_regex_split(input, llama3_regex, false); + const auto k2 = unicode_regex_split(input, k2_horizon_regex, false); + + if (llama3 != k2) { + fprintf(stderr, "K2-Horizon diverged from Llama 3 for ordinary input: %s\n", input.c_str()); + + fprintf(stderr, "Llama 3:"); + for (const auto & piece : llama3) { + fprintf(stderr, " [%s]", piece.c_str()); + } + + fprintf(stderr, "\nK2-Horizon:"); + for (const auto & piece : k2) { + fprintf(stderr, " [%s]", piece.c_str()); + } + + fprintf(stderr, "\n"); + return 1; + } + } + } + + + { + const std::vector k2_horizon_regex = { + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" + "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" + "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" + "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + + auto check_k2 = [&] ( + const char * name, + const std::string & input, + const std::vector & expected) { + const auto actual = unicode_regex_split(input, k2_horizon_regex, false); - n_fail += check("simple", " ~foo", - { "[~][A-Za-z]+| ?[\\p{S}]+|\\s+" }, - { " ~", "foo" }); + if (actual != expected) { + fprintf(stderr, "K2-Horizon reference mismatch for %s\n", name); - // K2-Horizon letter runs take marks and ZWNJ/ZWJ, so a ZWNJ must not end a run. - // "A" + "mi" + U+200C + "khaham" + " " + "1" - const std::string k2_text = "Aمی‌خواهم 1"; + fprintf(stderr, "expected:"); + for (const auto & piece : expected) { + fprintf(stderr, " [%s]", piece.c_str()); + } - n_fail += check("k2-horizon zwnj", k2_text, - { "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+" }, - { "Aمی‌خواهم", " ", "1" }); + fprintf(stderr, "\nactual:"); + for (const auto & piece : actual) { + fprintf(stderr, " [%s]", piece.c_str()); + } - return n_fail; + fprintf(stderr, "\n"); + return false; + } + + return true; + }; + + if (!check_k2( + "ZWNJ", + "ab\u200ccd", + { "ab\u200ccd" })) { + return 1; + } + + if (!check_k2( + "ZWJ", + "ef\u200dgh", + { "ef\u200dgh" })) { + return 1; + } + + if (!check_k2( + "combined ZWNJ/ZWJ", + "ab\u200ccd ef\u200dgh", + { "ab\u200ccd", " ef\u200dgh" })) { + return 1; + } + + // The reference tokenizer NFC-normalizes cafe + combining acute + // to the composed form before pre-tokenization. + if (!check_k2( + "NFC accent", + "caf\u00e9", + { "caf\u00e9" })) { + return 1; + } + + if (!check_k2( + "Persian ZWNJ", + "\u0645\u06cc\u200c\u0631\u0648\u0645", + { "\u0645\u06cc\u200c\u0631\u0648\u0645" })) { + return 1; + } + + if (!check_k2( + "Devanagari ZWJ", + "\u0915\u094d\u200d\u0937", + { "\u0915\u094d\u200d\u0937" })) { + return 1; + } + + if (!check_k2( + "contractions", + "can't won't we're they'll", + { "can", "'t", " won", "'t", " we", "'re", " they", "'ll" })) { + return 1; + } + + if (!check_k2( + "numbers", + "123 1234 1234567", + { "123", " ", "123", "4", " ", "123", "456", "7" })) { + return 1; + } + + if (!check_k2( + "punctuation/newline", + "Hello!!!\nNext line", + { "Hello", "!!!\n", "Next", " line" })) { + return 1; + } + } + + return 0; } From e78bd9435e3773974f3d080b635dcf17c52e7c00 Mon Sep 17 00:00:00 2001 From: aaryamonvikram Date: Thu, 17 Sep 2026 17:06:37 +0400 Subject: [PATCH 08/22] unicode: handle K2 Horizon case folding and empty input Assisted-by: Codex --- src/llama-vocab.cpp | 2 +- src/unicode.cpp | 11 ++++++++++- tests/test-unicode.cpp | 16 +++++++++++++++- 3 files changed, 26 insertions(+), 3 deletions(-) diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index 4d054db403fc..38e5623374c2 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -537,7 +537,7 @@ struct llm_tokenizer_bpe : llm_tokenizer { break; case LLAMA_VOCAB_PRE_TYPE_K2_HORIZON: regex_exprs = { - "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", }; break; case LLAMA_VOCAB_PRE_TYPE_WHITESPACE: diff --git a/src/unicode.cpp b/src/unicode.cpp index d2f166208f4a..10e8b408fbdd 100644 --- a/src/unicode.cpp +++ b/src/unicode.cpp @@ -520,6 +520,9 @@ static std::vector unicode_regex_split_custom_k2_horizon(const std::stri // regex: (?i:'s|'t|'re|'ve|'m|'ll|'d) // case insensitive if (cpt == '\'' && pos+1 < offset_end) { uint32_t cpt_next = unicode_tolower(_get_cpt(pos+1)); + if (cpt_next == 0x017F) { + cpt_next = 's'; // Unicode case-folding of long s + } if (cpt_next == 's' || cpt_next == 't' || cpt_next == 'm' || cpt_next == 'd') { pos += _add_token(pos+2); continue; @@ -1206,7 +1209,9 @@ static std::vector unicode_regex_split_custom(const std::string & text, } else if ( regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { bpe_offsets = unicode_regex_split_custom_qwen35(text, offsets); - } else if (regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { + } else if ( + regex_expr == "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+" || + regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { // K2-Horizon: llama3 splitter with marks + ZWNJ/ZWJ inside letter runs // (the generic std::regex fallback cannot parse \p{..} on MSVC) bpe_offsets = unicode_regex_split_custom_k2_horizon(text, offsets); @@ -1362,6 +1367,10 @@ bool unicode_cpt_is_han(uint32_t cpt) { } std::vector unicode_regex_split(const std::string & text, const std::vector & regex_exprs, bool byte_encode) { + if (text.empty()) { + return {}; + } + // unicode categories static const std::map k_ucat_enum = { { "\\p{N}", unicode_cpt_flags::NUMBER }, diff --git a/tests/test-unicode.cpp b/tests/test-unicode.cpp index 3ea2991fb46a..72a0f6ed0e09 100644 --- a/tests/test-unicode.cpp +++ b/tests/test-unicode.cpp @@ -106,7 +106,7 @@ int main() { { const std::vector k2_horizon_regex = { - "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" + "(?i:'s|'t|'re|'ve|'m|'ll|'d)|" "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", @@ -189,6 +189,20 @@ int main() { return 1; } + if (!check_k2( + "Unicode contraction", + "'\u017fa", + { "'\u017f", "a" })) { + return 1; + } + + if (!check_k2( + "empty input", + "", + { })) { + return 1; + } + if (!check_k2( "numbers", "123 1234 1234567", From ecf9741ea34e4d0bebe9b7fabea0cede48daa55b Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Sun, 27 Sep 2026 15:08:31 +0000 Subject: [PATCH 09/22] jinja : support sequence indices in selectattr and rejectattr Assisted-by: Codex --- common/jinja/value.cpp | 27 +++++++++++++++++++-------- tests/test-jinja.cpp | 38 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 57 insertions(+), 8 deletions(-) diff --git a/common/jinja/value.cpp b/common/jinja/value.cpp index 10130a6b0226..ce1d385df8e5 100644 --- a/common/jinja/value.cpp +++ b/common/jinja/value.cpp @@ -8,6 +8,7 @@ #include #include #include +#include #include #include #include @@ -261,6 +262,22 @@ static value tojson(const func_args & args) { return mk_val(json_str); } +// like jinja2, an all-digit attribute is an index into a sequence item, +// e.g. rejectattr('0', 'equalto', '$ref') on the (key, value) pairs of dict|items +static value get_attribute(const value & item, const value & attribute, value & default_val) { + const std::string attr = attribute->as_string().str(); + const bool is_index = !attr.empty() && attr.find_first_not_of("0123456789") == std::string::npos; + if (is_index && is_val(item)) { + int64_t index; + const auto result = std::from_chars(attr.data(), attr.data() + attr.size(), index); + return result.ec == std::errc() ? item->at(index, default_val) : default_val; + } + if (!is_val(item)) { + throw raised_exception("selectattr: item is not an object"); + } + return item->at(attribute, default_val); +} + template static value selectattr(const func_args & args) { args.ensure_count(2, 4); @@ -274,10 +291,7 @@ static value selectattr(const func_args & args) { if (args.count() == 2) { // example: array | selectattr("active") for (const auto & item : arr) { - if (!is_val(item)) { - throw raised_exception("selectattr: item is not an object"); - } - value attr_val = item->at(attribute, val_default); + value attr_val = get_attribute(item, attribute, val_default); bool is_selected = attr_val->as_bool(); if constexpr (is_reject) is_selected = !is_selected; if (is_selected) out->push_back(item); @@ -318,10 +332,7 @@ static value selectattr(const func_args & args) { } auto test_fn = it->second; for (const auto & item : arr) { - if (!is_val(item)) { - throw raised_exception("selectattr: item is not an object"); - } - value attr_val = item->at(attribute, val_default); + value attr_val = get_attribute(item, attribute, val_default); func_args test_args(args.ctx); test_args.push_back(attr_val); // attribute value test_args.push_back(extra_arg); // extra argument diff --git a/tests/test-jinja.cpp b/tests/test-jinja.cpp index 891b785c4f36..3d330235b130 100644 --- a/tests/test-jinja.cpp +++ b/tests/test-jinja.cpp @@ -1590,6 +1590,44 @@ static void test_array_methods(testing & t) { "b c " ); + test_template(t, "array|selectattr by index", + "{% for item in items|selectattr('1') %}{{ item[0] }} {% endfor %}", + {{"items", json::array({ + json::array({"a", false}), + json::array({"b", true}), + json::array({"c", true}) + })}}, + "b c " + ); + + test_template(t, "array|selectattr by index with operator", + "{% for item in items|selectattr('0', 'equalto', 'b') %}{{ item[1] }} {% endfor %}", + {{"items", json::array({ + json::array({"a", 1}), + json::array({"b", 2}), + json::array({"c", 3}) + })}}, + "2 " + ); + + test_template(t, "array|selectattr by index out of range", + "{% for item in items|selectattr('5') %}{{ item[0] }} {% endfor %}", + {{"items", json::array({json::array({"a", 1})})}}, + "" + ); + + test_template(t, "array|selectattr by index beyond int64", + "{% for item in items|selectattr('999999999999999999999999') %}{{ item[0] }}{% endfor %}", + {{"items", json::array({json::array({"a", 1})})}}, + "" + ); + + test_template(t, "dict|items|rejectattr by index", + "{% for k, v in obj|items|rejectattr('0', 'equalto', '$ref') %}{{ k }}={{ v }} {% endfor %}", + {{"obj", {{"$ref", "#/$defs/City"}, {"description", "origin"}}}}, + "description=origin " + ); + test_template(t, "array|tojson", "{{ arr|tojson }}", {{"arr", json::array({1, 2, 3})}}, From 63dded714c2a94474fa7bc4c570c89549fc12f31 Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Sun, 27 Sep 2026 15:09:21 +0000 Subject: [PATCH 10/22] model : add K2 Horizon dense and MoVA support Includes the K2 Horizon implementation from ifm-ai/llama.cpp with converter, tensor-parallel and model save/reload fixes. Assisted-by: Codex --- conversion/__init__.py | 2 + conversion/base.py | 8 + conversion/k2_horizon.py | 124 ++++++++++++ convert_hf_to_gguf_update.py | 3 + gguf-py/gguf/constants.py | 35 ++++ gguf-py/gguf/gguf_writer.py | 6 + gguf-py/gguf/tensor_mapping.py | 9 + src/llama-arch.cpp | 7 + src/llama-arch.h | 5 + src/llama-hparams.h | 2 + src/llama-model-saver.cpp | 2 + src/llama-model.cpp | 8 + src/llama-model.h | 4 + src/llama-vocab.cpp | 9 + src/llama-vocab.h | 1 + src/models/k2-horizon.cpp | 352 +++++++++++++++++++++++++++++++++ src/models/models.h | 14 ++ src/unicode.cpp | 157 +++++++++++++++ tests/test-llama-archs.cpp | 5 + tests/test-unicode.cpp | 224 +++++++++++++++++++-- 20 files changed, 964 insertions(+), 13 deletions(-) create mode 100644 conversion/k2_horizon.py create mode 100644 src/models/k2-horizon.cpp diff --git a/conversion/__init__.py b/conversion/__init__.py index 85db1d643183..7341ed6e8e5d 100644 --- a/conversion/__init__.py +++ b/conversion/__init__.py @@ -138,6 +138,8 @@ "JinaBertForMaskedLM": "bert", "JinaBertModel": "bert", "JinaEmbeddingsV5Model": "bert", + "K2HorizonForCausalLM": "k2_horizon", + "K2AuroraForCausalLM": "k2_horizon", # TODO: DELETE "KORMoForCausalLM": "qwen", "KimiK25ForConditionalGeneration": "deepseek", "KimiK3ForConditionalGeneration": "kimi_k3", diff --git a/conversion/base.py b/conversion/base.py index 5561481e7716..d257657bc3e8 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -234,6 +234,8 @@ def index_tensors(self, remote_hf_model_id: str | None = None) -> dict[str, Call prefix = "model" if not self.is_mistral_format else "consolidated" part_names: list[str] = ModelBase.get_model_part_names(self.dir_model, prefix, ".safetensors") + if not part_names and not self.is_mistral_format: + part_names = ModelBase.get_model_part_names(self.dir_model, "pytorch_model", ".safetensors") is_safetensors: bool = len(part_names) > 0 if not is_safetensors: part_names = ModelBase.get_model_part_names(self.dir_model, "pytorch_model", ".bin") @@ -1711,6 +1713,12 @@ def get_vocab_base_pre(self, tokenizer) -> str: if chkhsh == "0a766d034107bc736a3f2dc4968fd62e54a3570f1454443e0c5a4cc6bd7941ed": # ref: https://huggingface.co/XHToken/Spark-X2.5-1.7B res = "spark2_5" + if chkhsh == "1f9825a388f700a6b591722f17d470cbbcf10973ece35d2fd14239a14110ae1a": + # ref: https://huggingface.co/IFM/K2-Horizon-0.9B + res = "k2-horizon" + if chkhsh == "a9af07a84191f55098b248ae6f3dfe9e32d3190bebe8eafd91c1ddec9bc3449f": + # ref: https://huggingface.co/IFM/K2-Horizon-36B + res = "k2-horizon" if chkhsh == "0ef9807a4087ebef797fc749390439009c3b9eda9ad1a097abbe738f486c01e5": # ref: https://huggingface.co/meta-llama/Meta-Llama-3-8B res = "llama-bpe" diff --git a/conversion/k2_horizon.py b/conversion/k2_horizon.py new file mode 100644 index 000000000000..48f0b521fcb3 --- /dev/null +++ b/conversion/k2_horizon.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +import re +from collections.abc import Iterable +from typing import TYPE_CHECKING + +import torch + +if TYPE_CHECKING: + from torch import Tensor + +from .base import ModelBase, TextModel, gguf, logger + + +@ModelBase.register( + "K2HorizonForCausalLM", + "K2AuroraForCausalLM", # TODO: DELETE +) +@ModelBase.example("IFM/K2-Horizon-0.9B", "IFM/K2-Horizon-36B") +class K2HorizonModel(TextModel): + model_arch = gguf.MODEL_ARCH.K2HORIZON + + _experts: list[dict[str, Tensor]] | None = None + + def set_vocab(self): + super().set_vocab() + + # the 0.9B repo keeps an older chat_template.jinja next to the served chat_template_generation.jinja + tmpl_file = self.dir_model / "chat_template_generation.jinja" + if tmpl_file.is_file(): + self.gguf_writer.remove_key(gguf.Keys.Tokenizer.CHAT_TEMPLATE) + self.gguf_writer.add_chat_template(tmpl_file.read_text(encoding="utf-8")) + logger.info(f"gguf: using {tmpl_file.name} as the chat template") + + def set_gguf_parameters(self): + super().set_gguf_parameters() + hparams = self.hparams + + self.gguf_writer.add_group_norm_groups(int(hparams.get("layernorm_num_groups", 1))) + if (rope_head_dim := hparams.get("rope_head_dim")) is not None: + self.gguf_writer.add_rope_dimension_count(int(rope_head_dim)) + + if int(hparams.get("num_experts", 0)) > 0: + n_ff_exp = int(hparams["moe_intermediate_size"]) + n_shared = int(hparams.get("num_shared_experts", 0)) + + # the leading dense layers are the prefix of mlp_only_layers, unless given explicitly + n_dense = hparams.get("num_dense_layers") + if n_dense is None: + mlp_only_layers = {int(il) for il in hparams.get("mlp_only_layers", [])} + n_dense = 0 + while n_dense in mlp_only_layers: + n_dense += 1 + + gating_funcs = {"sigmoid": gguf.ExpertGatingFuncType.SIGMOID, "softmax": gguf.ExpertGatingFuncType.SOFTMAX} + router_func = hparams.get("router_score_func") + if router_func not in gating_funcs: + raise ValueError(f"Unsupported router_score_func: {router_func!r}") + + self.gguf_writer.add_expert_feed_forward_length(n_ff_exp) + self.gguf_writer.add_leading_dense_block_count(n_dense) + self.gguf_writer.add_moe_every_n_layers(int(hparams.get("decoder_sparse_step", 1))) + self.gguf_writer.add_expert_shared_count(n_shared) + self.gguf_writer.add_expert_weights_norm(bool(hparams.get("norm_topk_prob", False))) + if n_shared > 0: + self.gguf_writer.add_expert_shared_feed_forward_length(n_ff_exp * n_shared) + if (router_scale := hparams.get("router_scaling_factor")) is not None: + self.gguf_writer.add_expert_weights_scale(float(router_scale)) + self.gguf_writer.add_expert_gating_func(gating_funcs[router_func]) + + # MoVA + n_value_expert = int(hparams.get("mova_num_experts", 0)) + n_value_expert_used = int(hparams.get("mova_num_experts_per_tok", 0)) + if n_value_expert > 0 and n_value_expert_used > 0: + assert n_value_expert_used <= n_value_expert + self.gguf_writer.add_attention_value_expert_count(n_value_expert) + self.gguf_writer.add_attention_value_expert_used_count(n_value_expert_used) + + if (gate_func := hparams.get("attention_gate_func")) not in (None, "softplus"): + raise ValueError(f"Unsupported attention_gate_func: {gate_func!r}") + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + # the MoE router bias only selects experts + if name.endswith(".mlp.gate.bias"): + assert bid is not None + yield self.format_tensor_name(gguf.MODEL_TENSOR.FFN_EXP_PROBS_B, bid, ".bias"), data_torch + return + + if re.fullmatch(r"model\.layers\.\d+\.mlp\.experts\.\d+\.(down|gate|up)_proj\.weight", name): + yield from self._stack_experts(data_torch, name, bid, int(self.hparams["num_experts"]), + "model.layers.{bid}.mlp.experts.{xid}.{w}.weight", ("down_proj", "gate_proj", "up_proj")) + return + + if re.fullmatch(r"model\.layers\.\d+\.self_attn\.v_experts\.\d+\.weight", name): + yield from self._stack_experts(data_torch, name, bid, int(self.hparams["mova_num_experts"]), + "model.layers.{bid}.self_attn.v_experts.{xid}{w}.weight", ("",)) + return + + yield from super().modify_tensors(data_torch, name, bid) + + # collect the per-expert weights of a layer, then emit one stacked 3D tensor per projection + def _stack_experts(self, data_torch: Tensor, name: str, bid: int | None, n_experts: int, + fmt: str, projs: tuple[str, ...]) -> Iterable[tuple[str, Tensor]]: + assert bid is not None + if self._experts is None: + self._experts = [{} for _ in range(self.block_count)] + self._experts[bid][name] = data_torch + + names = {w: [fmt.format(bid=bid, xid=xid, w=w) for xid in range(n_experts)] for w in projs} + if not all(n in self._experts[bid] for ns in names.values() for n in ns): + return + + for w, ns in names.items(): + merged = torch.stack([self._experts[bid].pop(n) for n in ns], dim=0) + yield from super().modify_tensors(merged, fmt.replace(".{xid}", "").format(bid=bid, w=w), bid) + + def prepare_tensors(self): + super().prepare_tensors() + + if self._experts is not None: + # flatten the list of dicts + experts = [k for d in self._experts for k in d.keys()] + if len(experts) > 0: + raise ValueError(f"Unprocessed experts: {experts}") diff --git a/convert_hf_to_gguf_update.py b/convert_hf_to_gguf_update.py index 24b9bc075777..4067239da971 100755 --- a/convert_hf_to_gguf_update.py +++ b/convert_hf_to_gguf_update.py @@ -197,6 +197,9 @@ class TOKENIZER_TYPE(IntEnum): # no-op here); the gemma4 pre (escape ws, split on newlines only) matches it. {"name": "gemma4", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/danish-foundation-models/DFM-Mimir", "chkhsh": "846deafc5b0fa786186fa4ae6c7b49903cf2f1d1895bdb80b9120d60be135252"}, {"name": "spark2_5", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/XHToken/Spark-X2.5-1.7B", "chkhsh": "0a766d034107bc736a3f2dc4968fd62e54a3570f1454443e0c5a4cc6bd7941ed"}, + # K2 Horizon. 2 hashes because various sets of tokens depending on size + {"name": "k2-horizon", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/IFM/K2-Horizon-0.9B", "chkhsh": "1f9825a388f700a6b591722f17d470cbbcf10973ece35d2fd14239a14110ae1a"}, + {"name": "k2-horizon", "tokt": TOKENIZER_TYPE.BPE, "repo": "https://huggingface.co/IFM/K2-Horizon-36B", "chkhsh": "a9af07a84191f55098b248ae6f3dfe9e32d3190bebe8eafd91c1ddec9bc3449f"}, ] diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index 9a7a5e5bfaf6..41e3d06d8b9e 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -222,6 +222,8 @@ class Attention: RECURRENT_LAYERS = "{arch}.attention.recurrent_layers" TEMPERATURE_SCALE = "{arch}.attention.temperature_scale" ROPE_PATTERN = "{arch}.attention.rope_pattern" + VALUE_EXPERT_COUNT = "{arch}.attention.value_expert_count" + VALUE_EXPERT_USED_COUNT = "{arch}.attention.value_expert_used_count" class Indexer: HEAD_COUNT = "{arch}.attention.indexer.head_count" @@ -642,6 +644,7 @@ class MODEL_ARCH(IntEnum): NANBEIGE = auto() QWEN3TTS = auto() POCKETTTS = auto() + K2HORIZON = auto() class VISION_PROJECTOR_TYPE(IntEnum): @@ -908,6 +911,8 @@ class MODEL_TENSOR(IntEnum): INDEXER_COMPRESSOR_WGATE = auto() INDEXER_COMPRESSOR_APE = auto() INDEXER_COMPRESSOR_NORM = auto() + ATTN_V_GATE = auto() # k2-horizon MoVA router + ATTN_V_EXP = auto() # k2-horizon MoVA value experts # vision V_MMPROJ = auto() V_MMPROJ_FC = auto() @@ -1401,6 +1406,7 @@ class MODEL_TENSOR(IntEnum): MODEL_ARCH.NANBEIGE: "nanbeige", MODEL_ARCH.QWEN3TTS: "qwen3tts", MODEL_ARCH.POCKETTTS: "pockettts", + MODEL_ARCH.K2HORIZON: "k2-horizon", } VISION_PROJECTOR_TYPE_NAMES: dict[VISION_PROJECTOR_TYPE, str] = { @@ -1997,6 +2003,8 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.DFLASH_SELECTOR_NEXT: "selector_successor", MODEL_TENSOR.DFLASH_SELECTOR_HIDDEN: "selector_hidden", MODEL_TENSOR.D2T: "d2t", + MODEL_TENSOR.ATTN_V_GATE: "blk.{bid}.attn_v_gate", + MODEL_TENSOR.ATTN_V_EXP: "blk.{bid}.attn_v_exps", } MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = { @@ -5581,6 +5589,33 @@ class MODEL_TENSOR(IntEnum): MODEL_TENSOR.FFN_DOWN, MODEL_TENSOR.FFN_UP, ], + MODEL_ARCH.K2HORIZON: [ + MODEL_TENSOR.TOKEN_EMBD, + MODEL_TENSOR.OUTPUT_NORM, + MODEL_TENSOR.OUTPUT, + MODEL_TENSOR.ATTN_NORM, + MODEL_TENSOR.ATTN_Q, + MODEL_TENSOR.ATTN_Q_NORM, + MODEL_TENSOR.ATTN_K, + MODEL_TENSOR.ATTN_K_NORM, + MODEL_TENSOR.ATTN_V, + MODEL_TENSOR.ATTN_V_GATE, + MODEL_TENSOR.ATTN_V_EXP, + MODEL_TENSOR.ATTN_OUT, + MODEL_TENSOR.ATTN_GATE, + MODEL_TENSOR.FFN_NORM, + MODEL_TENSOR.FFN_GATE, + MODEL_TENSOR.FFN_UP, + MODEL_TENSOR.FFN_DOWN, + MODEL_TENSOR.FFN_GATE_INP, + MODEL_TENSOR.FFN_EXP_PROBS_B, + MODEL_TENSOR.FFN_GATE_EXP, + MODEL_TENSOR.FFN_UP_EXP, + MODEL_TENSOR.FFN_DOWN_EXP, + MODEL_TENSOR.FFN_GATE_SHEXP, + MODEL_TENSOR.FFN_UP_SHEXP, + MODEL_TENSOR.FFN_DOWN_SHEXP, + ], } # tensors that will not be serialized diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index cf7b367e520e..6a057ecc9b92 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -1591,6 +1591,12 @@ def add_xielu_beta(self, values: Sequence[float]): def add_xielu_eps(self, values: Sequence[float]): self.add_array(Keys.xIELU.EPS, values) + def add_attention_value_expert_count(self, count: int): + self.add_uint32(Keys.Attention.VALUE_EXPERT_COUNT.format(arch=self.arch), count) + + def add_attention_value_expert_used_count(self, count: int): + self.add_uint32(Keys.Attention.VALUE_EXPERT_USED_COUNT.format(arch=self.arch), count) + # diffusion models def add_diffusion_shift_logits(self, value: bool) -> None: diff --git a/gguf-py/gguf/tensor_mapping.py b/gguf-py/gguf/tensor_mapping.py index d2dfeece5952..89d39e5a19bb 100644 --- a/gguf-py/gguf/tensor_mapping.py +++ b/gguf-py/gguf/tensor_mapping.py @@ -394,6 +394,7 @@ class TensorNameMap: "model.layers.{bid}.self_attn.g_proj", # step3.5 head-wise attention gate "model.layers.{bid}.self_attn.output_gate", # minimax-01 "model.layers.{bid}.self_attn.linear_gate", # hy-v4 + "model.layers.{bid}.self_attn.attn_gate_proj", # k2-horizon ), # Feed-forward norm @@ -2757,6 +2758,14 @@ class TensorNameMap: MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM: ( "model.layers.{bid}.shared_head.norm", ), + + MODEL_TENSOR.ATTN_V_GATE: ( + "model.layers.{bid}.self_attn.v_router", # k2-horizon + ), + + MODEL_TENSOR.ATTN_V_EXP: ( + "model.layers.{bid}.self_attn.v_experts", # k2-horizon + ), } # architecture-specific block mappings diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index 8f1e239daef0..cb6424cd95da 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -158,6 +158,7 @@ static const std::map LLM_ARCH_NAMES = { { LLM_ARCH_NANBEIGE, "nanbeige" }, { LLM_ARCH_QWEN3TTS, "qwen3tts" }, { LLM_ARCH_POCKETTTS, "pockettts" }, + { LLM_ARCH_K2_HORIZON, "k2-horizon" }, { LLM_ARCH_UNKNOWN, "(unknown)" }, }; @@ -275,6 +276,8 @@ static const std::map LLM_KV_NAMES = { { LLM_KV_ATTENTION_SLIDING_WINDOW, "%s.attention.sliding_window" }, { LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, "%s.attention.sliding_window_pattern" }, { LLM_KV_ATTENTION_ROPE_PATTERN, "%s.attention.rope_pattern" }, + { LLM_KV_ATTENTION_VALUE_EXPERT_COUNT, "%s.attention.value_expert_count" }, + { LLM_KV_ATTENTION_VALUE_EXPERT_USED_COUNT, "%s.attention.value_expert_used_count" }, { LLM_KV_ATTENTION_SCALE, "%s.attention.scale" }, { LLM_KV_ATTENTION_OUTPUT_SCALE, "%s.attention.output_scale" }, @@ -706,6 +709,8 @@ static const std::map LLM_TENSOR_NAMES = { { LLM_TENSOR_DFLASH_SELECTOR_PREV, "selector_predecessor" }, { LLM_TENSOR_DFLASH_SELECTOR_NEXT, "selector_successor" }, { LLM_TENSOR_DFLASH_SELECTOR_HIDDEN, "selector_hidden" }, + { LLM_TENSOR_ATTN_V_GATE, "blk.%d.attn_v_gate" }, + { LLM_TENSOR_ATTN_V_EXPS, "blk.%d.attn_v_exps" }, }; // declare information about the model weight tensors: @@ -997,6 +1002,8 @@ static const std::map LLM_TENSOR_INFOS = { {LLM_TENSOR_DFLASH_SELECTOR_PREV, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_GET_ROWS}}, {LLM_TENSOR_DFLASH_SELECTOR_NEXT, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_GET_ROWS}}, {LLM_TENSOR_DFLASH_SELECTOR_HIDDEN, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ATTN_V_GATE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_ATTN_V_EXPS, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT_ID}}, }; LLM_KV::LLM_KV(llm_arch arch, const char * suffix) : arch(arch), suffix(suffix) {} diff --git a/src/llama-arch.h b/src/llama-arch.h index 23b6b38100c0..eaf1c2331674 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -163,6 +163,7 @@ enum llm_arch { LLM_ARCH_POCKETTTS, LLM_ARCH_MINIMAX_01, LLM_ARCH_HRM_TEXT, + LLM_ARCH_K2_HORIZON, LLM_ARCH_UNKNOWN, }; @@ -281,6 +282,8 @@ enum llm_kv { LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, LLM_KV_ATTENTION_SCALE, LLM_KV_ATTENTION_ROPE_PATTERN, + LLM_KV_ATTENTION_VALUE_EXPERT_COUNT, + LLM_KV_ATTENTION_VALUE_EXPERT_USED_COUNT, LLM_KV_ATTENTION_OUTPUT_SCALE, LLM_KV_ATTENTION_VALUE_SCALE, @@ -713,6 +716,8 @@ enum llm_tensor { LLM_TENSOR_DFLASH_SELECTOR_PREV, LLM_TENSOR_DFLASH_SELECTOR_NEXT, LLM_TENSOR_DFLASH_SELECTOR_HIDDEN, + LLM_TENSOR_ATTN_V_GATE, + LLM_TENSOR_ATTN_V_EXPS, }; diff --git a/src/llama-hparams.h b/src/llama-hparams.h index 73dffcc9f700..43e09bb11b66 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -71,6 +71,8 @@ struct llama_hparams { int32_t router_layer = -1; uint32_t n_expert = 0; uint32_t n_rel_attn_bkts = 0; + uint32_t n_value_expert = 0; // MoVA value experts (K2 Horizon) + uint32_t n_value_expert_used = 0; // TODO: this needs to be reworked int32_t n_layer_kv_from_start = -1; // if non-negative, the first n_layer_kv_from_start layers have KV cache diff --git a/src/llama-model-saver.cpp b/src/llama-model-saver.cpp index 6f58dd15092c..0ce0de86cda5 100644 --- a/src/llama-model-saver.cpp +++ b/src/llama-model-saver.cpp @@ -278,6 +278,8 @@ void llama_model_saver::add_kv_from_model() { add_kv(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); add_kv(LLM_KV_ATTENTION_GROUPNORM_EPS, hparams.f_norm_group_eps); add_kv(LLM_KV_ATTENTION_GROUPNORM_GROUPS, hparams.n_norm_groups); + add_kv(LLM_KV_ATTENTION_VALUE_EXPERT_COUNT, hparams.n_value_expert); + add_kv(LLM_KV_ATTENTION_VALUE_EXPERT_USED_COUNT, hparams.n_value_expert_used); add_kv(LLM_KV_ATTENTION_CAUSAL, hparams.causal_attn); add_kv(LLM_KV_ATTENTION_Q_LORA_RANK, hparams.n_lora_q); add_kv(LLM_KV_ATTENTION_KV_LORA_RANK, hparams.n_lora_kv); diff --git a/src/llama-model.cpp b/src/llama-model.cpp index ab5e744b5d10..f06e1de27900 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -344,6 +344,8 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params return new llama_model_step35(params); case LLM_ARCH_SPARK2_5: return new llama_model_spark2_5(params); + case LLM_ARCH_K2_HORIZON: + return new llama_model_k2_horizon(params); default: throw std::runtime_error(std::string("unsupported model architecture: '") + llm_arch_name(arch) + "'"); } @@ -381,6 +383,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str static const std::regex pattern_q_weight ("blk\\.\\d*\\.attn_q.weight"); static const std::regex pattern_kv_weight ("blk\\.\\d*\\.attn_(k|v).weight"); + static const std::regex pattern_v_exps_weight ("blk\\.\\d*\\.attn_v_exps.weight"); // K2 Horizon MoVA static const std::regex pattern_qkv_weight ("blk\\.\\d*\\.attn_qkv.weight"); static const std::regex pattern_q_bias ("blk\\.\\d*\\.attn_q\\.bias"); static const std::regex pattern_kv_bias ("blk\\.\\d*\\.attn_(k|v)\\.bias"); @@ -519,6 +522,10 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str if (std::regex_match(tensor_name, pattern_q_weight) || std::regex_match(tensor_name, pattern_kv_weight)) { return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1, "attn_output.weight", "ssm_out.weight"); } + // routed value experts {n_embd, n_embd_v_gqa, n_expert} produce V, so they split like attn_v + if (std::regex_match(tensor_name, pattern_v_exps_weight)) { + return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_1, "attn_output.weight"); + } if (std::regex_match(tensor_name, pattern_q_bias) || std::regex_match(tensor_name, pattern_kv_bias)) { return get_tensor_config_impl(GGML_BACKEND_SPLIT_AXIS_0, "attn_output.weight", "ssm_out.weight"); } @@ -3115,6 +3122,7 @@ llama_rope_type llama_model_rope_type(const llama_model * model) { case LLM_ARCH_STEP35: case LLM_ARCH_SPARK2_5: case LLM_ARCH_TALKIE: + case LLM_ARCH_K2_HORIZON: case LLM_ARCH_MELLUM: case LLM_ARCH_MAPLE: case LLM_ARCH_HRM_TEXT: diff --git a/src/llama-model.h b/src/llama-model.h index a0f9f11423e3..e478bd88558a 100644 --- a/src/llama-model.h +++ b/src/llama-model.h @@ -302,6 +302,10 @@ struct llama_layer { struct ggml_tensor * wv_enc = nullptr; struct ggml_tensor * wo_enc = nullptr; struct ggml_tensor * wqkv_gate = nullptr; + // K2 Horizon MoVA + struct ggml_tensor * attn_v_gate = nullptr; + struct ggml_tensor * attn_v_gate_b = nullptr; + struct ggml_tensor * attn_v_exps = nullptr; // relative position bias struct ggml_tensor * attn_rel_b = nullptr; diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index e038637ce707..80892e1c985e 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -551,6 +551,11 @@ struct llm_tokenizer_bpe : llm_tokenizer { "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}+| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", }; break; + case LLAMA_VOCAB_PRE_TYPE_K2_HORIZON: + regex_exprs = { + "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + break; case LLAMA_VOCAB_PRE_TYPE_WHITESPACE: // whitespace pre-tokenizer (jinaai/jina-embeddings-v2-base-zh) regex_exprs = { @@ -2419,6 +2424,10 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) { } else if ( tokenizer_pre == "mellum2") { pre_type = LLAMA_VOCAB_PRE_TYPE_MELLUM2; + } else if ( + tokenizer_pre == "k2-horizon") { + pre_type = LLAMA_VOCAB_PRE_TYPE_K2_HORIZON; + clean_spaces = false; } else { throw std::runtime_error(format("unknown pre-tokenizer type: '%s'", tokenizer_pre.c_str())); } diff --git a/src/llama-vocab.h b/src/llama-vocab.h index 3fb061f0ecf1..c132c82807d0 100644 --- a/src/llama-vocab.h +++ b/src/llama-vocab.h @@ -68,6 +68,7 @@ enum llama_vocab_pre_type { LLAMA_VOCAB_PRE_TYPE_HY_V4 = 57, LLAMA_VOCAB_PRE_TYPE_SPARK2_5 = 58, LLAMA_VOCAB_PRE_TYPE_UFAKZEKA = 59, + LLAMA_VOCAB_PRE_TYPE_K2_HORIZON = 60, }; struct LLM_KV; diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp new file mode 100644 index 000000000000..3c5991156788 --- /dev/null +++ b/src/models/k2-horizon.cpp @@ -0,0 +1,352 @@ +// K2 Horizon (MBZUAI IFM): grouped RMSNorm, optional per-head QK-norm and softplus +// attention output gate, DeepSeek-V3 style MoE (sigmoid router, selection bias, +// shared expert, leading dense layers) and MoVA: in MoE layers the V projection is +// replaced by routed value experts, V = sum_k w_k * silu(W_k x). + +#include "models.h" + +void llama_model_k2_horizon::load_arch_hparams(llama_model_loader & ml) { + ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); + ml.get_key(LLM_KV_ATTENTION_GROUPNORM_GROUPS, hparams.n_norm_groups, false); + if (hparams.n_norm_groups == 0) { + hparams.n_norm_groups = 1; + } + + if (hparams.n_expert > 0) { + ml.get_key_or_arr(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, hparams.n_ff_exp_arr, hparams.n_layer_all); + ml.get_key(LLM_KV_LEADING_DENSE_BLOCK_COUNT, hparams.n_layer_dense_lead, false); + ml.get_key(LLM_KV_MOE_EVERY_N_LAYERS, hparams.moe_every_n_layers, false); + ml.get_key(LLM_KV_EXPERT_SHARED_COUNT, hparams.n_expert_shared, false); + ml.get_key(LLM_KV_EXPERT_SHARED_FEED_FORWARD_LENGTH, hparams.n_ff_shexp, false); + ml.get_key(LLM_KV_EXPERT_WEIGHTS_SCALE, hparams.expert_weights_scale, false); + ml.get_key(LLM_KV_EXPERT_WEIGHTS_NORM, hparams.expert_weights_norm, false); + ml.get_key(LLM_KV_EXPERT_GATING_FUNC, hparams.expert_gating_func, false); + if (hparams.expert_gating_func == LLAMA_EXPERT_GATING_FUNC_TYPE_NONE) { + hparams.expert_gating_func = LLAMA_EXPERT_GATING_FUNC_TYPE_SIGMOID; + } + } + + // MoVA + ml.get_key(LLM_KV_ATTENTION_VALUE_EXPERT_COUNT, hparams.n_value_expert, false); + ml.get_key(LLM_KV_ATTENTION_VALUE_EXPERT_USED_COUNT, hparams.n_value_expert_used, false); + if (hparams.n_value_expert > 0) { + GGML_ASSERT(hparams.n_value_expert <= LLAMA_MAX_EXPERTS); + GGML_ASSERT(hparams.n_value_expert_used > 0); + GGML_ASSERT(hparams.n_value_expert_used <= hparams.n_value_expert); + } else { + GGML_ASSERT(hparams.n_value_expert_used == 0); + } + + if (hparams.n_layer() == 28 && hparams.n_embd == 1536) { + type = LLM_TYPE_1B; + } else if (hparams.n_layer() == 48 && hparams.n_embd == 2560) { + type = LLM_TYPE_36B; + } else { + type = LLM_TYPE_UNKNOWN; + } +} + +void llama_model_k2_horizon::load_arch_tensors(llama_model_loader &) { + LLAMA_LOAD_LOCALS; + + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0); + + // output + output_norm = create_tensor(tn(LLM_TENSOR_OUTPUT_NORM, "weight"), {n_embd}, 0); + output = create_tensor(tn(LLM_TENSOR_OUTPUT, "weight"), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED); + // if output is NULL, init from the input tok embed + if (output == NULL) { + output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, TENSOR_DUPLICATED); + } + + for (int i = 0; i < n_layer; ++i) { + auto & layer = layers[i]; + + const bool is_moe_layer = n_expert > 0 && (uint32_t) i >= hparams.n_layer_dense_lead; + const bool is_mova_layer = is_moe_layer && hparams.n_value_expert > 0; + + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, 0); + + layer.wq = create_tensor(tn(LLM_TENSOR_ATTN_Q, "weight", i), {n_embd, n_embd_head_k * n_head}, 0); + layer.wk = create_tensor(tn(LLM_TENSOR_ATTN_K, "weight", i), {n_embd, n_embd_k_gqa}, 0); + // one norm weight per head, stored flat; viewed as {head_dim, n_head} so it splits by head like Q/K + layer.attn_q_norm = create_tensor(tn(LLM_TENSOR_ATTN_Q_NORM, "weight", i), {n_embd_head_k, n_head}, TENSOR_NOT_REQUIRED | TENSOR_ALLOW_RESHAPE); + layer.attn_k_norm = create_tensor(tn(LLM_TENSOR_ATTN_K_NORM, "weight", i), {n_embd_head_k, n_head_kv}, TENSOR_NOT_REQUIRED | TENSOR_ALLOW_RESHAPE); + + if (is_mova_layer) { + layer.attn_v_gate = create_tensor(tn(LLM_TENSOR_ATTN_V_GATE, "weight", i), {n_embd, hparams.n_value_expert}, 0); + layer.attn_v_gate_b = create_tensor(tn(LLM_TENSOR_ATTN_V_GATE, "bias", i), {hparams.n_value_expert}, TENSOR_NOT_REQUIRED); + layer.attn_v_exps = create_tensor(tn(LLM_TENSOR_ATTN_V_EXPS, "weight", i), {n_embd, n_embd_v_gqa, hparams.n_value_expert}, 0); + } else { + layer.wv = create_tensor(tn(LLM_TENSOR_ATTN_V, "weight", i), {n_embd, n_embd_v_gqa}, 0); + } + + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), {n_embd_head_v * n_head, n_embd}, 0); + layer.wqkv_gate = create_tensor(tn(LLM_TENSOR_ATTN_GATE, "weight", i), {n_embd, n_embd_head_v * n_head}, TENSOR_NOT_REQUIRED); + + layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, 0); + + if (is_moe_layer) { + const int64_t n_ff_exp = hparams.n_ff_exp(i); + if (n_ff_exp == 0) { + throw std::runtime_error("K2 Horizon MoE layer requires expert_feed_forward_length"); + } + + 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_up_exps = create_tensor(tn(LLM_TENSOR_FFN_UP_EXPS, "weight", i), {n_embd, n_ff_exp, n_expert}, 0); + layer.ffn_gate_exps = create_tensor(tn(LLM_TENSOR_FFN_GATE_EXPS, "weight", 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}, 0); + + if (hparams.n_expert_shared > 0) { + const int64_t n_ff_shexp = hparams.n_ff_shexp > 0 ? hparams.n_ff_shexp : n_ff_exp * hparams.n_expert_shared; + + layer.ffn_up_shexp = create_tensor(tn(LLM_TENSOR_FFN_UP_SHEXP, "weight", i), {n_embd, n_ff_shexp}, 0); + layer.ffn_gate_shexp = create_tensor(tn(LLM_TENSOR_FFN_GATE_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); + } + } else { + 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}, 0); + layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), {n_ff, n_embd}, 0); + } + } +} + +std::unique_ptr llama_model_k2_horizon::build_arch_graph(const llm_graph_params & params) const { + return std::make_unique(*this, params); +} + +// RMS norm over n_groups equal slices of ne[0], then one full-width weight +static ggml_tensor * k2_horizon_group_rms_norm(ggml_context * ctx, ggml_tensor * cur, ggml_tensor * weight, int64_t n_groups, float eps) { + GGML_ASSERT(n_groups > 0 && cur->ne[0] % n_groups == 0); + + const int64_t n_embd = cur->ne[0]; + const int64_t n_tokens = cur->ne[1]; + + cur = ggml_reshape_3d(ctx, cur, n_embd / n_groups, n_groups, n_tokens); + cur = ggml_rms_norm(ctx, cur, eps); + cur = ggml_reshape_2d(ctx, cur, n_embd, n_tokens); + + return weight ? ggml_mul(ctx, cur, weight) : cur; +} + +// MoVA: route each token to n_value_expert_used value experts, V = sum_k w_k * silu(W_k x) +ggml_tensor * llama_model_k2_horizon::graph::build_routed_value(const llama_layer & layer, ggml_tensor * cur, int il) const { + const int64_t n_embd = cur->ne[0]; + const int64_t n_tokens = cur->ne[1]; + const int64_t n_embd_gqa = hparams.n_embd_v_gqa(il); + const int64_t n_values = hparams.n_value_expert; + const int64_t n_used = hparams.n_value_expert_used; + + ggml_tensor * logits = build_lora_mm(layer.attn_v_gate, cur); + ggml_tensor * probs = nullptr; + + switch ((llama_expert_gating_func_type) hparams.expert_gating_func) { + case LLAMA_EXPERT_GATING_FUNC_TYPE_SOFTMAX: probs = ggml_soft_max(ctx0, logits); break; + case LLAMA_EXPERT_GATING_FUNC_TYPE_SIGMOID: probs = ggml_sigmoid(ctx0, logits); break; + default: GGML_ABORT("unsupported K2 Horizon value-router gating function"); + } + + // the bias only affects which experts are selected, not their weights + ggml_tensor * selection_probs = probs; + if (layer.attn_v_gate_b) { + selection_probs = ggml_add(ctx0, probs, layer.attn_v_gate_b); + cb(selection_probs, "v_moe_probs_biased", il); + } + + ggml_tensor * selected_experts = ggml_argsort_top_k(ctx0, selection_probs, n_used); + + probs = ggml_reshape_3d(ctx0, probs, 1, n_values, n_tokens); + ggml_tensor * weights = ggml_get_rows(ctx0, probs, selected_experts); + + if (hparams.expert_weights_norm) { + weights = ggml_reshape_2d(ctx0, weights, n_used, n_tokens); + ggml_tensor * weights_sum = ggml_sum_rows(ctx0, weights); + weights_sum = ggml_clamp(ctx0, weights_sum, 6.103515625e-5f, INFINITY); + weights = ggml_div(ctx0, weights, weights_sum); + weights = ggml_reshape_3d(ctx0, weights, 1, n_used, n_tokens); + cb(weights, "v_moe_weights_norm", il); + } + + if (hparams.expert_weights_scale != 0.0f && hparams.expert_weights_scale != 1.0f) { + weights = ggml_scale(ctx0, weights, hparams.expert_weights_scale); + cb(weights, "v_moe_weights_scaled", il); + } + + cb(logits, "v_moe_logits", il); + cb(probs, "v_moe_probs", il); + cb(selected_experts->src[0], "v_moe_argsort", il); + cb(selected_experts, "v_moe_topk", il); + cb(weights, "v_moe_weights", il); + + ggml_tensor * values = build_lora_mm_id(layer.attn_v_exps, ggml_reshape_3d(ctx0, cur, n_embd, 1, n_tokens), selected_experts); + values = ggml_silu(ctx0, values); + values = ggml_mul(ctx0, values, weights); + cb(values, "v_moe_weighted", il); + + // sum the selected experts; 3D views of {n_embd_gqa, 1, n_tokens} keep the strides of values, + // which lets the tensor-parallel backend follow its split through the views + ggml_tensor * value_out = ggml_view_3d(ctx0, values, n_embd_gqa, 1, n_tokens, values->nb[1], values->nb[2], 0); + for (int64_t i = 1; i < n_used; ++i) { + value_out = ggml_add(ctx0, value_out, ggml_view_3d(ctx0, values, n_embd_gqa, 1, n_tokens, values->nb[1], values->nb[2], i * values->nb[1])); + } + if (n_used == 1) { + value_out = ggml_cont(ctx0, value_out); + } + cb(value_out, "Vcur_routed", il); + + return value_out; +} + +llama_model_k2_horizon::graph::graph(const llama_model & model, const llm_graph_params & params) : llm_graph_context(params) { + const int64_t n_embd_head = hparams.n_embd_head_v(); + GGML_ASSERT(n_embd_head == hparams.n_embd_head_k()); + + ggml_tensor * cur; + ggml_tensor * inpL; + + inpL = build_inp_embd(model.tok_embd); + + // inp_pos - contains the positions + ggml_tensor * inp_pos = build_inp_pos(); + + auto * inp_attn = build_attn_inp_kv(); + + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + const float kq_scale = 1.0f / sqrtf(float(n_embd_head)); + + for (int il = 0; il < n_layer; ++il) { + const auto & layer = model.layers[il]; + + res->t_layer_inp[il] = inpL; + + ggml_tensor * inpSA = inpL; + + const bool is_moe_layer = n_expert > 0 && (uint32_t) il >= hparams.n_layer_dense_lead; + const bool is_mova_layer = is_moe_layer && hparams.n_value_expert > 0; + + cur = k2_horizon_group_rms_norm(ctx0, inpL, layer.attn_norm, hparams.n_norm_groups, hparams.f_norm_rms_eps); + cb(cur, "attn_norm", il); + + // self-attention + { + ggml_tensor * attn_inp = cur; // saved for the output gate + + ggml_tensor * Qcur = build_lora_mm(layer.wq, cur, layer.wq_s); + ggml_tensor * Kcur = build_lora_mm(layer.wk, cur, layer.wk_s); + ggml_tensor * Vcur = is_mova_layer ? build_routed_value(layer, cur, il) : build_lora_mm(layer.wv, cur, layer.wv_s); + + Qcur = ggml_reshape_3d(ctx0, Qcur, n_embd_head, n_head, n_tokens); + Kcur = ggml_reshape_3d(ctx0, Kcur, n_embd_head, n_head_kv, n_tokens); + Vcur = ggml_reshape_3d(ctx0, Vcur, n_embd_head, n_head_kv, n_tokens); + + // per-head RMS norm with a separate weight for every head + if (layer.attn_q_norm) { + Qcur = build_norm(Qcur, layer.attn_q_norm, NULL, LLM_NORM_RMS, il); + } + if (layer.attn_k_norm) { + Kcur = build_norm(Kcur, layer.attn_k_norm, NULL, LLM_NORM_RMS, il); + } + + 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, "Qcur", il); + cb(Kcur, "Kcur", il); + cb(Vcur, "Vcur", il); + + // with an output gate, o_proj is applied after gating + const bool gated = layer.wqkv_gate != nullptr; + + cur = build_attn(inp_attn, + gated ? nullptr : layer.wo, gated ? nullptr : layer.wo_b, gated ? nullptr : layer.wo_s, + Qcur, Kcur, Vcur, nullptr, nullptr, nullptr, kq_scale, il); + + if (gated) { + // softplus with beta = ln(2): log2(1 + 2^x) + constexpr float ln2 = 0.6931471805599453f; + ggml_tensor * gate = build_lora_mm(layer.wqkv_gate, attn_inp, layer.wqkv_gate_s); + gate = ggml_scale(ctx0, gate, ln2); + gate = ggml_softplus(ctx0, gate); + gate = ggml_scale(ctx0, gate, 1.4426950408889634f); // 1 / ln(2) + + cur = ggml_mul(ctx0, cur, gate); + cur = build_lora_mm(layer.wo, cur, layer.wo_s); + if (layer.wo_b) { + cur = ggml_add(ctx0, cur, layer.wo_b); + } + } + } + + if (il == n_layer - 1 && inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids); + } + + ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpSA); + cb(ffn_inp, "ffn_inp", il); + + cur = k2_horizon_group_rms_norm(ctx0, ffn_inp, layer.ffn_norm, hparams.n_norm_groups, hparams.f_norm_rms_eps); + cb(cur, "ffn_norm", il); + + if (is_moe_layer) { + 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); + + if (layer.ffn_gate_shexp) { + ggml_tensor * ffn_shexp = build_ffn(cur, + layer.ffn_up_shexp, NULL, NULL, + layer.ffn_gate_shexp, NULL, NULL, + layer.ffn_down_shexp, NULL, NULL, + NULL, + LLM_FFN_SILU, LLM_FFN_PAR, il); + cur = ggml_add(ctx0, moe_out, ffn_shexp); + } else { + cur = moe_out; + } + } else { + cur = build_ffn(cur, + layer.ffn_up, NULL, NULL, + layer.ffn_gate, NULL, NULL, + layer.ffn_down, NULL, NULL, + NULL, + LLM_FFN_SILU, LLM_FFN_PAR, il); + } + cb(cur, "ffn_out", il); + + cur = ggml_add(ctx0, cur, ffn_inp); + cur = build_cvec(cur, il); + cb(cur, "l_out", il); + + // input for next layer + inpL = cur; + } + + cur = k2_horizon_group_rms_norm(ctx0, inpL, model.output_norm, hparams.n_norm_groups, hparams.f_norm_rms_eps); + cb(cur, "result_norm", -1); + res->t_embd = cur; + + // lm_head + cur = build_lora_mm(model.output, cur, model.output_s); + cb(cur, "result_output", -1); + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} diff --git a/src/models/models.h b/src/models/models.h index 3f9c67c63ca7..e0cb76d3f290 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -2660,3 +2660,17 @@ struct llama_model_spark2_5 : public llama_model_base { std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; }; + +struct llama_model_k2_horizon : public llama_model_base { + llama_model_k2_horizon(const struct llama_model_params & params) : llama_model_base(params) {} + void load_arch_hparams(llama_model_loader & ml) override; + void load_arch_tensors(llama_model_loader & ml) override; + + struct graph : public llm_graph_context { + graph(const llama_model & model, const llm_graph_params & params); + + ggml_tensor * build_routed_value(const llama_layer & layer, ggml_tensor * cur, int il) const; + }; + + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; +}; diff --git a/src/unicode.cpp b/src/unicode.cpp index 93996f9dd542..10e8b408fbdd 100644 --- a/src/unicode.cpp +++ b/src/unicode.cpp @@ -470,6 +470,153 @@ static std::vector unicode_regex_split_custom_llama3(const std::string & return bpe_offsets; } +static std::vector unicode_regex_split_custom_k2_horizon(const std::string & text, const std::vector & offsets) { + std::vector bpe_offsets; // store the offset of each word + bpe_offsets.reserve(offsets.size()); // Reserve memory for the approximate size + + const auto cpts = unicode_cpts_from_utf8(text); + + size_t start = 0; + for (auto offset : offsets) { + const size_t offset_ini = start; + const size_t offset_end = start + offset; + assert(offset_end <= cpts.size()); + start = offset_end; + + static const uint32_t OUT_OF_RANGE = 0xFFFFFFFF; + auto _get_cpt = [&] (const size_t pos) -> uint32_t { + return (offset_ini <= pos && pos < offset_end) ? cpts[pos] : OUT_OF_RANGE; + }; + + auto _get_flags = [&] (const size_t pos) -> unicode_cpt_flags { + return (offset_ini <= pos && pos < offset_end) ? unicode_cpt_flags_from_cpt(cpts[pos]) : unicode_cpt_flags{}; + }; + + // K2-Horizon: letter runs are (?:\p{L}|\p{M}|\u200C|\u200D)+ + auto _is_k2_letter = [&] (const size_t pos) -> bool { + const uint32_t c = _get_cpt(pos); + if (c == 0x200C || c == 0x200D) { + return true; + } + const auto f = _get_flags(pos); + return f.is_letter || f.is_accent_mark; + }; + + size_t _prev_end = offset_ini; + auto _add_token = [&] (const size_t end) -> size_t { + assert(_prev_end <= end && end <= offset_end); + size_t len = end - _prev_end; + if (len > 0) { + bpe_offsets.push_back(len); + } + _prev_end = end; + return len; + }; + + for (size_t pos = offset_ini; pos < offset_end; /*pos++*/ ) { + const uint32_t cpt = _get_cpt(pos); + const auto flags = _get_flags(pos); + + // regex: (?i:'s|'t|'re|'ve|'m|'ll|'d) // case insensitive + if (cpt == '\'' && pos+1 < offset_end) { + uint32_t cpt_next = unicode_tolower(_get_cpt(pos+1)); + if (cpt_next == 0x017F) { + cpt_next = 's'; // Unicode case-folding of long s + } + if (cpt_next == 's' || cpt_next == 't' || cpt_next == 'm' || cpt_next == 'd') { + pos += _add_token(pos+2); + continue; + } + if (pos+2 < offset_end) { + uint32_t cpt_next_next = unicode_tolower(_get_cpt(pos+2)); + if ((cpt_next == 'r' && cpt_next_next == 'e') || + (cpt_next == 'v' && cpt_next_next == 'e') || + (cpt_next == 'l' && cpt_next_next == 'l')) { + pos += _add_token(pos+3); + continue; + } + } + } + + // regex: [^\r\n\p{L}\p{N}]?(?:\p{L}|\p{M}|\u200C|\u200D)+ + if (!(cpt == '\r' || cpt == '\n' || flags.is_number)) { + if (_is_k2_letter(pos) || _is_k2_letter(pos+1)) { // one or more letters/marks/ZWNJ/ZWJ + pos++; + while (_is_k2_letter(pos)) { + pos++; + } + _add_token(pos); + continue; + } + } + + // regex: \p{N}{1,3} + if (flags.is_number) { + size_t ini = pos; + while (_get_flags(pos).is_number) { + if (++pos - ini >= 3 ) { + _add_token(pos); + ini = pos; + } + } + _add_token(pos); + continue; + } + + // regex: ?[^\s\p{L}\p{N}]+[\r\n]* + auto flags2 = (cpt == ' ' ? _get_flags(pos+1) : flags); + if (!(flags2.is_whitespace | flags2.is_letter | flags2.is_number) && flags.as_uint()) { + pos += (cpt == ' '); + while (!(flags2.is_whitespace | flags2.is_letter | flags2.is_number) && flags2.as_uint()) { + flags2 = _get_flags(++pos); + } + uint32_t cpt2 = _get_cpt(pos); + while (cpt2 == '\r' || cpt2 == '\n') { + cpt2 = _get_cpt(++pos); + } + _add_token(pos); + continue; + } + + size_t num_whitespaces = 0; + size_t last_end_r_or_n = 0; + while (_get_flags(pos+num_whitespaces).is_whitespace) { + uint32_t cpt2 = _get_cpt(pos+num_whitespaces); + if (cpt2 == '\r' || cpt2 == '\n') { + last_end_r_or_n = pos + num_whitespaces + 1; + } + num_whitespaces++; + } + + // regex: \s*[\r\n]+ + if (last_end_r_or_n > 0) { + pos = last_end_r_or_n; + _add_token(pos); + continue; + } + + // regex: \s+(?!\S) + if (num_whitespaces > 1 && _get_cpt(pos+num_whitespaces) != OUT_OF_RANGE) { + pos += num_whitespaces - 1; + _add_token(pos); + continue; + } + + // regex: \s+ + if (num_whitespaces > 0) { + pos += num_whitespaces; + _add_token(pos); + continue; + } + + // no matches + _add_token(++pos); + } + } + + return bpe_offsets; +} + // Qwen2 system regex: "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|\\p{N}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+" static std::vector unicode_regex_split_custom_qwen2(const std::string & text, const std::vector & offsets) { std::vector bpe_offsets; // store the offset of each word @@ -1062,6 +1209,12 @@ static std::vector unicode_regex_split_custom(const std::string & text, } else if ( regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { bpe_offsets = unicode_regex_split_custom_qwen35(text, offsets); + } else if ( + regex_expr == "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+" || + regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { + // K2-Horizon: llama3 splitter with marks + ZWNJ/ZWJ inside letter runs + // (the generic std::regex fallback cannot parse \p{..} on MSVC) + bpe_offsets = unicode_regex_split_custom_k2_horizon(text, offsets); } else if (regex_expr == "\\p{Han}+") { // K2's first pattern - handle all K2 patterns together bpe_offsets = unicode_regex_split_custom_kimi_k2(text, offsets); @@ -1214,6 +1367,10 @@ bool unicode_cpt_is_han(uint32_t cpt) { } std::vector unicode_regex_split(const std::string & text, const std::vector & regex_exprs, bool byte_encode) { + if (text.empty()) { + return {}; + } + // unicode categories static const std::map k_ucat_enum = { { "\\p{N}", unicode_cpt_flags::NUMBER }, diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp index 5a8a196c1f35..5be036ad0cb6 100644 --- a/tests/test-llama-archs.cpp +++ b/tests/test-llama-archs.cpp @@ -409,6 +409,10 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) { ms.add_kv(LLM_KV_EXPERT_GATING_FUNC, arch == LLM_ARCH_DEEPSEEK4 ? uint32_t(4) : uint32_t(2)); // sqrtsoftplus : sigmoid ms.add_kv(LLM_KV_EXPERT_GROUP_SCALE, 1.0f); ms.add_kv(LLM_KV_EXPERTS_PER_GROUP, uint32_t(1)); + if (arch == LLM_ARCH_K2_HORIZON) { + ms.add_kv(LLM_KV_ATTENTION_VALUE_EXPERT_COUNT, uint32_t(2)); + ms.add_kv(LLM_KV_ATTENTION_VALUE_EXPERT_USED_COUNT, uint32_t(2)); + } } ms.add_kv(LLM_KV_POSNET_EMBEDDING_LENGTH, n_embd); @@ -596,6 +600,7 @@ static bool moe_implemented(const llm_arch arch) { case LLM_ARCH_GRANITE_MOE: case LLM_ARCH_MISTRAL3: case LLM_ARCH_LLAMA_EMBED: + case LLM_ARCH_K2_HORIZON: return true; default: return false; diff --git a/tests/test-unicode.cpp b/tests/test-unicode.cpp index 2347d9000a8e..72a0f6ed0e09 100644 --- a/tests/test-unicode.cpp +++ b/tests/test-unicode.cpp @@ -1,23 +1,221 @@ #include "../src/unicode.h" #include +#include #include #include int main() { - const std::vector regex_exprs = { - "[~][A-Za-z]+| ?[\\p{S}]+|\\s+", - }; - const std::vector expected = { " ~", "foo" }; - const auto actual = unicode_regex_split(" ~foo", regex_exprs, false); - - if (actual != expected) { - fprintf(stderr, "unexpected split:"); - for (const auto & piece : actual) { - fprintf(stderr, " [%s]", piece.c_str()); - } - fprintf(stderr, "\n"); - return 1; + { + const std::vector regex_exprs = { + "[~][A-Za-z]+| ?[\\p{S}]+|\\s+", + }; + const std::vector expected = { " ~", "foo" }; + const auto actual = unicode_regex_split(" ~foo", regex_exprs, false); + + if (actual != expected) { + fprintf(stderr, "unexpected split:"); + for (const auto & piece : actual) { + fprintf(stderr, " [%s]", piece.c_str()); + } + fprintf(stderr, "\n"); + return 1; + } + } + + { + const std::vector regex_exprs = { + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" + "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" + "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" + "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + + const std::string input = "ab\u200ccd ef\u200dgh cafe\u0301"; + const std::vector expected = { + "ab\u200ccd", + " ef\u200dgh", + " cafe\u0301", + }; + + try { + const auto actual = unicode_regex_split(input, regex_exprs, false); + + if (actual != expected) { + fprintf(stderr, "unexpected K2-Horizon split:"); + for (const auto & piece : actual) { + fprintf(stderr, " [%s]", piece.c_str()); + } + fprintf(stderr, "\n"); + return 1; + } + } catch (const std::exception & e) { + fprintf(stderr, "K2-Horizon regex split threw exception: %s\n", e.what()); + return 1; + } + } + + + { + const std::vector llama3_regex = { + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" + "[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|" + "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" + "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + + const std::vector k2_horizon_regex = { + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" + "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" + "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" + "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + + const std::vector inputs = { + "Hello, world!", + "can't won't we're they'll", + "123 1234 1234567", + "Hello!!!\\nNext line", + "alpha beta gamma", + " café résumé", + }; + + for (const auto & input : inputs) { + const auto llama3 = unicode_regex_split(input, llama3_regex, false); + const auto k2 = unicode_regex_split(input, k2_horizon_regex, false); + + if (llama3 != k2) { + fprintf(stderr, "K2-Horizon diverged from Llama 3 for ordinary input: %s\n", input.c_str()); + + fprintf(stderr, "Llama 3:"); + for (const auto & piece : llama3) { + fprintf(stderr, " [%s]", piece.c_str()); + } + + fprintf(stderr, "\nK2-Horizon:"); + for (const auto & piece : k2) { + fprintf(stderr, " [%s]", piece.c_str()); + } + + fprintf(stderr, "\n"); + return 1; + } + } + } + + + { + const std::vector k2_horizon_regex = { + "(?i:'s|'t|'re|'ve|'m|'ll|'d)|" + "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" + "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" + "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + }; + + auto check_k2 = [&] ( + const char * name, + const std::string & input, + const std::vector & expected) { + const auto actual = unicode_regex_split(input, k2_horizon_regex, false); + + if (actual != expected) { + fprintf(stderr, "K2-Horizon reference mismatch for %s\n", name); + + fprintf(stderr, "expected:"); + for (const auto & piece : expected) { + fprintf(stderr, " [%s]", piece.c_str()); + } + + fprintf(stderr, "\nactual:"); + for (const auto & piece : actual) { + fprintf(stderr, " [%s]", piece.c_str()); + } + + fprintf(stderr, "\n"); + return false; + } + + return true; + }; + + if (!check_k2( + "ZWNJ", + "ab\u200ccd", + { "ab\u200ccd" })) { + return 1; + } + + if (!check_k2( + "ZWJ", + "ef\u200dgh", + { "ef\u200dgh" })) { + return 1; + } + + if (!check_k2( + "combined ZWNJ/ZWJ", + "ab\u200ccd ef\u200dgh", + { "ab\u200ccd", " ef\u200dgh" })) { + return 1; + } + + // The reference tokenizer NFC-normalizes cafe + combining acute + // to the composed form before pre-tokenization. + if (!check_k2( + "NFC accent", + "caf\u00e9", + { "caf\u00e9" })) { + return 1; + } + + if (!check_k2( + "Persian ZWNJ", + "\u0645\u06cc\u200c\u0631\u0648\u0645", + { "\u0645\u06cc\u200c\u0631\u0648\u0645" })) { + return 1; + } + + if (!check_k2( + "Devanagari ZWJ", + "\u0915\u094d\u200d\u0937", + { "\u0915\u094d\u200d\u0937" })) { + return 1; + } + + if (!check_k2( + "contractions", + "can't won't we're they'll", + { "can", "'t", " won", "'t", " we", "'re", " they", "'ll" })) { + return 1; + } + + if (!check_k2( + "Unicode contraction", + "'\u017fa", + { "'\u017f", "a" })) { + return 1; + } + + if (!check_k2( + "empty input", + "", + { })) { + return 1; + } + + if (!check_k2( + "numbers", + "123 1234 1234567", + { "123", " ", "123", "4", " ", "123", "456", "7" })) { + return 1; + } + + if (!check_k2( + "punctuation/newline", + "Hello!!!\nNext line", + { "Hello", "!!!\n", "Next", " line" })) { + return 1; + } } return 0; From 93819d9066717ec561551820719727329a0901a7 Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Sun, 27 Sep 2026 15:09:21 +0000 Subject: [PATCH 11/22] chat : support K2 Horizon reasoning and tool calls Assisted-by: Codex --- common/chat.cpp | 8 + common/jinja/caps.cpp | 79 +- common/parsers/k2-horizon.cpp | 195 +++++ common/parsers/parsers.h | 2 + common/parsers/sources.cmake | 1 + models/templates/IFM-K2-Horizon-7B.jinja | 994 +++++++++++++++++++++++ tests/test-chat.cpp | 81 ++ 7 files changed, 1330 insertions(+), 30 deletions(-) create mode 100644 common/parsers/k2-horizon.cpp create mode 100644 models/templates/IFM-K2-Horizon-7B.jinja diff --git a/common/chat.cpp b/common/chat.cpp index ed1942e15349..09d535a563b1 100644 --- a/common/chat.cpp +++ b/common/chat.cpp @@ -1133,6 +1133,14 @@ std::optional common_chat_try_specialized_template( return common_chat_params_init_kimi_k3(tmpl, params); } + // K2 Horizon - <|ifm|im_start|> turns, reasoning picked by reasoning_effort and + // sections; the three think tag pairs defeat the autoparser's reasoning detection + if (src.find("<|ifm|im_start|>") != std::string::npos && + src.find("") != std::string::npos) { + LOG_DBG("Using specialized template: K2 Horizon\n"); + return common_chat_params_init_k2_horizon(tmpl, params); + } + // Ling 3.0 / Bailing V3 - X sections with / tagged // tool calls. sections are unique to this family among the tagged-arg templates. if (src.find("ASSISTANT") != std::string::npos && diff --git a/common/jinja/caps.cpp b/common/jinja/caps.cpp index c5962ab77685..6ff17d10bb74 100644 --- a/common/jinja/caps.cpp +++ b/common/jinja/caps.cpp @@ -37,38 +37,57 @@ static void caps_try_execute(jinja::program & prog, const caps_ctx_fn & ctx_fn, const caps_json_fn & tools_fn, const caps_analyze_fn & analyze_fn) { - context ctx; - ctx.is_get_stats = true; - jinja::global_from_json(ctx, json{ - {"messages", messages_fn()}, - {"tools", tools_fn ? tools_fn() : json::array()}, - {"bos_token", ""}, - {"eos_token", ""}, - {"add_generation_prompt", true} - }, true); - - if (ctx_fn) { - ctx_fn(ctx); - } + json msgs = messages_fn(); + for (int attempt = 0; attempt < 2; attempt++) { + context ctx; + ctx.is_get_stats = true; + jinja::global_from_json(ctx, json{ + {"messages", msgs}, + {"tools", tools_fn ? tools_fn() : json::array()}, + {"bos_token", ""}, + {"eos_token", ""}, + {"add_generation_prompt", true} + }, true); + + if (ctx_fn) { + ctx_fn(ctx); + } - auto messages = ctx.get_val("messages"); - auto tools = ctx.get_val("tools"); - - bool success = false; - std::string result; - try { - jinja::runtime runtime(ctx); - auto results = runtime.execute(prog); - auto parts = jinja::runtime::gather_string_parts(results); - result = parts->as_string().str(); - success = true; - } catch (const std::exception & e) { - JJ_DEBUG("Exception during execution: %s", e.what()); - result = ""; - // ignore exceptions during capability analysis - } + auto messages = ctx.get_val("messages"); + auto tools = ctx.get_val("tools"); + + bool success = false; + std::string result; + try { + jinja::runtime runtime(ctx); + auto results = runtime.execute(prog); + auto parts = jinja::runtime::gather_string_parts(results); + result = parts->as_string().str(); + success = true; + } catch (const std::exception & e) { + JJ_DEBUG("Exception during execution: %s", e.what()); + result = ""; + // ignore exceptions during capability analysis + } + + // some templates require a thinking field on every assistant turn (e.g. K2 Horizon): + // retry once with an empty reasoning_content on the assistant turns that lack one + if (!success && attempt == 0) { + bool added = false; + for (auto & msg : msgs) { + if (msg.is_object() && msg.value("role", "") == "assistant" && !msg.contains("reasoning_content")) { + msg["reasoning_content"] = ""; + added = true; + } + } + if (added) { + continue; + } + } - analyze_fn(ctx, success, messages, tools, result); + analyze_fn(ctx, success, messages, tools, result); + return; + } } // for debugging only diff --git a/common/parsers/k2-horizon.cpp b/common/parsers/k2-horizon.cpp new file mode 100644 index 000000000000..4e0adf6cdd85 --- /dev/null +++ b/common/parsers/k2-horizon.cpp @@ -0,0 +1,195 @@ +#include "parsers.h" + +// K2 Horizon - reasoning effort picks one of three think tag pairs, tool calls are tagged: +// assistant := ... [content] +// [ {CALL} ] +// CALL (tool_call_format=xml, default) := name {k [t] +// v} +// CALL (tool_call_format=json) := {"name": name, "arguments": {...}} +// The generation prompt pre-opens the think block, so the model never emits the +// opening tag. Reasoning ends at the close tag or at a tool call section start. +common_chat_params common_chat_params_init_k2_horizon(const common_chat_template & tmpl, + const autoparser::generation_params & inputs) { + common_chat_params data; + + auto messages = inputs.messages; + for (auto & msg : messages) { + if (msg.value("role", "") == "assistant" && !msg.contains("reasoning_content")) { + msg["reasoning_content"] = ""; + } + } + + data.prompt = common_chat_template_direct_apply_impl(tmpl, inputs, messages); + data.generation_prompt = common_chat_template_generation_prompt_impl(tmpl, inputs, messages); + data.format = COMMON_CHAT_FORMAT_PEG_NATIVE; + data.supports_thinking = true; + + const std::string ROLE = "<|ifm|im_start|>assistant"; + const std::string TURN_END = "<|ifm|im_end|>"; + const std::string SECTION_START = ""; + const std::string SECTION_END = ""; + const std::string CALL_START = ""; + const std::string CALL_END = ""; + const std::string ARG_KEY = ""; + const std::string ARG_KEY_END = ""; + const std::string ARG_TYPE = ""; + const std::string ARG_TYPE_END = ""; + const std::string ARG_VAL = ""; + const std::string ARG_VAL_END = ""; + + // reasoning_effort high/medium/low opens //; + // the pair in use is the last one the generation prompt opened + std::string think = "ifm|think"; + size_t think_pos = std::string::npos; + for (const std::string tag : { "ifm|think", "ifm|think_fast", "ifm|think_faster" }) { + auto pos = data.generation_prompt.rfind("<" + tag + ">"); + if (pos != std::string::npos && (think_pos == std::string::npos || pos > think_pos)) { + think = tag; + think_pos = pos; + } + } + const std::string THINK_START = "<" + think + ">"; + const std::string THINK_END = ""; + + data.preserved_tokens = { + THINK_START, THINK_END, SECTION_START, SECTION_END, CALL_START, CALL_END, + ARG_KEY, ARG_KEY_END, ARG_TYPE, ARG_TYPE_END, ARG_VAL, ARG_VAL_END, TURN_END, + }; + + data.thinking_start_tag = THINK_START; + data.thinking_end_tags = { THINK_END, SECTION_START }; + + data.message_delimiters = { + { COMMON_CHAT_ROLE_ASSISTANT, "<|ifm|im_start|>assistant" }, + { COMMON_CHAT_ROLE_USER, "<|ifm|im_start|>user" }, + { COMMON_CHAT_ROLE_TOOL, "<|ifm|im_start|>tool" }, + { COMMON_CHAT_ROLE_SYSTEM, "<|ifm|im_start|>system" }, + }; + + // the turn ends with <|ifm|im_end|>, but only <|endoftext|> is EOG in the vocab + data.additional_stops = { TURN_END }; + + if (inputs.has_continuation()) { + const auto & msg = inputs.continue_msg; + + data.generation_prompt = ROLE + "\n" + THINK_START + "\n" + msg.reasoning_content; + if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) { + data.generation_prompt += THINK_END + msg.render_content(); + } + + data.prompt += data.generation_prompt; + } + + bool think_open = false; + if (inputs.has_continuation()) { + think_open = inputs.continue_final_message != COMMON_CHAT_CONTINUATION_CONTENT; + } else { + think_open = think_pos != std::string::npos && data.generation_prompt.find(THINK_END, think_pos) == std::string::npos; + } + + std::string call_format = "xml"; + if (inputs.extra_context.contains("tool_call_format") && inputs.extra_context.at("tool_call_format").is_string()) { + call_format = inputs.extra_context.at("tool_call_format"); + } + + auto has_tools = inputs.tools.is_array() && !inputs.tools.empty(); + auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE; + auto include_grammar = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE; + + auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) { + auto end = p.end(); + + // the effective parse input is generation_prompt + model output + auto opener = p.optional(p.literal(ROLE) + p.optional(p.space())); + + auto body_end = think_open ? p.until_one_of({ THINK_END, SECTION_START }) : p.until(THINK_END); + auto think_body = extract_reasoning ? p.reasoning(body_end) : p.content(body_end); + // the template writes "\n" and "\n"; those newlines are markup, not text + auto nl = p.optional(p.literal("\n")); + auto reasoning = p.optional(p.optional(p.literal(THINK_START) + nl) + think_body + p.optional(p.literal(THINK_END) + nl)); + + auto content = p.optional(p.content(p.until_one_of({ SECTION_START, TURN_END }))); + auto tail = p.optional(p.content(p.until(TURN_END))) + p.optional(p.literal(TURN_END)); + + if (!has_tools || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_NONE) { + return opener + reasoning + tail + end; + } + + auto tool_choices = p.choice(); + auto arg_close = p.tool_arg_close(p.literal(ARG_VAL_END)); + auto arg_string = p.rule("k2h-arg-string", p.tool_arg_string_value(p.until(ARG_VAL_END)) + arg_close); + auto arg_type = p.optional(p.optional(p.space()) + p.literal(ARG_TYPE) + p.until(ARG_TYPE_END) + p.literal(ARG_TYPE_END)); + + foreach_function(inputs.tools, [&](const json & tool) { + const auto & function = tool.at("function"); + std::string name = function.at("name"); + + if (call_format == "json") { + auto schema = common_chat_tool_parameters(function); + auto call = p.tool(p.tool_open(p.literal(CALL_START) + p.literal("{\"name\": \"") + p.tool_name(p.literal(name)) + + p.literal("\", \"arguments\": ")) + + p.tool_args(p.schema(p.json(), "k2h-tool-" + name + "-schema", schema)) + + p.tool_close(p.literal("}") + p.literal(CALL_END))); + tool_choices |= p.rule("k2h-tool-" + name, call); + return; + } + + // xml / xml_typed: strings are raw text up to the closing tag, other types are JSON + std::vector required_args; + std::vector optional_args; + foreach_parameter(function, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) { + auto rule_name = "k2h-arg-" + name + "-" + param.name; + auto types = param.schema->value_types(); + auto json_val = p.tool_arg_json_value(p.schema(p.json(), rule_name + "-schema", doc, *param.schema)) + arg_close; + + auto arg_value = types.is_only(common_chat_schema::TYPE_STRING) ? arg_string : + !types.has(common_chat_schema::TYPE_STRING) ? json_val : + p.gbnf(p.atomic(json_val) | arg_string, "k2h-arg-string"); + + auto arg = p.rule(rule_name, + p.optional(p.space()) + + p.tool_arg(p.tool_arg_open(p.literal(ARG_KEY) + p.tool_arg_name(p.literal(param.name)) + p.literal(ARG_KEY_END)) + + arg_type + p.optional(p.space()) + p.literal(ARG_VAL) + arg_value)); + + (param.required ? required_args : optional_args).push_back(arg); + }); + + auto args = p.permute("k2h-" + name + "-args", required_args); + if (!optional_args.empty()) { + args = args + p.zero_or_more(p.choice(optional_args)); + } + + auto call = p.tool(p.tool_open(p.literal(CALL_START) + p.tool_name(p.literal(name)) + p.optional(p.space())) + + p.tool_args(args) + + p.tool_close(p.optional(p.space()) + p.literal(CALL_END))); + tool_choices |= p.rule("k2h-tool-" + name, call); + }); + + auto calls = inputs.parallel_tool_calls ? tool_choices + p.zero_or_more(p.space() + tool_choices) : tool_choices; + + auto tools_section = p.trigger_rule("k2h-tool-call", + p.literal(SECTION_START) + p.space() + calls + p.space() + p.literal(SECTION_END)); + + // a required call follows the reasoning directly, as for gemma4 and gpt-oss + if (inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED) { + return opener + reasoning + p.optional(p.space()) + tools_section + tail + end; + } + + return opener + reasoning + content + p.optional(tools_section) + tail + end; + }); + + data.parser = parser.save(); + + if (include_grammar) { + data.grammar_lazy = inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED; + data.grammar = build_grammar([&](const common_grammar_builder & builder) { + parser.build_grammar(builder, data.grammar_lazy); + }); + + data.grammar_triggers = { + { COMMON_GRAMMAR_TRIGGER_TYPE_WORD, SECTION_START }, + }; + } + + return data; +} diff --git a/common/parsers/parsers.h b/common/parsers/parsers.h index f866360073bb..680420d06fa2 100644 --- a/common/parsers/parsers.h +++ b/common/parsers/parsers.h @@ -59,6 +59,8 @@ common_chat_params common_chat_params_init_gigachat_v3(const common_chat_templat common_chat_params common_chat_params_init_gpt_oss(const common_chat_template & tmpl, const autoparser::generation_params & inputs); +common_chat_params common_chat_params_init_k2_horizon(const common_chat_template & tmpl, const autoparser::generation_params & inputs); + common_chat_params common_chat_params_init_kimi_k2(const common_chat_template & tmpl, const autoparser::generation_params & inputs); common_chat_params common_chat_params_init_kimi_k3(const common_chat_template & tmpl, const autoparser::generation_params & inputs); diff --git a/common/parsers/sources.cmake b/common/parsers/sources.cmake index 70af84e25110..daaac29de490 100644 --- a/common/parsers/sources.cmake +++ b/common/parsers/sources.cmake @@ -9,6 +9,7 @@ set(LLAMA_CHAT_PARSERS_SOURCES ${CMAKE_CURRENT_LIST_DIR}/gemma4.cpp ${CMAKE_CURRENT_LIST_DIR}/gigachat-v3.cpp ${CMAKE_CURRENT_LIST_DIR}/gpt-oss.cpp + ${CMAKE_CURRENT_LIST_DIR}/k2-horizon.cpp ${CMAKE_CURRENT_LIST_DIR}/kimi-k2.cpp ${CMAKE_CURRENT_LIST_DIR}/kimi-k3.cpp ${CMAKE_CURRENT_LIST_DIR}/ling3.cpp diff --git a/models/templates/IFM-K2-Horizon-7B.jinja b/models/templates/IFM-K2-Horizon-7B.jinja new file mode 100644 index 000000000000..41dde1f4bc64 --- /dev/null +++ b/models/templates/IFM-K2-Horizon-7B.jinja @@ -0,0 +1,994 @@ +{{- bos_token }} +{%- if tool_presentation is defined -%} + {{- raise_exception("Unsupported argument: tool_presentation. Use tool_presentation_format with one of: json, xml, markdown.") -}} +{%- endif -%} +{%- if tool_calling_format is defined -%} + {{- raise_exception("Unsupported argument: tool_calling_format. Use tool_call_format with one of: json, xml, xml_typed.") -}} +{%- endif -%} +{%- if tool_format is defined -%} + {{- raise_exception("Unsupported argument: tool_format. Use tool_call_format with one of: json, xml, xml_typed.") -}} +{%- endif -%} +{%- set tool_presentation_fmt = tool_presentation_format | default('markdown') -%} +{%- set tool_call_fmt = tool_call_format | default('xml') -%} +{%- if tool_presentation_fmt != 'json' and tool_presentation_fmt != 'xml' and tool_presentation_fmt != 'markdown' -%} + {{- raise_exception("Unsupported tool_presentation_format: '" ~ tool_presentation_fmt ~ "'. Supported formats: json, xml, markdown.") -}} +{%- endif -%} +{%- if tool_call_fmt != 'json' and tool_call_fmt != 'xml' and tool_call_fmt != 'xml_typed' -%} + {{- raise_exception("Unsupported tool_call_format: '" ~ tool_call_fmt ~ "'. Supported formats: json, xml, xml_typed.") -}} +{%- endif -%} + +{#- Renderability state, computed during validate_tools (single walk, no extra -#} +{#- traversal at render time): ok = working flag for the tool being validated; -#} +{#- bad = pipe-delimited indices of tools that must render as verbatim JSON. -#} +{%- set RB = namespace(ok=true, bad='|') -%} + +{%- macro value_contains_mapping(v) -%} +{%- if v is mapping -%} +true +{%- elif v is sequence and v is not string -%} +{%- set f = namespace(x='false') -%} +{%- for c in v -%}{%- if value_contains_mapping(c) == 'true' -%}{%- set f.x = 'true' -%}{%- endif -%}{%- endfor -%} +{{- f.x -}} +{%- else -%} +false +{%- endif -%} +{%- endmacro -%} + +{#- $ref inlining state: defs = local $defs of the tool being rendered; seen = -#} +{#- pipe-delimited names already expanded for this tool (each def inlines at most -#} +{#- once; later references render by def name; cycles terminate immediately). -#} +{#- $ref-sibling annotations (description/default/...) merge OVER the def at -#} +{#- the inline site, so use-site annotations win and are never dropped. -#} +{%- set REFS = namespace(defs={}, seen='|') -%} + +{%- macro render_compact_type_name(type_name, spec) -%} +{%- if type_name == "array" -%} +array[{%- if 'items' in spec -%}{{ render_compact_type(spec['items']) }}{%- else -%}any{%- endif -%}] +{%- elif type_name -%} +{{- type_name -}} +{%- else -%} +any +{%- endif -%} +{%- endmacro -%} + +{%- macro render_compact_type(spec) -%} +{%- if spec is not mapping -%} +any +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string and spec.type | length > 0 -%} +{%- for type_name in spec.type -%}{{ render_compact_type_name(type_name, spec) }}{%- if not loop.last -%}|{%- endif -%}{%- endfor -%} +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string -%} +any +{%- elif spec.type -%} +{{- render_compact_type_name(spec.type, spec) -}} +{%- elif spec['$ref'] is string -%} +{{- spec['$ref'].split('/') | last -}} +{%- elif spec.oneOf -%} +oneOf[{%- for variant in spec.oneOf -%}{{ render_compact_type(variant) }}{%- if not loop.last -%}|{%- endif -%}{%- endfor -%}] +{%- elif spec.anyOf -%} +anyOf[{%- for variant in spec.anyOf -%}{{ render_compact_type(variant) }}{%- if not loop.last -%}|{%- endif -%}{%- endfor -%}] +{%- elif spec.properties -%} +object +{%- elif 'items' in spec -%} +array[{{ render_compact_type(spec['items']) }}] +{%- else -%} +any +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_type_name(type_name, spec) -%} +{%- if type_name == "array" -%} +array of {% if 'items' in spec %}{{ render_markdown_type(spec['items']) }}{% else %}any{% endif %} +{%- elif type_name -%} +{{- type_name -}} +{%- else -%} +any +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_type(spec) -%} +{%- if spec is sameas true -%} +True +{%- elif spec is sameas false -%} +False +{%- elif spec is not mapping -%} +any +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string and spec.type | length > 0 -%} +{%- for type_name in spec.type -%}{{ render_markdown_type_name(type_name, spec) }}{% if not loop.last %} or {% endif %}{%- endfor -%} +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string -%} +any +{%- elif spec.type -%} +{{- render_markdown_type_name(spec.type, spec) -}} +{%- elif spec['$ref'] is string -%} +{{- spec['$ref'].split('/') | last -}} +{%- elif spec.oneOf -%} +oneOf[{%- for variant in spec.oneOf -%}{{ render_markdown_type(variant) }}{% if not loop.last %} or {% endif %}{%- endfor -%}] +{%- elif spec.anyOf -%} +anyOf[{%- for variant in spec.anyOf -%}{{ render_markdown_type(variant) }}{% if not loop.last %} or {% endif %}{%- endfor -%}] +{%- elif spec.properties -%} +object +{%- elif 'items' in spec -%} +array of {{ render_markdown_type(spec['items']) }} +{%- else -%} +any +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_text(value) -%} +{{- value.split() | join(" ") -}} +{%- endmacro -%} + +{%- macro render_python_string(value) -%} +'{{- value.split() | join(" ") | replace("\\", "\\\\") | replace("'", "\\'") -}}' +{%- endmacro -%} + +{%- macro render_python_repr(value) -%} +{%- if value is string -%} +{{ render_python_string(value) }} +{%- elif value is sameas true -%} +True +{%- elif value is sameas false -%} +False +{%- elif value is none -%} +None +{%- elif value is mapping -%} +{{- "{" -}} +{%- for key, child in value | items -%} +{{ render_python_repr(key) }}: {{ render_python_repr(child) }}{%- if not loop.last -%}, {% endif -%} +{%- endfor -%} +{{- "}" -}} +{%- elif value is sequence -%} +{{- "[" -}} +{%- for child in value -%} +{{ render_python_repr(child) }}{%- if not loop.last -%}, {% endif -%} +{%- endfor -%} +{{- "]" -}} +{%- else -%} +{{- value -}} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_value(value) -%} +{%- if value is string -%}{{ render_xml_text(value) }}{%- else -%}{{ render_python_repr(value) }}{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_enum_value(value) -%} +{%- if value is string -%}"{{- value | replace("\\", "\\\\") | replace("\"", "\\\"") -}}"{%- else -%}"{{- render_python_repr(value) | replace("\\", "\\\\") | replace("\"", "\\\"") -}}"{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_enum(values) -%} +{%- for value in values -%}{{ render_xml_enum_value(value) }}{%- if not loop.last -%}|{%- endif -%}{%- endfor -%} +{%- endmacro -%} + +{%- macro render_xml_default_attr(value) -%} +{{- " default=" }}{%- if value is string -%}"{{- value | replace("\\", "\\\\") | replace("\"", "\\\"") -}}"{%- else -%}{{ render_xml_value(value) }}{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_attr(name, value) -%} +{{- " " + name + "=" }}{%- if value == "" -%}""{%- else -%}{{ render_xml_value(value) }}{%- endif -%} +{%- endmacro -%} + +{%- macro validate_schema(spec, path, lenient=false, classify=true, in_variant=false) -%} +{%- if spec is mapping -%} + {%- if not lenient -%} + {%- if spec.required is defined -%} + {%- if spec.required is string or spec.required is not sequence -%} + {{- raise_exception("Schema '" + path + "' has 'required' but it is not a list.") -}} + {%- endif -%} + {%- if spec.required | length > 0 and not spec.properties and not in_variant -%} + {{- raise_exception("Schema '" + path + "' has required fields but no properties object to define them.") -}} + {%- endif -%} + {%- if spec.properties -%} + {%- for required_name in spec.required -%} + {%- if required_name not in spec.properties -%} + {{- raise_exception("Schema '" + path + "' marks '" + required_name + "' as required, but that property is not defined in properties.") -}} + {%- endif -%} + {%- endfor -%} + {%- endif -%} + {%- endif -%} + {%- endif -%} + {#- renderability classification, piggybacking on this walk (no raises here): -#} + {#- constructs the pretty renderer does not fully handle flip RB.ok so the -#} + {#- tool falls back to verbatim JSON. Skipped entirely for json presentation. -#} + {%- if classify -%} + {%- for key, value in spec | items -%} + {%- if key == '$ref' -%} + {%- if value is not string -%}{%- set RB.ok = false -%} + {%- elif not (value.startswith('#/$defs/') or value.startswith('#/definitions/')) -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == '$defs' or key == 'definitions' -%} + {%- if value is mapping -%} + {%- for dk, dv in value | items -%} + {{- validate_schema(dv, path + ".$defs." + dk, true) -}} + {%- endfor -%} + {%- else -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == 'type' -%} + {%- if value is mapping -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == 'enum' -%} + {%- if value is string or value is mapping or value is not sequence -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == 'items' -%} + {#- any items shape renders: mapping structurally, others via repr detail -#} + {%- elif key == 'oneOf' or key == 'anyOf' -%} + {%- if value is mapping or value is string or value is not sequence -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif key == 'required' -%} + {%- if value and not spec.properties -%}{%- set RB.ok = false -%}{%- endif -%} + {%- elif ('|' ~ key ~ '|') in '|description|default|title|examples|properties|patternProperties|additionalProperties|returns|' -%} + {%- elif value is mapping -%} + {%- for uk, uv in value | items -%} + {%- if value_contains_mapping(uv) == 'true' -%}{%- set RB.ok = false -%}{%- endif -%} + {%- endfor -%} + {%- elif value is sequence and value is not string -%} + {%- if value_contains_mapping(value) == 'true' -%}{%- set RB.ok = false -%}{%- endif -%} + {%- endif -%} + {%- endfor -%} + {%- endif -%} + {%- if spec.properties -%} + {%- for child_name, child_spec in spec.properties | items -%} + {{- validate_schema(child_spec, path + "." + child_name, lenient, classify) -}} + {%- endfor -%} + {%- endif -%} + {%- if 'items' in spec -%}{{- validate_schema(spec['items'], path + "[]", lenient, classify) -}}{%- endif -%} + {%- if spec.oneOf -%} + {%- for variant in spec.oneOf -%}{{- validate_schema(variant, path + ".oneOf[" + (loop.index0 | string) + "]", lenient, classify, true) -}}{%- endfor -%} + {%- endif -%} + {%- if spec.anyOf -%} + {%- for variant in spec.anyOf -%}{{- validate_schema(variant, path + ".anyOf[" + (loop.index0 | string) + "]", lenient, classify, true) -}}{%- endfor -%} + {%- endif -%} + {%- if spec.additionalProperties is mapping -%}{{- validate_schema(spec.additionalProperties, path + ".additionalProperties", lenient, classify) -}}{%- endif -%} + {%- if spec.patternProperties is mapping -%} + {%- for pattern, pattern_spec in spec.patternProperties | items -%} + {{- validate_schema(pattern_spec, path + ".patternProperties[" + pattern + "]", lenient, classify) -}} + {%- endfor -%} + {%- endif -%} + {%- if spec.returns is mapping -%}{{- validate_schema(spec.returns, path + ".returns", lenient, classify) -}}{%- endif -%} +{%- endif -%} +{%- endmacro -%} + +{%- macro validate_tools(tools_list, classify=true) -%} +{%- set RB.bad = '|' -%} +{%- for tool in tools_list -%} + {%- set fn = tool.function if tool.function is defined else tool -%} + {%- set RB.ok = true -%} + {%- if fn.parameters is defined and fn.parameters is string -%} + {{- raise_exception("tool.function.parameters must be a dict, not a JSON string. Parse it before passing to the template.") -}} + {%- endif -%} + {%- if fn.parameters is not defined or fn.parameters is none -%} + {%- if fn.arguments is defined -%} + {{- raise_exception("Tool '" + fn.name + "' has 'arguments' instead of 'parameters'. Rename 'arguments' to 'parameters'.") -}} + {%- else -%} + {{- raise_exception("Tool '" + fn.name + "' is missing required 'parameters' field. Each tool must have a 'parameters' dict with 'type', 'properties', and 'required' keys.") -}} + {%- endif -%} + {%- endif -%} + {{- validate_schema(fn.parameters, "tool." + fn.name + ".parameters", false, classify) -}} + {%- if classify -%} + {%- if fn.parameters is mapping -%} + {#- unknown container-valued keys at the parameters ROOT are never rendered -#} + {#- by the pretty path (root extras are dropped) -> verbatim fallback. -#} + {%- for rk, rv in fn.parameters | items -%} + {%- if rk not in ['type', 'description', 'enum', 'default', 'properties', 'required', 'optional', 'title', 'items', 'oneOf', 'anyOf', 'additionalProperties', 'patternProperties', 'returns', 'examples', '$defs', 'definitions', '$ref'] -%} + {%- if rv is mapping or (rv is sequence and rv is not string) -%}{%- set RB.ok = false -%}{%- endif -%} + {%- endif -%} + {%- endfor -%} + {%- else -%} + {%- set RB.ok = false -%} + {%- endif -%} + {%- endif -%} + {%- if fn.returns is mapping -%}{{- validate_schema(fn.returns, "tool." + fn.name + ".returns", false, classify) -}}{%- endif -%} + {%- if classify and fn.returns is not defined and fn.response is mapping -%}{{- validate_schema(fn.response, "tool." + fn.name + ".response", true) -}}{%- endif -%} + {#- unknown container-valued keys at the FUNCTION level are never rendered -> fallback. -#} + {%- if classify -%} + {%- for fk, fv in fn | items -%} + {%- if fk not in ['name', 'description', 'parameters', 'returns', 'response', 'type', 'function'] -%} + {%- if fv is mapping or (fv is sequence and fv is not string) -%}{%- set RB.ok = false -%}{%- endif -%} + {%- endif -%} + {%- endfor -%} + {%- endif -%} + {%- if not RB.ok -%}{%- set RB.bad = RB.bad ~ loop.index0 ~ '|' -%}{%- endif -%} +{%- endfor -%} +{%- endmacro -%} + +{%- macro render_tools_json(tools_list) -%} +{{- "" }} +{%- for tool in tools_list %} +{{- "\n" }} +{{- tool | tojson }} +{%- endfor %} +{{- "\n" }} +{%- endmacro -%} + +{%- macro render_xml_schema_attrs(spec, include_value_attrs) -%} +{%- if spec is mapping -%} +{%- set structural_keys = ["type", "description", "enum", "default", "properties", "required", "items", "oneOf", "anyOf", "additionalProperties", "patternProperties", "returns"] -%} +{%- if include_value_attrs and spec.enum -%}{{- " enum=" }}{{ render_xml_enum(spec.enum) }}{%- endif -%} +{%- if include_value_attrs and spec.default is defined -%}{{ render_xml_default_attr(spec.default) }}{%- endif -%} +{%- if spec.additionalProperties is defined and spec.additionalProperties is not mapping -%}{{ render_xml_attr("additionalProperties", spec.additionalProperties) }}{%- endif -%} +{%- if spec.patternProperties is defined and spec.patternProperties is not mapping -%}{{ render_xml_attr("patternProperties", spec.patternProperties) }}{%- endif -%} +{%- for key, value in spec | items -%} + {%- if key not in structural_keys -%} +{{ render_xml_attr(key, value) }} + {%- endif -%} +{%- endfor -%} +{%- endif -%} +{%- endmacro -%} + +{%- macro xml_schema_has_children(spec, include_properties, include_description) -%} +{%- if spec is not mapping -%} +false +{%- elif (include_description and spec.description is defined) or (include_properties and spec.properties) or 'items' in spec or spec.oneOf or spec.anyOf or spec.additionalProperties is mapping or spec.patternProperties is mapping or spec.returns is defined -%} +true +{%- else -%} +false +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_schema_node(tag, spec, include_properties) -%} +{%- if spec is mapping and spec['$ref'] is string -%} + {%- set _r = spec['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- if _k is not none and ('|' + _k + '|') not in REFS.seen and REFS.defs[_k] is mapping -%} + {%- set spec = dict((REFS.defs[_k] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k + '|' -%} + {%- if spec['$ref'] is string -%} + {%- set _r2 = spec['$ref'] -%} + {%- set _k2 = _r2[8:] if _r2.startswith('#/$defs/') else (_r2[14:] if _r2.startswith('#/definitions/') else none) -%} + {%- if _k2 is not none and REFS.defs[_k2] is mapping -%} + {%- set spec = dict((REFS.defs[_k2] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k2 + '|' -%} + {%- endif -%} + {%- endif -%} + {%- endif -%} +{%- endif -%} +{%- if spec is mapping -%} +{{- "<" + tag + " type=" + render_compact_type(spec) }}{{ render_xml_schema_attrs(spec, true) }} +{%- if xml_schema_has_children(spec, include_properties, true) == 'true' -%} +{{- ">" }}{{ render_xml_schema_children(spec, include_properties, true) }}{{- "" }} +{%- else -%} +{{- "/>" }} +{%- endif -%} +{%- else -%} +{{- "<" + tag + ">" }}{{ render_xml_value(spec) }}{{- "" }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_pattern_property(pattern, spec) -%} +{%- if spec is mapping and spec['$ref'] is string -%} + {%- set _r = spec['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- if _k is not none and ('|' + _k + '|') not in REFS.seen and REFS.defs[_k] is mapping -%} + {%- set spec = dict((REFS.defs[_k] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k + '|' -%} + {%- if spec['$ref'] is string -%} + {%- set _r2 = spec['$ref'] -%} + {%- set _k2 = _r2[8:] if _r2.startswith('#/$defs/') else (_r2[14:] if _r2.startswith('#/definitions/') else none) -%} + {%- if _k2 is not none and REFS.defs[_k2] is mapping -%} + {%- set spec = dict((REFS.defs[_k2] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k2 + '|' -%} + {%- endif -%} + {%- endif -%} + {%- endif -%} +{%- endif -%} +{%- if spec is mapping -%} +{{- "" }}{{ render_xml_schema_children(spec, true, true) }}{{- "" }} +{%- else -%} +{{- "/>" }} +{%- endif -%} +{%- else -%} +{{- "" }}{{ render_xml_value(spec) }}{{- "" }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_schema_children(spec, include_properties, include_description) -%} +{%- if include_description and spec.description is defined -%}{{- "" }}{{ spec.description }}{{- "" }}{%- endif -%} +{%- if include_properties and spec.properties -%} +{%- for child_name, child_spec in spec.properties | items -%} +{{- render_xml_param(child_name, child_spec, spec.required or []) }} +{%- endfor -%} +{%- endif -%} +{%- if 'items' in spec -%}{{ render_xml_schema_node("items", spec['items'], true) }}{%- endif -%} +{%- if spec.oneOf -%} +{{- "" }} +{%- for variant in spec.oneOf -%}{{ render_xml_schema_node("variant", variant, true) }}{%- endfor -%} +{{- "" }} +{%- endif -%} +{%- if spec.anyOf -%} +{{- "" }} +{%- for variant in spec.anyOf -%}{{ render_xml_schema_node("variant", variant, true) }}{%- endfor -%} +{{- "" }} +{%- endif -%} +{%- if spec.additionalProperties is mapping -%}{{ render_xml_schema_node("additionalProperties", spec.additionalProperties, true) }}{%- endif -%} +{%- if spec.patternProperties is mapping -%} +{{- "" }} +{%- for pattern, pattern_spec in spec.patternProperties | items -%}{{ render_xml_pattern_property(pattern, pattern_spec) }}{%- endfor -%} +{{- "" }} +{%- elif spec.patternProperties is defined -%}{{ render_xml_value(spec.patternProperties) }}{%- endif -%} +{%- if spec.returns is mapping -%}{{ render_xml_schema_node("returns", spec.returns, true) }}{%- elif spec.returns is defined -%}{{ render_xml_value(spec.returns) }}{%- endif -%} +{%- endmacro -%} + +{%- macro render_xml_param(name, spec, required_list) -%} +{%- if spec is mapping and spec['$ref'] is string -%} + {%- set _r = spec['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- if _k is not none and ('|' + _k + '|') not in REFS.seen and REFS.defs[_k] is mapping -%} + {%- set spec = dict((REFS.defs[_k] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k + '|' -%} + {%- if spec['$ref'] is string -%} + {%- set _r2 = spec['$ref'] -%} + {%- set _k2 = _r2[8:] if _r2.startswith('#/$defs/') else (_r2[14:] if _r2.startswith('#/definitions/') else none) -%} + {%- if _k2 is not none and REFS.defs[_k2] is mapping -%} + {%- set spec = dict((REFS.defs[_k2] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k2 + '|' -%} + {%- endif -%} + {%- endif -%} + {%- endif -%} +{%- endif -%} +{{- "" }} +{%- if spec.description -%}{{ spec.description }}{%- endif -%} +{{- render_xml_schema_children(spec, true, false) }} +{{- "" }} +{%- else -%} +{{- "/>" }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_tools_xml(tools_list) -%} +{{- "" }} +{%- for tool in tools_list -%} + {%- set fn = tool.function if tool.function is defined else tool -%} + {%- set REFS.defs = fn.parameters['$defs'] if (fn.parameters is mapping and fn.parameters['$defs'] is mapping) else (fn.parameters['definitions'] if (fn.parameters is mapping and fn.parameters['definitions'] is mapping) else {}) -%} + {%- set REFS.seen = '|' -%} + {%- set fnp = namespace(p=fn.parameters) -%} + {%- if fnp.p is mapping and fnp.p['$ref'] is string -%} + {%- set _r = fnp.p['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- if _k is not none and REFS.defs[_k] is mapping -%} + {%- set fnp.p = dict((REFS.defs[_k] | items | list) + (fnp.p | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k + '|' -%} + {%- endif -%} + {%- endif -%} +{{- "\n" }} +{%- if fn.description -%} +{{- "" }}{{ fn.description }}{{- "" }} +{%- endif -%} +{{- "" }} +{%- if fnp.p and fnp.p.properties -%} + {%- for pname, pspec in fnp.p.properties | items -%} +{{- render_xml_param(pname, pspec, fnp.p.required or []) }} + {%- endfor -%} +{%- elif fnp.p is mapping and (fnp.p.oneOf or fnp.p.anyOf or 'items' in fnp.p) -%} +{{- render_xml_schema_children(fnp.p, true, false) }} +{%- endif -%} +{{- "" }} +{%- set fn_ret = fn.returns if fn.returns is defined else fn.response -%} +{%- if fn_ret is mapping -%}{{ render_xml_schema_node("returns", fn_ret, true) }}{%- elif fn_ret is defined -%}{{ render_xml_value(fn_ret) }}{%- endif -%} +{{- "" }} +{%- endfor -%} +{{- "\n" }} +{%- endmacro -%} + +{%- macro render_markdown_literal(value) -%} +{%- if value is string and value == "" -%}"" +{%- elif value is string -%}`{{ value | replace("\n", "\\n") }}` +{%- else -%}`{{ render_python_repr(value) }}` +{%- endif -%} +{%- endmacro -%} + +{%- macro render_allowed_values(values) -%} +{%- for value in values -%}{{ render_markdown_literal(value) }}{% if not loop.last %}, {% endif %}{%- endfor -%} +{%- endmacro -%} + +{%- macro render_markdown_value(value) -%} +{%- if value is string and value == "" -%}""{%- elif value is string -%}{{ value }}{%- else -%}{{ render_python_repr(value) }}{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_detail(indent, label, value) -%} +{{- "\n" + indent + " - " + label + ": " }}{{ render_markdown_value(value) }} +{%- endmacro -%} + +{%- macro render_markdown_metadata_detail(label, value) -%} +{{- "\n- " + label + ": " }}{{ render_markdown_value(value) }} +{%- endmacro -%} + +{%- macro render_markdown_schema_annotations(spec, indent, include_value_details) -%} +{%- if include_value_details and spec.description is defined -%}{{ render_markdown_detail(indent, "Description", spec.description | replace("\n", "\n" + indent + " ")) }}{%- endif -%} +{%- if include_value_details and spec.enum is defined -%}{{- "\n" + indent + " - Allowed values: " }}{{ render_allowed_values(spec.enum) }}{%- endif -%} +{%- if include_value_details and spec.default is defined -%}{{- "\n" + indent + " - Default: " }}{{ render_markdown_literal(spec.default) }}{%- endif -%} +{%- if spec.additionalProperties is defined -%} + {%- if spec.additionalProperties is mapping -%} +{{- "\n" + indent + " - Additional properties *(" + render_markdown_type(spec.additionalProperties) + ")*" }} +{{- render_markdown_schema_details(spec.additionalProperties, indent + " ", true) }} + {%- else -%} +{{ render_markdown_detail(indent, "Additional properties", spec.additionalProperties) }} + {%- endif -%} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_metadata_annotations(spec) -%} +{%- if spec.description is defined -%}{{ render_markdown_metadata_detail("Description", spec.description | replace("\n", "\n ")) }}{%- endif -%} +{%- if spec.enum is defined -%}{{- "\n- Allowed values: " }}{{ render_allowed_values(spec.enum) }}{%- endif -%} +{%- if spec.default is defined -%}{{- "\n- Default: " }}{{ render_markdown_literal(spec.default) }}{%- endif -%} +{%- if spec.additionalProperties is defined -%} + {%- if spec.additionalProperties is mapping -%} +{{- "\n- Additional properties *(" + render_markdown_type(spec.additionalProperties) + ")*" }} +{{- render_markdown_schema_details(spec.additionalProperties, "", true) }} + {%- else -%} +{{ render_markdown_metadata_detail("Additional properties", spec.additionalProperties) }} + {%- endif -%} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_schema_extras(spec, indent) -%} +{%- set rendered_keys = ["type", "description", "enum", "default", "properties", "required", "items", "oneOf", "anyOf", "additionalProperties", "patternProperties", "returns"] -%} +{%- for key, value in spec | items -%} + {%- if key not in rendered_keys -%} +{{- "\n" + indent + " - " + key + ": " }}{{ render_markdown_value(value) }} + {%- endif -%} +{%- endfor -%} +{%- endmacro -%} + +{%- macro render_markdown_metadata_extras(spec) -%} +{%- set rendered_keys = ["type", "description", "enum", "default", "properties", "required", "items", "oneOf", "anyOf", "additionalProperties", "patternProperties", "returns"] -%} +{%- for key, value in spec | items -%} + {%- if key not in rendered_keys -%} +{{- "\n- " + key + ": " }}{{ render_markdown_value(value) }} + {%- endif -%} +{%- endfor -%} +{%- endmacro -%} + +{%- macro markdown_schema_has_extra(spec) -%} +{%- set rendered_keys = ["type", "description", "enum", "default", "properties", "required", "items", "oneOf", "anyOf", "additionalProperties", "patternProperties", "returns"] -%} +{%- set found = namespace(value='false') -%} +{%- for key, value in spec | items -%} + {%- if key not in rendered_keys -%}{%- set found.value = 'true' -%}{%- endif -%} +{%- endfor -%} +{{- found.value -}} +{%- endmacro -%} + +{%- macro markdown_parameter_schema_has_details(spec) -%} +{%- if spec.description is defined or spec.enum is defined or spec.default is defined or spec.additionalProperties is defined or spec.patternProperties is defined or 'items' in spec or spec.oneOf or spec.anyOf or spec.returns is defined or markdown_schema_has_extra(spec) == 'true' -%} +true +{%- else -%} +false +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_schema_structure(spec, indent, include_properties) -%} +{%- if include_properties and spec.properties -%} + {%- for child_name, child_spec in spec.properties | items -%} +{{- render_markdown_param(child_name, child_spec, spec.required or [], indent + " ") }} + {%- endfor -%} +{%- endif -%} +{%- if 'items' in spec and spec['items'] is mapping -%} +{{- "\n" + indent + " - Items *(" + render_markdown_type(spec['items']) + ")*" }} +{{- render_markdown_schema_details(spec['items'], indent + " ", true) }} +{%- elif 'items' in spec -%} +{{ render_markdown_detail(indent, "Items", spec['items']) }} +{%- endif -%} +{%- if spec.oneOf -%} +{{- "\n" + indent + " - oneOf:" }} + {%- for variant in spec.oneOf -%} +{{- "\n" + indent + " - Variant " }}{{ loop.index }}{{- " *(" + render_markdown_type(variant) + ")*" }} +{{- render_markdown_schema_details(variant, indent + " ", true) }} + {%- endfor -%} +{%- endif -%} +{%- if spec.anyOf -%} +{{- "\n" + indent + " - anyOf:" }} + {%- for variant in spec.anyOf -%} +{{- "\n" + indent + " - Variant " }}{{ loop.index }}{{- " *(" + render_markdown_type(variant) + ")*" }} +{{- render_markdown_schema_details(variant, indent + " ", true) }} + {%- endfor -%} +{%- endif -%} +{%- if spec.patternProperties is mapping -%} +{{- "\n" + indent + " - Pattern properties:" }} + {%- for pattern, pattern_spec in spec.patternProperties | items -%} + {%- if pattern_spec is mapping -%} +{{- "\n" + indent + " - `" + pattern + "` *(" + render_markdown_type(pattern_spec) + ")*" }} +{{- render_markdown_schema_details(pattern_spec, indent + " ", true) }} + {%- else -%} +{{- "\n" + indent + " - `" + pattern + "`: " }}{{ render_markdown_value(pattern_spec) }} + {%- endif -%} + {%- endfor -%} +{%- elif spec.patternProperties is defined -%} +{{ render_markdown_detail(indent, "Pattern properties", spec.patternProperties) }} +{%- endif -%} +{%- if spec.returns is mapping -%} +{{- "\n" + indent + " - Returns *(" + render_markdown_type(spec.returns) + ")*" }} +{{- render_markdown_schema_details(spec.returns, indent + " ", true) }} +{%- elif spec.returns is defined -%} +{{ render_markdown_detail(indent, "Returns", spec.returns) }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_schema_details(spec, indent, include_value_details) -%} +{%- if spec is mapping and spec['$ref'] is string -%} + {%- set _r = spec['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- if _k is not none and ('|' + _k + '|') not in REFS.seen and REFS.defs[_k] is mapping -%} + {%- set spec = dict((REFS.defs[_k] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k + '|' -%} + {%- if spec['$ref'] is string -%} + {%- set _r2 = spec['$ref'] -%} + {%- set _k2 = _r2[8:] if _r2.startswith('#/$defs/') else (_r2[14:] if _r2.startswith('#/definitions/') else none) -%} + {%- if _k2 is not none and REFS.defs[_k2] is mapping -%} + {%- set spec = dict((REFS.defs[_k2] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k2 + '|' -%} + {%- endif -%} + {%- endif -%} + {%- endif -%} +{%- endif -%} +{%- if spec is mapping -%} +{{- render_markdown_schema_annotations(spec, indent, include_value_details) }} +{{- render_markdown_schema_structure(spec, indent, true) }} +{{- render_markdown_schema_extras(spec, indent) }} +{%- elif spec is not sameas true and spec is not sameas false -%} +{{- "\n" + indent + " - Value: " }}{{ render_markdown_literal(spec) }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_parameter_schema(spec) -%} +{%- if spec is mapping and spec['$ref'] is string -%} + {%- set _r = spec['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- if _k is not none and ('|' + _k + '|') not in REFS.seen and REFS.defs[_k] is mapping -%} + {%- set spec = dict((REFS.defs[_k] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k + '|' -%} + {%- if spec['$ref'] is string -%} + {%- set _r2 = spec['$ref'] -%} + {%- set _k2 = _r2[8:] if _r2.startswith('#/$defs/') else (_r2[14:] if _r2.startswith('#/definitions/') else none) -%} + {%- if _k2 is not none and REFS.defs[_k2] is mapping -%} + {%- set spec = dict((REFS.defs[_k2] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k2 + '|' -%} + {%- endif -%} + {%- endif -%} + {%- endif -%} +{%- endif -%} +{%- if spec is mapping -%} +{{- render_markdown_metadata_annotations(spec) }} +{%- if 'items' in spec and spec['items'] is mapping -%} +{{- "\n- Items *(" + render_markdown_type(spec['items']) + ")*" }} +{{- render_markdown_schema_details(spec['items'], "", true) }} +{%- elif 'items' in spec -%} +{{ render_markdown_metadata_detail("Items", spec['items']) }} +{%- endif -%} +{%- if spec.oneOf -%} +{{- "\n- oneOf:" }} + {%- for variant in spec.oneOf -%} +{{- "\n - Variant " }}{{ loop.index }}{{- " *(" + render_markdown_type(variant) + ")*" }} +{{- render_markdown_schema_details(variant, " ", true) }} + {%- endfor -%} +{%- endif -%} +{%- if spec.anyOf -%} +{{- "\n- anyOf:" }} + {%- for variant in spec.anyOf -%} +{{- "\n - Variant " }}{{ loop.index }}{{- " *(" + render_markdown_type(variant) + ")*" }} +{{- render_markdown_schema_details(variant, " ", true) }} + {%- endfor -%} +{%- endif -%} +{%- if spec.patternProperties is mapping -%} +{{- "\n- Pattern properties:" }} + {%- for pattern, pattern_spec in spec.patternProperties | items -%} + {%- if pattern_spec is mapping -%} +{{- "\n - `" + pattern + "` *(" + render_markdown_type(pattern_spec) + ")*" }} +{{- render_markdown_schema_details(pattern_spec, " ", true) }} + {%- else -%} +{{- "\n - `" + pattern + "`: " }}{{ render_markdown_value(pattern_spec) }} + {%- endif -%} + {%- endfor -%} +{%- elif spec.patternProperties is defined -%} +{{ render_markdown_metadata_detail("Pattern properties", spec.patternProperties) }} +{%- endif -%} +{%- if spec.returns is mapping -%} +{{- "\n- Returns *(" + render_markdown_type(spec.returns) + ")*" }} +{{- render_markdown_schema_details(spec.returns, "", true) }} +{%- elif spec.returns is defined -%} +{{ render_markdown_metadata_detail("Returns", spec.returns) }} +{%- endif -%} +{{- render_markdown_metadata_extras(spec) }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_markdown_param(name, spec, required_list, indent) -%} +{%- if spec is mapping and spec['$ref'] is string -%} + {%- set _r = spec['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- if _k is not none and ('|' + _k + '|') not in REFS.seen and REFS.defs[_k] is mapping -%} + {%- set spec = dict((REFS.defs[_k] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k + '|' -%} + {%- if spec['$ref'] is string -%} + {%- set _r2 = spec['$ref'] -%} + {%- set _k2 = _r2[8:] if _r2.startswith('#/$defs/') else (_r2[14:] if _r2.startswith('#/definitions/') else none) -%} + {%- if _k2 is not none and REFS.defs[_k2] is mapping -%} + {%- set spec = dict((REFS.defs[_k2] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k2 + '|' -%} + {%- endif -%} + {%- endif -%} + {%- endif -%} +{%- endif -%} +{{- "\n" + indent + "- `" + name + "` *(" + render_markdown_type(spec) }} +{%- if name in (required_list or []) -%}{{- ", required" }}{%- endif -%} +{{- ")*" }} +{%- if spec.description -%}{{- " - " + spec.description | replace("\n", "\n" + indent + " ") }}{%- endif -%} +{%- if spec.enum -%} +{{- "\n" + indent + " - Allowed values: " }}{{ render_allowed_values(spec.enum) }} +{%- endif -%} +{%- if spec.default is defined -%} +{{- "\n" + indent + " - Default: " }}{{ render_markdown_literal(spec.default) }} +{%- endif -%} +{{- render_markdown_schema_details(spec, indent, false) }} +{%- endmacro -%} + +{%- macro render_tools_markdown(tools_list) -%} +{{- "" }} +{%- for tool in tools_list -%} + {%- set fn = tool.function if tool.function is defined else tool -%} + {%- set REFS.defs = fn.parameters['$defs'] if (fn.parameters is mapping and fn.parameters['$defs'] is mapping) else (fn.parameters['definitions'] if (fn.parameters is mapping and fn.parameters['definitions'] is mapping) else {}) -%} + {%- set REFS.seen = '|' -%} + {%- set fnp = namespace(p=fn.parameters) -%} + {%- if fnp.p is mapping and fnp.p['$ref'] is string -%} + {%- set _r = fnp.p['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- if _k is not none and REFS.defs[_k] is mapping -%} + {%- set fnp.p = dict((REFS.defs[_k] | items | list) + (fnp.p | items | rejectattr('0', 'equalto', '$ref') | list)) -%} + {%- set REFS.seen = REFS.seen + _k + '|' -%} + {%- endif -%} + {%- endif -%} +{{- "\n## " + fn.name }} +{%- if fn.description -%} +{{- "\n" + fn.description }} +{%- endif -%} +{{- "\n\n**Parameters**" }} +{%- if fnp.p and fnp.p.properties -%} + {%- for pname, pspec in fnp.p.properties | items -%} +{{- render_markdown_param(pname, pspec, fnp.p.required or [], "") }} + {%- endfor -%} +{%- elif fnp.p is mapping and (fnp.p.oneOf or fnp.p.anyOf or 'items' in fnp.p) -%} +{{- render_markdown_parameter_schema(fnp.p) }} +{%- else -%} +{{- "\n- None" }} +{%- endif -%} +{%- set fn_ret = fn.returns if fn.returns is defined else fn.response -%} +{%- if fn_ret is mapping -%} +{{- "\n\n**Returns**" }} +{{- "\n- Return *(" + render_markdown_type(fn_ret) + ")*" }} +{{- render_markdown_schema_details(fn_ret, "", true) }} +{%- elif fn_ret is defined -%} +{{- "\n\n**Returns**\n- " }}{{ render_markdown_value(fn_ret) }} +{%- endif -%} +{%- if not loop.last -%}{{- "\n" }}{%- endif -%} +{%- endfor -%} +{{- "\n" }} +{%- endmacro -%} + +{%- macro render_tool_presentation(tools_list, fmt) -%} +{%- if fmt == 'json' -%} +{{- render_tools_json(tools_list) }} +{%- elif RB.bad != '|' -%} +{#- some tool uses constructs the pretty renderers cannot represent (verdicts -#} +{#- computed during validate_tools): render the WHOLE toolset exactly as the -#} +{#- json presentation would, so the block stays uniform and model-familiar. -#} +{{- render_tools_json(tools_list) }} +{%- elif fmt == 'xml' -%} +{{- render_tools_xml(tools_list) }} +{%- elif fmt == 'markdown' -%} +{{- render_tools_markdown(tools_list) }} +{%- else -%} +{{- raise_exception("Unsupported tool_presentation_format: '" + fmt + "'. Supported formats: json, xml, markdown.") }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_call_instructions(fmt) -%} +{%- if fmt == 'json' -%} +{{- "Wrap all tool calls in a single block. For each call, emit one JSON object with the function name and arguments on the same line inside tags:\n\n\n{\"name\": , \"arguments\": }\n" }} +{%- elif fmt == 'xml' -%} +{{- "Wrap all tool calls in a single block. For each call, write the function name at the start of , followed by paired and tags for each argument:\n\n\n$FUNCTION_NAME\n$PARAMETER_NAME\n$PARAMETER_VALUE\n...\n\n\n\nString and scalar parameters should be written as plain text. Array and object parameters should be written as JSON literals." }} +{%- elif fmt == 'xml_typed' -%} +{{- "Wrap all tool calls in a single block. For each call, write the function name at the start of , followed by , , and tags for each argument:\n\n\n$FUNCTION_NAME\n$PARAMETER_NAME\n$ARGUMENT_TYPE\n$PARAMETER_VALUE\n...\n\n\n\nUse the parameter type shown in the tool definition. If that type contains anyOf or oneOf, use the actual argument value type instead. String and scalar parameters should be written as plain text. Array and object parameters should be written as JSON literals." }} +{%- else -%} +{{- raise_exception("Unsupported tool_call_format: '" + fmt + "'. Supported formats: json, xml, xml_typed.") }} +{%- endif -%} +{%- endmacro -%} + +{%- macro render_system_with_tools(tools_list, system_content, presentation_fmt, call_fmt) -%} +{{- "<|ifm|im_start|>system\n# Tools\nYou may call one or more tools to assist with the user query.\n\nAvailable tools are:\n\n" }} +{{- render_tool_presentation(tools_list, presentation_fmt) }} +{{- "\n\nWhen calling tools, you MUST follow the tool-call format below:\n\n" }} +{{- render_call_instructions(call_fmt) }} +{%- if system_content -%} +{{- "\n\n" + system_content }} +{%- endif -%} +{{- "<|ifm|im_end|>" }} +{%- endmacro -%} + +{%- macro render_argument_value(value) -%} +{%- if value is string -%}{{- value -}}{%- else -%}{{- value | tojson -}}{%- endif -%} +{%- endmacro -%} + +{%- macro render_value_type(value) -%} +{%- if value is none -%}null +{%- elif value is boolean -%}boolean +{%- elif value is integer -%}integer +{%- elif value is number -%}number +{%- elif value is string -%}string +{%- elif value is mapping -%}object +{%- elif value is sequence -%}array +{%- else -%}any +{%- endif -%} +{%- endmacro -%} + +{%- macro schema_has_combinator(spec) -%} +{%- if spec.oneOf or spec.anyOf -%} +true +{%- elif spec.type is defined and spec.type is sequence and spec.type is not string and spec.type | length > 1 -%} +true +{%- elif spec.type == "array" and 'items' in spec -%} +{{- schema_has_combinator(spec['items']) -}} +{%- elif spec.properties -%} + {%- set found = namespace(value='false') -%} + {%- for child_name, child_spec in spec.properties | items -%} + {%- if schema_has_combinator(child_spec) == 'true' -%} + {%- set found.value = 'true' -%} + {%- endif -%} + {%- endfor -%} +{{- found.value -}} +{%- else -%} +false +{%- endif -%} +{%- endmacro -%} + +{%- macro render_arg_type(tools_list, tool_name, arg_name, value) -%} +{%- set found = namespace(type='any') -%} +{%- for tool in tools_list -%} + {%- set fn = tool.function if tool.function is defined else tool -%} + {%- if fn.name == tool_name and fn.parameters and fn.parameters.properties and arg_name in fn.parameters.properties -%} + {%- set spec = fn.parameters.properties[arg_name] -%} + {%- if spec is mapping and spec['$ref'] is string -%} + {%- set _r = spec['$ref'] -%} + {%- set _k = _r[8:] if _r.startswith('#/$defs/') else (_r[14:] if _r.startswith('#/definitions/') else none) -%} + {%- set _d = fn.parameters['$defs'] if fn.parameters['$defs'] is mapping else fn.parameters['definitions'] -%} + {%- set spec = dict((_d[_k] | items | list) + (spec | items | rejectattr('0', 'equalto', '$ref') | list)) if (_k is not none and _d is mapping and _d[_k] is mapping) else spec -%} + {%- endif -%} + {%- if schema_has_combinator(spec) == 'true' -%} + {%- set found.type = render_value_type(value) -%} + {%- else -%} + {%- set found.type = render_compact_type(spec) -%} + {%- endif -%} + {%- endif -%} +{%- endfor -%} +{{- found.type -}} +{%- endmacro -%} + +{%- macro render_tool_calls_block(tool_calls, fmt, tools_list) -%} +{{- "" }} +{%- for raw_tool_call in tool_calls -%} + {%- set tool_call = raw_tool_call.function if raw_tool_call.function else raw_tool_call -%} + {%- if tool_call.arguments is string -%} + {{- raise_exception("tool_call.arguments must be a dict, not a JSON string. Parse it before passing to the template.") -}} + {%- endif -%} + {%- if fmt == 'json' -%} +{{- "\n{\"name\": \"" + tool_call.name + "\", \"arguments\": " }}{{ tool_call.arguments | tojson }}{{- "}" }} + {%- elif fmt == 'xml' or fmt == 'xml_typed' -%} +{{- "\n" + tool_call.name + "\n" }} + {%- for key, value in tool_call.arguments | items -%} +{{- "" + key + "\n" }} +{%- if fmt == 'xml_typed' -%} +{{- "" + render_arg_type(tools_list, tool_call.name, key, value) + "\n" }} +{%- endif -%} +{{- "" }}{{ render_argument_value(value) }}{{- "\n" }} + {%- endfor -%} +{{- "" }} + {%- else -%} + {{- raise_exception("Unsupported tool_call_format: '" + fmt + "'. Supported formats: json, xml, xml_typed.") -}} + {%- endif -%} +{%- endfor -%} +{{- "\n" }} +{%- endmacro -%} + +{%- macro render_tool_response_messages(raw_content) -%} +{%- if raw_content is string -%} +{{- '<|ifm|im_start|>tool\n' + raw_content + '<|ifm|im_end|>' }} +{%- elif raw_content is sequence and raw_content is not string and raw_content is not mapping -%} + {%- if raw_content | length == 0 -%} + {{- raise_exception("tool message content list must not be empty.") -}} + {%- endif -%} +{{- '<|ifm|im_start|>tool\n' -}} + {%- for item in raw_content -%} + {%- if not loop.first -%}{{- '\n' -}}{%- endif -%} + {%- if item is string -%} +{{- item -}} + {%- elif item is mapping and item.text is string -%} +{{- item.text -}} + {%- else -%} +{{- (item | tojson) -}} + {%- endif -%} + {%- endfor -%} +{{- '<|ifm|im_end|>' -}} +{%- else -%} +{{- '<|ifm|im_start|>tool\n' }}{{ raw_content | tojson }}{{- '<|ifm|im_end|>' }} +{%- endif -%} +{%- endmacro -%} + +{%- set available_tools = tools if tools else [] -%} +{%- if (not available_tools) and messages[0].role == 'system' and messages[0].get('tools') -%} + {%- set available_tools = messages[0]['tools'] -%} +{%- endif -%} +{%- if available_tools -%} + {{- validate_tools(available_tools, tool_presentation_fmt != 'json') }} + {%- set system_content = '' -%} + {%- if messages[0].role == 'system' and messages[0].content -%} + {%- set system_content = messages[0].content -%} + {%- endif -%} + {{- render_system_with_tools(available_tools, system_content, tool_presentation_fmt, tool_call_fmt) }} +{%- else -%} + {%- if messages[0].role == 'system' -%} + {{- '<|ifm|im_start|>system\n' + messages[0].content + '<|ifm|im_end|>' }} + {%- endif -%} +{%- endif -%} + +{%- for message in messages -%} + {%- if message.content is string -%} + {%- set content = message.content -%} + {%- else -%} + {%- set content = '' -%} + {%- endif -%} + {%- if (message.role == "user") or (message.role == "system" and not loop.first) -%} + {{- '<|ifm|im_start|>' + message.role + '\n' + content + '<|ifm|im_end|>' }} + {%- elif message.role == "assistant" -%} + {%- set thinking_content = '' -%} + {%- set think_tag = 'ifm|think' -%} + {%- if message.think is defined and message.think is string -%} + {%- set thinking_content = message.think -%} + {%- set think_tag = 'ifm|think' -%} + {%- elif message.think_fast is defined and message.think_fast is string -%} + {%- set thinking_content = message.think_fast -%} + {%- set think_tag = 'ifm|think_fast' -%} + {%- elif message.think_faster is defined and message.think_faster is string -%} + {%- set thinking_content = message.think_faster -%} + {%- set think_tag = 'ifm|think_faster' -%} + {%- elif message.reasoning_content is defined and message.reasoning_content is string -%} + {%- set thinking_content = message.reasoning_content -%} + {%- set think_tag = 'ifm|think' -%} + {%- elif message.reasoning is defined and message.reasoning is string -%} + {%- set thinking_content = message.reasoning -%} + {%- set think_tag = 'ifm|think' -%} + {%- elif message.think is not defined and message.reasoning is not defined and message.reasoning_content is not defined and message.think_fast is not defined and message.think_faster is not defined -%} + {{- raise_exception("Assistant message is missing a thinking field. Provide one of: think, reasoning, reasoning_content, think_fast, think_faster.") -}} + {%- else -%} + {{- raise_exception("Assistant thinking fields must be strings. Provide one of: think, reasoning, reasoning_content, think_fast, think_faster as a string.") -}} + {%- endif -%} + {{- '<|ifm|im_start|>' + message.role }} + {% generation %} + {%- if think_tag -%} + {%- if thinking_content -%} + {{- '<' + think_tag + '>\n' + thinking_content + '' + content }} + {%- else -%} + {{- '<' + think_tag + '>\n' + content }} + {%- endif -%} + {%- else -%} + {{- content }} + {%- endif -%} + {%- if message.tool_calls -%} + {{- render_tool_calls_block(message.tool_calls, tool_call_fmt, available_tools) }} + {%- endif -%} + {{- '<|ifm|im_end|>' -}} + {%- endgeneration -%} + {%- elif message.role == "tool" -%} + {{- render_tool_response_messages(message.content) }} + {%- endif -%} +{%- endfor -%} +{%- if add_generation_prompt -%} + {%- set effort = reasoning_effort | default('high') -%} + {%- if effort == 'high' -%} + {{- '<|ifm|im_start|>assistant\n\n' }} + {%- elif effort == 'medium' -%} + {{- '<|ifm|im_start|>assistant\n\n' }} + {%- elif effort == 'low' -%} + {{- '<|ifm|im_start|>assistant\n\n' }} + {%- else -%} + {{- raise_exception("Unsupported reasoning_effort: '" + effort + "'. Supported values: high, medium, low.") -}} + {%- endif -%} +{%- endif -%} diff --git a/tests/test-chat.cpp b/tests/test-chat.cpp index f2728f7ccadd..526ffc52c1c5 100644 --- a/tests/test-chat.cpp +++ b/tests/test-chat.cpp @@ -4849,6 +4849,87 @@ static void test_template_output_peg_parsers(bool detailed_debug) { .run(); } + // K2 Horizon dedicated parser + { + auto tmpls = read_templates("models/templates/IFM-K2-Horizon-7B.jinja"); + const auto caps = common_chat_templates_get_caps(tmpls.get()); + GGML_ASSERT(caps.at("supports_parallel_tool_calls")); + GGML_ASSERT(caps.at("supports_object_arguments")); + assert_contains(common_chat_format_example(tmpls.get(), true, {}), "Hi there"); + + auto tst = peg_tester("models/templates/IFM-K2-Horizon-7B.jinja", detailed_debug); + + const std::string get_time_call = + "\n" + "get_time\n" + "city\n" + "Paris\n" + "\n" + ""; + + // The generation prompt pre-opens , so the model output starts inside it. + tst.test("Simple sum.\n\n51") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .expect_reasoning("Simple sum.\n") + .expect_content("51") + .run(); + + // The end-of-turn token must not leak into content. + tst.test("Simple sum.\n\n51<|ifm|im_end|>") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .expect_reasoning("Simple sum.\n") + .expect_content("51") + .run(); + + // A closed think block followed by a tool call section. + tst.test("I need the time.\n\n" + get_time_call) + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .tools({ get_time_tool }) + .expect_reasoning("I need the time.\n") + .expect_tool_calls({ { "get_time", R"({"city": "Paris"})", "" } }) + .run(); + + // A tool call section may start before the think block is closed. + tst.test("I need the time.\n" + get_time_call) + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .tools({ get_time_tool }) + .expect_reasoning("I need the time.\n") + .expect_tool_calls({ { "get_time", R"({"city": "Paris"})", "" } }) + .run(); + + // Non-string arguments parse as JSON, and required arguments may come in any order. + tst.test("\n\ntool_2req_4opt\n" + "req2\n7\n" + "req1\nhello\n" + "\n") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .tools({ tool_2req_4opt }) + .expect_tool_calls({ { "tool_2req_4opt", R"({"req2": 7, "req1": "hello"})", "" } }) + .run(); + + // Parallel tool calls share one section. + tst.test("\n\n" + "get_time\ncity\nParis\n\n" + "get_time\ncity\nRome\n\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .tools({ get_time_tool }) + .parallel_tool_calls(true) + .expect_tool_calls({ + { "get_time", R"({"city": "Paris"})", "" }, + { "get_time", R"({"city": "Rome"})", "" }, + }) + .run(); + + // reasoning_format=none keeps extracting tool calls. + tst.test("I need the time.\n\n" + get_time_call) + .reasoning_format(COMMON_REASONING_FORMAT_NONE) + .tools({ get_time_tool }) + .expect_content("I need the time.\n") + .expect_tool_calls({ { "get_time", R"({"city": "Paris"})", "" } }) + .run(); + } + // Kimi-K3 tests - custom parser // Unique feature: XTML tags built from <|open|>/<|close|>/<|sep|>, and a // generation prompt that leaves the think section already open. From f13798d897c7055150b2ac24e9be5e4274c99d1c Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Sun, 27 Sep 2026 16:14:41 +0000 Subject: [PATCH 12/22] conversion: remove obsolete K2 Aurora alias Assisted-by: Codex --- conversion/__init__.py | 1 - conversion/k2_horizon.py | 5 +---- 2 files changed, 1 insertion(+), 5 deletions(-) diff --git a/conversion/__init__.py b/conversion/__init__.py index 7341ed6e8e5d..1bd1c8f16405 100644 --- a/conversion/__init__.py +++ b/conversion/__init__.py @@ -139,7 +139,6 @@ "JinaBertModel": "bert", "JinaEmbeddingsV5Model": "bert", "K2HorizonForCausalLM": "k2_horizon", - "K2AuroraForCausalLM": "k2_horizon", # TODO: DELETE "KORMoForCausalLM": "qwen", "KimiK25ForConditionalGeneration": "deepseek", "KimiK3ForConditionalGeneration": "kimi_k3", diff --git a/conversion/k2_horizon.py b/conversion/k2_horizon.py index 48f0b521fcb3..811d66c6cb01 100644 --- a/conversion/k2_horizon.py +++ b/conversion/k2_horizon.py @@ -12,10 +12,7 @@ from .base import ModelBase, TextModel, gguf, logger -@ModelBase.register( - "K2HorizonForCausalLM", - "K2AuroraForCausalLM", # TODO: DELETE -) +@ModelBase.register("K2HorizonForCausalLM") @ModelBase.example("IFM/K2-Horizon-0.9B", "IFM/K2-Horizon-36B") class K2HorizonModel(TextModel): model_arch = gguf.MODEL_ARCH.K2HORIZON From 540e938c46f7e5887848fe43cfbefb926f2a3cb2 Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Mon, 28 Sep 2026 14:46:45 +0000 Subject: [PATCH 13/22] k2-horizon: enforce response schemas and load YaRN betas Constrain final JSON after reasoning, accept flexible JSON tool envelopes, enforce XML dialects, and handle repeated or alternate thinking markers. Load YaRN beta metadata instead of retaining the default values. Add schema, streaming, continuation, and model reload regressions. Validate CUDA and CPU builds and 0.9B, 4B, and MoVA conversation/tool round trips. Assisted-by: Codex --- common/parsers/k2-horizon.cpp | 92 ++++++++++++++++------ src/models/k2-horizon.cpp | 2 + tests/test-chat.cpp | 141 ++++++++++++++++++++++++++++++++++ tests/test-llama-archs.cpp | 10 +++ 4 files changed, 223 insertions(+), 22 deletions(-) diff --git a/common/parsers/k2-horizon.cpp b/common/parsers/k2-horizon.cpp index 4e0adf6cdd85..eb4930c34c6d 100644 --- a/common/parsers/k2-horizon.cpp +++ b/common/parsers/k2-horizon.cpp @@ -3,11 +3,11 @@ // K2 Horizon - reasoning effort picks one of three think tag pairs, tool calls are tagged: // assistant := ... [content] // [ {CALL} ] -// CALL (tool_call_format=xml, default) := name {k [t] -// v} +// CALL (tool_call_format=xml, default) := name {k v} +// CALL (tool_call_format=xml_typed) adds a required t before each value. // CALL (tool_call_format=json) := {"name": name, "arguments": {...}} -// The generation prompt pre-opens the think block, so the model never emits the -// opening tag. Reasoning ends at the close tag or at a tool call section start. +// The generation prompt pre-opens the think block; repeated opening tags are accepted. +// Reasoning ends at any think close tag or at a tool call section start. common_chat_params common_chat_params_init_k2_horizon(const common_chat_template & tmpl, const autoparser::generation_params & inputs) { common_chat_params data; @@ -52,12 +52,19 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template const std::string THINK_END = ""; data.preserved_tokens = { - THINK_START, THINK_END, SECTION_START, SECTION_END, CALL_START, CALL_END, + "", "", "", "", + "", "", SECTION_START, SECTION_END, CALL_START, CALL_END, ARG_KEY, ARG_KEY_END, ARG_TYPE, ARG_TYPE_END, ARG_VAL, ARG_VAL_END, TURN_END, }; data.thinking_start_tag = THINK_START; - data.thinking_end_tags = { THINK_END, SECTION_START }; + data.thinking_end_tags = { THINK_END }; + for (const std::string tag : { "", "", "" }) { + if (tag != THINK_END) { + data.thinking_end_tags.push_back(tag); + } + } + data.thinking_end_tags.push_back(SECTION_START); data.message_delimiters = { { COMMON_CHAT_ROLE_ASSISTANT, "<|ifm|im_start|>assistant" }, @@ -92,9 +99,10 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template call_format = inputs.extra_context.at("tool_call_format"); } - auto has_tools = inputs.tools.is_array() && !inputs.tools.empty(); - auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE; - auto include_grammar = has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE; + auto has_tools = inputs.tools.is_array() && !inputs.tools.empty(); + auto has_response_format = inputs.json_schema.is_object() && !inputs.json_schema.empty(); + auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE; + auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE); auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) { auto end = p.end(); @@ -102,11 +110,30 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template // the effective parse input is generation_prompt + model output auto opener = p.optional(p.literal(ROLE) + p.optional(p.space())); - auto body_end = think_open ? p.until_one_of({ THINK_END, SECTION_START }) : p.until(THINK_END); + auto think_ends = data.thinking_end_tags; + think_ends.pop_back(); + auto think_close = p.choice(); + for (const auto & tag : think_ends) { + think_close |= p.literal(tag); + } + auto body_ends = think_open ? data.thinking_end_tags : think_ends; + body_ends.push_back(TURN_END); + auto body_end = p.until_one_of(body_ends); auto think_body = extract_reasoning ? p.reasoning(body_end) : p.content(body_end); // the template writes "\n" and "\n"; those newlines are markup, not text - auto nl = p.optional(p.literal("\n")); - auto reasoning = p.optional(p.optional(p.literal(THINK_START) + nl) + think_body + p.optional(p.literal(THINK_END) + nl)); + auto nl = p.optional(p.literal("\n")); + auto think_start = p.one_or_more(p.literal(THINK_START) + nl); + auto reasoning = p.optional(p.optional(think_start) + think_body + p.optional(think_close + nl)); + + if (has_response_format) { + // The final answer must follow a closed reasoning block, including when the prompt pre-opens it. + // Do not inline reasoning into schema-constrained content when extraction is disabled. + auto schema_reasoning = extract_reasoning ? p.reasoning(body_end) : body_end; + auto closed_reasoning = p.optional(think_start + schema_reasoning + think_close + nl); + return opener + closed_reasoning + p.space() + + p.content(p.schema(p.json(), "k2h-response", inputs.json_schema)) + + p.space() + p.optional(p.literal(TURN_END)) + end; + } auto content = p.optional(p.content(p.until_one_of({ SECTION_START, TURN_END }))); auto tail = p.optional(p.content(p.until(TURN_END))) + p.optional(p.literal(TURN_END)); @@ -117,8 +144,7 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template auto tool_choices = p.choice(); auto arg_close = p.tool_arg_close(p.literal(ARG_VAL_END)); - auto arg_string = p.rule("k2h-arg-string", p.tool_arg_string_value(p.until(ARG_VAL_END)) + arg_close); - auto arg_type = p.optional(p.optional(p.space()) + p.literal(ARG_TYPE) + p.until(ARG_TYPE_END) + p.literal(ARG_TYPE_END)); + auto arg_string = p.rule("k2h-arg-string", p.tool_arg_string_value(p.until_one_of({ ARG_VAL_END, TURN_END })) + arg_close); foreach_function(inputs.tools, [&](const json & tool) { const auto & function = tool.at("function"); @@ -126,10 +152,14 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template if (call_format == "json") { auto schema = common_chat_tool_parameters(function); - auto call = p.tool(p.tool_open(p.literal(CALL_START) + p.literal("{\"name\": \"") + p.tool_name(p.literal(name)) + - p.literal("\", \"arguments\": ")) + - p.tool_args(p.schema(p.json(), "k2h-tool-" + name + "-schema", schema)) + - p.tool_close(p.literal("}") + p.literal(CALL_END))); + auto name_field = p.atomic(p.literal("\"name\"") + p.space() + p.literal(":") + p.space() + + p.literal("\"") + p.tool_name(p.literal(name)) + p.literal("\"")) + p.space(); + auto args_field = p.literal("\"arguments\"") + p.space() + p.literal(":") + p.space() + + p.tool_args(p.schema(p.json(), "k2h-tool-" + name + "-schema", schema)) + p.space(); + auto comma = p.literal(",") + p.space(); + auto call = p.tool(p.tool_open(p.literal(CALL_START) + p.space() + p.literal("{") + p.space()) + + ((name_field + comma + args_field) | (args_field + comma + name_field)) + + p.tool_close(p.literal("}") + p.space() + p.literal(CALL_END))); tool_choices |= p.rule("k2h-tool-" + name, call); return; } @@ -142,6 +172,22 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template auto types = param.schema->value_types(); auto json_val = p.tool_arg_json_value(p.schema(p.json(), rule_name + "-schema", doc, *param.schema)) + arg_close; + auto arg_type = p.eps(); + if (call_format == "xml_typed") { + std::string type_grammar; + for (auto type : { common_chat_schema::TYPE_NULL, common_chat_schema::TYPE_BOOLEAN, + common_chat_schema::TYPE_NUMBER, common_chat_schema::TYPE_INTEGER, + common_chat_schema::TYPE_STRING, common_chat_schema::TYPE_ARRAY, common_chat_schema::TYPE_OBJECT }) { + if (types.has(type)) { + type_grammar += (type_grammar.empty() ? "" : " | ") + gbnf_format_literal(common_chat_schema::type_name(type)); + } + } + // Parse compound labels too, but generate a schema type, never an argument value or markup. + auto type_text = p.chars("[^ \\t\\r\\n<]", 1, 1) + p.chars("[^<]", 0); + arg_type = p.space() + p.literal(ARG_TYPE) + p.space() + + p.gbnf(type_text, "(" + type_grammar + ")") + p.space() + p.literal(ARG_TYPE_END); + } + auto arg_value = types.is_only(common_chat_schema::TYPE_STRING) ? arg_string : !types.has(common_chat_schema::TYPE_STRING) ? json_val : p.gbnf(p.atomic(json_val) | arg_string, "k2h-arg-string"); @@ -181,14 +227,16 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template data.parser = parser.save(); if (include_grammar) { - data.grammar_lazy = inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED; + data.grammar_lazy = !has_response_format && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED; data.grammar = build_grammar([&](const common_grammar_builder & builder) { parser.build_grammar(builder, data.grammar_lazy); }); - data.grammar_triggers = { - { COMMON_GRAMMAR_TRIGGER_TYPE_WORD, SECTION_START }, - }; + if (data.grammar_lazy) { + data.grammar_triggers = { + { COMMON_GRAMMAR_TRIGGER_TYPE_WORD, SECTION_START }, + }; + } } return data; diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp index 3c5991156788..95de2877352a 100644 --- a/src/models/k2-horizon.cpp +++ b/src/models/k2-horizon.cpp @@ -6,6 +6,8 @@ #include "models.h" void llama_model_k2_horizon::load_arch_hparams(llama_model_loader & ml) { + ml.get_key(LLM_KV_ROPE_SCALING_YARN_BETA_FAST, hparams.yarn_beta_fast, false); + ml.get_key(LLM_KV_ROPE_SCALING_YARN_BETA_SLOW, hparams.yarn_beta_slow, false); ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); ml.get_key(LLM_KV_ATTENTION_GROUPNORM_GROUPS, hparams.n_norm_groups, false); if (hparams.n_norm_groups == 0) { diff --git a/tests/test-chat.cpp b/tests/test-chat.cpp index 526ffc52c1c5..f233f541994b 100644 --- a/tests/test-chat.cpp +++ b/tests/test-chat.cpp @@ -1527,6 +1527,11 @@ class peg_test_builder { return *this; } + peg_test_builder & chat_template_kwargs(const std::map & kwargs) { + tc_.params.chat_template_kwargs = kwargs; + return *this; + } + peg_test_builder & is_partial(bool val) { tc_.is_partial = val; return *this; @@ -4859,6 +4864,57 @@ static void test_template_output_peg_parsers(bool detailed_debug) { auto tst = peg_tester("models/templates/IFM-K2-Horizon-7B.jinja", detailed_debug); + const std::string answer_schema = R"({"type":"object","properties":{"answer":{"type":"integer","const":42}},"required":["answer"],"additionalProperties":false})"; + tst.test("Let me calculate.\n{\"answer\":42}<|ifm|im_end|>") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .json_schema(answer_schema) + .expect_reasoning("Let me calculate.") + .expect_content(R"({"answer":42})") + .run(); + + tst.test("Let me calculate.{\"answer\":42}") + .reasoning_format(COMMON_REASONING_FORMAT_NONE) + .json_schema(answer_schema) + .expect_content(R"({"answer":42})") + .run(); + + // Prefill advances the grammar through both reasoning and partial final content. + tst.test("42}") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .json_schema(answer_schema) + .messages({ message_user, simple_assist_msg("{\"answer\":", "Calculated.") }) + .continue_final_message(COMMON_CHAT_CONTINUATION_CONTENT) + .expect_reasoning("Calculated.") + .expect_content(R"({"answer":42})") + .run(); + + common_chat_templates_inputs schema_inputs; + schema_inputs.messages = { message_user }; + schema_inputs.add_generation_prompt = true; + schema_inputs.json_schema = answer_schema; + schema_inputs.tools = { get_time_tool }; + auto schema_params = common_chat_templates_apply(tmpls.get(), schema_inputs); + GGML_ASSERT(!schema_params.grammar.empty()); + GGML_ASSERT(!schema_params.grammar_lazy); + GGML_ASSERT(schema_params.grammar_triggers.empty()); + for (const std::string output : { + "{\"answer\":42}", + "```json\n{\"answer\":42}\n```", + "{\"answer\":\"42\"}", + "{\"answer\":41}", + "{\"wrong\":42}", + "{\"answer\":42,\"extra\":1}", + "{\"answer\":42} trailing text", + "Still thinking", "" }) { + auto grammar = build_grammar(schema_params.grammar); + GGML_ASSERT(match_string(schema_params.generation_prompt + output, grammar.get()) == + (output == "{\"answer\":42}")); + } + // A stop marker inside reasoning must be rejected, not accepted as an incomplete answer. + auto stop_grammar = build_grammar(schema_params.grammar); + auto stop_match = match_string_detailed(schema_params.generation_prompt + "<|ifm|im_end|>", stop_grammar.get()); + GGML_ASSERT(!stop_match.success && !stop_match.incomplete); + const std::string get_time_call = "\n" "get_time\n" @@ -4867,6 +4923,91 @@ static void test_template_output_peg_parsers(bool detailed_debug) { "\n" ""; + // JSON envelopes allow whitespace and either field order, including during streaming. + for (const std::string payload : { + R"({"name":"get_time","arguments":{"city":"Paris"}})", + R"({ "arguments" : {"city":"Paris"}, "name" : "get_time" })", + "{\n\t\"name\" : \"get_time\",\n\"arguments\" : {\"city\":\"Paris\"}\n}" }) { + tst.test("" + payload + "") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .tools({ get_time_tool }) + .chat_template_kwargs({ { "tool_call_format", R"("json")" } }) + .expect_tool_calls({ { "get_time", R"({"city":"Paris"})", "" } }) + .run(); + } + + // Do not emit the shorter name while a longer name is still being streamed. + auto longer_name_tool = get_time_tool; + longer_name_tool.name += "_extended"; + tst.test("" + "{\"name\":\"get_time_extended\",\"arguments\":{\"city\":\"Paris\"}}" + "") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .tools({ get_time_tool, longer_name_tool }) + .chat_template_kwargs({ { "tool_call_format", R"("json")" } }) + .expect_tool_calls({ { "get_time_extended", R"({"city":"Paris"})", "" } }) + .run(); + + const std::string typed_call = + "get_time" + "citystring" + "Paris"; + tst.test("" + typed_call) + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .tools({ get_time_tool }) + .chat_template_kwargs({ { "tool_call_format", R"("xml_typed")" } }) + .expect_tool_calls({ { "get_time", R"({"city":"Paris"})", "" } }) + .run(); + + // Wrong XML dialects are not complete calls and cannot be generated by the grammar. + for (const std::string format : { "xml", "xml_typed" }) { + common_chat_templates_inputs inputs; + inputs.messages = { message_user }; + inputs.tools = { get_time_tool }; + inputs.reasoning_format = COMMON_REASONING_FORMAT_DEEPSEEK; + inputs.chat_template_kwargs["tool_call_format"] = json(format).dump(); + auto parser = make_peg_parser(tmpls.get(), inputs); + const auto & invalid = format == "xml" ? typed_call : get_time_call; + GGML_ASSERT(parser.parse("" + invalid, false).tool_calls.empty()); + auto grammar = build_grammar(parser.params_.grammar); + GGML_ASSERT(!match_string(invalid, grammar.get())); + + const std::string arg_prefix = "get_timecity"; + const std::string unfinished = arg_prefix + (format == "xml" ? "Paris" : "string"); + auto stop_grammar = build_grammar(parser.params_.grammar); + auto stop_match = match_string_detailed(unfinished + "<|ifm|im_end|>", stop_grammar.get()); + GGML_ASSERT(!stop_match.success && !stop_match.incomplete); + if (format == "xml_typed") { + auto type_grammar = build_grammar(parser.params_.grammar); + auto type_match = match_string_detailed(arg_prefix + "17", type_grammar.get()); + GGML_ASSERT(!type_match.success && !type_match.incomplete); + } + } + + for (const std::string effort : { "high", "medium", "low" }) { + const auto tag = effort == "high" ? "ifm|think" : effort == "medium" ? "ifm|think_fast" : "ifm|think_faster"; + for (const std::string close : { "", "", "" }) { + tst.test("<" + std::string(tag) + ">Plan." + close + "42") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .chat_template_kwargs({ { "reasoning_effort", json(effort).dump() } }) + .expect_reasoning("Plan.") + .expect_content("42") + .run(); + tst.test("Plan." + close + "{\"answer\":42}") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .chat_template_kwargs({ { "reasoning_effort", json(effort).dump() } }) + .json_schema(answer_schema) + .expect_reasoning("Plan.") + .expect_content(R"({"answer":42})") + .run(); + } + } + + tst.test("Still thinking") + .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + .expect_reasoning("Still thinking") + .run(); + // The generation prompt pre-opens , so the model output starts inside it. tst.test("Simple sum.\n\n51") .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp index 5be036ad0cb6..bfb7ddef4103 100644 --- a/tests/test-llama-archs.cpp +++ b/tests/test-llama-archs.cpp @@ -9,6 +9,7 @@ // TODO: replace with #include "llama-ext.h" in the future #include "../src/llama-arch.h" +#include "../src/llama-model.h" #include "../src/llama-model-saver.h" #include @@ -184,6 +185,11 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) { ms.add_kv(LLM_KV_BLOCK_COUNT, n_layer); ms.add_kv(LLM_KV_LEADING_DENSE_BLOCK_COUNT, uint32_t(1)); + if (arch == LLM_ARCH_K2_HORIZON) { + ms.add_kv(LLM_KV_ROPE_SCALING_YARN_BETA_FAST, 128.0f); + ms.add_kv(LLM_KV_ROPE_SCALING_YARN_BETA_SLOW, 4.0f); + } + if (arch == LLM_ARCH_NEMOTRON_H || arch == LLM_ARCH_NEMOTRON_H_MOE) { std::vector n_ff_per_layer; n_ff_per_layer.reserve(n_layer); @@ -490,6 +496,10 @@ static std::pair get_model_and_ctx( if (!model) { throw std::runtime_error("failed to create llama model"); } + if (model->arch == LLM_ARCH_K2_HORIZON) { + GGML_ASSERT(model->hparams.yarn_beta_fast == 128.0f); + GGML_ASSERT(model->hparams.yarn_beta_slow == 4.0f); + } llama_context_ptr lctx(llama_init_from_model(model.get(), ctx_params)); if (!lctx) { throw std::runtime_error("failed to create llama context"); From cbee4a287407a515ddf6a5b1bf14b536982070b9 Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Wed, 30 Sep 2026 11:36:19 +0000 Subject: [PATCH 14/22] renaming template fixture --- .../{IFM-K2-Horizon-7B.jinja => IFM-K2-Horizon.jinja} | 0 tests/test-chat.cpp | 4 ++-- 2 files changed, 2 insertions(+), 2 deletions(-) rename models/templates/{IFM-K2-Horizon-7B.jinja => IFM-K2-Horizon.jinja} (100%) diff --git a/models/templates/IFM-K2-Horizon-7B.jinja b/models/templates/IFM-K2-Horizon.jinja similarity index 100% rename from models/templates/IFM-K2-Horizon-7B.jinja rename to models/templates/IFM-K2-Horizon.jinja diff --git a/tests/test-chat.cpp b/tests/test-chat.cpp index f233f541994b..9c333ca9eb19 100644 --- a/tests/test-chat.cpp +++ b/tests/test-chat.cpp @@ -4856,13 +4856,13 @@ static void test_template_output_peg_parsers(bool detailed_debug) { // K2 Horizon dedicated parser { - auto tmpls = read_templates("models/templates/IFM-K2-Horizon-7B.jinja"); + auto tmpls = read_templates("models/templates/IFM-K2-Horizon.jinja"); const auto caps = common_chat_templates_get_caps(tmpls.get()); GGML_ASSERT(caps.at("supports_parallel_tool_calls")); GGML_ASSERT(caps.at("supports_object_arguments")); assert_contains(common_chat_format_example(tmpls.get(), true, {}), "Hi there"); - auto tst = peg_tester("models/templates/IFM-K2-Horizon-7B.jinja", detailed_debug); + auto tst = peg_tester("models/templates/IFM-K2-Horizon.jinja", detailed_debug); const std::string answer_schema = R"({"type":"object","properties":{"answer":{"type":"integer","const":42}},"required":["answer"],"additionalProperties":false})"; tst.test("Let me calculate.\n{\"answer\":42}<|ifm|im_end|>") From 39ddb2b1fd35f2d15565bb204e51dbe480139996 Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Thu, 1 Oct 2026 20:49:43 +0000 Subject: [PATCH 15/22] adressing cisc follows ups --- conversion/base.py | 6 +- src/llama-vocab.cpp | 2 +- src/models/k2-horizon.cpp | 22 ++-- src/unicode.cpp | 1 - tests/test-llama-archs.cpp | 5 - tests/test-unicode.cpp | 224 +++---------------------------------- 6 files changed, 28 insertions(+), 232 deletions(-) diff --git a/conversion/base.py b/conversion/base.py index f3346809fbaf..e7337a7bcc13 100644 --- a/conversion/base.py +++ b/conversion/base.py @@ -1714,9 +1714,6 @@ def get_vocab_base_pre(self, tokenizer) -> str: if chkhsh == "1f9825a388f700a6b591722f17d470cbbcf10973ece35d2fd14239a14110ae1a": # ref: https://huggingface.co/IFM/K2-Horizon-0.9B res = "k2-horizon" - if chkhsh == "a9af07a84191f55098b248ae6f3dfe9e32d3190bebe8eafd91c1ddec9bc3449f": - # ref: https://huggingface.co/IFM/K2-Horizon-36B - res = "k2-horizon" if chkhsh == "0ef9807a4087ebef797fc749390439009c3b9eda9ad1a097abbe738f486c01e5": # ref: https://huggingface.co/meta-llama/Meta-Llama-3-8B res = "llama-bpe" @@ -1942,6 +1939,9 @@ def get_vocab_base_pre(self, tokenizer) -> str: if chkhsh == "653660222fb704f61cbf2b618a8ae6502b7f8b20c980f9a5de07ed78e13319cd": # ref: https://huggingface.co/ufakai/ufakzeka-1 res = "ufakzeka" + if chkhsh == "a9af07a84191f55098b248ae6f3dfe9e32d3190bebe8eafd91c1ddec9bc3449f": + # ref: https://huggingface.co/IFM/K2-Horizon-36B + res = "k2-horizon" if res is None: logger.warning("\n") diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index d77b3c253aba..c9144644fc29 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -553,7 +553,7 @@ struct llm_tokenizer_bpe : llm_tokenizer { break; case LLAMA_VOCAB_PRE_TYPE_K2_HORIZON: regex_exprs = { - "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", + "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", }; break; case LLAMA_VOCAB_PRE_TYPE_WHITESPACE: diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp index 95de2877352a..5c158ce77807 100644 --- a/src/models/k2-horizon.cpp +++ b/src/models/k2-horizon.cpp @@ -1,8 +1,3 @@ -// K2 Horizon (MBZUAI IFM): grouped RMSNorm, optional per-head QK-norm and softplus -// attention output gate, DeepSeek-V3 style MoE (sigmoid router, selection bias, -// shared expert, leading dense layers) and MoVA: in MoE layers the V projection is -// replaced by routed value experts, V = sum_k w_k * silu(W_k x). - #include "models.h" void llama_model_k2_horizon::load_arch_hparams(llama_model_loader & ml) { @@ -39,12 +34,17 @@ void llama_model_k2_horizon::load_arch_hparams(llama_model_loader & ml) { GGML_ASSERT(hparams.n_value_expert_used == 0); } - if (hparams.n_layer() == 28 && hparams.n_embd == 1536) { - type = LLM_TYPE_1B; - } else if (hparams.n_layer() == 48 && hparams.n_embd == 2560) { - type = LLM_TYPE_36B; - } else { - type = LLM_TYPE_UNKNOWN; + switch (hparams.n_layer()) { + case 28: type = LLM_TYPE_1B; break; + case 36: + switch (hparams.n_embd) { + case 2560: type = LLM_TYPE_4B; break; + case 4096: type = LLM_TYPE_7B; break; + default: type = LLM_TYPE_UNKNOWN; + } break; + case 48: type = LLM_TYPE_36B; break; + case 64: type = LLM_TYPE_32B; break; + default: type = LLM_TYPE_UNKNOWN; } } diff --git a/src/unicode.cpp b/src/unicode.cpp index 10e8b408fbdd..07b425f27f8d 100644 --- a/src/unicode.cpp +++ b/src/unicode.cpp @@ -1210,7 +1210,6 @@ static std::vector unicode_regex_split_custom(const std::string & text, regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?[\\p{L}\\p{M}]+|\\p{N}| ?[^\\s\\p{L}\\p{M}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { bpe_offsets = unicode_regex_split_custom_qwen35(text, offsets); } else if ( - regex_expr == "(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+" || regex_expr == "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+") { // K2-Horizon: llama3 splitter with marks + ZWNJ/ZWJ inside letter runs // (the generic std::regex fallback cannot parse \p{..} on MSVC) diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp index 49cf834594f5..c8f1aa2037ed 100644 --- a/tests/test-llama-archs.cpp +++ b/tests/test-llama-archs.cpp @@ -9,7 +9,6 @@ // TODO: replace with #include "llama-ext.h" in the future #include "../src/llama-arch.h" -#include "../src/llama-model.h" #include "../src/llama-model-saver.h" #include @@ -508,10 +507,6 @@ static std::pair get_model_and_ctx( if (!model) { throw std::runtime_error("failed to create llama model"); } - if (model->arch == LLM_ARCH_K2_HORIZON) { - GGML_ASSERT(model->hparams.yarn_beta_fast == 128.0f); - GGML_ASSERT(model->hparams.yarn_beta_slow == 4.0f); - } llama_context_ptr lctx(llama_init_from_model(model.get(), ctx_params)); if (!lctx) { throw std::runtime_error("failed to create llama context"); diff --git a/tests/test-unicode.cpp b/tests/test-unicode.cpp index 72a0f6ed0e09..2347d9000a8e 100644 --- a/tests/test-unicode.cpp +++ b/tests/test-unicode.cpp @@ -1,221 +1,23 @@ #include "../src/unicode.h" #include -#include #include #include int main() { - { - const std::vector regex_exprs = { - "[~][A-Za-z]+| ?[\\p{S}]+|\\s+", - }; - const std::vector expected = { " ~", "foo" }; - const auto actual = unicode_regex_split(" ~foo", regex_exprs, false); - - if (actual != expected) { - fprintf(stderr, "unexpected split:"); - for (const auto & piece : actual) { - fprintf(stderr, " [%s]", piece.c_str()); - } - fprintf(stderr, "\n"); - return 1; - } - } - - { - const std::vector regex_exprs = { - "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" - "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" - "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" - "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", - }; - - const std::string input = "ab\u200ccd ef\u200dgh cafe\u0301"; - const std::vector expected = { - "ab\u200ccd", - " ef\u200dgh", - " cafe\u0301", - }; - - try { - const auto actual = unicode_regex_split(input, regex_exprs, false); - - if (actual != expected) { - fprintf(stderr, "unexpected K2-Horizon split:"); - for (const auto & piece : actual) { - fprintf(stderr, " [%s]", piece.c_str()); - } - fprintf(stderr, "\n"); - return 1; - } - } catch (const std::exception & e) { - fprintf(stderr, "K2-Horizon regex split threw exception: %s\n", e.what()); - return 1; - } - } - - - { - const std::vector llama3_regex = { - "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" - "[^\\r\\n\\p{L}\\p{N}]?\\p{L}+|" - "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" - "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", - }; - - const std::vector k2_horizon_regex = { - "(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])|" - "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" - "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" - "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", - }; - - const std::vector inputs = { - "Hello, world!", - "can't won't we're they'll", - "123 1234 1234567", - "Hello!!!\\nNext line", - "alpha beta gamma", - " café résumé", - }; - - for (const auto & input : inputs) { - const auto llama3 = unicode_regex_split(input, llama3_regex, false); - const auto k2 = unicode_regex_split(input, k2_horizon_regex, false); - - if (llama3 != k2) { - fprintf(stderr, "K2-Horizon diverged from Llama 3 for ordinary input: %s\n", input.c_str()); - - fprintf(stderr, "Llama 3:"); - for (const auto & piece : llama3) { - fprintf(stderr, " [%s]", piece.c_str()); - } - - fprintf(stderr, "\nK2-Horizon:"); - for (const auto & piece : k2) { - fprintf(stderr, " [%s]", piece.c_str()); - } - - fprintf(stderr, "\n"); - return 1; - } - } - } - - - { - const std::vector k2_horizon_regex = { - "(?i:'s|'t|'re|'ve|'m|'ll|'d)|" - "[^\\r\\n\\p{L}\\p{N}]?(?:\\p{L}|\\p{M}|\\u200C|\\u200D)+|" - "\\p{N}{1,3}| ?[^\\s\\p{L}\\p{N}]+[\\r\\n]*|" - "\\s*[\\r\\n]+|\\s+(?!\\S)|\\s+", - }; - - auto check_k2 = [&] ( - const char * name, - const std::string & input, - const std::vector & expected) { - const auto actual = unicode_regex_split(input, k2_horizon_regex, false); - - if (actual != expected) { - fprintf(stderr, "K2-Horizon reference mismatch for %s\n", name); - - fprintf(stderr, "expected:"); - for (const auto & piece : expected) { - fprintf(stderr, " [%s]", piece.c_str()); - } - - fprintf(stderr, "\nactual:"); - for (const auto & piece : actual) { - fprintf(stderr, " [%s]", piece.c_str()); - } - - fprintf(stderr, "\n"); - return false; - } - - return true; - }; - - if (!check_k2( - "ZWNJ", - "ab\u200ccd", - { "ab\u200ccd" })) { - return 1; - } - - if (!check_k2( - "ZWJ", - "ef\u200dgh", - { "ef\u200dgh" })) { - return 1; - } - - if (!check_k2( - "combined ZWNJ/ZWJ", - "ab\u200ccd ef\u200dgh", - { "ab\u200ccd", " ef\u200dgh" })) { - return 1; - } - - // The reference tokenizer NFC-normalizes cafe + combining acute - // to the composed form before pre-tokenization. - if (!check_k2( - "NFC accent", - "caf\u00e9", - { "caf\u00e9" })) { - return 1; - } - - if (!check_k2( - "Persian ZWNJ", - "\u0645\u06cc\u200c\u0631\u0648\u0645", - { "\u0645\u06cc\u200c\u0631\u0648\u0645" })) { - return 1; - } - - if (!check_k2( - "Devanagari ZWJ", - "\u0915\u094d\u200d\u0937", - { "\u0915\u094d\u200d\u0937" })) { - return 1; - } - - if (!check_k2( - "contractions", - "can't won't we're they'll", - { "can", "'t", " won", "'t", " we", "'re", " they", "'ll" })) { - return 1; - } - - if (!check_k2( - "Unicode contraction", - "'\u017fa", - { "'\u017f", "a" })) { - return 1; - } - - if (!check_k2( - "empty input", - "", - { })) { - return 1; - } - - if (!check_k2( - "numbers", - "123 1234 1234567", - { "123", " ", "123", "4", " ", "123", "456", "7" })) { - return 1; - } - - if (!check_k2( - "punctuation/newline", - "Hello!!!\nNext line", - { "Hello", "!!!\n", "Next", " line" })) { - return 1; - } + const std::vector regex_exprs = { + "[~][A-Za-z]+| ?[\\p{S}]+|\\s+", + }; + const std::vector expected = { " ~", "foo" }; + const auto actual = unicode_regex_split(" ~foo", regex_exprs, false); + + if (actual != expected) { + fprintf(stderr, "unexpected split:"); + for (const auto & piece : actual) { + fprintf(stderr, " [%s]", piece.c_str()); + } + fprintf(stderr, "\n"); + return 1; } return 0; From 336cc79a684e1612bff68cc0d7683f67ed9a56db Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:15:21 +0000 Subject: [PATCH 16/22] desloppify the parser / adress aldehir comments --- common/parsers/k2-horizon.cpp | 262 ++++++++++------------- src/llama-vocab.cpp | 1 + tests/test-chat.cpp | 376 ++++++++++++++++++---------------- 3 files changed, 306 insertions(+), 333 deletions(-) diff --git a/common/parsers/k2-horizon.cpp b/common/parsers/k2-horizon.cpp index eb4930c34c6d..0348e5e72e9d 100644 --- a/common/parsers/k2-horizon.cpp +++ b/common/parsers/k2-horizon.cpp @@ -1,17 +1,15 @@ #include "parsers.h" -// K2 Horizon - reasoning effort picks one of three think tag pairs, tool calls are tagged: -// assistant := ... [content] -// [ {CALL} ] -// CALL (tool_call_format=xml, default) := name {k v} -// CALL (tool_call_format=xml_typed) adds a required t before each value. -// CALL (tool_call_format=json) := {"name": name, "arguments": {...}} -// The generation prompt pre-opens the think block; repeated opening tags are accepted. -// Reasoning ends at any think close tag or at a tool call section start. +// K2 Horizon format: +// - Reasoning: ..., or / for medium/low reasoning_effort +// - Tool calls: ......, one call per : +// xml (default): name k [t] v ... +// json: {"name": "...", "arguments": {...}} common_chat_params common_chat_params_init_k2_horizon(const common_chat_template & tmpl, const autoparser::generation_params & inputs) { common_chat_params data; + // The template requires a thinking field on every assistant message auto messages = inputs.messages; for (auto & msg : messages) { if (msg.value("role", "") == "assistant" && !msg.contains("reasoning_content")) { @@ -24,8 +22,18 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template data.format = COMMON_CHAT_FORMAT_PEG_NATIVE; data.supports_thinking = true; - const std::string ROLE = "<|ifm|im_start|>assistant"; - const std::string TURN_END = "<|ifm|im_end|>"; + const std::string effort = inputs.extra_context.value("reasoning_effort", "high"); + const std::string call_format = inputs.extra_context.value("tool_call_format", "xml"); + + // Templates that handle enable_thinking disable it with an empty block for every effort + const bool thinking_off = !inputs.enable_thinking && tmpl.source().find("enable_thinking") != std::string::npos; + const std::string think = thinking_off ? "ifm|think" : + effort == "medium" ? "ifm|think_fast" : + effort == "low" ? "ifm|think_faster" : "ifm|think"; + + const std::string GEN_PREFIX = "<|ifm|im_start|>assistant\n"; + const std::string THINK_START = "<" + think + ">"; + const std::string THINK_END = ""; const std::string SECTION_START = ""; const std::string SECTION_END = ""; const std::string CALL_START = ""; @@ -37,34 +45,18 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template const std::string ARG_VAL = ""; const std::string ARG_VAL_END = ""; - // reasoning_effort high/medium/low opens //; - // the pair in use is the last one the generation prompt opened - std::string think = "ifm|think"; - size_t think_pos = std::string::npos; - for (const std::string tag : { "ifm|think", "ifm|think_fast", "ifm|think_faster" }) { - auto pos = data.generation_prompt.rfind("<" + tag + ">"); - if (pos != std::string::npos && (think_pos == std::string::npos || pos > think_pos)) { - think = tag; - think_pos = pos; - } - } - const std::string THINK_START = "<" + think + ">"; - const std::string THINK_END = ""; - - data.preserved_tokens = { - "", "", "", "", - "", "", SECTION_START, SECTION_END, CALL_START, CALL_END, - ARG_KEY, ARG_KEY_END, ARG_TYPE, ARG_TYPE_END, ARG_VAL, ARG_VAL_END, TURN_END, - }; - data.thinking_start_tag = THINK_START; data.thinking_end_tags = { THINK_END }; - for (const std::string tag : { "", "", "" }) { - if (tag != THINK_END) { - data.thinking_end_tags.push_back(tag); - } + if (think != "ifm|think") { + // The 3.7B ends medium and low effort reasoning with + data.thinking_end_tags.push_back(""); } - data.thinking_end_tags.push_back(SECTION_START); + + data.preserved_tokens = data.thinking_end_tags; + data.preserved_tokens.insert(data.preserved_tokens.end(), { + THINK_START, SECTION_START, SECTION_END, CALL_START, CALL_END, + ARG_KEY, ARG_KEY_END, ARG_TYPE, ARG_TYPE_END, ARG_VAL, ARG_VAL_END, + }); data.message_delimiters = { { COMMON_CHAT_ROLE_ASSISTANT, "<|ifm|im_start|>assistant" }, @@ -73,13 +65,15 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template { COMMON_CHAT_ROLE_SYSTEM, "<|ifm|im_start|>system" }, }; - // the turn ends with <|ifm|im_end|>, but only <|endoftext|> is EOG in the vocab - data.additional_stops = { TURN_END }; + auto has_tools = inputs.tools.is_array() && !inputs.tools.empty(); + auto has_response_format = inputs.json_schema.is_object() && !inputs.json_schema.empty(); + auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE; + auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE); if (inputs.has_continuation()) { const auto & msg = inputs.continue_msg; - data.generation_prompt = ROLE + "\n" + THINK_START + "\n" + msg.reasoning_content; + data.generation_prompt = GEN_PREFIX + THINK_START + "\n" + msg.reasoning_content; if (inputs.continue_final_message == COMMON_CHAT_CONTINUATION_CONTENT) { data.generation_prompt += THINK_END + msg.render_content(); } @@ -87,147 +81,107 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template data.prompt += data.generation_prompt; } - bool think_open = false; - if (inputs.has_continuation()) { - think_open = inputs.continue_final_message != COMMON_CHAT_CONTINUATION_CONTENT; - } else { - think_open = think_pos != std::string::npos && data.generation_prompt.find(THINK_END, think_pos) == std::string::npos; - } - - std::string call_format = "xml"; - if (inputs.extra_context.contains("tool_call_format") && inputs.extra_context.at("tool_call_format").is_string()) { - call_format = inputs.extra_context.at("tool_call_format"); - } - - auto has_tools = inputs.tools.is_array() && !inputs.tools.empty(); - auto has_response_format = inputs.json_schema.is_object() && !inputs.json_schema.empty(); - auto extract_reasoning = inputs.reasoning_format != COMMON_REASONING_FORMAT_NONE; - auto include_grammar = has_response_format || (has_tools && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_NONE); - auto parser = build_chat_peg_parser([&](common_chat_peg_builder & p) { - auto end = p.end(); - - // the effective parse input is generation_prompt + model output - auto opener = p.optional(p.literal(ROLE) + p.optional(p.space())); + auto generation_prompt = p.literal(GEN_PREFIX); - auto think_ends = data.thinking_end_tags; - think_ends.pop_back(); - auto think_close = p.choice(); - for (const auto & tag : think_ends) { - think_close |= p.literal(tag); + auto think_end = p.choice(); + for (const auto & tag : data.thinking_end_tags) { + think_end |= p.literal(tag); } - auto body_ends = think_open ? data.thinking_end_tags : think_ends; - body_ends.push_back(TURN_END); - auto body_end = p.until_one_of(body_ends); - auto think_body = extract_reasoning ? p.reasoning(body_end) : p.content(body_end); - // the template writes "\n" and "\n"; those newlines are markup, not text - auto nl = p.optional(p.literal("\n")); - auto think_start = p.one_or_more(p.literal(THINK_START) + nl); - auto reasoning = p.optional(p.optional(think_start) + think_body + p.optional(think_close + nl)); + auto think_body = p.until_one_of(data.thinking_end_tags); + auto think_block = [&](const common_peg_parser & body) { + return p.optional(THINK_START + p.space() + p.ac(body + think_end, data.thinking_end_tags)); + }; + auto reasoning = extract_reasoning ? think_block(p.reasoning(think_body)) : p.eps(); if (has_response_format) { - // The final answer must follow a closed reasoning block, including when the prompt pre-opens it. - // Do not inline reasoning into schema-constrained content when extraction is disabled. - auto schema_reasoning = extract_reasoning ? p.reasoning(body_end) : body_end; - auto closed_reasoning = p.optional(think_start + schema_reasoning + think_close + nl); - return opener + closed_reasoning + p.space() + - p.content(p.schema(p.json(), "k2h-response", inputs.json_schema)) + - p.space() + p.optional(p.literal(TURN_END)) + end; + // The answer must be bare JSON, so the think block is consumed even when it is not extracted + auto thoughts = extract_reasoning ? reasoning : think_block(think_body); + return generation_prompt + (thoughts << p.content(p.schema(p.json(), "response-format", inputs.json_schema))); } - auto content = p.optional(p.content(p.until_one_of({ SECTION_START, TURN_END }))); - auto tail = p.optional(p.content(p.until(TURN_END))) + p.optional(p.literal(TURN_END)); - if (!has_tools || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_NONE) { - return opener + reasoning + tail + end; + return generation_prompt + (reasoning << p.content(p.rest())); } - auto tool_choices = p.choice(); - auto arg_close = p.tool_arg_close(p.literal(ARG_VAL_END)); - auto arg_string = p.rule("k2h-arg-string", p.tool_arg_string_value(p.until_one_of({ ARG_VAL_END, TURN_END })) + arg_close); - - foreach_function(inputs.tools, [&](const json & tool) { - const auto & function = tool.at("function"); - std::string name = function.at("name"); - - if (call_format == "json") { - auto schema = common_chat_tool_parameters(function); - auto name_field = p.atomic(p.literal("\"name\"") + p.space() + p.literal(":") + p.space() + - p.literal("\"") + p.tool_name(p.literal(name)) + p.literal("\"")) + p.space(); - auto args_field = p.literal("\"arguments\"") + p.space() + p.literal(":") + p.space() + - p.tool_args(p.schema(p.json(), "k2h-tool-" + name + "-schema", schema)) + p.space(); - auto comma = p.literal(",") + p.space(); - auto call = p.tool(p.tool_open(p.literal(CALL_START) + p.space() + p.literal("{") + p.space()) + - ((name_field + comma + args_field) | (args_field + comma + name_field)) + - p.tool_close(p.literal("}") + p.space() + p.literal(CALL_END))); - tool_choices |= p.rule("k2h-tool-" + name, call); - return; - } - - // xml / xml_typed: strings are raw text up to the closing tag, other types are JSON - std::vector required_args; - std::vector optional_args; - foreach_parameter(function, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) { - auto rule_name = "k2h-arg-" + name + "-" + param.name; - auto types = param.schema->value_types(); - auto json_val = p.tool_arg_json_value(p.schema(p.json(), rule_name + "-schema", doc, *param.schema)) + arg_close; - - auto arg_type = p.eps(); - if (call_format == "xml_typed") { - std::string type_grammar; - for (auto type : { common_chat_schema::TYPE_NULL, common_chat_schema::TYPE_BOOLEAN, - common_chat_schema::TYPE_NUMBER, common_chat_schema::TYPE_INTEGER, - common_chat_schema::TYPE_STRING, common_chat_schema::TYPE_ARRAY, common_chat_schema::TYPE_OBJECT }) { - if (types.has(type)) { - type_grammar += (type_grammar.empty() ? "" : " | ") + gbnf_format_literal(common_chat_schema::type_name(type)); + auto tool_choice = p.choice(); + if (call_format == "json") { + tool_choice = p.standard_json_tools(CALL_START, CALL_END, inputs.tools, false, true); + } else { + auto arg_close = p.tool_arg_close(p.literal(ARG_VAL_END)); + auto arg_string = p.rule("xml-arg-string", p.ac(p.tool_arg_string_value(p.until(ARG_VAL_END)) + arg_close, ARG_VAL_END)); + + // The models leave out even when asked for xml_typed + auto arg_type = call_format == "xml_typed" ? p.optional(ARG_TYPE + p.until(ARG_TYPE_END) + ARG_TYPE_END + p.space()) : p.eps(); + + foreach_function(inputs.tools, [&](const json & tool) { + const auto & function = tool.at("function"); + std::string name = function.at("name"); + + std::vector required_args; + std::vector optional_args; + foreach_parameter(function, [&](const common_chat_schema_property & param, const common_chat_schema_document_ptr & doc) { + auto rule_name = "tool-" + name + "-arg-" + param.name; + auto types = param.schema->value_types(); + auto arg_value = arg_string; + if (!types.has(common_chat_schema::TYPE_STRING)) { + arg_value = p.tool_arg_json_value(p.schema(p.json(), rule_name + "-schema", doc, *param.schema)) + arg_close; + } + if (types.has(common_chat_schema::TYPE_STRING) && !types.is_only(common_chat_schema::TYPE_STRING)) { + // The string alternative accepts any text, so only the parser needs the JSON alternatives. + auto json_value = p.choice(); + if (types.has(common_chat_schema::TYPE_OBJECT)) { + json_value |= p.json_object(); + } + if (types.has(common_chat_schema::TYPE_ARRAY)) { + json_value |= p.json_array(); + } + if (types.has(common_chat_schema::TYPE_NUMBER) || types.has(common_chat_schema::TYPE_INTEGER)) { + json_value |= p.json_number(); + } + if (types.has(common_chat_schema::TYPE_BOOLEAN)) { + json_value |= p.json_bool(); } + if (types.has(common_chat_schema::TYPE_NULL)) { + json_value |= p.json_null(); + } + arg_value = p.gbnf(p.atomic(p.tool_arg_json_value(json_value) + arg_close) | arg_string, "xml-arg-string"); } - // Parse compound labels too, but generate a schema type, never an argument value or markup. - auto type_text = p.chars("[^ \\t\\r\\n<]", 1, 1) + p.chars("[^<]", 0); - arg_type = p.space() + p.literal(ARG_TYPE) + p.space() + - p.gbnf(type_text, "(" + type_grammar + ")") + p.space() + p.literal(ARG_TYPE_END); - } - auto arg_value = types.is_only(common_chat_schema::TYPE_STRING) ? arg_string : - !types.has(common_chat_schema::TYPE_STRING) ? json_val : - p.gbnf(p.atomic(json_val) | arg_string, "k2h-arg-string"); + auto arg = p.space() + p.tool_arg(p.tool_arg_open(ARG_KEY + p.tool_arg_name(p.literal(param.name)) + ARG_KEY_END) << + arg_type + ARG_VAL + arg_value); + (param.required ? required_args : optional_args).push_back(p.rule(rule_name, arg)); + }); - auto arg = p.rule(rule_name, - p.optional(p.space()) + - p.tool_arg(p.tool_arg_open(p.literal(ARG_KEY) + p.tool_arg_name(p.literal(param.name)) + p.literal(ARG_KEY_END)) + - arg_type + p.optional(p.space()) + p.literal(ARG_VAL) + arg_value)); + auto args = p.permute("tool-" + name + "-args", required_args); + if (!optional_args.empty()) { + args = args + p.zero_or_more(p.choice(optional_args)); + } - (param.required ? required_args : optional_args).push_back(arg); + tool_choice |= p.rule("tool-" + name, p.tool( + p.tool_open(CALL_START + p.tool_name(p.literal(name)) + "\n") + p.tool_args(args) << p.tool_close(p.literal(CALL_END)))); }); + } - auto args = p.permute("k2h-" + name + "-args", required_args); - if (!optional_args.empty()) { - args = args + p.zero_or_more(p.choice(optional_args)); - } - - auto call = p.tool(p.tool_open(p.literal(CALL_START) + p.tool_name(p.literal(name)) + p.optional(p.space())) + - p.tool_args(args) + - p.tool_close(p.optional(p.space()) + p.literal(CALL_END))); - tool_choices |= p.rule("k2h-tool-" + name, call); - }); - - auto calls = inputs.parallel_tool_calls ? tool_choices + p.zero_or_more(p.space() + tool_choices) : tool_choices; - - auto tools_section = p.trigger_rule("k2h-tool-call", - p.literal(SECTION_START) + p.space() + calls + p.space() + p.literal(SECTION_END)); + auto required = inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED; + auto calls = inputs.parallel_tool_calls ? tool_choice + p.zero_or_more(p.space() + tool_choice) : tool_choice; + auto tool_calls = p.trigger_rule("tool-calls", p.repeat(SECTION_START << calls << SECTION_END, required ? 1 : 0, 1)); - // a required call follows the reasoning directly, as for gemma4 and gpt-oss - if (inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED) { - return opener + reasoning + p.optional(p.space()) + tools_section + tail + end; + // Keep thinking inline when required calls bypass the content parser. + if (required && !extract_reasoning) { + reasoning = p.content(think_block(think_body)); } - return opener + reasoning + content + p.optional(tools_section) + tail + end; + // A required call follows the reasoning directly, the models otherwise keep writing content + auto content = required ? p.eps() : p.content(p.until(SECTION_START)); + + return generation_prompt + (reasoning << content << tool_calls); }); data.parser = parser.save(); if (include_grammar) { - data.grammar_lazy = !has_response_format && inputs.tool_choice != COMMON_CHAT_TOOL_CHOICE_REQUIRED; + data.grammar_lazy = !(has_response_format || inputs.tool_choice == COMMON_CHAT_TOOL_CHOICE_REQUIRED); data.grammar = build_grammar([&](const common_grammar_builder & builder) { parser.build_grammar(builder, data.grammar_lazy); }); diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index c9144644fc29..0bcf386def14 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -2920,6 +2920,7 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) { || t.first == "<|tool_response>" // gemma4 || t.first == "<|end▁of▁sentence|>" // deepseek-ocr || t.first == "[e~[" // minimax-m2/m3 + || t.first == "<|ifm|im_end|>" // k2-horizon ) { special_eog_ids.insert(t.second); if ((attr & LLAMA_TOKEN_ATTR_CONTROL) == 0) { diff --git a/tests/test-chat.cpp b/tests/test-chat.cpp index eeecf56cfb3e..c8569d3945fb 100644 --- a/tests/test-chat.cpp +++ b/tests/test-chat.cpp @@ -4864,220 +4864,238 @@ static void test_template_output_peg_parsers(bool detailed_debug) { .run(); } - // K2 Horizon dedicated parser + // K2 Horizon - reasoning (/ for medium/low reasoning_effort), + // tool calls in an section as xml (default), xml_typed or json { - auto tmpls = read_templates("models/templates/IFM-K2-Horizon.jinja"); - const auto caps = common_chat_templates_get_caps(tmpls.get()); - GGML_ASSERT(caps.at("supports_parallel_tool_calls")); - GGML_ASSERT(caps.at("supports_object_arguments")); - assert_contains(common_chat_format_example(tmpls.get(), true, {}), "Hi there"); - auto tst = peg_tester("models/templates/IFM-K2-Horizon.jinja", detailed_debug); - const std::string answer_schema = R"({"type":"object","properties":{"answer":{"type":"integer","const":42}},"required":["answer"],"additionalProperties":false})"; - tst.test("Let me calculate.\n{\"answer\":42}<|ifm|im_end|>") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .json_schema(answer_schema) - .expect_reasoning("Let me calculate.") - .expect_content(R"({"answer":42})") + tst.test("I'm\nthinkingHello, world!\nWhat's up?") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .expect(message_assist_thoughts) + .expect_reconstruction() .run(); - tst.test("Let me calculate.{\"answer\":42}") + tst.test("I'm\nthinkingHello, world!\nWhat's up?") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .chat_template_kwargs({ { "reasoning_effort", R"("medium")" } }) + .expect(message_assist_thoughts) + .run(); + + tst.test("I'm\nthinkingHello, world!\nWhat's up?") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .chat_template_kwargs({ { "reasoning_effort", R"("low")" } }) + .expect(message_assist_thoughts) + .run(); + + // The 3.7B ends medium and low effort reasoning with + tst.test("I'm\nthinkingHello, world!\nWhat's up?") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .chat_template_kwargs({ { "reasoning_effort", R"("medium")" } }) + .expect(message_assist_thoughts) + .run(); + + tst.test("I'm\nthinking") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .expect_reasoning("I'm\nthinking") + .run(); + + tst.test("I'm\nthinkingHello, world!\nWhat's up?") .reasoning_format(COMMON_REASONING_FORMAT_NONE) - .json_schema(answer_schema) - .expect_content(R"({"answer":42})") + .expect_content("\nI'm\nthinkingHello, world!\nWhat's up?") .run(); - // Prefill advances the grammar through both reasoning and partial final content. - tst.test("42}") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .json_schema(answer_schema) - .messages({ message_user, simple_assist_msg("{\"answer\":", "Calculated.") }) - .continue_final_message(COMMON_CHAT_CONTINUATION_CONTENT) - .expect_reasoning("Calculated.") - .expect_content(R"({"answer":42})") + tst.test( + "I'm\nthinking\n" + "special_function\n" + "arg1\n" + "1\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ special_function_tool }) + .expect(message_assist_call_thoughts) + .expect_reconstruction() .run(); - common_chat_templates_inputs schema_inputs; - schema_inputs.messages = { message_user }; - schema_inputs.add_generation_prompt = true; - schema_inputs.json_schema = answer_schema; - schema_inputs.tools = { get_time_tool }; - auto schema_params = common_chat_templates_apply(tmpls.get(), schema_inputs); - GGML_ASSERT(!schema_params.grammar.empty()); - GGML_ASSERT(!schema_params.grammar_lazy); - GGML_ASSERT(schema_params.grammar_triggers.empty()); - for (const std::string output : { - "{\"answer\":42}", - "```json\n{\"answer\":42}\n```", - "{\"answer\":\"42\"}", - "{\"answer\":41}", - "{\"wrong\":42}", - "{\"answer\":42,\"extra\":1}", - "{\"answer\":42} trailing text", - "Still thinking", "" }) { - auto grammar = build_grammar(schema_params.grammar); - GGML_ASSERT(match_string(schema_params.generation_prompt + output, grammar.get()) == - (output == "{\"answer\":42}")); - } - // A stop marker inside reasoning must be rejected, not accepted as an incomplete answer. - auto stop_grammar = build_grammar(schema_params.grammar); - auto stop_match = match_string_detailed(schema_params.generation_prompt + "<|ifm|im_end|>", stop_grammar.get()); - GGML_ASSERT(!stop_match.success && !stop_match.incomplete); + tst.test( + "I'm\nthinking\n" + "special_function\n" + "arg1\n" + "integer\n" + "1\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ special_function_tool }) + .chat_template_kwargs({ { "tool_call_format", R"("xml_typed")" } }) + .expect(message_assist_call_thoughts) + .expect_reconstruction() + .run(); - const std::string get_time_call = - "\n" - "get_time\n" - "city\n" - "Paris\n" - "\n" - ""; - - // JSON envelopes allow whitespace and either field order, including during streaming. - for (const std::string payload : { - R"({"name":"get_time","arguments":{"city":"Paris"}})", - R"({ "arguments" : {"city":"Paris"}, "name" : "get_time" })", - "{\n\t\"name\" : \"get_time\",\n\"arguments\" : {\"city\":\"Paris\"}\n}" }) { - tst.test("" + payload + "") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .tools({ get_time_tool }) - .chat_template_kwargs({ { "tool_call_format", R"("json")" } }) - .expect_tool_calls({ { "get_time", R"({"city":"Paris"})", "" } }) - .run(); - } + // The models leave out the type even when asked for xml_typed + tst.test( + "I'm\nthinking\n" + "special_function\n" + "arg1\n" + "1\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ special_function_tool }) + .chat_template_kwargs({ { "tool_call_format", R"("xml_typed")" } }) + .expect(message_assist_call_thoughts) + .run(); - // Do not emit the shorter name while a longer name is still being streamed. - auto longer_name_tool = get_time_tool; - longer_name_tool.name += "_extended"; - tst.test("" - "{\"name\":\"get_time_extended\",\"arguments\":{\"city\":\"Paris\"}}" - "") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .tools({ get_time_tool, longer_name_tool }) + tst.test( + "I'm\nthinking\n" + "{\"name\": \"special_function\", \"arguments\": {\"arg1\": 1}}\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ special_function_tool }) .chat_template_kwargs({ { "tool_call_format", R"("json")" } }) - .expect_tool_calls({ { "get_time_extended", R"({"city":"Paris"})", "" } }) + .expect(message_assist_call_thoughts) + .expect_reconstruction() .run(); - const std::string typed_call = - "get_time" - "citystring" - "Paris"; - tst.test("" + typed_call) - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) + tst.test( + "\n" + "empty_args\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ empty_args_tool }) + .expect(simple_assist_msg("", "", "empty_args", "{}")) + .run(); + + tst.test( + "\n" + "get_time\n" + "city\n" + "Paris\n" + "\n" + "get_time\n" + "city\n" + "Rome\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .parallel_tool_calls(true) .tools({ get_time_tool }) - .chat_template_kwargs({ { "tool_call_format", R"("xml_typed")" } }) - .expect_tool_calls({ { "get_time", R"({"city":"Paris"})", "" } }) + .expect_tool_calls({ + { "get_time", R"({"city": "Paris"})", {} }, + { "get_time", R"({"city": "Rome"})", {} }, + }) .run(); - // Wrong XML dialects are not complete calls and cannot be generated by the grammar. - for (const std::string format : { "xml", "xml_typed" }) { - common_chat_templates_inputs inputs; - inputs.messages = { message_user }; - inputs.tools = { get_time_tool }; - inputs.reasoning_format = COMMON_REASONING_FORMAT_DEEPSEEK; - inputs.chat_template_kwargs["tool_call_format"] = json(format).dump(); - auto parser = make_peg_parser(tmpls.get(), inputs); - const auto & invalid = format == "xml" ? typed_call : get_time_call; - GGML_ASSERT(parser.parse("" + invalid, false).tool_calls.empty()); - auto grammar = build_grammar(parser.params_.grammar); - GGML_ASSERT(!match_string(invalid, grammar.get())); - - const std::string arg_prefix = "get_timecity"; - const std::string unfinished = arg_prefix + (format == "xml" ? "Paris" : "string"); - auto stop_grammar = build_grammar(parser.params_.grammar); - auto stop_match = match_string_detailed(unfinished + "<|ifm|im_end|>", stop_grammar.get()); - GGML_ASSERT(!stop_match.success && !stop_match.incomplete); - if (format == "xml_typed") { - auto type_grammar = build_grammar(parser.params_.grammar); - auto type_match = match_string_detailed(arg_prefix + "17", type_grammar.get()); - GGML_ASSERT(!type_match.success && !type_match.incomplete); - } - } + tst.test( + "\n" + "tool_2req_4opt\n" + "req2\n" + "7\n" + "req1\n" + "hello\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ tool_2req_4opt }) + .expect_tool_calls({ { "tool_2req_4opt", R"({"req2": 7, "req1": "hello"})", {} } }) + .run(); - for (const std::string effort : { "high", "medium", "low" }) { - const auto tag = effort == "high" ? "ifm|think" : effort == "medium" ? "ifm|think_fast" : "ifm|think_faster"; - for (const std::string close : { "", "", "" }) { - tst.test("<" + std::string(tag) + ">Plan." + close + "42") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .chat_template_kwargs({ { "reasoning_effort", json(effort).dump() } }) - .expect_reasoning("Plan.") - .expect_content("42") - .run(); - tst.test("Plan." + close + "{\"answer\":42}") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .chat_template_kwargs({ { "reasoning_effort", json(effort).dump() } }) - .json_schema(answer_schema) - .expect_reasoning("Plan.") - .expect_content(R"({"answer":42})") - .run(); - } + for (const std::string value : { "true", "42", "null", "[]", R"("quoted")", "{not valid json" }) { + tst.test( + "\n" + "set_union\n" + "value\n" + "" + value + "\n" + "amount\n" + "42\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ string_union_tool }) + .expect_tool_calls({ { "set_union", json({ { "value", value }, { "amount", 42 } }).dump(), {} } }) + .run(); } - tst.test("Still thinking") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .expect_reasoning("Still thinking") + tst.test( + "\n" + "set_union\n" + "value\n" + "{\"a\": 1}\n" + "amount\n" + "2 dollars\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ string_union_tool }) + .expect_tool_calls({ { "set_union", R"({"value": {"a": 1}, "amount": "2 dollars"})", {} } }) .run(); - // The generation prompt pre-opens , so the model output starts inside it. - tst.test("Simple sum.\n\n51") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .expect_reasoning("Simple sum.\n") - .expect_content("51") + tst.test( + "I'm\nthinking\n" + "special_function\n" + "arg1\n" + "1\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ special_function_tool }) + .tool_choice(COMMON_CHAT_TOOL_CHOICE_REQUIRED) + .expect(message_assist_call_thoughts) .run(); - // The end-of-turn token must not leak into content. - tst.test("Simple sum.\n\n51<|ifm|im_end|>") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .expect_reasoning("Simple sum.\n") - .expect_content("51") + tst.test( + "I'm\nthinking\n" + "special_function\n" + "arg1\n" + "1\n" + "\n" + "") + .reasoning_format(COMMON_REASONING_FORMAT_NONE) + .tools({ special_function_tool }) + .tool_choice(COMMON_CHAT_TOOL_CHOICE_REQUIRED) + .expect_content("\nI'm\nthinking") + .expect_tool_calls({ { "special_function", R"({"arg1": 1})", {} } }) .run(); - // A closed think block followed by a tool call section. - tst.test("I need the time.\n\n" + get_time_call) - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .tools({ get_time_tool }) - .expect_reasoning("I need the time.\n") - .expect_tool_calls({ { "get_time", R"({"city": "Paris"})", "" } }) + tst.test("I'm\nthinkingHello, world!\nWhat's up?") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .tools({ special_function_tool }) + .expect(message_assist_thoughts) .run(); - // A tool call section may start before the think block is closed. - tst.test("I need the time.\n" + get_time_call) - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .tools({ get_time_tool }) - .expect_reasoning("I need the time.\n") - .expect_tool_calls({ { "get_time", R"({"city": "Paris"})", "" } }) + const std::string answer_schema = R"({"type":"object","properties":{"answer":{"type":"integer"}},"required":["answer"]})"; + + tst.test("Let me calculate.{\"answer\":42}") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .json_schema(answer_schema) + .expect_reasoning("Let me calculate.") + .expect_content(R"({"answer":42})") .run(); - // Non-string arguments parse as JSON, and required arguments may come in any order. - tst.test("\n\ntool_2req_4opt\n" - "req2\n7\n" - "req1\nhello\n" - "\n") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .tools({ tool_2req_4opt }) - .expect_tool_calls({ { "tool_2req_4opt", R"({"req2": 7, "req1": "hello"})", "" } }) + tst.test("Let me calculate.{\"answer\":42}") + .reasoning_format(COMMON_REASONING_FORMAT_NONE) + .json_schema(answer_schema) + .expect_content(R"({"answer":42})") .run(); - // Parallel tool calls share one section. - tst.test("\n\n" - "get_time\ncity\nParis\n\n" - "get_time\ncity\nRome\n\n" - "") - .reasoning_format(COMMON_REASONING_FORMAT_DEEPSEEK) - .tools({ get_time_tool }) - .parallel_tool_calls(true) - .expect_tool_calls({ - { "get_time", R"({"city": "Paris"})", "" }, - { "get_time", R"({"city": "Rome"})", "" }, - }) + tst.test("42}") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .json_schema(answer_schema) + .messages({ message_user, simple_assist_msg("{\"answer\":", "Calculated.") }) + .add_generation_prompt(false) + .continue_final_message(COMMON_CHAT_CONTINUATION_CONTENT) + .expect_reasoning("Calculated.") + .expect_content(R"({"answer":42})") .run(); - // reasoning_format=none keeps extracting tool calls. - tst.test("I need the time.\n\n" + get_time_call) - .reasoning_format(COMMON_REASONING_FORMAT_NONE) - .tools({ get_time_tool }) - .expect_content("I need the time.\n") - .expect_tool_calls({ { "get_time", R"({"city": "Paris"})", "" } }) + tst.test(" thinkingHello, world!\nWhat's up?") + .reasoning_format(COMMON_REASONING_FORMAT_AUTO) + .messages({ message_user, message_assist_prefill_reasoning }) + .add_generation_prompt(false) + .continue_final_message(COMMON_CHAT_CONTINUATION_REASONING) + .expect_reasoning("I'm thinking") + .expect_content("Hello, world!\nWhat's up?") .run(); } From 59db596d0e3d3e79856b2078c530a4f1cf280d7b Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:27:01 +0000 Subject: [PATCH 17/22] clean test-chat --- tests/test-chat.cpp | 36 +----------------------------------- 1 file changed, 1 insertion(+), 35 deletions(-) diff --git a/tests/test-chat.cpp b/tests/test-chat.cpp index c8569d3945fb..4c1b980b99d7 100644 --- a/tests/test-chat.cpp +++ b/tests/test-chat.cpp @@ -4864,8 +4864,7 @@ static void test_template_output_peg_parsers(bool detailed_debug) { .run(); } - // K2 Horizon - reasoning (/ for medium/low reasoning_effort), - // tool calls in an section as xml (default), xml_typed or json + // K2 Horizon { auto tst = peg_tester("models/templates/IFM-K2-Horizon.jinja", detailed_debug); @@ -4875,25 +4874,6 @@ static void test_template_output_peg_parsers(bool detailed_debug) { .expect_reconstruction() .run(); - tst.test("I'm\nthinkingHello, world!\nWhat's up?") - .reasoning_format(COMMON_REASONING_FORMAT_AUTO) - .chat_template_kwargs({ { "reasoning_effort", R"("medium")" } }) - .expect(message_assist_thoughts) - .run(); - - tst.test("I'm\nthinkingHello, world!\nWhat's up?") - .reasoning_format(COMMON_REASONING_FORMAT_AUTO) - .chat_template_kwargs({ { "reasoning_effort", R"("low")" } }) - .expect(message_assist_thoughts) - .run(); - - // The 3.7B ends medium and low effort reasoning with - tst.test("I'm\nthinkingHello, world!\nWhat's up?") - .reasoning_format(COMMON_REASONING_FORMAT_AUTO) - .chat_template_kwargs({ { "reasoning_effort", R"("medium")" } }) - .expect(message_assist_thoughts) - .run(); - tst.test("I'm\nthinking") .reasoning_format(COMMON_REASONING_FORMAT_AUTO) .expect_reasoning("I'm\nthinking") @@ -4932,20 +4912,6 @@ static void test_template_output_peg_parsers(bool detailed_debug) { .expect_reconstruction() .run(); - // The models leave out the type even when asked for xml_typed - tst.test( - "I'm\nthinking\n" - "special_function\n" - "arg1\n" - "1\n" - "\n" - "") - .reasoning_format(COMMON_REASONING_FORMAT_AUTO) - .tools({ special_function_tool }) - .chat_template_kwargs({ { "tool_call_format", R"("xml_typed")" } }) - .expect(message_assist_call_thoughts) - .run(); - tst.test( "I'm\nthinking\n" "{\"name\": \"special_function\", \"arguments\": {\"arg1\": 1}}\n" From fbee59434f883bf1675ee7506f91eafb7f9dc75f Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Fri, 2 Oct 2026 19:49:35 +0000 Subject: [PATCH 18/22] remove fallback : model trained mostly on high anyway --- common/parsers/k2-horizon.cpp | 4 ---- 1 file changed, 4 deletions(-) diff --git a/common/parsers/k2-horizon.cpp b/common/parsers/k2-horizon.cpp index 0348e5e72e9d..687a893cee2e 100644 --- a/common/parsers/k2-horizon.cpp +++ b/common/parsers/k2-horizon.cpp @@ -47,10 +47,6 @@ common_chat_params common_chat_params_init_k2_horizon(const common_chat_template data.thinking_start_tag = THINK_START; data.thinking_end_tags = { THINK_END }; - if (think != "ifm|think") { - // The 3.7B ends medium and low effort reasoning with - data.thinking_end_tags.push_back(""); - } data.preserved_tokens = data.thinking_end_tags; data.preserved_tokens.insert(data.preserved_tokens.end(), { From c9b78f7dd06d630790d1ea43b226b1f85ad23e1f Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Sun, 4 Oct 2026 08:14:06 +0000 Subject: [PATCH 19/22] fix k2 attn_v_exp tn splitting and metal fusion baseline --- src/llama-model.cpp | 1 + tests/fusion/MTL.csv | 4 ++++ 2 files changed, 5 insertions(+) diff --git a/src/llama-model.cpp b/src/llama-model.cpp index 952df7646246..8f6da668ee7b 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -795,6 +795,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str // three stay in lockstep per device const int64_t granularity_v = (granularity_kv / hparams.n_embd_head_k(il)) * hparams.n_embd_head_v(il); if (std::regex_match(tensor_name, pattern_kv_weight) || + std::regex_match(tensor_name, pattern_v_exps_weight) || std::regex_match(tensor_name, pattern_kv_bias) || std::regex_match(tensor_name, pattern_kv_cache)) { GGML_ASSERT(segments.size() == 1); diff --git a/tests/fusion/MTL.csv b/tests/fusion/MTL.csv index 9d771ddce248..e839e20ce3f5 100644 --- a/tests/fusion/MTL.csv +++ b/tests/fusion/MTL.csv @@ -126,6 +126,10 @@ internlm2 ,0 ,any ,RMS_NORM+MUL , 5 jais ,0 ,any ,NORM+MUL+ADD , 5 jais2 ,0 ,any ,NORM+MUL+ADD , 5 jamba ,0 ,any ,RMS_NORM+MUL , 8 +k2-horizon ,0 ,any ,RMS_NORM+MUL , 4 +k2-horizon ,0 ,any ,ADD+ADD , 1 +k2-horizon ,0 ,any ,MUL+ADD , 2 +k2-horizon ,0 ,any ,RMS_NORM+MUL , 4 kimi-k3 ,0 ,any ,GATED_DELTA_NET+CPY , 1 kimi-k3 ,0 ,any ,MUL+ADD , 2 kimi-k3 ,0 ,any ,RMS_NORM+MUL , 17 From 5b9a73799a4acbcb16509bfdb7c70bd55aed46d7 Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Tue, 6 Oct 2026 10:06:43 +0000 Subject: [PATCH 20/22] k2-horizon : forward expand views before sums --- src/models/k2-horizon.cpp | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp index 5c158ce77807..2229527388ad 100644 --- a/src/models/k2-horizon.cpp +++ b/src/models/k2-horizon.cpp @@ -190,9 +190,17 @@ ggml_tensor * llama_model_k2_horizon::graph::build_routed_value(const llama_laye // sum the selected experts; 3D views of {n_embd_gqa, 1, n_tokens} keep the strides of values, // which lets the tensor-parallel backend follow its split through the views - ggml_tensor * value_out = ggml_view_3d(ctx0, values, n_embd_gqa, 1, n_tokens, values->nb[1], values->nb[2], 0); + // order the views before the adds so backends can fuse the sum + ggml_tensor * value_views[LLAMA_MAX_EXPERTS] = { nullptr }; + for (int64_t i = 0; i < n_used; ++i) { + value_views[i] = ggml_view_3d(ctx0, values, n_embd_gqa, 1, n_tokens, values->nb[1], values->nb[2], i * values->nb[1]); + ggml_build_forward_expand(gf, value_views[i]); + } + + ggml_tensor * value_out = value_views[0]; for (int64_t i = 1; i < n_used; ++i) { - value_out = ggml_add(ctx0, value_out, ggml_view_3d(ctx0, values, n_embd_gqa, 1, n_tokens, values->nb[1], values->nb[2], i * values->nb[1])); + value_out = ggml_add(ctx0, value_out, value_views[i]); + ggml_build_forward_expand(gf, value_out); } if (n_used == 1) { value_out = ggml_cont(ctx0, value_out); From ec9106d5ec4f2945154b6e720003961671893761 Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Tue, 6 Oct 2026 11:46:55 +0000 Subject: [PATCH 21/22] k2-horizon: copy embds before group norm to fix TP --- src/models/k2-horizon.cpp | 3 +++ 1 file changed, 3 insertions(+) diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp index 2229527388ad..c792bf3b2b0c 100644 --- a/src/models/k2-horizon.cpp +++ b/src/models/k2-horizon.cpp @@ -219,6 +219,9 @@ llama_model_k2_horizon::graph::graph(const llama_model & model, const llm_graph_ inpL = build_inp_embd(model.tok_embd); + // Copy the embeddings before group norm reshapes them, so the Meta backend doesn't get a view of CPU memory. + inpL = ggml_cont(ctx0, inpL); + // inp_pos - contains the positions ggml_tensor * inp_pos = build_inp_pos(); From 7f4fa00776d769fee3a77f982acab015a3e798d1 Mon Sep 17 00:00:00 2001 From: Bilal <60072763+bitalov@users.noreply.github.com> Date: Tue, 6 Oct 2026 15:29:25 +0000 Subject: [PATCH 22/22] disable tesnor parallelism --- src/llama-arch.cpp | 1 + src/models/k2-horizon.cpp | 3 --- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index 5e828b59d4e9..73d6899bcdb8 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -1217,6 +1217,7 @@ bool llm_arch_supports_sm_tensor(const llm_arch & arch) { case LLM_ARCH_KIMI_K3: case LLM_ARCH_GLM5_NEXT: case LLM_ARCH_QWEN3TTS: + case LLM_ARCH_K2_HORIZON: return false; default: return true; diff --git a/src/models/k2-horizon.cpp b/src/models/k2-horizon.cpp index c792bf3b2b0c..2229527388ad 100644 --- a/src/models/k2-horizon.cpp +++ b/src/models/k2-horizon.cpp @@ -219,9 +219,6 @@ llama_model_k2_horizon::graph::graph(const llama_model & model, const llm_graph_ inpL = build_inp_embd(model.tok_embd); - // Copy the embeddings before group norm reshapes them, so the Meta backend doesn't get a view of CPU memory. - inpL = ggml_cont(ctx0, inpL); - // inp_pos - contains the positions ggml_tensor * inp_pos = build_inp_pos();