diff --git a/.gitignore b/.gitignore index 737d8ef..5baed56 100644 --- a/.gitignore +++ b/.gitignore @@ -19,6 +19,7 @@ wandb/ *outputs/ plots/ *.out +*.err *.pt *.pth *.bin diff --git a/configs/gemma-3-1b-it/toy_exp/context_numbers_10.yaml b/configs/gemma-3-1b-it/toy_exp/context_numbers_10.yaml deleted file mode 100644 index 4d825f8..0000000 --- a/configs/gemma-3-1b-it/toy_exp/context_numbers_10.yaml +++ /dev/null @@ -1,57 +0,0 @@ -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: 64 -per_device_eval_batch_size: 128 -max_new_tokens: 64 -gen_per_device_eval_batch_size: 128 -max_val_samples_per_ds: 500 -# optim: schedule_free_adamw -learning_rate: 0.0001 -# lr_scheduler_type: "constant_with_warmup" -neftune_noise_alpha: 1 -weight_decay: 0.01 - -# LoRA -lora_r: 16 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- data/raw_datasets/context_numbers_2 -- data/raw_datasets/context_numbers_3 -- data/raw_datasets/context_numbers_4 -- data/raw_datasets/context_numbers_5 -- data/raw_datasets/context_numbers_6 -- data/raw_datasets/context_numbers_7 -- data/raw_datasets/context_numbers_8 -- data/raw_datasets/context_numbers_9 -- data/raw_datasets/context_numbers_10 - -val_ds_names: -- data/raw_datasets/context_numbers_2 -- data/raw_datasets/context_numbers_3 -- data/raw_datasets/context_numbers_4 -- data/raw_datasets/context_numbers_5 -- data/raw_datasets/context_numbers_6 -- data/raw_datasets/context_numbers_7 -- data/raw_datasets/context_numbers_8 -- data/raw_datasets/context_numbers_9 -- data/raw_datasets/context_numbers_10 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_128.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_128.yaml deleted file mode 100644 index 02f024c..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_128.yaml +++ /dev/null @@ -1,12 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- ctx_kv_64_128 - -val_ds_names: -- ctx_kv_64_128 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_128_self_gen.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_128_self_gen.yaml deleted file mode 100644 index bfe2203..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_128_self_gen.yaml +++ /dev/null @@ -1,12 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- self_gen/google/gemma-3-1b-it_temp_0.0_closed_qa_prob_0.0/ctx_kv_64_128 - -val_ds_names: -- ctx_kv_64_128 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_256.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_256.yaml deleted file mode 100644 index 4f281e0..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_kv_64_256.yaml +++ /dev/null @@ -1,14 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- ctx_kv_64_128 -- ctx_kv_128_256 - -val_ds_names: -- ctx_kv_64_128 -- ctx_kv_128_256 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_1024.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_1024.yaml deleted file mode 100644 index 4eae29c..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_1024.yaml +++ /dev/null @@ -1,21 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- ctx_numbers_64_128 -- ctx_numbers_128_256 -- ctx_numbers_256_512 -- ctx_numbers_512_768 -- ctx_numbers_768_1024 - - -val_ds_names: -- ctx_numbers_64_128 -- ctx_numbers_128_256 -- ctx_numbers_256_512 -- ctx_numbers_512_768 -- ctx_numbers_768_1024 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_128.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_128.yaml deleted file mode 100644 index 8dbc6b9..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_128.yaml +++ /dev/null @@ -1,12 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- ctx_numbers_64_128 - -val_ds_names: -- ctx_numbers_64_128 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_128_self_gen.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_128_self_gen.yaml deleted file mode 100644 index 05b8530..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_128_self_gen.yaml +++ /dev/null @@ -1,12 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- self_gen/google/gemma-3-1b-it_temp_0.0_closed_qa_prob_0.0/ctx_numbers_64_128 - -val_ds_names: -- ctx_numbers_64_128 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_2048.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_2048.yaml deleted file mode 100644 index c9a1108..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_2048.yaml +++ /dev/null @@ -1,29 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- ctx_numbers_64_128 -- ctx_numbers_128_256 -- ctx_numbers_256_512 -- ctx_numbers_512_768 -- ctx_numbers_768_1024 -- ctx_numbers_1024_1280 -- ctx_numbers_1280_1536 -- ctx_numbers_1536_1792 -- ctx_numbers_1792_2048 - - -val_ds_names: -- ctx_numbers_64_128 -- ctx_numbers_128_256 -- ctx_numbers_256_512 -- ctx_numbers_512_768 -- ctx_numbers_768_1024 -- ctx_numbers_1024_1280 -- ctx_numbers_1280_1536 -- ctx_numbers_1536_1792 -- ctx_numbers_1792_2048 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_256.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_256.yaml deleted file mode 100644 index d002fec..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_256.yaml +++ /dev/null @@ -1,14 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- ctx_numbers_64_128 -- ctx_numbers_128_256 - -val_ds_names: -- ctx_numbers_64_128 -- ctx_numbers_128_256 diff --git a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_512.yaml b/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_512.yaml deleted file mode 100644 index 6c5b5e6..0000000 --- a/configs/gemma-3-1b-it/toy_exp/ctx_numbers_64_512.yaml +++ /dev/null @@ -1,17 +0,0 @@ -# LoRA -lora_r: 8 -lora_dropout: 0.0 -target_modules: - - down_proj - -# data -train_ds_names: -- ctx_numbers_64_128 -- ctx_numbers_128_256 -- ctx_numbers_256_512 - - -val_ds_names: -- ctx_numbers_64_128 -- ctx_numbers_128_256 -- ctx_numbers_256_512 diff --git a/data/generate_ctx_magic_number.py b/data/generate_ctx_magic_number.py index f2a871b..23b05d0 100644 --- a/data/generate_ctx_magic_number.py +++ b/data/generate_ctx_magic_number.py @@ -168,6 +168,7 @@ def main(): tok_bins = [(32, 128), (128, 256), (256, 512), (512, 1024), (32, 1024)] + [ (1024 * i, 1024 * (i + 1)) for i in range(1, 16) ] + tok_bins += [(2**14 + 2**12 * (i), 2**14 + 2**12 * (i + 1)) for i in range(4)] if args.only_first_n_bins is not None: tok_bins = tok_bins[: args.only_first_n_bins] diff --git a/scripts/toy_exp.sh b/scripts/toy_exp.sh new file mode 100644 index 0000000..2cea0e2 --- /dev/null +++ b/scripts/toy_exp.sh @@ -0,0 +1,11 @@ +# data gen +uv run data/generate_ctx_magic_num.py + +# train +WANDB_PROJECT=ctx-magic-num srun --partition=aiscilow --gpus=1 --unbuffered uv run train.py configs/toy_exp/ctx_magic_number_32_256.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 --per_device_train_batch_size=-1 --gradient_accumulation_steps=32 --per_device_eval_batch_size=16 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj --num_blocks=8 --num_self_attn_per_block=0 --num_pre_head_layers=1 --lora_r=8 --eval_steps=100 --save_steps=1000 --learning_rate=4e-5 --lora_dropout=0.0 --neftune_noise_alpha=0 --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 --use_sequence_packing=True --max_packed_inp_len=2048 --max_packed_ctx_len=2048 --dataloader_num_workers=0 --dataloader_prefetch_factor=None --eval_on_start=False --ctx_encoder_type=early_exit --n_latent_queries=208 --use_kl_loss=False --eval_on_start=True --lora_r=8 --max_ctx_chunk_len=-1 --max_val_samples_per_ds=100 --seed=1 +WANDB_PROJECT=ctx-magic-num srun --partition=aiscilow --gpus=1 --unbuffered uv run train.py configs/toy_exp/ctx_magic_number_32_256.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 --per_device_train_batch_size=-1 --gradient_accumulation_steps=32 --per_device_eval_batch_size=16 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj --num_blocks=8 --num_self_attn_per_block=0 --num_pre_head_layers=1 --lora_r=8 --eval_steps=100 --save_steps=1000 --learning_rate=4e-5 --lora_dropout=0.0 --neftune_noise_alpha=0 --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 --use_sequence_packing=True --max_packed_inp_len=2048 --max_packed_ctx_len=2048 --dataloader_num_workers=0 --dataloader_prefetch_factor=None --eval_on_start=False --ctx_encoder_type=early_exit --n_latent_queries=208 --use_kl_loss=False --eval_on_start=True --lora_r=8 --max_ctx_chunk_len=512 --min_ctx_chunk_len=25 --num_chunk_probs='{"1":"0.5", "2":"0.5"}' --max_val_samples_per_ds=100 --seed=1 +WANDB_PROJECT=ctx-magic-num srun --partition=aiscilow --gpus=1 --unbuffered uv run train.py configs/toy_exp/ctx_magic_number_32_256.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 --per_device_train_batch_size=-1 --gradient_accumulation_steps=32 --per_device_eval_batch_size=16 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj --num_blocks=8 --num_self_attn_per_block=0 --num_pre_head_layers=1 --lora_r=8 --eval_steps=100 --save_steps=1000 --learning_rate=4e-5 --lora_dropout=0.0 --neftune_noise_alpha=0 --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 --use_sequence_packing=True --max_packed_inp_len=2048 --max_packed_ctx_len=2048 --dataloader_num_workers=0 --dataloader_prefetch_factor=None --eval_on_start=False --ctx_encoder_type=early_exit --n_latent_queries=208 --use_kl_loss=False --eval_on_start=True --lora_r=8 --max_ctx_chunk_len=512 --min_ctx_chunk_len=25 --num_chunk_probs='{"1":"0.5", "2":"0.25", "3":"0.125", "4":"0.125"}' --max_val_samples_per_ds=100 --seed=1 +WANDB_PROJECT=ctx-magic-num run uv run train.py configs/toy_exp/ctx_magic_number_32_256.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 --per_device_train_batch_size=-1 --gradient_accumulation_steps=32 --per_device_eval_batch_size=16 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj --num_blocks=8 --num_self_attn_per_block=0 --num_pre_head_layers=1 --lora_r=8 --eval_steps=100 --save_steps=1000 --learning_rate=4e-5 --lora_dropout=0.0 --neftune_noise_alpha=0 --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 --use_sequence_packing=True --max_packed_inp_len=2048 --max_packed_ctx_len=2048 --dataloader_num_workers=0 --dataloader_prefetch_factor=None --eval_on_start=False --ctx_encoder_type=early_exit --n_latent_queries=208 --use_kl_loss=False --eval_on_start=True --lora_r=8 --max_ctx_chunk_len=512 --min_ctx_chunk_len=25 --num_chunk_probs='{"1":"0.5", "2":"0.125", "3":"0.0625", "4":"0.0625", "5":"0.0625", "6":"0.0625", "7":"0.0625", "8":"0.0625"}' --max_val_samples_per_ds=100 --seed=1 + +# eval +WANDB_MODE=disabled srun --partition=aiscilow --gpus=1 --unbuffered uv run run_eval.py --checkpoint_path CHECKPOINT_PATH --datasets ctx_magic_number_32_1024 ctx_magic_number_1024_2048 ctx_magic_number_2048_3072 ctx_magic_number_3072_4096 ctx_magic_number_4096_5120 ctx_magic_number_5120_6144 ctx_magic_number_6144_7168 ctx_magic_number_7168_8192 ctx_magic_number_8192_9216 ctx_magic_number_9216_10240 ctx_magic_number_10240_11264 ctx_magic_number_11264_12288 ctx_magic_number_12288_13312 ctx_magic_number_13312_14336 ctx_magic_number_14336_15360 ctx_magic_number_15360_16384 ctx_magic_number_16384_20480 ctx_magic_number_20480_24576 ctx_magic_number_24576_28672 ctx_magic_number_28672_32768 --max_ctx_chunk_len=1024 --split test diff --git a/src/ctx_to_lora/data/definitions.py b/src/ctx_to_lora/data/definitions.py index b8eb625..6afaf15 100644 --- a/src/ctx_to_lora/data/definitions.py +++ b/src/ctx_to_lora/data/definitions.py @@ -698,6 +698,7 @@ tok_bins = [(64, 128), (128, 256), (256, 512)] + [ tok_bins += [(32, 128), (128, 256), (256, 512), (512, 1024), (32, 1024)] + [ (1024 * i, 1024 * (i + 1)) for i in range(1, 16) ] +tok_bins += [(2**14 + 2**12 * (i), 2**14 + 2**12 * (i + 1)) for i in range(4)] for toy_ds_name in ["ctx_numbers", "ctx_kv", "ctx_magic_number"]: for tok_bin in tok_bins: DS_KWARGS[f"{toy_ds_name}_{tok_bin[0]}_{tok_bin[1]}"] = dict(