"""
Simple interactive chat with a local safetensors model (streaming).
Usage: python chat_local.py --model_path C:\path\to\model\directory
"""

import argparse
import sys
import torch
from threading import Thread
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--model_path", type=str, required=True, help="Path to local model directory")
    parser.add_argument("--max_new_tokens", type=int, default=4096)
    parser.add_argument("--temperature", type=float, default=0.8)
    parser.add_argument("--top_p", type=float, default=0.9)
    parser.add_argument("--min_p", type=float, default=0.05)
    parser.add_argument("--repetition_penalty", type=float, default=1.1,
                        help="Penalty for repeated tokens. 1.0=off, 1.05-1.15=typical range")
    parser.add_argument("--no_repeat_ngram_size", type=int, default=4,
                        help="Block any n-gram from repeating. 0=off, 3-5=typical range")
    args = parser.parse_args()

    print(f"Loading model from {args.model_path} ...")
    tokenizer = AutoTokenizer.from_pretrained(args.model_path, local_files_only=True)
    model = AutoModelForCausalLM.from_pretrained(
        args.model_path,
        local_files_only=True,
        torch_dtype=torch.float16,
        device_map="auto",
    )
    print(f"Model loaded.")
    print(f"  max_new_tokens={args.max_new_tokens}  temperature={args.temperature}")
    print(f"  top_p={args.top_p}  min_p={args.min_p}")
    print(f"  repetition_penalty={args.repetition_penalty}  no_repeat_ngram_size={args.no_repeat_ngram_size}")
    print(f"Type 'quit' to exit, 'clear' to reset history.\n")

    messages = []

    while True:
        try:
            user_input = input("You: ").strip()
        except (EOFError, KeyboardInterrupt):
            print("\nBye.")
            break

        if not user_input:
            continue
        if user_input.lower() == "quit":
            break
        if user_input.lower() == "clear":
            messages.clear()
            print("-- history cleared --\n")
            continue

        messages.append({"role": "user", "content": user_input})

        prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
        inputs = tokenizer(prompt, return_tensors="pt").to(model.device)

        streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)

        generate_kwargs = dict(
            **inputs,
            max_new_tokens=args.max_new_tokens,
            temperature=args.temperature,
            top_p=args.top_p,
            min_p=args.min_p,
            do_sample=True,
            pad_token_id=tokenizer.eos_token_id,
            repetition_penalty=args.repetition_penalty,
            no_repeat_ngram_size=args.no_repeat_ngram_size,
            streamer=streamer,
        )

        thread = Thread(target=model.generate, kwargs=generate_kwargs)
        thread.start()

        sys.stdout.write("\nAssistant: ")
        sys.stdout.flush()

        reply_chunks = []
        for chunk in streamer:
            sys.stdout.write(chunk)
            sys.stdout.flush()
            reply_chunks.append(chunk)

        thread.join()

        reply = "".join(reply_chunks).strip()
        messages.append({"role": "assistant", "content": reply})
        print("\n")

if __name__ == "__main__":
    main()