lijn14
创建工程
c1a46f7
Raw
History Blame
2.36 kB
"""
训练入口脚本
使用方式:
# 从头训练 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()