diff --git a/common/common.cpp b/common/common.cpp index 4005990ce2..a368835577 100644 --- a/common/common.cpp +++ b/common/common.cpp @@ -3,6 +3,9 @@ #include "build-info.h" #include "common.h" + +#include "../src/llama-ext.h" + #include "fit.h" #include "log.h" #include "llama.h" @@ -1208,6 +1211,9 @@ common_decision_type common_get_decision_type(const struct llama_model * model) if (type == "laya") { return COMMON_DECISION_TYPE_LAYA; } + if (type == "clef") { + return COMMON_DECISION_TYPE_CLEF; + } return COMMON_DECISION_TYPE_UNKNOWN; } @@ -1263,9 +1269,10 @@ common_init_result::common_init_result(common_params & params, bool model_only) const llama_vocab * vocab = llama_model_get_vocab(model); - // this decision model returns a score for each token via the embeddings output + // these decision models return a score for each token via the embeddings output // TODO: maybe improve this in the future - if (common_get_decision_type(model) == COMMON_DECISION_TYPE_LAYA) { + const auto decision_type = common_get_decision_type(model); + if (decision_type == COMMON_DECISION_TYPE_LAYA || decision_type == COMMON_DECISION_TYPE_CLEF) { params.embedding = true; params.pooling_type = LLAMA_POOLING_TYPE_NONE; @@ -1274,7 +1281,7 @@ common_init_result::common_init_result(common_params & params, bool model_only) cparams.n_outputs_max = cparams.n_batch; cparams.n_outputs_max_per_seq = 1; - LOG_INF("%s", "laya decision model detected, enabling embedding mode\n"); + LOG_INF("%s", "decision model detected, enabling embedding mode\n"); } // load and optionally apply lora adapters @@ -2211,6 +2218,9 @@ llama_batch_ext * common_batch::get_sub_batch(int32_t off, int32_t n) { if (t.output) { llama_batch_ext_set_output_logits(res, idx, true); } + if (t.decision_order != 0) { + llama_batch_ext_set_decision_order(res, idx, t.decision_order); + } } return res; diff --git a/common/common.h b/common/common.h index 580a162fe5..95d75ce90b 100644 --- a/common/common.h +++ b/common/common.h @@ -950,6 +950,7 @@ enum common_decision_type { COMMON_DECISION_TYPE_UNKNOWN, // a decision model of a type that is not supported COMMON_DECISION_TYPE_OPENJEV, // logits of one label token per option, read at the last prompt token COMMON_DECISION_TYPE_LAYA, // score of one marker token per option, read from the embeddings output + COMMON_DECISION_TYPE_CLEF, // all questions in one prompt, score of option i read from the embeddings output at row i }; common_decision_type common_get_decision_type(const struct llama_model * model); @@ -1051,6 +1052,7 @@ struct common_batch { bool output; llama_embd embd; // non-owning view of the data passed to add_embd()/set_embd(), data == NULL if none std::vector seq_ids_extra; // see add_seq() + int32_t decision_order = 0; // see llama_batch_ext_set_decision_order() }; std::vector tokens; // mirror of the entries, tokens[i] describes batch index i diff --git a/conversion/__init__.py b/conversion/__init__.py index fac5d95a9f..4daaafa1e0 100644 --- a/conversion/__init__.py +++ b/conversion/__init__.py @@ -42,6 +42,7 @@ TEXT_MODEL_MAP: dict[str, str] = { "ChameleonForConditionalGeneration": "chameleon", "ChatGLMForConditionalGeneration": "chatglm", "ChatGLMModel": "chatglm", + "ClefModel": "clef", "CodeShellForCausalLM": "codeshell", "CogVLMForCausalLM": "cogvlm", "Cohere2MoeForCausalLM": "command_r", diff --git a/conversion/clef.py b/conversion/clef.py new file mode 100644 index 0000000000..58fa9284a2 --- /dev/null +++ b/conversion/clef.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +import json +import math + +from pathlib import Path +from typing import Any, Iterable, TYPE_CHECKING + +import torch + +if TYPE_CHECKING: + from torch import Tensor + +from .base import ModelBase, gguf, logger +from .qwen import Qwen3_5TextModel + + +def _is_clef_checkpoint(dir_model: Path) -> bool: + return (dir_model / "joint_head_config.json").is_file() and (dir_model / "config.json").is_file() + + +@ModelBase.register_hparams_loader(_is_clef_checkpoint) +def _load_clef_hparams(dir_model: Path) -> dict[str, Any]: + logger.info("gguf: detected Clef checkpoint") + hparams = ModelBase.load_hparams(dir_model, False, guess=False) + hparams["architectures"] = ["ClefModel"] + with open(dir_model / "joint_head_config.json", encoding="utf-8") as f: + hparams["decision"] = json.load(f) + return hparams + + +# TODO: image input needs token and embedding entries in the same batch, see https://github.com/ggml-org/llama.cpp/pull/29622 +@ModelBase.register("ClefModel") +class ClefModel(Qwen3_5TextModel): + model_arch = gguf.MODEL_ARCH.CLEF + no_mtp = True # the checkpoint has no MTP head + + # prompt follows joint_schema_model.py of the model repo + _SYSTEM_PROMPT = ( + "Read the complete state and schema. Decide every field jointly. Each answer " + "must be exactly one of that field's allowed options." + ) + # the pieces of the prompt are tokenized one by one, this separates them + _PIECE_SEP = "<>" + # start of a piece: state, question span, option span + _PIECE_STATE, _PIECE_QUESTION, _PIECE_OPTION = "<>", "<>", "<>" + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + head = self.hparams["decision"] + self._n_routing = head["routing_layers"] + # the head blocks are named dec.blk.N, routing blocks first + self.tensor_map = gguf.get_tensor_name_map(self.model_arch, max(self.block_count, self._n_routing + head["layers"])) + self._scales: dict[str, float] = {} + + def set_vocab(self): + super().set_vocab() + self.gguf_writer.add_chat_template([{"name": "systemone", "template": self._systemone_template()}]) + + @classmethod + def _systemone_template(cls) -> str: + def text(value: str) -> str: + return "{{ " + json.dumps(value) + " }}" + + sep = text(cls._PIECE_SEP) + # state, instructions and option text are given as strings + return ( + text(f"<|im_start|>system\n{cls._SYSTEM_PROMPT}<|im_end|>\n<|im_start|>user\nSTATE:\n") + + sep + text(cls._PIECE_STATE) + "{{ state }}" + + sep + text("\n\nSCHEMA FIELDS:\n") + + "{% for q in questions %}" + + sep + text("\nFIELD ") + "{{ loop.index }}" + text("\nID: ") + "{{ q.id }}" + + text("\nTYPE: ") + "{{ q.type }}" + text("\nINSTRUCTION: ") + + sep + text(cls._PIECE_QUESTION) + "{{ q.instructions }}" + + sep + text("\nALLOWED OPTIONS:\n") + + "{% for o in q.options %}" + + sep + text("OPTION ") + "{{ loop.index }}" + text(": ") + + sep + text(cls._PIECE_OPTION) + "{{ o.text }}" + + sep + text("\n") + + "{% endfor %}" + + sep + text("END FIELD\n") + + "{% endfor %}" + + sep + text("\n<|im_end|>\n<|im_start|>assistant\n\n\n\n\nJOINT SCHEMA DECISIONS:") + ) + + def set_gguf_parameters(self): + super().set_gguf_parameters() + head = self.hparams["decision"] + self.gguf_writer.add_decision_type(gguf.DecisionType.CLEF) + self.gguf_writer.add_decision_routing_block_count(head["routing_layers"]) + self.gguf_writer.add_decision_block_count(head["layers"]) + self.gguf_writer.add_decision_head_count(head["heads"]) + + def generate_extra_tensors(self) -> Iterable[tuple[str, Tensor]]: + yield from super().generate_extra_tensors() + from safetensors.torch import load_file + for name, data in load_file(self.dir_model / "joint_head.safetensors").items(): + yield "joint_head." + name, data + + def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]: + if not name.startswith("joint_head."): + yield from super().modify_tensors(data_torch, name, bid) + return + + parts = name.split(".") + + # learned scalars, stored as the values used at inference + if len(parts) == 2 and data_torch.ndim == 0: + value = float(data_torch) + if parts[1] == "residual_gate": + self._scales[parts[1]] = 1.0 / (1.0 + math.exp(-value)) + else: + self._scales[parts[1]] = math.exp(min(value, math.log(100.0))) + if len(self._scales) == 3: + scales = [self._scales[k] for k in ("prior_logit_scale", "joint_logit_scale", "residual_gate")] + yield self.format_tensor_name(gguf.MODEL_TENSOR.DECISION_SCALES, suffix=""), torch.tensor(scales, dtype=torch.float32) + return + + # routing blocks come first + if parts[1] == "layers": + parts[2] = str(int(parts[2]) + self._n_routing) + name = ".".join(parts) + + # nn.MultiheadAttention keeps q, k, v in one tensor + for suffix in ("weight", "bias"): + if name.endswith(".in_proj_" + suffix): + prefix = name[:-len("in_proj_" + suffix)] + for x, data in zip("qkv", data_torch.chunk(3, dim=0)): + yield self.map_tensor_name(prefix + x + "." + suffix), data + return + + yield self.map_tensor_name(name), data_torch diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index dab8445a9a..8cfd000ad1 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -324,6 +324,8 @@ class Keys: TYPE = "{arch}.decision.type" # note: single-use-case keys can be hard-coded in cpp code BLOCK_COUNT = "{arch}.decision.block_count" + ROUTING_BLOCK_COUNT = "{arch}.decision.routing_block_count" + HEAD_COUNT = "{arch}.decision.head_count" MAX_HEAD_TOKENS = "{arch}.decision.max_head_tokens" TEMPERATURE = "{arch}.decision.temperature.{name}" # name: "" or "." @@ -537,6 +539,7 @@ class MODEL_ARCH(IntEnum): QWEN3VLMOE = auto() QWEN35 = auto() QWEN35MOE = auto() + CLEF = auto() QWEN4EXP = auto() PHI2 = auto() PHI3 = auto() @@ -871,6 +874,18 @@ class MODEL_TENSOR(IntEnum): DEC_FFN_DOWN = auto() DEC_FFN_UP = auto() DEC_OUTPUT_NORM = auto() + DEC_CROSS_ATTN_NORM_KV = auto() + DECISION_HIDDEN_NORM = auto() + DECISION_PROJ_MEMORY = auto() + DECISION_PROJ_QUESTION = auto() + DECISION_PROJ_OPTION_QUESTION = auto() + DECISION_PROJ_GLOBAL = auto() + DECISION_PROJ_OPTION_CONTEXT = auto() + DECISION_PROJ_OPTION_LEXICAL = auto() + DECISION_OPTION_SUMMARY_NORM = auto() + DECISION_FIELD_NORM = auto() + DECISION_OPTION_NORM = auto() + DECISION_SCALES = auto() ENC_ATTN_NORM = auto() ENC_ATTN_Q = auto() ENC_ATTN_K = auto() @@ -1301,6 +1316,7 @@ MODEL_ARCH_NAMES: dict[MODEL_ARCH, str] = { MODEL_ARCH.QWEN3VLMOE: "qwen3vlmoe", MODEL_ARCH.QWEN35: "qwen35", MODEL_ARCH.QWEN35MOE: "qwen35moe", + MODEL_ARCH.CLEF: "clef", MODEL_ARCH.QWEN4EXP: "qwen4exp", MODEL_ARCH.PHI2: "phi2", MODEL_ARCH.PHI3: "phi3", @@ -1634,6 +1650,18 @@ TENSOR_NAMES: dict[MODEL_TENSOR, str] = { MODEL_TENSOR.DEC_FFN_DOWN: "dec.blk.{bid}.ffn_down", MODEL_TENSOR.DEC_FFN_UP: "dec.blk.{bid}.ffn_up", MODEL_TENSOR.DEC_OUTPUT_NORM: "dec.output_norm", + MODEL_TENSOR.DEC_CROSS_ATTN_NORM_KV: "dec.blk.{bid}.cross_attn_norm_kv", + MODEL_TENSOR.DECISION_HIDDEN_NORM: "decision.hidden_norm", + MODEL_TENSOR.DECISION_PROJ_MEMORY: "decision.proj_memory", + MODEL_TENSOR.DECISION_PROJ_QUESTION: "decision.proj_question", + MODEL_TENSOR.DECISION_PROJ_OPTION_QUESTION: "decision.proj_option_question", + MODEL_TENSOR.DECISION_PROJ_GLOBAL: "decision.proj_global", + MODEL_TENSOR.DECISION_PROJ_OPTION_CONTEXT: "decision.proj_option_context", + MODEL_TENSOR.DECISION_PROJ_OPTION_LEXICAL: "decision.proj_option_lexical", + MODEL_TENSOR.DECISION_OPTION_SUMMARY_NORM: "decision.option_summary_norm", + MODEL_TENSOR.DECISION_FIELD_NORM: "decision.field_norm", + MODEL_TENSOR.DECISION_OPTION_NORM: "decision.option_norm", + MODEL_TENSOR.DECISION_SCALES: "decision.scales", MODEL_TENSOR.ENC_ATTN_NORM: "enc.blk.{bid}.attn_norm", MODEL_TENSOR.ENC_ATTN_Q: "enc.blk.{bid}.attn_q", MODEL_TENSOR.ENC_ATTN_K: "enc.blk.{bid}.attn_k", @@ -2917,6 +2945,60 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = { MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD, MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM, ], + MODEL_ARCH.CLEF: [ + 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_OUT, + MODEL_TENSOR.ATTN_POST_NORM, + MODEL_TENSOR.ATTN_GATE, + MODEL_TENSOR.ATTN_QKV, + MODEL_TENSOR.FFN_GATE, + MODEL_TENSOR.FFN_DOWN, + MODEL_TENSOR.FFN_UP, + MODEL_TENSOR.SSM_A, + MODEL_TENSOR.SSM_CONV1D, + MODEL_TENSOR.SSM_DT, + MODEL_TENSOR.SSM_NORM, + MODEL_TENSOR.SSM_BETA, + MODEL_TENSOR.SSM_ALPHA, + MODEL_TENSOR.SSM_OUT, + # decision head + MODEL_TENSOR.TOKEN_TYPES, + MODEL_TENSOR.CLS, + MODEL_TENSOR.CLS_OUT, + MODEL_TENSOR.DEC_ATTN_NORM, + MODEL_TENSOR.DEC_ATTN_Q, + MODEL_TENSOR.DEC_ATTN_K, + MODEL_TENSOR.DEC_ATTN_V, + MODEL_TENSOR.DEC_ATTN_OUT, + MODEL_TENSOR.DEC_CROSS_ATTN_NORM, + MODEL_TENSOR.DEC_CROSS_ATTN_NORM_KV, + MODEL_TENSOR.DEC_CROSS_ATTN_Q, + MODEL_TENSOR.DEC_CROSS_ATTN_K, + MODEL_TENSOR.DEC_CROSS_ATTN_V, + MODEL_TENSOR.DEC_CROSS_ATTN_OUT, + MODEL_TENSOR.DEC_FFN_NORM, + MODEL_TENSOR.DEC_FFN_DOWN, + MODEL_TENSOR.DEC_FFN_UP, + MODEL_TENSOR.DECISION_HIDDEN_NORM, + MODEL_TENSOR.DECISION_PROJ_MEMORY, + MODEL_TENSOR.DECISION_PROJ_QUESTION, + MODEL_TENSOR.DECISION_PROJ_OPTION_QUESTION, + MODEL_TENSOR.DECISION_PROJ_GLOBAL, + MODEL_TENSOR.DECISION_PROJ_OPTION_CONTEXT, + MODEL_TENSOR.DECISION_PROJ_OPTION_LEXICAL, + MODEL_TENSOR.DECISION_OPTION_SUMMARY_NORM, + MODEL_TENSOR.DECISION_FIELD_NORM, + MODEL_TENSOR.DECISION_OPTION_NORM, + MODEL_TENSOR.DECISION_SCALES, + ], MODEL_ARCH.QWEN35MOE: [ MODEL_TENSOR.TOKEN_EMBD, MODEL_TENSOR.OUTPUT_NORM, @@ -5917,6 +5999,7 @@ class GGUFValueType(IntEnum): class DecisionType: LAYA = "laya" # head blocks + scorer on the hidden state of one marker token per option OPENJEV = "openjev" # logits of one label token per option + CLEF = "clef" # joint head over all questions, one score per option class VisionProjectorType: diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index 1dee3fe115..4f1f8e5a5d 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -1346,6 +1346,12 @@ class GGUFWriter: def add_decision_block_count(self, value: int) -> None: self.add_uint32(Keys.Decision.BLOCK_COUNT.format(arch=self.arch), value) + def add_decision_routing_block_count(self, value: int) -> None: + self.add_uint32(Keys.Decision.ROUTING_BLOCK_COUNT.format(arch=self.arch), value) + + def add_decision_head_count(self, value: int) -> None: + self.add_uint32(Keys.Decision.HEAD_COUNT.format(arch=self.arch), value) + def add_decision_max_head_tokens(self, value: int) -> None: self.add_uint32(Keys.Decision.MAX_HEAD_TOKENS.format(arch=self.arch), value) diff --git a/gguf-py/gguf/tensor_mapping.py b/gguf-py/gguf/tensor_mapping.py index c8d52f2fe2..d59e248357 100644 --- a/gguf-py/gguf/tensor_mapping.py +++ b/gguf-py/gguf/tensor_mapping.py @@ -49,6 +49,7 @@ class TensorNameMap: MODEL_TENSOR.TOKEN_TYPES: ( "embeddings.token_type_embeddings", # bert nomic-bert "type_emb", # laya + "joint_head.type_embedding", # clef ), # Normalization of token embeddings @@ -1181,22 +1182,27 @@ class TensorNameMap: MODEL_TENSOR.DEC_ATTN_NORM: ( "decoder.block.{bid}.layer.0.layer_norm", # t5 + "joint_head.layers.{bid}.norm1", # clef ), MODEL_TENSOR.DEC_ATTN_Q: ( "decoder.block.{bid}.layer.0.SelfAttention.q", # t5 + "joint_head.layers.{bid}.self_attn.q", # clef ), MODEL_TENSOR.DEC_ATTN_K: ( "decoder.block.{bid}.layer.0.SelfAttention.k", # t5 + "joint_head.layers.{bid}.self_attn.k", # clef ), MODEL_TENSOR.DEC_ATTN_V: ( "decoder.block.{bid}.layer.0.SelfAttention.v", # t5 + "joint_head.layers.{bid}.self_attn.v", # clef ), MODEL_TENSOR.DEC_ATTN_OUT: ( "decoder.block.{bid}.layer.0.SelfAttention.o", # t5 + "joint_head.layers.{bid}.self_attn.out_proj", # clef ), MODEL_TENSOR.DEC_ATTN_REL_B: ( @@ -1205,22 +1211,32 @@ class TensorNameMap: MODEL_TENSOR.DEC_CROSS_ATTN_NORM: ( "decoder.block.{bid}.layer.1.layer_norm", # t5 + "joint_head.layers.{bid}.norm2", # clef + "joint_head.evidence_layers.{bid}.query_norm", # clef ), MODEL_TENSOR.DEC_CROSS_ATTN_Q: ( "decoder.block.{bid}.layer.1.EncDecAttention.q", # t5 + "joint_head.layers.{bid}.multihead_attn.q", # clef + "joint_head.evidence_layers.{bid}.attention.q", # clef ), MODEL_TENSOR.DEC_CROSS_ATTN_K: ( "decoder.block.{bid}.layer.1.EncDecAttention.k", # t5 + "joint_head.layers.{bid}.multihead_attn.k", # clef + "joint_head.evidence_layers.{bid}.attention.k", # clef ), MODEL_TENSOR.DEC_CROSS_ATTN_V: ( "decoder.block.{bid}.layer.1.EncDecAttention.v", # t5 + "joint_head.layers.{bid}.multihead_attn.v", # clef + "joint_head.evidence_layers.{bid}.attention.v", # clef ), MODEL_TENSOR.DEC_CROSS_ATTN_OUT: ( "decoder.block.{bid}.layer.1.EncDecAttention.o", # t5 + "joint_head.layers.{bid}.multihead_attn.out_proj", # clef + "joint_head.evidence_layers.{bid}.attention.out_proj", # clef ), MODEL_TENSOR.DEC_CROSS_ATTN_REL_B: ( @@ -1229,6 +1245,8 @@ class TensorNameMap: MODEL_TENSOR.DEC_FFN_NORM: ( "decoder.block.{bid}.layer.2.layer_norm", # t5 + "joint_head.layers.{bid}.norm3", # clef + "joint_head.evidence_layers.{bid}.feedforward_norm", # clef ), MODEL_TENSOR.DEC_FFN_GATE: ( @@ -1238,16 +1256,64 @@ class TensorNameMap: MODEL_TENSOR.DEC_FFN_UP: ( "decoder.block.{bid}.layer.2.DenseReluDense.wi", # t5 "decoder.block.{bid}.layer.2.DenseReluDense.wi_1", # flan-t5 + "joint_head.layers.{bid}.linear1", # clef + "joint_head.evidence_layers.{bid}.feedforward.0", # clef ), MODEL_TENSOR.DEC_FFN_DOWN: ( "decoder.block.{bid}.layer.2.DenseReluDense.wo", # t5 + "joint_head.layers.{bid}.linear2", # clef + "joint_head.evidence_layers.{bid}.feedforward.3", # clef ), MODEL_TENSOR.DEC_OUTPUT_NORM: ( "decoder.final_layer_norm", # t5 ), + MODEL_TENSOR.DEC_CROSS_ATTN_NORM_KV: ( + "joint_head.evidence_layers.{bid}.memory_norm", # clef + ), + + MODEL_TENSOR.DECISION_HIDDEN_NORM: ( + "joint_head.hidden_norm", # clef + ), + + MODEL_TENSOR.DECISION_PROJ_MEMORY: ( + "joint_head.memory_projection", # clef + ), + + MODEL_TENSOR.DECISION_PROJ_QUESTION: ( + "joint_head.question_projection", # clef + ), + + MODEL_TENSOR.DECISION_PROJ_OPTION_QUESTION: ( + "joint_head.option_question_projection", # clef + ), + + MODEL_TENSOR.DECISION_PROJ_GLOBAL: ( + "joint_head.global_projection", # clef + ), + + MODEL_TENSOR.DECISION_PROJ_OPTION_CONTEXT: ( + "joint_head.option_context_projection", # clef + ), + + MODEL_TENSOR.DECISION_PROJ_OPTION_LEXICAL: ( + "joint_head.option_lexical_projection", # clef + ), + + MODEL_TENSOR.DECISION_OPTION_SUMMARY_NORM: ( + "joint_head.option_summary_norm", # clef + ), + + MODEL_TENSOR.DECISION_FIELD_NORM: ( + "joint_head.field_norm", # clef + ), + + MODEL_TENSOR.DECISION_OPTION_NORM: ( + "joint_head.option_norm", # clef + ), + MODEL_TENSOR.ENC_ATTN_NORM: ( "encoder.block.{bid}.layer.0.layer_norm", # t5 ), @@ -1449,11 +1515,13 @@ class TensorNameMap: "dense", # neobert "head.dense", # modern-bert "scorer.1", # laya + "joint_head.residual_scorer.0", # clef ), MODEL_TENSOR.CLS_OUT: ( "classifier.out_proj", # roberta "scorer.3", # laya + "joint_head.residual_scorer.3", # clef ), MODEL_TENSOR.CLS_NORM: ( diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index 9f205b9425..1b1dac0ff3 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -40,6 +40,7 @@ static const std::map LLM_ARCH_NAMES = { { LLM_ARCH_QWEN3VLMOE, "qwen3vlmoe" }, { LLM_ARCH_QWEN35, "qwen35" }, { LLM_ARCH_QWEN35MOE, "qwen35moe" }, + { LLM_ARCH_CLEF, "clef" }, { LLM_ARCH_QWEN4EXP, "qwen4exp" }, { LLM_ARCH_PHI2, "phi2" }, { LLM_ARCH_PHI3, "phi3" }, @@ -366,7 +367,9 @@ static const std::map LLM_KV_NAMES = { { LLM_KV_CLASSIFIER_OUTPUT_LABELS, "%s.classifier.output_labels" }, { LLM_KV_CLASSIFIER_POOLING_TYPE, "%s.classifier.pooling_type" }, - { LLM_KV_DECISION_BLOCK_COUNT, "%s.decision.block_count" }, + { LLM_KV_DECISION_BLOCK_COUNT, "%s.decision.block_count" }, + { LLM_KV_DECISION_ROUTING_BLOCK_COUNT, "%s.decision.routing_block_count" }, + { LLM_KV_DECISION_HEAD_COUNT, "%s.decision.head_count" }, { LLM_KV_TARGET_LAYERS, "%s.target_layers" }, { LLM_KV_TARGET_HIDDEN_SIZE, "%s.target_hidden_size" }, @@ -599,6 +602,18 @@ static const std::map LLM_TENSOR_NAMES = { { LLM_TENSOR_ATTN_SUB_NORM, "blk.%d.attn_sub_norm" }, { LLM_TENSOR_FFN_SUB_NORM, "blk.%d.ffn_sub_norm" }, { LLM_TENSOR_DEC_OUTPUT_NORM, "dec.output_norm" }, + { LLM_TENSOR_DEC_CROSS_ATTN_NORM_KV, "dec.blk.%d.cross_attn_norm_kv" }, + { LLM_TENSOR_DECISION_HIDDEN_NORM, "decision.hidden_norm" }, + { LLM_TENSOR_DECISION_PROJ_MEMORY, "decision.proj_memory" }, + { LLM_TENSOR_DECISION_PROJ_QUESTION, "decision.proj_question" }, + { LLM_TENSOR_DECISION_PROJ_OPTION_QUESTION, "decision.proj_option_question" }, + { LLM_TENSOR_DECISION_PROJ_GLOBAL, "decision.proj_global" }, + { LLM_TENSOR_DECISION_PROJ_OPTION_CONTEXT, "decision.proj_option_context" }, + { LLM_TENSOR_DECISION_PROJ_OPTION_LEXICAL, "decision.proj_option_lexical" }, + { LLM_TENSOR_DECISION_OPTION_SUMMARY_NORM, "decision.option_summary_norm" }, + { LLM_TENSOR_DECISION_FIELD_NORM, "decision.field_norm" }, + { LLM_TENSOR_DECISION_OPTION_NORM, "decision.option_norm" }, + { LLM_TENSOR_DECISION_SCALES, "decision.scales" }, { LLM_TENSOR_DEC_ATTN_NORM, "dec.blk.%d.attn_norm" }, { LLM_TENSOR_DEC_ATTN_Q, "dec.blk.%d.attn_q" }, { LLM_TENSOR_DEC_ATTN_K, "dec.blk.%d.attn_k" }, @@ -744,6 +759,18 @@ static const std::map LLM_TENSOR_INFOS = { {LLM_TENSOR_OUTPUT_NORM, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, {LLM_TENSOR_OUTPUT_NORM_LFM2, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, {LLM_TENSOR_DEC_OUTPUT_NORM, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_DEC_CROSS_ATTN_NORM_KV, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL}}, + {LLM_TENSOR_DECISION_HIDDEN_NORM, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_DECISION_PROJ_MEMORY, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_DECISION_PROJ_QUESTION, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_DECISION_PROJ_OPTION_QUESTION, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_DECISION_PROJ_GLOBAL, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_DECISION_PROJ_OPTION_CONTEXT, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_DECISION_PROJ_OPTION_LEXICAL, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL_MAT}}, + {LLM_TENSOR_DECISION_OPTION_SUMMARY_NORM, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_DECISION_FIELD_NORM, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_DECISION_OPTION_NORM, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, + {LLM_TENSOR_DECISION_SCALES, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, {LLM_TENSOR_ENC_OUTPUT_NORM, {LLM_TENSOR_LAYER_OUTPUT, GGML_OP_MUL}}, {LLM_TENSOR_ROPE_FREQS, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_ROPE}}, {LLM_TENSOR_ROPE_FACTORS_LONG, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_ROPE}}, diff --git a/src/llama-arch.h b/src/llama-arch.h index ca390fcb3d..4106027675 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -45,6 +45,7 @@ enum llm_arch { LLM_ARCH_QWEN3VLMOE, LLM_ARCH_QWEN35, LLM_ARCH_QWEN35MOE, + LLM_ARCH_CLEF, LLM_ARCH_QWEN4EXP, LLM_ARCH_PHI2, LLM_ARCH_PHI3, @@ -413,6 +414,8 @@ enum llm_kv { LLM_KV_CLASSIFIER_POOLING_TYPE, LLM_KV_DECISION_BLOCK_COUNT, + LLM_KV_DECISION_ROUTING_BLOCK_COUNT, + LLM_KV_DECISION_HEAD_COUNT, LLM_KV_TARGET_LAYERS, LLM_KV_TARGET_HIDDEN_SIZE, @@ -646,6 +649,18 @@ enum llm_tensor { LLM_TENSOR_DEC_FFN_DOWN, LLM_TENSOR_DEC_FFN_UP, LLM_TENSOR_DEC_OUTPUT_NORM, + LLM_TENSOR_DEC_CROSS_ATTN_NORM_KV, + LLM_TENSOR_DECISION_HIDDEN_NORM, + LLM_TENSOR_DECISION_PROJ_MEMORY, + LLM_TENSOR_DECISION_PROJ_QUESTION, + LLM_TENSOR_DECISION_PROJ_OPTION_QUESTION, + LLM_TENSOR_DECISION_PROJ_GLOBAL, + LLM_TENSOR_DECISION_PROJ_OPTION_CONTEXT, + LLM_TENSOR_DECISION_PROJ_OPTION_LEXICAL, + LLM_TENSOR_DECISION_OPTION_SUMMARY_NORM, + LLM_TENSOR_DECISION_FIELD_NORM, + LLM_TENSOR_DECISION_OPTION_NORM, + LLM_TENSOR_DECISION_SCALES, LLM_TENSOR_ENC_ATTN_NORM, LLM_TENSOR_ENC_ATTN_Q, LLM_TENSOR_ENC_ATTN_K, diff --git a/src/llama-batch.cpp b/src/llama-batch.cpp index 7a17b3a4e9..651932d54b 100644 --- a/src/llama-batch.cpp +++ b/src/llama-batch.cpp @@ -168,6 +168,16 @@ bool llama_batch_allocr::init( } } + for (int32_t i = 0; i < n_tok; ++i) { + if (batch_inp.tokens[i].decision_order != 0) { + decision_order.resize(n_tok, 0); + break; + } + } + for (size_t i = 0; i < decision_order.size(); ++i) { + decision_order[i] = batch_inp.tokens[i].decision_order; + } + // // set up the internal llama_batch to point to our owned arrays // @@ -765,6 +775,7 @@ void llama_batch_allocr::clear() { seq_id .clear(); seq_id_unq .clear(); output .clear(); + decision_order.clear(); for (auto & cur : seq_pos) { cur.clear(); @@ -826,6 +837,10 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector & idxs, u udata->n_seq_id[i] = batch.n_seq_id[idxs[i]]; udata->output[i] = batch.logits[idxs[i]]; + if (!decision_order.empty()) { + udata->decision_order.push_back(decision_order[idxs[i]]); + } + for (int s = 0; s < udata->n_seq_id[i]; ++s) { const llama_seq_id seq_id = batch.seq_id[idxs[i]][s]; @@ -870,6 +885,10 @@ llama_ubatch llama_batch_allocr::ubatch_add(const std::vector & idxs, u /*.data =*/ std::move(udata), }; + if (!res.data->decision_order.empty()) { + res.decision_order = res.data->decision_order.data(); + } + if (debug > 0) { LLAMA_LOG_DEBUG("%s: added ubatch to split:\n", __func__); @@ -1175,6 +1194,14 @@ bool llama_batch_ext::set_output(int32_t idx, bool output_last) { return true; } +bool llama_batch_ext::set_decision_order(int32_t idx, int32_t order) { + if (idx < 0 || idx >= (int32_t) tokens.size()) { + return false; + } + tokens[idx].decision_order = order; + return true; +} + // llama_batch_ext C API llama_batch_ext * llama_batch_ext_init(llama_context * ctx) { @@ -1243,6 +1270,10 @@ bool llama_batch_ext_set_output_logits(llama_batch_ext * batch, int32_t idx, boo return batch->set_output(idx, value); } +bool llama_batch_ext_set_decision_order(llama_batch_ext * batch, int32_t idx, int32_t order) { + return batch->set_decision_order(idx, order); +} + // llama_batch_compat void llama_batch_compat::init(llama_batch_ext & dst, const llama_batch & batch_inp, size_t n_embd_row) { diff --git a/src/llama-batch.h b/src/llama-batch.h index 18e8144c80..74d0aa3c7f 100644 --- a/src/llama-batch.h +++ b/src/llama-batch.h @@ -63,12 +63,16 @@ struct llama_ubatch { std::vector seq_idx; std::vector output; std::vector batch_idxs; // original batch index for each token + std::vector decision_order; std::vector seq_id_data; }; // the llama_ubatch pointers above point to this data if set. otherwise - point to external non-owning data std::shared_ptr data; + + // [n_tokens], see llama_batch_ext_set_decision_order(), NULL if no entry has one + int32_t * decision_order = nullptr; }; struct llama_hparams; @@ -96,6 +100,7 @@ struct llama_batch_ext { bool has_embd = false; // whether embd_off is set size_t embd_off = 0; // index offset in the embd array bool output = false; // TODO: have dedicated output flags + int32_t decision_order = 0; // see llama_batch_ext_set_decision_order() std::unordered_set seq_ids; std::array pos = {0, 0, 0, 0}; }; @@ -125,6 +130,7 @@ struct llama_batch_ext { bool set_token_embd(int32_t idx, llama_embd embd_in); bool set_token_pos(int32_t idx, const llama_pos * pos_in); bool set_output(int32_t idx, bool output_last); + bool set_decision_order(int32_t idx, int32_t order); }; // a helper for sanitizing, fulfilling and splitting a batch @@ -202,6 +208,7 @@ private: std::vector seq_id_unq; std::vector seq_idx; std::vector output; + std::vector decision_order; // empty if no entry has one using pos_set_t = std::set; using seq_cpl_t = std::vector; diff --git a/src/llama-context.cpp b/src/llama-context.cpp index 4f04db0ee6..96b5464e67 100644 --- a/src/llama-context.cpp +++ b/src/llama-context.cpp @@ -2406,6 +2406,7 @@ uint32_t llama_context::graph_max_nodes(uint32_t n_tokens) const { model.arch == LLM_ARCH_BAILINGMOE3 || model.arch == LLM_ARCH_QWEN35 || model.arch == LLM_ARCH_QWEN35MOE || + model.arch == LLM_ARCH_CLEF || model.arch == LLM_ARCH_QWEN4EXP || model.arch == LLM_ARCH_DEEPSEEK4 || (model.arch == LLM_ARCH_DFLASH && model.hparams.dsv4_hc_mult > 0) || diff --git a/src/llama-ext.h b/src/llama-ext.h index 92a759b7a0..3db728645d 100644 --- a/src/llama-ext.h +++ b/src/llama-ext.h @@ -100,6 +100,19 @@ LLAMA_API void llama_set_embeddings_nextn(struct llama_context * ctx, bool value // chain multiple trained NextN heads. Default 0 (first head). LLAMA_API void llama_set_nextn_layer_offset(struct llama_context * ctx, int32_t offset); +// Marks the entries that a joint decision head (clef) reads, the default is 0 +// A run of entries with the same value is one span, spans must be separated by entries with value 0 +// An option belongs to the last question before it +enum llama_decision_order { + LLAMA_DECISION_ORDER_NONE = 0, // not read by the head + LLAMA_DECISION_ORDER_QUESTION_NOUL = 1, // text of a question + LLAMA_DECISION_ORDER_QUESTION_CHOICE = 2, + LLAMA_DECISION_ORDER_QUESTION_SCORE = 3, + LLAMA_DECISION_ORDER_OPTION = 4, // text of an option +}; +// The embeddings output has one value per entry: row i is the score of option i +LLAMA_API bool llama_batch_ext_set_decision_order(struct llama_batch_ext * batch, int32_t idx, int32_t order); + // mirrors: // LLAMA_API float * llama_get_embeddings(struct llama_context * ctx); LLAMA_API float * llama_get_embeddings_nextn(struct llama_context * ctx); diff --git a/src/llama-model-saver.cpp b/src/llama-model-saver.cpp index eeecadac03..e4fd45ec4f 100644 --- a/src/llama-model-saver.cpp +++ b/src/llama-model-saver.cpp @@ -20,6 +20,7 @@ bool llama_model_saver_supports_arch(llm_arch arch) { case LLM_ARCH_T5: case LLM_ARCH_APERTUS: case LLM_ARCH_STEP35: + case LLM_ARCH_CLEF: // the head tensors are not saved return false; default: return true; diff --git a/src/llama-model.cpp b/src/llama-model.cpp index 404509888e..11d75fc829 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -326,6 +326,8 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params return new llama_model_qwen35(params); case LLM_ARCH_QWEN35MOE: return new llama_model_qwen35moe(params); + case LLM_ARCH_CLEF: + return new llama_model_clef(params); case LLM_ARCH_QWEN4EXP: return new llama_model_qwen4exp(params); case LLM_ARCH_MISTRAL3: @@ -607,7 +609,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str auto get_split_segments = [&](int axis, uint32_t il) -> std::vector> { // TODO: clarify why this is necessary specifically for these models // TODO: deduplicate condition [TAG_SPLIT_QGATE_QWEN] - if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE || + if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE || ud->model->arch == LLM_ARCH_CLEF || ud->model->arch == LLM_ARCH_QWEN4EXP) { // fused full attention layers with Q gate tensors that need n_embd doubled: @@ -762,7 +764,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str GGML_ASSERT(segments.size() == 1); // some models have Q gate tensors, for those cases the granularity needs to be doubled: // TODO: deduplicate condition [TAG_SPLIT_QGATE_QWEN] - if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE || + if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE || ud->model->arch == LLM_ARCH_CLEF || ud->model->arch == LLM_ARCH_QWEN4EXP) { return {std::lcm(2*n_embd_q, blck_size_perf)}; } @@ -795,7 +797,7 @@ struct ggml_backend_meta_split_state llama_meta_device_get_split_state(const str if (std::regex_match(tensor_name, pattern_qkv_weight) || std::regex_match(tensor_name, pattern_qkv_bias)) { // fused full attention layers need Q gate tensors handled like above: // TODO: deduplicate condition [TAG_SPLIT_QGATE_QWEN] - if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE || + if (ud->model->arch == LLM_ARCH_QWEN3NEXT || ud->model->arch == LLM_ARCH_QWEN35 || ud->model->arch == LLM_ARCH_QWEN35MOE || ud->model->arch == LLM_ARCH_CLEF || ud->model->arch == LLM_ARCH_QWEN4EXP) { return {std::lcm(2*n_embd_q, blck_size_perf), granularity_kv}; } @@ -1961,7 +1963,11 @@ bool llama_model_base::load_tensors(llama_model_loader & ml) { } ggml_tensor * llama_model_base::create_tensor(llama_model_loader & ml, const LLM_TN_IMPL & tn, const std::initializer_list & ne, int flags) { - const buft_list_t * buft_list_layer = tn.bid == -1 ? nullptr : pimpl->dev_layer.at(tn.bid).buft_list; + const buft_list_t * buft_list_layer = nullptr; + if (tn.bid != -1) { + // blocks that are not model layers (e.g. the blocks of a head) go with the output + buft_list_layer = (size_t) tn.bid < pimpl->dev_layer.size() ? pimpl->dev_layer.at(tn.bid).buft_list : pimpl->dev_output.buft_list; + } return ml.create_tensor( hparams, &pimpl->cpu_buft_list, pimpl->dev_input.buft_list, pimpl->dev_output.buft_list, buft_list_layer, tn, ne, flags); @@ -2369,6 +2375,7 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, case LLM_ARCH_LLADA: case LLM_ARCH_LLADA_MOE: case LLM_ARCH_RND1: + case LLM_ARCH_CLEF: { res = nullptr; } break; @@ -3200,6 +3207,7 @@ llama_rope_type llama_model_rope_type(const llama_model * model) { case LLM_ARCH_QWEN3VLMOE: case LLM_ARCH_QWEN35: case LLM_ARCH_QWEN35MOE: + case LLM_ARCH_CLEF: case LLM_ARCH_QWEN4EXP: case LLM_ARCH_QWEN3TTS: return LLAMA_ROPE_TYPE_IMROPE; diff --git a/src/models/clef.cpp b/src/models/clef.cpp new file mode 100644 index 0000000000..7bb846db16 --- /dev/null +++ b/src/models/clef.cpp @@ -0,0 +1,612 @@ +#include "models.h" + +#include + +// torch.nn.LayerNorm default +static const float CLEF_HEAD_NORM_EPS = 1e-5f; + +void llama_model_clef::load_arch_hparams(llama_model_loader & ml) { + llama_model_qwen35::load_arch_hparams(ml); + + ml.get_key(LLM_KV_DECISION_ROUTING_BLOCK_COUNT, n_layer_routing); + ml.get_key(LLM_KV_DECISION_BLOCK_COUNT, n_layer_joint); + ml.get_key(LLM_KV_DECISION_HEAD_COUNT, n_head_decision); + + if (n_head_decision == 0) { + throw std::runtime_error("invalid number of heads in the decision head"); + } + + hparams.f_norm_eps = CLEF_HEAD_NORM_EPS; + + // the output is one score per token, see llama_batch_ext_set_decision_order() + hparams.n_embd_out_impl = 1; +} + +void llama_model_clef::load_arch_tensors(llama_model_loader & ml) { + llama_model_qwen35::load_arch_tensors(ml); + + LLAMA_LOAD_LOCALS; + + const auto * w_memory = ml.get_weight(tn(LLM_TENSOR_DECISION_PROJ_MEMORY, "weight").str().c_str()); + const auto * w_ffn = ml.get_weight(tn(LLM_TENSOR_DEC_FFN_UP, "weight", 0).str().c_str()); + if (w_memory == nullptr || w_ffn == nullptr) { + throw std::runtime_error("the decision head is missing"); + } + const int64_t n_embd_h = w_memory->tensor->ne[1]; + const int64_t n_ff_h = w_ffn->tensor->ne[1]; + + if (n_embd_h % n_head_decision != 0) { + throw std::runtime_error("invalid width of the decision head"); + } + + auto load_norm = [&](norm & n, llm_tensor type, int64_t size, int il = -1) { + n.w = il < 0 ? create_tensor(tn(type, "weight"), {size}, 0) : create_tensor(tn(type, "weight", il), {size}, 0); + n.b = il < 0 ? create_tensor(tn(type, "bias"), {size}, 0) : create_tensor(tn(type, "bias", il), {size}, 0); + }; + + auto load_attn = [&](attn & a, llm_tensor q, llm_tensor k, llm_tensor v, llm_tensor o, int il) { + a.wq = create_tensor(tn(q, "weight", il), {n_embd_h, n_embd_h}, 0); + a.bq = create_tensor(tn(q, "bias", il), {n_embd_h}, 0); + a.wk = create_tensor(tn(k, "weight", il), {n_embd_h, n_embd_h}, 0); + a.bk = create_tensor(tn(k, "bias", il), {n_embd_h}, 0); + a.wv = create_tensor(tn(v, "weight", il), {n_embd_h, n_embd_h}, 0); + a.bv = create_tensor(tn(v, "bias", il), {n_embd_h}, 0); + a.wo = create_tensor(tn(o, "weight", il), {n_embd_h, n_embd_h}, 0); + a.bo = create_tensor(tn(o, "bias", il), {n_embd_h}, 0); + }; + + head_layers.resize(n_layer_routing + n_layer_joint); + for (int il = 0; il < (int) head_layers.size(); ++il) { + auto & layer = head_layers[il]; + + if (il < (int) n_layer_routing) { + load_norm(layer.cross_norm_kv, LLM_TENSOR_DEC_CROSS_ATTN_NORM_KV, n_embd_h, il); + } else { + load_norm(layer.self_norm, LLM_TENSOR_DEC_ATTN_NORM, n_embd_h, il); + load_attn(layer.self_attn, LLM_TENSOR_DEC_ATTN_Q, LLM_TENSOR_DEC_ATTN_K, LLM_TENSOR_DEC_ATTN_V, LLM_TENSOR_DEC_ATTN_OUT, il); + } + + load_norm(layer.cross_norm, LLM_TENSOR_DEC_CROSS_ATTN_NORM, n_embd_h, il); + load_attn(layer.cross_attn, LLM_TENSOR_DEC_CROSS_ATTN_Q, LLM_TENSOR_DEC_CROSS_ATTN_K, LLM_TENSOR_DEC_CROSS_ATTN_V, LLM_TENSOR_DEC_CROSS_ATTN_OUT, il); + + load_norm(layer.ffn_norm, LLM_TENSOR_DEC_FFN_NORM, n_embd_h, il); + layer.ffn_up = create_tensor(tn(LLM_TENSOR_DEC_FFN_UP, "weight", il), {n_embd_h, n_ff_h}, 0); + layer.ffn_up_b = create_tensor(tn(LLM_TENSOR_DEC_FFN_UP, "bias", il), {n_ff_h}, 0); + layer.ffn_down = create_tensor(tn(LLM_TENSOR_DEC_FFN_DOWN, "weight", il), {n_ff_h, n_embd_h}, 0); + layer.ffn_down_b = create_tensor(tn(LLM_TENSOR_DEC_FFN_DOWN, "bias", il), {n_embd_h}, 0); + } + + load_norm(hidden_norm, LLM_TENSOR_DECISION_HIDDEN_NORM, n_embd); + load_norm(option_summary_norm, LLM_TENSOR_DECISION_OPTION_SUMMARY_NORM, n_embd_h); + load_norm(field_norm, LLM_TENSOR_DECISION_FIELD_NORM, n_embd_h); + load_norm(option_norm, LLM_TENSOR_DECISION_OPTION_NORM, n_embd_h); + + proj_memory = create_tensor(tn(LLM_TENSOR_DECISION_PROJ_MEMORY, "weight"), {n_embd, n_embd_h}, 0); + proj_question = create_tensor(tn(LLM_TENSOR_DECISION_PROJ_QUESTION, "weight"), {n_embd, n_embd_h}, 0); + proj_option_question = create_tensor(tn(LLM_TENSOR_DECISION_PROJ_OPTION_QUESTION, "weight"), {n_embd, n_embd_h}, 0); + proj_global = create_tensor(tn(LLM_TENSOR_DECISION_PROJ_GLOBAL, "weight"), {n_embd, n_embd_h}, 0); + proj_option_context = create_tensor(tn(LLM_TENSOR_DECISION_PROJ_OPTION_CONTEXT, "weight"), {n_embd, n_embd_h}, 0); + proj_option_lexical = create_tensor(tn(LLM_TENSOR_DECISION_PROJ_OPTION_LEXICAL, "weight"), {n_embd, n_embd_h}, 0); + + scales = create_tensor(tn(LLM_TENSOR_DECISION_SCALES), {3}, 0); + type_embd = create_tensor(tn(LLM_TENSOR_TOKEN_TYPES, "weight"), {n_embd_h, 3}, 0); + cls = create_tensor(tn(LLM_TENSOR_CLS, "weight"), {4 * n_embd_h, n_embd_h}, 0); + cls_b = create_tensor(tn(LLM_TENSOR_CLS, "bias"), {n_embd_h}, 0); + cls_out = create_tensor(tn(LLM_TENSOR_CLS_OUT, "weight"), {n_embd_h, 1}, 0); + cls_out_b = create_tensor(tn(LLM_TENSOR_CLS_OUT, "bias"), {1}, 0); +} + +std::unique_ptr llama_model_clef::build_arch_graph(const llm_graph_params & params) const { + return std::make_unique(*this, params); +} + +// spans read by the head, [start, end) in ubatch token indices +struct clef_spans { + struct question { + int32_t type; // noul, choice, score + int32_t start; + int32_t end; + }; + struct option { + int32_t question; + int32_t start; + int32_t end; + }; + std::vector questions; + std::vector