diff --git a/README.md b/README.md index 77cba6c..78d9d4b 100644 --- a/README.md +++ b/README.md @@ -22,8 +22,8 @@ WANDB_MODE=disabled run python hyperlora/intx_sft.py configs/context_numbers_128 # python process_fineweb.py # python generate_fw_qa.py # python post_process_fw_qa.py - -vllm_model=mistralai/Mistral-Small-24B-Instruct-2501 run python generate_fw_qa_vllm.py 00_* 2 +conda activate vllm; +vllm_model=mistralai/Mistral-Small-3.1-24B-Instruct-2503 run python generate_fw_edu_qa_vllm.py "0*_*" 3 ``` diff --git a/configs/pretrain_all_xl_3_mini.yaml b/configs/pretrain_all_xl_3_mini.yaml new file mode 100644 index 0000000..13268e8 --- /dev/null +++ b/configs/pretrain_all_xl_3_mini.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_mini # ~ 267M + - 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/pyproject.toml b/pyproject.toml index f410132..c12dae7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,11 +6,11 @@ description = "" readme = "README.md" requires-python = ">= 3.10" dependencies = [ - "transformers==4.48.3", - "deepspeed==0.15.4", - "accelerate==1.2.1", + "transformers==4.51.3", + "deepspeed==0.16.7", + "accelerate==1.6.0", + "datasets==3.6.0", "peft", - "datasets", "bitsandbytes", "jupyter", "matplotlib", @@ -43,5 +43,13 @@ profile = "black" line-length = 89 [tool.pyright] -exclude = ["train_outputs", "data", "generated_tasks", "outputs", "plots", "tmp", "wandb"] +exclude = [ + "train_outputs", + "data", + "generated_tasks", + "outputs", + "plots", + "tmp", + "wandb", +] typeCheckingMode = "off" diff --git a/scripts/fw_qa_3/gemma_per_rank_fac2_per_layer.sh b/scripts/fw_qa_3/gemma_per_rank_fac2_per_layer.sh index 123bdf8..a824d2a 100644 --- a/scripts/fw_qa_3/gemma_per_rank_fac2_per_layer.sh +++ b/scripts/fw_qa_3/gemma_per_rank_fac2_per_layer.sh @@ -21,7 +21,7 @@ conda activate /home/rujikorn_sakana_ai/.conda/envs/ctx-to-lora accelerate launch --num_processes=4 --gradient_accumulation_steps=8 --gradient_clipping=1.0 \ --gpu_ids all --main_process_port 29562 intx_sft.py configs/pretrain_all_xl_3.yaml \ ---model_name_or_path=google/gemma-2-2b-it --num_train_epochs=5.1 --per_device_train_batch_size=32 \ +--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 --per_device_train_batch_size=32 \ --gradient_accumulation_steps=8 --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 \ diff --git a/src/ctx_to_lora/data_utils.py b/src/ctx_to_lora/data_utils.py index 5157f83..50f008a 100644 --- a/src/ctx_to_lora/data_utils.py +++ b/src/ctx_to_lora/data_utils.py @@ -410,11 +410,22 @@ def get_tokenized_dataset( ds_kwargs = get_ds_kwargs(ds_name, split) skip = ds_kwargs.pop("skip", None) take = ds_kwargs.pop("take", None) + # only load as stream if the dataset is local (parquet) + load_as_stream = ( + ("fw_qa" in ds_name) + and ("tiny" not in ds_name) + and ("mini" not in ds_name) + and streaming + and ds_kwargs["path"] == "parquet" + ) ds = load_dataset( **get_ds_kwargs(ds_name, split), trust_remote_code=True, - streaming=streaming, + streaming=load_as_stream, ) + if split == "train" and streaming and (not load_as_stream): + # if the dataset is hosted on HF, we load it first then convert to iterable + ds = ds.to_iterable_dataset(num_shards=128) if skip is not None: ds = ds.skip(skip) if take is not None: