Ctx-to-LoRA


--- [Project doc](https://docs.google.com/document/d/1RCQDzlVU7YGoTwR84gLQfhxTv0RFnW6srQlqfCC2bvQ/edit?usp=sharing) ## 🚀 API Usage [WIP] ```python from ctx_to_lora.modeling import ModulatedPretrainedModel model = ModulatedPretrainedModel.from_state_dict(...) ctx_info = "..." query = "..." ctx_ids = model.ctx_encoder.tokenize(ctx_info) input_ids = model.tokenize(query) outputs = model.generate(ctx_ids, input_ids) print(model.decode(outputs)) ``` ## 🏋️ Training ### 🔢 HyperLoRA w/ context_numbers_10 ```bash WANDB_MODE=disabled uv run intx_sft.py configs/context_numbers_10.yaml --model_name_or_path=google/gemma-3-1b-it --num_train_epochs=1 --per_device_train_batch_size=64 --gradient_accumulation_steps=1 --per_device_eval_batch_size=64 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj --num_self_attends_per_block=8 --num_latent_factor=2 --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 --use_sequence_packing=False --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 --use_sequence_packing=True --max_packed_inp_len=8000 --max_packed_ctx_len=16000 ``` ### Squad + Hotpot training (for testing/debugging) ```bash WANDB_MODE=disabled uv run intx_sft.py configs/squad.yaml --model_name_or_path=google/gemma-3-1b-it --num_train_epochs=5 --per_device_train_batch_size=64 --gradient_accumulation_steps=8 --per_device_eval_batch_size=64 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj --num_self_attends_per_block=8 --num_latent_factor=1 --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 --metric_for_best_model=eval_pwc_loss --add_negative_prompt=False --add_repeat_prompt=False --use_sequence_packing=True --max_packed_inp_len=16000 --max_packed_ctx_len=32000 --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 --logging_steps=10 ``` ### HyperLoRA w/ self-gen 3 mini ```bash # self-gen sft # 3_mini config uv run python data/self_generate_qa.py \ --vllm_model=google/gemma-2-2b-it --config=configs/self_gen_3_mini.yaml WANDB_MODE=disabled uv run python intx_sft.py configs/self_gen_3_mini.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=10 --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 --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 --gen_lora_l1_reg_coef=0.1 ``` ### HyperLoRA w/ fw-qa pretrain only ```bash WANDB_MODE=disabled run uv run python intx_sft.py configs/fw_qa_pretrain_only.yaml --model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 --per_device_train_batch_size=4 --gradient_accumulation_steps=8 --per_device_eval_batch_size=8 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj --num_self_attends_per_block=4 --num_latent_factor=1 --lora_r=8 --eval_steps=100 --save_steps=100 --learning_rate=4e-5 --lora_dropout=0.0 --neftune_noise_alpha=5 --use_light_weight_lora=False --add_negative_prompt=False --add_repeat_prompt=False --use_sequence_packing=True --per_rank_gen=True --per_layer_processing=True --gen_lora_l1_reg_coef=0.1 ``` ### Favourite numbers data generation ```bash uv run python data/generate_fav_num.py # or uv run python data/generate_fav_num_big.py ``` ### Data Generation (v2) Recursively generate more data! ```bash 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/*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/*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/*level_2" --n_qa_pairs=5 --vllm_model=google/gemma-3-12b-it ``` ### Continue from a checkpoint ```bash run python intx_sft.py configs/...yaml ... --from_pretrained_checkpoint=train_outputs/runs/May09_16-25-35_slurm0-a3nodeset-4_59459_ea85a571/checkpoint-10000/pytorch_model.bin --resume_from_checkpoint=train_outputs/runs/May09_16-25-35_slurm0-a3nodeset-4_59459_ea85a571/checkpoint-10000 ``` ### Evaluation LongBench ```bash # generative WANDB_MODE=disabled uv run python run_eval.py --checkpoint_path train_outputs/runs/.../pytorch_model.bin --datasets negative_nq triviaqa_retrieved squad longbench_e --split test # hypernet checkpoint WANDB_MODE=disabled uv run python run_eval.py --checkpoint_path train_outputs/runs/May08_13-56-31_slurm0-a3nodeset-5_59383_906acb28/checkpoint-105000/pytorch_model.bin --datasets negative_nq triviaqa_retrieved squad longbench_e --split test # base model WANDB_MODE=disabled uv run python run_eval.py --model_name_or_path google/gemma-2-2b-it --datasets negative_nq triviaqa_retrieved squad longbench_e --split test --eval_batch_size 2 # base model w/o context WANDB_MODE=disabled uv run python run_eval.py --model_name_or_path google/gemma-2-2b-it --datasets negative_nq triviaqa_retrieved squad longbench_e --split test --remove_context # # benchmark # cd LongBench/LongBench # run python pred_ctx_to_lora.py --checkpoint_path ../../train_outputs/runs/Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782/pytorch_model.bin # run python eval_ctx_to_lora.py --model_name Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782 --checkpoint_path ../../train_outputs/runs/Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782/pytorch_model.bin ``` ### LLM-comparator ```bash # install nvm curl -o- https://raw.githubusercontent.com/nvm-sh/nvm/v0.40.3/install.sh | bash nvm install 16 nvm use 16 git clone https://github.com/PAIR-code/llm-comparator.git cd llm-comparator npm install npm run build # running llm-comparator webui npm run serve # in another terminal # run http server for file fetching # cd back to root folder first # taken from https://stackoverflow.com/a/79135787 alias srv='echo -e "from sys import argv as a\nfrom http.server import HTTPServer as H, SimpleHTTPRequestHandler as HH, test as t\nclass C(HH):\n def end_headers (self):\n self.send_header(a[2],a[3])\n HH.end_headers(self)\nt(C,H,port=int(a[1]))" | /usr/bin/env python3 -- - 8001 "Access-Control-Allow-Origin" "*"' srv # copy-paste the relative path from the root # e.g., http://localhost:8001/train_outputs/runs/May08_13-56-31_slurm0-a3nodeset-5_59383_906acb28/eval-results-105000/comparator.json ```