diff --git a/ggml/src/ggml-cpu/repack.cpp b/ggml/src/ggml-cpu/repack.cpp index f18758f16b..9689ca3ced 100644 --- a/ggml/src/ggml-cpu/repack.cpp +++ b/ggml/src/ggml-cpu/repack.cpp @@ -2739,7 +2739,7 @@ static block_q8_0x4 make_block_q8_0x4(block_q8_0 * in, unsigned int blck_size_in return out; } -static block_q4_0x4 make_block_q4_0x4(block_q4_0 * in, unsigned int blck_size_interleave) { +static block_q4_0x4 make_block_q4_0x4(block_q4_0 * in, int blck_size_interleave) { block_q4_0x4 out; for (int i = 0; i < 4; i++) { diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index d0956df506..91b841b67b 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -213,7 +213,7 @@ typedef void * ggml_metal_rset_t; // a collection of residency sets (non-owning) typedef struct ggml_metal_rsets * ggml_metal_rsets_t; -ggml_metal_rsets_t ggml_metal_rsets_init(void); +ggml_metal_rsets_t ggml_metal_rsets_init(ggml_metal_device_t dev); void ggml_metal_rsets_free(ggml_metal_rsets_t rsets); // diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 4edd77c6f2..7d2a686850 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -557,7 +557,32 @@ struct ggml_metal_rsets { dispatch_group_t d_group; }; -ggml_metal_rsets_t ggml_metal_rsets_init(void) { +#if defined(GGML_METAL_HAS_RESIDENCY_SETS) +static void ggml_metal_dummy_work(ggml_metal_device_t dev) { + if (dev->mtl_queue == nil) { + return; + } + + @autoreleasepool { + // perform a minimal dummy operation on the GPU + id buf = [dev->mtl_device newBufferWithLength:1 options:MTLResourceStorageModePrivate]; + id cmd_buf = [dev->mtl_queue commandBuffer]; + + { + id encoder = [cmd_buf blitCommandEncoder]; + + [encoder fillBuffer:buf range:NSMakeRange(0, 1) value:0]; + + [encoder endEncoding]; + } + + [cmd_buf commit]; + [buf release]; + } +} +#endif + +ggml_metal_rsets_t ggml_metal_rsets_init(ggml_metal_device_t dev) { ggml_metal_rsets_t res = calloc(1, sizeof(struct ggml_metal_rsets)); res->lock = [[NSLock alloc] init]; @@ -610,6 +635,15 @@ ggml_metal_rsets_t ggml_metal_rsets_init(void) { #endif }); +#if defined(GGML_METAL_HAS_RESIDENCY_SETS) + if (@available(macOS 15.0, iOS 18.0, tvOS 18.0, visionOS 2.0, *)) { + // workaround for residency set memory not being released if no GPU operation occurs + // https://developer.apple.com/forums/thread/839089 + // https://github.com/ggml-org/llama.cpp/issues/25937 + ggml_metal_dummy_work(dev); + } +#endif + return res; } @@ -864,7 +898,7 @@ ggml_metal_device_t ggml_metal_device_init(int device) { } if (dev->props.use_residency_sets) { - dev->rsets = ggml_metal_rsets_init(); + dev->rsets = ggml_metal_rsets_init(dev); } else { dev->rsets = nil; } @@ -1484,6 +1518,7 @@ static void ggml_metal_buffer_rset_free(ggml_metal_buffer_t buf) { if (buf->rset) { [buf->rset endResidency]; [buf->rset removeAllAllocations]; + [buf->rset commit]; [buf->rset release]; } } diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index 6f123f4d1f..ac929cdc3b 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -4519,8 +4519,11 @@ MODEL_TENSORS: dict[MODEL_ARCH, list[MODEL_TENSOR]] = { MODEL_TENSOR.FFN_EXP_PROBS_B, MODEL_TENSOR.LAYER_OUT_NORM, MODEL_TENSOR.NEXTN_EH_PROJ, + MODEL_TENSOR.NEXTN_EMBED_TOKENS, MODEL_TENSOR.NEXTN_ENORM, MODEL_TENSOR.NEXTN_HNORM, + MODEL_TENSOR.NEXTN_SHARED_HEAD_HEAD, + MODEL_TENSOR.NEXTN_SHARED_HEAD_NORM, ], MODEL_ARCH.STEP35: [ MODEL_TENSOR.TOKEN_EMBD, diff --git a/src/llama-model.cpp b/src/llama-model.cpp index 76df995609..68c4cbafd4 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -2247,7 +2247,9 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, filter = [&](uint32_t il) { return il >= hparams.n_layer(); }; } - if ((arch == LLM_ARCH_STEP35 || arch == LLM_ARCH_HY_V3 || arch == LLM_ARCH_GLM_DSA) && hparams.n_layer_nextn > 0) { + if ((arch == LLM_ARCH_STEP35 || arch == LLM_ARCH_HY_V3 || arch == LLM_ARCH_GLM_DSA || + arch == LLM_ARCH_MIMO2) && + hparams.n_layer_nextn > 0) { if (params.ctx_type == LLAMA_CONTEXT_TYPE_MTP) { filter = [&](uint32_t il) { return il >= hparams.n_layer(); }; } else { diff --git a/src/models/mimo2.cpp b/src/models/mimo2.cpp index 8898916057..4080a934cb 100644 --- a/src/models/mimo2.cpp +++ b/src/models/mimo2.cpp @@ -25,9 +25,13 @@ void llama_model_mimo2::load_arch_hparams(llama_model_loader & ml) { } } -void llama_model_mimo2::load_arch_tensors(llama_model_loader &) { +void llama_model_mimo2::load_arch_tensors(llama_model_loader & ml) { LLAMA_LOAD_LOCALS; + const std::string mtp_probe = "blk." + std::to_string(n_layer) + ".nextn.eh_proj.weight"; + const bool trunk_only = (hparams.n_layer_nextn > 0) && (ml.get_weight(mtp_probe.c_str()) == nullptr); + const int mtp_flags = trunk_only ? TENSOR_NOT_REQUIRED : 0; + tok_embd = create_tensor(tn(LLM_TENSOR_TOKEN_EMBD, "weight"), {n_embd, n_vocab}, 0); // output @@ -40,41 +44,46 @@ void llama_model_mimo2::load_arch_tensors(llama_model_loader &) { uint32_t n_embd_v_gqa = hparams.n_embd_v_gqa(i); uint32_t n_head = hparams.n_head(i); - // NextN/MTP layers (the last n_nextn blocks) are preserved but disabled pending support const bool is_nextn = i >= n_layer; - const int skip = is_nextn ? TENSOR_SKIP : 0; + const int flags = is_nextn ? mtp_flags : 0; - create_tensor_qkv(layer, i, n_embd, n_embd_head_k * n_head, n_embd_k_gqa, n_embd_v_gqa, skip); - layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), { n_embd_head_v * n_head, n_embd }, skip); + create_tensor_qkv(layer, i, n_embd, n_embd_head_k * n_head, n_embd_k_gqa, n_embd_v_gqa, flags); + layer.wo = create_tensor(tn(LLM_TENSOR_ATTN_OUT, "weight", i), { n_embd_head_v * n_head, n_embd }, flags); - layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, skip); - layer.attn_sinks = create_tensor(tn(LLM_TENSOR_ATTN_SINKS, "weight", i), {n_head}, TENSOR_NOT_REQUIRED | skip); + layer.attn_norm = create_tensor(tn(LLM_TENSOR_ATTN_NORM, "weight", i), {n_embd}, flags); + layer.attn_sinks = create_tensor(tn(LLM_TENSOR_ATTN_SINKS, "weight", i), {n_head}, TENSOR_NOT_REQUIRED | flags); - layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, skip); + layer.ffn_norm = create_tensor(tn(LLM_TENSOR_FFN_NORM, "weight", i), {n_embd}, flags); // non-MoE branch - layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd, n_ff}, TENSOR_NOT_REQUIRED | skip); - layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), { n_ff, n_embd}, TENSOR_NOT_REQUIRED | skip); - layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, TENSOR_NOT_REQUIRED | skip); + layer.ffn_gate = create_tensor(tn(LLM_TENSOR_FFN_GATE, "weight", i), {n_embd, n_ff}, TENSOR_NOT_REQUIRED | flags); + layer.ffn_down = create_tensor(tn(LLM_TENSOR_FFN_DOWN, "weight", i), { n_ff, n_embd}, TENSOR_NOT_REQUIRED | flags); + layer.ffn_up = create_tensor(tn(LLM_TENSOR_FFN_UP, "weight", i), {n_embd, n_ff}, TENSOR_NOT_REQUIRED | flags); // MoE branch int64_t n_ff_exp = hparams.n_ff_exp; - layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), {n_embd, n_expert}, TENSOR_NOT_REQUIRED | skip); - layer.ffn_gate_exps = create_tensor(tn(LLM_TENSOR_FFN_GATE_EXPS, "weight", i), {n_embd, n_ff_exp, n_expert}, TENSOR_NOT_REQUIRED | skip); - layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), {n_ff_exp, n_embd, n_expert}, TENSOR_NOT_REQUIRED | skip); - layer.ffn_up_exps = create_tensor(tn(LLM_TENSOR_FFN_UP_EXPS, "weight", i), {n_embd, n_ff_exp, n_expert}, TENSOR_NOT_REQUIRED | skip); - layer.ffn_exp_probs_b = create_tensor(tn(LLM_TENSOR_FFN_EXP_PROBS_B, "bias", i), {n_expert}, TENSOR_NOT_REQUIRED | skip); + layer.ffn_gate_inp = create_tensor(tn(LLM_TENSOR_FFN_GATE_INP, "weight", i), {n_embd, n_expert}, TENSOR_NOT_REQUIRED | flags); + layer.ffn_gate_exps = create_tensor(tn(LLM_TENSOR_FFN_GATE_EXPS, "weight", i), {n_embd, n_ff_exp, n_expert}, TENSOR_NOT_REQUIRED | flags); + layer.ffn_down_exps = create_tensor(tn(LLM_TENSOR_FFN_DOWN_EXPS, "weight", i), {n_ff_exp, n_embd, n_expert}, TENSOR_NOT_REQUIRED | flags); + layer.ffn_up_exps = create_tensor(tn(LLM_TENSOR_FFN_UP_EXPS, "weight", i), {n_embd, n_ff_exp, n_expert}, TENSOR_NOT_REQUIRED | flags); + layer.ffn_exp_probs_b = create_tensor(tn(LLM_TENSOR_FFN_EXP_PROBS_B, "bias", i), {n_expert}, TENSOR_NOT_REQUIRED | flags); if (is_nextn) { - layer.nextn.eh_proj = create_tensor(tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", i), {2 * n_embd, n_embd}, skip); - layer.nextn.enorm = create_tensor(tn(LLM_TENSOR_NEXTN_ENORM, "weight", i), {n_embd}, skip); - layer.nextn.hnorm = create_tensor(tn(LLM_TENSOR_NEXTN_HNORM, "weight", i), {n_embd}, skip); - layer.layer_out_norm = create_tensor(tn(LLM_TENSOR_LAYER_OUT_NORM, "weight", i), {n_embd}, skip); + layer.nextn.eh_proj = create_tensor(tn(LLM_TENSOR_NEXTN_EH_PROJ, "weight", i), {2 * n_embd, n_embd}, flags); + layer.nextn.enorm = create_tensor(tn(LLM_TENSOR_NEXTN_ENORM, "weight", i), {n_embd}, flags); + layer.nextn.hnorm = create_tensor(tn(LLM_TENSOR_NEXTN_HNORM, "weight", i), {n_embd}, flags); + layer.nextn.embed_tokens = create_tensor(tn(LLM_TENSOR_NEXTN_EMBED_TOKENS, "weight", i), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED | flags); + layer.nextn.shared_head_head = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_HEAD, "weight", i), {n_embd, n_vocab}, TENSOR_NOT_REQUIRED | flags); + layer.nextn.shared_head_norm = create_tensor(tn(LLM_TENSOR_NEXTN_SHARED_HEAD_NORM, "weight", i), {n_embd}, TENSOR_NOT_REQUIRED | flags); + layer.layer_out_norm = create_tensor(tn(LLM_TENSOR_LAYER_OUT_NORM, "weight", i), {n_embd}, TENSOR_NOT_REQUIRED | flags); } } } std::unique_ptr llama_model_mimo2::build_arch_graph(const llm_graph_params & params) const { + if (params.gtype == LLM_GRAPH_TYPE_DECODER_MTP) { + return std::make_unique(*this, params); + } return std::make_unique(*this, params); } @@ -89,6 +98,8 @@ llama_model_mimo2::graph::graph(const llama_model & model, const llm_graph_param ggml_tensor * inp_out_ids = build_inp_out_ids(); const float v_scale = hparams.f_attn_value_scale; + const bool emit_h_nextn = cparams.embeddings_nextn; + const bool crop_last_layer = inp_out_ids && (!emit_h_nextn || cparams.embeddings_nextn_masked); for (int il = 0; il < n_layer; ++il) { ggml_tensor * inpSA = inpL; @@ -168,7 +179,7 @@ llama_model_mimo2::graph::graph(const llama_model & model, const llm_graph_param } } - if (il == n_layer - 1 && inp_out_ids) { + if (il == n_layer - 1 && crop_last_layer) { cur = ggml_get_rows(ctx0, cur, inp_out_ids); inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids); } @@ -218,6 +229,15 @@ llama_model_mimo2::graph::graph(const llama_model & model, const llm_graph_param cur = inpL; + if (emit_h_nextn) { + cb(cur, "h_nextn", -1); + res->t_h_nextn = cur; + + if (!cparams.embeddings_nextn_masked && inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + } + } + cur = build_norm(cur, model.output_norm, NULL, LLM_NORM_RMS, -1); @@ -233,3 +253,143 @@ llama_model_mimo2::graph::graph(const llama_model & model, const llm_graph_param ggml_build_forward_expand(gf, cur); } + +// Mirrors MiMo's appended NextN block: normalize and fuse token and hidden inputs, run the decoder block, +// expose its pre-head-norm state to the next draft step, then apply the shared output norm and LM head. +// Converted checkpoints may store that shared norm as layer_out_norm, so it remains in the fallback chain. +llama_model_mimo2::graph_mtp::graph_mtp(const llama_model & model, const llm_graph_params & params) + : llm_graph_context(params) { + GGML_ASSERT(hparams.n_layer_nextn > 0 && "MIMO2 MTP requires n_layer_nextn > 0"); + + const int il = hparams.n_layer() + cparams.nextn_layer_offset; + GGML_ASSERT(cparams.nextn_layer_offset >= 0 && + cparams.nextn_layer_offset < (int) hparams.n_layer_nextn && + "nextn_layer_offset out of range [0, n_layer_nextn)"); + + const auto & layer = model.layers[il]; + GGML_ASSERT(layer.nextn.eh_proj && "MIMO2 MTP block missing nextn.eh_proj"); + GGML_ASSERT(layer.nextn.enorm && "MIMO2 MTP block missing nextn.enorm"); + GGML_ASSERT(layer.nextn.hnorm && "MIMO2 MTP block missing nextn.hnorm"); + GGML_ASSERT(layer.wqkv && "MIMO2 MTP requires fused attn_qkv"); + + const uint32_t n_head_l = hparams.n_head(il); + const uint32_t n_head_kv_l = hparams.n_head_kv(il); + + const float freq_base_l = model.get_rope_freq_base(cparams, il); + const float freq_scale_l = model.get_rope_freq_scale(cparams, il); + const float v_scale = hparams.f_attn_value_scale; + + auto inp = std::make_unique(hparams.n_embd); + + inp->tokens = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_tokens); + ggml_set_input(inp->tokens); + + inp->embd = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, hparams.n_embd, n_tokens); + ggml_set_input(inp->embd); + ggml_set_name(inp->embd, "mtp_h_input"); + + ggml_tensor * tok_embd_w = layer.nextn.embed_tokens ? layer.nextn.embed_tokens : model.tok_embd; + ggml_tensor * h_input = inp->embd; + ggml_tensor * tok_embd = ggml_get_rows(ctx0, tok_embd_w, inp->tokens); + cb(tok_embd, "mtp_tok_embd", il); + + res->add_input(std::move(inp)); + + ggml_tensor * inp_pos = build_inp_pos(); + ggml_tensor * inp_out_ids = build_inp_out_ids(); + auto * inp_attn = build_attn_inp_kv_iswa(); + + ggml_tensor * h_norm = build_norm(h_input, layer.nextn.hnorm, nullptr, LLM_NORM_RMS, il); + cb(h_norm, "mtp_hnorm", il); + + ggml_tensor * e_norm = build_norm(tok_embd, layer.nextn.enorm, nullptr, LLM_NORM_RMS, il); + cb(e_norm, "mtp_enorm", il); + + ggml_tensor * concat = ggml_concat(ctx0, e_norm, h_norm, /*dim=*/ 0); + cb(concat, "mtp_concat", il); + + ggml_tensor * cur = build_lora_mm(layer.nextn.eh_proj, concat, layer.nextn.eh_proj_s); + cb(cur, "mtp_eh_proj", il); + + ggml_tensor * inpSA = cur; + + cur = build_norm(cur, layer.attn_norm, nullptr, LLM_NORM_RMS, il); + cb(cur, "mtp_attn_norm", il); + + ggml_tensor * qkv = build_lora_mm(layer.wqkv, cur, layer.wqkv_s); + cb(qkv, "mtp_wqkv", il); + + const size_t row_k = ggml_row_size(qkv->type, n_embd_head_k); + const size_t row_v = ggml_row_size(qkv->type, n_embd_head_v); + const size_t row_full = qkv->nb[1]; + const size_t k_off = row_k * n_head_l; + const size_t v_off = k_off + row_k * n_head_kv_l; + + ggml_tensor * Qcur = ggml_view_3d(ctx0, qkv, n_embd_head_k, n_head_l, n_tokens, row_k, row_full, 0); + ggml_tensor * Kcur = ggml_view_3d(ctx0, qkv, n_embd_head_k, n_head_kv_l, n_tokens, row_k, row_full, k_off); + ggml_tensor * Vcur = ggml_view_3d(ctx0, qkv, n_embd_head_v, n_head_kv_l, n_tokens, row_v, row_full, v_off); + + Qcur = ggml_rope_ext( + ctx0, Qcur, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base_l, freq_scale_l, + ext_factor, attn_factor, beta_fast, beta_slow); + + Kcur = ggml_rope_ext( + ctx0, Kcur, inp_pos, nullptr, + n_rot, rope_type, n_ctx_orig, freq_base_l, freq_scale_l, + ext_factor, attn_factor, beta_fast, beta_slow); + + cb(Qcur, "mtp_Qcur", il); + cb(Kcur, "mtp_Kcur", il); + cb(Vcur, "mtp_Vcur", il); + + cur = build_attn(inp_attn, + layer.wo, nullptr, layer.wo_s, + Qcur, Kcur, Vcur, nullptr, layer.attn_sinks, nullptr, + 1.0f / sqrtf(float(n_embd_head_k)), il); + cb(cur, "mtp_attn_out", il); + + if (v_scale) { + cur = ggml_scale(ctx0, cur, v_scale); + cb(cur, "mtp_attn_out_scaled", il); + } + + ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpSA); + cb(ffn_inp, "mtp_ffn_inp", il); + + cur = build_norm(ffn_inp, layer.ffn_norm, nullptr, LLM_NORM_RMS, il); + cb(cur, "mtp_ffn_norm", il); + + GGML_ASSERT(layer.ffn_gate && layer.ffn_down && layer.ffn_up && "MIMO2 MTP requires dense FFN tensors"); + cur = build_ffn(cur, + layer.ffn_up, layer.ffn_up_b, nullptr, + layer.ffn_gate, layer.ffn_gate_b, nullptr, + layer.ffn_down, layer.ffn_down_b, nullptr, + nullptr, + LLM_FFN_SILU, LLM_FFN_PAR, il); + cb(cur, "mtp_ffn_out", il); + + cur = ggml_add(ctx0, cur, ffn_inp); + cb(cur, "mtp_post_ffn", il); + + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + + cb(cur, "h_nextn", -1); + res->t_h_nextn = cur; + + ggml_tensor * head_norm_w = layer.nextn.shared_head_norm + ? layer.nextn.shared_head_norm + : (layer.layer_out_norm ? layer.layer_out_norm : model.output_norm); + GGML_ASSERT(head_norm_w && "MIMO2 MTP missing head norm fallback"); + cur = build_norm(cur, head_norm_w, nullptr, LLM_NORM_RMS, -1); + cb(cur, "mtp_shared_head_norm", -1); + + ggml_tensor * head_w = layer.nextn.shared_head_head ? layer.nextn.shared_head_head : model.output; + ggml_tensor * head_s = layer.nextn.shared_head_head ? layer.nextn.shared_head_head_s : model.output_s; + GGML_ASSERT(head_w && "MIMO2 MTP missing LM head fallback"); + cur = build_lora_mm(head_w, cur, head_s); + cb(cur, "result_output", -1); + + res->t_logits = cur; + ggml_build_forward_expand(gf, cur); +} diff --git a/src/models/minimax-m3.cpp b/src/models/minimax-m3.cpp index 6068fc6b87..3e7bada64b 100644 --- a/src/models/minimax-m3.cpp +++ b/src/models/minimax-m3.cpp @@ -2,7 +2,6 @@ #include "llama-kv-cache.h" #include #include -#include #include // MiniMax-M3: MiniMax-M2 style GQA (per-head QK-norm, partial rotary) with @@ -126,68 +125,6 @@ public: int64_t nblk; }; -// pooled score of a block with no visible token: -inf from the mask, or -FLT_MAX from the -// max-pool identity when every element of the block is -inf -static inline bool msa_score_masked(float x) { return x <= -1e30f; } - -// MSA block selection (batch regime) -// CPU custom op, the token-level expansion and the combination with the causal mask happen on the GPU. -static void msa_block_mask_op(struct ggml_tensor * dst, int ith, int nth, void * userdata) { - const struct ggml_tensor * bs = dst->src[0]; - const struct ggml_tensor * bias = dst->src[1]; - const msa_params * p = (const msa_params *) userdata; - - const int nblk = (int) bs->ne[0]; - const int Hd = (int) bs->ne[1]; - const int S = (int) bs->ne[2]; - - GGML_ASSERT(bs->type == GGML_TYPE_F32 && ggml_is_contiguous(bs)); - GGML_ASSERT(bias->type == GGML_TYPE_F32 && ggml_is_contiguous(bias)); - GGML_ASSERT(dst->type == GGML_TYPE_F16 && ggml_is_contiguous(dst)); - GGML_ASSERT(dst->ne[0] == nblk && dst->ne[1] == S && dst->ne[2] == Hd); - GGML_ASSERT(bias->ne[0] == nblk && bias->ne[1] == S); - - const int topk = p->topk_blocks < nblk ? p->topk_blocks : nblk; - - const ggml_fp16_t f16_zero = ggml_fp32_to_fp16(0.0f); - const ggml_fp16_t f16_ninf = ggml_fp32_to_fp16(-INFINITY); - - std::vector rank(nblk); - std::vector valid(nblk); - std::vector ord(nblk); - - ggml_fp16_t * out = (ggml_fp16_t *) dst->data; - - for (int i = ith; i < S; i += nth) { - const float * bias_col = (const float *) bias->data + (size_t) i * nblk; - for (int h = 0; h < Hd; ++h) { - const float * bs_col = (const float *) bs->data + ((size_t) i * Hd + h) * nblk; - - for (int bk = 0; bk < nblk; ++bk) { - // a block is selectable if it has a visible token or is locally forced - valid[bk] = !msa_score_masked(bs_col[bk]) || bias_col[bk] > 0.0f; - rank [bk] = bias_col[bk] > 0.0f ? bias_col[bk] : bs_col[bk]; - ord [bk] = bk; - } - - std::partial_sort(ord.begin(), ord.begin() + topk, ord.end(), - [&](int a, int b) { return rank[a] > rank[b]; }); - - ggml_fp16_t * dst_col = out + ((size_t) h * S + i) * nblk; - for (int bk = 0; bk < nblk; ++bk) { - dst_col[bk] = f16_ninf; - } - for (int t = 0; t < topk; ++t) { - const int bk = ord[t]; - if (!valid[bk]) { - break; // sorted desc: first invalid -> fewer than topk selectable blocks - } - dst_col[bk] = f16_zero; - } - } - } -} - // One FA call for all GQA groups (and at multi-stream decode, all streams) by mapping them onto the FA sequence dim (ne[3]) ggml_tensor * llama_model_minimax_m3::graph::build_attn_msa_fa( ggml_tensor * q_cur, // [D, HQ, T] @@ -433,8 +370,8 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ msa_mf->nb[1], msa_mf->nb[1], st*msa_mf->nb[3]); ggml_tensor * km_s = ggml_view_3d(ctx0, msa_kqm, n_kv, n_tps, 1, msa_kqm->nb[1], msa_kqm->nb[3], st*msa_kqm->nb[3]); - ggml_tensor * bias_s = ggml_view_2d(ctx0, msa_loc->bias, nblk, n_tps, - msa_loc->bias->nb[1], st*n_tps*msa_loc->bias->nb[1]); + ggml_tensor * bias_s = ggml_view_3d(ctx0, msa_loc->bias, nblk, 1, n_tps, + msa_loc->bias->nb[1], msa_loc->bias->nb[1], st*n_tps*msa_loc->bias->nb[1]); ggml_tensor * q_s = ggml_view_3d(ctx0, Qcur, D, n_head, n_tps, Qcur->nb[1], Qcur->nb[2], st*n_tps*Qcur->nb[2]); ggml_tensor * k_s = ggml_view_4d(ctx0, k, D, HKV, n_kv, 1, @@ -453,15 +390,27 @@ llama_model_minimax_m3::graph::graph(const llama_model & model, const llm_graph_ ggml_tensor * bs = ggml_pool_2d(ctx0, sc, GGML_OP_POOL_MAX, blk, 1, blk, 1, 0, 0); cb(bs, "msa_bs", il); - // block-level 0/-inf keep mask on the CPU, tiny transfer - ggml_tensor * srcs[2] = { bs, bias_s }; - ggml_tensor * bm = ggml_custom_4d(ctx0, GGML_TYPE_F16, - nblk, n_tps, Hd, 1, - srcs, 2, msa_block_mask_op, GGML_N_TASKS_MAX, - const_cast(&mm.msa_p)); + // bias the scores so locally-forced blocks always rank first + ggml_tensor * bsf = ggml_add(ctx0, bs, bias_s); // [nblk, Hd, n_tps] + cb(bsf, "msa_bsf", il); + + ggml_tensor * idx = ggml_top_k(ctx0, bsf, K); // [K, Hd, n_tps] i32 + + ggml_tensor * ninf = ggml_cast(ctx0, + ggml_scale_bias(ctx0, bias_s, 0.0f, -1e30f), + GGML_TYPE_F16); // [nblk, 1, n_tps] + ninf = ggml_repeat_4d(ctx0, ninf, nblk, Hd, n_tps, 1); + ggml_tensor * zero = ggml_scale(ctx0, + ggml_cast(ctx0, idx, GGML_TYPE_F32), 0.0f); + ggml_tensor * bm = ggml_set_rows(ctx0, + ggml_reshape_3d(ctx0, ninf, 1, nblk, Hd*n_tps), + ggml_reshape_3d(ctx0, zero, 1, K, Hd*n_tps), + ggml_reshape_2d(ctx0, idx, K, Hd*n_tps)); + bm = ggml_reshape_3d(ctx0, bm, nblk, Hd, n_tps); + bm = ggml_cont(ctx0, ggml_permute(ctx0, bm, 0, 2, 1, 3)); // [nblk, n_tps, Hd] cb(bm, "msa_block_mask", il); - // expand block -> token granularity on the GPU (j = bk*blk + t), + // expand block -> token granularity (j = bk*blk + t), // then combine with the causal mask in place ggml_tensor * bmx = ggml_repeat_4d(ctx0, ggml_reshape_3d(ctx0, bm, 1, nblk, n_tps*Hd), diff --git a/src/models/models.h b/src/models/models.h index c73136f3bd..bb372ece81 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -2127,6 +2127,10 @@ struct llama_model_mimo2 : public llama_model_base { graph(const llama_model & model, const llm_graph_params & params); }; + struct graph_mtp : public llm_graph_context { + graph_mtp(const llama_model & model, const llm_graph_params & params); + }; + std::unique_ptr build_arch_graph(const llm_graph_params & params) const override; }; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 7a93b19a07..805b744726 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -278,6 +278,9 @@ set_tests_properties(test-state-restore-fragmented PROPERTIES FIXTURES_REQUIRED llama_build_and_test(test-save-load-state.cpp LABEL "model" ARGS -m "${MODEL_DEST}") set_tests_properties(test-save-load-state PROPERTIES FIXTURES_REQUIRED test-download-model) +if (APPLE) + llama_build(test-rset-release.cpp get-model.cpp) +endif() if (NOT GGML_BACKEND_DL) # these tests use the backends directly and cannot be built with dynamic loading llama_build_and_test(test-barrier.cpp) diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index a5b660f47a..d0eb173a8f 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -1350,18 +1350,22 @@ struct test_case { // check if the backends support the ops bool supported = true; + std::string unsupported_str; for (ggml_backend_t backend : {backend1, backend2}) { for (ggml_tensor * t = ggml_get_first_tensor(ctx.get()); t != NULL; t = ggml_get_next_tensor(ctx.get(), t)) { if (!ggml_backend_supports_op(backend, t)) { supported = false; - break; + if (unsupported_str.empty()) { + unsupported_str = std::string(ggml_backend_name(backend)); + } else { + unsupported_str += ", " + std::string(ggml_backend_name(backend)); + } } } } if (!supported) { - // Create test result for unsupported operation - test_result result(ggml_backend_name(backend1), current_op_name, vars(), "test", + test_result result(unsupported_str, current_op_name, vars(), "test", false, false, "not supported"); print_test_result_locked(output_printer, result); diff --git a/tests/test-rset-release.cpp b/tests/test-rset-release.cpp new file mode 100644 index 0000000000..bf03c5e8bd --- /dev/null +++ b/tests/test-rset-release.cpp @@ -0,0 +1,53 @@ +// ref: https://github.com/ggml-org/llama.cpp/issues/25937 +// only works reliably when run with a large model that occupies 3GB+ of wired memory +// thus, this test is not run by default +// example model to run with: google/gemma-4-E4B-it-qat-q4_0-gguf + +#include +#include +#include +#include + +#include "llama.h" +#include "get-model.h" + +static uint64_t wired_memory() { + vm_statistics64_data_t vmstat; + mach_msg_type_number_t count = HOST_VM_INFO64_COUNT; + if (host_statistics64(mach_host_self(), HOST_VM_INFO64, (host_info64_t)&vmstat, &count) != KERN_SUCCESS) { + return UINT64_MAX; + } + return static_cast(vmstat.wire_count) * vm_kernel_page_size; +} + +int main(int argc, char ** argv) { + auto * model_path = get_model_or_exit(argc, argv); + + llama_backend_init(); + + const uint64_t wired_initial = wired_memory(); + + llama_model_params params = llama_model_default_params(); + params.load_mode = LLAMA_LOAD_MODE_NONE; + struct llama_model* model = llama_model_load_from_file(model_path, params); + + const uint64_t wired_loaded = wired_memory(); + const uint64_t wired_delta = wired_loaded - wired_initial; + // system memory fluctuates, so we need to allocate enough to reliably detect the release + GGML_ASSERT(wired_delta > 2'000'000'000); // 2GB + + llama_model_free(model); + + const uint64_t t_start_ms = ggml_time_ms(); + + // expect most of the allocated memory to be released within 10 seconds + // we allow for some tolerance due to system-wide memory fluctuations + while (wired_memory() > wired_loaded - 0.75 * wired_delta) { + GGML_ASSERT(ggml_time_ms() - t_start_ms < 10'000); + usleep(100'000); // 100ms + } + + llama_backend_free(); + + return 0; +} diff --git a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddDropdown.svelte b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddDropdown.svelte index 905c2fe6f0..f81dcf09c0 100644 --- a/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddDropdown.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatForm/ChatFormActions/ChatFormActionAdd/ChatFormActionAddDropdown.svelte @@ -48,6 +48,9 @@ }: Props = $props(); let dropdownOpen = $state(false); + // The system message action moves focus to the message editor, so the menu + // must not restore focus to the trigger on close + let suppressCloseAutoFocus = false; function handleMcpSettingsClick() { dropdownOpen = false; @@ -96,7 +99,16 @@ - + { + if (suppressCloseAutoFocus) { + suppressCloseAutoFocus = false; + e.preventDefault(); + } + }} + > @@ -148,7 +160,10 @@ { + suppressCloseAutoFocus = true; + onSystemPromptClick?.(); + }} > diff --git a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte index 8e8a14ac31..b8068f7907 100644 --- a/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatMessages/ChatMessage/ChatMessage.svelte @@ -2,6 +2,7 @@ import { goto } from '$app/navigation'; import { getChatActionsContext, setMessageEditContext } from '$lib/contexts'; import { chatStore, pendingEditMessageId } from '$lib/stores/chat.svelte'; + import { isMobile } from '$lib/stores/viewport.svelte'; import { conversationsStore } from '$lib/stores/conversations.svelte'; import { DatabaseService } from '$lib/services/database.service'; import { SYSTEM_MESSAGE_PLACEHOLDER } from '$lib/constants'; @@ -46,7 +47,14 @@ assistantMessages: number; messageTypes: string[]; } | null>(null); - let editedContent = $derived(message.content); + // The system message placeholder must never surface as editable content; keeping + // it in the derived (not just in handleEdit) guards against prop invalidation + // reverting the override while editing + let editedContent = $derived( + message.role === MessageRole.SYSTEM && message.content === SYSTEM_MESSAGE_PLACEHOLDER + ? '' + : message.content + ); let rawEditContent = $derived.by(() => { if (message.role !== MessageRole.ASSISTANT) return undefined; @@ -265,6 +273,12 @@ chatActions.navigateToSibling(siblingId); } + // After the system message flow ends, hand focus to the main chat form + function focusMainChatForm() { + if (isMobile.current) return; + document.querySelector('.chat-screen-form-wrapper textarea')?.focus(); + } + async function handleSaveEdit() { if (message.role === MessageRole.SYSTEM) { // System messages: update in place without branching @@ -276,6 +290,8 @@ isEditing = false; if (conversationDeleted) { goto(ROUTES.START); + } else { + focusMainChatForm(); } return; } @@ -285,6 +301,7 @@ if (index !== -1) { conversationsStore.updateMessageAtIndex(index, { content: newContent }); } + focusMainChatForm(); } else if (message.role === MessageRole.USER) { const finalExtras = await getMergedExtras(); chatActions.editWithBranching(message, editedContent.trim(), finalExtras); diff --git a/tools/ui/src/lib/components/app/chat/ChatScreen/ChatScreenForm.svelte b/tools/ui/src/lib/components/app/chat/ChatScreen/ChatScreenForm.svelte index 600180742a..8eb17eeae4 100644 --- a/tools/ui/src/lib/components/app/chat/ChatScreen/ChatScreenForm.svelte +++ b/tools/ui/src/lib/components/app/chat/ChatScreen/ChatScreenForm.svelte @@ -106,15 +106,23 @@ onFileRemove?.(fileId); } + // Auto-focus must not steal focus already claimed elsewhere (e.g. the system + // message editor opened just before a navigation) + function focusFormUnlessCaptured() { + const active = document.activeElement; + if (active instanceof HTMLTextAreaElement || active instanceof HTMLInputElement) return; + chatFormRef?.focus(); + } + onMount(() => { if (!isMobile.current) { - setTimeout(() => chatFormRef?.focus(), 100); + setTimeout(focusFormUnlessCaptured, 100); } }); afterNavigate((navigation) => { if (navigation?.from != null && !isMobile.current) { - setTimeout(() => chatFormRef?.focus(), 100); + setTimeout(focusFormUnlessCaptured, 100); } }); @@ -127,7 +135,7 @@ $effect(() => { if (previousIsLoading && !isLoading) { - setTimeout(() => chatFormRef?.focus(), 10); + setTimeout(focusFormUnlessCaptured, 10); } previousIsLoading = isLoading; diff --git a/tools/ui/src/lib/services/database.service.ts b/tools/ui/src/lib/services/database.service.ts index bc65caaca3..0a7c59b9c8 100644 --- a/tools/ui/src/lib/services/database.service.ts +++ b/tools/ui/src/lib/services/database.service.ts @@ -31,14 +31,19 @@ export class DatabaseService { * Creates a new conversation. * * @param name - Name of the conversation + * @param fields - Optional extra fields (e.g. reasoningEffort) * @returns The created conversation */ - static async createConversation(name: string): Promise { + static async createConversation( + name: string, + fields?: Partial> + ): Promise { const conversation: DatabaseConversation = { id: uuid(), name, lastModified: Date.now(), - currNode: '' + currNode: '', + ...fields }; await db[IDXDB_TABLES.conversations].add(conversation); @@ -137,7 +142,7 @@ export class DatabaseService { * @param systemPrompt - The system prompt content (must be non-empty) * @param parentId - Parent message ID (typically the root message) * @returns The created system message - * @throws Error if systemPrompt is empty + * @throws Error if systemPrompt is empty or the parent message does not exist */ static async createSystemMessage( convId: string, @@ -149,27 +154,30 @@ export class DatabaseService { throw new Error('Cannot create system message with empty content'); } - const systemMessage: DatabaseMessage = { - id: uuid(), - convId, - type: MessageRole.SYSTEM, - timestamp: Date.now(), - role: MessageRole.SYSTEM, - content: trimmedPrompt, - parent: parentId, - children: [] - }; + return await db.transaction('rw', db[IDXDB_TABLES.messages], async () => { + const parentMessage = await db[IDXDB_TABLES.messages].get(parentId); + if (!parentMessage) { + throw new Error(`Parent message ${parentId} not found`); + } - await db[IDXDB_TABLES.messages].add(systemMessage); + const systemMessage: DatabaseMessage = { + id: uuid(), + convId, + type: MessageRole.SYSTEM, + timestamp: Date.now(), + role: MessageRole.SYSTEM, + content: trimmedPrompt, + parent: parentId, + children: [] + }; - const parentMessage = await db[IDXDB_TABLES.messages].get(parentId); - if (parentMessage) { + await db[IDXDB_TABLES.messages].add(systemMessage); await db[IDXDB_TABLES.messages].update(parentId, { children: [...parentMessage.children, systemMessage.id] }); - } - return systemMessage; + return systemMessage; + }); } /** @@ -442,7 +450,8 @@ export class DatabaseService { } /** - * Updates a conversation. + * Updates a conversation. `lastModified` is never stamped implicitly; + * pass it in `updates` to bump the conversation in recency ordering. * * @param id - Conversation ID * @param updates - Partial updates to apply @@ -452,10 +461,7 @@ export class DatabaseService { id: string, updates: Partial> ): Promise { - await db[IDXDB_TABLES.conversations].update(id, { - ...updates, - lastModified: Date.now() - }); + await db[IDXDB_TABLES.conversations].update(id, updates); } /** @@ -473,7 +479,7 @@ export class DatabaseService { * @returns The new pinned status */ static async toggleConversationPin(id: string): Promise { - const conversation = await db.conversations.get(id); + const conversation = await db[IDXDB_TABLES.conversations].get(id); if (!conversation) { throw new Error(`Conversation ${id} not found`); } @@ -497,7 +503,6 @@ export class DatabaseService { const result = new Map(); if (cleanIds.length === 0) return result; - const now = Date.now(); await db.transaction('rw', db[IDXDB_TABLES.conversations], async () => { const convs = await db[IDXDB_TABLES.conversations].bulkGet(cleanIds); const updates: DatabaseConversation[] = []; @@ -505,7 +510,7 @@ export class DatabaseService { const conv = convs[i]; if (!conv) continue; const newPinned = !conv.pinned; - updates.push({ ...conv, pinned: newPinned, lastModified: now }); + updates.push({ ...conv, pinned: newPinned }); result.set(cleanIds[i], newPinned); } if (updates.length === 0) return; diff --git a/tools/ui/src/lib/stores/chat.svelte.ts b/tools/ui/src/lib/stores/chat.svelte.ts index 222723ab10..935ed8e165 100644 --- a/tools/ui/src/lib/stores/chat.svelte.ts +++ b/tools/ui/src/lib/stores/chat.svelte.ts @@ -1658,7 +1658,8 @@ class ChatStore { generateConversationTitle(newContent, Boolean(config().titleGenerationUseFirstLine)) ); const messagesToRemove = conversationsStore.activeMessages.slice(messageIndex + 1); - for (const message of messagesToRemove) await DatabaseService.deleteMessage(message.id); + if (messagesToRemove.length > 0) + await DatabaseService.deleteMessageCascading(activeConv.id, messagesToRemove[0].id); conversationsStore.sliceActiveMessages(messageIndex + 1); conversationsStore.updateConversationTimestamp(); this.setChatLoading(activeConv.id, true); @@ -1690,7 +1691,7 @@ class ChatStore { const { index: messageIndex } = result; try { const messagesToRemove = conversationsStore.activeMessages.slice(messageIndex); - for (const message of messagesToRemove) await DatabaseService.deleteMessage(message.id); + await DatabaseService.deleteMessageCascading(activeConv.id, messagesToRemove[0].id); conversationsStore.sliceActiveMessages(messageIndex); conversationsStore.updateConversationTimestamp(); this.setChatLoading(activeConv.id, true); @@ -2037,7 +2038,7 @@ class ChatStore { timings }); - conversationsStore.updateConversationTimestamp(); + conversationsStore.updateConversationTimestamp(msg.convId); this.setChatLoading(msg.convId, false); this.clearChatStreaming(msg.convId); diff --git a/tools/ui/src/lib/stores/conversations.svelte.ts b/tools/ui/src/lib/stores/conversations.svelte.ts index e467c8fad9..1a1157c72c 100644 --- a/tools/ui/src/lib/stores/conversations.svelte.ts +++ b/tools/ui/src/lib/stores/conversations.svelte.ts @@ -111,6 +111,9 @@ class ConversationsStore { | ((messageId: string, updates: Partial) => void) | null = null; + /** In-flight init run; shared by concurrent callers, reset on failure to allow retry */ + private initPromise: Promise | null = null; + /** * * @@ -121,19 +124,25 @@ class ConversationsStore { /** * Initialize the store by loading conversations from database. - * Must be called once after app startup. + * Safe to call multiple times: concurrent callers share a single run, + * and a failed run can be retried by calling again. */ - async init(): Promise { - if (!browser) return; - if (this.isInitialized) return; + init(): Promise { + if (!browser) return Promise.resolve(); + if (this.initPromise) return this.initPromise; - try { - await MigrationService.runAllMigrations(); - await this.loadConversations(); - this.isInitialized = true; - } catch (error) { - console.error('Failed to initialize conversations:', error); - } + this.initPromise = (async () => { + try { + await MigrationService.runAllMigrations(); + await this.loadConversations(); + this.isInitialized = true; + } catch (error) { + console.error('Failed to initialize conversations:', error); + this.initPromise = null; + } + })(); + + return this.initPromise; } /** @@ -237,15 +246,11 @@ class ConversationsStore { */ async createConversation(name?: string): Promise { const conversationName = name || `Chat ${new Date().toLocaleString()}`; - const conversation = await DatabaseService.createConversation(conversationName); // No MCP override list is seeded: getAllMcpServerOverrides resolves // servers without a per-conversation override to `mcpServers[i].enabled`, // and only explicit toggles are stored on the conversation. - - // Inherit the global reasoning default into the new conversation - conversation.reasoningEffort = this.pendingReasoningEffort; - await DatabaseService.updateConversation(conversation.id, { + const conversation = await DatabaseService.createConversation(conversationName, { reasoningEffort: this.pendingReasoningEffort }); @@ -358,10 +363,7 @@ class ConversationsStore { async deleteAll(): Promise { try { const allConversations = await DatabaseService.getAllConversations(); - - for (const conv of allConversations) { - await DatabaseService.deleteConversation(conv.id); - } + await DatabaseService.bulkDeleteConversations(allConversations.map((c) => c.id)); this.clearActiveConversation(); this.conversations = []; @@ -412,7 +414,9 @@ class ConversationsStore { } toast.success( - convIds.length === 1 ? 'Conversation deleted' : `${convIds.length} conversations deleted` + idsToRemove.size === 1 + ? 'Conversation deleted' + : `${idsToRemove.size} conversations deleted` ); } catch (error) { console.error('Failed to bulk delete conversations:', error); @@ -443,7 +447,6 @@ class ConversationsStore { const newPinned = updates.get(this.conversations[i].id); if (newPinned !== undefined) this.conversations[i].pinned = newPinned; } - this.conversations = [...this.conversations]; toast.success( convIds.length === 1 @@ -552,7 +555,6 @@ class ConversationsStore { if (convIndex !== -1) { this.conversations[convIndex].name = name; - this.conversations = [...this.conversations]; } if (this.activeConversation?.id === convId) { @@ -576,7 +578,6 @@ class ConversationsStore { if (convIndex !== -1) { this.conversations[convIndex].pinned = newPinnedState; - this.conversations = [...this.conversations]; } if (this.activeConversation?.id === convId) { @@ -591,18 +592,33 @@ class ConversationsStore { } /** - * Updates conversation lastModified timestamp and moves it to top of list + * Marks a conversation as recently active: stamps lastModified (persisted) + * and moves it to the top of the list. Only message-activity flows call + * this; metadata updates (rename, pin, settings) do not. + * + * @param convId - Conversation that produced the activity, defaults to the active one */ - updateConversationTimestamp(): void { - if (!this.activeConversation) return; + updateConversationTimestamp(convId?: string): void { + const targetId = convId ?? this.activeConversation?.id; + if (!targetId) return; - const chatIndex = this.conversations.findIndex((c) => c.id === this.activeConversation!.id); + const now = Date.now(); + + const chatIndex = this.conversations.findIndex((c) => c.id === targetId); if (chatIndex !== -1) { - this.conversations[chatIndex].lastModified = Date.now(); + this.conversations[chatIndex].lastModified = now; const updatedConv = this.conversations.splice(chatIndex, 1)[0]; this.conversations = [updatedConv, ...this.conversations]; } + + if (this.activeConversation?.id === targetId) { + this.activeConversation = { ...this.activeConversation, lastModified: now }; + } + + DatabaseService.updateConversation(targetId, { lastModified: now }).catch((error) => + console.error('Failed to update conversation timestamp:', error) + ); } /** @@ -773,7 +789,6 @@ class ConversationsStore { if (convIndex !== -1) { this.conversations[convIndex].mcpServerOverrides = newOverrides.length > 0 ? newOverrides : undefined; - this.conversations = [...this.conversations]; } } @@ -837,7 +852,6 @@ class ConversationsStore { const convIndex = this.conversations.findIndex((c) => c.id === this.activeConversation!.id); if (convIndex !== -1) { this.conversations[convIndex].reasoningEffort = effort; - this.conversations = [...this.conversations]; } } diff --git a/tools/ui/src/lib/utils/legacy-migration.ts b/tools/ui/src/lib/utils/legacy-migration.ts deleted file mode 100644 index 6b0890a363..0000000000 --- a/tools/ui/src/lib/utils/legacy-migration.ts +++ /dev/null @@ -1,364 +0,0 @@ -/** - * @deprecated Legacy migration utility — remove at some point in the future once all users have migrated to the new structured agentic message format. - * - * Converts old marker-based agentic messages to the new structured format - * with separate messages per turn. - * - * Old format: Single assistant message with markers in content: - * <<>>...<<>> - * <<>>...<<>> - * - * New format: Separate messages per turn: - * - assistant (content + reasoningContent + toolCalls) - * - tool (toolCallId + content) - * - assistant (next turn) - * - ... - */ - -import { LEGACY_AGENTIC_REGEX, LEGACY_REASONING_TAGS } from '$lib/constants'; -import { DatabaseService } from '$lib/services/database.service'; -import { MessageRole, MessageType } from '$lib/enums'; -import type { DatabaseMessage } from '$lib/types/database'; - -const MIGRATION_DONE_KEY = 'llama-ui-migration-v2-done'; -/** @deprecated Use {@link MIGRATION_DONE_KEY} instead */ -const DEPRECATED_MIGRATION_DONE_KEY = 'llama-webui-migration-v2-done'; - -/** - * @deprecated Part of legacy migration — remove with the migration module. - * Check if migration has been performed. - */ -export function isMigrationNeeded(): boolean { - try { - // Check new key first, fall back to deprecated old key - if (localStorage.getItem(MIGRATION_DONE_KEY)) return false; - if (localStorage.getItem(DEPRECATED_MIGRATION_DONE_KEY)) { - // Migrate to new key - try { - localStorage.setItem(MIGRATION_DONE_KEY, String(Date.now())); - localStorage.removeItem(DEPRECATED_MIGRATION_DONE_KEY); - } catch { - // Ignore storage errors - } - return false; - } - return true; - } catch { - return false; - } -} - -/** - * Mark migration as done. - */ -function markMigrationDone(): void { - try { - localStorage.setItem(MIGRATION_DONE_KEY, String(Date.now())); - } catch { - // Ignore localStorage errors - } -} - -/** - * Check if a message has legacy markers in its content. - */ -function hasLegacyMarkers(message: DatabaseMessage): boolean { - if (!message.content) return false; - return LEGACY_AGENTIC_REGEX.HAS_LEGACY_MARKERS.test(message.content); -} - -/** - * Extract reasoning content from legacy marker format. - */ -function extractLegacyReasoning(content: string): { reasoning: string; cleanContent: string } { - let reasoning = ''; - let cleanContent = content; - - // Extract all reasoning blocks - const re = new RegExp(LEGACY_AGENTIC_REGEX.REASONING_EXTRACT.source, 'g'); - let match; - while ((match = re.exec(content)) !== null) { - reasoning += match[1]; - } - - // Remove reasoning tags from content - cleanContent = cleanContent - .replace(new RegExp(LEGACY_AGENTIC_REGEX.REASONING_BLOCK.source, 'g'), '') - .replace(LEGACY_AGENTIC_REGEX.REASONING_OPEN, ''); - - return { reasoning, cleanContent }; -} - -/** - * Parse legacy content with tool call markers into structured turns. - */ -interface ParsedTurn { - textBefore: string; - toolCalls: Array<{ - name: string; - args: string; - result: string; - }>; -} - -function parseLegacyToolCalls(content: string): ParsedTurn[] { - const turns: ParsedTurn[] = []; - const regex = new RegExp(LEGACY_AGENTIC_REGEX.COMPLETED_TOOL_CALL.source, 'g'); - - let lastIndex = 0; - let currentTurn: ParsedTurn = { textBefore: '', toolCalls: [] }; - let match; - - while ((match = regex.exec(content)) !== null) { - const textBefore = content.slice(lastIndex, match.index).trim(); - - // If there's text between tool calls and we already have tool calls, - // that means a new turn started (text after tool results = new LLM turn) - if (textBefore && currentTurn.toolCalls.length > 0) { - turns.push(currentTurn); - currentTurn = { textBefore, toolCalls: [] }; - } else if (textBefore && currentTurn.toolCalls.length === 0) { - currentTurn.textBefore = textBefore; - } - - currentTurn.toolCalls.push({ - name: match[1], - args: match[2], - result: match[3].replace(/^\n+|\n+$/g, '') - }); - - lastIndex = match.index + match[0].length; - } - - // Any remaining text after the last tool call - const remainingText = content.slice(lastIndex).trim(); - - if (currentTurn.toolCalls.length > 0) { - turns.push(currentTurn); - } - - // If there's text after all tool calls, it's the final assistant response - if (remainingText) { - // Remove any partial/open markers - const cleanRemaining = remainingText - .replace(LEGACY_AGENTIC_REGEX.AGENTIC_TOOL_CALL_OPEN, '') - .trim(); - if (cleanRemaining) { - turns.push({ textBefore: cleanRemaining, toolCalls: [] }); - } - } - - // If no tool calls found at all, return the original content as a single turn - if (turns.length === 0) { - turns.push({ textBefore: content.trim(), toolCalls: [] }); - } - - return turns; -} - -/** - * Migrate a single conversation's messages from legacy format to new format. - */ -async function migrateConversation(convId: string): Promise { - const allMessages = await DatabaseService.getConversationMessages(convId); - let migratedCount = 0; - - for (const message of allMessages) { - if (message.role !== MessageRole.ASSISTANT) continue; - if (!hasLegacyMarkers(message)) { - // Still check for reasoning-only markers (no tool calls) - if (message.content?.includes(LEGACY_REASONING_TAGS.START)) { - const { reasoning, cleanContent } = extractLegacyReasoning(message.content); - await DatabaseService.updateMessage(message.id, { - content: cleanContent.trim(), - reasoningContent: reasoning || undefined - }); - migratedCount++; - } - continue; - } - - // Has agentic markers - full migration needed - const { reasoning, cleanContent } = extractLegacyReasoning(message.content); - const turns = parseLegacyToolCalls(cleanContent); - - // Parse existing toolCalls JSON to try to match IDs - let existingToolCalls: Array<{ - id: string; - function?: { name: string; arguments: string }; - }> = []; - if (message.toolCalls) { - try { - existingToolCalls = JSON.parse(message.toolCalls); - } catch { - // Ignore - } - } - - // First turn uses the existing message - const firstTurn = turns[0]; - if (!firstTurn) continue; - - // Match tool calls from the first turn to existing IDs - const firstTurnToolCalls = firstTurn.toolCalls.map((tc, i) => { - const existing = - existingToolCalls.find((e) => e.function?.name === tc.name) || existingToolCalls[i]; - return { - id: existing?.id || `legacy_tool_${i}`, - type: 'function' as const, - function: { name: tc.name, arguments: tc.args } - }; - }); - - // Update the existing message for the first turn - await DatabaseService.updateMessage(message.id, { - content: firstTurn.textBefore, - reasoningContent: reasoning || undefined, - toolCalls: firstTurnToolCalls.length > 0 ? JSON.stringify(firstTurnToolCalls) : '' - }); - - let currentParentId = message.id; - let toolCallIdCounter = existingToolCalls.length; - - // Create tool result messages for the first turn - for (let i = 0; i < firstTurn.toolCalls.length; i++) { - const tc = firstTurn.toolCalls[i]; - const toolCallId = firstTurnToolCalls[i]?.id || `legacy_tool_${i}`; - - const toolMsg = await DatabaseService.createMessageBranch( - { - convId, - type: MessageType.TEXT, - role: MessageRole.TOOL, - content: tc.result, - toolCallId, - timestamp: message.timestamp + i + 1, - toolCalls: '', - children: [] - }, - currentParentId - ); - currentParentId = toolMsg.id; - } - - // Create messages for subsequent turns - for (let turnIdx = 1; turnIdx < turns.length; turnIdx++) { - const turn = turns[turnIdx]; - - const turnToolCalls = turn.toolCalls.map((tc, i) => { - const idx = toolCallIdCounter + i; - const existing = existingToolCalls[idx]; - return { - id: existing?.id || `legacy_tool_${idx}`, - type: 'function' as const, - function: { name: tc.name, arguments: tc.args } - }; - }); - toolCallIdCounter += turn.toolCalls.length; - - // Create assistant message for this turn - const assistantMsg = await DatabaseService.createMessageBranch( - { - convId, - type: MessageType.TEXT, - role: MessageRole.ASSISTANT, - content: turn.textBefore, - timestamp: message.timestamp + turnIdx * 100, - toolCalls: turnToolCalls.length > 0 ? JSON.stringify(turnToolCalls) : '', - children: [], - model: message.model - }, - currentParentId - ); - currentParentId = assistantMsg.id; - - // Create tool result messages for this turn - for (let i = 0; i < turn.toolCalls.length; i++) { - const tc = turn.toolCalls[i]; - const toolCallId = turnToolCalls[i]?.id || `legacy_tool_${toolCallIdCounter + i}`; - - const toolMsg = await DatabaseService.createMessageBranch( - { - convId, - type: MessageType.TEXT, - role: MessageRole.TOOL, - content: tc.result, - toolCallId, - timestamp: message.timestamp + turnIdx * 100 + i + 1, - toolCalls: '', - children: [] - }, - currentParentId - ); - currentParentId = toolMsg.id; - } - } - - // Re-parent any children of the original message to the last created message - // (the original message's children list was the next user message or similar) - if (message.children.length > 0 && currentParentId !== message.id) { - for (const childId of message.children) { - // Skip children we just created (they were already properly parented) - const child = allMessages.find((m) => m.id === childId); - if (!child) continue; - // Only re-parent non-tool messages that were original children - if (child.role !== MessageRole.TOOL) { - await DatabaseService.updateMessage(childId, { parent: currentParentId }); - // Add to new parent's children - const newParent = await DatabaseService.getConversationMessages(convId).then((msgs) => - msgs.find((m) => m.id === currentParentId) - ); - if (newParent && !newParent.children.includes(childId)) { - await DatabaseService.updateMessage(currentParentId, { - children: [...newParent.children, childId] - }); - } - } - } - // Clear re-parented children from the original message - await DatabaseService.updateMessage(message.id, { children: [] }); - } - - migratedCount++; - } - - return migratedCount; -} - -/** - * @deprecated Part of legacy migration — remove with the migration module. - * Run the full migration across all conversations. - * This should be called once at app startup if migration is needed. - */ -export async function runLegacyMigration(): Promise { - if (!isMigrationNeeded()) return; - - if (import.meta.env.DEV && import.meta.env.VITE_DEBUG) - console.log('[Migration] Starting legacy message format migration...'); - - try { - const conversations = await DatabaseService.getAllConversations(); - let totalMigrated = 0; - - for (const conv of conversations) { - const count = await migrateConversation(conv.id); - totalMigrated += count; - } - - if (import.meta.env.DEV && import.meta.env.VITE_DEBUG) { - if (totalMigrated > 0) { - console.log( - `[Migration] Migrated ${totalMigrated} messages across ${conversations.length} conversations` - ); - } else { - console.log('[Migration] No legacy messages found, marking as done'); - } - } - - markMigrationDone(); - } catch (error) { - console.error('[Migration] Failed to migrate legacy messages:', error); - // Still mark as done to avoid infinite retry loops - markMigrationDone(); - } -}