Transformer Block - Multi-Head Attention */
| 274 | |
| 275 | /* Transformer Block - Multi-Head Attention */ |
| 276 | struct BlockMHA : public Block { |
| 277 | BlockMHA( |
| 278 | int layer_i, |
| 279 | const std::shared_ptr<Config> config, |
| 280 | const Tensor* rms_att_weight, |
| 281 | const Tensor* rms_q_a_weight, |
| 282 | const Tensor* rms_kv_a_weight, |
| 283 | const Tensor* rms_ffn_weight, |
| 284 | const Tensor* wq, |
| 285 | const Tensor* sq, |
| 286 | const Tensor* wq_a, |
| 287 | const Tensor* sq_a, |
| 288 | const Tensor* wkv_a, |
| 289 | const Tensor* skv_a, |
| 290 | const Tensor* wq_b, |
| 291 | const Tensor* sq_b, |
| 292 | const Tensor* wkv_b, |
| 293 | const Tensor* skv_b, |
| 294 | const Tensor* wo, |
| 295 | const Tensor* so, |
| 296 | const Tensor* w1, |
| 297 | const Tensor* s1, |
| 298 | const Tensor* w2, |
| 299 | const Tensor* s2, |
| 300 | const Tensor* w3, |
| 301 | const Tensor* s3, |
| 302 | const Tensor* shared_w1, |
| 303 | const Tensor* shared_s1, |
| 304 | const Tensor* shared_w2, |
| 305 | const Tensor* shared_s2, |
| 306 | const Tensor* shared_w3, |
| 307 | const Tensor* shared_s3, |
| 308 | const Tensor* moegate, |
| 309 | const Tensor* moegate_bias |
| 310 | ); |
| 311 | ~BlockMHA() override; |
| 312 | |
| 313 | float* rms_q_a_weight() const { return _rms_q_a_weight ? static_cast<float*>(_rms_q_a_weight->data) : nullptr; } |
| 314 | float* rms_kv_a_weight() const { return _rms_kv_a_weight ? static_cast<float*>(_rms_kv_a_weight->data) : nullptr; } |
| 315 | std::optional<QTensor> wq() const { return _wq; } |
| 316 | std::optional<QTensor> wq_a() const { return _wq_a; } |
| 317 | std::optional<QTensor> wq_b() const { return _wq_b; } |
| 318 | std::optional<QTensor> wkv_a() const { return _wkv_a; } |
| 319 | std::optional<QTensor> wkv_b() const { return _wkv_b; } |
| 320 | std::optional<QTensor> wo() const { return _wo; } |
| 321 | f16_t* key_cache() const { return _key_cache; } |
| 322 | f16_t* key_cache(int pos) const { return _key_cache + pos * _config->head_dim * _config->n_heads; } |
| 323 | f16_t* value_cache() const { return _value_cache; } |
| 324 | f16_t* value_cache(int pos) const { return _value_cache + pos * _config->v_head_dim * _config->n_heads; } |
| 325 | |
| 326 | double active_bytes(size_t pos) const override; |
| 327 | |
| 328 | protected: |
| 329 | void attention_impl( |
| 330 | InferenceState& s, int pos, int kv_sink, int kv_pos, int kv_len |
| 331 | ) const override; |
| 332 | |
| 333 | private: |
nothing calls this directly
no outgoing calls
no test coverage detected