From 5d89a48a50be575a8762c0d75514298870d73b4c Mon Sep 17 00:00:00 2001 From: Concedo <39025047+LostRuins@users.noreply.github.com> Date: Wed, 24 Sep 2025 18:14:59 +0800 Subject: [PATCH] add more rnn models supported --- gpttype_adapter.cpp | 8 ++++---- model_adapter.cpp | 8 ++------ model_adapter.h | 5 ++--- 3 files changed, 8 insertions(+), 13 deletions(-) diff --git a/gpttype_adapter.cpp b/gpttype_adapter.cpp index 8a79d1895..d8a95267e 100644 --- a/gpttype_adapter.cpp +++ b/gpttype_adapter.cpp @@ -484,8 +484,8 @@ void ContextRewind(std::vector &embd, std::vector ¤t_context_tok printf("\nWARNING: Don't use context rewind when in batch processing phase!\n"); return; } - bool is_recurrent = (file_format == FileFormat::GGUF_GENERIC && (file_format_meta.model_architecture==GGUFArch::ARCH_MAMBA - || file_format_meta.model_architecture==GGUFArch::ARCH_RWKV || file_format_meta.model_architecture==GGUFArch::ARCH_JAMBA)); + bool is_recurrent = (file_format == FileFormat::GGUF_GENERIC && (file_format_meta.model_architecture==GGUFArch::ARCH_MAMBALIKE + || file_format_meta.model_architecture==GGUFArch::ARCH_RWKV)); if(file_format == FileFormat::RWKV_1 || file_format==FileFormat::RWKV_2 || is_recurrent) { printf("\nWARNING: RNN models do not support context rewind!\n"); @@ -3747,8 +3747,8 @@ generation_outputs gpttype_generate(const generation_inputs inputs) printf("%s\n", RemoveBell(outstr).c_str()); } - bool is_recurrent = (file_format == FileFormat::GGUF_GENERIC && (file_format_meta.model_architecture==GGUFArch::ARCH_MAMBA - || file_format_meta.model_architecture==GGUFArch::ARCH_RWKV || file_format_meta.model_architecture==GGUFArch::ARCH_JAMBA)); + bool is_recurrent = (file_format == FileFormat::GGUF_GENERIC && (file_format_meta.model_architecture==GGUFArch::ARCH_MAMBALIKE + || file_format_meta.model_architecture==GGUFArch::ARCH_RWKV)); bool blank_prompt = (addedmemory=="" && kcpp_data->prompt==""); if (file_format == FileFormat::RWKV_1 || file_format==FileFormat::RWKV_2 || is_recurrent) diff --git a/model_adapter.cpp b/model_adapter.cpp index 49e2c8482..81022bc5c 100644 --- a/model_adapter.cpp +++ b/model_adapter.cpp @@ -367,13 +367,9 @@ std::string gguf_get_model_arch(const std::string & gguf_filename) { fileformatmeta->model_architecture = GGUFArch::ARCH_FALCON; } - else if(modelarch=="mamba") + else if(modelarch=="mamba" || modelarch=="mamba2" || modelarch=="nemotron_h" || modelarch=="jamba") //lazy approach, put all RNN models { - fileformatmeta->model_architecture = GGUFArch::ARCH_MAMBA; - } - else if(modelarch=="jamba") - { - fileformatmeta->model_architecture = GGUFArch::ARCH_JAMBA; + fileformatmeta->model_architecture = GGUFArch::ARCH_MAMBALIKE; } else if(modelarch=="llama" && freq_base_train==10000.0f && (n_tensors==435 || n_tensors==611)) { diff --git a/model_adapter.h b/model_adapter.h index adfee4942..bd3c9b81c 100644 --- a/model_adapter.h +++ b/model_adapter.h @@ -55,7 +55,7 @@ enum GGUFArch ARCH_DEFAULT = 0, //used for llama3 and other generic gguf ARCH_FALCON = 1, ARCH_PHI = 2, - ARCH_MAMBA = 3, + ARCH_MAMBALIKE = 3, ARCH_SOLAR = 4, ARCH_QWEN2 = 5, ARCH_RWKV = 6, @@ -63,8 +63,7 @@ enum GGUFArch ARCH_GEMMA3 = 8, ARCH_GLM4 = 9, ARCH_GEMMA3N = 10, - ARCH_JAMBA = 11, - ARCH_GPTOSS = 12, + ARCH_GPTOSS = 11, }; struct FileFormatExtraMeta