diff --git a/README.md b/README.md index af81de4..49df497 100644 --- a/README.md +++ b/README.md @@ -51,12 +51,15 @@ run python intx_sft.py configs/...yaml ... --from_pretrained_checkpoint=train_ou LongBench ```bash # generative -run python src/ctx_to_lora/eval.py --checkpoint_path train_outputs/runs/May08_13-56-31_slurm0-a3nodeset-5_59383_906acb28/checkpoint-105000/pytorch_model.bin +run python src/ctx_to_lora/eval.py --checkpoint_path train_outputs/runs/.../pytorch_model.bin --datasets negative_nq triviaqa_retrieved hotpot_qa squad longbench_e - +-split test -# benchmark -cd LongBench/LongBench +# run python src/ctx_to_lora/eval.py --checkpoint_path train_outputs/runs/May08_13-56-31_slurm0-a3nodeset-5_59383_906acb28/checkpoint-105000/pytorch_model.bin -run python pred_ctx_to_lora.py --checkpoint_path ../../train_outputs/runs/Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782/pytorch_model.bin +# # benchmark +# cd LongBench/LongBench -run python eval_ctx_to_lora.py --model_name Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782 --checkpoint_path ../../train_outputs/runs/Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782/pytorch_model.bin +# run python pred_ctx_to_lora.py --checkpoint_path ../../train_outputs/runs/Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782/pytorch_model.bin + +# run python eval_ctx_to_lora.py --model_name Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782 --checkpoint_path ../../train_outputs/runs/Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782/pytorch_model.bin ``` \ No newline at end of file diff --git a/generate_negative_natural_questions.py b/generate_negative_natural_questions.py new file mode 100644 index 0000000..428f141 --- /dev/null +++ b/generate_negative_natural_questions.py @@ -0,0 +1,44 @@ +import os +from tqdm import tqdm +from datasets import load_dataset, Dataset + + +if __name__ == "__main__": + ds = load_dataset("google-research-datasets/natural_questions") + ds = ds.shuffle(seed=42) + for split in ds: + out = [] + for i, sample in enumerate(tqdm(ds[split])): + ctx = " ".join( + [ + token + for is_html, token in zip( + sample["document"]["tokens"]["is_html"], + sample["document"]["tokens"]["token"], + ) + if not is_html + ] + ).strip() + if not ctx: + continue + answers = [] + for answer in sample["annotations"]["short_answers"]: + answers += answer["text"] + if not answers: + continue + answer = answers[0].strip() + if not answer: + continue + prompt = sample["question"]["text"].capitalize().strip() + " ?" + out.append(dict(prompt=prompt, context=ctx, answer=answer.capitalize())) + # if i >= 10: + # break + shifted_ctxs = [d["context"] for d in out] + shifted_ctxs = shifted_ctxs[1:] + shifted_ctxs[:1] + for d, shifted_ctx in zip(out, shifted_ctxs): + d["context"] = shifted_ctx + new_ds = Dataset.from_list(out) + print(new_ds) + save_dir = "data/raw_datasets/negative_natural_questions" + os.makedirs(save_dir, exist_ok=True) + new_ds.to_json(f"{save_dir}/{split}.jsonl") diff --git a/src/ctx_to_lora/data_definitions.py b/src/ctx_to_lora/data_definitions.py index 40d67fc..0efc52c 100644 --- a/src/ctx_to_lora/data_definitions.py +++ b/src/ctx_to_lora/data_definitions.py @@ -299,8 +299,14 @@ DS_KWARGS = { split="train", ), ), - # TODO: for negative training and evaluating - "natural_questions": ..., + # for negative training and evaluating + "negative_nq": dict( + test=dict( + path="json", + data_files="data/raw_datasets/negative_natural_questions/validation.jsonl", + split="train", + ), + ), } # LongBench kwargs @@ -341,6 +347,7 @@ CLOSED_QA_INTX_TEMPLATES = [ EVAL_INTX_TEMPLATES = { + "negative_nq": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}", "triviaqa_retrieved": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}", "hotpot_qa": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}", "squad": "Answer the following question. Output only the answer and do not output any other words.\n\nQuestion: {input}", diff --git a/src/ctx_to_lora/data_utils.py b/src/ctx_to_lora/data_utils.py index 2806fa2..38deb21 100644 --- a/src/ctx_to_lora/data_utils.py +++ b/src/ctx_to_lora/data_utils.py @@ -192,25 +192,19 @@ def get_preprocessing_fn( "response": sample["answers"][0], } - elif "natural_questions" in ds_name: - pass - # ctx = " ".join( - # [ - # token - # for is_html, token in zip( - # ds[0]["document"]["tokens"]["is_html"], - # ds[0]["document"]["tokens"]["token"], - # ) - # if not is_html - # ] - # ) - elif ds_name == "triviaqa_retrieved": - # maybe needed for training (negative sample for training) - # def f(sample): - # ctx = sample["entity_page"]["wiki_context"] - # if not ctx: - # return None + elif ds_name == "negative_nq": + def f(sample): + q = sample["prompt"] + prompt = closed_qa_prompting(q) if not is_eval else q + return { + "context": sample["context"], + "prompt": prompt, + "response": sample["answer"], + } + + elif ds_name == "triviaqa_retrieved": + # only used for eval def f(sample): return { "context": sample["context"], diff --git a/src/ctx_to_lora/eval.py b/src/ctx_to_lora/eval.py index 50bb7e0..1290cd3 100644 --- a/src/ctx_to_lora/eval.py +++ b/src/ctx_to_lora/eval.py @@ -95,6 +95,7 @@ CLOSED_QA_DATASETS = { "longbench/musique", "squad", "triviaqa_retrieved", + "negative_nq", } for ds_name in list(CLOSED_QA_DATASETS): @@ -452,7 +453,13 @@ def evaluate( # if generative: # model = model.to(torch.bfloat16) else: - model = get_model(model_name_or_path, train=False, requires_grad=False) + model_kwargs = dict(attn_implementation="flash_attention_2") + model = get_model( + model_name_or_path, + train=False, + requires_grad=False, + model_kwargs=model_kwargs, + ) # NOTE: there is still some randomness in the eval result # despite all the deterministic settings if is_liger_kernel_available(): @@ -681,7 +688,7 @@ if __name__ == "__main__": args.output_dir = f"{run_dir}/eval-results-{cur_it}" args.logging_dir = f"{run_dir}/eval-results-{cur_it}" args.run_name = run_dir.split("/")[-1] - args.remove_context = True + args.remove_context = False # modulated model doesn't see ctx by default else: args = Namespace( model_name_or_path=cli_args.model_name_or_path, diff --git a/src/ctx_to_lora/training_utils.py b/src/ctx_to_lora/training_utils.py index 3e4ca2a..b76858d 100644 --- a/src/ctx_to_lora/training_utils.py +++ b/src/ctx_to_lora/training_utils.py @@ -124,7 +124,7 @@ def train_model( compute_metrics=compute_metrics, ) if is_modulated_model: - logger.info(f"Training with modulated model. Using CustomTrainer.") + logger.info("Training with modulated model. Using CustomTrainer.") trainer_kwargs["gen_lora_l1_reg_coef"] = training_args.gen_lora_l1_reg_coef del training_args.gen_lora_l1_reg_coef diff --git a/webui/app.py b/webui/app.py index 4ccb8c3..22052d6 100644 --- a/webui/app.py +++ b/webui/app.py @@ -116,51 +116,134 @@ def get_run_data(run_path): return grouped_data -def get_generated_text_data(run_path): +def get_generated_text_data(run_path: str) -> dict: """ - Loads generated text data from all *_generated_text.jsonl files. + Loads generated text data from all *_generated_text.jsonl files from the current run's output. + It specifically excludes any files ending with _no_context_generated_text.jsonl, + as those are handled by get_base_model_generated_text. Also recursively searches for generated text files in subfolders. Args: - run_path: Path to the directory containing the .jsonl files. + run_path: Path to the directory containing the .jsonl files for the current run. Returns: - A dictionary containing the generated text data for each split found. + A dictionary containing the generated text data for each primary split found. """ generated_data = {} - files = sorted(glob(f"{run_path}/*_generated_text.jsonl")) - print(files) - # Find all *_generated_text.jsonl files in the current directory + # Find all *_generated_text.jsonl files in the current directory, excluding _no_context_ variants + files = sorted( + [ + f + for f in glob(os.path.join(run_path, "*_generated_text.jsonl")) + if "_no_context_generated_text.jsonl" not in os.path.basename(f) + ] + ) + print(f"Processing generated text files from {run_path}: {files}") + for filename in files: - # Extract split name from filename (remove _generated_text.jsonl) - split = filename.split("/")[-1].split("_generated_text.jsonl")[0] + split = os.path.basename(filename).replace("_generated_text.jsonl", "") try: with open(filename) as f: - lines = f.readlines() + lines = f.readlines(100_000) data = [json.loads(line) for line in lines] generated_data[split] = data except FileNotFoundError: generated_data[split] = None + except json.JSONDecodeError as e: + print(f"Error decoding JSON from {filename}: {e}") + generated_data[split] = None # Recursively search for *_generated_text.jsonl files in subdirectories - for subdir in [d for d in glob(os.path.join(run_path, "*")) if os.path.isdir(d)]: - subdir_name = os.path.basename(subdir) - subdir_files = sorted(glob(f"{subdir}/*_generated_text.jsonl")) + for subdir_path in [ + d for d in glob(os.path.join(run_path, "*")) if os.path.isdir(d) + ]: + subdir_name = os.path.basename(subdir_path) + subdir_files = sorted( + [ + f + for f in glob(os.path.join(subdir_path, "*_generated_text.jsonl")) + if "_no_context_generated_text.jsonl" not in os.path.basename(f) + ] + ) + print( + f"Processing generated text files from subdirectory {subdir_path}: {subdir_files}" + ) for filename in subdir_files: - # Extract split name from filename and include subfolder name - base_split = filename.split("/")[-1].split("_generated_text.jsonl")[0] + base_split = os.path.basename(filename).replace("_generated_text.jsonl", "") + # Create a unique split key including the subdirectory name split = f"{subdir_name}_{base_split}" try: with open(filename) as f: - lines = f.readlines() + lines = f.readlines(100_000) data = [json.loads(line) for line in lines] generated_data[split] = data except FileNotFoundError: generated_data[split] = None - + except json.JSONDecodeError as e: + print(f"Error decoding JSON from {filename}: {e}") + generated_data[split] = None return generated_data +def get_base_model_generated_text(model_name, generated_data): + """ + Loads generated text data from the base model for comparison. + Also loads "no context" base model data if available. + + Args: + model_name: The name of the base model + generated_data: Dictionary of generated data from the fine-tuned model + + Returns: + Tuple: (base_model_data, base_model_no_context_data) + Dictionaries containing base model generated texts for matching splits + """ + if not model_name: + return {}, {} + + base_model_data = {} + base_model_no_context_data = {} + + # # Create normalized model name for directory lookup + # normalized_model_name = model_name.replace("/", "_") + + for split in generated_data: + # Skip if no data for this split + if not generated_data[split]: + base_model_data[split] = None + base_model_no_context_data[split] = None + continue + + # Construct path to the base model's output for this split + base_model_path = f"eval_results/{model_name}/{split}_generated_text.jsonl" + base_model_no_context_path = ( + f"eval_results/{model_name}/{split}_no_context_generated_text.jsonl" + ) + + try: + with open(base_model_path) as f: + lines = f.readlines(100_000) + data = [json.loads(line) for line in lines] + base_model_data[split] = data + except FileNotFoundError: + print(f"Base model output not found for {split} at {base_model_path}") + # Store None to indicate we tried but didn't find matching data + base_model_data[split] = None + + try: + with open(base_model_no_context_path) as f: + lines = f.readlines(100_000) + data = [json.loads(line) for line in lines] + base_model_no_context_data[split] = data + except FileNotFoundError: + print( + f"Base model (no context) output not found for {split} at {base_model_no_context_path}" + ) + base_model_no_context_data[split] = None + + return base_model_data, base_model_no_context_data + + def get_available_checkpoints(run): """ Finds all checkpoint directories in a run folder. @@ -292,6 +375,31 @@ def visualize(run): except FileNotFoundError: model_name = None + # If config.yaml doesn't have the model name, try args.yaml + if not model_name: + args_path = os.path.join(logdir, "args.yaml") + try: + with open(args_path) as f: + yaml.add_constructor("!", lambda loader, node: None) + yaml.add_multi_constructor( + "tag:yaml.org,2002:python/object", + lambda loader, suffix, node: None, + Loader=yaml.SafeLoader, + ) + args = yaml.safe_load(f) + model_name = args.get("model_name_or_path") + except FileNotFoundError: + model_name = None + + # Get base model outputs if we have the model name and modulated outputs + base_model_data = {} + base_model_no_context_data = {} + print(generated_data.keys()) + if model_name and generated_data: + base_model_data, base_model_no_context_data = get_base_model_generated_text( + model_name, generated_data + ) + return render_template( "visualize.html", run=run, @@ -299,6 +407,8 @@ def visualize(run): selected_eval_folder=selected_eval_folder, data=data, generated_data=generated_data, + base_model_data=base_model_data, + base_model_no_context_data=base_model_no_context_data, model_name=model_name, checkpoints=checkpoint_names, ) diff --git a/webui/templates/visualize.html b/webui/templates/visualize.html index 67fd98a..a81623e 100644 --- a/webui/templates/visualize.html +++ b/webui/templates/visualize.html @@ -452,6 +452,43 @@ background-color: #bbdefb; cursor: not-allowed; } + + .model-comparison { + display: flex; + flex-wrap: wrap; + gap: 20px; + margin-bottom: 15px; + } + + .model-output { + flex: 1; + min-width: 300px; + border: 1px solid #e0e0e0; + border-radius: 5px; + padding: 10px; + background-color: #f9f9f9; + } + + .model-output p { + margin: 0; + } + + .model-output span { + display: block; + padding: 10px; + background-color: white; + border: 1px solid #eee; + border-radius: 4px; + min-height: 80px; + max-height: 150px; + overflow-y: auto; + white-space: pre-wrap; + margin-top: 5px; + } + + .highlight-diff { + background-color: #fff8e1; + } {% endblock %} @@ -536,7 +573,22 @@ Context:
Input:
-Generated:
+ + +Modulated Model:
+Base Model ({{ model_name }}): Not + available
+Base Model (No Context - {{ model_name }}): Not available
+Label: