From 3b798c670c8f5ef493704a56a345c87f5c4dbbe1 Mon Sep 17 00:00:00 2001 From: 51616 Date: Wed, 1 Oct 2025 13:11:54 +0000 Subject: [PATCH] python api + remove `generate_multi_lora` + rename config path + test eval scripts --- README.md | 93 +++++----- .../self_gen_lv1_closed_qa_1_and_lv3_l2l.yaml | 32 ++++ .../self_gen_lv1_closed_qa_1_l2l.yaml | 10 - .../ctx_magic_number_32_256.yaml | 0 data/sakana_wiki.txt | 8 + data/self_generate_qa.py | 70 ++----- examples/python_api.py | 43 +++++ install.sh | 8 +- scripts/main_exp/0-download_data.py | 11 ++ scripts/main_exp/{train.sh => 1-train.sh} | 0 scripts/main_exp/README.md | 25 +++ scripts/main_exp/eval/base_model.sh | 2 +- scripts/main_exp/eval/base_model_test.sh | 8 + scripts/main_exp/eval/cd.sh | 4 +- scripts/main_exp/eval/cd_oracle.sh | 4 +- scripts/main_exp/eval/cd_oracle_test.sh | 6 + scripts/main_exp/eval/cd_test.sh | 6 + scripts/main_exp/eval/d2l.sh | 10 +- scripts/main_exp/eval/d2l_test.sh | 13 ++ scripts/main_exp/eval/llmlingua_test.sh | 5 + scripts/main_exp/eval/t2l_test.sh | 4 + scripts/main_exp/gen_data.sh | 4 +- scripts/main_exp/gen_data_test.sh | 12 +- src/ctx_to_lora/data/preprocessing_fn.py | 24 +++ src/ctx_to_lora/eval_utils.py | 3 - src/ctx_to_lora/modeling/hypernet.py | 171 ++++++++---------- webui/app.py | 2 +- 27 files changed, 345 insertions(+), 233 deletions(-) create mode 100644 configs/main_exp/self_gen_lv1_closed_qa_1_and_lv3_l2l.yaml rename configs/{toy_exp => niah_exp}/ctx_magic_number_32_256.yaml (100%) create mode 100644 data/sakana_wiki.txt create mode 100644 examples/python_api.py create mode 100644 scripts/main_exp/0-download_data.py rename scripts/main_exp/{train.sh => 1-train.sh} (100%) create mode 100644 scripts/main_exp/README.md create mode 100755 scripts/main_exp/eval/base_model_test.sh create mode 100755 scripts/main_exp/eval/cd_oracle_test.sh create mode 100755 scripts/main_exp/eval/cd_test.sh create mode 100755 scripts/main_exp/eval/d2l_test.sh create mode 100755 scripts/main_exp/eval/llmlingua_test.sh create mode 100755 scripts/main_exp/eval/t2l_test.sh diff --git a/README.md b/README.md index 56bc8a6..b57cd52 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,5 @@
-

Ctx-to-LoRA

+

Doc-to-LoRA


@@ -7,56 +7,57 @@ --- -## ๐Ÿš€ API Usage [WIP] +## ๐Ÿš€ Python API Usage ```python -from ctx_to_lora.modeling import ModulatedPretrainedModel -model = ModulatedPretrainedModel.from_state_dict(...) +# caveat: this interface only supports non-batched inputs +# for batched inference please see `src/ctx_to_lora/modeling/hypernet.py` +import torch -ctx_info = "..." -query = "..." +from ctx_to_lora.model_loading import get_tokenizer +from ctx_to_lora.modeling.hypernet import ModulatedPretrainedModel -ctx_ids = model.ctx_encoder.tokenize(ctx_info) -input_ids = model.tokenize(query) -outputs = model.generate(ctx_ids, input_ids) -print(model.decode(outputs)) +# model loading +checkpoint_path = ... +state_dict = torch.load(checkpoint_path, weights_only=False) +model = ModulatedPretrainedModel.from_state_dict( + state_dict, train=False, use_sequence_packing=False +) +model.reset() +tokenizer = get_tokenizer(model.base_model.name_or_path) + +# prepare data +doc = open("data/sakana_wiki.txt", "r").read() +chat = [{"role": "user", "content": "Summarize what Sakana AI does."}] +chat_ids = tokenizer.apply_chat_template( + chat, + add_special_tokens=False, + return_attention_mask=False, + add_generation_prompt=True, + return_tensors="pt", +).to(model.device) + + +# calls after internalization will be influenced by internalized info +model.internalize(doc) + +outputs = model.generate(input_ids=chat_ids, max_new_tokens=256) +print(tokenizer.decode(outputs[0])) + + +# remove internalized info +model.reset() + +outputs = model.generate(input_ids=chat_ids, max_new_tokens=256) +print(tokenizer.decode(outputs[0])) ``` - -***Generate data from scratch*** -# 0. download fineweb_edu to `data/raw_datasets/fineweb_edu -uv run data/download_fineweb_edu.py -``` - -1. Recursively generate more data! (depends on step 0) +### ๐ŸŽฎ Interactive Demo [WIP] ```bash -# run from 000 to 0013 -run uv run 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 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; +uv run webui/demo.py ``` -2. Self-generated response QA data (depends on step 0 and 1) -```bash -# Example commands using gemma-2-2b-it -# self-gen data for fw_qa_v2 -uv run data/self_generate_qa.py --vllm_model google/gemma-2-2b-it --glob_pattern 'data/raw_datasets/fw_qa_v2/min_0_to_2000/013*_level_3*' --closed_qa_prob 1.0 # or 0.0 - -# val split -uv run 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_short_ctx_self_gen_no_fw_qa.yaml -uv run data/self_generate_qa.py --vllm_model google/gemma-2-2b-it --config configs/qa_short_ctx_self_gen_no_fw_qa.yaml -``` - -### Evaluation -See [eval scripts](scripts/eval/). -```bash -WANDB_MODE=disabled uv run run_eval.py --checkpoint_path train_outputs/runs/$RUN_NAME/pytorch_model.bin --datasets squad --split test --max_ctx_chunk_len 8192 --eval_batch_size_gen 4 - - -# base model -WANDB_MODE=disabled uv run 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 run_eval.py --model_name_or_path google/gemma-2-2b-it --datasets negative_nq triviaqa_retrieved squad longbench_e --split test --remove_context - - +### ๐Ÿงช Experimental Scripts +| Experiment | Data prep | Training | Evaluation | Notes | +| ------------------------------------ | ----------------------------------------------------------------------------------------------------------------------------------- | ----------------------------------------- | ---------------------------------------------------------------------------------------------------------------------------------------- | ----------------------------------------------------------------------------------------------------------------------------------- | +| [Main experiment](scripts/main_exp/) | `uv run bash scripts/main_exp/0-download_data.sh`
`uv run bash scripts/main_exp/gen_data.sh` *(optional, rebuilds from scratch)* | `uv run bash scripts/main_exp/1-train.sh` | `uv run bash scripts/main_exp/eval/