"""muP initialisation (same as ``optimfactory.mup_init`` / ``mup_init_output``).""" import math import torch def mup_init(params, is_output: bool = False) -> None: for param in params: if param.ndim == 1: continue fan_in = math.prod(param.shape[1:]) std = (1 / fan_in) ** (1 if is_output else 0.5) torch.nn.init.normal_(param, mean=0.0, std=std) def mup_init_output(param: torch.Tensor) -> None: mup_init([param], is_output=True)