tiny configs + fix dataset gen (remove repeated response + qa template for compact ds) + add ropes and drop eval ds

This commit is contained in:
51616 2025-06-30 16:10:52 +09:00
parent e4cb1ffa55
commit 28eb85a2f3
19 changed files with 626 additions and 36 deletions

View file

@ -0,0 +1,58 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# short ctx data
train_ds_names:
- squad_compact
- pwc_compact
- drop_compact
- ropes_compact
# the contexts are kinda long-ish
# - narrativeqa
# - quoref
# - synthetic_convqa
val_ds_names:
- fw_qa_v2_2k_len_level_0
- squad
- drop
- ropes
- pwc

View file

@ -39,7 +39,7 @@ target_modules:
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it/fw_qa_v2_2k_len_level_3
- self_gen/google/gemma-2-2b-it/fw_qa_v2/min_0_to_2000
- self_gen/google/gemma-2-2b-it/squad_compact
- self_gen/google/gemma-2-2b-it/pwc_compact
- self_gen/google/gemma-2-2b-it/ropes_compact

View file

@ -0,0 +1,52 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.0/fw_qa_v2_2k_len_level_0_tiny
- self_gen/google/gemma-2-2b-it_temp_0.0/squad_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/pwc_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/ropes_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/drop_compact
val_ds_names:
- squad
- pwc
- drop
- ropes

View file

@ -0,0 +1,52 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.0/fw_qa_v2_2k_len_level_1_tiny
- self_gen/google/gemma-2-2b-it_temp_0.0/squad_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/pwc_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/ropes_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/drop_compact
val_ds_names:
- squad
- pwc
- drop
- ropes

View file

@ -0,0 +1,52 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.0/fw_qa_v2_2k_len_level_2_tiny
- self_gen/google/gemma-2-2b-it_temp_0.0/squad_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/pwc_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/ropes_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/drop_compact
val_ds_names:
- squad
- pwc
- drop
- ropes

View file

@ -0,0 +1,52 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.0/fw_qa_v2_2k_len_level_3_tiny
- self_gen/google/gemma-2-2b-it_temp_0.0/squad_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/pwc_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/ropes_compact
- self_gen/google/gemma-2-2b-it_temp_0.0/drop_compact
val_ds_names:
- squad
- pwc
- drop
- ropes

View file

@ -0,0 +1,52 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.3/fw_qa_v2_2k_len_level_3_tiny
- self_gen/google/gemma-2-2b-it_temp_0.3/squad_compact
- self_gen/google/gemma-2-2b-it_temp_0.3/pwc_compact
- self_gen/google/gemma-2-2b-it_temp_0.3/ropes_compact
- self_gen/google/gemma-2-2b-it_temp_0.3/drop_compact
val_ds_names:
- squad
- pwc
- drop
- ropes

View file

@ -0,0 +1,52 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.5/fw_qa_v2_2k_len_level_3_tiny
- self_gen/google/gemma-2-2b-it_temp_0.5/squad_compact
- self_gen/google/gemma-2-2b-it_temp_0.5/pwc_compact
- self_gen/google/gemma-2-2b-it_temp_0.5/ropes_compact
- self_gen/google/gemma-2-2b-it_temp_0.5/drop_compact
val_ds_names:
- squad
- pwc
- drop
- ropes

View file

@ -0,0 +1,52 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_0.7/fw_qa_v2_2k_len_level_3_tiny
- self_gen/google/gemma-2-2b-it_temp_0.7/squad_compact
- self_gen/google/gemma-2-2b-it_temp_0.7/pwc_compact
- self_gen/google/gemma-2-2b-it_temp_0.7/ropes_compact
- self_gen/google/gemma-2-2b-it_temp_0.7/drop_compact
val_ds_names:
- squad
- pwc
- drop
- ropes

View file

@ -0,0 +1,52 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: google/gemma-3-1b-it
label_names: ["labels"]
# eval_on_start: True
# eval_strategy: "steps"
# eval_steps: 500
# save_strategy: "no"
# # save_steps: 500
# logging_strategy: "steps"
# logging_steps: 100
# use_liger_kernel: true
# remove_unused_columns: false
# needed to avoid OOM by compute the metrics batch by batch
# w/o this the trainer stores logits of all sample in memory...
# batch_eval_metrics: true
per_device_train_batch_size: 8
per_device_eval_batch_size: 8
max_val_samples_per_ds: 1000
# optim: schedule_free_adamw
learning_rate: 0.00004
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
warmup_steps: 100
dataloader_prefetch_factor: 16
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- self_gen/google/gemma-2-2b-it_temp_1.0/fw_qa_v2_2k_len_level_3_tiny
- self_gen/google/gemma-2-2b-it_temp_1.0/squad_compact
- self_gen/google/gemma-2-2b-it_temp_1.0/pwc_compact
- self_gen/google/gemma-2-2b-it_temp_1.0/ropes_compact
- self_gen/google/gemma-2-2b-it_temp_1.0/drop_compact
val_ds_names:
- squad
- pwc
- drop
- ropes

View file

@ -3,6 +3,8 @@ import gc
from datasets import Dataset, load_dataset
from tqdm import tqdm
from ctx_to_lora.data.processing import closed_qa_prompting
if __name__ == "__main__":
ds_name = "ucinlp/drop"
@ -14,7 +16,7 @@ if __name__ == "__main__":
ctx = sample["passage"]
if ctx not in ctx_qa_dict:
ctx_qa_dict[ctx] = {"prompts": [], "responses": []}
question = sample["question"]
question = closed_qa_prompting(sample["question"])
answer = ", ".join(set(sample["answers_spans"]["spans"]))
ctx_qa_dict[ctx]["prompts"].append(question)
ctx_qa_dict[ctx]["responses"].append(answer)

View file

@ -3,6 +3,8 @@ import gc
from datasets import Dataset, load_dataset
from tqdm import tqdm
from ctx_to_lora.data.processing import closed_qa_prompting
if __name__ == "__main__":
ds_name = "sggetao/PwC"
@ -14,7 +16,7 @@ if __name__ == "__main__":
ctx = sample["input"]
if ctx not in ctx_qa_dict:
ctx_qa_dict[ctx] = {"prompts": [], "responses": []}
question = sample["prompt"]
question = closed_qa_prompting(sample["prompt"])
answer = sample["answer"]
ctx_qa_dict[ctx]["prompts"].append(question)
ctx_qa_dict[ctx]["responses"].append(answer)

View file

@ -3,6 +3,8 @@ import gc
from datasets import Dataset, load_dataset
from tqdm import tqdm
from ctx_to_lora.data.processing import closed_qa_prompting
if __name__ == "__main__":
ds_name = "allenai/ropes"
@ -16,7 +18,7 @@ if __name__ == "__main__":
bg_txt = sample["background"]
situation_txt = sample["situation"]
ctx = ctx_template.format(background=bg_txt, situation=situation_txt)
q = sample["question"]
q = closed_qa_prompting(sample["question"])
if ctx not in ctx_qa_dict:
ctx_qa_dict[ctx] = {"prompts": [], "responses": []}
ctx_qa_dict[ctx]["prompts"].append(q)

View file

@ -3,6 +3,8 @@ import gc
from datasets import Dataset, load_dataset
from tqdm import tqdm
from ctx_to_lora.data.processing import closed_qa_prompting
if __name__ == "__main__":
ds_name = "rajpurkar/squad"
@ -14,7 +16,7 @@ if __name__ == "__main__":
ctx = sample["context"]
if ctx not in ctx_qa_dict:
ctx_qa_dict[ctx] = {"prompts": [], "responses": []}
question = sample["question"]
question = closed_qa_prompting(sample["question"])
answer = sample["answers"]["text"][0]
ctx_qa_dict[ctx]["prompts"].append(question)
ctx_qa_dict[ctx]["responses"].append(answer)

View file

@ -7,6 +7,10 @@ import pandas as pd
from datasets import Dataset, load_dataset
from vllm import LLM, SamplingParams
STOP_STRINGS = {
"google/gemma-2-2b-it": ["<eos>", "<end_of_turn>"],
}
SYSTEM_TEMPLATE = (
"You are a creative and helpful assistant.\n"
"You are tasked with generating questions for reading comprehension tests.\n"
@ -57,6 +61,14 @@ def get_prompt(context, n_qa_pairs):
return prompt
def check_should_skip(txt: str, vllm_model: str) -> bool:
"""Check if the response should be skipped based on stop strings."""
for stop in STOP_STRINGS[vllm_model]:
if stop in txt[-len(stop) :]:
return (txt.split(stop)[0], False) # Found a valid stop string
return (txt, True) # No valid stop string found, skip this response
def postprocess_qa_pairs(res_txt: str):
"""
Postprocesses the QA pairs from the response text.
@ -83,10 +95,21 @@ def postprocess_qa_pairs(res_txt: str):
out_q = []
out_a = []
n_skips = 0
if (len(questions) > 0) and (len(answers) > 0):
for i in range(min(len(questions), len(answers))):
n_gen_pairs = min(len(questions), len(answers))
has_left_over = n_gen_pairs < len(questions) or n_gen_pairs < len(answers)
for i in range(n_gen_pairs):
response = answers[i].strip()
if (not has_left_over) and (i == n_gen_pairs - 1):
response, skip = check_should_skip(response, vllm_model)
if skip:
print(f"Skipping due to missing stop string")
n_skips += 1
continue
out_q.append(questions[i].strip())
out_a.append(answers[i].strip())
print(f"Skipped {n_skips} responses due to missing stop strings")
return out_q, out_a
@ -193,7 +216,10 @@ if __name__ == "__main__":
messages,
sampling_params=SamplingParams(
max_tokens=2048,
temperature=0.7,
temperature=0.0,
# needed for checking if stop tokens are present
skip_special_tokens=False,
include_stop_str_in_output=True,
),
)
samples = []

View file

@ -7,6 +7,10 @@ from glob import glob
from datasets import load_dataset
from vllm import LLM, SamplingParams
STOP_STRINGS = {
"google/gemma-2-2b-it": ["<eos>", "<end_of_turn>"],
}
SYSTEM_TEMPLATE = (
"You are a creative and helpful assistant.\n"
"You are tasked with generating questions for reading comprehension tests.\n"
@ -64,6 +68,14 @@ def get_prompt(context, example_qa_pairs, n_qa_pairs):
return prompt
def check_should_skip(txt: str, vllm_model: str) -> bool:
"""Check if the response should be skipped based on stop strings."""
for stop in STOP_STRINGS[vllm_model]:
if stop in txt[-len(stop) :]:
return (txt.split(stop)[0], False) # Found a valid stop string
return (txt, True) # No valid stop string found, skip this response
def postprocess_qa_pairs(res_txt: str):
"""
Postprocesses the QA pairs from the response text.
@ -90,10 +102,21 @@ def postprocess_qa_pairs(res_txt: str):
out_q = []
out_a = []
n_skips = 0
if (len(questions) > 0) and (len(answers) > 0):
for i in range(min(len(questions), len(answers))):
n_gen_pairs = min(len(questions), len(answers))
has_left_over = n_gen_pairs < len(questions) or n_gen_pairs < len(answers)
for i in range(n_gen_pairs):
response = answers[i].strip()
if (not has_left_over) and (i == n_gen_pairs - 1):
response, skip = check_should_skip(response, vllm_model)
if skip:
print(f"Skipping due to missing stop string")
n_skips += 1
continue
out_q.append(questions[i].strip())
out_a.append(answers[i].strip())
print(f"Skipped {n_skips} responses due to missing stop strings")
return out_q, out_a
@ -217,8 +240,10 @@ if __name__ == "__main__":
completions = llm.chat(
messages,
sampling_params=SamplingParams(
max_tokens=2048,
temperature=0.7,
temperature=0.0,
# needed for checking if stop tokens are present
skip_special_tokens=False,
include_stop_str_in_output=True,
),
)
samples = []

View file

@ -1,5 +1,6 @@
import argparse
import os
import random
from glob import glob
import pandas as pd
@ -15,6 +16,10 @@ from ctx_to_lora.data.processing import (
)
from ctx_to_lora.utils import clear_gpu
STOP_STRINGS = {
"google/gemma-2-2b-it": ["<eos>", "<end_of_turn>"],
}
SYSTEM_TEMPLATE = (
"### SYSTEM INSTRUCTIONS ###\n"
"You are a creative and helpful assistant.\n"
@ -42,6 +47,14 @@ def load_config(config_path: str) -> dict:
return config
def check_should_skip(txt: str, vllm_model: str) -> bool:
"""Check if the response should be skipped based on stop strings."""
for stop in STOP_STRINGS[vllm_model]:
if stop in txt[-len(stop) :]:
return (txt.split(stop)[0], False) # Found a valid stop string
return (txt, True) # No valid stop string found, skip this response
def get_dataset_configs(
ds_names: list[str] | None,
config: dict | None,
@ -147,6 +160,10 @@ def self_generate(
print(f"Loading dataset: {ds_name} with split: {split}")
ds = load_and_process_dataset(**kwargs, streaming=False, num_proc=8)
temp = 0.0
if "_temp_" in ds_name:
temp = float(ds_name.split("_temp_")[-1].split("/")[0])
if args.debug:
ds = ds.take(10)
@ -163,19 +180,38 @@ def self_generate(
completions = llm.chat(
messages,
sampling_params=SamplingParams(
max_tokens=2048,
temperature=1.0, # TODO: lower the temp
max_tokens=1024,
temperature=temp,
# needed for checking if stop tokens are present
skip_special_tokens=False,
include_stop_str_in_output=True,
),
)
self_gen_data = {ctx: {"prompts": [], "responses": []} for ctx in ctxs}
c = 0
n_skips = 0
for ctx, q_list in zip(ctxs, questions):
for i, q in enumerate(q_list):
response = completions[c + i].outputs[0].text
response, skip = check_should_skip(response, args.vllm_model)
# skip = True
# for stop in STOP_STRINGS[args.vllm_model]:
# if stop == response[-len(stop) :]:
# # Check if response ends with stop string
# response = response.split(stop)[0]
# skip = False
# break
if skip:
print(f"Skipping due to missing stop string")
n_skips += 1
continue
self_gen_data[ctx]["prompts"].append(q)
self_gen_data[ctx]["responses"].append(completions[c + i].outputs[0].text)
self_gen_data[ctx]["responses"].append(response)
c += i + 1
print(f"Skipped {n_skips} responses due to missing stop strings")
samples = [
{
"context": ctx,
@ -193,20 +229,21 @@ def self_generate(
print("=" * 80)
print(f"Generated {len(samples)} samples")
# random.shuffle(samples)
random.shuffle(samples)
# Save results
df = pd.DataFrame(samples)
ds_out = Dataset.from_pandas(df)
fpath = f"{SELF_GEN_DATA_DIR}/{args.vllm_model}/{ds_name}/{split}/ds{shard_name}"
os.makedirs(os.path.dirname(fpath), exist_ok=True)
fpath = f"{fpath}.parquet"
ds_out.to_parquet(fpath)
print(f"Saved to {fpath}")
if not args.debug:
df = pd.DataFrame(samples)
ds_out = Dataset.from_pandas(df)
fpath = f"{SELF_GEN_DATA_DIR}/{args.vllm_model}_temp_{temp}/{ds_name}/{split}/ds{shard_name}"
os.makedirs(os.path.dirname(fpath), exist_ok=True)
fpath = f"{fpath}.parquet"
ds_out.to_parquet(fpath)
print(f"Saved to {fpath}")
# Cleanup
del samples, df, ds_out, completions, messages, ctxs, questions
clear_gpu()
# Cleanup
del samples, df, ds_out, completions, messages, ctxs, questions
clear_gpu()
def parse_args() -> argparse.Namespace:

View file

@ -149,8 +149,10 @@ DS_KWARGS = {
"fw_qa_v2_2k_len_level_0_tiny": dict(
train=dict(
path="parquet",
data_files="data/raw_datasets/fw_qa_v2/min_0_to_2000/000_00000_level_0.parquet",
split="train",
data_files=glob(
"data/raw_datasets/fw_qa_v2/min_0_to_2000/013*level_0.parquet"
),
split="train[:400000]",
),
),
"fw_qa_v2_2k_len_level_0": dict(
@ -181,6 +183,15 @@ DS_KWARGS = {
# split="train",
# ),
),
"fw_qa_v2_2k_len_level_1_tiny": dict(
train=dict(
path="parquet",
data_files=glob(
"data/raw_datasets/fw_qa_v2/min_0_to_2000/013*level_1.parquet"
),
split="train[:200000]",
),
),
"fw_qa_v2_2k_len_level_2": dict(
train=dict(
path="parquet",
@ -195,11 +206,13 @@ DS_KWARGS = {
# split="train",
# ),
),
"fw_qa_v2_2k_len_level_3_tiny": dict(
"fw_qa_v2_2k_len_level_2_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",
data_files=glob(
"data/raw_datasets/fw_qa_v2/min_0_to_2000/013*level_2.parquet"
),
split="train[:100000]",
),
),
"fw_qa_v2_2k_len_level_3": dict(
@ -216,13 +229,15 @@ DS_KWARGS = {
# split="train",
# ),
),
# "fw_qa_3_mini": dict(
# train=dict(
# path="parquet",
# data_files=glob("data/raw_datasets/fw_qa_3/000_00001.parquet"),
# split="train[:100000]",
# ),
# ),
"fw_qa_v2_2k_len_level_3_tiny": dict(
train=dict(
path="parquet",
data_files=glob(
"data/raw_datasets/fw_qa_v2/min_0_to_2000/013*level_3.parquet"
),
split="train[:50000]",
),
),
# "fw_qa_3_small": dict(
# train=dict(
# path="parquet",
@ -491,6 +506,8 @@ CLOSED_QA_DATASETS = {
"squad",
"triviaqa_retrieved",
"negative_nq",
"ropes",
"drop",
}
MULTI_ANSWER_DATASETS = {
@ -543,6 +560,8 @@ CLOSED_QA_INTX_TEMPLATES = [
EVAL_INTX_TEMPLATES = {
"ropes": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}",
"drop": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}",
"negative_nq": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}",
"triviaqa_retrieved": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}",
"hotpot_qa": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}",

View file

@ -185,6 +185,7 @@ def get_preprocessing_fn(
situation_txt = sample["situation"]
ctx = ctx_template.format(background=bg_txt, situation=situation_txt)
q = sample["question"]
q = closed_qa_prompting(q) if not is_eval else q
return {"context": ctx, "prompt": q, "response": response}
elif ds_name in ["narrativeqa", "quoref"]: # , "ropes"]: