Skip to content

Commit 95dfc14

Browse files
committed
Merge branch 'main' of https://github.com/OpenMLRL/CoMLRL
2 parents cd708cc + d2e7a5e commit 95dfc14

4 files changed

Lines changed: 949 additions & 91 deletions

File tree

comlrl/trainers/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,3 +3,4 @@
33
from .mareinforce import MAREINFORCEConfig, MAREINFORCETrainer
44
from .maremax import MAReMaxTrainer
55
from .marloo import MARLOOTrainer
6+
from .maac import MAACConfig, MAACTrainer

comlrl/trainers/iac.py

Lines changed: 115 additions & 78 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)