4949
5050import argparse
5151import sys
52+ import time
5253from typing import List , Optional
5354
5455
@@ -92,6 +93,7 @@ def _generate_and_print(
9293 tokenizer ,
9394 new_tokens : List [int ],
9495 max_tokens : int ,
96+ max_response_tokens : int = 0 ,
9597) -> int :
9698 """Drive one append + generate cycle. Streams tokens to stdout
9799 as they arrive, returns the count emitted. The generator's
@@ -103,36 +105,74 @@ def _generate_and_print(
103105 print ("kakeya> " , end = "" , flush = True )
104106 n = 0
105107 accumulated = []
108+ started = time .perf_counter ()
109+ server_elapsed = 0.0
110+ stop_reason = "unknown"
106111 try :
107- for token_id in session .generate (max_tokens = max_tokens ):
108- n += 1
109- accumulated .append (token_id )
110- # Decode incrementally — tokenizer.decode on the running
111- # buffer gives the right text including BPE merges that
112- # span multiple tokens. We re-decode the full buffer
113- # each time (Qwen3-family tokenizers re-decode in <1ms
114- # for a 64-token buffer; per-token decoding loses some
115- # whitespace correctness on the tokenizer level).
116- text_so_far = tokenizer .decode (
117- accumulated , skip_special_tokens = True ,
112+ while True :
113+ before = n
114+ remaining = (
115+ min (max_tokens , max_response_tokens - n )
116+ if max_response_tokens > 0 else max_tokens
118117 )
119- # Print only the suffix that's new since last frame.
120- if hasattr (_generate_and_print , "_last_text" ):
121- last = _generate_and_print ._last_text
122- else :
123- last = ""
124- new_text = text_so_far [len (last ):]
125- print (new_text , end = "" , flush = True )
126- _generate_and_print ._last_text = text_so_far
118+ if remaining <= 0 :
119+ stop_reason = "client_safety_limit"
120+ break
121+ for token_id in session .generate (max_tokens = remaining ):
122+ n += 1
123+ accumulated .append (token_id )
124+ # Decode incrementally — tokenizer.decode on the running
125+ # buffer gives the right text including BPE merges that
126+ # span multiple tokens.
127+ text_so_far = tokenizer .decode (
128+ accumulated , skip_special_tokens = True ,
129+ )
130+ if hasattr (_generate_and_print , "_last_text" ):
131+ last = _generate_and_print ._last_text
132+ else :
133+ last = ""
134+ new_text = text_so_far [len (last ):]
135+ print (new_text , end = "" , flush = True )
136+ _generate_and_print ._last_text = text_so_far
137+ server_elapsed += float (
138+ getattr (session , "last_total_duration_seconds" , 0.0 ) or 0.0
139+ )
140+ stop_reason = {
141+ 1 : "max_tokens" ,
142+ 2 : "eos" ,
143+ 3 : "cancelled" ,
144+ 4 : "truncated" ,
145+ }.get (getattr (session , "last_stop_reason" , None ), "unknown" )
146+ if stop_reason != "max_tokens" :
147+ break
148+ if n == before :
149+ stop_reason = "no_progress"
150+ break
127151 except KeyboardInterrupt :
128152 print ("\n [interrupted]" , file = sys .stderr )
153+ stop_reason = "interrupted"
129154 finally :
130155 # Reset the per-call decoder state so the next turn starts
131156 # fresh.
132157 if hasattr (_generate_and_print , "_last_text" ):
133158 del _generate_and_print ._last_text
134159
135160 print () # final newline
161+ elapsed = max (time .perf_counter () - started , 1e-9 )
162+ measured = server_elapsed or elapsed
163+ print (
164+ f"[{ n } tokens · { measured :.2f} s · { n / measured :.2f} tok/s "
165+ f"· stop={ stop_reason } ]" ,
166+ file = sys .stderr ,
167+ flush = True ,
168+ )
169+ if stop_reason == "client_safety_limit" :
170+ print (
171+ f"[response reached optional --max-response-tokens "
172+ f"{ max_response_tokens } ]" ,
173+ file = sys .stderr ,
174+ flush = True ,
175+ )
136176 return n
137177
138178
@@ -161,7 +201,11 @@ def main() -> int:
161201 )
162202 ap .add_argument (
163203 "--max-tokens" , type = int , default = 64 ,
164- help = "max_tokens per turn" ,
204+ help = "tokens per streaming Generate RPC; max_tokens continues automatically" ,
205+ )
206+ ap .add_argument (
207+ "--max-response-tokens" , type = int , default = 0 ,
208+ help = "optional client safety cap per answer; 0 means continue until EOS" ,
165209 )
166210 ap .add_argument (
167211 "--system-prompt" , default = "You are a helpful assistant." ,
@@ -250,6 +294,7 @@ def _make_session(client):
250294 tokenizer = tokenizer ,
251295 new_tokens = new_tokens ,
252296 max_tokens = args .max_tokens ,
297+ max_response_tokens = args .max_response_tokens ,
253298 )
254299 except KakeyaError as exc :
255300 print (f"[runtime error: { exc } ]" , file = sys .stderr )
0 commit comments