File size: 4,578 Bytes
26d5b81 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 | #ifndef NOMINMAX
#define NOMINMAX
#endif
#include <cstdlib>
#include <filesystem>
#include <iostream>
#include <string>
#include <vector>
#include "neuroflow/sft.hpp"
using neuroflow::SFTTrainConfig;
using neuroflow::SFTTrainer;
using neuroflow::validate_path;
static void print_sft_usage() {
std::cerr << "用法: neuroflow_sft --data_path <path> --ckpt_path <path> --tokenizer_path <path> --output_dir <dir>"
<< " [--lr <float>] [--epochs <N>] [--max_seq_len <N>] [--warmup_ratio <float>]"
<< " [--weight_decay <float>] [--grad_clip <float>] [--save_interval <N>]"
<< " [--log_interval <N>] [--seed <N>]" << std::endl;
}
static SFTTrainConfig parse_sft_args(int argc, char* argv[]) {
SFTTrainConfig cfg;
for (int i = 1; i < argc; ++i) {
std::string arg = argv[i];
if (arg == "--data_path" && i + 1 < argc) { cfg.data_path = argv[++i]; }
else if (arg == "--ckpt_path" && i + 1 < argc) { cfg.ckpt_path = argv[++i]; }
else if (arg == "--tokenizer_path" && i + 1 < argc) { cfg.tokenizer_path = argv[++i]; }
else if (arg == "--output_dir" && i + 1 < argc) { cfg.output_dir = argv[++i]; }
else if (arg == "--lr" && i + 1 < argc) { cfg.learning_rate = std::stof(argv[++i]); }
else if (arg == "--epochs" && i + 1 < argc) { cfg.epochs = std::stoi(argv[++i]); }
else if (arg == "--max_seq_len" && i + 1 < argc) { cfg.max_seq_len = std::stoul(argv[++i]); }
else if (arg == "--warmup_ratio" && i + 1 < argc) { cfg.warmup_ratio = std::stof(argv[++i]); }
else if (arg == "--weight_decay" && i + 1 < argc) { cfg.weight_decay = std::stof(argv[++i]); }
else if (arg == "--grad_clip" && i + 1 < argc) { cfg.grad_clip = std::stof(argv[++i]); }
else if (arg == "--adam_beta1" && i + 1 < argc) { cfg.adam_beta1 = std::stof(argv[++i]); }
else if (arg == "--adam_beta2" && i + 1 < argc) { cfg.adam_beta2 = std::stof(argv[++i]); }
else if (arg == "--save_interval" && i + 1 < argc) { cfg.save_interval = std::stoul(argv[++i]); }
else if (arg == "--log_interval" && i + 1 < argc) { cfg.log_interval = std::stoul(argv[++i]); }
else if (arg == "--seed" && i + 1 < argc) { cfg.seed = static_cast<unsigned>(std::stoul(argv[++i])); }
else if (arg == "--help" || arg == "-h") { print_sft_usage(); std::exit(0); }
else { std::cerr << "未知参数: " << arg << std::endl; }
}
return cfg;
}
int main(int argc, char* argv[]) {
setvbuf(stderr, nullptr, _IONBF, 0);
setvbuf(stdout, nullptr, _IONBF, 0);
std::ios::sync_with_stdio(false);
SFTTrainConfig cfg = parse_sft_args(argc, argv);
if (cfg.data_path.empty() || cfg.ckpt_path.empty() ||
cfg.tokenizer_path.empty() || cfg.output_dir.empty()) {
std::cerr << "错误: 缺少必填参数" << std::endl;
print_sft_usage();
return 1;
}
if (!validate_path(cfg.data_path) || !validate_path(cfg.ckpt_path) ||
!validate_path(cfg.tokenizer_path) || !validate_path(cfg.output_dir)) {
std::cerr << "错误: 路径包含非法字符(..)" << std::endl;
return 1;
}
if (!std::filesystem::exists(cfg.data_path)) {
std::cerr << "错误: 数据文件不存在: " << cfg.data_path << std::endl;
return 1;
}
if (!std::filesystem::exists(cfg.ckpt_path)) {
std::cerr << "错误: Checkpoint文件不存在: " << cfg.ckpt_path << std::endl;
return 1;
}
if (!std::filesystem::exists(cfg.tokenizer_path)) {
std::cerr << "错误: Tokenizer文件不存在: " << cfg.tokenizer_path << std::endl;
return 1;
}
std::cerr << "======================================" << std::endl;
std::cerr << "NeuroFlow SFT Training" << std::endl;
std::cerr << "======================================" << std::endl;
std::cerr << "数据: " << cfg.data_path << std::endl;
std::cerr << "Checkpoint: " << cfg.ckpt_path << std::endl;
std::cerr << "Tokenizer: " << cfg.tokenizer_path << std::endl;
std::cerr << "输出: " << cfg.output_dir << std::endl;
std::cerr << "Epochs: " << cfg.epochs << " LR: " << cfg.learning_rate
<< " MaxSeqLen: " << cfg.max_seq_len << std::endl;
try {
SFTTrainer trainer(cfg);
trainer.train();
} catch (const std::exception& e) {
std::cerr << "SFT训练异常: " << e.what() << std::endl;
return 1;
}
return 0;
} |