v1.0.0 update

This commit is contained in:
hanzhn
2024-05-27 13:15:48 +08:00
parent 8076aae7da
commit c70ef0fc47
186 changed files with 7505 additions and 3117 deletions
@@ -4,7 +4,6 @@ from abc import ABCMeta
import torch
import torch.nn as nn
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.model.base_model import BaseModel
from scepter.modules.utils.config import dict_to_yaml
-1
View File
@@ -6,7 +6,6 @@ import cv2
import numpy as np
import torch
from PIL import Image
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
-1
View File
@@ -6,7 +6,6 @@ import cv2
import numpy as np
import torch
from PIL import Image
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
+15 -12
View File
@@ -11,13 +11,14 @@ from abc import ABCMeta
import cv2
import numpy as np
import torch
import torchvision
from einops import rearrange
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.distribute import we
from scepter.modules.utils.file_system import FS
from torchvision.transforms import InterpolationMode
def nms(x, t, s):
@@ -123,6 +124,8 @@ class HedAnnotator(BaseAnnotator, metaclass=ABCMeta):
if len(image.shape) == 3:
image = rearrange(image, 'h w c -> 1 c h w')
B, C, H, W = image.shape
elif len(image.shape) == 4:
B, C, H, W = image.shape
else:
raise "Unsurpport input image's shape"
elif isinstance(image, np.ndarray):
@@ -130,22 +133,22 @@ class HedAnnotator(BaseAnnotator, metaclass=ABCMeta):
if len(image.shape) == 3:
image = rearrange(image, 'h w c -> 1 c h w')
B, C, H, W = image.shape
elif len(image.shape) == 4:
B, C, H, W = image.shape
else:
raise "Unsurpport input image's shape"
else:
raise "Unsurpport input image's type"
transform = torchvision.transforms.Resize(
(H, W), interpolation=InterpolationMode.BILINEAR, antialias=True)
edges = self.netNetwork(image.to(we.device_id))
edges = [
e.detach().cpu().numpy().astype(np.float32)[0, 0] for e in edges
]
edges = [
cv2.resize(e, (W, H), interpolation=cv2.INTER_LINEAR)
for e in edges
]
edges = np.stack(edges, axis=2)
edge = 1 / (1 + np.exp(-np.mean(edges, axis=2).astype(np.float64)))
edge = 255 - (edge * 255.0).clip(0, 255).astype(np.uint8)
return edge[..., None].repeat(3, 2)
edges = [transform(e) for e in edges]
edges = torch.cat(edges, dim=1)
edges = 1 / (1 +
torch.exp(-torch.mean(edges, dim=1).type(torch.float)))
edges = edges.cpu().numpy()
edges = 255 - (edges * 255.0).clip(0, 255).astype(np.uint8)
return edges[..., None].repeat(3, -1)
@staticmethod
def get_config_template():
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# based on https://github.com/isl-org/MiDaS
import cv2
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
import torch.nn as nn
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
import torch.nn as nn
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
"""MidashNet: Network for monocular depth estimation trained by mixing several datasets.
This file contains code that is adapted from
https://github.com/thomasjpfan/pytorch_refinenet/blob/master/pytorch_refinenet/refinenet/refinenet_4cascade.py
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
"""MidashNet: Network for monocular depth estimation trained by mixing several datasets.
This file contains code that is adapted from
https://github.com/thomasjpfan/pytorch_refinenet/blob/master/pytorch_refinenet/refinenet/refinenet_4cascade.py
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import cv2
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
"""Utils for monoDepth."""
import re
import sys
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import math
import types
-1
View File
@@ -9,7 +9,6 @@ import numpy as np
import torch
from einops import rearrange
from PIL import Image
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.midas.api import MiDaSInference
from scepter.modules.annotator.registry import ANNOTATORS
@@ -0,0 +1,2 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
import torch.nn as nn
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
import torch
import torch.nn as nn
import torch.utils.model_zoo as model_zoo
+1
View File
@@ -1,4 +1,5 @@
# -*- coding: utf-8 -*-
# Copyright (c) Alibaba, Inc. and its affiliates.
# modified by lihaoweicv
# pytorch version
-1
View File
@@ -11,7 +11,6 @@ import cv2
import numpy as np
import torch
from PIL import Image
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.mlsd.mbv2_mlsd_large import MobileV2_MLSD_Large
from scepter.modules.annotator.mlsd.utils import pred_lines
+2 -3
View File
@@ -15,13 +15,12 @@ import numpy as np
import torch
import torch.nn as nn
from PIL import Image
from scipy.ndimage.filters import gaussian_filter
from skimage.measure import label
from scepter.modules.annotator.base_annotator import BaseAnnotator
from scepter.modules.annotator.registry import ANNOTATORS
from scepter.modules.utils.config import dict_to_yaml
from scepter.modules.utils.file_system import FS
from scipy.ndimage.filters import gaussian_filter
from skimage.measure import label
os.environ['KMP_DUPLICATE_LIB_OK'] = 'TRUE'
+5 -5
View File
@@ -37,30 +37,30 @@ class AnnotatorProcessor():
hed_cfg = {
'NAME': 'HedAnnotator',
'PRETRAINED_MODEL':
'ms://damo/scepter_scedit@annotator/ckpts/ControlNetHED.pth',
'ms://iic/scepter_scedit@annotator/ckpts/ControlNetHED.pth',
'INPUT_KEYS': ['img'],
'OUTPUT_KEYS': ['hed']
}
openpose_cfg = {
'NAME': 'OpenposeAnnotator',
'BODY_MODEL_PATH':
'ms://damo/scepter_scedit@annotator/ckpts/body_pose_model.pth',
'ms://iic/scepter_scedit@annotator/ckpts/body_pose_model.pth',
'HAND_MODEL_PATH':
'ms://damo/scepter_scedit@annotator/ckpts/hand_pose_model.pth',
'ms://iic/scepter_scedit@annotator/ckpts/hand_pose_model.pth',
'INPUT_KEYS': ['img'],
'OUTPUT_KEYS': ['openpose']
}
midas_cfg = {
'NAME': 'MidasDetector',
'PRETRAINED_MODEL':
'ms://damo/scepter_scedit@annotator/ckpts/dpt_hybrid-midas-501f0c75.pt',
'ms://iic/scepter_scedit@annotator/ckpts/dpt_hybrid-midas-501f0c75.pt',
'INPUT_KEYS': ['img'],
'OUTPUT_KEYS': ['depth']
}
mlsd_cfg = {
'NAME': 'MLSDdetector',
'PRETRAINED_MODEL':
'ms://damo/scepter_scedit@annotator/ckpts/mlsd_large_512_fp32.pth',
'ms://iic/scepter_scedit@annotator/ckpts/mlsd_large_512_fp32.pth',
'INPUT_KEYS': ['img'],
'OUTPUT_KEYS': ['mlsd']
}