forward a batch of data
(self, data)
| 184 | self.scheduler.step() |
| 185 | |
| 186 | def _forward_batch(self, data): |
| 187 | """forward a batch of data""" |
| 188 | pts = data["pts"] |
| 189 | pred = self.net(self.input_grid, pts) |
| 190 | |
| 191 | pred_sdf = pred[..., :1] |
| 192 | gt_sdf = data["sdf"] |
| 193 | if self.sdf_loss_type == "l1": |
| 194 | sdf_loss = F.l1_loss(pred_sdf, gt_sdf) |
| 195 | elif self.sdf_loss_type == "weightedl1": |
| 196 | lamb = 0.5 |
| 197 | weight = (1 + lamb * torch.sign(gt_sdf) * torch.sign(gt_sdf - pred_sdf)) |
| 198 | sdf_loss = ((pred_sdf - gt_sdf).abs() * weight).mean() |
| 199 | elif self.sdf_loss_type == "weightedl1_clamp": |
| 200 | lamb = 0.5 |
| 201 | gt_sdf = torch.clamp(gt_sdf, -self.sdf_threshold, self.sdf_threshold) |
| 202 | pred_sdf = torch.clamp(pred_sdf, -self.sdf_threshold, self.sdf_threshold) |
| 203 | weight = (1 + lamb * torch.sign(gt_sdf) * torch.sign(gt_sdf - pred_sdf)) |
| 204 | sdf_loss = ((pred_sdf - gt_sdf).abs() * weight).mean() |
| 205 | else: |
| 206 | raise NotImplementedError |
| 207 | loss_dict = {"sdf_loss": sdf_loss} |
| 208 | |
| 209 | if self.data_type != "sdf": |
| 210 | pred_tex = pred[..., 1:] |
| 211 | gt_tex = data["tex"] |
| 212 | if self.sdf_renorm: |
| 213 | mask = gt_sdf.squeeze(1).abs() < 1.0 * self.tex_threshold_ratio |
| 214 | else: |
| 215 | mask = gt_sdf.squeeze(1).abs() < self.sdf_threshold * self.tex_threshold_ratio |
| 216 | if self.data_type == "sdftex": |
| 217 | if self.tex_loss_type == "l1": |
| 218 | tex_loss = F.l1_loss(pred_tex[mask], gt_tex[mask]) * self.tex_weight |
| 219 | elif self.tex_loss_type == "l2": |
| 220 | tex_loss = F.mse_loss(pred_tex[mask], gt_tex[mask]) * self.tex_weight |
| 221 | elif self.tex_loss_type == "huber": |
| 222 | tex_loss = F.huber_loss(pred_tex[mask], gt_tex[mask], delta=0.1) * self.tex_weight |
| 223 | else: |
| 224 | raise NotImplementedError |
| 225 | loss_dict["tex_loss"] = tex_loss |
| 226 | elif self.data_type == "sdfpbr": |
| 227 | if self.tex_loss_type == "l1": |
| 228 | rgb_loss = F.l1_loss(pred_tex[mask, :3], gt_tex[mask, :3]) * self.tex_weight |
| 229 | mr_loss = F.l1_loss(pred_tex[mask, 3:5], gt_tex[mask, 3:5]) * self.tex_weight |
| 230 | normal_loss = F.l1_loss(pred_tex[mask, 5:], gt_tex[mask, 5:]) * self.tex_weight |
| 231 | loss_dict["rgb_loss"] = rgb_loss |
| 232 | loss_dict["mr_loss"] = mr_loss |
| 233 | loss_dict["normal_loss"] = normal_loss |
| 234 | else: |
| 235 | raise NotImplementedError |
| 236 | |
| 237 | return pred, loss_dict |
| 238 | |
| 239 | def train(self, data_path): |
| 240 | # load data |