Team Ai
Modelpublic

cwenzi/neuroflow-cpp

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
1likes
sft_train.cpp336 linesDownload Raw Back to src
1#include "neuroflow/sft.hpp"
2
3#include <algorithm>
4#include <chrono>
5#include <cmath>
6#include <cstring>
7#include <filesystem>
8#include <fstream>
9#include <iostream>
10#include <numeric>
11#include <sstream>
12
13namespace neuroflow {
14
15SFTDataLoader::SFTDataLoader(const std::string& jsonl_path, size_t max_samples) {
16    std::ifstream ifs(jsonl_path);
17    if (!ifs) {
18        std::cerr << "SFT数据文件无法打开: " << jsonl_path << std::endl;
19        return;
20    }
21    std::string line;
22    while (std::getline(ifs, line)) {
23        if (line.empty() || line[0] == '#') continue;
24        std::string instruction = extract_json_string(line, "instruction");
25        std::string response = extract_json_string(line, "response");
26        unescape_json(instruction);
27        unescape_json(response);
28        if (instruction.empty() || response.empty()) {
29            invalid_count_++;
30            continue;
31        }
32        samples_.push_back({instruction, response});
33        if (max_samples > 0 && samples_.size() >= max_samples) break;
34    }
35    std::cerr << "SFT数据加载: " << samples_.size() << " 样本, "
36              << invalid_count_ << " 无效" << std::endl;
37}
38
39bool SFTDataLoader::has_next() const {
40    return cursor_ < samples_.size();
41}
42
43SFTSample SFTDataLoader::next() {
44    return samples_[cursor_++];
45}
46
47void SFTDataLoader::reset() {
48    cursor_ = 0;
49}
50
51void SFTDataLoader::shuffle(std::mt19937& rng) {
52    std::shuffle(samples_.begin(), samples_.end(), rng);
53}
54
55SFTTrainingTensors build_sft_training_tensors(const SFTSample& sample,
56                                                BPETokenizer& tokenizer,
57                                                size_t max_seq_len) {
58    SFTTrainingTensors result;
59    std::string prompt = sample.instruction + "\n";
60    std::string full_text = prompt + sample.response;
61
62    std::vector<size_t> prompt_ids = tokenizer.encode(prompt, max_seq_len);
63    std::vector<size_t> full_ids = tokenizer.encode(full_text, max_seq_len);
64
65    if (full_ids.size() > max_seq_len) {
66        full_ids.resize(max_seq_len);
67    }
68    if (full_ids.size() < 2) {
69        result.instruction_len = 0;
70        return result;
71    }
72
73    result.input_ids.assign(full_ids.begin(), full_ids.end() - 1);
74    result.target_ids.assign(full_ids.begin() + 1, full_ids.end());
75    result.instruction_len = std::min(prompt_ids.size(), result.input_ids.size());
76
77    result.loss_mask.resize(result.target_ids.size(), 0.0f);
78    for (size_t i = result.instruction_len; i < result.target_ids.size(); ++i) {
79        result.loss_mask[i] = 1.0f;
80    }
81
82    return result;
83}
84
85MaskedCEOutput masked_cross_entropy(const Tensor& logits, size_t vocab_size,
86                                     const std::vector<size_t>& target_ids,
87                                     const std::vector<float>& loss_mask) {
88    MaskedCEOutput output;
89    output.loss = 0.0f;
90    output.valid_token_count = 0;
91
92    size_t seq_len = target_ids.size();
93    if (seq_len == 0 || logits.numel() == 0) {
94        output.logits_grad = Tensor({1, vocab_size}, QuantType::FP32);
95        memset(output.logits_grad.as_fp32(), 0, output.logits_grad.data_size_);
96        return output;
97    }
98
99    output.logits_grad = Tensor({seq_len, vocab_size}, QuantType::FP32);
100    memset(output.logits_grad.as_fp32(), 0, output.logits_grad.data_size_);
101    float* lg = output.logits_grad.as_fp32();
102
103    for (size_t t = 0; t < seq_len; ++t) {
104        if (loss_mask[t] < 0.5f) continue;
105
106        size_t target = target_ids[t];
107        if (target >= vocab_size) target = 1;
108
109        const float* pred = logits.as_fp32() + t * vocab_size;
110
111        float max_val = -1e30f;
112        for (size_t j = 0; j < vocab_size; ++j) {
113            if (pred[j] > max_val) max_val = pred[j];
114        }
115        float sum_exp = 0.0f;
116        for (size_t j = 0; j < vocab_size; ++j) {
117            sum_exp += std::exp(pred[j] - max_val);
118        }
119        float log_sum_exp = max_val + std::log(sum_exp);
120        float token_loss = -(pred[target] - log_sum_exp);
121
122        if (!std::isfinite(token_loss)) continue;
123
124        output.loss += token_loss;
125        output.valid_token_count++;
126
127        float* grad_row = lg + t * vocab_size;
128        for (size_t j = 0; j < vocab_size; ++j) {
129            float softmax_val = std::exp(pred[j] - max_val) / sum_exp;
130            grad_row[j] = softmax_val;
131            if (j == target) grad_row[j] -= 1.0f;
132        }
133    }
134
135    if (output.valid_token_count > 0) {
136        output.loss /= static_cast<float>(output.valid_token_count);
137        float inv_n = 1.0f / static_cast<float>(output.valid_token_count);
138        for (size_t i = 0; i < seq_len * vocab_size; ++i) {
139            lg[i] *= inv_n;
140        }
141    }
142
143    return output;
144}
145
146SFTTrainer::SFTTrainer(const SFTTrainConfig& cfg) : config(cfg) {
147    CausalLMConfig lm_config;
148    lm_config.vocab_size = 128000;
149    lm_config.d_model = 512;
150    lm_config.max_seq_len = cfg.max_seq_len;
151    lm_config.num_attn_layers = 4;
152    lm_config.num_attn_heads = 8;
153    lm_config.n_kv_heads = 2;
154    lm_config.use_rope = true;
155    lm_config.use_qk_norm = true;
156    lm_config.use_swiglu = true;
157    lm_config.use_bridge = true;
158    lm_config.weight_tying = true;
159    lm_config.pooling = "last";
160
161    model_ = std::make_unique<CausalLMHead>(lm_config);
162    if (!cfg.ckpt_path.empty()) {
163        load_lm_checkpoint(*model_, cfg.ckpt_path);
164    }
165
166    tokenizer_ = std::make_unique<BPETokenizer>(cfg.tokenizer_path);
167
168    size_t total_steps = 0;
169    {
170        SFTDataLoader tmp_loader(cfg.data_path);
171        total_steps = tmp_loader.total_samples() * cfg.epochs;
172    }
173
174    optimizer_ = std::make_unique<AdamW>(cfg.learning_rate, cfg.adam_beta1,
175                                          cfg.adam_beta2, cfg.adam_eps,
176                                          cfg.weight_decay);
177
178    model_->register_trainable_params(*optimizer_, cfg.learning_rate, cfg.weight_decay);
179
180    scheduler_ = std::make_unique<CosineScheduler>(cfg.learning_rate, total_steps,
181                                                     0.1f, cfg.warmup_ratio);
182}
183
184void SFTTrainer::train() {
185    SFTDataLoader loader(config.data_path);
186    if (loader.total_samples() == 0) {
187        std::cerr << "SFT训练: 无有效样本" << std::endl;
188        return;
189    }
190
191    std::cerr << "SFT训练开始: " << loader.total_samples() << " 样本, "
192              << config.epochs << " epochs" << std::endl;
193
194    model_->train();
195
196    size_t global_step = 0;
197    auto train_start = std::chrono::steady_clock::now();
198
199    for (int epoch = 0; epoch < config.epochs; ++epoch) {
200        auto epoch_start = std::chrono::steady_clock::now();
201        loader.reset();
202        std::mt19937 shuffle_rng(config.seed + epoch);
203        loader.shuffle(shuffle_rng);
204
205        float epoch_loss = 0.0f;
206        size_t step_count = 0;
207
208        while (loader.has_next()) {
209            SFTSample sample = loader.next();
210            float lr = scheduler_->get_lr(global_step);
211            optimizer_->set_lr(lr);
212            float sample_loss = train_on_sample(sample);
213            global_step++;
214
215            if (std::isfinite(sample_loss)) {
216                epoch_loss += sample_loss;
217                step_count++;
218            }
219
220            if (config.log_interval > 0 && global_step % config.log_interval == 0) {
221                auto now = std::chrono::steady_clock::now();
222                float elapsed = static_cast<float>(
223                    std::chrono::duration<double>(now - train_start).count());
224                std::cerr << "[SFT] step=" << global_step
225                          << " epoch=" << (epoch + 1)
226                          << " loss=" << sample_loss
227                          << " lr=" << optimizer_->get_lr()
228                          << " elapsed=" << elapsed << "s" << std::endl;
229            }
230
231            if (config.save_interval > 0 && global_step % config.save_interval == 0) {
232                std::string cdir = config.output_dir + "/checkpoint_step" + std::to_string(global_step);
233                std::filesystem::create_directories(cdir);
234                save_lm_checkpoint(*model_, cdir + "/lm_head.nfv1");
235                std::cerr << "[SFT] Checkpoint: step=" << global_step << std::endl;
236            }
237        }
238
239        float avg_loss = (step_count > 0) ? epoch_loss / static_cast<float>(step_count) : 0.0f;
240        auto epoch_end = std::chrono::steady_clock::now();
241        float epoch_elapsed = static_cast<float>(
242            std::chrono::duration<double>(epoch_end - epoch_start).count());
243        std::cerr << "[SFT] Epoch " << (epoch + 1) << "/" << config.epochs
244                  << " avg_loss=" << avg_loss
245                  << " elapsed=" << epoch_elapsed << "s" << std::endl;
246
247        std::string cdir = config.output_dir + "/checkpoint_epoch" + std::to_string(epoch + 1);
248        std::filesystem::create_directories(cdir);
249        save_lm_checkpoint(*model_, cdir + "/lm_head.nfv1");
250    }
251
252    std::filesystem::create_directories(config.output_dir);
253    save_lm_checkpoint(*model_, config.output_dir + "/lm_head_sft_final.nfv1");
254    std::cerr << "SFT训练完成, 模型已保存: " << config.output_dir << "/lm_head_sft_final.nfv1" << std::endl;
255}
256
257float SFTTrainer::train_on_sample(const SFTSample& sample) {
258    SFTTrainingTensors tensors = build_sft_training_tensors(sample, *tokenizer_, config.max_seq_len);
259    if (tensors.input_ids.empty() || tensors.instruction_len == 0) {
260        return 0.0f;
261    }
262
263    size_t seq_len = tensors.input_ids.size();
264    size_t vocab_size = model_->config_.vocab_size;
265
266    float total_loss = 0.0f;
267    float total_grad_norm = 0.0f;
268
269    for (size_t t = 0; t < seq_len; ++t) {
270        std::vector<size_t> input_prefix(tensors.input_ids.begin(),
271                                          tensors.input_ids.begin() + t + 1);
272
273        Tensor logits = model_->forward_for_training(input_prefix);
274
275        size_t target_id = tensors.target_ids[t];
276        if (target_id >= vocab_size) target_id = 1;
277
278        const float* pred = logits.as_fp32();
279        float max_val = -1e30f;
280        for (size_t j = 0; j < vocab_size; ++j) {
281            if (pred[j] > max_val) max_val = pred[j];
282        }
283        float sum_exp = 0.0f;
284        for (size_t j = 0; j < vocab_size; ++j) {
285            sum_exp += std::exp(pred[j] - max_val);
286        }
287        float log_sum_exp = max_val + std::log(sum_exp);
288        float token_loss = -(pred[target_id] - log_sum_exp);
289
290        if (tensors.loss_mask[t] < 0.5f) {
291            continue;
292        }
293
294        if (!std::isfinite(token_loss)) {
295            std::cerr << "[SFT WARN] NaN/Inf loss at token " << t << ", skipping" << std::endl;
296            continue;
297        }
298
299        Tensor logits_grad({1, vocab_size}, QuantType::FP32);
300        float* lg = logits_grad.as_fp32();
301        float grad_norm = 0.0f;
302        for (size_t j = 0; j < vocab_size; ++j) {
303            float softmax_val = std::exp(pred[j] - max_val) / sum_exp;
304            lg[j] = softmax_val;
305            if (j == target_id) lg[j] -= 1.0f;
306            grad_norm += lg[j] * lg[j];
307        }
308
309        float gn = std::sqrt(grad_norm);
310        float clip_scale = 1.0f;
311        if (!std::isfinite(gn) || (gn > config.grad_clip && config.grad_clip > 0.0f)) {
312            clip_scale = config.grad_clip / gn;
313        }
314        if (clip_scale < 1.0f) {
315            float* lg2 = logits_grad.as_fp32();
316            for (size_t j = 0; j < vocab_size; ++j) lg2[j] *= clip_scale;
317        }
318
319        auto lm_grads = model_->backward_from_logits(logits_grad);
320
321        model_->assign_grads_to_optimizer(*optimizer_, lm_grads);
322        optimizer_->step();
323
324        total_loss += token_loss;
325        total_grad_norm += grad_norm;
326    }
327
328    size_t valid_count = 0;
329    for (size_t i = 0; i < tensors.loss_mask.size(); ++i) {
330        if (tensors.loss_mask[i] > 0.5f) valid_count++;
331    }
332
333    return (valid_count > 0) ? total_loss / static_cast<float>(valid_count) : 0.0f;
334}
335
336} // namespace neuroflow