MCPcopy Create free account
hub / github.com/andrewkchan/deepseek.cpp / BlockMHA

Class BlockMHA

src/model.h:276–363  ·  view source on GitHub ↗

Transformer Block - Multi-Head Attention */

Source from the content-addressed store, hash-verified

274
275/* Transformer Block - Multi-Head Attention */
276struct 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
328protected:
329 void attention_impl(
330 InferenceState& s, int pos, int kv_sink, int kv_pos, int kv_len
331 ) const override;
332
333private:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected