Team Ai
Apppublic

KernelPilot/KernelPilot-Optimization

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
app.py151 linesDownload Raw Back to root
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)