diff --git a/conversion/glm.py b/conversion/glm.py index 895cefc22b89..9068df8e169b 100644 --- a/conversion/glm.py +++ b/conversion/glm.py @@ -237,6 +237,12 @@ def set_gguf_parameters(self): self.gguf_writer.add_indexer_head_count(self.hparams["index_n_heads"]) self.gguf_writer.add_indexer_key_length(self.hparams["index_head_dim"]) self.gguf_writer.add_indexer_top_k(self.hparams["index_topk"]) + if (indexer_types := self.hparams.get("indexer_types")) is not None: + if len(indexer_types) != self.hparams["num_hidden_layers"]: + raise ValueError("indexer_types must contain one entry per target layer") + if invalid := set(indexer_types) - {"full", "shared"}: + raise ValueError(f"unsupported indexer_types values: {sorted(invalid)}") + self.gguf_writer.add_indexer_types([indexer_type == "full" for indexer_type in indexer_types]) @ModelBase.register("SolarOpenForCausalLM") diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index 2745a7dbb29e..2c897ef1e5f0 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -1161,6 +1161,75 @@ void ggml_set_f32_nd(const struct ggml_tensor * tensor, int i0, int i1, int i2, // ggml_compute_forward_mul_mat +static bool ggml_compute_forward_mul_mat_precise_ref_type(enum ggml_type type) { + return ggml_is_quantized(type) || type == GGML_TYPE_F16 || type == GGML_TYPE_BF16; +} + +static float ggml_compute_forward_mul_mat_precise_ref_dot(const struct ggml_type_traits * traits, + const void * weight, + const float * activation, + float * row, + int64_t n_input) { + traits->to_float(weight, row, n_input); + + double sum = 0.0; + for (int64_t input = 0; input < n_input; ++input) { + sum += (double) row[input] * (double) activation[input]; + } + return (float) sum; +} + +static bool ggml_compute_forward_mul_mat_precise_ref(const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + + if (!params->use_ref || src1->type != GGML_TYPE_F32 || !ggml_compute_forward_mul_mat_precise_ref_type(src0->type)) { + return false; + } + + GGML_ASSERT(src0->ne[0] == src1->ne[0]); + GGML_ASSERT(dst->ne[0] == src0->ne[1]); + GGML_ASSERT(dst->ne[1] == src1->ne[1]); + GGML_ASSERT(dst->ne[2] == src1->ne[2]); + GGML_ASSERT(dst->ne[3] == src1->ne[3]); + GGML_ASSERT(src1->ne[2] % src0->ne[2] == 0); + GGML_ASSERT(src1->ne[3] % src0->ne[3] == 0); + + const struct ggml_type_traits * traits = ggml_get_type_traits(src0->type); + GGML_ASSERT(traits->to_float != NULL); + + const int64_t n_input = src0->ne[0]; + const int64_t n_output = dst->ne[0]; + const int64_t n_columns = dst->ne[1]; + const int64_t n_tasks = ggml_nelements(dst); + const int64_t r2 = src1->ne[2] / src0->ne[2]; + const int64_t r3 = src1->ne[3] / src0->ne[3]; + + float * row = (float *) malloc(n_input * sizeof(float)); + GGML_ASSERT(row != NULL); + + for (int64_t task = params->ith; task < n_tasks; task += params->nth) { + const int64_t output = task % n_output; + int64_t rest = task / n_output; + const int64_t column = rest % n_columns; + rest /= n_columns; + const int64_t i2 = rest % dst->ne[2]; + const int64_t i3 = rest / dst->ne[2]; + + const void * weight = + (const char *) src0->data + output * src0->nb[1] + (i2 / r2) * src0->nb[2] + (i3 / r3) * src0->nb[3]; + const float * activation = + (const float *) ((const char *) src1->data + column * src1->nb[1] + i2 * src1->nb[2] + i3 * src1->nb[3]); + float * result = (float *) ((char *) dst->data + output * dst->nb[0] + column * dst->nb[1] + i2 * dst->nb[2] + + i3 * dst->nb[3]); + *result = ggml_compute_forward_mul_mat_precise_ref_dot(traits, weight, activation, row, n_input); + } + + free(row); + return true; +} + static void ggml_compute_forward_mul_mat_one_chunk( const struct ggml_compute_params * params, struct ggml_tensor * dst, @@ -1254,6 +1323,9 @@ static void ggml_compute_forward_mul_mat_one_chunk( void ggml_compute_forward_mul_mat( const struct ggml_compute_params * params, struct ggml_tensor * dst) { + if (ggml_compute_forward_mul_mat_precise_ref(params, dst)) { + return; + } const struct ggml_tensor * src0 = dst->src[0]; const struct ggml_tensor * src1 = dst->src[1]; @@ -1523,6 +1595,56 @@ static void ggml_compute_forward_mul_mat_id_one_chunk( } } +static bool ggml_compute_forward_mul_mat_id_precise_ref(const struct ggml_compute_params * params, + struct ggml_tensor * dst) { + const struct ggml_tensor * src0 = dst->src[0]; + const struct ggml_tensor * src1 = dst->src[1]; + const struct ggml_tensor * ids = dst->src[2]; + + if (!params->use_ref || src1->type != GGML_TYPE_F32 || !ggml_compute_forward_mul_mat_precise_ref_type(src0->type)) { + return false; + } + + GGML_ASSERT(src0->ne[0] == src1->ne[0]); + GGML_ASSERT(dst->ne[0] == src0->ne[1]); + GGML_ASSERT(dst->ne[1] == ids->ne[0]); + GGML_ASSERT(dst->ne[2] == ids->ne[1]); + GGML_ASSERT(src1->ne[1] == 1 || src1->ne[1] == ids->ne[0]); + GGML_ASSERT(src1->ne[2] == ids->ne[1]); + GGML_ASSERT(src0->ne[3] == 1 && src1->ne[3] == 1 && dst->ne[3] == 1); + + const struct ggml_type_traits * traits = ggml_get_type_traits(src0->type); + GGML_ASSERT(traits->to_float != NULL); + + const int64_t n_input = src0->ne[0]; + const int64_t n_output = src0->ne[1]; + const int64_t n_ids = ids->ne[0]; + const int64_t n_tokens = ids->ne[1]; + const int64_t n_tasks = n_output * n_ids * n_tokens; + + float * row = (float *) malloc(n_input * sizeof(float)); + GGML_ASSERT(row != NULL); + + for (int64_t task = params->ith; task < n_tasks; task += params->nth) { + const int64_t output = task % n_output; + const int64_t rest = task / n_output; + const int64_t id = rest % n_ids; + const int64_t token = rest / n_ids; + const int64_t src1_id = src1->ne[1] == 1 ? 0 : id; + const int32_t expert = *(const int32_t *) ((const char *) ids->data + id * ids->nb[0] + token * ids->nb[1]); + GGML_ASSERT(expert >= 0 && expert < src0->ne[2]); + + const void * weight = (const char *) src0->data + output * src0->nb[1] + expert * src0->nb[2]; + const float * activation = + (const float *) ((const char *) src1->data + src1_id * src1->nb[1] + token * src1->nb[2]); + float * result = (float *) ((char *) dst->data + output * dst->nb[0] + id * dst->nb[1] + token * dst->nb[2]); + *result = ggml_compute_forward_mul_mat_precise_ref_dot(traits, weight, activation, row, n_input); + } + + free(row); + return true; +} + static void * incr_ptr_aligned(void ** p, size_t size, size_t align) { void * ptr = *p; @@ -1534,6 +1656,9 @@ static void * incr_ptr_aligned(void ** p, size_t size, size_t align) { static void ggml_compute_forward_mul_mat_id( const struct ggml_compute_params * params, struct ggml_tensor * dst) { + if (ggml_compute_forward_mul_mat_id_precise_ref(params, dst)) { + return; + } const struct ggml_tensor * src0 = dst->src[0]; const struct ggml_tensor * src1 = dst->src[1]; diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 15290c3d1091..8da1649e3ba5 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -638,6 +638,24 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_gated_delta_net( return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer(ggml_metal_library_t lib, + const ggml_tensor * op) { + GGML_ASSERT(op->op == GGML_OP_LIGHTNING_INDEXER); + + char base[256]; + char name[256]; + + snprintf(base, sizeof(base), "kernel_lightning_indexer_%s", ggml_type_name(op->src[1]->type)); + snprintf(name, sizeof(name), "%s", base); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr); + } + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_solve_tri(ggml_metal_library_t lib, const ggml_tensor * op) { char base[256]; char name[256]; diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index 9d4aca121595..832c95623188 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -129,6 +129,8 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_rwkv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_gated_delta_net (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer(ggml_metal_library_t lib, + const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_solve_tri (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mv_ext (ggml_metal_library_t lib, const struct ggml_tensor * op, int nsg, int nxpsg, int r1ptg); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_mul_mm (ggml_metal_library_t lib, const struct ggml_tensor * op); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 5d29250f654b..106527737955 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1272,6 +1272,28 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te return true; case GGML_OP_GATED_DELTA_NET: return has_simdgroup_reduction && op->src[2]->ne[0] % 32 == 0; + case GGML_OP_LIGHTNING_INDEXER: + { + const enum ggml_type k_type = op->src[1]->type; + const bool k_type_supported = + k_type == GGML_TYPE_F32 || + k_type == GGML_TYPE_F16 || + (k_type == GGML_TYPE_BF16 && has_bfloat) || + k_type == GGML_TYPE_Q8_0 || + k_type == GGML_TYPE_Q5_1 || + k_type == GGML_TYPE_Q5_0 || + k_type == GGML_TYPE_Q4_1 || + k_type == GGML_TYPE_Q4_0 || + k_type == GGML_TYPE_IQ4_NL; + + return op->type == GGML_TYPE_F32 && + op->src[0]->type == GGML_TYPE_F32 && + k_type_supported && + op->src[2]->type == GGML_TYPE_F32 && + op->src[3]->type == GGML_TYPE_F16 && + ggml_is_contiguous_rows(op->src[0]) && + ggml_is_contiguous_rows(op->src[1]); + } case GGML_OP_SOLVE_TRI: case GGML_OP_MUL_MAT: case GGML_OP_MUL_MAT_ID: diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index d6761023b76c..c644493c02a9 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -927,6 +927,49 @@ typedef struct { uint64_t nb3; } ggml_metal_kargs_gated_delta_net; +typedef struct { + int32_t ne00; + int32_t ne01; + int32_t ne02; + int32_t ne03; + uint64_t nb00; + uint64_t nb01; + uint64_t nb02; + uint64_t nb03; + int32_t ne10; + int32_t ne11; + int32_t ne12; + int32_t ne13; + uint64_t nb10; + uint64_t nb11; + uint64_t nb12; + uint64_t nb13; + int32_t ne20; + int32_t ne21; + int32_t ne22; + int32_t ne23; + uint64_t nb20; + uint64_t nb21; + uint64_t nb22; + uint64_t nb23; + int32_t ne30; + int32_t ne31; + int32_t ne32; + int32_t ne33; + uint64_t nb30; + uint64_t nb31; + uint64_t nb32; + uint64_t nb33; + int32_t ne0; + int32_t ne1; + int32_t ne2; + int32_t ne3; + uint64_t nb0; + uint64_t nb1; + uint64_t nb2; + uint64_t nb3; +} ggml_metal_kargs_lightning_indexer; + typedef struct { int32_t ne00; int32_t ne01; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 45909c4777b5..1ba320a832d8 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -337,6 +337,11 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) { { n_fuse = ggml_metal_op_gated_delta_net(ctx, idx); } break; + case GGML_OP_LIGHTNING_INDEXER: + { + n_fuse = ggml_metal_op_lightning_indexer(ctx, idx); + } + break; case GGML_OP_SOLVE_TRI: { n_fuse = ggml_metal_op_solve_tri(ctx, idx); @@ -1674,6 +1679,83 @@ int ggml_metal_op_gated_delta_net(ggml_metal_op_t ctx, int idx) { return 1; } +int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { + ggml_tensor * op = ctx->node(idx); + + ggml_metal_library_t lib = ctx->lib; + ggml_metal_encoder_t enc = ctx->enc; + + GGML_TENSOR_LOCALS(int32_t, ne0, op->src[0], ne); + GGML_TENSOR_LOCALS(uint64_t, nb0, op->src[0], nb); + GGML_TENSOR_LOCALS(int32_t, ne1, op->src[1], ne); + GGML_TENSOR_LOCALS(uint64_t, nb1, op->src[1], nb); + GGML_TENSOR_LOCALS(int32_t, ne2, op->src[2], ne); + GGML_TENSOR_LOCALS(uint64_t, nb2, op->src[2], nb); + GGML_TENSOR_LOCALS(int32_t, ne3, op->src[3], ne); + GGML_TENSOR_LOCALS(uint64_t, nb3, op->src[3], nb); + GGML_TENSOR_LOCALS(int32_t, ne, op, ne); + GGML_TENSOR_LOCALS(uint64_t, nb, op, nb); + + ggml_metal_kargs_lightning_indexer args = { + /*.ne00 =*/ne00, + /*.ne01 =*/ne01, + /*.ne02 =*/ne02, + /*.ne03 =*/ne03, + /*.nb00 =*/nb00, + /*.nb01 =*/nb01, + /*.nb02 =*/nb02, + /*.nb03 =*/nb03, + /*.ne10 =*/ne10, + /*.ne11 =*/ne11, + /*.ne12 =*/ne12, + /*.ne13 =*/ne13, + /*.nb10 =*/nb10, + /*.nb11 =*/nb11, + /*.nb12 =*/nb12, + /*.nb13 =*/nb13, + /*.ne20 =*/ne20, + /*.ne21 =*/ne21, + /*.ne22 =*/ne22, + /*.ne23 =*/ne23, + /*.nb20 =*/nb20, + /*.nb21 =*/nb21, + /*.nb22 =*/nb22, + /*.nb23 =*/nb23, + /*.ne30 =*/ne30, + /*.ne31 =*/ne31, + /*.ne32 =*/ne32, + /*.ne33 =*/ne33, + /*.nb30 =*/nb30, + /*.nb31 =*/nb31, + /*.nb32 =*/nb32, + /*.nb33 =*/nb33, + /*.ne0 =*/ne0, + /*.ne1 =*/ne1, + /*.ne2 =*/ne2, + /*.ne3 =*/ne3, + /*.nb0 =*/nb0, + /*.nb1 =*/nb1, + /*.nb2 =*/nb2, + /*.nb3 =*/nb3, + }; + + auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(lib, op); + const int nth = std::min(64, ggml_metal_pipeline_max_theads_per_threadgroup(pipeline)); + int ida = 0; + + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), ida++); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[0]), ida++); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[1]), ida++); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[2]), ida++); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op->src[3]), ida++); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), ida++); + + ggml_metal_encoder_dispatch_threadgroups(enc, (ne0 + nth - 1) / nth, ne1, ne3, nth, 1, 1); + + return 1; +} + int ggml_metal_op_solve_tri(ggml_metal_op_t ctx, int idx) { ggml_tensor * op = ctx->node(idx); diff --git a/ggml/src/ggml-metal/ggml-metal-ops.h b/ggml/src/ggml-metal/ggml-metal-ops.h index 0bebd836a185..61df873daa06 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.h +++ b/ggml/src/ggml-metal/ggml-metal-ops.h @@ -59,6 +59,7 @@ int ggml_metal_op_ssm_conv (ggml_metal_op_t ctx, int idx); int ggml_metal_op_ssm_scan (ggml_metal_op_t ctx, int idx); int ggml_metal_op_rwkv (ggml_metal_op_t ctx, int idx); int ggml_metal_op_gated_delta_net (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx); int ggml_metal_op_solve_tri (ggml_metal_op_t ctx, int idx); int ggml_metal_op_set (ggml_metal_op_t ctx, int idx); int ggml_metal_op_cpy (ggml_metal_op_t ctx, int idx); diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal index 6b6f9fd870c2..ca31701fe2df 100644 --- a/ggml/src/ggml-metal/ggml-metal.metal +++ b/ggml/src/ggml-metal/ggml-metal.metal @@ -2819,6 +2819,116 @@ template [[host_name("kernel_gated_delta_net_f32_2")]] kernel kernel_gated_delta template [[host_name("kernel_gated_delta_net_f32_4")]] kernel kernel_gated_delta_net_t kernel_gated_delta_net_impl; #endif +static inline float lightning_indexer_mask( + constant ggml_metal_kargs_lightning_indexer & args, + device const char * mask, + int i_kv, + int i_batch, + int i_stream) { + const int i_mask_stream = i_stream % args.ne33; + device const half * value = (device const half *) ( + mask + i_kv * args.nb30 + i_batch * args.nb31 + i_mask_stream * args.nb33); + return float(*value); +} + +template +kernel void kernel_lightning_indexer_impl( + constant ggml_metal_kargs_lightning_indexer & args, + device const char * q, + device const char * k, + device const char * weights, + device const char * mask, + device char * dst, + uint3 gid [[thread_position_in_grid]]) { + const int i_kv = gid.x; + const int i_batch = gid.y; + const int i_stream = gid.z; + + if (i_kv >= args.ne0 || i_batch >= args.ne1 || i_stream >= args.ne3) { + return; + } + + float score = 0.0f; + for (int i_head = 0; i_head < args.ne01; ++i_head) { + float qk = 0.0f; + for (int i_embd = 0; i_embd < args.ne00; ++i_embd) { + device const float * q_ptr = (device const float *) ( + q + i_embd * args.nb00 + i_head * args.nb01 + i_batch * args.nb02 + i_stream * args.nb03); + device const K * k_ptr = (device const K *) ( + k + i_embd * args.nb10 + i_kv * args.nb12 + i_stream * args.nb13); + qk += *q_ptr * float(*k_ptr); + } + + device const float * weight_ptr = (device const float *) ( + weights + i_head * args.nb20 + i_batch * args.nb21 + i_stream * args.nb23); + score += max(qk, 0.0f) * *weight_ptr; + } + + device float * dst_ptr = (device float *) ( + dst + i_kv * args.nb0 + i_batch * args.nb1 + i_stream * args.nb3); + *dst_ptr = score + lightning_indexer_mask(args, mask, i_kv, i_batch, i_stream); +} + +template +kernel void kernel_lightning_indexer_quantized_impl( + constant ggml_metal_kargs_lightning_indexer & args, + device const char * q, + device const char * k, + device const char * weights, + device const char * mask, + device char * dst, + uint3 gid [[thread_position_in_grid]]) { + const int i_kv = gid.x; + const int i_batch = gid.y; + const int i_stream = gid.z; + + if (i_kv >= args.ne0 || i_batch >= args.ne1 || i_stream >= args.ne3) { + return; + } + + device const char * k_row = k + i_kv * args.nb12 + i_stream * args.nb13; + + float score = 0.0f; + for (int i_head = 0; i_head < args.ne01; ++i_head) { + float qk = 0.0f; + for (int i_embd = 0; i_embd < args.ne00; i_embd += 4) { + device const K * k_block = (device const K *) ( + k_row + (i_embd / epb) * args.nb10); + float4 k_values; + deq_t4(k_block, short((i_embd % epb) / 4), k_values); + + device const float4 * q_values = (device const float4 *) ( + q + i_embd * args.nb00 + i_head * args.nb01 + + i_batch * args.nb02 + i_stream * args.nb03); + qk += dot(*q_values, k_values); + } + + device const float * weight_ptr = (device const float *) ( + weights + i_head * args.nb20 + i_batch * args.nb21 + i_stream * args.nb23); + score += max(qk, 0.0f) * *weight_ptr; + } + + device float * dst_ptr = (device float *) ( + dst + i_kv * args.nb0 + i_batch * args.nb1 + i_stream * args.nb3); + *dst_ptr = score + lightning_indexer_mask(args, mask, i_kv, i_batch, i_stream); +} + +typedef decltype(kernel_lightning_indexer_impl) kernel_lightning_indexer_t; +typedef decltype(kernel_lightning_indexer_quantized_impl) + kernel_lightning_indexer_quantized_t; + +template [[host_name("kernel_lightning_indexer_f32")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer_impl; +template [[host_name("kernel_lightning_indexer_f16")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer_impl; +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_lightning_indexer_bf16")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer_impl; +#endif +template [[host_name("kernel_lightning_indexer_q8_0")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized_impl; +template [[host_name("kernel_lightning_indexer_q5_1")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized_impl; +template [[host_name("kernel_lightning_indexer_q5_0")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized_impl; +template [[host_name("kernel_lightning_indexer_q4_1")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized_impl; +template [[host_name("kernel_lightning_indexer_q4_0")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized_impl; +template [[host_name("kernel_lightning_indexer_iq4_nl")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized_impl; + constant short FC_solve_tri_nsg [[function_constant(FC_SOLVE_TRI + 0)]]; constant short FC_solve_tri_n [[function_constant(FC_SOLVE_TRI + 1)]]; constant short FC_solve_tri_k [[function_constant(FC_SOLVE_TRI + 2)]]; diff --git a/gguf-py/gguf/constants.py b/gguf-py/gguf/constants.py index 869e436acd5c..7038950e0daa 100644 --- a/gguf-py/gguf/constants.py +++ b/gguf-py/gguf/constants.py @@ -200,6 +200,7 @@ class Indexer: HEAD_COUNT = "{arch}.attention.indexer.head_count" KEY_LENGTH = "{arch}.attention.indexer.key_length" TOP_K = "{arch}.attention.indexer.top_k" + TYPES = "{arch}.attention.indexer.types" class HyperConnection: COUNT = "{arch}.hyper_connection.count" diff --git a/gguf-py/gguf/gguf_writer.py b/gguf-py/gguf/gguf_writer.py index 1e277f0687c5..bb21596701d4 100644 --- a/gguf-py/gguf/gguf_writer.py +++ b/gguf-py/gguf/gguf_writer.py @@ -793,6 +793,10 @@ def add_indexer_key_length(self, length: int) -> None: def add_indexer_top_k(self, top_k: int) -> None: self.add_uint32(Keys.Attention.Indexer.TOP_K.format(arch=self.arch), top_k) + def add_indexer_types(self, value: Sequence[bool]) -> None: + key = Keys.Attention.Indexer.TYPES.format(arch=self.arch) + self.add_array(key, value) + def add_max_alibi_bias(self, bias: float) -> None: self.add_float32(Keys.Attention.MAX_ALIBI_BIAS.format(arch=self.arch), bias) diff --git a/src/llama-arch.cpp b/src/llama-arch.cpp index b890e66fcf6e..58bd043d25fe 100644 --- a/src/llama-arch.cpp +++ b/src/llama-arch.cpp @@ -251,6 +251,7 @@ static const std::map LLM_KV_NAMES = { { LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, "%s.attention.indexer.head_count" }, { LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, "%s.attention.indexer.key_length" }, { LLM_KV_ATTENTION_INDEXER_TOP_K, "%s.attention.indexer.top_k" }, + { LLM_KV_ATTENTION_INDEXER_TYPES, "%s.attention.indexer.types" }, { LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT, "%s.attention.output_group_count" }, { LLM_KV_ATTENTION_OUTPUT_LORA_RANK, "%s.attention.output_lora_rank" }, { LLM_KV_ATTENTION_COMPRESS_ROPE_FREQ_BASE, "%s.attention.compress_rope_freq_base" }, diff --git a/src/llama-arch.h b/src/llama-arch.h index a4f5091e7170..081f1a391e37 100644 --- a/src/llama-arch.h +++ b/src/llama-arch.h @@ -256,6 +256,7 @@ enum llm_kv { LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, LLM_KV_ATTENTION_INDEXER_TOP_K, + LLM_KV_ATTENTION_INDEXER_TYPES, LLM_KV_ATTENTION_OUTPUT_GROUP_COUNT, LLM_KV_ATTENTION_OUTPUT_LORA_RANK, LLM_KV_ATTENTION_COMPRESS_ROPE_FREQ_BASE, diff --git a/src/llama-graph.cpp b/src/llama-graph.cpp index c8ecb0a2854c..fd0407e3a36f 100644 --- a/src/llama-graph.cpp +++ b/src/llama-graph.cpp @@ -2820,37 +2820,54 @@ ggml_tensor * llm_graph_context::build_attn( const auto & kq_mask = inp->get_kq_mask_mla(); - // prepare new kq mask - starts filled with -INFINITY - ggml_tensor * kq_mask_all = ggml_fill(ctx0, kq_mask, -INFINITY); + ggml_tensor * q = q_cur; + ggml_tensor * k = mctx_cur->get_k(ctx0, il); + ggml_tensor * v = + ggml_view_4d(ctx0, k, v_cur->ne[0], k->ne[1], k->ne[2], k->ne[3], k->nb[1], k->nb[2], k->nb[3], 0); - // reshape KQ mask into tensor with rows of size 1: - // [n_kv, n_batch, 1, n_stream] -> [1, n_kv, n_batch, n_stream] - kq_mask_all = ggml_view_4d(ctx0, kq_mask_all, 1, kq_mask_all->ne[0], kq_mask_all->ne[1], kq_mask_all->ne[3], kq_mask_all->nb[0], kq_mask_all->nb[1], kq_mask_all->nb[2], 0); + ggml_tensor * attn_mask = nullptr; - // reshape top_k indices: [n_top_k, n_batch, 1, n_stream] -> [n_top_k, n_batch, n_stream, 1] - ggml_tensor * top_k_3d = ggml_view_4d(ctx0, top_k, top_k->ne[0], top_k->ne[1], top_k->ne[3], 1, top_k->nb[1], top_k->nb[2], top_k->ne[3]*top_k->nb[3], 0); + const int64_t n_stream = k->ne[3]; + const bool compact_decode = cparams.flash_attn && top_k->ne[1] == 1 && q_cur->ne[2] == n_stream; + if (compact_decode) { + GGML_ASSERT(k->ne[1] == 1 && v->ne[1] == 1); + GGML_ASSERT(kq_mask->type == GGML_TYPE_F16); - // prepare zero-filled tensor with rows of size 1: [1, n_top_k, n_batch, n_stream] - // this will be our source of zero values for unmasking top k mask elements - ggml_tensor * zeros = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, top_k_3d->ne[0], top_k_3d->ne[1], top_k_3d->ne[2]); - zeros = ggml_fill(ctx0, zeros, 0.0f); + ggml_tensor * top_k_rows = ggml_reshape_3d(ctx0, top_k, top_k->ne[0], 1, n_stream); - // modify KQ mask by unmasking elements that are in top_k indices - // ggml_set_rows([1, n_kv, n_batch, n_stream], [1, n_top_k, n_batch, n_stream], [n_top_k, n_batch, n_stream, 1]) - ggml_tensor * kq_mask_top_k = ggml_set_rows(ctx0, kq_mask_all, zeros, top_k_3d); + k = ggml_get_rows(ctx0, ggml_permute(ctx0, k, 0, 2, 1, 3), top_k_rows); + v = ggml_get_rows(ctx0, ggml_permute(ctx0, v, 0, 2, 1, 3), top_k_rows); + cb(k, "k_selected", il); + cb(v, "v_selected", il); - // reshape to restore the original shape of KQ mask: - // [1, n_kv, n_batch, n_stream] -> [n_kv, n_batch, 1, n_stream] - kq_mask_top_k = ggml_view_4d(ctx0, kq_mask_top_k, kq_mask_top_k->ne[1], kq_mask_top_k->ne[2], 1, kq_mask_top_k->ne[3], kq_mask_top_k->nb[2], kq_mask_top_k->nb[3], kq_mask_top_k->nb[3], 0); + k = ggml_reshape_4d(ctx0, k, k->ne[0], 1, k->ne[1], n_stream); + v = ggml_reshape_4d(ctx0, v, v->ne[0], 1, v->ne[1], n_stream); - // combine with the original kq mask - kq_mask_top_k = ggml_add(ctx0, kq_mask_top_k, kq_mask); + ggml_tensor * mask_rows = ggml_reshape_4d(ctx0, kq_mask, 1, kq_mask->ne[0], 1, n_stream); + attn_mask = ggml_get_rows(ctx0, mask_rows, top_k_rows); + attn_mask = ggml_reshape_4d(ctx0, attn_mask, top_k->ne[0], 1, 1, n_stream); + attn_mask = ggml_cast(ctx0, attn_mask, GGML_TYPE_F16); + cb(attn_mask, "kq_mask_selected", il); + } else { + // Keep the dense path for prefill and non-flash attention. + ggml_tensor * kq_mask_all = ggml_fill(ctx0, kq_mask, -INFINITY); + kq_mask_all = ggml_view_4d(ctx0, kq_mask_all, 1, kq_mask_all->ne[0], kq_mask_all->ne[1], kq_mask_all->ne[3], + kq_mask_all->nb[0], kq_mask_all->nb[1], kq_mask_all->nb[2], 0); - ggml_tensor * q = q_cur; - ggml_tensor * k = mctx_cur->get_k(ctx0, il); - ggml_tensor * v = ggml_view_4d(ctx0, k, v_cur->ne[0], k->ne[1], k->ne[2], k->ne[3], k->nb[1], k->nb[2], k->nb[3], 0); + ggml_tensor * top_k_3d = ggml_view_4d(ctx0, top_k, top_k->ne[0], top_k->ne[1], top_k->ne[3], 1, top_k->nb[1], + top_k->nb[2], top_k->ne[3] * top_k->nb[3], 0); + + ggml_tensor * zeros = + ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, 1, top_k_3d->ne[0], top_k_3d->ne[1], top_k_3d->ne[2]); + zeros = ggml_fill(ctx0, zeros, 0.0f); + + attn_mask = ggml_set_rows(ctx0, kq_mask_all, zeros, top_k_3d); + attn_mask = ggml_view_4d(ctx0, attn_mask, attn_mask->ne[1], attn_mask->ne[2], 1, attn_mask->ne[3], + attn_mask->nb[2], attn_mask->nb[3], attn_mask->nb[3], 0); + attn_mask = ggml_add(ctx0, attn_mask, kq_mask); + } - ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, kq_mask_top_k, sinks, v_mla, kq_scale, il); + ggml_tensor * cur = build_attn_mha(q, k, v, kq_b, attn_mask, sinks, v_mla, kq_scale, il); cb(cur, "kqv_out", il); if (wo) { diff --git a/src/llama-hparams.cpp b/src/llama-hparams.cpp index 9d0683d2fec4..846d4c69a626 100644 --- a/src/llama-hparams.cpp +++ b/src/llama-hparams.cpp @@ -248,6 +248,14 @@ bool llama_hparams::is_mla() const { return n_embd_head_k_mla_impl != 0 && n_embd_head_v_mla_impl != 0; } +bool llama_hparams::is_indexer_full(uint32_t il) const { + if (il < n_layer()) { + return is_indexer_full_impl[il]; + } + + GGML_ABORT("%s: il (%u) out of bounds (n_layer: %u)\n", __func__, il, n_layer()); +} + uint32_t llama_hparams::n_embd_head_k_mla() const { return is_mla() ? n_embd_head_k_mla_impl : n_embd_head_k(); } diff --git a/src/llama-hparams.h b/src/llama-hparams.h index 8be5f28f39e6..747754fc0d0b 100644 --- a/src/llama-hparams.h +++ b/src/llama-hparams.h @@ -227,6 +227,10 @@ struct llama_hparams { uint32_t indexer_head_size = 0; uint32_t indexer_top_k = 0; + // Indexer is "full" (1) or "shared" (0) + // Shared indexers reuse top-k from previous full layer + std::array is_indexer_full_impl; + // DeepSeek-V4 uint32_t dsv4_o_group_count = 0; uint32_t dsv4_o_lora_rank = 0; @@ -302,6 +306,8 @@ struct llama_hparams { bool is_swa(uint32_t il) const; + bool is_indexer_full(uint32_t il) const; + void set_recr_pattern(uint32_t n_pattern, bool dense_first = false); // whether or not the given layer is recurrent (for hybrid models) diff --git a/src/llama-kv-cache-dsa.cpp b/src/llama-kv-cache-dsa.cpp index 241c50365a13..fa8c8f0cb9be 100644 --- a/src/llama-kv-cache-dsa.cpp +++ b/src/llama-kv-cache-dsa.cpp @@ -25,12 +25,18 @@ llama_kv_cache_dsa::llama_kv_cache_dsa( llama_swa_type swa_type, const layer_filter_cb & filter, const layer_reuse_cb & reuse) : - hparams_lid(model.hparams), n_stream(unified ? 1 : n_seq_max) { + hparams_mla(model.hparams), hparams_lid(model.hparams), n_stream(unified ? 1 : n_seq_max) { + + if (model.arch == LLM_ARCH_GLM_DSA) { + // GLM caches one compressed MLA key per token, not one key per attention head. + std::fill(hparams_mla.n_head_kv_arr.begin(), hparams_mla.n_head_kv_arr.end(), 1); + hparams_mla.n_embd_head_k_full = model.hparams.n_lora_kv + model.hparams.n_rot(); + } LLAMA_LOG_INFO("%s: creating main KV cache, size = %u cells\n", __func__, kv_size); kv_mla = std::make_unique( - model, model.hparams, type_k, type_v, + model, hparams_mla, type_k, type_v, v_trans, offload, unified, kv_size, n_seq_max, n_pad, n_swa, swa_type, nullptr, filter, reuse, nullptr); @@ -42,7 +48,7 @@ llama_kv_cache_dsa::llama_kv_cache_dsa( // DSA lightning indexer uses MQA with single key head std::fill(hparams_lid.n_head_kv_arr.begin(), hparams_lid.n_head_kv_arr.end(), 1); hparams_lid.n_embd_head_k_full = model.hparams.indexer_head_size; - hparams_lid.rope_type = LLAMA_ROPE_TYPE_NEOX; + hparams_lid.rope_type = model.arch == LLM_ARCH_GLM_DSA ? LLAMA_ROPE_TYPE_NORM : LLAMA_ROPE_TYPE_NEOX; LLAMA_LOG_INFO("%s: creating indexer KV cache, size = %u cells\n", __func__, kv_size); diff --git a/src/llama-kv-cache-dsa.h b/src/llama-kv-cache-dsa.h index e2b330993b84..77c18e5e274d 100644 --- a/src/llama-kv-cache-dsa.h +++ b/src/llama-kv-cache-dsa.h @@ -72,7 +72,8 @@ class llama_kv_cache_dsa : public llama_memory_i { llama_kv_cache * get_lid() const; private: - // we keep indexer KV cache hparams instance here as llama_kv_cache stores only reference to it + // We keep private hparams instances because llama_kv_cache stores references to them. + llama_hparams hparams_mla; llama_hparams hparams_lid; const uint32_t n_stream = 1; diff --git a/src/llama-kv-cache.cpp b/src/llama-kv-cache.cpp index e70583e64152..1e209f22311f 100644 --- a/src/llama-kv-cache.cpp +++ b/src/llama-kv-cache.cpp @@ -322,9 +322,9 @@ llama_kv_cache::llama_kv_cache( ggml_is_quantized(type_k) && hparams.n_embd_head_k() % 64 == 0; - // always create Hadamard rotation tensors for DeepSeek lightning indexers - if ((model.arch == LLM_ARCH_DEEPSEEK32 || model.arch == LLM_ARCH_DEEPSEEK4) && - hparams.n_embd_head_k_full == hparams.indexer_head_size) { + // always create Hadamard rotation tensors for DSA lightning indexers + if ((model.arch == LLM_ARCH_DEEPSEEK32 || model.arch == LLM_ARCH_DEEPSEEK4 || model.arch == LLM_ARCH_GLM_DSA) && + hparams.n_embd_head_k_full == hparams.indexer_head_size) { attn_rot_k = true; } diff --git a/src/llama-model-saver.cpp b/src/llama-model-saver.cpp index a3928523ba8d..cf568e999152 100644 --- a/src/llama-model-saver.cpp +++ b/src/llama-model-saver.cpp @@ -11,6 +11,7 @@ #include #include +#include bool llama_model_saver_supports_arch(llm_arch arch) { switch (arch) { @@ -280,6 +281,11 @@ void llama_model_saver::add_kv_from_model() { add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, hparams.indexer_n_head); add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, hparams.indexer_head_size); add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K, hparams.indexer_top_k); + if (model->arch == LLM_ARCH_GLM_DSA) { + const std::vector indexer_types( + hparams.is_indexer_full_impl.begin(), hparams.is_indexer_full_impl.begin() + hparams.n_layer()); + add_kv(LLM_KV_ATTENTION_INDEXER_TYPES, indexer_types); + } add_kv(LLM_KV_ATTENTION_RECURRENT_LAYERS, hparams.is_recr_impl, true); const float rope_scaling_factor = hparams.rope_freq_scale_train == 1.0f ? 0.0f : 1.0f/hparams.rope_freq_scale_train; diff --git a/src/llama-model.cpp b/src/llama-model.cpp index d87481381e46..5e253d13d529 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -2048,6 +2048,7 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, res = nullptr; } break; case LLM_ARCH_DEEPSEEK32: + case LLM_ARCH_GLM_DSA: { res = new llama_kv_cache_dsa( *this, diff --git a/src/models/glm-dsa.cpp b/src/models/glm-dsa.cpp index 32fe6def6f3c..f1a6439be800 100644 --- a/src/models/glm-dsa.cpp +++ b/src/models/glm-dsa.cpp @@ -1,5 +1,55 @@ +#include "llama-kv-cache-dsa.h" #include "models.h" +// https://huggingface.co/zai-org/GLM-5.2/blob/main/config.json#L26 +const std::array GLM_DSA_DEFAULT_INDEXER_TYPES = { + 1, 1, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, + 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, +}; + +static void load_indexer_types(llama_model_loader & ml, llama_hparams & hparams) { + hparams.is_indexer_full_impl = GLM_DSA_DEFAULT_INDEXER_TYPES; + + const std::string key = ml.llm_kv(LLM_KV_ATTENTION_INDEXER_TYPES); + const int key_id = gguf_find_key(ml.metadata, key.c_str()); + if (key_id < 0) { + return; + } + if (gguf_get_kv_type(ml.metadata, key_id) != GGUF_TYPE_ARRAY) { + throw std::runtime_error(key + " must be an array"); + } + + const enum gguf_type element_type = gguf_get_arr_type(ml.metadata, key_id); + if (element_type == GGUF_TYPE_STRING) { + std::vector values; + ml.get_arr(LLM_KV_ATTENTION_INDEXER_TYPES, values); + if (values.size() != hparams.n_layer()) { + throw std::runtime_error(format("%s has wrong array length; expected %u, got %u", key.c_str(), + hparams.n_layer(), static_cast(values.size()))); + } + for (uint32_t il = 0; il < hparams.n_layer(); ++il) { + if (values[il] == "full") { + hparams.is_indexer_full_impl[il] = 1; + } else if (values[il] == "shared") { + hparams.is_indexer_full_impl[il] = 0; + } else { + throw std::runtime_error(format("%s[%u] must be 'full' or 'shared'", key.c_str(), il)); + } + } + return; + } + + if (element_type != GGUF_TYPE_BOOL && element_type != GGUF_TYPE_UINT32 && element_type != GGUF_TYPE_INT32) { + throw std::runtime_error(key + " must contain strings, bools, or 32-bit integers"); + } + ml.get_key_or_arr(LLM_KV_ATTENTION_INDEXER_TYPES, hparams.is_indexer_full_impl, hparams.n_layer(), false); + for (uint32_t il = 0; il < hparams.n_layer(); ++il) { + if (hparams.is_indexer_full_impl[il] > 1) { + throw std::runtime_error(format("%s[%u] must be 0 or 1", key.c_str(), il)); + } + } +} + void llama_model_glm_dsa::load_arch_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_EXPERT_FEED_FORWARD_LENGTH, hparams.n_ff_exp); ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); @@ -36,8 +86,13 @@ void llama_model_glm_dsa::load_arch_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_NEXTN_PREDICT_LAYERS, hparams.n_layer_nextn, false); GGML_ASSERT(hparams.n_layer_nextn < hparams.n_layer_all && "n_layer_nextn must be < n_layer_impl"); + load_indexer_types(ml, hparams); + if (!hparams.is_indexer_full(0)) { + throw std::runtime_error("GLM-DSA indexer schedule must start with a full layer"); + } + switch (hparams.n_layer()) { - case 79: type = LLM_TYPE_744B_A40B; break; + case 78: type = LLM_TYPE_744B_A40B; break; default: type = LLM_TYPE_UNKNOWN; } } @@ -150,3 +205,356 @@ std::unique_ptr llama_model_glm_dsa::build_arch_graph(const l return std::make_unique(*this, params); } +llama_model_glm_dsa::graph::graph(const llama_model & model, const llm_graph_params & params) : + llm_graph_context(params) { + const bool is_mla = hparams.is_mla(); + GGML_ASSERT(is_mla); + + // note: these are the actual head sizes you get when treating as MHA or after "decompression" using wv_b for MLA + const int64_t n_embd_head_k = hparams.n_embd_head_k_mla(); + const int64_t n_embd_head_v = hparams.n_embd_head_v_mla(); + GGML_UNUSED(n_embd_head_v); + + const int64_t n_embd_head_qk_rope = hparams.n_rot(); + const int64_t n_embd_head_qk_nope = n_embd_head_k - n_embd_head_qk_rope; + + const int64_t n_indexer_head = hparams.indexer_n_head; + const int64_t n_embd_indexer_head = hparams.indexer_head_size; + const int64_t n_embd_indexer_head_rope = hparams.n_rot(); + const int64_t n_embd_indexer_head_nope = n_embd_indexer_head - n_embd_indexer_head_rope; + const uint32_t n_indexer_top_k = hparams.indexer_top_k; + + const uint32_t kv_lora_rank = hparams.n_lora_kv; + + // We have to pre-scale kq_scale and attn_factor to make the YaRN RoPE work correctly. + // See https://github.com/ggml-org/llama.cpp/discussions/7416 for detailed explanation. + // And also: https://github.com/ggml-org/llama.cpp/pull/17945 [TAG_DEEPSEEK2_YARN_LOG_MUL_FIX] + + // first cancel the adjustment from llama_hparams::yarn_attn_factor_adjust to get the original attn_factor + GGML_ASSERT(ext_factor >= 0.0f); + const float attn_factor_org = attn_factor * (1.0f + 0.1f * logf(1.0f / freq_scale)); + + // use the original attn_factor to pre-scale the kq_scale + const float mscale = attn_factor_org * (1.0f + 0.1f * hparams.rope_yarn_log_mul * logf(1.0f / freq_scale)); + const float kq_scale = 1.0f * mscale * mscale / sqrtf(float(n_embd_head_k)); + + ggml_tensor * cur; + ggml_tensor * inpL; + + // {n_embd, n_tokens} + inpL = build_inp_embd(model.tok_embd); + + // inp_pos - contains the positions + ggml_tensor * inp_pos = build_inp_pos(); + + llm_graph_input_attn_k_dsa * inp_attn_dsa = build_attn_inp_k_dsa(); + + ggml_tensor * inp_out_ids = build_inp_out_ids(); + + // Difference vs Deepseek 3.2: shared indexer layers reuse the top_k from the previous full indexer layers + // See https://huggingface.co/zai-org/GLM-5.2/blob/main/config.json#L30 + ggml_tensor * prev_top_k = nullptr; + for (int il = 0; il < n_layer; ++il) { + ggml_tensor * inpSA = inpL; + + // norm + cur = build_norm(inpL, model.layers[il].attn_norm, NULL, LLM_NORM_RMS, il); + cb(cur, "attn_norm", il); + + // self_attention + { + ggml_tensor * qr = ggml_mul_mat(ctx0, model.layers[il].wq_a, cur); + cb(qr, "qr", il); + + qr = build_norm(qr, model.layers[il].attn_q_a_norm, nullptr, LLM_NORM_RMS, il); + cb(qr, "qr", il); + + ggml_tensor * top_k = nullptr; + + // lightning indexer + if (hparams.is_indexer_full(il)) { + // "full" layer + ggml_tensor * indexer_q = ggml_mul_mat(ctx0, model.layers[il].indexer_attn_q_b, qr); + cb(indexer_q, "indexer_q", il); + + // split into {n_embd_indexer_head_rope, n_indexer_head, n_tokens} + ggml_tensor * indexer_q_pe = + ggml_view_3d(ctx0, indexer_q, n_embd_indexer_head_rope, n_indexer_head, n_tokens, + ggml_row_size(indexer_q->type, n_embd_indexer_head), + ggml_row_size(indexer_q->type, n_embd_indexer_head) * n_indexer_head, 0); + cb(indexer_q_pe, "indexer_q_pe", il); + + // and {n_embd_indexer_head_nope, n_indexer_head, n_tokens} + ggml_tensor * indexer_q_nope = + ggml_view_3d(ctx0, indexer_q, n_embd_indexer_head_nope, n_indexer_head, n_tokens, + ggml_row_size(indexer_q->type, n_embd_indexer_head), + ggml_row_size(indexer_q->type, n_embd_indexer_head) * n_indexer_head, + ggml_row_size(indexer_q->type, n_embd_indexer_head_nope)); + cb(indexer_q_nope, "indexer_q_nope", il); + + indexer_q_pe = + ggml_rope_ext(ctx0, indexer_q_pe, inp_pos, nullptr, n_rot, LLAMA_ROPE_TYPE_NORM, n_ctx_orig, + freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow); + cb(indexer_q_pe, "indexer_q_pe", il); + + // {n_embd_indexer_head_rope + n_embd_indexer_head_nope, n_head, n_tokens} + indexer_q = ggml_concat(ctx0, indexer_q_pe, indexer_q_nope, 0); + cb(indexer_q, "indexer_q", il); + + ggml_tensor * indexer_k = ggml_mul_mat(ctx0, model.layers[il].indexer_attn_k, cur); + cb(indexer_k, "indexer_k", il); + + indexer_k = build_norm(indexer_k, model.layers[il].indexer_k_norm, model.layers[il].indexer_k_norm_b, + LLM_NORM, il); + cb(indexer_k, "indexer_k", il); + + // split into {n_embd_indexer_head_rope, 1, n_tokens} + ggml_tensor * indexer_k_pe = ggml_view_3d(ctx0, indexer_k, n_embd_indexer_head_rope, 1, n_tokens, + ggml_row_size(indexer_k->type, n_embd_indexer_head), + ggml_row_size(indexer_k->type, n_embd_indexer_head) * 1, 0); + cb(indexer_k_pe, "indexer_k_pe", il); + + // and {n_embd_indexer_head_nope, 1, n_tokens} + ggml_tensor * indexer_k_nope = ggml_view_3d(ctx0, indexer_k, n_embd_indexer_head_nope, 1, n_tokens, + ggml_row_size(indexer_k->type, n_embd_indexer_head), + ggml_row_size(indexer_k->type, n_embd_indexer_head) * 1, + ggml_row_size(indexer_k->type, n_embd_indexer_head_nope)); + cb(indexer_k_nope, "indexer_k_nope", il); + + indexer_k_pe = + ggml_rope_ext(ctx0, indexer_k_pe, inp_pos, nullptr, n_rot, LLAMA_ROPE_TYPE_NORM, n_ctx_orig, + freq_base, freq_scale, ext_factor, attn_factor, beta_fast, beta_slow); + cb(indexer_k_pe, "indexer_k_pe", il); + + // {n_embd_indexer_head_rope + n_embd_indexer_head_nope, 1, n_tokens} + indexer_k = ggml_concat(ctx0, indexer_k_pe, indexer_k_nope, 0); + cb(indexer_k, "indexer_k", il); + + // perform Hadamard transform on indexer q and k + indexer_q = ggml_mul_mat(ctx0, inp_attn_dsa->self_k_rot_lid, indexer_q); + cb(indexer_q, "indexer_q", il); + indexer_k = ggml_mul_mat(ctx0, inp_attn_dsa->self_k_rot_lid, indexer_k); + cb(indexer_k, "indexer_k", il); + + // store indexer keys to KV cache + const auto * mctx_lid = inp_attn_dsa->mctx->get_lid(); + const auto & k_idxs_lid = inp_attn_dsa->get_k_idxs_lid(); + ggml_build_forward_expand(gf, mctx_lid->cpy_k(ctx0, indexer_k, k_idxs_lid, il)); + + // prepare indexer weights + ggml_tensor * indexer_weights = ggml_mul_mat(ctx0, model.layers[il].indexer_proj, cur); + cb(indexer_weights, "indexer_weights", il); + + // get cached indexer keys + indexer_k = mctx_lid->get_k(ctx0, il); + + // split the batch into streams if needed + const auto n_stream = indexer_k->ne[3]; + indexer_q = + ggml_view_4d(ctx0, indexer_q, indexer_q->ne[0], indexer_q->ne[1], indexer_q->ne[2] / n_stream, + n_stream, indexer_q->nb[1], indexer_q->nb[2], indexer_q->nb[3] / n_stream, 0); + indexer_weights = + ggml_view_4d(ctx0, indexer_weights, indexer_weights->ne[0], indexer_weights->ne[1] / n_stream, + indexer_weights->ne[2], n_stream, indexer_weights->nb[1], + indexer_weights->nb[2] / n_stream, indexer_weights->nb[3] / n_stream, 0); + + // pre-scale weights to avoid scaling operations on huge indexer_score tensor + indexer_weights = + ggml_scale(ctx0, indexer_weights, 1.0f / sqrtf(float(n_embd_indexer_head * n_indexer_head))); + cb(indexer_weights, "indexer_weights", il); + + ggml_tensor * indexer_score = nullptr; + if (cparams.fused_lid) { + indexer_score = ggml_lightning_indexer(ctx0, indexer_q, indexer_k, indexer_weights, + inp_attn_dsa->get_kq_mask_lid()); + cb(indexer_score, "indexer_score", il); + res->add_fused_node({ LLM_FUSED_OP_LIGHTNING_INDEXER, indexer_score, il }); + } else { + // calculate indexer kq + indexer_q = ggml_permute(ctx0, indexer_q, 0, 2, 1, 3); + cb(indexer_q, "indexer_q", il); + indexer_k = ggml_permute(ctx0, indexer_k, 0, 2, 1, 3); + cb(indexer_k, "indexer_k", il); + + ggml_tensor * indexer_kq = ggml_mul_mat(ctx0, indexer_k, indexer_q); + cb(indexer_kq, "indexer_kq", il); + + // ReLU requires contiguous tensors + indexer_kq = ggml_cont(ctx0, ggml_permute(ctx0, indexer_kq, 2, 1, 0, 3)); + cb(indexer_kq, "indexer_kq", il); + + // apply ReLU + indexer_score = ggml_relu(ctx0, indexer_kq); + cb(indexer_score, "indexer_score", il); + + // multiply scores by indexer weights + indexer_score = ggml_mul(ctx0, indexer_score, indexer_weights); + cb(indexer_score, "indexer_score", il); + + // sum by q n_indexer_head dimension + indexer_score = ggml_sum_rows(ctx0, indexer_score); + cb(indexer_score, "indexer_score", il); + + // permute result to match KQ mask + indexer_score = ggml_cont(ctx0, ggml_permute(ctx0, indexer_score, 2, 1, 0, 3)); + cb(indexer_score, "indexer_score", il); + + // mask indexer scores + ggml_tensor * indexer_kq_mask = inp_attn_dsa->get_kq_mask_lid(); + if (indexer_kq_mask->type != indexer_score->type) { + indexer_kq_mask = ggml_cast(ctx0, indexer_kq_mask, indexer_score->type); + } + indexer_score = ggml_add(ctx0, indexer_score, indexer_kq_mask); + cb(indexer_score, "indexer_score", il); + } + + // get indices of top k indexer scores + uint32_t n_top_k = indexer_score->ne[0] < n_indexer_top_k ? indexer_score->ne[0] : n_indexer_top_k; + top_k = ggml_cont(ctx0, ggml_top_k(ctx0, indexer_score, n_top_k)); + prev_top_k = top_k; + cb(top_k, "top_k", il); + } else { + // "shared" indexer layer - reuse from previous + top_k = prev_top_k; + cb(top_k, "top_k", il); + } + + ggml_tensor * q = ggml_mul_mat(ctx0, model.layers[il].wq_b, qr); + cb(q, "q", il); + + // split into {n_embd_head_qk_nope, n_head, n_tokens} + ggml_tensor * q_nope = + ggml_view_3d(ctx0, q, n_embd_head_qk_nope, n_head, n_tokens, ggml_row_size(q->type, n_embd_head_k), + ggml_row_size(q->type, n_embd_head_k) * n_head, 0); + cb(q_nope, "q_nope", il); + + // and {n_embd_head_qk_rope, n_head, n_tokens} + ggml_tensor * q_pe = ggml_view_3d( + ctx0, q, n_embd_head_qk_rope, n_head, n_tokens, ggml_row_size(q->type, n_embd_head_k), + ggml_row_size(q->type, n_embd_head_k) * n_head, ggml_row_size(q->type, n_embd_head_qk_nope)); + cb(q_pe, "q_pe", il); + + ggml_tensor * kv_cmpr_pe = ggml_mul_mat(ctx0, model.layers[il].wkv_a_mqa, cur); + cb(kv_cmpr_pe, "kv_cmpr_pe", il); + + // split into {kv_lora_rank, n_tokens} + ggml_tensor * kv_cmpr = + ggml_view_2d(ctx0, kv_cmpr_pe, kv_lora_rank, n_tokens, + ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope), 0); + cb(kv_cmpr, "kv_cmpr", il); + + // and {n_embd_head_qk_rope, 1, n_tokens} + ggml_tensor * k_pe = ggml_view_3d(ctx0, kv_cmpr_pe, n_embd_head_qk_rope, 1, n_tokens, + ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope), + ggml_row_size(kv_cmpr_pe->type, kv_lora_rank + n_embd_head_qk_rope), + ggml_row_size(kv_cmpr_pe->type, kv_lora_rank)); + cb(k_pe, "k_pe", il); + + q_pe = ggml_rope_ext(ctx0, q_pe, inp_pos, nullptr, n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + cb(q_pe, "q_pe", il); + + k_pe = ggml_rope_ext(ctx0, k_pe, inp_pos, nullptr, n_rot, rope_type, n_ctx_orig, freq_base, freq_scale, + ext_factor, attn_factor, beta_fast, beta_slow); + cb(k_pe, "k_pe", il); + + kv_cmpr = build_norm(kv_cmpr, model.layers[il].attn_kv_a_norm, nullptr, LLM_NORM_RMS, il); + cb(kv_cmpr, "kv_cmpr", il); + + // MLA attention + { + // {n_embd_head_qk_nope, n_tokens, n_head} + q_nope = ggml_permute(ctx0, q_nope, 0, 2, 1, 3); + cb(q_nope, "q_nope_perm", il); + + // {n_embd_head_qk_nope, kv_lora_rank, n_head} x {n_embd_head_qk_nope, n_tokens, n_head} + ggml_tensor * q_nope_absorbed = ggml_mul_mat(ctx0, model.layers[il].wk_b, q_nope); + cb(q_nope_absorbed, "q_nope_absorbed", il); + + // {kv_lora_rank, n_head, n_tokens} + q_nope_absorbed = ggml_permute(ctx0, q_nope_absorbed, 0, 2, 1, 3); + cb(q_nope_absorbed, "q_nope_absorbed_perm", il); + + // {n_embd_head_qk_rope + kv_lora_rank, n_head, n_tokens} + // note: rope must go first for in-place context shifting in build_rope_shift() + ggml_tensor * Qcur = ggml_concat(ctx0, q_nope_absorbed, q_pe, 0); + cb(Qcur, "Qcur", il); + + kv_cmpr = ggml_reshape_3d(ctx0, kv_cmpr, kv_lora_rank, 1, n_tokens); + cb(kv_cmpr, "kv_cmpr_reshape", il); + + // {n_embd_head_qk_rope + kv_lora_rank, 1, n_tokens} + ggml_tensor * Kcur = ggml_concat(ctx0, kv_cmpr, k_pe, 0); + cb(Kcur, "Kcur", il); + + // {kv_lora_rank, 1, n_tokens} + ggml_tensor * Vcur = kv_cmpr; + cb(Vcur, "Vcur", il); + + // note: MLA with the absorption optimization converts into MQA (ie: GQA with 1 group) + cur = build_attn(inp_attn_dsa, model.layers[il].wo, NULL, model.layers[il].wo_s, Qcur, Kcur, Vcur, + nullptr, nullptr, model.layers[il].wv_b, top_k, kq_scale, il); + } + } + if (il == n_layer - 1 && inp_out_ids) { + cur = ggml_get_rows(ctx0, cur, inp_out_ids); + inpSA = ggml_get_rows(ctx0, inpSA, inp_out_ids); + } + ggml_tensor * ffn_inp = ggml_add(ctx0, cur, inpSA); + cb(ffn_inp, "ffn_inp", il); + + cur = build_norm(ffn_inp, model.layers[il].ffn_norm, NULL, LLM_NORM_RMS, il); + cb(cur, "ffn_norm", il); + + if ((uint32_t) il < hparams.n_layer_dense_lead) { + cur = build_ffn(cur, model.layers[il].ffn_up, NULL, model.layers[il].ffn_up_s, model.layers[il].ffn_gate, + NULL, model.layers[il].ffn_gate_s, model.layers[il].ffn_down, NULL, + model.layers[il].ffn_down_s, NULL, LLM_FFN_SILU, LLM_FFN_PAR, il); + cb(cur, "ffn_out", il); + } else { + // MoE branch + ggml_tensor * moe_out = build_moe_ffn( + cur, model.layers[il].ffn_gate_inp, model.layers[il].ffn_up_exps, model.layers[il].ffn_gate_exps, + model.layers[il].ffn_down_exps, model.layers[il].ffn_exp_probs_b, n_expert, n_expert_used, LLM_FFN_SILU, + hparams.expert_weights_norm, hparams.expert_weights_scale, + (llama_expert_gating_func_type) hparams.expert_gating_func, il, nullptr, + model.layers[il].ffn_gate_up_exps, model.layers[il].ffn_up_exps_s, model.layers[il].ffn_gate_exps_s, + model.layers[il].ffn_down_exps_s); + cb(moe_out, "ffn_moe_out", il); + + // FFN shared expert + { + ggml_tensor * ffn_shexp = + build_ffn(cur, model.layers[il].ffn_up_shexp, NULL, model.layers[il].ffn_up_shexp_s, + model.layers[il].ffn_gate_shexp, NULL, model.layers[il].ffn_gate_shexp_s, + model.layers[il].ffn_down_shexp, NULL, model.layers[il].ffn_down_shexp_s, NULL, + LLM_FFN_SILU, LLM_FFN_PAR, il); + cb(ffn_shexp, "ffn_shexp", il); + + cur = ggml_add(ctx0, moe_out, ffn_shexp); + cb(cur, "ffn_out", il); + } + } + cur = ggml_add(ctx0, cur, ffn_inp); + + cur = build_cvec(cur, il); + cb(cur, "l_out", il); + + // input for next layer + inpL = cur; + } + cur = inpL; + + cur = build_norm(cur, model.output_norm, NULL, LLM_NORM_RMS, -1); + + cb(cur, "result_norm", -1); + res->t_embd = cur; + + // lm_head + cur = ggml_mul_mat(ctx0, model.output, cur); + + cb(cur, "result_output", -1); + res->t_logits = cur; + + ggml_build_forward_expand(gf, cur); +} diff --git a/src/models/models.h b/src/models/models.h index 7a52e7bc1ab7..78984a3effda 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -1216,7 +1216,9 @@ struct llama_model_glm_dsa : public llama_model_base { void load_arch_hparams(llama_model_loader & ml) override; void load_arch_tensors(llama_model_loader & ml) override; - using graph = llama_model_deepseek2::graph; + struct graph : public llm_graph_context { + graph(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 24da69780d21..785c804ab95b 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -193,6 +193,9 @@ if (NOT WIN32 OR NOT BUILD_SHARED_LIBS) # llama_build_and_test(test-double-float.cpp) # SLOW llama_build_and_test(test-llama-archs.cpp) + llama_build_and_test(test-glm-dsa.cpp) + target_sources(test-glm-dsa PRIVATE test-glm-dsa-greedy.cpp test-glm-dsa-moe.cpp test-glm-dsa-stability.cpp) + target_include_directories(test-glm-dsa PRIVATE ${PROJECT_SOURCE_DIR}/src) endif() llama_build_and_test(test-chat-peg-parser.cpp peg-parser/simple-tokenize.cpp) diff --git a/tests/make-glm-dsa-layer-fixture.py b/tests/make-glm-dsa-layer-fixture.py new file mode 100644 index 000000000000..eb96acc49b6f --- /dev/null +++ b/tests/make-glm-dsa-layer-fixture.py @@ -0,0 +1,176 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import argparse +import json +from pathlib import Path +import sys + +sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "gguf-py")) +import gguf + + +GLM_PREFIX = "glm-dsa." + + +def field_value(reader: gguf.GGUFReader, key: str): + field = reader.get_field(key) + if field is None: + raise ValueError(f"missing required GGUF field: {key}") + return field.contents() + + +def add_metadata( + reader: gguf.GGUFReader, + writer: gguf.GGUFWriter, + layer_count: int, + top_k: int, + context_length: int, + vocab_size: int, +) -> None: + indexer_types = field_value(reader, f"{GLM_PREFIX}attention.indexer.types") + if len(indexer_types) < layer_count: + raise ValueError("fixture requests more layers than the GLM indexer schedule contains") + + dense_layer_count = int(field_value(reader, f"{GLM_PREFIX}leading_dense_block_count")) + overrides = { + f"{GLM_PREFIX}vocab_size": (gguf.GGUFValueType.UINT32, vocab_size, None), + f"{GLM_PREFIX}block_count": (gguf.GGUFValueType.UINT32, layer_count, None), + f"{GLM_PREFIX}context_length": (gguf.GGUFValueType.UINT32, context_length, None), + f"{GLM_PREFIX}leading_dense_block_count": ( + gguf.GGUFValueType.UINT32, + min(dense_layer_count, layer_count), + None, + ), + f"{GLM_PREFIX}nextn_predict_layers": (gguf.GGUFValueType.UINT32, 0, None), + f"{GLM_PREFIX}attention.indexer.top_k": (gguf.GGUFValueType.UINT32, top_k, None), + f"{GLM_PREFIX}attention.indexer.types": ( + gguf.GGUFValueType.ARRAY, + indexer_types[:layer_count], + gguf.GGUFValueType.STRING, + ), + } + + for field in reader.fields.values(): + if field.name.startswith("GGUF.") or field.name == gguf.Keys.General.ARCHITECTURE: + continue + if not field.name.startswith(("general.", "tokenizer.", GLM_PREFIX)): + continue + if field.name in overrides: + continue + + value_type = field.types[0] + sub_type = field.types[-1] if value_type == gguf.GGUFValueType.ARRAY else None + writer.add_key_value(field.name, field.contents(), value_type, sub_type=sub_type) + + for key, (value_type, value, sub_type) in overrides.items(): + writer.add_key_value(key, value, value_type, sub_type=sub_type) + + +def fixture_tensor_data(tensor: gguf.ReaderTensor, vocab_size: int): + if tensor.name in {"token_embd.weight", "output.weight"}: + if tensor.data.ndim != 2 or tensor.data.shape[0] < vocab_size: + raise ValueError(f"cannot trim padded vocabulary rows from {tensor.name}") + return tensor.data[:vocab_size] + return tensor.data + + +def add_tensor_info(writer: gguf.GGUFWriter, readers: list[gguf.GGUFReader], vocab_size: int) -> None: + for reader in readers: + for tensor in reader.tensors: + data = fixture_tensor_data(tensor, vocab_size) + writer.add_tensor_info( + tensor.name, + data.shape, + data.dtype, + data.nbytes, + tensor.tensor_type, + ) + + +def write_tensor_data(writer: gguf.GGUFWriter, readers: list[gguf.GGUFReader], vocab_size: int) -> None: + for reader in readers: + for tensor in reader.tensors: + data = fixture_tensor_data(tensor, vocab_size) + print(f"writing {tensor.name}: {data.nbytes / (1024 * 1024):.2f} MiB", flush=True) + writer.write_tensor_data(data, tensor_endianess=reader.endianess) + + +def package_component_paths(package: Path, layer_count: int) -> tuple[Path, list[Path]]: + manifest_path = package / "model-package.json" + if not manifest_path.is_file(): + raise FileNotFoundError(f"missing model package manifest: {manifest_path}") + + manifest = json.loads(manifest_path.read_text()) + if manifest.get("format") != "layer-package": + raise ValueError(f"unsupported model package format: {manifest.get('format')}") + + layers = {entry["layer_index"]: package / entry["path"] for entry in manifest["layers"]} + missing_layers = [layer for layer in range(layer_count) if layer not in layers] + if missing_layers: + raise ValueError(f"model package is missing requested layers: {missing_layers}") + + metadata = package / "shared" / "metadata.gguf" + paths = [package / "shared" / "embeddings.gguf"] + paths.extend(layers[layer] for layer in range(layer_count)) + paths.append(package / "shared" / "output.gguf") + return metadata, paths + + +def make_fixture( + package: Path, + output: Path, + layer_count: int, + top_k: int, + context_length: int, +) -> None: + metadata_path, paths = package_component_paths(package, layer_count) + missing = [path for path in paths if not path.is_file()] + if not metadata_path.is_file(): + missing.append(metadata_path) + if missing: + raise FileNotFoundError(f"missing GGUF component files: {missing}") + if output.exists(): + raise FileExistsError(f"output already exists: {output}") + + metadata = gguf.GGUFReader(metadata_path, "r") + readers = [gguf.GGUFReader(path, "r") for path in paths] + architecture = field_value(metadata, gguf.Keys.General.ARCHITECTURE) + vocab_size = len(field_value(metadata, gguf.Keys.Tokenizer.LIST)) + writer = gguf.GGUFWriter(output, arch=architecture, endianess=metadata.endianess) + + alignment = metadata.get_field(gguf.Keys.General.ALIGNMENT) + if alignment is not None: + writer.data_alignment = alignment.contents() + + add_metadata(metadata, writer, layer_count, top_k, context_length, vocab_size) + add_tensor_info(writer, readers, vocab_size) + writer.write_header_to_file() + writer.write_kv_data_to_file() + writer.write_ti_data_to_file() + write_tensor_data(writer, readers, vocab_size) + writer.close() + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Build a real-weight GLM-DSA parity fixture from a layer package." + ) + parser.add_argument("package", type=Path) + parser.add_argument("output", type=Path) + parser.add_argument("--layers", type=int, default=7) + parser.add_argument("--top-k", type=int, default=8) + parser.add_argument("--context-length", type=int, default=131072) + args = parser.parse_args() + + if args.layers < 1: + parser.error("--layers must be positive") + if args.top_k < 1: + parser.error("--top-k must be positive") + if args.context_length < 1: + parser.error("--context-length must be positive") + make_fixture(args.package, args.output, args.layers, args.top_k, args.context_length) + + +if __name__ == "__main__": + main() diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 084344fb25d7..d9d84177221f 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -8829,6 +8829,11 @@ static std::vector> make_test_cases_eval() { } } + // GLM-5.2 routed experts at single-token decode dimensions. Eight stored + // experts exercise all selected lanes without allocating all 256 experts. + test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q2_K, GGML_TYPE_F32, 8, 8, false, 2048, 1, 6144)); + test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_Q3_K, GGML_TYPE_F32, 8, 8, false, 6144, 1, 2048)); + for (int bs : {1, 4, 512}) { for (ggml_type type_a : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_Q4_0, GGML_TYPE_Q4_K}) { for (ggml_type type_b : {GGML_TYPE_F32}) { @@ -9668,6 +9673,11 @@ static std::vector> make_test_cases_perf() { // Qwen3-VL-8B https://github.com/ggml-org/llama.cpp/issues/17012 test_cases.emplace_back(new test_flash_attn_ext(72, 72, 16, {1, 1}, 5776, 5776, false, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // GLM-5.2 compact DSA decode: 64 query heads share one MLA KV head. + for (int kv : { 2048, 8192, 32768, 131072 }) { + test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {64, 1}, kv, 1, false, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + } + test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_Q4_0)); diff --git a/tests/test-glm-dsa-greedy.cpp b/tests/test-glm-dsa-greedy.cpp new file mode 100644 index 000000000000..ddbe504be25b --- /dev/null +++ b/tests/test-glm-dsa-greedy.cpp @@ -0,0 +1,565 @@ +#include "test-glm-dsa-greedy.h" + +#include "common.h" +#include "ggml-backend.h" +#include "ggml-cpp.h" +#include "ggml.h" +#include "llama-context.h" +#include "llama-cpp.h" +#include "llama.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#if defined(__APPLE__) +# include +#endif + +namespace { + +constexpr int k_decode_steps = 256; +constexpr int k_warmup_steps = 16; +constexpr int k_route_probe_step = 130; +constexpr size_t k_layer_count = 7; + +static void require(bool condition, const char * message) { + if (!condition) { + throw std::runtime_error(message); + } +} + +static uint64_t resident_bytes() { +#if defined(__APPLE__) + mach_task_basic_info_data_t info{}; + mach_msg_type_number_t count = MACH_TASK_BASIC_INFO_COUNT; + const kern_return_t status = + task_info(mach_task_self(), MACH_TASK_BASIC_INFO, reinterpret_cast(&info), &count); + return status == KERN_SUCCESS ? info.resident_size : 0; +#else + return 0; +#endif +} + +static bool is_metal_backend(ggml_backend_t backend) { + if (!backend) { + return false; + } + const char * backend_name = ggml_backend_name(backend); + const char * device_name = ggml_backend_dev_name(ggml_backend_get_device(backend)); + return (backend_name && std::strncmp(backend_name, "Metal", 5) == 0) || + (device_name && std::strncmp(device_name, "MTL", 3) == 0); +} + +static bool is_named_get_rows(const ggml_tensor * tensor, const char * name) { + return tensor->op == GGML_OP_GET_ROWS && std::strstr(tensor->name, name); +} + +struct op_signature { + int lightning_indexer = 0; + int selected_k = 0; + int selected_v = 0; + int dense_mask = 0; + int rope = 0; + int flash_attention = 0; + int heavy = 0; + int heavy_metal = 0; +}; + +static bool same_shape(const op_signature & lhs, const op_signature & rhs) { + return lhs.lightning_indexer == rhs.lightning_indexer && lhs.selected_k == rhs.selected_k && + lhs.selected_v == rhs.selected_v && lhs.dense_mask == rhs.dense_mask && lhs.rope == rhs.rope && + lhs.flash_attention == rhs.flash_attention && lhs.heavy == rhs.heavy && lhs.heavy_metal == rhs.heavy_metal; +} + +struct route_capture { + std::array, k_layer_count> top_k; + std::array, k_layer_count> scores; +}; + +struct decode_observer { + llama_context * context = nullptr; + bool active = false; + bool capture_routes = false; + op_signature current; + std::vector non_metal_ops; + route_capture routes; + + void begin_step() { + active = true; + current = {}; + non_metal_ops.clear(); + } + + op_signature finish_step() { + active = false; + return current; + } +}; + +static int tensor_layer(const ggml_tensor * tensor, const char * prefix) { + const size_t prefix_length = std::strlen(prefix); + if (std::strncmp(tensor->name, prefix, prefix_length) != 0 || tensor->name[prefix_length] != '-') { + return -1; + } + + char * end = nullptr; + const long layer = std::strtol(tensor->name + prefix_length + 1, &end, 10); + return end != tensor->name + prefix_length + 1 && *end == '\0' && layer >= 0 && layer < (long) k_layer_count ? + layer : + -1; +} + +static bool should_capture_route(const ggml_tensor * tensor) { + return tensor_layer(tensor, "ffn_moe_topk") >= 0 || tensor_layer(tensor, "ffn_moe_probs_biased") >= 0; +} + +static void capture_route_tensor(decode_observer & observer, ggml_tensor * tensor) { + if (const int layer = tensor_layer(tensor, "ffn_moe_topk"); layer >= 0) { + require(tensor->type == GGML_TYPE_I32, "captured GLM expert route is not i32"); + observer.routes.top_k[layer].resize(ggml_nelements(tensor)); + ggml_backend_tensor_get(tensor, observer.routes.top_k[layer].data(), 0, + observer.routes.top_k[layer].size() * sizeof(int32_t)); + return; + } + + if (const int layer = tensor_layer(tensor, "ffn_moe_probs_biased"); layer >= 0) { + require(tensor->type == GGML_TYPE_F32, "captured GLM expert scores are not f32"); + observer.routes.scores[layer].resize(ggml_nelements(tensor)); + ggml_backend_tensor_get(tensor, observer.routes.scores[layer].data(), 0, + observer.routes.scores[layer].size() * sizeof(float)); + } +} + +static bool is_dense_sparse_mask(const ggml_tensor * tensor) { + return tensor->op == GGML_OP_SET_ROWS && tensor->ne[0] == 1 && tensor->ne[1] > 1; +} + +static bool observe_decode(ggml_tensor * tensor, bool ask, void * user_data) { + auto * observer = static_cast(user_data); + if (!ask) { + capture_route_tensor(*observer, tensor); + return true; + } + if (!observer->active) { + return false; + } + + const bool selected_k = is_named_get_rows(tensor, "k_selected"); + const bool selected_v = is_named_get_rows(tensor, "v_selected"); + const bool dense_mask = is_dense_sparse_mask(tensor); + observer->current.lightning_indexer += tensor->op == GGML_OP_LIGHTNING_INDEXER; + observer->current.selected_k += selected_k; + observer->current.selected_v += selected_v; + observer->current.dense_mask += dense_mask; + observer->current.rope += tensor->op == GGML_OP_ROPE; + observer->current.flash_attention += tensor->op == GGML_OP_FLASH_ATTN_EXT; + + const bool heavy = tensor->op == GGML_OP_LIGHTNING_INDEXER || tensor->op == GGML_OP_ROPE || + tensor->op == GGML_OP_FLASH_ATTN_EXT || selected_k || selected_v; + if (heavy) { + ++observer->current.heavy; + ggml_backend_t backend = ggml_backend_sched_get_tensor_backend(observer->context->get_sched(), tensor); + if (is_metal_backend(backend)) { + ++observer->current.heavy_metal; + } else { + const char * backend_name = backend ? ggml_backend_name(backend) : "unassigned"; + observer->non_metal_ops.emplace_back(std::string(tensor->name) + " (" + backend_name + ")"); + } + } + return observer->capture_routes && should_capture_route(tensor); +} + +static bool silent_model_load_progress(float, void *) { + return true; +} + +static llama_model_ptr load_model(const char * model_path, ggml_backend_dev_t device) { + std::array devices = { device, nullptr }; + llama_model_params model_params = llama_model_default_params(); + model_params.devices = devices.data(); + model_params.n_gpu_layers = -1; + model_params.split_mode = LLAMA_SPLIT_MODE_NONE; + model_params.progress_callback = silent_model_load_progress; + llama_model_ptr model(llama_model_load_from_file(model_path, model_params)); + require(model != nullptr, "failed to load GLM long-parity fixture"); + return model; +} + +static llama_context_ptr make_context(llama_model * model, + uint32_t context_length, + decode_observer & observer, + bool native_metal) { + llama_context_params params = llama_context_default_params(); + params.n_ctx = context_length; + params.n_batch = 32; + params.n_ubatch = 32; + params.n_threads = 8; + params.n_threads_batch = 8; + params.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_ENABLED; + params.cb_eval = observe_decode; + params.cb_eval_user_data = &observer; + + llama_context_ptr context(llama_init_from_model(model, params)); + require(context != nullptr, "failed to create GLM long-parity context"); + observer.context = context.get(); + + llama_cparams & cparams = const_cast(context->get_cparams()); + cparams.fused_lid = native_metal; + cparams.auto_flid = false; + cparams.flash_attn = native_metal; + cparams.auto_fa = false; + return context; +} + +static void enable_cpu_reference(llama_context * context) { + ggml_backend_sched_t scheduler = context->get_sched(); + for (int i = 0; i < ggml_backend_sched_get_n_backends(scheduler); ++i) { + ggml_backend_t backend = ggml_backend_sched_get_backend(scheduler, i); + if (ggml_backend_dev_type(ggml_backend_get_device(backend)) != GGML_BACKEND_DEVICE_TYPE_CPU) { + continue; + } + using set_use_ref_fn = void (*)(ggml_backend_t, bool); + ggml_backend_reg_t reg = ggml_backend_dev_backend_reg(ggml_backend_get_device(backend)); + auto set_use_ref = + reinterpret_cast(ggml_backend_reg_get_proc_address(reg, "ggml_backend_cpu_set_use_ref")); + require(set_use_ref != nullptr, "CPU backend has no reference-mode control"); + set_use_ref(backend, true); + } +} + +struct batch_guard { + llama_batch value; + + batch_guard() : value(llama_batch_init(16, 0, 1)) {} + + ~batch_guard() { llama_batch_free(value); } +}; + +static const float * decode_tokens(llama_context * context, + batch_guard & batch, + const llama_token * tokens, + size_t count, + int32_t position, + decode_observer * observer, + op_signature * signature) { + common_batch_clear(batch.value); + for (size_t i = 0; i < count; ++i) { + common_batch_add(batch.value, tokens[i], position + i, { 0 }, i + 1 == count); + } + if (observer) { + observer->begin_step(); + } + require(llama_decode(context, batch.value) == 0, "GLM long-parity decode failed"); + if (observer && signature) { + *signature = observer->finish_step(); + } + const float * logits = llama_get_logits_ith(context, batch.value.n_tokens - 1); + require(logits != nullptr, "GLM long-parity decode produced no logits"); + return logits; +} + +static llama_token greedy_token(const float * logits, int32_t n_vocab) { + llama_token result = -1; + float best_value = -std::numeric_limits::infinity(); + for (llama_token token = 0; token < n_vocab; ++token) { + if (!std::isnan(logits[token]) && (result < 0 || logits[token] > best_value)) { + result = token; + best_value = logits[token]; + } + } + require(result >= 0, "GLM long-parity logits contain no finite argmax"); + return result; +} + +static double logit_nmse(const float * expected, const float * actual, int32_t n_vocab) { + double squared_error = 0.0; + double squared_reference = 0.0; + for (int32_t i = 0; i < n_vocab; ++i) { + if (!std::isfinite(expected[i]) || !std::isfinite(actual[i])) { + if (std::isinf(expected[i]) && std::isinf(actual[i]) && + std::signbit(expected[i]) == std::signbit(actual[i])) { + continue; + } + return std::numeric_limits::infinity(); + } + const double difference = static_cast(expected[i]) - actual[i]; + squared_error += difference * difference; + squared_reference += static_cast(expected[i]) * expected[i]; + } + return squared_reference == 0.0 ? squared_error : squared_error / squared_reference; +} + +struct nmse_accumulator { + long double squared_error = 0.0; + long double squared_reference = 0.0; + + void add(const float * expected, const float * actual, int32_t count) { + for (int32_t i = 0; i < count; ++i) { + if (!std::isfinite(expected[i]) || !std::isfinite(actual[i])) { + continue; + } + const long double difference = static_cast(expected[i]) - actual[i]; + squared_error += difference * difference; + squared_reference += static_cast(expected[i]) * expected[i]; + } + } + + double value() const { + return squared_reference == 0.0 ? static_cast(squared_error) : + static_cast(squared_error / squared_reference); + } +}; + +static bool same_indices(const std::vector & lhs, const std::vector & rhs) { + std::vector lhs_sorted = lhs; + std::vector rhs_sorted = rhs; + std::sort(lhs_sorted.begin(), lhs_sorted.end()); + std::sort(rhs_sorted.begin(), rhs_sorted.end()); + return lhs_sorted == rhs_sorted; +} + +static bool route_difference_is_near_tie(const std::vector & expected, + const std::vector & actual, + const std::vector & expected_scores, + const std::vector & actual_scores, + double & max_score_error, + double & max_reference_gap) { + require(expected.size() == actual.size(), "GLM expert route widths differ"); + require(expected_scores.size() == actual_scores.size(), "GLM expert score widths differ"); + + max_score_error = 0.0; + for (size_t i = 0; i < expected_scores.size(); ++i) { + max_score_error = + std::max(max_score_error, std::fabs(static_cast(expected_scores[i]) - actual_scores[i])); + } + + std::vector expected_only; + std::vector actual_only; + for (int32_t expert : expected) { + if (std::find(actual.begin(), actual.end(), expert) == actual.end()) { + expected_only.push_back(expert); + } + } + for (int32_t expert : actual) { + if (std::find(expected.begin(), expected.end(), expert) == expected.end()) { + actual_only.push_back(expert); + } + } + require(expected_only.size() == actual_only.size() && !expected_only.empty(), + "GLM expert route difference is malformed"); + + max_reference_gap = 0.0; + for (int32_t selected : expected_only) { + require(selected >= 0 && static_cast(selected) < expected_scores.size(), + "CPU GLM expert route is out of range"); + for (int32_t replacement : actual_only) { + require(replacement >= 0 && static_cast(replacement) < expected_scores.size(), + "Metal GLM expert route is out of range"); + max_reference_gap = std::max(max_reference_gap, + static_cast(expected_scores[selected]) - expected_scores[replacement]); + } + } + + return max_reference_gap <= 2.0 * max_score_error + 1e-6; +} + +static bool validate_route_probe(const route_capture & expected, const route_capture & actual) { + bool divergence_seen = false; + for (size_t layer = 3; layer < k_layer_count; ++layer) { + require(!expected.top_k[layer].empty() && !actual.top_k[layer].empty(), + "GLM route probe is missing expert indices"); + require(!expected.scores[layer].empty() && !actual.scores[layer].empty(), + "GLM route probe is missing expert scores"); + if (same_indices(expected.top_k[layer], actual.top_k[layer])) { + continue; + } + + if (!divergence_seen) { + double max_score_error = 0.0; + double max_reference_gap = 0.0; + require(route_difference_is_near_tie(expected.top_k[layer], actual.top_k[layer], expected.scores[layer], + actual.scores[layer], max_score_error, max_reference_gap), + "first CPU/Metal GLM expert route difference is not explained by a score tie"); + std::printf("GLM route probe: first difference at layer %zu, score gap %.3e, max score error %.3e\n", layer, + max_reference_gap, max_score_error); + } + divergence_seen = true; + } + return divergence_seen; +} + +static void prefill(llama_context * context, batch_guard & batch) { + std::array prompt{}; + std::iota(prompt.begin(), prompt.end(), 1); + decode_tokens(context, batch, prompt.data(), prompt.size(), 0, nullptr, nullptr); +} + +struct reference_trace { + int32_t n_vocab = 0; + std::vector greedy_tokens; + std::vector logits; + route_capture routes; +}; + +static reference_trace run_reference(const char * model_path, uint32_t context_length) { + ggml_backend_dev_t cpu_device = ggml_backend_dev_by_type(GGML_BACKEND_DEVICE_TYPE_CPU); + require(cpu_device != nullptr, "CPU device unavailable for GLM long parity"); + llama_model_ptr model = load_model(model_path, cpu_device); + decode_observer observer; + llama_context_ptr context = make_context(model.get(), context_length, observer, false); + enable_cpu_reference(context.get()); + batch_guard batch; + prefill(context.get(), batch); + + reference_trace trace; + trace.n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model.get())); + trace.greedy_tokens.reserve(k_decode_steps); + trace.logits.reserve(static_cast(trace.n_vocab) * k_decode_steps); + llama_token current = 17; + op_signature first_signature; + const auto started = std::chrono::steady_clock::now(); + for (int step = 0; step < k_decode_steps; ++step) { + observer.capture_routes = step == k_route_probe_step; + op_signature signature; + const float * logits = decode_tokens(context.get(), batch, ¤t, 1, 16 + step, &observer, &signature); + observer.capture_routes = false; + if (step == 0) { + first_signature = signature; + require(signature.lightning_indexer == 0 && signature.selected_k == 0 && signature.selected_v == 0 && + signature.dense_mask == 7, + "CPU reference did not retain the generic dense GLM-DSA graph"); + } else { + require(same_shape(first_signature, signature), "CPU reference graph shape changed during decode"); + } + current = greedy_token(logits, trace.n_vocab); + trace.greedy_tokens.push_back(current); + trace.logits.insert(trace.logits.end(), logits, logits + trace.n_vocab); + if ((step + 1) % 64 == 0) { + std::printf("GLM CPU reference: ctx=%u step=%d/%d\n", context_length, step + 1, k_decode_steps); + std::fflush(stdout); + } + } + const auto elapsed = std::chrono::duration(std::chrono::steady_clock::now() - started); + std::printf("GLM CPU reference: ctx=%u %.3f ms/step\n", context_length, elapsed.count() / k_decode_steps); + trace.routes = std::move(observer.routes); + return trace; +} + +static void validate_native_signature(const op_signature & signature, const decode_observer & observer) { + if (!observer.non_metal_ops.empty()) { + for (const std::string & op : observer.non_metal_ops) { + std::fprintf(stderr, "non-Metal GLM long-parity op: %s\n", op.c_str()); + } + } + require(signature.lightning_indexer == 4, "Metal decode did not execute four Full indexers"); + require(signature.selected_k == 7 && signature.selected_v == 7, + "Metal decode did not execute seven compact K/V gathers"); + require(signature.dense_mask == 0, "Metal decode materialized a dense sparse-attention mask"); + require(signature.flash_attention == 7, "Metal decode did not execute seven flash-attention operations"); + require(signature.rope > 0, "Metal decode did not execute RoPE"); + require(signature.heavy > 0 && signature.heavy == signature.heavy_metal, + "Metal decode assigned a heavy GLM-DSA operation outside Metal"); +} + +static void run_metal(const char * model_path, uint32_t context_length, const reference_trace & reference) { + ggml_backend_dev_t metal_device = ggml_backend_dev_by_name("MTL0"); + require(metal_device != nullptr, "Metal device unavailable for GLM long parity"); + llama_model_ptr model = load_model(model_path, metal_device); + decode_observer observer; + llama_context_ptr context = make_context(model.get(), context_length, observer, true); + batch_guard batch; + prefill(context.get(), batch); + + require(reference.greedy_tokens.size() == k_decode_steps, "GLM reference token trace is incomplete"); + require(reference.logits.size() == static_cast(reference.n_vocab) * k_decode_steps, + "GLM reference logit trace is incomplete"); + op_signature first_signature; + double max_nmse = 0.0; + int max_nmse_step = 0; + int first_bad_nmse_step = -1; + int first_bad_token_step = -1; + uint64_t rss_baseline = 0; + uint64_t rss_peak = 0; + nmse_accumulator aggregate_nmse; + const auto started = std::chrono::steady_clock::now(); + for (int step = 0; step < k_decode_steps; ++step) { + const llama_token input = step == 0 ? 17 : reference.greedy_tokens[step - 1]; + observer.capture_routes = step == k_route_probe_step; + op_signature signature; + const float * logits = decode_tokens(context.get(), batch, &input, 1, 16 + step, &observer, &signature); + observer.capture_routes = false; + validate_native_signature(signature, observer); + if (step == 0) { + first_signature = signature; + } else { + require(same_shape(first_signature, signature), "Metal GLM-DSA graph shape changed during decode"); + } + + const float * expected_logits = reference.logits.data() + static_cast(step) * reference.n_vocab; + const double nmse = logit_nmse(expected_logits, logits, reference.n_vocab); + aggregate_nmse.add(expected_logits, logits, reference.n_vocab); + if (nmse > max_nmse) { + max_nmse = nmse; + max_nmse_step = step; + } + const llama_token actual = greedy_token(logits, reference.n_vocab); + if (actual != reference.greedy_tokens[step] && first_bad_token_step < 0) { + first_bad_token_step = step; + std::fprintf(stderr, "GLM greedy mismatch: ctx=%u step=%d CPU=%d Metal=%d NMSE=%.3e\n", context_length, + step, reference.greedy_tokens[step], actual, nmse); + } + if (nmse > 2e-4 && first_bad_nmse_step < 0) { + first_bad_nmse_step = step; + std::fprintf(stderr, "GLM logit NMSE crossed tolerance: ctx=%u step=%d NMSE=%.3e\n", context_length, step, + nmse); + } + if (step + 1 == k_warmup_steps) { + llama_synchronize(context.get()); + rss_baseline = resident_bytes(); + rss_peak = rss_baseline; + } else if (step + 1 > k_warmup_steps && (step + 1) % 16 == 0) { + llama_synchronize(context.get()); + rss_peak = std::max(rss_peak, resident_bytes()); + } + if ((step + 1) % 64 == 0) { + std::printf("GLM Metal parity: ctx=%u step=%d/%d max-NMSE=%.3e\n", context_length, step + 1, k_decode_steps, + max_nmse); + std::fflush(stdout); + } + } + llama_synchronize(context.get()); + const auto elapsed = std::chrono::duration(std::chrono::steady_clock::now() - started); + const uint64_t rss_growth = rss_peak > rss_baseline ? rss_peak - rss_baseline : 0; + const double aggregate_error = aggregate_nmse.value(); + const bool route_tie = validate_route_probe(reference.routes, observer.routes); + require(rss_baseline == 0 || rss_growth <= 32ULL * 1024 * 1024, "Metal GLM context grew after long-parity warmup"); + std::printf( + "GLM long parity: ctx=%u tokens=%d aggregate-NMSE=%.3e max-NMSE=%.3e@%d Metal=%.3f ms/step " + "RSS-growth=%.1f MiB\n", + context_length, k_decode_steps, aggregate_error, max_nmse, max_nmse_step, elapsed.count() / k_decode_steps, + rss_growth / (1024.0 * 1024.0)); + require(first_bad_token_step < 0, "CPU and Metal greedy token traces differ"); + require(aggregate_error <= 2e-4, "CPU and Metal aggregate logit NMSE exceeds tolerance"); + require(first_bad_nmse_step < 0 || (first_bad_nmse_step == k_route_probe_step && route_tie && max_nmse <= 5e-3), + "CPU and Metal logit divergence is not limited to the characterized expert-route tie"); +} + +} // namespace + +void test_glm_dsa_long_greedy_parity(const char * model_path) { + const reference_trace reference = run_reference(model_path, 2048); + for (uint32_t context_length : std::array{ 2048, 32768, 131072 }) { + run_metal(model_path, context_length, reference); + } +} diff --git a/tests/test-glm-dsa-greedy.h b/tests/test-glm-dsa-greedy.h new file mode 100644 index 000000000000..aa7341daaac5 --- /dev/null +++ b/tests/test-glm-dsa-greedy.h @@ -0,0 +1,3 @@ +#pragma once + +void test_glm_dsa_long_greedy_parity(const char * model_path); diff --git a/tests/test-glm-dsa-moe.cpp b/tests/test-glm-dsa-moe.cpp new file mode 100644 index 000000000000..920cca4bcdc8 --- /dev/null +++ b/tests/test-glm-dsa-moe.cpp @@ -0,0 +1,158 @@ +#include "test-glm-dsa-moe.h" + +#include "ggml-backend.h" +#include "ggml-cpp.h" +#include "ggml.h" + +#include +#include +#include +#include +#include + +namespace { + +struct moe_shape { + ggml_type type; + int64_t n_input; + int64_t n_output; +}; + +static void require(bool condition, const char * message) { + if (!condition) { + throw std::runtime_error(message); + } +} + +static float fixture_value(uint64_t index, uint32_t salt) { + uint32_t value = static_cast(index) ^ salt; + value ^= value >> 16; + value *= 0x7feb352dU; + value ^= value >> 15; + value *= 0x846ca68bU; + value ^= value >> 16; + return static_cast(value & 0xffffU) / 32768.0f - 1.0f; +} + +static std::vector make_values(size_t count, uint32_t salt) { + std::vector values(count); + for (size_t i = 0; i < count; ++i) { + values[i] = fixture_value(i, salt); + } + return values; +} + +static std::vector make_quantized_weights(const moe_shape & shape) { + const ggml_type_traits * traits = ggml_get_type_traits(shape.type); + require(traits && traits->from_float_ref && traits->to_float, "missing quantization traits"); + + const size_t row_size = ggml_row_size(shape.type, shape.n_input); + std::vector weights(row_size * shape.n_output); + std::vector row(shape.n_input); + for (int64_t output = 0; output < shape.n_output; ++output) { + for (int64_t input = 0; input < shape.n_input; ++input) { + row[input] = fixture_value(static_cast(output) * shape.n_input + input, 0x51a7d3e2U); + } + traits->from_float_ref(row.data(), weights.data() + output * row_size, shape.n_input); + } + return weights; +} + +static std::vector run_mul_mat_id(ggml_backend_t backend, + const moe_shape & shape, + const std::vector & weights, + const std::vector & input) { + ggml_init_params params = { + /* .mem_size = */ 2 * 1024 * 1024, + /* .mem_buffer = */ nullptr, + /* .no_alloc = */ true, + }; + ggml_context_ptr context(ggml_init(params)); + require(context != nullptr, "failed to create MoE precision context"); + + ggml_tensor * matrix = ggml_new_tensor_3d(context.get(), shape.type, shape.n_input, shape.n_output, 1); + ggml_tensor * vector = ggml_new_tensor_3d(context.get(), GGML_TYPE_F32, shape.n_input, 1, 1); + ggml_tensor * ids = ggml_new_tensor_2d(context.get(), GGML_TYPE_I32, 1, 1); + ggml_tensor * output = ggml_mul_mat_id(context.get(), matrix, vector, ids); + + ggml_cgraph * graph = ggml_new_graph_custom(context.get(), 32, false); + ggml_build_forward_expand(graph, output); + require(ggml_backend_supports_op(backend, output), "backend does not support GLM MoE matvec"); + + ggml_backend_buffer_ptr buffer(ggml_backend_alloc_ctx_tensors(context.get(), backend)); + require(buffer != nullptr, "failed to allocate MoE precision tensors"); + require(ggml_nbytes(matrix) == weights.size(), "quantized MoE weight size differs"); + ggml_backend_tensor_set(matrix, weights.data(), 0, weights.size()); + ggml_backend_tensor_set(vector, input.data(), 0, input.size() * sizeof(float)); + const int32_t expert = 0; + ggml_backend_tensor_set(ids, &expert, 0, sizeof(expert)); + + require(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS, "GLM MoE precision graph failed"); + std::vector values(ggml_nelements(output)); + ggml_backend_tensor_get(output, values.data(), 0, values.size() * sizeof(float)); + return values; +} + +static std::vector precise_reference(const moe_shape & shape, + const std::vector & weights, + const std::vector & input) { + const ggml_type_traits * traits = ggml_get_type_traits(shape.type); + const size_t row_size = ggml_row_size(shape.type, shape.n_input); + std::vector row(shape.n_input); + std::vector output(shape.n_output); + for (int64_t i = 0; i < shape.n_output; ++i) { + traits->to_float(weights.data() + i * row_size, row.data(), shape.n_input); + double sum = 0.0; + for (int64_t j = 0; j < shape.n_input; ++j) { + sum += static_cast(row[j]) * input[j]; + } + output[i] = static_cast(sum); + } + return output; +} + +static double normalized_mean_squared_error(const std::vector & reference, const std::vector & actual) { + require(reference.size() == actual.size(), "MoE precision result sizes differ"); + double squared_error = 0.0; + double reference_sum = 0.0; + for (size_t i = 0; i < reference.size(); ++i) { + const double difference = static_cast(reference[i]) - actual[i]; + squared_error += difference * difference; + reference_sum += static_cast(reference[i]) * reference[i]; + } + return squared_error / reference_sum; +} + +static void test_shape(ggml_backend_t cpu, ggml_backend_t metal, const moe_shape & shape) { + const std::vector weights = make_quantized_weights(shape); + const std::vector input = make_values(shape.n_input, 0x9c6d812fU); + const std::vector reference = precise_reference(shape, weights, input); + const std::vector cpu_out = run_mul_mat_id(cpu, shape, weights, input); + const std::vector metal_out = run_mul_mat_id(metal, shape, weights, input); + + const double cpu_nmse = normalized_mean_squared_error(reference, cpu_out); + const double metal_nmse = normalized_mean_squared_error(reference, metal_out); + const double cpu_metal_nmse = normalized_mean_squared_error(cpu_out, metal_out); + std::printf("GLM MoE %s [%lldx%lld]: CPU/reference %.3e, Metal/reference %.3e, CPU/Metal %.3e\n", + ggml_type_name(shape.type), static_cast(shape.n_output), + static_cast(shape.n_input), cpu_nmse, metal_nmse, cpu_metal_nmse); + + require(metal_nmse <= 2e-4, "Metal GLM MoE matvec exceeds the numerical parity threshold"); + require(metal_nmse <= cpu_nmse, "Metal GLM MoE matvec is less accurate than the CPU Q8_K path"); +} + +} // namespace + +void test_glm_dsa_moe_precision() { + ggml_backend_dev_t metal_device = ggml_backend_dev_by_name("MTL0"); + if (!metal_device) { + return; + } + + ggml_backend_ptr cpu(ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr)); + ggml_backend_ptr metal(ggml_backend_dev_init(metal_device, nullptr)); + require(cpu != nullptr && metal != nullptr, "failed to initialize MoE precision backends"); + + test_shape(cpu.get(), metal.get(), { GGML_TYPE_Q2_K, 6144, 2048 }); + test_shape(cpu.get(), metal.get(), { GGML_TYPE_Q3_K, 2048, 6144 }); +} diff --git a/tests/test-glm-dsa-moe.h b/tests/test-glm-dsa-moe.h new file mode 100644 index 000000000000..f8f0f0e4cc96 --- /dev/null +++ b/tests/test-glm-dsa-moe.h @@ -0,0 +1,3 @@ +#pragma once + +void test_glm_dsa_moe_precision(); diff --git a/tests/test-glm-dsa-stability.cpp b/tests/test-glm-dsa-stability.cpp new file mode 100644 index 000000000000..3186267ae4af --- /dev/null +++ b/tests/test-glm-dsa-stability.cpp @@ -0,0 +1,235 @@ +#include "test-glm-dsa-stability.h" + +#include "ggml-backend.h" +#include "ggml-cpp.h" +#include "ggml.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +#if defined(__APPLE__) +# include +#endif + +namespace { + +constexpr int64_t k_head_size_k = 576; +constexpr int64_t k_head_size_v = 512; +constexpr int64_t k_attention_heads = 64; +constexpr int64_t k_indexer_head_size = 128; +constexpr int64_t k_indexer_heads = 32; +constexpr int64_t k_indexer_top_k = 2048; +constexpr int k_warmup_steps = 16; +constexpr int k_decode_steps = 256; + +static void require(bool condition, const char * message) { + if (!condition) { + throw std::runtime_error(message); + } +} + +static uint64_t resident_bytes() { +#if defined(__APPLE__) + mach_task_basic_info_data_t info{}; + mach_msg_type_number_t count = MACH_TASK_BASIC_INFO_COUNT; + const kern_return_t status = + task_info(mach_task_self(), MACH_TASK_BASIC_INFO, reinterpret_cast(&info), &count); + return status == KERN_SUCCESS ? info.resident_size : 0; +#else + return 0; +#endif +} + +static std::vector fixture_values(size_t count, int period) { + std::vector values(count); + for (size_t i = 0; i < count; ++i) { + values[i] = static_cast(static_cast(i % period) - period / 2) / period; + } + return values; +} + +struct stability_graph { + ggml_context_ptr context; + ggml_backend_buffer_ptr buffer; + ggml_cgraph * graph = nullptr; + ggml_tensor * position = nullptr; + ggml_tensor * lid_scores = nullptr; + ggml_tensor * attention = nullptr; + size_t buffer_size = 0; + int node_count = 0; +}; + +static void set_f32(ggml_tensor * tensor, const std::vector & values) { + require(ggml_nelements(tensor) == static_cast(values.size()), "wrong stability f32 size"); + ggml_backend_tensor_set(tensor, values.data(), 0, values.size() * sizeof(float)); +} + +static stability_graph make_stability_graph(ggml_backend_t backend, int64_t context_length) { + ggml_init_params params = { + /* .mem_size = */ 4 * 1024 * 1024, + /* .mem_buffer = */ nullptr, + /* .no_alloc = */ true, + }; + ggml_context_ptr context(ggml_init(params)); + require(context != nullptr, "failed to create GLM stability context"); + + ggml_tensor * q_raw = ggml_new_tensor_4d(context.get(), GGML_TYPE_F32, k_head_size_k, k_attention_heads, 1, 1); + ggml_tensor * position = ggml_new_tensor_1d(context.get(), GGML_TYPE_I32, 1); + ggml_tensor * q_rope = ggml_rope_ext(context.get(), q_raw, position, nullptr, 64, GGML_ROPE_TYPE_NORMAL, 131072, + 10000.0f, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + ggml_tensor * q = ggml_permute(context.get(), q_rope, 0, 2, 1, 3); + + ggml_tensor * indexer_q = + ggml_new_tensor_4d(context.get(), GGML_TYPE_F32, k_indexer_head_size, k_indexer_heads, 1, 1); + ggml_tensor * indexer_k = + ggml_new_tensor_4d(context.get(), GGML_TYPE_F16, k_indexer_head_size, 1, context_length, 1); + ggml_tensor * indexer_weights = ggml_new_tensor_4d(context.get(), GGML_TYPE_F32, k_indexer_heads, 1, 1, 1); + ggml_tensor * indexer_mask = ggml_new_tensor_4d(context.get(), GGML_TYPE_F16, context_length, 1, 1, 1); + ggml_tensor * lid_scores = + ggml_lightning_indexer(context.get(), indexer_q, indexer_k, indexer_weights, indexer_mask); + ggml_set_name(lid_scores, "stability_lid_scores"); + + const int64_t selected_rows = std::min(context_length, k_indexer_top_k); + ggml_tensor * top_k = ggml_new_tensor_3d(context.get(), GGML_TYPE_I32, selected_rows, 1, 1); + ggml_tensor * k_cache = ggml_new_tensor_4d(context.get(), GGML_TYPE_F16, k_head_size_k, 1, context_length, 1); + ggml_tensor * v_cache = ggml_new_tensor_4d(context.get(), GGML_TYPE_F16, k_head_size_v, 1, context_length, 1); + ggml_tensor * causal_mask = ggml_new_tensor_4d(context.get(), GGML_TYPE_F16, context_length, 1, 1, 1); + + ggml_tensor * k_full = ggml_permute(context.get(), k_cache, 0, 2, 1, 3); + ggml_tensor * v_full = ggml_permute(context.get(), v_cache, 0, 2, 1, 3); + ggml_tensor * k_selected = ggml_get_rows(context.get(), k_full, top_k); + ggml_tensor * v_selected = ggml_get_rows(context.get(), v_full, top_k); + ggml_set_name(k_selected, "stability_k_selected"); + ggml_set_name(v_selected, "stability_v_selected"); + require(k_selected->ne[1] == selected_rows && v_selected->ne[1] == selected_rows, + "stability graph did not preserve compact selected rows"); + k_selected = ggml_cast(context.get(), k_selected, GGML_TYPE_F16); + v_selected = ggml_cast(context.get(), v_selected, GGML_TYPE_F16); + + ggml_tensor * mask_rows = ggml_reshape_4d(context.get(), causal_mask, 1, context_length, 1, 1); + ggml_tensor * mask_selected = ggml_get_rows(context.get(), mask_rows, top_k); + mask_selected = ggml_reshape_4d(context.get(), mask_selected, selected_rows, 1, 1, 1); + mask_selected = ggml_cast(context.get(), mask_selected, GGML_TYPE_F16); + + ggml_tensor * attention = ggml_flash_attn_ext(context.get(), q, k_selected, v_selected, mask_selected, + 1.0f / std::sqrt(static_cast(k_head_size_k)), 0.0f, 0.0f); + ggml_flash_attn_ext_set_prec(attention, GGML_PREC_F32); + ggml_set_name(attention, "stability_flash_attention"); + + ggml_cgraph * graph = ggml_new_graph_custom(context.get(), 128, false); + ggml_build_forward_expand(graph, lid_scores); + ggml_build_forward_expand(graph, attention); + + int lid_count = 0; + int gather_count = 0; + int rope_count = 0; + int flash_count = 0; + int dense_mask_count = 0; + for (int i = 0; i < ggml_graph_n_nodes(graph); ++i) { + ggml_tensor * node = ggml_graph_node(graph, i); + require(ggml_backend_supports_op(backend, node), "stability graph contains an op unsupported by Metal"); + lid_count += node->op == GGML_OP_LIGHTNING_INDEXER; + gather_count += node->op == GGML_OP_GET_ROWS; + rope_count += node->op == GGML_OP_ROPE; + flash_count += node->op == GGML_OP_FLASH_ATTN_EXT; + dense_mask_count += node->op == GGML_OP_SET_ROWS; + } + require(lid_count == 1 && gather_count == 3 && rope_count == 1 && flash_count == 1, + "stability graph is missing a native GLM sparse-decode operation"); + require(dense_mask_count == 0, "stability graph materialized a dense sparse-attention mask"); + + ggml_backend_buffer_ptr buffer(ggml_backend_alloc_ctx_tensors(context.get(), backend)); + require(buffer != nullptr, "failed to allocate GLM stability tensors"); + ggml_backend_tensor_memset(indexer_k, 0, 0, ggml_nbytes(indexer_k)); + ggml_backend_tensor_memset(indexer_mask, 0, 0, ggml_nbytes(indexer_mask)); + ggml_backend_tensor_memset(k_cache, 0, 0, ggml_nbytes(k_cache)); + ggml_backend_tensor_memset(v_cache, 0, 0, ggml_nbytes(v_cache)); + ggml_backend_tensor_memset(causal_mask, 0, 0, ggml_nbytes(causal_mask)); + set_f32(q_raw, fixture_values(ggml_nelements(q_raw), 31)); + set_f32(indexer_q, fixture_values(ggml_nelements(indexer_q), 29)); + set_f32(indexer_weights, fixture_values(ggml_nelements(indexer_weights), 17)); + + std::vector indices(selected_rows); + for (int64_t i = 0; i < selected_rows; ++i) { + indices[i] = static_cast((i * context_length) / selected_rows); + } + ggml_backend_tensor_set(top_k, indices.data(), 0, indices.size() * sizeof(int32_t)); + + stability_graph result; + result.context = std::move(context); + result.buffer = std::move(buffer); + result.graph = graph; + result.position = position; + result.lid_scores = lid_scores; + result.attention = attention; + result.buffer_size = ggml_backend_buffer_get_size(result.buffer.get()); + result.node_count = ggml_graph_n_nodes(graph); + return result; +} + +static void run_context(ggml_backend_t backend, int64_t context_length) { + stability_graph fixture = make_stability_graph(backend, context_length); + const int32_t first_position = static_cast(context_length - k_decode_steps); + for (int step = 0; step < k_warmup_steps; ++step) { + const int32_t position = first_position + step; + ggml_backend_tensor_set(fixture.position, &position, 0, sizeof(position)); + require(ggml_backend_graph_compute(backend, fixture.graph) == GGML_STATUS_SUCCESS, + "GLM stability warmup failed"); + } + ggml_backend_synchronize(backend); + + const uint64_t rss_baseline = resident_bytes(); + uint64_t rss_peak = rss_baseline; + const auto started = std::chrono::steady_clock::now(); + for (int step = 0; step < k_decode_steps; ++step) { + const int32_t position = first_position + step; + ggml_backend_tensor_set(fixture.position, &position, 0, sizeof(position)); + require(ggml_backend_graph_compute(backend, fixture.graph) == GGML_STATUS_SUCCESS, + "GLM stability decode failed"); + if ((step + 1) % 16 == 0) { + ggml_backend_synchronize(backend); + rss_peak = std::max(rss_peak, resident_bytes()); + } + } + ggml_backend_synchronize(backend); + const auto elapsed = std::chrono::duration(std::chrono::steady_clock::now() - started); + + require(ggml_graph_n_nodes(fixture.graph) == fixture.node_count, "GLM stability graph changed across decode"); + require(ggml_backend_buffer_get_size(fixture.buffer.get()) == fixture.buffer_size, + "GLM stability buffer grew across decode"); + const uint64_t rss_growth = rss_peak > rss_baseline ? rss_peak - rss_baseline : 0; + require(rss_baseline == 0 || rss_growth <= 32ULL * 1024 * 1024, "GLM stability RSS grew after Metal warmup"); + + float lid_value = 0.0f; + float attention_value = 0.0f; + ggml_backend_tensor_get(fixture.lid_scores, &lid_value, 0, sizeof(lid_value)); + ggml_backend_tensor_get(fixture.attention, &attention_value, 0, sizeof(attention_value)); + require(std::isfinite(lid_value) && std::isfinite(attention_value), "GLM stability output is not finite"); + + std::printf( + "GLM sparse stability: ctx=%lld top-k=%lld steps=%d nodes=%d buffer=%.1f MiB " + "%.3f ms/step RSS-growth=%.1f MiB\n", + static_cast(context_length), static_cast(std::min(context_length, k_indexer_top_k)), + k_decode_steps, fixture.node_count, fixture.buffer_size / (1024.0 * 1024.0), elapsed.count() / k_decode_steps, + rss_growth / (1024.0 * 1024.0)); +} + +} // namespace + +void test_glm_dsa_decode_stability() { + ggml_backend_dev_t metal_device = ggml_backend_dev_by_name("MTL0"); + if (!metal_device) { + return; + } + ggml_backend_ptr metal(ggml_backend_dev_init(metal_device, nullptr)); + require(metal != nullptr, "failed to initialize Metal stability backend"); + for (int64_t context_length : std::array{ 2048, 32768, 131072 }) { + run_context(metal.get(), context_length); + } +} diff --git a/tests/test-glm-dsa-stability.h b/tests/test-glm-dsa-stability.h new file mode 100644 index 000000000000..eabd404b37e6 --- /dev/null +++ b/tests/test-glm-dsa-stability.h @@ -0,0 +1,3 @@ +#pragma once + +void test_glm_dsa_decode_stability(); diff --git a/tests/test-glm-dsa.cpp b/tests/test-glm-dsa.cpp new file mode 100644 index 000000000000..716a45da2a0f --- /dev/null +++ b/tests/test-glm-dsa.cpp @@ -0,0 +1,1058 @@ +#include "common.h" +#include "ggml-backend.h" +#include "ggml-cpp.h" +#include "ggml.h" +#include "llama-context.h" +#include "llama-cpp.h" +#include "llama.h" +#include "test-glm-dsa-greedy.h" +#include "test-glm-dsa-moe.h" +#include "test-glm-dsa-stability.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +static void require(bool condition, const char * message) { + if (!condition) { + throw std::runtime_error(message); + } +} + +static ggml_context_ptr make_context() { + ggml_init_params params = { + /*.mem_size =*/8 * 1024 * 1024, + /*.mem_buffer =*/nullptr, + /*.no_alloc =*/true, + }; + ggml_context_ptr ctx(ggml_init(params)); + require(ctx != nullptr, "failed to initialize ggml context"); + return ctx; +} + +static ggml_backend_ptr make_cpu_backend() { + ggml_backend_ptr backend(ggml_backend_init_by_type(GGML_BACKEND_DEVICE_TYPE_CPU, nullptr)); + require(backend != nullptr, "failed to initialize CPU backend"); + return backend; +} + +static void set_f32(ggml_tensor * tensor, const std::vector & values) { + require(ggml_nelements(tensor) == static_cast(values.size()), "wrong f32 fixture size"); + ggml_backend_tensor_set(tensor, values.data(), 0, values.size() * sizeof(float)); +} + +static void set_f16(ggml_tensor * tensor, const std::vector & values) { + require(ggml_nelements(tensor) == static_cast(values.size()), "wrong f16 fixture size"); + std::vector converted(values.size()); + std::transform(values.begin(), values.end(), converted.begin(), ggml_fp32_to_fp16); + ggml_backend_tensor_set(tensor, converted.data(), 0, converted.size() * sizeof(ggml_fp16_t)); +} + +static void set_i32(ggml_tensor * tensor, const std::vector & values) { + require(ggml_nelements(tensor) == static_cast(values.size()), "wrong i32 fixture size"); + ggml_backend_tensor_set(tensor, values.data(), 0, values.size() * sizeof(int32_t)); +} + +static std::vector get_f32(const ggml_tensor * tensor) { + std::vector values(ggml_nelements(tensor)); + ggml_backend_tensor_get(tensor, values.data(), 0, values.size() * sizeof(float)); + return values; +} + +static std::vector get_float_values(const ggml_tensor * tensor) { + const size_t count = ggml_nelements(tensor); + std::vector values(count); + if (tensor->type == GGML_TYPE_F32) { + ggml_backend_tensor_get(tensor, values.data(), 0, values.size() * sizeof(float)); + return values; + } + + if (tensor->type == GGML_TYPE_F16) { + std::vector data(count); + ggml_backend_tensor_get(tensor, data.data(), 0, data.size() * sizeof(ggml_fp16_t)); + std::transform(data.begin(), data.end(), values.begin(), ggml_fp16_to_fp32); + return values; + } + + if (tensor->type == GGML_TYPE_BF16) { + std::vector data(count); + ggml_backend_tensor_get(tensor, data.data(), 0, data.size() * sizeof(ggml_bf16_t)); + std::transform(data.begin(), data.end(), values.begin(), ggml_bf16_to_fp32); + return values; + } + + throw std::runtime_error("unsupported captured activation type"); +} + +static std::vector get_i32_values(const ggml_tensor * tensor) { + require(tensor->type == GGML_TYPE_I32, "captured index tensor is not i32"); + std::vector values(ggml_nelements(tensor)); + ggml_backend_tensor_get(tensor, values.data(), 0, values.size() * sizeof(int32_t)); + return values; +} + +static void check_close(float actual, float expected, float tolerance, const char * message) { + if (std::isinf(expected)) { + require(std::isinf(actual) && std::signbit(actual) == std::signbit(expected), message); + return; + } + require(std::fabs(actual - expected) <= tolerance, message); +} + +static std::vector reference_indexer_scores(const std::vector & q, + const std::vector & k, + const std::vector & weights, + const std::vector & mask, + int64_t head_size, + int64_t n_head, + int64_t n_kv, + int64_t n_query) { + std::vector scores(n_kv * n_query); + const float scale = 1.0f / std::sqrt(static_cast(head_size * n_head)); + + // Matches GlmMoeDsaIndexer.forward: weighted head-wise relu(q dot k), scaling, then the causal mask. + for (int64_t iq = 0; iq < n_query; ++iq) { + for (int64_t ik = 0; ik < n_kv; ++ik) { + float score = 0.0f; + for (int64_t ih = 0; ih < n_head; ++ih) { + float dot = 0.0f; + for (int64_t id = 0; id < head_size; ++id) { + dot += q[(iq * n_head + ih) * head_size + id] * k[ik * head_size + id]; + } + score += std::max(dot, 0.0f) * weights[iq * n_head + ih] * scale; + } + scores[iq * n_kv + ik] = score + mask[iq * n_kv + ik]; + } + } + return scores; +} + +static void test_indexer_scores_and_top_k() { + constexpr int64_t head_size = 4; + constexpr int64_t n_head = 2; + constexpr int64_t n_kv = 6; + constexpr int64_t n_query = 2; + constexpr int64_t n_top_k = 3; + + const std::vector q = { + 1.0f, -1.0f, 0.5f, 2.0f, -0.5f, 1.0f, 1.5f, -1.0f, 0.25f, 2.0f, -1.0f, 0.5f, 1.0f, 0.5f, -0.75f, 1.25f, + }; + const std::vector k = { + 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, + 0.0f, 0.0f, 0.0f, 1.0f, 1.0f, 1.0f, 1.0f, 1.0f, -1.0f, 0.5f, 1.0f, -0.5f, + }; + const std::vector weights = { 0.7f, -0.2f, -0.3f, 0.9f }; + const float neg_inf = -std::numeric_limits::infinity(); + const std::vector mask = { + 0.0f, 0.0f, 0.0f, neg_inf, neg_inf, neg_inf, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, + }; + const std::vector official_scores = { + 0.247487373f, -0.0707106781f, 0.0176776695f, neg_inf, neg_inf, neg_inf, + 0.291681547f, -0.0530330086f, 0.0f, 0.344714556f, 0.450780573f, 0.0f, + }; + const std::vector official_top_k = { 0, 2, 1, 4, 3, 0 }; + + ggml_context_ptr ctx = make_context(); + ggml_tensor * q_tensor = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F32, head_size, n_head, n_query, 1); + ggml_tensor * k_tensor = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F32, head_size, 1, n_kv, 1); + ggml_tensor * w_tensor = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F32, n_head, n_query, 1, 1); + ggml_tensor * m_tensor = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F16, n_kv, n_query, 1, 1); + + const float scale = 1.0f / std::sqrt(static_cast(head_size * n_head)); + ggml_tensor * scaled_weights = ggml_scale(ctx.get(), w_tensor, scale); + ggml_tensor * scores = ggml_lightning_indexer(ctx.get(), q_tensor, k_tensor, scaled_weights, m_tensor); + ggml_tensor * top_k = ggml_top_k(ctx.get(), scores, n_top_k); + + ggml_cgraph * graph = ggml_new_graph(ctx.get()); + ggml_build_forward_expand(graph, top_k); + + ggml_backend_ptr backend = make_cpu_backend(); + ggml_backend_buffer_ptr buffer(ggml_backend_alloc_ctx_tensors(ctx.get(), backend.get())); + require(buffer != nullptr, "failed to allocate indexer tensors"); + + set_f32(q_tensor, q); + set_f32(k_tensor, k); + set_f32(w_tensor, weights); + std::vector mask_f16(mask.size()); + std::transform(mask.begin(), mask.end(), mask_f16.begin(), ggml_fp32_to_fp16); + ggml_backend_tensor_set(m_tensor, mask_f16.data(), 0, mask_f16.size() * sizeof(ggml_fp16_t)); + + require(ggml_backend_graph_compute(backend.get(), graph) == GGML_STATUS_SUCCESS, + "indexer graph computation failed"); + + const std::vector expected_scores = + reference_indexer_scores(q, k, weights, mask, head_size, n_head, n_kv, n_query); + const std::vector actual_scores = get_f32(scores); + for (size_t i = 0; i < expected_scores.size(); ++i) { + check_close(expected_scores[i], official_scores[i], 1e-7f, + "C++ reference score differs from the Transformers golden"); + check_close(actual_scores[i], official_scores[i], 1e-6f, "indexer score differs from the Transformers golden"); + } + + std::vector actual_top_k(ggml_nelements(top_k)); + ggml_backend_tensor_get(top_k, actual_top_k.data(), 0, actual_top_k.size() * sizeof(int32_t)); + // ggml_top_k need not rank output; K, V, and mask rows are gathered in the same order. + for (int64_t iq = 0; iq < n_query; ++iq) { + auto actual_begin = actual_top_k.begin() + iq * n_top_k; + auto expected_begin = official_top_k.begin() + iq * n_top_k; + std::sort(actual_begin, actual_begin + n_top_k); + std::vector expected(expected_begin, expected_begin + n_top_k); + std::sort(expected.begin(), expected.end()); + require(std::equal(actual_begin, actual_begin + n_top_k, expected.begin()), + "indexer top-k selection differs from reference"); + } +} + +static void run_generic_indexer_mask_cast(ggml_backend_t backend, int64_t n_kv, int64_t n_query) { + ggml_context_ptr ctx = make_context(); + ggml_tensor * scores = ggml_new_tensor_2d(ctx.get(), GGML_TYPE_F32, n_kv, n_query); + ggml_tensor * mask = ggml_new_tensor_2d(ctx.get(), GGML_TYPE_F16, n_kv, n_query); + ggml_tensor * mask_f32 = ggml_cast(ctx.get(), mask, scores->type); + ggml_tensor * masked_score = ggml_add(ctx.get(), scores, mask_f32); + + ggml_cgraph * graph = ggml_new_graph_custom(ctx.get(), 8, false); + ggml_build_forward_expand(graph, masked_score); + + require(mask_f32->type == GGML_TYPE_F32, "generic indexer mask was not normalized to score precision"); + require(masked_score->src[0]->type == GGML_TYPE_F32 && masked_score->src[1]->type == GGML_TYPE_F32, + "generic indexer add retained mixed input precision"); + require(ggml_backend_supports_op(backend, mask_f32), "backend does not support the generic indexer mask cast"); + require(ggml_backend_supports_op(backend, masked_score), "backend does not support the generic indexer mask add"); + + ggml_backend_buffer_ptr buffer(ggml_backend_alloc_ctx_tensors(ctx.get(), backend)); + require(buffer != nullptr, "failed to allocate generic indexer mask tensors"); + + const size_t count = static_cast(n_kv * n_query); + std::vector scores_data(count); + std::vector mask_data(count); + std::vector mask_f16(count); + for (size_t i = 0; i < count; ++i) { + scores_data[i] = static_cast(static_cast(i % 31) - 15) / 16.0f; + mask_data[i] = i % 7 == 0 ? -std::numeric_limits::infinity() : 0.0f; + mask_f16[i] = ggml_fp32_to_fp16(mask_data[i]); + } + ggml_backend_tensor_set(scores, scores_data.data(), 0, scores_data.size() * sizeof(float)); + ggml_backend_tensor_set(mask, mask_f16.data(), 0, mask_f16.size() * sizeof(ggml_fp16_t)); + + require(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS, + "generic indexer mask graph computation failed"); + const std::vector actual = get_f32(masked_score); + for (size_t i = 0; i < count; ++i) { + const float expected = scores_data[i] + mask_data[i]; + check_close(actual[i], expected, 0.0f, "generic indexer mask result differs"); + } +} + +static void test_generic_indexer_mask_cast() { + ggml_backend_ptr cpu = make_cpu_backend(); + run_generic_indexer_mask_cast(cpu.get(), 256, 19); + run_generic_indexer_mask_cast(cpu.get(), 256, 2304); + + if (ggml_backend_dev_t metal_device = ggml_backend_dev_by_name("MTL0")) { + ggml_backend_ptr metal(ggml_backend_dev_init(metal_device, nullptr)); + require(metal != nullptr, "failed to initialize Metal for generic indexer mask regression"); + run_generic_indexer_mask_cast(metal.get(), 256, 19); + run_generic_indexer_mask_cast(metal.get(), 256, 2304); + } +} + +static void test_interleaved_rope() { + constexpr int64_t n_dims = 4; + constexpr int64_t n_tokens = 2; + constexpr float freq_base = 10000.0f; + const std::vector input = { 1.0f, 2.0f, 3.0f, 4.0f, -1.0f, 0.5f, 2.0f, -0.25f }; + const std::array positions = { 1, 3 }; + + ggml_context_ptr ctx = make_context(); + ggml_tensor * input_tensor = ggml_new_tensor_3d(ctx.get(), GGML_TYPE_F32, n_dims, 1, n_tokens); + ggml_tensor * position_tensor = ggml_new_tensor_1d(ctx.get(), GGML_TYPE_I32, n_tokens); + ggml_tensor * output = ggml_rope_ext(ctx.get(), input_tensor, position_tensor, nullptr, n_dims, + GGML_ROPE_TYPE_NORMAL, 128, freq_base, 1.0f, 0.0f, 1.0f, 0.0f, 0.0f); + + ggml_cgraph * graph = ggml_new_graph(ctx.get()); + ggml_build_forward_expand(graph, output); + + ggml_backend_ptr backend = make_cpu_backend(); + ggml_backend_buffer_ptr buffer(ggml_backend_alloc_ctx_tensors(ctx.get(), backend.get())); + require(buffer != nullptr, "failed to allocate RoPE tensors"); + set_f32(input_tensor, input); + ggml_backend_tensor_set(position_tensor, positions.data(), 0, sizeof(positions)); + require(ggml_backend_graph_compute(backend.get(), graph) == GGML_STATUS_SUCCESS, "RoPE graph computation failed"); + + const std::vector actual = get_f32(output); + for (int64_t it = 0; it < n_tokens; ++it) { + for (int64_t ip = 0; ip < n_dims / 2; ++ip) { + const float theta = positions[it] * std::pow(freq_base, -2.0f * ip / n_dims); + const float x0 = input[it * n_dims + 2 * ip]; + const float x1 = input[it * n_dims + 2 * ip + 1]; + const float expected0 = x0 * std::cos(theta) - x1 * std::sin(theta); + const float expected1 = x0 * std::sin(theta) + x1 * std::cos(theta); + check_close(actual[it * n_dims + 2 * ip], expected0, 1e-5f, + "GGML_ROPE_TYPE_NORMAL even component differs from reference"); + check_close(actual[it * n_dims + 2 * ip + 1], expected1, 1e-5f, + "GGML_ROPE_TYPE_NORMAL odd component differs from reference"); + } + } +} + +struct sparse_attention_outputs { + std::vector dense; + std::vector compact; +}; + +static std::vector fixture_values(size_t count, int period) { + std::vector values(count); + for (size_t i = 0; i < count; ++i) { + values[i] = static_cast(static_cast(i % period) - period / 2) / period; + } + return values; +} + +static sparse_attention_outputs run_sparse_attention(ggml_backend_t backend, bool require_native_support) { + constexpr int64_t head_size_k = 576; + constexpr int64_t head_size_v = 512; + constexpr int64_t n_head = 8; + constexpr int64_t n_kv = 512; + constexpr int64_t n_top_k = 8; + constexpr int64_t n_stream = 2; + + ggml_context_ptr ctx = make_context(); + ggml_tensor * q = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F32, head_size_k, 1, n_head, n_stream); + ggml_tensor * k_cache = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F16, head_size_k, 1, n_kv, n_stream); + ggml_tensor * v_cache = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F16, head_size_v, 1, n_kv, n_stream); + ggml_tensor * causal_mask = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F16, n_kv, 1, 1, n_stream); + ggml_tensor * dense_mask = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_F16, n_kv, 1, 1, n_stream); + ggml_tensor * top_k = ggml_new_tensor_4d(ctx.get(), GGML_TYPE_I32, n_top_k, 1, 1, n_stream); + + ggml_tensor * k_full = ggml_permute(ctx.get(), k_cache, 0, 2, 1, 3); + ggml_tensor * v_full = ggml_permute(ctx.get(), v_cache, 0, 2, 1, 3); + ggml_tensor * dense = ggml_flash_attn_ext(ctx.get(), q, k_full, v_full, dense_mask, + 1.0f / std::sqrt(static_cast(head_size_k)), 0.0f, 0.0f); + ggml_flash_attn_ext_set_prec(dense, GGML_PREC_F32); + + ggml_tensor * top_k_rows = ggml_reshape_3d(ctx.get(), top_k, n_top_k, 1, n_stream); + ggml_tensor * k_selected = ggml_get_rows(ctx.get(), k_full, top_k_rows); + ggml_tensor * v_selected = ggml_get_rows(ctx.get(), v_full, top_k_rows); + require(k_selected->ne[1] == n_top_k && v_selected->ne[1] == n_top_k, + "compact attention did not reduce K/V rows to top-k"); + require(k_full->ne[1] == 64 * k_selected->ne[1] && v_full->ne[1] == 64 * v_selected->ne[1], + "compact attention fixture did not exercise 64x selected-row scaling"); + + k_selected = ggml_cast(ctx.get(), k_selected, GGML_TYPE_F16); + v_selected = ggml_cast(ctx.get(), v_selected, GGML_TYPE_F16); + ggml_tensor * mask_rows = ggml_reshape_4d(ctx.get(), causal_mask, 1, n_kv, 1, n_stream); + ggml_tensor * mask_selected = ggml_get_rows(ctx.get(), mask_rows, top_k_rows); + mask_selected = ggml_reshape_4d(ctx.get(), mask_selected, n_top_k, 1, 1, n_stream); + mask_selected = ggml_cast(ctx.get(), mask_selected, GGML_TYPE_F16); + + ggml_tensor * compact = ggml_flash_attn_ext(ctx.get(), q, k_selected, v_selected, mask_selected, + 1.0f / std::sqrt(static_cast(head_size_k)), 0.0f, 0.0f); + ggml_flash_attn_ext_set_prec(compact, GGML_PREC_F32); + + ggml_cgraph * graph = ggml_new_graph_custom(ctx.get(), GGML_DEFAULT_GRAPH_SIZE, false); + ggml_build_forward_expand(graph, dense); + ggml_build_forward_expand(graph, compact); + + if (require_native_support) { + for (int i = 0; i < ggml_graph_n_nodes(graph); ++i) { + require(ggml_backend_supports_op(backend, ggml_graph_node(graph, i)), + "compact sparse attention graph contains an op unsupported by Metal"); + } + } + + ggml_backend_buffer_ptr buffer(ggml_backend_alloc_ctx_tensors(ctx.get(), backend)); + require(buffer != nullptr, "failed to allocate sparse attention tensors"); + + std::vector indices(n_top_k * n_stream); + for (int64_t is = 0; is < n_stream; ++is) { + for (int64_t ik = 0; ik < n_top_k; ++ik) { + indices[is * n_top_k + ik] = is * 31 + ik * (n_kv / n_top_k); + } + } + + std::vector causal_values(n_kv * n_stream, 0.0f); + std::vector dense_values(n_kv * n_stream, -std::numeric_limits::infinity()); + for (int64_t is = 0; is < n_stream; ++is) { + for (int64_t ik = 0; ik < n_top_k; ++ik) { + dense_values[is * n_kv + indices[is * n_top_k + ik]] = 0.0f; + } + } + + set_f32(q, fixture_values(ggml_nelements(q), 29)); + set_f16(k_cache, fixture_values(ggml_nelements(k_cache), 31)); + set_f16(v_cache, fixture_values(ggml_nelements(v_cache), 37)); + set_f16(causal_mask, causal_values); + set_f16(dense_mask, dense_values); + set_i32(top_k, indices); + + require(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS, + "sparse attention graph computation failed"); + + return { get_f32(dense), get_f32(compact) }; +} + +static void test_compact_sparse_attention() { + ggml_backend_ptr cpu = make_cpu_backend(); + const sparse_attention_outputs cpu_outputs = run_sparse_attention(cpu.get(), false); + require(cpu_outputs.dense.size() == cpu_outputs.compact.size(), "sparse attention output size differs"); + for (size_t i = 0; i < cpu_outputs.dense.size(); ++i) { + check_close(cpu_outputs.compact[i], cpu_outputs.dense[i], 1e-5f, + "compact sparse attention differs from dense reference"); + } + + ggml_backend_dev_t metal_device = ggml_backend_dev_by_name("MTL0"); + if (!metal_device) { + return; + } + + ggml_backend_ptr metal(ggml_backend_dev_init(metal_device, nullptr)); + require(metal != nullptr, "failed to initialize Metal backend"); + const sparse_attention_outputs metal_outputs = run_sparse_attention(metal.get(), true); + for (size_t i = 0; i < cpu_outputs.compact.size(); ++i) { + check_close(metal_outputs.compact[i], cpu_outputs.compact[i], 5e-4f, + "Metal compact sparse attention differs from CPU"); + } +} + +static constexpr size_t real_layer_count = 7; + +struct real_step_capture { + std::array, real_layer_count> top_k; + std::array, real_layer_count> indexer_scores; + std::array, real_layer_count> hidden_states; +}; + +struct real_decode_observer { + int lightning_indexer_count = 0; + int k_selected_count = 0; + int v_selected_count = 0; + int dense_sparse_mask_count = 0; + int rope_count = 0; + int flash_attention_count = 0; + int heavy_op_count = 0; + int heavy_op_metal_count = 0; + int64_t selected_rows = 0; + + llama_context * context = nullptr; + bool capture_values = false; + size_t next_dense_layer = 0; + + std::vector non_metal_heavy_ops; + std::vector steps; + std::unordered_map dense_layer_by_tensor; + + void reset() { + lightning_indexer_count = 0; + k_selected_count = 0; + v_selected_count = 0; + dense_sparse_mask_count = 0; + rope_count = 0; + flash_attention_count = 0; + heavy_op_count = 0; + heavy_op_metal_count = 0; + selected_rows = 0; + capture_values = false; + next_dense_layer = 0; + non_metal_heavy_ops.clear(); + steps.clear(); + dense_layer_by_tensor.clear(); + } + + void begin_decode_step() { + capture_values = true; + next_dense_layer = 0; + dense_layer_by_tensor.clear(); + steps.emplace_back(); + } +}; + +static int tensor_layer(const ggml_tensor * tensor, const char * prefix) { + const size_t prefix_length = std::strlen(prefix); + if (std::strncmp(tensor->name, prefix, prefix_length) != 0 || tensor->name[prefix_length] != '-') { + return -1; + } + + char * end = nullptr; + const long il = std::strtol(tensor->name + prefix_length + 1, &end, 10); + if (end == tensor->name + prefix_length + 1 || *end != '\0' || il < 0 || il >= (long) real_layer_count) { + return -1; + } + return il; +} + +static bool is_metal_backend(ggml_backend_t backend) { + if (!backend) { + return false; + } + const char * backend_name = ggml_backend_name(backend); + const char * device_name = ggml_backend_dev_name(ggml_backend_get_device(backend)); + return (backend_name && std::strncmp(backend_name, "Metal", 5) == 0) || + (device_name && std::strncmp(device_name, "MTL", 3) == 0); +} + +static bool is_dense_sparse_mask(const ggml_tensor * tensor) { + return tensor->op == GGML_OP_SET_ROWS && tensor->ne[0] == 1 && tensor->ne[1] > 1; +} + +static bool is_k_selected(const ggml_tensor * tensor) { + return tensor->op == GGML_OP_GET_ROWS && std::strstr(tensor->name, "k_selected"); +} + +static bool is_v_selected(const ggml_tensor * tensor) { + return tensor->op == GGML_OP_GET_ROWS && std::strstr(tensor->name, "v_selected"); +} + +static bool is_heavy_glm_op(const ggml_tensor * tensor) { + return tensor->op == GGML_OP_LIGHTNING_INDEXER || tensor->op == GGML_OP_ROPE || + tensor->op == GGML_OP_FLASH_ATTN_EXT || is_k_selected(tensor) || is_v_selected(tensor); +} + +static bool is_final_indexer_score(const ggml_tensor * tensor) { + return tensor_layer(tensor, "indexer_score") >= 0 && + (tensor->op == GGML_OP_LIGHTNING_INDEXER || tensor->op == GGML_OP_ADD); +} + +static void observe_heavy_backend(real_decode_observer & observer, ggml_tensor * tensor) { + if (!observer.context || !is_heavy_glm_op(tensor)) { + return; + } + + ++observer.heavy_op_count; + ggml_backend_t backend = ggml_backend_sched_get_tensor_backend(observer.context->get_sched(), tensor); + if (is_metal_backend(backend)) { + ++observer.heavy_op_metal_count; + return; + } + + const char * backend_name = backend ? ggml_backend_name(backend) : "unassigned"; + observer.non_metal_heavy_ops.emplace_back(std::string(tensor->name) + " (" + backend_name + ")"); +} + +static bool should_capture_real_tensor(const ggml_tensor * tensor) { + return tensor_layer(tensor, "l_out") >= 0 || is_final_indexer_score(tensor) || is_k_selected(tensor) || + is_dense_sparse_mask(tensor); +} + +static void capture_real_tensor(real_decode_observer & observer, ggml_tensor * tensor) { + require(!observer.steps.empty(), "real GLM tensor capture has no active decode step"); + real_step_capture & capture = observer.steps.back(); + + if (const int il = tensor_layer(tensor, "l_out"); il >= 0) { + capture.hidden_states[il] = get_float_values(tensor); + return; + } + + if (const int il = tensor_layer(tensor, "indexer_score"); il >= 0 && is_final_indexer_score(tensor)) { + capture.indexer_scores[il] = get_float_values(tensor); + return; + } + + if (is_k_selected(tensor)) { + const int il = tensor_layer(tensor, "k_selected"); + require(il >= 0 && tensor->src[1] != nullptr, "selected K tensor has no layer or index source"); + capture.top_k[il] = get_i32_values(tensor->src[1]); + return; + } + + const auto dense_layer = observer.dense_layer_by_tensor.find(tensor); + require(dense_layer != observer.dense_layer_by_tensor.end(), "dense sparse mask has no layer mapping"); + require(tensor->src[1] != nullptr, "dense sparse mask has no index source"); + capture.top_k[dense_layer->second] = get_i32_values(tensor->src[1]); +} + +static bool observe_real_decode(ggml_tensor * tensor, bool ask, void * user_data) { + auto * observer = static_cast(user_data); + if (ask) { + if (tensor->op == GGML_OP_LIGHTNING_INDEXER) { + ++observer->lightning_indexer_count; + } else if (is_k_selected(tensor)) { + ++observer->k_selected_count; + observer->selected_rows = tensor->ne[1]; + } else if (is_v_selected(tensor)) { + ++observer->v_selected_count; + } else if (is_dense_sparse_mask(tensor)) { + ++observer->dense_sparse_mask_count; + if (observer->capture_values) { + require(observer->next_dense_layer < real_layer_count, "too many dense sparse masks in one decode"); + observer->dense_layer_by_tensor[tensor] = observer->next_dense_layer++; + } + } + + if (tensor->op == GGML_OP_ROPE) { + ++observer->rope_count; + } else if (tensor->op == GGML_OP_FLASH_ATTN_EXT) { + ++observer->flash_attention_count; + } + observe_heavy_backend(*observer, tensor); + return observer->capture_values && should_capture_real_tensor(tensor); + } + + capture_real_tensor(*observer, tensor); + return true; +} + +static bool silent_model_load_progress(float, void *) { + return true; +} + +struct real_execution_config { + const char * name; + bool fused_indexer; + bool compact_decode; +}; + +static void append_real_logits(llama_context * context, + const llama_model * model, + const std::vector & tokens, + int32_t position, + std::vector & logits, + real_decode_observer * observer, + bool capture_values) { + if (capture_values) { + require(observer != nullptr, "real GLM decode capture requires an observer"); + observer->begin_decode_step(); + } + + llama_batch batch = llama_batch_init(tokens.size(), 0, 1); + for (size_t i = 0; i < tokens.size(); ++i) { + common_batch_add(batch, tokens[i], position + i, { 0 }, i + 1 == tokens.size()); + } + + require(llama_decode(context, batch) == 0, "real GLM layer decode failed"); + const float * batch_logits = llama_get_logits_ith(context, batch.n_tokens - 1); + require(batch_logits != nullptr, "real GLM layer produced no logits"); + const int32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model)); + logits.insert(logits.end(), batch_logits, batch_logits + n_vocab); + llama_batch_free(batch); +} + +struct real_model_result { + std::vector logits; + real_decode_observer prefill_observer; + real_decode_observer decode_observer; +}; + +static real_model_result run_real_model(const char * model_path, + ggml_backend_dev_t device, + const real_execution_config & config) { + std::array devices = { device, nullptr }; + llama_model_params model_params = llama_model_default_params(); + model_params.devices = devices.data(); + model_params.n_gpu_layers = -1; + model_params.split_mode = LLAMA_SPLIT_MODE_NONE; + model_params.progress_callback = silent_model_load_progress; + + llama_model_ptr model(llama_model_load_from_file(model_path, model_params)); + require(model != nullptr, "failed to load real GLM layer fixture"); + + real_decode_observer observer; + llama_context_params context_params = llama_context_default_params(); + context_params.n_ctx = 32; + context_params.n_batch = 32; + context_params.n_ubatch = 32; + context_params.n_threads = 8; + context_params.n_threads_batch = 8; + context_params.flash_attn_type = LLAMA_FLASH_ATTN_TYPE_ENABLED; + context_params.cb_eval = observe_real_decode; + context_params.cb_eval_user_data = &observer; + + llama_context_ptr context(llama_init_from_model(model.get(), context_params)); + require(context != nullptr, "failed to create real GLM layer context"); + + const char * device_name = ggml_backend_dev_name(device); + if (ggml_backend_dev_type(device) == GGML_BACKEND_DEVICE_TYPE_CPU || + (device_name && std::strcmp(device_name, "BLAS") == 0)) { + ggml_backend_sched_t scheduler = context->get_sched(); + for (int i = 0; i < ggml_backend_sched_get_n_backends(scheduler); ++i) { + ggml_backend_t backend = ggml_backend_sched_get_backend(scheduler, i); + if (ggml_backend_dev_type(ggml_backend_get_device(backend)) != GGML_BACKEND_DEVICE_TYPE_CPU) { + continue; + } + + using set_use_ref_fn = void (*)(ggml_backend_t, bool); + ggml_backend_reg_t reg = ggml_backend_dev_backend_reg(ggml_backend_get_device(backend)); + auto set_use_ref = reinterpret_cast( + ggml_backend_reg_get_proc_address(reg, "ggml_backend_cpu_set_use_ref")); + require(set_use_ref != nullptr, "CPU backend has no reference-mode control"); + set_use_ref(backend, true); + } + } + + // Test-local controls let one loaded model exercise reference and native graph shapes. + llama_cparams & test_cparams = const_cast(context->get_cparams()); + test_cparams.fused_lid = config.fused_indexer; + test_cparams.auto_flid = false; + test_cparams.flash_attn = config.compact_decode; + test_cparams.auto_fa = false; + + // Context initialization probes fused operations. Reset before inference. + observer.context = context.get(); + observer.reset(); + + std::vector tokens(18); + std::iota(tokens.begin(), tokens.end(), 1); + std::vector logits; + append_real_logits(context.get(), model.get(), std::vector(tokens.begin(), tokens.begin() + 16), 0, + logits, &observer, false); + real_decode_observer prefill_observer = observer; + prefill_observer.context = nullptr; + prefill_observer.dense_layer_by_tensor.clear(); + require(prefill_observer.k_selected_count == 0 && prefill_observer.v_selected_count == 0, + "real GLM prefill unexpectedly used compact selected K/V rows"); + require(prefill_observer.dense_sparse_mask_count > 0, + "real GLM prefill did not retain the dense sparse-attention fallback"); + + observer.reset(); + append_real_logits(context.get(), model.get(), { tokens[16] }, 16, logits, &observer, true); + const real_decode_observer first_decode_observer = observer; + if (config.compact_decode) { + require(first_decode_observer.k_selected_count == (int) real_layer_count && + first_decode_observer.v_selected_count == (int) real_layer_count, + "first real GLM decode did not use one compact selected K/V path per layer"); + require(first_decode_observer.selected_rows == 8, "first real GLM decode did not preserve the fixture top-k"); + require(first_decode_observer.dense_sparse_mask_count == 0, + "first real GLM decode materialized a dense sparse-attention mask"); + } else { + require(first_decode_observer.k_selected_count == 0 && first_decode_observer.v_selected_count == 0, + "dense reference decode unexpectedly selected compact K/V rows"); + require(first_decode_observer.dense_sparse_mask_count == (int) real_layer_count, + "dense reference decode did not materialize one sparse mask per layer"); + } + + append_real_logits(context.get(), model.get(), { tokens[17] }, 17, logits, &observer, true); + require(observer.k_selected_count == 2 * first_decode_observer.k_selected_count && + observer.v_selected_count == 2 * first_decode_observer.v_selected_count, + "repeated real GLM decode changed selected K/V graph shape"); + require(observer.lightning_indexer_count == 2 * first_decode_observer.lightning_indexer_count, + "repeated real GLM decode changed Lightning Indexer placement"); + require(observer.dense_sparse_mask_count == 2 * first_decode_observer.dense_sparse_mask_count, + "repeated real GLM decode changed dense sparse-mask graph shape"); + if (config.compact_decode) { + require(observer.selected_rows == first_decode_observer.selected_rows, + "repeated real GLM decode changed the selected-row count"); + } + + observer.context = nullptr; + observer.dense_layer_by_tensor.clear(); + std::printf("real GLM mode %s: LID=%d selected-K/V=%d/%d dense-mask=%d Metal-heavy=%d/%d\n", config.name, + observer.lightning_indexer_count, observer.k_selected_count, observer.v_selected_count, + observer.dense_sparse_mask_count, observer.heavy_op_metal_count, observer.heavy_op_count); + return { std::move(logits), std::move(prefill_observer), std::move(observer) }; +} + +static double normalized_mean_squared_error(const std::vector & expected, const std::vector & actual) { + require(expected.size() == actual.size(), "real GLM logit vector sizes differ"); + double squared_error = 0.0; + double squared_reference = 0.0; + for (size_t i = 0; i < expected.size(); ++i) { + if (!std::isfinite(expected[i]) || !std::isfinite(actual[i])) { + if (std::isinf(expected[i]) && std::isinf(actual[i]) && + std::signbit(expected[i]) == std::signbit(actual[i])) { + continue; + } + return std::numeric_limits::infinity(); + } + const double difference = static_cast(expected[i]) - actual[i]; + squared_error += difference * difference; + squared_reference += static_cast(expected[i]) * expected[i]; + } + return squared_reference == 0.0 ? squared_error : squared_error / squared_reference; +} + +static double normalized_mean_squared_error(const std::vector & expected, + const std::vector & actual, + size_t begin, + size_t end) { + double squared_error = 0.0; + double squared_reference = 0.0; + for (size_t i = begin; i < end; ++i) { + const double difference = static_cast(expected[i]) - actual[i]; + squared_error += difference * difference; + squared_reference += static_cast(expected[i]) * expected[i]; + } + return squared_reference == 0.0 ? squared_error : squared_error / squared_reference; +} + +static std::vector sorted_indices(std::vector values) { + std::sort(values.begin(), values.end()); + return values; +} + +static size_t full_indexer_source(size_t layer) { + if (layer >= 3 && layer <= 5) { + return 2; + } + return layer; +} + +static void validate_real_captures(const real_model_result & result, const char * mode_name) { + require(result.decode_observer.steps.size() == 2, "real GLM capture did not contain two decode steps"); + for (size_t step = 0; step < result.decode_observer.steps.size(); ++step) { + const real_step_capture & capture = result.decode_observer.steps[step]; + for (size_t layer = 0; layer < real_layer_count; ++layer) { + require(!capture.top_k[layer].empty(), "real GLM capture is missing a top-k set"); + require(!capture.hidden_states[layer].empty(), "real GLM capture is missing a hidden state"); + } + for (const size_t full_layer : { 0U, 1U, 2U, 6U }) { + require(!capture.indexer_scores[full_layer].empty(), "real GLM capture is missing Full indexer scores"); + } + + const std::vector shared = sorted_indices(capture.top_k[2]); + for (size_t layer = 3; layer <= 5; ++layer) { + require(sorted_indices(capture.top_k[layer]) == shared, + "Shared GLM-DSA layer did not reuse the preceding Full top-k set"); + } + } + std::printf("real GLM mode %s captured top-k, Full scores, and seven hidden states for two decode steps\n", + mode_name); +} + +static bool top_k_difference_is_score_tie(const real_step_capture & reference, + size_t layer, + const std::vector & expected, + const std::vector & actual) { + const size_t source_layer = full_indexer_source(layer); + const auto & scores = reference.indexer_scores[source_layer]; + if (scores.empty() || expected.empty()) { + return false; + } + + float boundary = std::numeric_limits::infinity(); + for (int32_t index : expected) { + if (index < 0 || (size_t) index >= scores.size()) { + return false; + } + boundary = std::min(boundary, scores[index]); + } + + const float tolerance = 2e-5f * std::max(1.0f, std::fabs(boundary)); + for (int32_t index : expected) { + if (!std::binary_search(actual.begin(), actual.end(), index) && + std::fabs(scores[index] - boundary) > tolerance) { + return false; + } + } + for (int32_t index : actual) { + if (index < 0 || (size_t) index >= scores.size()) { + return false; + } + if (!std::binary_search(expected.begin(), expected.end(), index) && + std::fabs(scores[index] - boundary) > tolerance) { + return false; + } + } + return true; +} + +static void compare_real_captures(const real_model_result & reference, + const real_model_result & actual, + const char * mode_name, + double tolerance) { + require(reference.decode_observer.steps.size() == actual.decode_observer.steps.size(), + "real GLM capture step counts differ"); + + double max_hidden_nmse = 0.0; + double max_score_nmse = 0.0; + size_t tie_count = 0; + size_t max_hidden_step = 0; + size_t max_hidden_layer = 0; + size_t max_score_step = 0; + size_t max_score_layer = 0; + for (size_t step = 0; step < reference.decode_observer.steps.size(); ++step) { + const auto & expected_step = reference.decode_observer.steps[step]; + const auto & actual_step = actual.decode_observer.steps[step]; + for (size_t layer = 0; layer < real_layer_count; ++layer) { + const std::vector expected_top_k = sorted_indices(expected_step.top_k[layer]); + const std::vector actual_top_k = sorted_indices(actual_step.top_k[layer]); + if (expected_top_k != actual_top_k) { + require(top_k_difference_is_score_tie(expected_step, layer, expected_top_k, actual_top_k), + "real GLM top-k sets differ without a boundary-score tie"); + ++tie_count; + } + + const double hidden_nmse = + normalized_mean_squared_error(expected_step.hidden_states[layer], actual_step.hidden_states[layer]); + if (hidden_nmse > max_hidden_nmse) { + max_hidden_nmse = hidden_nmse; + max_hidden_step = step; + max_hidden_layer = layer; + } + + if (!expected_step.indexer_scores[layer].empty()) { + require(!actual_step.indexer_scores[layer].empty(), "real GLM Full indexer score capture differs"); + const double score_nmse = normalized_mean_squared_error(expected_step.indexer_scores[layer], + actual_step.indexer_scores[layer]); + if (score_nmse > max_score_nmse) { + max_score_nmse = score_nmse; + max_score_step = step; + max_score_layer = layer; + } + } + } + } + + std::printf( + "real GLM mode %s: max hidden NMSE %.3e (step %zu layer %zu), " + "max indexer-score NMSE %.3e (step %zu layer %zu), top-k ties %zu\n", + mode_name, max_hidden_nmse, max_hidden_step, max_hidden_layer, max_score_nmse, max_score_step, max_score_layer, + tie_count); + require(max_hidden_nmse <= tolerance, "real GLM hidden-state NMSE exceeds tolerance"); + require(max_score_nmse <= tolerance, "real GLM indexer-score NMSE exceeds tolerance"); +} + +static void require_native_metal_residency(const real_model_result & result) { + const real_decode_observer & observer = result.decode_observer; + if (!observer.non_metal_heavy_ops.empty()) { + for (const std::string & op : observer.non_metal_heavy_ops) { + std::fprintf(stderr, "non-Metal GLM-DSA heavy op: %s\n", op.c_str()); + } + } + require(observer.heavy_op_count > 0, "native Metal decode observed no heavy GLM-DSA operations"); + require(observer.heavy_op_count == observer.heavy_op_metal_count, + "native Metal decode assigned a heavy GLM-DSA operation outside Metal"); + require(observer.lightning_indexer_count == 8, + "native Metal decode did not execute four Full Lightning Indexers per step"); + require(observer.k_selected_count == 14 && observer.v_selected_count == 14, + "native Metal decode did not execute seven compact K/V gathers per step"); + require(observer.flash_attention_count == 14, + "native Metal decode did not execute seven flash-attention operations per step"); + require(observer.rope_count > 0, "native Metal decode did not execute RoPE operations"); +} + +static void compare_real_logits(const real_model_result & expected, + const real_model_result & actual, + const char * device_name, + double tolerance) { + const double error = normalized_mean_squared_error(expected.logits, actual.logits); + std::printf("real GLM layer %s logit NMSE: %.3e\n", device_name, error); + + constexpr size_t n_steps = 3; + require(expected.logits.size() % n_steps == 0, "real GLM logit fixture has an unexpected step count"); + const size_t n_vocab = expected.logits.size() / n_steps; + double max_step_error = 0.0; + bool top_1_matches = true; + for (size_t step = 0; step < n_steps; ++step) { + const size_t begin = step * n_vocab; + const size_t end = begin + n_vocab; + const double step_error = normalized_mean_squared_error(expected.logits, actual.logits, begin, end); + const auto expected_top = std::max_element(expected.logits.begin() + begin, expected.logits.begin() + end); + const auto actual_top = std::max_element(actual.logits.begin() + begin, actual.logits.begin() + end); + std::printf(" step %zu: NMSE %.3e, top-1 %zu/%zu\n", step, step_error, + static_cast(expected_top - expected.logits.begin()) - begin, + static_cast(actual_top - actual.logits.begin()) - begin); + max_step_error = std::max(max_step_error, step_error); + top_1_matches &= expected_top - expected.logits.begin() == actual_top - actual.logits.begin(); + } + require(error <= tolerance && max_step_error <= tolerance, "real GLM layer logits differ from CPU reference"); + require(top_1_matches, "real GLM layer top-1 logit differs from CPU reference"); +} + +static void test_real_layer_logits(const char * model_path) { + static constexpr real_execution_config reference_config = { + "CPU generic-indexer+dense-mask", + false, + false, + }; + static constexpr real_execution_config metal_reference_config = { + "Metal generic-indexer+dense-mask", + false, + false, + }; + static constexpr real_execution_config metal_fused_dense_config = { + "Metal fused-indexer+dense-mask", + true, + false, + }; + static constexpr real_execution_config metal_generic_compact_config = { + "Metal generic-indexer+compact-flash", + false, + true, + }; + static constexpr real_execution_config metal_native_config = { + "Metal fused-indexer+compact-flash", + true, + true, + }; + + ggml_backend_dev_t cpu_device = ggml_backend_dev_by_name("CPU"); + require(cpu_device != nullptr, "CPU backend device is unavailable"); + const real_model_result cpu = run_real_model(model_path, cpu_device, reference_config); + validate_real_captures(cpu, reference_config.name); + require(cpu.decode_observer.lightning_indexer_count == 0, + "CPU reference unexpectedly retained the fused Lightning Indexer"); + require(cpu.decode_observer.dense_sparse_mask_count == 14, + "CPU reference did not build seven dense sparse masks per decode step"); + + if (ggml_backend_dev_t blas_device = ggml_backend_dev_by_name("BLAS")) { + const real_model_result blas = run_real_model(model_path, blas_device, reference_config); + validate_real_captures(blas, "BLAS generic-indexer+dense-mask"); + compare_real_logits(cpu, blas, "BLAS fallback", 1e-7); + compare_real_captures(cpu, blas, "BLAS fallback", 1e-7); + } + + if (ggml_backend_dev_t metal_device = ggml_backend_dev_by_name("MTL0")) { + const real_model_result metal_reference = run_real_model(model_path, metal_device, metal_reference_config); + validate_real_captures(metal_reference, metal_reference_config.name); + + const real_model_result metal_fused_dense = run_real_model(model_path, metal_device, metal_fused_dense_config); + validate_real_captures(metal_fused_dense, metal_fused_dense_config.name); + compare_real_logits(metal_reference, metal_fused_dense, "Metal fused-indexer A/B", 2e-4); + compare_real_captures(metal_reference, metal_fused_dense, "Metal fused-indexer A/B", 2e-4); + + const real_model_result metal_generic_compact = + run_real_model(model_path, metal_device, metal_generic_compact_config); + validate_real_captures(metal_generic_compact, metal_generic_compact_config.name); + compare_real_logits(metal_reference, metal_generic_compact, "Metal compact-flash A/B", 2e-4); + compare_real_captures(metal_reference, metal_generic_compact, "Metal compact-flash A/B", 2e-4); + + const real_model_result metal_native = run_real_model(model_path, metal_device, metal_native_config); + validate_real_captures(metal_native, metal_native_config.name); + compare_real_logits(metal_generic_compact, metal_native, "Metal fused-indexer compact A/B", 2e-4); + compare_real_captures(metal_generic_compact, metal_native, "Metal fused-indexer compact A/B", 2e-4); + require_native_metal_residency(metal_native); + + compare_real_captures(cpu, metal_reference, metal_reference_config.name, 2e-4); + compare_real_logits(cpu, metal_reference, metal_reference_config.name, 2e-4); + compare_real_captures(cpu, metal_native, "Metal native", 2e-4); + compare_real_logits(cpu, metal_native, "Metal native", 2e-4); + } +} + +int main(int argc, char ** argv) { + llama_backend_init(); + try { + ggml_backend_load_all(); + test_indexer_scores_and_top_k(); + test_generic_indexer_mask_cast(); + test_interleaved_rope(); + test_compact_sparse_attention(); + test_glm_dsa_moe_precision(); + if (argc == 2 && std::strcmp(argv[1], "--stability") == 0) { + test_glm_dsa_decode_stability(); + } else if (argc == 3 && std::strcmp(argv[1], "--long-parity") == 0) { + test_glm_dsa_long_greedy_parity(argv[2]); + } else if (argc == 3 && std::strcmp(argv[1], "--real-model") == 0) { + test_real_layer_logits(argv[2]); + } else if (argc != 1) { + throw std::runtime_error( + "usage: test-glm-dsa [--stability | --long-parity model.gguf | --real-model model.gguf]"); + } + std::printf("GLM-DSA reference tests passed\n"); + llama_backend_free(); + return 0; + } catch (const std::exception & error) { + std::fprintf(stderr, "GLM-DSA reference test failed: %s\n", error.what()); + llama_backend_free(); + return 1; + } +} diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp index f39abe773fc6..678691820a06 100644 --- a/tests/test-llama-archs.cpp +++ b/tests/test-llama-archs.cpp @@ -10,7 +10,10 @@ // TODO: replace with #include "llama-ext.h" in the future #include "../src/llama-arch.h" #include "../src/llama-model-saver.h" +#include "../src/llama-model.h" +#include +#include #include #include #include @@ -77,7 +80,7 @@ static std::vector get_tokens(const uint32_t n_tokens, const uint32 return ret; } -static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) { +static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe, const bool glm_string_indexer_types = false) { gguf_context_ptr ret(gguf_init_empty()); llama_model_saver ms(arch, ret.get()); const uint32_t n_ctx = 128; @@ -101,12 +104,16 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) { n_layer = 22; // hparams.n_layer_kv_from_start = 20 is hardcoded } else if (arch == LLM_ARCH_DEEPSEEK2 || arch == LLM_ARCH_DEEPSEEK32 - || arch == LLM_ARCH_GLM_DSA || arch == LLM_ARCH_KIMI_LINEAR || arch == LLM_ARCH_MISTRAL4) { n_embd = 128; n_head = 1; n_ff = 192; + } else if (arch == LLM_ARCH_GLM_DSA) { + n_embd = 128; + n_head = 1; + n_ff = 192; + n_layer = 7; // cover the default Full/Shared cadence through the second full indexer } else if (arch == LLM_ARCH_NEMOTRON_H || arch == LLM_ARCH_NEMOTRON_H_MOE) { n_layer = 3; } else if (arch == LLM_ARCH_CHAMELEON) { @@ -199,6 +206,10 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) { ms.add_kv(LLM_KV_ATTENTION_INDEXER_HEAD_COUNT, uint32_t(1)); ms.add_kv(LLM_KV_ATTENTION_INDEXER_KEY_LENGTH, uint32_t(64)); ms.add_kv(LLM_KV_ATTENTION_INDEXER_TOP_K, uint32_t(8)); + if (arch == LLM_ARCH_GLM_DSA && glm_string_indexer_types) { + ms.add_kv(LLM_KV_ATTENTION_INDEXER_TYPES, + std::vector({ "full", "full", "full", "shared", "shared", "shared", "full" })); + } ms.add_kv(LLM_KV_ROPE_DIMENSION_SECTIONS, std::vector({n_embd_head/4, n_embd_head/4, n_embd_head/4, n_embd_head/4})); ms.add_kv(LLM_KV_TOKENIZER_MODEL, "no_vocab"); // ms.add_kv(LLM_KV_DENSE_2_FEAT_OUT, n_embd); @@ -252,9 +263,55 @@ static bool silent_model_load_progress(float /*progress*/, void * /*user_data*/) return true; } +struct glm_decode_observer { + int lightning_indexer_count = 0; + int k_selected_count = 0; + int v_selected_count = 0; + int64_t selected_rows = 0; + uint64_t lightning_indexer_layers = 0; +}; + +static bool observe_glm_decode(ggml_tensor * tensor, bool ask, void * user_data) { + if (!ask) { + return true; + } + + auto * observer = static_cast(user_data); + if (tensor->op == GGML_OP_LIGHTNING_INDEXER) { + ++observer->lightning_indexer_count; + uint32_t il = 0; + if (sscanf(tensor->name, "indexer_score-%u", &il) == 1 && il < 64) { + observer->lightning_indexer_layers |= 1ULL << il; + } + } else if (tensor->op == GGML_OP_GET_ROWS && strstr(tensor->name, "k_selected")) { + ++observer->k_selected_count; + observer->selected_rows = tensor->ne[1]; + } else if (tensor->op == GGML_OP_GET_ROWS && strstr(tensor->name, "v_selected")) { + ++observer->v_selected_count; + } + return false; +} + +static void require_glm_decode_observed(const glm_decode_observer & observer) { + constexpr uint64_t full_indexer_layers = (1ULL << 0) | (1ULL << 1) | (1ULL << 2) | (1ULL << 6); + GGML_ASSERT(observer.lightning_indexer_count >= 12); + GGML_ASSERT(observer.lightning_indexer_count % 4 == 0); + GGML_ASSERT(observer.lightning_indexer_layers == full_indexer_layers); + GGML_ASSERT(observer.k_selected_count == 14); + GGML_ASSERT(observer.v_selected_count == 14); + GGML_ASSERT(observer.selected_rows == 8); +} + +static bool has_metal_device(const std::vector & devices) { + return std::any_of(devices.begin(), devices.end(), [](ggml_backend_dev_t device) { + return strncmp(ggml_backend_dev_name(device), "MTL", 3) == 0; + }); +} + static std::pair get_model_and_ctx( struct gguf_context * gguf_ctx, FILE * file, const size_t seed, const std::vector & devs, - const llama_split_mode split_mode = LLAMA_SPLIT_MODE_LAYER, bool encode = false) { + const llama_split_mode split_mode = LLAMA_SPLIT_MODE_LAYER, bool encode = false, + glm_decode_observer * decode_observer = nullptr) { GGML_ASSERT((gguf_ctx == nullptr) != (file == nullptr)); llama_model_params model_params = llama_model_default_params(); model_params.progress_callback = silent_model_load_progress; @@ -270,6 +327,10 @@ static std::pair get_model_and_ctx( if (!encode) { ctx_params.n_ubatch = 64; } + if (decode_observer) { + ctx_params.cb_eval = observe_glm_decode; + ctx_params.cb_eval_user_data = decode_observer; + } size_t tmp = seed; llama_model_ptr model(gguf_ctx != nullptr ? @@ -278,24 +339,60 @@ static std::pair get_model_and_ctx( if (!model) { throw std::runtime_error("failed to create llama model"); } + if (model->arch == LLM_ARCH_GLM_DSA) { + static constexpr std::array expected = { true, true, true, false, false, false, true }; + GGML_ASSERT(model->hparams.n_layer() == expected.size()); + for (uint32_t il = 0; il < expected.size(); ++il) { + GGML_ASSERT(model->hparams.is_indexer_full(il) == expected[il]); + } + } llama_context_ptr lctx(llama_init_from_model(model.get(), ctx_params)); if (!lctx) { throw std::runtime_error("failed to create llama context"); } + if (decode_observer) { + *decode_observer = {}; + } return std::make_pair(std::move(model), std::move(lctx)); } -static std::vector get_logits( - llama_model * model, llama_context * lctx, const std::vector & tokens, bool encode = false) { - const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model)); - const uint32_t n_ctx = llama_n_ctx(lctx); - const uint32_t n_tokens = tokens.size(); +static void test_glm_indexer_metadata_compatibility(const size_t seed) { + { + gguf_context_ptr gguf_ctx = get_gguf_ctx(LLM_ARCH_GLM_DSA, true); + auto model_and_ctx = get_model_and_ctx(gguf_ctx.get(), nullptr, seed, {}); + GGML_ASSERT(model_and_ctx.first); + GGML_ASSERT(model_and_ctx.second); + } + + { + gguf_context_ptr gguf_ctx = get_gguf_ctx(LLM_ARCH_GLM_DSA, true); + static constexpr std::array indexer_types = { true, true, true, false, false, false, true }; + const std::string key = std::string(llm_arch_name(LLM_ARCH_GLM_DSA)) + ".attention.indexer.types"; + gguf_set_arr_data(gguf_ctx.get(), key.c_str(), GGUF_TYPE_BOOL, indexer_types.data(), indexer_types.size()); + auto model_and_ctx = get_model_and_ctx(gguf_ctx.get(), nullptr, seed, {}); + GGML_ASSERT(model_and_ctx.first); + GGML_ASSERT(model_and_ctx.second); + } + + { + gguf_context_ptr gguf_ctx = get_gguf_ctx(LLM_ARCH_GLM_DSA, true, true); + auto model_and_ctx = get_model_and_ctx(gguf_ctx.get(), nullptr, seed, {}); + GGML_ASSERT(model_and_ctx.first); + GGML_ASSERT(model_and_ctx.second); + } +} + +static void append_logits( + llama_model * model, llama_context * lctx, const std::vector & tokens, + uint32_t begin, uint32_t end, bool encode, std::vector & logits) { + const uint32_t n_vocab = llama_vocab_n_tokens(llama_model_get_vocab(model)); + const uint32_t n_ctx = llama_n_ctx(lctx); llama_batch batch = llama_batch_init(n_ctx, 0, 1); - GGML_ASSERT(n_tokens <= n_ctx); - for (uint32_t pos = 0; pos < n_tokens; pos++) { + GGML_ASSERT(begin < end && end <= n_ctx); + for (uint32_t pos = begin; pos < end; pos++) { common_batch_add(batch, tokens[pos], pos, {0}, true); } - batch.n_tokens = n_tokens; + batch.n_tokens = end - begin; if (encode) { if (llama_encode(lctx, batch)) { llama_batch_free(batch); @@ -307,15 +404,29 @@ static std::vector get_logits( throw std::runtime_error("failed to decode batch"); } - std::vector ret; - ret.reserve(n_tokens*n_vocab); - for (uint32_t i = 0; i < n_tokens; i++) { + for (int32_t i = 0; i < batch.n_tokens; i++) { const float * logits_ith = llama_get_logits_ith(lctx, i); for (uint32_t j = 0; j < n_vocab; j++) { - ret.push_back(logits_ith[j]); + logits.push_back(logits_ith[j]); } } llama_batch_free(batch); +} + +static std::vector get_logits( + llama_model * model, llama_context * lctx, const std::vector & tokens, bool encode = false) { + const uint32_t n_tokens = tokens.size(); + std::vector ret; + ret.reserve(n_tokens*llama_vocab_n_tokens(llama_model_get_vocab(model))); + + const bool dsa_decode = model->arch == LLM_ARCH_GLM_DSA || model->arch == LLM_ARCH_DEEPSEEK32; + if (dsa_decode && !encode && n_tokens >= 3) { + append_logits(model, lctx, tokens, 0, n_tokens - 2, false, ret); + append_logits(model, lctx, tokens, n_tokens - 2, n_tokens - 1, false, ret); + append_logits(model, lctx, tokens, n_tokens - 1, n_tokens, false, ret); + } else { + append_logits(model, lctx, tokens, 0, n_tokens, encode, ret); + } return ret; } @@ -497,6 +608,10 @@ static int test_backends(const llm_arch target_arch, const size_t seed, const gg ud->original_logger.callback(level_eff, text, ud->original_logger.user_data); }, &ud); + if (target_arch == LLM_ARCH_UNKNOWN || target_arch == LLM_ARCH_GLM_DSA) { + test_glm_indexer_metadata_compatibility(seed); + } + const std::vector tokens = get_tokens(128, 128, seed); struct device_config { @@ -573,9 +688,10 @@ static int test_backends(const llm_arch target_arch, const size_t seed, const gg continue; } const std::string config_name = moe ? "MoE" : "Dense"; - gguf_context_ptr gguf_ctx = get_gguf_ctx(arch, moe); + gguf_context_ptr gguf_ctx = get_gguf_ctx(arch, moe, arch == LLM_ARCH_GLM_DSA); std::pair model_and_ctx_cpu; std::vector logits_cpu; + glm_decode_observer decode_observer_cpu; for (device_config & dc : dev_configs) { // print test config first; should anything fail during model loading or inference, at least we know which test case caused it printf(template_row_cfg.c_str(), @@ -584,6 +700,7 @@ static int test_backends(const llm_arch target_arch, const size_t seed, const gg std::pair model_and_ctx_dev; std::vector logits_dev; + glm_decode_observer decode_observer_dev; std::string status_nmse = "\033[1;33mSKIP\033[0m"; std::string status_roundtrip = "\033[1;33mSKIP\033[0m"; char nmse_str[12] = {0}; @@ -593,12 +710,23 @@ static int test_backends(const llm_arch target_arch, const size_t seed, const gg #endif // GGML_USE_WEBGPU if (!skip) { if (logits_cpu.empty()) { - model_and_ctx_cpu = get_model_and_ctx(gguf_ctx.get(), nullptr, seed, {}, LLAMA_SPLIT_MODE_LAYER, encode); + model_and_ctx_cpu = get_model_and_ctx( + gguf_ctx.get(), nullptr, seed, {}, LLAMA_SPLIT_MODE_LAYER, encode, + arch == LLM_ARCH_GLM_DSA ? &decode_observer_cpu : nullptr); logits_cpu = get_logits(model_and_ctx_cpu.first.get(), model_and_ctx_cpu.second.get(), tokens, encode); + if (arch == LLM_ARCH_GLM_DSA) { + require_glm_decode_observed(decode_observer_cpu); + } } if (dc.split_mode != LLAMA_SPLIT_MODE_TENSOR || llm_arch_supports_sm_tensor(arch)) { - model_and_ctx_dev = get_model_and_ctx(gguf_ctx.get(), nullptr, seed, dc.devs, dc.split_mode, encode); + const bool expect_compact_decode = arch == LLM_ARCH_GLM_DSA && has_metal_device(dc.devs); + model_and_ctx_dev = get_model_and_ctx( + gguf_ctx.get(), nullptr, seed, dc.devs, dc.split_mode, encode, + expect_compact_decode ? &decode_observer_dev : nullptr); logits_dev = get_logits(model_and_ctx_dev.first.get(), model_and_ctx_dev.second.get(), tokens, encode); + if (expect_compact_decode) { + require_glm_decode_observed(decode_observer_dev); + } const double nmse_val = nmse(logits_cpu, logits_dev); snprintf(nmse_str, sizeof(nmse_str), "(%.2e)", nmse_val); status_nmse = "\033[1;32mOK\033[0m"; @@ -619,9 +747,16 @@ static int test_backends(const llm_arch target_arch, const size_t seed, const gg ms.save(file); rewind(file); - auto model_and_ctx_roundtrip = get_model_and_ctx(nullptr, file, seed, dc.devs, dc.split_mode, encode); + glm_decode_observer decode_observer_roundtrip; + const bool expect_compact_decode = arch == LLM_ARCH_GLM_DSA && has_metal_device(dc.devs); + auto model_and_ctx_roundtrip = get_model_and_ctx( + nullptr, file, seed, dc.devs, dc.split_mode, encode, + expect_compact_decode ? &decode_observer_roundtrip : nullptr); const std::vector logits_roundtrip = get_logits( model_and_ctx_roundtrip.first.get(), model_and_ctx_roundtrip.second.get(), tokens, encode); + if (expect_compact_decode) { + require_glm_decode_observed(decode_observer_roundtrip); + } status_roundtrip = "\033[1;32mOK\033[0m"; GGML_ASSERT(logits_roundtrip.size() == logits_dev.size()); for (size_t i = 0; i < logits_roundtrip.size(); i++) {