MCPcopy Create free account
hub / github.com/Sin3DM/Sin3DM / _forward_batch

Method _forward_batch

src/encoding/model.py:186–237  ·  view source on GitHub ↗

forward a batch of data

(self, data)

Source from the content-addressed store, hash-verified

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

Callers 1

trainMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected