replit/replit-code-v1_5-3b
316238
1import math2import warnings3from collections.abc import Sequence4from functools import partial5from typing import Any, Callable, Optional, Tuple, Union6import torch7from torch import nn8from .fc import FC_CLASS_REGISTRY9from .norm import NORM_CLASS_REGISTRY10try:11 import transformer_engine.pytorch as te12except:13 te = None14 15def torch_default_param_init_fn_(module: nn.Module, **kwargs: Any) -> None:16 del kwargs17 if hasattr(module, 'reset_parameters') and isinstance(module.reset_parameters, Callable):18 module.reset_parameters()19 20def fused_init_helper_(module: nn.Module, init_fn_: Callable) -> None:21 _fused = getattr(module, '_fused', None)22 if _fused is None:23 raise RuntimeError(f'Internal logic error')24 assert isinstance(module.weight, torch.Tensor)25 (dim, splits) = _fused26 splits = (0, *splits, module.weight.size(dim))27 for (s, e) in zip(splits[:-1], splits[1:]):28 slice_indices = [slice(None)] * module.weight.ndim29 slice_indices[dim] = slice(s, e)30 init_fn_(module.weight[slice_indices])31 32def generic_param_init_fn_(module: nn.Module, init_fn_: Callable, n_layers: int, d_model: Optional[int]=None, init_div_is_residual: Union[int, float, str, bool]=True, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, **kwargs: Any) -> None:33 del kwargs34 init_div_is_residual = init_div_is_residual35 if init_div_is_residual is False:36 div_is_residual = 1.037 elif init_div_is_residual is True:38 div_is_residual = math.sqrt(2 * n_layers)39 elif isinstance(init_div_is_residual, float) or isinstance(init_div_is_residual, int):40 div_is_residual = init_div_is_residual41 elif init_div_is_residual.isnumeric():42 div_is_residual = float(init_div_is_residual)43 else:44 div_is_residual = 1.045 raise ValueError(f'Expected init_div_is_residual to be boolean or numeric, got {init_div_is_residual}')46 if isinstance(module, tuple(set(FC_CLASS_REGISTRY.values()))):47 if hasattr(module, '_fused'):48 fused_init_helper_(module, init_fn_)49 else:50 init_fn_(module.weight)51 if module.bias is not None:52 assert isinstance(module.bias, torch.Tensor)53 torch.nn.init.zeros_(module.bias)54 if init_div_is_residual is not False and getattr(module, '_is_residual', False):55 with torch.no_grad():56 module.weight.div_(div_is_residual)57 elif isinstance(module, nn.Embedding):58 if emb_init_std is not None:59 std = emb_init_std60 if std == 0:61 warnings.warn(f'Embedding layer initialized to 0.')62 emb_init_fn_ = partial(torch.nn.init.normal_, mean=0.0, std=std)63 elif emb_init_uniform_lim is not None:64 lim = emb_init_uniform_lim65 if isinstance(lim, Sequence):66 if len(lim) > 2:67 raise ValueError(f'Uniform init requires a min and a max limit. User input: {lim}.')68 if lim[0] == lim[1]:69 warnings.warn(f'Embedding layer initialized to {lim[0]}.')70 else:71 if lim == 0:72 warnings.warn(f'Embedding layer initialized to 0.')73 lim = [-lim, lim]74 (a, b) = lim75 emb_init_fn_ = partial(torch.nn.init.uniform_, a=a, b=b)76 else:77 emb_init_fn_ = init_fn_78 emb_init_fn_(module.weight)79 elif isinstance(module, tuple(set(NORM_CLASS_REGISTRY.values()))):80 if hasattr(module, 'weight') and isinstance(module.weight, torch.Tensor):81 torch.nn.init.ones_(module.weight)82 if hasattr(module, 'bias') and isinstance(module.bias, torch.Tensor):83 torch.nn.init.zeros_(module.bias)84 elif isinstance(module, nn.MultiheadAttention):85 if module._qkv_same_embed_dim:86 assert module.in_proj_weight is not None87 assert module.q_proj_weight is None and module.k_proj_weight is None and (module.v_proj_weight is None)88 assert d_model is not None89 _d = d_model90 splits = (0, _d, 2 * _d, 3 * _d)91 for (s, e) in zip(splits[:-1], splits[1:]):92 init_fn_(module.in_proj_weight[s:e])93 else:94 assert module.q_proj_weight is not None and module.k_proj_weight is not None and (module.v_proj_weight is not None)95 assert module.in_proj_weight is None96 init_fn_(module.q_proj_weight)97 init_fn_(module.k_proj_weight)98 init_fn_(module.v_proj_weight)99 if module.in_proj_bias is not None:100 torch.nn.init.zeros_(module.in_proj_bias)101 if module.bias_k is not None:102 torch.nn.init.zeros_(module.bias_k)103 if module.bias_v is not None:104 torch.nn.init.zeros_(module.bias_v)105 init_fn_(module.out_proj.weight)106 if init_div_is_residual is not False and getattr(module.out_proj, '_is_residual', False):107 with torch.no_grad():108 module.out_proj.weight.div_(div_is_residual)109 if module.out_proj.bias is not None:110 torch.nn.init.zeros_(module.out_proj.bias)111 elif te is not None and isinstance(module, te.LayerNormMLP):112 if isinstance(module.layer_norm_weight, torch.Tensor):113 torch.nn.init.ones_(module.layer_norm_weight)114 if isinstance(module.layer_norm_bias, torch.Tensor):115 torch.nn.init.zeros_(module.layer_norm_bias)116 init_fn_(module.fc1_weight)117 if module.fc1_bias is not None:118 assert isinstance(module.fc1_bias, torch.Tensor)119 torch.nn.init.zeros_(module.fc1_bias)120 init_fn_(module.fc2_weight)121 if module.fc2_bias is not None:122 assert isinstance(module.fc2_bias, torch.Tensor)123 torch.nn.init.zeros_(module.fc2_bias)124 with torch.no_grad():125 module.fc2_weight.div_(div_is_residual)126 else:127 for _ in module.parameters(recurse=False):128 raise NotImplementedError(f'{module.__class__.__name__} parameters are not initialized by param_init_fn.')129 130def _normal_init_(std: float, mean: float=0.0) -> Callable:131 return partial(torch.nn.init.normal_, mean=mean, std=std)132 133def _normal_param_init_fn_(module: nn.Module, std: float, n_layers: int, d_model: Optional[int]=None, init_div_is_residual: Union[int, float, str, bool]=True, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, **kwargs: Any) -> None:134 del kwargs135 init_fn_ = _normal_init_(std=std)136 generic_param_init_fn_(module=module, init_fn_=init_fn_, d_model=d_model, n_layers=n_layers, init_div_is_residual=init_div_is_residual, emb_init_std=emb_init_std, emb_init_uniform_lim=emb_init_uniform_lim)137 138def baseline_param_init_fn_(module: nn.Module, init_std: Optional[float], n_layers: int, d_model: Optional[int]=None, init_div_is_residual: Union[int, float, str, bool]=True, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, **kwargs: Any) -> None:139 del kwargs140 if init_std is None:141 raise ValueError("You must set model.init_config['init_std'] to a float value to use the default initialization scheme.")142 _normal_param_init_fn_(module=module, std=init_std, d_model=d_model, n_layers=n_layers, init_div_is_residual=init_div_is_residual, emb_init_std=emb_init_std, emb_init_uniform_lim=emb_init_uniform_lim)143 144def small_param_init_fn_(module: nn.Module, n_layers: int, d_model: int, init_div_is_residual: Union[int, float, str, bool]=True, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, **kwargs: Any) -> None:145 del kwargs146 std = math.sqrt(2 / (5 * d_model))147 _normal_param_init_fn_(module=module, std=std, d_model=d_model, n_layers=n_layers, init_div_is_residual=init_div_is_residual, emb_init_std=emb_init_std, emb_init_uniform_lim=emb_init_uniform_lim)148 149def neox_param_init_fn_(module: nn.Module, n_layers: int, d_model: int, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, **kwargs: Any) -> None:150 """From section 2.3.1 of GPT-NeoX-20B:151 152 An Open-Source AutoregressiveLanguage Model — Black et. al. (2022)153 see https://github.com/EleutherAI/gpt-neox/blob/9610391ab319403cef079b438edd016a2443af54/megatron/model/init_functions.py#L151154 and https://github.com/EleutherAI/gpt-neox/blob/main/megatron/model/transformer.py155 """156 del kwargs157 residual_div = n_layers / math.sqrt(10)158 small_param_init_fn_(module=module, d_model=d_model, n_layers=n_layers, init_div_is_residual=residual_div, emb_init_std=emb_init_std, emb_init_uniform_lim=emb_init_uniform_lim)159 160def kaiming_uniform_param_init_fn_(module: nn.Module, n_layers: int, d_model: Optional[int]=None, init_div_is_residual: Union[int, float, str, bool]=True, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, init_gain: float=0, fan_mode: str='fan_in', init_nonlinearity: str='leaky_relu', **kwargs: Any) -> None:161 del kwargs162 kaiming_uniform_ = partial(nn.init.kaiming_uniform_, a=init_gain, mode=fan_mode, nonlinearity=init_nonlinearity)163 generic_param_init_fn_(module=module, init_fn_=kaiming_uniform_, d_model=d_model, n_layers=n_layers, init_div_is_residual=init_div_is_residual, emb_init_std=emb_init_std, emb_init_uniform_lim=emb_init_uniform_lim)164 165def kaiming_normal_param_init_fn_(module: nn.Module, n_layers: int, d_model: Optional[int]=None, init_div_is_residual: Union[int, float, str, bool]=True, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, init_gain: float=0, fan_mode: str='fan_in', init_nonlinearity: str='leaky_relu', **kwargs: Any) -> None:166 del kwargs167 kaiming_normal_ = partial(torch.nn.init.kaiming_normal_, a=init_gain, mode=fan_mode, nonlinearity=init_nonlinearity)168 generic_param_init_fn_(module=module, init_fn_=kaiming_normal_, d_model=d_model, n_layers=n_layers, init_div_is_residual=init_div_is_residual, emb_init_std=emb_init_std, emb_init_uniform_lim=emb_init_uniform_lim)169 170def xavier_uniform_param_init_fn_(module: nn.Module, n_layers: int, d_model: Optional[int]=None, init_div_is_residual: Union[int, float, str, bool]=True, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, init_gain: float=0, **kwargs: Any) -> None:171 del kwargs172 xavier_uniform_ = partial(torch.nn.init.xavier_uniform_, gain=init_gain)173 generic_param_init_fn_(module=module, init_fn_=xavier_uniform_, d_model=d_model, n_layers=n_layers, init_div_is_residual=init_div_is_residual, emb_init_std=emb_init_std, emb_init_uniform_lim=emb_init_uniform_lim)174 175def xavier_normal_param_init_fn_(module: nn.Module, n_layers: int, d_model: Optional[int]=None, init_div_is_residual: Union[int, float, str, bool]=True, emb_init_std: Optional[float]=None, emb_init_uniform_lim: Optional[Union[Tuple[float, float], float]]=None, init_gain: float=0, **kwargs: Any) -> None:176 del kwargs177 xavier_normal_ = partial(torch.nn.init.xavier_normal_, gain=init_gain)178 generic_param_init_fn_(module=module, init_fn_=xavier_normal_, d_model=d_model, n_layers=n_layers, init_div_is_residual=init_div_is_residual, emb_init_std=emb_init_std, emb_init_uniform_lim=emb_init_uniform_lim)179MODEL_INIT_REGISTRY = {'default_': torch_default_param_init_fn_, 'baseline_': baseline_param_init_fn_, 'kaiming_uniform_': kaiming_uniform_param_init_fn_, 'kaiming_normal_': kaiming_normal_param_init_fn_, 'neox_init_': neox_param_init_fn_, 'small_init_': small_param_init_fn_, 'xavier_uniform_': xavier_uniform_param_init_fn_, 'xavier_normal_': xavier_normal_param_init_fn_}