File size: 2,260 Bytes
80c3430
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
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"
        )