7e4ef6f98b
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
101 lines
5.5 KiB
Python
101 lines
5.5 KiB
Python
"""Export HandFlow with its own MANO convention and camera-to-world transform."""
|
|
import argparse
|
|
import json
|
|
from pathlib import Path
|
|
import sys
|
|
|
|
import numpy as np
|
|
import torch
|
|
from scipy.spatial.transform import Rotation, Slerp
|
|
|
|
ROOT = Path(__file__).resolve().parents[1]
|
|
sys.path.insert(0, str(ROOT))
|
|
from utils.mano_utils import MANOForwardKinematics
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument('--input', type=Path, required=True)
|
|
parser.add_argument('--output', type=Path, required=True)
|
|
parser.add_argument('--gap-policy', choices=['interpolate', 'raw'], default='interpolate')
|
|
parser.add_argument('--min-detection-run', type=int, default=3,
|
|
help='Reject isolated detection runs shorter than this before gap filling')
|
|
args = parser.parse_args()
|
|
torch.set_num_threads(4)
|
|
source = np.load(args.input)
|
|
assert str(source['side']) == 'right', 'This exporter targets the L20 right hand'
|
|
pose = torch.tensor(source['pose'], dtype=torch.float32)
|
|
count = len(pose)
|
|
betas = torch.tensor(source['betas'], dtype=torch.float32)
|
|
if betas.ndim == 1:
|
|
betas = betas[None].expand(count, -1)
|
|
elif len(betas) == 1:
|
|
betas = betas.expand(count, -1)
|
|
assert betas.shape == (count, 10)
|
|
trans = torch.tensor(source['trans'], dtype=torch.float32)
|
|
c2w = source['c2w']
|
|
assert len(c2w) == count and len(source['verts_cam']) == count
|
|
fk = MANOForwardKinematics(str(ROOT / 'third_party/hamer/_DATA/data/mano'), torch.device('cpu'))
|
|
local_pose = pose.clone()
|
|
local_pose[:, :3] = 0
|
|
local, camera, verts = [], [], []
|
|
for start in range(0, count, 64):
|
|
end = min(count, start + 64)
|
|
sides = ['right'] * (end - start)
|
|
local.append(fk.joints(local_pose[start:end], betas[start:end], torch.zeros_like(trans[start:end]), sides).numpy() / 1000)
|
|
camera.append(fk.joints(pose[start:end], betas[start:end], trans[start:end], sides).numpy() / 1000)
|
|
verts.append(fk.verts(pose[start:end], betas[start:end], trans[start:end], sides).numpy() / 1000)
|
|
local = np.concatenate(local)
|
|
camera = np.concatenate(camera)
|
|
verts = np.concatenate(verts)
|
|
local -= local[:, :1].copy()
|
|
world = np.einsum('tij,tkj->tki', c2w[:, :3, :3], camera) + c2w[:, None, :3, 3]
|
|
root_R = c2w[:, :3, :3] @ Rotation.from_rotvec(pose[:, :3].numpy()).as_matrix()
|
|
reconstructed = np.einsum('tij,tkj->tki', root_R, local) + world[:, :1]
|
|
vertex_error = float(np.abs(verts - source['verts_cam']).max())
|
|
rigid_error = float(np.abs(reconstructed - world).max())
|
|
world_vertex_error = float(np.abs(np.einsum('tij,tkj->tki', c2w[:, :3, :3], verts) + c2w[:, None, :3, 3] - source['verts_world']).max())
|
|
assert np.isfinite(world).all() and np.isfinite(local).all()
|
|
assert max(vertex_error, rigid_error, world_vertex_error) < 1e-5
|
|
valid = np.asarray(source['pred_valid'], dtype=bool)
|
|
observed_valid = valid.copy()
|
|
valid_ids = np.flatnonzero(valid)
|
|
for group in np.split(valid_ids, np.flatnonzero(np.diff(valid_ids) > 1) + 1):
|
|
if len(group) < args.min_detection_run:
|
|
valid[group] = False
|
|
local_raw, wrist_raw, root_raw = local.copy(), world[:, 0].copy(), root_R.copy()
|
|
wrist = wrist_raw.copy()
|
|
if a_policy := (args.gap_policy == 'interpolate' and not valid.all()):
|
|
observed = np.flatnonzero(valid)
|
|
assert len(observed) >= 2, 'Too few detected frames for motion export'
|
|
times = np.arange(count)
|
|
local = np.stack([np.interp(times, observed, local[:, i, j][observed])
|
|
for i in range(21) for j in range(3)], axis=1).reshape(count,21,3)
|
|
wrist = np.stack([np.interp(times, observed, wrist[:, j][observed]) for j in range(3)],axis=1)
|
|
root_R = Slerp(observed, Rotation.from_matrix(root_R[observed]))(np.clip(times,observed[0],observed[-1])).as_matrix()
|
|
args.output.parent.mkdir(parents=True, exist_ok=True)
|
|
np.savez_compressed(args.output, joints=local, wrist_world=wrist,
|
|
root_orient=Rotation.from_matrix(root_R).as_rotvec(), fps=source['fps'],
|
|
detection_valid=valid, detection_valid_raw=observed_valid, source=str(args.input.resolve()),
|
|
source_model='HandFlow/manopth flat_hand_mean=True', joints_world_raw=world,
|
|
joints_raw=local_raw, wrist_world_raw=wrist_raw, root_world_R_raw=root_raw,
|
|
gap_policy=args.gap_policy)
|
|
report = dict(frames=count, fps=float(source['fps']), finite=True,
|
|
detected_frames=int(source['pred_valid'].sum()),
|
|
usable_detection_frames=int(valid.sum()),
|
|
rejected_isolated_detection_frames=np.flatnonzero(observed_valid & ~valid).tolist(),
|
|
min_detection_run=args.min_detection_run,
|
|
camera_vertex_reproduction_max_error_m=vertex_error,
|
|
world_vertex_reproduction_max_error_m=world_vertex_error,
|
|
local_to_world_joint_max_error_m=rigid_error,
|
|
source=str(args.input.resolve()),
|
|
gap_policy=args.gap_policy, filled_frames=int((~valid).sum()) if a_policy else 0,
|
|
gap_note='Retargeting holds endpoint detections and interpolates internal gaps; raw inference preserved',
|
|
convention='HandFlow manopth, flat_hand_mean=True; mm converted to m; c2w applied once')
|
|
args.output.with_suffix('.validation.json').write_text(json.dumps(report, indent=2))
|
|
print(json.dumps(report, indent=2))
|
|
|
|
|
|
if __name__ == '__main__':
|
|
main()
|