Skip to content

Commit 4ab5a6e

Browse files
committed
(fix): fix val setting.
(cherry picked from commit beadb24b6065c8dbff89cae0e591e751d30b5076)
1 parent ea2ea84 commit 4ab5a6e

1 file changed

Lines changed: 17 additions & 6 deletions

File tree

roll/configs/base_config.py

Lines changed: 17 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)