Spaces:
Running
Running
File size: 8,027 Bytes
17f1f54 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 | # coding: utf-8
"""
PINN regularizers for Sign-IDD (mask-aware, stable, fast)
---------------------------------------------------------
Inputs:
- skel: (B, T, J, 3) or (B, T, 3J)
- mask: (B, T) boolean/0-1, True for valid frames
Loss terms:
1) bone length consistency: enforce per-sequence rest-length consistency across time
2) smoothness: velocity + acceleration
3) FK consistency (meaningful): reconstruct child using parent + unit_dir * rest_len
Notes:
- FK in your reference was identity => always 0. This file fixes that.
- Rest length is detached by default to avoid trivial shrink-to-zero solutions.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Dict, Optional
import torch
import torch.nn as nn
@dataclass
class PINNConfig:
lambda_bone: float = 1.0
lambda_vel: float = 0.1
lambda_acc: float = 0.05
lambda_fk: float = 0.5
eps: float = 1e-8
dt: float = 1.0
# rest length estimation
rest_from: str = "first_valid" # "first_valid" or "mean_valid"
detach_rest: bool = True # important for stability
# robust penalty
use_huber: bool = True
huber_delta: float = 1.0
def _safe_norm(x: torch.Tensor, dim: int = -1, eps: float = 1e-8) -> torch.Tensor:
return torch.sqrt(torch.clamp((x * x).sum(dim=dim), min=eps))
def _masked_mean(x: torch.Tensor, mask: Optional[torch.Tensor], eps: float = 1e-8) -> torch.Tensor:
if mask is None:
return x.mean()
# broadcast mask to x
m = mask
while m.ndim < x.ndim:
m = m.unsqueeze(-1)
m = m.to(dtype=x.dtype)
num = (x * m).sum()
den = m.sum().clamp_min(eps)
return num / den
def _huber(x: torch.Tensor, delta: float) -> torch.Tensor:
ax = x.abs()
q = torch.minimum(ax, torch.tensor(delta, device=x.device, dtype=x.dtype))
l = ax - q
return 0.5 * q * q + delta * l
class PINNLoss(nn.Module):
def __init__(self, parents: Dict[int, int], num_joints: int = 50, cfg: PINNConfig = PINNConfig()):
super().__init__()
self.parents = parents
self.num_joints = num_joints
self.cfg = cfg
child, parent = [], []
for j in range(num_joints):
p = parents.get(j, -1)
if p is None or p == -1 or p == j:
continue
child.append(j)
parent.append(p)
if len(child) == 0:
raise ValueError("PINNLoss: no valid bones from parents dict.")
self.register_buffer("_child", torch.tensor(child, dtype=torch.long), persistent=False)
self.register_buffer("_parent", torch.tensor(parent, dtype=torch.long), persistent=False)
def _ensure_shape(self, skel: torch.Tensor) -> torch.Tensor:
if skel.ndim == 3:
B, T, D = skel.shape
exp = 3 * self.num_joints
if D != exp:
raise ValueError(f"PINNLoss: expected last dim {exp}, got {D}")
return skel.view(B, T, self.num_joints, 3)
if skel.ndim == 4:
if skel.shape[2] != self.num_joints or skel.shape[3] != 3:
raise ValueError(f"PINNLoss: expected (B,T,{self.num_joints},3), got {tuple(skel.shape)}")
return skel
raise ValueError(f"PINNLoss: unsupported shape {tuple(skel.shape)}")
def _bones(self, skel: torch.Tensor) -> torch.Tensor:
# (B,T,Nb,3)
return skel[:, :, self._child, :] - skel[:, :, self._parent, :]
def _rest_lengths(self, bone_len: torch.Tensor, mask: Optional[torch.Tensor]) -> torch.Tensor:
"""
bone_len: (B,T,Nb)
returns rest: (B,1,Nb)
"""
cfg = self.cfg
B, T, Nb = bone_len.shape
if mask is None:
if cfg.rest_from == "first_valid":
rest = bone_len[:, :1, :]
else:
rest = bone_len.mean(dim=1, keepdim=True)
else:
m = mask.to(dtype=bone_len.dtype) # (B,T)
if cfg.rest_from == "first_valid":
idx = (m > 0.5).float().argmax(dim=1) # (B,)
rest = bone_len[torch.arange(B, device=bone_len.device), idx, :].unsqueeze(1) # (B,1,Nb)
else:
mt = m.unsqueeze(-1) # (B,T,1)
rest = (bone_len * mt).sum(dim=1, keepdim=True) / mt.sum(dim=1, keepdim=True).clamp_min(cfg.eps)
if cfg.detach_rest:
rest = rest.detach()
return rest
def bone_length_loss(self, skel: torch.Tensor, mask: Optional[torch.Tensor]) -> torch.Tensor:
cfg = self.cfg
bones = self._bones(skel) # (B,T,Nb,3)
bl = _safe_norm(bones, dim=-1, eps=cfg.eps) # (B,T,Nb)
rest = self._rest_lengths(bl, mask) # (B,1,Nb)
diff = bl - rest # (B,T,Nb)
if cfg.use_huber:
per = _huber(diff, cfg.huber_delta)
else:
per = diff * diff
return _masked_mean(per, mask, eps=cfg.eps)
def velocity_loss(self, skel: torch.Tensor, mask: Optional[torch.Tensor]) -> torch.Tensor:
cfg = self.cfg
vel = (skel[:, 1:] - skel[:, :-1]) / cfg.dt # (B,T-1,J,3)
vm = _safe_norm(vel, dim=-1, eps=cfg.eps) # (B,T-1,J)
if cfg.use_huber:
vm = _huber(vm, cfg.huber_delta)
m = None
if mask is not None:
m = (mask[:, 1:] & mask[:, :-1]).to(dtype=torch.bool) # (B,T-1)
return _masked_mean(vm, m, eps=cfg.eps)
def acceleration_loss(self, skel: torch.Tensor, mask: Optional[torch.Tensor]) -> torch.Tensor:
cfg = self.cfg
vel = (skel[:, 1:] - skel[:, :-1]) / cfg.dt # (B,T-1,J,3)
acc = (vel[:, 1:] - vel[:, :-1]) / cfg.dt # (B,T-2,J,3)
am = _safe_norm(acc, dim=-1, eps=cfg.eps) # (B,T-2,J)
if cfg.use_huber:
am = _huber(am, cfg.huber_delta)
m = None
if mask is not None:
m = (mask[:, 2:] & mask[:, 1:-1] & mask[:, :-2]).to(dtype=torch.bool) # (B,T-2)
return _masked_mean(am, m, eps=cfg.eps)
def forward_kinematics_loss(self, skel: torch.Tensor, mask: Optional[torch.Tensor]) -> torch.Tensor:
"""
Reconstruct child: parent + unit_dir * rest_len
Penalize ||child - recon||^2
"""
cfg = self.cfg
bones = self._bones(skel) # (B,T,Nb,3)
bl = _safe_norm(bones, dim=-1, eps=cfg.eps) # (B,T,Nb)
rest = self._rest_lengths(bl, mask) # (B,1,Nb)
unit = bones / bl.unsqueeze(-1).clamp_min(cfg.eps) # (B,T,Nb,3)
parent_pos = skel[:, :, self._parent, :] # (B,T,Nb,3)
recon = parent_pos + unit * rest.unsqueeze(-1) # (B,T,Nb,3)
true_child = skel[:, :, self._child, :] # (B,T,Nb,3)
diff = true_child - recon # (B,T,Nb,3)
per = (diff * diff).sum(dim=-1) # (B,T,Nb)
if cfg.use_huber:
per = _huber(per, cfg.huber_delta)
return _masked_mean(per, mask, eps=cfg.eps)
def forward(self, skel: torch.Tensor, mask: Optional[torch.Tensor] = None):
skel = self._ensure_shape(skel)
cfg = self.cfg
L_bone = self.bone_length_loss(skel, mask)
L_vel = self.velocity_loss(skel, mask) if cfg.lambda_vel != 0 else skel.new_tensor(0.0)
L_acc = self.acceleration_loss(skel, mask) if cfg.lambda_acc != 0 else skel.new_tensor(0.0)
L_fk = self.forward_kinematics_loss(skel, mask) if cfg.lambda_fk != 0 else skel.new_tensor(0.0)
total = (
cfg.lambda_bone * L_bone
+ cfg.lambda_vel * L_vel
+ cfg.lambda_acc * L_acc
+ cfg.lambda_fk * L_fk
)
return {
"total": total,
"bone": L_bone,
"velocity": L_vel,
"acceleration": L_acc,
"fk": L_fk,
}
|