Buckets:
| 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; | |
| } |
Xet Storage Details
- Size:
- 4.58 kB
- Xet hash:
- 53614a881cc550641bd602b447fb37e2f2e7fa7e395cc1dfcaa63af9e229d3f6
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.