Skip to content
Closed
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
8 changes: 5 additions & 3 deletions kernels/src/kernels/importer.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,13 +79,15 @@ def _import_from_path(
if not file_path.exists():
raise FileNotFoundError(f"No kernel module found at: `{variant_path}`")

spec = importlib.util.spec_from_file_location(metadata.id, file_path)
# Hub revisions can share a build ID but must have separate Python submodules.
import_name = f"{metadata.id}_{repo_info.revision}" if repo_info is not None else metadata.id

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

does repo info always include a revision, e.g. is a version resolved to a revision?

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.

Yup:

def resolve_kernel_version(repo_id: str, version: KernelVersion, *, local_files_only: bool) -> Oid:
"""Resolve a kernel version to the commit it refers to.
A `KernelVersion` can either be a version number or a revision (branch,
tag, or commit). This function resolves the version or revision into a
full Git commit SHA.
"""
if isinstance(version, KernelVersion.Version):
# `name` rather than `ref`: the cache names its refs `v1`, not
# `refs/heads/v1`.
version_ref = resolve_version_spec_as_ref(repo_id, version.version, local_files_only=local_files_only)
ref, commit = version_ref.name, Oid.from_str(version_ref.target_commit)
elif isinstance(version, KernelVersion.Revision):
ref, commit = version.revision, _resolve_ref(repo_id, version.revision, local_files_only=local_files_only)
else:
raise ValueError(f"Invalid version type: {version}")
if not local_files_only:
# Kernels are fetched by commit, since we need the commit hash for receipt
# validation, etc. However, that means that snapshot downloads do not create
# refs in the cache. This causes a kernel fetched by version/ref not to be
# found in offline mode. To work around this problem, create a ref ourselves.
#
# Note that this can create the situation where the ref exists, but no
# snapshot or an incomplete snapshot. However, this is fine for
# huggingface_hub, since it also writes the ref before downloading the
# snapshot:
#
# https://github.com/huggingface/huggingface_hub/blob/5a9cdda63f231a1b57a05eab88dc4357c790ba87/src/huggingface_hub/_snapshot_download.py#L426
_record_ref_in_cache(repo_id, ref, str(commit))
return commit

For local kernel, this is handled differently, i.e., no repo_info, of course.

@danieldk danieldk Oct 9, 2026 •

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.

We shouldn't do this. The identifier from the metadata should be unique. This will still not solve it in contexts where we do not have a revision.

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.

But it doesn't seem to be 👀

spec = importlib.util.spec_from_file_location(import_name, file_path)
if spec is None:
raise ImportError(f"Cannot load spec for {module_name} from {file_path}")
module = importlib.util.module_from_spec(spec)
if module is None:
raise ImportError(f"Cannot load module {module_name} from spec")
sys.modules[metadata.id] = module
sys.modules[import_name] = module

# Avoid an import cycle.
from kernels.deps import use_kernel_deps
Expand All @@ -96,7 +98,7 @@ def _import_from_path(
except Exception as e:
# Remove the partially initialized module, so that a retry
# imports from scratch.
sys.modules.pop(metadata.id, None)
sys.modules.pop(import_name, None)
if hasattr(e, "add_note"):
origin = f"({repo_info.repo_id}, revision: {repo_info.revision})" if repo_info else ""
e.add_note(f"while importing kernel '{metadata.name}', variant '{variant_path.name}' {origin}")
Expand Down
47 changes: 43 additions & 4 deletions kernels/tests/test_importer.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@

import pytest

from kernels._rust import Oid
from kernels.hf_hub import RepoInfo
from kernels.importer import _import_from_path, _loaded_kernels


Expand All @@ -22,16 +24,53 @@ def _write_variant(tmp_path):
return variant_dir


def test_failed_import_cleans_up_sys_modules(tmp_path):
@pytest.mark.parametrize("revision", [None, "a" * 40])
def test_failed_import_cleans_up_sys_modules(tmp_path, revision):
variant_dir = _write_variant(tmp_path)
repo_info = RepoInfo("test/broken", Oid.from_str(revision)) if revision else None
import_name = f"broken_1_cuda_{revision}" if revision else "broken_1_cuda"
_loaded_kernels.pop(variant_dir, None)
try:
with pytest.raises(RuntimeError, match="kernel is broken") as exc_info:
_import_from_path(variant_dir, deps={})
assert "broken_1_cuda" not in sys.modules
_import_from_path(variant_dir, deps={}, repo_info=repo_info)
assert import_name not in sys.modules
assert variant_dir not in _loaded_kernels
if sys.version_info >= (3, 11):
assert any("broken" in note for note in exc_info.value.__notes__)
finally:
_loaded_kernels.pop(variant_dir, None)
sys.modules.pop("broken_1_cuda", None)
sys.modules.pop(import_name, None)


def test_revisions_with_same_metadata_id_have_separate_submodules(tmp_path):
variants = []
modules = []
try:
for revision in ("a" * 40, "b" * 40):
variant_dir = _write_variant(tmp_path / revision)
variants.append(variant_dir)
(variant_dir / "__init__.py").write_text(
"from . import layers\ndef load_lazy():\n from . import lazy\n return lazy\n"
)
for filename in ("layers.py", "lazy.py"):
(variant_dir / filename).write_text(f"revision = {revision!r}\n")
repo_info = RepoInfo("test/broken", Oid.from_str(revision))
module = _import_from_path(variant_dir, deps={}, repo_info=repo_info)
modules.append(module)
assert _import_from_path(variant_dir, deps={}, repo_info=repo_info) is module

old, new = modules
assert old is not new
assert old.layers is not new.layers
assert old.layers.revision == "a" * 40
assert new.layers.revision == "b" * 40
# Imports made after both revisions are loaded must remain isolated too.
assert old.load_lazy().revision == "a" * 40
assert new.load_lazy().revision == "b" * 40
finally:
for variant_dir in variants:
_loaded_kernels.pop(variant_dir, None)
for module in modules:
for name in list(sys.modules):
if name == module.__name__ or name.startswith(f"{module.__name__}."):
sys.modules.pop(name, None)
Loading