#include "neuroflow/dpo.hpp" #include #include #include #include #include #include #include #include #include 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& 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 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(lm_config); if (!cfg.sft_ckpt_path.empty()) { load_lm_checkpoint(*policy_, cfg.sft_ckpt_path); } reference_ = std::make_unique(lm_config); if (!cfg.sft_ckpt_path.empty()) { load_lm_checkpoint(*reference_, cfg.sft_ckpt_path); } reference_->eval(); tokenizer_ = std::make_unique(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(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(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(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( std::chrono::duration(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(step_count) : 0.0f; auto epoch_end = std::chrono::steady_clock::now(); float epoch_elapsed = static_cast( std::chrono::duration(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 chosen_ids = tokenizer_->encode(chosen_text, config.max_seq_len); std::vector rejected_ids = tokenizer_->encode(rejected_text, config.max_seq_len); std::vector 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& ids, size_t p_len, float scale) { for (size_t t = p_len; t < ids.size(); ++t) { std::vector 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; } } // namespace neuroflow