diff --git a/run_eval.py b/run_eval.py index 0ce8c62..a10a3f0 100644 --- a/run_eval.py +++ b/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)) diff --git a/src/ctx_to_lora/eval_utils.py b/src/ctx_to_lora/eval_utils.py index 0edf9cb..8ddeb56 100644 --- a/src/ctx_to_lora/eval_utils.py +++ b/src/ctx_to_lora/eval_utils.py @@ -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)}") diff --git a/src/ctx_to_lora/modeling/context_distillation.py b/src/ctx_to_lora/modeling/context_distillation.py index fb54373..84a6307 100644 --- a/src/ctx_to_lora/modeling/context_distillation.py +++ b/src/ctx_to_lora/modeling/context_distillation.py @@ -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, ) diff --git a/src/ctx_to_lora/modeling/hypernet.py b/src/ctx_to_lora/modeling/hypernet.py index 4f61e7c..da0aefa 100644 --- a/src/ctx_to_lora/modeling/hypernet.py +++ b/src/ctx_to_lora/modeling/hypernet.py @@ -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): diff --git a/src/ctx_to_lora/modeling/llm_lingua.py b/src/ctx_to_lora/modeling/llm_lingua.py new file mode 100644 index 0000000..f2b33d6 --- /dev/null +++ b/src/ctx_to_lora/modeling/llm_lingua.py @@ -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}")