from __future__ import annotations import math from torch import nn from .model_config import AdditionModelConfig def initialize_module(module: nn.Module, config: AdditionModelConfig) -> None: init_mode = str(config.init_mode) if init_mode == "normal": _initialize_normal(module) elif init_mode == "orthogonal": _initialize_orthogonal(module) else: raise ValueError(f"Unsupported init mode: {init_mode}") def _initialize_normal(module: nn.Module) -> None: for child in module.modules(): if isinstance(child, nn.Embedding): nn.init.normal_(child.weight, mean=0.0, std=1.0 / math.sqrt(2.0)) elif isinstance(child, nn.Linear): nn.init.normal_(child.weight, mean=0.0, std=1.0 / math.sqrt(child.in_features)) _zero_bias(child) def _initialize_orthogonal(module: nn.Module) -> None: for child in module.modules(): if isinstance(child, nn.Embedding): nn.init.normal_(child.weight, mean=0.0, std=1.0 / math.sqrt(2.0)) elif isinstance(child, nn.Linear): gain = math.sqrt(child.out_features / child.in_features) if child.out_features > child.in_features else 1.0 nn.init.orthogonal_(child.weight, gain=gain) _zero_bias(child) def _zero_bias(module: nn.Linear) -> None: if module.bias is not None: nn.init.zeros_(module.bias)