TTVidT / mup.py
KBlueLeaf's picture
TT-VidT TT3D encoder (structured resample), transformers remote code
742c169 verified
Raw History Blame Contribute Delete
488 Bytes
"""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)