Team Ai
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
1likes
train_v2.cpp1877 linesDownload Raw Back to src
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;

Showing the first 1,200 of 1877 lines. Download the file for the rest.