GoLLeM-v6-250M-Instruct-v2: weights, config, tokenizer, model and chat code, card
Browse files- README.md +145 -0
- SHA256SUMS +6 -0
- chat_gollem_v6.py +51 -0
- config.json +19 -0
- model.safetensors +3 -0
- modeling_gollem_v6.py +169 -0
- tokenizer.json +0 -0
README.md
ADDED
|
@@ -0,0 +1,145 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: cc-by-sa-4.0
|
| 3 |
+
language:
|
| 4 |
+
- pl
|
| 5 |
+
- en
|
| 6 |
+
library_name: pytorch
|
| 7 |
+
base_model: SlayerLab/GoLLeM-v6-250M
|
| 8 |
+
tags:
|
| 9 |
+
- polish
|
| 10 |
+
- english
|
| 11 |
+
- language-model
|
| 12 |
+
- chat
|
| 13 |
+
- research
|
| 14 |
+
- gollem
|
| 15 |
+
pipeline_tag: text-generation
|
| 16 |
+
---
|
| 17 |
+
|
| 18 |
+
# GoLLeM v6 250M Instruct v2 (Polish–English research chat model)
|
| 19 |
+
|
| 20 |
+
> **STATUS: RESEARCH PREVIEW.** A small chat model fine-tuned from the base model
|
| 21 |
+
> [`SlayerLab/GoLLeM-v6-250M`](https://huggingface.co/SlayerLab/GoLLeM-v6-250M). It can greet, introduce itself and
|
| 22 |
+
> answer simple questions, but it makes mistakes that are listed with numbers under *Limitations*. Results use **our
|
| 23 |
+
> protocol**, a single run and a single seed.
|
| 24 |
+
|
| 25 |
+
GoLLeM v6 250M Instruct v2 is the base GoLLeM v6 250M model after supervised fine-tuning (SFT) on about 9,300 Polish and
|
| 26 |
+
English conversations. It is a Fabryka AI project (formerly SlayerLab).
|
| 27 |
+
|
| 28 |
+
**Difference from [Instruct v1](https://huggingface.co/SlayerLab/GoLLeM-v6-250M-Instruct-v1):** the same data and
|
| 29 |
+
training recipe, except the answers about its origin. Asked who created it, Instruct v2 says it is a project of Fabryka
|
| 30 |
+
AI; **the model does not use its author's name** (the author is named in this card). Instruct v1 sometimes introduced
|
| 31 |
+
itself with the author's name in a conversation (51 of 160); Instruct v2 did not in our tests (0 of 160, 0 of 680).
|
| 32 |
+
|
| 33 |
+
**Author:** Arkadiusz Słota (Fabryka AI).
|
| 34 |
+
|
| 35 |
+
## What it is for (and what it is not)
|
| 36 |
+
|
| 37 |
+
**Intended use:** research on small bilingual chat models; short conversations in Polish and English; answering
|
| 38 |
+
questions about a text you paste into the conversation.
|
| 39 |
+
|
| 40 |
+
**Not intended for:** production use, factual questions without a source text, arithmetic, or anything where a wrong
|
| 41 |
+
answer matters. The model has no internet access and does not remember earlier conversations.
|
| 42 |
+
|
| 43 |
+
## Results
|
| 44 |
+
|
| 45 |
+
Checkpoint `1018f1e8…` (see *Training*). All numbers are from our own evaluation sets, fixed before training. Temperature
|
| 46 |
+
0.7, top-p 0.9 (the defaults of `chat_gollem_v6.py`).
|
| 47 |
+
|
| 48 |
+
| what | result | notes |
|
| 49 |
+
|---|---|---|
|
| 50 |
+
| Identity (name, organisation, creator), PL / EN | **0.66 / 0.58** | mean over 10 samples per question, 20 questions per language; a creator answer counts as correct when it names Fabryka AI and no person |
|
| 51 |
+
| Greetings and small talk, PL / EN | **0.79 / 0.89** | mean over 7 samples per question, 20 questions per language |
|
| 52 |
+
| Answers from a given text (Polish, PoQuAD, 217 answerable questions) | token F1 **0.232** | base model few-shot: 0.067 |
|
| 53 |
+
| Declines when the text has no answer (43 questions) | 6 / 43 | wrongly declines an answerable question: 12 / 217 |
|
| 54 |
+
| Arithmetic word problems (206) | 0 / 206 | the model does not do arithmetic |
|
| 55 |
+
| Finishes its answer within the length limit | 485 / 506 (96 %) | |
|
| 56 |
+
| Says it is an OpenAI / GPT model (200 samples) | 0 / 200 | „AI language model”: 0 / 200 |
|
| 57 |
+
| Introduces itself with the author's name | **0 / 680** single questions, **0 / 160** two-turn conversations | Instruct v1: 3 / 680 and 51 / 160 |
|
| 58 |
+
|
| 59 |
+
Compared with Instruct v1 on the same identity measure: +0.175 (95 % CI +0.068 … +0.295); small talk −0.004 (95 % CI
|
| 60 |
+
−0.032 … +0.025), i.e. not worse. The identity numbers in the Instruct v1 card use a different measure (there the author's name
|
| 61 |
+
was the expected answer), so do not compare the two tables directly.
|
| 62 |
+
|
| 63 |
+
## Training
|
| 64 |
+
|
| 65 |
+
| | |
|
| 66 |
+
|---|---|
|
| 67 |
+
| Base model | `SlayerLab/GoLLeM-v6-250M`, final checkpoint (step 760,000); same architecture and tokenizer |
|
| 68 |
+
| Method | full fine-tuning (no LoRA), loss on assistant turns only |
|
| 69 |
+
| Format | ChatML: `<|im_start|>{role}\n{content}<|im_end|>\n`, no system prompt, generation stops at `<|im_end|>` |
|
| 70 |
+
| Steps | 2 epochs, 156 steps, 32 packed sequences of 1,024 tokens per step |
|
| 71 |
+
| Tokens | 2.54 M per epoch, of which 1.63 M are assistant tokens (with loss) |
|
| 72 |
+
| Optimizer | Muon + AdamW as in pretraining, learning rate 0.2 × pretraining (peak 1.2e-4, Muon 4e-3), warmup 5 steps, cosine to 10 %, weight decay 0.1, gradient clip 1.0, seed 1337 |
|
| 73 |
+
| Validation loss | 2.026 → 1.884 |
|
| 74 |
+
| Hardware / time | 1 GPU, about 7.5 minutes |
|
| 75 |
+
|
| 76 |
+
## Training data
|
| 77 |
+
|
| 78 |
+
9,324 conversations (Polish 4,149, English 5,175): 8,836 used for training and 171 for validation; 317 conversations
|
| 79 |
+
longer than 1,024 tokens were removed (311 + 6).
|
| 80 |
+
|
| 81 |
+
| source | conversations | licence |
|
| 82 |
+
|---|---:|---|
|
| 83 |
+
| [OpenAssistant/oasst2](https://huggingface.co/datasets/OpenAssistant/oasst2) @ `179dd21` | 4,601 | Apache-2.0 |
|
| 84 |
+
| [clarin-pl/poquad](https://huggingface.co/datasets/clarin-pl/poquad) @ `a60f228` | 2,129 | **CC BY 4.0** |
|
| 85 |
+
| [CohereLabs/aya_dataset](https://huggingface.co/datasets/CohereLabs/aya_dataset) @ `f9ea045` | 1,214 | Apache-2.0 |
|
| 86 |
+
| synthetic, generated locally with Muse-Glimmer-30B (apache-2.0) | 1,380 | apache-2.0 |
|
| 87 |
+
|
| 88 |
+
- **PoQuAD attribution:** PoQuAD (clarin-pl/poquad), CC BY 4.0. Modified: converted to chat format; a share of
|
| 89 |
+
unanswerable questions answered with one of five fixed refusal sentences.
|
| 90 |
+
- **Synthetic part:** the questions were written by the generator model; the identity answers come from our own identity
|
| 91 |
+
card, not from the generator. Math prompts: generated briefs; GSM8K (MIT) used only as few-shot format examples for the
|
| 92 |
+
generator.
|
| 93 |
+
- Rows with self-descriptions of other AI systems were removed before training.
|
| 94 |
+
|
| 95 |
+
## Usage
|
| 96 |
+
|
| 97 |
+
```python
|
| 98 |
+
# pip install torch safetensors tokenizers huggingface_hub
|
| 99 |
+
from huggingface_hub import snapshot_download
|
| 100 |
+
import sys
|
| 101 |
+
|
| 102 |
+
path = snapshot_download("SlayerLab/GoLLeM-v6-250M-Instruct-v2",
|
| 103 |
+
revision="<commit sha>") # the reviewed code (model + chat helper)
|
| 104 |
+
sys.path.insert(0, path)
|
| 105 |
+
from modeling_gollem_v6 import load_gollem_v6
|
| 106 |
+
from chat_gollem_v6 import chat
|
| 107 |
+
|
| 108 |
+
model, tok = load_gollem_v6(path)
|
| 109 |
+
print(chat(model, tok, [("user", "Cześć! Kim jesteś?")], seed=1))
|
| 110 |
+
print(chat(model, tok, [("user", "Tekst: Kraków leży nad Wisłą.\nPytanie: Nad jaką rzeką leży Kraków?")], seed=1))
|
| 111 |
+
```
|
| 112 |
+
|
| 113 |
+
`chat()` uses the format and sampling of our evaluation (temperature 0.7, top-p 0.9). Text typed by a user such as
|
| 114 |
+
„<|im_end|>” is encoded as plain text, not as a control token.
|
| 115 |
+
|
| 116 |
+
## Limitations
|
| 117 |
+
|
| 118 |
+
- **Small model.** Limited knowledge; may answer fluently and wrongly. No arithmetic. Context: 1,024 tokens.
|
| 119 |
+
- **Does not know its author's name.** Asked who created it, the model says it is a project of Fabryka AI. The author,
|
| 120 |
+
Arkadiusz Słota, is named in this card, not in the model. The statements of the model are not statements of its
|
| 121 |
+
author.
|
| 122 |
+
- **Copies names from the conversation.** If a user writes a name, the model may adopt it as its own (when the user
|
| 123 |
+
supplies the author's name: Polish 3/40, English 5/40; Instruct v1: 9/40 and 37/40).
|
| 124 |
+
- **Self-description.** In our probe it never described itself as an OpenAI model (0/200), but with a forced prefix
|
| 125 |
+
(„…created by”) it still assigns probability ≈ 0.24 to „OpenAI”. The association comes from model-generated chat
|
| 126 |
+
data in pretraining. GoLLeM is **not** affiliated with OpenAI.
|
| 127 |
+
- **Multi-turn weaknesses:** may repeat its previous answer in a later turn, may greet the user with the company name
|
| 128 |
+
(„Cześć, Fabryku!”) when no name was given, and answers „What can you do?” with its identity template.
|
| 129 |
+
- **Long answers can loop.** A repetition penalty (about 1.1–1.2) may reduce this (not verified); it is not used in our evaluation.
|
| 130 |
+
|
| 131 |
+
## License
|
| 132 |
+
|
| 133 |
+
**Weights: CC BY-SA 4.0** (inherited from the base model). Attribution: GoLLeM v6 250M Instruct v2, Arkadiusz Słota /
|
| 134 |
+
Fabryka AI, link to this repository; derivative weights under the same licence. Fine-tuning data keep their licences;
|
| 135 |
+
see *Training data* (PoQuAD: CC BY 4.0, attribution above).
|
| 136 |
+
|
| 137 |
+
## Po polsku (skrót)
|
| 138 |
+
|
| 139 |
+
GoLLeM v6 250M Instruct v2 to mały model do rozmowy po polsku i angielsku, dostrojony (SFT) z bazowego GoLLeM v6 250M
|
| 140 |
+
na ok. 9,3 tys. rozmów. Wersja badawcza (research preview): wita się, przedstawia i odpowiada na pytania do podanego
|
| 141 |
+
tekstu, ale się myli, nie liczy i nie ma dostępu do internetu. Od Instruct v1 różni się tym, że nie używa nazwiska
|
| 142 |
+
autora: na pytanie o twórcę odpowiada, że jest projektem Fabryki AI (wcześniej SlayerLab). Autora podaje ta karta.
|
| 143 |
+
|
| 144 |
+
---
|
| 145 |
+
*GoLLeM v6 250M Instruct v2 — Fabryka AI. **Author: Arkadiusz Słota.** Research preview.*
|
SHA256SUMS
ADDED
|
@@ -0,0 +1,6 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
6e97419dd8aea96a37f9efbbb3d577ad9da8c1735deb07963eda7eafc4f49aaa README.md
|
| 2 |
+
82a57236ca89151046ed16a1c463d60710a4b2c677708176c286156ef3bea412 chat_gollem_v6.py
|
| 3 |
+
61d8594034918d7e49253f3801da79bae9f14f57ba871fccde6d307a01c73a8f config.json
|
| 4 |
+
67ce9524e4f2f5716ab5322d867d81010563e2a41efd0344f3ad31152636418d model.safetensors
|
| 5 |
+
ee2ea8c4e957413f3a7c1e8fe93a0a0b4431a2bcabaa6cc5ceeaea103fcd58ee modeling_gollem_v6.py
|
| 6 |
+
0640d3bd3674d7a6e59c540d945a2ad275a12e1bcaeb6007c37a90751f9815dd tokenizer.json
|
chat_gollem_v6.py
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""GoLLeM v6 Instruct: chat helper (ChatML), matching the format and sampling used in training and evaluation.
|
| 2 |
+
|
| 3 |
+
import sys; sys.path.insert(0, path)
|
| 4 |
+
from modeling_gollem_v6 import load_gollem_v6
|
| 5 |
+
from chat_gollem_v6 import chat
|
| 6 |
+
model, tok = load_gollem_v6(path)
|
| 7 |
+
print(chat(model, tok, [("user", "Kim jesteś?")], seed=1))
|
| 8 |
+
|
| 9 |
+
Format: for every message `<|im_start|>{role}\n{content}<|im_end|>\n`, then `<|im_start|>assistant\n`; no system prompt;
|
| 10 |
+
generation stops at `<|im_end|>`. Message text is encoded as plain text (the strings "<|im_start|>"/"<|im_end|>" typed by
|
| 11 |
+
a user are NOT turned into control tokens). Default sampling: temperature 0.7, nucleus top-p 0.9 (as in our evaluation).
|
| 12 |
+
"""
|
| 13 |
+
from __future__ import annotations
|
| 14 |
+
|
| 15 |
+
import torch
|
| 16 |
+
|
| 17 |
+
IM_START, IM_END = 1, 2
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def encode_chat(tok, messages):
|
| 21 |
+
"""Token ids of a ChatML conversation ending with an open assistant turn. messages: [(role, text), ...]."""
|
| 22 |
+
prev = tok.encode_special_tokens
|
| 23 |
+
tok.encode_special_tokens = True
|
| 24 |
+
try:
|
| 25 |
+
text = lambda s: tok.encode(s, add_special_tokens=False).ids # noqa: E731
|
| 26 |
+
ids = []
|
| 27 |
+
for role, content in messages:
|
| 28 |
+
ids += [IM_START] + text(role + "\n") + text(content) + [IM_END] + text("\n")
|
| 29 |
+
return ids + [IM_START] + text("assistant\n")
|
| 30 |
+
finally:
|
| 31 |
+
tok.encode_special_tokens = prev
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
@torch.no_grad()
|
| 35 |
+
def chat(model, tok, messages, max_new_tokens=256, temperature=0.7, top_p=0.9, seed=None):
|
| 36 |
+
"""Assistant reply to `messages` ([(role, text), ...], roles "user"/"assistant"). Returns the reply text."""
|
| 37 |
+
g = torch.Generator().manual_seed(seed) if seed is not None else None
|
| 38 |
+
dev = next(model.parameters()).device
|
| 39 |
+
x = torch.tensor([encode_chat(tok, messages)], device=dev)
|
| 40 |
+
out = []
|
| 41 |
+
for _ in range(max_new_tokens):
|
| 42 |
+
logits = model(x[:, -model.block:])[0, -1].float()
|
| 43 |
+
p = torch.softmax(logits / temperature, -1)
|
| 44 |
+
sp, si = torch.sort(p, descending=True)
|
| 45 |
+
sp = sp * (torch.cumsum(sp, 0) - sp < top_p)
|
| 46 |
+
nxt = int(si[torch.multinomial((sp / sp.sum()).cpu(), 1, generator=g)])
|
| 47 |
+
if nxt == IM_END:
|
| 48 |
+
break
|
| 49 |
+
out.append(nxt)
|
| 50 |
+
x = torch.cat([x, torch.tensor([[nxt]], device=dev)], 1)
|
| 51 |
+
return tok.decode(out)
|
config.json
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"block": 1024,
|
| 3 |
+
"ffn": "swiglu",
|
| 4 |
+
"ffn_mult": 2.667,
|
| 5 |
+
"logit_cap": 0.0,
|
| 6 |
+
"n_embd": 960,
|
| 7 |
+
"n_head": 15,
|
| 8 |
+
"n_layer": 20,
|
| 9 |
+
"norm": "rmsnorm",
|
| 10 |
+
"norm_eps": 1e-06,
|
| 11 |
+
"pos": "rope",
|
| 12 |
+
"qk_norm": true,
|
| 13 |
+
"rope_theta": 100000.0,
|
| 14 |
+
"tied_weights": {
|
| 15 |
+
"tok.weight": "head.weight"
|
| 16 |
+
},
|
| 17 |
+
"value_residual": true,
|
| 18 |
+
"vocab": 32768
|
| 19 |
+
}
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:67ce9524e4f2f5716ab5322d867d81010563e2a41efd0344f3ad31152636418d
|
| 3 |
+
size 1136891540
|
modeling_gollem_v6.py
ADDED
|
@@ -0,0 +1,169 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""GoLLeM v6 model code (inference only), extracted from the training script used for this run.
|
| 2 |
+
|
| 3 |
+
The module definitions are those of the training script; only the weight initialisation and the training loss were
|
| 4 |
+
removed and comments translated. Weights load with strict=True. Requirements: torch, safetensors, tokenizers.
|
| 5 |
+
|
| 6 |
+
from modeling_gollem_v6 import load_gollem_v6, generate
|
| 7 |
+
model, tok = load_gollem_v6(".") # directory with config.json, model.safetensors, tokenizer.json
|
| 8 |
+
print(generate(model, tok, "Najwyższym szczytem Polski jest", max_new_tokens=40))
|
| 9 |
+
|
| 10 |
+
Notes: RoPE uses the interleaved convention (pairs of adjacent channels), not rotate_half; the output head is tied to
|
| 11 |
+
the token embedding; context length is 1,024 tokens.
|
| 12 |
+
"""
|
| 13 |
+
from __future__ import annotations
|
| 14 |
+
|
| 15 |
+
import json
|
| 16 |
+
import os
|
| 17 |
+
import types
|
| 18 |
+
|
| 19 |
+
import torch
|
| 20 |
+
import torch.nn as nn
|
| 21 |
+
import torch.nn.functional as F
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class RMSNorm(nn.Module):
|
| 25 |
+
"""RMSNorm with fp32 compute for stability."""
|
| 26 |
+
def __init__(self, d, eps=1e-6):
|
| 27 |
+
super().__init__()
|
| 28 |
+
self.weight = nn.Parameter(torch.ones(d))
|
| 29 |
+
self.eps = eps
|
| 30 |
+
|
| 31 |
+
def forward(self, x):
|
| 32 |
+
return x * torch.rsqrt(x.float().pow(2).mean(-1, keepdim=True) + self.eps).to(x.dtype) * self.weight
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def make_norm(d, cfg):
|
| 36 |
+
return RMSNorm(d, cfg.norm_eps) if cfg.norm == "rmsnorm" else nn.LayerNorm(d)
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def apply_rope(x, base=100000.0):
|
| 40 |
+
"""Parameter-free RoPE on [B, H, T, D], interleaved convention (even/odd channel pairs)."""
|
| 41 |
+
_, _, T, dim = x.shape
|
| 42 |
+
pos = torch.arange(T, device=x.device, dtype=torch.float32)
|
| 43 |
+
freq = 1.0 / (base ** (torch.arange(0, dim, 2, device=x.device, dtype=torch.float32) / dim))
|
| 44 |
+
ang = torch.outer(pos, freq)
|
| 45 |
+
cos, sin = ang.cos().to(x.dtype)[None, None], ang.sin().to(x.dtype)[None, None]
|
| 46 |
+
even, odd = x[..., ::2], x[..., 1::2]
|
| 47 |
+
return torch.stack((even * cos - odd * sin, even * sin + odd * cos), dim=-1).flatten(-2)
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
class SwiGLU(nn.Module):
|
| 51 |
+
"""Gated MLP: down(silu(gate(x)) * up(x))."""
|
| 52 |
+
def __init__(self, d, hidden):
|
| 53 |
+
super().__init__()
|
| 54 |
+
self.gate = nn.Linear(d, hidden, bias=False)
|
| 55 |
+
self.up = nn.Linear(d, hidden, bias=False)
|
| 56 |
+
self.down = nn.Linear(hidden, d, bias=False)
|
| 57 |
+
|
| 58 |
+
def forward(self, x):
|
| 59 |
+
return self.down(F.silu(self.gate(x)) * self.up(x))
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
class Block(nn.Module):
|
| 63 |
+
def __init__(self, d, nh, block, cfg, is_first=False):
|
| 64 |
+
super().__init__()
|
| 65 |
+
self.ln1 = make_norm(d, cfg)
|
| 66 |
+
self.ln2 = make_norm(d, cfg)
|
| 67 |
+
self.qkv = nn.Linear(d, 3 * d)
|
| 68 |
+
self.proj = nn.Linear(d, d)
|
| 69 |
+
if cfg.ffn == "swiglu":
|
| 70 |
+
self.mlp = SwiGLU(d, int(round(cfg.ffn_mult * d)))
|
| 71 |
+
else:
|
| 72 |
+
self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d))
|
| 73 |
+
self.nh = nh
|
| 74 |
+
self.d = d
|
| 75 |
+
self.cfg = cfg
|
| 76 |
+
self.is_first = is_first
|
| 77 |
+
if cfg.value_residual and not is_first:
|
| 78 |
+
self.vr_lambda = nn.Parameter(torch.zeros(1))
|
| 79 |
+
if cfg.qk_norm:
|
| 80 |
+
hd = d // nh
|
| 81 |
+
self.q_norm = RMSNorm(hd, cfg.norm_eps)
|
| 82 |
+
self.k_norm = RMSNorm(hd, cfg.norm_eps)
|
| 83 |
+
|
| 84 |
+
def forward(self, x, v0=None):
|
| 85 |
+
B, T, D = x.size()
|
| 86 |
+
h = self.ln1(x)
|
| 87 |
+
q, k, v = self.qkv(h).split(self.d, dim=2)
|
| 88 |
+
hd = D // self.nh
|
| 89 |
+
q = q.view(B, T, self.nh, hd).transpose(1, 2)
|
| 90 |
+
k = k.view(B, T, self.nh, hd).transpose(1, 2)
|
| 91 |
+
v = v.view(B, T, self.nh, hd).transpose(1, 2)
|
| 92 |
+
if self.cfg.qk_norm:
|
| 93 |
+
q = self.q_norm(q)
|
| 94 |
+
k = self.k_norm(k)
|
| 95 |
+
if self.cfg.pos == "rope":
|
| 96 |
+
q = apply_rope(q, self.cfg.rope_theta)
|
| 97 |
+
k = apply_rope(k, self.cfg.rope_theta)
|
| 98 |
+
if self.cfg.value_residual:
|
| 99 |
+
if self.is_first:
|
| 100 |
+
v0 = v
|
| 101 |
+
else:
|
| 102 |
+
v = v + self.vr_lambda * v0
|
| 103 |
+
y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
|
| 104 |
+
y = y.transpose(1, 2).contiguous().view(B, T, D)
|
| 105 |
+
x = x + self.proj(y)
|
| 106 |
+
x = x + self.mlp(self.ln2(x))
|
| 107 |
+
return x, v0
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
class GPT(nn.Module):
|
| 111 |
+
def __init__(self, vocab, n_layer, n_embd, n_head, block, cfg):
|
| 112 |
+
super().__init__()
|
| 113 |
+
self.cfg = cfg
|
| 114 |
+
self.tok = nn.Embedding(vocab, n_embd)
|
| 115 |
+
self.use_rope = cfg.pos == "rope"
|
| 116 |
+
if not self.use_rope:
|
| 117 |
+
self.pos = nn.Embedding(block, n_embd)
|
| 118 |
+
self.blocks = nn.ModuleList([Block(n_embd, n_head, block, cfg, is_first=(i == 0)) for i in range(n_layer)])
|
| 119 |
+
self.lnf = make_norm(n_embd, cfg)
|
| 120 |
+
self.head = nn.Linear(n_embd, vocab, bias=False)
|
| 121 |
+
self.head.weight = self.tok.weight # tied
|
| 122 |
+
self.block = block
|
| 123 |
+
|
| 124 |
+
def forward(self, idx):
|
| 125 |
+
B, T = idx.size()
|
| 126 |
+
x = self.tok(idx)
|
| 127 |
+
if not self.use_rope:
|
| 128 |
+
pos = torch.arange(T, device=idx.device)
|
| 129 |
+
x = x + self.pos(pos)[None]
|
| 130 |
+
v0 = None
|
| 131 |
+
for b in self.blocks:
|
| 132 |
+
x, v0 = b(x, v0)
|
| 133 |
+
logits = self.head(self.lnf(x))
|
| 134 |
+
cap = getattr(self.cfg, "logit_cap", 0.0)
|
| 135 |
+
if cap and cap > 0:
|
| 136 |
+
logits = cap * torch.tanh(logits / cap)
|
| 137 |
+
return logits
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
def load_gollem_v6(path, device="cpu"):
|
| 141 |
+
"""Model (eval mode) and tokenizer from a directory with config.json, model.safetensors and tokenizer.json."""
|
| 142 |
+
from safetensors.torch import load_file
|
| 143 |
+
from tokenizers import Tokenizer
|
| 144 |
+
|
| 145 |
+
with open(os.path.join(path, "config.json"), encoding="utf-8") as f:
|
| 146 |
+
cfg = types.SimpleNamespace(**json.load(f))
|
| 147 |
+
model = GPT(cfg.vocab, cfg.n_layer, cfg.n_embd, cfg.n_head, cfg.block, cfg)
|
| 148 |
+
model.load_state_dict(load_file(os.path.join(path, "model.safetensors")), strict=True)
|
| 149 |
+
return model.to(device).eval(), Tokenizer.from_file(os.path.join(path, "tokenizer.json"))
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
@torch.no_grad()
|
| 153 |
+
def generate(model, tok, prompt, max_new_tokens=64, temperature=0.0, top_k=None, seed=None):
|
| 154 |
+
"""Continue `prompt`. temperature 0 = greedy; otherwise sampling, optionally top-k."""
|
| 155 |
+
g = torch.Generator().manual_seed(seed) if seed is not None else None
|
| 156 |
+
dev = next(model.parameters()).device
|
| 157 |
+
x = torch.tensor([tok.encode(prompt).ids], device=dev)
|
| 158 |
+
for _ in range(max_new_tokens):
|
| 159 |
+
logits = model(x[:, -model.block:])[:, -1].float().cpu()
|
| 160 |
+
if temperature <= 0:
|
| 161 |
+
nxt = logits.argmax(-1, keepdim=True)
|
| 162 |
+
else:
|
| 163 |
+
logits = logits / temperature
|
| 164 |
+
if top_k:
|
| 165 |
+
kth = torch.topk(logits, top_k).values[:, -1:]
|
| 166 |
+
logits = logits.masked_fill(logits < kth, float("-inf"))
|
| 167 |
+
nxt = torch.multinomial(torch.softmax(logits, -1), 1, generator=g)
|
| 168 |
+
x = torch.cat([x, nxt.to(dev)], dim=1)
|
| 169 |
+
return tok.decode(x[0].tolist())
|
tokenizer.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|