Hydra是一个强大的Python配置管理框架,特别适用于需要复杂配置和实验管理的机器学习项目。以下是一个全面的使用指南。
pip install hydra-core
# 可选:颜色输出支持
pip install hydra-colorlog
# app.py
import hydra
from omegaconf import DictConfig
@hydra.main(version_base=None, config_path="configs", config_name="config")
def my_app(cfg: DictConfig):
print(f"Database host: {cfg.database.host}")
print(f"Database port: {cfg.database.port}")
print(f"Model name: {cfg.model.name}")
if __name__ == "__main__":
my_app()
# configs/config.yaml
database:
host: localhost
port: 5432
model:
name: resnet50
learning_rate: 0.001
batch_size: 32
运行:
python app.py
configs/
├── config.yaml # 主配置
├── database/
│ ├── mysql.yaml # MySQL配置
│ └── postgres.yaml # PostgreSQL配置
└── model/
├── cnn.yaml # CNN模型配置
└── transformer.yaml # Transformer配置
# configs/config.yaml
defaults:
- database: mysql # 默认使用mysql配置
- model: cnn # 默认使用cnn配置
project:
name: "my_project"
version: 1.0
# configs/database/mysql.yaml
host: "localhost"
port: 3306
username: "root"
password: "secret"
database: "myapp"
# configs/model/cnn.yaml
name: "cnn"
layers: 5
filters: [32, 64, 128]
activation: "relu"
dropout: 0.5
# 覆盖单个参数
python app.py database.host=127.0.0.1
# 覆盖多个参数
python app.py database.host=127.0.0.1 database.port=3307 model.learning_rate=0.01
# 切换配置文件
python app.py database=postgres model=transformer
# 使用+添加配置项(不覆盖默认配置)
python app.py +optimizer.weight_decay=0.01
# 删除配置项
python app.py ~model.dropout
# 对多个参数进行网格搜索
python app.py --multirun model=cnn,transformer dataset=mnist,cifar10
# 指定范围
python app.py --multirun model.lr=0.001,0.01,0.1
# 使用区间语法
python app.py --multirun seed=1,2,3 batch_size=32,64,128
# configs/sweep.yaml
defaults:
- override /model: cnn
- override /optimizer: adam
hydra:
sweeper:
params:
model.lr: 0.001,0.01,0.1
model.dropout: 0.2,0.3,0.4
dataset.batch_size: 32,64,128
# config.yaml
user:
name: "john"
email: "${user.name}@example.com"
home_dir: "/home/${user.name}"
paths:
data: "${user.home_dir}/data"
logs: "${user.home_dir}/logs/${now:%Y-%m-%d}"
@hydra.main(version_base=None, config_path="configs", config_name="config")
def my_app(cfg: DictConfig):
# 访问插值
print(cfg.user.email) # john@example.com
# 类型转换
port = int(cfg.database.port)
# 使用OmegaConf API
from omegaconf import OmegaConf
print(OmegaConf.to_yaml(cfg))
# 更新配置
OmegaConf.update(cfg, "new.key", "value")
# config.yaml
hydra:
run:
dir: outputs/${hydra.job.name}/${now:%Y-%m-%d_%H-%M-%S}
job_logging:
formatters:
simple:
format: '[%(asctime)s][%(levelname)s] - %(message)s'
handlers:
console:
class: logging.StreamHandler
formatter: simple
stream: ext://sys.stdout
import hydra
from hydra import compose, initialize
from omegaconf import OmegaConf
# 初始化配置
initialize(config_path="conf", version_base="1.1")
cfg = compose(config_name="config")
# 或者使用配置文件初始化
@hydra.main(config_path="conf", config_name="config", version_base="1.1")
def my_app(cfg):
pass
import hydra
from omegaconf import DictConfig
from hydra.core.config_store import ConfigStore
# 定义配置类
class DatabaseConfig:
def __init__(self, host: str = "localhost", port: int = 5432):
self.host = host
self.port = port
# 注册配置
cs = ConfigStore.instance()
cs.store(name="database_config", node=DatabaseConfig)
@hydra.main(config_name="config", version_base=None)
def app(cfg: DictConfig):
print(f"Database: {cfg.host}:{cfg.port}")
# config.yaml
database:
host: ${env:DB_HOST,localhost}
port: ${env:DB_PORT,5432}
logging:
level: ${env:LOG_LEVEL,INFO}
my_project/
├── src/
│ ├── __init__.py
│ ├── train.py
│ └── utils.py
├── conf/
│ ├── config.yaml
│ ├── dataset/
│ ├── model/
│ └── optimizer/
├── outputs/ # Hydra自动生成
│ └── 2023-10-01_14-30-00/
│ ├── .hydra/
│ └── main.log
├── .gitignore
└── requirements.txt
from dataclasses import dataclass
from typing import List
from hydra.core.config_store import ConfigStore
@dataclass
class ModelConfig:
name: str
layers: int
activation: str = "relu"
dropout: float = 0.0
@dataclass
class Config:
model: ModelConfig
learning_rate: float = 0.001
batch_size: int = 32
# 注册
cs = ConfigStore.instance()
cs.store(name="base_config", node=Config)
@hydra.main(config_name="base_config", version_base=None)
def train(cfg: Config):
# 现在cfg具有类型提示
pass
确保:@hydra.main中的config_path正确,配置文件存在
使用OmegaConf.to_object(cfg)将DictConfig转换为对象
使用--multirun时注意内存使用,考虑使用作业调度器
Hydra提供了强大的配置管理功能:
对于复杂的机器学习项目,Hydra能显著提高配置管理的效率和可维护性。