@@ -119,6 +119,14 @@ class BaseConfig:
119119 default = None ,
120120 metadata = {"help" : "The maximum length of the sequence to be padded." },
121121 )
122+ val_prompt_length : Optional [int ] = field (
123+ default = None ,
124+ metadata = {"help" : "The maximum length of a prompt to be padded." },
125+ )
126+ val_sequence_length : Optional [int ] = field (
127+ default = None ,
128+ metadata = {"help" : "The maximum length of the sequence to be padded." },
129+ )
122130 alive_check_interval : int = field (
123131 default = 10 ,
124132 metadata = {"help" : "The interval of worker alive check." }
@@ -159,16 +167,19 @@ def __post_init__(self):
159167
160168 if self .sequence_length is None :
161169 self .sequence_length = self .response_length + self .prompt_length
162- logger .warning (
163- f"sequence_length is not set, use response_length + prompt_length as sequence_length: { self .sequence_length } "
164- )
165170
166171 if self .response_length is not None :
167- logger .warning (
168- f"response_length is deprecated, use sequence_length instead, sequence_length is { self .sequence_length } "
169- )
170172 self .response_length = None
171173
174+ if self .val_prompt_length is None :
175+ assert self .val_sequence_length is None , "val_prompt_length and val_sequence_length must be set simultaneously"
176+ self .val_prompt_length = self .prompt_length
177+ self .val_sequence_length = self .sequence_length
178+
179+ if self .val_prompt_length is not None :
180+ assert self .val_sequence_length , "val_prompt_length and val_sequence_length must be set simultaneously"
181+
182+
172183 if self .track_with == "tensorboard" :
173184 self .tracker_kwargs ["log_dir" ] = os .path .join (
174185 self .tracker_kwargs .get ("log_dir" , self .output_dir ), self .exp_name , datetime .now ().strftime ("%Y%m%d-%H%M%S" )
0 commit comments