1#!/usr/bin/env python3
2"""
3Tiny Mistral REPL demo — streaming tokens (TextStreamer if available, else manual sampling).
4Commands: :quit, :help, :show, :set <param> <value> (max_new_tokens, temperature, top_p, full_output)
5"""
6from __future__ import annotations
7import shlex
8import time
9import torch
10from typing import Optional
11
12from transformers import AutoTokenizer, MistralForCausalLM
13
14# --------- CONFIG ----------
15MODEL_DIR = "Harley-ml/TinyWord-134k"
16TOKENIZER_DIR = MODEL_DIR
17DEFAULT_MAX_NEW_TOKENS = 8 # I don't reccomend going higher than this
18DEFAULT_TEMPERATURE = 0.4
19DEFAULT_TOP_P = 0.9
20DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
21PROMPT = ">>> "
22# ---------------------------
23
24def load_tokenizer(path: str):
25 print("Loading tokenizer...", path)
26 tok = AutoTokenizer.from_pretrained(path, use_fast=True, local_files_only=False)
27 if tok.pad_token is None:
28 if getattr(tok, "eos_token", None) is not None:
29 tok.add_special_tokens({"pad_token": tok.eos_token})
30 else:
31 tok.add_special_tokens({"pad_token": "<pad>", "eos_token": "</s>"})
32 print("Tokenizer ready. vocab_size=", getattr(tok, "vocab_size", "N/A"))
33 return tok
34
35def load_model(path: str, device: str):
36 print("Loading model...", path)
37 model = None
38 try:
39 desired_dtype = torch.float16 if device.startswith("cuda") else torch.float32
40 model = MistralForCausalLM.from_pretrained(path, local_files_only=False, dtype=desired_dtype)
41 print("Loaded with dtype arg.")
42 except TypeError:
43 model = MistralForCausalLM.from_pretrained(path, local_files_only=False)
44 print("Loaded without dtype; will convert.")
45 except Exception as e:
46 print("Load warning, retrying without dtype:", e)
47 model = MistralForCausalLM.from_pretrained(path, local_files_only=False)
48
49 try:
50 model.to(device)
51 if device.startswith("cuda") and next(model.parameters()).dtype != torch.float16:
52 model.half()
53 if not device.startswith("cuda") and next(model.parameters()).dtype != torch.float32:
54 model.to(torch.float32)
55 except Exception as e:
56 print("Model move/convert warning:", e)
57
58 model.config.pad_token_id = getattr(model.config, "pad_token_id", None)
59 model.eval()
60 return model
61
62# Simple nucleus/top-p filtering for a single logits vector
63def top_p_filtering(logits: torch.Tensor, top_p: float, min_keep: int = 1) -> torch.Tensor:
64 if top_p <= 0 or top_p >= 1.0:
65 return logits
66 sorted_logits, sorted_idx = torch.sort(logits, descending=True)
67 probs = torch.softmax(sorted_logits, dim=-1)
68 cumprobs = torch.cumsum(probs, dim=-1)
69 cutoff = (cumprobs > top_p).nonzero(as_tuple=False)
70 if cutoff.numel() > 0:
71 idx = int(cutoff[0].item())
72 cutoff_idx = max(idx + 1, min_keep)
73 else:
74 cutoff_idx = sorted_logits.size(-1)
75 mask = torch.ones_like(sorted_logits, dtype=torch.bool)
76 mask[cutoff_idx:] = False
77 filtered = sorted_logits.masked_fill(~mask, -float("inf"))
78 return torch.empty_like(filtered).scatter_(0, sorted_idx, filtered)
79
80# Manual streaming generator (single-batch)
81def manual_stream_generate(model, tokenizer, prompt: str, device: str,
82 max_new_tokens: int = 64, temperature: float = 1.0, top_p: float = 0.9,
83 eos_token_id: Optional[int] = None):
84 inputs = tokenizer(prompt, return_tensors="pt", add_special_tokens=False)
85 input_ids = inputs["input_ids"].to(device)
86 attention_mask = inputs.get("attention_mask", None)
87 if attention_mask is not None:
88 attention_mask = attention_mask.to(device)
89
90 past = None
91 with torch.no_grad():
92 out = model(input_ids=input_ids, attention_mask=attention_mask, use_cache=True)
93 past = getattr(out, "past_key_values", None)
94
95 # start sampling tokens
96 next_input = input_ids[:, -1:].to(device) if past is not None else input_ids.to(device)
97 for _ in range(max_new_tokens):
98 with torch.no_grad():
99 out = model(input_ids=next_input, past_key_values=past, use_cache=True)
100 logits = out.logits[:, -1, :] # (batch, vocab)
101 past = getattr(out, "past_key_values", past)
102
103 if temperature != 1.0:
104 logits = logits / max(temperature, 1e-8)
105
106 filtered = top_p_filtering(logits[0].cpu(), top_p).to(device)
107 probs = torch.nn.functional.softmax(filtered.unsqueeze(0), dim=-1)
108 next_token = torch.multinomial(probs, num_samples=1)
109 token_id = int(next_token[0, 0].item())
110
111 token_text = tokenizer.decode([token_id], clean_up_tokenization_spaces=False)
112 yield token_id, token_text
113
114 if eos_token_id is not None and token_id == eos_token_id:
115 break
116 next_input = torch.tensor([[token_id]], dtype=torch.long, device=device)
117
118def has_text_streamer():
119 try:
120 from transformers import TextStreamer # type: ignore
121 return True
122 except Exception:
123 return False
124
125# tiny REPL state
126class State:
127 def __init__(self):
128 self.max_new_tokens = DEFAULT_MAX_NEW_TOKENS
129 self.temperature = DEFAULT_TEMPERATURE
130 self.top_p = DEFAULT_TOP_P
131 self.full_output = False
132 self.stream = True
133
134def handle_generation(model, tokenizer, prompt: str, device: str, state: State):
135 eos = getattr(tokenizer, "eos_token_id", None)
136 try:
137 if has_text_streamer():
138 from transformers import TextStreamer
139 streamer = TextStreamer(tokenizer, skip_prompt=not state.full_output, skip_special_tokens=True)
140 inputs = tokenizer(prompt, return_tensors="pt", truncation=True, add_special_tokens=False)
141 inputs = {k: v.to(device) for k, v in inputs.items() if isinstance(v, torch.Tensor)}
142 inputs.pop("token_type_ids", None)
143 model.generate(**inputs,
144 max_new_tokens=state.max_new_tokens,
145 do_sample=True,
146 temperature=state.temperature,
147 top_p=state.top_p,
148 pad_token_id=tokenizer.pad_token_id,
149 eos_token_id=tokenizer.eos_token_id,
150 streamer=streamer)
151 print("") # newline after streamer
152 return
153 # fallback: manual streaming
154 gen = manual_stream_generate(model, tokenizer, prompt, device,
155 max_new_tokens=state.max_new_tokens,
156 temperature=state.temperature,
157 top_p=state.top_p,
158 eos_token_id=eos)
159 if state.full_output:
160 print("PROMPT:", prompt)
161 print("GENERATING:", end=" ", flush=True)
162 else:
163 print("GENERATING:", end=" ", flush=True)
164
165 count = 0
166 t0 = time.time()
167 for _tok_id, tok_text in gen:
168 count += 1
169 print(tok_text, end="", flush=True)
170 print()
171 print(f"(generated {count} tokens in {time.time()-t0:.2f}s)")
172 except KeyboardInterrupt:
173 print("\n[interrupted] Generation aborted by user.")
174 except Exception as e:
175 print("Generation error:", e)
176
177def repl(model, tokenizer, device):
178 state = State()
179 help_text = (
180 "Commands:\n"
181 " :quit\n"
182 " :help\n"
183 " :show\n"
184 " :set <param> <value> # params: max_new_tokens, temperature, top_p, full_output, stream\n"
185 " (blank line repeats last prompt)\n"
186 )
187 print("Tiny Mistral REPL — device:", device)
188 print(help_text)
189 last = ""
190 while True:
191 try:
192 raw = input(PROMPT).strip()
193 except (EOFError, KeyboardInterrupt):
194 print("\nExiting.")
195 break
196 if not raw:
197 raw = last
198 if not raw:
199 continue
200
201 if raw.startswith(":"):
202 toks = shlex.split(raw)
203 cmd = toks[0].lower()
204 if cmd == ":quit":
205 print("bye.")
206 break
207 if cmd == ":help":
208 print(help_text); continue
209 if cmd == ":show":
210 print(f"max_new_tokens={state.max_new_tokens}, temperature={state.temperature}, top_p={state.top_p}, full_output={state.full_output}, stream={state.stream}")
211 continue
212 if cmd == ":set":
213 if len(toks) < 3:
214 print("usage: :set <param> <value>"); continue
215 k, v = toks[1], toks[2]
216 try:
217 if k == "max_new_tokens":
218 state.max_new_tokens = int(v)
219 elif k == "temperature":
220 state.temperature = float(v)
221 elif k == "top_p":
222 state.top_p = float(v)
223 elif k in ("full_output", "full"):
224 state.full_output = v.lower() in ("1", "true", "yes", "y")
225 elif k == "stream":
226 state.stream = v.lower() in ("1", "true", "yes", "y")
227 else:
228 print("unknown param:", k)
229 continue
230 print("OK.")
231 except Exception as e:
232 print("set error:", e)
233 continue
234 print("unknown command")
235 continue
236
237 last = raw
238 if state.stream:
239 handle_generation(model, tokenizer, raw, device, state)
240 else:
241 # non-streaming generate
242 try:
243 inputs = tokenizer(raw, return_tensors="pt", truncation=True, add_special_tokens=False)
244 inputs = {k: v.to(device) for k, v in inputs.items() if isinstance(v, torch.Tensor)}
245 inputs.pop("token_type_ids", None)
246 out = model.generate(**inputs,
247 max_new_tokens=state.max_new_tokens,
248 do_sample=True,
249 temperature=state.temperature,
250 top_p=state.top_p,
251 pad_token_id=tokenizer.pad_token_id,
252 eos_token_id=tokenizer.eos_token_id)
253 seq = out[0]
254 input_len = inputs["input_ids"].shape[1] if "input_ids" in inputs else 0
255 text = tokenizer.decode(seq if state.full_output else seq[input_len:], skip_special_tokens=True)
256 print("\nOUTPUT\n", text)
257 except Exception as e:
258 print("Generation failed:", e)
259
260def main():
261 device = DEVICE
262 tokenizer = load_tokenizer(TOKENIZER_DIR)
263 model = load_model(MODEL_DIR, device)
264 repl(model, tokenizer, device)
265
266if __name__ == "__main__":
267 main()