PertMind / tui_chat.py
lukatang
Publish PertMind model
14a19cc
Raw
History Blame Contribute Delete
5.11 kB
#!/usr/bin/env python3
"""A small terminal chat UI for PertMind."""
from __future__ import annotations
import argparse
from dataclasses import dataclass
DEFAULT_SYSTEM_PROMPT = (
"You are PertMind, a biomedical assistant. For biomedical prediction, "
"screen-ranking, or gene-set interpretation tasks, answer first and then "
"provide a concise explanation. Use this style when applicable:\n"
"Final Answer: <answer>\nExplanation: <brief explanation>"
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--model", default=".", help="Path or Hugging Face model id.")
parser.add_argument("--backend", choices=["vllm", "transformers"], default="vllm")
parser.add_argument("--system-prompt", default=DEFAULT_SYSTEM_PROMPT)
parser.add_argument("--max-new-tokens", type=int, default=768)
parser.add_argument("--temperature", type=float, default=0.0)
parser.add_argument("--top-p", type=float, default=0.95)
parser.add_argument("--max-model-len", type=int, default=12288)
parser.add_argument("--gpu-memory-utilization", type=float, default=0.85)
return parser.parse_args()
def render_chat(tokenizer, messages: list[dict[str, str]]) -> str:
try:
return tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
except TypeError:
return tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
@dataclass
class VllmBackend:
model: str
max_model_len: int
gpu_memory_utilization: float
def __post_init__(self) -> None:
from vllm import LLM
self.llm = LLM(
model=self.model,
trust_remote_code=True,
dtype="bfloat16",
max_model_len=self.max_model_len,
gpu_memory_utilization=self.gpu_memory_utilization,
)
self.tokenizer = self.llm.get_tokenizer()
def generate(self, messages: list[dict[str, str]], max_new_tokens: int, temperature: float, top_p: float) -> str:
from vllm import SamplingParams
prompt = render_chat(self.tokenizer, messages)
params = SamplingParams(temperature=temperature, top_p=top_p, max_tokens=max_new_tokens)
return self.llm.generate([prompt], params)[0].outputs[0].text.strip()
@dataclass
class TransformersBackend:
model: str
def __post_init__(self) -> None:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
self.torch = torch
self.tokenizer = AutoTokenizer.from_pretrained(self.model, trust_remote_code=True)
self.llm = AutoModelForCausalLM.from_pretrained(
self.model,
torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
device_map="auto",
trust_remote_code=True,
)
def generate(self, messages: list[dict[str, str]], max_new_tokens: int, temperature: float, top_p: float) -> str:
prompt = render_chat(self.tokenizer, messages)
inputs = self.tokenizer(prompt, return_tensors="pt").to(self.llm.device)
do_sample = temperature > 0
outputs = self.llm.generate(
**inputs,
max_new_tokens=max_new_tokens,
do_sample=do_sample,
temperature=temperature if do_sample else None,
top_p=top_p if do_sample else None,
pad_token_id=self.tokenizer.eos_token_id,
)
generated = outputs[0, inputs["input_ids"].shape[-1] :]
return self.tokenizer.decode(generated, skip_special_tokens=True).strip()
def print_panel(title: str, text: str) -> None:
line = "=" * min(88, max(20, len(title) + 8))
print(f"\n{line}\n{title}\n{line}\n{text}\n")
def main() -> int:
args = parse_args()
if args.backend == "vllm":
backend = VllmBackend(args.model, args.max_model_len, args.gpu_memory_utilization)
else:
backend = TransformersBackend(args.model)
messages: list[dict[str, str]] = [{"role": "system", "content": args.system_prompt}]
print_panel("PertMind TUI", "Type your question and press Enter. Commands: /reset, /exit")
while True:
try:
user_text = input("You> ").strip()
except (EOFError, KeyboardInterrupt):
print()
break
if not user_text:
continue
if user_text.lower() in {"/exit", "exit", "quit", "/quit"}:
break
if user_text.lower() == "/reset":
messages = [{"role": "system", "content": args.system_prompt}]
print_panel("PertMind", "Conversation reset.")
continue
messages.append({"role": "user", "content": user_text})
answer = backend.generate(messages, args.max_new_tokens, args.temperature, args.top_p)
messages.append({"role": "assistant", "content": answer})
print_panel("PertMind", answer)
return 0
if __name__ == "__main__":
raise SystemExit(main())