-
Notifications
You must be signed in to change notification settings - Fork 297
Expand file tree
/
Copy path__init__.py
More file actions
41 lines (34 loc) · 1.23 KB
/
Copy path__init__.py
File metadata and controls
41 lines (34 loc) · 1.23 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import importlib
import sys
from typing import Any
# Implementations live with their owning operation families. Expose them here
# lazily so importing this experimental compatibility package does not pull in
# PyTorch until a specific operation is requested.
_LAZY_ALIASES = {
"moe_grouped_matmul": "cudnn.gemm.ops.moe_grouped_matmul",
"swiglu_mlp": "cudnn.gemm.ops.swiglu_mlp",
"rms_norm": "cudnn.ops.norm.rmsnorm",
"layer_norm": "cudnn.ops.norm.layernorm",
}
def __getattr__(name: str) -> Any:
try:
target = _LAZY_ALIASES[name]
except KeyError:
raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from None
try:
module = importlib.import_module(target)
except ImportError as error:
from cudnn import _optional_dependency_message
raise ImportError(_optional_dependency_message(name, error)) from error
sys.modules[f"{__name__}.{name}"] = module
value = getattr(module, name)
globals()[name] = value
return value
__all__ = [
"moe_grouped_matmul",
"swiglu_mlp",
"rms_norm",
"layer_norm",
]