Files
hand-motion-pipeline/scripts/export_handflow_dex.py
T

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()