| 154 | |
| 155 | |
| 156 | class MVSNet(nn.Module): |
| 157 | def __init__(self, ndepths, depth_interval_ratio, cr_base_chs=None, fea_mode="fpn", agg_mode="variance", depth_mode="regression",winner_take_all_to_generate_depth=True,inverse_depth=False): |
| 158 | super(MVSNet, self).__init__() |
| 159 | |
| 160 | if cr_base_chs is None: |
| 161 | cr_base_chs = [8] * len(ndepths) |
| 162 | self.ndepths = ndepths |
| 163 | self.depth_interval_ratio = depth_interval_ratio |
| 164 | self.fea_mode = fea_mode |
| 165 | self.cr_base_chs = cr_base_chs |
| 166 | self.num_stage = len(ndepths) |
| 167 | self.inverse_depth=inverse_depth |
| 168 | |
| 169 | print("netphs:", ndepths) |
| 170 | print("depth_intervals_ratio:", depth_interval_ratio) |
| 171 | print("cr_base_chs:", cr_base_chs) |
| 172 | print("fea_mode:", fea_mode) |
| 173 | print("agg_mode:", agg_mode) |
| 174 | print("depth_mode:", depth_mode) |
| 175 | |
| 176 | assert len(ndepths) == len(depth_interval_ratio) |
| 177 | |
| 178 | self.feature = FeatureNet(base_channels=8, stride=4, num_stage=self.num_stage, mode=self.fea_mode) |
| 179 | self.cost_aggregation = CostAgg(agg_mode, self.feature.out_channels) |
| 180 | |
| 181 | self.cost_regularization = nn.ModuleList( |
| 182 | [CostRegNet(in_channels=2, base_channels=self.cr_base_chs[i],stage=i) for i in range(self.num_stage)]) |
| 183 | self.cost_regularization_refine = nn.ModuleList( |
| 184 | [CostRegNet_refine(in_channels=2, base_channels=self.cr_base_chs[i],stage=i) for i in range(self.num_stage)]) |
| 185 | |
| 186 | self.DepthNet = DepthNet(depth_mode) |
| 187 | |
| 188 | def forward(self, imgs, proj_matrices, depth_values): |
| 189 | """ |
| 190 | :param is_flip: augment only for 3D-UNet |
| 191 | :param imgs: (b, nview, c, h, w) |
| 192 | :param proj_matrices: |
| 193 | :param depth_values: |
| 194 | :return: |
| 195 | """ |
| 196 | depth_interval = (depth_values[0, -1] - depth_values[0, 0]) / depth_values.size(1) |
| 197 | |
| 198 | # step 1. feature extraction |
| 199 | features = [] |
| 200 | for nview_idx in range(imgs.size(1)): # imgs shape (B, N, C, H, W) |
| 201 | img = imgs[:, nview_idx] |
| 202 | features.append(self.feature(img)) |
| 203 | |
| 204 | ori_shape = imgs[:, 0].shape[2:] # (H, W) |
| 205 | |
| 206 | outputs = {} |
| 207 | last_depth = None |
| 208 | for stage_idx in range(self.num_stage): |
| 209 | # print("*********************stage{}*********************".format(stage_idx + 1)) |
| 210 | # stage feature, proj_mats, scales |
| 211 | features_stage = [feat["stage{}".format(stage_idx + 1)] for feat in features] |
| 212 | proj_matrices_stage = proj_matrices["stage{}".format(stage_idx + 1)] |
| 213 | # stage1: 1/4, stage2: 1/2, stage3: 1 |