cwenzi/neuroflow-cpp
1
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