eval unpack + fw_qa_level_3 data def + fast tokenizer

This commit is contained in:
51616 2025-06-24 21:51:17 +09:00
parent d264b2c76f
commit 1ec757f43a
3 changed files with 55 additions and 21 deletions

View file

@ -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",

View file

@ -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,

View file

@ -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])