1-layer-addition / initialization.py
melephant's picture
Publish addition-transformer run s85nnxtf
58223a8 verified
Raw
History Blame Contribute Delete
1.4 kB
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)