mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
3.4 KiB
3.4 KiB
Doc-to-LoRA
🚀 Python API Usage
# caveat: this interface only supports non-batched inputs
# for batched inference please see `src/ctx_to_lora/modeling/hypernet.py`
import torch
from ctx_to_lora.model_loading import get_tokenizer
from ctx_to_lora.modeling.hypernet import ModulatedPretrainedModel
# 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]))
🎮 Interactive Demo [WIP]
uv run webui/demo.py
🧪 Experimental Scripts
| Experiment | Data prep | Training | Evaluation | Notes |
|---|---|---|---|---|
| Main experiment | uv run bash scripts/main_exp/0-download_data.shuv 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/<script>.sh (pick from base_model.sh, cd.sh, cd_oracle.sh, d2l.sh, llmlingua.sh, t2l.sh) |
Downloading data is fastest; regenerate only if you need fresh synthetic data. Evaluation scripts reproduce the main paper metrics. |
| NIAH | uv run bash scripts/niah/0-gen_data.sh |
uv run bash scripts/niah/1-train.sh |
uv run bash scripts/niah/2-eval.sh |
Run the scripts in order; data generation only needs to happen once |