CoolFace
Apppublic

ethanrom/retrieval-compression

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py125 linesDownload Raw Back to root
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, "")