File size: 5,105 Bytes
14a19cc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
135
136
137
138
139
#!/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())