Law_Slm / scripts /chat_run.py
Amitkumar001's picture
Upload folder using huggingface_hub
d7228c8 verified
Raw History Blame Contribute Delete
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)