KernelPilot/KernelPilot-Optimization
0
1import os, tempfile, time2import gradio as gr3from tool.test import run_autotune_pipeline, DATA_DIR4 5# ---------- Core callback ----------6 7def get_test_text(test_file, test_data_input):8 if test_file is not None:9 if hasattr(test_file, "read"):10 return test_file.read().decode("utf-8")11 elif hasattr(test_file, "data"):12 return test_file.data if isinstance(test_file.data, str) else test_file.data.decode("utf-8")13 elif hasattr(test_file, "name") and os.path.exists(test_file.name):14 with open(test_file.name, "r", encoding="utf-8") as f:15 return f.read()16 # fallback to textbox17 return test_data_input or ""18 19def generate_kernel(text_input, test_data_input, test_file, n_iters, progress=gr.Progress()):20 """21 text_input : string from textbox (NL description or base CUDA code)22 test_data_input: test data (variable name, data)23 file_input : gr.File upload object (or None)24 Returns : (kernel_code_str, downloadable_file_path)25 """26 progress((0, n_iters), desc="Initializing...")27 # 1) Select input source28 29 if not text_input.strip():30 return "⚠️ Please paste a description or baseline CUDA code."31 32 # td = tempfile.mkdtemp(prefix="auto_")33 34 # # ------- select test data source -------35 # if test_file is not None and test_file.size > 0:36 # test_text = test_file.read().decode("utf-8")37 # elif test_data_input.strip():38 # test_text = test_data_input39 # else:40 # return "Test data required: either fill Test Data Input or upload a .txt file.", "", None41 42 # src_path = os.path.join(td, f"input_{int(time.time())}.txt")43 # test_path = os.path.join(td, f"test_data_{int(time.time())}.txt")44 45 # with open(src_path, "w") as f:46 # f.write(text_input)47 48 # with open(test_path, "w") as f:49 # f.write(test_data_input or "")50 51 # if test_file is not None:52 # test_text = test_file.read().decode("utf-8")53 # else:54 # test_text = test_data_input55 56 test_text = get_test_text(test_file, test_data_input)57 58 if not test_text.strip():59 return "⚠️ Test data required."60 61 62 best_code = ""63 for info in run_autotune_pipeline(64 input_code=text_input,65 test_data_input=test_text,66 test_file=None,67 bin_dir=DATA_DIR,68 max_iterations=int(n_iters)69 ):70 # 1) update progress bar (if iteration known)71 if info["iteration"] is not None:72 # print(f"Iteration {info['iteration']} / {n_iters}: {info['message']}")73 progress((info["iteration"], n_iters), desc=info["message"])74 75 # 3) kernel output only when we get new code76 if info["code"]:77 best_code = info["code"]78 79 # TBD: download button80 return best_code81 82 83# ---------- Gradio UI ----------84with gr.Blocks(85 title="KernelPilot", 86 theme=gr.themes.Soft(87 text_size="lg", 88 font=[89 "system-ui",90 "-apple-system",91 "BlinkMacSystemFont",92 "Segoe UI",93 "Roboto",94 "Helvetica Neue",95 "Arial",96 "Noto Sans",97 "sans-serif"98 ])) as demo:99 gr.Markdown(100 """# 🚀 KernelPilot Optimizer 101Enter a code, test data, then click **Generate** to obtain the optimized kernel function."""102 )103 104 with gr.Row():105 txt_input = gr.Textbox(106 label="📝 Input",107 lines=10,108 placeholder="Enter the code",109 scale=3110 )111 level = gr.Number(112 label="Optimazation Level",113 minimum=1,114 maximum=5,115 value=5,116 step=1,117 scale=1118 )119 120 with gr.Row():121 test_data_input = gr.Textbox(122 label="Test Data Input",123 lines=10,124 placeholder="<number_of_test_cases>\n<number_of_variables>\n\n<variable_1_name>\n<variable_1_testcase_1_data>\n<variable_1_testcase_2_data>\n...\n<variable_1_testcase_N_data>\n\n<variable_2_name>\n<variable_2_testcase_1_data>\n...\n<variable_2_testcase_N_data>\n\n...",125 scale=2126 )127 test_file = gr.File(128 label="Upload Test Data (.txt)",129 file_types=["text"],130 scale=1131 )132 133 gen_btn = gr.Button("⚡ Generate")134 135 kernel_output = gr.Code(136 label="🎯 Tuned CUDA Kernel",137 language="cpp"138 )139 140 gen_btn.click(141 fn=generate_kernel,142 inputs=[txt_input, test_data_input, test_file, level],143 outputs=[kernel_output],144 queue=True, # keeps requests queued145 show_progress=True, # show progress bar146 show_progress_on=kernel_output # update log box with progress147 )148 149if __name__ == "__main__":150 demo.queue(default_concurrency_limit=1, max_size=50)151 demo.launch(server_name="0.0.0.0", server_port=7860)