cwenzi/neuroflow-cpp
1
1#ifndef NOMINMAX
2#define NOMINMAX
3#endif
4
5// C 系统头文件
6#ifndef _WIN32
7#include <fcntl.h>
8#include <sys/mman.h>
9#include <sys/types.h>
10#include <unistd.h>
11#else
12#include <windows.h>
13#endif
14
15// C++ 标准库
16#include <algorithm>
17#include <chrono>
18#include <cmath>
19#include <cstdio>
20#include <cstdlib>
21#include <cstring>
22#include <filesystem>
23#include <fstream>
24#include <iostream>
25#include <numeric>
26#include <sstream>
27#include <string>
28#include <vector>
29
30// 第三方库
31#ifdef _OPENMP
32#include <omp.h>
33#endif
34
35// 项目头文件
36#include "neuroflow/backprop.hpp"
37#include "neuroflow/generative.hpp"
38#include "neuroflow/model.hpp"
39#include "neuroflow/online_learning.hpp"
40#include "weight_io.hpp"
41
42using neuroflow::BPETokenizer;
43using neuroflow::CausalLMConfig;
44using neuroflow::CausalLMHead;
45using neuroflow::FullBackpropEngine;
46using neuroflow::InitStrategy;
47using neuroflow::NeuroFlowModel;
48using neuroflow::QuantType;
49using neuroflow::Tensor;
50using neuroflow::WeightInitializer;
51
52namespace {
53
54Tensor lm_linear_backward_input(const Tensor& output_grad, const Tensor& weight) {
55 size_t batch = output_grad.shape_[0];
56 size_t out_f = output_grad.shape_[1];
57 size_t in_f = weight.shape_[1];
58 Tensor input_grad({batch, in_f}, QuantType::FP32);
59
60#ifdef USE_CUDA
61 if (CudaContext::instance().is_available() && output_grad.is_on_gpu()) {
62 input_grad.to_gpu();
63 CudaContext::instance().sgemm_rowmajor(false, false,
64 static_cast<int>(batch), static_cast<int>(in_f), static_cast<int>(out_f),
65 1.0f, output_grad.as_gpu_fp32(), static_cast<int>(out_f),
66 weight.as_gpu_fp32(), static_cast<int>(in_f),
67 0.0f, input_grad.as_gpu_fp32(), static_cast<int>(in_f));
68 input_grad.gpu_dirty_ = true;
69 return input_grad;
70 }
71#endif
72
73 memset(input_grad.as_fp32(), 0, input_grad.data_size_);
74
75 if (out_f > 8192 && batch == 1) {
76 const float* og = output_grad.as_fp32();
77 const float* w = weight.as_fp32();
78 float* ig = input_grad.as_fp32();
79 float threshold = 1e-4f;
80 size_t n_active = 0;
81 for (size_t i = 0; i < out_f; ++i) {
82 if (std::abs(og[i]) >= threshold) n_active++;
83 }
84 if (n_active < out_f / 4) {
85 for (size_t i = 0; i < out_f; ++i) {
86 float g = og[i];
87 if (std::abs(g) < threshold) continue;
88 const float* w_row = w + i * in_f;
89 for (size_t j = 0; j < in_f; ++j) {
90 ig[j] += g * w_row[j];
91 }
92 }
93 return input_grad;
94 }
95 }
96
97#ifdef USE_CBLAS
98 cblas_sgemm(CblasRowMajor, CblasNoTrans, CblasNoTrans,
99 batch, in_f, out_f,
100 1.0f, output_grad.as_fp32(), out_f, weight.as_fp32(), in_f,
101 0.0f, input_grad.as_fp32(), in_f);
102#else
103 const float* og2 = output_grad.as_fp32();
104 const float* w2 = weight.as_fp32();
105 float* ig2 = input_grad.as_fp32();
106 for (size_t b = 0; b < batch; ++b) {
107 for (size_t j = 0; j < in_f; ++j) {
108 float sum = 0.0f;
109 for (size_t i = 0; i < out_f; ++i) {
110 sum += og2[b * out_f + i] * w2[i * in_f + j];
111 }
112 ig2[b * in_f + j] = sum;
113 }
114 }
115#endif
116 return input_grad;
117}
118
119Tensor lm_linear_backward_weight(const Tensor& input, const Tensor& output_grad) {
120 size_t batch = input.shape_[0];
121 size_t in_f = input.shape_[1];
122 size_t out_f = output_grad.shape_[1];
123 Tensor weight_grad({out_f, in_f}, QuantType::FP32);
124#ifdef USE_CUDA
125 if (CudaContext::instance().is_available() && output_grad.is_on_gpu()) {
126 weight_grad.to_gpu();
127 CudaContext::instance().sgemm_rowmajor(true, false,
128 static_cast<int>(out_f), static_cast<int>(in_f), static_cast<int>(batch),
129 1.0f / batch, output_grad.as_gpu_fp32(), static_cast<int>(out_f),
130 input.as_gpu_fp32(), static_cast<int>(in_f),
131 0.0f, weight_grad.as_gpu_fp32(), static_cast<int>(in_f));
132 weight_grad.gpu_dirty_ = true;
133 return weight_grad;
134 }
135#endif
136#ifdef USE_CBLAS
137 cblas_sgemm(CblasRowMajor, CblasTrans, CblasNoTrans,
138 out_f, in_f, batch,
139 1.0f / batch, output_grad.as_fp32(), out_f, input.as_fp32(), in_f,
140 0.0f, weight_grad.as_fp32(), in_f);
141#else
142 const float* inp = input.as_fp32();
143 const float* og = output_grad.as_fp32();
144 float* wg = weight_grad.as_fp32();
145 for (size_t i = 0; i < out_f; ++i) {
146 for (size_t j = 0; j < in_f; ++j) {
147 float sum = 0.0f;
148 for (size_t b = 0; b < batch; ++b) {
149 sum += og[b * out_f + i] * inp[b * in_f + j];
150 }
151 wg[i * in_f + j] = sum / batch;
152 }
153 }
154#endif
155 return weight_grad;
156}
157
158}
159
160struct TrainConfig {
161 std::string config_path;
162 std::string tokenizer_path;
163 std::string data_path;
164 std::string output_dir = "./output";
165 int epochs = 10;
166 float learning_rate = 1e-4f;
167 uint32_t seed = 42;
168 std::string resume_path = "";
169 std::string init_strategy = "xavier";
170 int batch_size = 32;
171 int log_interval = 50;
172 int save_interval = 5000;
173 float grad_clip = 1.0f;
174 bool use_adam = false;
175 int grad_accum_steps = 4;
176 std::string vocab_mask_path = "";
177 int replay_buffer_size = 10000;
178 float replay_ratio = 0.25f;
179 bool use_cuda = false;
180 bool verbose = false;
181};
182
183TrainConfig parse_args(int argc, char* argv[]) {
184 TrainConfig cfg;
185 for (int i = 1; i < argc; ++i) {
186 std::string arg = argv[i];
187 if (arg == "--config" && i + 1 < argc) cfg.config_path = argv[++i];
188 else if (arg == "--tokenizer" && i + 1 < argc) cfg.tokenizer_path = argv[++i];
189 else if (arg == "--data" && i + 1 < argc) cfg.data_path = argv[++i];
190 else if (arg == "--output" && i + 1 < argc) cfg.output_dir = argv[++i];
191 else if (arg == "--epochs" && i + 1 < argc) cfg.epochs = std::atoi(argv[++i]);
192 else if (arg == "--lr" && i + 1 < argc) cfg.learning_rate = std::atof(argv[++i]);
193 else if (arg == "--seed" && i + 1 < argc) cfg.seed = static_cast<uint32_t>(std::atoi(argv[++i]));
194 else if (arg == "--resume" && i + 1 < argc) cfg.resume_path = argv[++i];
195 else if (arg == "--init-weights" && i + 1 < argc) cfg.init_strategy = argv[++i];
196 else if (arg == "--batch-size" && i + 1 < argc) cfg.batch_size = std::atoi(argv[++i]);
197 else if (arg == "--log-interval" && i + 1 < argc) cfg.log_interval = std::atoi(argv[++i]);
198 else if (arg == "--save-interval" && i + 1 < argc) cfg.save_interval = std::atoi(argv[++i]);
199 else if (arg == "--grad-clip" && i + 1 < argc) cfg.grad_clip = std::atof(argv[++i]);
200 else if (arg == "--adam") cfg.use_adam = true;
201 else if (arg == "--grad-accum" && i + 1 < argc) cfg.grad_accum_steps = std::atoi(argv[++i]);
202 else if (arg == "--vocab-mask" && i + 1 < argc) cfg.vocab_mask_path = argv[++i];
203 else if (arg == "--replay-buffer" && i + 1 < argc) cfg.replay_buffer_size = std::atoi(argv[++i]);
204 else if (arg == "--replay-ratio" && i + 1 < argc) cfg.replay_ratio = std::atof(argv[++i]);
205 else if (arg == "--use-cuda") cfg.use_cuda = true;
206 else if (arg == "--verbose") cfg.verbose = true;
207 }
208 return cfg;
209}
210
211namespace {
212
213std::string read_file(const std::string& path) {
214 std::ifstream ifs(path);
215 if (!ifs) return "";
216 return std::string((std::istreambuf_iterator<char>(ifs)), std::istreambuf_iterator<char>());
217}
218
219std::string trim(const std::string& s) {
220 size_t start = s.find_first_not_of(" \t\r\n");
221 if (start == std::string::npos) return "";
222 size_t end = s.find_last_not_of(" \t\r\n");
223 return s.substr(start, end - start + 1);
224}
225
226std::string extract_json_string(const std::string& json, const std::string& key) {
227 std::string search = "\"" + key + "\"";
228 size_t pos = json.find(search);
229 if (pos == std::string::npos) return "";
230 pos = json.find(':', pos + search.size());
231 if (pos == std::string::npos) return "";
232 pos++;
233 while (pos < json.size() && (json[pos] == ' ' || json[pos] == '\t')) pos++;
234 if (pos >= json.size() || json[pos] != '"') return "";
235 size_t end = json.find('"', pos + 1);
236 if (end == std::string::npos) return "";
237 return json.substr(pos + 1, end - pos - 1);
238}
239
240size_t extract_json_number(const std::string& json, const std::string& key, size_t default_val = 0) {
241 std::string search = "\"" + key + "\"";
242 size_t pos = json.find(search);
243 if (pos == std::string::npos) return default_val;
244 pos = json.find(':', pos + search.size());
245 if (pos == std::string::npos) return default_val;
246 pos++;
247 while (pos < json.size() && (json[pos] == ' ' || json[pos] == '\t')) pos++;
248 size_t end = pos;
249 while (end < json.size() && (json[end] >= '0' && json[end] <= '9')) end++;
250 if (end == pos) return default_val;
251 return std::stoul(json.substr(pos, end - pos));
252}
253
254bool extract_json_bool(const std::string& json, const std::string& key, bool default_val = false) {
255 std::string search = "\"" + key + "\"";
256 size_t pos = json.find(search);
257 if (pos == std::string::npos) return default_val;
258 pos = json.find(':', pos + search.size());
259 if (pos == std::string::npos) return default_val;
260 pos++;
261 while (pos < json.size() && (json[pos] == ' ' || json[pos] == '\t')) pos++;
262 if (pos + 3 < json.size() && json.substr(pos, 4) == "true") return true;
263 if (pos + 4 < json.size() && json.substr(pos, 5) == "false") return false;
264 return default_val;
265}
266
267}
268
269NeuroFlowModel::Config load_model_config(const std::string& json_path) {
270 NeuroFlowModel::Config cfg;
271 std::string json = read_file(json_path);
272 if (json.empty()) {
273 std::cerr << "配置加载: " << json_path << " (使用默认配置)" << std::endl;
274 return cfg;
275 }
276
277 cfg.input_dim = extract_json_number(json, "d_model", cfg.input_dim);
278 cfg.hidden_dim = extract_json_number(json, "hidden_dim", cfg.hidden_dim);
279 cfg.output_dim = extract_json_number(json, "output_dim", cfg.output_dim);
280 cfg.memory_dim = extract_json_number(json, "memory_dim", cfg.memory_dim);
281 cfg.memory_slots = extract_json_number(json, "memory_slots", cfg.memory_slots);
282 cfg.num_layers = extract_json_number(json, "num_layers", cfg.num_layers);
283 cfg.num_associations = extract_json_number(json, "num_associations", cfg.num_associations);
284 cfg.use_quantization = extract_json_bool(json, "use_quantization", cfg.use_quantization);
285 cfg.use_mla = extract_json_bool(json, "use_mla", cfg.use_mla);
286 cfg.use_causal_lm = extract_json_bool(json, "use_causal_lm", cfg.use_causal_lm);
287 cfg.mla_latent_dim = extract_json_number(json, "mla_latent_dim", cfg.mla_latent_dim);
288 cfg.vocab_size = extract_json_number(json, "vocab_size", cfg.vocab_size);
289 cfg.max_seq_len = extract_json_number(json, "max_seq_len", cfg.max_seq_len);
290 cfg.causal_window_size = extract_json_number(json, "causal_window_size", cfg.causal_window_size);
291 cfg.lm_num_attn_layers = extract_json_number(json, "lm_num_attn_layers", cfg.lm_num_attn_layers);
292 cfg.lm_num_attn_heads = extract_json_number(json, "lm_num_attn_heads", cfg.lm_num_attn_heads);
293 cfg.lm_n_kv_heads = extract_json_number(json, "lm_n_kv_heads", cfg.lm_n_kv_heads);
294 cfg.lm_use_rope = extract_json_bool(json, "lm_use_rope", cfg.lm_use_rope);
295 cfg.lm_use_qk_norm = extract_json_bool(json, "lm_use_qk_norm", cfg.lm_use_qk_norm);
296 cfg.lm_use_swiglu = extract_json_bool(json, "lm_use_swiglu", cfg.lm_use_swiglu);
297 cfg.lm_use_bridge = extract_json_bool(json, "lm_use_bridge", cfg.lm_use_bridge);
298 cfg.lm_pooling = extract_json_string(json, "lm_pooling");
299 if (cfg.lm_pooling.empty()) cfg.lm_pooling = "last";
300
301 std::cerr << "配置加载: " << json_path << std::endl;
302 std::cerr << " d_model=" << cfg.input_dim << " hidden_dim=" << cfg.hidden_dim
303 << " output_dim=" << cfg.output_dim << " vocab_size=" << cfg.vocab_size << std::endl;
304 return cfg;
305}
306
307struct ResumeState {
308 bool found = false;
309 size_t step = 0;
310 int epoch = 0;
311 float loss = 0.0f;
312 float lr = 0.0f;
313 size_t config_d_model = 0;
314 size_t config_output_dim = 0;
315};
316
317ResumeState load_checkpoint(NeuroFlowModel& model, const std::string& ckpt_path) {
318 ResumeState state;
319
320 std::string model_path, state_path;
321 if (ckpt_path.size() >= 5 && ckpt_path.substr(ckpt_path.size() - 5) == ".nfv1") {
322 model_path = ckpt_path;
323 state_path = ckpt_path.substr(0, ckpt_path.size() - 5) + "/../training_state.json";
324 std::string dir = ckpt_path.substr(0, ckpt_path.find_last_of("/\\"));
325 state_path = dir + "/training_state.json";
326 } else {
327 model_path = ckpt_path + "/model.nfv1";
328 state_path = ckpt_path + "/training_state.json";
329 }
330
331 std::ifstream ifs(model_path, std::ios::binary);
332 if (!ifs) return state;
333 char magic[5] = {0};
334 ifs.read(magic, 4);
335 if (std::string(magic) != "NFv1") return state;
336 ifs.close();
337
338 std::ifstream state_ifs(state_path);
339 if (state_ifs) {
340 std::string content((std::istreambuf_iterator<char>(state_ifs)),
341 std::istreambuf_iterator<char>());
342 state.step = extract_json_number(content, "step", 0);
343 state.epoch = static_cast<int>(extract_json_number(content, "epoch", 0));
344 state.loss = std::atof(extract_json_string(content, "loss").c_str());
345 state.lr = std::atof(extract_json_string(content, "learning_rate").c_str());
346 state.config_d_model = extract_json_number(content, "config_d_model", 0);
347 state.config_output_dim = extract_json_number(content, "config_output_dim", 0);
348 }
349
350 if (state.config_d_model > 0 && state.config_d_model != model.config.input_dim) {
351 std::cerr << "警告: checkpoint的d_model=" << state.config_d_model
352 << " 与当前config的d_model=" << model.config.input_dim
353 << " 不匹配,可能导致维度错误!" << std::endl;
354 }
355 if (state.config_output_dim > 0 && state.config_output_dim != model.config.output_dim) {
356 std::cerr << "警告: checkpoint的output_dim=" << state.config_output_dim
357 << " 与当前config的output_dim=" << model.config.output_dim
358 << " 不匹配!" << std::endl;
359 }
360
361 model.load(model_path);
362 state.found = true;
363 return state;
364}
365
366bool load_lm_head(CausalLMHead& lm_head, const std::string& lm_path) {
367 std::ifstream ifs(lm_path, std::ios::binary);
368 if (!ifs) {
369 std::cerr << "LM Head checkpoint未找到: " << lm_path << std::endl;
370 return false;
371 }
372
373 char magic[5] = {0};
374 ifs.read(magic, 4);
375 if (std::string(magic) != "LMH2" && std::string(magic) != "LMH1") {
376 std::cerr << "LM Head格式不匹配: " << std::string(magic) << " (期望LMH2/LMH1)" << std::endl;
377 ifs.close();
378 return false;
379 }
380
381 auto read_named_tensor = [&]() -> std::pair<std::string, Tensor> {
382 uint32_t nl = 0; ifs.read((char*)&nl, 4);
383 if (nl == 0 || ifs.eof()) return {"", Tensor()};
384 std::string name(nl, '\0'); ifs.read(&name[0], nl);
385 uint32_t nd = 0; ifs.read((char*)&nd, 4);
386 std::vector<size_t> shape(nd);
387 for (size_t i = 0; i < nd; ++i) { uint32_t d = 0; ifs.read((char*)&d, 4); shape[i] = d; }
388 uint32_t ds = 0; ifs.read((char*)&ds, 4);
389 Tensor t(shape, QuantType::FP32);
390 if (t.data_size_ == ds) {
391 ifs.read((char*)t.data_.get(), ds);
392 } else {
393 ifs.seekg(ds, std::ios::cur);
394 t = Tensor();
395 }
396 return {name, std::move(t)};
397 };
398
399 std::unordered_map<std::string, Tensor*> tensor_map;
400 tensor_map["w_embed"] = &lm_head.w_embed_;
401 tensor_map["w_pos"] = &lm_head.w_pos_;
402 tensor_map["dw_kernel"] = &lm_head.dw_kernel_;
403 tensor_map["pw_conv.weight"] = &lm_head.pw_conv_->weight;
404 tensor_map["sae_encode.weight"] = &lm_head.sae_w_encode_->weight;
405 tensor_map["sae_decode.weight"] = &lm_head.sae_w_decode_->weight;
406 tensor_map["ntm_read.weight"] = &lm_head.ntm_w_read_->weight;
407 tensor_map["ntm_write.weight"] = &lm_head.ntm_w_write_->weight;
408 tensor_map["ntm_erase.weight"] = &lm_head.ntm_w_erase_->weight;
409 tensor_map["ntm_memory"] = &lm_head.ntm_memory_;
410 tensor_map["w_proj.weight"] = &lm_head.w_proj_->weight;
411 tensor_map["w_proj.bias"] = &lm_head.w_proj_->bias;
412 if (lm_head.bridge_) {
413 tensor_map["bridge.weight"] = &lm_head.bridge_->weight;
414 tensor_map["bridge.bias"] = &lm_head.bridge_->bias;
415 }
416 tensor_map["w_out.weight"] = &lm_head.w_out_->weight;
417 tensor_map["w_out.bias"] = &lm_head.w_out_->bias;
418 tensor_map["ln.weight"] = &lm_head.ln_->weight;
419 tensor_map["ln.bias"] = &lm_head.ln_->bias;
420 for (size_t i = 0; i < lm_head.attn_layers_.size(); ++i) {
421 std::string p = "attn" + std::to_string(i) + ".";
422 tensor_map[p + "w_q.weight"] = &lm_head.attn_layers_[i]->w_q->weight;
423 tensor_map[p + "w_q.bias"] = &lm_head.attn_layers_[i]->w_q->bias;
424 tensor_map[p + "w_k.weight"] = &lm_head.attn_layers_[i]->w_k->weight;
425 tensor_map[p + "w_k.bias"] = &lm_head.attn_layers_[i]->w_k->bias;
426 tensor_map[p + "w_v.weight"] = &lm_head.attn_layers_[i]->w_v->weight;
427 tensor_map[p + "w_v.bias"] = &lm_head.attn_layers_[i]->w_v->bias;
428 tensor_map[p + "w_out.weight"] = &lm_head.attn_layers_[i]->w_out->weight;
429 tensor_map[p + "w_out.bias"] = &lm_head.attn_layers_[i]->w_out->bias;
430 tensor_map[p + "norm.weight"] = &lm_head.attn_layers_[i]->norm->weight;
431 tensor_map[p + "norm.bias"] = &lm_head.attn_layers_[i]->norm->bias;
432 }
433
434 size_t loaded = 0;
435 while (ifs) {
436 auto [name, tensor] = read_named_tensor();
437 if (name.empty()) break;
438 auto it = tensor_map.find(name);
439 if (it != tensor_map.end() && tensor.numel() > 0) {
440 if (it->second->shape_ == tensor.shape_) {
441 memcpy(it->second->data_.get(), tensor.data_.get(), tensor.data_size_);
442 loaded++;
443 } else {
444 std::cerr << " 跳过 '" << name << "': 形状不匹配" << std::endl;
445 }
446 } else if (name.find("w_qkv.") != std::string::npos) {
447 std::string base = name.substr(0, name.find("w_qkv."));
448 std::string suffix = name.substr(name.find("w_qkv.") + 6);
449 for (size_t i = 0; i < lm_head.attn_layers_.size(); ++i) {
450 std::string p = "attn" + std::to_string(i) + ".";
451 if (base != p) continue;
452 size_t d_model = lm_head.config_.d_model;
453 size_t head_dim = d_model / lm_head.config_.num_attn_heads;
454 size_t n_q = lm_head.attn_layers_[i]->n_q_heads_;
455 size_t n_kv = lm_head.attn_layers_[i]->n_kv_heads_;
456 if (suffix == "weight" && tensor.shape_.size() == 2 && tensor.shape_[0] == 3 * d_model) {
457 const float* src = tensor.as_fp32();
458 float* dq = lm_head.attn_layers_[i]->w_q->weight.as_fp32();
459 float* dk = lm_head.attn_layers_[i]->w_k->weight.as_fp32();
460 float* dv = lm_head.attn_layers_[i]->w_v->weight.as_fp32();
461 for (size_t r = 0; r < d_model; ++r) {
462 memcpy(dq + r * n_q * head_dim, src + r * 3 * d_model, n_q * head_dim * sizeof(float));
463 memcpy(dk + r * n_kv * head_dim, src + r * 3 * d_model + d_model, n_kv * head_dim * sizeof(float));
464 memcpy(dv + r * n_kv * head_dim, src + r * 3 * d_model + 2 * d_model, n_kv * head_dim * sizeof(float));
465 }
466 loaded++;
467 std::cerr << " 拆分旧格式 '" << name << "' -> w_q/w_k/w_v" << std::endl;
468 } else if (suffix == "bias" && tensor.shape_[0] == 3 * d_model) {
469 const float* src = tensor.as_fp32();
470 float* bq = lm_head.attn_layers_[i]->w_q->bias.as_fp32();
471 float* bk = lm_head.attn_layers_[i]->w_k->bias.as_fp32();
472 float* bv = lm_head.attn_layers_[i]->w_v->bias.as_fp32();
473 memcpy(bq, src, n_q * head_dim * sizeof(float));
474 memcpy(bk, src + d_model, n_kv * head_dim * sizeof(float));
475 memcpy(bv, src + 2 * d_model, n_kv * head_dim * sizeof(float));
476 loaded++;
477 std::cerr << " 拆分旧格式 '" << name << "' -> w_q/w_k/w_v bias" << std::endl;
478 }
479 }
480 }
481 }
482
483 ifs.close();
484 if (lm_head.config_.weight_tying) lm_head.tie_weights();
485 std::cerr << "LM Head已恢复: " << loaded << " 个张量从 " << lm_path << std::endl;
486 return loaded > 0;
487}
488
489NeuroFlowModel build_model(const NeuroFlowModel::Config& cfg, const TrainConfig& train_cfg,
490 ResumeState& resume_state) {
491 NeuroFlowModel model(cfg);
492 if (!train_cfg.resume_path.empty()) {
493 resume_state = load_checkpoint(model, train_cfg.resume_path);
494 if (resume_state.found) {
495 std::cerr << "从checkpoint恢复: " << train_cfg.resume_path
496 << " (step=" << resume_state.step
497 << ", epoch=" << resume_state.epoch
498 << ", loss=" << resume_state.loss << ")" << std::endl;
499 } else {
500 throw std::runtime_error("无法加载checkpoint: " + train_cfg.resume_path);
501 }
502 } else {
503 InitStrategy strategy = InitStrategy::XAVIER_UNIFORM;
504 if (train_cfg.init_strategy == "kaiming") strategy = InitStrategy::KAIMING_NORMAL;
505 else if (train_cfg.init_strategy == "zeros") strategy = InitStrategy::ZEROS;
506 WeightInitializer::init_model_weights(model, strategy, train_cfg.seed);
507 std::cerr << "权重初始化: " << train_cfg.init_strategy << std::endl;
508 }
509 return model;
510}
511
512static void unescape_json(std::string& s) {
513 size_t pos = 0;
514 while ((pos = s.find('\\', pos)) != std::string::npos) {
515 if (pos + 1 < s.size()) {
516 char c = s[pos + 1];
517 if (c == 'n') { s.replace(pos, 2, "\n"); pos++; }
518 else if (c == 't') { s.replace(pos, 2, "\t"); pos++; }
519 else if (c == 'r') { s.replace(pos, 2, "\r"); pos++; }
520 else if (c == '"') { s.replace(pos, 2, "\""); pos++; }
521 else if (c == '\\') { s.replace(pos, 2, "\\"); pos++; }
522 else if (c == '/') { s.replace(pos, 2, "/"); pos++; }
523 else if (c == 'u' && pos + 5 < s.size()) { pos += 6; }
524 else { pos += 2; }
525 } else { pos++; }
526 }
527}
528
529struct LMSample {
530 std::vector<size_t> token_ids;
531};
532
533struct TextSpan {
534 size_t file_index;
535 size_t offset;
536 size_t length;
537};
538
539struct TrainingSample {
540 std::vector<float> input;
541 std::vector<float> target;
542};
543
544class StreamingDataLoader {
545public:
546 StreamingDataLoader(const std::string& path, int batch_size,
547 BPETokenizer* tokenizer, size_t max_seq_len = 128,
548 size_t max_samples = 500000)
549 : batch_size_(batch_size), tokenizer_(tokenizer),
550 max_seq_len_(max_seq_len), max_samples_(max_samples),
551 current_idx_(0) {
552 index_files(path);
553 std::cerr << "StreamingDataLoader: " << spans_.size() << " 文本段已索引"
554 << " (" << file_paths_.size() << " 文件)" << std::endl;
555 }
556
557 void index_files(const std::string& path) {
558 namespace fs = std::filesystem;
559 fs::path p(path);
560
561 if (fs::is_directory(p)) {
562 std::vector<fs::path> files;
563 for (auto& entry : fs::recursive_directory_iterator(p, fs::directory_options::follow_directory_symlink)) {
564 if (!entry.is_regular_file()) continue;
565 auto ext = entry.path().extension().string();
566 if (ext == ".txt" || ext == ".json" || ext == ".jsonl" || ext == ".csv" || ext == ".tsv") {
567 files.push_back(entry.path());
568 }
569 }
570 std::sort(files.begin(), files.end());
571 for (auto& f : files) index_single_file(f.string());
572 } else {
573 index_single_file(p.string());
574 }
575 }
576
577 void index_single_file(const std::string& file_path) {
578 namespace fs = std::filesystem;
579 std::ifstream ifs(file_path, std::ios::binary | std::ios::ate);
580 if (!ifs) return;
581
582 size_t file_size = static_cast<size_t>(ifs.tellg());
583 ifs.seekg(0, std::ios::beg);
584
585 std::string ext = fs::path(file_path).extension().string();
586 size_t file_idx = file_paths_.size();
587 file_paths_.push_back(file_path);
588
589 if (ext == ".jsonl" || ext == ".json") {
590 std::string line;
591 size_t line_offset = 0;
592 while (std::getline(ifs, line) && spans_.size() < max_samples_) {
593 if (line.size() < 12) { line_offset += line.size() + 1; continue; }
594 std::string text = extract_text_from_line(line);
595 if (text.size() >= 10) {
596 text_cache_.push_back(std::move(text));
597 spans_.push_back(TextSpan{SIZE_MAX, text_cache_.size() - 1, text_cache_.back().size()});
598 }
599 line_offset += line.size() + 1;
600 }
601 } else if (ext == ".txt") {
602 std::string line;
603 while (std::getline(ifs, line) && spans_.size() < max_samples_) {
604 if (line.size() < 10) continue;
605 text_cache_.push_back(std::move(line));
606 spans_.push_back(TextSpan{SIZE_MAX, text_cache_.size() - 1, text_cache_.back().size()});
607 }
608 } else if (ext == ".csv" || ext == ".tsv") {
609 char delim = (ext == ".tsv") ? '\t' : ',';
610 std::string line;
611 while (std::getline(ifs, line) && spans_.size() < max_samples_) {
612 size_t pos = line.find(delim);
613 if (pos == std::string::npos) continue;
614 std::string text = line.substr(pos + 1);
615 if (text.size() >= 10) {
616 text_cache_.push_back(std::move(text));
617 spans_.push_back(TextSpan{SIZE_MAX, text_cache_.size() - 1, text_cache_.back().size()});
618 }
619 }
620 }
621 }
622
623 std::string extract_text_from_line(const std::string& line) {
624 const char* fields[] = {"\"text\"", "\"content\"", "\"title\"", "\"question\"", "\"answer\""};
625 for (auto& field : fields) {
626 size_t pos = line.find(field);
627 if (pos == std::string::npos) continue;
628 size_t colon = line.find(':', pos + strlen(field));
629 if (colon == std::string::npos) continue;
630 size_t q1 = line.find('"', colon);
631 if (q1 == std::string::npos) continue;
632 size_t q2 = line.find('"', q1 + 1);
633 if (q2 == std::string::npos) continue;
634 std::string text = line.substr(q1 + 1, q2 - q1 - 1);
635 unescape_json(text);
636 return text;
637 }
638 return "";
639 }
640
641 LMSample get(size_t idx) const {
642 if (idx >= spans_.size()) return {};
643 const auto& span = spans_[idx];
644 const std::string& text = text_cache_[span.offset];
645 auto ids = tokenizer_->encode(text, max_seq_len_);
646 return LMSample{std::move(ids)};
647 }
648
649 bool has_next() const { return current_idx_ < spans_.size(); }
650
651 std::vector<LMSample> next_lm_batch_raw() {
652 std::vector<LMSample> batch;
653 for (int i = 0; i < batch_size_ && current_idx_ < spans_.size(); ++i) {
654 batch.push_back(get(current_idx_++));
655 }
656 return batch;
657 }
658
659 void reset() { current_idx_ = 0; }
660
661 void shuffle(std::mt19937& rng) {
662 std::shuffle(spans_.begin(), spans_.end(), rng);
663 }
664
665 size_t total_samples() const { return spans_.size(); }
666
667 size_t memory_usage_bytes() const {
668 size_t total = spans_.capacity() * sizeof(TextSpan);
669 for (auto& t : text_cache_) total += t.capacity();
670 return total;
671 }
672
673private:
674 int batch_size_;
675 BPETokenizer* tokenizer_;
676 size_t max_seq_len_;
677 size_t max_samples_;
678 size_t current_idx_;
679 std::vector<std::string> file_paths_;
680 std::vector<TextSpan> spans_;
681 std::vector<std::string> text_cache_;
682};
683
684class DataLoader {
685public:
686 enum class Mode { LM, CLASSIFICATION, REGRESSION };
687
688 DataLoader(const std::string& path, int batch_size, size_t input_dim, size_t output_dim,
689 Mode mode = Mode::LM, BPETokenizer* tokenizer = nullptr,
690 size_t max_seq_len = 128, size_t max_samples = 500000,
691 size_t max_memory_mb = 4096)
692 : path_(path), batch_size_(batch_size), input_dim_(input_dim), output_dim_(output_dim),
693 current_idx_(0), mode_(mode), tokenizer_(tokenizer), max_seq_len_(max_seq_len),
694 max_samples_(max_samples), max_memory_bytes_(max_memory_mb * 1024ULL * 1024ULL) {
695 load_data();
696 }
697
698 void load_data() {
699 if (path_.empty()) {
700 throw std::runtime_error("DataLoader: 数据路径为空,请指定--data参数");
701 }
702
703 namespace fs = std::filesystem;
704 fs::path p(path_);
705
706 if (!fs::exists(p)) {
707 throw std::runtime_error("DataLoader: 路径不存在: " + path_);
708 }
709
710 if (fs::is_directory(p)) {
711 load_directory(p.string());
712 } else {
713 load_file(p.string());
714 }
715
716 if (lm_samples_.empty() && samples_.empty()) {
717 throw std::runtime_error(
718 "DataLoader: 未能从 " + path_ + " 加载任何训练数据。"
719 "请检查文件格式(支持txt/json/jsonl/csv/tsv)和内容是否有效。");
720 }
721
722 if (!lm_samples_.empty()) {
723 std::cerr << "DataLoader: " << lm_samples_.size() << " LM样本已加载" << std::endl;
724 } else {
725 std::cerr << "DataLoader: " << samples_.size() << " 数值样本已加载" << std::endl;
726 }
727 }
728
729 void load_directory(const std::string& dir_path) {
730 namespace fs = std::filesystem;
731 size_t file_count = 0;
732 size_t skipped = 0;
733
734 std::cerr << "DataLoader: 扫描目录 " << dir_path << " ..." << std::endl;
735 std::vector<fs::path> files;
736 for (auto& entry : fs::recursive_directory_iterator(dir_path, fs::directory_options::follow_directory_symlink)) {
737 if (!entry.is_regular_file()) continue;
738 files.push_back(entry.path());
739 }
740 std::cerr << "DataLoader: 发现 " << files.size() << " 个文件" << std::endl;
741 std::sort(files.begin(), files.end());
742
743 for (size_t fi = 0; fi < files.size(); ++fi) {
744 auto& fpath = files[fi];
745 std::string ext = fpath.extension().string();
746 std::transform(ext.begin(), ext.end(), ext.begin(), ::tolower);
747
748 if (ext == ".json" || ext == ".jsonl" || ext == ".txt" ||
749 ext == ".csv" || ext == ".tsv" || ext == ".tok1" || ext == ".bin") {
750 size_t file_size_mb = 0;
751 try { file_size_mb = fs::file_size(fpath) / 1024 / 1024; } catch (...) {}
752 std::cerr << "[" << (fi + 1) << "/" << files.size() << "] 加载: "
753 << fpath.filename().string()
754 << " (" << file_size_mb << " MB)" << std::endl;
755
756 auto load_start = std::chrono::steady_clock::now();
757 try {
758 load_file(fpath.string());
759 file_count++;
760 } catch (const std::exception& e) {
761 std::cerr << " 跳过: " << e.what() << std::endl;
762 skipped++;
763 }
764 auto load_end = std::chrono::steady_clock::now();
765 double load_sec = std::chrono::duration<double>(load_end - load_start).count();
766
767 size_t total = lm_samples_.size() + samples_.size();
768 std::cerr << " 完成: " << load_sec << "s, 累计 "
769 << total << " 样本" << std::endl;
770
771 if (total >= max_samples_) {
772 std::cerr << "DataLoader: 已达采样上限 " << max_samples_
773 << ",停止遍历" << std::endl;
774 break;
775 }
776 } else if (ext == ".parquet") {
777 skipped++;
778 }
779 }
780 std::cerr << "DataLoader: 遍历 " << file_count << " 个文件, 跳过 " << skipped << std::endl;
781 }
782
783 void load_file(const std::string& file_path) {
784 namespace fs = std::filesystem;
785 std::string ext = fs::path(file_path).extension().string();
786 std::transform(ext.begin(), ext.end(), ext.begin(), ::tolower);
787
788 {
789 std::ifstream ifs(file_path, std::ios::binary);
790 if (ifs) {
791 char magic[5] = {0};
792 ifs.read(magic, 4);
793 ifs.close();
794 if (std::string(magic) == "TOK1") {
795 load_tok1(file_path);
796 return;
797 }
798 }
799 }
800
801 if (mode_ == Mode::LM && tokenizer_) {
802 if (ext == ".json") {
803 size_t file_size = fs::file_size(file_path);
804 if (file_size > 512 * 1024 * 1024) {
805 load_json_mmap(file_path);
806 } else {
807 load_json(file_path);
808 }
809 }
810 else if (ext == ".jsonl") load_jsonl(file_path);
811 else if (ext == ".txt") load_txt(file_path);
812 else if (ext == ".csv" || ext == ".tsv") load_txt(file_path);
813 else if (ext == ".tok1" || ext == ".bin") load_tok1(file_path);
814 else if (ext == ".parquet") {
815 throw std::runtime_error("Parquet格式需Python预处理: python3 scripts/preprocess_corpus.py split --input <path>");
816 } else {
817 throw std::runtime_error("不支持的文件格式: " + ext);
818 }
819 } else {
820 if (ext == ".csv" || ext == ".tsv") load_csv(file_path);
821 else if (ext == ".tok1" || ext == ".bin") load_tok1(file_path);
822 else {
823 throw std::runtime_error("数值模式仅支持csv/tsv/tok1格式,收到: " + ext);
824 }
825 }
826 }
827
828 void load_tok1(const std::string& file_path) {
829 std::ifstream ifs(file_path, std::ios::binary);
830 if (!ifs) throw std::runtime_error("无法打开TOK1文件: " + file_path);
831
832 char magic[5] = {0};
833 ifs.read(magic, 4);
834 if (std::string(magic) != "TOK1") {
835 throw std::runtime_error("TOK1格式Magic不匹配: " + std::string(magic));
836 }
837
838 uint16_t version = 0;
839 ifs.read(reinterpret_cast<char*>(&version), 2);
840
841 uint32_t vocab_size = 0, max_seq = 0, total_samples = 0;
842 ifs.read(reinterpret_cast<char*>(&vocab_size), 4);
843 ifs.read(reinterpret_cast<char*>(&max_seq), 4);
844 ifs.read(reinterpret_cast<char*>(&total_samples), 4);
845
846 std::cerr << "TOK1: version=" << version << " vocab=" << vocab_size
847 << " max_seq=" << max_seq << " samples=" << total_samples << std::endl;
848
849 size_t loaded = 0;
850 for (size_t i = 0; i < total_samples; ++i) {
851 uint16_t seq_len = 0;
852 ifs.read(reinterpret_cast<char*>(&seq_len), 2);
853 if (!ifs) break;
854
855 LMSample sample;
856 sample.token_ids.resize(seq_len);
857 ifs.read(reinterpret_cast<char*>(sample.token_ids.data()), seq_len * 4);
858 if (!ifs) break;
859
860 lm_samples_.push_back(std::move(sample));
861 loaded++;
862
863 if (lm_samples_.size() >= max_samples_) {
864 std::cerr << "TOK1: 达到采样上限 " << max_samples_ << std::endl;
865 break;
866 }
867 }
868
869 std::cerr << "TOK1: 加载 " << loaded << " 样本" << std::endl;
870 }
871
872 void load_json_mmap(const std::string& file_path) {
873#ifdef _WIN32
874 load_json_chunked(file_path);
875#else
876 int fd = open(file_path.c_str(), O_RDONLY);
877 if (fd < 0) throw std::runtime_error("无法mmap: " + file_path);
878
879 size_t file_size = lseek(fd, 0, SEEK_END);
880 lseek(fd, 0, SEEK_SET);
881
882 void* addr = mmap(nullptr, file_size, PROT_READ, MAP_PRIVATE, fd, 0);
883 if (addr == MAP_FAILED) {
884 close(fd);
885 throw std::runtime_error("mmap失败: " + file_path);
886 }
887
888 const char* data = static_cast<const char*>(addr);
889 std::cerr << "mmap: 映射 " << file_path << " (" << file_size / 1024 / 1024 << " MB)" << std::endl;
890
891 size_t pos = 0;
892 size_t brace_depth = 0;
893 std::string record;
894 const char* fields[] = {"\"text\"", "\"content\"", "\"title\"", "\"question\"", "\"answer\""};
895
896 while (pos < file_size && lm_samples_.size() < max_samples_) {
897 char ch = data[pos];
898 if (ch == '{') {
899 if (brace_depth == 0) record.clear();
900 brace_depth++;
901 } else if (ch == '}') {
902 brace_depth--;
903 if (brace_depth == 0 && !record.empty()) {
904 record += '}';
905 extract_text_from_json_record(record);
906 record.clear();
907 }
908 }
909 if (brace_depth > 0) record += ch;
910 pos++;
911
912 if (lm_samples_.size() % 100000 == 0 && lm_samples_.size() > 0 && pos % (100 * 1024 * 1024) == 0) {
913 std::cerr << " mmap进度: " << pos * 100 / file_size << "%, "
914 << lm_samples_.size() << " 样本" << std::endl;
915 }
916 }
917
918 munmap(addr, file_size);
919 close(fd);
920 std::cerr << "mmap: 完成, " << lm_samples_.size() << " 样本" << std::endl;
921#endif
922 }
923
924 void load_json_chunked(const std::string& file_path) {
925 const size_t chunk_size = 256 * 1024 * 1024;
926 std::ifstream ifs(file_path, std::ios::binary);
927 if (!ifs) throw std::runtime_error("无法打开: " + file_path);
928
929 size_t file_size = std::filesystem::file_size(file_path);
930 std::cerr << "分块读取: " << file_path << " (" << file_size / 1024 / 1024 << " MB)" << std::endl;
931
932 size_t offset = 0;
933 size_t brace_depth = 0;
934 std::string record;
935 std::string overlap;
936
937 while (offset < file_size && lm_samples_.size() < max_samples_) {
938 size_t read_size = std::min(chunk_size, file_size - offset);
939 std::string buffer;
940 buffer.resize(read_size);
941 ifs.read(&buffer[0], read_size);
942 if (!ifs) break;
943
944 if (!overlap.empty()) {
945 buffer = overlap + buffer;
946 overlap.clear();
947 }
948
949 size_t pos = 0;
950 while (pos < buffer.size() && lm_samples_.size() < max_samples_) {
951 char ch = buffer[pos];
952 if (ch == '{') {
953 if (brace_depth == 0) record.clear();
954 brace_depth++;
955 } else if (ch == '}') {
956 brace_depth--;
957 if (brace_depth == 0 && !record.empty()) {
958 record += '}';
959 extract_text_from_json_record(record);
960 record.clear();
961 }
962 }
963 if (brace_depth > 0 && record.size() < 1048576) record += ch;
964 pos++;
965 }
966
967 if (brace_depth > 0) {
968 overlap = record;
969 }
970
971 offset += read_size;
972 std::cerr << " 分块进度: " << offset * 100 / file_size << "%, "
973 << lm_samples_.size() << " 样本" << std::endl;
974 }
975
976 if (!record.empty() && brace_depth == 0) {
977 extract_text_from_json_record(record);
978 }
979 }
980
981 void load_json(const std::string& file_path) {
982 std::ifstream ifs(file_path);
983 if (!ifs) throw std::runtime_error("无法打开: " + file_path);
984
985 std::string line;
986 size_t brace_depth = 0;
987 std::string record;
988 size_t records_found = 0;
989
990 while (std::getline(ifs, line)) {
991 for (size_t i = 0; i < line.size(); ++i) {
992 if (line[i] == '{') {
993 if (brace_depth == 0) record.clear();
994 brace_depth++;
995 } else if (line[i] == '}') {
996 brace_depth--;
997 if (brace_depth == 0 && !record.empty()) {
998 record += '}';
999 extract_text_from_json_record(record);
1000 record.clear();
1001 records_found++;
1002 if (lm_samples_.size() >= max_samples_) return;
1003 }
1004 }
1005 }
1006 if (brace_depth > 0) {
1007 if (record.size() < 1048576) {
1008 record += line;
1009 record += '\n';
1010 }
1011 }
1012 }
1013
1014 if (records_found == 0) {
1015 std::cerr << " 警告: 流式record解析未找到数据,尝试逐行文本提取" << std::endl;
1016 ifs.clear();
1017 ifs.seekg(0);
1018 std::string line;
1019 while (std::getline(ifs, line)) {
1020 if (line.size() < 20) continue;
1021 size_t q1 = line.find('"');
1022 if (q1 == std::string::npos) continue;
1023 size_t q2 = line.rfind('"');
1024 if (q2 <= q1 + 1) continue;
1025 std::string text = line.substr(q1 + 1, q2 - q1 - 1);
1026 unescape_json(text);
1027 if (text.size() >= 10) add_lm_sample(text);
1028 if (lm_samples_.size() >= max_samples_) return;
1029 }
1030 }
1031 }
1032
1033 void extract_text_from_json_record(const std::string& record) {
1034 const char* fields[] = {"\"text\"", "\"content\"", "\"title\"", "\"question\"", "\"answer\""};
1035 for (auto& field : fields) {
1036 size_t pos = record.find(field);
1037 if (pos == std::string::npos) continue;
1038
1039 size_t colon = record.find(':', pos + strlen(field));
1040 if (colon == std::string::npos) continue;
1041 colon++;
1042 while (colon < record.size() && record[colon] != '"') colon++;
1043 if (colon >= record.size()) continue;
1044 colon++;
1045
1046 size_t end = colon;
1047 while (end < record.size() && record[end] != '"') {
1048 if (record[end] == '\\') end++;
1049 end++;
1050 }
1051
1052 std::string text = record.substr(colon, end - colon);
1053 unescape_json(text);
1054 if (text.size() >= 10) add_lm_sample(text);
1055 }
1056 }
1057
1058 void extract_all_text_fields(const std::string& content) {
1059 size_t pos = 0;
1060 const char* fields[] = {"\"text\"", "\"content\"", "\"title\"", "\"question\"", "\"answer\""};
1061 while (pos < content.size()) {
1062 size_t best_pos = std::string::npos;
1063 for (auto& field : fields) {
1064 size_t p = content.find(field, pos);
1065 if (p != std::string::npos && (best_pos == std::string::npos || p < best_pos)) {
1066 best_pos = p;
1067 }
1068 }
1069 if (best_pos == std::string::npos) break;
1070
1071 size_t colon = content.find(':', best_pos);
1072 if (colon == std::string::npos) break;
1073 colon++;
1074 while (colon < content.size() && content[colon] != '"') colon++;
1075 if (colon >= content.size()) break;
1076 colon++;
1077
1078 size_t end = colon;
1079 while (end < content.size() && content[end] != '"') {
1080 if (content[end] == '\\') end++;
1081 end++;
1082 }
1083 if (end >= content.size()) break;
1084
1085 std::string text = content.substr(colon, end - colon);
1086 unescape_json(text);
1087 add_lm_sample(text);
1088
1089 pos = end + 1;
1090 if (lm_samples_.size() >= max_samples_) return;
1091 }
1092 }
1093
1094 void load_jsonl(const std::string& file_path) {
1095 std::ifstream ifs(file_path);
1096 if (!ifs) throw std::runtime_error("无法打开: " + file_path);
1097 std::string line;
1098 size_t line_num = 0;
1099 while (std::getline(ifs, line)) {
1100 line_num++;
1101 if (line.empty() || line[0] == '#') continue;
1102
1103 const char* fields[] = {"\"text\"", "\"content\"", "\"title\"", "\"question\"", "\"answer\""};
1104 bool found = false;
1105 for (auto& field : fields) {
1106 size_t text_pos = line.find(field);
1107 if (text_pos == std::string::npos) continue;
1108
1109 size_t colon = line.find(':', text_pos);
1110 if (colon == std::string::npos) continue;
1111 colon++;
1112 while (colon < line.size() && line[colon] != '"') colon++;
1113 if (colon >= line.size()) continue;
1114 colon++;
1115
1116 size_t end = colon;
1117 while (end < line.size() && line[end] != '"') {
1118 if (line[end] == '\\') end++;
1119 end++;
1120 }
1121
1122 std::string text = line.substr(colon, end - colon);
1123 unescape_json(text);
1124 add_lm_sample(text);
1125 found = true;
1126 break;
1127 }
1128 if (!found && line_num <= 5) {
1129 std::cerr << " jsonl行" << line_num << ": 未找到text/content字段" << std::endl;
1130 }
1131 if (lm_samples_.size() >= max_samples_) return;
1132 }
1133 }
1134
1135 void load_txt(const std::string& file_path) {
1136 std::ifstream ifs(file_path);
1137 if (!ifs) throw std::runtime_error("无法打开: " + file_path);
1138 std::string line;
1139 std::string paragraph;
1140 size_t line_count = 0;
1141
1142 while (std::getline(ifs, line)) {
1143 line_count++;
1144 if (line.empty()) {
1145 if (paragraph.size() >= 10) {
1146 add_lm_sample(paragraph);
1147 paragraph.clear();
1148 if (lm_samples_.size() >= max_samples_) return;
1149 }
1150 } else {
1151 if (!paragraph.empty()) paragraph += '\n';
1152 paragraph += line;
1153 if (paragraph.size() > 10000) {
1154 add_lm_sample(paragraph);
1155 paragraph.clear();
1156 if (lm_samples_.size() >= max_samples_) return;
1157 }
1158 }
1159 if (line_count % 500000 == 0) {
1160 std::cerr << " load_txt: " << line_count << " 行, "
1161 << lm_samples_.size() << " 样本" << std::endl;
1162 }
1163 }
1164 if (paragraph.size() >= 10) add_lm_sample(paragraph);
1165 }
1166
1167 void load_csv(const std::string& file_path) {
1168 std::ifstream ifs(file_path);
1169 if (!ifs) throw std::runtime_error("无法打开: " + file_path);
1170 std::string line;
1171 while (std::getline(ifs, line)) {
1172 if (line.empty() || line[0] == '#') continue;
1173 TrainingSample s;
1174 std::stringstream ss(line);
1175 std::string val;
1176 char delim = ',';
1177 if (line.find('\t') != std::string::npos && line.find(',') == std::string::npos) {
1178 delim = '\t';
1179 }
1180 bool in_target = false;
1181 while (std::getline(ss, val, delim)) {
1182 float v = std::atof(val.c_str());
1183 if (!in_target) {
1184 s.input.push_back(v);
1185 if (s.input.size() == input_dim_) in_target = true;
1186 } else {
1187 s.target.push_back(v);
1188 }
1189 }
1190 if (!s.input.empty() && !s.target.empty()) {
1191 samples_.push_back(std::move(s));
1192 }
1193 if (samples_.size() >= max_samples_) return;
1194 }
1195 }
1196
1197 void add_lm_sample(const std::string& text) {
1198 if (text.size() < 10) return;
1199 if (!tokenizer_) return;
1200 if (lm_samples_.size() >= max_samples_) return;
