Team Ai
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
1likes
infer_v2.cpp384 linesDownload Raw Back to src
1#ifndef NOMINMAX
2#define NOMINMAX
3#endif
4#include <iostream>
5#include <string>
6#include <vector>
7#include <fstream>
8#include "neuroflow/generative.hpp"
9#include "neuroflow/model.hpp"
10
11using namespace neuroflow;
12
13int main(int argc, char* argv[]) {
14    std::string config_path = "configs/config_distill_small.json";
15    std::string tokenizer_path = "configs/tokenizer_cn_013.json";
16    std::string lm_head_path = "";
17    int max_new_tokens = 64;
18    float temperature = 0.8f;
19    int top_k = 40;
20    float top_p = 0.9f;
21    float repetition_penalty = 1.0f;
22    bool use_cuda = false;
23    std::string strategy_name = "top_k";
24    int yarn_max_seq_len = -1;
25
26    for (int i = 1; i < argc; ++i) {
27        std::string arg = argv[i];
28        if (arg == "--config" && i + 1 < argc) config_path = argv[++i];
29        else if (arg == "--tokenizer" && i + 1 < argc) tokenizer_path = argv[++i];
30        else if (arg == "--lm-head" && i + 1 < argc) lm_head_path = argv[++i];
31        else if (arg == "--max-tokens" && i + 1 < argc) max_new_tokens = std::atoi(argv[++i]);
32        else if (arg == "--temperature" && i + 1 < argc) temperature = std::atof(argv[++i]);
33        else if (arg == "--top-k" && i + 1 < argc) top_k = std::atoi(argv[++i]);
34        else if (arg == "--top-p" && i + 1 < argc) top_p = std::atof(argv[++i]);
35        else if (arg == "--repetition-penalty" && i + 1 < argc) repetition_penalty = std::atof(argv[++i]);
36        else if (arg == "--strategy" && i + 1 < argc) strategy_name = argv[++i];
37        else if (arg == "--max-seq-len" && i + 1 < argc) yarn_max_seq_len = std::atoi(argv[++i]);
38        else if (arg == "--use-cuda") use_cuda = true;
39    }
40
41    if (lm_head_path.empty()) {
42        std::cerr << "用法: neuroflow_infer --lm-head <path> [--config <path>] [--tokenizer <path>]"
43                  << " [--max-tokens <N>] [--temperature <f>] [--top-k <N>] [--use-cuda]" << std::endl;
44        return 1;
45    }
46
47    NeuroFlowModel::Config model_cfg;
48    {
49        std::ifstream jf(config_path);
50        if (jf) {
51            std::string json_str((std::istreambuf_iterator<char>(jf)), std::istreambuf_iterator<char>());
52            auto ex_num = [&](const std::string& key, size_t def) -> size_t {
53                std::string pat = "\"" + key + "\"";
54                size_t pos = json_str.find(pat);
55                if (pos == std::string::npos) return def;
56                pos = json_str.find(':', pos + pat.size());
57                if (pos == std::string::npos) return def;
58                pos++;
59                while (pos < json_str.size() && (json_str[pos] == ' ' || json_str[pos] == '\t')) pos++;
60                return std::stoull(json_str.substr(pos));
61            };
62            auto ex_str = [&](const std::string& key, const std::string& def) -> std::string {
63                std::string pat = "\"" + key + "\"";
64                size_t pos = json_str.find(pat);
65                if (pos == std::string::npos) return def;
66                pos = json_str.find(':', pos + pat.size());
67                if (pos == std::string::npos) return def;
68                size_t q1 = json_str.find('"', pos);
69                if (q1 == std::string::npos) return def;
70                size_t q2 = json_str.find('"', q1 + 1);
71                if (q2 == std::string::npos) return def;
72                return json_str.substr(q1 + 1, q2 - q1 - 1);
73            };
74            model_cfg.input_dim = ex_num("d_model", model_cfg.input_dim);
75            model_cfg.hidden_dim = ex_num("hidden_dim", model_cfg.hidden_dim);
76            model_cfg.output_dim = ex_num("output_dim", model_cfg.output_dim);
77            model_cfg.vocab_size = ex_num("vocab_size", model_cfg.vocab_size);
78            model_cfg.max_seq_len = ex_num("max_seq_len", model_cfg.max_seq_len);
79            model_cfg.causal_window_size = ex_num("causal_window_size", model_cfg.causal_window_size);
80            model_cfg.sae_k = ex_num("sae_k", model_cfg.sae_k);
81            model_cfg.ntm_memory_slots = ex_num("ntm_memory_slots", model_cfg.ntm_memory_slots);
82            model_cfg.lm_num_attn_layers = ex_num("lm_num_attn_layers", model_cfg.lm_num_attn_layers);
83            model_cfg.lm_pooling = ex_str("lm_pooling", "mean");
84        }
85    }
86
87    CausalLMConfig lm_cfg;
88    lm_cfg.vocab_size = model_cfg.vocab_size;
89    lm_cfg.d_model = model_cfg.hidden_dim;
90    lm_cfg.max_seq_len = model_cfg.max_seq_len;
91    lm_cfg.causal_window_size = model_cfg.causal_window_size;
92    lm_cfg.sae_k = model_cfg.sae_k;
93    lm_cfg.ntm_memory_slots = model_cfg.ntm_memory_slots;
94    lm_cfg.use_mla = model_cfg.use_mla;
95    lm_cfg.mla_latent_dim = model_cfg.mla_latent_dim;
96    lm_cfg.mla_n_heads = model_cfg.mla_n_heads;
97    lm_cfg.mla_max_cache_len = 4096;
98    lm_cfg.weight_tying = true;
99    lm_cfg.num_attn_layers = model_cfg.lm_num_attn_layers;
100    lm_cfg.num_attn_heads = 4;
101    lm_cfg.pooling = model_cfg.lm_pooling;
102
103    CausalLMHead lm_head(lm_cfg);
104    if (lm_cfg.weight_tying) lm_head.tie_weights();
105
106    std::cerr << "加载 LM Head: " << lm_head_path << std::endl;
107    {
108        std::ifstream ifs(lm_head_path, std::ios::binary);
109        if (!ifs) {
110            std::cerr << "错误: 无法打开 " << lm_head_path << std::endl;
111            return 1;
112        }
113        char magic[5] = {0};
114        ifs.read(magic, 4);
115        if (std::string(magic) != "LMH2" && std::string(magic) != "LMH1") {
116            std::cerr << "错误: 格式不匹配 " << std::string(magic) << std::endl;
117            return 1;
118        }
119
120        std::unordered_map<std::string, Tensor*> tensor_map;
121        tensor_map["w_embed"] = &lm_head.w_embed_;
122        tensor_map["w_pos"] = &lm_head.w_pos_;
123        tensor_map["dw_kernel"] = &lm_head.dw_kernel_;
124        tensor_map["pw_conv.weight"] = &lm_head.pw_conv_->weight;
125        tensor_map["sae_encode.weight"] = &lm_head.sae_w_encode_->weight;
126        tensor_map["sae_decode.weight"] = &lm_head.sae_w_decode_->weight;
127        tensor_map["ntm_read.weight"] = &lm_head.ntm_w_read_->weight;
128        tensor_map["ntm_write.weight"] = &lm_head.ntm_w_write_->weight;
129        tensor_map["ntm_erase.weight"] = &lm_head.ntm_w_erase_->weight;
130        tensor_map["ntm_memory"] = &lm_head.ntm_memory_;
131        tensor_map["w_proj.weight"] = &lm_head.w_proj_->weight;
132        tensor_map["w_proj.bias"] = &lm_head.w_proj_->bias;
133        tensor_map["w_out.weight"] = &lm_head.w_out_->weight;
134        tensor_map["w_out.bias"] = &lm_head.w_out_->bias;
135        tensor_map["ln.weight"] = &lm_head.ln_->weight;
136        tensor_map["ln.bias"] = &lm_head.ln_->bias;
137        for (size_t i = 0; i < lm_head.attn_layers_.size(); ++i) {
138            std::string p = "attn" + std::to_string(i) + ".";
139            tensor_map[p + "w_q.weight"] = &lm_head.attn_layers_[i]->w_q->weight;
140            tensor_map[p + "w_q.bias"] = &lm_head.attn_layers_[i]->w_q->bias;
141            tensor_map[p + "w_k.weight"] = &lm_head.attn_layers_[i]->w_k->weight;
142            tensor_map[p + "w_k.bias"] = &lm_head.attn_layers_[i]->w_k->bias;
143            tensor_map[p + "w_v.weight"] = &lm_head.attn_layers_[i]->w_v->weight;
144            tensor_map[p + "w_v.bias"] = &lm_head.attn_layers_[i]->w_v->bias;
145            tensor_map[p + "w_out.weight"] = &lm_head.attn_layers_[i]->w_out->weight;
146            tensor_map[p + "w_out.bias"] = &lm_head.attn_layers_[i]->w_out->bias;
147            tensor_map[p + "norm.weight"] = &lm_head.attn_layers_[i]->norm->weight;
148            tensor_map[p + "norm.bias"] = &lm_head.attn_layers_[i]->norm->bias;
149        }
150
151        size_t loaded = 0;
152        while (ifs) {
153            uint32_t nl = 0;
154            ifs.read((char*)&nl, 4);
155            if (nl == 0 || ifs.eof()) break;
156            if (nl > 256) { std::cerr << "异常名称长度: " << nl << ", 可能文件损坏" << std::endl; break; }
157            std::string name(nl, '\0'); ifs.read(&name[0], nl);
158            uint32_t nd = 0; ifs.read((char*)&nd, 4);
159            if (nd > 10) { std::cerr << "异常维度: " << nd << ", name=" << name << std::endl; break; }
160            std::vector<size_t> shape(nd);
161            for (size_t i = 0; i < nd; ++i) { uint32_t d = 0; ifs.read((char*)&d, 4); shape[i] = d; }
162            uint32_t ds = 0; ifs.read((char*)&ds, 4);
163            Tensor t(shape, QuantType::FP32);
164            if (t.data_size_ == ds) {
165                ifs.read((char*)t.data_.get(), ds);
166            } else {
167                ifs.seekg(ds, std::ios::cur);
168                continue;
169            }
170            auto it = tensor_map.find(name);
171            if (it != tensor_map.end() && it->second->shape_ == t.shape_) {
172                memcpy(it->second->data_.get(), t.data_.get(), t.data_size_);
173                loaded++;
174            } else if (name.find("w_qkv.") != std::string::npos) {
175                std::string base = name.substr(0, name.find("w_qkv."));
176                std::string suffix = name.substr(name.find("w_qkv.") + 6);
177                for (size_t i = 0; i < lm_head.attn_layers_.size(); ++i) {
178                    std::string p = "attn" + std::to_string(i) + ".";
179                    if (base != p) continue;
180                    size_t d_model = lm_head.config_.d_model;
181                    size_t head_dim = d_model / lm_head.config_.num_attn_heads;
182                    size_t n_q = lm_head.attn_layers_[i]->n_q_heads_;
183                    size_t n_kv = lm_head.attn_layers_[i]->n_kv_heads_;
184                    if (suffix == "weight" && t.shape_.size() == 2 && t.shape_[0] == 3 * d_model) {
185                        const float* src = t.as_fp32();
186                        float* dq = lm_head.attn_layers_[i]->w_q->weight.as_fp32();
187                        float* dk = lm_head.attn_layers_[i]->w_k->weight.as_fp32();
188                        float* dv = lm_head.attn_layers_[i]->w_v->weight.as_fp32();
189                        for (size_t r = 0; r < d_model; ++r) {
190                            memcpy(dq + r * n_q * head_dim, src + r * 3 * d_model, n_q * head_dim * sizeof(float));
191                            memcpy(dk + r * n_kv * head_dim, src + r * 3 * d_model + d_model, n_kv * head_dim * sizeof(float));
192                            memcpy(dv + r * n_kv * head_dim, src + r * 3 * d_model + 2 * d_model, n_kv * head_dim * sizeof(float));
193                        }
194                        loaded++;
195                    } else if (suffix == "bias" && t.shape_[0] == 3 * d_model) {
196                        const float* src = t.as_fp32();
197                        float* bq = lm_head.attn_layers_[i]->w_q->bias.as_fp32();
198                        float* bk = lm_head.attn_layers_[i]->w_k->bias.as_fp32();
199                        float* bv = lm_head.attn_layers_[i]->w_v->bias.as_fp32();
200                        memcpy(bq, src, n_q * head_dim * sizeof(float));
201                        memcpy(bk, src + d_model, n_kv * head_dim * sizeof(float));
202                        memcpy(bv, src + 2 * d_model, n_kv * head_dim * sizeof(float));
203                        loaded++;
204                    }
205                }
206            }
207        }
208        ifs.close();
209        if (lm_cfg.weight_tying) lm_head.tie_weights();
210        std::cerr << "已加载 " << loaded << " 个张量" << std::endl;
211    }
212
213    BPETokenizer tokenizer(tokenizer_path);
214    std::cerr << "词表: " << tokenizer.vocab_size() << " tokens" << std::endl;
215
216#ifdef USE_CUDA
217    bool cuda_active = false;
218    if (use_cuda) {
219        if (CudaContext::instance().initialize(0)) {
220            cuda_active = true;
221            std::cerr << "GPU后端: 已启用" << std::endl;
222
223            lm_head.w_embed_.to_gpu();
224            lm_head.w_pos_.to_gpu();
225            lm_head.dw_kernel_.to_gpu();
226            lm_head.pw_conv_->weight.to_gpu();
227            lm_head.sae_w_encode_->weight.to_gpu();
228            lm_head.sae_w_decode_->weight.to_gpu();
229            lm_head.ntm_w_read_->weight.to_gpu();
230            lm_head.ntm_w_write_->weight.to_gpu();
231            lm_head.ntm_w_erase_->weight.to_gpu();
232            lm_head.ntm_memory_.to_gpu();
233            lm_head.w_proj_->weight.to_gpu();
234            lm_head.w_proj_->bias.to_gpu();
235            lm_head.w_out_->weight.to_gpu();
236            if (lm_head.w_out_->bias.data_) lm_head.w_out_->bias.to_gpu();
237            lm_head.ln_->weight.to_gpu();
238            lm_head.ln_->bias.to_gpu();
239            for (auto& attn : lm_head.attn_layers_) {
240                attn->w_q->weight.to_gpu();
241                attn->w_q->bias.to_gpu();
242                attn->w_k->weight.to_gpu();
243                attn->w_k->bias.to_gpu();
244                attn->w_v->weight.to_gpu();
245                attn->w_v->bias.to_gpu();
246                attn->w_out->weight.to_gpu();
247                attn->w_out->bias.to_gpu();
248                attn->norm->weight.to_gpu();
249                attn->norm->bias.to_gpu();
250            }
251            CudaContext::instance().synchronize();
252            size_t free_mem = CudaContext::instance().free_memory();
253            size_t total_mem = CudaContext::instance().total_memory();
254            std::cerr << "GPU显存: " << (total_mem - free_mem) / (1024*1024)
255                      << " MB 已用 / " << total_mem / (1024*1024) << " MB 总计" << std::endl;
256        } else {
257            std::cerr << "[CUDA WARNING] GPU初始化失败,回退CPU后端" << std::endl;
258            use_cuda = false;
259        }
260    } else {
261        std::cerr << "GPU后端: 未启用 (使用--use-cuda启用)" << std::endl;
262    }
263#else
264    if (use_cuda) {
265        std::cerr << "[CUDA WARNING] 编译时未启用CUDA支持 (NEUROFLOW_USE_CUDA=OFF),回退CPU后端" << std::endl;
266        use_cuda = false;
267    }
268#endif
269
270    lm_head.eval();
271
272    if (yarn_max_seq_len > 0 && static_cast<size_t>(yarn_max_seq_len) > lm_cfg.max_seq_len) {
273        float scale_factor = static_cast<float>(yarn_max_seq_len) / static_cast<float>(lm_cfg.max_seq_len);
274        lm_head.set_yarn_scale(scale_factor);
275    }
276
277    std::cerr << "\n=== NeuroFlow 推理模式 ===" << std::endl;
278    std::cerr << "输入提示(Ctrl+C退出):" << std::endl;
279
280    std::string line;
281    while (std::getline(std::cin, line)) {
282        if (line.empty()) continue;
283        if (line == "quit" || line == "exit") break;
284
285        auto token_ids = tokenizer.encode(line, lm_cfg.max_seq_len);
286        if (token_ids.empty()) {
287            std::cerr << "(空token序列)" << std::endl;
288            continue;
289        }
290
291        lm_head.clear_cache();
292
293        std::vector<size_t> prefix(token_ids.begin(), token_ids.end() - 1);
294        Tensor logits;
295        if (prefix.size() > 0) {
296            logits = lm_head.forward(prefix);
297#ifdef USE_CUDA
298            if (cuda_active && logits.is_on_gpu()) {
299                logits.to_cpu();
300            }
301#endif
302        }
303
304        std::mt19937 rng(42);
305        std::vector<size_t> generated;
306        size_t last_id = token_ids.back();
307
308        GenerateConfig gen_cfg;
309        gen_cfg.max_new_tokens = static_cast<size_t>(max_new_tokens);
310        gen_cfg.temperature = temperature;
311        gen_cfg.top_k = static_cast<size_t>(top_k);
312        gen_cfg.top_p = top_p;
313        gen_cfg.repetition_penalty = repetition_penalty;
314        gen_cfg.eos_id = 0;
315
316        std::unique_ptr<SamplingStrategy> sampler;
317        if (strategy_name == "greedy") {
318            sampler = std::make_unique<GreedyDecoding>();
319        } else if (strategy_name == "top_p") {
320            sampler = std::make_unique<TopPSampling>();
321        } else if (strategy_name == "top_k_top_p") {
322            sampler = std::make_unique<TopKTopPSampling>();
323        } else {
324            sampler = std::make_unique<TopKSampling>();
325        }
326
327        for (int step = 0; step < max_new_tokens; ++step) {
328            size_t pos = token_ids.size() - 1 + step;
329            logits = lm_head.forward_step(last_id, pos);
330
331#ifdef USE_CUDA
332            if (cuda_active && logits.is_on_gpu()) {
333                int* d_sampled_token = nullptr;
334                int h_sampled_token = 0;
335                cudaError_t alloc_err = cudaMalloc(reinterpret_cast<void**>(&d_sampled_token), sizeof(int));
336
337                if (alloc_err == cudaSuccess) {
338                    unsigned int seed = static_cast<unsigned int>(step * 7919 + 42);
339                    bool ok = launch_topk_topp_sampling(
340                        logits.as_gpu_fp32(), d_sampled_token,
341                        static_cast<int>(lm_cfg.vocab_size),
342                        top_k, top_p, temperature, seed,
343                        CudaContext::instance().stream());
344
345                    if (ok) {
346                        CudaContext::instance().synchronize();
347                        cudaMemcpy(&h_sampled_token, d_sampled_token, sizeof(int), cudaMemcpyDeviceToHost);
348                        cudaFree(d_sampled_token);
349
350                        size_t chosen = static_cast<size_t>(h_sampled_token);
351                        generated.push_back(chosen);
352                        last_id = chosen;
353
354                        if (chosen == 0 || chosen == 1) break;
355                        continue;
356                    } else {
357                        cudaFree(d_sampled_token);
358                        std::cerr << "[INFER WARNING] GPU sampling failed, falling back to CPU sampling" << std::endl;
359                    }
360                } else {
361                    std::cerr << "[INFER WARNING] cudaMalloc failed: "
362                              << cudaGetErrorString(alloc_err)
363                              << ", falling back to CPU sampling" << std::endl;
364                }
365
366                logits.to_cpu();
367            }
368#endif
369
370            Tensor probs = sampler->apply(std::move(logits), gen_cfg, generated);
371            size_t chosen = sampler->sample(probs, rng);
372
373            generated.push_back(chosen);
374            last_id = chosen;
375
376            if (chosen == 0 || chosen == 1) break;
377        }
378
379        std::string output = tokenizer.decode(generated);
380        std::cout << output << std::endl;
381    }
382
383    return 0;
384}