File size: 4,028 Bytes
49d36c0
 
 
 
 
ef90f84
49d36c0
 
 
 
 
 
 
ef90f84
 
 
 
 
 
 
49d36c0
 
 
 
 
 
 
 
 
 
0b8c5ca
 
 
 
 
 
 
 
 
 
49d36c0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import argparse
import json
import shutil
import sys
from pathlib import Path
from types import SimpleNamespace

import cv2
import torch
from huggingface_hub import snapshot_download

PROJECT_ROOT = Path(__file__).resolve().parent
SOMA_SOURCE = PROJECT_ROOT / "third_party" / "soma"
for source_dir in (PROJECT_ROOT, SOMA_SOURCE):
    source = str(source_dir)
    if source not in sys.path:
        sys.path.insert(0, source)


def ensure_soma_assets() -> Path:
    target = Path("inputs/soma_assets")
    required = target / "SOMA_neutral.npz"
    if not required.exists():
        snapshot_download(
            repo_id="nvidia/soma-x",
            repo_type="model",
            local_dir=target,
        )
    # Spaces can omit `.pt` files during repository uploads. Deployment stores
    # these weights as `.bin`, then restores the names SOMA-X expects.
    aliases = (
        (target / "correctives_model.bin", target / "correctives_model.pt"),
        (target / "MHR" / "mhr_model_lod6.bin", target / "MHR" / "mhr_model_lod6.pt"),
    )
    for source, destination in aliases:
        if source.exists() and not destination.exists():
            destination.parent.mkdir(parents=True, exist_ok=True)
            shutil.copyfile(source, destination)
    return target


def run(args: argparse.Namespace) -> dict:
    ensure_soma_assets()

    from gem.utils.kp2d_utils import render_2d_keypoints
    from scripts.demo import demo_soma_onnx as pipeline
    from space_bvh import export_both_bvhs

    demo_args = SimpleNamespace(
        video=str(Path(args.video).resolve()),
        output_root=str(Path(args.output_root).resolve()),
        static_cam=args.static_camera,
        verbose=False,
        ckpt=None,
        exp="gem_soma_regression",
    )
    cfg = pipeline._build_cfg(demo_args)
    no_imgfeat = True

    pipeline.run_preprocess_fast(cfg, force_pytorch=args.ddim, no_imgfeat=no_imgfeat)
    data = pipeline.load_data_dict(cfg, no_imgfeat=no_imgfeat)
    pred = pipeline.run_inference_fast(
        cfg,
        data,
        force_pytorch=args.ddim,
        no_imgfeat=no_imgfeat,
        use_ddim=args.ddim,
    )
    raw_path = Path(cfg.paths.hpe_results)
    raw_path.parent.mkdir(parents=True, exist_ok=True)
    torch.save(pred, raw_path)

    cap = cv2.VideoCapture(str(cfg.video_path))
    fps = float(cap.get(cv2.CAP_PROP_FPS)) or 30.0
    frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
    cap.release()

    output_dir = Path(cfg.output_dir)
    preview = output_dir / "gemx_77_keypoints.mp4"
    render_2d_keypoints(
        video_path=cfg.video_path,
        vitpose_path=cfg.paths.vitpose,
        bbx_path=cfg.paths.bbx,
        output_path=str(preview),
        fps=fps,
    )

    body_params = pred.get("body_params_global") or pred.get("pred_body_params_global")
    if body_params is None:
        raise RuntimeError("GEM-X output has no world-space SOMA body parameters.")

    bvh77 = output_dir / "soma77_blender.bvh"
    bvh78 = output_dir / "soma78_virtual_root.bvh"
    export_both_bvhs(body_params, fps, bvh77, bvh78)

    archive_base = Path(args.output_root) / "gemx_motion_results"
    archive_path = Path(
        shutil.make_archive(
            str(archive_base),
            "zip",
            root_dir=output_dir,
        )
    )
    return {
        "preview": str(preview),
        "bvh77": str(bvh77),
        "bvh78": str(bvh78),
        "raw": str(raw_path),
        "archive": str(archive_path),
        "fps": fps,
        "frames": frames,
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser()
    parser.add_argument("--video", required=True)
    parser.add_argument("--output-root", required=True)
    parser.add_argument("--static-camera", action="store_true")
    parser.add_argument("--ddim", action="store_true")
    return parser.parse_args()


if __name__ == "__main__":
    result = run(parse_args())
    print("GEMX_RESULT=" + json.dumps(result, ensure_ascii=False), flush=True)