Spaces:
Runtime error
Runtime error
| 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>"] | |
| 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() | |