diff --git a/src/lmflow/args.py b/src/lmflow/args.py index 7cb44f86d..9ba9dabce 100644 --- a/src/lmflow/args.py +++ b/src/lmflow/args.py @@ -1344,6 +1344,9 @@ class DPOAlignerArguments: run_name: Optional[str] = field( default="dpo", metadata={"help": "The name of the run."} ) + eval_dataset_path: Optional[str] = field( + default=None, metadata={"help": "The path of the eval dataset."} + ) @dataclass diff --git a/src/lmflow/pipeline/dpo_aligner.py b/src/lmflow/pipeline/dpo_aligner.py index 19aed2c2d..42eee808a 100644 --- a/src/lmflow/pipeline/dpo_aligner.py +++ b/src/lmflow/pipeline/dpo_aligner.py @@ -71,6 +71,8 @@ def __init__(self, model_args, data_args, aligner_args): self.model_args = model_args self.data_args = data_args self.aligner_args = aligner_args + self.train_dataset = None + self.eval_dataset = None def _initialize_trainer(self, model, tokenizer): peft_config = LoraConfig( @@ -118,7 +120,7 @@ def _initialize_trainer(self, model, tokenizer): args=training_args, beta=self.aligner_args.beta, train_dataset=self.train_dataset, - eval_dataset=self.eval_dataset, + eval_dataset=self.eval_dataset if self.eval_dataset else None, tokenizer=tokenizer, peft_config=peft_config, max_prompt_length=self.aligner_args.beta, @@ -136,13 +138,14 @@ def _load_dataset(self): and len(x["prompt"]) + len(x["rejected"]) <= self.aligner_args.max_length ) # load evaluation set - self.eval_dataset = get_paired_dataset(data_root=self.data_args.dataset_path, - data_dir="test", - sanity_check=True) - self.eval_dataset = self.eval_dataset.filter( - lambda x: len(x["prompt"]) + len(x["chosen"]) <= self.aligner_args.max_length - and len(x["prompt"]) + len(x["rejected"]) <= self.aligner_args.max_length - ) + if self.aligner_args.eval_dataset_path: + self.eval_dataset = get_paired_dataset(data_root=self.aligner_args.eval_dataset_path, + data_dir="test", + sanity_check=True) + self.eval_dataset = self.eval_dataset.filter( + lambda x: len(x["prompt"]) + len(x["chosen"]) <= self.aligner_args.max_length + and len(x["prompt"]) + len(x["rejected"]) <= self.aligner_args.max_length + ) def align(self, model, dataset, reward_model): tokenizer = model.get_tokenizer()