CoolFace
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
1likes
test_generative.cpp227 linesDownload Raw Back to tests
1#include <iostream>
2#include <cassert>
3#include <cmath>
4#include "neuroflow/generative.hpp"
5
6using namespace neuroflow;
7
8void test_tokenizer() {
9    std::cout << "=== Tokenizer Test ===" << std::endl;
10
11    BPETokenizer tok;
12    tok.add_vocab("你", 4);
13    tok.add_vocab("好", 5);
14    tok.add_vocab("世", 6);
15    tok.add_vocab("界", 7);
16    tok.add_vocab("hello", 8);
17    tok.add_vocab(" ", 9);
18    tok.add_vocab("world", 10);
19    tok.set_vocab_size(11);
20
21    auto ids = tok.encode("你好世界");
22    std::cout << "encode('你好世界') = [";
23    for (size_t i = 0; i < ids.size(); ++i) {
24        std::cout << ids[i];
25        if (i < ids.size() - 1) std::cout << ", ";
26    }
27    std::cout << "]" << std::endl;
28
29    assert(ids[0] == tok.bos_id());
30    assert(ids.back() == tok.eos_id());
31    std::cout << "  BOS/EOS check: PASS" << std::endl;
32
33    auto decoded = tok.decode(ids);
34    std::cout << "decode result: '" << decoded << "'" << std::endl;
35    std::cout << "  Tokenizer test PASSED" << std::endl;
36}
37
38void test_causal_lm_head() {
39    std::cout << "\n=== CausalLMHead Test ===" << std::endl;
40
41    CausalLMConfig config;
42    config.vocab_size = 100;
43    config.d_model = 32;
44    config.max_seq_len = 64;
45    config.causal_window_size = 8;
46    config.sae_k = 16;
47    config.ntm_memory_slots = 4;
48    config.use_mla = false;
49
50    CausalLMHead lm(config);
51    std::cout << "  CausalLMHead constructed: vocab=" << config.vocab_size
52              << " d_model=" << config.d_model << std::endl;
53
54    std::vector<size_t> ids = {2, 5, 10, 20, 3};
55    Tensor logits = lm.forward(ids);
56    std::cout << "  forward() output shape: [" << logits.shape_[0] << ", " << logits.shape_[1] << "]" << std::endl;
57    assert(logits.shape_[0] == 1);
58    assert(logits.shape_[1] == config.vocab_size);
59
60    float max_logit = *std::max_element(logits.as_fp32(), logits.as_fp32() + logits.numel());
61    float min_logit = *std::min_element(logits.as_fp32(), logits.as_fp32() + logits.numel());
62    std::cout << "  logits range: [" << min_logit << ", " << max_logit << "]" << std::endl;
63    assert(!std::isnan(max_logit) && !std::isnan(min_logit));
64    std::cout << "  NaN check: PASS" << std::endl;
65
66    lm.clear_cache();
67    Tensor step_logits = lm.forward_step(5, 0);
68    std::cout << "  forward_step() output shape: [" << step_logits.shape_[0] << ", " << step_logits.shape_[1] << "]" << std::endl;
69    assert(step_logits.shape_[1] == config.vocab_size);
70    std::cout << "  CausalLMHead test PASSED" << std::endl;
71}
72
73void test_sampling_strategies() {
74    std::cout << "\n=== Sampling Strategy Test ===" << std::endl;
75
76    std::mt19937 rng(42);
77
78    Tensor logits({1, 10}, QuantType::FP32);
79    float* data = logits.as_fp32();
80    for (size_t i = 0; i < 10; ++i) data[i] = static_cast<float>(i) * 0.5f;
81
82    GenerateConfig config;
83    config.temperature = 1.0f;
84    config.top_k = 5;
85    config.top_p = 0.9f;
86    config.repetition_penalty = 1.0f;
87
88    GreedyDecoding greedy;
89    Tensor greedy_probs = greedy.apply(logits.clone(), config, {});
90    size_t greedy_id = greedy.sample(greedy_probs, rng);
91    std::cout << "  Greedy: selected token " << greedy_id << " (expected 9)" << std::endl;
92    assert(greedy_id == 9);
93
94    rng.seed(42);
95    TopKSampling topk;
96    Tensor topk_probs = topk.apply(logits.clone(), config, {});
97    size_t topk_id = topk.sample(topk_probs, rng);
98    std::cout << "  Top-K(K=5): selected token " << topk_id << std::endl;
99    assert(topk_id >= 5);
100
101    rng.seed(42);
102    TopPSampling topp;
103    Tensor topp_probs = topp.apply(logits.clone(), config, {});
104    size_t topp_id = topp.sample(topp_probs, rng);
105    std::cout << "  Top-P(P=0.9): selected token " << topp_id << std::endl;
106
107    config.temperature = 0.0f;
108    Tensor temp0_probs = topk.apply(logits.clone(), config, {});
109    size_t temp0_id = topk.sample(temp0_probs, rng);
110    std::cout << "  Temperature=0 (greedy fallback): selected token " << temp0_id << std::endl;
111    assert(temp0_id == 9);
112
113    std::cout << "  Sampling strategy test PASSED" << std::endl;
114}
115
116void test_generative_model() {
117    std::cout << "\n=== GenerativeModel Test ===" << std::endl;
118
119    CausalLMConfig lm_config;
120    lm_config.vocab_size = 200;
121    lm_config.d_model = 64;
122    lm_config.max_seq_len = 64;
123    lm_config.causal_window_size = 8;
124    lm_config.sae_k = 16;
125    lm_config.ntm_memory_slots = 4;
126    lm_config.use_mla = false;
127
128    auto tokenizer = std::make_unique<BPETokenizer>();
129    tokenizer->add_vocab("你", 4);
130    tokenizer->add_vocab("好", 5);
131    tokenizer->add_vocab("世", 6);
132    tokenizer->add_vocab("界", 7);
133    tokenizer->add_vocab("测", 8);
134    tokenizer->add_vocab("试", 9);
135    tokenizer->add_vocab("生", 10);
136    tokenizer->add_vocab("成", 11);
137    tokenizer->set_vocab_size(200);
138
139    GenerativeModel model(lm_config, std::move(tokenizer));
140    std::cout << "  GenerativeModel constructed" << std::endl;
141
142    GenerateConfig gen_config;
143    gen_config.max_new_tokens = 10;
144    gen_config.temperature = 0.8f;
145    gen_config.top_k = 20;
146    gen_config.random_seed = 12345;
147    gen_config.eos_id = 3;
148
149    GenerateOutput output = model.generate("你好", gen_config);
150    std::cout << "  Generated text: '" << output.text << "'" << std::endl;
151    std::cout << "  Generated " << output.token_ids.size() << " tokens" << std::endl;
152    std::cout << "  Finish reason: " << static_cast<int>(output.finish_reason) << std::endl;
153    std::cout << "  Cache stats: len=" << output.cache_stats.cache_len
154              << " mem=" << output.cache_stats.memory_bytes << " bytes" << std::endl;
155
156    assert(!output.token_ids.empty());
157
158    model.set_strategy(SamplingStrategyType::GREEDY);
159    gen_config.random_seed = 42;
160    GenerateOutput greedy_out = model.generate("测试", gen_config);
161    std::cout << "  Greedy output: '" << greedy_out.text << "'" << std::endl;
162
163    model.set_strategy(SamplingStrategyType::TOP_P);
164    gen_config.temperature = 1.0f;
165    gen_config.random_seed = 99;
166    GenerateOutput topp_out = model.generate("生成", gen_config);
167    std::cout << "  Top-P output: '" << topp_out.text << "'" << std::endl;
168
169    std::cout << "  GenerativeModel test PASSED" << std::endl;
170}
171
172void test_repetition_penalty() {
173    std::cout << "\n=== Repetition Penalty Test ===" << std::endl;
174
175    CausalLMConfig config;
176    config.vocab_size = 50;
177    config.d_model = 16;
178    config.max_seq_len = 32;
179    config.causal_window_size = 4;
180    config.sae_k = 8;
181    config.ntm_memory_slots = 2;
182    config.use_mla = false;
183
184    auto tokenizer = std::make_unique<BPETokenizer>();
185    tokenizer->set_vocab_size(50);
186
187    GenerativeModel model(config, std::move(tokenizer));
188
189    GenerateConfig gen_config;
190    gen_config.max_new_tokens = 15;
191    gen_config.temperature = 0.8f;
192    gen_config.top_k = 10;
193    gen_config.repetition_penalty = 1.5f;
194    gen_config.random_seed = 42;
195
196    GenerateOutput output = model.generate("测试", gen_config);
197
198    std::unordered_map<size_t, size_t> counts;
199    for (auto id : output.token_ids) counts[id]++;
200    size_t max_repeat = 0;
201    for (auto& [id, cnt] : counts) max_repeat = std::max(max_repeat, cnt);
202    std::cout << "  Max repetition count: " << max_repeat << std::endl;
203    std::cout << "  Repetition penalty test PASSED" << std::endl;
204}
205
206int main() {
207    std::cout << "========================================" << std::endl;
208    std::cout << "NeuroFlow Generative Model Test Suite" << std::endl;
209    std::cout << "========================================" << std::endl;
210
211    try {
212        test_tokenizer();
213        test_causal_lm_head();
214        test_sampling_strategies();
215        test_generative_model();
216        test_repetition_penalty();
217
218        std::cout << "\n========================================" << std::endl;
219        std::cout << "ALL TESTS PASSED!" << std::endl;
220        std::cout << "========================================" << std::endl;
221    } catch (const std::exception& e) {
222        std::cerr << "TEST FAILED: " << e.what() << std::endl;
223        return 1;
224    }
225
226    return 0;
227}