File size: 1,342 Bytes
d8bfe4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# train.py
import os
import importlib
import hydra
from omegaconf import DictConfig, OmegaConf
from core.tasks.task_registry import get_task


def auto_import_tasks():
    tasks_dir = os.path.join(os.path.dirname(__file__), "../core", "tasks")
    for filename in os.listdir(tasks_dir):
        if filename.endswith(".py") and filename not in (
            "base_task.py",
            "task_registry.py",
        ):
            importlib.import_module(f"core.tasks.{filename[:-3]}")


@hydra.main(
    version_base=None, config_path=None, config_name=None  # 配置路径由 CLI 指定
)
def main(cfg: DictConfig):
    """
    cfg: Hydra 合并后的 DictConfig
    """

    # ====== 调试用:第一次强烈建议保留 ======
    print("========== Final Config ==========")
    print(OmegaConf.to_yaml(cfg))
    print("==================================")

    # ====== 自动导入 tasks,触发注册 ======
    auto_import_tasks()

    # ====== 获取 task_name ======
    if "task_name" not in cfg:
        raise ValueError("配置中必须包含 task_name")

    task_name = cfg.task_name

    # ====== 构建并运行 Task ======
    TaskClass = get_task(task_name)
    if not TaskClass:
        raise ValueError(f"未注册的任务: {task_name}")

    task = TaskClass(cfg)
    task.run()

if __name__ == "__main__":
    main()