Skip to content
Merged
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
14 changes: 9 additions & 5 deletions emle/models/_mace.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,8 +34,8 @@
from typing import List, Dict, Optional

from ._emle import EMLE as _EMLE
from ._emle import _has_nnpops
from ._utils import _get_neighbor_pairs
from ._utils import _has_neighbor_pairs

from torch import Tensor

Expand Down Expand Up @@ -172,8 +172,10 @@ def __init__(
)
if not _has_e3nn:
raise ImportError("e3nn is required to compile the MACEmodel.")
if not _has_nnpops:
raise ImportError("NNPOps is required to use the MACEEMLE model.")
if not _has_neighbor_pairs:
raise ImportError(
"NNPOps.neighbors.getNeighborPairs is required to use the MACEEMLE model."
)

if device is not None:
if not isinstance(device, _torch.device):
Expand Down Expand Up @@ -805,8 +807,10 @@ def __init__(
)
if not _has_e3nn:
raise ImportError("e3nn is required to compile the MACEmodel.")
if not _has_nnpops:
raise ImportError("NNPOps is required to use the MACEEMLE model.")
if not _has_neighbor_pairs:
raise ImportError(
"NNPOps.neighbors.getNeighborPairs is required to use the MACEEMLE model."
)

if device is not None:
if not isinstance(device, _torch.device):
Expand Down
6 changes: 4 additions & 2 deletions emle/models/_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,9 +32,11 @@
from typing import Optional, Tuple

try:
import NNPOps.neighbors.getNeighborPairs as _getNeighborPairs
from NNPOps.neighbors import getNeighborPairs as _getNeighborPairs

_has_neighbor_pairs = True
except:
pass
_has_neighbor_pairs = False


_DEPRECATED_ALPHA_MODES = {"species": "fixed", "reference": "flexible"}
Expand Down
Loading