1# Tweakable parameters
2# MODEL_PATH = "./Qwen3-BEST" # local run
3MODEL_PATH = "g023/Qwen3-1.77B-g023"
4MAX_NEW_TOKENS = 8192
5TEMPERATURE = 0.7
6DO_SAMPLE = True
7TOP_P = 0.9
8TOP_K = 50
9REPETITION_PENALTY = 1.1
10STREAMING = True # Set to True for streaming inference
11INPUT_MESSAGE = "You are completing the next step in a task to create an arcade game in javascript. Your available tools are rationalize, red_green_tdd, and create_plan. Synthesize their output when reasoning. "
12
13from transformers import AutoModelForCausalLM, AutoTokenizer, TextStreamer
14import time
15
16def load_model():
17 print("Loading model...")
18 model = AutoModelForCausalLM.from_pretrained(
19 MODEL_PATH,
20 device_map="auto",
21 )
22 tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH)
23 print("Model loaded.")
24 return model, tokenizer
25
26def inference_non_streaming(model, tokenizer, messages):
27 text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, enable_thinking=True)
28 inputs = tokenizer(text, return_tensors="pt").to(model.device)
29 outputs = model.generate(
30 **inputs,
31 max_new_tokens=MAX_NEW_TOKENS,
32 temperature=TEMPERATURE,
33 do_sample=DO_SAMPLE,
34 top_p=TOP_P,
35 top_k=TOP_K,
36 repetition_penalty=REPETITION_PENALTY,
37 )
38 response = tokenizer.decode(outputs[0][inputs["input_ids"].shape[1]:], skip_special_tokens=True)
39 print("Response:", response)
40 return response
41
42def inference_streaming(model, tokenizer, messages):
43 final_response = ""
44 text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True, enable_thinking=True)
45 inputs = tokenizer(text, return_tensors="pt").to(model.device)
46 streamer = TextStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
47 outputs = model.generate(
48 **inputs,
49 max_new_tokens=MAX_NEW_TOKENS,
50 temperature=TEMPERATURE,
51 do_sample=DO_SAMPLE,
52 top_p=TOP_P,
53 top_k=TOP_K,
54 repetition_penalty=REPETITION_PENALTY,
55 streamer=streamer,
56 )
57
58
59 # return a final str
60 return final_response
61
62def llm_stream(model, tokenizer, conversation):
63 import time
64 start_time = time.time()
65 text = tokenizer.apply_chat_template(conversation, tokenize=False, add_generation_prompt=True, enable_thinking=True)
66 inputs = tokenizer(text, return_tensors="pt").to(model.device)
67 from io import StringIO
68 buffer = StringIO()
69 class CapturingTextStreamer(TextStreamer):
70 def __init__(self, tokenizer, buffer):
71 super().__init__(tokenizer, skip_prompt=True, skip_special_tokens=True)
72 self.buffer = buffer
73 def on_finalized_text(self, text, stream_end=False):
74 self.buffer.write(text)
75 print(text, end="", flush=True)
76 streamer = CapturingTextStreamer(tokenizer, buffer)
77 outputs = model.generate(
78 **inputs,
79 max_new_tokens=MAX_NEW_TOKENS,
80 temperature=TEMPERATURE,
81 do_sample=DO_SAMPLE,
82 top_p=TOP_P,
83 top_k=TOP_K,
84 repetition_penalty=REPETITION_PENALTY,
85 streamer=streamer,
86 )
87 response = buffer.getvalue()
88
89 if "</think>" in response:
90 parts = response.rsplit("</think>", 1)
91 reasoning = parts[0].strip()
92 content = parts[1].strip()
93 else:
94 reasoning = ""
95 content = response.strip()
96 char_per_token = 3.245
97 reasoning_tokens = round(len(reasoning) / char_per_token)
98 content_tokens = round(len(content) / char_per_token)
99 total_tokens = reasoning_tokens + content_tokens
100 time_taken = time.time() - start_time
101 ret_dict = {
102 "reasoning": reasoning,
103 "content": content,
104 "usage": {
105 "reasoning_tokens": reasoning_tokens,
106 "content_tokens": content_tokens,
107 "total_tokens": total_tokens,
108 },
109 "time_taken": time_taken,
110 }
111 return ret_dict
112
113if __name__ == "__main__":
114 model, tokenizer = load_model()
115 messages = [{"role": "user", "content": INPUT_MESSAGE}]
116 ret = llm_stream(model, tokenizer, messages)
117 print("Result dict:", ret)
118
119 # output tokens per second by taking total_tokens and time_taken
120 if ret["usage"]["total_tokens"] > 0 and ret["time_taken"] > 0:
121 tps = ret["usage"]["total_tokens"] / ret["time_taken"]
122 print(f"Tokens per second: {tps:.2f}")