diff --git a/README.md b/README.md
index ea30c61..01f8f99 100644
--- a/README.md
+++ b/README.md
@@ -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.
-

+
We welcome computational sponsorship to support the continued development of CoMLRL. If you are interested in supporting this project, please contact us.
diff --git a/comlrl/trainers/actor_critic/ac_base.py b/comlrl/trainers/actor_critic/ac_base.py
index 1ae6523..9ff3dce 100644
--- a/comlrl/trainers/actor_critic/ac_base.py
+++ b/comlrl/trainers/actor_critic/ac_base.py
@@ -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):
diff --git a/comlrl/trainers/actor_critic/iac.py b/comlrl/trainers/actor_critic/iac.py
index 8ed72ea..5e32b04 100644
--- a/comlrl/trainers/actor_critic/iac.py
+++ b/comlrl/trainers/actor_critic/iac.py
@@ -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
@@ -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:
@@ -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
@@ -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.")
@@ -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(
@@ -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 = []
@@ -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 = (
@@ -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,
@@ -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:
@@ -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()
@@ -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,
},
)
@@ -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()
@@ -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)
@@ -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
@@ -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)
@@ -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)
diff --git a/comlrl/trainers/actor_critic/maac.py b/comlrl/trainers/actor_critic/maac.py
index 57b4099..43d5eed 100644
--- a/comlrl/trainers/actor_critic/maac.py
+++ b/comlrl/trainers/actor_critic/maac.py
@@ -15,6 +15,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
@@ -54,6 +63,9 @@ class MAACConfig:
eval_num_samples: int = 4
eval_batch_size: int = 1
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:
@@ -96,6 +108,7 @@ def __post_init__(self) -> None:
"parallel_training='mp' requires explicit critic_devices."
)
self.parallel_training = mode
+ validate_reference_kl_config(self, self.num_agents)
class MAACTrainer(ActorCriticTrainerBase):
@@ -198,6 +211,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("agent_model must be provided for MAAC.")
@@ -211,12 +227,9 @@ def __init__(
value_head_hidden_dim=None,
)
else:
- model_kwargs = self._filter_model_kwargs(
- self.model_config.get("model_kwargs", {})
- )
try:
base = AutoModelForCausalLM.from_pretrained(
- actor_source, **model_kwargs
+ actor_source, **actor_model_kwargs
)
except (OSError, ValueError) as exc:
raise ValueError(
@@ -272,12 +285,37 @@ def __init__(
critic_model_instance.to(self.critic_device)
self.critics: List[CausalLMWithValueHead] = [critic_model_instance]
+ 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):
- apply_tokenizer_specials(tok, [self.agents[idx]])
+ models = [self.agents[idx]]
+ if idx < len(self.reference_models):
+ models.append(self.reference_models[idx])
+ apply_tokenizer_specials(tok, models)
apply_tokenizer_specials(self.tokenizers[0], [self.critics[0]])
else:
- apply_tokenizer_specials(self.tokenizer, [*self.agents, self.critics[0]])
+ apply_tokenizer_specials(
+ self.tokenizer, [*self.agents, self.critics[0], *self.reference_models]
+ )
self.formatters = build_formatters(formatters, self.args.num_agents)
try:
@@ -340,6 +378,9 @@ def _init_wandb(self) -> None:
"max_new_tokens": self.args.max_new_tokens,
"num_generations": self.args.num_generations,
"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 = (
@@ -490,16 +531,72 @@ def _generate(self, agent_model, prompt: str, agent_idx: int) -> Dict[str, Any]:
completion_texts.append(
tokenizer.decode(seq[:resp_len], skip_special_tokens=True)
)
+ full_attention_mask = torch.ones_like(sequences, device=agent_device)
+ reference_kls = self._reference_kl_values(
+ agent_idx,
+ agent_model,
+ sequences,
+ full_attention_mask,
+ prompt_len,
+ response_lens,
+ )
return {
"prompt": prompt,
"prompt_len": prompt_len,
"sequences": sequences,
- "attention_mask": torch.ones_like(sequences, device=agent_device),
+ "attention_mask": full_attention_mask,
"response_lens": response_lens,
+ "reference_kls": reference_kls,
"completions": 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 _collect_rollouts(self, item: Dict[str, Any]) -> List[RolloutSample]:
num_turns = max(1, int(getattr(self.args, "num_turns", 1)))
if num_turns > 1:
@@ -522,6 +619,7 @@ def _generate_agent(agent_idx: int) -> Dict[str, Any]:
"sequences": gen["sequences"],
"attention_mask": gen["attention_mask"],
"response_lens": gen["response_lens"],
+ "reference_kls": gen["reference_kls"],
"completion_texts": gen["completions"],
}
@@ -583,6 +681,7 @@ def _generate_agent(agent_idx: int) -> Dict[str, Any]:
attn = data["attention_mask"][i]
resp_len = data["response_lens"][i]
reward = float(rewards_matrix[agent_idx][i])
+ shaped_reward, kl_meta = self._kl_shaped_reward(reward, data, i)
logprob, _ = self._policy_eval(
self.agents[agent_idx],
@@ -603,6 +702,7 @@ def _generate_agent(agent_idx: int) -> Dict[str, Any]:
joint_len = int(critic_pack["prompt_len"])
value = critic_pack["value"].detach().cpu()
reward_cpu = torch.tensor([reward], dtype=torch.float32)
+ shaped_reward_cpu = torch.tensor([shaped_reward], dtype=torch.float32)
logprob_cpu = logprob.detach().cpu()
rollouts.append(
RolloutSample(
@@ -619,7 +719,7 @@ def _generate_agent(agent_idx: int) -> Dict[str, Any]:
old_logprob=logprob_cpu,
old_value=value,
reward=reward_cpu,
- returns=reward_cpu.clone(),
+ returns=shaped_reward_cpu.clone(),
advantage=torch.zeros_like(reward_cpu),
normalized_advantage=None,
metadata={
@@ -627,13 +727,14 @@ def _generate_agent(agent_idx: int) -> Dict[str, Any]:
"joint_attention_mask": joint_mask.detach().cpu(),
"joint_prompt_len": joint_len,
"turn_idx": 0,
- "adv_target": reward_cpu,
+ "adv_target": shaped_reward_cpu,
+ **kl_meta,
},
)
)
for sample in rollouts:
- r = float(sample.reward.view(-1)[0].item())
+ r = float(sample.returns.view(-1)[0].item())
sample.metadata["value_target"] = torch.tensor([r]).cpu()
if self.metrics_callback is not None:
@@ -703,6 +804,7 @@ def _generate_agent_turn(agent_idx: int) -> Dict[str, Any]:
"sequences": gen["sequences"],
"attention_mask": gen["attention_mask"],
"response_lens": gen["response_lens"],
+ "reference_kls": gen["reference_kls"],
"completion_texts": gen["completions"],
}
@@ -744,6 +846,8 @@ def _generate_agent_turn(agent_idx: int) -> Dict[str, Any]:
resp_len = data["response_lens"][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)
logprob, _ = self._policy_eval(
self.agents[agent_idx],
@@ -768,7 +872,7 @@ def _generate_agent_turn(agent_idx: int) -> Dict[str, Any]:
old_logprob=logprob_cpu,
old_value=value,
reward=reward_cpu,
- returns=reward_cpu.clone(),
+ returns=shaped_reward_cpu.clone(),
advantage=torch.zeros_like(reward_cpu),
normalized_advantage=None,
metadata={
@@ -776,6 +880,7 @@ def _generate_agent_turn(agent_idx: int) -> Dict[str, Any]:
"joint_attention_mask": joint_mask.detach().cpu(),
"joint_prompt_len": joint_len,
"turn_idx": turn_idx,
+ **kl_meta,
},
)
rollouts.append(sample)
@@ -792,7 +897,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
@@ -804,7 +909,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)
diff --git a/comlrl/trainers/reinforce/magrpo.py b/comlrl/trainers/reinforce/magrpo.py
index e27f200..1c91ce7 100644
--- a/comlrl/trainers/reinforce/magrpo.py
+++ b/comlrl/trainers/reinforce/magrpo.py
@@ -21,6 +21,15 @@
)
from comlrl.utils.formatters import build_formatters
from comlrl.utils.model_loading import infer_model_name, 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
from comlrl.utils.tokenizer_utils import (
apply_tokenizer_specials,
@@ -61,6 +70,9 @@ class MAGRPOConfig:
train_batch_size: Optional[int] = None
advantage_normalization: bool = True
advantage_mode: str = "mean"
+ 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.num_train_epochs < 1:
@@ -95,6 +107,7 @@ def __post_init__(self) -> None:
if mode == "mp" and self.agent_devices is None:
raise ValueError("parallel_training='mp' requires explicit agent_devices.")
self.parallel_training = mode
+ validate_reference_kl_config(self, self.num_agents)
@dataclass
@@ -228,17 +241,17 @@ def __init__(
self.agent_devices = [single_device] * self.num_agents
self.device = self.agent_devices[0]
self.dist_env = local_context(self.device)
+ model_kwargs = {}
+ torch_dtype = None
+ if isinstance(self.model_config, dict):
+ torch_dtype = self.model_config.get("torch_dtype") or self.model_config.get(
+ "dtype"
+ )
+ if torch_dtype is not None:
+ model_kwargs["torch_dtype"] = torch_dtype
if actor_sources and all(isinstance(src, str) for src in actor_sources):
from transformers import AutoModelForCausalLM
- model_kwargs = {}
- torch_dtype = None
- if isinstance(self.model_config, dict):
- torch_dtype = self.model_config.get(
- "torch_dtype"
- ) or self.model_config.get("dtype")
- if torch_dtype is not None:
- model_kwargs["torch_dtype"] = torch_dtype
self.agents = [
AutoModelForCausalLM.from_pretrained(name, **model_kwargs)
for name in actor_sources
@@ -262,6 +275,31 @@ def __init__(
for tok in self.tokenizers:
tok.add_special_tokens(special_tokens)
+ 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.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=model_kwargs,
+ )
+ else:
+ self.reference_models = clone_reference_models(
+ self.agents,
+ devices=self.reference_devices,
+ )
+ if self.tokenizers and len(self.tokenizers) == len(self.reference_models):
+ for idx, tok in enumerate(self.tokenizers):
+ apply_tokenizer_specials(tok, [self.reference_models[idx]])
+ else:
+ apply_tokenizer_specials(self.tokenizer, self.reference_models)
+
# Allow single-agent as a special case (GRPO)
if self.num_agents < 1:
raise ValueError("num_agents must be >= 1")
@@ -393,6 +431,9 @@ def _init_wandb(self):
"algorithm": self.algorithm_name,
"advantage_mode": self.advantage_mode,
"advantage_normalization": self.args.advantage_normalization,
+ "reference_kl_enabled": reference_kl_enabled(self.args),
+ "reference_kl_coef": reference_kl_coef(self.args),
+ "reference_devices": getattr(self.args, "reference_devices", None),
"agent_learning_rate": self.args.agent_learning_rate,
"num_train_epochs": self.args.num_train_epochs,
"num_generations": self.args.num_generations,
@@ -1245,15 +1286,23 @@ def _generate_completions(
batch_completions = []
batch_completion_tokens = []
+ batch_response_lens = []
end_idx = min(num_return_sequences, total_sequences)
for s in range(end_idx):
completion_tokens = completion_input_ids[s, prompt_len:]
+ pad_positions = (completion_tokens == tokenizer.pad_token_id).nonzero()
+ response_len = (
+ pad_positions[0].item()
+ if pad_positions.shape[0] > 0
+ else completion_tokens.shape[0]
+ )
batch_completion_tokens.append(completion_tokens)
+ batch_response_lens.append(int(response_len))
completion_text = tokenizer.decode(
- completion_tokens, skip_special_tokens=True
+ completion_tokens[:response_len], skip_special_tokens=True
)
batch_completions.append(completion_text)
@@ -1270,6 +1319,14 @@ def _generate_completions(
logits = (
generation_output.scores if hasattr(generation_output, "scores") else []
)
+ reference_kls = self._reference_kl_values(
+ agent_idx,
+ agent_module,
+ completion_input_ids[:end_idx],
+ torch.ones_like(completion_input_ids[:end_idx], device=device),
+ prompt_len,
+ batch_response_lens,
+ )
return {
"prompts": prompts,
@@ -1279,6 +1336,8 @@ def _generate_completions(
"completions": completions,
"completion_input_ids": completion_tokens_list,
"completion_attention_mask": completion_attention_masks,
+ "response_lens": batch_response_lens,
+ "reference_kls": reference_kls,
"logits": logits,
}
@@ -1419,8 +1478,56 @@ def _pack_completions_for_buffer(
return {
"prompt_input_ids": prompt_ids,
"completion_input_ids": packed_completion_ids,
+ "reference_kls": list(completions_data.get("reference_kls") or []),
}
+ def _reference_kl_values(
+ self,
+ agent_idx: int,
+ policy_model: Any,
+ 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 _apply_reference_kl_to_returns(
+ self, returns_tensor: torch.Tensor, completions_data: Dict[str, Any]
+ ) -> torch.Tensor:
+ if not reference_kl_enabled(self.args):
+ return returns_tensor
+ raw_kls = list(completions_data.get("reference_kls") or [])
+ if not raw_kls:
+ return returns_tensor
+ kls = torch.zeros_like(returns_tensor)
+ usable = min(len(raw_kls), returns_tensor.numel())
+ if usable > 0:
+ kls[:usable] = torch.tensor(
+ raw_kls[:usable],
+ dtype=returns_tensor.dtype,
+ device=returns_tensor.device,
+ )
+ return returns_tensor - reference_kl_coef(self.args) * kls
+
def _should_log_train(self, step: int) -> bool:
interval = int(getattr(self.args, "logging_steps", 1))
if interval <= 1:
@@ -1458,6 +1565,18 @@ def _process_buffer(
batch_log[prefix + "expected_return"] = float(
np.mean([s.node_mean_return for s in samples])
)
+ reference_kls = [
+ float(kl)
+ for sample in samples
+ for kl in sample.completions_data.get("reference_kls", [])
+ ]
+ if reference_kls:
+ batch_log[prefix + "reference_kl_mean"] = float(
+ np.mean(reference_kls)
+ )
+ batch_log[prefix + "reference_kl_penalty_mean"] = float(
+ reference_kl_coef(self.args) * np.mean(reference_kls)
+ )
step = max(s.node_env_step for s in samples)
log_entries.append(
{
@@ -1560,7 +1679,10 @@ def _compute_loss_with_gradients(self, agent, completions_data, returns):
# Convert returns to tensor
returns_tensor = torch.tensor(returns, dtype=torch.float, device=device)
- advantages = self._compute_advantages(returns_tensor)
+ effective_returns = self._apply_reference_kl_to_returns(
+ returns_tensor, completions_data
+ )
+ advantages = self._compute_advantages(effective_returns)
if self.args.advantage_normalization and advantages.numel() > 1:
mean = advantages.mean()
std = advantages.std(unbiased=False).clamp(min=1e-6)
diff --git a/comlrl/utils/__init__.py b/comlrl/utils/__init__.py
index bca4c0d..0d78c03 100644
--- a/comlrl/utils/__init__.py
+++ b/comlrl/utils/__init__.py
@@ -1,5 +1,14 @@
from .formatters import build_formatters
from .model_loading import infer_model_name, resolve_model_sources
+from .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 .reward_processor import RewardProcessors
from .reward_utils import call_reward_function, normalize_reward_lengths
from .distributed import (
@@ -27,6 +36,13 @@
"resolve_tokenizers",
"infer_model_name",
"resolve_model_sources",
+ "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",
"RewardProcessors",
"call_reward_function",
"normalize_reward_lengths",
diff --git a/comlrl/utils/reference_kl.py b/comlrl/utils/reference_kl.py
new file mode 100644
index 0000000..8f946dd
--- /dev/null
+++ b/comlrl/utils/reference_kl.py
@@ -0,0 +1,160 @@
+from __future__ import annotations
+
+import copy
+from typing import Any, Mapping, List, Sequence
+
+import torch
+import torch.nn.functional as F
+from transformers import AutoModelForCausalLM
+
+from comlrl.schedulers import DeviceScheduler
+from comlrl.utils.distributed import unwrap_model
+
+
+def reference_kl_enabled(args: Any) -> bool:
+ return bool(getattr(args, "reference_kl_enabled", False))
+
+
+def reference_kl_coef(args: Any) -> float:
+ return float(getattr(args, "reference_kl_coef", 0.1))
+
+
+def validate_reference_kl_config(args: Any, expected_count: int) -> None:
+ coef = reference_kl_coef(args)
+ if coef < 0:
+ raise ValueError("reference_kl_coef must be >= 0.")
+ reference_devices = getattr(args, "reference_devices", None)
+ if reference_kl_enabled(args) and reference_devices is not None:
+ DeviceScheduler.resolve_devices(
+ reference_devices,
+ expected_count,
+ kind="reference_devices",
+ )
+
+
+def resolve_reference_devices(
+ args: Any,
+ fallback_devices: Sequence[torch.device],
+ expected_count: int,
+) -> List[torch.device]:
+ reference_devices = getattr(args, "reference_devices", None)
+ if reference_devices is None:
+ return list(fallback_devices)
+ return DeviceScheduler.resolve_devices(
+ reference_devices,
+ expected_count,
+ kind="reference_devices",
+ )
+
+
+def clone_reference_models(
+ policy_models: Sequence[Any],
+ *,
+ devices: Sequence[torch.device],
+) -> List[Any]:
+ references: List[Any] = []
+ for idx, policy_model in enumerate(policy_models):
+ reference_model = copy.deepcopy(policy_model)
+ reference_model.to(devices[idx])
+ references.append(_freeze_reference_model(reference_model))
+ return references
+
+
+def load_reference_models_from_sources(
+ model_sources: Sequence[str],
+ *,
+ devices: Sequence[torch.device],
+ model_kwargs: Mapping[str, Any] | None = None,
+) -> List[Any]:
+ references: List[Any] = []
+ kwargs = dict(model_kwargs or {})
+ for idx, model_source in enumerate(model_sources):
+ reference_model = AutoModelForCausalLM.from_pretrained(model_source, **kwargs)
+ reference_model.to(devices[idx])
+ references.append(_freeze_reference_model(reference_model))
+ return references
+
+
+def _freeze_reference_model(model: Any) -> Any:
+ model.eval()
+ for param in model.parameters():
+ param.requires_grad = False
+ return model
+
+
+def response_token_logprobs(
+ model: Any,
+ sequences: torch.Tensor,
+ attention_mask: torch.Tensor,
+ prompt_len: int,
+ response_len: int,
+) -> torch.Tensor:
+ module = unwrap_model(model)
+ try:
+ outputs = module(
+ input_ids=sequences,
+ attention_mask=attention_mask,
+ output_values=False,
+ )
+ except TypeError:
+ outputs = module(
+ input_ids=sequences,
+ attention_mask=attention_mask,
+ return_dict=True,
+ )
+ logits = outputs.logits
+ shifted_logits = logits[:, :-1, :]
+ shifted_targets = sequences[:, 1:]
+ log_probs = F.log_softmax(shifted_logits, dim=-1)
+ token_log_probs = log_probs.gather(
+ dim=-1, index=shifted_targets.unsqueeze(-1)
+ ).squeeze(-1)
+ start_index = max(int(prompt_len) - 1, 0)
+ end_index = start_index + int(response_len)
+ return token_log_probs[:, start_index:end_index]
+
+
+def reference_kl_for_sequence(
+ policy_model: Any,
+ reference_model: Any,
+ sequences: torch.Tensor,
+ attention_mask: torch.Tensor,
+ prompt_len: int,
+ response_len: int,
+) -> torch.Tensor:
+ """
+ Return a non-negative sampled KL estimate for generated response tokens.
+
+ Uses Schulman's k3 estimator per token:
+ exp(log p_ref - log p_policy) - (log p_ref - log p_policy) - 1.
+ """
+ policy_module = unwrap_model(policy_model)
+ reference_module = unwrap_model(reference_model)
+ policy_device = next(policy_module.parameters()).device
+ reference_device = next(reference_module.parameters()).device
+ policy_seq = sequences.to(policy_device)
+ policy_mask = attention_mask.to(policy_device)
+ reference_seq = sequences.to(reference_device)
+ reference_mask = attention_mask.to(reference_device)
+ policy_training = bool(policy_module.training)
+ reference_training = bool(reference_module.training)
+ policy_module.eval()
+ reference_module.eval()
+ try:
+ with torch.no_grad():
+ policy_logps = response_token_logprobs(
+ policy_model, policy_seq, policy_mask, prompt_len, response_len
+ ).to(reference_device)
+ reference_logps = response_token_logprobs(
+ reference_model,
+ reference_seq,
+ reference_mask,
+ prompt_len,
+ response_len,
+ )
+ log_ratio_ref_policy = reference_logps - policy_logps
+ token_kl = torch.exp(log_ratio_ref_policy) - log_ratio_ref_policy - 1.0
+ finally:
+ policy_module.train(policy_training)
+ reference_module.train(reference_training)
+ return token_kl.sum(dim=-1).detach().cpu()
diff --git a/docs/assets/sponsors.jpg b/docs/assets/sponsors.jpg
deleted file mode 100644
index 9f4be2b..0000000
Binary files a/docs/assets/sponsors.jpg and /dev/null differ
diff --git a/docs/assets/sponsors.png b/docs/assets/sponsors.png
new file mode 100644
index 0000000..de3a45d
Binary files /dev/null and b/docs/assets/sponsors.png differ
diff --git a/tests/test_reference_kl.py b/tests/test_reference_kl.py
new file mode 100644
index 0000000..b57f17e
--- /dev/null
+++ b/tests/test_reference_kl.py
@@ -0,0 +1,252 @@
+from types import SimpleNamespace
+
+import pytest
+import torch
+from transformers import GPT2Config, GPT2LMHeadModel
+
+from comlrl.trainers.actor_critic import IACTrainer
+from comlrl.trainers.actor_critic.iac import IACConfig
+from comlrl.trainers.actor_critic.maac import MAACConfig, MAACTrainer
+from comlrl.trainers.reinforce import MAGRPOTrainer
+from comlrl.trainers.reinforce.magrpo import MAGRPOConfig
+from comlrl.utils.reference_kl import (
+ reference_kl_for_sequence,
+ resolve_reference_devices,
+)
+
+
+def _tiny_model(vocab_size: int = 32) -> GPT2LMHeadModel:
+ cfg = GPT2Config(
+ vocab_size=vocab_size,
+ n_positions=32,
+ n_ctx=32,
+ n_embd=16,
+ n_layer=1,
+ n_head=1,
+ )
+ return GPT2LMHeadModel(cfg)
+
+
+def _dummy_tokenizer():
+ return SimpleNamespace(
+ pad_token="",
+ eos_token="",
+ pad_token_id=0,
+ eos_token_id=1,
+ )
+
+
+def _reward_func(*_args, **_kwargs):
+ return [0.0]
+
+
+def test_reference_kl_config_defaults_off():
+ magrpo = MAGRPOConfig()
+ iac = IACConfig()
+ maac = MAACConfig()
+
+ assert magrpo.reference_kl_enabled is False
+ assert iac.reference_kl_enabled is False
+ assert maac.reference_kl_enabled is False
+ assert magrpo.reference_devices is None
+ assert iac.reference_devices is None
+ assert maac.reference_devices is None
+
+
+@pytest.mark.parametrize("config_cls", [MAGRPOConfig, IACConfig, MAACConfig])
+def test_reference_kl_rejects_negative_coef(config_cls):
+ with pytest.raises(ValueError, match="reference_kl_coef"):
+ config_cls(reference_kl_coef=-0.1)
+
+
+@pytest.mark.parametrize("config_cls", [MAGRPOConfig, IACConfig, MAACConfig])
+def test_reference_kl_rejects_mismatched_reference_devices(config_cls):
+ kwargs = {
+ "num_agents": 2,
+ "reference_kl_enabled": True,
+ "reference_devices": ["cpu", "cpu", "cpu"],
+ }
+ if config_cls is MAGRPOConfig:
+ kwargs["num_generations"] = 2
+ with pytest.raises(ValueError, match="reference_devices length"):
+ config_cls(**kwargs)
+
+
+def test_reference_devices_can_be_resolved_separately_from_agents():
+ args = SimpleNamespace(reference_devices=["cuda:1"])
+ fallback = [torch.device("cuda:0")]
+
+ devices = resolve_reference_devices(args, fallback, expected_count=1)
+
+ assert devices == [torch.device("cuda:1")]
+
+
+def test_reference_kl_enabled_uses_self_reference_by_default():
+ cfg = MAGRPOConfig(
+ num_agents=1,
+ num_turns=1,
+ num_generations=2,
+ agent_devices="cpu",
+ reference_kl_enabled=True,
+ )
+ trainer = MAGRPOTrainer(
+ agents=[_tiny_model()],
+ tokenizer=_dummy_tokenizer(),
+ reward_func=_reward_func,
+ args=cfg,
+ )
+
+ assert len(trainer.reference_models) == 1
+ assert trainer.reference_models[0] is not trainer.agents[0]
+ assert all(not p.requires_grad for p in trainer.reference_models[0].parameters())
+
+
+def test_reference_model_loads_from_actor_source(tmp_path):
+ model_dir = tmp_path / "tiny-model"
+ _tiny_model().save_pretrained(model_dir)
+
+ trainer = IACTrainer(
+ agent_model=str(model_dir),
+ tokenizer=_dummy_tokenizer(),
+ reward_func=_reward_func,
+ args=IACConfig(
+ num_agents=1,
+ num_turns=1,
+ use_separate_critic=False,
+ agent_devices="cpu",
+ reference_kl_enabled=True,
+ reference_devices="cpu",
+ ),
+ )
+
+ assert len(trainer.reference_models) == 1
+ assert isinstance(trainer.reference_models[0], GPT2LMHeadModel)
+ assert trainer.reference_models[0] is not trainer.agents[0]
+ assert all(not p.requires_grad for p in trainer.reference_models[0].parameters())
+
+
+def test_reference_kl_for_identical_model_is_zero():
+ model = _tiny_model()
+ sequences = torch.tensor([[2, 3, 4, 5]], dtype=torch.long)
+ attention_mask = torch.ones_like(sequences)
+
+ kl = reference_kl_for_sequence(
+ model,
+ model,
+ sequences,
+ attention_mask,
+ prompt_len=2,
+ response_len=2,
+ )
+
+ assert torch.allclose(kl, torch.zeros_like(kl), atol=1e-6)
+
+
+def test_reference_kl_disabled_actor_critic_shaping_is_noop():
+ iac = IACTrainer(
+ agents=[_tiny_model()],
+ tokenizer=_dummy_tokenizer(),
+ reward_func=_reward_func,
+ args=IACConfig(
+ num_agents=1,
+ num_turns=1,
+ use_separate_critic=False,
+ agent_devices="cpu",
+ critic_devices="cpu",
+ ),
+ )
+ shaped_reward, metadata = iac._kl_shaped_reward(1.0, {"reference_kls": [0.4]}, 0)
+ assert shaped_reward == pytest.approx(1.0)
+ assert metadata == {}
+
+ maac = MAACTrainer(
+ agents=[_tiny_model()],
+ critics=[_tiny_model()],
+ tokenizer=_dummy_tokenizer(),
+ reward_func=_reward_func,
+ args=MAACConfig(
+ num_agents=1,
+ num_turns=1,
+ agent_devices="cpu",
+ critic_devices="cpu",
+ ),
+ )
+ shaped_reward, metadata = maac._kl_shaped_reward(1.0, {"reference_kls": [0.4]}, 0)
+ assert shaped_reward == pytest.approx(1.0)
+ assert metadata == {}
+
+
+def test_magrpo_applies_reference_kl_to_returns():
+ args = MAGRPOConfig(
+ num_agents=1,
+ num_turns=1,
+ num_generations=2,
+ agent_devices="cpu",
+ reference_kl_enabled=True,
+ reference_kl_coef=0.5,
+ )
+ trainer = MAGRPOTrainer(
+ agents=[_tiny_model()],
+ tokenizer=_dummy_tokenizer(),
+ reward_func=_reward_func,
+ args=args,
+ )
+
+ returns = torch.tensor([1.0, 2.0])
+ adjusted = trainer._apply_reference_kl_to_returns(
+ returns, {"reference_kls": [0.2, 0.4]}
+ )
+
+ assert torch.allclose(adjusted, torch.tensor([0.9, 1.8]))
+
+
+def test_iac_uses_reference_kl_shaped_reward():
+ args = IACConfig(
+ num_agents=1,
+ num_turns=1,
+ use_separate_critic=False,
+ agent_devices="cpu",
+ critic_devices="cpu",
+ reference_kl_enabled=True,
+ reference_kl_coef=0.25,
+ )
+ trainer = IACTrainer(
+ agents=[_tiny_model()],
+ tokenizer=_dummy_tokenizer(),
+ reward_func=_reward_func,
+ args=args,
+ )
+
+ shaped_reward, metadata = trainer._kl_shaped_reward(
+ 1.0, {"reference_kls": [0.4]}, 0
+ )
+
+ assert shaped_reward == pytest.approx(0.9)
+ assert metadata["reference_kl"] == pytest.approx(0.4)
+ assert metadata["reference_kl_penalty"] == pytest.approx(0.1)
+
+
+def test_maac_uses_reference_kl_shaped_reward():
+ args = MAACConfig(
+ num_agents=1,
+ num_turns=1,
+ agent_devices="cpu",
+ critic_devices="cpu",
+ reference_kl_enabled=True,
+ reference_kl_coef=0.25,
+ )
+ trainer = MAACTrainer(
+ agents=[_tiny_model()],
+ critics=[_tiny_model()],
+ tokenizer=_dummy_tokenizer(),
+ reward_func=_reward_func,
+ args=args,
+ )
+
+ shaped_reward, metadata = trainer._kl_shaped_reward(
+ 1.0, {"reference_kls": [0.4]}, 0
+ )
+
+ assert shaped_reward == pytest.approx(0.9)
+ assert metadata["reference_kl"] == pytest.approx(0.4)
+ assert metadata["reference_kl_penalty"] == pytest.approx(0.1)