Transformer Block - Multi-Latent Attention */
| 364 | |
| 365 | /* Transformer Block - Multi-Latent Attention */ |
| 366 | struct BlockMLA : public Block { |
| 367 | BlockMLA( |
| 368 | int layer_i, |
| 369 | const std::shared_ptr<Config> config, |
| 370 | const Tensor* rms_att_weight, |
| 371 | const Tensor* rms_q_a_weight, |
| 372 | const Tensor* rms_kv_a_weight, |
| 373 | const Tensor* rms_ffn_weight, |
| 374 | const Tensor* wq_a, |
| 375 | const Tensor* sq_a, |
| 376 | const Tensor* wkv_a, |
| 377 | const Tensor* skv_a, |
| 378 | const Tensor* wo, |
| 379 | const Tensor* so, |
| 380 | const Tensor* wc, |
| 381 | const Tensor* sc, |
| 382 | const Tensor* wq_rope_b, |
| 383 | const Tensor* sq_rope_b, |
| 384 | const Tensor* wv_b, |
| 385 | const Tensor* sv_b, |
| 386 | const Tensor* w1, |
| 387 | const Tensor* s1, |
| 388 | const Tensor* w2, |
| 389 | const Tensor* s2, |
| 390 | const Tensor* w3, |
| 391 | const Tensor* s3, |
| 392 | const Tensor* shared_w1, |
| 393 | const Tensor* shared_s1, |
| 394 | const Tensor* shared_w2, |
| 395 | const Tensor* shared_s2, |
| 396 | const Tensor* shared_w3, |
| 397 | const Tensor* shared_s3, |
| 398 | const Tensor* moegate, |
| 399 | const Tensor* moegate_bias |
| 400 | ); |
| 401 | ~BlockMLA() override; |
| 402 | |
| 403 | float* rms_q_a_weight() const { return _rms_q_a_weight ? static_cast<float*>(_rms_q_a_weight->data) : nullptr; } |
| 404 | float* rms_kv_a_weight() const { return _rms_kv_a_weight ? static_cast<float*>(_rms_kv_a_weight->data) : nullptr; } |
| 405 | std::optional<QTensor> wq_a() const { return _wq_a; } |
| 406 | std::optional<QTensor> wkv_a() const { return _wkv_a; } |
| 407 | std::optional<QTensor> wo() const { return _wo; } |
| 408 | std::optional<QTensor> wc() const { return _wc; } |
| 409 | std::optional<QTensor> wq_rope_b() const { return _wq_rope_b; } |
| 410 | std::optional<QTensor> wv_b() const { return _wv_b; } |
| 411 | f16_t* kv_nope_cache() const { return _kv_nope_cache; } |
| 412 | f16_t* kv_nope_cache(int pos) const { return _kv_nope_cache + pos * _config->kv_lora_rank; } |
| 413 | f16_t* kv_rope_cache() const { return _kv_rope_cache; } |
| 414 | f16_t* kv_rope_cache(int pos) const { return _kv_rope_cache + pos * _config->qk_rope_head_dim; } |
| 415 | |
| 416 | double active_bytes(size_t pos) const override; |
| 417 | |
| 418 | protected: |
| 419 | void attention_impl( |
| 420 | InferenceState& s, int pos, int kv_sink, int kv_pos, int kv_len |
| 421 | ) const override; |
| 422 | private: |
| 423 | template <typename T> |
nothing calls this directly
no outgoing calls
no test coverage detected