update 1.4.0
This commit is contained in:
@@ -1,23 +1,60 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) Alibaba, Inc. and its affiliates.
|
||||
from scepter.modules.annotator.base_annotator import GeneralAnnotator
|
||||
from scepter.modules.annotator.canny import CannyAnnotator
|
||||
from scepter.modules.annotator.color import ColorAnnotator
|
||||
from scepter.modules.annotator.degradation import DegradationAnnotator
|
||||
from scepter.modules.annotator.doodle import DoodleAnnotator
|
||||
from scepter.modules.annotator.gray import GrayAnnotator
|
||||
from scepter.modules.annotator.hed import HedAnnotator
|
||||
from scepter.modules.annotator.identity import IdentityAnnotator
|
||||
from scepter.modules.annotator.informative_drawing import (
|
||||
InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
|
||||
InfoDrawOpenSketchAnnotator)
|
||||
from scepter.modules.annotator.inpainting import InpaintingAnnotator
|
||||
from scepter.modules.annotator.invert import InvertAnnotator
|
||||
from scepter.modules.annotator.midas_op import MidasDetector
|
||||
from scepter.modules.annotator.mlsd_op import MLSDdetector
|
||||
from scepter.modules.annotator.openpose import OpenposeAnnotator
|
||||
from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
|
||||
from scepter.modules.annotator.pidinet import PiDiAnnotator
|
||||
from scepter.modules.annotator.segmentation import ESAMAnnotator
|
||||
from scepter.modules.annotator.sketch import SketchAnnotator
|
||||
from scepter.modules.annotator.lama import LamaAnnotator
|
||||
from typing import TYPE_CHECKING
|
||||
from scepter.modules.utils.import_utils import LazyImportModule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from scepter.modules.annotator.base_annotator import GeneralAnnotator
|
||||
from scepter.modules.annotator.canny import CannyAnnotator
|
||||
from scepter.modules.annotator.color import ColorAnnotator
|
||||
from scepter.modules.annotator.degradation import DegradationAnnotator
|
||||
from scepter.modules.annotator.doodle import DoodleAnnotator
|
||||
from scepter.modules.annotator.gray import GrayAnnotator
|
||||
from scepter.modules.annotator.hed import HedAnnotator
|
||||
from scepter.modules.annotator.identity import IdentityAnnotator
|
||||
from scepter.modules.annotator.informative_drawing import (
|
||||
InfoDrawAnimeAnnotator, InfoDrawContourAnnotator,
|
||||
InfoDrawOpenSketchAnnotator)
|
||||
from scepter.modules.annotator.inpainting import InpaintingAnnotator
|
||||
from scepter.modules.annotator.invert import InvertAnnotator
|
||||
from scepter.modules.annotator.midas_op import MidasDetector
|
||||
from scepter.modules.annotator.mlsd_op import MLSDdetector
|
||||
from scepter.modules.annotator.openpose import OpenposeAnnotator
|
||||
from scepter.modules.annotator.outpainting import OutpaintingAnnotator, OutpaintingResize
|
||||
from scepter.modules.annotator.pidinet import PiDiAnnotator
|
||||
from scepter.modules.annotator.segmentation import ESAMAnnotator
|
||||
from scepter.modules.annotator.sketch import SketchAnnotator
|
||||
from scepter.modules.annotator.lama import LamaAnnotator
|
||||
else:
|
||||
_import_structure = {
|
||||
'base_annotator': ['GeneralAnnotator'],
|
||||
'canny': ['CannyAnnotator'],
|
||||
'color': ['ColorAnnotator'],
|
||||
'degradation': ['DegradationAnnotator'],
|
||||
'doodle': ['DoodleAnnotator'],
|
||||
'gray': ['GrayAnnotator'],
|
||||
'hed': ['HedAnnotator'],
|
||||
'identity': ['IdentityAnnotator'],
|
||||
'informative_drawing': ['InfoDrawAnimeAnnotator',
|
||||
'InfoDrawContourAnnotator',
|
||||
'InfoDrawOpenSketchAnnotator'],
|
||||
'inpainting': ['InpaintingAnnotator'],
|
||||
'invert': ['InvertAnnotator'],
|
||||
'midas_op': ['MidasDetector'],
|
||||
'mlsd_op': ['MLSDdetector'],
|
||||
'openpose': ['OpenposeAnnotator'],
|
||||
'outpainting': ['OutpaintingAnnotator', 'OutpaintingResize'],
|
||||
'pidinet': ['PiDiAnnotator'],
|
||||
'segmentation': ['ESAMAnnotator'],
|
||||
'sketch': ['SketchAnnotator'],
|
||||
'lama': ['LamaAnnotator'],
|
||||
}
|
||||
|
||||
import sys
|
||||
sys.modules[__name__] = LazyImportModule(
|
||||
__name__,
|
||||
globals()['__file__'],
|
||||
_import_structure,
|
||||
module_spec=__spec__,
|
||||
extra_objects={},
|
||||
)
|
||||
|
||||
@@ -114,7 +114,7 @@ class HedAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
self.netNetwork.load_state_dict(torch.load(local_path))
|
||||
self.netNetwork.load_state_dict(torch.load(local_path, weights_only=True))
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
|
||||
@@ -120,7 +120,7 @@ class InfoDrawContourAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
self.model = ContourInference(input_nc, output_nc, n_residual_blocks,
|
||||
sigmoid)
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
self.model.load_state_dict(torch.load(local_path))
|
||||
self.model.load_state_dict(torch.load(local_path, weights_only=True))
|
||||
self.model = self.model.eval().requires_grad_(False).to(we.device_id)
|
||||
|
||||
@torch.no_grad()
|
||||
|
||||
@@ -10,7 +10,7 @@ class BaseModel(torch.nn.Module):
|
||||
Args:
|
||||
path (str): file path
|
||||
"""
|
||||
parameters = torch.load(path, map_location=torch.device('cpu'))
|
||||
parameters = torch.load(path, map_location=torch.device('cpu'), weights_only=True)
|
||||
|
||||
if 'optimizer' in parameters:
|
||||
parameters = parameters['model']
|
||||
|
||||
@@ -29,7 +29,7 @@ class MLSDdetector(BaseAnnotator, metaclass=ABCMeta):
|
||||
pretrained_model = cfg.get('PRETRAINED_MODEL', None)
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
model.load_state_dict(torch.load(local_path), strict=True)
|
||||
model.load_state_dict(torch.load(local_path, weights_only=True), strict=True)
|
||||
self.model = model.eval()
|
||||
self.thr_v = cfg.get('THR_V', 0.1)
|
||||
self.thr_d = cfg.get('THR_D', 0.1)
|
||||
|
||||
@@ -423,7 +423,7 @@ class Hand(object):
|
||||
self.model = handpose_model()
|
||||
if torch.cuda.is_available():
|
||||
self.model = self.model.to(device)
|
||||
model_dict = transfer(self.model, torch.load(model_path))
|
||||
model_dict = transfer(self.model, torch.load(model_path, weights_only=True))
|
||||
self.model.load_state_dict(model_dict)
|
||||
self.model.eval()
|
||||
self.device = device
|
||||
@@ -503,7 +503,7 @@ class Body(object):
|
||||
self.model = bodypose_model()
|
||||
if torch.cuda.is_available():
|
||||
self.model = self.model.to(device)
|
||||
model_dict = transfer(self.model, torch.load(model_path))
|
||||
model_dict = transfer(self.model, torch.load(model_path, weights_only=True))
|
||||
self.model.load_state_dict(model_dict)
|
||||
self.model.eval()
|
||||
self.device = device
|
||||
|
||||
@@ -882,7 +882,7 @@ class PiDiAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
state = torch.load(local_path,
|
||||
map_location='cpu')['state_dict']
|
||||
map_location='cpu', weights_only=True)['state_dict']
|
||||
if vanilla_cnn:
|
||||
state = convert_pidinet(state, 'carv4')
|
||||
state = {
|
||||
|
||||
@@ -10,7 +10,11 @@ import torchvision.transforms as T
|
||||
from PIL import Image
|
||||
from pycocotools import mask as mask_utils
|
||||
from scipy import ndimage
|
||||
from sklearn.cluster import KMeans
|
||||
try:
|
||||
from sklearn.cluster import KMeans
|
||||
except:
|
||||
import warnings
|
||||
warnings.warn("ignore sklearn import, please pip install scikit-learn.")
|
||||
from torchvision.ops.boxes import batched_nms
|
||||
|
||||
from scepter.modules.annotator.base_annotator import BaseAnnotator
|
||||
|
||||
@@ -86,7 +86,7 @@ class SketchAnnotator(BaseAnnotator, metaclass=ABCMeta):
|
||||
std=0.0858381272736797).eval()
|
||||
if pretrained_model:
|
||||
with FS.get_from(pretrained_model, wait_finish=True) as local_path:
|
||||
state = torch.load(local_path, map_location='cpu')
|
||||
state = torch.load(local_path, map_location='cpu', weights_only=True)
|
||||
self.model.load_state_dict(state)
|
||||
|
||||
@torch.no_grad()
|
||||
|
||||
Reference in New Issue
Block a user