jburtoft's picture
docs: add revision=v1.0.0 + trust_remote_code=True to get_kernel() calls
2a17ff8
Raw
History Blame Contribute Delete
1.26 kB
"""Shared loader helper for the examples.
Provides `load_kernel()` which returns the mamba3 kernel package. Uses the
local clone (examples/../build/torch-neuron/) when running from a git clone,
or falls back to the HF Hub via `get_kernel("jburtoft/mamba3-neuron-kernels")`
when installed via the `kernels` library.
"""
import os
import sys
def load_kernel():
"""Return the mamba3 kernel module."""
local_path = os.path.abspath(
os.path.join(os.path.dirname(__file__), "..", "build", "torch-neuron", "__init__.py")
)
if os.path.exists(local_path):
# print(f"[loader] using local clone: {os.path.dirname(local_path)}")
import importlib.util
spec = importlib.util.spec_from_file_location(
"mamba3", local_path,
submodule_search_locations=[os.path.dirname(local_path)],
)
m = importlib.util.module_from_spec(spec)
sys.modules["mamba3"] = m
spec.loader.exec_module(m)
return m
else:
# print("[loader] using HF Hub: jburtoft/mamba3-neuron-kernels")
from kernels import get_kernel
return get_kernel(
"jburtoft/mamba3-neuron-kernels",
revision="v1.0.0",
trust_remote_code=True,
)