| """ |
| 训练入口脚本 |
| |
| 使用方式: |
| # 从头训练 Transformer |
| python scripts/train.py --config configs/default_config.yaml |
| |
| # 微调 NLLB |
| python scripts/train.py --config configs/default_config.yaml model.type=finetune_nllb |
| |
| # 命令行覆盖参数 |
| python scripts/train.py --config configs/default_config.yaml training.optimizer.lr=1e-4 |
| |
| # 分布式训练 |
| torchrun --nproc_per_node=4 scripts/train.py --config configs/default_config.yaml |
| """ |
|
|
| import argparse |
| import sys |
| from pathlib import Path |
|
|
| |
| sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src")) |
|
|
|
|
| def parse_args(): |
| parser = argparse.ArgumentParser(description="EasyTranslate Training") |
| parser.add_argument("--config", type=str, default="configs/default_config.yaml", help="配置文件路径") |
| parser.add_argument("--resume", type=str, default=None, help="检查点路径 (断点续训)") |
| args, unknown = parser.parse_known_args() |
| return args, unknown |
|
|
|
|
| def main(): |
| """ |
| 训练主流程。 |
| |
| TODO [整合阶段 - 所有人协作]: |
| 1. 解析命令行参数 |
| 2. 加载配置 (config_from_cli) |
| 3. set_seed |
| 4. setup_logging |
| 5. 根据 model.type 选择训练模式: |
| |
| A) model.type == "transformer_scratch": |
| a. 加载数据集 (load_wmt_dataset / load_opus_dataset) |
| b. 预处理 (preprocess_pipeline) |
| c. 训练分词器 (train_bpe_tokenizer) 或加载已有分词器 |
| d. 构建 TranslationDataset + DataLoader |
| e. 构建 TransformerTranslationModel |
| f. 构建 Trainer 并开始训练 |
| |
| B) model.type == "finetune_nllb" / "finetune_mbart": |
| a. 加载预训练模型和分词器 (load_pretrained_model) |
| b. 配置 LoRA (setup_lora) |
| c. 加载数据集,使用预训练分词器处理 |
| d. 构建 Trainer 并开始训练 |
| |
| 6. 训练完成后保存最终模型 |
| 7. 在测试集上评估 |
| """ |
| args, cli_overrides = parse_args() |
|
|
| print("=" * 60) |
| print(" EasyTranslate - English to Chinese Translation") |
| print("=" * 60) |
|
|
| |
| raise NotImplementedError( |
| "TODO: 整合阶段实现训练主流程\n" |
| "请在所有模块完成后,协作完成此脚本" |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|