juanwisz/modernbert-python-code-retrieval
1
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 