evex-2: 正規化記号を損失から外して記号だけの返答を 38.5% → 0%
Browse files- README.md +182 -0
- config.json +15 -0
- model.safetensors +3 -0
- modeling_evex.py +222 -0
- speakers.json +242 -0
- tok.model +3 -0
README.md
ADDED
|
@@ -0,0 +1,182 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
language:
|
| 3 |
+
- ja
|
| 4 |
+
license: mit
|
| 5 |
+
library_name: pytorch
|
| 6 |
+
pipeline_tag: text-generation
|
| 7 |
+
tags:
|
| 8 |
+
- causal-lm
|
| 9 |
+
- japanese
|
| 10 |
+
- from-scratch
|
| 11 |
+
- tiny
|
| 12 |
+
---
|
| 13 |
+
|
| 14 |
+
# evex-2
|
| 15 |
+
|
| 16 |
+
ある Discord サーバーの過去ログ **だけ** をゼロから学習した 5.87M パラメータの言語モデル。
|
| 17 |
+
既存モデルからの継続学習ではなく、tokenizer も含めて全部そのログから作った。
|
| 18 |
+
**このモデルにとって世界はそのログだけで、一般知識は一切持っていない。**
|
| 19 |
+
|
| 20 |
+
[evex-1](https://huggingface.co/tako080614/evex-1) の第2世代。コーパスも語彙も構成も
|
| 21 |
+
同じで、変えたのは2つだけ。
|
| 22 |
+
|
| 23 |
+
## evex-1 から変えたこと
|
| 24 |
+
|
| 25 |
+
### 1. 学習率 3e-4 → 1e-3
|
| 26 |
+
|
| 27 |
+
evex-1 は 10 epoch 通して val loss が単調に下がり続けていた = **収束していなかった**。
|
| 28 |
+
上げて 12 epoch 回したら val 4.2404 → 4.0384 (同じ尺度で測った値)。
|
| 29 |
+
|
| 30 |
+
### 2. 正規化記号を損失から外した
|
| 31 |
+
|
| 32 |
+
evex-1 の一番大きい問題は、**返答の 38.5% が `<url>` や `<file>` だけ**だったこと。
|
| 33 |
+
正規化で作った記号なので発言ではなく、bot は画像もリンクも実際には出せない。
|
| 34 |
+
|
| 35 |
+
原因は密度ではなく1トークンであること。添付だけの発言 (コーパスの 3.6%) が `<file>`
|
| 36 |
+
1個に潰れていたので、発言の先頭で確率が1点に集中していた。
|
| 37 |
+
|
| 38 |
+
そこで `<file>` `<url>` `<mention>` `<channel>` `<time>` を**損失から外して**学習した。
|
| 39 |
+
文脈からは消していない — 「誰かが画像を貼って、他の人が反応する」流れは学ばせたいので、
|
| 40 |
+
入力には残したまま、その位置の予測だけ評価しない。
|
| 41 |
+
|
| 42 |
+
`<nl>` と `<code>` は外していない。改行もコードブロックもモデルが書いて良いもの。
|
| 43 |
+
|
| 44 |
+
## 効果
|
| 45 |
+
|
| 46 |
+
同じチェックポイント同士を**同じ尺度で測り直した**結果 (12 epoch / 96 生成 /
|
| 47 |
+
プロンプトと乱数を固定):
|
| 48 |
+
|
| 49 |
+
| | 素 val | マスク val | 記号だけの返答 |
|
| 50 |
+
|---|---|---|---|
|
| 51 |
+
| マスクなし | **4.0384** | 4.0465 | **38.5%** |
|
| 52 |
+
| マスクあり (これ) | 4.0970 | **4.0330** | **0.0%** |
|
| 53 |
+
|
| 54 |
+
- **素 val** … `<url>` `<file>` も予測対象に含める尺度
|
| 55 |
+
- **マスク val** … 記号を損失から外す尺度 = 実際に読まれる語だけの尺度
|
| 56 |
+
|
| 57 |
+
素の尺度で 0.059 負けているのは、**その尺度が `<url>` を当てることを点数にしている**
|
| 58 |
+
から。出してほしくないトークンなので、そこで負けるのは払って良いコスト。実際に読まれる
|
| 59 |
+
語だけで測れば勝っていて、記号だけの返答は消えた。
|
| 60 |
+
|
| 61 |
+
**推論時に記号を禁止する必要がなくなった。** evex-1 では `ban_ids` で潰して 38% → 12%
|
| 62 |
+
に抑えるしかなかったが、あれは高確率のトークンを削って再正規化するので出てくる第二候補が
|
| 63 |
+
歪む。損失から外せば確率の質量が最初から実際の語に乗る。
|
| 64 |
+
|
| 65 |
+
```
|
| 66 |
+
evex-1 '<mention>' / '<file>' / 'AGPLだったら'
|
| 67 |
+
evex-2 'Cloudflare Codexがバグってる?' / 'おぉ'
|
| 68 |
+
'RadeonはCloudflare Cloudflare Pagesですらにゃい'
|
| 69 |
+
```
|
| 70 |
+
|
| 71 |
+
## 使い方
|
| 72 |
+
|
| 73 |
+
`transformers` の Auto クラスには当てはまらない構成なので、同梱の `modeling_evex.py` を使う。
|
| 74 |
+
**tokenizer は evex-1 と同一** (`tok.model` の md5 が一致) なので、既に evex-1 を持って
|
| 75 |
+
いるならそちらを使い回せる。
|
| 76 |
+
|
| 77 |
+
```bash
|
| 78 |
+
pip install torch safetensors sentencepiece huggingface_hub
|
| 79 |
+
hf download tako080614/evex-2 --local-dir evex-2
|
| 80 |
+
```
|
| 81 |
+
|
| 82 |
+
`git clone` で取るなら git-lfs が必要。入れずに clone すると `model.safetensors` が
|
| 83 |
+
133 バイトのポインタになり、読み込もうとしても壊れる。`ls -la` して 22MB あるか確かめる。
|
| 84 |
+
|
| 85 |
+
```python
|
| 86 |
+
import json, sys, torch, sentencepiece as spm
|
| 87 |
+
from safetensors.torch import load_file
|
| 88 |
+
|
| 89 |
+
sys.path.insert(0, "evex-2")
|
| 90 |
+
from modeling_evex import Config, MicroLM
|
| 91 |
+
|
| 92 |
+
cfg_json = json.load(open("evex-2/config.json"))
|
| 93 |
+
cfg = Config(
|
| 94 |
+
vocab_size=cfg_json["vocab_size"], n_layers=cfg_json["n_layers"],
|
| 95 |
+
d_model=cfg_json["d_model"], n_heads=cfg_json["n_heads"],
|
| 96 |
+
context=cfg_json["context"], dropout=0.0, attn_dropout=0.0,
|
| 97 |
+
)
|
| 98 |
+
|
| 99 |
+
model = MicroLM(cfg)
|
| 100 |
+
state = load_file("evex-2/model.safetensors")
|
| 101 |
+
state["head.weight"] = state["embed.weight"] # weight tying を結び直す
|
| 102 |
+
model.load_state_dict(state)
|
| 103 |
+
model.eval()
|
| 104 |
+
|
| 105 |
+
sp = spm.SentencePieceProcessor(model_file="evex-2/tok.model")
|
| 106 |
+
end_id = sp.piece_to_id("<|end|>")
|
| 107 |
+
|
| 108 |
+
prompt = "<|conv|><|s3|>これバグってる?<|other|>"
|
| 109 |
+
ids = torch.tensor([sp.encode(prompt, out_type=int)])
|
| 110 |
+
out = model.generate(ids, max_new_tokens=60, temperature=0.9, top_k=40, stop_id=end_id)
|
| 111 |
+
print(sp.decode(out[0].tolist()))
|
| 112 |
+
```
|
| 113 |
+
|
| 114 |
+
`head.weight` は入っていない。weight tying で `embed.weight` と同じテンソルを指しており、
|
| 115 |
+
safetensors はストレージを共有したテンソルを保存できないので落としてある。上のよ��に結び直す。
|
| 116 |
+
|
| 117 |
+
`ban_ids` は要らない。渡しても害は無いが、記号だけの返答は既に出ない。
|
| 118 |
+
|
| 119 |
+
### プロンプトの形
|
| 120 |
+
|
| 121 |
+
学習データと同じ直列化でないと、モデルは一度も見ていない形を受け取って崩れる。
|
| 122 |
+
|
| 123 |
+
```
|
| 124 |
+
<|conv|><|s3|>今日ひま?<|s7|><|re|>ひま<|end|>
|
| 125 |
+
```
|
| 126 |
+
|
| 127 |
+
| トークン | 意味 |
|
| 128 |
+
|---|---|
|
| 129 |
+
| `<\|conv\|>` | 会話の開始 |
|
| 130 |
+
| `<\|end\|>` | 会話の終了 |
|
| 131 |
+
| `<\|s0\|>`〜`<\|s47\|>` | 話者。発言数の多い上位48人 (人間の発言の 85.3% を被覆) |
|
| 132 |
+
| `<\|other\|>` | それ以外の 2,599 人 |
|
| 133 |
+
| `<\|re\|>` | 直前の誰かへの返信 |
|
| 134 |
+
| `<nl>` | 発言内の改行 |
|
| 135 |
+
| `<url>` `<mention>` `<channel>` `<time>` `<file>` | 正規化した URL / メンション / チャンネル / 時刻 / 添付。**このモデルはこれらを出さない** |
|
| 136 |
+
| `<code>` `</code>` | コードブロック |
|
| 137 |
+
|
| 138 |
+
末尾に話者トークンを置くと、その話者として続きを書く。`speakers.json` に各話者の
|
| 139 |
+
発言数が入っている (実アカウントとの対応は公開していない)。
|
| 140 |
+
|
| 141 |
+
**同じ話者トークンが続けて出ることがある。** 学習データでは話者の塊のうち 27.4% が
|
| 142 |
+
2連続以上なので、モデルもそう書く。1発言だけ欲しいなら最初の話者トークンで切る。
|
| 143 |
+
|
| 144 |
+
絵文字・`草`・`www`・顔文字は正規化せず残してあるので、そのまま出る。
|
| 145 |
+
|
| 146 |
+
## 数字
|
| 147 |
+
|
| 148 |
+
| | evex-2 | evex-1 |
|
| 149 |
+
|---|---|---|
|
| 150 |
+
| パラメータ | 5,868,800 | 5,868,800 |
|
| 151 |
+
| 学習トークン | 6,685,152 | 6,685,152 |
|
| 152 |
+
| 語彙 | 4,096 (SentencePiece BPE / byte fallback) | 同じ |
|
| 153 |
+
| context | 512 | 512 |
|
| 154 |
+
| 構成 | decoder-only / 6 層 / d_model 256 / 4 head / d_ff 704<br>RoPE + RMSNorm + SwiGLU + weight tying | 同じ |
|
| 155 |
+
| 学習 | 12 epoch / lr 1e-3 / AdamW / cosine / **CPU のみ** | 10 epoch / lr 3e-4 |
|
| 156 |
+
| 損失マスク | `<file>` `<url>` `<mention>` `<channel>` `<time>` | なし |
|
| 157 |
+
| train / val loss | 3.5535 / 4.0559 (マスク尺度) | 3.8685 / 4.2404 (素) |
|
| 158 |
+
| 記号だけの返答 | **0.0%** | 38.5% |
|
| 159 |
+
|
| 160 |
+
**データは足りていない。** Chinchilla 最適 (20 トークン/パラメータ) はこの規模だと
|
| 161 |
+
33万パラメータで、5.87M は最適の 6%。val は 12 epoch でもまだ下がり続けているので、
|
| 162 |
+
学習の余地は残っている。
|
| 163 |
+
|
| 164 |
+
## できること / できないこと
|
| 165 |
+
|
| 166 |
+
**できる**: チャットの口調、短い応答、ネットスラング、そのサーバー特有の語彙と話題、
|
| 167 |
+
話者ごとの癖 (ある話者は「〜にゃい」と書くが、それを再現する)
|
| 168 |
+
|
| 169 |
+
**できない**: 一般常識、推論、数学、コード生成、長い整合性、知らない話題への応答
|
| 170 |
+
|
| 171 |
+
固有名詞は形だけ真似て中身が合わない。`Cloudflare Codex` `AGP Core MT6589` のように、
|
| 172 |
+
見たことのある語を組み替えたものが出る。**事実として読むものではない。**
|
| 173 |
+
|
| 174 |
+
## 出どころと制限
|
| 175 |
+
|
| 176 |
+
学習データは**同意を明示的に取っていない実在の人物の会話**。
|
| 177 |
+
逐語での再生は 20 文字以上の完全一致で 0 箇所だったが、669万トークンを12周している
|
| 178 |
+
ので**実際の発言に近いものが出る可能性は残る**。
|
| 179 |
+
|
| 180 |
+
- 生成物を事実として扱ってはいけない
|
| 181 |
+
- 特定の人物の発言として扱ってはいけない
|
| 182 |
+
- 学習に使ったログそのもの、話者と実アカウントの対応表は公開していない
|
config.json
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"vocab_size": 4096,
|
| 3 |
+
"n_layers": 6,
|
| 4 |
+
"d_model": 256,
|
| 5 |
+
"n_heads": 4,
|
| 6 |
+
"context": 512,
|
| 7 |
+
"dropout": 0.1,
|
| 8 |
+
"attn_dropout": 0.0,
|
| 9 |
+
"architecture": "decoder-only transformer (RoPE + RMSNorm + SwiGLU)",
|
| 10 |
+
"tie_word_embeddings": true,
|
| 11 |
+
"trained_epoch": 12,
|
| 12 |
+
"train_loss": 3.5534612417221068,
|
| 13 |
+
"val_loss": 4.055904150009155,
|
| 14 |
+
"params": 5868800
|
| 15 |
+
}
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:57585ec95d58e7f57b451f5a22f915e055287fe2f6c75e1f625201f3cdf34f78
|
| 3 |
+
size 23479240
|
modeling_evex.py
ADDED
|
@@ -0,0 +1,222 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""evex-1 のモデル定義 (単体で動く)。
|
| 2 |
+
|
| 3 |
+
from modeling_evex import Config, MicroLM
|
| 4 |
+
|
| 5 |
+
学習に使ったものと同じコード。HF のどの Auto クラスにも当てはまらない構成なので、
|
| 6 |
+
transformers ではなくこれを直接使う。
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
"""Decoder-only Transformer。RoPE + RMSNorm + SwiGLU + weight tying。
|
| 10 |
+
|
| 11 |
+
669万トークンしか無いので、パラメータは意図的に小さく取る (Chinchilla 最適は 33万)。
|
| 12 |
+
d_model / n_layers は環境変数で振れるようにしてある。
|
| 13 |
+
|
| 14 |
+
.venv-llm/bin/python scripts/llm/model.py # パラメータ数と過学習テスト
|
| 15 |
+
"""
|
| 16 |
+
|
| 17 |
+
import math
|
| 18 |
+
from dataclasses import dataclass
|
| 19 |
+
|
| 20 |
+
import torch
|
| 21 |
+
import torch.nn.functional as F
|
| 22 |
+
from torch import nn
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
@dataclass
|
| 26 |
+
class Config:
|
| 27 |
+
vocab_size: int = 4096
|
| 28 |
+
n_layers: int = 6
|
| 29 |
+
d_model: int = 256
|
| 30 |
+
n_heads: int = 4
|
| 31 |
+
context: int = 512
|
| 32 |
+
dropout: float = 0.0
|
| 33 |
+
# アテンション内の dropout は既定で切る。
|
| 34 |
+
#
|
| 35 |
+
# dropout_p > 0 を渡すと scaled_dot_product_attention は融合カーネルを使えず、
|
| 36 |
+
# B×H×T×T のアテンション行列を実体化する math 経路に落ちる
|
| 37 |
+
# (24×4×512×512 で 1 層あたり 100MB。6 層ぶんの往復でメモリ帯域を食い潰す)。
|
| 38 |
+
# 正則化は残差側の dropout で足りるので、ここは 0 にして融合経路に乗せる。
|
| 39 |
+
attn_dropout: float = 0.0
|
| 40 |
+
|
| 41 |
+
@property
|
| 42 |
+
def d_ff(self):
|
| 43 |
+
# SwiGLU は行列が 3 つなので、4*d_model 相当に合わせて 2/3 に縮める。
|
| 44 |
+
# 64 の倍数に丸めて行列積を素直にする。
|
| 45 |
+
raw = int(self.d_model * 4 * 2 / 3)
|
| 46 |
+
return (raw + 63) // 64 * 64
|
| 47 |
+
|
| 48 |
+
@property
|
| 49 |
+
def d_head(self):
|
| 50 |
+
return self.d_model // self.n_heads
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
class RMSNorm(nn.Module):
|
| 54 |
+
"""LayerNorm から平均を引く処理を落としたもの。小さいモデルでは差が出ないが安い。"""
|
| 55 |
+
|
| 56 |
+
def __init__(self, dim, eps=1e-6):
|
| 57 |
+
super().__init__()
|
| 58 |
+
self.weight = nn.Parameter(torch.ones(dim))
|
| 59 |
+
self.eps = eps
|
| 60 |
+
|
| 61 |
+
def forward(self, x):
|
| 62 |
+
norm = x.float().pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
|
| 63 |
+
return (x.float() * norm).type_as(x) * self.weight
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
def rope_cache(context, d_head, device, base=10000.0):
|
| 67 |
+
"""RoPE の cos/sin を先に作っておく。学習中は使い回すだけ。"""
|
| 68 |
+
inv = 1.0 / (base ** (torch.arange(0, d_head, 2, device=device).float() / d_head))
|
| 69 |
+
pos = torch.arange(context, device=device).float()
|
| 70 |
+
freqs = torch.outer(pos, inv)
|
| 71 |
+
return freqs.cos(), freqs.sin()
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
def apply_rope(x, cos, sin):
|
| 75 |
+
# x: (B, heads, T, d_head)
|
| 76 |
+
t = x.shape[2]
|
| 77 |
+
cos = cos[:t].view(1, 1, t, -1)
|
| 78 |
+
sin = sin[:t].view(1, 1, t, -1)
|
| 79 |
+
|
| 80 |
+
even, odd = x[..., 0::2], x[..., 1::2]
|
| 81 |
+
rotated = torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1)
|
| 82 |
+
return rotated.flatten(-2)
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
class Attention(nn.Module):
|
| 86 |
+
def __init__(self, cfg):
|
| 87 |
+
super().__init__()
|
| 88 |
+
self.cfg = cfg
|
| 89 |
+
self.qkv = nn.Linear(cfg.d_model, cfg.d_model * 3, bias=False)
|
| 90 |
+
self.proj = nn.Linear(cfg.d_model, cfg.d_model, bias=False)
|
| 91 |
+
self.dropout = cfg.attn_dropout
|
| 92 |
+
|
| 93 |
+
def forward(self, x, cos, sin):
|
| 94 |
+
b, t, _ = x.shape
|
| 95 |
+
h, dh = self.cfg.n_heads, self.cfg.d_head
|
| 96 |
+
|
| 97 |
+
q, k, v = self.qkv(x).split(self.cfg.d_model, dim=2)
|
| 98 |
+
q = q.view(b, t, h, dh).transpose(1, 2)
|
| 99 |
+
k = k.view(b, t, h, dh).transpose(1, 2)
|
| 100 |
+
v = v.view(b, t, h, dh).transpose(1, 2)
|
| 101 |
+
|
| 102 |
+
q = apply_rope(q, cos, sin)
|
| 103 |
+
k = apply_rope(k, cos, sin)
|
| 104 |
+
|
| 105 |
+
# is_causal で三角マスクは自前で持たない (CPU でも flash 経路に乗る)
|
| 106 |
+
out = F.scaled_dot_product_attention(
|
| 107 |
+
q, k, v, is_causal=True, dropout_p=self.dropout if self.training else 0.0
|
| 108 |
+
)
|
| 109 |
+
return self.proj(out.transpose(1, 2).contiguous().view(b, t, self.cfg.d_model))
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
class SwiGLU(nn.Module):
|
| 113 |
+
def __init__(self, cfg):
|
| 114 |
+
super().__init__()
|
| 115 |
+
self.gate = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
|
| 116 |
+
self.up = nn.Linear(cfg.d_model, cfg.d_ff, bias=False)
|
| 117 |
+
self.down = nn.Linear(cfg.d_ff, cfg.d_model, bias=False)
|
| 118 |
+
|
| 119 |
+
def forward(self, x):
|
| 120 |
+
return self.down(F.silu(self.gate(x)) * self.up(x))
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
class Block(nn.Module):
|
| 124 |
+
def __init__(self, cfg):
|
| 125 |
+
super().__init__()
|
| 126 |
+
self.n1 = RMSNorm(cfg.d_model)
|
| 127 |
+
self.attn = Attention(cfg)
|
| 128 |
+
self.n2 = RMSNorm(cfg.d_model)
|
| 129 |
+
self.ff = SwiGLU(cfg)
|
| 130 |
+
self.drop = nn.Dropout(cfg.dropout)
|
| 131 |
+
|
| 132 |
+
def forward(self, x, cos, sin):
|
| 133 |
+
x = x + self.drop(self.attn(self.n1(x), cos, sin))
|
| 134 |
+
return x + self.drop(self.ff(self.n2(x)))
|
| 135 |
+
|
| 136 |
+
|
| 137 |
+
class MicroLM(nn.Module):
|
| 138 |
+
def __init__(self, cfg):
|
| 139 |
+
super().__init__()
|
| 140 |
+
self.cfg = cfg
|
| 141 |
+
self.embed = nn.Embedding(cfg.vocab_size, cfg.d_model)
|
| 142 |
+
self.drop = nn.Dropout(cfg.dropout)
|
| 143 |
+
self.blocks = nn.ModuleList(Block(cfg) for _ in range(cfg.n_layers))
|
| 144 |
+
self.norm = RMSNorm(cfg.d_model)
|
| 145 |
+
self.head = nn.Linear(cfg.d_model, cfg.vocab_size, bias=False)
|
| 146 |
+
|
| 147 |
+
# weight tying。669万トークンで語彙 4096 ぶんの出力行列を別に学ぶ余裕はない
|
| 148 |
+
self.head.weight = self.embed.weight
|
| 149 |
+
|
| 150 |
+
cos, sin = rope_cache(cfg.context, cfg.d_head, torch.device("cpu"))
|
| 151 |
+
self.register_buffer("cos", cos, persistent=False)
|
| 152 |
+
self.register_buffer("sin", sin, persistent=False)
|
| 153 |
+
|
| 154 |
+
self.apply(self._init)
|
| 155 |
+
# 残差の出口だけ層数でスケールを落とす (深くしたときに発散させない)
|
| 156 |
+
for name, param in self.named_parameters():
|
| 157 |
+
if name.endswith("proj.weight") or name.endswith("down.weight"):
|
| 158 |
+
nn.init.normal_(param, mean=0.0, std=0.02 / math.sqrt(2 * cfg.n_layers))
|
| 159 |
+
|
| 160 |
+
@staticmethod
|
| 161 |
+
def _init(module):
|
| 162 |
+
if isinstance(module, nn.Linear):
|
| 163 |
+
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
| 164 |
+
elif isinstance(module, nn.Embedding):
|
| 165 |
+
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
| 166 |
+
|
| 167 |
+
def forward(self, idx, targets=None):
|
| 168 |
+
x = self.drop(self.embed(idx))
|
| 169 |
+
for block in self.blocks:
|
| 170 |
+
x = block(x, self.cos, self.sin)
|
| 171 |
+
logits = self.head(self.norm(x))
|
| 172 |
+
|
| 173 |
+
if targets is None:
|
| 174 |
+
return logits, None
|
| 175 |
+
|
| 176 |
+
loss = F.cross_entropy(
|
| 177 |
+
logits.view(-1, logits.size(-1)), targets.reshape(-1), ignore_index=-1
|
| 178 |
+
)
|
| 179 |
+
return logits, loss
|
| 180 |
+
|
| 181 |
+
def parameter_count(self):
|
| 182 |
+
# tying しているので head は数えない (embed と同じテンソル)
|
| 183 |
+
seen = set()
|
| 184 |
+
total = 0
|
| 185 |
+
for param in self.parameters():
|
| 186 |
+
if id(param) in seen:
|
| 187 |
+
continue
|
| 188 |
+
seen.add(id(param))
|
| 189 |
+
total += param.numel()
|
| 190 |
+
return total
|
| 191 |
+
|
| 192 |
+
@torch.no_grad()
|
| 193 |
+
def generate(self, idx, max_new_tokens, temperature=0.9, top_k=40, stop_id=None,
|
| 194 |
+
ban_ids=None, min_new_tokens=0):
|
| 195 |
+
"""ban_ids: 絶対に出させないトークン。min_new_tokens: それまでは stop_id も出させない。
|
| 196 |
+
|
| 197 |
+
チャットに使うと `<url>` や `<file>` だけを吐いて終わることが多い
|
| 198 |
+
(実測 38%)。あれは正規化が作った記号で発言ではないので、
|
| 199 |
+
呼び出し側から外せるようにしてある。
|
| 200 |
+
"""
|
| 201 |
+
self.eval()
|
| 202 |
+
for step in range(max_new_tokens):
|
| 203 |
+
window = idx[:, -self.cfg.context:]
|
| 204 |
+
logits, _ = self(window)
|
| 205 |
+
logits = logits[:, -1, :] / max(temperature, 1e-5)
|
| 206 |
+
|
| 207 |
+
if ban_ids:
|
| 208 |
+
logits[:, ban_ids] = float("-inf")
|
| 209 |
+
# 何か言う前に終わらせない
|
| 210 |
+
if stop_id is not None and step < min_new_tokens:
|
| 211 |
+
logits[:, stop_id] = float("-inf")
|
| 212 |
+
|
| 213 |
+
if top_k:
|
| 214 |
+
kth = torch.topk(logits, min(top_k, logits.size(-1))).values[:, -1:]
|
| 215 |
+
logits = logits.masked_fill(logits < kth, float("-inf"))
|
| 216 |
+
|
| 217 |
+
nxt = torch.multinomial(F.softmax(logits, dim=-1), num_samples=1)
|
| 218 |
+
idx = torch.cat((idx, nxt), dim=1)
|
| 219 |
+
|
| 220 |
+
if stop_id is not None and int(nxt) == stop_id:
|
| 221 |
+
break
|
| 222 |
+
return idx
|
speakers.json
ADDED
|
@@ -0,0 +1,242 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
[
|
| 2 |
+
{
|
| 3 |
+
"token": "<|s0|>",
|
| 4 |
+
"rank": 0,
|
| 5 |
+
"messages": 55316
|
| 6 |
+
},
|
| 7 |
+
{
|
| 8 |
+
"token": "<|s1|>",
|
| 9 |
+
"rank": 1,
|
| 10 |
+
"messages": 38384
|
| 11 |
+
},
|
| 12 |
+
{
|
| 13 |
+
"token": "<|s2|>",
|
| 14 |
+
"rank": 2,
|
| 15 |
+
"messages": 30599
|
| 16 |
+
},
|
| 17 |
+
{
|
| 18 |
+
"token": "<|s3|>",
|
| 19 |
+
"rank": 3,
|
| 20 |
+
"messages": 26538
|
| 21 |
+
},
|
| 22 |
+
{
|
| 23 |
+
"token": "<|s4|>",
|
| 24 |
+
"rank": 4,
|
| 25 |
+
"messages": 26056
|
| 26 |
+
},
|
| 27 |
+
{
|
| 28 |
+
"token": "<|s5|>",
|
| 29 |
+
"rank": 5,
|
| 30 |
+
"messages": 24511
|
| 31 |
+
},
|
| 32 |
+
{
|
| 33 |
+
"token": "<|s6|>",
|
| 34 |
+
"rank": 6,
|
| 35 |
+
"messages": 23817
|
| 36 |
+
},
|
| 37 |
+
{
|
| 38 |
+
"token": "<|s7|>",
|
| 39 |
+
"rank": 7,
|
| 40 |
+
"messages": 22628
|
| 41 |
+
},
|
| 42 |
+
{
|
| 43 |
+
"token": "<|s8|>",
|
| 44 |
+
"rank": 8,
|
| 45 |
+
"messages": 21982
|
| 46 |
+
},
|
| 47 |
+
{
|
| 48 |
+
"token": "<|s9|>",
|
| 49 |
+
"rank": 9,
|
| 50 |
+
"messages": 20735
|
| 51 |
+
},
|
| 52 |
+
{
|
| 53 |
+
"token": "<|s10|>",
|
| 54 |
+
"rank": 10,
|
| 55 |
+
"messages": 17102
|
| 56 |
+
},
|
| 57 |
+
{
|
| 58 |
+
"token": "<|s11|>",
|
| 59 |
+
"rank": 11,
|
| 60 |
+
"messages": 16132
|
| 61 |
+
},
|
| 62 |
+
{
|
| 63 |
+
"token": "<|s12|>",
|
| 64 |
+
"rank": 12,
|
| 65 |
+
"messages": 13383
|
| 66 |
+
},
|
| 67 |
+
{
|
| 68 |
+
"token": "<|s13|>",
|
| 69 |
+
"rank": 13,
|
| 70 |
+
"messages": 12956
|
| 71 |
+
},
|
| 72 |
+
{
|
| 73 |
+
"token": "<|s14|>",
|
| 74 |
+
"rank": 14,
|
| 75 |
+
"messages": 10498
|
| 76 |
+
},
|
| 77 |
+
{
|
| 78 |
+
"token": "<|s15|>",
|
| 79 |
+
"rank": 15,
|
| 80 |
+
"messages": 8806
|
| 81 |
+
},
|
| 82 |
+
{
|
| 83 |
+
"token": "<|s16|>",
|
| 84 |
+
"rank": 16,
|
| 85 |
+
"messages": 8699
|
| 86 |
+
},
|
| 87 |
+
{
|
| 88 |
+
"token": "<|s17|>",
|
| 89 |
+
"rank": 17,
|
| 90 |
+
"messages": 8459
|
| 91 |
+
},
|
| 92 |
+
{
|
| 93 |
+
"token": "<|s18|>",
|
| 94 |
+
"rank": 18,
|
| 95 |
+
"messages": 8369
|
| 96 |
+
},
|
| 97 |
+
{
|
| 98 |
+
"token": "<|s19|>",
|
| 99 |
+
"rank": 19,
|
| 100 |
+
"messages": 7698
|
| 101 |
+
},
|
| 102 |
+
{
|
| 103 |
+
"token": "<|s20|>",
|
| 104 |
+
"rank": 20,
|
| 105 |
+
"messages": 7586
|
| 106 |
+
},
|
| 107 |
+
{
|
| 108 |
+
"token": "<|s21|>",
|
| 109 |
+
"rank": 21,
|
| 110 |
+
"messages": 7184
|
| 111 |
+
},
|
| 112 |
+
{
|
| 113 |
+
"token": "<|s22|>",
|
| 114 |
+
"rank": 22,
|
| 115 |
+
"messages": 6746
|
| 116 |
+
},
|
| 117 |
+
{
|
| 118 |
+
"token": "<|s23|>",
|
| 119 |
+
"rank": 23,
|
| 120 |
+
"messages": 6308
|
| 121 |
+
},
|
| 122 |
+
{
|
| 123 |
+
"token": "<|s24|>",
|
| 124 |
+
"rank": 24,
|
| 125 |
+
"messages": 5715
|
| 126 |
+
},
|
| 127 |
+
{
|
| 128 |
+
"token": "<|s25|>",
|
| 129 |
+
"rank": 25,
|
| 130 |
+
"messages": 5308
|
| 131 |
+
},
|
| 132 |
+
{
|
| 133 |
+
"token": "<|s26|>",
|
| 134 |
+
"rank": 26,
|
| 135 |
+
"messages": 4541
|
| 136 |
+
},
|
| 137 |
+
{
|
| 138 |
+
"token": "<|s27|>",
|
| 139 |
+
"rank": 27,
|
| 140 |
+
"messages": 4419
|
| 141 |
+
},
|
| 142 |
+
{
|
| 143 |
+
"token": "<|s28|>",
|
| 144 |
+
"rank": 28,
|
| 145 |
+
"messages": 3834
|
| 146 |
+
},
|
| 147 |
+
{
|
| 148 |
+
"token": "<|s29|>",
|
| 149 |
+
"rank": 29,
|
| 150 |
+
"messages": 3649
|
| 151 |
+
},
|
| 152 |
+
{
|
| 153 |
+
"token": "<|s30|>",
|
| 154 |
+
"rank": 30,
|
| 155 |
+
"messages": 3628
|
| 156 |
+
},
|
| 157 |
+
{
|
| 158 |
+
"token": "<|s31|>",
|
| 159 |
+
"rank": 31,
|
| 160 |
+
"messages": 3138
|
| 161 |
+
},
|
| 162 |
+
{
|
| 163 |
+
"token": "<|s32|>",
|
| 164 |
+
"rank": 32,
|
| 165 |
+
"messages": 2834
|
| 166 |
+
},
|
| 167 |
+
{
|
| 168 |
+
"token": "<|s33|>",
|
| 169 |
+
"rank": 33,
|
| 170 |
+
"messages": 2650
|
| 171 |
+
},
|
| 172 |
+
{
|
| 173 |
+
"token": "<|s34|>",
|
| 174 |
+
"rank": 34,
|
| 175 |
+
"messages": 2527
|
| 176 |
+
},
|
| 177 |
+
{
|
| 178 |
+
"token": "<|s35|>",
|
| 179 |
+
"rank": 35,
|
| 180 |
+
"messages": 2504
|
| 181 |
+
},
|
| 182 |
+
{
|
| 183 |
+
"token": "<|s36|>",
|
| 184 |
+
"rank": 36,
|
| 185 |
+
"messages": 2302
|
| 186 |
+
},
|
| 187 |
+
{
|
| 188 |
+
"token": "<|s37|>",
|
| 189 |
+
"rank": 37,
|
| 190 |
+
"messages": 2280
|
| 191 |
+
},
|
| 192 |
+
{
|
| 193 |
+
"token": "<|s38|>",
|
| 194 |
+
"rank": 38,
|
| 195 |
+
"messages": 2243
|
| 196 |
+
},
|
| 197 |
+
{
|
| 198 |
+
"token": "<|s39|>",
|
| 199 |
+
"rank": 39,
|
| 200 |
+
"messages": 2229
|
| 201 |
+
},
|
| 202 |
+
{
|
| 203 |
+
"token": "<|s40|>",
|
| 204 |
+
"rank": 40,
|
| 205 |
+
"messages": 2049
|
| 206 |
+
},
|
| 207 |
+
{
|
| 208 |
+
"token": "<|s41|>",
|
| 209 |
+
"rank": 41,
|
| 210 |
+
"messages": 2047
|
| 211 |
+
},
|
| 212 |
+
{
|
| 213 |
+
"token": "<|s42|>",
|
| 214 |
+
"rank": 42,
|
| 215 |
+
"messages": 1936
|
| 216 |
+
},
|
| 217 |
+
{
|
| 218 |
+
"token": "<|s43|>",
|
| 219 |
+
"rank": 43,
|
| 220 |
+
"messages": 1932
|
| 221 |
+
},
|
| 222 |
+
{
|
| 223 |
+
"token": "<|s44|>",
|
| 224 |
+
"rank": 44,
|
| 225 |
+
"messages": 1925
|
| 226 |
+
},
|
| 227 |
+
{
|
| 228 |
+
"token": "<|s45|>",
|
| 229 |
+
"rank": 45,
|
| 230 |
+
"messages": 1896
|
| 231 |
+
},
|
| 232 |
+
{
|
| 233 |
+
"token": "<|s46|>",
|
| 234 |
+
"rank": 46,
|
| 235 |
+
"messages": 1885
|
| 236 |
+
},
|
| 237 |
+
{
|
| 238 |
+
"token": "<|s47|>",
|
| 239 |
+
"rank": 47,
|
| 240 |
+
"messages": 1880
|
| 241 |
+
}
|
| 242 |
+
]
|
tok.model
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:3b420040c24f8b3ef739d8f2230820300e8a17d8371333fd731cc19444b0f2f3
|
| 3 |
+
size 58498
|