Text Generation
PyTorch
Safetensors
English
tinyllm
tinystories
small-language-model
File size: 4,357 Bytes
515688a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
90
91
92
93
94
95
96
97
98
99
100
101
"""๋ฐฐํฌ์šฉ ๋…๋ฆฝ ๋ชจ๋ธ ์ •์˜.

์ด ํŒŒ์ผ ํ•˜๋‚˜์™€ `config.json`, `model.safetensors`, `tokenizer.json` ๋งŒ ์žˆ์œผ๋ฉด
๊ฐ€์ค‘์น˜๋ฅผ ์—ด ์ˆ˜ ์žˆ๋‹ค. ํ•™์Šต ํ•˜๋„ค์Šค(deeptool)๋Š” ํ•„์š” ์—†๋‹ค -- ๊ฐ€์ค‘์น˜ ํ•˜๋‚˜ ์—ด์ž๊ณ 
ํ•™์Šต์šฉ ๋ผ์ด๋ธŒ๋Ÿฌ๋ฆฌ๋ฅผ ์„ค์น˜ํ•˜๊ฒŒ ๋งŒ๋“ค ์ด์œ ๊ฐ€ ์—†๋‹ค.

`src/tinyllm/model.py` ์˜ TinyLM ๊ณผ ๋ชจ๋“ˆ ์ด๋ฆ„์ด ๊ฐ™์•„์•ผ state_dict ํ‚ค๊ฐ€ ๋งž๋Š”๋‹ค.
๊ทธ ๋™์ผ์„ฑ์€ tests/test_modeling.py ๊ฐ€ ์‹ค์ œ ์ฒดํฌํฌ์ธํŠธ๋กœ ๊ฒ€์ฆํ•œ๋‹ค.

    from modeling_tinyllm import TinyLM
    model = TinyLM.from_pretrained(".")
"""

import json
from pathlib import Path

import torch
from torch import nn


class TinyLM(nn.Module):
    """GPT-2 ํ˜•ํƒœ์˜ decoder-only LM, torch ๋‚ด์žฅ ๋ ˆ์ด์–ด๋กœ ์กฐ๋ฆฝ.

    decoder-only ๋ฅผ `TransformerEncoderLayer` ๋กœ ๋งŒ๋“œ๋Š” ์ด์œ ๋Š”
    `TransformerDecoderLayer` ๊ฐ€ cross-attention ์šฉ `memory` ๋ฅผ ํ•„์ˆ˜๋กœ ์š”๊ตฌํ•˜๊ธฐ
    ๋•Œ๋ฌธ์ด๋‹ค. causal mask ๋ฅผ ๋„˜๊ธฐ๋ฉด encoder layer ๊ฐ€ ๊ณง GPT ๋ธ”๋ก์ด๋‹ค.
    """

    def __init__(self, vocab_size=8000, d_model=512, n_head=8, n_layer=8,
                 d_ff=2048, block_size=512, dropout=0.0):
        super().__init__()
        self.vocab_size = vocab_size
        self.block_size = block_size

        self.tok = nn.Embedding(vocab_size, d_model)
        self.pos = nn.Embedding(block_size, d_model)
        self.drop = nn.Dropout(dropout)
        layer = nn.TransformerEncoderLayer(
            d_model, n_head, d_ff, dropout,
            activation="gelu", norm_first=True, batch_first=True,
        )
        self.blocks = nn.TransformerEncoder(
            layer, n_layer, norm=nn.RMSNorm(d_model), enable_nested_tensor=False,
        )
        self.head = nn.Linear(d_model, vocab_size, bias=False)
        self.head.weight = self.tok.weight   # weight tying

    def forward(self, ids):
        """ids: (B, T) int64 -> logits (B, T, vocab_size)"""
        T = ids.size(1)
        pos = torch.arange(T, device=ids.device)
        h = self.drop(self.tok(ids) + self.pos(pos))
        mask = nn.Transformer.generate_square_subsequent_mask(T, device=ids.device)
        return self.head(self.blocks(h, mask=mask, is_causal=True))

    @torch.no_grad()
    def generate(self, ids, max_new_tokens, temperature=0.8, top_k=40,
                 stop_id=None):
        """ids: (B, T) ํ”„๋กฌํ”„ํŠธ -> (B, T + n). stop_id ๊ฐ€ ์ „ ๋ฐฐ์น˜์— ๋‚˜์˜ค๋ฉด ์กฐ๊ธฐ ์ข…๋ฃŒ.

        KV ์บ์‹œ๋Š” ์“ฐ์ง€ ์•Š๋Š”๋‹ค. 29M ยท 512 ํ† ํฐ์ด๋ฉด forward ๊ฐ€ ์ˆ˜ ms ๋ผ ์บ์‹œ๊ฐ€ ์—†์–ด๋„
        100 ํ† ํฐ ์ƒ์„ฑ์ด 1 ์ดˆ ๋ฏธ๋งŒ์ด๋‹ค.
        """
        was_training = self.training
        self.eval()
        for _ in range(max_new_tokens):
            window = ids[:, -self.block_size:]
            logits = self(window)[:, -1] / temperature
            if top_k:
                kth = logits.topk(min(top_k, logits.size(-1)), dim=-1).values[:, -1:]
                logits = logits.masked_fill(logits < kth, float("-inf"))
            nxt = torch.multinomial(logits.softmax(-1), num_samples=1)
            ids = torch.cat([ids, nxt], dim=1)
            if stop_id is not None and bool((nxt == stop_id).all()):
                break
        self.train(was_training)
        return ids

    @classmethod
    def from_pretrained(cls, path, device="cpu"):
        """`config.json` ๊ณผ `model.safetensors` ๊ฐ€ ์žˆ๋Š” ํด๋”์—์„œ ๋ชจ๋ธ์„ ์„ธ์šด๋‹ค.

        ์ €์žฅ๋ณธ์—๋Š” `head.weight` ๊ฐ€ ์—†๋‹ค. tying ๋•Œ๋ฌธ์— `tok.weight` ์™€ ์ €์žฅ์†Œ๋ฅผ
        ๊ณต์œ ํ•˜๋Š”๋ฐ safetensors ๋Š” ๊ณต์œ  ์ €์žฅ์†Œ๋ฅผ ๊ฑฐ๋ถ€ํ•˜๊ธฐ ๋•Œ๋ฌธ์ด๋‹ค. ์ƒ์„ฑ์ž๊ฐ€ ๋‹ค์‹œ
        ๋ฌถ์œผ๋ฏ€๋กœ `strict=False` ๋กœ ๋„ฃ์–ด๋„ head ๊ฐ€ ๋น„์ง€ ์•Š๋Š”๋‹ค.
        """
        from safetensors.torch import load_file

        path = Path(path)
        config = json.loads((path / "config.json").read_text())
        model = cls(**{k: config[k] for k in
                       ("vocab_size", "d_model", "n_head", "n_layer",
                        "d_ff", "block_size")})
        state = load_file(path / "model.safetensors")
        missing, unexpected = model.load_state_dict(state, strict=False)
        if unexpected:
            raise ValueError(f"์ €์žฅ๋ณธ์— ๋ชจ๋ฅด๋Š” ํ…์„œ๊ฐ€ ์žˆ๋‹ค: {unexpected}")
        if missing != ["head.weight"]:
            raise ValueError(f"๋น ์ง„ ํ…์„œ๊ฐ€ head.weight ๋งŒ์ด ์•„๋‹ˆ๋‹ค: {missing}")
        return model.to(device).eval()