add more rnn models supported

This commit is contained in:
Concedo
2025-09-24 18:14:59 +08:00
parent c7a1eec4e4
commit 5d89a48a50
3 changed files with 8 additions and 13 deletions
+4 -4
View File
@@ -484,8 +484,8 @@ void ContextRewind(std::vector<int> &embd, std::vector<int> &current_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)
+2 -6
View File
@@ -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))
{
+2 -3
View File
@@ -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