| """
|
| 训练入口脚本
|
|
|
| 使用方式:
|
| # 从头训练 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()
|
|
|