From 610750e6fb05f85de1d95e816867a66dc6e8f861 Mon Sep 17 00:00:00 2001 From: Rujikorn Charakorn Date: Wed, 30 Jul 2025 18:38:51 +0900 Subject: [PATCH] refactor after distillation (#6) --- README.md | 54 +- configs/context_numbers_10.yaml | 2 - configs/context_numbers_10_self_gen.yaml | 2 - .../qa_short_ctx_self_gen_lv3_tiny.yaml | 10 +- data/self_generate_qa.py | 13 +- intx_sft.py | 105 +- run_eval.py | 20 +- .../short_ctx/gemma_qa_short_ctx_old_conf.sh | 8 +- .../gemma_qa_short_ctx_old_conf_distill.sh | 6 +- ...gemma_qa_short_ctx_old_conf_l2l_distill.sh | 42 + src/ctx_to_lora/configs.py | 16 +- src/ctx_to_lora/data/collator.py | 136 +- src/ctx_to_lora/data/packing.py | 107 +- src/ctx_to_lora/data/preprocessing_fn.py | 280 ++++ src/ctx_to_lora/data/processing.py | 1222 ++--------------- src/ctx_to_lora/eval_utils.py | 7 +- src/ctx_to_lora/modeling/hypernet.py | 186 +-- src/ctx_to_lora/modeling/lora_layer.py | 1 - src/ctx_to_lora/trainer.py | 122 +- webui/app.py | 3 +- 20 files changed, 694 insertions(+), 1648 deletions(-) create mode 100644 scripts/short_ctx/gemma_qa_short_ctx_old_conf_l2l_distill.sh create mode 100644 src/ctx_to_lora/data/preprocessing_fn.py diff --git a/README.md b/README.md index 2f0dc8e..a47a1a3 100644 --- a/README.md +++ b/README.md @@ -31,12 +31,12 @@ print(model.decode(outputs)) ## 🏋️ Training ### 🔢 HyperLoRA w/ context_numbers_10 ```bash -WANDB_MODE=disabled run uv run intx_sft.py configs/context_numbers_10.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=3 --per_device_train_batch_size=128 --gradient_accumulation_steps=1 --per_device_eval_batch_size=64 --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=1000 --save_steps=1000 --learning_rate=4e-5 --lora_dropout=0.0 --neftune_noise_alpha=5 --use_light_weight_lora=False --load_best_model_at_end=False --add_negative_prompt=False --add_repeat_prompt=False --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 --use_sequence_packing=True --max_packed_inp_len=4096 --max_packed_ctx_len=4096 --dataloader_num_workers=0 --dataloader_prefetch_factor=None --eval_on_start=False --ctx_encoder_type=early_exit --n_latent_queries=208 +WANDB_MODE=disabled run uv run intx_sft.py configs/context_numbers_10.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=3 --per_device_train_batch_size=128 --gradient_accumulation_steps=1 --per_device_eval_batch_size=64 --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=1000 --save_steps=1000 --learning_rate=4e-5 --lora_dropout=0.0 --neftune_noise_alpha=5 --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 --use_sequence_packing=True --max_packed_inp_len=4096 --max_packed_ctx_len=4096 --dataloader_num_workers=0 --dataloader_prefetch_factor=None --eval_on_start=False --ctx_encoder_type=early_exit --n_latent_queries=208 ``` KL loss ```bash - WANDB_MODE=disabled run uv run intx_sft.py configs/context_numbers_10.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=3 --per_device_train_batch_size=128 --gradient_accumulation_steps=1 --per_device_eval_batch_size=64 --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=1000 --save_steps=1000 --learning_rate=4e-5 --lora_dropout=0.0 --neftune_noise_alpha=5 --use_light_weight_lora=False --load_best_model_at_end=False --add_negative_prompt=False --add_repeat_prompt=False --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.0 --use_sequence_packing=True --max_packed_inp_len=4096 --max_packed_ctx_len=4096 --dataloader_num_workers=0 --dataloader_prefetch_factor=None --eval_on_start=False --ctx_encoder_type=early_exit --n_latent_queries=208 --use_kl_loss=True --eval_steps=99999999 --use_liger_kernel=False +WANDB_MODE=disabled run uv run intx_sft.py configs/context_numbers_10_self_gen.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=3 --per_device_train_batch_size=128 --gradient_accumulation_steps=2 --per_device_eval_batch_size=64 --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=5 --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=True ``` ### FWQA-v2 Level-0 Tiny ```bash @@ -76,6 +76,11 @@ uv run python data/generate_fav_num.py uv run python data/generate_fav_num_big.py ``` +Self-gen for the number toy dataset +```bash +run uv run python data/self_generate_qa.py --vllm_model google/gemma-2-2b-it --ds_names context_numbers_2_10 --split train +``` + ### SQuAD ```bash # for some reason download directly through `load_dataset` does not work @@ -145,7 +150,10 @@ uv run data/download_fineweb_edu.py 1. Recursively generate more data! (depends on step 0) ```bash # run from 000 to 0013 -run uv run python data/generate_fw_edu_qa_v2.py --shard_pattern "000_00000" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it --max_length=2000 --max_model_length=2048; run uv run python data/generate_fw_edu_qa_v2_repeat.py --shard_pattern "min_0_to_2000/000*level_0*" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it; run uv run python data/generate_fw_edu_qa_v2_repeat.py --shard_pattern "min_0_to_2000/000*level_1*" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it; run uv run python data/generate_fw_edu_qa_v2_repeat.py --shard_pattern "min_0_to_2000/000*level_2*" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it +run uv run python data/generate_fw_edu_qa_v2.py --shard_pattern "000_00000" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it --max_length=2000 --max_model_length=2048; +run uv run python data/generate_fw_edu_qa_v2_repeat.py --shard_pattern "min_0_to_2000/000*level_0*" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it; +run uv run python data/generate_fw_edu_qa_v2_repeat.py --shard_pattern "min_0_to_2000/000*level_1*" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it; +run uv run python data/generate_fw_edu_qa_v2_repeat.py --shard_pattern "min_0_to_2000/000*level_2*" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it ``` 2. Self-generated response QA data (depends on step 0 and 1) @@ -157,47 +165,11 @@ uv run python data/self_generate_qa.py --vllm_model google/gemma-2-2b-it --glob_ # val split uv run python data/self_generate_qa.py --vllm_model google/gemma-2-2b-it --glob_pattern 'data/raw_datasets/fw_qa_v2/min_0_to_2000/*_level_0_val.parquet' -# self-gen data for other ds listed in qa_no_fw.yaml -uv run python data/self_generate_qa.py --vllm_model google/gemma-2-2b-it --config configs/self_gen_qa_short_ctx_no_fw_qa.yaml +# self-gen data for other ds listed in qa_short_ctx_self_gen_no_fw_qa.yaml +uv run python data/self_generate_qa.py --vllm_model google/gemma-2-2b-it --config configs/qa_short_ctx_self_gen_no_fw_qa.yaml ``` - -