Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 8 additions & 16 deletions src/llama-model-saver.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -15,26 +15,11 @@

bool llama_model_saver_supports_arch(llm_arch arch) {
switch (arch) {
case LLM_ARCH_PLAMO3:
case LLM_ARCH_GEMMA3:
case LLM_ARCH_GEMMA3N:
case LLM_ARCH_COHERE2:
case LLM_ARCH_COHERE2MOE:
case LLM_ARCH_OLMO2:
case LLM_ARCH_BITNET:
case LLM_ARCH_T5:
case LLM_ARCH_EXAONE_MOE:
case LLM_ARCH_AFMOE:
case LLM_ARCH_APERTUS:
case LLM_ARCH_MIMO2:
case LLM_ARCH_STEP35:
case LLM_ARCH_SPARK2_5:
case LLM_ARCH_MUSE_GLIMMER:
case LLM_ARCH_MELLUM:
case LLM_ARCH_LAGUNA:
case LLM_ARCH_GRANITE_SWA:
case LLM_ARCH_DOTS3NOTE: // TODO: need to handle SWA pattern and MLA+SWA config
case LLM_ARCH_MAPLE:
return false;
default:
return true;
Expand Down Expand Up @@ -290,7 +275,11 @@ void llama_model_saver::add_kv_from_model() {
add_kv(LLM_KV_ATTENTION_RELATIVE_BUCKETS_COUNT, hparams.n_rel_attn_bkts);
add_kv(LLM_KV_ATTENTION_ROPE_PATTERN, hparams.rope_pattern, true);
add_kv(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa);
// add_kv(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, ???);
if (hparams.swa_type != LLAMA_SWA_TYPE_NONE) {
// never collapsed to a scalar: the loaders read a scalar as a period
add_kv(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, std::vector<uint32_t>(
hparams.is_swa_impl.begin(), hparams.is_swa_impl.begin() + hparams.n_layer_all));
}
add_kv(LLM_KV_ATTENTION_SCALE, hparams.f_attention_scale);
add_kv(LLM_KV_ATTENTION_OUTPUT_SCALE, hparams.f_attn_out_scale);
add_kv(LLM_KV_ATTENTION_VALUE_SCALE, hparams.f_attn_value_scale);
Expand All @@ -300,6 +289,9 @@ void llama_model_saver::add_kv_from_model() {
add_kv(LLM_KV_ATTENTION_VALUE_LENGTH_MLA, hparams.n_embd_head_v_mla_impl);
add_kv(LLM_KV_ATTENTION_KEY_LENGTH_SWA, hparams.n_embd_head_k_swa);
add_kv(LLM_KV_ATTENTION_VALUE_LENGTH_SWA, hparams.n_embd_head_v_swa);
add_kv(LLM_KV_ATTENTION_KEY_LENGTH_MLA_SWA, hparams.n_embd_head_k_mla_swa);
add_kv(LLM_KV_ATTENTION_VALUE_LENGTH_MLA_SWA, hparams.n_embd_head_v_mla_swa);
add_kv(LLM_KV_ATTENTION_KV_LORA_RANK_SWA, hparams.n_lora_kv_swa);
add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, hparams.indexer_n_head);
add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, hparams.indexer_head_size);
add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K, hparams.indexer_top_k);
Expand Down
9 changes: 9 additions & 0 deletions src/llama-model.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3305,6 +3305,15 @@ void llama_model_base::create_tensor_qkv(llama_layer & layer, int bid,
}
}

void llama_model_base::load_swa_pattern(llama_model_loader & ml, uint32_t n_pattern, bool dense_first) {
if (ml.get_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.is_swa_impl, false)) {
return;
}

ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, n_pattern, false);
hparams.set_swa_pattern(n_pattern, dense_first);
}

const int32_t * llama_model_target_layer_ids(const struct llama_model * model) {
const auto & v = model->target_layer_ids;
return v.empty() ? nullptr : v.data();
Expand Down
3 changes: 3 additions & 0 deletions src/llama-model.h
Original file line number Diff line number Diff line change
Expand Up @@ -814,6 +814,9 @@ struct llama_model_base : public llama_model {
int64_t n_embd_, int64_t n_embd_q_, int64_t n_embd_k_, int64_t n_embd_v_,
int flags);

// helper: read the SWA pattern as one flag per layer, or as a period expanded by set_swa_pattern
void load_swa_pattern(llama_model_loader & ml, uint32_t n_pattern, bool dense_first = false);

void load_stats (llama_model_loader & ml) override;
void load_hparams(llama_model_loader & ml) override;
void load_vocab (llama_model_loader & ml) override;
Expand Down
4 changes: 1 addition & 3 deletions src/models/afmoe.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,9 +14,7 @@ void llama_model_afmoe::load_arch_hparams(llama_model_loader & ml) {
// Pattern: 3 sliding - 1 full (global_attn_every_n_layers = 4)
if (hparams.n_swa > 0) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
uint32_t swa_period = 4;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 4);

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
4 changes: 1 addition & 3 deletions src/models/cohere2.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,7 @@

void llama_model_cohere2::load_arch_hparams(llama_model_loader & ml) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
uint32_t swa_period = 4;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 4);

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
7 changes: 1 addition & 6 deletions src/models/cohere2moe.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -25,12 +25,7 @@ void llama_model_cohere2moe::load_arch_hparams(llama_model_loader & ml) {
}

hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
uint32_t swa_period = 4;
if (ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false)) {
hparams.set_swa_pattern(swa_period, true);
} else {
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.is_swa_impl, hparams.n_layer());
}
load_swa_pattern(ml, 4, true);

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
4 changes: 1 addition & 3 deletions src/models/exaone-moe.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,7 @@
void llama_model_exaone_moe::load_arch_hparams(llama_model_loader & ml) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
hparams.n_swa = 128;
uint32_t swa_period = 4;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 4);
hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;

Expand Down
4 changes: 1 addition & 3 deletions src/models/exaone4.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,7 @@ void llama_model_exaone4::load_arch_hparams(llama_model_loader & ml) {
if (hparams.n_layer() == 64) { // 32B
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
hparams.n_swa = 4096;
uint32_t swa_period = 4;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 4);

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
4 changes: 1 addition & 3 deletions src/models/gemma-embedding.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,9 +2,7 @@

void llama_model_gemma_embedding::load_arch_hparams(llama_model_loader & ml) {
hparams.swa_type = LLAMA_SWA_TYPE_SYMMETRIC;
uint32_t swa_period = 6;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 6);

hparams.causal_attn = false; // embeddings do not use causal attention

Expand Down
4 changes: 1 addition & 3 deletions src/models/gemma2.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,7 @@
void llama_model_gemma2::load_arch_hparams(llama_model_loader & ml) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
hparams.n_swa = 4096; // default value of gemma 2
uint32_t swa_period = 2;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 2);
hparams.attn_soft_cap = true;
hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
4 changes: 1 addition & 3 deletions src/models/gemma3.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,7 @@ void llama_model_gemma3::load_arch_hparams(llama_model_loader & ml) {
const bool found_swa = ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa, false);
if (found_swa && hparams.n_swa > 0) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
uint32_t swa_period = 6;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 6);

ml.get_key(LLM_KV_ROPE_FREQ_BASE_SWA, hparams.rope_freq_base_train_swa, false);
} else {
Expand Down
4 changes: 1 addition & 3 deletions src/models/gemma3n.cpp
Original file line number Diff line number Diff line change
@@ -1,10 +1,8 @@
#include "models.h"

void llama_model_gemma3n::load_arch_hparams(llama_model_loader & ml) {
uint32_t swa_period = 5;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 5);

hparams.n_layer_kv_from_start = 20;
hparams.f_attention_scale = 1.0f;
Expand Down
4 changes: 1 addition & 3 deletions src/models/laguna.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -36,9 +36,7 @@ void llama_model_laguna::load_arch_hparams(llama_model_loader & ml) {
if (hparams.n_swa > 0) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;

uint32_t swa_period = 4;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period, /*dense_first=*/true); // XS.2: FULL at il%4==0
load_swa_pattern(ml, 4, /*dense_first=*/true); // XS.2: FULL at il%4==0

// Per-layer-type RoPE: full layers use YaRN θ=500000 over 64 dims;
// SWA layers use default RoPE θ=10000 over 128 dims. Base load_hparams
Expand Down
4 changes: 1 addition & 3 deletions src/models/llama4.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,9 +16,7 @@ void llama_model_llama4::load_arch_hparams(llama_model_loader & ml) {
hparams.f_attn_temp_scale = 0.1f;
hparams.f_attn_temp_offset = 1.0f;

uint32_t swa_period = 4; // pattern: 3 chunked - 1 full
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 4); // pattern: 3 chunked - 1 full

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
8 changes: 1 addition & 7 deletions src/models/mellum.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,13 +8,7 @@ void llama_model_mellum::load_arch_hparams(llama_model_loader & ml) {
if (hparams.n_swa > 0) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;

uint32_t swa_period = 4;
const auto res = ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
if (res) {
hparams.set_swa_pattern(swa_period);
} else {
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.is_swa_impl, hparams.n_layer());
}
load_swa_pattern(ml, 4);

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
4 changes: 1 addition & 3 deletions src/models/modern-bert.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,9 +5,7 @@ void llama_model_modern_bert::load_arch_hparams(llama_model_loader & ml) {
if (found_swa && hparams.n_swa > 0) {
hparams.swa_type = LLAMA_SWA_TYPE_SYMMETRIC;
ml.get_key(LLM_KV_ROPE_FREQ_BASE_SWA, hparams.rope_freq_base_train_swa, false);
uint32_t swa_period = 3;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period, true);
load_swa_pattern(ml, 3, true);
} else {
hparams.swa_type = LLAMA_SWA_TYPE_NONE;
}
Expand Down
7 changes: 1 addition & 6 deletions src/models/muse-glimmer.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,12 +10,7 @@ void llama_model_muse_glimmer::load_arch_hparams(llama_model_loader & ml) {
ml.get_key(LLM_KV_ROPE_FREQ_BASE_SWA, hparams.rope_freq_base_train_swa, false);

hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
uint32_t swa_period = 4;
if (ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false)) {
hparams.set_swa_pattern(swa_period);
} else {
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, hparams.is_swa_impl, hparams.n_layer());
}
load_swa_pattern(ml, 4);

switch (hparams.n_layer()) {
case 52: type = LLM_TYPE_30B; break;
Expand Down
4 changes: 1 addition & 3 deletions src/models/olmo2.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,7 @@ void llama_model_olmo2::load_arch_hparams(llama_model_loader & ml) {
const bool found_swa = ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa, false);
if (found_swa && hparams.n_swa > 0) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
uint32_t swa_period = 4;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 4);

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = 1.0; // See olmo2.cpp
Expand Down
4 changes: 1 addition & 3 deletions src/models/openai-moe.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,7 @@ void llama_model_openai_moe::load_arch_hparams(llama_model_loader & ml) {
ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa);

hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
uint32_t swa_period = 2;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 2);

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
4 changes: 1 addition & 3 deletions src/models/plamo3.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,7 @@ void llama_model_plamo3::load_arch_hparams(llama_model_loader & ml) {
if (found_swa && hparams.n_swa > 0) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
ml.get_key(LLM_KV_ROPE_FREQ_BASE_SWA, hparams.rope_freq_base_train_swa, false);
uint32_t swa_period = 8;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period);
load_swa_pattern(ml, 8);
} else {
hparams.swa_type = LLAMA_SWA_TYPE_NONE;
}
Expand Down
4 changes: 1 addition & 3 deletions src/models/smallthinker.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,7 @@ void llama_model_smallthinker::load_arch_hparams(llama_model_loader & ml) {
if (found_swa && hparams.n_swa > 0) {
hparams.swa_type = LLAMA_SWA_TYPE_STANDARD;
hparams.n_swa = 4096;
uint32_t swa_period = 4;
ml.get_key_or_arr(LLM_KV_ATTENTION_SLIDING_WINDOW_PATTERN, swa_period, false);
hparams.set_swa_pattern(swa_period, true);
load_swa_pattern(ml, 4, true);

hparams.rope_freq_base_train_swa = hparams.rope_freq_base_train;
hparams.rope_freq_scale_train_swa = hparams.rope_freq_scale_train;
Expand Down
Loading