File size: 755 Bytes
41ff959
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
from dataclasses import dataclass
from typing import Type, TypeVar

from dacite import from_dict
from omegaconf import DictConfig, OmegaConf

from src.model.decoder import DecoderCfg
from src.model.encoder import EncoderCfg


@dataclass
class ModelCfg:
    decoder: DecoderCfg
    encoder: EncoderCfg


@dataclass
class RootCfg:
    model: ModelCfg


T = TypeVar("T")


def load_typed_config(cfg: DictConfig, data_class: Type[T]) -> T:
    """Convert one resolved Hydra config into its inference dataclass."""
    return from_dict(data_class, OmegaConf.to_container(cfg, resolve=True))


def load_typed_root_config(cfg: DictConfig) -> RootCfg:
    """Load the typed root configuration used by demo inference."""
    return load_typed_config(cfg, RootCfg)