| # 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]}") | |
| 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() | |