Team Ai
Apppublic

juanwisz/modernbert-python-code-retrieval

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
1likes
app.py90 linesDownload Raw Back to root
1import gradio as gr2import torch3import numpy as np4 5from sentence_transformers import SentenceTransformer, util6 7# 1. Load your fine-tuned retrieval model (on CodeSearchNet - Python)8#    This is the model you pushed to the Hugging Face Hub after training.9model_name = "juanwisz/modernbert-python-code-retrieval"10device = "cuda" if torch.cuda.is_available() else "cpu"11 12# SentenceTransformer automatically handles tokenizer + embedding13embedding_model = SentenceTransformer(model_name, device=device)14 15# 2. Define a function to:16#    - Parse code snippets from the text box (split by "---")17#    - Compute embeddings for the user’s query and each snippet18#    - Return the top 3 most relevant code snippets based on cosine similarity19def retrieve_top_snippets(query, code_input):20    # Split the code snippets by "---"21    # Each snippet is trimmed for cleanliness22    snippets = [s.strip() for s in code_input.split("---") if s.strip()]23 24    # Edge-case: if user provided no code, just return25    if len(snippets) == 0:26        return "No code snippets detected (make sure to separate them with ---)."27 28    # Embed the query and code snippets29    query_emb = embedding_model.encode(query, convert_to_tensor=True)30    snippets_emb = embedding_model.encode(snippets, convert_to_tensor=True)31 32    # Compute cosine similarities [batch_size x 1] with all code snippets33    cos_scores = util.cos_sim(query_emb, snippets_emb)[0]34 35    # Sort results by decreasing score36    # argsort(descending) means the first indices are the most relevant37    top_indices = torch.topk(cos_scores, k=min(3, len(snippets))).indices38 39    # Prepare text output with top 3 matches40    results = []41    for idx in top_indices:42        score = cos_scores[idx].item()43        snippet_text = snippets[idx]44        results.append(f"**Score**: {score:.4f}\n```python\n{snippet_text}\n```")45 46    # Join all results nicely47    return "\n\n".join(results)48 49 50#####################51### Gradio Layout ###52#####################53css = """54#container {55    margin: 0 auto;56    max-width: 700px;57}58"""59 60with gr.Blocks(css=css) as demo:61    gr.Markdown("# Code Retrieval using ModernBERT\n"62                "Enter a natural language query and paste multiple Python code snippets, "63                "delimited by `---`. We'll return the top 3 matches.")64 65    with gr.Column(elem_id="container"):66        with gr.Row():67            query_input = gr.Textbox(68                label="Natural Language Query",69                placeholder="What does your function do? e.g., 'Parse JSON from a string'"70            )71 72        code_snippets_input = gr.Textbox(73            label="Paste Python functions (delimited by ---)",74            lines=10,75            placeholder="Example:\n---\ndef parse_json(data):\n    return json.loads(data)\n---\ndef add_numbers(a, b):\n    return a + b\n---"76        )77 78        search_btn = gr.Button("Search", variant="primary")79        results_output = gr.Markdown(label="Top 3 Matches")80 81        # On click, run our retrieval function82        search_btn.click(83            fn=retrieve_top_snippets,84            inputs=[query_input, code_snippets_input],85            outputs=results_output86        )87 88if __name__ == "__main__":89    demo.launch()90