Adapt Muse-Glimmer loading (#2314)

This commit is contained in:
Kawrakow 2026-08-14 07:55:11 +02:00 committed by GitHub
parent 981e5ea0d7
commit 43afea46c2
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 33 additions and 4 deletions

View File

@ -2002,11 +2002,14 @@ void llm_load_hparams(
hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
ml.get_key(LLM_KV_ROPE_FREQ_BASE_SWA, hparams.rope_freq_base_train_swa, false);
if (!ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.swa_layers, hparams.n_layer, false)) {
uint32_t swa_period = 4;
ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period);
if (uint32_t swa_period; ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false) && swa_period > 0) {
for (int il = 0; il < hparams.n_layer; ++il) {
hparams.swa_layers[il] = (swa_period == 0 || (il % swa_period < swa_period - 1));
hparams.swa_layers[il] = (il % swa_period < swa_period - 1);
}
} else if (!ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.swa_layers, hparams.n_layer)) {
LLAMA_LOG_WARN("================================ No attention.sliding_window_pattern key found! Assuming a period of 4\n");
for (int il = 0; il < hparams.n_layer; ++il) {
hparams.swa_layers[il] = (il % 4 < 3);
}
}

View File

@ -846,6 +846,30 @@ bool llama_model_loader::get_key_or_arr(const enum llm_kv kid, T & result, uint3
return get_key_or_arr(llm_kv(kid), result, n, required);
}
bool llama_model_loader::get_key_or_arr(enum llm_kv kid, uint32_t & result, bool required) {
const std::string key = llm_kv(kid);
const int id = gguf_find_key(meta, key.c_str());
if (id < 0) {
if (required) {
throw std::runtime_error(format("key not found in model: %s", key.c_str()));
}
return false;
}
// throw and error if type is an array
if (gguf_get_kv_type(meta, id) == GGUF_TYPE_ARRAY) {
if (required) {
throw std::runtime_error(format("expected scalar, found array for key: %s", key.c_str()));
}
return false;
}
return get_key(key, result, required);
}
const char * llama_model_loader::get_tensor_name(int i) const {
return weights.at(i).tensor->name;
}

View File

@ -127,6 +127,8 @@ struct llama_model_loader {
template<typename T>
bool get_key_or_arr(const enum llm_kv kid, T & result, uint32_t n, const bool required = true);
bool get_key_or_arr(enum llm_kv kid, uint32_t & result, bool required = true);
const std::string& get_arch_name() const { return arch_name; }
enum llm_arch get_arch() const { return llm_kv.arch; }