| 307 | const modules::LinearModule q_proj(binding::linear_config(config.dim, config.heads * config.dim_head, false)); |
| 308 | const modules::LinearModule k_proj(binding::linear_config(config.dim, config.heads * config.dim_head, false)); |
| 309 | const modules::LinearModule v_proj(binding::linear_config(config.dim, config.heads * config.dim_head, false)); |
| 310 | const modules::LinearModule gate_proj(binding::linear_config(config.dim, config.heads, true)); |
| 311 | const modules::LinearModule out_proj(binding::linear_config(config.heads * config.dim_head, config.dim, false)); |
| 312 | |
| 313 | auto x = build_reference_rms_norm(ctx, input, config.dim, weights.norm); |
| 314 | core::TensorValue q; |
| 315 | core::TensorValue k; |
| 316 | core::TensorValue v; |
| 317 | if (config.fused_qkv) { |
| 318 | const int64_t inner = config.heads * config.dim_head; |
| 319 | const modules::LinearModule qkv_proj( |
| 320 | binding::linear_config(config.dim, 3 * inner, false)); |
| 321 | auto qkv = qkv_proj.build( |
| 322 | ctx, |
| 323 | x, |
| 324 | binding::linear_data( |
| 325 | ctx, weights.qkv.weight, weights.qkv.bias)); |
| 326 | q = modules::SliceModule({2, 0, inner}).build(ctx, qkv); |
no test coverage detected