CoolFace
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
1likes
test_gqa.cpp115 linesDownload Raw Back to tests
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