File size: 2,357 Bytes
c1a46f7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
训练入口脚本

使用方式:
    # 从头训练 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

# 将 src 目录加入 Python 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)

    # TODO: 实现训练主流程
    raise NotImplementedError(
        "TODO: 整合阶段实现训练主流程\n"
        "请在所有模块完成后,协作完成此脚本"
    )


if __name__ == "__main__":
    main()