初始化:换发型/换发色/训练发型服务
包含: - hair_service_sd: 主服务(换发型/换发色/生发,端口8801) - photo_service: LoRA调度+训练(端口32678) - hair_grow_service: 调试测试页(端口8888,含4个测试页) - 批量训练脚本(batch_train_hairstyles.py) - 发际线mask自动识别(hairline_mask.py,4种方案) - 手绘mask换发型(hair_swap_manual.py) - 文档:README.md + LARGE_FILES.md + docs/ 大文件(模型权重200G、训练数据123G)已排除,见 LARGE_FILES.md OSS/COS密钥已脱敏为环境变量,原文件备份在本地
This commit is contained in:
@@ -0,0 +1,75 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
def get_norm(norm, out_channels=None):
|
||||
"""
|
||||
Args:
|
||||
norm (str or callable):
|
||||
|
||||
Returns:
|
||||
nn.Module or None: the normalization layer
|
||||
"""
|
||||
if isinstance(norm, str):
|
||||
if len(norm) == 0:
|
||||
return None
|
||||
norm = {
|
||||
"BN": nn.BatchNorm2d,
|
||||
"IN": nn.InstanceNorm2d,
|
||||
"GN": lambda channels: nn.GroupNorm(32, channels),
|
||||
"nnSyncBN": nn.SyncBatchNorm, # keep for debugging
|
||||
}[norm]
|
||||
if out_channels is not None:
|
||||
return norm(out_channels)
|
||||
else:
|
||||
return norm
|
||||
|
||||
class Conv2d(torch.nn.Conv2d):
|
||||
"""
|
||||
A wrapper around :class:`torch.nn.Conv2d` to support zero-size tensor and more features.
|
||||
"""
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
"""
|
||||
Extra keyword arguments supported in addition to those in `torch.nn.Conv2d`:
|
||||
|
||||
Args:
|
||||
norm (nn.Module, optional): a normalization layer
|
||||
activation (callable(Tensor) -> Tensor): a callable activation function
|
||||
|
||||
It assumes that norm layer is used before activation.
|
||||
"""
|
||||
norm = kwargs.pop("norm", None)
|
||||
activation = kwargs.pop("activation", None)
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
self.norm = norm
|
||||
self.activation = activation
|
||||
|
||||
def forward(self, x):
|
||||
x = super().forward(x)
|
||||
if self.norm is not None:
|
||||
x = self.norm(x)
|
||||
if self.activation is not None:
|
||||
x = self.activation(x)
|
||||
return x
|
||||
|
||||
class Backbone(nn.Module):
|
||||
"""
|
||||
Abstract base class for network backbones.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
"""
|
||||
The `__init__` method of any subclass can specify its own set of arguments.
|
||||
"""
|
||||
super().__init__()
|
||||
|
||||
def forward(self):
|
||||
"""
|
||||
Subclasses must override this method, but adhere to the same return type.
|
||||
|
||||
Returns:
|
||||
dict[str: Tensor]: mapping from feature name (e.g., "res2") to tensor
|
||||
"""
|
||||
pass
|
||||
@@ -0,0 +1,268 @@
|
||||
import math
|
||||
import utils.weight_init as weight_init
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from bodyseg.backbone.backbone import Backbone, get_norm, Conv2d
|
||||
from bodyseg.backbone.resnet import build_resnet_backbone
|
||||
|
||||
class FPN(Backbone):
|
||||
"""
|
||||
This module implements Feature Pyramid Network.
|
||||
It creates pyramid features built on top of some input feature maps.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self, bottom_up, in_features, out_channels, norm="", top_block=None, fuse_type="sum"
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
bottom_up (Backbone): module representing the bottom up subnetwork.
|
||||
Must be a subclass of :class:`Backbone`. The multi-scale feature
|
||||
maps generated by the bottom up network, and listed in `in_features`,
|
||||
are used to generate FPN levels.
|
||||
in_features (list[str]): names of the input feature maps coming
|
||||
from the backbone to which FPN is attached. For example, if the
|
||||
backbone produces ["res2", "res3", "res4"], any *contiguous* sublist
|
||||
of these may be used; order must be from high to low resolution.
|
||||
out_channels (int): number of channels in the output feature maps.
|
||||
norm (str): the normalization to use.
|
||||
top_block (nn.Module or None): if provided, an extra operation will
|
||||
be performed on the output of the last (smallest resolution)
|
||||
FPN output, and the result will extend the result list. The top_block
|
||||
further downsamples the feature map. It must have an attribute
|
||||
"num_levels", meaning the number of extra FPN levels added by
|
||||
this block, and "in_feature", which is a string representing
|
||||
its input feature (e.g., p5).
|
||||
fuse_type (str): types for fusing the top down features and the lateral
|
||||
ones. It can be "sum" (default), which sums up element-wise; or "avg",
|
||||
which takes the element-wise mean of the two.
|
||||
"""
|
||||
super(FPN, self).__init__()
|
||||
assert isinstance(bottom_up, Backbone)
|
||||
|
||||
# Feature map strides and channels from the bottom up network (e.g. ResNet)
|
||||
in_strides = [bottom_up._out_feature_strides[f] for f in in_features]
|
||||
in_channels = [bottom_up._out_feature_channels[f] for f in in_features]
|
||||
|
||||
_assert_strides_are_log2_contiguous(in_strides)
|
||||
lateral_convs = []
|
||||
output_convs = []
|
||||
|
||||
use_bias = norm == ""
|
||||
for idx, in_channels in enumerate(in_channels):
|
||||
lateral_norm = get_norm(norm, out_channels)
|
||||
output_norm = get_norm(norm, out_channels)
|
||||
|
||||
lateral_conv = Conv2d(
|
||||
in_channels, out_channels, kernel_size=1, bias=use_bias, norm=lateral_norm
|
||||
)
|
||||
output_conv = Conv2d(
|
||||
out_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
bias=use_bias,
|
||||
norm=output_norm,
|
||||
)
|
||||
weight_init.c2_xavier_fill(lateral_conv)
|
||||
weight_init.c2_xavier_fill(output_conv)
|
||||
stage = int(math.log2(in_strides[idx]))
|
||||
|
||||
lateral_convs.append(lateral_conv)
|
||||
output_convs.append(output_conv)
|
||||
# Place convs into top-down order (from low to high resolution)
|
||||
# to make the top-down computation in forward clearer.
|
||||
self.lateral_convs = nn.ModuleList(lateral_convs[::-1])
|
||||
self.output_convs = nn.ModuleList(output_convs[::-1])
|
||||
self.top_block = top_block
|
||||
self.in_features = in_features
|
||||
self.bottom_up = bottom_up
|
||||
# Return feature names are "p<stage>", like ["p2", "p3", ..., "p6"]
|
||||
self._out_feature_strides = {"p{}".format(int(math.log2(s))): s for s in in_strides}
|
||||
# top block output feature maps.
|
||||
if self.top_block is not None:
|
||||
for s in range(stage, stage + self.top_block.num_levels):
|
||||
self._out_feature_strides["p{}".format(s + 1)] = 2 ** (s + 1)
|
||||
|
||||
self._out_features = list(self._out_feature_strides.keys())
|
||||
self._out_feature_channels = {k: out_channels for k in self._out_features}
|
||||
assert fuse_type in {"avg", "sum"}
|
||||
self._fuse_type = fuse_type
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
Args:
|
||||
input (dict[str: Tensor]): mapping feature map name (e.g., "res5") to
|
||||
feature map tensor for each feature level in high to low resolution order.
|
||||
|
||||
Returns:
|
||||
dict[str: Tensor]:
|
||||
mapping from feature map name to FPN feature map tensor
|
||||
in high to low resolution order. Returned feature names follow the FPN
|
||||
paper convention: "p<stage>", where stage has stride = 2 ** stage e.g.,
|
||||
["p2", "p3", ..., "p6"].
|
||||
"""
|
||||
# Reverse feature maps into top-down order (from low to high resolution)
|
||||
bottom_up_features = self.bottom_up(x)
|
||||
x = [bottom_up_features[f] for f in self.in_features[::-1]]
|
||||
results = []
|
||||
prev_features = self.lateral_convs[0](x[0])
|
||||
results.append(self.output_convs[0](prev_features))
|
||||
for features, lateral_conv, output_conv in zip(
|
||||
x[1:], self.lateral_convs[1:], self.output_convs[1:]
|
||||
):
|
||||
top_down_features = F.interpolate(prev_features, scale_factor=2, mode="nearest")
|
||||
lateral_features = lateral_conv(features)
|
||||
prev_features = lateral_features + top_down_features
|
||||
if self._fuse_type == "avg":
|
||||
prev_features /= 2
|
||||
results.insert(0, output_conv(prev_features))
|
||||
|
||||
if self.top_block is not None:
|
||||
top_block_in_feature = bottom_up_features.get(self.top_block.in_feature, None)
|
||||
if top_block_in_feature is None:
|
||||
top_block_in_feature = results[self._out_features.index(self.top_block.in_feature)]
|
||||
results.extend(self.top_block(top_block_in_feature))
|
||||
assert len(self._out_features) == len(results)
|
||||
return dict(zip(self._out_features, results))
|
||||
|
||||
def _assert_strides_are_log2_contiguous(strides):
|
||||
"""
|
||||
Assert that each stride is 2x times its preceding stride, i.e. "contiguous in log2".
|
||||
"""
|
||||
for i, stride in enumerate(strides[1:], 1):
|
||||
assert stride == 2 * strides[i - 1], "Strides {} {} are not log2 contiguous".format(
|
||||
stride, strides[i - 1]
|
||||
)
|
||||
|
||||
|
||||
class LastLevelMaxPool(nn.Module):
|
||||
"""
|
||||
This module is used in the original FPN to generate a downsampled
|
||||
P6 feature from P5.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.num_levels = 1
|
||||
self.in_feature = "p5"
|
||||
|
||||
def forward(self, x):
|
||||
return [F.max_pool2d(x, kernel_size=1, stride=2, padding=0)]
|
||||
|
||||
|
||||
class LastLevelP6P7(nn.Module):
|
||||
"""
|
||||
This module is used in RetinaNet to generate extra layers, P6 and P7 from
|
||||
C5 feature.
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels, out_channels):
|
||||
super().__init__()
|
||||
self.num_levels = 2
|
||||
self.in_feature = "res5"
|
||||
self.p6 = nn.Conv2d(in_channels, out_channels, 3, 2, 1)
|
||||
self.p7 = nn.Conv2d(out_channels, out_channels, 3, 2, 1)
|
||||
for module in [self.p6, self.p7]:
|
||||
weight_init.c2_xavier_fill(module)
|
||||
|
||||
def forward(self, c5):
|
||||
p6 = self.p6(c5)
|
||||
p7 = self.p7(F.relu(p6))
|
||||
return [p6, p7]
|
||||
|
||||
|
||||
def build_resnet_fpn_backbone(in_channels=3):
|
||||
"""
|
||||
Args:
|
||||
cfg: a detectron2 CfgNode
|
||||
|
||||
Returns:
|
||||
backbone (Backbone): backbone module, must be a subclass of :class:`Backbone`.
|
||||
"""
|
||||
bottom_up = build_resnet_backbone(in_channels)
|
||||
in_features = ["res2", "res3", "res4"]
|
||||
out_channels = 256
|
||||
backbone = FPN(
|
||||
bottom_up=bottom_up,
|
||||
in_features=in_features,
|
||||
out_channels=out_channels,
|
||||
norm="BN",
|
||||
# top_block=LastLevelMaxPool(),
|
||||
top_block=None,
|
||||
fuse_type="sum",
|
||||
)
|
||||
return backbone
|
||||
|
||||
def build_retinanet_resnet_fpn_backbone(cfg, in_channels=3):
|
||||
"""
|
||||
Args:
|
||||
cfg: a detectron2 CfgNode
|
||||
|
||||
Returns:
|
||||
backbone (Backbone): backbone module, must be a subclass of :class:`Backbone`.
|
||||
"""
|
||||
bottom_up = build_resnet_backbone(cfg, in_channels)
|
||||
in_features = cfg.MODEL.FPN.IN_FEATURES
|
||||
out_channels = cfg.MODEL.FPN.OUT_CHANNELS
|
||||
in_channels_p6p7 = bottom_up._out_feature_channels["res5"]
|
||||
backbone = FPN(
|
||||
bottom_up=bottom_up,
|
||||
in_features=in_features,
|
||||
out_channels=out_channels,
|
||||
norm=cfg.MODEL.FPN.NORM,
|
||||
top_block=LastLevelP6P7(in_channels_p6p7, out_channels),
|
||||
fuse_type=cfg.MODEL.FPN.FUSE_TYPE,
|
||||
)
|
||||
return backbone
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
from config.default import get_cfg
|
||||
|
||||
def setup(args):
|
||||
"""
|
||||
Create configs and perform basic setups.
|
||||
"""
|
||||
cfg = get_cfg()
|
||||
cfg.merge_from_file(args.cfg)
|
||||
cfg.merge_from_list(args.opts)
|
||||
cfg.freeze()
|
||||
return cfg
|
||||
|
||||
parser = argparse.ArgumentParser(description='Train ImageNet network')
|
||||
# general
|
||||
parser.add_argument('--cfg',
|
||||
help='experiment configure file name',
|
||||
required=True,
|
||||
type=str)
|
||||
|
||||
parser.add_argument('opts',
|
||||
help="Modify config options using the command-line",
|
||||
default=None,
|
||||
nargs=argparse.REMAINDER)
|
||||
|
||||
args = parser.parse_args()
|
||||
cfg = setup(args)
|
||||
print(cfg)
|
||||
|
||||
model = build_resnet_fpn_backbone(cfg, 3)
|
||||
# model = build_retinanet_resnet_fpn_backbone(cfg, 3)
|
||||
print(model)
|
||||
# model = torch.nn.DataParallel(model, list(range(2))).cuda()
|
||||
dummy_input = torch.randn(4, 3, 512, 512)
|
||||
|
||||
out = model(dummy_input)
|
||||
|
||||
for k, v in out.items():
|
||||
print(k, v.shape)
|
||||
|
||||
# torch.onnx.export(model, dummy_input, "tmp.onnx", verbose=True,
|
||||
# input_names=['input'],
|
||||
# output_names=['output'])
|
||||
|
||||
pass
|
||||
|
||||
@@ -0,0 +1,298 @@
|
||||
import numpy as np
|
||||
import utils.weight_init as weight_init
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from bodyseg.backbone.backbone import Backbone, get_norm, Conv2d
|
||||
|
||||
class BasicStem(nn.Module):
|
||||
def __init__(self, in_channels=3, out_channels=64, norm="BN"):
|
||||
"""
|
||||
Args:
|
||||
norm (str or callable): a callable that takes the number of
|
||||
channels and return a `nn.Module`, or a pre-defined string
|
||||
(one of {"FrozenBN", "BN", "GN"}).
|
||||
"""
|
||||
super().__init__()
|
||||
self.conv1 = Conv2d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=7,
|
||||
stride=2,
|
||||
padding=3,
|
||||
bias=False,
|
||||
norm=get_norm(norm, out_channels),
|
||||
)
|
||||
weight_init.c2_msra_fill(self.conv1)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv1(x)
|
||||
x = F.relu_(x)
|
||||
x = F.max_pool2d(x, kernel_size=3, stride=2, padding=1)
|
||||
return x
|
||||
|
||||
@property
|
||||
def out_channels(self):
|
||||
return self.conv1.out_channels
|
||||
|
||||
@property
|
||||
def stride(self):
|
||||
return 4 # = stride 2 conv -> stride 2 max pool
|
||||
|
||||
class ResNetBlockBase(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, stride):
|
||||
"""
|
||||
The `__init__` method of any subclass should also contain these arguments.
|
||||
|
||||
Args:
|
||||
in_channels (int):
|
||||
out_channels (int):
|
||||
stride (int):
|
||||
"""
|
||||
super().__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.stride = stride
|
||||
|
||||
class BottleneckBlock(ResNetBlockBase):
|
||||
def __init__(
|
||||
self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
*,
|
||||
bottleneck_channels,
|
||||
stride=1,
|
||||
num_groups=1,
|
||||
norm="BN",
|
||||
stride_in_1x1=False,
|
||||
dilation=1,
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
norm (str or callable): a callable that takes the number of
|
||||
channels and return a `nn.Module`, or a pre-defined string
|
||||
(one of {"FrozenBN", "BN", "GN"}).
|
||||
stride_in_1x1 (bool): when stride==2, whether to put stride in the
|
||||
first 1x1 convolution or the bottleneck 3x3 convolution.
|
||||
"""
|
||||
super().__init__(in_channels, out_channels, stride)
|
||||
|
||||
if in_channels != out_channels:
|
||||
self.shortcut = Conv2d(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=1,
|
||||
stride=stride,
|
||||
bias=False,
|
||||
norm=get_norm(norm, out_channels),
|
||||
)
|
||||
else:
|
||||
self.shortcut = None
|
||||
|
||||
# The original MSRA ResNet models have stride in the first 1x1 conv
|
||||
# The subsequent fb.torch.resnet and Caffe2 ResNe[X]t implementations have
|
||||
# stride in the 3x3 conv
|
||||
stride_1x1, stride_3x3 = (stride, 1) if stride_in_1x1 else (1, stride)
|
||||
|
||||
self.conv1 = Conv2d(
|
||||
in_channels,
|
||||
bottleneck_channels,
|
||||
kernel_size=1,
|
||||
stride=stride_1x1,
|
||||
bias=False,
|
||||
norm=get_norm(norm, bottleneck_channels),
|
||||
)
|
||||
|
||||
self.conv2 = Conv2d(
|
||||
bottleneck_channels,
|
||||
bottleneck_channels,
|
||||
kernel_size=3,
|
||||
stride=stride_3x3,
|
||||
padding=1 * dilation,
|
||||
bias=False,
|
||||
groups=num_groups,
|
||||
dilation=dilation,
|
||||
norm=get_norm(norm, bottleneck_channels),
|
||||
)
|
||||
|
||||
self.conv3 = Conv2d(
|
||||
bottleneck_channels,
|
||||
out_channels,
|
||||
kernel_size=1,
|
||||
bias=False,
|
||||
norm=get_norm(norm, out_channels),
|
||||
)
|
||||
|
||||
for layer in [self.conv1, self.conv2, self.conv3, self.shortcut]:
|
||||
if layer is not None: # shortcut can be None
|
||||
weight_init.c2_msra_fill(layer)
|
||||
|
||||
# Zero-initialize the last normalization in each residual branch,
|
||||
# so that at the beginning, the residual branch starts with zeros,
|
||||
# and each residual block behaves like an identity.
|
||||
# See Sec 5.1 in "Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour":
|
||||
# "For BN layers, the learnable scaling coefficient γ is initialized
|
||||
# to be 1, except for each residual block's last BN
|
||||
# where γ is initialized to be 0."
|
||||
|
||||
# nn.init.constant_(self.conv3.norm.weight, 0)
|
||||
# TODO this somehow hurts performance when training GN models from scratch.
|
||||
# Add it as an option when we need to use this code to train a backbone.
|
||||
|
||||
def forward(self, x):
|
||||
out = self.conv1(x)
|
||||
out = F.relu_(out)
|
||||
|
||||
out = self.conv2(out)
|
||||
out = F.relu_(out)
|
||||
|
||||
out = self.conv3(out)
|
||||
|
||||
if self.shortcut is not None:
|
||||
shortcut = self.shortcut(x)
|
||||
else:
|
||||
shortcut = x
|
||||
|
||||
out += shortcut
|
||||
out = F.relu_(out)
|
||||
return out
|
||||
|
||||
def make_stage(block_class, num_blocks, first_stride, **kwargs):
|
||||
"""
|
||||
Create a resnet stage by creating many blocks.
|
||||
Args:
|
||||
block_class (class): a subclass of ResNetBlockBase
|
||||
num_blocks (int):
|
||||
first_stride (int): the stride of the first block. The other blocks will have stride=1.
|
||||
A `stride` argument will be passed to the block constructor.
|
||||
kwargs: other arguments passed to the block constructor.
|
||||
|
||||
Returns:
|
||||
list[nn.Module]: a list of block module.
|
||||
"""
|
||||
blocks = []
|
||||
for i in range(num_blocks):
|
||||
blocks.append(block_class(stride=first_stride if i == 0 else 1, **kwargs))
|
||||
kwargs["in_channels"] = kwargs["out_channels"]
|
||||
return blocks
|
||||
|
||||
class ResNet(Backbone):
|
||||
def __init__(self, stem, stages, num_classes=None, out_features=None):
|
||||
"""
|
||||
Args:
|
||||
stem (nn.Module): a stem module
|
||||
stages (list[list[ResNetBlock]]): several (typically 4) stages,
|
||||
each contains multiple :class:`ResNetBlockBase`.
|
||||
num_classes (None or int): if None, will not perform classification.
|
||||
out_features (list[str]): name of the layers whose outputs should
|
||||
be returned in forward. Can be anything in "stem", "linear", or "res2" ...
|
||||
If None, will return the output of the last layer.
|
||||
"""
|
||||
super(ResNet, self).__init__()
|
||||
self.stem = stem
|
||||
self.num_classes = num_classes
|
||||
|
||||
current_stride = self.stem.stride
|
||||
self._out_feature_strides = {"stem": current_stride}
|
||||
self._out_feature_channels = {"stem": self.stem.out_channels}
|
||||
|
||||
self.stages = []
|
||||
self.names = []
|
||||
for i, blocks in enumerate(stages):
|
||||
for block in blocks:
|
||||
assert isinstance(block, ResNetBlockBase), block
|
||||
curr_channels = block.out_channels
|
||||
stage = nn.Sequential(*blocks)
|
||||
name = "res" + str(i + 2)
|
||||
self.stages.append(stage)
|
||||
self.names.append(name)
|
||||
self._out_feature_strides[name] = current_stride = int(
|
||||
current_stride * np.prod([k.stride for k in blocks])
|
||||
)
|
||||
self._out_feature_channels[name] = blocks[-1].out_channels
|
||||
self.stages = nn.ModuleList(self.stages)
|
||||
|
||||
if num_classes is not None:
|
||||
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
||||
self.linear = nn.Linear(curr_channels, num_classes)
|
||||
|
||||
# Sec 5.1 in "Accurate, Large Minibatch SGD: Training ImageNet in 1 Hour":
|
||||
# "The 1000-way fully-connected layer is initialized by
|
||||
# drawing weights from a zero-mean Gaussian with standard deviation of 0.01."
|
||||
nn.init.normal_(self.linear.weight, stddev=0.01)
|
||||
name = "linear"
|
||||
|
||||
if out_features is None:
|
||||
out_features = [name]
|
||||
self._out_features = out_features
|
||||
assert len(self._out_features)
|
||||
for out_feature in self._out_features:
|
||||
assert out_feature in self.names, "Available children: {}".format(", ".join(self.names))
|
||||
|
||||
def forward(self, x):
|
||||
outputs = {}
|
||||
x = self.stem(x)
|
||||
if "stem" in self._out_features:
|
||||
outputs["stem"] = x
|
||||
for ix, stage in enumerate(self.stages):
|
||||
name = self.names[ix]
|
||||
x = stage(x)
|
||||
if name in self._out_features:
|
||||
outputs[name] = x
|
||||
if self.num_classes is not None:
|
||||
x = self.avgpool(x)
|
||||
x = self.linear(x)
|
||||
if "linear" in self._out_features:
|
||||
outputs["linear"] = x
|
||||
return outputs
|
||||
|
||||
def build_resnet_backbone(in_channels=3):
|
||||
norm = "BN"
|
||||
stem = BasicStem(
|
||||
in_channels=in_channels,
|
||||
out_channels=64,
|
||||
norm=norm,
|
||||
)
|
||||
|
||||
# fmt: off
|
||||
out_features = ["res2", "res3", "res4"]
|
||||
depth = 101
|
||||
num_groups = 1
|
||||
bottleneck_channels = 64
|
||||
in_channels = 64
|
||||
out_channels = 256
|
||||
stride_in_1x1 = True
|
||||
res5_dilation = 1
|
||||
# fmt: on
|
||||
assert res5_dilation in {1, 2}, "res5_dilation cannot be {}.".format(res5_dilation)
|
||||
|
||||
num_blocks_per_stage = {50: [3, 4, 6, 3], 101: [3, 4, 23, 3], 152: [3, 8, 36, 3]}[depth]
|
||||
|
||||
stages = []
|
||||
|
||||
# Avoid creating variables without gradients
|
||||
# It consumes extra memory and may cause allreduce to fail
|
||||
out_stage_idx = [{"res2": 2, "res3": 3, "res4": 4, "res5": 5}[f] for f in out_features]
|
||||
max_stage_idx = max(out_stage_idx)
|
||||
for idx, stage_idx in enumerate(range(2, max_stage_idx + 1)):
|
||||
dilation = res5_dilation if stage_idx == 5 else 1
|
||||
first_stride = 1 if idx == 0 or (stage_idx == 5 and dilation == 2) else 2
|
||||
stage_kargs = dict()
|
||||
stage_kargs.update({
|
||||
"num_blocks": num_blocks_per_stage[idx],
|
||||
"first_stride": first_stride,
|
||||
"in_channels": in_channels,
|
||||
"bottleneck_channels": bottleneck_channels,
|
||||
"out_channels": out_channels,
|
||||
"num_groups": num_groups,
|
||||
"norm": norm,
|
||||
"stride_in_1x1": stride_in_1x1,
|
||||
"dilation": dilation,
|
||||
})
|
||||
stage_kargs["block_class"] = BottleneckBlock
|
||||
blocks = make_stage(**stage_kargs)
|
||||
in_channels = out_channels
|
||||
out_channels *= 2
|
||||
bottleneck_channels *= 2
|
||||
stages.append(blocks)
|
||||
return ResNet(stem, stages, out_features=out_features)
|
||||
@@ -0,0 +1,244 @@
|
||||
import math
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.model_zoo as model_zoo
|
||||
from bodyseg.backbone.backbone import Backbone
|
||||
|
||||
def fixed_padding(inputs, kernel_size, dilation):
|
||||
kernel_size_effective = kernel_size + (kernel_size - 1) * (dilation - 1)
|
||||
pad_total = kernel_size_effective - 1
|
||||
pad_beg = pad_total // 2
|
||||
pad_end = pad_total - pad_beg
|
||||
padded_inputs = F.pad(inputs, (pad_beg, pad_end, pad_beg, pad_end))
|
||||
return padded_inputs
|
||||
|
||||
class SeparableConv2d(nn.Module):
|
||||
def __init__(self, inplanes, planes, kernel_size=3, stride=1, dilation=1, bias=False, BatchNorm=None):
|
||||
super(SeparableConv2d, self).__init__()
|
||||
|
||||
self.conv1 = nn.Conv2d(inplanes, inplanes, kernel_size, stride, 0, dilation,
|
||||
groups=inplanes, bias=bias)
|
||||
self.bn = BatchNorm(inplanes)
|
||||
self.pointwise = nn.Conv2d(inplanes, planes, 1, 1, 0, 1, 1, bias=bias)
|
||||
|
||||
def forward(self, x):
|
||||
x = fixed_padding(x, self.conv1.kernel_size[0], dilation=self.conv1.dilation[0])
|
||||
x = self.conv1(x)
|
||||
x = self.bn(x)
|
||||
x = self.pointwise(x)
|
||||
return x
|
||||
|
||||
class Block(nn.Module):
|
||||
def __init__(self, inplanes, planes, reps, stride=1, dilation=1, BatchNorm=None,
|
||||
start_with_relu=True, grow_first=True, is_last=False):
|
||||
super(Block, self).__init__()
|
||||
|
||||
if planes != inplanes or stride != 1:
|
||||
self.skip = nn.Conv2d(inplanes, planes, 1, stride=stride, bias=False)
|
||||
self.skipbn = BatchNorm(planes)
|
||||
else:
|
||||
self.skip = None
|
||||
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
rep = []
|
||||
|
||||
filters = inplanes
|
||||
if grow_first:
|
||||
rep.append(self.relu)
|
||||
rep.append(SeparableConv2d(inplanes, planes, 3, 1, dilation, BatchNorm=BatchNorm))
|
||||
rep.append(BatchNorm(planes))
|
||||
filters = planes
|
||||
|
||||
for i in range(reps - 1):
|
||||
rep.append(self.relu)
|
||||
rep.append(SeparableConv2d(filters, filters, 3, 1, dilation, BatchNorm=BatchNorm))
|
||||
rep.append(BatchNorm(filters))
|
||||
|
||||
if not grow_first:
|
||||
rep.append(self.relu)
|
||||
rep.append(SeparableConv2d(inplanes, planes, 3, 1, dilation, BatchNorm=BatchNorm))
|
||||
rep.append(BatchNorm(planes))
|
||||
|
||||
if stride != 1:
|
||||
rep.append(self.relu)
|
||||
rep.append(SeparableConv2d(planes, planes, 3, 2, BatchNorm=BatchNorm))
|
||||
rep.append(BatchNorm(planes))
|
||||
|
||||
if stride == 1 and is_last:
|
||||
rep.append(self.relu)
|
||||
rep.append(SeparableConv2d(planes, planes, 3, 1, BatchNorm=BatchNorm))
|
||||
rep.append(BatchNorm(planes))
|
||||
|
||||
if not start_with_relu:
|
||||
rep = rep[1:]
|
||||
|
||||
self.rep = nn.Sequential(*rep)
|
||||
|
||||
def forward(self, inp):
|
||||
x = self.rep(inp)
|
||||
|
||||
if self.skip is not None:
|
||||
skip = self.skip(inp)
|
||||
skip = self.skipbn(skip)
|
||||
else:
|
||||
skip = inp
|
||||
|
||||
x = x + skip
|
||||
|
||||
return x
|
||||
|
||||
class AlignedXception(Backbone):
|
||||
"""
|
||||
Modified Alighed Xception
|
||||
"""
|
||||
def __init__(self, output_stride, BatchNorm):
|
||||
super(AlignedXception, self).__init__()
|
||||
|
||||
if output_stride == 16:
|
||||
entry_block3_stride = 2
|
||||
middle_block_dilation = 1
|
||||
exit_block_dilations = (1, 2)
|
||||
elif output_stride == 8:
|
||||
entry_block3_stride = 1
|
||||
middle_block_dilation = 2
|
||||
exit_block_dilations = (2, 4)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
# Entry flow
|
||||
self.conv1 = nn.Conv2d(3, 32, 3, stride=2, padding=1, bias=False)
|
||||
self.bn1 = BatchNorm(32)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
self.conv2 = nn.Conv2d(32, 64, 3, stride=1, padding=1, bias=False)
|
||||
self.bn2 = BatchNorm(64)
|
||||
|
||||
self.block1 = Block(64, 128, reps=2, stride=2, BatchNorm=BatchNorm, start_with_relu=False)
|
||||
self.block2 = Block(128, 256, reps=2, stride=2, BatchNorm=BatchNorm, start_with_relu=False,
|
||||
grow_first=True)
|
||||
self.block3 = Block(256, 728, reps=2, stride=entry_block3_stride, BatchNorm=BatchNorm,
|
||||
start_with_relu=True, grow_first=True, is_last=True)
|
||||
|
||||
# Middle flow
|
||||
self.block4 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block5 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block6 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block7 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block8 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block9 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block10 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block11 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block12 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block13 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block14 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block15 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block16 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block17 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block18 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
self.block19 = Block(728, 728, reps=3, stride=1, dilation=middle_block_dilation,
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=True)
|
||||
|
||||
# Exit flow
|
||||
self.block20 = Block(728, 1024, reps=2, stride=1, dilation=exit_block_dilations[0],
|
||||
BatchNorm=BatchNorm, start_with_relu=True, grow_first=False, is_last=True)
|
||||
|
||||
self.conv3 = SeparableConv2d(1024, 1536, 3, stride=1, dilation=exit_block_dilations[1], BatchNorm=BatchNorm)
|
||||
self.bn3 = BatchNorm(1536)
|
||||
|
||||
self.conv4 = SeparableConv2d(1536, 1536, 3, stride=1, dilation=exit_block_dilations[1], BatchNorm=BatchNorm)
|
||||
self.bn4 = BatchNorm(1536)
|
||||
|
||||
self.conv5 = SeparableConv2d(1536, 2048, 3, stride=1, dilation=exit_block_dilations[1], BatchNorm=BatchNorm)
|
||||
self.bn5 = BatchNorm(2048)
|
||||
|
||||
# Init weights
|
||||
self._init_weight()
|
||||
|
||||
def forward(self, x):
|
||||
# Entry flow
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
|
||||
x = self.conv2(x)
|
||||
x = self.bn2(x)
|
||||
x = self.relu(x)
|
||||
|
||||
x = self.block1(x)
|
||||
# add relu here
|
||||
x = self.relu(x)
|
||||
low_level_feat = x
|
||||
x = self.block2(x)
|
||||
x = self.block3(x)
|
||||
|
||||
# Middle flow
|
||||
x = self.block4(x)
|
||||
x = self.block5(x)
|
||||
x = self.block6(x)
|
||||
x = self.block7(x)
|
||||
x = self.block8(x)
|
||||
x = self.block9(x)
|
||||
x = self.block10(x)
|
||||
x = self.block11(x)
|
||||
x = self.block12(x)
|
||||
x = self.block13(x)
|
||||
x = self.block14(x)
|
||||
x = self.block15(x)
|
||||
x = self.block16(x)
|
||||
x = self.block17(x)
|
||||
x = self.block18(x)
|
||||
x = self.block19(x)
|
||||
|
||||
# Exit flow
|
||||
x = self.block20(x)
|
||||
x = self.relu(x)
|
||||
x = self.conv3(x)
|
||||
x = self.bn3(x)
|
||||
x = self.relu(x)
|
||||
|
||||
x = self.conv4(x)
|
||||
x = self.bn4(x)
|
||||
x = self.relu(x)
|
||||
|
||||
x = self.conv5(x)
|
||||
x = self.bn5(x)
|
||||
x = self.relu(x)
|
||||
|
||||
return x, low_level_feat
|
||||
|
||||
def _init_weight(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
m.weight.data.normal_(0, math.sqrt(2. / n))
|
||||
elif isinstance(m, nn.SyncBatchNorm):
|
||||
m.weight.data.fill_(1)
|
||||
m.bias.data.zero_()
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
m.weight.data.fill_(1)
|
||||
m.bias.data.zero_()
|
||||
|
||||
if __name__ == "__main__":
|
||||
import torch
|
||||
model = AlignedXception(BatchNorm=nn.BatchNorm2d, output_stride=16)
|
||||
input = torch.rand(1, 3, 512, 512)
|
||||
output, low_level_feat = model(input)
|
||||
print(output.size())
|
||||
print(low_level_feat.size())
|
||||
@@ -0,0 +1,330 @@
|
||||
import sys
|
||||
# sys.path.append("/Users/momo/human_seg_train")
|
||||
# print(sys.path)
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from bodyseg.backbone.backbone import get_norm
|
||||
|
||||
class ConvBNReLU(nn.Sequential):
|
||||
def __init__(self, in_planes, out_planes, kernel_size=3, stride=1, groups=1, norm_layer=None):
|
||||
padding = (kernel_size - 1) // 2
|
||||
if norm_layer is None:
|
||||
norm_layer = nn.BatchNorm2d
|
||||
super(ConvBNReLU, self).__init__(
|
||||
nn.Conv2d(in_planes, out_planes, kernel_size, stride, padding, groups=groups, bias=False),
|
||||
norm_layer(out_planes),
|
||||
nn.ReLU6(inplace=True)
|
||||
)
|
||||
|
||||
class InvertedResidual(nn.Module):
|
||||
def __init__(self, inp, oup, stride, expand_ratio, norm_layer=None):
|
||||
super(InvertedResidual, self).__init__()
|
||||
self.stride = stride
|
||||
assert stride in [1, 2]
|
||||
|
||||
if norm_layer is None:
|
||||
norm_layer = nn.BatchNorm2d
|
||||
|
||||
hidden_dim = int(round(inp * expand_ratio))
|
||||
self.use_res_connect = self.stride == 1 and inp == oup
|
||||
|
||||
layers = []
|
||||
if expand_ratio != 1:
|
||||
# pw
|
||||
layers.append(ConvBNReLU(inp, hidden_dim, kernel_size=1, norm_layer=norm_layer))
|
||||
layers.extend([
|
||||
# dw
|
||||
ConvBNReLU(hidden_dim, hidden_dim, stride=stride, groups=hidden_dim, norm_layer=norm_layer),
|
||||
# pw-linear
|
||||
nn.Conv2d(hidden_dim, oup, 1, 1, 0, bias=False),
|
||||
norm_layer(oup),
|
||||
])
|
||||
self.conv = nn.Sequential(*layers)
|
||||
|
||||
def forward(self, x):
|
||||
if self.use_res_connect:
|
||||
return x + self.conv(x)
|
||||
else:
|
||||
return self.conv(x)
|
||||
|
||||
class UpSampleBlock(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, expand_ratio=6):
|
||||
super(UpSampleBlock, self).__init__()
|
||||
self.refine = InvertedResidual(in_channels, out_channels, 1, expand_ratio)
|
||||
|
||||
def forward(self, x0, x1):
|
||||
x = torch.cat([x0, x1], dim=1)
|
||||
x = self.refine(x)
|
||||
return x
|
||||
|
||||
class BodySegNet_32_thin_sigmod(nn.Module):
|
||||
def __init__(self, input_channels=3, class_nums=1, output_onnx=False):
|
||||
super(BodySegNet_32_thin_sigmod, self).__init__()
|
||||
self.class_nums = class_nums
|
||||
self.output_onnx = output_onnx
|
||||
|
||||
self.stage_1 = nn.Sequential(
|
||||
ConvBNReLU(input_channels, 16, kernel_size=3, stride=2)
|
||||
)
|
||||
self.stage_2 = nn.Sequential(
|
||||
ConvBNReLU(16, 16, kernel_size=3, stride=2, groups=16),
|
||||
ConvBNReLU(16, 16, kernel_size=1, stride=1),
|
||||
)
|
||||
self.stage_3 = nn.Sequential(
|
||||
InvertedResidual(16, 24, stride=2, expand_ratio=6),
|
||||
InvertedResidual(24, 24, stride=1, expand_ratio=6),
|
||||
InvertedResidual(24, 24, stride=1, expand_ratio=6),
|
||||
)
|
||||
self.stage_4 = nn.Sequential(
|
||||
InvertedResidual(24, 32, stride=2, expand_ratio=6),
|
||||
InvertedResidual(32, 32, stride=1, expand_ratio=6),
|
||||
InvertedResidual(32, 32, stride=1, expand_ratio=6),
|
||||
InvertedResidual(32, 32, stride=1, expand_ratio=6),
|
||||
)
|
||||
self.stage_5 = nn.Sequential(
|
||||
InvertedResidual(32, 48, stride=2, expand_ratio=6),
|
||||
InvertedResidual(48, 48, stride=1, expand_ratio=6),
|
||||
InvertedResidual(48, 48, stride=1, expand_ratio=6),
|
||||
InvertedResidual(48, 48, stride=1, expand_ratio=6)
|
||||
)
|
||||
self.up_to_4 = UpSampleBlock(48 + 32, 16)
|
||||
self.up_to_3 = UpSampleBlock(16 + 24, 16)
|
||||
self.up_to_2 = UpSampleBlock(16 + 16, 16)
|
||||
self.up_to_1 = UpSampleBlock(16 + 16, 16)
|
||||
self.last_layer = nn.Sequential(
|
||||
ConvBNReLU(16, 16, kernel_size=1, stride=1),
|
||||
nn.Conv2d(16, self.class_nums, kernel_size=1, stride=1)
|
||||
)
|
||||
self._initialize_weights()
|
||||
|
||||
def forward(self, x):
|
||||
feature_S = []
|
||||
x1 = self.stage_1(x)
|
||||
x2 = self.stage_2(x1)
|
||||
x3 = self.stage_3(x2)
|
||||
x4 = self.stage_4(x3)
|
||||
feature = self.stage_5(x4)
|
||||
|
||||
feature = F.interpolate(feature, size=x4.size()[2:], mode='bilinear', align_corners=True)
|
||||
feature = self.up_to_4(x4, feature)
|
||||
feature = F.interpolate(feature, size=x3.size()[2:], mode='bilinear', align_corners=True)
|
||||
feature = self.up_to_3(x3, feature)
|
||||
feature = F.interpolate(feature, size=x2.size()[2:], mode='bilinear', align_corners=True)
|
||||
feature = self.up_to_2(x2, feature)
|
||||
feature = F.interpolate(feature, size=x1.size()[2:], mode='bilinear', align_corners=True)
|
||||
feature = self.up_to_1(x1, feature)
|
||||
feature_S.append(feature)
|
||||
feature = self.last_layer(feature)
|
||||
feature_S.append(feature)
|
||||
output = F.interpolate(feature, size=x.size()[2:], mode='bilinear', align_corners=True)
|
||||
# output = torch.sigmoid(output)
|
||||
if self.output_onnx:
|
||||
output = torch.argmax(output, dim=1).to(torch.float32)
|
||||
|
||||
return output, feature_S
|
||||
|
||||
def _initialize_weights(self):
|
||||
for name, m in self.named_modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
if 'first' in name:
|
||||
nn.init.normal_(m.weight, 0, 0.01)
|
||||
else:
|
||||
nn.init.normal_(m.weight, 0, 1.0 / m.weight.shape[1])
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
nn.init.constant_(m.weight, 1)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0.0001)
|
||||
nn.init.constant_(m.running_mean, 0)
|
||||
elif isinstance(m, nn.BatchNorm1d):
|
||||
nn.init.constant_(m.weight, 1)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0.0001)
|
||||
nn.init.constant_(m.running_mean, 0)
|
||||
elif isinstance(m, nn.Linear):
|
||||
nn.init.normal_(m.weight, 0, 0.01)
|
||||
if m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
|
||||
class _ASPPModule(nn.Module):
|
||||
def __init__(self, inplanes, planes, kernel_size, padding, dilation, BatchNorm):
|
||||
super(_ASPPModule, self).__init__()
|
||||
self.atrous_conv = nn.Conv2d(inplanes, planes, kernel_size=kernel_size,
|
||||
stride=1, padding=padding, dilation=dilation, bias=False)
|
||||
self.bn = BatchNorm(planes)
|
||||
self.relu = nn.ReLU()
|
||||
|
||||
self._init_weight()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.atrous_conv(x)
|
||||
x = self.bn(x)
|
||||
|
||||
return self.relu(x)
|
||||
|
||||
def _init_weight(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
torch.nn.init.kaiming_normal_(m.weight)
|
||||
|
||||
class ASPP(nn.Module):
|
||||
def __init__(self, backbone, output_stride, BatchNorm):
|
||||
super(ASPP, self).__init__()
|
||||
if backbone == 'drn':
|
||||
inplanes = 512
|
||||
elif backbone == 'mobilenet':
|
||||
inplanes = 320
|
||||
elif backbone == 'resnet_fpn':
|
||||
inplanes = 256
|
||||
else:
|
||||
inplanes = 2048
|
||||
if output_stride == 16:
|
||||
dilations = [1, 6, 12, 18]
|
||||
elif output_stride == 8:
|
||||
dilations = [1, 12, 24, 36]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
self.aspp1 = _ASPPModule(inplanes, 256, 1, padding=0, dilation=dilations[0], BatchNorm=BatchNorm)
|
||||
self.aspp2 = _ASPPModule(inplanes, 256, 3, padding=dilations[1], dilation=dilations[1], BatchNorm=BatchNorm)
|
||||
self.aspp3 = _ASPPModule(inplanes, 256, 3, padding=dilations[2], dilation=dilations[2], BatchNorm=BatchNorm)
|
||||
self.aspp4 = _ASPPModule(inplanes, 256, 3, padding=dilations[3], dilation=dilations[3], BatchNorm=BatchNorm)
|
||||
|
||||
self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
|
||||
nn.Conv2d(inplanes, 256, 1, stride=1, bias=False),
|
||||
BatchNorm(256),
|
||||
nn.ReLU())
|
||||
self.conv1 = nn.Conv2d(1280, 256, 1, bias=False)
|
||||
self.bn1 = BatchNorm(256)
|
||||
self.relu = nn.ReLU()
|
||||
self.dropout = nn.Dropout(0.5)
|
||||
self._init_weight()
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.aspp1(x)
|
||||
x2 = self.aspp2(x)
|
||||
x3 = self.aspp3(x)
|
||||
x4 = self.aspp4(x)
|
||||
x5 = self.global_avg_pool(x)
|
||||
x5 = F.interpolate(x5, size=x4.size()[2:], mode='bilinear', align_corners=True)
|
||||
# x5 = F.interpolate(x5, size=x4.size()[2:], mode='nearest')
|
||||
x = torch.cat((x1, x2, x3, x4, x5), dim=1)
|
||||
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
|
||||
return self.dropout(x)
|
||||
|
||||
def _init_weight(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
torch.nn.init.kaiming_normal_(m.weight)
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self,num_classes, backbone, BatchNorm):
|
||||
super(Decoder, self).__init__()
|
||||
if backbone == 'resnet_fpn' or backbone == 'drn':
|
||||
low_level_inplanes = 256
|
||||
elif backbone == 'xception':
|
||||
low_level_inplanes = 128
|
||||
elif backbone == 'mobilenet':
|
||||
low_level_inplanes = 24
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
self.conv1 = nn.Conv2d(low_level_inplanes, 16, 1, bias=False)
|
||||
self.bn1 = BatchNorm(16)
|
||||
self.relu = nn.ReLU()
|
||||
self.last_conv = nn.Sequential(nn.Conv2d(304, 256, kernel_size=3, stride=1, padding=1, bias=False),
|
||||
BatchNorm(256),
|
||||
nn.ReLU(),
|
||||
nn.Dropout(0.5),
|
||||
nn.Conv2d(256, 256, kernel_size=3, stride=1, padding=1, bias=False),
|
||||
BatchNorm(256),
|
||||
nn.ReLU(),
|
||||
nn.Dropout(0.1),
|
||||
nn.Conv2d(256, num_classes, kernel_size=1, stride=1))
|
||||
|
||||
self.up_to_1 = UpSampleBlock(16 + 16, 16)
|
||||
self.last_layer = nn.Sequential(
|
||||
ConvBNReLU(16, 16, kernel_size=1, stride=1),
|
||||
nn.Conv2d(16, 1, kernel_size=1, stride=1)
|
||||
)
|
||||
self._init_weight()
|
||||
|
||||
# x(1,256,8,6) low(1,256,32,24) -》 x(1,1,32,24)
|
||||
def forward(self, x, low_level_feat):
|
||||
feature_T = []
|
||||
# deeplab part
|
||||
#(1,256,32,24) -> (1,16,64,48)
|
||||
low_level_feat = self.conv1(low_level_feat)
|
||||
low_level_feat = self.bn1(low_level_feat)
|
||||
low_level_feat = self.relu(low_level_feat)
|
||||
|
||||
# x(1, 256, 8, 6)-> (1,16,32,24)
|
||||
x = F.interpolate(x, size=low_level_feat.size()[2:], mode='bilinear', align_corners=True)
|
||||
low_level_feat = F.interpolate(low_level_feat,
|
||||
size=[low_level_feat.size()[2] * 2, low_level_feat.size()[3] * 2],
|
||||
mode='bilinear', align_corners=True)
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
|
||||
# bodyseg part
|
||||
feature = F.interpolate(x, size=[x.size()[2] * 2,x.size()[3] * 2], mode='bilinear', align_corners=True)
|
||||
# 输入up_to_1 x(low)(1,16,64,48),上采样2倍后的feature(1,16,64,48)
|
||||
feature = self.up_to_1(low_level_feat, feature)
|
||||
feature_T.append(feature)
|
||||
# (1,1,64,48)
|
||||
feature = self.last_layer(feature)
|
||||
feature_T.append(feature)
|
||||
# (1,1,128,96)
|
||||
output = F.interpolate(feature, size=[128,96], mode='bilinear', align_corners=True)
|
||||
# output = torch.sigmoid(output)
|
||||
|
||||
return output, feature_T
|
||||
|
||||
def _init_weight(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
torch.nn.init.kaiming_normal_(m.weight)
|
||||
|
||||
def build_backbone(backbone, output_stride, BatchNorm, input_channel=3):
|
||||
from bodyseg.backbone.fpn import build_resnet_fpn_backbone
|
||||
from bodyseg.backbone.xception import AlignedXception
|
||||
|
||||
return build_resnet_fpn_backbone(input_channel)
|
||||
|
||||
def build_aspp(backbone, output_stride, BatchNorm):
|
||||
return ASPP(backbone, output_stride, BatchNorm)
|
||||
|
||||
def build_decoder(num_classes, backbone, BatchNorm):
|
||||
return Decoder(num_classes, backbone, BatchNorm)
|
||||
|
||||
class DeepLab(nn.Module):
|
||||
def __init__(self, input_channel=3, class_num=1):
|
||||
super(DeepLab, self).__init__()
|
||||
|
||||
BatchNorm = get_norm("BN")
|
||||
|
||||
self.backbone = build_backbone("resnet_fpn", 16, BatchNorm, input_channel=input_channel)
|
||||
self.aspp = build_aspp("resnet_fpn", 16, BatchNorm)
|
||||
self.decoder = build_decoder(class_num, "resnet_fpn", BatchNorm)
|
||||
|
||||
def forward(self, input):
|
||||
# input(1,3,128,96)
|
||||
#output: "p2"(1,256,32,24), "p3"(1,256,16,12), "p4"(1,256,8,6)
|
||||
output = self.backbone(input)
|
||||
# "p4"(1,256,8,6) "p2"(1,256,32,24)
|
||||
x, low_level_feat = output['p4'], output['p2']
|
||||
# x(1,256,8,6)-》(1,256,8,6)
|
||||
x = self.aspp(x)
|
||||
# x(1,256,8,6) low(1,256,32,24) -》 x(1,1,32,24)
|
||||
x, feature_T = self.decoder(x, low_level_feat)
|
||||
# x(1,1,128,96)
|
||||
x = F.interpolate(x, size=input.size()[2:], mode='bilinear', align_corners=True)
|
||||
# x = F.interpolate(x, size=input.size()[2:], mode='nearest')
|
||||
return x, feature_T
|
||||
Reference in New Issue
Block a user