Minor delta-net tweaks (#1429)

This commit is contained in:
Kawrakow 2026-03-16 13:59:55 +01:00 committed by GitHub
parent fe92e30d1e
commit e3f5f3d823
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
3 changed files with 41 additions and 34 deletions

View File

@ -17576,23 +17576,24 @@ static void ggml_compute_forward_l2_norm_f32(
GGML_ASSERT(eps >= 0.0f);
for (int64_t i03 = 0; i03 < ne03; i03++) {
for (int64_t i02 = 0; i02 < ne02; i02++) {
for (int64_t i01 = ith; i01 < ne01; i01 += nth) {
const float * x = (const float *) ((const char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03);
int nrows = ne01*ne02*ne03;
int nrows_per_thread = (nrows + nth - 1)/nth;
int first = ith*nrows_per_thread;
int last = MIN(first + nrows_per_thread, nrows);
ggml_float sum = 0.0;
for (int64_t i00 = 0; i00 < ne00; i00++) {
sum += (ggml_float) (x[i00] * x[i00]);
}
float * y = (float *) ((char *) dst->data + i01*nb1 + i02*nb2 + i03*nb3);
memcpy(y, x, ne00 * sizeof(float));
const float scale = 1.0f/fmaxf(sqrtf(sum), eps);
ggml_vec_scale_f32(ne00, y, scale);
}
for (int ir = first; ir < last; ++ir) {
int i03 = ir/(ne01*ne02);
int i02 = (ir - i03*ne01*ne02)/ne01;
int i01 = ir - i03*ne01*ne02 - i02*ne01;
const float * x = (const float *) ((const char *) src0->data + i01*nb01 + i02*nb02 + i03*nb03);
float * y = (float *) ((char *) dst->data + i01*nb1 + i02*nb2 + i03*nb3);
ggml_float sum = 0.0;
for (int64_t i00 = 0; i00 < ne00; i00++) {
sum += (ggml_float) (x[i00] * x[i00]);
}
const float scale = 1.0f/fmaxf(sqrtf(sum), eps);
for (int j = 0; j < (int)ne00; ++j) {
y[j] = scale * x[j];
}
}
}

View File

@ -444,15 +444,17 @@ void iqk_mul_multi_add(struct ggml_tensor * dst, int ith, int nth) {
for (int ir = first; ir < last; ++ir) {
auto c0 = (const char *)src0->data + ir*src0->nb[2];
auto c1 = (const char *)src1->data + ir*src1->nb[2];
auto cy = ( char *) dst->data + ir* dst->nb[1];
std::memset(cy, 0, ne00*sizeof(float));
for (int j = 0; j < ne01; ++j) {
auto x0 = (const float *)c0;
auto x1 = (const float *)c1;
auto y = ( float *)cy;
for (int k = 0; k < ne00; ++k) y[k] += x0[k] * x1[0];
auto cy = ( char *)dst->data + ir* dst->nb[1];
auto y = ( float *)cy;
auto x0 = (const float *)c0;
auto x1 = (const float *)c1;
for (int k = 0; k < ne00; ++k) y[k] = x0[k] * x1[0];
for (int j = 1; j < ne01; ++j) {
c0 += src0->nb[1];
c1 += src1->nb[1];
x0 = (const float *)c0;
x1 = (const float *)c1;
for (int k = 0; k < ne00; ++k) y[k] += x0[k] * x1[0];
}
}
}

View File

@ -80,17 +80,16 @@ std::pair<ggml_tensor *, ggml_tensor *> delta_net::build_fused_delta_net(ggml_co
int il, const llm_build_cb & cb, int repeat_type) {
const int64_t S_k = q->ne[0];
const int64_t H_k = q->ne[1];
const int64_t n_tokens = q->ne[2];
const int64_t H_k = q->ne[2];
const int64_t n_tokens = q->ne[1];
const int64_t n_seqs = q->ne[3];
const int64_t S_v = v->ne[0];
const int64_t H_v = v->ne[1];
GGML_ASSERT(q->ne[0] == S_k && q->ne[1] == H_k && q->ne[2] == n_tokens && q->ne[3] == n_seqs);
GGML_ASSERT(k->ne[0] == S_k && k->ne[1] == H_k && k->ne[2] == n_tokens && k->ne[3] == n_seqs);
GGML_ASSERT(q->ne[0] == S_k && q->ne[2] == H_k && q->ne[1] == n_tokens && q->ne[3] == n_seqs);
GGML_ASSERT(k->ne[0] == S_k && k->ne[2] == H_k && k->ne[1] == n_tokens && k->ne[3] == n_seqs);
GGML_ASSERT(v->ne[2] == n_tokens);
GGML_ASSERT(k->ne[2] == n_tokens);
GGML_ASSERT(g->ne[0] == H_v && g->ne[1] == n_tokens && g->ne[2] == n_seqs);
GGML_ASSERT(beta->ne[0] == H_v && beta->ne[2] == n_tokens && beta->ne[3] == n_seqs);
GGML_ASSERT(state->ne[0] == S_v && state->ne[1] == S_v && state->ne[2] == H_v && state->ne[3] == n_seqs);
@ -104,14 +103,10 @@ std::pair<ggml_tensor *, ggml_tensor *> delta_net::build_fused_delta_net(ggml_co
cb(g, "g_in", il);
cb(state,"state_in", il);
q = ggml_permute(ctx0, q, 0, 2, 1, 3);
k = ggml_permute(ctx0, k, 0, 2, 1, 3);
v = ggml_permute(ctx0, v, 0, 2, 1, 3);
g = ggml_permute(ctx0, g, 2, 0, 3, 1);
beta = ggml_permute(ctx0, beta, 2, 0, 1, 3);
if (n_seqs > 1 || n_tokens > 1) {
q = ggml_cont_4d(ctx0, q, S_k, n_tokens, H_k, n_seqs);
k = ggml_cont_4d(ctx0, k, S_k, n_tokens, H_k, n_seqs);
v = ggml_cont_4d(ctx0, v, S_v, n_tokens, H_v, n_seqs);
g = ggml_cont_4d(ctx0, g, n_tokens, 1, H_v, n_seqs);
beta = ggml_cont_4d(ctx0, beta, 1, n_tokens, H_v, n_seqs);
@ -363,8 +358,17 @@ ggml_tensor * delta_net::build_qkv(ggml_context * ctx0, ggml_tensor * state_stor
cb(k_conv, "k_conv", il);
cb(v_conv, "v_conv", il);
q_conv = ggml_l2_norm(ctx0, q_conv, eps_norm);
k_conv = ggml_l2_norm(ctx0, k_conv, eps_norm);
if (n_seq_tokens > 1) {
q_conv = ggml_permute(ctx0, q_conv, 0, 2, 1, 3);
k_conv = ggml_permute(ctx0, k_conv, 0, 2, 1, 3);
q_conv = ggml_l2_norm(ctx0, q_conv, eps_norm);
k_conv = ggml_l2_norm(ctx0, k_conv, eps_norm);
} else {
q_conv = ggml_l2_norm(ctx0, q_conv, eps_norm);
k_conv = ggml_l2_norm(ctx0, k_conv, eps_norm);
q_conv = ggml_permute(ctx0, q_conv, 0, 2, 1, 3);
k_conv = ggml_permute(ctx0, k_conv, 0, 2, 1, 3);
}
cb(q_conv, "q_conv_normed", il);
cb(k_conv, "k_conv_normed", il);