cwenzi/neuroflow-cpp
1
1#include "test_framework.hpp"
2#include "neuroflow/causal_lm.hpp"
3#include <cmath>
4
5using namespace neuroflow;
6
7TEST(GQA, MHAWhenQEqualsKV) {
8 CausalLMConfig config;
9 config.d_model = 64;
10 config.num_attn_heads = 4;
11 config.n_kv_heads = 4;
12 config.vocab_size = 100;
13 config.max_seq_len = 32;
14 config.use_rope = false;
15 config.use_qk_norm = false;
16
17 CausalLMHead lm(config);
18 lm.eval();
19
20 std::vector<size_t> ids = {1, 2, 3, 4};
21 Tensor logits = lm.forward(ids);
22 EXPECT_EQ(logits.shape_[0], 1u);
23 EXPECT_EQ(logits.shape_[1], 100u);
24}
25
26TEST(GQA, GQAReducedKVHeads) {
27 CausalLMConfig config;
28 config.d_model = 64;
29 config.num_attn_heads = 4;
30 config.n_kv_heads = 2;
31 config.vocab_size = 100;
32 config.max_seq_len = 32;
33 config.use_rope = false;
34 config.use_qk_norm = false;
35
36 CausalLMHead lm(config);
37 lm.eval();
38
39 std::vector<size_t> ids = {1, 2, 3, 4};
40 Tensor logits = lm.forward(ids);
41 EXPECT_EQ(logits.shape_[0], 1u);
42 EXPECT_EQ(logits.shape_[1], 100u);
43
44 EXPECT_FALSE(std::isnan(logits.as_fp32()[0]));
45}
46
47TEST(GQA, InvalidRatioThrows) {
48 EXPECT_THROW({
49 CausalSelfAttention attn(64, 5, 2, false, 32, false);
50 }, std::invalid_argument);
51}
52
53TEST(GQA, TrainingBackwardWithGQA) {
54 CausalLMConfig config;
55 config.d_model = 64;
56 config.num_attn_heads = 4;
57 config.n_kv_heads = 2;
58 config.vocab_size = 100;
59 config.max_seq_len = 32;
60 config.use_rope = false;
61 config.use_qk_norm = false;
62
63 CausalLMHead lm(config);
64 lm.train();
65
66 std::vector<size_t> ids = {1, 2, 3, 4};
67 Tensor logits = lm.forward_for_training(ids);
68
69 Tensor grad({1, 100}, QuantType::FP32);
70 float* gp = grad.as_fp32();
71 for (size_t i = 0; i < 100; ++i) gp[i] = 0.01f;
72
73 auto grads = lm.backward_from_logits(grad);
74 EXPECT_GT(grads.attn_grads.size(), 0u);
75 EXPECT_GT(grads.attn_grads[0].w_q_weight_grad.numel(), 0u);
76 EXPECT_GT(grads.attn_grads[0].w_k_weight_grad.numel(), 0u);
77 EXPECT_GT(grads.attn_grads[0].w_v_weight_grad.numel(), 0u);
78}
79
80TEST(GQA, KVParamsSmallerWithGQA) {
81 CausalLMConfig config_mha;
82 config_mha.d_model = 64;
83 config_mha.num_attn_heads = 4;
84 config_mha.n_kv_heads = 4;
85 config_mha.vocab_size = 100;
86 config_mha.max_seq_len = 32;
87 config_mha.use_rope = false;
88 config_mha.use_qk_norm = false;
89
90 CausalLMConfig config_gqa;
91 config_gqa.d_model = 64;
92 config_gqa.num_attn_heads = 4;
93 config_gqa.n_kv_heads = 2;
94 config_gqa.vocab_size = 100;
95 config_gqa.max_seq_len = 32;
96 config_gqa.use_rope = false;
97 config_gqa.use_qk_norm = false;
98
99 CausalLMHead lm_mha(config_mha);
100 CausalLMHead lm_gqa(config_gqa);
101
102 size_t mha_kv_params = 0;
103 size_t gqa_kv_params = 0;
104 for (auto& attn : lm_mha.attn_layers_) {
105 mha_kv_params += attn->w_k->weight.numel() + attn->w_v->weight.numel();
106 }
107 for (auto& attn : lm_gqa.attn_layers_) {
108 gqa_kv_params += attn->w_k->weight.numel() + attn->w_v->weight.numel();
109 }
110
111 EXPECT_LT(gqa_kv_params, mha_kv_params);
112}
113
114int main() { RUN_ALL_TESTS(); }
115 