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