Skip to content
Open
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
15 changes: 15 additions & 0 deletions kernels/src/kernels/importer.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import importlib
import logging
import sys
from dataclasses import dataclass
from pathlib import Path
Expand All @@ -7,6 +8,8 @@
from kernels._rust import Metadata
from kernels.hf_hub import RepoInfo

logger = logging.getLogger(__name__)


@dataclass(frozen=True)
class LoadedKernel:
Expand Down Expand Up @@ -71,6 +74,18 @@ def _import_from_path(
return loaded_kernel.module

metadata = Metadata.read_from_file(variant_path / "metadata.json")

# Kernel ids are unique per build: if this build was already imported
# reuse it instead of executing it again.
if (module := sys.modules.get(metadata.id)) is not None:
logging.debug(f"Kernel already loaded, skipping: {metadata.id}")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I think the debug message emitting behaviour should also be tested?

@danieldk danieldk Oct 9, 2026 •

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why? I'm not sure it's that important. I mean the debug message is already something that is purely informative. It's not something we really have to rely on that it is there, there is a lot of caching in a lot of places that we don't even log?

Testing logging seems mostly important for error/critical and maaaybe warning.

_loaded_kernels[variant_path] = LoadedKernel(
metadata=metadata,
module=module,
repo_info=repo_info,
)
return module
Comment thread
danieldk marked this conversation as resolved.

module_name = metadata.name.python_name

file_path = variant_path / "__init__.py"
Expand Down
62 changes: 62 additions & 0 deletions kernels/tests/test_importer.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,14 @@
import json
import sys
import types

import pytest

from kernels.importer import _import_from_path, _loaded_kernels

_EXEC_LOG_MODULE = "_kernels_test_exec_log"
_COUNTING_ID = "counting_1_cuda"


def _write_variant(tmp_path):
variant_dir = tmp_path / "build" / "torch28-cxx11-cu128-x86_64-linux"
Expand Down Expand Up @@ -35,3 +39,61 @@ def test_failed_import_cleans_up_sys_modules(tmp_path):
finally:
_loaded_kernels.pop(variant_dir, None)
sys.modules.pop("broken_1_cuda", None)


def _write_counting_variant(base_path, kernel_id):
"""Write a kernel variant that records each execution of its module."""
variant_dir = base_path / "build" / "torch28-cxx11-cu128-x86_64-linux"
variant_dir.mkdir(parents=True)
metadata = {
"id": kernel_id,
"name": "counting",
"version": 1,
"license": "Apache-2.0",
"python-depends": ["torch"],
"backend": {"type": "cuda"},
}
(variant_dir / "metadata.json").write_text(json.dumps(metadata))
(variant_dir / "__init__.py").write_text(f"import {_EXEC_LOG_MODULE}\n{_EXEC_LOG_MODULE}.calls.append(__file__)\n")
return variant_dir


@pytest.fixture
def exec_log(monkeypatch):
log = types.ModuleType(_EXEC_LOG_MODULE)
log.calls = [] # type: ignore[attr-defined]
monkeypatch.setitem(sys.modules, _EXEC_LOG_MODULE, log)
return log.calls # type: ignore[attr-defined]


def test_same_id_different_path_is_not_reloaded(tmp_path, exec_log):
first_dir = _write_counting_variant(tmp_path / "first", _COUNTING_ID)
second_dir = _write_counting_variant(tmp_path / "second", _COUNTING_ID)
try:
first = _import_from_path(first_dir, deps={})
second = _import_from_path(second_dir, deps={})

assert first is second

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Wish we could also do assert 2 + 2 is 5.

assert len(exec_log) == 1
assert _loaded_kernels[first_dir].module is first
assert _loaded_kernels[second_dir].module is first
finally:
_loaded_kernels.pop(first_dir, None)
_loaded_kernels.pop(second_dir, None)
sys.modules.pop(_COUNTING_ID, None)


def test_already_imported_kernel_is_reregistered(tmp_path, exec_log):
variant_dir = _write_counting_variant(tmp_path, _COUNTING_ID)
try:
first = _import_from_path(variant_dir, deps={})
_loaded_kernels.pop(variant_dir)

second = _import_from_path(variant_dir, deps={})

assert first is second
assert len(exec_log) == 1
assert _loaded_kernels[variant_dir].module is first
finally:
_loaded_kernels.pop(variant_dir, None)
sys.modules.pop(_COUNTING_ID, None)
Loading