kernels 0.14.0.dev0__tar.gz → 0.14.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.
- {kernels-0.14.0.dev0 → kernels-0.14.1}/PKG-INFO +2 -1
- {kernels-0.14.0.dev0 → kernels-0.14.1}/pyproject.toml +2 -1
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/__init__.py +7 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/func.py +17 -13
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/kernelize.py +6 -5
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/layer.py +19 -13
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/utils.py +158 -83
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels.egg-info/PKG-INFO +2 -1
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels.egg-info/SOURCES.txt +1 -1
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels.egg-info/requires.txt +1 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_basic.py +46 -91
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_deps.py +4 -3
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_func.py +3 -3
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_kernel_locking.py +4 -4
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_layer.py +19 -47
- kernels-0.14.1/tests/test_loaded_kernels.py +98 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_tvm_ffi.py +2 -2
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_user_agent.py +13 -0
- kernels-0.14.0.dev0/src/kernels/metadata.py +0 -44
- {kernels-0.14.0.dev0 → kernels-0.14.1}/README.md +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/setup.cfg +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/_system.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/_versions.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/_windows.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/backends.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/benchmark.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/benchmarks/__init__.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/benchmarks/activation.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/benchmarks/attention.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/benchmarks/layer_norm.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/cli/__init__.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/cli/benchmark.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/cli/benchmark_graphics.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/cli/check.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/cli/versions.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/compat.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/deps.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/__init__.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/_interval_tree.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/device.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/globals.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/mode.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/layer/repos.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/lockfile.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/python_depends.json +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/status.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels/variants.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels.egg-info/dependency_links.txt +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels.egg-info/entry_points.txt +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/src/kernels.egg-info/top_level.txt +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_benchmarks.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_doctest.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_interval_tree.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_status.py +0 -0
- {kernels-0.14.0.dev0 → kernels-0.14.1}/tests/test_variants.py +0 -0
|
@@ -1,12 +1,13 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: kernels
|
|
3
|
-
Version: 0.14.
|
|
3
|
+
Version: 0.14.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"
|
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
[project]
|
|
2
2
|
name = "kernels"
|
|
3
|
-
version = "0.14.
|
|
3
|
+
version = "0.14.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'",
|
|
@@ -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 (
|
|
@@ -22,6 +24,8 @@ from kernels.layer import (
|
|
|
22
24
|
use_kernel_mapping,
|
|
23
25
|
)
|
|
24
26
|
from kernels.utils import (
|
|
27
|
+
LoadedKernel,
|
|
28
|
+
RepoInfo,
|
|
25
29
|
get_kernel,
|
|
26
30
|
get_loaded_kernels,
|
|
27
31
|
get_local_kernel,
|
|
@@ -40,11 +44,14 @@ __all__ = [
|
|
|
40
44
|
"Device",
|
|
41
45
|
"FuncRepository",
|
|
42
46
|
"LayerRepository",
|
|
47
|
+
"LoadedKernel",
|
|
43
48
|
"LocalFuncRepository",
|
|
44
49
|
"LocalLayerRepository",
|
|
45
50
|
"LockedFuncRepository",
|
|
46
51
|
"LockedLayerRepository",
|
|
52
|
+
"Metadata",
|
|
47
53
|
"Mode",
|
|
54
|
+
"RepoInfo",
|
|
48
55
|
"get_kernel",
|
|
49
56
|
"get_loaded_kernels",
|
|
50
57
|
"get_local_kernel",
|
|
@@ -64,12 +64,14 @@ class FuncRepository:
|
|
|
64
64
|
func_name: str,
|
|
65
65
|
revision: str | None = None,
|
|
66
66
|
version: int | None = None,
|
|
67
|
+
trust_remote_code: bool | list[str] = False,
|
|
67
68
|
):
|
|
68
69
|
if revision is not None and version is not None:
|
|
69
70
|
raise ValueError("Either a revision or a version must be specified, not both.")
|
|
70
71
|
|
|
71
72
|
self._repo_id = repo_id
|
|
72
73
|
self.func_name = func_name
|
|
74
|
+
self._trust_remote_code = trust_remote_code
|
|
73
75
|
|
|
74
76
|
# We are going to resolve these lazily, since we do not want
|
|
75
77
|
# to do a network request for every registered FuncRepository.
|
|
@@ -85,7 +87,9 @@ class FuncRepository:
|
|
|
85
87
|
)
|
|
86
88
|
|
|
87
89
|
def load(self) -> Type["nn.Module"]:
|
|
88
|
-
kernel = get_kernel(
|
|
90
|
+
kernel = get_kernel(
|
|
91
|
+
self._repo_id, revision=self._resolve_revision(), trust_remote_code=self._trust_remote_code
|
|
92
|
+
)
|
|
89
93
|
return _get_kernel_func(self, kernel)
|
|
90
94
|
|
|
91
95
|
def __eq__(self, other):
|
|
@@ -95,10 +99,11 @@ class FuncRepository:
|
|
|
95
99
|
and self._repo_id == other._repo_id
|
|
96
100
|
and self._revision == other._revision
|
|
97
101
|
and self._version == other._version
|
|
102
|
+
and self._trust_remote_code == other._trust_remote_code
|
|
98
103
|
)
|
|
99
104
|
|
|
100
105
|
def __hash__(self):
|
|
101
|
-
return hash((self.func_name, self._repo_id, self._revision, self._version))
|
|
106
|
+
return hash((self.func_name, self._repo_id, self._revision, self._version, self._trust_remote_code))
|
|
102
107
|
|
|
103
108
|
def __str__(self) -> str:
|
|
104
109
|
return f"`{self._repo_id}` (revision: {self._resolve_revision()}), function `{self.func_name}`"
|
|
@@ -111,8 +116,6 @@ class LocalFuncRepository:
|
|
|
111
116
|
Args:
|
|
112
117
|
repo_path (`Path`):
|
|
113
118
|
The local repository containing the layer.
|
|
114
|
-
package_name (`str`):
|
|
115
|
-
Package name of the kernel.
|
|
116
119
|
func_name (`str`):
|
|
117
120
|
The name of the function within the kernel repository.
|
|
118
121
|
|
|
@@ -125,7 +128,6 @@ class LocalFuncRepository:
|
|
|
125
128
|
# Reference a specific layer by revision
|
|
126
129
|
layer_repo = LocalFuncRepository(
|
|
127
130
|
repo_path=Path("/home/daniel/kernels/activation"),
|
|
128
|
-
package_name="activation",
|
|
129
131
|
func_name="silu_and_mul",
|
|
130
132
|
)
|
|
131
133
|
```
|
|
@@ -135,15 +137,13 @@ class LocalFuncRepository:
|
|
|
135
137
|
self,
|
|
136
138
|
repo_path: Path,
|
|
137
139
|
*,
|
|
138
|
-
package_name: str,
|
|
139
140
|
func_name: str,
|
|
140
141
|
):
|
|
141
142
|
self._repo_path = repo_path
|
|
142
|
-
self._package_name = package_name
|
|
143
143
|
self.func_name = func_name
|
|
144
144
|
|
|
145
145
|
def load(self) -> Type["nn.Module"]:
|
|
146
|
-
kernel = get_local_kernel(self._repo_path
|
|
146
|
+
kernel = get_local_kernel(self._repo_path)
|
|
147
147
|
return _get_kernel_func(self, kernel)
|
|
148
148
|
|
|
149
149
|
def __eq__(self, other):
|
|
@@ -151,14 +151,13 @@ class LocalFuncRepository:
|
|
|
151
151
|
isinstance(other, LocalFuncRepository)
|
|
152
152
|
and self.func_name == other.func_name
|
|
153
153
|
and self._repo_path == other._repo_path
|
|
154
|
-
and self._package_name == other._package_name
|
|
155
154
|
)
|
|
156
155
|
|
|
157
156
|
def __hash__(self):
|
|
158
|
-
return hash((self.func_name, self._repo_path
|
|
157
|
+
return hash((self.func_name, self._repo_path))
|
|
159
158
|
|
|
160
159
|
def __str__(self) -> str:
|
|
161
|
-
return f"`{self._repo_path}` (
|
|
160
|
+
return f"`{self._repo_path}` (layer `{self.func_name}`"
|
|
162
161
|
|
|
163
162
|
|
|
164
163
|
def use_kernel_func_from_hub(func_name: str):
|
|
@@ -230,6 +229,7 @@ class LockedFuncRepository:
|
|
|
230
229
|
*,
|
|
231
230
|
lockfile: Path | None = None,
|
|
232
231
|
func_name: str,
|
|
232
|
+
trust_remote_code: bool | list[str] = False,
|
|
233
233
|
):
|
|
234
234
|
"""
|
|
235
235
|
Construct a function repository.
|
|
@@ -239,10 +239,13 @@ class LockedFuncRepository:
|
|
|
239
239
|
lockfile (`Path`, *optional*): Path to the lockfile. If not provided,
|
|
240
240
|
the lockfile will be inferred from the caller's context.
|
|
241
241
|
func_name (`str`): The name of the function within the kernel repository.
|
|
242
|
+
trust_remote_code (`bool`, *optional*, defaults to `False`):
|
|
243
|
+
Whether to allow loading kernels from untrusted organisations.
|
|
242
244
|
"""
|
|
243
245
|
self._repo_id = repo_id
|
|
244
246
|
self._lockfile = lockfile
|
|
245
247
|
self.func_name = func_name
|
|
248
|
+
self._trust_remote_code = trust_remote_code
|
|
246
249
|
self._revision = self._resolve_revision()
|
|
247
250
|
|
|
248
251
|
def _resolve_revision(self) -> str:
|
|
@@ -258,7 +261,7 @@ class LockedFuncRepository:
|
|
|
258
261
|
return locked_sha
|
|
259
262
|
|
|
260
263
|
def load(self) -> Type["nn.Module"]:
|
|
261
|
-
kernel = get_kernel(repo_id=self._repo_id, revision=self._revision)
|
|
264
|
+
kernel = get_kernel(repo_id=self._repo_id, revision=self._revision, trust_remote_code=self._trust_remote_code)
|
|
262
265
|
return _get_kernel_func(self, kernel)
|
|
263
266
|
|
|
264
267
|
def __eq__(self, other):
|
|
@@ -267,10 +270,11 @@ class LockedFuncRepository:
|
|
|
267
270
|
and self.func_name == other.func_name
|
|
268
271
|
and self._repo_id == other._repo_id
|
|
269
272
|
and self._revision == other._revision
|
|
273
|
+
and self._trust_remote_code == other._trust_remote_code
|
|
270
274
|
)
|
|
271
275
|
|
|
272
276
|
def __hash__(self):
|
|
273
|
-
return hash((self.func_name, self._repo_id, self._revision))
|
|
277
|
+
return hash((self.func_name, self._repo_id, self._revision, self._trust_remote_code))
|
|
274
278
|
|
|
275
279
|
def __str__(self) -> str:
|
|
276
280
|
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="
|
|
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="
|
|
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,
|
|
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
|
-
|
|
235
|
+
with use_kernel_mapping(mapping):
|
|
236
|
+
kernelized_model = kernelize(model, mode=Mode.TRAINING | Mode.TORCH_COMPILE)
|
|
236
237
|
```
|
|
237
238
|
"""
|
|
238
239
|
|
|
@@ -42,6 +42,10 @@ class LayerRepository:
|
|
|
42
42
|
The specific revision (branch, tag, or commit) to download. Cannot be used together with `version`.
|
|
43
43
|
version (`int`, *optional*):
|
|
44
44
|
The kernel version to download. Cannot be used together with `revision`.
|
|
45
|
+
trust_remote_code (`bool | list[str]`, *optional*, defaults to `False`):
|
|
46
|
+
Whether to allow loading kernels from untrusted organisations. A list
|
|
47
|
+
of signing identities can be provided for future verification support;
|
|
48
|
+
until then it warns and falls back to the default trust check.
|
|
45
49
|
|
|
46
50
|
Example:
|
|
47
51
|
```python
|
|
@@ -63,12 +67,14 @@ class LayerRepository:
|
|
|
63
67
|
layer_name: str,
|
|
64
68
|
revision: str | None = None,
|
|
65
69
|
version: int | None = None,
|
|
70
|
+
trust_remote_code: bool | list[str] = False,
|
|
66
71
|
):
|
|
67
72
|
if revision is not None and version is not None:
|
|
68
73
|
raise ValueError("Either a revision or a version must be specified, not both.")
|
|
69
74
|
|
|
70
75
|
self._repo_id = repo_id
|
|
71
76
|
self.layer_name = layer_name
|
|
77
|
+
self._trust_remote_code = trust_remote_code
|
|
72
78
|
|
|
73
79
|
# We are going to resolve these lazily, since we do not want
|
|
74
80
|
# to do a network request for every registered LayerRepository.
|
|
@@ -84,7 +90,9 @@ class LayerRepository:
|
|
|
84
90
|
)
|
|
85
91
|
|
|
86
92
|
def load(self) -> Type["nn.Module"]:
|
|
87
|
-
kernel = get_kernel(
|
|
93
|
+
kernel = get_kernel(
|
|
94
|
+
self._repo_id, revision=self._resolve_revision(), trust_remote_code=self._trust_remote_code
|
|
95
|
+
)
|
|
88
96
|
return _get_kernel_layer(self, kernel)
|
|
89
97
|
|
|
90
98
|
def __eq__(self, other):
|
|
@@ -94,10 +102,11 @@ class LayerRepository:
|
|
|
94
102
|
and self._repo_id == other._repo_id
|
|
95
103
|
and self._revision == other._revision
|
|
96
104
|
and self._version == other._version
|
|
105
|
+
and self._trust_remote_code == other._trust_remote_code
|
|
97
106
|
)
|
|
98
107
|
|
|
99
108
|
def __hash__(self):
|
|
100
|
-
return hash((self.layer_name, self._repo_id, self._revision, self._version))
|
|
109
|
+
return hash((self.layer_name, self._repo_id, self._revision, self._version, self._trust_remote_code))
|
|
101
110
|
|
|
102
111
|
def __str__(self) -> str:
|
|
103
112
|
return f"`{self._repo_id}` (revision: {self._resolve_revision()}), layer `{self.layer_name}`"
|
|
@@ -110,8 +119,6 @@ class LocalLayerRepository:
|
|
|
110
119
|
Args:
|
|
111
120
|
repo_path (`Path`):
|
|
112
121
|
The local repository containing the layer.
|
|
113
|
-
package_name (`str`):
|
|
114
|
-
Package name of the kernel.
|
|
115
122
|
layer_name (`str`):
|
|
116
123
|
The name of the layer within the kernel repository.
|
|
117
124
|
|
|
@@ -124,7 +131,6 @@ class LocalLayerRepository:
|
|
|
124
131
|
# Reference a specific layer by revision
|
|
125
132
|
layer_repo = LocalLayerRepository(
|
|
126
133
|
repo_path=Path("/home/daniel/kernels/activation"),
|
|
127
|
-
package_name="activation",
|
|
128
134
|
layer_name="SiluAndMul",
|
|
129
135
|
)
|
|
130
136
|
```
|
|
@@ -134,15 +140,13 @@ class LocalLayerRepository:
|
|
|
134
140
|
self,
|
|
135
141
|
repo_path: Path,
|
|
136
142
|
*,
|
|
137
|
-
package_name: str,
|
|
138
143
|
layer_name: str,
|
|
139
144
|
):
|
|
140
145
|
self._repo_path = repo_path
|
|
141
|
-
self._package_name = package_name
|
|
142
146
|
self.layer_name = layer_name
|
|
143
147
|
|
|
144
148
|
def load(self) -> Type["nn.Module"]:
|
|
145
|
-
kernel = get_local_kernel(self._repo_path
|
|
149
|
+
kernel = get_local_kernel(self._repo_path)
|
|
146
150
|
return _get_kernel_layer(self, kernel)
|
|
147
151
|
|
|
148
152
|
def __eq__(self, other):
|
|
@@ -150,14 +154,13 @@ class LocalLayerRepository:
|
|
|
150
154
|
isinstance(other, LocalLayerRepository)
|
|
151
155
|
and self.layer_name == other.layer_name
|
|
152
156
|
and self._repo_path == other._repo_path
|
|
153
|
-
and self._package_name == other._package_name
|
|
154
157
|
)
|
|
155
158
|
|
|
156
159
|
def __hash__(self):
|
|
157
|
-
return hash((self.layer_name, self._repo_path
|
|
160
|
+
return hash((self.layer_name, self._repo_path))
|
|
158
161
|
|
|
159
162
|
def __str__(self) -> str:
|
|
160
|
-
return f"`{self._repo_path}` (
|
|
163
|
+
return f"`{self._repo_path}` (layer `{self.layer_name}`"
|
|
161
164
|
|
|
162
165
|
|
|
163
166
|
class LockedLayerRepository:
|
|
@@ -174,6 +177,7 @@ class LockedLayerRepository:
|
|
|
174
177
|
*,
|
|
175
178
|
lockfile: Path | None = None,
|
|
176
179
|
layer_name: str,
|
|
180
|
+
trust_remote_code: bool | list[str] = False,
|
|
177
181
|
):
|
|
178
182
|
"""
|
|
179
183
|
Construct a layer repository.
|
|
@@ -184,6 +188,7 @@ class LockedLayerRepository:
|
|
|
184
188
|
self._repo_id = repo_id
|
|
185
189
|
self._lockfile = lockfile
|
|
186
190
|
self.layer_name = layer_name
|
|
191
|
+
self._trust_remote_code = trust_remote_code
|
|
187
192
|
self._revision = self._resolve_revision()
|
|
188
193
|
|
|
189
194
|
def _resolve_revision(self) -> str:
|
|
@@ -199,7 +204,7 @@ class LockedLayerRepository:
|
|
|
199
204
|
return locked_sha
|
|
200
205
|
|
|
201
206
|
def load(self) -> Type["nn.Module"]:
|
|
202
|
-
kernel = get_kernel(repo_id=self._repo_id, revision=self._revision)
|
|
207
|
+
kernel = get_kernel(repo_id=self._repo_id, revision=self._revision, trust_remote_code=self._trust_remote_code)
|
|
203
208
|
return _get_kernel_layer(self, kernel)
|
|
204
209
|
|
|
205
210
|
def __eq__(self, other):
|
|
@@ -208,10 +213,11 @@ class LockedLayerRepository:
|
|
|
208
213
|
and self.layer_name == other.layer_name
|
|
209
214
|
and self._repo_id == other._repo_id
|
|
210
215
|
and self._revision == other._revision
|
|
216
|
+
and self._trust_remote_code == other._trust_remote_code
|
|
211
217
|
)
|
|
212
218
|
|
|
213
219
|
def __hash__(self):
|
|
214
|
-
return hash((self.layer_name, self._repo_id, self._revision))
|
|
220
|
+
return hash((self.layer_name, self._repo_id, self._revision, self._trust_remote_code))
|
|
215
221
|
|
|
216
222
|
def __str__(self) -> str:
|
|
217
223
|
return f"`{self._repo_id}` (revision: {self._revision}), layer `{self.layer_name}`"
|