kernel-lens 0.1.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- kernel_lens-0.1.0/LICENSE +21 -0
- kernel_lens-0.1.0/PKG-INFO +47 -0
- kernel_lens-0.1.0/README.md +0 -0
- kernel_lens-0.1.0/kernel_lens/__init__.py +11 -0
- kernel_lens-0.1.0/kernel_lens/backends/__init__.py +0 -0
- kernel_lens-0.1.0/kernel_lens/backends/builder.py +129 -0
- kernel_lens-0.1.0/kernel_lens/backends/ort_gen.py +266 -0
- kernel_lens-0.1.0/kernel_lens/backends/trt_gen.py +403 -0
- kernel_lens-0.1.0/kernel_lens/compiler/__init__.py +0 -0
- kernel_lens-0.1.0/kernel_lens/compiler/ast_analyzer.py +50 -0
- kernel_lens-0.1.0/kernel_lens/compiler/core.py +141 -0
- kernel_lens-0.1.0/kernel_lens/compiler/interaction.py +41 -0
- kernel_lens-0.1.0/kernel_lens/compiler/manifest.py +72 -0
- kernel_lens-0.1.0/kernel_lens/compiler/onnx_exporter.py +101 -0
- kernel_lens-0.1.0/kernel_lens/compiler/tracer.py +182 -0
- kernel_lens-0.1.0/kernel_lens/runtime/__init__.py +0 -0
- kernel_lens-0.1.0/kernel_lens/runtime/engine.py +167 -0
- kernel_lens-0.1.0/kernel_lens/utils/__init__.py +0 -0
- kernel_lens-0.1.0/kernel_lens/utils/deployment.py +61 -0
- kernel_lens-0.1.0/kernel_lens/utils/end_check.py +43 -0
- kernel_lens-0.1.0/kernel_lens.egg-info/PKG-INFO +47 -0
- kernel_lens-0.1.0/kernel_lens.egg-info/SOURCES.txt +25 -0
- kernel_lens-0.1.0/kernel_lens.egg-info/dependency_links.txt +1 -0
- kernel_lens-0.1.0/kernel_lens.egg-info/requires.txt +13 -0
- kernel_lens-0.1.0/kernel_lens.egg-info/top_level.txt +1 -0
- kernel_lens-0.1.0/pyproject.toml +40 -0
- kernel_lens-0.1.0/setup.cfg +4 -0
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2026 Duc Justin
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|
|
@@ -0,0 +1,47 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: kernel-lens
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: A production-grade, multi-backend compiler for Triton kernels (ONNX & TensorRT)
|
|
5
|
+
Author: Compiler Engineer
|
|
6
|
+
License: MIT License
|
|
7
|
+
|
|
8
|
+
Copyright (c) 2026 Duc Justin
|
|
9
|
+
|
|
10
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
11
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
12
|
+
in the Software without restriction, including without limitation the rights
|
|
13
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
14
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
15
|
+
furnished to do so, subject to the following conditions:
|
|
16
|
+
|
|
17
|
+
The above copyright notice and this permission notice shall be included in all
|
|
18
|
+
copies or substantial portions of the Software.
|
|
19
|
+
|
|
20
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
21
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
22
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
23
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
24
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
25
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
26
|
+
SOFTWARE.
|
|
27
|
+
|
|
28
|
+
Project-URL: Source Code, https://github.com/yourusername/kernel-lens
|
|
29
|
+
Project-URL: Bug Tracker, https://github.com/yourusername/kernel-lens/issues
|
|
30
|
+
Classifier: Programming Language :: Python :: 3
|
|
31
|
+
Classifier: License :: OSI Approved :: MIT License
|
|
32
|
+
Classifier: Operating System :: POSIX :: Linux
|
|
33
|
+
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
|
|
34
|
+
Requires-Python: >=3.9
|
|
35
|
+
Description-Content-Type: text/markdown
|
|
36
|
+
License-File: LICENSE
|
|
37
|
+
Requires-Dist: torch>=2.1.0
|
|
38
|
+
Requires-Dist: triton>=2.1.0
|
|
39
|
+
Requires-Dist: onnx>=1.15.0
|
|
40
|
+
Provides-Extra: ort
|
|
41
|
+
Requires-Dist: onnxruntime-gpu>=1.17.0; extra == "ort"
|
|
42
|
+
Provides-Extra: trt
|
|
43
|
+
Requires-Dist: tensorrt>=8.6.0; extra == "trt"
|
|
44
|
+
Provides-Extra: all
|
|
45
|
+
Requires-Dist: onnxruntime-gpu>=1.17.0; extra == "all"
|
|
46
|
+
Requires-Dist: tensorrt>=8.6.0; extra == "all"
|
|
47
|
+
Dynamic: license-file
|
|
File without changes
|
|
File without changes
|
|
@@ -0,0 +1,129 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import subprocess
|
|
3
|
+
import urllib.request
|
|
4
|
+
import tarfile
|
|
5
|
+
import onnxruntime
|
|
6
|
+
|
|
7
|
+
def build_ort_plugin(ort_plugins_dir: str, cache_dir: str):
|
|
8
|
+
"""
|
|
9
|
+
Natively compiles the generated C++ files into a Shared Library (.so)
|
|
10
|
+
replacing the need for an external bash script.
|
|
11
|
+
"""
|
|
12
|
+
# 1. Get exact ORT version from the current Python environment
|
|
13
|
+
ort_version = onnxruntime.__version__.split('+')[0]
|
|
14
|
+
# print(f" [Builder] Detected ONNX Runtime v{ort_version}")
|
|
15
|
+
|
|
16
|
+
ort_release_dir = os.path.join(cache_dir, f"onnxruntime-linux-x64-gpu-{ort_version}")
|
|
17
|
+
ort_tgz = f"{ort_release_dir}.tgz"
|
|
18
|
+
|
|
19
|
+
# 2. Download exact matching C++ Developer Release if missing
|
|
20
|
+
if not os.path.exists(ort_release_dir):
|
|
21
|
+
url = f"https://github.com/microsoft/onnxruntime/releases/download/v{ort_version}/onnxruntime-linux-x64-gpu-{ort_version}.tgz"
|
|
22
|
+
# print(f" [Builder] Downloading ORT C++ headers from {url}...")
|
|
23
|
+
urllib.request.urlretrieve(url, ort_tgz)
|
|
24
|
+
with tarfile.open(ort_tgz, "r:gz") as tar:
|
|
25
|
+
tar.extractall(path=cache_dir)
|
|
26
|
+
|
|
27
|
+
ort_inc = os.path.join(ort_release_dir, "include")
|
|
28
|
+
ort_lib = os.path.join(ort_release_dir, "lib")
|
|
29
|
+
|
|
30
|
+
# # 3. Dynamically find CUDA paths via nvcc
|
|
31
|
+
# print(" [Builder] Querying system for CUDA configuration...")
|
|
32
|
+
try:
|
|
33
|
+
nvcc_path = subprocess.check_output(["which", "nvcc"]).decode().strip()
|
|
34
|
+
cuda_home = os.path.dirname(os.path.dirname(nvcc_path))
|
|
35
|
+
except Exception:
|
|
36
|
+
# Fallback to standard Linux path
|
|
37
|
+
cuda_home = "/usr/local/cuda"
|
|
38
|
+
|
|
39
|
+
cuda_inc = os.path.join(cuda_home, "include")
|
|
40
|
+
cuda_lib = os.path.join(cuda_home, "lib64")
|
|
41
|
+
# print(f" [Builder] Detected CUDA at {cuda_home}")
|
|
42
|
+
|
|
43
|
+
# 4. Compile the .cu files into object files
|
|
44
|
+
cu_files = [f for f in os.listdir(ort_plugins_dir) if f.endswith(".cu")]
|
|
45
|
+
obj_files = []
|
|
46
|
+
|
|
47
|
+
for cu_file in cu_files:
|
|
48
|
+
cu_path = os.path.join(ort_plugins_dir, cu_file)
|
|
49
|
+
obj_path = os.path.join(ort_plugins_dir, cu_file.replace(".cu", ".o"))
|
|
50
|
+
obj_files.append(obj_path)
|
|
51
|
+
|
|
52
|
+
cmd = [
|
|
53
|
+
"nvcc", "-c", cu_path, "-o", obj_path, "-O3", "-Xcompiler", "-fPIC",
|
|
54
|
+
f"-I{ort_inc}", f"-I{cuda_inc}"
|
|
55
|
+
]
|
|
56
|
+
subprocess.run(cmd, check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
|
|
57
|
+
|
|
58
|
+
# 5. Compile the registrar (register_ops.cpp)
|
|
59
|
+
reg_cpp = os.path.join(ort_plugins_dir, "register_ops.cpp")
|
|
60
|
+
reg_obj = os.path.join(ort_plugins_dir, "register_ops.o")
|
|
61
|
+
obj_files.append(reg_obj)
|
|
62
|
+
|
|
63
|
+
cmd = [
|
|
64
|
+
"g++", "-c", reg_cpp, "-o", reg_obj, "-O3", "-fPIC",
|
|
65
|
+
f"-I{ort_inc}", f"-I{cuda_inc}"
|
|
66
|
+
]
|
|
67
|
+
subprocess.run(cmd, check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
|
|
68
|
+
|
|
69
|
+
# 6. Link everything into the final .so library (with RPATH baked in)
|
|
70
|
+
so_path = os.path.join(ort_plugins_dir, "libtriton_ort_plugins.so")
|
|
71
|
+
abs_ort_lib = os.path.abspath(ort_lib)
|
|
72
|
+
|
|
73
|
+
cmd = [
|
|
74
|
+
"g++", "-shared", "-o", so_path
|
|
75
|
+
] + obj_files + [
|
|
76
|
+
f"-L{ort_lib}", "-lonnxruntime",
|
|
77
|
+
f"-L{cuda_lib}", "-lcuda", "-lcudart",
|
|
78
|
+
f"-Wl,-rpath,{abs_ort_lib}"
|
|
79
|
+
]
|
|
80
|
+
subprocess.run(cmd, check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
|
|
81
|
+
|
|
82
|
+
# print(f" [Builder] Compilation successful! Plugin saved to {so_path}")
|
|
83
|
+
|
|
84
|
+
def build_trt_plugin(trt_plugins_dir: str, cache_dir: str):
|
|
85
|
+
"""
|
|
86
|
+
Natively compiles the generated C++ files into a TensorRT Shared Library (.so).
|
|
87
|
+
"""
|
|
88
|
+
# print(" [Builder] Querying system for CUDA configuration...")
|
|
89
|
+
try:
|
|
90
|
+
nvcc_path = subprocess.check_output(["which", "nvcc"]).decode().strip()
|
|
91
|
+
cuda_home = os.path.dirname(os.path.dirname(nvcc_path))
|
|
92
|
+
except Exception:
|
|
93
|
+
cuda_home = "/usr/local/cuda"
|
|
94
|
+
|
|
95
|
+
cuda_inc = os.path.join(cuda_home, "include")
|
|
96
|
+
cuda_lib = os.path.join(cuda_home, "lib64")
|
|
97
|
+
|
|
98
|
+
# In TRT 8.6+, the library is often split into nvinfer and nvinfer_plugin
|
|
99
|
+
# We will assume standard system paths for TRT (/usr/lib/x86_64-linux-gnu or LD_LIBRARY_PATH)
|
|
100
|
+
|
|
101
|
+
cu_files = [f for f in os.listdir(trt_plugins_dir) if f.endswith(".cu")]
|
|
102
|
+
obj_files = []
|
|
103
|
+
|
|
104
|
+
for cu_file in cu_files:
|
|
105
|
+
cu_path = os.path.join(trt_plugins_dir, cu_file)
|
|
106
|
+
obj_path = os.path.join(trt_plugins_dir, cu_file.replace(".cu", ".o"))
|
|
107
|
+
obj_files.append(obj_path)
|
|
108
|
+
|
|
109
|
+
cmd = [
|
|
110
|
+
"nvcc", "-c", cu_path, "-o", obj_path, "-O3", "-Xcompiler", "-fPIC",
|
|
111
|
+
f"-I{cuda_inc}", "-Wno-deprecated-gpu-targets"
|
|
112
|
+
]
|
|
113
|
+
subprocess.run(cmd, check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
|
|
114
|
+
|
|
115
|
+
so_path = os.path.join(trt_plugins_dir, "libtriton_trt_plugins.so")
|
|
116
|
+
|
|
117
|
+
cmd = [
|
|
118
|
+
"g++", "-shared", "-o", so_path
|
|
119
|
+
] + obj_files + [
|
|
120
|
+
f"-L{cuda_lib}", "-lcuda", "-lcudart", "-lnvinfer"
|
|
121
|
+
]
|
|
122
|
+
|
|
123
|
+
try:
|
|
124
|
+
subprocess.run(cmd, check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
|
|
125
|
+
except subprocess.CalledProcessError as e:
|
|
126
|
+
print("\n[ERROR] TensorRT compilation failed. Ensure TensorRT is installed and in your LD_LIBRARY_PATH.")
|
|
127
|
+
raise e
|
|
128
|
+
|
|
129
|
+
# print(f" [Builder] TRT Compilation successful! Plugin saved to {so_path}")
|
|
@@ -0,0 +1,266 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import textwrap
|
|
3
|
+
import re
|
|
4
|
+
|
|
5
|
+
class ORTGenerator:
|
|
6
|
+
def __init__(self, manifests, ops_package="triton_custom"):
|
|
7
|
+
self.manifests = manifests
|
|
8
|
+
self.ops_package = ops_package
|
|
9
|
+
|
|
10
|
+
def _generate_kernel_h(self, manifest) -> str:
|
|
11
|
+
op_name = f"{manifest.kernel_name}Op"
|
|
12
|
+
kernel_name = f"{manifest.kernel_name}Kernel"
|
|
13
|
+
|
|
14
|
+
inputs_to_node = manifest.arguments
|
|
15
|
+
outputs_from_node = [a for a in manifest.arguments if a.kind == 'output']
|
|
16
|
+
|
|
17
|
+
input_types_cpp = []
|
|
18
|
+
mem_types_cpp = []
|
|
19
|
+
|
|
20
|
+
for a in inputs_to_node:
|
|
21
|
+
if a.kind == 'scalar':
|
|
22
|
+
# --- NEW: Keep Scalars on the CPU so we can safely read them! ---
|
|
23
|
+
mem_types_cpp.append("OrtMemTypeCPUInput")
|
|
24
|
+
if 'float' in a.dtype.lower():
|
|
25
|
+
input_types_cpp.append("ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE")
|
|
26
|
+
else:
|
|
27
|
+
input_types_cpp.append("ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64")
|
|
28
|
+
else:
|
|
29
|
+
# --- Keep Tensors on the GPU ---
|
|
30
|
+
mem_types_cpp.append("OrtMemTypeDefault")
|
|
31
|
+
input_types_cpp.append("ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT")
|
|
32
|
+
|
|
33
|
+
input_types_str = ",\n ".join(input_types_cpp)
|
|
34
|
+
mem_types_str = ",\n ".join(mem_types_cpp)
|
|
35
|
+
|
|
36
|
+
tpl = f'''#pragma once
|
|
37
|
+
// Unlock the Custom Op Initialization API
|
|
38
|
+
#define ORT_API_MANUAL_INIT
|
|
39
|
+
#include <onnxruntime_cxx_api.h>
|
|
40
|
+
#include <cuda.h>
|
|
41
|
+
|
|
42
|
+
namespace custom {{
|
|
43
|
+
|
|
44
|
+
struct {kernel_name} {{
|
|
45
|
+
void Compute(OrtKernelContext* context);
|
|
46
|
+
}};
|
|
47
|
+
|
|
48
|
+
struct {op_name} : Ort::CustomOpBase<{op_name}, {kernel_name}> {{
|
|
49
|
+
{op_name}() = default;
|
|
50
|
+
|
|
51
|
+
void* CreateKernel(const OrtApi& api, const OrtKernelInfo* info) const {{
|
|
52
|
+
return new {kernel_name}();
|
|
53
|
+
}}
|
|
54
|
+
|
|
55
|
+
const char* GetName() const {{ return "{manifest.kernel_name}"; }}
|
|
56
|
+
const char* GetExecutionProviderType() const {{ return "CUDAExecutionProvider"; }}
|
|
57
|
+
|
|
58
|
+
size_t GetInputTypeCount() const {{ return {len(inputs_to_node)}; }}
|
|
59
|
+
ONNXTensorElementDataType GetInputType(size_t index) const {{
|
|
60
|
+
static const ONNXTensorElementDataType types[] = {{
|
|
61
|
+
{input_types_str}
|
|
62
|
+
}};
|
|
63
|
+
return types[index];
|
|
64
|
+
}}
|
|
65
|
+
|
|
66
|
+
size_t GetOutputTypeCount() const {{ return {max(1, len(outputs_from_node))}; }}
|
|
67
|
+
ONNXTensorElementDataType GetOutputType(size_t index) const {{
|
|
68
|
+
return ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT;
|
|
69
|
+
}}
|
|
70
|
+
|
|
71
|
+
// --- THE MAGIC FIX: Dynamic Memory Placement ---
|
|
72
|
+
OrtMemType GetInputMemoryType(size_t index) const {{
|
|
73
|
+
static const OrtMemType types[] = {{
|
|
74
|
+
{mem_types_str}
|
|
75
|
+
}};
|
|
76
|
+
return types[index];
|
|
77
|
+
}}
|
|
78
|
+
|
|
79
|
+
OrtCustomOpInputOutputCharacteristic GetInputCharacteristic(size_t index) const {{
|
|
80
|
+
return OrtCustomOpInputOutputCharacteristic::INPUT_OUTPUT_REQUIRED;
|
|
81
|
+
}}
|
|
82
|
+
}};
|
|
83
|
+
|
|
84
|
+
}} // namespace custom
|
|
85
|
+
'''
|
|
86
|
+
return textwrap.dedent(tpl).strip()
|
|
87
|
+
|
|
88
|
+
def _generate_kernel_cu(self, manifest) -> str:
|
|
89
|
+
kernel_name = f"{manifest.kernel_name}Kernel"
|
|
90
|
+
|
|
91
|
+
arg_setup_lines = []
|
|
92
|
+
arg_setup_lines.append("Ort::KernelContext ctx(context);")
|
|
93
|
+
arg_setup_lines.append("auto ref_in = ctx.GetInput(0);")
|
|
94
|
+
arg_setup_lines.append("auto info = ref_in.GetTensorTypeAndShapeInfo();")
|
|
95
|
+
arg_setup_lines.append("std::vector<int64_t> dim_values = info.GetShape();")
|
|
96
|
+
|
|
97
|
+
import re
|
|
98
|
+
|
|
99
|
+
# --- ROBUST GRID EVALUATION ---
|
|
100
|
+
grid_strs = []
|
|
101
|
+
if hasattr(manifest, '_sym_grid_asts') and manifest._sym_grid_asts:
|
|
102
|
+
for g in manifest._sym_grid_asts:
|
|
103
|
+
# Always extract pure SymPy string FIRST
|
|
104
|
+
expr = str(g.node.expr) if hasattr(g, 'node') else str(g)
|
|
105
|
+
expr = re.sub(r'([a-zA-Z0-9_]+)\*\*([a-zA-Z0-9_]+)', r'std::pow(\1, \2)', expr)
|
|
106
|
+
expr = expr.replace("//", "/")
|
|
107
|
+
expr = re.sub(r'floor\((.*?)\)', r'(\1)', expr)
|
|
108
|
+
|
|
109
|
+
# Replace SymPy symbols with ORT C++
|
|
110
|
+
# Use regex to avoid replacing s0 inside s01
|
|
111
|
+
expr = re.sub(r'\bs0\b', "(int64_t)dim_values[0]", expr)
|
|
112
|
+
expr = re.sub(r'\bs1\b', "(int64_t)(dim_values.size() > 1 ? dim_values[1] : 1)", expr)
|
|
113
|
+
expr = re.sub(r'\bs2\b', "(int64_t)(dim_values.size() > 2 ? dim_values[2] : 1)", expr)
|
|
114
|
+
grid_strs.append(expr)
|
|
115
|
+
|
|
116
|
+
while len(grid_strs) < 3:
|
|
117
|
+
grid_strs.append("1")
|
|
118
|
+
|
|
119
|
+
grid_eval_lines = [
|
|
120
|
+
f"unsigned int grid_x = std::max(1u, (unsigned int)({grid_strs[0]}));",
|
|
121
|
+
f"unsigned int grid_y = std::max(1u, (unsigned int)({grid_strs[1]}));",
|
|
122
|
+
f"unsigned int grid_z = std::max(1u, (unsigned int)({grid_strs[2]}));"
|
|
123
|
+
]
|
|
124
|
+
|
|
125
|
+
arg_setup_lines.append("// Stable memory addresses for CUDA kernel launch")
|
|
126
|
+
for slot_idx, arg in enumerate(manifest.arguments):
|
|
127
|
+
if arg.kind == 'input':
|
|
128
|
+
arg_setup_lines.append(f"const void* arg_{slot_idx} = nullptr;")
|
|
129
|
+
elif arg.kind == 'scalar':
|
|
130
|
+
if 'float' in arg.dtype.lower():
|
|
131
|
+
arg_setup_lines.append(f"float arg_{slot_idx} = 0.0f;")
|
|
132
|
+
else:
|
|
133
|
+
arg_setup_lines.append(f"int32_t arg_{slot_idx} = 0;")
|
|
134
|
+
elif arg.kind == 'output':
|
|
135
|
+
arg_setup_lines.append(f"void* arg_{slot_idx} = nullptr;")
|
|
136
|
+
|
|
137
|
+
padding_idx = len(manifest.arguments)
|
|
138
|
+
arg_setup_lines.append(f"int64_t pad_{padding_idx} = 0;")
|
|
139
|
+
|
|
140
|
+
arg_setup_lines.append("std::vector<void*> kp;")
|
|
141
|
+
|
|
142
|
+
input_idx = 0
|
|
143
|
+
output_idx = 0
|
|
144
|
+
|
|
145
|
+
for slot_idx, arg in enumerate(manifest.arguments):
|
|
146
|
+
if arg.kind == 'input':
|
|
147
|
+
arg_setup_lines.append(f"auto input_{slot_idx} = ctx.GetInput({input_idx});")
|
|
148
|
+
arg_setup_lines.append(f"arg_{slot_idx} = (const void*)input_{slot_idx}.GetTensorData<float>();")
|
|
149
|
+
arg_setup_lines.append(f"kp.push_back(&arg_{slot_idx});")
|
|
150
|
+
input_idx += 1
|
|
151
|
+
elif arg.kind == 'scalar':
|
|
152
|
+
arg_setup_lines.append(f"auto input_{slot_idx} = ctx.GetInput({input_idx});")
|
|
153
|
+
# NOW THIS IS 100% SAFE because ONNX stored it in CPU memory!
|
|
154
|
+
if 'float' in arg.dtype.lower():
|
|
155
|
+
arg_setup_lines.append(f"arg_{slot_idx} = (float)(*input_{slot_idx}.GetTensorData<double>());")
|
|
156
|
+
else:
|
|
157
|
+
arg_setup_lines.append(f"arg_{slot_idx} = (int32_t)(*input_{slot_idx}.GetTensorData<int64_t>());")
|
|
158
|
+
|
|
159
|
+
arg_name = getattr(arg, 'name', '')
|
|
160
|
+
if 'stride' in arg_name.lower():
|
|
161
|
+
arg_setup_lines.append(f"if (arg_{slot_idx} != 1) kp.push_back(&arg_{slot_idx});")
|
|
162
|
+
else:
|
|
163
|
+
arg_setup_lines.append(f"kp.push_back(&arg_{slot_idx});")
|
|
164
|
+
input_idx += 1
|
|
165
|
+
elif arg.kind == 'output':
|
|
166
|
+
arg_setup_lines.append(f"// Skip PyTorch's dummy input")
|
|
167
|
+
arg_setup_lines.append(f"auto dummy_in_{slot_idx} = ctx.GetInput({input_idx});")
|
|
168
|
+
input_idx += 1
|
|
169
|
+
|
|
170
|
+
arg_setup_lines.append(f"auto output_{slot_idx} = ctx.GetOutput({output_idx}, dim_values.data(), dim_values.size());")
|
|
171
|
+
arg_setup_lines.append(f"arg_{slot_idx} = (void*)output_{slot_idx}.GetTensorMutableData<float>();")
|
|
172
|
+
arg_setup_lines.append(f"kp.push_back(&arg_{slot_idx});")
|
|
173
|
+
output_idx += 1
|
|
174
|
+
|
|
175
|
+
arg_setup_lines.append(f"kp.push_back(&pad_{padding_idx});")
|
|
176
|
+
|
|
177
|
+
dynamic_args_cpp = "\n ".join(arg_setup_lines)
|
|
178
|
+
grid_cpp = "\n ".join(grid_eval_lines)
|
|
179
|
+
|
|
180
|
+
tpl = f'''#include "{manifest.kernel_name}Op.h"
|
|
181
|
+
#include <stdexcept>
|
|
182
|
+
#include <iostream>
|
|
183
|
+
#include <vector>
|
|
184
|
+
#include <cmath>
|
|
185
|
+
#include <algorithm>
|
|
186
|
+
|
|
187
|
+
namespace custom {{
|
|
188
|
+
|
|
189
|
+
static const char* PTX_CODE = R"ptx(
|
|
190
|
+
{manifest.ptx}
|
|
191
|
+
)ptx";
|
|
192
|
+
|
|
193
|
+
void {kernel_name}::Compute(OrtKernelContext* context) {{
|
|
194
|
+
{dynamic_args_cpp}
|
|
195
|
+
|
|
196
|
+
static CUmodule mModule = nullptr;
|
|
197
|
+
static CUfunction mKernel = nullptr;
|
|
198
|
+
|
|
199
|
+
if (mModule == nullptr) {{
|
|
200
|
+
CUresult res = cuModuleLoadDataEx(&mModule, PTX_CODE, 0, nullptr, nullptr);
|
|
201
|
+
if (res != CUDA_SUCCESS) throw std::runtime_error("Failed to load PTX module");
|
|
202
|
+
res = cuModuleGetFunction(&mKernel, mModule, "{manifest.kernel_name}");
|
|
203
|
+
if (res != CUDA_SUCCESS) throw std::runtime_error("Failed to extract function");
|
|
204
|
+
}}
|
|
205
|
+
|
|
206
|
+
cudaStream_t stream = reinterpret_cast<cudaStream_t>(ctx.GetGPUComputeStream());
|
|
207
|
+
|
|
208
|
+
{grid_cpp}
|
|
209
|
+
unsigned int block_x = 128;
|
|
210
|
+
|
|
211
|
+
cuLaunchKernel(mKernel, grid_x, grid_y, grid_z, block_x, 1, 1, {manifest.shared_memory_bytes}, stream, kp.data(), nullptr);
|
|
212
|
+
}}
|
|
213
|
+
|
|
214
|
+
}} // namespace custom
|
|
215
|
+
'''
|
|
216
|
+
return textwrap.dedent(tpl).strip()
|
|
217
|
+
|
|
218
|
+
def generate(self, output_dir: str):
|
|
219
|
+
if not os.path.exists(output_dir):
|
|
220
|
+
os.makedirs(output_dir)
|
|
221
|
+
|
|
222
|
+
for m in self.manifests:
|
|
223
|
+
base_path = os.path.join(output_dir, f"{m.kernel_name}Op")
|
|
224
|
+
with open(f"{base_path}.h", "w") as f:
|
|
225
|
+
f.write(self._generate_kernel_h(m))
|
|
226
|
+
with open(f"{base_path}.cu", "w") as f:
|
|
227
|
+
f.write(self._generate_kernel_cu(m))
|
|
228
|
+
|
|
229
|
+
registrar_code = '''#define ORT_API_MANUAL_INIT
|
|
230
|
+
#include <onnxruntime_cxx_api.h>
|
|
231
|
+
|
|
232
|
+
'''
|
|
233
|
+
for m in self.manifests:
|
|
234
|
+
registrar_code += f'#include "{m.kernel_name}Op.h"\n'
|
|
235
|
+
|
|
236
|
+
registrar_code += '''
|
|
237
|
+
#ifndef ORT_EXPORT
|
|
238
|
+
#ifdef _WIN32
|
|
239
|
+
#define ORT_EXPORT __declspec(dllexport)
|
|
240
|
+
#else
|
|
241
|
+
#define ORT_EXPORT __attribute__((visibility("default")))
|
|
242
|
+
#endif
|
|
243
|
+
#endif
|
|
244
|
+
|
|
245
|
+
extern "C" {
|
|
246
|
+
ORT_EXPORT OrtStatus* ORT_API_CALL RegisterCustomOps(OrtSessionOptions* options, const OrtApiBase* api_base) {
|
|
247
|
+
Ort::InitApi(api_base->GetApi(ORT_API_VERSION));
|
|
248
|
+
static Ort::CustomOpDomain custom_domain("triton_custom");
|
|
249
|
+
'''
|
|
250
|
+
for m in self.manifests:
|
|
251
|
+
registrar_code += f' static custom::{m.kernel_name}Op c_{m.kernel_name};\n'
|
|
252
|
+
registrar_code += f' custom_domain.Add(&c_{m.kernel_name});\n'
|
|
253
|
+
|
|
254
|
+
registrar_code += '''
|
|
255
|
+
Ort::UnownedSessionOptions sess_options(options);
|
|
256
|
+
sess_options.Add(custom_domain);
|
|
257
|
+
return nullptr;
|
|
258
|
+
}
|
|
259
|
+
}
|
|
260
|
+
'''
|
|
261
|
+
with open(os.path.join(output_dir, "register_ops.cpp"), "w") as f:
|
|
262
|
+
f.write(registrar_code)
|
|
263
|
+
|
|
264
|
+
def generate_ort_bindings(manifests, output_path: str):
|
|
265
|
+
gen = ORTGenerator(manifests)
|
|
266
|
+
gen.generate(output_path)
|