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. Email 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)