Files
modelscope-scepter/scepter/modules/utils/import_utils.py
T
2025-02-03 13:36:44 +08:00

334 lines
12 KiB
Python

# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import functools
import importlib
import logging
import os
import sys
from collections import OrderedDict
from importlib import import_module
from itertools import chain
from types import ModuleType
from typing import Any
import scepter
from scepter.modules.utils.ast_utils import (INDEX_KEY,
MODULE_KEY,
REQUIREMENT_KEY,
load_index)
from scepter.modules.utils.error import *
from scepter.modules.utils.logger import get_logger
if sys.version_info < (3, 8):
import importlib_metadata
else:
import importlib.metadata as importlib_metadata
logger = get_logger()
def get_dirname():
return os.path.dirname(scepter.__file__)
def import_modules(imports, allow_failed_imports=False):
"""Import modules from the given list of strings.
Args:
imports (list | str | None): The given module names to be imported.
allow_failed_imports (bool): If True, the failed imports will return
None. Otherwise, an ImportError is raise. Default: False.
Returns:
list[module] | module | None: The imported modules.
Examples:
>>> osp, sys = import_modules(
... ['os.path', 'sys'])
>>> import os.path as osp_
>>> import sys as sys_
>>> assert osp == osp_
>>> assert sys == sys_
"""
if not imports:
return
single_import = False
if isinstance(imports, str):
single_import = True
imports = [imports]
if not isinstance(imports, list):
raise TypeError(
f'custom_imports must be a list but got type {type(imports)}')
imported = []
for imp in imports:
if not isinstance(imp, str):
raise TypeError(
f'{imp} is of type {type(imp)} and cannot be imported.')
try:
imported_tmp = import_module(imp)
except ImportError:
if allow_failed_imports:
logger.warning(f'{imp} failed to import and is ignored.')
imported_tmp = None
else:
raise ImportError
imported.append(imported_tmp)
if single_import:
imported = imported[0]
return imported
# following code borrows implementation from huggingface/transformers
ENV_VARS_TRUE_VALUES = {'1', 'ON', 'YES', 'TRUE'}
ENV_VARS_TRUE_AND_AUTO_VALUES = ENV_VARS_TRUE_VALUES.union({'AUTO'})
USE_TORCH = os.environ.get('USE_TORCH', 'AUTO').upper()
_torch_version = 'N/A'
if USE_TORCH in ENV_VARS_TRUE_AND_AUTO_VALUES:
_torch_available = importlib.util.find_spec('torch') is not None
if _torch_available:
try:
_torch_version = importlib_metadata.version('torch')
logger.info(f'PyTorch version {_torch_version} Found.')
except importlib_metadata.PackageNotFoundError:
_torch_available = False
else:
logger.info('Disabling PyTorch because USE_TF is set')
_torch_available = False
def is_torchvision_available():
return importlib.util.find_spec('torchvision') is not None
def is_sentencepiece_available():
return importlib.util.find_spec('sentencepiece') is not None
def is_scepter_available():
return importlib.util.find_spec('scepter') is not None
def is_torch_available():
return _torch_available
def is_torch_cuda_available():
if is_torch_available():
import torch
return torch.cuda.is_available()
else:
return False
def is_swift_available():
return importlib.util.find_spec('swift') is not None
def is_opencv_available():
return importlib.util.find_spec('cv2') is not None
def is_pillow_available():
return importlib.util.find_spec('PIL.Image') is not None
def _is_package_available_fn(pkg_name):
return importlib.util.find_spec(pkg_name) is not None
def is_package_available(pkg_name):
return functools.partial(_is_package_available_fn, pkg_name)
def is_flash_attn_available():
return importlib.util.find_spec('flash-attn') is not None
def is_transformers_available():
return importlib.util.find_spec('transformers') is not None
REQUIREMENTS_MAAPING = OrderedDict([
('scepter', (is_scepter_available(), SCEPTER_IMPORT_ERROR)),
('torch', (is_torch_available, PYTORCH_IMPORT_ERROR)),
('torchvision', (is_torchvision_available(), TORCHVISION_IMPORT_ERROR)),
('cv2', (is_opencv_available, OPENCV_IMPORT_ERROR)),
('PIL', (is_pillow_available, PILLOW_IMPORT_ERROR)),
('modelscope', (is_package_available('modelscope'), MODELSCOPE_IMPORT_ERROR)),
('flash-attn', (is_flash_attn_available, FLASH_ATTN_IMPORT_ERROR)),
('xformers', (is_package_available('funasr'), XFORMERS_IMPORT_ERROR)),
('albumentations', (is_package_available('albumentations'), ALBUMENTATIONS_IMPORT_ERROR)),
('decord', (is_package_available('decord'), DECORD_IMPORT_ERROR)),
('beautifulsoup4', (is_package_available('beautifulsoup4'), BEAUTIFULSOUP4_IMPORT_ERROR)),
('bezier', (is_package_available('bezier'), BEZIER_IMPORT_ERROR)),
('einops', (is_package_available('einops'), EINOPS_IMPORT_ERROR)),
('numpy', (is_package_available('numpy'), NUMPY_IMPORT_ERROR)),
('oss2', (is_package_available('oss2'), OSS2_IMPORT_ERROR)),
('pycocotools', (is_package_available('pycocotools'), PYCOCOTOOLS_IMPORT_ERROR)),
('open_clip', (is_package_available('open_clip'), OPENCLIP_IMPORT_ERROR)),
('pyyaml', (is_package_available('pyyaml'), PYYAML_IMPORT_ERROR)),
('transformers', (is_package_available('transformers'), TRANSFORMERS_IMPORT_ERROR)),
('ms-swift', (is_package_available('ms-swift'), SWIFT_IMPORT_ERROR)),
('gradio', (is_package_available('gradio'), SWIFT_IMPORT_ERROR)),
('scikit-image', (is_package_available('scikit-image'), SCIKIT_IMAGE_IMPORT_ERROR)),
('scikit-learn', (is_package_available('scikit-learn'), SCIKIT_LEARN_IMPORT_ERROR)),
('sentencepiece', (is_package_available('sentencepiece'), SENTENCEPIECE_IMPORT_ERROR)),
('torchsde', (is_package_available('torchsde'), TORCHSDE_IMPORT_ERROR)),
('bitsandbytes', (is_package_available('bitsandbytes'), BITSANDBYTES_IMPORT_ERROR)),
('gradio_imageslider', (is_package_available('gradio_imageslider'), GRADIO_IMAGESLIDER_IMPORT_ERROR)),
('imagehash', (is_package_available('imagehash'), IMAGEHASH_IMPORT_ERROR)),
('psutil', (is_package_available('psutil'), PSUTIL_IMPORT_ERROR)),
('tiktoken', (is_package_available('tiktoken'), TIKTOKEN_IMPORT_ERROR))
])
SYSTEM_PACKAGE = set(['os', 'sys', 'typing'])
def requires(obj, requirements):
if not isinstance(requirements, (list, tuple)):
requirements = [requirements]
if isinstance(obj, str):
name = obj
else:
name = obj.__name__ if hasattr(obj,
'__name__') else obj.__class__.__name__
checks = []
for req in requirements:
if req == '' or req in SYSTEM_PACKAGE:
continue
if req in REQUIREMENTS_MAAPING:
check = REQUIREMENTS_MAAPING[req]
else:
check_fn = is_package_available(req)
err_msg = GENERAL_IMPORT_ERROR.replace('REQ', req)
check = (check_fn, err_msg)
checks.append(check)
failed = [msg.format(name) for available, msg in checks if not available]
if failed:
raise ImportError(''.join(failed))
def torch_required(func):
# Chose a different decorator name than in tests so it's clear they are not the same.
@functools.wraps(func)
def wrapper(*args, **kwargs):
if is_torch_available():
return func(*args, **kwargs)
else:
raise ImportError(f'Method `{func.__name__}` requires PyTorch.')
return wrapper
class LazyImportModule(ModuleType):
_AST_INDEX = None
def __init__(self,
name,
module_file,
import_structure,
module_spec=None,
extra_objects=None,
try_to_pre_import=False):
super().__init__(name)
self._modules = set(import_structure.keys())
self._class_to_module = {}
for key, values in import_structure.items():
for value in values:
self._class_to_module[value] = key
# Needed for autocompletion in an IDE
self.__all__ = list(import_structure.keys()) + list(
chain(*import_structure.values()))
self.__file__ = module_file
self.__spec__ = module_spec
self.__path__ = [os.path.dirname(module_file)]
self._objects = {} if extra_objects is None else extra_objects
self._name = name
self._import_structure = import_structure
if try_to_pre_import:
self._try_to_import()
def _try_to_import(self):
for sub_module in self._class_to_module.keys():
try:
getattr(self, sub_module)
except Exception as e:
logger.warning(
f'pre load module {sub_module} error, please check {e}')
def __dir__(self):
result = super().__dir__()
for attr in self.__all__:
if attr not in result:
result.append(attr)
return result
def __getattr__(self, name: str) -> Any:
if name in self._objects:
return self._objects[name]
if name in self._modules:
value = self._get_module(name)
elif name in self._class_to_module.keys():
module = self._get_module(self._class_to_module[name])
value = getattr(module, name)
else:
raise AttributeError(
f'module {self.__name__} has no attribute {name}')
setattr(self, name, value)
return value
def _get_module(self, module_name: str):
try:
module_name_full = self.__name__ + '.' + module_name
if not any(
module_name_full.startswith(f'scepter.{prefix}')
for prefix in ['modules', 'studio', 'version', 'tools', 'workflow']):
# check requirements before module import
requirements = self.get_requirements()
if module_name_full in requirements:
requires(module_name_full, requirements)
return importlib.import_module('.' + module_name, self.__name__)
except Exception as e:
raise RuntimeError(
f'Failed to import {self.__name__}.{module_name} because of the following error '
f'(look up to see its traceback):\n{e}') from e
def __reduce__(self):
return self.__class__, (self._name, self.__file__,
self._import_structure)
@staticmethod
def get_ast_index():
if LazyImportModule._AST_INDEX is None:
LazyImportModule._AST_INDEX = load_index()
return LazyImportModule._AST_INDEX
@staticmethod
def import_module(signature):
""" import a lazy import module using signature
Args:
signature (tuple): a tuple of str, (registry_name, class_name)
"""
ast_index = LazyImportModule.get_ast_index()
if signature in ast_index[INDEX_KEY]:
mod_index = ast_index[INDEX_KEY][signature]
module_name = mod_index[MODULE_KEY]
if module_name in ast_index[REQUIREMENT_KEY]:
requirements = ast_index[REQUIREMENT_KEY][module_name]
requires(module_name, requirements)
importlib.import_module(module_name)
else:
logger.warning(f'{signature} not found in ast index file')
@staticmethod
def get_module_type(module):
ast_index = LazyImportModule.get_ast_index()
if module in ast_index[INDEX_KEY]:
return True
else:
return False