From 49db8a3bf340f93f83a46b469bc9d4691a51590f Mon Sep 17 00:00:00 2001 From: N!no Date: Tue, 30 Jun 2026 15:57:23 -0400 Subject: [PATCH 1/2] add kl --- train_ac.py | 2 ++ train_grpo.py | 2 ++ train_iac.py | 2 ++ train_maac.py | 2 ++ train_magrpo.py | 2 ++ 5 files changed, 10 insertions(+) diff --git a/train_ac.py b/train_ac.py index 6488f8d..8cf5323 100644 --- a/train_ac.py +++ b/train_ac.py @@ -301,6 +301,8 @@ def main() -> None: eval_num_samples=ac_cfg.get("eval_num_samples", 4), eval_batch_size=ac_cfg.get("eval_batch_size", 1), logging_steps=ac_cfg.get("logging_steps", 1), + reference_kl_enabled=ac_cfg.get("reference_kl_enabled", False), + reference_kl_coef=ac_cfg.get("reference_kl_coef", 0.1), ), train_dataset=train_dataset, eval_dataset=eval_dataset, diff --git a/train_grpo.py b/train_grpo.py index 2a3a903..9571204 100644 --- a/train_grpo.py +++ b/train_grpo.py @@ -235,6 +235,8 @@ def main(): eval_batch_size=grpo_cfg.get("eval_batch_size", 1), train_batch_size=grpo_cfg.get("train_batch_size"), advantage_normalization=grpo_cfg.get("advantage_normalization", True), + reference_kl_enabled=grpo_cfg.get("reference_kl_enabled", False), + reference_kl_coef=grpo_cfg.get("reference_kl_coef", 0.1), ) import rewards.arxiv_rewards as arxiv_rewards diff --git a/train_iac.py b/train_iac.py index 7c3736d..a58294b 100644 --- a/train_iac.py +++ b/train_iac.py @@ -328,6 +328,8 @@ def main() -> None: eval_num_samples=iac_cfg.get("eval_num_samples", 4), eval_batch_size=iac_cfg.get("eval_batch_size", 1), logging_steps=iac_cfg.get("logging_steps", 50), + reference_kl_enabled=iac_cfg.get("reference_kl_enabled", False), + reference_kl_coef=iac_cfg.get("reference_kl_coef", 0.1), ), train_dataset=train_dataset, eval_dataset=eval_dataset, diff --git a/train_maac.py b/train_maac.py index 167621c..1dab718 100644 --- a/train_maac.py +++ b/train_maac.py @@ -323,6 +323,8 @@ def main() -> None: eval_num_samples=maac_cfg.get("eval_num_samples", 4), eval_batch_size=maac_cfg.get("eval_batch_size", 1), logging_steps=maac_cfg.get("logging_steps", 50), + reference_kl_enabled=maac_cfg.get("reference_kl_enabled", False), + reference_kl_coef=maac_cfg.get("reference_kl_coef", 0.1), ), train_dataset=train_dataset, eval_dataset=eval_dataset, diff --git a/train_magrpo.py b/train_magrpo.py index 9367113..f186d83 100644 --- a/train_magrpo.py +++ b/train_magrpo.py @@ -332,6 +332,8 @@ def main(): eval_interval=magrpo_cfg.get("eval_interval", 20), eval_num_samples=magrpo_cfg.get("eval_num_samples", 4), eval_batch_size=magrpo_cfg.get("eval_batch_size", 1), + reference_kl_enabled=magrpo_cfg.get("reference_kl_enabled", False), + reference_kl_coef=magrpo_cfg.get("reference_kl_coef", 0.1), ) import rewards.arxiv_rewards as arxiv_rewards From fe228f5512d9c43cd4843fd9f900e34cd9e840d8 Mon Sep 17 00:00:00 2001 From: N!no Date: Tue, 30 Jun 2026 17:30:41 -0400 Subject: [PATCH 2/2] allow ref on separate devices --- train_ac.py | 1 + train_grpo.py | 1 + train_iac.py | 1 + train_maac.py | 1 + train_magrpo.py | 1 + 5 files changed, 5 insertions(+) diff --git a/train_ac.py b/train_ac.py index 8cf5323..41c8f29 100644 --- a/train_ac.py +++ b/train_ac.py @@ -303,6 +303,7 @@ def main() -> None: logging_steps=ac_cfg.get("logging_steps", 1), reference_kl_enabled=ac_cfg.get("reference_kl_enabled", False), reference_kl_coef=ac_cfg.get("reference_kl_coef", 0.1), + reference_devices=ac_cfg.get("reference_devices", None), ), train_dataset=train_dataset, eval_dataset=eval_dataset, diff --git a/train_grpo.py b/train_grpo.py index 9571204..7a4a910 100644 --- a/train_grpo.py +++ b/train_grpo.py @@ -237,6 +237,7 @@ def main(): advantage_normalization=grpo_cfg.get("advantage_normalization", True), reference_kl_enabled=grpo_cfg.get("reference_kl_enabled", False), reference_kl_coef=grpo_cfg.get("reference_kl_coef", 0.1), + reference_devices=grpo_cfg.get("reference_devices", None), ) import rewards.arxiv_rewards as arxiv_rewards diff --git a/train_iac.py b/train_iac.py index a58294b..80142e7 100644 --- a/train_iac.py +++ b/train_iac.py @@ -330,6 +330,7 @@ def main() -> None: logging_steps=iac_cfg.get("logging_steps", 50), reference_kl_enabled=iac_cfg.get("reference_kl_enabled", False), reference_kl_coef=iac_cfg.get("reference_kl_coef", 0.1), + reference_devices=iac_cfg.get("reference_devices", None), ), train_dataset=train_dataset, eval_dataset=eval_dataset, diff --git a/train_maac.py b/train_maac.py index 1dab718..770658e 100644 --- a/train_maac.py +++ b/train_maac.py @@ -325,6 +325,7 @@ def main() -> None: logging_steps=maac_cfg.get("logging_steps", 50), reference_kl_enabled=maac_cfg.get("reference_kl_enabled", False), reference_kl_coef=maac_cfg.get("reference_kl_coef", 0.1), + reference_devices=maac_cfg.get("reference_devices", None), ), train_dataset=train_dataset, eval_dataset=eval_dataset, diff --git a/train_magrpo.py b/train_magrpo.py index f186d83..68138ea 100644 --- a/train_magrpo.py +++ b/train_magrpo.py @@ -334,6 +334,7 @@ def main(): eval_batch_size=magrpo_cfg.get("eval_batch_size", 1), reference_kl_enabled=magrpo_cfg.get("reference_kl_enabled", False), reference_kl_coef=magrpo_cfg.get("reference_kl_coef", 0.1), + reference_devices=magrpo_cfg.get("reference_devices", None), ) import rewards.arxiv_rewards as arxiv_rewards