QC67_cosmo / architecture /cosmos-arch.patch
phera-ra's picture
Add LLM_ARCH_COSMOS for llama.cpp; correct the section-3 retraction
220ff1c verified
Raw
History Blame Contribute Delete
4.86 kB
diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp
index 72968607d..8a1927243 100644
--- a/src/llama-arch.cpp
+++ b/src/llama-arch.cpp
@@ -141,6 +141,7 @@ static const std::map<llm_arch, const char *> LLM_ARCH_NAMES = {
{ LLM_ARCH_KIMI_LINEAR, "kimi-linear" },
{ LLM_ARCH_TALKIE, "talkie" },
{ LLM_ARCH_MELLUM, "mellum" },
+ { LLM_ARCH_COSMOS, "cosmos" },
{ LLM_ARCH_UNKNOWN, "(unknown)" },
};
@@ -397,6 +398,7 @@ static const std::map<llm_tensor, const char *> LLM_TENSOR_NAMES = {
{ LLM_TENSOR_ATTN_Q_NORM, "blk.%d.attn_q_norm" },
{ LLM_TENSOR_ATTN_K_NORM, "blk.%d.attn_k_norm" },
{ LLM_TENSOR_ATTN_GATE, "blk.%d.attn_gate" },
+ { LLM_TENSOR_ATTN_54, "blk.%d.attn_54" },
{ LLM_TENSOR_FFN_POST_NORM, "blk.%d.post_ffw_norm" },
{ LLM_TENSOR_FFN_POST_NORM_1, "blk.%d.post_ffw_norm_1" },
{ LLM_TENSOR_FFN_POST_NORM_2, "blk.%d.post_ffw_norm_2" },
@@ -640,6 +642,7 @@ static const std::map<llm_tensor, llm_tensor_info> LLM_TENSOR_INFOS = {
{LLM_TENSOR_ATTN_QKV, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
{LLM_TENSOR_ATTN_OUT, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
{LLM_TENSOR_ATTN_GATE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
+ {LLM_TENSOR_ATTN_54, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
{LLM_TENSOR_FFN_GATE, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
{LLM_TENSOR_FFN_DOWN, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
{LLM_TENSOR_FFN_UP, {LLM_TENSOR_LAYER_REPEATING, GGML_OP_MUL_MAT}},
diff --git a/src/llama-arch.h b/src/llama-arch.h
index b74d53af4..036f4e30f 100644
--- a/src/llama-arch.h
+++ b/src/llama-arch.h
@@ -146,6 +146,7 @@ enum llm_arch {
LLM_ARCH_MELLUM,
LLM_ARCH_EAGLE3,
LLM_ARCH_DFLASH,
+ LLM_ARCH_COSMOS,
LLM_ARCH_UNKNOWN,
};
@@ -400,6 +401,7 @@ enum llm_tensor {
LLM_TENSOR_ATTN_ROT_EMBD,
LLM_TENSOR_ATTN_SINKS,
LLM_TENSOR_ATTN_GATE,
+ LLM_TENSOR_ATTN_54,
LLM_TENSOR_FFN_GATE_INP,
LLM_TENSOR_FFN_GATE_INP_SHEXP,
LLM_TENSOR_FFN_NORM,
diff --git a/src/llama-model.cpp b/src/llama-model.cpp
index 4c10e4126..1cb328fa4 100644
--- a/src/llama-model.cpp
+++ b/src/llama-model.cpp
@@ -298,6 +298,8 @@ static llama_model * llama_model_mapping(llm_arch arch, const llama_model_params
return new llama_model_eagle3(params);
case LLM_ARCH_DFLASH:
return new llama_model_dflash(params);
+ case LLM_ARCH_COSMOS:
+ return new llama_model_cosmos(params);
case LLM_ARCH_MIMO2:
return new llama_model_mimo2(params);
case LLM_ARCH_KIMI_LINEAR:
@@ -2444,6 +2446,7 @@ llama_rope_type llama_model_rope_type(const llama_model * model) {
case LLM_ARCH_NEMOTRON_H:
case LLM_ARCH_NEMOTRON_H_MOE:
case LLM_ARCH_KIMI_LINEAR:
+ case LLM_ARCH_COSMOS:
return LLAMA_ROPE_TYPE_NONE;
// use what we call a normal RoPE, operating on pairs of consecutive head values
diff --git a/src/llama-model.h b/src/llama-model.h
index 45b054ced..a5c54f058 100644
--- a/src/llama-model.h
+++ b/src/llama-model.h
@@ -466,6 +466,10 @@ struct llama_layer {
// openai-moe
struct ggml_tensor * attn_sinks = nullptr;
+ // cosmos: 54D mixture-of-states Hebbian attention
+ struct ggml_tensor * attn_54 = nullptr;
+ struct ggml_tensor * attn_gate = nullptr;
+
// DeepSeek-V4
struct ggml_tensor * attn_kv_norm = nullptr;
struct ggml_tensor * hc_attn_fn = nullptr;
diff --git a/src/models/models.h b/src/models/models.h
index a86ae05aa..91c7a512d 100644
--- a/src/models/models.h
+++ b/src/models/models.h
@@ -437,6 +437,18 @@ struct llama_model_qwen : public llama_model_base {
};
+struct llama_model_cosmos : public llama_model_base {
+ llama_model_cosmos(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);
+ };
+
+ std::unique_ptr<llm_graph_context> build_arch_graph(const llm_graph_params & params) const override;
+};
+
struct llama_model_qwen2 : public llama_model_base {
llama_model_qwen2(const struct llama_model_params & params) : llama_model_base(params) {}
void load_arch_hparams(llama_model_loader & ml) override;