Spaces:
Runtime error
Runtime error
File size: 3,274 Bytes
3781007 | 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 | import argparse
from pathlib import Path
import torch
from superlillm.model import ModelConfig, SuperLilLM
from superlillm.tokenizer import WordTokenizer
CHECKPOINT = Path("checkpoints/superlillm.pt")
TOKENIZER = Path("checkpoints/tokenizer.json")
class ChatEngine:
def __init__(self, checkpoint=CHECKPOINT, tokenizer_path=TOKENIZER):
self.device = "mps" if torch.backends.mps.is_available() else "cuda" if torch.cuda.is_available() else "cpu"
self.tokenizer = WordTokenizer.load(tokenizer_path)
payload = torch.load(checkpoint, map_location=self.device)
config = ModelConfig(**payload["config"])
self.model = SuperLilLM(config).to(self.device)
self.model.load_state_dict(payload["model_state"])
self.model.eval()
self.eos_id = self.tokenizer.token_to_id["<eos>"]
@torch.no_grad()
def generate(self, message, max_new_tokens=80, temperature=0.45, top_k=1):
message = message.replace("’", "'").replace("“", '"').replace("”", '"').strip()
prompt = f"User: {message}\nAssistant:"
ids = self.tokenizer.encode(prompt, add_bos=True)
generated = ids[:]
generated_token_ids = []
for _ in range(max_new_tokens):
context = generated[-self.model.config.block_size :]
x = torch.tensor([context], dtype=torch.long, device=self.device)
logits, _ = self.model(x)
logits = logits[0, -1] / max(temperature, 0.05)
if top_k:
values, _ = torch.topk(logits, min(top_k, logits.numel()))
logits[logits < values[-1]] = -float("inf")
probs = torch.softmax(logits, dim=-1)
next_id = int(torch.multinomial(probs, num_samples=1).item())
generated.append(next_id)
generated_token_ids.append(next_id)
if next_id == self.eos_id:
break
answer_ids = generated[len(ids) :]
answer = self.tokenizer.decode(answer_ids)
if " user:" in answer:
answer = answer.split(" user:", 1)[0]
if " assistant:" in answer:
answer = answer.split(" assistant:", 1)[0]
answer = answer.replace("<eos>", "").strip() or "I am not sure yet, but I can try a simpler answer."
tokens = [
{
"id": token_id,
"text": self.tokenizer.id_to_token.get(token_id, "<unk>"),
}
for token_id in generated_token_ids
if token_id != self.eos_id
]
return {"reply": answer, "tokens": tokens}
def reply(self, message, **kwargs):
return self.generate(message, **kwargs)["reply"]
def main():
parser = argparse.ArgumentParser()
parser.add_argument("message", nargs="*", help="Message to send. Leave empty for interactive chat.")
args = parser.parse_args()
engine = ChatEngine()
if args.message:
result = engine.generate(" ".join(args.message))
print(result["reply"])
return
print("SuperLilLM chat. Type 'exit' to stop.")
while True:
text = input("you> ").strip()
if text.lower() in {"exit", "quit"}:
break
print("bot>", engine.reply(text))
if __name__ == "__main__":
main()
|