| #include "neuroflow/dpo.hpp"
|
|
|
| #include <algorithm>
|
| #include <chrono>
|
| #include <cmath>
|
| #include <cstring>
|
| #include <filesystem>
|
| #include <fstream>
|
| #include <iostream>
|
| #include <numeric>
|
| #include <sstream>
|
|
|
| namespace neuroflow {
|
|
|
| DPODataLoader::DPODataLoader(const std::string& jsonl_path, size_t max_samples) {
|
| std::ifstream ifs(jsonl_path);
|
| if (!ifs) {
|
| std::cerr << "DPO数据文件无法打开: " << jsonl_path << std::endl;
|
| return;
|
| }
|
| std::string line;
|
| while (std::getline(ifs, line)) {
|
| if (line.empty() || line[0] == '#') continue;
|
| std::string instruction = extract_json_string(line, "instruction");
|
| std::string chosen = extract_json_string(line, "chosen");
|
| std::string rejected = extract_json_string(line, "rejected");
|
| unescape_json(instruction);
|
| unescape_json(chosen);
|
| unescape_json(rejected);
|
| if (instruction.empty() || chosen.empty() || rejected.empty()) {
|
| invalid_count_++;
|
| continue;
|
| }
|
| if (chosen == rejected) {
|
| invalid_count_++;
|
| continue;
|
| }
|
| samples_.push_back({instruction, chosen, rejected});
|
| if (max_samples > 0 && samples_.size() >= max_samples) break;
|
| }
|
| std::cerr << "DPO数据加载: " << samples_.size() << " 样本, "
|
| << invalid_count_ << " 无效" << std::endl;
|
| }
|
|
|
| bool DPODataLoader::has_next() const {
|
| return cursor_ < samples_.size();
|
| }
|
|
|
| DPOSample DPODataLoader::next() {
|
| return samples_[cursor_++];
|
| }
|
|
|
| void DPODataLoader::reset() {
|
| cursor_ = 0;
|
| }
|
|
|
| void DPODataLoader::shuffle(std::mt19937& rng) {
|
| std::shuffle(samples_.begin(), samples_.end(), rng);
|
| }
|
|
|
| float compute_log_prob(CausalLMHead& model, const std::vector<size_t>& token_ids,
|
| size_t prompt_len, size_t vocab_size) {
|
| if (token_ids.size() < 2) return 0.0f;
|
|
|
| float total_log_prob = 0.0f;
|
| size_t valid_count = 0;
|
|
|
| for (size_t t = prompt_len; t < token_ids.size(); ++t) {
|
| std::vector<size_t> input_prefix(token_ids.begin(), token_ids.begin() + t);
|
| size_t target_id = token_ids[t];
|
| if (target_id >= vocab_size) target_id = 1;
|
|
|
| Tensor logits = model.forward(input_prefix);
|
|
|
| const float* pred = logits.as_fp32();
|
| float max_val = -1e30f;
|
| for (size_t j = 0; j < vocab_size; ++j) {
|
| if (pred[j] > max_val) max_val = pred[j];
|
| }
|
| float sum_exp = 0.0f;
|
| for (size_t j = 0; j < vocab_size; ++j) {
|
| sum_exp += std::exp(pred[j] - max_val);
|
| }
|
| float log_sum_exp = max_val + std::log(sum_exp);
|
| float lp = pred[target_id] - log_sum_exp;
|
|
|
| if (std::isfinite(lp)) {
|
| total_log_prob += lp;
|
| valid_count++;
|
| }
|
| }
|
|
|
| return (valid_count > 0) ? total_log_prob : 0.0f;
|
| }
|
|
|
| DPOLossOutput compute_dpo_loss(float log_prob_chosen_policy,
|
| float log_prob_rejected_policy,
|
| float log_prob_chosen_ref,
|
| float log_prob_rejected_ref,
|
| float beta) {
|
| DPOLossOutput output;
|
|
|
| float reward_chosen = beta * (log_prob_chosen_policy - log_prob_chosen_ref);
|
| float reward_rejected = beta * (log_prob_rejected_policy - log_prob_rejected_ref);
|
|
|
| output.reward_chosen = reward_chosen;
|
| output.reward_rejected = reward_rejected;
|
|
|
| float diff = reward_chosen - reward_rejected;
|
|
|
| float sigmoid_val;
|
| if (diff > 20.0f) {
|
| sigmoid_val = 1.0f;
|
| } else if (diff < -20.0f) {
|
| sigmoid_val = 0.0f;
|
| } else {
|
| sigmoid_val = 1.0f / (1.0f + std::exp(-diff));
|
| }
|
|
|
| output.loss = -std::log(sigmoid_val + 1e-10f);
|
| output.alpha = sigmoid_val;
|
|
|
| return output;
|
| }
|
|
|
| DPOTrainer::DPOTrainer(const DPOTrainConfig& cfg) : config(cfg) {
|
| CausalLMConfig lm_config;
|
| lm_config.vocab_size = 128000;
|
| lm_config.d_model = 512;
|
| lm_config.max_seq_len = cfg.max_seq_len;
|
| lm_config.num_attn_layers = 4;
|
| lm_config.num_attn_heads = 8;
|
| lm_config.n_kv_heads = 2;
|
| lm_config.use_rope = true;
|
| lm_config.use_qk_norm = true;
|
| lm_config.use_swiglu = true;
|
| lm_config.use_bridge = true;
|
| lm_config.weight_tying = true;
|
| lm_config.pooling = "last";
|
|
|
| policy_ = std::make_unique<CausalLMHead>(lm_config);
|
| if (!cfg.sft_ckpt_path.empty()) {
|
| load_lm_checkpoint(*policy_, cfg.sft_ckpt_path);
|
| }
|
|
|
| reference_ = std::make_unique<CausalLMHead>(lm_config);
|
| if (!cfg.sft_ckpt_path.empty()) {
|
| load_lm_checkpoint(*reference_, cfg.sft_ckpt_path);
|
| }
|
| reference_->eval();
|
|
|
| tokenizer_ = std::make_unique<BPETokenizer>(cfg.tokenizer_path);
|
|
|
| size_t total_steps = 0;
|
| {
|
| DPODataLoader tmp_loader(cfg.data_path);
|
| total_steps = tmp_loader.total_samples() * cfg.epochs;
|
| }
|
|
|
| optimizer_ = std::make_unique<AdamW>(cfg.learning_rate, cfg.adam_beta1,
|
| cfg.adam_beta2, cfg.adam_eps,
|
| cfg.weight_decay);
|
|
|
| policy_->register_trainable_params(*optimizer_, cfg.learning_rate, cfg.weight_decay);
|
|
|
| scheduler_ = std::make_unique<CosineScheduler>(cfg.learning_rate, total_steps,
|
| 0.1f, cfg.warmup_ratio);
|
| }
|
|
|
| float DPOTrainer::compute_w_embed_checksum() {
|
| if (!reference_ || reference_->w_embed_.numel() == 0) return 0.0f;
|
| const float* data = reference_->w_embed_.as_fp32();
|
| float sum = 0.0f;
|
| size_t n = std::min(reference_->w_embed_.numel(), static_cast<size_t>(1000));
|
| for (size_t i = 0; i < n; ++i) {
|
| sum += data[i];
|
| }
|
| return sum;
|
| }
|
|
|
| void DPOTrainer::train() {
|
| DPODataLoader loader(config.data_path);
|
| if (loader.total_samples() == 0) {
|
| std::cerr << "DPO训练: 无有效样本" << std::endl;
|
| return;
|
| }
|
|
|
| std::cerr << "DPO训练开始: " << loader.total_samples() << " 样本, "
|
| << config.epochs << " epochs, beta=" << config.beta << std::endl;
|
|
|
| policy_->train();
|
| reference_->eval();
|
|
|
| float ref_checksum = compute_w_embed_checksum();
|
|
|
| size_t global_step = 0;
|
| auto train_start = std::chrono::steady_clock::now();
|
|
|
| for (int epoch = 0; epoch < config.epochs; ++epoch) {
|
| auto epoch_start = std::chrono::steady_clock::now();
|
| loader.reset();
|
| std::mt19937 shuffle_rng(config.seed + epoch);
|
| loader.shuffle(shuffle_rng);
|
|
|
| float epoch_loss = 0.0f;
|
| size_t step_count = 0;
|
|
|
| while (loader.has_next()) {
|
| DPOSample sample = loader.next();
|
| float lr = scheduler_->get_lr(global_step);
|
| optimizer_->set_lr(lr);
|
| float sample_loss = train_on_sample(sample);
|
| global_step++;
|
|
|
| if (std::isfinite(sample_loss)) {
|
| epoch_loss += sample_loss;
|
| step_count++;
|
| }
|
|
|
| if (config.log_interval > 0 && global_step % config.log_interval == 0) {
|
| auto now = std::chrono::steady_clock::now();
|
| float elapsed = static_cast<float>(
|
| std::chrono::duration<double>(now - train_start).count());
|
| std::cerr << "[DPO] step=" << global_step
|
| << " epoch=" << (epoch + 1)
|
| << " loss=" << sample_loss
|
| << " lr=" << optimizer_->get_lr()
|
| << " elapsed=" << elapsed << "s" << std::endl;
|
| }
|
|
|
| if (config.save_interval > 0 && global_step % config.save_interval == 0) {
|
| std::string cdir = config.output_dir + "/checkpoint_step" + std::to_string(global_step);
|
| std::filesystem::create_directories(cdir);
|
| save_lm_checkpoint(*policy_, cdir + "/lm_head.nfv1");
|
| std::cerr << "[DPO] Checkpoint: step=" << global_step << std::endl;
|
| }
|
| }
|
|
|
| float current_checksum = compute_w_embed_checksum();
|
| if (std::abs(current_checksum - ref_checksum) > 1e-3f) {
|
| std::cerr << "[DPO 警告] 参考模型权重校验和不匹配! "
|
| << "预期=" << ref_checksum << " 实际=" << current_checksum
|
| << " (参考模型可能被意外修改)" << std::endl;
|
| }
|
|
|
| float avg_loss = (step_count > 0) ? epoch_loss / static_cast<float>(step_count) : 0.0f;
|
| auto epoch_end = std::chrono::steady_clock::now();
|
| float epoch_elapsed = static_cast<float>(
|
| std::chrono::duration<double>(epoch_end - epoch_start).count());
|
| std::cerr << "[DPO] Epoch " << (epoch + 1) << "/" << config.epochs
|
| << " avg_loss=" << avg_loss
|
| << " elapsed=" << epoch_elapsed << "s" << std::endl;
|
|
|
| std::string cdir = config.output_dir + "/checkpoint_epoch" + std::to_string(epoch + 1);
|
| std::filesystem::create_directories(cdir);
|
| save_lm_checkpoint(*policy_, cdir + "/lm_head.nfv1");
|
| }
|
|
|
| std::filesystem::create_directories(config.output_dir);
|
| save_lm_checkpoint(*policy_, config.output_dir + "/lm_head_dpo_final.nfv1");
|
| std::cerr << "DPO训练完成, 模型已保存: " << config.output_dir << "/lm_head_dpo_final.nfv1" << std::endl;
|
| }
|
|
|
| float DPOTrainer::train_on_sample(const DPOSample& sample) {
|
| std::string prompt = sample.instruction + "\n";
|
| std::string chosen_text = prompt + sample.chosen;
|
| std::string rejected_text = prompt + sample.rejected;
|
|
|
| std::vector<size_t> chosen_ids = tokenizer_->encode(chosen_text, config.max_seq_len);
|
| std::vector<size_t> rejected_ids = tokenizer_->encode(rejected_text, config.max_seq_len);
|
| std::vector<size_t> prompt_ids = tokenizer_->encode(prompt, config.max_seq_len);
|
|
|
| if (chosen_ids.size() < 2 || rejected_ids.size() < 2) return 0.0f;
|
|
|
| size_t prompt_len = std::min(prompt_ids.size(),
|
| std::min(chosen_ids.size(), rejected_ids.size()));
|
| size_t vocab_size = policy_->config_.vocab_size;
|
|
|
| reference_->eval();
|
|
|
| float log_prob_chosen_ref = compute_log_prob(*reference_, chosen_ids, prompt_len, vocab_size);
|
| float log_prob_rejected_ref = compute_log_prob(*reference_, rejected_ids, prompt_len, vocab_size);
|
|
|
| policy_->train();
|
| float log_prob_chosen_policy = compute_log_prob(*policy_, chosen_ids, prompt_len, vocab_size);
|
| float log_prob_rejected_policy = compute_log_prob(*policy_, rejected_ids, prompt_len, vocab_size);
|
|
|
| DPOLossOutput dpo_out = compute_dpo_loss(log_prob_chosen_policy, log_prob_rejected_policy,
|
| log_prob_chosen_ref, log_prob_rejected_ref,
|
| config.beta);
|
|
|
| if (!std::isfinite(dpo_out.loss)) {
|
| std::cerr << "[DPO WARN] NaN/Inf DPO loss, skipping" << std::endl;
|
| return 0.0f;
|
| }
|
|
|
| float alpha = dpo_out.alpha;
|
|
|
| auto compute_and_apply_grad = [&](const std::vector<size_t>& ids, size_t p_len, float scale) {
|
| for (size_t t = p_len; t < ids.size(); ++t) {
|
| std::vector<size_t> input_prefix(ids.begin(), ids.begin() + t);
|
| size_t target_id = ids[t];
|
| if (target_id >= vocab_size) target_id = 1;
|
|
|
| Tensor logits = policy_->forward_for_training(input_prefix);
|
|
|
| const float* pred = logits.as_fp32();
|
| float max_val = -1e30f;
|
| for (size_t j = 0; j < vocab_size; ++j) {
|
| if (pred[j] > max_val) max_val = pred[j];
|
| }
|
| float sum_exp = 0.0f;
|
| for (size_t j = 0; j < vocab_size; ++j) {
|
| sum_exp += std::exp(pred[j] - max_val);
|
| }
|
|
|
| Tensor logits_grad({1, vocab_size}, QuantType::FP32);
|
| float* lg = logits_grad.as_fp32();
|
| float grad_norm = 0.0f;
|
| for (size_t j = 0; j < vocab_size; ++j) {
|
| float softmax_val = std::exp(pred[j] - max_val) / sum_exp;
|
| lg[j] = softmax_val;
|
| if (j == target_id) lg[j] -= 1.0f;
|
| lg[j] *= scale;
|
| grad_norm += lg[j] * lg[j];
|
| }
|
|
|
| float gn = std::sqrt(grad_norm);
|
| float clip_scale = 1.0f;
|
| if (!std::isfinite(gn) || (gn > config.grad_clip && config.grad_clip > 0.0f)) {
|
| clip_scale = config.grad_clip / gn;
|
| }
|
| if (clip_scale < 1.0f) {
|
| float* lg2 = logits_grad.as_fp32();
|
| for (size_t j = 0; j < vocab_size; ++j) lg2[j] *= clip_scale;
|
| }
|
|
|
| auto lm_grads = policy_->backward_from_logits(logits_grad);
|
|
|
| policy_->assign_grads_to_optimizer(*optimizer_, lm_grads);
|
| optimizer_->step();
|
| }
|
| };
|
|
|
| float grad_scale = config.beta * (1.0f - alpha);
|
| compute_and_apply_grad(chosen_ids, prompt_len, -grad_scale);
|
| compute_and_apply_grad(rejected_ids, prompt_len, grad_scale);
|
|
|
| return dpo_out.loss;
|
| }
|
|
|
| } |