ethanrom/retrieval-compression
0
1import streamlit as st2import pickle3import os4from langchain.embeddings import HuggingFaceEmbeddings5from langchain.vectorstores import FAISS6from langchain.retrievers import EnsembleRetriever, BM25Retriever, ContextualCompressionRetriever7from langchain.document_transformers import EmbeddingsRedundantFilter8from langchain.retrievers.document_compressors import EmbeddingsFilter, DocumentCompressorPipeline9from langchain.text_splitter import CharacterTextSplitter10 11 12from analysis import calculate_word_overlaps, calculate_duplication_rate, cosine_similarity_score, jaccard_similarity_score, display_similarity_results13 14 15 16with open("docs_data.pkl", "rb") as file:17 docs = pickle.load(file)18 19metadata_list = []20unique_metadata_list = []21seen = set()22 23embeddings = HuggingFaceEmbeddings()24vectorstore = FAISS.load_local("faiss_index", embeddings)25retriever = vectorstore.as_retriever(search_type="similarity")26 27splitter = CharacterTextSplitter(chunk_size=300, chunk_overlap=0, separator=". ")28redundant_filter = EmbeddingsRedundantFilter(embeddings=embeddings)29relevant_filter = EmbeddingsFilter(embeddings=embeddings, similarity_threshold=0.5)30pipeline_compressor = DocumentCompressorPipeline(31 transformers=[splitter, redundant_filter, relevant_filter]32)33 34bm25_retriever = BM25Retriever.from_texts(docs)35 36st.title("Document Retrieval App")37 38vecotstore_k = st.number_input("Set k value for Dense Retriever:", value=5, min_value=1, step=1)39bm25_k = st.number_input("Set k value for sparse Retriever:", value=2, min_value=1, step=1)40 41retriever.search_kwargs["k"] = vecotstore_k42bm25_retriever.k = bm25_k43 44compressed_retriever = ContextualCompressionRetriever(base_compressor=pipeline_compressor, base_retriever=retriever)45bm25_compression_retriever = ContextualCompressionRetriever(base_compressor=pipeline_compressor, base_retriever=bm25_retriever)46 47query = st.text_input("Enter a query:", "what is a horizontal conflict")48 49if st.button("Retrieve Documents"):50 51 compressed_ensemble_retriever = EnsembleRetriever(retrievers=[compressed_retriever, bm25_compression_retriever], weights=[0.5, 0.5])52 ensemble_retriever = EnsembleRetriever(retrievers=[retriever, bm25_retriever], weights=[0.5, 0.5])53 54 with st.expander("Retrieved Documents"):55 col1, col2 = st.columns(2)56 57 with col1:58 st.header("Without Compression")59 normal_results = ensemble_retriever.get_relevant_documents(query)60 for doc in normal_results:61 st.write(doc.page_content)62 st.write("---")63 64 with col2:65 st.header("With Compression")66 compressed_results = compressed_ensemble_retriever.get_relevant_documents(query)67 for doc in compressed_results:68 st.write(doc.page_content)69 st.write("---")70 71 if hasattr(doc, 'metadata'):72 metadata = doc.metadata73 metadata_list.append(metadata)74 75 for metadata in metadata_list:76 metadata_tuple = tuple(metadata.items())77 if metadata_tuple not in seen:78 unique_metadata_list.append(metadata)79 seen.add(metadata_tuple)80 81 st.write(unique_metadata_list)82 83 with st.expander("Analysis"):84 st.write("Analysis of Retrieval Results")85 86 total_words_normal = sum(len(doc.page_content.split()) for doc in normal_results)87 total_words_compressed = sum(len(doc.page_content.split()) for doc in compressed_results)88 reduction_percentage = ((total_words_normal - total_words_compressed) / total_words_normal) * 10089 90 col1, col2 = st.columns(2)91 92 93 st.write(f"Total words in documents (Normal): {total_words_normal}")94 st.write(f"Total words in documents (Compressed): {total_words_compressed}")95 st.write(f"Reduction Percentage: {reduction_percentage:.2f}%")96 97 average_word_overlap_normal = calculate_word_overlaps([doc.page_content for doc in normal_results], query)98 average_word_overlap_compressed = calculate_word_overlaps([doc.page_content for doc in compressed_results], query)99 100 duplication_rate_normal = calculate_duplication_rate([doc.page_content for doc in normal_results])101 duplication_rate_compressed = calculate_duplication_rate([doc.page_content for doc in compressed_results])102 103 cosine_scores_normal = cosine_similarity_score([doc.page_content for doc in normal_results], query)104 jaccard_scores_normal = jaccard_similarity_score([doc.page_content for doc in normal_results], query)105 106 cosine_scores_compressed = cosine_similarity_score([doc.page_content for doc in compressed_results], query)107 jaccard_scores_compressed = jaccard_similarity_score([doc.page_content for doc in compressed_results], query)108 109 with col1:110 st.subheader("Normal")111 112 st.write(f"Average Word Overlap: {average_word_overlap_normal:.2f}")113 st.write(f"Duplication Rate: {duplication_rate_normal:.2%}")114 115 st.write("Results without Compression:")116 display_similarity_results(cosine_scores_normal, jaccard_scores_normal, "") 117 118 with col2:119 st.subheader("Compressed")120 121 st.write(f"Average Word Overlap: {average_word_overlap_compressed:.2f}")122 st.write(f"Duplication Rate: {duplication_rate_compressed:.2%}") 123 124 st.write("Results with Compression:")125 display_similarity_results(cosine_scores_compressed, jaccard_scores_compressed, "")