MustaqiLLM / layers.py
kmamaroziqov's picture
MilliyLM-5B: instruction-tuned Uzbek chat model (SFT of NeuronAI-5B-Base)
80c3430 verified
Raw
History Blame Contribute Delete
2.26 kB
from __future__ import annotations
import torch
import torch.nn.functional as F
from torch import Tensor, nn
__all__ = [
"RMSNorm",
"SwiGLU",
]
class RMSNorm(nn.RMSNorm):
def __init__(
self,
hidden_size: int,
eps: float = 1e-5,
*,
device: torch.device | str | None = None,
dtype: torch.dtype | None = None,
) -> None:
if hidden_size <= 0:
raise ValueError(f"hidden_size must be positive, got {hidden_size}")
if eps <= 0.0:
raise ValueError(f"eps must be positive, got {eps}")
super().__init__(
normalized_shape=hidden_size,
eps=eps,
elementwise_affine=True,
device=device,
dtype=dtype,
)
self.hidden_size = hidden_size
class SwiGLU(nn.Module):
def __init__(
self,
hidden_size: int,
intermediate_size: int,
*,
device: torch.device | str | None = None,
dtype: torch.dtype | None = None,
) -> None:
super().__init__()
if hidden_size <= 0:
raise ValueError(f"hidden_size must be positive, got {hidden_size}")
if intermediate_size <= 0:
raise ValueError(
f"intermediate_size must be positive, got {intermediate_size}"
)
self.hidden_size = hidden_size
self.intermediate_size = intermediate_size
self.gate_up_proj = nn.Linear(
in_features=hidden_size,
out_features=2 * intermediate_size,
bias=False,
device=device,
dtype=dtype,
)
self.down_proj = nn.Linear(
in_features=intermediate_size,
out_features=hidden_size,
bias=False,
device=device,
dtype=dtype,
)
def forward(self, hidden_states: Tensor) -> Tensor:
gate, up = self.gate_up_proj(hidden_states).chunk(2, dim=-1)
hidden_states = F.silu(gate) * up
return self.down_proj(hidden_states)
def extra_repr(self) -> str:
return (
f"hidden_size={self.hidden_size}, "
f"intermediate_size={self.intermediate_size}, "
"bias=False"
)