diff --git a/examples/vllm_multiturn/config/tool_config/search_tool_config.yaml b/examples/vllm_multiturn/config/tool_config/search_tool_config.yaml index 4050f8847f6..cd9f5b9c8f5 100644 --- a/examples/vllm_multiturn/config/tool_config/search_tool_config.yaml +++ b/examples/vllm_multiturn/config/tool_config/search_tool_config.yaml @@ -2,9 +2,9 @@ tools: - class_name: verl.tools.janv2_tool.web_search_tool.WebSearchTool config: type: native - rag_server_url: "http://10.220.108.31:3030" + rag_server_url: "http://localhost:3030" num_results: 10 - topk_retrieval: 30 + topk_retrieval: 200 num_workers: 64 rate_limit: 100000 timeout: 600 @@ -29,7 +29,7 @@ tools: - class_name: verl.tools.janv2_tool.scrape_tool.ScrapeTool config: type: native - rag_server_url: "http://10.220.108.31:3030" + rag_server_url: "http://localhost:3030" num_workers: 50 rate_limit: 100000 timeout: 600 diff --git a/examples/vllm_multiturn/rag_setup/bm25_index.py b/examples/vllm_multiturn/rag_setup/bm25_index.py index 831bdf679c0..c4419bba00c 100644 --- a/examples/vllm_multiturn/rag_setup/bm25_index.py +++ b/examples/vllm_multiturn/rag_setup/bm25_index.py @@ -1,7 +1,12 @@ import bm25s import json, re +import gzip +import os import datasets -import Stemmer +import Stemmer +import logging + +logger = logging.getLogger(__name__) def load_corpus(corpus_path: str): """Load corpus using datasets library""" @@ -20,35 +25,72 @@ def load_docs(corpus, doc_idxs): return results class BM25RetrieverLunce: - def __init__(self, corpus_path: str): + def __init__(self, corpus_path_or_corpus, is_corpus=False, cache_dir=None): + """ + Args: + corpus_path_or_corpus: Either a path string or a pre-loaded corpus dataset + is_corpus: If True, corpus_path_or_corpus is a pre-loaded corpus + cache_dir: Directory to save/load BM25 index. If None, uses corpus_path + "_bm25_cache" + """ + if is_corpus: + self.corpus = corpus_path_or_corpus + self.cache_dir = cache_dir + else: + logger.info("BM25: Loading corpus...") + self.corpus = load_corpus(corpus_path=corpus_path_or_corpus) + # Default cache dir based on corpus path + if cache_dir is None: + self.cache_dir = corpus_path_or_corpus + "_bm25_cache" + else: + self.cache_dir = cache_dir - self.retriever = self._build_index(corpus_path) - self.corpus = load_corpus(corpus_path=corpus_path) - - def _build_index(self, corpus_path): - with open(corpus_path,"r") as file: - lines = file.readlines() - self.raw_data = [] - for line in lines: - try: - data = json.loads(line) - self.raw_data.append(data) - except: - print(f"error when loading: {data}") - corpus = [re.sub(r'[^\w\s]', '', data["contents"]) for data in self.raw_data] self.stemmer = Stemmer.Stemmer("english") - retriever = bm25s.BM25() #corpus=corpus - retriever.index(bm25s.tokenize(corpus, stopwords="en", stemmer=self.stemmer)) + self.retriever = self._load_or_build_index() + + def _load_or_build_index(self): + """Load index from cache if exists, otherwise build and save""" + if self.cache_dir and os.path.exists(self.cache_dir): + logger.info(f"BM25: Loading cached index from {self.cache_dir}...") + retriever = bm25s.BM25.load(self.cache_dir, load_corpus=False) + logger.info("BM25: Cached index loaded successfully!") + return retriever + else: + logger.info(f"BM25: No cached index found, building new index...") + retriever = self._build_index() + + # Save to cache + if self.cache_dir: + logger.info(f"BM25: Saving index to {self.cache_dir}...") + os.makedirs(self.cache_dir, exist_ok=True) + retriever.save(self.cache_dir) + logger.info("BM25: Index saved to cache!") + + return retriever + + def _build_index(self): + logger.info(f"BM25: Building index for {len(self.corpus)} documents...") + + # Extract texts directly from corpus (more efficient) + corpus_texts = [re.sub(r'[^\w\s]', '', doc["contents"]) for doc in self.corpus] + + logger.info("BM25: Tokenizing corpus...") + tokens = bm25s.tokenize(corpus_texts, stopwords="en", stemmer=self.stemmer) + + logger.info("BM25: Indexing tokens...") + retriever = bm25s.BM25() + retriever.index(tokens) + + logger.info("BM25: Index built successfully!") return retriever - def _search(self,query: str, num: int): + def _search(self, query: str, num: int): results, scores = self.retriever.retrieve(bm25s.tokenize(query, stopwords="en", stemmer=self.stemmer), k=num) return results[0], scores[0] if __name__ == "__main__": - bm25_ = BM25Retriever("/mnt/nas/alex/deep-research/src/rag_setup/data/corpus/corpus.jsonl") + bm25_ = BM25RetrieverLunce("/mnt/nas/alex/deep-research/src/rag_setup/data/corpus/corpus.jsonl") print(bm25_._search("Mc Donald", 5)) result = bm25_._search(" Donald", 5) print(load_docs(bm25_.corpus, result[0][0])) \ No newline at end of file diff --git a/examples/vllm_multiturn/rag_setup/flashrag_server.py b/examples/vllm_multiturn/rag_setup/flashrag_server.py index 7e04b96a814..85ab5b0cc9a 100644 --- a/examples/vllm_multiturn/rag_setup/flashrag_server.py +++ b/examples/vllm_multiturn/rag_setup/flashrag_server.py @@ -1,8 +1,13 @@ -""" -RAG Server -""" - import os + +# Set CUDA device before importing torch - must be done first! +cuda_device = os.environ.get("CUDA_DEVICE", None) +if cuda_device is not None: + os.environ["CUDA_VISIBLE_DEVICES"] = cuda_device +else: + import random + os.environ["CUDA_VISIBLE_DEVICES"] = str(random.randint(0, 1)) + import json import logging import argparse @@ -25,14 +30,16 @@ from tqdm import tqdm from collections import defaultdict from bm25_index import BM25RetrieverLunce +from utils import load_index + # Configure logging logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) -TOP_BM25_RETRIEVAL = 1000 +TOP_BM25_RETRIEVAL = 2000 + # Suppress some warnings -import random warnings.filterwarnings("ignore", category=UserWarning) -os.environ["CUDA_VISIBLE_DEVICES"] = str(random.randint(0,1)) + def load_corpus(corpus_path: str): """Load corpus using datasets library""" corpus = datasets.load_dataset( @@ -162,69 +169,6 @@ def search(self, *args, **kwargs): def batch_search(self, *args, **kwargs): return self._batch_search(*args, **kwargs) -class BM25Retriever(BaseRetriever): - """BM25 retriever based on pre-built pyserini index.""" - - def __init__(self, config): - super().__init__(config) - from pyserini.search.lucene import LuceneSearcher - self.searcher = LuceneSearcher(self.index_path) - self.contain_doc = self._check_contain_doc() - if not self.contain_doc: - self.corpus = load_corpus(self.corpus_path) - self.max_process_num = 8 - - def _check_contain_doc(self): - """Check if the index contains document content""" - return self.searcher.doc(0).raw() is not None - - def _search(self, query: str, num: int = None, return_score: bool = False) -> List[Dict[str, str]]: - if num is None: - num = self.topk - - hits = self.searcher.search(query, num) - if len(hits) < 1: - if return_score: - return [], [] - else: - return [] - - scores = [hit.score for hit in hits] - if len(hits) < num: - warnings.warn('Not enough documents retrieved!') - else: - hits = hits[:num] - - if self.contain_doc: - all_contents = [json.loads(self.searcher.doc(hit.docid).raw())['contents'] for hit in hits] - results = [ - { - 'title': content.split("\n")[0].strip("\""), - 'text': "\n".join(content.split("\n")[1:]), - 'contents': content - } - for content in all_contents - ] - else: - results = load_docs(self.corpus, [hit.docid for hit in hits]) - - if return_score: - return results, scores - else: - return results - - def _batch_search(self, query_list, num: int = None, return_score: bool = False): - results = [] - scores = [] - for query in query_list: - item_result, item_score = self._search(query, num, True) - results.append(item_result) - scores.append(item_score) - - if return_score: - return results, scores - else: - return results class DenseRetriever(BaseRetriever): """Dense retriever based on pre-built faiss index.""" @@ -233,10 +177,14 @@ def __init__(self, config): super().__init__(config) self.index = faiss.read_index(self.index_path) if config.faiss_gpu: - co = faiss.GpuMultipleClonerOptions() - co.useFloat16 = True - co.shard = True - self.index = faiss.index_cpu_to_all_gpus(self.index, co=co) + # Check if faiss-gpu is available + if hasattr(faiss, 'GpuMultipleClonerOptions'): + co = faiss.GpuMultipleClonerOptions() + co.useFloat16 = True + co.shard = True + self.index = faiss.index_cpu_to_all_gpus(self.index, co=co) + else: + logger.warning("faiss-gpu not installed, falling back to CPU. Install faiss-gpu for GPU support.") self.corpus = load_corpus(self.corpus_path) self.encoder = Encoder( @@ -305,11 +253,10 @@ def _batch_search(self, query_list: List[str], num: int = None, return_score: bo return results def get_retriever(config): - """Automatically select retriever class based on config's retrieval method""" - if config.retrieval_method == "bm25": - return BM25Retriever(config) - else: - return DenseRetriever(config) + """Get dense retriever (BM25 is handled separately via BM25RetrieverLunce)""" + # Note: BM25 is now handled by BM25RetrieverLunce (pure Python, no Java) + # This function only returns DenseRetriever + return DenseRetriever(config) class BaseCrossEncoder: def __init__(self, model, batch_size=32, device="cuda"): @@ -404,7 +351,7 @@ class RetrieverConfig: retrieval_query_max_length: int = field(default=256) retrieval_use_fp16: bool = field(default=True) retrieval_batch_size: int = field(default=128) - retrieval_topk: int = field(default=10) + retrieval_topk: int = field(default=200) # Get 35 for reranking index_path: str = field(default="indexes/dense/e5_base_v2.index") corpus_path: str = field(default="data/corpus/processed_corpus.jsonl") faiss_gpu: bool = field(default=True) @@ -413,7 +360,7 @@ class RetrieverConfig: class RerankerConfig: """Configuration for reranker (from rerank_server.py)""" max_length: int = field(default=512) - rerank_topk: int = field(default=3) + rerank_topk: int = field(default=10) # Return top 10 after reranking rerank_model_name_or_path: str = field(default="cross-encoder/ms-marco-MiniLM-L12-v2") batch_size: int = field(default=32) reranker_type: str = field(default="sentence_transformer") @@ -434,8 +381,8 @@ def convert_title_format(text): class SearchRequest(BaseModel): queries: List[str] - topk_retrieval: Optional[int] = 10 - topk_rerank: Optional[int] = 3 + topk_retrieval: Optional[int] = 200 # Dense retrieval candidates for reranking + topk_rerank: Optional[int] = 10 # Final results after reranking return_scores: bool = False class SearchResponse(BaseModel): @@ -444,6 +391,7 @@ class SearchResponse(BaseModel): class VisitRequest(BaseModel): url: str + class HealthResponse(BaseModel): status: str pipeline_loaded: bool @@ -466,37 +414,53 @@ class StatsResponse(BaseModel): reranker_config = None bm25_retriever = None +# Initialize configurations from environment variables retriever_config = RetrieverConfig( - retrieval_method="e5", - index_path="data/corpus/e5_Flat.index", - corpus_path="data/corpus/corpus.jsonl", - retrieval_topk=25, - faiss_gpu=False, - retrieval_model_path="intfloat/e5-base-v2", - retrieval_pooling_method="mean", - retrieval_query_max_length=256, - retrieval_use_fp16=True, - retrieval_batch_size=128, - ) - + retrieval_method=os.environ.get("RETRIEVER_NAME", "e5"), + index_path=os.environ.get("INDEX_PATH", "data/corpus/e5_Flat.index"), + corpus_path=os.environ.get("CORPUS_PATH", "data/corpus/corpus.jsonl"), + retrieval_topk=int(os.environ.get("RETRIEVAL_TOPK", "35")), # 35 candidates for reranking + faiss_gpu=os.environ.get("FAISS_GPU", "true").lower() == "true", + retrieval_model_path=os.environ.get("RETRIEVER_MODEL", "intfloat/e5-base-v2"), + retrieval_pooling_method=os.environ.get("RETRIEVAL_POOLING_METHOD", "mean"), + retrieval_query_max_length=int(os.environ.get("RETRIEVAL_QUERY_MAX_LENGTH", "256")), + retrieval_use_fp16=os.environ.get("RETRIEVAL_USE_FP16", "true").lower() == "true", + retrieval_batch_size=int(os.environ.get("RETRIEVAL_BATCH_SIZE", "128")), +) + reranker_config = RerankerConfig( - rerank_topk=10, - rerank_model_name_or_path="cross-encoder/ms-marco-MiniLM-L12-v2", - batch_size=64, - ) + rerank_topk=int(os.environ.get("RERANKING_TOPK", "10")), # Top 10 after reranking + rerank_model_name_or_path=os.environ.get("RERANKER_MODEL", "cross-encoder/ms-marco-MiniLM-L12-v2"), + batch_size=int(os.environ.get("RERANKER_BATCH_SIZE", "64")), +) from contextlib import asynccontextmanager @asynccontextmanager async def lifespan(app: FastAPI): # Load the ML model - global retriever, reranker, retriever_config, reranker_config, bm25_retriever + global retriever, reranker, retriever_config, reranker_config, bm25_retriever, doc_store logger.info("Loading retrieval and reranking components...") + logger.info(f"Index path: {retriever_config.index_path}") + logger.info(f"Corpus path: {retriever_config.corpus_path}") + logger.info(f"CUDA device: {os.environ.get('CUDA_VISIBLE_DEVICES', 'not set')}") - # These configs will be set from command line arguments + # load json file mapping doc_id to full content + doc_store = load_index("./saved_index_data") + + logger.info("Loading dense retriever...") retriever = get_retriever(retriever_config) + + # Load reranker + logger.info("Loading reranker...") reranker = get_reranker(reranker_config) - bm25_retriever = BM25RetrieverLunce(retriever_config.corpus_path) + + bm25_index_path = os.environ.get("BM25_INDEX_PATH", None) + + bm25_cache_dir = os.environ.get("BM25_CACHE_DIR", retriever_config.corpus_path + "_bm25_cache") + logger.info(f"BM25 cache dir: {bm25_cache_dir}") + logger.info("Loading/Building BM25 index with bm25s...") + bm25_retriever = BM25RetrieverLunce(retriever.corpus, is_corpus=True, cache_dir=bm25_cache_dir) logger.info("FlashRAG pipeline initialized successfully") yield @@ -522,14 +486,11 @@ async def startup_event(): @app.post("/visit") async def visit(request: VisitRequest): - doc = retriever._get_doc(request.url.split("_")[-1]) - content = doc["contents"] - title = content.split("\n")[0] - text = "\n".join(content.split("\n")[1:]) + id = request.url.split("_")[-1] return { "result": [[{ - "title": title, - "text": text + "title": doc_store.get(str(id))['title'], + "text": doc_store.get(str(id))['full_contents'], }]] } @@ -545,49 +506,83 @@ async def search_endpoint(request: SearchRequest): try: start_time = time.time() - # step 0: get 1000 examples - + # get candidates using BM25 (using bm25s - pure Python) results_ids, _ = bm25_retriever._search(request.queries[0], TOP_BM25_RETRIEVAL) - # Step 1: Retrieve documents + # Retrieve documents retrieved_docs = retriever.batch_search( - query_list=request.queries, - num=request.topk_retrieval, - return_score=False, - search_indices = results_ids - ) + query_list=request.queries, + num=request.topk_retrieval, + return_score=False, + search_indices=results_ids + ) - # Step 2: Rerank documents + # Rerank documents reranked = reranker.rerank(request.queries, retrieved_docs) - # Step 3: Format response (following retrieval_rerank_server.py pattern) + + # Format response response = [] for i, doc_scores in reranked.items(): - doc_scores = doc_scores[:request.topk_rerank] - if request.return_scores: - combined = [] - for doc, score, doc_id in doc_scores: - # Convert the document format - converted_doc = convert_title_format(doc) - # Parse back to structured format - lines = converted_doc.split('\n', 1) - title = lines[0].strip('"') if lines else "No title" - text = lines[1] if len(lines) > 1 else "" + combined = [] + seen_titles = set() + + for doc, score, doc_id in doc_scores: + if len(combined) >= request.topk_rerank: + break + + converted_doc = convert_title_format(doc) + lines = converted_doc.split('\n', 1) + title = lines[0].strip('"') if lines else "No title" + + if title in seen_titles: + continue + + seen_titles.add(title) + + text = lines[1] if len(lines) > 1 else "" + + doc_dict = { + "doc_id": doc_id, + "title": title, + "text": text, + "contents": converted_doc + } + + if request.return_scores: + doc_dict["score"] = float(score) + + combined.append(doc_dict) + + # Safe check: if we don't have enough documents after reranking, retrieve more + if len(combined) < request.topk_rerank: + additional_needed = request.topk_rerank - len(combined) + additional_topk = request.topk_retrieval + (additional_needed * 10) + + # Retrieve additional documents + additional_docs = retriever.batch_search( + query_list=[request.queries[i]], + num=additional_topk, + return_score=False, + search_indices=results_ids + ) + + # Rerank the additional documents + additional_reranked = reranker.rerank([request.queries[i]], additional_docs) + + # Add additional documents until we reach topk_rerank + for doc, score, doc_id in additional_reranked.get(0, []): + if len(combined) >= request.topk_rerank: + break - doc_dict = { - "doc_id": doc_id, - "title": title, - "text": text, - "contents": converted_doc, - "score": float(score) - } - combined.append(doc_dict) - response.append(combined) - else: - formatted_docs = [] - for doc, _, doc_id in doc_scores: converted_doc = convert_title_format(doc) lines = converted_doc.split('\n', 1) title = lines[0].strip('"') if lines else "No title" + + if title in seen_titles: + continue + + seen_titles.add(title) + text = lines[1] if len(lines) > 1 else "" doc_dict = { @@ -596,8 +591,13 @@ async def search_endpoint(request: SearchRequest): "text": text, "contents": converted_doc } - formatted_docs.append(doc_dict) - response.append(formatted_docs) + + if request.return_scores: + doc_dict["score"] = float(score) + + combined.append(doc_dict) + + response.append(combined) processing_time = time.time() - start_time @@ -640,50 +640,75 @@ async def get_stats(): raise HTTPException(status_code=500, detail=str(e)) # ============================================================================ -# Main Function +# Main Function (for running with python directly) # ============================================================================ if __name__ == "__main__": parser = argparse.ArgumentParser(description="Launch FlashRAG-style server") + # CUDA device argument + parser.add_argument("--cuda_device", type=str, default=None, help="CUDA device to use (e.g., '0', '1', '0,1')") + # Retriever arguments - parser.add_argument("--index_path", type=str, default="data/corpus/e5_Flat.index", help="Corpus indexing file.") - parser.add_argument("--corpus_path", type=str, default="data/corpus/corpus.jsonl", help="Local corpus file.") - parser.add_argument("--retrieval_topk", type=int, default=25, help="Number of retrieved passages for one query.") - parser.add_argument("--retriever_name", type=str, default="e5", help="Name of the retriever model.") - parser.add_argument("--retriever_model", type=str, default="intfloat/e5-base-v2", help="Path of the retriever model.") - parser.add_argument('--faiss_gpu', action='store_true', default=True, help='Use GPU for computation') + parser.add_argument("--index_path", type=str, default=None, help="Corpus indexing file.") + parser.add_argument("--corpus_path", type=str, default=None, help="Local corpus file.") + parser.add_argument("--retrieval_topk", type=int, default=None, help="Number of retrieved passages for one query.") + parser.add_argument("--retriever_name", type=str, default=None, help="Name of the retriever model.") + parser.add_argument("--retriever_model", type=str, default=None, help="Path of the retriever model.") + parser.add_argument('--faiss_gpu', action='store_true', default=None, help='Use GPU for computation') # Reranker arguments - parser.add_argument("--reranking_topk", type=int, default=10, help="Number of reranked passages for one query.") - parser.add_argument("--reranker_model", type=str, default="cross-encoder/ms-marco-MiniLM-L12-v2", help="Path of the reranker model.") - parser.add_argument("--reranker_batch_size", type=int, default=64, help="Batch size for the reranker inference.") + parser.add_argument("--reranking_topk", type=int, default=None, help="Number of reranked passages for one query.") + parser.add_argument("--reranker_model", type=str, default=None, help="Path of the reranker model.") + parser.add_argument("--reranker_batch_size", type=int, default=None, help="Batch size for the reranker inference.") # Server arguments parser.add_argument("--host", type=str, default="0.0.0.0", help="Server host") parser.add_argument("--port", type=int, default=2223, help="Server port") - args = parser.parse_args() + cmd_args = parser.parse_args() + + # Override environment variables with command line arguments if provided + if cmd_args.cuda_device is not None: + os.environ["CUDA_VISIBLE_DEVICES"] = cmd_args.cuda_device + if cmd_args.index_path is not None: + os.environ["INDEX_PATH"] = cmd_args.index_path + if cmd_args.corpus_path is not None: + os.environ["CORPUS_PATH"] = cmd_args.corpus_path + if cmd_args.retrieval_topk is not None: + os.environ["RETRIEVAL_TOPK"] = str(cmd_args.retrieval_topk) + if cmd_args.retriever_name is not None: + os.environ["RETRIEVER_NAME"] = cmd_args.retriever_name + if cmd_args.retriever_model is not None: + os.environ["RETRIEVER_MODEL"] = cmd_args.retriever_model + if cmd_args.faiss_gpu is not None: + os.environ["FAISS_GPU"] = str(cmd_args.faiss_gpu).lower() + if cmd_args.reranking_topk is not None: + os.environ["RERANKING_TOPK"] = str(cmd_args.reranking_topk) + if cmd_args.reranker_model is not None: + os.environ["RERANKER_MODEL"] = cmd_args.reranker_model + if cmd_args.reranker_batch_size is not None: + os.environ["RERANKER_BATCH_SIZE"] = str(cmd_args.reranker_batch_size) - # Initialize configurations + # Reinitialize configs with updated environment variables retriever_config = RetrieverConfig( - retrieval_method=args.retriever_name, - index_path=args.index_path, - corpus_path=args.corpus_path, - retrieval_topk=args.retrieval_topk, - faiss_gpu=args.faiss_gpu, - retrieval_model_path=args.retriever_model, - retrieval_pooling_method="mean", - retrieval_query_max_length=256, - retrieval_use_fp16=True, - retrieval_batch_size=128, + retrieval_method=os.environ.get("RETRIEVER_NAME", "e5"), + index_path=os.environ.get("INDEX_PATH", "data/corpus/e5_Flat.index"), + corpus_path=os.environ.get("CORPUS_PATH", "data/corpus/corpus.jsonl"), + retrieval_topk=int(os.environ.get("RETRIEVAL_TOPK", "35")), # 35 candidates for reranking + faiss_gpu=os.environ.get("FAISS_GPU", "true").lower() == "true", + retrieval_model_path=os.environ.get("RETRIEVER_MODEL", "intfloat/e5-base-v2"), + retrieval_pooling_method=os.environ.get("RETRIEVAL_POOLING_METHOD", "mean"), + retrieval_query_max_length=int(os.environ.get("RETRIEVAL_QUERY_MAX_LENGTH", "256")), + retrieval_use_fp16=os.environ.get("RETRIEVAL_USE_FP16", "true").lower() == "true", + retrieval_batch_size=int(os.environ.get("RETRIEVAL_BATCH_SIZE", "128")), ) reranker_config = RerankerConfig( - rerank_topk=args.reranking_topk, - rerank_model_name_or_path=args.reranker_model, - batch_size=args.reranker_batch_size, + rerank_topk=int(os.environ.get("RERANKING_TOPK", "10")), # Top 10 after reranking + rerank_model_name_or_path=os.environ.get("RERANKER_MODEL", "cross-encoder/ms-marco-MiniLM-L12-v2"), + batch_size=int(os.environ.get("RERANKER_BATCH_SIZE", "64")), ) # Launch the server - uvicorn.run(app, host=args.host, port=args.port) + uvicorn.run(app, host=cmd_args.host, port=cmd_args.port) \ No newline at end of file diff --git a/examples/vllm_multiturn/rag_setup/rag_server.sh b/examples/vllm_multiturn/rag_setup/rag_server.sh index fdcbda0dabe..33765f4314d 100644 --- a/examples/vllm_multiturn/rag_setup/rag_server.sh +++ b/examples/vllm_multiturn/rag_setup/rag_server.sh @@ -1,7 +1,7 @@ -corpus_file=data/corpus/corpus.jsonl # jsonl -save_dir=data/corpus -retriever_name=e5 # this is for indexing naming -retriever_model=intfloat/e5-base-v2 +# corpus_file=/mnt/nas/bachvd/Code-Agent/verl/data/searchR1_processed_direct/data/wiki-18.jsonl # jsonl +# save_dir=data/corpus +# retriever_name=e5 # this is for indexing naming +# retriever_model=intfloat/e5-base-v2 echo "Starting FlashRAG server..." # python flashrag_server.py \ @@ -18,12 +18,12 @@ echo "Starting FlashRAG server..." # --faiss_gpu \ # --workers 64 -export INDEX_PATH=$save_dir/${retriever_name}_Flat.index -export CORPUS_PATH=$corpus_file -export RETRIEVER_NAME=$retriever_name -export RETRIEVER_MODEL=$retriever_model +export CUDA_DEVICE=0 +export INDEX_PATH="/mnt/nas/bachvd/Code-Agent/verl/data/janv2_searchr1/data/wiki-18_e5.index" +export CORPUS_PATH="/mnt/nas/bachvd/Code-Agent/verl/data/janv2_searchr1/data/wiki-18.jsonl" +export RETRIEVAL_TOPK=200 +export RERANKING_TOPK=10 +export BM25_CACHE_DIR="/mnt/nas/bachvd/Code-Agent/verl/data/janv2_searchr1/data/bm25_cache" +export FAISS_GPU=false -uvicorn flashrag_server:app \ - --host 0.0.0.0 \ - --port 3030 \ - --workers 64 \ No newline at end of file +uvicorn flashrag_server:app --host 0.0.0.0 --port 3030 --workers 1 \ No newline at end of file diff --git a/examples/vllm_multiturn/rag_setup/utils.py b/examples/vllm_multiturn/rag_setup/utils.py index cdeb5c304d5..f687a8060b4 100644 --- a/examples/vllm_multiturn/rag_setup/utils.py +++ b/examples/vllm_multiturn/rag_setup/utils.py @@ -1,7 +1,35 @@ # preprocess_corpus.py import json import re +import os from pathlib import Path +import shutil +import datasets +import logging +import uuid +from collections import defaultdict +import pickle + +logger = logging.getLogger(__name__) + + +class DocumentStore: + def __init__(self, id_to_title, title_to_content): + self.id_to_title = id_to_title + self.title_to_content = title_to_content + + def get(self, chunk_id): + title = self.id_to_title.get(chunk_id) + if title is None: + return None + + return { + 'title': title, + 'full_contents': self.title_to_content.get(title, "") + } + + def __contains__(self, chunk_id): + return chunk_id in self.id_to_title def clean_text(text): """Clean and normalize text""" @@ -39,5 +67,125 @@ def preprocess_corpus(input_file, output_file): print(f"Error parsing line {line_num}") continue +def build_and_save_index(corpus_path: str, save_dir: str, temp_dir: str = "temp_shards_storage", num_proc: int = 64): + + def create_title(batch): + titles = [] + for content in batch['contents']: + title = content.split("\n")[0].strip('"') + titles.append(title) + return {'title': titles} + + def clean_content(batch): + contents = [] + for content in batch['contents']: + parts = content.split("\n") + if len(parts) > 1: + content_new = "\n".join(parts[1:]) + else: + content_new = "" + contents.append(content_new) + return {'contents': contents} + + def create_group_shards(batch, rank): + local_group = defaultdict(list) + iterator = zip(batch['id'], batch['title'], batch['contents']) + + for id, title, content in iterator: + local_group[title].append((id, content)) + + unique_name = f"{rank}_{uuid.uuid4().hex}.json" + if not os.path.exists(temp_dir): + os.makedirs(temp_dir, exist_ok=True) + + file_path = os.path.join(temp_dir, unique_name) + + with open(file_path, 'w') as f: + json.dump(local_group, f) + + return batch + + print(f"Building index from: {corpus_path}") + + db = datasets.load_dataset('json', data_files=corpus_path, split="train", num_proc=16) + db = db.map(create_title, batch_size=10000, num_proc=num_proc, batched=True) + db = db.map(clean_content, batch_size=10000, num_proc=num_proc, batched=True) + + if not os.path.exists(temp_dir): + os.makedirs(temp_dir) + + db.map( + create_group_shards, + batched=True, + batch_size=10000, + num_proc=num_proc, + with_rank=True + ) + + raw_groups = defaultdict(list) + shard_files = os.listdir(temp_dir) + + for filename in shard_files: + file_path = os.path.join(temp_dir, filename) + try: + with open(file_path, 'r') as f: + shard_data = json.load(f) + + for title, items in shard_data.items(): + raw_groups[title].extend(items) + except Exception as e: + logger.error(f"Error reading shard {filename}: {e}") + + id_to_title = {} + title_to_content = {} + + for title, chunks in raw_groups.items(): + chunks.sort(key=lambda x: x[0]) + full_contents = " ".join([item[1].strip() for item in chunks]) + title_to_content[title] = full_contents + + for chunk_id, _ in chunks: + id_to_title[chunk_id] = title + + if not os.path.exists(save_dir): + os.makedirs(save_dir) + + print(f"Saving artifacts to {save_dir}...") + + with open(os.path.join(save_dir, "id_to_title.pkl"), "wb") as f: + pickle.dump(id_to_title, f) + + with open(os.path.join(save_dir, "title_to_content.pkl"), "wb") as f: + pickle.dump(title_to_content, f) + + print("Build complete.") + +def load_index(save_dir: str): + print(f"Loading index from {save_dir}...") + + with open(os.path.join(save_dir, "id_to_title.pkl"), "rb") as f: + id_to_title = pickle.load(f) + + with open(os.path.join(save_dir, "title_to_content.pkl"), "rb") as f: + title_to_content = pickle.load(f) + + print(f"Loaded. IDs: {len(id_to_title)}, Documents: {len(title_to_content)}") + return DocumentStore(id_to_title, title_to_content) + if __name__ == "__main__": - preprocess_corpus('data/raw_corpus.jsonl', 'data/corpus/processed_corpus.jsonl') \ No newline at end of file + DATA_PATH = "/mnt/nas/bachvd/Code-Agent/verl/data/janv2_searchr1/data/wiki-18.jsonl" + INDEX_DIR = "./saved_index_data" + + # build_and_save_index(DATA_PATH, INDEX_DIR, num_proc=64) + + if os.path.exists(INDEX_DIR): + doc_store = load_index(INDEX_DIR) + + # Test + test_id = "10" + result = doc_store.get(test_id) + print(result) + if result: + print(f"Retrieved: {result['title']}") + else: + print("Index not found. Please run build_and_save_index first.") \ No newline at end of file