Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -112,7 +112,7 @@ For new contributors, please see [contributing guidelines](./CONTRIBUTING.md) on

CoMLRL was developed using substantial computational resources. Its growth has been made possible by the generous support of the following organizations and institutions.

<p><img src="docs/assets/sponsors.jpg" width="500px;" alt=""/></p>
<p><img src="docs/assets/sponsors.png" width="500px;" alt=""/></p>

We welcome computational sponsorship to support the continued development of CoMLRL. If you are interested in supporting this project, please contact us.
<a href="mailto:liu.shuo2@northeastern.edu"><img src="docs/assets/email.svg" width="22px" alt="Email"/></a>
Expand Down
20 changes: 20 additions & 0 deletions comlrl/trainers/actor_critic/ac_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,26 @@ def _summarize_rollout_metrics(self, rollouts: List[Any]) -> Dict[str, float]:
if target_vals.numel() > 0 and torch.isfinite(target_vals).all():
metrics["value_target_mean"] = float(target_vals.mean().item())

reference_kls = [
float(sample.metadata.get("reference_kl", 0.0))
for sample in rollouts
if "reference_kl" in sample.metadata
]
if reference_kls:
vals = torch.tensor(reference_kls, dtype=torch.float32)
if torch.isfinite(vals).all():
metrics["reference_kl_mean"] = float(vals.mean().item())

reference_penalties = [
float(sample.metadata.get("reference_kl_penalty", 0.0))
for sample in rollouts
if "reference_kl_penalty" in sample.metadata
]
if reference_penalties:
vals = torch.tensor(reference_penalties, dtype=torch.float32)
if torch.isfinite(vals).all():
metrics["reference_kl_penalty_mean"] = float(vals.mean().item())

return metrics

def _iter_dataloader(self, dataloader, epoch: int, total_epochs: int):
Expand Down
137 changes: 128 additions & 9 deletions comlrl/trainers/actor_critic/iac.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,15 @@
from comlrl.utils.distributed import local_context, unwrap_model
from comlrl.utils.formatters import build_formatters
from comlrl.utils.model_loading import resolve_model_sources
from comlrl.utils.reference_kl import (
clone_reference_models,
load_reference_models_from_sources,
reference_kl_coef,
reference_kl_enabled,
reference_kl_for_sequence,
resolve_reference_devices,
validate_reference_kl_config,
)
from comlrl.utils.reward_utils import call_reward_function, normalize_reward_lengths
from comlrl.utils.tokenizer_utils import apply_tokenizer_specials, resolve_tokenizers
from .ac_base import ActorCriticTrainerBase
Expand Down Expand Up @@ -58,6 +67,9 @@ class IACConfig:
eval_batch_size: int = 1
early_termination_threshold: Optional[float] = -0.2
logging_steps: int = 1
reference_kl_enabled: bool = False
reference_kl_coef: float = 0.1
reference_devices: Optional[Union[str, Sequence[str]]] = None

def __post_init__(self) -> None:
if self.rollout_buffer_size < 1:
Expand Down Expand Up @@ -103,6 +115,7 @@ def __post_init__(self) -> None:
"use_separate_critic=True."
)
self.parallel_training = mode
validate_reference_kl_config(self, self.num_agents)


@dataclass
Expand Down Expand Up @@ -230,6 +243,9 @@ def __init__(
expected_count=self.args.num_agents,
model_label="agent_model",
)
actor_model_kwargs = self._filter_model_kwargs(
self.model_config.get("model_kwargs", {})
)
for idx, actor_source in enumerate(actor_sources):
if actor_source is None:
raise ValueError("A policy model identifier or instance is required.")
Expand All @@ -244,12 +260,9 @@ def __init__(
attach_value_head=attach_value,
)
else:
model_kwargs = self._filter_model_kwargs(
self.model_config.get("model_kwargs", {})
)
try:
base_model = AutoModelForCausalLM.from_pretrained(
actor_source, **model_kwargs
actor_source, **actor_model_kwargs
)
except (OSError, ValueError) as exc:
raise ValueError(
Expand Down Expand Up @@ -308,14 +321,38 @@ def __init__(
else:
self.critics = []

self.reference_models: List[Any] = []
self.reference_devices: List[torch.device] = []
if reference_kl_enabled(self.args):
self.reference_devices = resolve_reference_devices(
self.args,
self.agent_devices,
self.args.num_agents,
)
if actor_sources and all(isinstance(src, str) for src in actor_sources):
self.reference_models = load_reference_models_from_sources(
actor_sources,
devices=self.reference_devices,
model_kwargs=actor_model_kwargs,
)
else:
self.reference_models = clone_reference_models(
self.agents,
devices=self.reference_devices,
)

if self.tokenizers and len(self.tokenizers) == len(self.agents):
for idx, tok in enumerate(self.tokenizers):
models = [self.agents[idx]]
if idx < len(self.critics):
models.append(self.critics[idx])
if idx < len(self.reference_models):
models.append(self.reference_models[idx])
apply_tokenizer_specials(tok, models)
else:
apply_tokenizer_specials(self.tokenizer, [*self.agents, *self.critics])
apply_tokenizer_specials(
self.tokenizer, [*self.agents, *self.critics, *self.reference_models]
)

self.agent_optimizers = []
self.critic_optimizers = []
Expand Down Expand Up @@ -390,6 +427,9 @@ def _init_wandb(self) -> None:
"max_new_tokens": self.args.max_new_tokens,
"use_separate_critic": self.args.use_separate_critic,
"critic_type": self.args.critic_type,
"reference_kl_enabled": reference_kl_enabled(self.args),
"reference_kl_coef": reference_kl_coef(self.args),
"reference_devices": getattr(self.args, "reference_devices", None),
}

sections = (
Expand Down Expand Up @@ -549,6 +589,14 @@ def _generate_rollout(
output_values=False,
)
logprobs.append(lp.squeeze(0))
reference_kls = self._reference_kl_values(
agent_idx,
agent_model,
sequences,
full_attention_mask,
prompt_len,
response_lens,
)

return {
"prompt": prompt,
Expand All @@ -557,11 +605,58 @@ def _generate_rollout(
"attention_mask": full_attention_mask,
"response_lens": response_lens,
"logprobs": logprobs,
"reference_kls": reference_kls,
"values": value,
"completions": completion_texts,
"char_lengths": [len(txt) for txt in completion_texts],
}

def _reference_kl_values(
self,
agent_idx: int,
policy_model: CausalLMWithValueHead,
sequences: torch.Tensor,
attention_mask: torch.Tensor,
prompt_len: int,
response_lens: Sequence[int],
) -> List[float]:
if not reference_kl_enabled(self.args):
return []
if not getattr(self, "reference_models", None):
raise RuntimeError(
"Reference KL is enabled but reference models are missing."
)
reference_model = self.reference_models[agent_idx]
values: List[float] = []
for seq, attn, response_len in zip(sequences, attention_mask, response_lens):
kl_value = reference_kl_for_sequence(
policy_model,
reference_model,
seq.unsqueeze(0),
attn.unsqueeze(0),
prompt_len,
int(response_len),
)
values.append(float(kl_value.view(-1)[0].item()))
return values

def _kl_shaped_reward(
self, reward: float, data: Dict[str, Any], index: int
) -> Tuple[float, Dict[str, float]]:
if not reference_kl_enabled(self.args):
return reward, {}
reference_kl = 0.0
penalty = 0.0
raw_kls = data.get("reference_kls") or []
if index < len(raw_kls):
reference_kl = float(raw_kls[index])
penalty = reference_kl_coef(self.args) * reference_kl
return reward - penalty, {
"reference_kl": reference_kl,
"reference_kl_penalty": penalty,
"environment_reward": reward,
}

def _expand_rewards(self, rewards: List[float], num_ret: int) -> List[List[float]]:
num_agents = self.args.num_agents
if len(rewards) == 1:
Expand Down Expand Up @@ -628,8 +723,10 @@ def _generate_agent(agent_idx: int) -> Dict[str, Any]:
value = data["values"][i]
reward = float(rewards_matrix[agent_idx][i])
reward_cpu = torch.tensor([reward], dtype=torch.float32)
shaped_reward, kl_meta = self._kl_shaped_reward(reward, data, i)
shaped_reward_cpu = torch.tensor([shaped_reward], dtype=torch.float32)
value_cpu = value.detach().cpu()
returns_cpu = reward_cpu.clone()
returns_cpu = shaped_reward_cpu.clone()
advantage_cpu = returns_cpu - value_cpu.to(dtype=returns_cpu.dtype)
logprob_cpu = logprob.detach().cpu()

Expand All @@ -652,6 +749,7 @@ def _generate_agent(agent_idx: int) -> Dict[str, Any]:
advantage=advantage_cpu,
metadata={
"char_length": data["char_lengths"][i],
**kl_meta,
"value_target": returns_cpu,
},
)
Expand Down Expand Up @@ -743,6 +841,8 @@ def _generate_agent_turn(agent_idx: int) -> Dict[str, Any]:
value = data["values"][0]
reward_val = float(rewards_matrix[agent_idx][0])
reward_cpu = torch.tensor([reward_val], dtype=torch.float32)
shaped_reward, kl_meta = self._kl_shaped_reward(reward_val, data, 0)
shaped_reward_cpu = torch.tensor([shaped_reward], dtype=torch.float32)
value_cpu = value.detach().cpu()
logprob_cpu = logprob.detach().cpu()

Expand All @@ -758,12 +858,13 @@ def _generate_agent_turn(agent_idx: int) -> Dict[str, Any]:
old_logprob=logprob_cpu,
old_value=value_cpu,
reward=reward_cpu,
returns=reward_cpu.clone(),
returns=shaped_reward_cpu.clone(),
advantage=torch.zeros_like(reward_cpu),
normalized_advantage=None,
metadata={
"char_length": data["char_lengths"][0],
"turn_idx": turn_idx,
**kl_meta,
},
)
rollouts.append(sample)
Expand All @@ -780,7 +881,7 @@ def _generate_agent_turn(agent_idx: int) -> Dict[str, Any]:
for agent_idx in range(self.args.num_agents):
traj = per_agent_samples[agent_idx]
for t, sample in enumerate(traj):
r = float(sample.reward.view(-1)[0].item())
r = float(sample.returns.view(-1)[0].item())
if t < len(traj) - 1:
next_v = float(traj[t + 1].old_value.view(-1)[0].item())
target = r + gamma * next_v
Expand All @@ -792,7 +893,7 @@ def _generate_agent_turn(agent_idx: int) -> Dict[str, Any]:
for agent_idx in range(self.args.num_agents):
future = 0.0
for sample in reversed(per_agent_samples[agent_idx]):
immediate = float(sample.reward.view(-1)[0].item())
immediate = float(sample.returns.view(-1)[0].item())
future = immediate + gamma * future
sample.returns = torch.tensor([future], dtype=torch.float32)
sample.advantage = torch.zeros_like(sample.returns)
Expand Down Expand Up @@ -1085,6 +1186,24 @@ def _update(
metrics["reward_mean"].append(rewards.mean().item())
if returns_raw.numel() > 0 and torch.isfinite(returns_raw).all():
metrics["expected_return"].append(returns_raw.mean().item())
reference_kls = [
float(sample.metadata.get("reference_kl", 0.0))
for sample in rollouts
if "reference_kl" in sample.metadata
]
if reference_kls:
vals = torch.tensor(reference_kls, dtype=torch.float32)
if torch.isfinite(vals).all():
metrics["reference_kl_mean"].append(vals.mean().item())
reference_penalties = [
float(sample.metadata.get("reference_kl_penalty", 0.0))
for sample in rollouts
if "reference_kl_penalty" in sample.metadata
]
if reference_penalties:
vals = torch.tensor(reference_penalties, dtype=torch.float32)
if torch.isfinite(vals).all():
metrics["reference_kl_penalty_mean"].append(vals.mean().item())

if self.metrics_callback is not None:
extra = self.metrics_callback(rollouts)
Expand Down
Loading
Loading