File size: 4,953 Bytes
d7228c8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
"""
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)