MCPcopy Create free account
hub / github.com/AllentDan/LibtorchSegmentation / PSPNetImpl

Method PSPNetImpl

src/architectures/PSPNet.cpp:3–27  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1#include "PSPNet.h"
2
3PSPNetImpl::PSPNetImpl(int _num_classes, std::string encoder_name, std::string pretrained_path, int _encoder_depth,
4 int psp_out_channels, bool psp_use_batchnorm, float psp_dropout, double upsampling) {
5 num_classes = _num_classes;
6 encoder_depth = _encoder_depth;
7
8 auto encoder_param = encoder_params();
9 std::vector<int> encoder_channels = encoder_param[encoder_name]["out_channels"];
10 if (!encoder_param.contains(encoder_name))
11 std::cout<< "encoder name must in {resnet18, resnet34, resnet50, resnet101, resnet150, \
12 resnext50_32x4d, resnext101_32x8d, vgg11, vgg11_bn, vgg13, vgg13_bn, \
13 vgg16, vgg16_bn, vgg19, vgg19_bn,}";
14 if (encoder_param[encoder_name]["class_type"] == "resnet")
15 encoder = new ResNetImpl(encoder_param[encoder_name]["layers"], 1000, encoder_name);
16 else if (encoder_param[encoder_name]["class_type"] == "vgg")
17 encoder = new VGGImpl(encoder_param[encoder_name]["cfg"], 1000, encoder_param[encoder_name]["batch_norm"]);
18 else std::cout<< "unknown error in backbone initialization";
19
20 encoder->load_pretrained(pretrained_path);
21 decoder = PSPDecoder(encoder_channels, psp_out_channels, psp_dropout, psp_use_batchnorm);
22 segmentation_head = SegmentationHead(psp_out_channels, num_classes, 3, upsampling);
23
24 register_module("encoder", std::shared_ptr<Backbone>(encoder));
25 register_module("decoder", decoder);
26 register_module("segmentation_head", segmentation_head);
27}
28
29torch::Tensor PSPNetImpl::forward(torch::Tensor x) {
30 std::vector<torch::Tensor> features = encoder->features(x, encoder_depth);

Callers

nothing calls this directly

Calls 2

encoder_paramsFunction · 0.85
load_pretrainedMethod · 0.45

Tested by

no test coverage detected