| """Reconstruction losses for DWPose sign-language motion. |
| |
| Differs from T2M-GPT's ReConsLoss in two ways that matter for this data: |
| * masked -- undetected keypoints (low DWPose score / NaN) carry no gradient; |
| * keypoint-weighted -- hands dominate meaning in sign language, so they are |
| weighted above body, and face is down-weighted. |
| The velocity term is computed on the same weighted/masked coordinates, which for |
| a raw-coordinate representation is just the first temporal difference. |
| """ |
| import torch |
| import torch.nn as nn |
|
|
| from dataset.dataset_vsl import kp_weights |
|
|
|
|
| def _base_loss(name): |
| if name == "l1": |
| return lambda a, b: (a - b).abs() |
| if name == "l2": |
| return lambda a, b: (a - b) ** 2 |
| if name == "l1_smooth": |
| return lambda a, b: nn.functional.smooth_l1_loss(a, b, reduction="none", beta=0.1) |
| raise ValueError(f"unknown recons loss {name}") |
|
|
|
|
| class VSLReConsLoss(nn.Module): |
| def __init__(self, recons_loss="l1_smooth", w_body=1.0, w_face=0.5, w_hand=3.0, |
| layout=None, w_finger=None, w_fingertip=None): |
| super().__init__() |
| self.fn = _base_loss(recons_loss) |
| w = torch.from_numpy(kp_weights(w_body, w_face, w_hand, layout=layout, |
| finger=w_finger, fingertip=w_fingertip)) |
| self.register_buffer("w", w.view(1, 1, -1)) |
|
|
| def _reduce(self, err, mask): |
| wm = mask * self.w |
| return (err * wm).sum() / wm.sum().clamp(min=1.0) |
|
|
| def forward(self, pred, gt, mask): |
| """pred/gt/mask: [B,T,256]. mask is 1 where the keypoint is usable.""" |
| return self._reduce(self.fn(pred, gt), mask) |
|
|
| def forward_vel(self, pred, gt, mask): |
| dp = pred[:, 1:] - pred[:, :-1] |
| dg = gt[:, 1:] - gt[:, :-1] |
| m = mask[:, 1:] * mask[:, :-1] |
| return self._reduce(self.fn(dp, dg), m) |
|
|
|
|
| @torch.no_grad() |
| def mpjpe_groups(pred_xy, gt_xy, valid, groups=None): |
| """Mean per-joint position error in raw frame-normalized units, per group. |
| |
| pred_xy / gt_xy: [B,T,2*n_kpts] un-normalized coordinates. valid: [B,T,n_kpts]. |
| `groups` comes from Layout.metric_groups(); defaults to the full 128 layout. |
| Returns a dict of scalars (all / body / face / hands). |
| """ |
| B, T, D = pred_xy.shape |
| nk = D // 2 |
| p = pred_xy.view(B, T, nk, 2) |
| g = gt_xy.view(B, T, nk, 2) |
| d = torch.linalg.norm(p - g, dim=-1) |
| out = {} |
| if groups is None: |
| groups = {"all": (0, 128), "body": (0, 18), "face": (18, 86), "hands": (86, 128)} |
| for name, (a, b) in groups.items(): |
| m = valid[..., a:b] |
| out[name] = float((d[..., a:b] * m).sum() / m.sum().clamp(min=1.0)) |
| return out |
|
|