Download scripts/chat_run.py from Amitkumar001/Law_Slm: direct link, hf CLI and curl.
- Browser
- Download file 4.95 kB
-
https://huggingface.co/Amitkumar001/Law_Slm/resolve/main/scripts/chat_run.py
- Command line
-
hf download hf://Amitkumar001/Law_Slm/scripts/chat_run.py
-
curl -L -o chat_run.py https://huggingface.co/Amitkumar001/Law_Slm/resolve/main/scripts/chat_run.py
4.95 kB
| """ | |
| Interactive REPL Chat script for asking questions and chatting with the Small Language Model. | |
| """ | |
| import os | |
| import sys | |
| import glob | |
| from typing import Optional | |
| sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))) | |
| import torch | |
| from slm.config.model_config import ModelConfig | |
| from slm.model.transformer_lm import SLMForCausalLM | |
| from slm.tokenizer.bpe import BPETokenizer | |
| from slm.sampling.generator import TextGenerator | |
| from slm.checkpoint.manager import CheckpointManager | |
| from slm.utils.logger import get_logger | |
| logger = get_logger("slm.chat") | |
| def resolve_checkpoint_path(target_path: Optional[str]) -> Optional[str]: | |
| """Resolves checkpoint file path from file path, directory, or search folders.""" | |
| if target_path: | |
| if os.path.isfile(target_path): | |
| return target_path | |
| if os.path.isdir(target_path): | |
| pts = sorted(glob.glob(os.path.join(target_path, "*.pt")), key=os.path.getmtime) | |
| if pts: | |
| return pts[-1] | |
| if os.path.isfile("checkpoints/best_model.pt"): | |
| return "checkpoints/best_model.pt" | |
| search_dirs = ["checkpoints", "checkpoints_nano", "checkpoints_pipeline", "checkpoints_micro", "checkpoints_base"] | |
| for sdir in search_dirs: | |
| if os.path.exists(sdir): | |
| pts = sorted(glob.glob(os.path.join(sdir, "*.pt")), key=os.path.getmtime) | |
| if pts: | |
| return pts[-1] | |
| return None | |
| def start_chat(checkpoint_path: Optional[str] = None) -> None: | |
| """ | |
| Launches an interactive console terminal chat interface. | |
| """ | |
| resolved_path = resolve_checkpoint_path(checkpoint_path) | |
| if resolved_path: | |
| ckpt_dir = os.path.dirname(resolved_path) | |
| logger.info(f"Loading model checkpoint from {resolved_path}...") | |
| try: | |
| ckpt_data = torch.load(resolved_path, map_location="cpu", weights_only=False) | |
| except Exception: | |
| ckpt_data = torch.load(resolved_path, map_location="cpu") | |
| if isinstance(ckpt_data, dict) and "model_config" in ckpt_data: | |
| config = ModelConfig.from_dict(ckpt_data["model_config"]) | |
| else: | |
| config = ModelConfig(vocab_size=2000, d_model=128, n_heads=4, n_layers=2) | |
| model = SLMForCausalLM(config) | |
| manager = CheckpointManager(output_dir=ckpt_dir) | |
| manager.load_checkpoint(resolved_path, model) | |
| tok_dir = os.path.join(ckpt_dir, "tokenizer") | |
| if os.path.exists(tok_dir): | |
| tokenizer = BPETokenizer.load(tok_dir) | |
| else: | |
| logger.warning(f"Tokenizer directory not found at {tok_dir}. Training fallback tokenizer...") | |
| tokenizer = BPETokenizer() | |
| tokenizer.train_on_texts(["Interactive chat training text sample for tokenizer setup."], vocab_size=config.vocab_size) | |
| else: | |
| logger.warning("No checkpoint file found in workspace! Initializing active SLM model for demo chat session...") | |
| config = ModelConfig(vocab_size=2000, d_model=128, n_heads=4, n_layers=2, d_ff=512) | |
| model = SLMForCausalLM(config) | |
| tokenizer = BPETokenizer() | |
| corpus = [ | |
| "User: What is a Small Language Model?\nSLM: A Small Language Model is an efficient decoder-only transformer network.", | |
| "User: How does self-attention work?\nSLM: Self-attention computes scaled dot-product matrix operations over queries, keys, and values.", | |
| "User: Hello!\nSLM: Hello! How can I assist you with language modeling today?" | |
| ] * 10 | |
| tokenizer.train_on_texts(corpus, vocab_size=2000) | |
| generator = TextGenerator(model, tokenizer) | |
| print("\n" + "=" * 65) | |
| print(" LawSLM INTERACTIVE ASSISTANT (Built Completely From Scratch)") | |
| print(" Role: Legal Information, General AI, Programming & Analysis") | |
| print("=" * 65) | |
| print("Type your question/prompt below. Type 'exit', 'quit', or 'q' to end session.") | |
| print("=" * 65 + "\n") | |
| while True: | |
| try: | |
| user_input = input("\nUser > ").strip() | |
| if not user_input: | |
| continue | |
| if user_input.lower() in ("exit", "quit", "q"): | |
| print("\nEnding chat session. Goodbye!") | |
| break | |
| prompt = f"User: {user_input}\nSLM:" | |
| print("SLM > ", end="", flush=True) | |
| def stream_callback(token_str: str): | |
| print(token_str, end="", flush=True) | |
| generator.generate( | |
| prompt=prompt, | |
| max_new_tokens=100, | |
| temperature=0.0, | |
| top_k=1, | |
| top_p=1.0, | |
| repetition_penalty=1.05, | |
| stream_callback=stream_callback | |
| ) | |
| print() | |
| except KeyboardInterrupt: | |
| print("\nChat session interrupted. Goodbye!") | |
| break | |
| if __name__ == "__main__": | |
| ckpt = sys.argv[1] if len(sys.argv) > 1 else None | |
| start_chat(ckpt) | |