kernels 0.14.0.dev1__tar.gz → 0.15.1__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.
Files changed (56) hide show
  1. {kernels-0.14.0.dev1 → kernels-0.15.1}/PKG-INFO +12 -3
  2. {kernels-0.14.0.dev1 → kernels-0.15.1}/README.md +10 -0
  3. {kernels-0.14.0.dev1 → kernels-0.15.1}/pyproject.toml +19 -2
  4. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/__init__.py +17 -0
  5. kernels-0.15.1/src/kernels/_versions.py +110 -0
  6. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/cli/__init__.py +7 -35
  7. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/cli/benchmark.py +1 -1
  8. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/cli/versions.py +3 -13
  9. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/__init__.py +2 -1
  10. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/device.py +4 -1
  11. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/func.py +22 -14
  12. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/kernelize.py +6 -5
  13. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/layer.py +38 -15
  14. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/utils.py +285 -144
  15. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/variants.py +207 -114
  16. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels.egg-info/PKG-INFO +12 -3
  17. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels.egg-info/SOURCES.txt +0 -2
  18. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels.egg-info/requires.txt +1 -3
  19. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_basic.py +119 -107
  20. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_benchmarks.py +1 -1
  21. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_deps.py +4 -3
  22. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_func.py +9 -3
  23. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_kernel_locking.py +4 -4
  24. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_layer.py +53 -45
  25. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_loaded_kernels.py +17 -21
  26. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_tvm_ffi.py +2 -2
  27. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_user_agent.py +13 -0
  28. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_variants.py +200 -29
  29. kernels-0.14.0.dev1/src/kernels/_versions.py +0 -72
  30. kernels-0.14.0.dev1/src/kernels/cli/check.py +0 -149
  31. kernels-0.14.0.dev1/src/kernels/metadata.py +0 -44
  32. {kernels-0.14.0.dev1 → kernels-0.15.1}/setup.cfg +0 -0
  33. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/_system.py +0 -0
  34. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/_windows.py +0 -0
  35. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/backends.py +0 -0
  36. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/benchmark.py +0 -0
  37. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/benchmarks/__init__.py +0 -0
  38. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/benchmarks/activation.py +0 -0
  39. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/benchmarks/attention.py +0 -0
  40. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/benchmarks/layer_norm.py +0 -0
  41. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/cli/benchmark_graphics.py +0 -0
  42. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/compat.py +0 -0
  43. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/deps.py +0 -0
  44. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/_interval_tree.py +0 -0
  45. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/globals.py +0 -0
  46. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/mode.py +0 -0
  47. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/layer/repos.py +0 -0
  48. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/lockfile.py +0 -0
  49. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/python_depends.json +0 -0
  50. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels/status.py +0 -0
  51. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels.egg-info/dependency_links.txt +0 -0
  52. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels.egg-info/entry_points.txt +0 -0
  53. {kernels-0.14.0.dev1 → kernels-0.15.1}/src/kernels.egg-info/top_level.txt +0 -0
  54. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_doctest.py +0 -0
  55. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_interval_tree.py +0 -0
  56. {kernels-0.14.0.dev1 → kernels-0.15.1}/tests/test_status.py +0 -0
@@ -1,18 +1,17 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: kernels
3
- Version: 0.14.0.dev1
3
+ Version: 0.15.1
4
4
  Summary: Download compute kernels
5
5
  Author-email: Daniel de Kok <daniel@huggingface.co>, David Holtz <david@huggingface.co>
6
6
  License: Apache-2.0
7
7
  Requires-Python: >=3.10
8
8
  Description-Content-Type: text/markdown
9
9
  Requires-Dist: huggingface-hub>=1.10.0
10
+ Requires-Dist: kernels-data>=0.14.0.dev1
10
11
  Requires-Dist: packaging>=20.0
11
12
  Requires-Dist: pyyaml>=6
12
13
  Requires-Dist: tomli>=2.0; python_version < "3.11"
13
14
  Requires-Dist: tomlkit>=0.13.3
14
- Provides-Extra: abi-check
15
- Requires-Dist: kernel-abi-check<0.7.0,>=0.6.2; extra == "abi-check"
16
15
  Provides-Extra: benchmark
17
16
  Requires-Dist: matplotlib>=3.7.0; extra == "benchmark"
18
17
  Requires-Dist: numpy>=2.0.2; extra == "benchmark"
@@ -73,3 +72,13 @@ the Hub.
73
72
  ## 📚 Documentation
74
73
 
75
74
  Read the [documentation of kernels](https://huggingface.co/docs/kernels/).
75
+
76
+ ## Test coverage
77
+
78
+ To reproduce the coverage number reported on PRs locally:
79
+
80
+ ```bash
81
+ uv run pytest --cov=kernels --cov-report=term-missing tests
82
+ ```
83
+
84
+ CI measures coverage on a single canonical matrix cell (Python 3.10 / Torch 2.12.0) and posts a sticky comment on the PR; the threshold is 80% (warn-only — the check stays green either way).
@@ -48,3 +48,13 @@ the Hub.
48
48
  ## 📚 Documentation
49
49
 
50
50
  Read the [documentation of kernels](https://huggingface.co/docs/kernels/).
51
+
52
+ ## Test coverage
53
+
54
+ To reproduce the coverage number reported on PRs locally:
55
+
56
+ ```bash
57
+ uv run pytest --cov=kernels --cov-report=term-missing tests
58
+ ```
59
+
60
+ CI measures coverage on a single canonical matrix cell (Python 3.10 / Torch 2.12.0) and posts a sticky comment on the PR; the threshold is 80% (warn-only — the check stays green either way).
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "kernels"
3
- version = "0.14.0.dev1"
3
+ version = "0.15.1"
4
4
  description = "Download compute kernels"
5
5
  authors = [
6
6
  { name = "Daniel de Kok", email = "daniel@huggingface.co" },
@@ -11,6 +11,7 @@ readme = "README.md"
11
11
  requires-python = ">= 3.10"
12
12
  dependencies = [
13
13
  "huggingface-hub>=1.10.0",
14
+ "kernels-data>=0.14.0.dev1",
14
15
  "packaging>=20.0",
15
16
  "pyyaml>=6",
16
17
  "tomli>=2.0; python_version<'3.11'",
@@ -25,9 +26,11 @@ build-backend = "setuptools.build_meta"
25
26
  dev = [
26
27
  "mktestdocs>=0.2.5",
27
28
  "mypy>=1.15.0",
29
+ "pre-commit",
28
30
  "pytest>=8",
29
31
  # Whatever version is compatible with pytest.
30
32
  "pytest-benchmark",
33
+ "pytest-cov>=5",
31
34
  "torch>=2.5",
32
35
  "apache-tvm-ffi>=0.1.9,<0.2.0",
33
36
  "types-pyyaml",
@@ -35,7 +38,6 @@ dev = [
35
38
  ]
36
39
 
37
40
  [project.optional-dependencies]
38
- abi-check = ["kernel-abi-check>=0.6.2,<0.7.0"]
39
41
  benchmark = [
40
42
  "matplotlib>=3.7.0",
41
43
  "numpy>=2.0.2",
@@ -87,3 +89,18 @@ lint.ignore = ["E501"]
87
89
  lint.select = ["E", "F", "I", "W"]
88
90
 
89
91
  [tool.ruff.format]
92
+
93
+ [tool.coverage.run]
94
+ source = ["kernels"]
95
+ branch = false
96
+ relative_files = true
97
+ omit = [
98
+ "*/benchmarks/*",
99
+ "*/benchmark.py",
100
+ "*/cli/*",
101
+ "*/_windows.py",
102
+ ]
103
+
104
+ [tool.coverage.report]
105
+ show_missing = true
106
+ skip_empty = true
@@ -2,6 +2,8 @@ import importlib.metadata
2
2
 
3
3
  __version__ = importlib.metadata.version("kernels")
4
4
 
5
+ from kernels_data import Metadata
6
+
5
7
  from kernels._windows import _add_additional_dll_paths
6
8
  from kernels.benchmark import Benchmark
7
9
  from kernels.layer import (
@@ -14,6 +16,7 @@ from kernels.layer import (
14
16
  LockedFuncRepository,
15
17
  LockedLayerRepository,
16
18
  Mode,
19
+ ROCMProperties,
17
20
  kernelize,
18
21
  register_kernel_mapping,
19
22
  replace_kernel_forward_from_hub,
@@ -22,7 +25,10 @@ from kernels.layer import (
22
25
  use_kernel_mapping,
23
26
  )
24
27
  from kernels.utils import (
28
+ LoadedKernel,
29
+ RepoInfo,
25
30
  get_kernel,
31
+ get_kernel_variants,
26
32
  get_loaded_kernels,
27
33
  get_local_kernel,
28
34
  get_locked_kernel,
@@ -30,6 +36,10 @@ from kernels.utils import (
30
36
  install_kernel,
31
37
  load_kernel,
32
38
  )
39
+ from kernels.variants import (
40
+ VariantAccepted,
41
+ VariantRejected,
42
+ )
33
43
 
34
44
  _add_additional_dll_paths()
35
45
 
@@ -38,14 +48,21 @@ __all__ = [
38
48
  "Benchmark",
39
49
  "CUDAProperties",
40
50
  "Device",
51
+ "ROCMProperties",
41
52
  "FuncRepository",
42
53
  "LayerRepository",
54
+ "LoadedKernel",
43
55
  "LocalFuncRepository",
44
56
  "LocalLayerRepository",
45
57
  "LockedFuncRepository",
46
58
  "LockedLayerRepository",
59
+ "Metadata",
47
60
  "Mode",
61
+ "RepoInfo",
62
+ "VariantAccepted",
63
+ "VariantRejected",
48
64
  "get_kernel",
65
+ "get_kernel_variants",
49
66
  "get_loaded_kernels",
50
67
  "get_local_kernel",
51
68
  "get_locked_kernel",
@@ -0,0 +1,110 @@
1
+ import logging
2
+ import os
3
+ from pathlib import Path
4
+
5
+ from huggingface_hub import constants
6
+ from huggingface_hub.file_download import repo_folder_name
7
+ from huggingface_hub.hf_api import GitRefInfo
8
+
9
+ logger = logging.getLogger(__name__)
10
+
11
+
12
+ def _get_available_versions(repo_id: str) -> dict[int, GitRefInfo]:
13
+ """Get kernel versions that are available in the repository."""
14
+ from kernels.utils import _get_hf_api
15
+
16
+ if constants.HF_HUB_OFFLINE:
17
+ return _get_available_versions_from_cache(repo_id)
18
+
19
+ refs = _get_hf_api().list_repo_refs(repo_id=repo_id, repo_type="kernel")
20
+
21
+ versions = {}
22
+ for branch in refs.branches:
23
+ if not branch.name.startswith("v"):
24
+ continue
25
+ try:
26
+ versions[int(branch.name[1:])] = branch
27
+ except ValueError:
28
+ continue
29
+
30
+ return versions
31
+
32
+
33
+ def _get_available_versions_from_cache(repo_id: str) -> dict[int, GitRefInfo]:
34
+ """Get kernel versions from the local Hugging Face cache."""
35
+ cache_dir = os.environ.get("KERNELS_CACHE") or constants.HF_HUB_CACHE
36
+
37
+ versions: dict[int, GitRefInfo] = {}
38
+ # Tolerate both layouts: the "kernel" repo type used by newer
39
+ # huggingface_hub, and the legacy "model" prefix that older caches use.
40
+ for repo_type in ("kernel", "model"):
41
+ refs_dir = Path(cache_dir) / repo_folder_name(repo_id=repo_id, repo_type=repo_type) / "refs"
42
+ if not refs_dir.is_dir():
43
+ continue
44
+ for ref_path in refs_dir.iterdir():
45
+ if not ref_path.is_file():
46
+ continue
47
+ ref_name = ref_path.name
48
+ if not ref_name.startswith("v"):
49
+ continue
50
+ try:
51
+ version = int(ref_name[1:])
52
+ except ValueError:
53
+ continue
54
+ try:
55
+ commit = ref_path.read_text().strip()
56
+ except OSError:
57
+ continue
58
+ versions[version] = GitRefInfo(name=ref_name, ref=ref_name, target_commit=commit)
59
+
60
+ return versions
61
+
62
+
63
+ def resolve_version_spec_as_ref(repo_id: str, version_spec: int) -> GitRefInfo:
64
+ """
65
+ Get the ref for a kernel with the given version.
66
+ """
67
+ versions = _get_available_versions(repo_id)
68
+
69
+ ref = versions.get(version_spec, None)
70
+ if ref is None:
71
+ if constants.HF_HUB_OFFLINE and not versions:
72
+ raise ValueError(
73
+ f"Version {version_spec} of '{repo_id}' is not available in the local cache "
74
+ "and Hugging Face Hub is in offline mode. Download the kernel "
75
+ "while online first, or pass an explicit `revision=<commit>`."
76
+ )
77
+ raise ValueError(
78
+ f"Version {version_spec} not found, available versions: {', '.join(str(v) for v in sorted(versions.keys()))}"
79
+ )
80
+
81
+ latest_version = max(versions.keys())
82
+ if version_spec < latest_version:
83
+ logger.warning(
84
+ "You are using version %d of '%s', but version %d is available.",
85
+ version_spec,
86
+ repo_id,
87
+ latest_version,
88
+ )
89
+
90
+ return ref
91
+
92
+
93
+ def select_revision_or_version(
94
+ repo_id: str,
95
+ *,
96
+ revision: str | None,
97
+ version: int | None,
98
+ ) -> str:
99
+ if revision is not None and version is not None:
100
+ raise ValueError("Only one of `revision` or `version` must be specified.")
101
+ elif revision is not None:
102
+ return revision
103
+ elif version is not None:
104
+ return resolve_version_spec_as_ref(repo_id, version).target_commit
105
+ else:
106
+ raise ValueError(
107
+ "A kernel version or revision must be specified. "
108
+ "Use `version=<major>` for a stable kernel API version or `revision=<branch/tag/commit>` "
109
+ "for an explicit Hub revision. See: https://huggingface.co/docs/kernels/migration"
110
+ )
@@ -18,25 +18,8 @@ def main():
18
18
  subparsers = parser.add_subparsers(required=True)
19
19
 
20
20
  check_parser = subparsers.add_parser("check", help="Check a kernel for compliance")
21
- check_parser.add_argument("repo_id", type=str, help="The kernel repo ID")
22
- check_parser.add_argument(
23
- "--revision",
24
- type=str,
25
- default="main",
26
- help="The kernel revision (branch, tag, or commit SHA, defaults to 'main')",
27
- )
28
- check_parser.add_argument("--macos", type=str, help="macOS version", default="15.0")
29
- check_parser.add_argument("--manylinux", type=str, help="Manylinux version", default="manylinux_2_28")
30
- check_parser.add_argument("--python-abi", type=str, help="Python ABI version", default="3.9")
31
- check_parser.set_defaults(
32
- func=lambda args: check_kernel(
33
- macos=args.macos,
34
- manylinux=args.manylinux,
35
- python_abi=args.python_abi,
36
- repo_id=args.repo_id,
37
- revision=args.revision,
38
- )
39
- )
21
+ check_parser.add_argument("repo_id", type=str, nargs="?")
22
+ check_parser.set_defaults(func=_check_moved)
40
23
 
41
24
  download_parser = subparsers.add_parser("download", help="Download locked kernels")
42
25
  download_parser.add_argument(
@@ -168,23 +151,12 @@ class _JSONEncoder(json.JSONEncoder):
168
151
  return super().default(o)
169
152
 
170
153
 
171
- def check_kernel(*, macos: str, manylinux: str, python_abi: str, repo_id: str, revision: str):
172
- try:
173
- from kernels.cli import check
174
- except ImportError:
175
- print(
176
- "`kernels check` requires the `kernel-abi-check` package: pip install kernel-abi-check",
177
- file=sys.stderr,
178
- )
179
- sys.exit(1)
180
-
181
- check.check_kernel(
182
- macos=macos,
183
- manylinux=manylinux,
184
- python_abi=python_abi,
185
- repo_id=repo_id,
186
- revision=revision,
154
+ def _check_moved(_args):
155
+ print(
156
+ "`kernels check` has moved to `kernel-builder check-abi`",
157
+ file=sys.stderr,
187
158
  )
159
+ sys.exit(1)
188
160
 
189
161
 
190
162
  def run_benchmark(args):
@@ -464,7 +464,7 @@ def run_benchmark_class(
464
464
  from kernels import get_kernel, get_local_kernel
465
465
 
466
466
  if is_local:
467
- kernel = get_local_kernel(Path(repo_id), "activation")
467
+ kernel = get_local_kernel(Path(repo_id))
468
468
  else:
469
469
  kernel = get_kernel(repo_id, revision=revision)
470
470
 
@@ -3,6 +3,7 @@ from kernels.utils import _get_hf_api
3
3
  from kernels.variants import (
4
4
  get_variants,
5
5
  resolve_variants,
6
+ variants_trace_str,
6
7
  )
7
8
 
8
9
 
@@ -15,16 +16,5 @@ def print_kernel_versions(repo_id: str):
15
16
 
16
17
  for version, ref in sorted(versions.items(), key=lambda x: x[0]):
17
18
  variants = get_variants(api, repo_id=repo_id, revision=ref.ref)
18
- resolved = resolve_variants(variants, None)
19
- best = resolved[0] if resolved else None
20
- resolved_set = set(resolved)
21
- print(f"Version {version}: ", end="")
22
- variant_strs = [
23
- (
24
- f"✅ {variant.variant_str} ({'compatible, preferred' if variant == best else 'compatible'})"
25
- if variant in resolved_set
26
- else f"{variant.variant_str}"
27
- )
28
- for variant in variants
29
- ]
30
- print(", ".join(variant_strs))
19
+ _, status = resolve_variants(variants, None)
20
+ print(f"Version {version}:\n\n{variants_trace_str(status)}")
@@ -1,4 +1,4 @@
1
- from .device import CUDAProperties, Device
1
+ from .device import CUDAProperties, Device, ROCMProperties
2
2
  from .func import (
3
3
  FuncRepository,
4
4
  LocalFuncRepository,
@@ -22,6 +22,7 @@ from .mode import Mode
22
22
  __all__ = [
23
23
  "CUDAProperties",
24
24
  "Device",
25
+ "ROCMProperties",
25
26
  "FuncRepository",
26
27
  "LayerRepository",
27
28
  "LocalFuncRepository",
@@ -141,12 +141,15 @@ class Device:
141
141
  """
142
142
 
143
143
  type: str
144
- properties: CUDAProperties | None = None
144
+ properties: CUDAProperties | ROCMProperties | None = None
145
145
 
146
146
  def __post_init__(self):
147
147
  if self.properties is not None and isinstance(self.properties, CUDAProperties):
148
148
  if self.type != "cuda":
149
149
  raise ValueError("CUDAProperties is only supported for 'cuda' devices.")
150
+ if self.properties is not None and isinstance(self.properties, ROCMProperties):
151
+ if self.type != "rocm":
152
+ raise ValueError("ROCMProperties is only supported for 'rocm' devices.")
150
153
 
151
154
  def __eq__(self, other):
152
155
  if not isinstance(other, Device):
@@ -33,10 +33,11 @@ class FuncRepository:
33
33
  The Hub repository containing the layer.
34
34
  func_name (`str`):
35
35
  The name of the function within the kernel repository.
36
- revision (`str`, *optional*, defaults to `"main"`):
36
+ revision (`str`, *optional*):
37
37
  The specific revision (branch, tag, or commit) to download. Cannot be used together with `version`.
38
38
  version (`int`, *optional*):
39
39
  The kernel version to download. Cannot be used together with `revision`.
40
+ Either `version` or `revision` must be specified.
40
41
 
41
42
  Example:
42
43
  ```python
@@ -46,6 +47,7 @@ class FuncRepository:
46
47
  layer_repo = FuncRepository(
47
48
  repo_id="kernels-community/activation",
48
49
  func_name="silu_and_mul",
50
+ revision="main",
49
51
  )
50
52
 
51
53
  # Reference a layer by version
@@ -64,12 +66,16 @@ class FuncRepository:
64
66
  func_name: str,
65
67
  revision: str | None = None,
66
68
  version: int | None = None,
69
+ trust_remote_code: bool | list[str] = False,
67
70
  ):
68
71
  if revision is not None and version is not None:
69
72
  raise ValueError("Either a revision or a version must be specified, not both.")
73
+ if revision is None and version is None:
74
+ raise ValueError("Either a revision or a version must be specified.")
70
75
 
71
76
  self._repo_id = repo_id
72
77
  self.func_name = func_name
78
+ self._trust_remote_code = trust_remote_code
73
79
 
74
80
  # We are going to resolve these lazily, since we do not want
75
81
  # to do a network request for every registered FuncRepository.
@@ -85,7 +91,9 @@ class FuncRepository:
85
91
  )
86
92
 
87
93
  def load(self) -> Type["nn.Module"]:
88
- kernel = get_kernel(self._repo_id, revision=self._resolve_revision())
94
+ kernel = get_kernel(
95
+ self._repo_id, revision=self._resolve_revision(), trust_remote_code=self._trust_remote_code
96
+ )
89
97
  return _get_kernel_func(self, kernel)
90
98
 
91
99
  def __eq__(self, other):
@@ -95,10 +103,11 @@ class FuncRepository:
95
103
  and self._repo_id == other._repo_id
96
104
  and self._revision == other._revision
97
105
  and self._version == other._version
106
+ and self._trust_remote_code == other._trust_remote_code
98
107
  )
99
108
 
100
109
  def __hash__(self):
101
- return hash((self.func_name, self._repo_id, self._revision, self._version))
110
+ return hash((self.func_name, self._repo_id, self._revision, self._version, self._trust_remote_code))
102
111
 
103
112
  def __str__(self) -> str:
104
113
  return f"`{self._repo_id}` (revision: {self._resolve_revision()}), function `{self.func_name}`"
@@ -111,8 +120,6 @@ class LocalFuncRepository:
111
120
  Args:
112
121
  repo_path (`Path`):
113
122
  The local repository containing the layer.
114
- package_name (`str`):
115
- Package name of the kernel.
116
123
  func_name (`str`):
117
124
  The name of the function within the kernel repository.
118
125
 
@@ -125,7 +132,6 @@ class LocalFuncRepository:
125
132
  # Reference a specific layer by revision
126
133
  layer_repo = LocalFuncRepository(
127
134
  repo_path=Path("/home/daniel/kernels/activation"),
128
- package_name="activation",
129
135
  func_name="silu_and_mul",
130
136
  )
131
137
  ```
@@ -135,15 +141,13 @@ class LocalFuncRepository:
135
141
  self,
136
142
  repo_path: Path,
137
143
  *,
138
- package_name: str,
139
144
  func_name: str,
140
145
  ):
141
146
  self._repo_path = repo_path
142
- self._package_name = package_name
143
147
  self.func_name = func_name
144
148
 
145
149
  def load(self) -> Type["nn.Module"]:
146
- kernel = get_local_kernel(self._repo_path, self._package_name)
150
+ kernel = get_local_kernel(self._repo_path)
147
151
  return _get_kernel_func(self, kernel)
148
152
 
149
153
  def __eq__(self, other):
@@ -151,14 +155,13 @@ class LocalFuncRepository:
151
155
  isinstance(other, LocalFuncRepository)
152
156
  and self.func_name == other.func_name
153
157
  and self._repo_path == other._repo_path
154
- and self._package_name == other._package_name
155
158
  )
156
159
 
157
160
  def __hash__(self):
158
- return hash((self.func_name, self._repo_path, self._package_name))
161
+ return hash((self.func_name, self._repo_path))
159
162
 
160
163
  def __str__(self) -> str:
161
- return f"`{self._repo_path}` (package: {self._package_name}), layer `{self.func_name}`"
164
+ return f"`{self._repo_path}` (layer `{self.func_name}`"
162
165
 
163
166
 
164
167
  def use_kernel_func_from_hub(func_name: str):
@@ -230,6 +233,7 @@ class LockedFuncRepository:
230
233
  *,
231
234
  lockfile: Path | None = None,
232
235
  func_name: str,
236
+ trust_remote_code: bool | list[str] = False,
233
237
  ):
234
238
  """
235
239
  Construct a function repository.
@@ -239,10 +243,13 @@ class LockedFuncRepository:
239
243
  lockfile (`Path`, *optional*): Path to the lockfile. If not provided,
240
244
  the lockfile will be inferred from the caller's context.
241
245
  func_name (`str`): The name of the function within the kernel repository.
246
+ trust_remote_code (`bool`, *optional*, defaults to `False`):
247
+ Whether to allow loading kernels from untrusted organisations.
242
248
  """
243
249
  self._repo_id = repo_id
244
250
  self._lockfile = lockfile
245
251
  self.func_name = func_name
252
+ self._trust_remote_code = trust_remote_code
246
253
  self._revision = self._resolve_revision()
247
254
 
248
255
  def _resolve_revision(self) -> str:
@@ -258,7 +265,7 @@ class LockedFuncRepository:
258
265
  return locked_sha
259
266
 
260
267
  def load(self) -> Type["nn.Module"]:
261
- kernel = get_kernel(repo_id=self._repo_id, revision=self._revision)
268
+ kernel = get_kernel(repo_id=self._repo_id, revision=self._revision, trust_remote_code=self._trust_remote_code)
262
269
  return _get_kernel_func(self, kernel)
263
270
 
264
271
  def __eq__(self, other):
@@ -267,10 +274,11 @@ class LockedFuncRepository:
267
274
  and self.func_name == other.func_name
268
275
  and self._repo_id == other._repo_id
269
276
  and self._revision == other._revision
277
+ and self._trust_remote_code == other._trust_remote_code
270
278
  )
271
279
 
272
280
  def __hash__(self):
273
- return hash((self.func_name, self._repo_id, self._revision))
281
+ return hash((self.func_name, self._repo_id, self._revision, self._trust_remote_code))
274
282
 
275
283
  def __str__(self) -> str:
276
284
  return f"`{self._repo_id}` (revision: {self._revision}), function `{self.func_name}`"
@@ -139,12 +139,12 @@ def register_kernel_mapping(
139
139
  "MultiHeadAttention": {
140
140
  "cuda": {
141
141
  Mode.TRAINING: LayerRepository(
142
- repo_id="username/training-kernels",
142
+ repo_id="kernels-community/training-kernels",
143
143
  layer_name="TrainingAttention",
144
144
  version=1,
145
145
  ),
146
146
  Mode.INFERENCE: LayerRepository(
147
- repo_id="username/inference-kernels",
147
+ repo_id="kernels-community/inference-kernels",
148
148
  layer_name="FastAttention",
149
149
  version=1,
150
150
  ),
@@ -206,7 +206,7 @@ def kernelize(
206
206
  import torch
207
207
  import torch.nn as nn
208
208
 
209
- from kernels import kernelize, Mode, register_kernel_mapping, LayerRepository
209
+ from kernels import kernelize, Mode, use_kernel_mapping, LayerRepository
210
210
  from kernels import use_kernel_forward_from_hub
211
211
 
212
212
  @use_kernel_forward_from_hub("SiluAndMul")
@@ -220,10 +220,10 @@ def kernelize(
220
220
  "cuda": LayerRepository(
221
221
  repo_id="kernels-community/activation",
222
222
  layer_name="SiluAndMul",
223
+ version=1,
223
224
  )
224
225
  }
225
226
  }
226
- register_kernel_mapping(mapping)
227
227
 
228
228
  # Create and kernelize a model
229
229
  model = nn.Sequential(
@@ -232,7 +232,8 @@ def kernelize(
232
232
  )
233
233
 
234
234
  # Kernelize for inference
235
- kernelized_model = kernelize(model, mode=Mode.TRAINING | Mode.TORCH_COMPILE)
235
+ with use_kernel_mapping(mapping):
236
+ kernelized_model = kernelize(model, mode=Mode.TRAINING | Mode.TORCH_COMPILE)
236
237
  ```
237
238
  """
238
239