t2m-gpt-vsl-code / utils /losses_vsl.py
Tri1's picture
T2M-GPT VSL adaptation: Python sources only (82 files, no checkpoints or data)
8e5456b verified
Raw
History Blame Contribute Delete
2.75 kB
"""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)) # [1,1,dim]
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] # a velocity is valid only if both frames are
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) # [B,T,nk]
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