Team Ai
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
1likes
alignment_common.hpp320 linesDownload Raw Back to neuroflow
1#ifndef NEUROFLOW_ALIGNMENT_COMMON_HPP
2#define NEUROFLOW_ALIGNMENT_COMMON_HPP
3
4#include <algorithm>
5#include <cmath>
6#include <cstring>
7#include <fstream>
8#include <iostream>
9#include <string>
10#include <unordered_map>
11#include <vector>
12
13#include "causal_lm.hpp"
14#include "tensor.hpp"
15
16namespace neuroflow {
17
18struct SFTSample {
19    std::string instruction;
20    std::string response;
21};
22
23struct DPOSample {
24    std::string instruction;
25    std::string chosen;
26    std::string rejected;
27};
28
29struct SFTTrainingTensors {
30    std::vector<size_t> input_ids;
31    std::vector<size_t> target_ids;
32    std::vector<float> loss_mask;
33    size_t instruction_len;
34};
35
36struct DPOTrainingTensors {
37    std::vector<size_t> chosen_ids;
38    std::vector<size_t> rejected_ids;
39    size_t chosen_prompt_len;
40    size_t rejected_prompt_len;
41};
42
43inline std::string extract_json_string(const std::string& json, const std::string& key) {
44    std::string search = "\"" + key + "\"";
45    size_t pos = json.find(search);
46    if (pos == std::string::npos) return "";
47    pos = json.find(':', pos + search.size());
48    if (pos == std::string::npos) return "";
49    pos++;
50    while (pos < json.size() && (json[pos] == ' ' || json[pos] == '\t')) pos++;
51    if (pos >= json.size() || json[pos] != '"') return "";
52    size_t end = pos + 1;
53    while (end < json.size()) {
54        if (json[end] == '"' && json[end - 1] != '\\') break;
55        if (json[end - 1] == '\\') { end++; continue; }
56        end++;
57    }
58    if (end >= json.size()) return "";
59    return json.substr(pos + 1, end - pos - 1);
60}
61
62inline size_t extract_json_number(const std::string& json, const std::string& key, size_t default_val = 0) {
63    std::string search = "\"" + key + "\"";
64    size_t pos = json.find(search);
65    if (pos == std::string::npos) return default_val;
66    pos = json.find(':', pos + search.size());
67    if (pos == std::string::npos) return default_val;
68    pos++;
69    while (pos < json.size() && (json[pos] == ' ' || json[pos] == '\t')) pos++;
70    size_t end = pos;
71    while (end < json.size() && (json[end] >= '0' && json[end] <= '9')) end++;
72    if (end == pos) return default_val;
73    return std::stoul(json.substr(pos, end - pos));
74}
75
76inline bool extract_json_bool(const std::string& json, const std::string& key, bool default_val = false) {
77    std::string search = "\"" + key + "\"";
78    size_t pos = json.find(search);
79    if (pos == std::string::npos) return default_val;
80    pos = json.find(':', pos + search.size());
81    if (pos == std::string::npos) return default_val;
82    pos++;
83    while (pos < json.size() && (json[pos] == ' ' || json[pos] == '\t')) pos++;
84    if (pos + 3 < json.size() && json.substr(pos, 4) == "true") return true;
85    if (pos + 4 < json.size() && json.substr(pos, 5) == "false") return false;
86    return default_val;
87}
88
89inline void unescape_json(std::string& s) {
90    size_t pos = 0;
91    while ((pos = s.find('\\', pos)) != std::string::npos) {
92        if (pos + 1 < s.size()) {
93            char c = s[pos + 1];
94            if (c == 'n') { s.replace(pos, 2, "\n"); pos++; }
95            else if (c == 't') { s.replace(pos, 2, "\t"); pos++; }
96            else if (c == 'r') { s.replace(pos, 2, "\r"); pos++; }
97            else if (c == '"') { s.replace(pos, 2, "\""); pos++; }
98            else if (c == '\\') { s.replace(pos, 2, "\\"); pos++; }
99            else if (c == '/') { s.replace(pos, 2, "/"); pos++; }
100            else if (c == 'u' && pos + 5 < s.size()) { pos += 6; }
101            else { pos += 2; }
102        } else { pos++; }
103    }
104}
105
106inline bool validate_path(const std::string& path) {
107    if (path.find("..") != std::string::npos) return false;
108    return true;
109}
110
111inline void save_lm_checkpoint(const CausalLMHead& model, const std::string& path) {
112    CausalLMHead& lm_head = const_cast<CausalLMHead&>(model);
113#ifdef USE_CUDA
114    auto sync_to_cpu = [](const Tensor& t) {
115        if (t.is_on_gpu()) { const_cast<Tensor&>(t).to_cpu(); }
116    };
117    sync_to_cpu(lm_head.w_embed_);
118    sync_to_cpu(lm_head.w_pos_);
119    sync_to_cpu(lm_head.dw_kernel_);
120    sync_to_cpu(lm_head.pw_conv_->weight);
121    sync_to_cpu(lm_head.sae_w_encode_->weight);
122    sync_to_cpu(lm_head.sae_w_decode_->weight);
123    sync_to_cpu(lm_head.ntm_w_read_->weight);
124    sync_to_cpu(lm_head.ntm_w_write_->weight);
125    sync_to_cpu(lm_head.ntm_w_erase_->weight);
126    sync_to_cpu(lm_head.ntm_memory_);
127    sync_to_cpu(lm_head.w_proj_->weight);
128    sync_to_cpu(lm_head.w_proj_->bias);
129    if (lm_head.bridge_) {
130        sync_to_cpu(lm_head.bridge_->weight);
131        sync_to_cpu(lm_head.bridge_->bias);
132    }
133    sync_to_cpu(lm_head.w_out_->weight);
134    sync_to_cpu(lm_head.w_out_->bias);
135    sync_to_cpu(lm_head.ln_->weight);
136    sync_to_cpu(lm_head.ln_->bias);
137    for (auto& attn : lm_head.attn_layers_) {
138        sync_to_cpu(attn->w_q->weight);
139        sync_to_cpu(attn->w_q->bias);
140        sync_to_cpu(attn->w_k->weight);
141        sync_to_cpu(attn->w_k->bias);
142        sync_to_cpu(attn->w_v->weight);
143        sync_to_cpu(attn->w_v->bias);
144        sync_to_cpu(attn->w_out->weight);
145        sync_to_cpu(attn->w_out->bias);
146        sync_to_cpu(attn->norm->weight);
147        sync_to_cpu(attn->norm->bias);
148    }
149#endif
150    auto sl = [](std::ofstream& o, const std::string& n, const Tensor& t) {
151        uint32_t nl = n.size(); o.write((char*)&nl, 4); o.write(n.data(), nl);
152        uint32_t nd = t.shape_.size(); o.write((char*)&nd, 4);
153        for (auto d : t.shape_) { uint32_t dd = d; o.write((char*)&dd, 4); }
154        uint32_t ds = t.data_size_; o.write((char*)&ds, 4);
155        o.write((char*)t.data_.get(), ds);
156    };
157    std::ofstream o(path, std::ios::binary);
158    o.write("LMH2", 4);
159    sl(o, "w_embed", lm_head.w_embed_);
160    sl(o, "w_pos", lm_head.w_pos_);
161    sl(o, "dw_kernel", lm_head.dw_kernel_);
162    sl(o, "pw_conv.weight", lm_head.pw_conv_->weight);
163    sl(o, "sae_encode.weight", lm_head.sae_w_encode_->weight);
164    sl(o, "sae_decode.weight", lm_head.sae_w_decode_->weight);
165    sl(o, "ntm_read.weight", lm_head.ntm_w_read_->weight);
166    sl(o, "ntm_write.weight", lm_head.ntm_w_write_->weight);
167    sl(o, "ntm_erase.weight", lm_head.ntm_w_erase_->weight);
168    sl(o, "ntm_memory", lm_head.ntm_memory_);
169    sl(o, "w_proj.weight", lm_head.w_proj_->weight);
170    sl(o, "w_proj.bias", lm_head.w_proj_->bias);
171    if (lm_head.bridge_) {
172        sl(o, "bridge.weight", lm_head.bridge_->weight);
173        sl(o, "bridge.bias", lm_head.bridge_->bias);
174    }
175    sl(o, "w_out.weight", lm_head.w_out_->weight);
176    if (lm_head.w_out_->bias.data_) sl(o, "w_out.bias", lm_head.w_out_->bias);
177    sl(o, "ln.weight", lm_head.ln_->weight);
178    sl(o, "ln.bias", lm_head.ln_->bias);
179    for (size_t i = 0; i < lm_head.attn_layers_.size(); ++i) {
180        std::string p = "attn" + std::to_string(i) + ".";
181        sl(o, p + "w_q.weight", lm_head.attn_layers_[i]->w_q->weight);
182        sl(o, p + "w_q.bias", lm_head.attn_layers_[i]->w_q->bias);
183        sl(o, p + "w_k.weight", lm_head.attn_layers_[i]->w_k->weight);
184        sl(o, p + "w_k.bias", lm_head.attn_layers_[i]->w_k->bias);
185        sl(o, p + "w_v.weight", lm_head.attn_layers_[i]->w_v->weight);
186        sl(o, p + "w_v.bias", lm_head.attn_layers_[i]->w_v->bias);
187        sl(o, p + "w_out.weight", lm_head.attn_layers_[i]->w_out->weight);
188        sl(o, p + "w_out.bias", lm_head.attn_layers_[i]->w_out->bias);
189        sl(o, p + "norm.weight", lm_head.attn_layers_[i]->norm->weight);
190        sl(o, p + "norm.bias", lm_head.attn_layers_[i]->norm->bias);
191    }
192    uint32_t z = 0; o.write((char*)&z, 4); o.close();
193}
194
195inline bool load_lm_checkpoint(CausalLMHead& lm_head, const std::string& lm_path) {
196    std::ifstream ifs(lm_path, std::ios::binary);
197    if (!ifs) {
198        std::cerr << "LM Head checkpoint未找到: " << lm_path << std::endl;
199        return false;
200    }
201
202    char magic[5] = {0};
203    ifs.read(magic, 4);
204    if (std::string(magic) != "LMH2" && std::string(magic) != "LMH1") {
205        std::cerr << "LM Head格式不匹配: " << std::string(magic) << " (期望LMH2/LMH1)" << std::endl;
206        ifs.close();
207        return false;
208    }
209
210    auto read_named_tensor = [&]() -> std::pair<std::string, Tensor> {
211        uint32_t nl = 0; ifs.read((char*)&nl, 4);
212        if (nl == 0 || ifs.eof()) return {"", Tensor()};
213        std::string name(nl, '\0'); ifs.read(&name[0], nl);
214        uint32_t nd = 0; ifs.read((char*)&nd, 4);
215        std::vector<size_t> shape(nd);
216        for (size_t i = 0; i < nd; ++i) { uint32_t d = 0; ifs.read((char*)&d, 4); shape[i] = d; }
217        uint32_t ds = 0; ifs.read((char*)&ds, 4);
218        Tensor t(shape, QuantType::FP32);
219        if (t.data_size_ == ds) {
220            ifs.read((char*)t.data_.get(), ds);
221        } else {
222            ifs.seekg(ds, std::ios::cur);
223            t = Tensor();
224        }
225        return {name, std::move(t)};
226    };
227
228    std::unordered_map<std::string, Tensor*> tensor_map;
229    tensor_map["w_embed"] = &lm_head.w_embed_;
230    tensor_map["w_pos"] = &lm_head.w_pos_;
231    tensor_map["dw_kernel"] = &lm_head.dw_kernel_;
232    tensor_map["pw_conv.weight"] = &lm_head.pw_conv_->weight;
233    tensor_map["sae_encode.weight"] = &lm_head.sae_w_encode_->weight;
234    tensor_map["sae_decode.weight"] = &lm_head.sae_w_decode_->weight;
235    tensor_map["ntm_read.weight"] = &lm_head.ntm_w_read_->weight;
236    tensor_map["ntm_write.weight"] = &lm_head.ntm_w_write_->weight;
237    tensor_map["ntm_erase.weight"] = &lm_head.ntm_w_erase_->weight;
238    tensor_map["ntm_memory"] = &lm_head.ntm_memory_;
239    tensor_map["w_proj.weight"] = &lm_head.w_proj_->weight;
240    tensor_map["w_proj.bias"] = &lm_head.w_proj_->bias;
241    if (lm_head.bridge_) {
242        tensor_map["bridge.weight"] = &lm_head.bridge_->weight;
243        tensor_map["bridge.bias"] = &lm_head.bridge_->bias;
244    }
245    tensor_map["w_out.weight"] = &lm_head.w_out_->weight;
246    tensor_map["w_out.bias"] = &lm_head.w_out_->bias;
247    tensor_map["ln.weight"] = &lm_head.ln_->weight;
248    tensor_map["ln.bias"] = &lm_head.ln_->bias;
249    for (size_t i = 0; i < lm_head.attn_layers_.size(); ++i) {
250        std::string p = "attn" + std::to_string(i) + ".";
251        tensor_map[p + "w_q.weight"] = &lm_head.attn_layers_[i]->w_q->weight;
252        tensor_map[p + "w_q.bias"] = &lm_head.attn_layers_[i]->w_q->bias;
253        tensor_map[p + "w_k.weight"] = &lm_head.attn_layers_[i]->w_k->weight;
254        tensor_map[p + "w_k.bias"] = &lm_head.attn_layers_[i]->w_k->bias;
255        tensor_map[p + "w_v.weight"] = &lm_head.attn_layers_[i]->w_v->weight;
256        tensor_map[p + "w_v.bias"] = &lm_head.attn_layers_[i]->w_v->bias;
257        tensor_map[p + "w_out.weight"] = &lm_head.attn_layers_[i]->w_out->weight;
258        tensor_map[p + "w_out.bias"] = &lm_head.attn_layers_[i]->w_out->bias;
259        tensor_map[p + "norm.weight"] = &lm_head.attn_layers_[i]->norm->weight;
260        tensor_map[p + "norm.bias"] = &lm_head.attn_layers_[i]->norm->bias;
261    }
262
263    size_t loaded = 0;
264    while (ifs) {
265        auto [name, tensor] = read_named_tensor();
266        if (name.empty()) break;
267        auto it = tensor_map.find(name);
268        if (it != tensor_map.end() && tensor.numel() > 0) {
269            if (it->second->shape_ == tensor.shape_) {
270                memcpy(it->second->data_.get(), tensor.data_.get(), tensor.data_size_);
271                loaded++;
272            } else {
273                std::cerr << "  跳过 '" << name << "': 形状不匹配" << std::endl;
274            }
275        } else if (name.find("w_qkv.") != std::string::npos) {
276            std::string base = name.substr(0, name.find("w_qkv."));
277            std::string suffix = name.substr(name.find("w_qkv.") + 6);
278            for (size_t i = 0; i < lm_head.attn_layers_.size(); ++i) {
279                std::string p = "attn" + std::to_string(i) + ".";
280                if (base != p) continue;
281                size_t d_model = lm_head.config_.d_model;
282                size_t head_dim = d_model / lm_head.config_.num_attn_heads;
283                size_t n_q = lm_head.attn_layers_[i]->n_q_heads_;
284                size_t n_kv = lm_head.attn_layers_[i]->n_kv_heads_;
285                if (suffix == "weight" && tensor.shape_.size() == 2 && tensor.shape_[0] == 3 * d_model) {
286                    const float* src = tensor.as_fp32();
287                    float* dq = lm_head.attn_layers_[i]->w_q->weight.as_fp32();
288                    float* dk = lm_head.attn_layers_[i]->w_k->weight.as_fp32();
289                    float* dv = lm_head.attn_layers_[i]->w_v->weight.as_fp32();
290                    for (size_t r = 0; r < d_model; ++r) {
291                        memcpy(dq + r * n_q * head_dim, src + r * 3 * d_model, n_q * head_dim * sizeof(float));
292                        memcpy(dk + r * n_kv * head_dim, src + r * 3 * d_model + d_model, n_kv * head_dim * sizeof(float));
293                        memcpy(dv + r * n_kv * head_dim, src + r * 3 * d_model + 2 * d_model, n_kv * head_dim * sizeof(float));
294                    }
295                    loaded++;
296                    std::cerr << "  拆分旧格式 '" << name << "' -> w_q/w_k/w_v" << std::endl;
297                } else if (suffix == "bias" && tensor.shape_[0] == 3 * d_model) {
298                    const float* src = tensor.as_fp32();
299                    float* bq = lm_head.attn_layers_[i]->w_q->bias.as_fp32();
300                    float* bk = lm_head.attn_layers_[i]->w_k->bias.as_fp32();
301                    float* bv = lm_head.attn_layers_[i]->w_v->bias.as_fp32();
302                    memcpy(bq, src, n_q * head_dim * sizeof(float));
303                    memcpy(bk, src + d_model, n_kv * head_dim * sizeof(float));
304                    memcpy(bv, src + 2 * d_model, n_kv * head_dim * sizeof(float));
305                    loaded++;
306                    std::cerr << "  拆分旧格式 '" << name << "' -> w_q/w_k/w_v bias" << std::endl;
307                }
308            }
309        }
310    }
311
312    ifs.close();
313    if (lm_head.config_.weight_tying) lm_head.tie_weights();
314    std::cerr << "LM Head已恢复: " << loaded << " 个张量从 " << lm_path << std::endl;
315    return loaded > 0;
316}
317
318} // namespace neuroflow
319
320#endif // NEUROFLOW_ALIGNMENT_COMMON_HPP