File size: 488 Bytes
742c169
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
"""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)