mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
add val split for ds w/o one
This commit is contained in:
parent
71f7a664c6
commit
a6104c2c68
1 changed files with 8 additions and 4 deletions
|
|
@ -284,10 +284,7 @@ def main(output_dir):
|
|||
}
|
||||
|
||||
train_ds = tokenized_ds["train"]
|
||||
val_train_indices = np.random.permutation(len(train_ds))[:500]
|
||||
val_ds = {
|
||||
"train": tokenized_ds["train"].select(val_train_indices),
|
||||
}
|
||||
val_ds = dict()
|
||||
if "validation" in tokenized_ds:
|
||||
for ds_name, ds in tokenized_ds["validation"].items():
|
||||
val_ds[ds_name] = ds
|
||||
|
|
@ -296,6 +293,13 @@ def main(output_dir):
|
|||
: data_args.max_val_samples_per_ds
|
||||
]
|
||||
val_ds[ds_name] = val_ds[ds_name].select(val_indices)
|
||||
else:
|
||||
# take some samples from train_ds
|
||||
n_val_samples = data_args.max_val_samples_per_ds
|
||||
train_ds = train_ds.skip(n_val_samples)
|
||||
val_ds["train_unseen"] = train_ds.take(n_val_samples)
|
||||
val_train_indices = np.random.permutation(len(train_ds))[:500]
|
||||
val_ds["train"] = train_ds.select(val_train_indices)
|
||||
test_ds = tokenized_ds.get("test", None)
|
||||
|
||||
logger.info(f"train_ds: {train_ds}")
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue