Skip to content
Open
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
10 changes: 5 additions & 5 deletions common/arg.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3877,7 +3877,7 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
}
).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_N_GPU_LAYERS_DRAFT"));
add_opt(common_arg(
{"--spec-draft-model", "-md", "--model-draft"}, "FNAME",
{"--spec-draft-model", "-md", "--model-draft", "--mtp-head"}, "FNAME",
"draft model for speculative decoding (default: unused)",
[](common_params & params, const std::string & value) {
params.speculative.draft.mparams.path = value;
Expand Down Expand Up @@ -4023,10 +4023,10 @@ common_params_context common_params_parser_init(common_params & params, llama_ex
//

add_opt(common_arg(
{"--draft", "--draft-n", "--draft-max"}, "N",
"the argument has been removed. use --spec-draft-n-max or --spec-ngram-mod-n-max",
[](common_params & /*params*/, int /*value*/) {
arg_removed("use --spec-draft-n-max or --spec-ngram-mod-n-max");
{"--draft", "--draft-n", "--draft-max", "--draft-block-size"}, "N",
"alias for --spec-draft-n-max: max number of tokens to draft per step",
[](common_params & params, int value) {
params.speculative.draft.n_max = value;
}
).set_spec().set_examples({LLAMA_EXAMPLE_SPECULATIVE, LLAMA_EXAMPLE_LOOKUP, LLAMA_EXAMPLE_SERVER, LLAMA_EXAMPLE_CLI}).set_env("LLAMA_ARG_DRAFT_MAX"));
add_opt(common_arg(
Expand Down
1 change: 1 addition & 0 deletions common/speculative.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ const std::map<std::string, common_speculative_type> common_speculative_type_fro
{"draft-simple", COMMON_SPECULATIVE_TYPE_DRAFT_SIMPLE},
{"draft-eagle3", COMMON_SPECULATIVE_TYPE_DRAFT_EAGLE3},
{"draft-mtp", COMMON_SPECULATIVE_TYPE_DRAFT_MTP},
{"mtp", COMMON_SPECULATIVE_TYPE_DRAFT_MTP},
{"draft-dflash", COMMON_SPECULATIVE_TYPE_DRAFT_DFLASH},
{"ngram-simple", COMMON_SPECULATIVE_TYPE_NGRAM_SIMPLE},
{"ngram-map-k", COMMON_SPECULATIVE_TYPE_NGRAM_MAP_K},
Expand Down
6 changes: 3 additions & 3 deletions src/llama-arch.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@ static const std::map<llm_arch, const char *> LLM_ARCH_NAMES = {
{ LLM_ARCH_GEMMA3, "gemma3" },
{ LLM_ARCH_GEMMA3N, "gemma3n" },
{ LLM_ARCH_GEMMA4, "gemma4" },
{ LLM_ARCH_GEMMA4_ASSISTANT, "gemma4-assistant" },
{ LLM_ARCH_GEMMA4_ASSISTANT, "gemma4_assistant" },
{ LLM_ARCH_GEMMA_EMBEDDING, "gemma-embedding" },
{ LLM_ARCH_STARCODER2, "starcoder2" },
{ LLM_ARCH_MAMBA, "mamba" },
Expand Down Expand Up @@ -504,8 +504,8 @@ static const std::map<llm_tensor, const char *> LLM_TENSOR_NAMES = {
{ LLM_TENSOR_FFN_NORM_EXPS, "blk.%d.ffn_norm_exps" },
{ LLM_TENSOR_ATTN_K_B, "blk.%d.attn_k_b" },
{ LLM_TENSOR_ATTN_V_B, "blk.%d.attn_v_b" },
{ LLM_TENSOR_NEXTN_PROJ_PRE, "nextn.pre_projection" },
{ LLM_TENSOR_NEXTN_PROJ_POST, "nextn.post_projection" },
{ LLM_TENSOR_NEXTN_PROJ_PRE, "mtp.pre_projection" },
{ LLM_TENSOR_NEXTN_PROJ_POST, "mtp.post_projection" },
{ LLM_TENSOR_NEXTN_EH_PROJ, "blk.%d.nextn.eh_proj" },
{ LLM_TENSOR_NEXTN_EMBED_TOKENS, "blk.%d.nextn.embed_tokens" },
{ LLM_TENSOR_NEXTN_ENORM, "blk.%d.nextn.enorm" },
Expand Down
2 changes: 1 addition & 1 deletion src/llama-model.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1493,7 +1493,7 @@ bool llama_model_base::load_tensors(llama_model_loader & ml) {
}
}
}
ml.done_getting_tensors();
ml.done_getting_tensors(arch == LLM_ARCH_GEMMA4_ASSISTANT);

// Tied NVFP4 output is valid when no separate LM-head scale tensors are present.
// If sidecar scales exist, the output weight must be an actual output tensor.
Expand Down
20 changes: 14 additions & 6 deletions src/models/gemma4-assistant.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,9 @@ void llama_model_gemma4_assistant::load_arch_hparams(llama_model_loader & ml) {
hparams.f_attention_scale = 1.0f;

ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.n_layer_nextn, false);
GGML_ASSERT(hparams.n_layer_nextn == hparams.n_layer_all && "n_layer_nextn must be == n_layer_impl");
if (hparams.n_layer_nextn == 0) {
hparams.n_layer_nextn = hparams.n_layer();
}

ml.get_key(LLM_KV_ROPE_FREQ_BASE_SWA, hparams.rope_freq_base_train_swa, false);
ml.get_key(LLM_KV_ATTENTION_SLIDING_WINDOW, hparams.n_swa);
Expand All @@ -21,7 +23,7 @@ void llama_model_gemma4_assistant::load_arch_hparams(llama_model_loader & ml) {
ml.get_key(LLM_KV_ATTENTION_VALUE_LENGTH_SWA, hparams.n_embd_head_v_swa);
}

void llama_model_gemma4_assistant::load_arch_tensors(llama_model_loader &) {
void llama_model_gemma4_assistant::load_arch_tensors(llama_model_loader & ml) {
LLAMA_LOAD_LOCALS;

if (n_embd_head_k != n_embd_head_v) {
Expand All @@ -30,9 +32,6 @@ void llama_model_gemma4_assistant::load_arch_tensors(llama_model_loader &) {
if (hparams.n_embd_head_k_swa != hparams.n_embd_head_v_swa) {
throw std::runtime_error("Gemma 4 assistant requires n_embd_head_k_swa == n_embd_head_v_swa");
}
if (hparams.n_embd_out() == n_embd) {
throw std::runtime_error("Gemma 4 assistant requires embedding_length_out to carry the target hidden size");
}

tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), { n_embd, n_vocab }, 0);
output = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), { n_embd, n_vocab }, TENSOR_DUPLICATED);
Expand All @@ -42,7 +41,16 @@ void llama_model_gemma4_assistant::load_arch_tensors(llama_model_loader &) {
create_tensor(tn(LLM_TENSOR_MASKED_EMBD_CENTROIDS, "weight"), {}, TENSOR_NOT_REQUIRED);
create_tensor(tn(LLM_TENSOR_MASKED_EMBD_ORDERING), {}, TENSOR_NOT_REQUIRED);

const int64_t n_embd_backbone = hparams.n_embd_inp();
// Determine backbone hidden size from projection tensor shape
int64_t n_embd_backbone = hparams.n_embd_inp();
{
auto * meta = ml.get_tensor_meta(tn(LLM_TENSOR_NEXTN_PROJ_POST, "weight").str().c_str());
if (meta && meta->ne[1] > 0) {
n_embd_backbone = meta->ne[1];
}
}
hparams.n_embd_inp_impl = n_embd_backbone;
hparams.n_embd_out_impl = n_embd_backbone;
nextn_proj_post = create_tensor(tn(LLM_TENSOR_NEXTN_PROJ_POST, "weight"), { n_embd, n_embd_backbone }, 0);

int rope_freqs_flag = 0;
Expand Down