From 2ea0683a9114531e621f6ac84a500f1f7a01d09a Mon Sep 17 00:00:00 2001 From: N!no Date: Tue, 30 Jun 2026 15:57:37 -0400 Subject: [PATCH 1/2] add kl --- bridge_build/utils/trainer_args.py | 12 ++++++++++++ house_build/utils/trainer_args.py | 12 ++++++++++++ str_build/utils/trainer_args.py | 12 ++++++++++++ 3 files changed, 36 insertions(+) diff --git a/bridge_build/utils/trainer_args.py b/bridge_build/utils/trainer_args.py index 901b9d2..8ed52d5 100644 --- a/bridge_build/utils/trainer_args.py +++ b/bridge_build/utils/trainer_args.py @@ -190,6 +190,10 @@ def get_trainer_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MA "external_prompt_passthrough": _as_bool( ext.get("external_prompt_passthrough", False), False ), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } ) @@ -248,6 +252,10 @@ def get_maac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MAACC "eval_num_samples": _as_int(tr.get("eval_num_samples", 2), 2), "eval_batch_size": _as_int(tr.get("eval_batch_size", 1), 1), "logging_steps": _as_int(tr.get("logging_steps", 20), 20), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } try: @@ -309,6 +317,10 @@ def get_iac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> IACCon "eval_num_samples": _as_int(tr.get("eval_num_samples", 2), 2), "eval_batch_size": _as_int(tr.get("eval_batch_size", 1), 1), "logging_steps": _as_int(tr.get("logging_steps", 20), 20), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } try: diff --git a/house_build/utils/trainer_args.py b/house_build/utils/trainer_args.py index 8ff6a66..fff0de9 100644 --- a/house_build/utils/trainer_args.py +++ b/house_build/utils/trainer_args.py @@ -190,6 +190,10 @@ def get_trainer_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MA "external_prompt_passthrough": _as_bool( ext.get("external_prompt_passthrough", False), False ), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } ) @@ -248,6 +252,10 @@ def get_maac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MAACC "eval_num_samples": _as_int(tr.get("eval_num_samples", 2), 2), "eval_batch_size": _as_int(tr.get("eval_batch_size", 1), 1), "logging_steps": _as_int(tr.get("logging_steps", 40), 40), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } try: @@ -309,6 +317,10 @@ def get_iac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> IACCon "eval_num_samples": _as_int(tr.get("eval_num_samples", 2), 2), "eval_batch_size": _as_int(tr.get("eval_batch_size", 1), 1), "logging_steps": _as_int(tr.get("logging_steps", 40), 40), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } try: diff --git a/str_build/utils/trainer_args.py b/str_build/utils/trainer_args.py index 901b9d2..8ed52d5 100644 --- a/str_build/utils/trainer_args.py +++ b/str_build/utils/trainer_args.py @@ -190,6 +190,10 @@ def get_trainer_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MA "external_prompt_passthrough": _as_bool( ext.get("external_prompt_passthrough", False), False ), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } ) @@ -248,6 +252,10 @@ def get_maac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MAACC "eval_num_samples": _as_int(tr.get("eval_num_samples", 2), 2), "eval_batch_size": _as_int(tr.get("eval_batch_size", 1), 1), "logging_steps": _as_int(tr.get("logging_steps", 20), 20), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } try: @@ -309,6 +317,10 @@ def get_iac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> IACCon "eval_num_samples": _as_int(tr.get("eval_num_samples", 2), 2), "eval_batch_size": _as_int(tr.get("eval_batch_size", 1), 1), "logging_steps": _as_int(tr.get("logging_steps", 20), 20), + "reference_kl_enabled": _as_bool( + tr.get("reference_kl_enabled", False), False + ), + "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), } try: From ac655b2829a2722d10c1b11f7bef1f8a00ad3311 Mon Sep 17 00:00:00 2001 From: N!no Date: Tue, 30 Jun 2026 17:30:35 -0400 Subject: [PATCH 2/2] allow ref on separate devices --- bridge_build/utils/trainer_args.py | 3 +++ house_build/utils/trainer_args.py | 3 +++ str_build/utils/trainer_args.py | 3 +++ 3 files changed, 9 insertions(+) diff --git a/bridge_build/utils/trainer_args.py b/bridge_build/utils/trainer_args.py index 8ed52d5..7fc3f30 100644 --- a/bridge_build/utils/trainer_args.py +++ b/bridge_build/utils/trainer_args.py @@ -194,6 +194,7 @@ def get_trainer_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MA tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } ) @@ -256,6 +257,7 @@ def get_maac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MAACC tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: @@ -321,6 +323,7 @@ def get_iac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> IACCon tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: diff --git a/house_build/utils/trainer_args.py b/house_build/utils/trainer_args.py index fff0de9..3accd3f 100644 --- a/house_build/utils/trainer_args.py +++ b/house_build/utils/trainer_args.py @@ -194,6 +194,7 @@ def get_trainer_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MA tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } ) @@ -256,6 +257,7 @@ def get_maac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MAACC tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: @@ -321,6 +323,7 @@ def get_iac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> IACCon tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: diff --git a/str_build/utils/trainer_args.py b/str_build/utils/trainer_args.py index 8ed52d5..7fc3f30 100644 --- a/str_build/utils/trainer_args.py +++ b/str_build/utils/trainer_args.py @@ -194,6 +194,7 @@ def get_trainer_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MA tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } ) @@ -256,6 +257,7 @@ def get_maac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> MAACC tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: @@ -321,6 +323,7 @@ def get_iac_args(cfg: Dict[str, Any], *, sampling_cfg: Dict[str, Any]) -> IACCon tr.get("reference_kl_enabled", False), False ), "reference_kl_coef": _as_float(tr.get("reference_kl_coef", 0.1), 0.1), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: