medium fw qa 3 + fallback to manual cache dataset

This commit is contained in:
51616 2025-05-10 15:31:08 +00:00
parent 47edd5df1e
commit 0ff762cdff
5 changed files with 227 additions and 31 deletions

View file

@ -0,0 +1,60 @@
output_dir: "" # just a placeholder
bf16: true
model_name_or_path: meta-llama/Llama-3.2-1B-Instruct
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: 5
weight_decay: 0.01
#
warmup_steps: 100
dataloader_prefetch_factor: 8
dataloader_num_workers: 8
# LoRA
lora_r: 8
lora_dropout: 0.0
target_modules:
- down_proj
# data
train_ds_names:
- fw_qa_3_medium # ~ 130M?
- ctx_qa # 300k
- pwc # 240k
- hotpot_qa # 90k
- squad # 90k
- drop # 77k
- narrativeqa # 40k
- quoref # 11k
- ropes # 11k
- synthetic_convqa # 40k
val_ds_names:
- fw_qa_3
- fw_qa_xl
- ctx_qa
- pwc
- hotpot_qa
- squad
load_best_model_at_end: true
metric_for_best_model: eval_pwc_loss

View file

@ -37,7 +37,7 @@ target_modules:
- down_proj
# data
train_ds_names:
- fw_qa_3_mini # ~ 267M
- fw_qa_3_mini # 100k
- ctx_qa # 300k
- pwc # 240k
- hotpot_qa # 90k

View file

@ -3,6 +3,7 @@ import os
import random
import string
import time
from multiprocess import set_start_method
from math import ceil
from collections import defaultdict
from copy import copy, deepcopy
@ -284,7 +285,7 @@ def main():
training_args.lr_scheduler_type == "cosine_with_min_lr"
and training_args.lr_scheduler_kwargs is None
):
training_args.lr_scheduler_kwargs = {"min_lr": 1e-8}
training_args.lr_scheduler_kwargs = {"min_lr": 1e-7}
args = {
**vars(deepcopy(data_args)),
**vars(deepcopy(ctx_args)),
@ -316,7 +317,7 @@ def main():
)
if "Llama" in ctx_name and "Vision" in ctx_name:
ctx_encoder_model_config = ctx_encoder_model_config.text_config
ctx_tokenizer = get_tokenizer(ctx_name)
ctx_tokenizer = get_tokenizer(ctx_name, train=True)
else:
ctx_name = model.base_model.config.name_or_path
ctx_encoder_model_config = model.config
@ -353,11 +354,11 @@ def main():
train=True,
use_flash_attn=model_args.use_flash_attn,
)
tokenizer = get_tokenizer(model.base_model.config.name_or_path)
tokenizer = get_tokenizer(model.base_model.config.name_or_path, train=True)
ctx_name = model.ctx_encoder_args.ctx_encoder_model_name_or_path
if ctx_name is None:
ctx_name = model.base_model.config.name_or_path
ctx_tokenizer = get_tokenizer(ctx_name)
ctx_tokenizer = get_tokenizer(ctx_name, train=True)
if len([p for p in model.ctx_encoder.parameters() if p.requires_grad]):
raise ValueError("ctx_encoder contains trainable parameters")
if len([p for p in model.base_model.parameters() if p.requires_grad]):
@ -573,10 +574,12 @@ if __name__ == "__main__":
os.environ["WANDB_WATCH"] = "" # "all"
os.environ["WANDB_CONSOLE"] = "off"
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
os.environ["KMP_AFFINITY"] = "disabled" # fixing iterable dataset stuck
os.environ["OMP_NUM_THREADS"] = "16"
# os.environ["KMP_AFFINITY"] = "disabled" # fixing iterable dataset stuck
# os.environ["OMP_NUM_THREADS"] = "16"
# os.environ["HF_DATASETS_IN_MEMORY_MAX_SIZE"] = "137438953472" # 128 TB
if os.getenv("DEBUG", False):
disable_caching()
# randomly sleep to avoid run_name collision
# time.sleep(random.random() * 13)
# set_start_method("spawn", force=True)
main()

View file

@ -0,0 +1,33 @@
#!/bin/bash
#SBATCH --job-name=ctxlora_medium
#SBATCH --partition=a3
#SBATCH --nodes=1
#SBATCH --gpus=8
#SBATCH --output=outputs/%x-%j.out
#SBATCH --error=outputs/%x-%j.out
# module load
# module load cuda/12.1
# module load cudnn/8.9.7
# module load nccl/cuda-12.1/2.18.3
# module load hpcx/2.20
# export OMP_NUM_THREADS=24
# export TRITON_CACHE_DIR=/tmp/.triton/
. ~/miniconda3/etc/profile.d/conda.sh
conda activate /home/rujikorn_sakana_ai/.conda/envs/ctx-to-lora
# eval "$@"
accelerate launch --num_processes=8 --gradient_accumulation_steps=4 --gradient_clipping=1.0 \
--gpu_ids all --main_process_port 29568 intx_sft.py configs/pretrain_all_xl_3_medium.yaml \
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 --per_device_train_batch_size=32 \
--gradient_accumulation_steps=4 --per_device_eval_batch_size=32 --exp_setup=hyper_lora --aggregator_type=perceiver \
--target_modules=down_proj \
--num_self_attends_per_block=8 --num_latent_factor=2 \
--lora_r=8 \
--eval_steps=5000 --save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
--neftune_noise_alpha=5 --use_light_weight_lora=False \
--load_best_model_at_end=True --metric_for_best_model=pwc_loss --add_negative_prompt=False \
--add_repeat_prompt=False \
--use_sequence_packing=True --per_rank_gen=True \
--per_layer_processing=True

View file

@ -1,9 +1,13 @@
import logging
from os import path
import numpy as np
import hashlib
import json
from os import path
from glob import glob
from typing import Any, Callable, Iterator, Optional
import datasets
from datasets import load_dataset, IterableDataset
from transformers import PreTrainedTokenizerBase
@ -16,10 +20,13 @@ FW_QA_PATHS = [
f"data/raw_datasets/fw_qa/{i:05d}.parquet" for i in [0, 1, 6, 7, 8, 10, 22, 30, 35]
]
TRANSFORMED_DATA_DIR = "data/processed_datasets"
# approximate length of the datasets
# needed for streaming datasets
DS_LEN = {
"fw_qa_3_mini": 100_000,
"fw_qa_3_medium": 121_000_000,
"fw_qa_3": 270_000_000,
"fw_qa_xl": 27_000_000,
"ctx_qa": 300_000,
@ -124,6 +131,13 @@ DS_KWARGS = {
split="train[:100000]",
),
),
"fw_qa_3_medium": dict(
train=dict(
path="parquet",
data_files=glob("data/raw_datasets/fw_qa_3/00[0-5]*[!val].parquet"),
split="train",
),
),
"fw_qa_3": dict(
train=dict(
path="parquet",
@ -389,23 +403,14 @@ def filter_none(samples):
return out
def get_tokenized_dataset(
def _load_and_process_dataset(
ds_name: str,
split: str,
tokenizer: PreTrainedTokenizerBase,
tokenizer_kwargs: dict[str, Any],
ctx_tokenizer: PreTrainedTokenizerBase,
ctx_tokenizer_kwargs: dict[str, Any],
add_ctx_to_chat: bool,
add_repeat_prompt: bool,
add_negative_prompt: bool,
use_kl_loss: bool,
set_format: Optional[str] = None,
streaming: bool = False,
) -> dict[str, Any]:
logger.debug(f"Loading dataset {ds_name} with split {split}...")
need_ctx_ids = not add_ctx_to_chat
add_repeat_prompt: bool,
streaming: bool,
ds_path: str,
):
try:
ds_kwargs = get_ds_kwargs(ds_name, split)
skip = ds_kwargs.pop("skip", None)
@ -438,15 +443,74 @@ def get_tokenized_dataset(
cols_to_remove = [
col for col in ds.column_names if col not in ["context", "prompt", "response"]
]
ds = ds.map(get_preprocessing_fn(ds_name))
ds = ds.remove_columns(cols_to_remove)
ds = ds.filter(filter_none, batched=True)
ds = ds.filter(filter_long_samples, batched=True)
ds = ds.map(
get_preprocessing_fn(ds_name),
remove_columns=cols_to_remove,
num_proc=16,
)
# ds = ds.remove_columns(cols_to_remove)
ds = ds.filter(filter_none, batched=True, num_proc=16)
ds = ds.filter(filter_long_samples, batched=True, num_proc=16)
if split == "train":
if add_negative_prompt:
ds = ds.map(add_negative_prompt_fn, batched=True)
ds = ds.map(
add_negative_prompt_fn,
batched=True,
batch_size=100_000,
num_proc=16,
)
if add_repeat_prompt and "context_numbers" not in ds_name:
ds = ds.map(add_repeat_prompt_fn, batched=True)
ds = ds.map(
add_repeat_prompt_fn,
batched=True,
batch_size=100_000,
num_proc=16,
)
ds.save_to_disk(ds_path, num_proc=16)
return ds
def get_tokenized_dataset(
ds_name: str,
split: str,
tokenizer: PreTrainedTokenizerBase,
tokenizer_kwargs: dict[str, Any],
ctx_tokenizer: PreTrainedTokenizerBase,
ctx_tokenizer_kwargs: dict[str, Any],
add_ctx_to_chat: bool,
add_repeat_prompt: bool,
add_negative_prompt: bool,
use_kl_loss: bool,
set_format: Optional[str] = None,
streaming: bool = False,
) -> dict[str, Any]:
assert not use_kl_loss, "KL loss is deprecated"
logger.debug(f"Loading dataset {ds_name} with split {split}...")
need_ctx_ids = not add_ctx_to_chat
load_and_process_kwargs = dict(
ds_name=ds_name,
split=split,
add_negative_prompt=add_negative_prompt,
add_repeat_prompt=add_repeat_prompt,
streaming=streaming,
)
ds_hash = hashlib.sha256(json.dumps(load_and_process_kwargs).encode()).hexdigest()
ds_path = f"{TRANSFORMED_DATA_DIR}/{ds_hash}"
if path.exists(ds_path):
# load the cached ds
logger.info(f"Loaded processed dataset from {ds_path}")
else:
logger.info(f"Loading dataset {ds_name} with split {split}...")
_load_and_process_dataset(
**load_and_process_kwargs,
ds_path=ds_path,
)
ds = datasets.load_from_disk(ds_path)
tokenized_ds = construct_and_tokenize_ctx_qa(
tokenizer,
tokenizer_kwargs,
@ -472,16 +536,47 @@ def construct_and_tokenize_ctx_qa(
ds,
set_format=None,
):
kwargs = dict(
tokenizer=repr(tokenizer),
tokenizer_kwargs=json.dumps(tokenizer_kwargs),
ctx_tokenizer=repr(ctx_tokenizer),
ctx_tokenizer_kwargs=json.dumps(ctx_tokenizer_kwargs),
add_ctx_to_chat=add_ctx_to_chat,
use_kl_loss=use_kl_loss,
need_ctx_ids=need_ctx_ids,
ds=ds._fingerprint,
set_format=set_format,
)
kwargs_str = json.dumps(kwargs)
logger.debug(f"Tokenizing dataset with kwargs: {kwargs_str}")
ds_hash = hashlib.sha256(kwargs_str.encode()).hexdigest()
ds_path = f"{TRANSFORMED_DATA_DIR}/{ds_hash}"
if path.exists(ds_path):
# load the cached ds
logger.info(f"Loaded tokenized dataset from {ds_path}")
ds = datasets.load_from_disk(ds_path)
return ds
# for sft + chat_model, we need to convert the dataset to chat format
# add "messages" field
ds = ds.map(
convert_ctx_prompt_response_to_messages,
fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat},
num_proc=16,
)
# add "chat" field
ds = ds.map(get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer))
ds = ds.map(
get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer),
batched=True,
batch_size=100_000,
num_proc=16,
)
# tokenize the chat + mask the assistant inputs
ds = ds.filter(filter_long_chat, batched=True)
ds = ds.filter(
filter_long_chat,
batched=True,
batch_size=100_000,
num_proc=16,
)
# add "input_ids", "attention_mask", "labels"
tokenized_ds = ds.map(
@ -491,6 +586,7 @@ def construct_and_tokenize_ctx_qa(
"mask_assistant_inputs": True,
"tokenizer_kwargs": tokenizer_kwargs,
},
num_proc=16,
)
# for use_kl_loss, we need "chat_ids" and "chat_attn_mask"
@ -521,6 +617,8 @@ def construct_and_tokenize_ctx_qa(
tokenize_ctx_text,
fn_kwargs={"tokenizer": ctx_tokenizer},
batched=True,
batch_size=100_000,
num_proc=16,
)
tokenized_ds = tokenized_ds.remove_columns(
@ -532,7 +630,9 @@ def construct_and_tokenize_ctx_qa(
# # the columns are unknown when using streaming dataset
# tokenized_ds = tokenized_ds._resolve_features()
# validate_columns(tokenized_ds)
return tokenized_ds
tokenized_ds.save_to_disk(ds_path, num_proc=16)
del tokenized_ds
return datasets.load_from_disk(ds_path)
def get_sft_prompt_formatting_fn(