mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
medium fw qa 3 + fallback to manual cache dataset
This commit is contained in:
parent
47edd5df1e
commit
0ff762cdff
5 changed files with 227 additions and 31 deletions
60
configs/pretrain_all_xl_3_medium.yaml
Normal file
60
configs/pretrain_all_xl_3_medium.yaml
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
15
intx_sft.py
15
intx_sft.py
|
|
@ -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()
|
||||
|
|
|
|||
33
scripts/fw_qa_3/medium_gemma_per_rank_fac2_per_layer.sh
Normal file
33
scripts/fw_qa_3/medium_gemma_per_rank_fac2_per_layer.sh
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue