| 1 | #include "PSPNet.h" |
| 2 | |
| 3 | PSPNetImpl::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 | |
| 29 | torch::Tensor PSPNetImpl::forward(torch::Tensor x) { |
| 30 | std::vector<torch::Tensor> features = encoder->features(x, encoder_depth); |
nothing calls this directly
no test coverage detected