Text Generation
Safetensors
Japanese
japanese
discord
from-scratch
File size: 19,420 Bytes
75f8e79
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
"""Decoder-only Transformer。RoPE + RMSNorm + SwiGLU + weight tying。

669万トークンしか無いので、パラメータは意図的に小さく取る (Chinchilla 最適は 33万)。
d_model / n_layers は環境変数で振れるようにしてある。

    .venv-llm/bin/python scripts/llm/model.py     # パラメータ数と過学習テスト
"""

import math
import os
from dataclasses import dataclass

import torch
import torch.nn.functional as F
from torch import nn


@dataclass
class Config:
    vocab_size: int = 4096
    n_layers: int = int(os.environ.get("LLM_LAYERS", 6))
    d_model: int = int(os.environ.get("LLM_DMODEL", 256))
    n_heads: int = int(os.environ.get("LLM_HEADS", 4))
    context: int = int(os.environ.get("LLM_CONTEXT", 512))
    dropout: float = float(os.environ.get("LLM_DROPOUT", 0.1))
    # アテンション内の dropout は既定で切る。
    #
    # dropout_p > 0 を渡すと scaled_dot_product_attention は融合カーネルを使えず、
    # B×H×T×T のアテンション行列を実体化する math 経路に落ちる
    # (24×4×512×512 で 1 層あたり 100MB。6 層ぶんの往復でメモリ帯域を食い潰す)。
    # 正則化は残差側の dropout で足りるので、ここは 0 にして融合経路に乗せる。
    attn_dropout: float = float(os.environ.get("LLM_ATTN_DROPOUT", 0.0))

    # --- evex-5 で足したもの。**既定は False で、旧世代の重みがそのまま読める** ---
    #
    # PLE (Per-Layer Embeddings / Gemma 3n)。トークンごと・層ごとの補助ベクトルを
    # 引いて各層の残差に足す。**行列積ではなく引き算**なので、容量は増えるのに
    # 計算量はほぼ増えない (d_ple=64 で パラメータ +34% / FLOP +1.04%)。
    #
    # evex は CPU 推論で行列積律速なので、この交換比は理屈が合う。
    ple: bool = os.environ.get("LLM_PLE", "0") == "1"
    d_ple: int = int(os.environ.get("LLM_DPLE", 64))
    # q/k を RMSNorm してから RoPE を掛ける (Gemma 3)。ほぼ 0 パラメータで
    # 学習が安定し、学習率を上げられる
    qk_norm: bool = os.environ.get("LLM_QK_NORM", "0") == "1"

    @property
    def d_ff(self):
        # SwiGLU は行列が 3 つなので、4*d_model 相当に合わせて 2/3 に縮める。
        # 64 の倍数に丸めて行列積を素直にする。
        raw = int(self.d_model * 4 * 2 / 3)
        return (raw + 63) // 64 * 64

    @property
    def d_head(self):
        return self.d_model // self.n_heads


class RMSNorm(nn.Module):
    """LayerNorm から平均を引く処理を落としたもの。小さいモデルでは差が出ないが安い。"""

    def __init__(self, dim, eps=1e-6):
        super().__init__()
        self.weight = nn.Parameter(torch.ones(dim))
        self.eps = eps

    def forward(self, x):
        norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
        return (x.float() * norm).type_as(x) * self.weight


def rope_cache(context, d_head, device, base=10000.0):
    """RoPE の cos/sin を先に作っておく。学習中は使い回すだけ。"""
    inv = 1.0 / (base ** (torch.arange(0, d_head, 2, device=device).float() / d_head))
    pos = torch.arange(context, device=device).float()
    freqs = torch.outer(pos, inv)
    return freqs.cos(), freqs.sin()


def apply_rope(x, cos, sin, offset=0):
    """offset は「この列が何トークン目から始まるか」。

    **KV キャッシュを使うときに要る。**2 トークン目以降は 1 個ずつ入れるので、
    そのままだと毎回「0 トークン目」として回してしまい、位置が壊れる。
    """
    # x: (B, heads, T, d_head)
    t = x.shape[2]
    cos = cos[offset:offset + t].view(1, 1, t, -1)
    sin = sin[offset:offset + t].view(1, 1, t, -1)

    # **回転は fp32 で計算して、最後に x の型に戻す。**
    #
    # cos/sin は fp32 の buffer なので、半精度の x と掛けると結果だけ fp32 に
    # 昇格する。そのまま返すと q/k が fp32・v が半精度で
    # scaled_dot_product_attention に入り、型が揃わない。
    # 精度も落とさずに済むので、計算は fp32 のまま最後に揃える。
    even, odd = x[..., 0::2].float(), x[..., 1::2].float()
    rotated = torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1)
    return rotated.flatten(-2).to(x.dtype)


class Attention(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        self.qkv = nn.Linear(cfg.d_model, cfg.d_model * 3, bias=False)
        self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
        self.dropout = cfg.attn_dropout
        # **RoPE の前に掛ける。**後だと回転した向きごと正規化してしまう
        self.qn = RMSNorm(cfg.d_head) if cfg.qk_norm else None
        self.kn = RMSNorm(cfg.d_head) if cfg.qk_norm else None

    def forward(self, x, cos, sin, attn_mask=None, cache=None, offset=0):
        """cache に (k, v) を渡すと、そこに継ぎ足して使う (生成用)。

        返すのは出力だけ。**新しい cache は self.last_cache に置く** —
        Block と MicroLM の戻り値の形を変えると学習側まで書き換えになる。
        """
        b, t, _ = x.shape
        h, dh = self.cfg.n_heads, self.cfg.d_head

        q, k, v = self.qkv(x).split(self.cfg.d_model, dim=2)
        q = q.view(b, t, h, dh).transpose(1, 2)
        k = k.view(b, t, h, dh).transpose(1, 2)
        v = v.view(b, t, h, dh).transpose(1, 2)

        if self.qn is not None:
            q, k = self.qn(q).type_as(v), self.kn(k).type_as(v)

        q = apply_rope(q, cos, sin, offset)
        k = apply_rope(k, cos, sin, offset)

        if cache is not None:
            past_k, past_v = cache
            if past_k is not None:
                k = torch.cat((past_k, k), dim=2)
                v = torch.cat((past_v, v), dim=2)
            self.last_cache = (k, v)

        # is_causal で三角マスクは自前で持たない (CPU でも flash 経路に乗る)。
        #
        # **文書内マスクを渡すと融合カーネルから落ちる。**呼ぶ側が既に因果性を
        # 含めたマスクを組んでいるので、そのときは is_causal を外す
        # **キャッシュを使って 1 トークンだけ入れるときは is_causal を外す。**
        # 問い合わせが 1 個で鍵が過去全部なので、三角マスクを掛けると
        # 自分より前を全部隠してしまう
        causal = attn_mask is None and q.shape[2] == k.shape[2]
        out = F.scaled_dot_product_attention(
            q, k, v,
            attn_mask=attn_mask,
            is_causal=causal,
            dropout_p=self.dropout if self.training else 0.0
        )
        return self.proj(out.transpose(1, 2).contiguous().view(b, t, self.cfg.d_model))


class SwiGLU(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.gate = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
        self.up = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
        self.down = nn.Linear(cfg.d_ff, cfg.d_model, bias=False)

    def forward(self, x):
        return self.down(F.silu(self.gate(x)) * self.up(x))


class Block(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.n1 = RMSNorm(cfg.d_model)
        self.attn = Attention(cfg)
        self.n2 = RMSNorm(cfg.d_model)
        self.ff = SwiGLU(cfg)
        self.drop = nn.Dropout(cfg.dropout)

        # PLE の注入 (Gemma 3n の per-layer input)。その層ぶんの補助ベクトル p を
        # ゲートで混ぜて残差に足す。**行列は d_model×d_ple の 2 枚だけ**なので、
        # 引いてきた容量に対して計算量はほとんど増えない
        if cfg.ple:
            self.ple_gate = nn.Linear(cfg.d_model, cfg.d_ple, bias=False)
            self.ple_out = nn.Linear(cfg.d_ple, cfg.d_model, bias=False)
        else:
            self.ple_gate = self.ple_out = None

    def forward(self, x, cos, sin, per_layer=None, attn_mask=None, cache=None, offset=0):
        x = x + self.drop(self.attn(self.n1(x), cos, sin, attn_mask, cache, offset))
        x = x + self.drop(self.ff(self.n2(x)))

        if self.ple_out is not None and per_layer is not None:
            x = x + self.ple_out(F.silu(self.ple_gate(x)) * per_layer)
        return x


class MicroLM(nn.Module):
    def __init__(self, cfg):
        super().__init__()
        self.cfg = cfg
        self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model)
        self.drop = nn.Dropout(cfg.dropout)
        self.blocks = nn.ModuleList(Block(cfg) for _ in range(cfg.n_layers))
        self.norm = RMSNorm(cfg.d_model)
        self.head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)

        # PLE の 2 系統 (Gemma 3n と同じ組み合わせ):
        #   ple_table  トークン同一性。語彙 × 層 × d_ple の引き表
        #   ple_proj   文脈側。入力埋め込みから層ぶんを作る (per_layer_model_projection)
        # 足して RMSNorm したものが、その層の補助ベクトルになる
        if cfg.ple:
            self.ple_table = nn.Embedding(cfg.vocab_size, cfg.n_layers * cfg.d_ple)
            self.ple_proj = nn.Linear(cfg.d_model, cfg.n_layers * cfg.d_ple, bias=False)
            self.ple_norm = RMSNorm(cfg.d_ple)
        else:
            self.ple_table = self.ple_proj = self.ple_norm = None

        # weight tying。669万トークンで語彙 4096 ぶんの出力行列を別に学ぶ余裕はない
        self.head.weight = self.embed.weight

        cos, sin = rope_cache(cfg.context, cfg.d_head, torch.device("cpu"))
        self.register_buffer("cos", cos, persistent=False)
        self.register_buffer("sin", sin, persistent=False)

        self.apply(self._init)
        # 残差の出口だけ層数でスケールを落とす (深くしたときに発散させない)
        for name, param in self.named_parameters():
            if name.endswith("proj.weight") or name.endswith("down.weight"):
                nn.init.normal_(param, mean=0.0, std=0.02 / math.sqrt(2 * cfg.n_layers))

    @staticmethod
    def _init(module):
        if isinstance(module, nn.Linear):
            nn.init.normal_(module.weight, mean=0.0, std=0.02)
        elif isinstance(module, nn.Embedding):
            nn.init.normal_(module.weight, mean=0.0, std=0.02)

    def per_layer_inputs(self, idx, x):
        """各層に配る補助ベクトル (B, T, n_layers, d_ple)。PLE が無ければ None。"""
        if self.ple_table is None:
            return None
        b, t = idx.shape
        shape = (b, t, self.cfg.n_layers, self.cfg.d_ple)
        # 表の側だけ sqrt(d_ple) で持ち上げる (Gemma 3n と同じ。射影側と桁を揃える)
        table = self.ple_table(idx).view(shape) * math.sqrt(self.cfg.d_ple)
        return self.ple_norm(table + self.ple_proj(x).view(shape)).type_as(x)

    def forward(self, idx, targets=None, attn_mask=None, chunks=1, z_loss=0.0,
                caches=None, offset=0):
        """targets を渡すと (None, loss)、渡さないと (logits, None)。

        **targets があるとき logits は返さない。**呼ぶ側は全部捨てているうえ、
        logits は batch×context×語彙 (48×1024×12288×4B = 2.25GB) あって、
        これが 2 回 OOM を出した張本人。時間方向に割って足し合わせれば、
        一度に実体化する量が 1/chunks になる。
        """
        x = self.drop(self.embed(idx))
        per_layer = self.per_layer_inputs(idx, x)

        for i, block in enumerate(self.blocks):
            p = per_layer[:, :, i] if per_layer is not None else None
            x = block(x, self.cos, self.sin, p, attn_mask,
                      None if caches is None else caches[i], offset)
        h = self.norm(x)

        if targets is None:
            return self.head(h), None

        # 損失に入る位置の数。**割った塊ごとに平均すると、塊で個数が違うときに
        # 重みがずれる。**先に総数を出して足し込む
        counted = (targets != -1).sum().clamp(min=1)
        size = max(1, h.size(1) // max(1, chunks))
        loss = h.new_zeros((), dtype=torch.float32)

        for hc, tc in zip(h.split(size, dim=1), targets.split(size, dim=1)):
            logits = self.head(hc).float()
            flat, flat_t = logits.view(-1, logits.size(-1)), tc.reshape(-1)
            loss = loss + F.cross_entropy(
                flat, flat_t, ignore_index=-1, reduction="sum"
            )
            # z-loss: logsumexp を 0 に寄せて logits が膨らむのを抑える。
            # fp16 で回すときの安定化 (PaLM / Gemma と同じ)
            if z_loss:
                keep = flat_t != -1
                if keep.any():
                    z = torch.logsumexp(flat[keep], dim=-1)
                    loss = loss + z_loss * z.pow(2).sum()

        return None, loss / counted

    def parameter_count(self):
        # tying しているので head は数えない (embed と同じテンソル)
        seen = set()
        total = 0
        for param in self.parameters():
            if id(param) in seen:
                continue
            seen.add(id(param))
            total += param.numel()
        return total

    @torch.no_grad()
    def generate(self, idx, max_new_tokens, temperature=0.9, top_k=40, stop_id=None,
                 ban_ids=None, min_new_tokens=0, min_p=0.0, repetition_penalty=1.0):
        """ban_ids: 絶対に出させないトークン。min_new_tokens: それまでは stop_id も出させない。

        チャットに使うと `<url>` や `<file>` だけを吐いて終わることが多い
        (実測 38%)。あれは正規化が作った記号で発言ではないので、
        呼び出し側から外せるようにしてある。

        min_p と repetition_penalty は **evex-ft (transformers) 側と同じ手を
        こちらでも使えるようにするため**に足した。世代を読み比べるときに、
        サンプリングが違うと差がモデル由来かハーネス由来か分からなくなる。
        """
        self.eval()

        # **KV キャッシュ。**無いと 1 トークン出すたびに文脈全体を 8 層へ通し直す。
        # 文脈 300 で 40 トークン生成すると 12,000 トークン分の計算になり、
        # 本来の 340 に対して 35 倍の無駄 (bot の応答が遅い原因はこれ)。
        #
        # 最初に prompt をまとめて通し (prefill)、そのあとは 1 トークンずつ。
        caches = [(None, None) for _ in self.blocks]
        fed = 0

        for step in range(max_new_tokens):
            window = idx[:, -self.cfg.context:]
            if fed == 0:
                piece, offset = window, 0            # prefill
            else:
                piece, offset = idx[:, -1:], fed     # 1 トークンだけ
            logits, _ = self(piece, caches=caches, offset=offset)
            fed = offset + piece.shape[1]

            # 次の周のために、各層が置いた新しい鍵と値を拾う
            caches = [block.attn.last_cache for block in self.blocks]
            # 文脈からあふれたら古い方を捨てる (窓と同じ長さに保つ)
            if fed > self.cfg.context:
                drop = fed - self.cfg.context
                caches = [(k[:, :, drop:], v[:, :, drop:]) for k, v in caches]
                fed = self.cfg.context

            logits = logits[:, -1, :]

            # 繰り返しペナルティは温度より前に掛ける (transformers と同じ順序)。
            # 既に出したトークンの確率を割る。負の logit は掛ける方が下がるので分ける
            if repetition_penalty and repetition_penalty != 1.0:
                for row in range(idx.size(0)):
                    seen = torch.unique(idx[row])
                    picked = logits[row, seen]
                    logits[row, seen] = torch.where(
                        picked > 0, picked / repetition_penalty, picked * repetition_penalty
                    )

            logits = logits / max(temperature, 1e-5)

            if ban_ids:
                logits[:, ban_ids] = float("-inf")
            # 何か言う前に終わらせない
            if stop_id is not None and step < min_new_tokens:
                logits[:, stop_id] = float("-inf")

            if top_k:
                kth = torch.topk(logits, min(top_k, logits.size(-1))).values[:, -1:]
                logits = logits.masked_fill(logits < kth, float("-inf"))

            probs = F.softmax(logits, dim=-1)

            # min_p: 最大確率の min_p 倍を下回る候補を切る。top_k だけより崩れが減る。
            # 分布が尖っているときは強く絞り、平らなときは緩む
            if min_p and min_p > 0:
                floor = probs.max(dim=-1, keepdim=True).values * min_p
                probs = torch.where(probs < floor, torch.zeros_like(probs), probs)
                probs = probs / probs.sum(dim=-1, keepdim=True)

            nxt = torch.multinomial(probs, num_samples=1)
            idx = torch.cat((idx, nxt), dim=1)

            if stop_id is not None and int(nxt) == stop_id:
                break
        return idx


if __name__ == "__main__":
    cfg = Config()
    model = MicroLM(cfg)
    params = model.parameter_count()

    print(f"layers {cfg.n_layers} / d_model {cfg.d_model} / heads {cfg.n_heads} "
          f"/ d_ff {cfg.d_ff} / context {cfg.context}")
    print(f"パラメータ {params:,} ({params / 1e6:.2f}M)")

    embed = cfg.vocab_size * cfg.d_model
    print(f"  うち埋め込み {embed:,} ({embed / params * 100:.0f}%)")
    print(f"669万トークンに対して {6_685_152 / params:.1f} トークン/パラメータ"
          f" (Chinchilla 最適の {6_685_152 / params / 20 * 100:.0f}%)")

    # --- 実装が正しいかの確認 ---
    #
    # 小さい切片を暗記させる。ここで loss が落ちないならモデルかデータの配線が
    # 壊れているので、本番を一晩回す意味がない。
    #
    # dropout は切る。見ているのは「暗記できるか = 配線が通っているか」で、
    # 正則化が効いていると当然落ちきらない (0.1 のままだと 400 step で 0.62 止まり、
    # 切れば 200 step で 0.012 まで落ちる)。
    torch.manual_seed(0)
    model = MicroLM(Config(dropout=0.0))
    data = torch.randint(0, cfg.vocab_size, (4, 65))
    opt = torch.optim.AdamW(model.parameters(), lr=3e-3)

    model.train()
    losses = []
    for step in range(200):
        _, loss = model(data[:, :-1], data[:, 1:])
        opt.zero_grad(set_to_none=True)
        loss.backward()
        opt.step()
        losses.append(loss.item())

    print(f"暗記テスト loss {losses[0]:.3f} -> {losses[-1]:.4f}")
    assert losses[-1] < 0.1, f"暗記できていない (loss {losses[-1]:.3f})。配線が壊れている"

    # 生成が止まること
    out = model.generate(torch.zeros((1, 1), dtype=torch.long), max_new_tokens=8, stop_id=None)
    assert out.shape == (1, 9), out.shape

    print("model ok")