Team Ai
Apppublic

SwastikM/Embedding-Quantization

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes
app.py95 linesDownload Raw Back to root
1 2import gradio as gr3from datasets import load_from_disk4import pandas as pd5from sentence_transformers import SentenceTransformer6from sentence_transformers.quantization import quantize_embeddings7import faiss8from usearch.index import Index9import numpy as np10import os11 12base_path = os.getcwd()13full_path = os.path.join(base_path, 'conala')14conala_dataset = load_from_disk(full_path)15 16int8_view = Index.restore(os.path.join(base_path, 'conala_int8_usearch.index'), view=True)17binary_index: faiss.IndexBinaryFlat = faiss.read_index_binary(os.path.join(base_path, 'conala.index'))18 19model = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2")20 21def search(query, top_k: int = 20):22    # 1. Embed the query as float3223    query_embedding = model.encode(query)24 25    # 2. Quantize the query to ubinary. To perform actual search with faiss26    query_embedding_ubinary = quantize_embeddings(query_embedding.reshape(1, -1), "ubinary")27 28 29    # 3. Search the binary index 30    index =  binary_index31    _scores, binary_ids = index.search(query_embedding_ubinary, top_k)32    binary_ids = binary_ids[0]33 34 35    # 4. Load the corresponding int8 embeddings. To perform rescoring to calculate score of fetched documents.36    int8_embeddings = int8_view[binary_ids].astype(int)37 38    # 5. Rescore the top_k * rescore_multiplier using the float32 query embedding and the int8 document embeddings39    scores = query_embedding @ int8_embeddings.T40 41    # 6. Sort the scores and return the top_k42    indices = scores.argsort()[::-1][:top_k]43    top_k_indices = binary_ids[indices]44    top_k_scores = scores[indices]45 46    top_k_codes = conala_dataset[top_k_indices]47 48    return top_k_codes49 50 51def response_generator(user_prompt):52    top_k_outputs = search(user_prompt)53    probs = top_k_outputs['prob']54    snippets = top_k_outputs['snippet']55    idx = np.argsort(probs)[::-1]56    results = np.array(snippets)[idx]57    filtered_results = []58    for item in results:59        if len(filtered_results)<3:60            if item not in filtered_results:61                filtered_results.append(item)62 63    output_template = "User Query: {user_query}\nBelow are some examples of previous conversations.\nQuery: {query1} Solution: {solution1}\nQuery: {query2} Solution: {solution2}\nYou may use the above examples for reference only. Create your own solution and provide only the solution"64    output_template = "The top three most relevant code snippets from the database are:\n\n1. {snippet1}\n\n2. {snippet2}\n\n3. {snippet3}"65    output = f'{output_template.format(snippet1=filtered_results[0],snippet2=filtered_results[1],snippet3=filtered_results[2])}'66 67    return {output_box:output}  68 69 70with gr.Blocks() as demo:71    72    gr.Markdown(73    """74    # Embedding Quantization75 76    ## Quantized Semantic Search77 78    - ***Embedding:*** all-MiniLM-L6-v279    - ***Vetor DB:*** faiss, USearch80    - ***Vector_DB Size:*** `5,93,891`81 82    """)83 84    state_var = gr.State([])85 86 87    input_box = gr.Textbox(autoscroll=True,visible=True,label='User',info="Enter a query.",value="How to extract the n-th elements from a list of tuples in python?")88    output_box = gr.Textbox(autoscroll=True,max_lines=30,value="Output",label='Assistant')89    gr.Interface(fn=response_generator, inputs=[input_box], outputs=[output_box],90                 delete_cache=(20,10),91                 allow_flagging='never')92    93demo.queue()94demo.launch()95