InMecha's picture
download
raw
3.55 kB
"""
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()

Xet Storage Details

Size:
3.55 kB
·
Xet hash:
674c70ae66240d95402879dcbd681b567a98b6c259b068bfc83a474f9d128f89

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.