mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
eval unpack + fw_qa_level_3 data def + fast tokenizer
This commit is contained in:
parent
d264b2c76f
commit
1ec757f43a
3 changed files with 55 additions and 21 deletions
|
|
@ -150,6 +150,27 @@ DS_KWARGS = {
|
|||
split="train",
|
||||
),
|
||||
),
|
||||
"fw_qa_v2_2k_len_level_3_tiny": dict(
|
||||
train=dict(
|
||||
path="parquet",
|
||||
data_files="data/raw_datasets/fw_qa_v2/min_0_to_2000/000_00000_level_3.parquet",
|
||||
split="train",
|
||||
),
|
||||
),
|
||||
"fw_qa_v2_2k_len_level_3": dict(
|
||||
train=dict(
|
||||
path="parquet",
|
||||
data_files="data/raw_datasets/fw_qa_v2/min_0_to_2000/*level_3.parquet",
|
||||
split="train",
|
||||
),
|
||||
# validation=dict(
|
||||
# path="parquet",
|
||||
# data_files=glob(
|
||||
# "data/raw_datasets/fw_qa_v2/min_0_to_2000/*level_3_val.parquet"
|
||||
# ),
|
||||
# split="train",
|
||||
# ),
|
||||
),
|
||||
# "fw_qa_3_mini": dict(
|
||||
# train=dict(
|
||||
# path="parquet",
|
||||
|
|
|
|||
|
|
@ -250,7 +250,7 @@ if __name__ == "__main__":
|
|||
tokenizer_kwargs = {"max_length": base_model_max_len} # not used
|
||||
ctx_tokenizer_kwargs = {"max_length": base_model_max_len} # not used for now
|
||||
ds = get_tokenized_dataset(
|
||||
ds_name="squad",
|
||||
ds_name="fw_qa_v2_2k_len_level_0_tiny",
|
||||
split="train",
|
||||
base_model_max_len=model.base_model.config.max_position_embeddings,
|
||||
tokenizer=tokenizer,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import random
|
||||
from collections.abc import Callable
|
||||
from glob import glob
|
||||
|
|
@ -70,27 +71,27 @@ def get_preprocessing_fn(
|
|||
|
||||
if "fw_qa_v2" in ds_name and "level" in ds_name:
|
||||
|
||||
def f_train(sample):
|
||||
def f(sample):
|
||||
# get questions/answers from all levels in the ds
|
||||
q_cols = [col for col in sample.keys() if col.startswith("prompts_level")]
|
||||
a_cols = [col for col in sample.keys() if col.startswith("answers_level")]
|
||||
r_cols = [col for col in sample.keys() if col.startswith("responses_level")]
|
||||
questions = concat_list([sample[col] for col in q_cols])
|
||||
answers = concat_list([sample[col] for col in a_cols])
|
||||
responses = concat_list([sample[col] for col in r_cols])
|
||||
min_len = min(len(questions), len(responses))
|
||||
|
||||
if min_len == 0:
|
||||
return {
|
||||
"context": None,
|
||||
"prompts": None,
|
||||
"responses": None,
|
||||
}
|
||||
|
||||
return {
|
||||
"context": sample["context"],
|
||||
"prompts": questions,
|
||||
"responses": answers,
|
||||
"prompts": questions[:min_len],
|
||||
"responses": responses[:min_len],
|
||||
}
|
||||
|
||||
def f_eval(sample):
|
||||
return {
|
||||
"context": sample["context"],
|
||||
"prompt": sample["prompt"],
|
||||
"response": sample["response"],
|
||||
}
|
||||
|
||||
f = f_eval if is_eval else f_train
|
||||
|
||||
elif ds_name.startswith("longbench"):
|
||||
|
||||
def f(sample):
|
||||
|
|
@ -262,7 +263,10 @@ def get_preprocessing_fn(
|
|||
def eval_intx_decorator(f):
|
||||
def g(sample):
|
||||
sample = f(sample)
|
||||
sample["prompt"] = prompt_template.format(input=sample["prompt"])
|
||||
prompt_field = "prompt" if "prompt" in sample else "prompts"
|
||||
sample[prompt_field] = prompt_template.format(
|
||||
input=sample[prompt_field]
|
||||
)
|
||||
return sample
|
||||
|
||||
return g
|
||||
|
|
@ -693,6 +697,7 @@ def construct_and_tokenize_ctx_qa(
|
|||
remove_columns=[col for col in ds.column_names if col != "context"],
|
||||
)
|
||||
# add input_ids, attention_mask, labels to "data"
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "true"
|
||||
tokenized_ds = ds.map(
|
||||
get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer),
|
||||
batched=True,
|
||||
|
|
@ -746,6 +751,7 @@ def construct_and_tokenize_ctx_qa(
|
|||
|
||||
if need_ctx_ids:
|
||||
# tokenize the ctx_text to get ctx_ids and ctx_attn_mask
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "true"
|
||||
tokenized_ds = tokenized_ds.map(
|
||||
tokenize_ctx_text,
|
||||
fn_kwargs={"tokenizer": ctx_tokenizer},
|
||||
|
|
@ -767,13 +773,16 @@ def construct_and_tokenize_ctx_qa(
|
|||
tokenized_ds = tokenized_ds.map(
|
||||
unpack_data,
|
||||
num_proc=num_proc,
|
||||
remove_columns=["data"],
|
||||
remove_columns=["data", "context"],
|
||||
)
|
||||
else:
|
||||
cols_to_remove = ["data", "context"]
|
||||
if "ctx_ids" in tokenized_ds.column_names:
|
||||
cols_to_remove += ["ctx_ids", "ctx_attn_mask"]
|
||||
tokenized_ds = tokenized_ds.map(
|
||||
unpack_data_eval,
|
||||
num_proc=num_proc,
|
||||
remove_columns=["data"],
|
||||
remove_columns=cols_to_remove,
|
||||
batched=True,
|
||||
batch_size=100_000,
|
||||
)
|
||||
|
|
@ -808,8 +817,6 @@ def construct_and_tokenize_ctx_qa(
|
|||
tokenized_ds = tokenized_ds.remove_columns(["context", "text"])
|
||||
if not is_paraphrased:
|
||||
tokenized_ds = tokenized_ds.remove_columns(["qas"])
|
||||
else:
|
||||
tokenized_ds = tokenized_ds.remove_columns(["context"])
|
||||
|
||||
if set_format:
|
||||
tokenized_ds.set_format(type=set_format)
|
||||
|
|
@ -955,12 +962,18 @@ def unpack_data_eval(samples):
|
|||
# n_queries always == 1 for eval
|
||||
data = samples["data"]
|
||||
out = dict(input_ids=[], attention_mask=[], labels=[])
|
||||
for d in data:
|
||||
if "ctx_ids" in samples:
|
||||
out["ctx_ids"] = []
|
||||
out["ctx_attn_mask"] = []
|
||||
for i, d in enumerate(data):
|
||||
for tokens in zip(
|
||||
d["input_ids"],
|
||||
d["attention_mask"],
|
||||
d["labels"],
|
||||
):
|
||||
if "ctx_ids" in samples:
|
||||
out["ctx_ids"].append(samples["ctx_ids"][i])
|
||||
out["ctx_attn_mask"].append(samples["ctx_attn_mask"][i])
|
||||
out["input_ids"].append(tokens[0])
|
||||
out["attention_mask"].append(tokens[1])
|
||||
out["labels"].append(tokens[2])
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue