mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
add t2l
This commit is contained in:
parent
3105fac6dc
commit
8a6c2cfd42
7 changed files with 1471 additions and 23 deletions
10
run_eval.py
10
run_eval.py
|
|
@ -118,6 +118,16 @@ if __name__ == "__main__":
|
|||
default=0.9,
|
||||
help="Compression rate for LLMLingua",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_t2l",
|
||||
action="store_true",
|
||||
help="Use Text-to-LoRA model for evaluation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--add_ctx_to_input",
|
||||
action="store_true",
|
||||
help="Add ctx to base model's input",
|
||||
)
|
||||
|
||||
cli_args = vars(parser.parse_args())
|
||||
# setup_logging(output_dir, debug=os.getenv("DEBUG", False))
|
||||
|
|
|
|||
|
|
@ -225,7 +225,9 @@ def get_tokenized_dataset(
|
|||
)
|
||||
logger.info(f"Loading dataset {ds_name} with split {split}...")
|
||||
# TODO: fix this to allow using both ctx_ids and add_ctx_ids
|
||||
need_ctx_ids = not add_ctx_to_chat and bool(ctx_model_max_len)
|
||||
need_ctx_ids = (
|
||||
ctx_model_max_len is not None
|
||||
) # not add_ctx_to_chat and bool(ctx_model_max_len)
|
||||
|
||||
load_and_process_kwargs = dict(
|
||||
ds_name=ds_name,
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ import yaml
|
|||
from datasets import disable_caching
|
||||
from peft import get_peft_model
|
||||
from transformers import (
|
||||
PreTrainedModel,
|
||||
Seq2SeqTrainer,
|
||||
Seq2SeqTrainingArguments,
|
||||
Trainer,
|
||||
|
|
@ -53,6 +54,7 @@ from ctx_to_lora.modeling import hypernet
|
|||
from ctx_to_lora.modeling.context_distillation import CtxDistillModel
|
||||
from ctx_to_lora.modeling.hypernet import ModulatedPretrainedModel
|
||||
from ctx_to_lora.modeling.llm_lingua import LLMLinguaModel
|
||||
from ctx_to_lora.modeling.text_to_lora import TextToLoRA
|
||||
from ctx_to_lora.tracker.tracker import (
|
||||
add_tracker,
|
||||
print_global_tracker_stats,
|
||||
|
|
@ -764,6 +766,7 @@ def evaluate(
|
|||
|
||||
use_cd = False
|
||||
ctx_model_max_len = None
|
||||
base_model = None
|
||||
|
||||
if model_name_or_path is None:
|
||||
try:
|
||||
|
|
@ -785,14 +788,14 @@ def evaluate(
|
|||
add_tracker(model.combine_lora, "combine_lora")
|
||||
add_tracker(model.apply_lora_to_layers, "apply_lora_to_layers")
|
||||
else:
|
||||
model = get_model(
|
||||
model = base_model = get_model(
|
||||
model_name_or_path,
|
||||
train=False,
|
||||
requires_grad=False,
|
||||
model_kwargs=model_kwargs,
|
||||
use_flash_attn=True,
|
||||
)
|
||||
add_tracker(model.generate, "generate")
|
||||
add_tracker(base_model.generate, "generate")
|
||||
if use_cd := getattr(args, "use_cd", False):
|
||||
peft_config = get_lora_config(
|
||||
model_name_or_path,
|
||||
|
|
@ -801,7 +804,7 @@ def evaluate(
|
|||
target_modules=["down_proj"],
|
||||
)
|
||||
peft_config.lora_alpha = 16
|
||||
peft_model = get_peft_model(model, peft_config)
|
||||
peft_model = get_peft_model(base_model, peft_config)
|
||||
sep_seq = (
|
||||
tokenizer(
|
||||
SELF_QA_INTX.strip("\n"),
|
||||
|
|
@ -809,16 +812,17 @@ def evaluate(
|
|||
return_tensors="pt",
|
||||
)
|
||||
.input_ids[0]
|
||||
.to(model.device)
|
||||
.to(base_model.device)
|
||||
)
|
||||
ctx_distill_kwargs = dict(
|
||||
prefix_tokens=torch.tensor(
|
||||
CTX_AFFIXES[model_name_or_path]["prefix"], device=model.device
|
||||
CTX_AFFIXES[model_name_or_path]["prefix"], device=base_model.device
|
||||
),
|
||||
ctx_inp_sep_seq=sep_seq,
|
||||
pad_token_id=tokenizer.pad_token_id,
|
||||
update_iterations=args.cd_update_iterations,
|
||||
tokenizer=tokenizer,
|
||||
reprompt_ctx=args.add_ctx_to_input,
|
||||
)
|
||||
if args.cd_use_gen_q:
|
||||
q_model, q_tokenizer = get_model_and_tokenizer(
|
||||
|
|
@ -835,18 +839,27 @@ def evaluate(
|
|||
add_tracker(model.teacher_generate, "teacher_generate")
|
||||
add_tracker(model.student_generate, "student_generate")
|
||||
elif use_llmlingua := getattr(args, "use_llmlingua", False):
|
||||
model = LLMLinguaModel(model, tokenizer, args.llmlingua_compression_rate)
|
||||
model = LLMLinguaModel(
|
||||
base_model, tokenizer, args.llmlingua_compression_rate
|
||||
)
|
||||
ctx_model_max_len = model.base_model.config.max_position_embeddings
|
||||
add_tracker(model.base_model.generate, "generate")
|
||||
add_tracker(model.compress, "prompt_compress")
|
||||
|
||||
base_model = (
|
||||
model.base_model
|
||||
if isinstance(model, ModulatedPretrainedModel)
|
||||
or isinstance(model, CtxDistillModel)
|
||||
or isinstance(model, LLMLinguaModel)
|
||||
else model
|
||||
)
|
||||
elif use_t2l := getattr(args, "use_t2l", False):
|
||||
model = TextToLoRA(
|
||||
base_model.name_or_path,
|
||||
prefix_tokens=torch.tensor(
|
||||
CTX_AFFIXES[model_name_or_path]["prefix"], device=base_model.device
|
||||
),
|
||||
device=base_model.device,
|
||||
)
|
||||
ctx_model_max_len = model.base_model.config.max_position_embeddings
|
||||
add_tracker(model.base_model.generate, "base_model.generate")
|
||||
add_tracker(model.generate_weights, "generate_weights")
|
||||
|
||||
if base_model is None:
|
||||
base_model = model.base_model
|
||||
base_model.config.pad_token_id = tokenizer.pad_token_id
|
||||
base_model.generation_config.pad_token_id = tokenizer.pad_token_id
|
||||
|
||||
|
|
@ -857,12 +870,10 @@ def evaluate(
|
|||
ctx_tokenizer.pad_token_id = ctx_tokenizer.eos_token_id
|
||||
|
||||
add_ctx_to_chat = (
|
||||
not (
|
||||
isinstance(model, ModulatedPretrainedModel)
|
||||
or isinstance(model, LLMLinguaModel)
|
||||
)
|
||||
and not args.remove_context
|
||||
) or isinstance(model, CtxDistillModel)
|
||||
(isinstance(model, PreTrainedModel) and not args.remove_context)
|
||||
or isinstance(model, CtxDistillModel)
|
||||
or args.add_ctx_to_input
|
||||
)
|
||||
|
||||
_get_tokenized_dataset = partial(
|
||||
get_tokenized_dataset,
|
||||
|
|
@ -1038,13 +1049,20 @@ def run_eval(
|
|||
use_iterative_mode: bool = False,
|
||||
use_llmlingua: bool = False,
|
||||
llmlingua_compression_rate: float = 0.9,
|
||||
use_t2l: bool = False,
|
||||
add_ctx_to_input: bool = False,
|
||||
) -> None:
|
||||
"""Run evaluation with the specified parameters."""
|
||||
assert bool(model_name_or_path) ^ bool(checkpoint_path), (
|
||||
"Either --model_name_or_path or --checkpoint_path must be provided"
|
||||
)
|
||||
if (use_cd or use_llmlingua) and eval_batch_size != 1:
|
||||
raise ValueError("When using context distillation, eval_batch_size must be 1.")
|
||||
if (use_cd or use_llmlingua or use_t2l) and eval_batch_size != 1:
|
||||
raise ValueError("When using a baseline method, eval_batch_size must be 1.")
|
||||
|
||||
if use_llmlingua and add_ctx_to_input:
|
||||
raise ValueError(
|
||||
"LLMLingua always adds compressed context to input by default."
|
||||
)
|
||||
|
||||
disable_caching()
|
||||
set_seed(42)
|
||||
|
|
@ -1099,10 +1117,13 @@ def run_eval(
|
|||
if use_llmlingua:
|
||||
args.use_llmlingua = use_llmlingua
|
||||
args.llmlingua_compression_rate = llmlingua_compression_rate
|
||||
if use_t2l:
|
||||
args.use_t2l = use_t2l
|
||||
if max_val_samples_per_ds > 0:
|
||||
args.max_val_samples_per_ds = max_val_samples_per_ds
|
||||
if max_test_samples_per_ds > 0:
|
||||
args.max_test_samples_per_ds = max_test_samples_per_ds
|
||||
args.add_ctx_to_input = add_ctx_to_input
|
||||
setup_logging(args.logging_dir)
|
||||
logger.debug(f"CMD: {' '.join(os.sys.argv)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ LENGTH_BINS = [
|
|||
(2**10, 2**11 - 1),
|
||||
(2**11, 2**12 - 1),
|
||||
(2**12, 2**13 - 1),
|
||||
(0, 2**13 - 1),
|
||||
(2**13, 2**14 - 1),
|
||||
(2**14, 2**15 - 1),
|
||||
(2**15, float("inf")),
|
||||
|
|
|
|||
|
|
@ -1129,7 +1129,6 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
@torch.inference_mode()
|
||||
def generate(
|
||||
self,
|
||||
# TODO: allow more than one LoRA per sample (multi-lora)
|
||||
ctx_ids: Integer[Tensor, "n_chunks ctx_length"] | None = None,
|
||||
ctx_attn_mask: Integer[Tensor, "n_chunks ctx_length"] | None = None,
|
||||
ctx_position_ids: Integer[Tensor, "n_chunks ctx_length"] | None = None,
|
||||
|
|
|
|||
160
src/ctx_to_lora/modeling/text_to_lora.py
Normal file
160
src/ctx_to_lora/modeling/text_to_lora.py
Normal file
|
|
@ -0,0 +1,160 @@
|
|||
from functools import partial
|
||||
|
||||
import torch
|
||||
from peft import PeftConfig
|
||||
from torch import nn
|
||||
|
||||
from ctx_to_lora.modeling.lora_layer import apply_lora_to_layers, lora_forward
|
||||
from ctx_to_lora.modeling.text_to_lora_impl import (
|
||||
embed_texts,
|
||||
get_layers,
|
||||
get_peft_config,
|
||||
load_hypermod,
|
||||
)
|
||||
from ctx_to_lora.utils import get_peft_modules
|
||||
|
||||
|
||||
class TextToLoRA(nn.Module):
|
||||
def __init__(self, model_name_or_path, prefix_tokens, device):
|
||||
assert model_name_or_path == "google/gemma-2-2b-it"
|
||||
super().__init__()
|
||||
hypermod_dir = "trained_t2l/gemma_2b_t2l"
|
||||
peft_config = get_peft_config(
|
||||
PeftConfig.from_json_file(f"{hypermod_dir}/adapter_config.json")
|
||||
)
|
||||
|
||||
# ours lora forward pass uses alpha directly
|
||||
peft_config.lora_alpha = peft_config.lora_alpha / peft_config.r
|
||||
self.prefix_tokens = prefix_tokens
|
||||
self.device = device
|
||||
(
|
||||
_,
|
||||
self.t2l_model,
|
||||
self.base_model,
|
||||
self.tokenizer,
|
||||
self.emb_model,
|
||||
self.emb_tokenizer,
|
||||
self.task_desc_format_fn,
|
||||
self.pooling_fn,
|
||||
) = load_hypermod(hypermod_dir, device)
|
||||
layer_indices = range(len(get_layers(self.base_model)))
|
||||
|
||||
self.layer_indices = torch.tensor(
|
||||
layer_indices, dtype=torch.long, device=device
|
||||
)
|
||||
# patch base model forward pass to use lora
|
||||
layers = get_layers(self.base_model)
|
||||
lora_forward_fn = lora_forward
|
||||
|
||||
for layer_idx in self.layer_indices:
|
||||
for module_info in get_peft_modules(layers[layer_idx], peft_config):
|
||||
module = module_info["module"]
|
||||
module.forward = partial(
|
||||
lora_forward_fn,
|
||||
self=module,
|
||||
lora_dropout_p=peft_config.lora_dropout,
|
||||
scaling=peft_config.lora_alpha,
|
||||
)
|
||||
|
||||
@property
|
||||
def generation_config(self):
|
||||
return self.base_model.generation_config
|
||||
|
||||
def generate_weights(self, ctx_txt: str):
|
||||
# generate loras
|
||||
ctx_emb = embed_texts(
|
||||
[ctx_txt],
|
||||
self.emb_model,
|
||||
self.emb_tokenizer,
|
||||
self.task_desc_format_fn,
|
||||
self.pooling_fn,
|
||||
self.device,
|
||||
)
|
||||
encoder_out = self.t2l_model.task_encoder(ctx_emb)
|
||||
encoded_task_emb = encoder_out["encoded_task_emb"].detach()
|
||||
|
||||
lora_A, lora_B = dict(), dict()
|
||||
lora_dict = dict()
|
||||
for target_module in self.t2l_model.target_modules:
|
||||
factorized_delta_w = self.t2l_model.get_delta_weights(
|
||||
self.layer_indices,
|
||||
target_module,
|
||||
encoded_task_emb.expand(self.layer_indices.shape[0], -1),
|
||||
factorized=True,
|
||||
)
|
||||
# lora_A[target_module]: [n_layers, r, d_in]
|
||||
# lora_A[target_module]: [n_layers, d_out, r]
|
||||
lora_A[target_module], lora_B[target_module] = factorized_delta_w
|
||||
|
||||
# convert to lora format used by lora_forward
|
||||
# dict of {module:
|
||||
# {A: [bs, n_layers, r, d_inim],
|
||||
# B: [bs, n_layers, r, d_outim]}}
|
||||
lora_dict[target_module] = dict(
|
||||
A=lora_A[target_module].unsqueeze(0),
|
||||
B=lora_B[target_module].transpose(-1, -2).unsqueeze(0),
|
||||
)
|
||||
return lora_dict
|
||||
|
||||
def generate(self, *args, **kwargs):
|
||||
ctx_ids_full = kwargs["ctx_ids"]
|
||||
ctx_txt = self.tokenizer.decode(
|
||||
ctx_ids_full[0, len(self.prefix_tokens) :], skip_special_tokens=True
|
||||
)
|
||||
generated_loras = self.generate_weights(ctx_txt)
|
||||
apply_lora_to_layers(
|
||||
self.base_model,
|
||||
self.layer_indices,
|
||||
generated_loras,
|
||||
n_qs=torch.tensor([1], device=self.device),
|
||||
position_ids=None,
|
||||
)
|
||||
kwargs.pop("ctx_ids", None)
|
||||
kwargs.pop("ctx_attn_mask", None)
|
||||
kwargs.pop("n_ctx_chunks", None)
|
||||
return self.base_model.generate(*args, **kwargs)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from ctx_to_lora.data.definitions import CTX_AFFIXES
|
||||
from ctx_to_lora.data.processing import load_and_process_dataset
|
||||
|
||||
model_name = "google/gemma-2-2b-it"
|
||||
tokenizer = AutoTokenizer.from_pretrained(model_name)
|
||||
|
||||
ds = load_and_process_dataset("pwc", split="train", num_proc=8)
|
||||
ctx = ds[0]["context"]
|
||||
inp = ds[1]["prompts"][0]
|
||||
ctx_ids = tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": ctx}], return_tensors="pt", return_dict=True
|
||||
)
|
||||
input_ids = tokenizer.apply_chat_template(
|
||||
[{"role": "user", "content": ctx + "\n\n" + inp}],
|
||||
return_tensors="pt",
|
||||
return_dict=True,
|
||||
)
|
||||
ctx_ids = {k: v.to("cuda") for k, v in ctx_ids.items()}
|
||||
input_ids = {k: v.to("cuda") for k, v in input_ids.items()}
|
||||
|
||||
prefix_tokens = CTX_AFFIXES[model_name]["prefix"]
|
||||
prefix_tokens = torch.tensor(prefix_tokens, dtype=torch.long)
|
||||
|
||||
t2l_model = TextToLoRA(
|
||||
model_name,
|
||||
prefix_tokens,
|
||||
device="cuda",
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
for _ in range(1):
|
||||
outputs = t2l_model.generate(
|
||||
**input_ids,
|
||||
ctx_ids=ctx_ids["input_ids"],
|
||||
max_new_tokens=256,
|
||||
do_sample=False,
|
||||
)
|
||||
print(
|
||||
f"Student response: {tokenizer.batch_decode(outputs, skip_special_tokens=False)}"
|
||||
)
|
||||
1255
src/ctx_to_lora/modeling/text_to_lora_impl.py
Normal file
1255
src/ctx_to_lora/modeling/text_to_lora_impl.py
Normal file
File diff suppressed because it is too large
Load diff
Loading…
Add table
Add a link
Reference in a new issue