diff --git a/README.md b/README.md index 49df497..76d095e 100644 --- a/README.md +++ b/README.md @@ -1,20 +1,41 @@ +# `Ctx-to-LoRA`: Compressing Context Information to LoRAs 🧠 +
+ +
+ + +--- + [Project doc](https://docs.google.com/document/d/1RCQDzlVU7YGoTwR84gLQfhxTv0RFnW6srQlqfCC2bvQ/edit?usp=sharing) -### Finetuning the base model with LoRA adaptor + + +## 🚀 `Ctx-to-LoRA` API +```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)) ``` - - -### HyperLoRA w/ context_numbers_10 +## 🏋️ Training +### 🔢 HyperLoRA w/ context_numbers_10 ```bash -WANDB_MODE=disabled run python hyperlora/intx_sft.py configs/context_numbers_10.yaml --model_name_or_path=meta-llama/Llama-3.2-1B-Instruct --num_train_epochs=100 --per_device_train_batch_size=64 --per_device_eval_batch_size=64 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj,up_proj +WANDB_MODE=disabled uv run python intx_sft.py configs/context_numbers_10.yaml --model_name_or_path=meta-llama/Llama-3.2-1B-Instruct --num_train_epochs=100 --per_device_train_batch_size=64 --per_device_eval_batch_size=64 --exp_setup=hyper_lora --aggregator_type=perceiver --target_modules=down_proj --per_rank_gen=True --per_layer_processing=True --decoder_depth=2 --seed=1 --gen_lora_l1_reg_coef=0 --use_token_mixing=False --lora_r=8 --bf16=True --tf32=True --dataloader_num_workers=8 --dataloader_prefetch_factor=8 ``` -### HyperLoRA w/ context_numbers_128 + ### Generate fineweb qa ```bash @@ -26,7 +47,7 @@ conda activate vllm; vllm_model=mistralai/Mistral-Small-3.1-24B-Instruct-2503 run python generate_fw_edu_qa_vllm.py "0*_*" 3 ``` - + ### Continue from a checkpoint ```bash @@ -51,8 +72,16 @@ run python intx_sft.py configs/...yaml ... --from_pretrained_checkpoint=train_ou LongBench ```bash # generative -run python src/ctx_to_lora/eval.py --checkpoint_path train_outputs/runs/.../pytorch_model.bin --datasets negative_nq triviaqa_retrieved hotpot_qa squad longbench_e - --split test +run python src/ctx_to_lora/eval.py --checkpoint_path train_outputs/runs/.../pytorch_model.bin --datasets negative_nq triviaqa_retrieved squad longbench_e --split test + +# hypernet checkpoint +uv run python 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 +uv run python eval.py --model_name_or_path google/gemma-2-2b-it --datasets longbench_e --split test + +# base model w/o context +run uv run python eval.py --model_name_or_path google/gemma-2-2b-it --datasets longbench_e --split test --remove_context # run python src/ctx_to_lora/eval.py --checkpoint_path train_outputs/runs/May08_13-56-31_slurm0-a3nodeset-5_59383_906acb28/checkpoint-105000/pytorch_model.bin diff --git a/assets/cover.png b/assets/cover.png new file mode 100644 index 0000000..17fe20f Binary files /dev/null and b/assets/cover.png differ