diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..3da31c6 --- /dev/null +++ b/.gitignore @@ -0,0 +1,148 @@ +# Byte-compiled / optimized / DLL files +__pycache__/ +*.py[cod] +*$py.class + +# C extensions +*.so + +# Distribution / packaging +.Python +build/ +develop-eggs/ +dist/ +downloads/ +eggs/ +.eggs/ +lib/ +lib64/ +parts/ +sdist/ +var/ +wheels/ +pip-wheel-metadata/ +share/python-wheels/ +*.egg-info/ +.installed.cfg +*.egg +MANIFEST + +# PyInstaller +# Usually these files are written by a python script from a template +# before PyInstaller builds the exe, so as to inject date/other infos into it. +*.manifest +*.spec + +# Installer logs +pip-log.txt +pip-delete-this-directory.txt + +# Unit test / coverage reports +htmlcov/ +.tox/ +.nox/ +.coverage +.coverage.* +.cache +nosetests.xml +coverage.xml +*.cover +*.py,cover +.hypothesis/ +.pytest_cache/ +**/cache/ +**/output/ +**/cache.pkl + +# Translations +*.mo +*.pot + +# Django stuff: +*.log +local_settings.py +db.sqlite3 +db.sqlite3-journal + +# Flask stuff: +instance/ +.webassets-cache + +# Scrapy stuff: +.scrapy + +# Sphinx documentation +docs/_build/ + +# PyBuilder +target/ + +# Jupyter Notebook +.ipynb_checkpoints + +# IPython +profile_default/ +ipython_config.py + +# pyenv +.python-version + +# pipenv +# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control. +# However, in case of collaboration, if having platform-specific dependencies or dependencies +# having no cross-platform support, pipenv may install dependencies that don't work, or not +# install all needed dependencies. +#Pipfile.lock + +# PEP 582; used by e.g. github.com/David-OConnor/pyflow +__pypackages__/ + +# Celery stuff +celerybeat-schedule +celerybeat.pid + +# SageMath parsed files +*.sage.py + +# Environments +.env +.venv +env/ +venv/ +ENV/ +env.bak/ +venv.bak/ + +# Spyder project settings +.spyderproject +.spyproject + +# Rope project settings +.ropeproject + +# mkdocs documentation +/site + +# mypy +.mypy_cache/ +.dmypy.json +dmypy.json + +# Pyre type checker +.pyre/ +.idea +volumes + +# macOS +.DS_Store + +*.npz +*.csv +*.json +*.png +*.jpg +*.txt +*.arrow +*.pdf +*.md +*.prof diff --git a/.whitesource b/.whitesource new file mode 100644 index 0000000..26e9c47 --- /dev/null +++ b/.whitesource @@ -0,0 +1,3 @@ +{ + "settingsInheritedFrom": "ibm-mend-config/mend-config@main" +} \ No newline at end of file diff --git a/bench_latency_vidore2.py b/bench_latency_vidore2.py new file mode 100644 index 0000000..b68c25f --- /dev/null +++ b/bench_latency_vidore2.py @@ -0,0 +1,203 @@ +import argparse +import json +import os +import time +import statistics as stats +from functools import partial + +import torch + +from dataset_configs import Datasets, RagDataset, DataSplit +from embedding_configs import Embedders, all_embedders +from fusion_methods import ( + reciprocal_rank_fusion, + average_ranking_fusion, + sim_score_fusion, + normalize_min_max, + normalize_softmax, +) +from query_optimizations import scores_feedback, all_optimization_funcs +from retriever import Retriever +from utils import get_device, set_seed, on_ccc +from query_optimizations import OptimizationFunctions, kl_divergence + + +def sync_accel(): + if torch.cuda.is_available(): + torch.cuda.synchronize() + if torch.backends.mps.is_available(): + torch.mps.synchronize() + + +class SimpleRunner: + def __init__(self, values=5, warmups=1): + self.values = values + self.warmups = warmups + self.results = [] + + def bench(self, name, fn, *, inner_loops=1): + for _ in range(self.warmups): + fn() + sync_accel() + samples = [] + for _ in range(self.values): + t0 = time.perf_counter() + fn() + sync_accel() + dt = time.perf_counter() - t0 + samples.append(dt / inner_loops) + r = { + "name": name, + "values": samples, + "mean": stats.fmean(samples) if samples else None, + "median": stats.median(samples) if samples else None, + "min": min(samples) if samples else None, + "stdev": stats.pstdev(samples) if len(samples) > 1 else 0.0, + } + self.results.append(r) + return r + + def dump(self, path, extra_meta=None): + out = {"benchmarks": self.results, "metadata": extra_meta or {}} + with open(path, "w") as f: + json.dump(out, f, indent=2) + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--embedders", nargs="+", choices=all_embedders, + default=[Embedders.colnomic, Embedders.linq, Embedders.qwen_text]) + ap.add_argument("--warmups", type=int, default=1) + ap.add_argument("--values", type=int, default=3) + ap.add_argument("--opt_func", nargs="+", choices=all_optimization_funcs, + default=[OptimizationFunctions.union_with_search.name]) + ap.add_argument("--opt_steps", nargs="+", default=[10, 25, 50]) + ap.add_argument("--seed", type=int, default=42) + ap.add_argument("--out_dir", default="latency_results_vidore2") + args = ap.parse_args() + + device = get_device() + set_seed(args.seed) + prefix = "/proj/omri/" if on_ccc() else "" + dataset = RagDataset(Datasets.economics_reports_v2, prefix=prefix) + + retrievers = [Retriever(m) for m in args.embedders] + for r in retrievers: + r.load_embs(dataset) + + os.makedirs(args.out_dir, exist_ok=True) + runner = SimpleRunner(values=args.values, warmups=args.warmups) + + meta = { + "device": str(device), + "torch": torch.__version__, + "cuda": getattr(torch.version, "cuda", None), + "mps": torch.backends.mps.is_available(), + "dataset": dataset.id, + "embedders": " ".join([ret.id for ret in retrievers]), + "values": args.values, + "warmups": args.warmups, + } + + results = [] + + # per-model retrieval timing + for i in range(len(retrievers) - 1): + mod1 = retrievers[i] + res = runner.bench( + f"run_retrieval[{dataset.id}][{mod1.id}]", + partial(mod1.run_retrieval, dataset) + ) + results.append(res) + r1, top_idx_1 = mod1.run_retrieval(dataset) + + for j in range(i + 1, len(retrievers)): + mod2 = retrievers[j] + + res = runner.bench( + f"run_retrieval[{dataset.id}][{mod2.id}]", + partial(mod2.run_retrieval, dataset) + ) + results.append(res) + r2, top_idx_2 = mod2.run_retrieval(dataset) + + if not mod1.is_sparse: + runner.bench( + f"fusion[scores_feedback][{dataset.id}][{mod1.id}<-{mod2.id}]", + partial(scores_feedback, + main_model=mod1, feedback_model=mod2, dataset=dataset, top_k_idxs=top_idx_1) + ) + if not mod2.is_sparse: + runner.bench( + f"fusion[scores_feedback][{dataset.id}][{mod2.id}<-{mod1.id}]", + partial(scores_feedback, + main_model=mod2, feedback_model=mod1, dataset=dataset, top_k_idxs=top_idx_2) + ) + + runner.bench( + f"fusion[RRF][{dataset.id}][{mod1.id}+{mod2.id}]", + partial(reciprocal_rank_fusion, r1=r1, r2=r2) + ) + runner.bench( + f"fusion[average][{dataset.id}][{mod1.id}+{mod2.id}]", + partial(average_ranking_fusion, r1=r1, r2=r2) + ) + + runner.bench( + f"fusion[sim_score_minmax][{dataset.id}][{mod1.id}+{mod2.id}]", + partial(sim_score_fusion, r1=r1, r2=r2, normal_func=normalize_min_max) + ) + runner.bench( + f"fusion[sim_score_softmax][{dataset.id}][{mod1.id}+{mod2.id}]", + partial(sim_score_fusion, r1=r1, r2=r2, normal_func=normalize_softmax) + ) + for opt_func in args.opt_func: + for n_steps in args.opt_steps: + if not mod1.is_sparse: + runner.bench( + f"query_opt[{opt_func} n={n_steps}]" + f"[{dataset.id}][{mod1.id}<-{mod2.id}]", + partial(OptimizationFunctions[opt_func].value, + main_model=mod1, feedback_model=mod2, dataset=dataset, + device=device, + k=10, + lr=0.001, + n_steps=n_steps, + T=1, + mixture_alpha=0.5, + loss_func=kl_divergence, + optimizer=torch.optim.Adam, + split=DataSplit.TEST + ) + ) + + if not mod2.is_sparse: + runner.bench( + f"query_opt[{opt_func} n={n_steps}]" + f"[{dataset.id}][{mod1.id}<-{mod2.id}]", + partial(OptimizationFunctions[opt_func].value, + main_model=mod2, feedback_model=mod1, dataset=dataset, + device=device, + k=10, + lr=0.001, + n_steps=n_steps, + T=1, + mixture_alpha=0.5, + loss_func=kl_divergence, + optimizer=torch.optim.Adam, + split=DataSplit.TEST + ) + ) + + for r in retrievers: + r.clean_embs(dataset) + if torch.cuda.is_available(): torch.cuda.empty_cache() + sync_accel() + + out_json = os.path.join(args.out_dir, "suite.json") + runner.dump(out_json, extra_meta=meta) + print(f"Wrote {out_json}") + + +if __name__ == "__main__": + main() diff --git a/calc_embeddings.py b/calc_embeddings.py new file mode 100644 index 0000000..e5d75af --- /dev/null +++ b/calc_embeddings.py @@ -0,0 +1,57 @@ +import argparse +import gc +import io + +import torch +from natsort import natsorted, ns + +from dataset_configs import RagDataset +from embedding_configs import Embedders, Modality, EncoderConfig, all_embedders + + +def load_documents(path): + docs = [] + for child in natsorted(path.iterdir(), alg=ns.IGNORECASE): + if not child.is_file(): + continue + file_bytes = child.read_bytes() + stream = io.BytesIO(file_bytes) + docs.append(stream) + return docs + + +def main(args): + configs: list[EncoderConfig] = [getattr(Embedders, model) for model in args.models] + + with torch.no_grad(): + for config in configs: + embedder = config.embedder() + for dataset_name in args.datasets: + dataset = RagDataset(dataset_name, prefix=args.datasets_path_prefix) + print(f"Running for {config.model_id}") + if config.modality == Modality.VISION: + folder = dataset.path / "images" + img_docs = natsorted(folder.iterdir(), alg=ns.IGNORECASE) + docs = [f for f in img_docs if f.is_file() and not f.name == ".DS_Store"] + else: + folder = dataset.path / "texts" + text_docs = load_documents(folder) + docs = [c.read().decode() for c in text_docs] + benchmark = dataset.get_benchmark_obj() + queries = [q for q in benchmark['question']] + embedder.calc_and_save_embeddings(dataset.path, docs, queries, config.model_id, + custom_out_path=args.out_path) + gc.collect() + torch.cuda.empty_cache() + + +if __name__ == "__main__": + from dataset_configs import REAL_MM_RAG_DATASETS, VIDORE1_DATASETS, VIDORE2_DATASETS + parser = argparse.ArgumentParser() + parser.add_argument('--datasets', nargs='+', default=VIDORE2_DATASETS) + parser.add_argument('--models', nargs='+', required=True, choices=all_embedders) + parser.add_argument('--datasets_path_prefix', default="/proj/omri/") + parser.add_argument('--out_path') + + main_args = parser.parse_args() + main(main_args) diff --git a/dataset_configs.py b/dataset_configs.py new file mode 100644 index 0000000..86ae623 --- /dev/null +++ b/dataset_configs.py @@ -0,0 +1,117 @@ +import ast +from enum import Enum +from pathlib import Path + +import numpy as np +import pandas as pd +import pytrec_eval +import torch + + +class DataSplit(Enum): + TEST = "test" + DEV = "dev" + + +class Datasets: + arxivqa = "Vidore1/arxivqa_test_subsampled_beir" + docvqa = "Vidore1/docvqa_test_subsampled_beir" + infovqa = "Vidore1/infovqa_test_subsampled_beir" + tabfquad = "Vidore1/tabfquad_test_subsampled_beir" + tatdqa = "Vidore1/tatdqa_test_beir" + shiftproject = "Vidore1/shiftproject_test_beir" + artificial_intelligence = "Vidore1/syntheticDocQA_artificial_intelligence_test_beir" + energy_test = "Vidore1/syntheticDocQA_energy_test_beir" + government_reports = "Vidore1/syntheticDocQA_government_reports_test_beir" + healthcare = "Vidore1/syntheticDocQA_healthcare_industry_test_beir" + + esg_reports_v2 = 'Vidore2/esg_reports_v2' + biomedical_lectures_v2 = 'Vidore2/biomedical_lectures_v2' + economics_reports_v2 = 'Vidore2/economics_reports_v2' + esg_reports_human_labeled_v2 = 'Vidore2/esg_reports_human_labeled_v2' + + FinReport = 'REAL-MM-RAG/FinReport' + FinSlides = 'REAL-MM-RAG/FinSlides' + TechReport = 'REAL-MM-RAG/TechReport' + TechSlides = 'REAL-MM-RAG/TechSlides/' + + +VIDORE1_DATASETS = [ + Datasets.arxivqa, + Datasets.docvqa, + Datasets.infovqa, + Datasets.tabfquad, + Datasets.tatdqa, + Datasets.shiftproject, + Datasets.artificial_intelligence, + Datasets.energy_test, + Datasets.government_reports, + Datasets.healthcare +] + + +VIDORE2_DATASETS = [ + Datasets.esg_reports_v2, + Datasets.biomedical_lectures_v2, + Datasets.economics_reports_v2, + Datasets.esg_reports_human_labeled_v2 +] + + +REAL_MM_RAG_DATASETS = [ + Datasets.FinReport, + Datasets.FinSlides, + Datasets.TechReport, + Datasets.TechSlides +] + + +class RagDataset: + + def __init__(self, path, prefix, dev_size=0.1): + self.path = Path(prefix + path) + self.id = self.path.stem + self.dev_size = dev_size + + self.idx = None + + self.benchmark = self.get_benchmark_obj() + self.num_queries = len(self.benchmark) + + def get_benchmark_obj(self) -> pd.DataFrame: + csv_fp = self.path / "images" / "benchmark" / "benchmark.csv" + df = pd.read_csv( + csv_fp, + converters={"correct_answer_document_ids": ast.literal_eval}, + ) + return df + + def get_queries_indices(self, split=DataSplit.TEST): + if self.idx is None: + self.idx = torch.randperm(self.num_queries) + n = int(self.num_queries * self.dev_size) + return self.idx[n:] if split == DataSplit.TEST else self.idx[:n] + + def evaluate(self, results, return_raw_results=True): + b = self.benchmark[self.benchmark['question'].isin(list(results.keys()))] + + qs = list(b['question']) + + def _qrels_entry(rels): + return {str(d): v for d, v in rels.items()} + + qrels = {str(q): _qrels_entry(rels) for q, rels in zip(b['question'], b['correct_answer_document_ids'])} + run = {str(q): {str(d): float(s) for d, s in results[q]} for q in qs} + ks = (1, 5, 10, 100) + metrics = {f'ndcg_cut.{",".join(map(str, ks))}', f'recall.{",".join(map(str, ks))}'} + scores = pytrec_eval.RelevanceEvaluator(qrels, metrics).evaluate(run) + + out = {} + for k in ks: + ndcg = [scores.get(str(q), {}).get(f'ndcg_cut_{k}', np.nan) for q in qs] + rec = [scores.get(str(q), {}).get(f'recall_{k}', np.nan) for q in qs] + out[f'ndcg@{k}'] = round(float(np.nanmean(ndcg)), 3) + out[f'recall@{k}'] = round(float(np.nanmean(rec)), 3) + if return_raw_results: + out["raw_results"] = run + return out diff --git a/embedding_configs.py b/embedding_configs.py new file mode 100644 index 0000000..6c9af45 --- /dev/null +++ b/embedding_configs.py @@ -0,0 +1,332 @@ +import importlib +import math +from dataclasses import dataclass +from enum import Enum +from pathlib import Path + +import numpy as np +import torch +from fastembed import SparseTextEmbedding +from PIL import Image +from sentence_transformers import SentenceTransformer +from tqdm import tqdm +from transformers import AutoModel +from transformers.utils.import_utils import is_flash_attn_2_available + +from utils import get_device, batch + + +class Embedder: + max_length: int = None + batch_size: int = None + dtype: np.floating = None + is_multi_vector: bool = False + is_sparse: bool = False + file_suffix = "_embeddings" + + def calc_and_save_embeddings(self, dataset_path: Path, docs: list[str], queries: list[str], model_id: str, + custom_out_path=None): + model_name = model_id.split("/")[1] + if custom_out_path: + out_path = Path(custom_out_path) / dataset_path.as_posix().split("/")[-1] + else: + out_path = dataset_path + out_path = out_path / (model_name + self.file_suffix) + out_path.mkdir(parents=True, exist_ok=True) + + docs_file = out_path / "docs.npz" + queries_file = out_path / "queries.npz" + + # TODO check if the embeddings are already there + print(f"Encoding documents - {model_id}") + docs_embeddings, query_embeddings = self.calc_embeddings(model_id, docs, queries) + print(f"Writing {len(docs_embeddings)} encoded documents and {len(query_embeddings)} encoded queries to disk") + np.savez_compressed(queries_file, query_embeddings) + np.savez_compressed(docs_file, docs_embeddings) + + def calc_embeddings(self, model_id: str, docs: list[str], queries: list[str]): + raise NotImplementedError() + + +class JinaEmbedder(Embedder): + dtype = np.float32 + + def calc_embeddings(self, model_id, docs, queries): + m = importlib.import_module("transformers.modeling_flash_attention_utils") + + try: + from flash_attn.layers.rotary import apply_rotary_emb as _apply_rotary_emb + except Exception: + _apply_rotary_emb = None + try: + from flash_attn import flash_attn_varlen_func as _favf + except Exception: + _favf = None + + setattr(m, "apply_rotary_emb", getattr(m, "apply_rotary_emb", _apply_rotary_emb)) + setattr(m, "flash_attn_varlen_func", getattr(m, "flash_attn_varlen_func", _favf)) + + model = AutoModel.from_pretrained(model_id, trust_remote_code=True) + model.to(get_device()) + + with torch.no_grad(): + query_embeddings = model.encode_text( + texts=queries, + task="retrieval", + prompt_name="query", + return_multivector=self.is_multi_vector, + ) + document_embeddings = self._encode_documents(model, docs) + query_embeddings = self._convert_to_output_format(query_embeddings) + document_embeddings = self._convert_to_output_format(document_embeddings) + + return document_embeddings, query_embeddings + + def _convert_to_output_format(self, embeddings): + if self.is_multi_vector: + if len(set(e.shape for e in embeddings)) > 1: # embedding matrices have different shapes + return np.array([e.cpu().numpy().astype(self.dtype) for e in embeddings], dtype=object) + else: # all embedding matrices have the same shape + return np.array([e.cpu().numpy().astype(self.dtype) for e in embeddings]) + else: + return [e.cpu() for e in embeddings] + + def _encode_documents(self, model, docs): + raise NotImplementedError() + + +class JinaImageEmbedder(JinaEmbedder): + def _encode_documents(self, model, docs): + images = [Image.open(path) for path in docs] + return model.encode_image( + images=images, + task="retrieval", + return_multivector=self.is_multi_vector, + ) + + +class JinaTextEmbedder(JinaEmbedder): + file_suffix = "_text_embeddings" + + def _encode_documents(self, model, docs): + return model.encode_text( + texts=docs, + task="retrieval", + prompt_name="passage", + return_multivector=self.is_multi_vector, + ) + + +class JinaImageMultiEmbedder(JinaImageEmbedder): + is_multi_vector = True + file_suffix = "_multi_embeddings" + + +class JinaTextMultiEmbedder(JinaTextEmbedder): + is_multi_vector = True + file_suffix = "_multi_txt_embeddings" + + +class SentenceTransformersEmbedder(Embedder): + query_extra_kwargs: dict + doc_extra_kwargs: dict + + def calc_embeddings(self, model_id, docs, queries): + device = get_device() + model = SentenceTransformer(model_id) + model.max_seq_length = self.max_length + model.eval().half().to(device) + + query_embeddings = self._encode(queries, model, device, extra_kwargs=self.query_extra_kwargs) + document_embeddings = self._encode(docs, model, device, extra_kwargs=self.doc_extra_kwargs) + del model + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return document_embeddings, query_embeddings + + def _encode(self, texts: list[str], model: SentenceTransformer, device: str, extra_kwargs: dict = None): + extra_kwargs = extra_kwargs if extra_kwargs else {} + out = [] + for batch_texts in tqdm(batch(texts, self.batch_size), + desc=f"Calculating embeddings (Batch size={self.batch_size})", + total=math.ceil(len(texts) / self.batch_size)): + with torch.inference_mode(): + embs = model.encode( + batch_texts, + batch_size=len(batch_texts), + convert_to_numpy=True, + show_progress_bar=False, + device=device, + **extra_kwargs, + ).astype(self.dtype, copy=False) + out.append(embs) + if torch.cuda.is_available(): + torch.cuda.empty_cache() + return np.concatenate(out, axis=0) + + +class QwenEmbedder(SentenceTransformersEmbedder): + max_length = 512 + batch_size = 8 + dtype = np.float32 + + query_extra_kwargs = {"prompt_name": "query"} + doc_extra_kwargs = {"prompt_name": None} + + +class LinqMistralEmbedder(SentenceTransformersEmbedder): + max_length = 512 + batch_size = 8 + dtype = np.float32 + + task = "Given a question, retrieve Wikipedia passages that answer the question" + prompt = f"Instruct: {task}\nQuery: " + + query_extra_kwargs = {"prompt": prompt} + doc_extra_kwargs = {"prompt": None} + + +class BatchedMultiEmbedder(Embedder): + image_batch_size: int + text_batch_size: int + is_multi_vector = True + file_suffix = "_multi_embeddings" + + @staticmethod + def to_list(x): + return list(x) if isinstance(x, (list, tuple)) else [x] + + @staticmethod + def pad_for_concat(arrs, pad_axis=1): + if not arrs: + return arrs + target = max(a.shape[pad_axis] for a in arrs) + out = [] + for a in arrs: + if a.shape[pad_axis] == target: + out.append(a) + else: + pad = [(0, 0)] * a.ndim + pad[pad_axis] = (0, target - a.shape[pad_axis]) + out.append(np.pad(a, pad, mode="constant")) + return out + + +class NvidiaEmbedder(BatchedMultiEmbedder): + image_batch_size = 64 + text_batch_size = 256 + + def calc_embeddings(self, model_id, docs, queries): + model = AutoModel.from_pretrained( + model_id, + device_map='cuda', + trust_remote_code=True, + torch_dtype=torch.bfloat16, + attn_implementation="flash_attention_2" if is_flash_attn_2_available() else None, + revision='50c36f4d5271c6851aa08bd26d69f6e7ca8b870c' # TODO what is this? + ).eval() + + img_embs = [] + with torch.no_grad(): + for batch_paths in batch(docs, self.image_batch_size): + imgs = [Image.open(p) for p in batch_paths] + emb = model.forward_passages(imgs, batch_size=len(imgs)) + img_embs.extend([t.detach().to(torch.float32).cpu().numpy() for t in self.to_list(emb)]) + for im in imgs: + im.close() + + q_embs = [] + with torch.no_grad(): + for batch_q in batch(queries, self.text_batch_size): + emb = model.forward_queries(batch_q, batch_size=len(batch_q)) + q_embs.extend([t.detach().to(torch.float32).cpu().numpy() for t in self.to_list(emb)]) + + print("concatenate imgs") + document_embeddings = np.concatenate(self.pad_for_concat(img_embs, pad_axis=1), axis=0) + print("concatenate queries") + query_embeddings = np.concatenate(self.pad_for_concat(q_embs, pad_axis=1), axis=0) + + return document_embeddings, query_embeddings + + +class NomicEmbedder(BatchedMultiEmbedder): + image_batch_size = 32 + text_batch_size = 256 + + def calc_embeddings(self, model_id, docs, queries): + from colpali_engine.models import ColQwen2_5, ColQwen2_5_Processor + model = ColQwen2_5.from_pretrained( + model_id, + torch_dtype=torch.bfloat16, + device_map="cuda:0", + attn_implementation="flash_attention_2" if is_flash_attn_2_available() else None, + ).eval() + processor = ColQwen2_5_Processor.from_pretrained(model_id) + + img_embs = [] + with torch.no_grad(): + for batch_paths in batch(docs, self.image_batch_size): + imgs = [Image.open(p) for p in batch_paths] + processed_images = processor.process_images(imgs).to(model.device) + emb = model(**processed_images) + img_embs.extend([t.detach().to(torch.float32).cpu().numpy() for t in self.to_list(emb)]) + for im in imgs: + im.close() + + q_embs = [] + with torch.no_grad(): + for batch_q in batch(queries, self.text_batch_size): + processed_queries = processor.process_queries(batch_q).to(model.device) + emb = model(**processed_queries) + q_embs.extend([t.detach().to(torch.float32).cpu().numpy() for t in self.to_list(emb)]) + + document_embeddings = np.concatenate(self.pad_for_concat(img_embs, pad_axis=1), axis=0) + query_embeddings = np.concatenate(self.pad_for_concat(q_embs, pad_axis=1), axis=0) + + return document_embeddings, query_embeddings + + +class SparseEmbedder(Embedder): + is_sparse = True + + def calc_embeddings(self, model_id: str, docs: list[str], queries: list[str]): + model = SparseTextEmbedding(model_name=model_id) + document_embeddings = list(model.embed(docs)) + query_embeddings = list(model.embed(queries)) + return document_embeddings, query_embeddings + + +class Modality(Enum): + TEXT = "text" + VISION = "vision" + + +@dataclass +class EncoderConfig: + model_id: str + modality: Modality + embedder: type[Embedder] + + +class Embedders: + qwen_text = EncoderConfig(model_id="Qwen/Qwen3-Embedding-4B", modality=Modality.TEXT, + embedder=QwenEmbedder) + linq = EncoderConfig(model_id="Linq-AI-Research/Linq-Embed-Mistral", modality=Modality.TEXT, + embedder=LinqMistralEmbedder) + nvidia = EncoderConfig(model_id="nvidia/llama-nemoretriever-colembed-3b-v1", modality=Modality.VISION, + embedder=NvidiaEmbedder) + jina_multi = EncoderConfig(model_id="jinaai/jina-embeddings-v4", modality=Modality.VISION, + embedder=JinaImageMultiEmbedder) + jina_single = EncoderConfig(model_id="jinaai/jina-embeddings-v4", modality=Modality.VISION, + embedder=JinaImageEmbedder) + jina_text_multi = EncoderConfig(model_id="jinaai/jina-embeddings-v4", modality=Modality.TEXT, + embedder=JinaTextMultiEmbedder) + jina_text_single = EncoderConfig(model_id="jinaai/jina-embeddings-v4", modality=Modality.TEXT, + embedder=JinaTextEmbedder) + colnomic = EncoderConfig(model_id="nomic-ai/colnomic-embed-multimodal-7b", modality=Modality.VISION, + embedder=NomicEmbedder) + bm25 = EncoderConfig(model_id="Qdrant/bm25", modality=Modality.TEXT, + embedder=SparseEmbedder) + + +all_embedders = [x for x in Embedders.__dict__.keys() if not x.startswith("_")] diff --git a/fusion_methods.py b/fusion_methods.py new file mode 100644 index 0000000..c4a7fa1 --- /dev/null +++ b/fusion_methods.py @@ -0,0 +1,103 @@ +import numpy as np +import torch + + +def reciprocal_rank_fusion(r1, r2, k=60): + r = {} + for q in r1.keys(): + curr_r1 = r1[q] + curr_r2 = r2[q] + rrf_scores = {} + for i, ((doc1, _), (doc2, _)) in enumerate(zip(curr_r1, curr_r2)): + rrf_scores[doc1] = rrf_scores.get(doc1, 0) + 1 / (k + i) + rrf_scores[doc2] = rrf_scores.get(doc2, 0) + 1 / (k + i) + r[q] = [(doc, score) + for doc, score in sorted(rrf_scores.items(), key=lambda x: x[1], reverse=True)][:len(curr_r1)] + return r + + +def average_ranking_fusion(r1, r2): + + def get_rank(doc, curr_r): + for i, (d, s) in enumerate(curr_r): + if d == doc: + return i + return len(curr_r) + + r = {} + for q in r1.keys(): + curr_r1 = r1[q] + curr_r2 = r2[q] + union = set([doc for doc, _ in curr_r1]).union(set([doc for doc, _ in curr_r2])) + raw_rank_averages = {} + for doc in union: + avg_rank = (get_rank(doc, curr_r1) + get_rank(doc, curr_r2))/2 + raw_rank_averages[doc] = avg_rank + r[q] = [(doc, 1/(score+1e-6)) + for doc, score in sorted(raw_rank_averages.items(), key=lambda x: x[1])][:len(curr_r1)] + return r + + +def normalize_softmax(rankings): + scores = torch.from_numpy(np.array([score for _, score in rankings])) + scores = torch.softmax(scores, dim=0).tolist() + new_rankings = [] + for i, ranking in enumerate(rankings): + new_rankings.append((ranking[0], scores[i])) + return new_rankings + + +def normalize_min_max(rankings): + s = torch.tensor([b for _, b in rankings], dtype=torch.float32) + s = (s - s.min()) / (s.max() - s.min() + 1e-12) + return [(a, v.item()) for (a, _), v in zip(rankings, s)] + + +def sim_score_fusion(r1, r2, normal_func=normalize_min_max, alpha=0.5): + def get_score(doc, curr_r): + for d, s in curr_r: + if d == doc: + return s + return 0 + + r = {} + for q in r1.keys(): + curr_r1 = r1[q] + curr_r2 = r2[q] + union = set([doc for doc, _ in curr_r1]).union(set([doc for doc, _ in curr_r2])) + curr_r1, curr_r2 = normal_func(curr_r1), normal_func(curr_r2) + raw_score_averages = {} + for doc in union: + agg_score = alpha*get_score(doc, curr_r1) + (1-alpha)*get_score(doc, curr_r2) + raw_score_averages[doc] = agg_score + r[q] = [(doc, score) + for doc, score in sorted(raw_score_averages.items(), key=lambda x: x[1], reverse=True)][:len(curr_r1)] + return r + + +# def oracle_fusion(rankings_1, rankings_2, question, top_k=40): +# gold_truth_ids = question["ground_truths_context_ids"] +# ndcg1 = ndcg_at_k([vars(r) for r in rankings_1],gold_truth_ids) +# ndcg2 = ndcg_at_k([vars(r) for r in rankings_2],gold_truth_ids) +# if ndcg1 > ndcg2: +# return [r.metadata["document_id"] for r in rankings_1] +# return [r.metadata["document_id"] for r in rankings_2] + +# def fuse_rankings(rag_results, model1_results, model2_results, fusion_func, fusion_factor=200, trimm_factor=200): +# q_ids = set(rag_results.get_qids()) +# for q_id in q_ids: +# context_model1 = model1_results[q_id]["context"][:fusion_factor] +# context_model2 = model2_results[q_id]["context"][:fusion_factor] +# if fusion_func == oracle_fusion: +# docs_ids = oracle_fusion(context_model1, context_model2, rag_results[q_id]) +# else: +# docs_ids = fusion_func(context_model1, context_model2, top_k=trimm_factor) +# unique_docs = [] +# seen_ids = set() +# for doc_id in docs_ids: +# doc = get_doc(context_model1, context_model2, doc_id) +# if doc_id not in seen_ids: +# seen_ids.add(doc_id) +# unique_docs.append(doc) +# +# rag_results[q_id]["context"] = unique_docs diff --git a/query_optimizations.py b/query_optimizations.py new file mode 100644 index 0000000..497481f --- /dev/null +++ b/query_optimizations.py @@ -0,0 +1,221 @@ +from enum import Enum +from functools import partial +from typing import Callable + +import torch +import torch.nn.functional as F + +from dataset_configs import DataSplit +from retriever import get_topk, Retriever +from utils import get_device, slice_sparse_coo_tensor + + +def js_divergence(p, q, eps=1e-8): + m = .5 * (p + q) + return .5 * (torch.sum(p * torch.log((p + eps) / (m + eps))) + torch.sum(q * torch.log((q + eps) / (m + eps)))) + + +def kl_divergence(p, q, eps=1e-8): + return torch.sum(p * torch.log((p + eps) / (q + eps))) + + +def optimize_queries_with_search(main_model: Retriever, feedback_model: Retriever, dataset, device, split: DataSplit, + optimize_queries_func: Callable, **kwargs): + optimized_queries, _, _ = optimize_queries_func( + main_model=main_model, feedback_model=feedback_model, dataset=dataset, device=device, split=split, **kwargs) + r, _ = main_model.run_retrieval(dataset, q=optimized_queries.to(device), split=split) + return r + + +def optimize_queries_no_search(main_model: Retriever, feedback_model: Retriever, dataset, device, split: DataSplit, + optimize_queries_func: Callable, **kwargs): + optimized_queries, doc_indices_per_query, docs_per_query = optimize_queries_func( + main_model=main_model, feedback_model=feedback_model, dataset=dataset, device=device, split=split, **kwargs) + + questions = dataset.benchmark['question'] + q_indices = dataset.get_queries_indices(split) + questions = questions[q_indices.numpy()] + r = {} + for question, optimized_query, doc_indices, docs in zip( + questions, optimized_queries, doc_indices_per_query, docs_per_query): + scores = main_model.compute_scores(optimized_query.unsqueeze(0).to(device), docs.to(device)) + r[question] = [(idx, val) for idx, val in zip(doc_indices.cpu().tolist(), scores.cpu().tolist())] + return r + + +def optimize_queries_main(main_model, feedback_model, dataset, k, lr, n_steps, T, mixture_alpha, loss_func, device, + optimizer, split: DataSplit): + f_d, f_q_dict = feedback_model.load_embs(dataset) + f_q = f_q_dict[split] + d, q = main_model.d.to(device), main_model.q_dict[split].to(device) + Q = q.size(0) + out_q = [] + doc_indices_per_query = [] + docs_per_query = [] + + qs = [q[i].unsqueeze(0).clone().detach().requires_grad_(True) for i in range(Q)] + opt = optimizer(qs, lr=lr) # one optimizer for all queries + top_k = get_topk(d, q, k, is_multi=main_model.is_multi)[1].squeeze() + for q_id in range(Q): + q1 = qs[q_id] + q2 = f_q[q_id].unsqueeze(0).to(device) + set1 = d[top_k[q_id]] + doc_indices_per_query.append(top_k[q_id].detach().cpu()) + docs_per_query.append(set1.detach().cpu()) + + if f_d.is_sparse: + set2 = slice_sparse_coo_tensor(f_d, slice_indices=top_k[q_id]).to(device) + else: + set2 = f_d[(top_k[q_id]).to(f_d.device)].to(device) + d2 = F.softmax(feedback_model.compute_scores(q2, set2) / T, -1).to(device) + d1 = F.softmax(main_model.compute_scores(q1, set1) / T, -1).to(device) + mixture_d = (1-mixture_alpha)*d1 + mixture_alpha*d2 + for step in range(n_steps): + d1 = F.softmax(main_model.compute_scores(q1, set1) / T, -1).to(device) + loss = loss_func(mixture_d, d1) + opt.zero_grad(set_to_none=True) + loss.backward(retain_graph=True) + opt.step() + out_q.append(q1.detach().squeeze()) + + return torch.stack(out_q), doc_indices_per_query, docs_per_query + + +def optimize_queries_union(main_model, feedback_model, dataset, k, lr, n_steps, T, mixture_alpha, loss_func, device, + optimizer, split: DataSplit): + f_d, f_q_dict = feedback_model.load_embs(dataset) + d, q, f_d, f_q = main_model.d.to(device), main_model.q_dict[split].to(device), f_d.to(device), f_q_dict[split].to(device) + Q = q.size(0) + out_q = [] + doc_indices_per_query = [] + docs_per_query = [] + + qs = [q[i].unsqueeze(0).clone().detach().requires_grad_(True) for i in range(Q)] + opt = optimizer(qs, lr=lr) # one optimizer for all queries + + sim1 = main_model.compute_scores(q, d) + sim2 = feedback_model.compute_scores(f_q, f_d) + + _, top_idx1 = torch.topk(sim1, k=k, dim=-1) + _, top_idx2 = torch.topk(sim2, k=k, dim=-1) + + for i in range(Q): + u = torch.unique(torch.cat([top_idx1[i], top_idx2[i]])) + docs_u = d.index_select(0, u.to(d.device)) + doc_indices_per_query.append(u.detach().cpu()) + docs_per_query.append(docs_u.detach().cpu()) + + d1 = torch.softmax(sim1[i, u] / T, dim=-1).to(device) + d2 = torch.softmax(sim2[i, u] / T, dim=-1).to(device) + mixture_d = (1-mixture_alpha)*d1 + mixture_alpha*d2 + for step in range(n_steps): + d1 = F.softmax(main_model.compute_scores(qs[i], docs_u) / T, -1).to(device) + loss = loss_func(mixture_d, d1) + opt.zero_grad(set_to_none=True) + loss.backward(retain_graph=True) + opt.step() + out_q.append(qs[i].detach().squeeze()) + + return torch.stack(out_q), doc_indices_per_query, docs_per_query + + +def optimize_queries_union_sample(main_model, feedback_model, dataset, k, lr, n_steps, T, mixture_alpha, loss_func, device, + optimizer, split): + f_d, f_q_dict = feedback_model.load_embs(dataset) + d, q, f_d, f_q = main_model.d.to(device), main_model.q_dict[split].to(device), f_d.to(device), f_q_dict[split].to(device) + Q = q.size(0) + out_q = [] + doc_indices_per_query = [] + docs_per_query = [] + + qs = [q[i].unsqueeze(0).clone().detach().requires_grad_(True) for i in range(Q)] + opt = optimizer(qs, lr=lr) # one optimizer for all queries + + sim1 = main_model.compute_scores(q, d) + sim2 = feedback_model.compute_scores(f_q, f_d) + + _, top_idx1 = torch.topk(sim1, k=k, dim=-1) + _, top_idx2 = torch.topk(sim2, k=k, dim=-1) + + for i in range(Q): + u = torch.unique(torch.cat([top_idx1[i], top_idx2[i]])) + docs_u = d.index_select(0, u.to(d.device)) + doc_indices_per_query.append(u.detach().cpu()) + docs_per_query.append(docs_u.detach().cpu()) + + d1 = sim1[i, u].to(device) + d2 = sim2[i, u].to(device) + mixture_d = (1-mixture_alpha)*d1 + mixture_alpha*d2 + for step in range(n_steps): + indices_to_sample = torch.randperm(u.size()[0])[:10] + sample_docs_u = d.index_select(0, u[indices_to_sample].to(d.device)) + d1 = F.softmax(main_model.compute_scores(qs[i], sample_docs_u) / T, -1).to(device) + other_d = (mixture_d[indices_to_sample] / T).softmax(-1) + loss = loss_func(other_d, d1) + opt.zero_grad(set_to_none=True) + loss.backward(retain_graph=True) + opt.step() + out_q.append(qs[i].detach().squeeze()) + return torch.stack(out_q), doc_indices_per_query, docs_per_query + + +def optimize_queries_dynamic(mod_1, mod_2, dataset, k, lr, n_steps, T, mixture_alpha, loss_func): + # TODO not functional currently + device = get_device() + for _ in range(5): + optimized_queries_1 = mod_1.optimize_queries(mod_2, dataset, k=k, lr=lr, n_steps=n_steps, T=T, + mixture_alpha=mixture_alpha, loss_func=loss_func) + r3, _ = mod_1.run_retrieval(dataset, q=optimized_queries_1.to(device)) + print(dataset.evaluate(r3)) + mod_1.set_queries(optimized_queries_1) + + optimized_queries_2 = mod_2.optimize_queries(mod_1, dataset, k=k, lr=lr, n_steps=n_steps, T=T, + mixture_alpha=mixture_alpha, loss_func=loss_func) + r4, _ = mod_2.run_retrieval(dataset, q=optimized_queries_2.to(device)) + print(dataset.evaluate(r4)) + mod_2.set_queries(optimized_queries_2) + + +class OptimizationFunctions(Enum): + main_with_search = partial(optimize_queries_with_search, + optimize_queries_func=optimize_queries_main) + union_with_search = partial(optimize_queries_with_search, + optimize_queries_func=optimize_queries_union) + union_sample_with_search = partial(optimize_queries_with_search, + optimize_queries_func=optimize_queries_union_sample) + + main_no_search = partial(optimize_queries_no_search, + optimize_queries_func=optimize_queries_main) + union_no_search = partial(optimize_queries_no_search, + optimize_queries_func=optimize_queries_union) + union_sample_no_search = partial(optimize_queries_no_search, + optimize_queries_func=optimize_queries_union_sample) + + +all_optimization_funcs = [x for x in OptimizationFunctions.__dict__.keys() if not x.startswith("_")] + + +def scores_feedback(main_model, feedback_model, dataset, top_k_idxs, alpha=0.5, split=DataSplit.TEST): + device = get_device() + f_d, f_q_dict = feedback_model.load_embs(dataset) + d, q, f_d, f_q = main_model.d.to(device), main_model.q_dict[split].to(device), f_d.to(device), f_q_dict[split].to(device) + top_k = top_k_idxs.squeeze() + results = {} + questions = dataset.benchmark['question'] + q_indices = dataset.get_queries_indices(split) + questions = questions[q_indices.numpy()] + for q_id, question in enumerate(questions): + q1 = q[q_id].unsqueeze(0) + q2 = f_q[q_id].unsqueeze(0) + set1 = d[top_k[q_id]] + if f_d.is_sparse: + set2 = slice_sparse_coo_tensor(f_d, slice_indices=top_k[q_id]) + else: + set2 = f_d[top_k[q_id]] + d2 = F.softmax(feedback_model.compute_scores(q2, set2), -1).to(device) + d1 = F.softmax(main_model.compute_scores(q1, set1), -1).to(device) + mixture_d = alpha * d1 + (1 - alpha) * d2 + results[question] = sorted([(idx.item(), val.item()) for (idx, val) in zip(top_k[q_id], mixture_d)], + key=lambda x: x[1], reverse=True) + + return results diff --git a/retriever.py b/retriever.py new file mode 100644 index 0000000..fc4f6b8 --- /dev/null +++ b/retriever.py @@ -0,0 +1,218 @@ +import numpy as np +import torch +import torch.nn.functional as F + +from dataset_configs import DataSplit, RagDataset +from embedding_configs import EncoderConfig +from utils import get_device + + +def score_multi_vector( + emb_queries: torch.Tensor | list[torch.Tensor], + emb_passages: torch.Tensor | list[torch.Tensor], + batch_size: int, +) -> torch.Tensor: + """ + Evaluate the similarity scores using the MaxSim scoring function. + + Inputs: + - emb_queries: List of query embeddings, each of shape (n_seq, emb_dim). + - emb_passages: List of document embeddings, each of shape (n_seq, emb_dim). + - batch_size: Batch size for the similarity computation. + """ + if len(emb_queries) == 0: + raise ValueError("No queries provided") + if len(emb_passages) == 0: + raise ValueError("No passages provided") + + if emb_queries[0].device != emb_passages[0].device: + raise ValueError("Queries and passages must be on the same device") + + if emb_queries[0].dtype != emb_passages[0].dtype: + raise ValueError("Queries and passages must have the same dtype") + + scores: list[torch.Tensor] = [] + + for i in range(0, len(emb_queries), batch_size): + batch_scores = [] + qs_batch = torch.nn.utils.rnn.pad_sequence( + emb_queries[i: i + batch_size], + batch_first=True, + padding_value=0, + ) + for j in range(0, len(emb_passages), batch_size): + ps_batch = torch.nn.utils.rnn.pad_sequence( + emb_passages[j: j + batch_size], + batch_first=True, + padding_value=0, + ) + batch_scores.append(torch.einsum("bnd,csd->bcns", qs_batch, ps_batch).max(dim=3)[0].sum(dim=2)) + batch_scores = torch.cat(batch_scores, dim=1) + scores.append(batch_scores) + + return torch.cat(scores, dim=0) + + +def get_topk(d_embs, q_embs, k, is_multi=False): + device = get_device() + if is_multi: + similarity = score_multi_vector(q_embs, d_embs, batch_size=8) + else: + similarity = torch.matmul(q_embs.to(device), d_embs.T.to(device)) + if similarity.is_sparse: + similarity = similarity.to_dense() + top_vals, top_idx = torch.topk(similarity, k=k, dim=-1) + return top_vals, top_idx + + +class Retriever: + def __init__(self, encoder_config: EncoderConfig): + file_suffix = encoder_config.embedder.file_suffix + self.id = encoder_config.model_id.split("/")[-1] + file_suffix.replace("_embeddings", "") + self.modality = encoder_config.modality + self.is_multi = encoder_config.embedder.is_multi_vector + self.is_sparse = encoder_config.embedder.is_sparse + + self.d: torch.Tensor | None = None + self.q_dict: dict[DataSplit, torch.Tensor] | None = None + self._emb_cache = {} + + def load_embs(self, dataset: RagDataset) -> tuple[torch.Tensor, dict[DataSplit, torch.Tensor]]: + key = (str(dataset.path.resolve()), self.id) + if key in self._emb_cache: + self.d, self.q_dict = self._emb_cache[key] + return self.d, self.q_dict + + def load_to_torch(embs_obj: np.ndarray) -> torch.Tensor | list[torch.Tensor]: + if self.is_sparse: + return [torch.sparse_coo_tensor(indices=torch.from_numpy(arr.indices).unsqueeze(0), values=arr.values) + for arr in list(embs_obj)] + elif embs_obj.dtype == np.object_: # multi-vector embedding tensors with varying sizes, which we treat as an object/list + return [torch.from_numpy(arr) for arr in list(embs_obj)] + else: # a single tensor with all the (single- or multi-vector) embeddings + return torch.from_numpy(embs_obj) + + p = dataset.path / (self.id + "_embeddings") + k = self.get_embs_key() + d = load_to_torch( + np.load(p / "docs.npz", allow_pickle=True)[k]) + q = load_to_torch( + np.load(p / "queries.npz", allow_pickle=True)[k]) + + if self.is_multi: + d, q = self.pad_multi_vectors(d), self.pad_multi_vectors(q) + elif self.is_sparse: + d, q = self.join_sparse_vectors(d, q) + else: + d, q = self.normalize(d), self.normalize(q) + + self.d = d + assert q.size(0) == dataset.num_queries, \ + "the number of query embeddings doesn't match the the number of queries in the dataset" + dev_q = q[dataset.get_queries_indices(split=DataSplit.DEV)] + test_q = q[dataset.get_queries_indices(split=DataSplit.TEST)] + self.q_dict = {DataSplit.DEV: dev_q, DataSplit.TEST: test_q} + self._emb_cache[key] = (self.d, self.q_dict) + return self.d, self.q_dict + + def clean_embs(self, dataset): + key = (str(dataset.path.resolve()), self.id) + if key in self._emb_cache: + del self._emb_cache[key] + + def set_queries(self, new_q): + self.q_dict = new_q + + def get_embs_key(self): + return "arr_0" + + @staticmethod + def normalize(t: torch.Tensor): + return F.normalize(t.float(), dim=-1) + + @staticmethod + def pad_multi_vectors(t: torch.Tensor): + return torch.nn.utils.rnn.pad_sequence(t, batch_first=True, padding_value=0) + + @staticmethod + def join_sparse_vectors(a: list[torch.Tensor], b: list[torch.Tensor]): + def build_sparse(indices_list, values_list, rows): + if not indices_list: + return torch.sparse_coo_tensor(torch.empty((2, 0), dtype=torch.long), + torch.tensor([], dtype=torch.float32), + size=(rows, max_cols), dtype=torch.float32) + indices = torch.cat(indices_list, dim=1) + values = torch.cat(values_list) + return torch.sparse_coo_tensor(indices, values, size=(rows, max_cols), dtype=torch.float32) + + all_indices_a = [] + all_values_a = [] + all_indices_b = [] + all_values_b = [] + max_cols = 0 + + for i, t in enumerate(a): + t = t.coalesce() + idx = t.indices() + val = t.values() + + if idx.shape[0] != 1: + raise ValueError("Expected 1D sparse tensor") + + col_idx = idx[0] + row_idx = torch.full((col_idx.size(0),), i, dtype=torch.long) + indices = torch.stack([row_idx, col_idx]) + + all_indices_a.append(indices) + all_values_a.append(val) + + max_cols = max(max_cols, t.shape[0]) + + for i, t in enumerate(b): + t = t.coalesce() + idx = t.indices() + val = t.values() + + if idx.shape[0] != 1: + raise ValueError("Expected 1D sparse tensor") + + col_idx = idx[0] + row_idx = torch.full((col_idx.size(0),), i, dtype=torch.long) + indices = torch.stack([row_idx, col_idx]) + + all_indices_b.append(indices) + all_values_b.append(val) + + max_cols = max(max_cols, t.shape[0]) + + A = build_sparse(all_indices_a, all_values_a, len(a)) + B = build_sparse(all_indices_b, all_values_b, len(b)) + return A, B + + def run_retrieval(self, dataset, q=None, d=None, split=DataSplit.TEST): + device = get_device() + if q is None: + q = self.q_dict[split].to(device) + if d is None: + d = self.d.to(device) + top_vals, top_idx = get_topk(d, q, k=50, is_multi=self.is_multi) + results = {} + questions = dataset.benchmark['question'] + q_indices = dataset.get_queries_indices(split) + questions = questions[q_indices.numpy()] + for q, val_row, idx_row in zip( + questions, top_vals.cpu().tolist(), top_idx.cpu().tolist() + ): + results[q] = [(idx, val) for (idx, val) in zip(idx_row, val_row)] + return results, top_idx + + def compute_scores(self, q, d): + if self.is_multi: + similarity = score_multi_vector(q, d, batch_size=16) + similarity = similarity / q.size(1) + else: + similarity = torch.matmul(q, d.T) + + if similarity.is_sparse: + similarity = similarity.to_dense() + return similarity.squeeze() diff --git a/run_experiments.py b/run_experiments.py new file mode 100644 index 0000000..63092b3 --- /dev/null +++ b/run_experiments.py @@ -0,0 +1,397 @@ +import argparse +import ast +import gc +import json +import os +from collections import defaultdict +from itertools import product +from functools import partial +from multiprocessing import Pool, cpu_count +from multiprocessing.pool import ThreadPool +from typing import Callable + +import numpy as np +import pandas as pd +import torch +import tqdm + +from dataset_configs import Datasets, DataSplit, RagDataset +from embedding_configs import Embedders +from fusion_methods import average_ranking_fusion, normalize_softmax, normalize_min_max, reciprocal_rank_fusion, \ + sim_score_fusion +from query_optimizations import OptimizationFunctions, scores_feedback, kl_divergence, js_divergence +from retriever import Retriever +from utils import get_device, get_run_hash, set_seed, on_ccc + + +def oracle_retriever(dataset: RagDataset): + b = dataset.benchmark + oracle_res = {q: [(k, v) for k, v in r.items()] + for q, r in zip(b['question'], b['correct_answer_document_ids'])} + return oracle_res + + +def run_baselines(dataset: RagDataset, mod_1, mod_2, r1, r2, top_idx_1, top_idx_2, target_metric='ndcg@5'): + results = [] + info_dict = { + "lr": 0.0, "k": 0, "n_steps": 0, "temp": 0.0, + "mixture_alpha": 0, "loss_func": "baseline", + "weight": 0, + "optimization_func": "", "optimizer": "", + "dataset": dataset.id, + } + + results.append({ + **info_dict, + "run_id": mod_1.id, + "main_model": mod_1.id, + "feedback_model": "", + "metrics": dataset.evaluate(r1) + }) + + results.append({ + **info_dict, + "run_id": mod_2.id, + "main_model": mod_2.id, + "feedback_model": "", + "metrics": dataset.evaluate(r2) + }) + + oracle_r = oracle_retriever(dataset) + results.append({ + **info_dict, + "run_id": "oracle", + "main_model": "", + "feedback_model": "", + "metrics": dataset.evaluate(oracle_r) + }) + + late_pipelines = { + "RRF": reciprocal_rank_fusion, + "average": average_ranking_fusion, + } + + tunable_pipelines = { + "sim_score_minmax": partial(sim_score_fusion, normal_func=normalize_min_max), + "sim_score_softmax": partial(sim_score_fusion, normal_func=normalize_softmax), + f"scores_feedback_{mod_1.id}-{mod_2.id}": partial(scores_feedback, main_model=mod_1, feedback_model=mod_2, + dataset=dataset, top_k_idxs=top_idx_1), + f"scores_feedback_{mod_2.id}-{mod_1.id}": partial(scores_feedback, main_model=mod_2, feedback_model=mod_1, + dataset=dataset, top_k_idxs=top_idx_2) + } + + if args.tune: + r1_dev, _ = mod_1.run_retrieval(dataset, split=DataSplit.DEV) + r2_dev, _ = mod_2.run_retrieval(dataset, split=DataSplit.DEV) + + for pipeline, fusion_func in tunable_pipelines.items(): + best_val, best_alpha = 0, 0 + for alpha in [0.1*i for i in range(1, 10)]: + if "scores_feedback" in pipeline: + r = fusion_func(alpha=alpha, split=DataSplit.DEV) + else: + r = fusion_func(r1_dev, r2_dev, alpha=alpha) + + eval = dataset.evaluate(r) + if eval[target_metric] > best_val: + best_val = eval[target_metric] + best_alpha = alpha + + if "scores_feedback" in pipeline: + r = fusion_func(alpha=best_alpha) + run_id = pipeline + else: + r = fusion_func(r1, r2, alpha=best_alpha) + run_id = f"{pipeline}-{mod_1.id}-{mod_2.id}" + + results.append({ + **info_dict, + "run_id": run_id, + "weight": best_alpha, + "main_model": "", + "feedback_model": "", + "metrics": dataset.evaluate(r) + }) + else: + for pipeline, fusion_func in tunable_pipelines.items(): + for alpha in [0.1*i for i in range(1, 10)]: + late_pipelines[f"{pipeline}_{alpha:.2f}-{1-alpha:.2f}"] = partial(fusion_func, alpha=alpha) + + for pipeline, fusion_func in late_pipelines.items(): + if "scores_feedback" in pipeline: + r = fusion_func() + else: + r = fusion_func(r1, r2) + results.append({ + **info_dict, + "run_id": f"{pipeline}-{mod_1.id}-{mod_2.id}", + "main_model": "", + "feedback_model": "", + "metrics": dataset.evaluate(r) + }) + + return results + + +def tune_hyper(dataset, mod_1, mod_2, device, input_params, weight_for_feedback_model=0.5, target_metric='ndcg@5'): + if any(x[4] == "dynamic" for x in input_params): + input_params = [p[:4] + (weight_for_feedback_model,) + p[5:] + if p[4] == "dynamic" else p for p in input_params] + dev_results = [] + for run_args in tqdm.tqdm(input_params): + dev_results.append(run_query_optimizations(dataset, mod_1, mod_2, device, *run_args, split=DataSplit.DEV)) + max_val, best_params = 0, None + for results, params in zip(dev_results, input_params): + val = results['metrics'][target_metric] + if val > max_val: + max_val = val + best_params = params + print(f'best params: {best_params} the {target_metric} is {max_val}') + return best_params + + +def run_query_optimizations(dataset, mod_1: Retriever, mod_2: Retriever, device, lr, k, n, t, mixture_alpha, + loss_func: Callable, optimizer: torch.optim.Optimizer, optimization_func_name: str, + split=DataSplit.TEST): + if mod_1.is_sparse: + result_dict = {} + else: + optimization_func = OptimizationFunctions[optimization_func_name].value + r = optimization_func( + mod_1, mod_2, dataset, device=device, + k=k, lr=lr, n_steps=n, T=t, mixture_alpha=mixture_alpha, loss_func=loss_func, + optimizer=optimizer, split=split) + + result_dict = {"run_id": f"{mod_1.id}-feedback-from-{mod_2.id}", + "main_model": mod_1.id, + "feedback_model": mod_2.id, + "metrics": dataset.evaluate(r)} + + return result_dict + + +def main(args): + device = get_device() + prefix = "/proj/omri/" if on_ccc() else "" + + # benchmarks + vidore1 = [Datasets.arxivqa, Datasets.docvqa, Datasets.infovqa, Datasets.tabfquad, Datasets.tatdqa, + Datasets.shiftproject, Datasets.artificial_intelligence, Datasets.energy_test, + Datasets.government_reports, Datasets.healthcare] + vidore2 = [ + Datasets.esg_reports_v2, + Datasets.biomedical_lectures_v2, + Datasets.economics_reports_v2, + Datasets.esg_reports_human_labeled_v2 + ] + real_mm_rag = [ + Datasets.FinReport, + Datasets.FinSlides, + Datasets.TechReport, + Datasets.TechSlides + ] + + benchmarks = {"vidore1": vidore1, "vidore2": vidore2, "real_mm_rag": real_mm_rag} + + models_in_experiment = [ + # Embedders.nvidia, + # Embedders.jina_multi, + # Embedders.jina_single, + Embedders.colnomic, + + # Embedders.jina_text_single, + # Embedders.jina_text_multi, + Embedders.linq, + Embedders.qwen_text, + # Embedders.bm25, + ] + + datasets_in_experiment = [] + for k in args.benchmarks: + datasets_in_experiment += benchmarks[k] + + lrs = [ + 5e-6, + 1e-5, + 3e-5, + 5e-5, + 1e-4, + 3e-4, + 5e-4, + 1e-3, + 3e-3, + 5e-3, + ] + ks = [ + 10, + # 20 + # 50, + ] + n_steps = [ + # 10, + # 25, + 50, + # 100 + ] + Ts = [1] + mixture = [ + "dynamic" + ] + loss_funcs = [ + kl_divergence, + ] + optimization_funcs = [ + # OptimizationFunctions.main_no_search.name, + OptimizationFunctions.union_no_search.name, + # OptimizationFunctions.union_sample_no_search.name, + ] + + optimizers = [ + torch.optim.Adam, + # torch.optim.AdamW, + # torch.optim.Adagrad, + # torch.optim.SGD, + # torch.optim.RMSprop + ] + + models_in_experiment = [Retriever(m) for m in models_in_experiment] + h = get_run_hash(models_in_experiment, datasets_in_experiment, lrs, ks, n_steps, Ts, mixture, + loss_funcs, optimizers, optimization_funcs) + out_dir = f"output/results-{h}{args.out_dir_suffix}" + os.makedirs(out_dir, exist_ok=True) + print(f"Results in {out_dir}") + + rows = [] + for dataset_name in datasets_in_experiment: + set_seed() + dataset = RagDataset(dataset_name, prefix=prefix) + for i in range(len(models_in_experiment) - 1): + mod_1 = models_in_experiment[i] + + # the loaded embeddings and retrieval results of mod_1 (i) are reused for every mod_2 (j) + mod_1.load_embs(dataset) + r1, top_idx_1 = mod_1.run_retrieval(dataset) + + for j in range(i + 1, len(models_in_experiment)): + mod_2 = models_in_experiment[j] + + print(f"\n{mod_1.id}-{mod_2.id}-{dataset.id}") + mod_2.load_embs(dataset) + r2, top_idx_2 = mod_2.run_retrieval(dataset) + + baselines = run_baselines(dataset, mod_1, mod_2, + r1, r2, top_idx_1, top_idx_2) + rows += baselines + mod_1_weight = 0.5 + if args.tune: + for res_dict in baselines: + if "sim_score_softmax" in res_dict["run_id"]: + assert f"{mod_1.id}-{mod_2.id}" in res_dict["run_id"] + mod_1_weight = res_dict["weight"] + + exp_results = [] + input_params = [(lr, k, n, t, mixture_alpha, loss_func, optimizer, optimization_func) + for lr, k, n, t, mixture_alpha, loss_func, optimizer, optimization_func, + in product(lrs, ks, n_steps, Ts, mixture, loss_funcs, optimizers, optimization_funcs)] + + if args.tune and len(input_params) > 1: + mod_1_best_params = tune_hyper(dataset, mod_1, mod_2, device, input_params, + weight_for_feedback_model=1-mod_1_weight) + mod_2_best_params = tune_hyper(dataset, mod_2, mod_1, device, input_params, + weight_for_feedback_model=mod_1_weight) + input_params = [(dataset, mod_1, mod_2, device, *mod_1_best_params), + (dataset, mod_2, mod_1, device, *mod_2_best_params)] + else: + input_params = [(dataset, mod_1, mod_2, device, *params) for params in input_params] + \ + [(dataset, mod_2, mod_1, device, *params) for params in input_params] + + description = f"Running all query optimizations for pair {mod_1.id},{mod_2.id} (parallelization={args.use_parallelization})" + if args.use_parallelization: + all_results = [] + pbar = tqdm.tqdm(total=len(input_params), desc=description) + pool = ThreadPool(4) if on_ccc() else Pool(cpu_count()) + for run_args in input_params: + all_results.append( + pool.apply_async(run_query_optimizations, run_args, callback=lambda _: pbar.update(1))) + pool.close() + pool.join() + for process_result in all_results: + exp_results.append(process_result.get()) + pbar.close() + else: + for run_args in tqdm.tqdm(input_params, desc=description): + exp_results.append(run_query_optimizations(*run_args)) + + for result_dict, (_, _, _, _, lr, k, n, t, mixture_alpha, loss_func, optimizer, optimization_func) in zip( + exp_results, input_params): + if len(result_dict) == 0: + continue + rows.append({ + **result_dict, + "weight": 0, + "lr": lr, + "k": k, + "n_steps": n, + "temp": t, + "mixture_alpha": mixture_alpha, + "loss_func": loss_func.__name__, + "optimization_func": optimization_func, + "optimizer": optimizer.__name__, + "dataset": dataset.id, + }) + + if (j == (len(models_in_experiment) - 1) + and i < (len(models_in_experiment) - 2) + and (models_in_experiment[i + 1].id == mod_2.id or models_in_experiment[i + 2].id == mod_2.id)): + pass # mod_2 embeddings can be reused + else: + mod_2.clean_embs(dataset) + + gc.collect() + torch.cuda.empty_cache() + + mod_1.clean_embs(dataset) + + # after every dataset, save all results collected so far + df = pd.DataFrame(rows) + metrics = sorted({k for m in df["metrics"] for k in m.keys()}) + for metric in metrics: + tmp = df.copy() + if metric == "raw_results": + tmp["raw_results"] = tmp["metrics"].apply(lambda m: m.get(metric, np.nan)) + raw_results = defaultdict(defaultdict) + for _, row in tmp.iterrows(): + raw_results[row["run_id"]][row["dataset"]] = row.to_dict() + with open(os.path.join(out_dir, "raw_results.json"), "w") as f: + json.dump(raw_results, f) + continue + + def merge_values(x): + non_nan = x.dropna().unique() + if len(non_nan) > 1: + return non_nan.tolist() + else: + return non_nan[0] + + id_cols = [col for col in df.columns if col not in {"dataset", "metrics"}] + tmp["metric_value"] = tmp["metrics"].apply(lambda m: m.get(metric, np.nan)) + wide = tmp.pivot_table(index=id_cols, columns="dataset", values="metric_value", aggfunc="first").reset_index() + if args.tune: + os.makedirs(os.path.join(out_dir, "tune"), exist_ok=True) + wide.to_csv(os.path.join(out_dir, "tune", f"{metric}.csv"), index=False) + wide = wide.groupby("run_id", as_index=False).agg(merge_values) + dataset_cols = [col for col in wide.columns if col not in id_cols] + wide["average"] = wide[dataset_cols].mean(axis=1, numeric_only=True) + wide.to_csv(os.path.join(out_dir, f"{metric}.csv"), index=False) + return h + + +if __name__ == "__main__": + parser = argparse.ArgumentParser() + parser.add_argument('--benchmarks', nargs='+', required=True) + parser.add_argument('-p', '--use_parallelization', type=ast.literal_eval, default=False) + parser.add_argument('-t', '--tune', type=ast.literal_eval, default=False) + parser.add_argument('-o', '--out_dir_suffix', default='') + + args = parser.parse_args() + main(args) diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..cb886fc --- /dev/null +++ b/utils.py @@ -0,0 +1,66 @@ +import os +import random + +import numpy as np +import torch +import hashlib +import json + + +def get_device(): + if torch.backends.mps.is_available(): + return "mps" # mac GPU + elif torch.cuda.is_available(): + return "cuda" + else: + return "cpu" + + +def on_ccc(): + return os.path.isdir("/proj/") + + +def get_doc(context_1, context_2, doc_id): + for doc in context_1: + if doc.metadata["document_id"] == doc_id: + return doc + for doc in context_2: + if doc.metadata["document_id"] == doc_id: + return doc + + +def batch(x, n): + for i in range(0, len(x), n): + yield x[i:i + n] + + +def get_run_hash(*args): + s = json.dumps(args, sort_keys=True, default=str) + return hashlib.blake2b(s.encode(), digest_size=8).hexdigest() + + +def slice_sparse_coo_tensor(t: torch.Tensor, slice_indices: torch.Tensor): + def ainb(a, b): + indices = torch.zeros_like(a, dtype=torch.uint8) + for elem in b: + indices = indices | (a == elem) + + return indices.type(torch.bool) + + new_shape = (slice_indices.shape[0], t.shape[1]) + t = t.coalesce() + mask = ainb(a=t.indices()[0], b=slice_indices.to(t.device)) + sliced_tensor = torch.sparse_coo_tensor(t.indices()[:, mask], + t.values()[mask], + size=new_shape) + return sliced_tensor + + +def set_seed(seed: int = 42): + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + torch.backends.cudnn.deterministic = True + torch.backends.cudnn.benchmark = False