add llmlingua + different q model for cd

This commit is contained in:
51616 2025-09-13 23:52:07 +09:00
parent 0e05644d6a
commit 0e46a7c568
5 changed files with 265 additions and 65 deletions

View file

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

View file

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

View file

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

View file

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

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