zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
1.34 kB
# 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()