from __future__ import annotations from dataclasses import asdict, dataclass @dataclass(frozen=True) class ATIConfig: name: str low_rank: bool = False cross_attention: bool = False anchored: bool = True lambda_interaction: float = 1e-3 lambda_mask: float = 0.0 rank: int = 4 hidden: int = 64 gru_hidden_per_direction: int = 32 attention_heads: int = 4 attention_ffn: int = 128 eta_init: float = 0.1 def to_dict(self) -> dict[str, object]: return asdict(self) CONFIGS: dict[str, ATIConfig] = { "A0": ATIConfig(name="A0_main_effects"), "A1": ATIConfig(name="A1_low_rank_pairs", low_rank=True), "A2": ATIConfig(name="A2_anchored_pairwise", low_rank=True, cross_attention=True), "A3": ATIConfig( name="A3_pairwise_mask_aux", low_rank=True, cross_attention=True, lambda_mask=0.05 ), "D0": ATIConfig( name="D0_unanchored_diagnostic", low_rank=True, cross_attention=True, anchored=False ), }