@@ -41,6 +41,7 @@ def _infer(
4141 output_tokens : int ,
4242 get_stats ,
4343 on_token = None ,
44+ max_response_tokens = None ,
4445):
4546 before = get_stats ()
4647 started = time .perf_counter ()
@@ -50,12 +51,30 @@ def _infer(
5051 append_done = time .perf_counter ()
5152 first_at = None
5253 generated = []
53- for token in s .generate (max_tokens = output_tokens ):
54- generated .append (int (token ))
55- if on_token is not None :
56- on_token (generated )
57- if first_at is None :
58- first_at = time .perf_counter ()
54+ response_limit = int (max_response_tokens or output_tokens )
55+ stop_reason = "unknown"
56+ while len (generated ) < response_limit :
57+ before_count = len (generated )
58+ chunk = min (output_tokens , response_limit - len (generated ))
59+ for token in s .generate (max_tokens = chunk ):
60+ generated .append (int (token ))
61+ if on_token is not None :
62+ on_token (generated )
63+ if first_at is None :
64+ first_at = time .perf_counter ()
65+ stop_reason = {
66+ 1 : "max_tokens" ,
67+ 2 : "eos" ,
68+ 3 : "cancelled" ,
69+ 4 : "truncated" ,
70+ }.get (s .last_stop_reason , "unknown" )
71+ if stop_reason != "max_tokens" :
72+ break
73+ if len (generated ) == before_count :
74+ stop_reason = "no_progress"
75+ break
76+ if len (generated ) >= response_limit and stop_reason == "max_tokens" :
77+ stop_reason = "client_safety_limit"
5978 done = time .perf_counter ()
6079 after = get_stats ()
6180 first_at = first_at or done
@@ -67,6 +86,8 @@ def _infer(
6786 "decode_s" : done - append_done ,
6887 "e2e_s" : done - started ,
6988 "delta" : _delta (before , after ),
89+ "stop_reason" : stop_reason ,
90+ "complete" : stop_reason == "eos" ,
7091 }
7192
7293
@@ -85,6 +106,7 @@ def main() -> int:
85106 "reliably; larger values require more worker memory or timeout." ,
86107 )
87108 parser .add_argument ("--output-tokens" , type = int , default = 64 )
109+ parser .add_argument ("--max-response-tokens" , type = int , default = 512 )
88110 parser .add_argument ("--report" , default = "/tmp/kakeya-agent-gan-demo.json" )
89111 parser .add_argument ("--skip-ensure" , action = "store_true" )
90112 args = parser .parse_args ()
@@ -113,7 +135,9 @@ def main() -> int:
113135 "role" : "system" ,
114136 "content" : (
115137 "You are the Generator agent. Propose a technically precise "
116- "architecture improvement. Respond with actionable reasoning."
138+ "architecture improvement. Respond with actionable reasoning. "
139+ "For open or unsolved problems, state the accepted boundary "
140+ "honestly and never fabricate a proof."
117141 ),
118142 },
119143 {"role" : "user" , "content" : task },
@@ -123,7 +147,10 @@ def main() -> int:
123147 "content" : (
124148 "You are the Critic/Discriminator agent. Attack the proposal, "
125149 "identify false assumptions and bottlenecks, score it from 0 to "
126- "10, and demand specific corrections."
150+ "10, and demand specific corrections. Do not call a response "
151+ "incomplete merely because it refuses to fabricate a solution to "
152+ "an open problem. Claim truncation only when completion_status is "
153+ "not EOS or the text is syntactically cut off."
127154 ),
128155 }]
129156
@@ -139,6 +166,7 @@ def main() -> int:
139166 "agents" : ["generator" , "critic" ],
140167 "rounds" : args .rounds ,
141168 "output_tokens" : args .output_tokens ,
169+ "max_response_tokens" : args .max_response_tokens ,
142170 },
143171 },
144172 )
@@ -164,6 +192,7 @@ def execute_agent(client, name, round_index, history):
164192 token_ids ,
165193 args .output_tokens ,
166194 get_stats ,
195+ max_response_tokens = args .max_response_tokens ,
167196 )
168197 text = tokenizer .decode (generated , skip_special_tokens = True )
169198 delta = actual ["delta" ]
@@ -174,7 +203,7 @@ def execute_agent(client, name, round_index, history):
174203 "agent" : name ,
175204 "round" : round_index ,
176205 "hit_source" : "primary_hot" if delta ["local_hits" ] else "unknown" ,
177- "ok" : ok ,
206+ "ok" : ok and actual [ "complete" ] ,
178207 "warmup_prefix_tokens" : warm ["prefix_tokens" ],
179208 "warmup_tokens_reused" : (
180209 warm ["delta" ]["tokens_reused" ]
@@ -194,7 +223,7 @@ def execute_agent(client, name, round_index, history):
194223 )
195224 all_stages .append (stage )
196225 print (f"\n [{ name .upper ()} round { round_index } ]\n { text } \n " , flush = True )
197- return text
226+ return text , stage
198227
199228 try :
200229 with Client (args .address ) as client :
@@ -208,17 +237,19 @@ def execute_agent(client, name, round_index, history):
208237 + critic_feedback
209238 ),
210239 })
211- proposal = execute_agent (
240+ proposal , generator_stage = execute_agent (
212241 client , "generator" , round_index , generator_history ,
213242 )
214243 generator_history .append ({"role" : "assistant" , "content" : proposal })
215244 critic_history .append ({
216245 "role" : "user" ,
217246 "content" : (
218247 f"Architecture task:\n { task } \n \n Generator proposal:\n { proposal } "
248+ f"\n \n completion_status={ generator_stage ['stop_reason' ]} ; "
249+ f"complete={ generator_stage ['complete' ]} "
219250 ),
220251 })
221- critic_feedback = execute_agent (
252+ critic_feedback , _critic_stage = execute_agent (
222253 client , "critic" , round_index , critic_history ,
223254 )
224255 critic_history .append ({
0 commit comments