| 73 | } |
| 74 | |
| 75 | FPNDecoderImpl::FPNDecoderImpl(std::vector<int> encoder_channels, int encoder_depth, int pyramid_channels, int segmentation_channels, |
| 76 | float dropout_, std::string merge_policy) |
| 77 | { |
| 78 | out_channels = merge_policy == "add"? segmentation_channels :segmentation_channels * 4; |
| 79 | if(encoder_depth<3) std::cout<< "Encoder depth for FPN decoder cannot be less than 3"; |
| 80 | std::reverse(std::begin(encoder_channels),std::end(encoder_channels)); |
| 81 | encoder_channels = std::vector<int> (encoder_channels.begin(),encoder_channels.begin()+encoder_depth+1); |
| 82 | p5 = torch::nn::Conv2d(conv_options(encoder_channels[0], pyramid_channels, 1)); |
| 83 | p4 = FPNBlock(pyramid_channels, encoder_channels[1]); |
| 84 | p3 = FPNBlock(pyramid_channels, encoder_channels[2]); |
| 85 | p2 = FPNBlock(pyramid_channels, encoder_channels[3]); |
| 86 | for(int i = 3; i>=0; i--){ |
| 87 | seg_blocks->push_back(SegmentationBlock(pyramid_channels, segmentation_channels, i)); |
| 88 | } |
| 89 | merge = MergeBlock(merge_policy); |
| 90 | dropout = torch::nn::Dropout2d(torch::nn::Dropout2dOptions().p(dropout_).inplace(true)); |
| 91 | |
| 92 | register_module("p5",p5); |
| 93 | register_module("p4",p4); |
| 94 | register_module("p3",p3); |
| 95 | register_module("p2",p2); |
| 96 | register_module("seg_blocks",seg_blocks); |
| 97 | register_module("merge",merge); |
| 98 | } |
| 99 | |
| 100 | torch::Tensor FPNDecoderImpl::forward(std::vector<torch::Tensor> features){ |
| 101 | int features_len = features.size(); |
nothing calls this directly
no test coverage detected