diff --git a/configs/context_numbers_10.yaml b/configs/context_numbers_10.yaml index 87cdeb4..4106235 100644 --- a/configs/context_numbers_10.yaml +++ b/configs/context_numbers_10.yaml @@ -20,6 +20,7 @@ per_device_train_batch_size: 64 per_device_eval_batch_size: 128 max_new_tokens: 64 gen_per_device_eval_batch_size: 128 +max_val_samples_per_ds: 500 # optim: schedule_free_adamw learning_rate: 0.0001 # lr_scheduler_type: "constant_with_warmup" diff --git a/configs/context_numbers_128.yaml b/configs/context_numbers_128.yaml index 8c3c06c..37e29eb 100644 --- a/configs/context_numbers_128.yaml +++ b/configs/context_numbers_128.yaml @@ -18,6 +18,7 @@ label_names: ["labels"] per_device_train_batch_size: 64 per_device_eval_batch_size: 64 +max_val_samples_per_ds: 50 # optim: schedule_free_adamw learning_rate: 0.0001 # lr_scheduler_type: "constant_with_warmup" diff --git a/configs/context_numbers_256.yaml b/configs/context_numbers_256.yaml index 2f4b079..3e639b1 100644 --- a/configs/context_numbers_256.yaml +++ b/configs/context_numbers_256.yaml @@ -18,6 +18,7 @@ label_names: ["labels"] per_device_train_batch_size: 32 per_device_eval_batch_size: 32 +max_val_samples_per_ds: 20 # optim: schedule_free_adamw learning_rate: 0.0001 # lr_scheduler_type: "constant_with_warmup" diff --git a/hyperlora/configs.py b/hyperlora/configs.py index 752a4a0..bfb4b99 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -243,9 +243,9 @@ class DataArguments: default=None, metadata={"help": "Test dataset names."}, ) - max_val_samples: Optional[int] = field( + max_val_samples_per_ds: Optional[int] = field( default=5000, - metadata={"help": "Maximum number of validation samples."}, + metadata={"help": "Maximum number of validation samples per dataset."}, ) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 2b72c42..dbb5c88 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -278,9 +278,15 @@ def main(output_dir): ): if ds_names is None: continue - tokenized_ds[split] = concatenate_datasets( - [_get_tokenized_dataset(ds_name, split) for ds_name in ds_names] - ) + if split == "train": + tokenized_ds[split] = concatenate_datasets( + [_get_tokenized_dataset(ds_name, split) for ds_name in ds_names] + ) + else: + tokenized_ds[split] = { + os.path.basename(ds_name): _get_tokenized_dataset(ds_name, split) + for ds_name in ds_names + } train_ds = tokenized_ds["train"] val_train_indices = np.random.permutation(len(train_ds))[:500] @@ -288,10 +294,13 @@ def main(output_dir): "train": tokenized_ds["train"].select(val_train_indices), } if "validation" in tokenized_ds: - val_ds["val"] = tokenized_ds["validation"] - val_ds_size = len(val_ds["val"]) - val_indices = np.random.permutation(val_ds_size)[: data_args.max_val_samples] - val_ds["val"] = val_ds["val"].select(val_indices) + for ds_name, ds in tokenized_ds["validation"].items(): + val_ds[ds_name] = ds + val_ds_size = len(val_ds[ds_name]) + val_indices = np.random.permutation(val_ds_size)[ + : data_args.max_val_samples_per_ds + ] + val_ds[ds_name] = val_ds[ds_name].select(val_indices) test_ds = tokenized_ds.get("test", None) logger.info(f"train_ds: {train_ds}") @@ -406,6 +415,7 @@ def main(output_dir): # compute_metrics, # preprocess_logits_for_metrics, ) + logger.info(f"Training run finished and saved to {output_dir}") if __name__ == "__main__": diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 6990be5..f47af7f 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -74,27 +74,32 @@ def decode_test_result(test_dataset, test_result, tokenizer): def eval_generation(eval_trainer, tokenizer, dataset, split, gen_kwargs): - eval_result = eval_trainer.predict( - dataset, - metric_key_prefix=split, - **gen_kwargs, - ) + if not isinstance(dataset, dict): + dataset = {"": dataset} - decoded_txts = decode_test_result(dataset, eval_result, tokenizer) - rouge_metrics = compute_rouge( - [txt["generated"] for txt in decoded_txts], - [txt["label"] for txt in decoded_txts], - ) - for k, v in rouge_metrics.items(): - eval_result.metrics[f"{split}_{k}"] = v - eval_trainer.log_metrics("eval" if split == "val" else split, eval_result.metrics) - eval_trainer.save_metrics("eval" if split == "val" else split, eval_result.metrics) + for ds_name, ds in dataset.items(): + split_name = f"{split}_{ds_name}" if ds_name else split + eval_result = eval_trainer.predict( + ds, + metric_key_prefix=split_name, + **gen_kwargs, + ) + decoded_txts = decode_test_result(ds, eval_result, tokenizer) + rouge_metrics = compute_rouge( + [txt["generated"] for txt in decoded_txts], + [txt["label"] for txt in decoded_txts], + ) + for k, v in rouge_metrics.items(): + eval_result.metrics[f"{split}_{k}"] = v - save_generated_text( - decoded_txts, - split=split, - output_dir=eval_trainer.args.output_dir, - ) + save_generated_text( + decoded_txts, + split=split_name, + output_dir=eval_trainer.args.output_dir, + ) + eval_trainer.log_metrics(split_name, eval_result.metrics) + eval_trainer.save_metrics(split_name, eval_result.metrics) + clear_gpu() # def per_sample_loss_avg_fn(outputs, labels, num_items_in_batch): @@ -230,12 +235,11 @@ def train_model( data_collator=generation_collator, ) - # TODO: log different datasets separately - if val_dataset is not None: - if isinstance(val_dataset, dict): - val_dataset = val_dataset["val"] - eval_generation(eval_trainer, tokenizer, val_dataset, "val", gen_kwargs) + for split, ds in zip(["eval", "test"], [val_dataset, test_dataset]): + if ds is None: + continue + eval_generation(eval_trainer, tokenizer, ds, split, gen_kwargs) clear_gpu() - if test_dataset is not None: - eval_generation(eval_trainer, tokenizer, test_dataset, "test", gen_kwargs) + # if test_dataset is not None: + # eval_generation(eval_trainer, tokenizer, test_dataset, "test", gen_kwargs)