| 4418 | } |
| 4419 | |
| 4420 | static struct ggml_tensor * llm_build_norm( |
| 4421 | struct ggml_context * ctx, |
| 4422 | struct ggml_tensor * cur, |
| 4423 | const llama_hparams & hparams, |
| 4424 | struct ggml_tensor * mw, |
| 4425 | struct ggml_tensor * mb, |
| 4426 | llm_norm_type type, |
| 4427 | const llm_build_cb & cb, |
| 4428 | int il) { |
| 4429 | switch (type) { |
| 4430 | case LLM_NORM: cur = ggml_norm (ctx, cur, hparams.f_norm_eps); break; |
| 4431 | case LLM_NORM_RMS: cur = ggml_rms_norm(ctx, cur, hparams.f_norm_rms_eps); break; |
| 4432 | } |
| 4433 | |
| 4434 | if (mw || mb) { |
| 4435 | cb(cur, "norm", il); |
| 4436 | } |
| 4437 | |
| 4438 | if (mw) { |
| 4439 | cur = ggml_mul(ctx, cur, mw); |
| 4440 | if (mb) { |
| 4441 | cb(cur, "norm_w", il); |
| 4442 | } |
| 4443 | } |
| 4444 | |
| 4445 | if (mb) { |
| 4446 | cur = ggml_add(ctx, cur, mb); |
| 4447 | } |
| 4448 | |
| 4449 | return cur; |
| 4450 | } |
| 4451 | |
| 4452 | static struct ggml_tensor * llm_build_ffn( |
| 4453 | struct ggml_context * ctx, |