diff --git a/bridge_build/utils/trainer_args.py b/bridge_build/utils/trainer_args.py index 901b9d2..7fc3f30 100644 --- a/bridge_build/utils/trainer_args.py +++ b/bridge_build/utils/trainer_args.py @@ -190,6 +190,11 @@ 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), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } ) @@ -248,6 +253,11 @@ 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), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: @@ -309,6 +319,11 @@ 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), + "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 8ff6a66..3accd3f 100644 --- a/house_build/utils/trainer_args.py +++ b/house_build/utils/trainer_args.py @@ -190,6 +190,11 @@ 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), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } ) @@ -248,6 +253,11 @@ 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), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: @@ -309,6 +319,11 @@ 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), + "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 901b9d2..7fc3f30 100644 --- a/str_build/utils/trainer_args.py +++ b/str_build/utils/trainer_args.py @@ -190,6 +190,11 @@ 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), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } ) @@ -248,6 +253,11 @@ 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), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: @@ -309,6 +319,11 @@ 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), + "reference_devices": _as_device_spec(tr.get("reference_devices", None)), } try: