1- import json
21import time
32from datetime import datetime
43from typing import TYPE_CHECKING , Any
@@ -64,13 +63,13 @@ class ArborReinforceJob(ReinforceJob):
6463 "max_completion_length" : None ,
6564 "gradient_checkpointing_kwargs" : None ,
6665 "bf16" : False ,
66+ "fp16" : True ,
6767 "scale_rewards" : True ,
6868 "max_grad_norm" : 1.0 ,
6969 "report_to" : "none" ,
7070 "log_completions" : True ,
7171 "logging_steps" : 10 ,
72- # By default, none is the model's max context length
73- "max_context_length" : None ,
72+ "max_seq_len" : None ,
7473 "lora_config" : None ,
7574 "loss_type" : "dapo" ,
7675 "soft_completion_penalty_length" : None ,
@@ -135,6 +134,7 @@ def initialize(self):
135134 self .DEFAULT_TRAIN_KWARGS ["mask_truncated_completions" ],
136135 )
137136 bf16 = self .train_kwargs .get ("bf16" , self .DEFAULT_TRAIN_KWARGS ["bf16" ])
137+ fp16 = self .train_kwargs .get ("fp16" , self .DEFAULT_TRAIN_KWARGS ["fp16" ])
138138 scale_rewards = self .train_kwargs .get (
139139 "scale_rewards" , self .DEFAULT_TRAIN_KWARGS ["scale_rewards" ]
140140 )
@@ -154,8 +154,8 @@ def initialize(self):
154154 logging_steps = self .train_kwargs .get (
155155 "logging_steps" , self .DEFAULT_TRAIN_KWARGS ["logging_steps" ]
156156 )
157- max_context_length = self .train_kwargs .get (
158- "max_context_length " , self .DEFAULT_TRAIN_KWARGS ["max_context_length " ]
157+ max_seq_len = self .train_kwargs .get (
158+ "max_seq_len " , self .DEFAULT_TRAIN_KWARGS ["max_seq_len " ]
159159 )
160160 max_steps = self .train_kwargs .get ("max_steps" , 500 )
161161 num_training_gpus = self .train_kwargs .get ("num_training_gpus" , 1 )
@@ -186,14 +186,14 @@ def initialize(self):
186186 "soft_completion_penalty_length" : soft_completion_penalty_length ,
187187 "mask_truncated_completions" : mask_truncated_completions ,
188188 "bf16" : bf16 ,
189+ "fp16" : fp16 ,
189190 "scale_rewards" : scale_rewards ,
190191 "gradient_checkpointing_kwargs" : gradient_checkpointing_kwargs ,
191192 "max_grad_norm" : max_grad_norm ,
192193 "report_to" : report_to ,
193194 "log_completions" : log_completions ,
194195 "logging_steps" : logging_steps ,
195- # "max_context_length": max_context_length,
196- # "max_seq_len": max_context_length,
196+ "max_seq_len" : max_seq_len ,
197197 "max_steps" : max_steps ,
198198 "loss_type" : loss_type ,
199199 "num_training_gpus" : num_training_gpus ,
@@ -202,7 +202,7 @@ def initialize(self):
202202 },
203203 "inference_config" : {
204204 "model" : finetune_model ,
205- "max_context_length " : max_context_length ,
205+ "max_seq_len " : max_seq_len ,
206206 },
207207 "gpu_config" : {
208208 "type" : "multi" ,
@@ -214,12 +214,21 @@ def initialize(self):
214214 }
215215 url = urljoin (api_base , "fine_tuning/grpo/initialize" )
216216 headers = {"Content-Type" : "application/json" }
217- response = requests .post (url = url , headers = headers , json = data )
218- print (json .dumps (response .json (), indent = 2 ))
219- response .raise_for_status ()
220- response = response .json ()
221- self .lm .model = ArborProvider ._add_provider_prefix (response ["current_model" ])
222- self .provider_job_id = response .get ("job_id" )
217+ try :
218+ response = requests .post (url = url , headers = headers , json = data )
219+ response .raise_for_status ()
220+ response = response .json ()
221+ self .lm .model = ArborProvider ._add_provider_prefix (
222+ response ["current_model" ]
223+ )
224+ self .provider_job_id = response .get ("job_id" )
225+ except KeyboardInterrupt :
226+ print (
227+ "Keyboard interrupt received. Stopping reinforcement learning job initialization."
228+ )
229+ except Exception as err :
230+ print (f"Error initializing reinforcement learning job: { err } " )
231+ raise err
223232
224233 def _run_grpo_step_one_group (
225234 self ,
0 commit comments