diff --git a/configs/pretrain_all_xl_3_medium.yaml b/configs/pretrain_all_xl_3_medium.yaml new file mode 100644 index 0000000..d8b3c5d --- /dev/null +++ b/configs/pretrain_all_xl_3_medium.yaml @@ -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 diff --git a/configs/pretrain_all_xl_3_mini.yaml b/configs/pretrain_all_xl_3_mini.yaml index 13268e8..4ede74e 100644 --- a/configs/pretrain_all_xl_3_mini.yaml +++ b/configs/pretrain_all_xl_3_mini.yaml @@ -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 diff --git a/intx_sft.py b/intx_sft.py index ede5cca..8781977 100755 --- a/intx_sft.py +++ b/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() diff --git a/scripts/fw_qa_3/medium_gemma_per_rank_fac2_per_layer.sh b/scripts/fw_qa_3/medium_gemma_per_rank_fac2_per_layer.sh new file mode 100644 index 0000000..08b2e04 --- /dev/null +++ b/scripts/fw_qa_3/medium_gemma_per_rank_fac2_per_layer.sh @@ -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 diff --git a/src/ctx_to_lora/data_utils.py b/src/ctx_to_lora/data_utils.py index 50f008a..ffee7b1 100644 --- a/src/ctx_to_lora/data_utils.py +++ b/src/ctx_to_lora/data_utils.py @@ -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(