File size: 2,512 Bytes
b56bb3a
 
 
 
 
 
 
 
 
 
ca99cc3
b56bb3a
 
 
 
 
 
 
 
 
 
db4d31a
b56bb3a
ca99cc3
b56bb3a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
db4d31a
b56bb3a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
import torch.nn as nn
import torch.utils.checkpoint

from transformers import PreTrainedModel
from transformers.modeling_outputs import BaseModelOutputWithPast

from .transformer_block import TransformerBlock
from .positional_embeddings import RoPESplitHalf
from .layer_normalization import RMSNorm
from .configuration_bettergpt import BetterGPTConfig
from .logger import get_logger

logger = get_logger(__name__)


class BetterGPTModel(PreTrainedModel):
    """BetterGPT model with rotary position embeddings and pre-norm transformer blocks.

    This model is designed for efficient training and inference, supporting gradient checkpointing"""

    config_class = BetterGPTConfig
    base_model_prefix = "model"
    supports_gradient_checkpointing = True
    keys_to_ignore_at_inference = ["past_key_values"]

    def __init__(self, config: BetterGPTConfig):
        super().__init__(config)
        self.gradient_checkpointing = False

        hid = int((8 * config.emb_dim) // 3)
        hid_dim = config.ffn_multiple * (
            (hid + config.ffn_multiple - 1) // config.ffn_multiple
        )
        head_dim = config.emb_dim // config.head_count

        self.emb_layer = nn.Embedding(config.vocab_size, config.emb_dim)
        self.rmsnorm = RMSNorm(config.emb_dim, eps=config.rmsnorm_eps)
        self.rope = RoPESplitHalf(head_dim=head_dim, base=config.rope_base)

        self.transformer_block = nn.ModuleList(
            [
                TransformerBlock(
                    head_count=config.head_count,
                    head_dim=head_dim,
                    emb_dim=config.emb_dim,
                    hid_dim=hid_dim,
                    eps=config.rmsnorm_eps,
                )
                for _ in range(config.num_blocks)
            ]
        )
        self.post_init()

    def forward(self, input_ids=None, attention_mask=None, **kwargs):
        x = self.emb_layer(input_ids)
        cos, sin = self.rope(x, x.shape[1])

        for block in self.transformer_block:
            if self.gradient_checkpointing and self.training:
                x = torch.utils.checkpoint.checkpoint(
                    block, x, sin, cos, attention_mask, use_reentrant=False
                )
            else:
                x = block(x, sin, cos, attention_mask)

        x = self.rmsnorm(x)

        # Base model returns the raw hidden states
        return BaseModelOutputWithPast(last_hidden_state=x)