This commit is contained in:
51616 2025-09-15 04:04:32 +09:00
parent 3105fac6dc
commit 8a6c2cfd42
7 changed files with 1471 additions and 23 deletions

View file

@ -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))

View file

@ -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,

View file

@ -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)}")

View file

@ -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")),

View file

@ -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,

View 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)}"
)

File diff suppressed because it is too large Load diff