@@ -61,6 +61,7 @@ class IACConfig:
6161 num_agents : int = 1
6262 num_turns : int = 1
6363 reward_norm_eps : float = 1e-3
64+ num_return_sequences : int = 1
6465
6566 def __post_init__ (self ) -> None :
6667 if self .rollout_buffer_size < 1 :
@@ -79,6 +80,8 @@ def __post_init__(self) -> None:
7980 )
8081 if self .critic_learning_rate is None :
8182 self .critic_learning_rate = self .actor_learning_rate
83+ if self .num_return_sequences < 1 :
84+ raise ValueError ("num_return_sequences must be >= 1." )
8285
8386
8487@dataclass
@@ -451,6 +454,9 @@ def _call_with_args():
451454 return processed * num_agents
452455 if len (processed ) == num_agents :
453456 return processed
457+ num_ret = int (getattr (self .args , "num_return_sequences" , 1 ))
458+ if len (processed ) == num_ret :
459+ return processed
454460 raise ValueError (
455461 f"Reward function must return either 1 or { num_agents } values per prompt for multi-agent IAC."
456462 )
@@ -460,8 +466,9 @@ def _call_with_args():
460466 # --------------------------------------------------------------------- #
461467 def _collect_rollouts (self , item : Dict [str , Any ]) -> List [RolloutSample ]:
462468 prompts : List [str ] = []
463- completions : List [str ] = []
469+ completions_per_agent : List [List [ str ] ] = []
464470 rollout_data : List [Dict [str , Any ]] = []
471+ num_ret = int (getattr (self .args , "num_return_sequences" , 1 ))
465472
466473 for agent_idx , actor_model in enumerate (self .actor_models ):
467474 prompt = self ._format_prompt (item , agent_idx )
@@ -478,6 +485,8 @@ def _collect_rollouts(self, item: Dict[str, Any]) -> List[RolloutSample]:
478485 "temperature" : self .args .temperature ,
479486 "top_p" : self .args .top_p ,
480487 "pad_token_id" : self .args .pad_token_id ,
488+ "num_return_sequences" : num_ret ,
489+ "num_beams" : 1 ,
481490 }
482491 if self .args .top_k is not None :
483492 generation_kwargs ["top_k" ] = self .args .top_k
@@ -487,93 +496,121 @@ def _collect_rollouts(self, item: Dict[str, Any]) -> List[RolloutSample]:
487496 raise RuntimeError ("Model produced an empty completion during rollout." )
488497
489498 response_tokens = sequences [:, prompt_len :]
490- completion_text = self .tokenizer .decode (
491- response_tokens [0 ], skip_special_tokens = True
492- )
493- response_len = response_tokens .size (1 )
494- response_char_length = len (completion_text )
499+ pad_id = self .args .pad_token_id
500+ response_lens : List [int ] = []
501+ completion_texts : List [str ] = []
502+ for seq in response_tokens :
503+ if pad_id is not None :
504+ pad_positions = (seq == pad_id ).nonzero (as_tuple = False )
505+ resp_len = (
506+ pad_positions [0 ].item ()
507+ if pad_positions .numel () > 0
508+ else seq .size (0 )
509+ )
510+ else :
511+ resp_len = seq .size (0 )
512+ response_lens .append (resp_len )
513+ completion_texts .append (
514+ self .tokenizer .decode (seq [:resp_len ], skip_special_tokens = True )
515+ )
495516
517+ completions_per_agent .append (completion_texts )
496518 full_attention_mask = torch .ones_like (sequences , device = self .device )
497519
498520 with torch .no_grad ():
499- # Policy log-prob uses full prompt+response; value uses prompt only.
500- logprob , _ = self ._policy_eval (
501- actor_model ,
502- sequences ,
503- full_attention_mask ,
504- prompt_len ,
505- response_len ,
506- output_values = False ,
507- )
508521 if self .args .use_separate_critic :
509522 critic_model = self .critic_models [agent_idx ]
510523 if critic_model is None :
511524 raise RuntimeError ("Critic model missing for agent." )
512- value = self ._critic_eval (
513- critic_model ,
514- sequences ,
515- full_attention_mask ,
516- prompt_len ,
517- response_len ,
525+ value = self ._value_on_prompt_only (
526+ critic_model , sequences , full_attention_mask , prompt_len
518527 )
519528 else :
520529 value = self ._value_on_prompt_only (
521530 actor_model , sequences , full_attention_mask , prompt_len
522531 )
523532
524- prompts .append (prompt )
525- completions .append (completion_text )
533+ logprobs = []
534+ for seq , attn , resp_len in zip (
535+ sequences , full_attention_mask , response_lens
536+ ):
537+ lp , _ = self ._policy_eval (
538+ actor_model ,
539+ seq .unsqueeze (0 ),
540+ attn .unsqueeze (0 ),
541+ prompt_len ,
542+ resp_len ,
543+ output_values = False ,
544+ )
545+ logprobs .append (lp .squeeze (0 ))
546+
526547 rollout_data .append (
527548 {
528549 "agent_idx" : agent_idx ,
529550 "prompt" : prompt ,
530- "completion " : completion_text ,
551+ "prompt_len " : prompt_len ,
531552 "sequences" : sequences ,
532553 "attention_mask" : full_attention_mask ,
533- "prompt_len" : prompt_len ,
534- "response_len" : response_len ,
535- "logprob" : logprob ,
536- "value" : value ,
537- "char_length" : response_char_length ,
554+ "response_lens" : response_lens ,
555+ "logprobs" : logprobs ,
556+ "values" : value ,
557+ "char_lengths" : [len (txt ) for txt in completion_texts ],
538558 }
539559 )
560+ prompts .append (prompt )
540561
541- completion_lists = [[ text ] for text in completions ]
542- rewards = self ._call_reward_func ( prompts , completion_lists )
562+ rewards = self . _call_reward_func ( prompts , completions_per_agent )
563+ num_agents = self .args . num_agents
543564
544- if len (rewards ) != len (rollout_data ):
545- if len (rewards ) == 1 :
546- rewards = rewards * len (rollout_data )
547- else :
548- raise ValueError (
549- "Reward function returned unexpected number of values."
550- )
565+ # Normalize rewards to a per-agent x per-sample matrix for downstream use.
566+ if len (rewards ) == 1 :
567+ rewards_matrix = [[rewards [0 ]] * num_ret for _ in range (num_agents )]
568+ elif len (rewards ) == num_ret :
569+ rewards_matrix = [list (rewards ) for _ in range (num_agents )]
570+ elif len (rewards ) == num_agents :
571+ rewards_matrix = [[rewards [a ]] * num_ret for a in range (num_agents )]
572+ else :
573+ raise ValueError (
574+ "Reward function must return 1 value, num_return_sequences values, "
575+ "or num_agents values."
576+ )
551577
552578 rollouts : List [RolloutSample ] = []
553- for data , reward in zip (rollout_data , rewards ):
554- reward_tensor = torch .tensor (
555- [reward ], device = self .device , dtype = torch .float32
556- )
557- returns = reward_tensor .clone ()
558- advantage = returns - data ["value" ]
559-
560- rollouts .append (
561- RolloutSample (
562- agent_idx = data ["agent_idx" ],
563- prompt = data ["prompt" ],
564- completion = data ["completion" ],
565- full_input_ids = data ["sequences" ].squeeze (0 ).detach ().cpu (),
566- attention_mask = data ["attention_mask" ].squeeze (0 ).detach ().cpu (),
567- prompt_len = data ["prompt_len" ],
568- response_len = data ["response_len" ],
569- old_logprob = data ["logprob" ].detach ().cpu (),
570- old_value = data ["value" ].detach ().cpu (),
571- reward = reward_tensor .detach ().cpu (),
572- returns = returns .detach ().cpu (),
573- advantage = advantage .detach ().cpu (),
574- metadata = {"char_length" : data ["char_length" ]},
579+ for data in rollout_data :
580+ agent_idx = data ["agent_idx" ]
581+ for i in range (num_ret ):
582+ seq = data ["sequences" ][i ]
583+ attn = data ["attention_mask" ][i ]
584+ resp_len = data ["response_lens" ][i ]
585+ logprob = data ["logprobs" ][i ]
586+ value = data ["values" ][i ]
587+ reward = float (rewards_matrix [agent_idx ][i ])
588+ reward_tensor = torch .tensor (
589+ [reward ], device = self .device , dtype = torch .float32
590+ )
591+ returns = reward_tensor .clone ()
592+ advantage = returns - value
593+
594+ rollouts .append (
595+ RolloutSample (
596+ agent_idx = agent_idx ,
597+ prompt = data ["prompt" ],
598+ completion = self .tokenizer .decode (
599+ seq [data ["prompt_len" ] : data ["prompt_len" ] + resp_len ],
600+ skip_special_tokens = True ,
601+ ),
602+ full_input_ids = seq .detach ().cpu (),
603+ attention_mask = attn .detach ().cpu (),
604+ prompt_len = data ["prompt_len" ],
605+ response_len = resp_len ,
606+ old_logprob = logprob .detach ().cpu (),
607+ old_value = value .detach ().cpu (),
608+ reward = reward_tensor .detach ().cpu (),
609+ returns = returns .detach ().cpu (),
610+ advantage = advantage .detach ().cpu (),
611+ metadata = {"char_length" : data ["char_lengths" ][i ]},
612+ )
575613 )
576- )
577614
578615 return rollouts
579616
@@ -893,24 +930,12 @@ def train(self) -> None:
893930 buffer = self .rollout_buffers [agent_idx ]
894931 buffer .append (sample )
895932 if len (buffer ) >= self .args .rollout_buffer_size :
896- metrics = self ._update (agent_idx , buffer )
897- buffer .clear ()
898- tagged = self ._tag_metrics (metrics , agent_idx )
899- self ._log_metrics (tagged )
900- self .global_step += 1
901- for key , value in tagged .items ():
902- epoch_metrics [key ].append (value )
933+ self ._process_buffer (agent_idx , buffer , epoch_metrics )
903934
904935 for agent_idx , buffer in enumerate (self .rollout_buffers ):
905936 if not buffer :
906937 continue
907- metrics = self ._update (agent_idx , buffer )
908- buffer .clear ()
909- tagged = self ._tag_metrics (metrics , agent_idx )
910- self ._log_metrics (tagged )
911- self .global_step += 1
912- for key , value in tagged .items ():
913- epoch_metrics [key ].append (value )
938+ self ._process_buffer (agent_idx , buffer , epoch_metrics )
914939
915940 summary = {
916941 key : float (sum (values ) / len (values ))
@@ -926,16 +951,28 @@ def train(self) -> None:
926951 def _tag_metrics (
927952 self , metrics : Dict [str , float ], agent_idx : int
928953 ) -> Dict [str , float ]:
929- if self .args .num_agents == 1 :
930- return metrics
931- return {f"agent_{ agent_idx } /{ key } " : value for key , value in metrics .items ()}
954+ return {f"turn_1/{ key } " : value for key , value in metrics .items ()}
932955
933956 def _log_metrics (self , metrics : Dict [str , float ]) -> None :
934957 if not metrics :
935958 return
936959 if self .wandb_initialized and wandb is not None :
937960 wandb .log (metrics , step = self .global_step )
938961
962+ def _process_buffer (
963+ self ,
964+ agent_idx : int ,
965+ buffer : List [RolloutSample ],
966+ epoch_metrics : Dict [str , List [float ]],
967+ ) -> None :
968+ metrics = self ._update (agent_idx , buffer )
969+ buffer .clear ()
970+ tagged = self ._tag_metrics (metrics , agent_idx )
971+ self ._log_metrics (tagged )
972+ self .global_step += 1
973+ for key , value in tagged .items ():
974+ epoch_metrics [key ].append (value )
975+
939976 def save_model (self , output_dir : str ) -> None :
940977 os .makedirs (output_dir , exist_ok = True )
941978 if self .args .num_agents == 1 :
0 commit comments