Hydra框架中如何从代码或配置内获取所用模型配置文件名
实现方案
你可以根据自己的使用场景选择以下任意一种方案实现:
方案1:手动给每个模型配置加标识(最简单,无需额外开发)
直接在每个模型配置文件里新增名称字段即可,维护成本极低:model_a.yaml示例:
dropout: true dense_layers: 128 model_name: model_a
后续你既可以在其他配置里通过${model.model_name}引用,也可以在Python代码里直接通过config.model.model_name取值,默认值和命令行覆盖的取值都可以正常获取。
方案2:自定义插值器实现配置内自动获取(无需每个配置手动加字段)
如果不想手动给每个模型配置重复加字段,可以通过注册OmegaConf自定义插值器实现${__filename__}的效果:
步骤1:提前注册自定义插值函数
在Hydra初始化前(主函数执行前)运行以下代码:
from omegaconf import OmegaConf from hydra.core.hydra_config import HydraConfig def get_model_filename(_root_): # 优先读取命令行覆盖的模型配置 overrides = HydraConfig.get().overrides.task model_override = next((o for o in overrides if o.startswith("model=")), None) if model_override: return model_override.split("=")[1] # 无覆盖时取默认配置值 return next(d["model"] for d in _root_.defaults if isinstance(d, dict) and "model" in d) OmegaConf.register_new_resolver("__filename__", get_model_filename)
步骤2:在模型配置中直接使用
之后你就可以在任意模型配置里直接写:
dropout: true dense_layers: 128 model_name: ${__filename__:}
方案3:直接在Python代码中读取运行时元数据(无需修改任何配置文件)
如果不需要在配置文件里引用,仅需要在Python代码中获取模型文件名,可以直接读取Hydra的运行时配置:
import hydra from omegaconf import DictConfig from hydra.core.hydra_config import HydraConfig @hydra.main(version_base=None, config_path=".", config_name="config") def main(config: DictConfig): # 优先取命令行覆盖值 overrides = HydraConfig.get().overrides.task model_override = next((o for o in overrides if o.startswith("model=")), None) if model_override: model_name = model_override.split("=")[1] else: # 无覆盖时取默认配置值 model_name = next(d["model"] for d in config.defaults if isinstance(d, dict) and "model" in d) # 后续直接使用model_name即可 print(model_name)
内容的提问来源于stack exchange,提问作者miccio
相关产品推荐
相关产品推荐

