KernelPilot/KernelPilot-V1-Server2
4
1import os, tempfile, time2import gradio as gr3from tool.testv3 import run_autotune_pipeline4 5# ---------- Core callback ----------6def generate_kernel(text_input, n_iters, progress=gr.Progress()):7 """8 text_input : string from textbox (NL description or base CUDA code)9 file_input : gr.File upload object (or None)10 Returns : (kernel_code_str, downloadable_file_path)11 """12 progress((0, n_iters), desc="Initializing...")13 # 1) Select input source14 15 if not text_input.strip():16 return "⚠️ Please paste a description or baseline CUDA code.", "", None17 18 td = tempfile.mkdtemp(prefix="auto_")19 src_path = os.path.join(td, f"input_{int(time.time())}.txt")20 with open(src_path, "w") as f:21 f.write(text_input)22 23 best_code = ""24 for info in run_autotune_pipeline(src_path, n_iters):25 # 1) update progress bar (if iteration known)26 if info["iteration"] is not None:27 # print(f"Iteration {info['iteration']} / {n_iters}: {info['message']}")28 progress((info["iteration"], n_iters), desc=info["message"])29 30 # 3) kernel output only when we get new code31 if info["code"]:32 best_code = info["code"]33 34 35 # last yield enables the download button36 return best_code37 38 39# ---------- Gradio UI ----------40with gr.Blocks(title="KernelPilot", theme=gr.themes.Soft(text_size="lg", font=[41 "system-ui",42 "-apple-system",43 "BlinkMacSystemFont",44 "Segoe UI",45 "Roboto",46 "Helvetica Neue",47 "Arial",48 "Noto Sans",49 "sans-serif"50 ])) as demo:51 gr.Markdown(52 """# 🚀 KernelPilot 53Enter a natural‑language description, 54then click **Generate** to obtain the kernel function."""55 )56 57 with gr.Row():58 txt_input = gr.Textbox(59 label="📝 Input",60 lines=10,61 placeholder="Describe the kernel",62 scale=363 )64 level = gr.Number(65 label="Optimization Level",66 minimum=1,67 maximum=5,68 value=2,69 step=1,70 scale=171 )72 73 74 gen_btn = gr.Button("⚡ Generate")75 76 kernel_output = gr.Code(77 label="🎯 Tuned CUDA Kernel",78 language="cpp"79 )80 81 gen_btn.click(82 fn=generate_kernel,83 inputs=[txt_input, level],84 outputs=[kernel_output],85 queue=True, # keeps requests queued86 show_progress=True, # show progress bar87 show_progress_on=kernel_output # update log box with progress88 )89 90if __name__ == "__main__":91 demo.queue(default_concurrency_limit=1, max_size=50)92 demo.launch(server_name="0.0.0.0", server_port=7860)