Merge branch 'master' into xsn/qwen3-tts

This commit is contained in:
Xuan Son Nguyen
2026-07-30 17:04:49 +02:00
18 changed files with 440 additions and 531 deletions
+1 -1
View File
@@ -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++) {
+1 -1
View File
@@ -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);
//
+37 -2
View File
@@ -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<MTLBuffer> buf = [dev->mtl_device newBufferWithLength:1 options:MTLResourceStorageModePrivate];
id<MTLCommandBuffer> cmd_buf = [dev->mtl_queue commandBuffer];
{
id<MTLBlitCommandEncoder> 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];
}
}
+3
View File
@@ -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,
+3 -1
View File
@@ -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 {
+181 -21
View File
@@ -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<llm_graph_context> llama_model_mimo2::build_arch_graph(const llm_graph_params & params) const {
if (params.gtype == LLM_GRAPH_TYPE_DECODER_MTP) {
return std::make_unique<graph_mtp>(*this, params);
}
return std::make_unique<graph>(*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<llm_graph_input_embd>(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);
}
+21 -72
View File
@@ -2,7 +2,6 @@
#include "llama-kv-cache.h"
#include <cmath>
#include <vector>
#include <algorithm>
#include <cstdint>
// 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<float> rank(nblk);
std::vector<char> valid(nblk);
std::vector<int> 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<msa_params *>(&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),
+4
View File
@@ -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<llm_graph_context> build_arch_graph(const llm_graph_params & params) const override;
};
+3
View File
@@ -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)
+7 -3
View File
@@ -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);
+53
View File
@@ -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 <cstdint>
#include <mach/mach.h>
#include <mach/mach_host.h>
#include <unistd.h>
#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<uint64_t>(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;
}
@@ -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 @@
</Tooltip.Content>
</Tooltip.Root>
<DropdownMenu.Content align="start" class="w-52">
<DropdownMenu.Content
align="start"
class="w-52"
onCloseAutoFocus={(e) => {
if (suppressCloseAutoFocus) {
suppressCloseAutoFocus = false;
e.preventDefault();
}
}}
>
<ChatFormActionAddReasoningSubmenu />
<DropdownMenu.Separator />
@@ -148,7 +160,10 @@
<DropdownMenu.Item
class="flex cursor-pointer items-center gap-2"
onclick={onSystemPromptClick}
onclick={() => {
suppressCloseAutoFocus = true;
onSystemPromptClick?.();
}}
>
<MessageSquare class={ICON_CLASS_DEFAULT} />
@@ -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<HTMLTextAreaElement>('.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);
@@ -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;
+31 -26
View File
@@ -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<DatabaseConversation> {
static async createConversation(
name: string,
fields?: Partial<Omit<DatabaseConversation, 'id' | 'name' | 'lastModified'>>
): Promise<DatabaseConversation> {
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<Omit<DatabaseConversation, 'id'>>
): Promise<void> {
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<boolean> {
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<string, boolean>();
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;
+4 -3
View File
@@ -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);
+45 -31
View File
@@ -111,6 +111,9 @@ class ConversationsStore {
| ((messageId: string, updates: Partial<DatabaseMessage>) => void)
| null = null;
/** In-flight init run; shared by concurrent callers, reset on failure to allow retry */
private initPromise: Promise<void> | 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<void> {
if (!browser) return;
if (this.isInitialized) return;
init(): Promise<void> {
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<string> {
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<void> {
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];
}
}
-364
View File
@@ -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:
* <<<reasoning_content_start>>>...<<<reasoning_content_end>>>
* <<<AGENTIC_TOOL_CALL_START>>>...<<<AGENTIC_TOOL_CALL_END>>>
*
* 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<number> {
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<void> {
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();
}
}