Skip to content

Commit bf87817

Browse files
committed
add checking for imports
1 parent 75db907 commit bf87817

5 files changed

Lines changed: 65 additions & 4 deletions

File tree

src/diffusers/__init__.py

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,18 @@
11
# flake8: noqa
22
# There's no way to ignore "F401 '...' imported but unused" warnings in this
33
# module, but to preserve other warnings. So, don't check this module at all.
4-
from .utils import is_inflect_available, is_transformers_available, is_unidecode_available
4+
from .utils import (
5+
is_inflect_available,
6+
is_torch_geometric_available,
7+
is_transformers_available,
8+
is_unidecode_available,
9+
)
510

611

712
__version__ = "0.1.1"
813

914
from .modeling_utils import ModelMixin
10-
from .models import AutoencoderKL, MoleculeGNN, UNet2DConditionModel, UNet2DModel, VQModel
15+
from .models import AutoencoderKL, UNet2DConditionModel, UNet2DModel, VQModel
1116
from .pipeline_utils import DiffusionPipeline
1217
from .pipelines import DDIMPipeline, DDPMPipeline, LDMPipeline, PNDMPipeline, ScoreSdeVePipeline
1318
from .schedulers import DDIMScheduler, DDPMScheduler, PNDMScheduler, SchedulerMixin, ScoreSdeVeScheduler
@@ -17,3 +22,8 @@
1722
from .pipelines import LDMTextToImagePipeline
1823
else:
1924
from .utils.dummy_transformers_objects import *
25+
26+
if is_torch_geometric_available():
27+
from .models import MoleculeGNN
28+
else:
29+
from .utils.dummy_torch_geometric_objects import *

src/diffusers/models/__init__.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,12 @@
1616
# See the License for the specific language governing permissions and
1717
# limitations under the License.
1818

19-
from .molecule_gnn import MoleculeGNN
19+
from ..utils import is_torch_geometric_available
20+
21+
22+
if is_torch_geometric_available():
23+
from .molecule_gnn import MoleculeGNN
24+
2025
from .unet_2d import UNet2DModel
2126
from .unet_2d_condition import UNet2DConditionModel
2227
from .vae import AutoencoderKL, VQModel

src/diffusers/utils/__init__.py

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -68,6 +68,20 @@
6868
except importlib_metadata.PackageNotFoundError:
6969
_modelcards_available = False
7070

71+
_torch_scatter_available = importlib.util.find_spec("torch_scatter") is not None
72+
try:
73+
_torch_scatter_version = importlib_metadata.version("torch_scatter")
74+
logger.debug(f"Successfully imported torch_scatter version {_torch_scatter_version}")
75+
except importlib_metadata.PackageNotFoundError:
76+
_torch_scatter_available = False
77+
78+
_torch_scatter_available = importlib.util.find_spec("torch_geometric") is not None
79+
try:
80+
_torch_geometric_version = importlib_metadata.version("torch_geometric")
81+
logger.debug(f"Successfully imported torch_geometric version {_torch_geometric_version}")
82+
except importlib_metadata.PackageNotFoundError:
83+
_torch_geometric_available = False
84+
7185

7286
def is_transformers_available():
7387
return _transformers_available
@@ -85,6 +99,16 @@ def is_modelcards_available():
8599
return _modelcards_available
86100

87101

102+
def is_torch_scatter_available():
103+
return _torch_scatter_available
104+
105+
106+
def is_torch_geometric_available():
107+
# the model source of the Molecule Generation GNN requires a specific torch geometric version
108+
# for more info, see the original repo https://github.com/MinkaiXu/GeoDiff or our colab in readme
109+
return _torch_geometric_version == "1.7.2"
110+
111+
88112
class RepositoryNotFoundError(HTTPError):
89113
"""
90114
Raised when trying to access a hf.co URL with an invalid repository name, or with a private repo name the user does
@@ -117,6 +141,11 @@ class RevisionNotFoundError(HTTPError):
117141
inflect`
118142
"""
119143

144+
TORCH_GEOMETRIC_IMPORT_ERROR = """
145+
{0} requires version 1.7.2 of torch_geometric but it was not found in your environment. You can install it with conda:
146+
`conda install -c rusty1s pytorch-geometric=1.7.2`, given pytorch 1.8
147+
"""
148+
120149

121150
BACKENDS_MAPPING = OrderedDict(
122151
[
Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
# This file is autogenerated by the command `make fix-copies`, do not edit.
2+
# flake8: noqa
3+
from ..utils import DummyObject, requires_backends
4+
5+
6+
class MoleculeGNN(metaclass=DummyObject):
7+
_backends = ["torch_geometric"]
8+
9+
def __init__(self, *args, **kwargs):
10+
requires_backends(self, ["torch_geometric"])

tests/test_modeling_utils.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,14 +31,21 @@
3131
DDPMScheduler,
3232
LDMPipeline,
3333
LDMTextToImagePipeline,
34-
MoleculeGNN,
3534
PNDMPipeline,
3635
PNDMScheduler,
3736
ScoreSdeVePipeline,
3837
ScoreSdeVeScheduler,
3938
UNet2DModel,
4039
VQModel,
4140
)
41+
from diffusers.utils import is_torch_geometric_available
42+
43+
44+
if is_torch_geometric_available():
45+
from diffusers import MoleculeGNN
46+
else:
47+
from diffusers.utils.dummy_torch_geometric_objects import *
48+
4249
from diffusers.configuration_utils import ConfigMixin, register_to_config
4350
from diffusers.pipeline_utils import DiffusionPipeline
4451
from diffusers.testing_utils import floats_tensor, slow, torch_device

0 commit comments

Comments
 (0)