mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
add llmlingua + different q model for cd
This commit is contained in:
parent
0e05644d6a
commit
0e46a7c568
5 changed files with 265 additions and 65 deletions
20
run_eval.py
20
run_eval.py
|
|
@ -60,6 +60,15 @@ if __name__ == "__main__":
|
|||
"If -1, uses values from checkpoint config."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_test_samples_per_ds",
|
||||
type=int,
|
||||
default=1000,
|
||||
help=(
|
||||
"Maximum number of validation samples per dataset. "
|
||||
"If -1, uses values from checkpoint config."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_ctx_chunk_len",
|
||||
type=int,
|
||||
|
|
@ -104,6 +113,17 @@ if __name__ == "__main__":
|
|||
action="store_true",
|
||||
help="Use iterative mode LoRA layer-by-layer generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_llmlingua",
|
||||
action="store_true",
|
||||
help="Use LLMLingua compression for evaluation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--llmlingua_compression_rate",
|
||||
type=float,
|
||||
default=0.9,
|
||||
help="Compression rate for LLMLingua",
|
||||
)
|
||||
|
||||
cli_args = vars(parser.parse_args())
|
||||
# setup_logging(output_dir, debug=os.getenv("DEBUG", False))
|
||||
|
|
|
|||
|
|
@ -43,10 +43,16 @@ from ctx_to_lora.metrics import (
|
|||
compute_prefix_matching,
|
||||
compute_rouge,
|
||||
)
|
||||
from ctx_to_lora.model_loading import get_lora_config, get_model, get_tokenizer
|
||||
from ctx_to_lora.model_loading import (
|
||||
get_lora_config,
|
||||
get_model,
|
||||
get_model_and_tokenizer,
|
||||
get_tokenizer,
|
||||
)
|
||||
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.tracker.tracker import (
|
||||
add_tracker,
|
||||
print_global_tracker_stats,
|
||||
|
|
@ -757,6 +763,7 @@ def evaluate(
|
|||
tokenizer.pad_token_id = tokenizer.eos_token_id
|
||||
|
||||
use_cd = False
|
||||
ctx_model_max_len = None
|
||||
|
||||
if model_name_or_path is None:
|
||||
try:
|
||||
|
|
@ -771,6 +778,7 @@ def evaluate(
|
|||
use_flash_attn=True,
|
||||
use_sequence_packing=False, # for generation
|
||||
)
|
||||
ctx_model_max_len = model.ctx_encoder.config.max_position_embeddings
|
||||
model.enable_iterative_mode(args.use_iterative_mode)
|
||||
add_tracker(model.base_model.generate, "generate")
|
||||
add_tracker(model.generate_weights, "generate_weights")
|
||||
|
|
@ -794,29 +802,50 @@ def evaluate(
|
|||
)
|
||||
peft_config.lora_alpha = 16
|
||||
peft_model = get_peft_model(model, peft_config)
|
||||
q_model = get_model(
|
||||
"google/gemma-3-4b-it", train=False, requires_grad=False
|
||||
sep_seq = (
|
||||
tokenizer(
|
||||
SELF_QA_INTX.strip("\n"),
|
||||
add_special_tokens=False,
|
||||
return_tensors="pt",
|
||||
)
|
||||
.input_ids[0]
|
||||
.to(model.device)
|
||||
)
|
||||
model = CtxDistillModel(
|
||||
peft_model,
|
||||
ctx_distill_kwargs = dict(
|
||||
prefix_tokens=torch.tensor(
|
||||
CTX_AFFIXES[model_name_or_path]["prefix"], device=model.device
|
||||
),
|
||||
ctx_inp_sep_seq=sep_seq,
|
||||
pad_token_id=tokenizer.pad_token_id,
|
||||
update_iterations=args.cd_update_iterations,
|
||||
q_model=q_model, # peft_model if args.cd_use_gen_q else None,
|
||||
num_gen_q=args.num_gen_q,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
if args.cd_use_gen_q:
|
||||
q_model, q_tokenizer = get_model_and_tokenizer(
|
||||
"google/gemma-3-4b-it",
|
||||
train=False,
|
||||
requires_grad=False,
|
||||
)
|
||||
ctx_distill_kwargs["q_model"] = q_model
|
||||
ctx_distill_kwargs["q_tokenizer"] = q_tokenizer
|
||||
ctx_distill_kwargs["num_gen_q"] = args.num_gen_q
|
||||
model = CtxDistillModel(peft_model, **ctx_distill_kwargs)
|
||||
|
||||
add_tracker(model._distill_context, "distill_context")
|
||||
add_tracker(model.generate_questions, "generate_questions")
|
||||
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)
|
||||
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
|
||||
)
|
||||
base_model.config.pad_token_id = tokenizer.pad_token_id
|
||||
|
|
@ -829,13 +858,12 @@ def evaluate(
|
|||
ctx_tokenizer.pad_token_id = ctx_tokenizer.eos_token_id
|
||||
|
||||
add_ctx_to_chat = (
|
||||
not isinstance(model, ModulatedPretrainedModel) and not args.remove_context
|
||||
not (
|
||||
isinstance(model, ModulatedPretrainedModel)
|
||||
or isinstance(model, LLMLinguaModel)
|
||||
)
|
||||
and not args.remove_context
|
||||
) or isinstance(model, CtxDistillModel)
|
||||
ctx_model_max_len = (
|
||||
model.ctx_encoder.config.max_position_embeddings
|
||||
if isinstance(model, ModulatedPretrainedModel)
|
||||
else None
|
||||
)
|
||||
|
||||
_get_tokenized_dataset = partial(
|
||||
get_tokenized_dataset,
|
||||
|
|
@ -878,8 +906,17 @@ def evaluate(
|
|||
if ds_name in answers:
|
||||
answers[ds_name] = answers[ds_name].select(val_indices)
|
||||
|
||||
print(f"Datasets: {datasets}")
|
||||
print(f"Answers: {answers}")
|
||||
max_test_samples_per_ds = getattr(args, "max_test_samples_per_ds", 0)
|
||||
if split == "test" and max_test_samples_per_ds > 0:
|
||||
print(f"Truncating all test ds to {max_test_samples_per_ds} samples")
|
||||
for ds_name, ds in datasets.items():
|
||||
test_indices = np.random.permutation(len(ds))[:max_test_samples_per_ds]
|
||||
datasets[ds_name] = ds.select(test_indices)
|
||||
if ds_name in answers:
|
||||
answers[ds_name] = answers[ds_name].select(test_indices)
|
||||
|
||||
print(f"Datasets: {datasets}")
|
||||
print(f"Answers: {answers}")
|
||||
|
||||
gen_kwargs = dict(
|
||||
do_sample=False,
|
||||
|
|
@ -924,19 +961,13 @@ def evaluate(
|
|||
if max_ctx_chunk_len > 0:
|
||||
model.generate = model.generate_with_multi_loras
|
||||
|
||||
if isinstance(model, CtxDistillModel):
|
||||
sep_seq = (
|
||||
tokenizer(
|
||||
SELF_QA_INTX.strip("\n"), add_special_tokens=False, return_tensors="pt"
|
||||
)
|
||||
.input_ids[0]
|
||||
.to(model.device)
|
||||
)
|
||||
model.generate = partial(
|
||||
model.generate,
|
||||
ctx_inp_sep_seq=sep_seq,
|
||||
reset=True,
|
||||
)
|
||||
# if isinstance(model, CtxDistillModel):
|
||||
|
||||
# model.generate = partial(
|
||||
# model.generate,
|
||||
# ctx_inp_sep_seq=sep_seq,
|
||||
# reset=True,
|
||||
# )
|
||||
|
||||
trainer_kwargs = {
|
||||
"model": model,
|
||||
|
|
@ -997,6 +1028,7 @@ def run_eval(
|
|||
split: str = "validation",
|
||||
eval_batch_size: int = 8,
|
||||
max_val_samples_per_ds: int = -1,
|
||||
max_test_samples_per_ds: int = -1,
|
||||
max_ctx_chunk_len: int = -1,
|
||||
remove_context: bool = False,
|
||||
max_new_tokens: int = 256,
|
||||
|
|
@ -1006,12 +1038,14 @@ def run_eval(
|
|||
cd_use_gen_q: bool = False,
|
||||
num_gen_q: int = 20,
|
||||
use_iterative_mode: bool = False,
|
||||
use_llmlingua: bool = False,
|
||||
llmlingua_compression_rate: float = 0.9,
|
||||
) -> 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 and eval_batch_size != 1:
|
||||
if (use_cd or use_llmlingua) and eval_batch_size != 1:
|
||||
raise ValueError("When using context distillation, eval_batch_size must be 1.")
|
||||
|
||||
disable_caching()
|
||||
|
|
@ -1065,8 +1099,13 @@ def run_eval(
|
|||
args.cd_update_iterations = cd_update_iterations
|
||||
args.cd_use_gen_q = cd_use_gen_q
|
||||
args.num_gen_q = num_gen_q
|
||||
if use_llmlingua:
|
||||
args.use_llmlingua = use_llmlingua
|
||||
args.llmlingua_compression_rate = llmlingua_compression_rate
|
||||
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
|
||||
setup_logging(args.logging_dir)
|
||||
logger.debug(f"CMD: {' '.join(os.sys.argv)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -130,11 +130,15 @@ class CtxDistillModel(nn.Module):
|
|||
self,
|
||||
base_model: PeftModel,
|
||||
prefix_tokens: Integer[Tensor, "n"],
|
||||
ctx_inp_sep_seq: Integer[Tensor, "m"],
|
||||
pad_token_id: int,
|
||||
update_iterations: int,
|
||||
q_model: PreTrainedModel | None = None,
|
||||
num_gen_q: int | None = None,
|
||||
reset: bool = True,
|
||||
tokenizer=None,
|
||||
q_model: PreTrainedModel | None = None,
|
||||
q_tokenizer=None,
|
||||
num_gen_q: int | None = None,
|
||||
reprompt_ctx: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.register_module("base_model", base_model)
|
||||
|
|
@ -145,9 +149,13 @@ class CtxDistillModel(nn.Module):
|
|||
)
|
||||
self.num_gen_q = num_gen_q
|
||||
self.register_buffer("prefix_tokens", prefix_tokens)
|
||||
self.register_buffer("ctx_inp_sep_seq", ctx_inp_sep_seq)
|
||||
self.tokenizer = tokenizer
|
||||
self.q_tokenizer = q_tokenizer
|
||||
self.pad_token_id = pad_token_id
|
||||
self.update_iterations = update_iterations
|
||||
self.reprompt_ctx = reprompt_ctx
|
||||
self.reset = reset
|
||||
self.device = base_model.device
|
||||
self.to(self.device)
|
||||
|
||||
|
|
@ -177,7 +185,7 @@ class CtxDistillModel(nn.Module):
|
|||
)
|
||||
|
||||
def reset_lora(self):
|
||||
print("Resetiing LoRA")
|
||||
print("Resetting LoRA")
|
||||
for layer in get_peft_layers(self.base_model, self.peft_config):
|
||||
layer.reset_lora_parameters(self.adapter_name, init_lora_weights=True)
|
||||
self._init_optim()
|
||||
|
|
@ -265,8 +273,6 @@ class CtxDistillModel(nn.Module):
|
|||
# n_ctx_chunks: Integer[Tensor, "n_ctx"] | None = None,
|
||||
# n_queries: Integer[Tensor, "n_ctx"] | None = None,
|
||||
*model_inputs_args: Any,
|
||||
ctx_inp_sep_seq: Integer[Tensor, "l"],
|
||||
reset: bool,
|
||||
**model_inputs_kwargs: dict[str, Any],
|
||||
):
|
||||
# where to get the questions???
|
||||
|
|
@ -275,17 +281,15 @@ class CtxDistillModel(nn.Module):
|
|||
|
||||
# update peft module with CD (if labels provided)
|
||||
|
||||
if reset:
|
||||
if self.reset:
|
||||
self.reset_lora()
|
||||
|
||||
# teacher tokens
|
||||
if model_inputs_args:
|
||||
ctx_inp_ids = model_inputs_args[0]
|
||||
else:
|
||||
ctx_inp_ids = model_inputs_kwargs.pop("input_ids")
|
||||
orig_ctx_inp_ids = model_inputs_kwargs.pop("input_ids")
|
||||
ctx_inp_ids = orig_ctx_inp_ids.clone()
|
||||
_, orig_inp_ids = ctx_inp_split(
|
||||
ctx_inp_ids,
|
||||
ctx_inp_sep_seq,
|
||||
self.ctx_inp_sep_seq,
|
||||
self.pad_token_id,
|
||||
self.prefix_tokens,
|
||||
padding_side="left",
|
||||
|
|
@ -295,7 +299,7 @@ class CtxDistillModel(nn.Module):
|
|||
if self.q_model is not None:
|
||||
# Extract context-only portion after separator (remove prefix tokens from first row)
|
||||
ctx_ids_full, _ = ctx_inp_split(
|
||||
ctx_inp_ids, ctx_inp_sep_seq, self.pad_token_id
|
||||
ctx_inp_ids, self.ctx_inp_sep_seq, self.pad_token_id
|
||||
) # [bs, var_len]
|
||||
ctx_ids = ctx_ids_full[0, len(self.prefix_tokens) :]
|
||||
ctx_txt = self.tokenizer.decode(ctx_ids, skip_special_tokens=True)
|
||||
|
|
@ -303,7 +307,7 @@ class CtxDistillModel(nn.Module):
|
|||
messages_list = [
|
||||
build_messages(ctx_txt, 1, 1, 1, 1) for _ in range(self.num_gen_q)
|
||||
]
|
||||
q_inputs = self.tokenizer.apply_chat_template(
|
||||
q_inputs = self.q_tokenizer.apply_chat_template(
|
||||
messages_list,
|
||||
tokenize=True,
|
||||
add_special_tokens=False,
|
||||
|
|
@ -322,11 +326,13 @@ class CtxDistillModel(nn.Module):
|
|||
max_new_tokens=256,
|
||||
do_sample=True,
|
||||
top_p=0.95,
|
||||
temperature=1.0, # high temp for diverse questions
|
||||
temperature=2.0, # high temp for diverse questions
|
||||
)
|
||||
# Slice off the prompt portion
|
||||
gen_only = question_outputs[:, q_inputs["input_ids"].shape[-1] :]
|
||||
questions = self.tokenizer.batch_decode(gen_only, skip_special_tokens=True)
|
||||
questions = self.q_tokenizer.batch_decode(
|
||||
gen_only, skip_special_tokens=True
|
||||
)
|
||||
questions = [q.split("Message:")[-1].strip() for q in questions]
|
||||
|
||||
ctx_inp_messages = [
|
||||
|
|
@ -347,8 +353,6 @@ class CtxDistillModel(nn.Module):
|
|||
ctx_inp_ids = encoded_ctx_inp["input_ids"]
|
||||
ctx_inp_attention_mask = encoded_ctx_inp["attention_mask"]
|
||||
|
||||
# TODO: check labels + loss calculation + padding
|
||||
|
||||
# sample responses first
|
||||
ctx_inp_res_ids = self.teacher_generate(
|
||||
ctx_inp_ids,
|
||||
|
|
@ -373,7 +377,7 @@ class CtxDistillModel(nn.Module):
|
|||
# student tokens
|
||||
_, inp_res_ids = ctx_inp_split(
|
||||
ctx_inp_res_ids,
|
||||
ctx_inp_sep_seq,
|
||||
self.ctx_inp_sep_seq,
|
||||
self.pad_token_id,
|
||||
self.prefix_tokens,
|
||||
padding_side="left",
|
||||
|
|
@ -415,12 +419,18 @@ class CtxDistillModel(nn.Module):
|
|||
# # inp_ids = torch.cat([self.prefix_tokens.expand(bs, -1), inp_ids], dim=-1)
|
||||
# # inp_ids = inp_res_ids[:, :-res_len]
|
||||
# print(self.tokenizer.batch_decode(inp_ids))
|
||||
inp_attention_mask = torch.where(orig_inp_ids != self.pad_token_id, 1, 0).long()
|
||||
model_inputs_kwargs.pop("attention_mask", None)
|
||||
model_inputs_kwargs.pop("input_ids", None)
|
||||
model_outputs = self.student_generate(
|
||||
orig_inp_ids, attention_mask=inp_attention_mask, **model_inputs_kwargs
|
||||
)
|
||||
if self.reprompt_ctx:
|
||||
attention_mask = torch.where(orig_ctx_inp_ids != self.pad_token_id, 1, 0)
|
||||
model_outputs = self.student_generate(
|
||||
orig_ctx_inp_ids, attention_mask=attention_mask, **model_inputs_kwargs
|
||||
)
|
||||
else:
|
||||
attention_mask = torch.where(orig_inp_ids != self.pad_token_id, 1, 0).long()
|
||||
model_outputs = self.student_generate(
|
||||
orig_inp_ids, attention_mask=attention_mask, **model_inputs_kwargs
|
||||
)
|
||||
return model_outputs
|
||||
|
||||
|
||||
|
|
@ -429,6 +439,7 @@ if __name__ == "__main__":
|
|||
from ctx_to_lora.model_loading import get_lora_config, get_model_and_tokenizer
|
||||
|
||||
model_name = "google/gemma-2-2b-it"
|
||||
q_model_name = "google/gemma-3-4b-it"
|
||||
peft_config = get_lora_config(
|
||||
model_name, r=8, target_modules=["down_proj"], lora_dropout=0.0
|
||||
)
|
||||
|
|
@ -439,6 +450,12 @@ if __name__ == "__main__":
|
|||
requires_grad=False,
|
||||
peft_config=peft_config,
|
||||
)
|
||||
q_model, q_tokenizer = get_model_and_tokenizer(
|
||||
q_model_name,
|
||||
train=False,
|
||||
requires_grad=False,
|
||||
peft_config=peft_config,
|
||||
)
|
||||
|
||||
ds = load_and_process_dataset("pwc", split="train", num_proc=8)
|
||||
ctx = ds[0]["context"]
|
||||
|
|
@ -459,26 +476,29 @@ if __name__ == "__main__":
|
|||
prefix_tokens = CTX_AFFIXES[model_name]["prefix"]
|
||||
prefix_tokens = torch.tensor(prefix_tokens, dtype=torch.long)
|
||||
|
||||
cd_model = CtxDistillModel(
|
||||
base_model=model,
|
||||
prefix_tokens=prefix_tokens,
|
||||
pad_token_id=tokenizer.pad_token_id,
|
||||
update_iterations=200,
|
||||
q_model=model,
|
||||
num_gen_q=20,
|
||||
tokenizer=tokenizer,
|
||||
)
|
||||
|
||||
sep_ids = (
|
||||
tokenizer(sep_text.strip("\n"), add_special_tokens=False, return_tensors="pt")
|
||||
.input_ids[0]
|
||||
.to(model.device)
|
||||
)
|
||||
|
||||
cd_model = CtxDistillModel(
|
||||
base_model=model,
|
||||
prefix_tokens=prefix_tokens,
|
||||
ctx_inp_sep_seq=sep_ids,
|
||||
pad_token_id=tokenizer.pad_token_id,
|
||||
update_iterations=200,
|
||||
q_model=q_model,
|
||||
q_tokenizer=q_tokenizer,
|
||||
num_gen_q=20,
|
||||
tokenizer=tokenizer,
|
||||
reprompt_ctx=True,
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
for _ in range(1):
|
||||
base_model_res = model.generate(
|
||||
encoded["input_ids"],
|
||||
input_ids=encoded["input_ids"],
|
||||
attention_mask=encoded["attention_mask"],
|
||||
max_new_tokens=256,
|
||||
do_sample=False,
|
||||
|
|
@ -488,10 +508,8 @@ if __name__ == "__main__":
|
|||
)
|
||||
|
||||
outputs = cd_model.generate(
|
||||
encoded["input_ids"],
|
||||
input_ids=encoded["input_ids"],
|
||||
attention_mask=encoded["attention_mask"],
|
||||
ctx_inp_sep_seq=sep_ids,
|
||||
reset=True,
|
||||
max_new_tokens=256,
|
||||
do_sample=False,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -225,6 +225,7 @@ class HyperLoRA(nn.Module):
|
|||
# or via a perceiver w/ bottleneck size = n_modules * n_layers
|
||||
self.config = config
|
||||
logger.debug(f"HyperLoRA config: {self.config}")
|
||||
self.iterative_mode = False
|
||||
self._init_model()
|
||||
|
||||
def _init_model(self):
|
||||
|
|
|
|||
122
src/ctx_to_lora/modeling/llm_lingua.py
Normal file
122
src/ctx_to_lora/modeling/llm_lingua.py
Normal file
|
|
@ -0,0 +1,122 @@
|
|||
import torch
|
||||
from llmlingua import PromptCompressor
|
||||
from torch import nn
|
||||
|
||||
from ctx_to_lora.data.definitions import CTX_AFFIXES
|
||||
|
||||
|
||||
class LLMLinguaModel(nn.Module):
|
||||
def __init__(self, model, tokenizer, compression_rate):
|
||||
super().__init__()
|
||||
self.base_model = model
|
||||
self.compressor = PromptCompressor(
|
||||
model_name="microsoft/llmlingua-2-xlm-roberta-large-meetingbank",
|
||||
use_llmlingua2=True, # Whether to use llmlingua-2
|
||||
)
|
||||
model_name = self.base_model.name_or_path
|
||||
self.register_buffer("prefix", torch.tensor(CTX_AFFIXES[model_name]["prefix"]))
|
||||
self.register_buffer("suffix", torch.tensor(CTX_AFFIXES[model_name]["suffix"]))
|
||||
self.len_prefix = len(self.prefix)
|
||||
self.len_suffix = len(self.suffix)
|
||||
self.tokenizer = tokenizer
|
||||
self.compression_rate = compression_rate
|
||||
|
||||
@property
|
||||
def generation_config(self):
|
||||
return self.base_model.generation_config
|
||||
|
||||
def compress(self, prompt_txt: str, rate: float):
|
||||
return self.compressor.compress_prompt(
|
||||
prompt_txt, rate=rate, force_tokens=["\n", "?"]
|
||||
)
|
||||
|
||||
def generate(self, *args, **kwargs):
|
||||
# take ctx_ids
|
||||
# strip prefix and suffix
|
||||
# ctx_ids is left padded
|
||||
ctx_ids = kwargs["ctx_ids"][:, self.len_prefix : -self.len_suffix]
|
||||
# decode ctx_ids to ctx_txt
|
||||
ctx_txt = self.tokenizer.batch_decode(ctx_ids)
|
||||
# 4x compression
|
||||
compressed_ctx_txt = self.compress(ctx_txt, rate=self.compression_rate)
|
||||
compressed_ctx_ids = self.tokenizer(
|
||||
compressed_ctx_txt["compressed_prompt"] + "\n\n",
|
||||
return_attention_mask=False,
|
||||
add_special_tokens=False,
|
||||
return_tensors="pt",
|
||||
).to(self.base_model.device)
|
||||
|
||||
bs = ctx_ids.shape[0]
|
||||
ctx_inp_ids = torch.cat(
|
||||
[
|
||||
self.prefix.expand(bs, -1),
|
||||
compressed_ctx_ids["input_ids"],
|
||||
kwargs["input_ids"][:, self.len_prefix :],
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
attn_mask = torch.ones_like(ctx_inp_ids)
|
||||
for k in [
|
||||
"ctx_ids",
|
||||
"ctx_attn_mask",
|
||||
"n_ctx_chunks",
|
||||
"input_ids",
|
||||
"attention_mask",
|
||||
]:
|
||||
kwargs.pop(k, None)
|
||||
return self.base_model.generate(ctx_inp_ids, attention_mask=attn_mask, **kwargs)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from ctx_to_lora.model_loading import get_model_and_tokenizer
|
||||
|
||||
model, tokenizer = get_model_and_tokenizer(
|
||||
"google/gemma-2-2b-it",
|
||||
train=False,
|
||||
requires_grad=False,
|
||||
)
|
||||
|
||||
# Demo: build wrapper, create a toy context + prompt, run compression + generation.
|
||||
device = "cuda"
|
||||
llm = LLMLinguaModel(model, tokenizer).to(device)
|
||||
|
||||
# Toy context and user prompt
|
||||
context_text = (
|
||||
"This is a short illustrative context about large language models and compression. "
|
||||
"They can reduce prompt length while preserving meaning."
|
||||
)
|
||||
user_prompt = (
|
||||
"Summarize the context in one concise sentence." # what we want model to do
|
||||
)
|
||||
|
||||
# Tokenize raw context (core) without special tokens
|
||||
core_ctx_ids = tokenizer.apply_chat_template(
|
||||
[[{"role": "user", "content": context_text}]],
|
||||
tokenize=True,
|
||||
add_generation_prompt=True,
|
||||
return_attention_mask=False,
|
||||
padding=False,
|
||||
truncation=False,
|
||||
return_tensors="pt",
|
||||
add_special_tokens=False,
|
||||
).to(device)
|
||||
|
||||
# Input prompt tokens (what follows the contextual block)
|
||||
input_ids = tokenizer.apply_chat_template(
|
||||
[[{"role": "user", "content": user_prompt}]],
|
||||
tokenize=True,
|
||||
add_generation_prompt=True,
|
||||
return_attention_mask=False,
|
||||
padding=False,
|
||||
truncation=False,
|
||||
return_tensors="pt",
|
||||
add_special_tokens=False,
|
||||
).to(device)
|
||||
|
||||
print("Original context length (chars):", len(context_text))
|
||||
|
||||
# Run generation (may vary depending on model capabilities)
|
||||
output_ids = llm.generate(ctx_ids=core_ctx_ids, input_ids=input_ids)
|
||||
# Decode only the tail beyond supplied input for readability
|
||||
generated_text = tokenizer.decode(output_ids[0], skip_special_tokens=False)
|
||||
print(f"\nFull generated text:\n{generated_text}")
|
||||
Loading…
Add table
Add a link
Reference in a new issue