Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion humming/jit/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,7 @@ def prepare(self):
disable_fast_math=self.disable_fast_math,
postprocess_cubin=self.postprocess_cubin,
)
kernel_name = jit_utils.find_kernel_name_in_cubin(kernel_filename, self.name)
kernel_name = jit_utils.find_kernel_name_in_cubin_cached(kernel_filename, self.name)
self.kernel_name = kernel_name
self.kernel_filename = kernel_filename
if threading.current_thread() is threading.main_thread():
Expand Down
28 changes: 28 additions & 0 deletions humming/utils/jit.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,34 @@ def read_symbol_value(filename, symbol_name, default_value=None):
return symbol_value


def find_kernel_name_in_cubin_cached(filename, func_keyword):
# Pure-python ELF symbol iteration takes seconds per cubin and holds the
# GIL, so warm-cache kernel loading is slower than cold compilation when
# many kernels load concurrently. Persist the resolved name next to the
# cubin and reuse it on later runs.
cache_filename = filename + ".name"
try:
with open(cache_filename) as f:
cached = f.read().strip()
except OSError:
cached = ""
if cached and re.match(f"^_Z\\d+{func_keyword}", cached):
return cached

kernel_name = find_kernel_name_in_cubin(filename, func_keyword)
tmp_filename = f"{cache_filename}.{os.getpid()}.{threading.get_ident()}.tmp"
try:
with open(tmp_filename, "w") as f:
f.write(kernel_name)
os.replace(tmp_filename, cache_filename)
except OSError:
try:
os.remove(tmp_filename)
except OSError:
pass
return kernel_name


def find_kernel_name_in_cubin(filename, func_keyword):
with open(filename, "rb") as f:
elffile = ELFFile(f)
Expand Down
Loading