Maggio33 commited on
Commit
085b7b5
·
verified ·
1 Parent(s): a10874f

GoLLeM-v6-250M-Instruct-v2: weights, config, tokenizer, model and chat code, card

Browse files
Files changed (7) hide show
  1. README.md +145 -0
  2. SHA256SUMS +6 -0
  3. chat_gollem_v6.py +51 -0
  4. config.json +19 -0
  5. model.safetensors +3 -0
  6. modeling_gollem_v6.py +169 -0
  7. 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