# 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()