diff --git a/.gitattributes b/.gitattributes
index 565d0bcb2930768808d9fa9e6f407e40d27aefbc..24aaf2304a343c1877873363bdc90a4a17e38fc5 100644
--- a/.gitattributes
+++ b/.gitattributes
@@ -101,3 +101,4 @@ kimodo/assets/skeletons/g1skel34/meshes/g1/waist_constraint_L.STL filter=lfs dif
kimodo/assets/skeletons/g1skel34/meshes/g1/waist_constraint_R.STL filter=lfs diff=lfs merge=lfs -text
kimodo/assets/skeletons/g1skel34/meshes/g1/waist_support_link.STL filter=lfs diff=lfs merge=lfs -text
kimodo/assets/skeletons/g1skel34/meshes/g1/waist_yaw_link.STL filter=lfs diff=lfs merge=lfs -text
+kimodo/assets/skeletons/g1skel34/meshes/g1/waist_yaw_link_rev_1_0.STL filter=lfs diff=lfs merge=lfs -text
diff --git a/kimodo/assets.py b/kimodo/assets.py
new file mode 100644
index 0000000000000000000000000000000000000000..91facad0faed373c9ed1ad4667980cf19788b093
--- /dev/null
+++ b/kimodo/assets.py
@@ -0,0 +1,19 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from pathlib import Path
+
+PACKAGE_ROOT = Path(__file__).resolve().parent
+ASSETS_ROOT = PACKAGE_ROOT / "assets"
+DEMO_ASSETS_ROOT = ASSETS_ROOT / "demo"
+DEMO_EXAMPLES_ROOT = DEMO_ASSETS_ROOT / "examples"
+SKELETONS_ROOT = ASSETS_ROOT / "skeletons"
+SOMA_ASSETS_ROOT = ASSETS_ROOT / "SOMA"
+
+
+def skeleton_asset_path(*parts: str) -> Path:
+ return SKELETONS_ROOT.joinpath(*parts)
+
+
+def demo_asset_path(*parts: str) -> Path:
+ return DEMO_ASSETS_ROOT.joinpath(*parts)
diff --git a/kimodo/assets/skeletons/g1skel34/meshes/g1/waist_yaw_link_rev_1_0.STL b/kimodo/assets/skeletons/g1skel34/meshes/g1/waist_yaw_link_rev_1_0.STL
new file mode 100644
index 0000000000000000000000000000000000000000..dc628fbdda04129fda29140243d148cbd2c81c83
--- /dev/null
+++ b/kimodo/assets/skeletons/g1skel34/meshes/g1/waist_yaw_link_rev_1_0.STL
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:ec6db442b11f25eed898b5add07940c85d804f300de24dcbd264ccd8be7d554c
+size 619984
diff --git a/kimodo/assets/skeletons/g1skel34/rest_pose_local_rot.p b/kimodo/assets/skeletons/g1skel34/rest_pose_local_rot.p
new file mode 100644
index 0000000000000000000000000000000000000000..d441fc0a0b15a434c8455e4ead8cd37753742546
Binary files /dev/null and b/kimodo/assets/skeletons/g1skel34/rest_pose_local_rot.p differ
diff --git a/kimodo/assets/skeletons/g1skel34/xml/g1.xml b/kimodo/assets/skeletons/g1skel34/xml/g1.xml
new file mode 100644
index 0000000000000000000000000000000000000000..36231a5b97ab1784ad2e2c889b974b086639c120
--- /dev/null
+++ b/kimodo/assets/skeletons/g1skel34/xml/g1.xml
@@ -0,0 +1,413 @@
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+ >
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
+
diff --git a/kimodo/assets/skeletons/smplx22/beta.npy b/kimodo/assets/skeletons/smplx22/beta.npy
new file mode 100644
index 0000000000000000000000000000000000000000..d46183ecf1060a78cf4b2b3bcccd13ae98f51551
--- /dev/null
+++ b/kimodo/assets/skeletons/smplx22/beta.npy
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:abfc1cc9e819d017580739e4a65b39458ae840808f01819ad54260283575fa11
+size 1328
diff --git a/kimodo/assets/skeletons/smplx22/joints.p b/kimodo/assets/skeletons/smplx22/joints.p
new file mode 100644
index 0000000000000000000000000000000000000000..fbba652bf10ab0df1dc858fbf7ada72787a56810
Binary files /dev/null and b/kimodo/assets/skeletons/smplx22/joints.p differ
diff --git a/kimodo/assets/skeletons/smplx22/mean_hands.npy b/kimodo/assets/skeletons/smplx22/mean_hands.npy
new file mode 100644
index 0000000000000000000000000000000000000000..acd2d561fca5d3d3f873ac11f6e2c96919f1221e
--- /dev/null
+++ b/kimodo/assets/skeletons/smplx22/mean_hands.npy
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:eff97ce9d98bfb1495501d8f00a883c08bfdfc05079037e9832f5d64b58d0e39
+size 848
diff --git a/kimodo/assets/skeletons/somaskel30/joints.p b/kimodo/assets/skeletons/somaskel30/joints.p
new file mode 100644
index 0000000000000000000000000000000000000000..9e319f6fa3cb4959ed3221040b0ae945ec9646c7
Binary files /dev/null and b/kimodo/assets/skeletons/somaskel30/joints.p differ
diff --git a/kimodo/assets/skeletons/somaskel30/soma_base_fit_mhr_params.npz b/kimodo/assets/skeletons/somaskel30/soma_base_fit_mhr_params.npz
new file mode 100644
index 0000000000000000000000000000000000000000..25ea8aa2ee5d40a79c7cf2761b6b3a28bd7c1133
--- /dev/null
+++ b/kimodo/assets/skeletons/somaskel30/soma_base_fit_mhr_params.npz
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:61f0b96c8386a19d44823faf4399816642b7e20691ed5b398c7402921090096e
+size 1877
diff --git a/kimodo/assets/skeletons/somaskel77/bvh_joints.p b/kimodo/assets/skeletons/somaskel77/bvh_joints.p
new file mode 100644
index 0000000000000000000000000000000000000000..3e3f2195fe6e094af5943cd98007344defa0c426
Binary files /dev/null and b/kimodo/assets/skeletons/somaskel77/bvh_joints.p differ
diff --git a/kimodo/assets/skeletons/somaskel77/joints.p b/kimodo/assets/skeletons/somaskel77/joints.p
new file mode 100644
index 0000000000000000000000000000000000000000..8c5eb2e67e1314d8765c821ae8dd4e0ff1351a3b
Binary files /dev/null and b/kimodo/assets/skeletons/somaskel77/joints.p differ
diff --git a/kimodo/assets/skeletons/somaskel77/relaxed_hands_rest_pose.npy b/kimodo/assets/skeletons/somaskel77/relaxed_hands_rest_pose.npy
new file mode 100644
index 0000000000000000000000000000000000000000..74dc084f8e62e534a64e840141165cd56f6d85e5
--- /dev/null
+++ b/kimodo/assets/skeletons/somaskel77/relaxed_hands_rest_pose.npy
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:64a3828e0d1ef1f1de8228c74eba8040c0810d898169f1892d8997a213b2b64c
+size 2900
diff --git a/kimodo/assets/skeletons/somaskel77/skin_standard.npz b/kimodo/assets/skeletons/somaskel77/skin_standard.npz
new file mode 100644
index 0000000000000000000000000000000000000000..4a82cd790d42a6a48a858b13bf8f6594bac30241
--- /dev/null
+++ b/kimodo/assets/skeletons/somaskel77/skin_standard.npz
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:90ee2cdf50f168382a7dd7e9c88b118298aba0aca4b72022080c89c4bab0ceb2
+size 531434
diff --git a/kimodo/assets/skeletons/somaskel77/somaskel77_standard_tpose.bvh b/kimodo/assets/skeletons/somaskel77/somaskel77_standard_tpose.bvh
new file mode 100644
index 0000000000000000000000000000000000000000..2998cb69ade063fc38d329288de7a14255e46b65
--- /dev/null
+++ b/kimodo/assets/skeletons/somaskel77/somaskel77_standard_tpose.bvh
@@ -0,0 +1,395 @@
+HIERARCHY
+ROOT Root
+{
+ OFFSET 0.0 0.0 0.0
+ CHANNELS 6 Xposition Yposition Zposition Zrotation Yrotation Xrotation
+ JOINT Hips
+ {
+ OFFSET 0.0 100.0 0.0
+ CHANNELS 6 Xposition Yposition Zposition Zrotation Yrotation Xrotation
+ JOINT Spine1
+ {
+ OFFSET -0.013727 5.003763 -0.053727
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT Spine2
+ {
+ OFFSET -0.0 7.125301 -0.029825
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT Chest
+ {
+ OFFSET -1e-06 7.550063 -0.815971
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT Neck1
+ {
+ OFFSET -0.181677 26.311295 -0.553348
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT Neck2
+ {
+ OFFSET -3e-06 7.709397 2.302585
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT Head
+ {
+ OFFSET -5e-06 6.128916 1.953709
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT HeadEnd
+ {
+ OFFSET 0.003598 16.065403 -1.835379
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ JOINT Jaw
+ {
+ OFFSET 0.002637 0.475592 3.094941
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ JOINT LeftEye
+ {
+ OFFSET 3.206381 5.380205 7.586883
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ JOINT RightEye
+ {
+ OFFSET -3.22244 5.361869 7.558234
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ JOINT LeftShoulder
+ {
+ OFFSET 1.621652 23.237164 5.113413
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftArm
+ {
+ OFFSET 14.919846 2e-06 -5.502326
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftForeArm
+ {
+ OFFSET 28.739307 0.0 -0.002588
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHand
+ {
+ OFFSET 27.093981 -1e-06 0.002609
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandThumb1
+ {
+ OFFSET 2.276482 -1.392045 3.191413
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandThumb2
+ {
+ OFFSET 4.012836 -1.828127 1.641654
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandThumb3
+ {
+ OFFSET 2.798515 0.0 -3e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandThumbEnd
+ {
+ OFFSET 3.180793 -4e-06 4e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ JOINT LeftHandIndex1
+ {
+ OFFSET 3.247555 -0.531998 2.296169
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandIndex2
+ {
+ OFFSET 6.364578 0.01206 0.1786
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandIndex3
+ {
+ OFFSET 3.662364 0.0 0.0
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandIndex4
+ {
+ OFFSET 2.329242 4e-06 4e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandIndexEnd
+ {
+ OFFSET 2.759615 -0.180537 -0.113024
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ JOINT LeftHandMiddle1
+ {
+ OFFSET 3.163495 0.240981 1.000332
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandMiddle2
+ {
+ OFFSET 6.19078 -0.259278 -1.002548
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandMiddle3
+ {
+ OFFSET 4.35652 -4e-06 -1e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandMiddle4
+ {
+ OFFSET 2.996877 -8e-06 0.0
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandMiddleEnd
+ {
+ OFFSET 2.304287 -0.294569 -0.031741
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ JOINT LeftHandRing1
+ {
+ OFFSET 2.882643 -0.053652 -0.322543
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandRing2
+ {
+ OFFSET 5.854541 -0.486202 -1.373841
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandRing3
+ {
+ OFFSET 4.350578 0.0 3e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandRing4
+ {
+ OFFSET 2.651321 7e-06 2e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandRingEnd
+ {
+ OFFSET 1.936105 0.077687 -7.1e-05
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ JOINT LeftHandPinky1
+ {
+ OFFSET 2.8655 -0.310005 -1.600378
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandPinky2
+ {
+ OFFSET 5.087849 -1.331141 -1.77123
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandPinky3
+ {
+ OFFSET 3.070974 4e-06 0.0
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandPinky4
+ {
+ OFFSET 1.549672 0.0 1e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftHandPinkyEnd
+ {
+ OFFSET 1.944893 -0.157802 0.057219
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ JOINT RightShoulder
+ {
+ OFFSET -1.380118 23.180309 5.214158
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightArm
+ {
+ OFFSET -15.037196 1.2e-05 -5.545604
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightForeArm
+ {
+ OFFSET -28.736639 2e-06 -0.002597
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHand
+ {
+ OFFSET -27.133619 -0.0 0.002613
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandThumb1
+ {
+ OFFSET -2.274032 -1.383988 3.163127
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandThumb2
+ {
+ OFFSET -4.011429 -1.827466 1.640914
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandThumb3
+ {
+ OFFSET -2.794935 -4e-06 -3e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandThumbEnd
+ {
+ OFFSET -3.183852 4e-06 1e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ JOINT RightHandIndex1
+ {
+ OFFSET -3.253266 -0.520057 2.282866
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandIndex2
+ {
+ OFFSET -6.341917 0.012471 0.178266
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandIndex3
+ {
+ OFFSET -3.654871 -8e-06 -0.0
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandIndex4
+ {
+ OFFSET -2.327586 0.0 1e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandIndexEnd
+ {
+ OFFSET -2.76179 -0.180656 -0.113078
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ JOINT RightHandMiddle1
+ {
+ OFFSET -3.168106 0.246593 1.00103
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandMiddle2
+ {
+ OFFSET -6.180828 -0.258836 -1.000895
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandMiddle3
+ {
+ OFFSET -4.348901 0.0 -0.0
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandMiddle4
+ {
+ OFFSET -3.00024 -4e-06 -2e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandMiddleEnd
+ {
+ OFFSET -2.30252 -0.29437 -0.031706
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ JOINT RightHandRing1
+ {
+ OFFSET -2.88569 -0.067952 -0.308858
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandRing2
+ {
+ OFFSET -5.854198 -0.48613 -1.373731
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandRing3
+ {
+ OFFSET -4.33881 -4e-06 -0.0
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandRing4
+ {
+ OFFSET -2.654903 -4e-06 4e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandRingEnd
+ {
+ OFFSET -1.933568 0.077527 -5.2e-05
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ JOINT RightHandPinky1
+ {
+ OFFSET -2.866425 -0.342796 -1.584145
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandPinky2
+ {
+ OFFSET -5.091371 -1.332055 -1.772385
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandPinky3
+ {
+ OFFSET -3.062664 -4e-06 1e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandPinky4
+ {
+ OFFSET -1.546529 4e-06 -2e-06
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightHandPinkyEnd
+ {
+ OFFSET -1.945119 -0.157718 0.057211
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ }
+ JOINT LeftLeg
+ {
+ OFFSET 10.043214 -8.434526 2.595655
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftShin
+ {
+ OFFSET -1e-06 -43.221752 -0.802913
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftFoot
+ {
+ OFFSET 1e-06 -42.155094 -3.481523
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftToeBase
+ {
+ OFFSET 0.0 -5.059472 13.231529
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT LeftToeEnd
+ {
+ OFFSET -0.009607 -1.647619 6.513017
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ JOINT RightLeg
+ {
+ OFFSET -10.047278 -8.29526 2.620317
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightShin
+ {
+ OFFSET 1e-06 -43.362206 -0.805556
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightFoot
+ {
+ OFFSET 2e-06 -42.117393 -3.478398
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightToeBase
+ {
+ OFFSET -0.0 -5.079609 13.284196
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ JOINT RightToeEnd
+ {
+ OFFSET 0.009532 -1.634378 6.460591
+ CHANNELS 3 Zrotation Yrotation Xrotation
+ }
+ }
+ }
+ }
+ }
+ }
+}
+MOTION
+Frames: 1
+Frame Time: 0.03333333333333333
+0.0 0.0 0.0 0.0 0.0 0.0 0.0 100.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0 0.0
diff --git a/kimodo/assets/skeletons/somaskel77/standard_t_pose_global_offsets_rots.p b/kimodo/assets/skeletons/somaskel77/standard_t_pose_global_offsets_rots.p
new file mode 100644
index 0000000000000000000000000000000000000000..b6d68af1e73df0fc69c2cd9a59f2f6ed6aebb425
Binary files /dev/null and b/kimodo/assets/skeletons/somaskel77/standard_t_pose_global_offsets_rots.p differ
diff --git a/kimodo/constraints.py b/kimodo/constraints.py
new file mode 100644
index 0000000000000000000000000000000000000000..7accd150dc5ccf8481442fc885562e92ed269765
--- /dev/null
+++ b/kimodo/constraints.py
@@ -0,0 +1,625 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Constraint sets for conditioning motion generation (root 2D, full body, end-effectors)."""
+
+from typing import Optional, Union
+
+import torch
+from torch import Tensor
+
+from kimodo.motion_rep.feature_utils import compute_heading_angle
+from kimodo.skeleton import SkeletonBase, SOMASkeleton30, SOMASkeleton77
+from kimodo.tools import ensure_batched, load_json, save_json
+
+from .geometry import axis_angle_to_matrix, matrix_to_axis_angle
+
+
+def _convert_constraint_local_rots_to_skeleton(local_rot_mats: Tensor, skeleton: SkeletonBase) -> Tensor:
+ """Convert loaded local rotation matrices to match the skeleton's joint count.
+
+ Handles SOMA 30↔77: constraint files may have been saved with 30 or 77 joints while the session
+ skeleton (e.g. from the SOMA30 model) uses SOMASkeleton77.
+ """
+ n_joints = local_rot_mats.shape[-3]
+ skeleton_joints = skeleton.nbjoints
+ if n_joints == skeleton_joints:
+ return local_rot_mats
+ if n_joints == 77 and skeleton_joints == 30 and isinstance(skeleton, SOMASkeleton30):
+ return skeleton.from_SOMASkeleton77(local_rot_mats)
+ if n_joints == 30 and skeleton_joints == 77 and isinstance(skeleton, SOMASkeleton77):
+ skel30 = SOMASkeleton30()
+ return skel30.to_SOMASkeleton77(local_rot_mats)
+ raise ValueError(
+ f"Constraint joint count ({n_joints}) does not match skeleton joint count "
+ f"({skeleton_joints}). Only SOMA 30↔77 conversion is supported."
+ )
+
+
+def create_pairs(tensor_A: Tensor, tensor_B: Tensor) -> Tensor:
+ """Form all (a, b) pairs from two 1D tensors; output shape (len(A)*len(B), 2)."""
+ pairs = torch.stack(
+ (
+ tensor_A[:, None].expand(-1, len(tensor_B)),
+ tensor_B.expand(len(tensor_A), -1),
+ ),
+ dim=-1,
+ ).reshape(-1, 2)
+ return pairs
+
+
+def compute_global_heading(global_joints_positions: Tensor, skeleton: SkeletonBase) -> Tensor:
+ """Compute global root heading (cos, sin) from global joint positions using skeleton."""
+ root_heading_angle = compute_heading_angle(global_joints_positions, skeleton)
+ global_root_heading = torch.stack([torch.cos(root_heading_angle), torch.sin(root_heading_angle)], dim=-1)
+ return global_root_heading
+
+
+def _tensor_to(
+ t: Tensor,
+ device: Optional[Union[str, torch.device]] = None,
+ dtype: Optional[torch.dtype] = None,
+) -> Tensor:
+ """Move tensor to device and/or dtype.
+
+ Returns same tensor if no args.
+ """
+ if device is not None and dtype is not None:
+ return t.to(device=device, dtype=dtype)
+ if device is not None:
+ return t.to(device=device)
+ if dtype is not None:
+ return t.to(dtype=dtype)
+ return t
+
+
+class Root2DConstraintSet:
+ """Constraint set fixing root (x, z) trajectory and optionally global heading on given
+ frames."""
+
+ name = "root2d"
+
+ def __init__(
+ self,
+ skeleton: SkeletonBase,
+ frame_indices: Tensor,
+ smooth_root_2d: Tensor,
+ to_crop: bool = False,
+ global_root_heading: Optional[Tensor] = None,
+ ) -> None:
+ self.skeleton = skeleton
+
+ # if we pass the full smooth root 3D as input
+ if smooth_root_2d.shape[-1] == 3:
+ smooth_root_2d = smooth_root_2d[..., [0, 1]]
+
+ if to_crop:
+ smooth_root_2d = smooth_root_2d[frame_indices]
+ if global_root_heading is not None:
+ global_root_heading = global_root_heading[frame_indices]
+ else:
+ assert len(smooth_root_2d) == len(
+ frame_indices
+ ), "The number of smooth root 2d should be match the number of frames"
+ if global_root_heading is not None:
+ assert len(global_root_heading) == len(
+ frame_indices
+ ), "The number of global root heading should be match the number of frames"
+
+ self.smooth_root_2d = smooth_root_2d
+ self.global_root_heading = global_root_heading
+ self.frame_indices = frame_indices
+
+ def update_constraints(self, data_dict: dict, index_dict: dict) -> None:
+ """Append this constraint's smooth_root_2d (and optional global_root_heading) to data/index
+ dicts."""
+ data_dict["smooth_root_2d"].append(self.smooth_root_2d)
+ index_dict["smooth_root_2d"].append(self.frame_indices)
+
+ if self.global_root_heading is not None:
+ # constraint the global heading
+ data_dict["global_root_heading"].append(self.global_root_heading)
+ index_dict["global_root_heading"].append(self.frame_indices)
+
+ def crop_move(self, start: int, end: int) -> "Root2DConstraintSet":
+ """Return a new constraint set for the cropped frame range [start, end)."""
+ mask = (self.frame_indices >= start) & (self.frame_indices < end)
+
+ if self.global_root_heading is not None:
+ masked_global_root_heading = self.global_root_heading[mask]
+ else:
+ masked_global_root_heading = None
+
+ return Root2DConstraintSet(
+ self.skeleton,
+ self.frame_indices[mask] - start,
+ self.smooth_root_2d[mask],
+ global_root_heading=masked_global_root_heading,
+ )
+
+ def get_save_info(self) -> dict:
+ """Return a dict suitable for JSON serialization (frame_indices, smooth_root_2d, optional
+ global_root_heading)."""
+ out = {
+ "type": self.name,
+ "frame_indices": self.frame_indices,
+ "smooth_root_2d": self.smooth_root_2d,
+ }
+ if self.global_root_heading is not None:
+ out["global_root_heading"] = self.global_root_heading
+ return out
+
+ def to(
+ self,
+ device: Optional[Union[str, torch.device]] = None,
+ dtype: Optional[torch.dtype] = None,
+ ) -> "Root2DConstraintSet":
+ self.smooth_root_2d = _tensor_to(self.smooth_root_2d, device, dtype)
+ self.frame_indices = _tensor_to(self.frame_indices, device, dtype)
+ if self.global_root_heading is not None:
+ self.global_root_heading = _tensor_to(self.global_root_heading, device, dtype)
+ if device is not None and hasattr(self.skeleton, "to"):
+ self.skeleton = self.skeleton.to(device)
+ return self
+
+ @classmethod
+ def from_dict(cls, skeleton: SkeletonBase, dico: dict) -> "Root2DConstraintSet":
+ """Build a Root2DConstraintSet from a dict (e.g. loaded from JSON)."""
+ device = skeleton.device if hasattr(skeleton, "device") else "cpu"
+
+ if "global_root_heading" in dico:
+ global_root_heading = torch.tensor(dico["global_root_heading"], device=device)
+ else:
+ global_root_heading = None
+
+ return cls(
+ skeleton,
+ frame_indices=torch.tensor(dico["frame_indices"]),
+ smooth_root_2d=torch.tensor(dico["smooth_root_2d"], device=device),
+ global_root_heading=global_root_heading,
+ )
+
+
+class FullBodyConstraintSet:
+ """Constraint set fixing full-body global positions and rotations on given keyframes."""
+
+ name = "fullbody"
+
+ def __init__(
+ self,
+ skeleton: SkeletonBase,
+ frame_indices: Tensor,
+ global_joints_positions: Tensor,
+ global_joints_rots: Tensor,
+ smooth_root_2d: Optional[Tensor] = None,
+ to_crop: bool = False,
+ ):
+ self.skeleton = skeleton
+ self.frame_indices = frame_indices
+
+ # if we pass the full smooth root 3D as input
+ if smooth_root_2d is not None and smooth_root_2d.shape[-1] == 3:
+ smooth_root_2d = smooth_root_2d[..., [0, 1]]
+
+ if to_crop:
+ global_joints_positions = global_joints_positions[frame_indices]
+ global_joints_rots = global_joints_rots[frame_indices]
+ if smooth_root_2d is not None:
+ smooth_root_2d = smooth_root_2d[frame_indices]
+ else:
+ assert len(global_joints_positions) == len(
+ frame_indices
+ ), "The number of global positions should be match the number of frames"
+ assert len(global_joints_rots) == len(
+ frame_indices
+ ), "The number of global joint rotations should be match the number of frames"
+
+ if smooth_root_2d is not None:
+ assert len(smooth_root_2d) == len(
+ frame_indices
+ ), "The number of smooth root 2d (if specified) should be match the number of frames"
+
+ if smooth_root_2d is None:
+ # substitute the smooth root 2d with the real root
+ smooth_root_2d = global_joints_positions[:, skeleton.root_idx, [0, 2]]
+
+ # root y: from smooth or pelvis is the same
+ self.root_y_pos = global_joints_positions[:, skeleton.root_idx, 1]
+
+ self.global_joints_positions = global_joints_positions
+ self.global_joints_rots = global_joints_rots
+ self.global_root_heading = compute_global_heading(global_joints_positions, skeleton)
+ self.smooth_root_2d = smooth_root_2d
+
+ def update_constraints(self, data_dict: dict, index_dict: dict) -> None:
+ """Append global positions, smooth root 2D, root y, and global heading to data/index
+ dicts."""
+ nbjoints = self.skeleton.nbjoints
+ indices_lst = create_pairs(
+ self.frame_indices,
+ torch.arange(nbjoints, device=self.frame_indices.device),
+ )
+ data_dict["global_joints_positions"].append(
+ self.global_joints_positions.reshape(-1, 3)
+ ) # flatten the global positions
+ index_dict["global_joints_positions"].append(indices_lst)
+
+ # global rotations are not used here
+
+ # as we use smooth root, also constraint the smooth root to get the same full body
+ # maybe keep storing the hips offset, if we smooth it ourselves
+ data_dict["smooth_root_2d"].append(self.smooth_root_2d)
+ index_dict["smooth_root_2d"].append(self.frame_indices)
+
+ # constraint the y pos of the root
+ data_dict["root_y_pos"].append(self.root_y_pos)
+ index_dict["root_y_pos"].append(self.frame_indices)
+
+ # constraint the global heading
+ data_dict["global_root_heading"].append(self.global_root_heading)
+ index_dict["global_root_heading"].append(self.frame_indices)
+
+ def crop_move(self, start: int, end: int) -> "FullBodyConstraintSet":
+ """Return a new FullBodyConstraintSet for the cropped frame range [start, end)."""
+ mask = (self.frame_indices >= start) & (self.frame_indices < end)
+ return FullBodyConstraintSet(
+ self.skeleton,
+ self.frame_indices[mask] - start,
+ self.global_joints_positions[mask],
+ self.global_joints_rots[mask],
+ self.smooth_root_2d[mask],
+ )
+
+ def get_save_info(self) -> dict:
+ """Return a dict for JSON save: type, frame_indices, local_joints_rot, root_positions, smooth_root_2d."""
+ local_joints_rot = self.skeleton.global_rots_to_local_rots(self.global_joints_rots)
+ if isinstance(self.skeleton, SOMASkeleton30):
+ local_joints_rot = self.skeleton.to_SOMASkeleton77(local_joints_rot)
+ local_joints_rot = matrix_to_axis_angle(local_joints_rot)
+
+ root_positions = self.global_joints_positions[:, self.skeleton.root_idx]
+ return {
+ "type": self.name,
+ "frame_indices": self.frame_indices,
+ "local_joints_rot": local_joints_rot,
+ "root_positions": root_positions,
+ "smooth_root_2d": self.smooth_root_2d,
+ }
+
+ def to(
+ self,
+ device: Optional[Union[str, torch.device]] = None,
+ dtype: Optional[torch.dtype] = None,
+ ) -> "FullBodyConstraintSet":
+ self.frame_indices = _tensor_to(self.frame_indices, device, dtype)
+ self.global_joints_positions = _tensor_to(self.global_joints_positions, device, dtype)
+ self.global_joints_rots = _tensor_to(self.global_joints_rots, device, dtype)
+ self.root_y_pos = _tensor_to(self.root_y_pos, device, dtype)
+ self.global_root_heading = _tensor_to(self.global_root_heading, device, dtype)
+ self.smooth_root_2d = _tensor_to(self.smooth_root_2d, device, dtype)
+ if device is not None and hasattr(self.skeleton, "to"):
+ self.skeleton = self.skeleton.to(device)
+ return self
+
+ @classmethod
+ def from_dict(cls, skeleton: SkeletonBase, dico: dict) -> "FullBodyConstraintSet":
+ """Build a FullBodyConstraintSet from a dict (e.g. loaded from JSON)."""
+ frame_indices = torch.tensor(dico["frame_indices"])
+ device = skeleton.device if hasattr(skeleton, "device") else "cpu"
+ local_rot = torch.tensor(dico["local_joints_rot"], device=device)
+ local_rot_mats = axis_angle_to_matrix(local_rot)
+ local_rot_mats = _convert_constraint_local_rots_to_skeleton(local_rot_mats, skeleton)
+ global_joints_rots, global_joints_positions, _ = skeleton.fk(
+ local_rot_mats,
+ torch.tensor(dico["root_positions"], device=device),
+ )
+ smooth_root_2d = None
+ if "smooth_root_2d" in dico:
+ smooth_root_2d = torch.tensor(dico["smooth_root_2d"], device=device)
+
+ return cls(
+ skeleton,
+ frame_indices=frame_indices,
+ global_joints_positions=global_joints_positions,
+ global_joints_rots=global_joints_rots,
+ smooth_root_2d=smooth_root_2d,
+ )
+
+
+class EndEffectorConstraintSet:
+ """Constraint set fixing selected end-effector positions and rotations on given frames."""
+
+ name = "end-effector"
+
+ def __init__(
+ self,
+ skeleton: SkeletonBase,
+ frame_indices: Tensor,
+ global_joints_positions: Tensor,
+ global_joints_rots: Tensor,
+ smooth_root_2d: Optional[Tensor],
+ *,
+ joint_names: list[str],
+ to_crop: bool = False,
+ ) -> None:
+ self.skeleton = skeleton
+ self.frame_indices = frame_indices
+ self.joint_names = joint_names
+
+ # joint_names are constant for all the frames
+ rot_joint_names, pos_joint_names = self.skeleton.expand_joint_names(self.joint_names)
+ # indexing works for motion_rep with smooth root only (contains pelvis index)
+ self.pos_indices = torch.tensor([self.skeleton.bone_index[jname] for jname in pos_joint_names])
+ self.rot_indices = torch.tensor([self.skeleton.bone_index[jname] for jname in rot_joint_names])
+
+ # if we pass the full smooth root 3D as input
+ if smooth_root_2d is not None and smooth_root_2d.shape[-1] == 3:
+ smooth_root_2d = smooth_root_2d[..., [0, 1]]
+
+ if to_crop:
+ global_joints_positions = global_joints_positions[frame_indices]
+ global_joints_rots = global_joints_rots[frame_indices]
+ if smooth_root_2d is not None:
+ smooth_root_2d = smooth_root_2d[frame_indices]
+ else:
+ assert len(global_joints_positions) == len(
+ frame_indices
+ ), "The number of global positions should be match the number of frames"
+ assert len(global_joints_rots) == len(
+ frame_indices
+ ), "The number of global joint rotations should be match the number of frames"
+ if smooth_root_2d is not None:
+ assert len(smooth_root_2d) == len(
+ frame_indices
+ ), "The number of smooth root 2d (if specified) should be match the number of frames"
+
+ if smooth_root_2d is None:
+ # substitute the smooth root 2d with the real root
+ smooth_root_2d = global_joints_positions[:, skeleton.root_idx, [0, 2]]
+
+ # root y: from smooth or pelvis is the same
+ self.root_y_pos = global_joints_positions[:, skeleton.root_idx, 1]
+
+ self.global_joints_positions = global_joints_positions
+ self.global_root_heading = compute_global_heading(global_joints_positions, skeleton)
+ self.global_joints_rots = global_joints_rots
+ self.smooth_root_2d = smooth_root_2d
+
+ def update_constraints(self, data_dict: dict, index_dict: dict) -> None:
+ """Append constrained joint positions/rots, smooth root 2D, root y, and heading to
+ data/index dicts."""
+ crop_frames_indexing = torch.arange(len(self.frame_indices), device=self.frame_indices.device)
+
+ # constraint positions
+ pos_indices_real = create_pairs(
+ self.frame_indices,
+ self.pos_indices,
+ )
+ pos_indices_crop = create_pairs(
+ crop_frames_indexing,
+ self.pos_indices,
+ )
+ data_dict["global_joints_positions"].append(self.global_joints_positions[tuple(pos_indices_crop.T)])
+ index_dict["global_joints_positions"].append(pos_indices_real)
+
+ # constraint rotations
+ rot_indices_real = create_pairs(
+ self.frame_indices,
+ self.rot_indices,
+ )
+ rot_indices_crop = create_pairs(
+ crop_frames_indexing,
+ self.rot_indices,
+ )
+ data_dict["global_joints_rots"].append(self.global_joints_rots[tuple(rot_indices_crop.T)])
+ index_dict["global_joints_rots"].append(rot_indices_real)
+
+ # as we use smooth root, also constraint the smooth root to get the same full body
+ # maybe keep storing the hips offset, if we smooth it ourselves
+ data_dict["smooth_root_2d"].append(self.smooth_root_2d)
+ index_dict["smooth_root_2d"].append(self.frame_indices)
+
+ # constraint the y pos of the root
+ data_dict["root_y_pos"].append(self.root_y_pos)
+ index_dict["root_y_pos"].append(self.frame_indices)
+
+ # constraint the global heading
+ data_dict["global_root_heading"].append(self.global_root_heading)
+ index_dict["global_root_heading"].append(self.frame_indices)
+
+ def crop_move(self, start: int, end: int) -> "EndEffectorConstraintSet":
+ """Return a new EndEffectorConstraintSet for the cropped frame range [start, end)."""
+ mask = (self.frame_indices >= start) & (self.frame_indices < end)
+
+ cls = type(self)
+ kwargs = {}
+ if not hasattr(cls, "joint_names"):
+ kwargs["joint_names"] = self.joint_names
+
+ return cls(
+ self.skeleton,
+ self.frame_indices[mask] - start,
+ self.global_joints_positions[mask],
+ self.global_joints_rots[mask],
+ self.smooth_root_2d[mask],
+ **kwargs,
+ )
+
+ def get_save_info(self) -> dict:
+ """Return a dict for JSON save: type, frame_indices, local_joints_rot, root_positions, smooth_root_2d, joint_names."""
+ local_joints_rot = self.skeleton.global_rots_to_local_rots(self.global_joints_rots)
+ if isinstance(self.skeleton, SOMASkeleton30):
+ local_joints_rot = self.skeleton.to_SOMASkeleton77(local_joints_rot)
+ local_joints_rot = matrix_to_axis_angle(local_joints_rot)
+
+ root_positions = self.global_joints_positions[:, self.skeleton.root_idx]
+ output = {
+ "type": self.name,
+ "frame_indices": self.frame_indices,
+ "local_joints_rot": local_joints_rot,
+ "root_positions": root_positions,
+ "smooth_root_2d": self.smooth_root_2d,
+ }
+ if not hasattr(self.__class__, "joint_names"):
+ # save the joint_names for this base class
+ # but not for children
+ output["joint_names"] = self.joint_names
+ return output
+
+ def to(
+ self,
+ device: Optional[Union[str, torch.device]] = None,
+ dtype: Optional[torch.dtype] = None,
+ ) -> "EndEffectorConstraintSet":
+ self.frame_indices = _tensor_to(self.frame_indices, device, dtype)
+ self.pos_indices = _tensor_to(self.pos_indices, device, dtype)
+ self.rot_indices = _tensor_to(self.rot_indices, device, dtype)
+ self.root_y_pos = _tensor_to(self.root_y_pos, device, dtype)
+ self.global_joints_positions = _tensor_to(self.global_joints_positions, device, dtype)
+ self.global_root_heading = _tensor_to(self.global_root_heading, device, dtype)
+ self.global_joints_rots = _tensor_to(self.global_joints_rots, device, dtype)
+ self.smooth_root_2d = _tensor_to(self.smooth_root_2d, device, dtype)
+ if device is not None and hasattr(self.skeleton, "to"):
+ self.skeleton = self.skeleton.to(device)
+ return self
+
+ @classmethod
+ def from_dict(cls, skeleton: SkeletonBase, dico: dict) -> "EndEffectorConstraintSet":
+ """Build an EndEffectorConstraintSet from a dict (e.g. loaded from JSON)."""
+ frame_indices = torch.tensor(dico["frame_indices"])
+ device = skeleton.device if hasattr(skeleton, "device") else "cpu"
+ local_rot = torch.tensor(dico["local_joints_rot"], device=device)
+ local_rot_mats = axis_angle_to_matrix(local_rot)
+ local_rot_mats = _convert_constraint_local_rots_to_skeleton(local_rot_mats, skeleton)
+ global_joints_rots, global_joints_positions, _ = skeleton.fk(
+ local_rot_mats,
+ torch.tensor(dico["root_positions"], device=device),
+ )
+ smooth_root_2d = None
+ if "smooth_root_2d" in dico:
+ smooth_root_2d = torch.tensor(dico["smooth_root_2d"], device=device)
+
+ kwargs = {}
+ if not hasattr(cls, "joint_names"):
+ kwargs["joint_names"] = dico["joint_names"]
+
+ return cls(
+ skeleton,
+ frame_indices=frame_indices,
+ global_joints_positions=global_joints_positions,
+ global_joints_rots=global_joints_rots,
+ smooth_root_2d=smooth_root_2d,
+ **kwargs,
+ )
+
+
+class LeftHandConstraintSet(EndEffectorConstraintSet):
+ """End-effector constraint for the left hand only."""
+
+ name = "left-hand"
+ joint_names: list[str] = ["LeftHand"]
+
+ def __init__(self, *args, **kwargs: dict):
+ super().__init__(*args, joint_names=self.joint_names, **kwargs)
+
+
+class RightHandConstraintSet(EndEffectorConstraintSet):
+ """End-effector constraint for the right hand only."""
+
+ name = "right-hand"
+ joint_names: list[str] = ["RightHand"]
+
+ def __init__(self, *args, **kwargs: dict):
+ super().__init__(*args, joint_names=self.joint_names, **kwargs)
+
+
+class LeftFootConstraintSet(EndEffectorConstraintSet):
+ """End-effector constraint for the left foot only."""
+
+ name = "left-foot"
+ joint_names: list[str] = ["LeftFoot"]
+
+ def __init__(self, *args, **kwargs: dict):
+ super().__init__(*args, joint_names=self.joint_names, **kwargs)
+
+
+class RightFootConstraintSet(EndEffectorConstraintSet):
+ """End-effector constraint for the right foot only."""
+
+ name = "right-foot"
+ joint_names: list[str] = ["RightFoot"]
+
+ def __init__(self, *args, **kwargs: dict):
+ super().__init__(*args, joint_names=self.joint_names, **kwargs)
+
+
+TYPE_TO_CLASS = {
+ "root2d": Root2DConstraintSet,
+ "fullbody": FullBodyConstraintSet,
+ "left-hand": LeftHandConstraintSet,
+ "right-hand": RightHandConstraintSet,
+ "left-foot": LeftFootConstraintSet,
+ "right-foot": RightFootConstraintSet,
+ "end-effector": EndEffectorConstraintSet,
+}
+
+
+def load_constraints_lst(
+ path_or_data: str | list,
+ skeleton: SkeletonBase,
+ device: Optional[Union[str, torch.device]] = None,
+ dtype: Optional[torch.dtype] = None,
+):
+ """Load a list of constraints from JSON path or list of dicts.
+
+ Args:
+ path_or_data: Path to constraints.json or list of constraint dicts.
+ skeleton: Skeleton instance (used for from_dict).
+ device: If set, move all constraint tensors and skeleton to this device.
+ dtype: If set, cast constraint tensors to this dtype.
+ """
+ if isinstance(path_or_data, str):
+ saved = load_json(path_or_data)
+ else:
+ saved = path_or_data
+
+ constraints_lst = []
+ for el in saved:
+ cls = TYPE_TO_CLASS[el["type"]]
+ c = cls.from_dict(skeleton, el)
+ if device is not None or dtype is not None:
+ c.to(device=device, dtype=dtype)
+ constraints_lst.append(c)
+ return constraints_lst
+
+
+def save_constraints_lst(path: str, constraints_lst: list) -> list | None:
+ """Save a list of constraint sets to a JSON file.
+
+ Returns None if list is empty.
+ """
+ if not constraints_lst:
+ print("The constraints lst is empty. Skip saving")
+ return
+
+ to_save = []
+
+ def tensor_to_list(obj):
+ """Recursively convert tensors to lists for JSON serialization."""
+ if isinstance(obj, Tensor):
+ return obj.cpu().tolist()
+ elif isinstance(obj, dict):
+ return {k: tensor_to_list(v) for k, v in obj.items()}
+ elif isinstance(obj, list):
+ return [tensor_to_list(v) for v in obj]
+ else:
+ return obj
+
+ for constraint in constraints_lst:
+ constraint_info = constraint.get_save_info()
+ # Convert all tensors to lists for JSON serialization
+ constraint_info = tensor_to_list(constraint_info)
+ to_save.append(constraint_info)
+
+ save_json(path, to_save)
+ print(f"Saved constraints to {path}")
+ return to_save
diff --git a/kimodo/demo/__init__.py b/kimodo/demo/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..967579971a3b88b9979b8416245693c15051c8ef
--- /dev/null
+++ b/kimodo/demo/__init__.py
@@ -0,0 +1,29 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+# ruff: noqa: I001
+import argparse
+
+from kimodo.model import DEFAULT_MODEL
+from kimodo.model.registry import resolve_model_name
+
+from .app import Demo
+
+
+def main() -> None:
+ parser = argparse.ArgumentParser(description="Run the kimodo demo UI.")
+ parser.add_argument(
+ "--model",
+ type=str,
+ default=DEFAULT_MODEL,
+ help="Default model to load (e.g. Kimodo-SOMA-RP-v1, kimodo-soma-rp, or SOMA).",
+ )
+ args = parser.parse_args()
+
+ resolved = resolve_model_name(args.model, "Kimodo")
+ demo = Demo(default_model_name=resolved)
+ demo.run()
+
+
+if __name__ == "__main__":
+ main()
diff --git a/kimodo/demo/__main__.py b/kimodo/demo/__main__.py
new file mode 100644
index 0000000000000000000000000000000000000000..444b3563f3c1d6ee659edb044c2dbb69f7066932
--- /dev/null
+++ b/kimodo/demo/__main__.py
@@ -0,0 +1,8 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Entry point for `python -m kimodo.demo`."""
+
+from kimodo.demo import main
+
+if __name__ == "__main__":
+ main()
diff --git a/kimodo/demo/app.py b/kimodo/demo/app.py
new file mode 100644
index 0000000000000000000000000000000000000000..a84fea15d685e43378fdbefe0dfc729c2d93ca8d
--- /dev/null
+++ b/kimodo/demo/app.py
@@ -0,0 +1,690 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+import base64
+import os
+import shutil
+import threading
+import time
+from typing import Optional
+
+import numpy as np
+import torch
+
+import viser
+from kimodo.assets import DEMO_ASSETS_ROOT
+from kimodo.model.load_model import load_model
+from kimodo.model.registry import resolve_model_name
+from kimodo.skeleton import SkeletonBase, SOMASkeleton30
+from kimodo.tools import load_json
+from kimodo.viz import viser_utils
+from kimodo.viz.viser_utils import (
+ Character,
+ CharacterMotion,
+ EEJointsKeyframeSet,
+ FullbodyKeyframeSet,
+ RootKeyframe2DSet,
+)
+from viser.theme import TitlebarButton, TitlebarConfig, TitlebarImage
+
+from . import generation, ui
+from .config import (
+ DARK_THEME,
+ DEFAULT_CUR_DURATION,
+ DEFAULT_MODEL,
+ DEFAULT_PLAYBACK_SPEED,
+ DEFAULT_PROMPT,
+ DEMO_UI_QUICK_START_MODAL_MD,
+ EXAMPLES_ROOT_DIR,
+ HF_MODE,
+ LIGHT_THEME,
+ MAX_ACTIVE_USERS,
+ MAX_DURATION,
+ MAX_SESSION_MINUTES,
+ MIN_DURATION,
+ MODEL_EXAMPLES_DIRS,
+ MODEL_NAMES,
+ SERVER_NAME,
+ SERVER_PORT,
+)
+from .embedding_cache import CachedTextEncoder
+from .queue_manager import QueueManager, UserQueue
+from .state import ClientSession, ModelBundle
+
+
+class Demo:
+ def __init__(self, default_model_name: str = DEFAULT_MODEL):
+ self.device = "cuda:0" if torch.cuda.is_available() else "cpu"
+ print(f"Using device: {self.device}")
+ self.models: dict[str, ModelBundle] = {}
+ self._text_encoder = None
+ resolved = resolve_model_name(default_model_name, "Kimodo")
+ if resolved not in MODEL_NAMES:
+ raise ValueError(f"Unknown model '{default_model_name}'. Expected one of: {MODEL_NAMES}")
+ self.default_model_name = resolved
+ self.ensure_examples_layout()
+ self.load_model(self.default_model_name)
+
+ # Serialize GPU-bound generation across all clients
+ self._generation_lock = threading.Lock()
+ self._cuda_healthy = True
+
+ # Per-client sessions
+ self.client_sessions: dict[int, ClientSession] = {}
+ self.start_direction_markers: dict[int, viser_utils.WaypointMesh] = {}
+ self.grid_handles: dict[int, viser.GridHandle] = {}
+
+ self.server = viser.ViserServer(
+ host=SERVER_NAME,
+ port=SERVER_PORT,
+ label="Kimodo",
+ enable_camera_keyboard_controls=False, # don't move the camera with the arrow keys
+ )
+ self.server.scene.world_axes.visible = False # used for debugging
+ self.server.scene.set_up_direction("+y")
+
+ # Register callbacks for session handling
+ self.server.on_client_connect(self.on_client_connect)
+ self.server.on_client_disconnect(self.on_client_disconnect)
+
+ # HF mode: queue and session limit
+ if HF_MODE:
+ self.user_queue = UserQueue(MAX_ACTIVE_USERS, MAX_SESSION_MINUTES)
+ self.queue_manager = QueueManager(
+ queue=self.user_queue,
+ server=self.server,
+ setup_demo_for_client=self._setup_demo_for_client,
+ cleanup_session=self._cleanup_session_for_client,
+ )
+ else:
+ self.user_queue = None
+ self.queue_manager = None
+
+ # create grid and floor
+ self.floor_len = 20.0 # meters
+
+ def ensure_examples_layout(self) -> None:
+ os.makedirs(EXAMPLES_ROOT_DIR, exist_ok=True)
+ for model_dir in MODEL_EXAMPLES_DIRS.values():
+ os.makedirs(model_dir, exist_ok=True)
+
+ for entry in os.listdir(EXAMPLES_ROOT_DIR):
+ if entry in MODEL_EXAMPLES_DIRS:
+ continue
+ src = os.path.join(EXAMPLES_ROOT_DIR, entry)
+ if not os.path.isdir(src):
+ continue
+ dst = os.path.join(
+ MODEL_EXAMPLES_DIRS.get(DEFAULT_MODEL, next(iter(MODEL_EXAMPLES_DIRS.values()))),
+ entry,
+ )
+ if not os.path.exists(dst):
+ shutil.move(src, dst)
+
+ def get_examples_base_dir(self, model_name: str, absolute: bool = True) -> str:
+ return MODEL_EXAMPLES_DIRS[model_name]
+
+ def load_model(self, model_name: str) -> ModelBundle:
+ if model_name in self.models:
+ return self.models[model_name]
+
+ print(f"Loading model {model_name}...")
+ try:
+ model = load_model(
+ modelname=model_name,
+ device=self.device,
+ text_encoder=self._text_encoder,
+ )
+ except Exception as e:
+ print(f"Error loading model: {e}\nMake sure text encoder server is running!")
+ raise e
+
+ if hasattr(model, "text_encoder"):
+ if self._text_encoder is None:
+ self._text_encoder = model.text_encoder
+ model.text_encoder = CachedTextEncoder(model.text_encoder, model_name=model_name)
+
+ skeleton = model.motion_rep.skeleton
+ if isinstance(skeleton, SOMASkeleton30):
+ skeleton = skeleton.somaskel77.to(model.device)
+ bundle = ModelBundle(
+ model=model,
+ motion_rep=model.motion_rep,
+ skeleton=skeleton,
+ model_fps=model.motion_rep.fps,
+ )
+ self.models[model_name] = bundle
+ print(f"Model {model_name} loaded successfully")
+ self.prewarm_embedding_cache(model_name, bundle.model)
+ return bundle
+
+ def prewarm_embedding_cache(self, model_name: str, model: object) -> None:
+ encoder = getattr(model, "text_encoder", None)
+ if not isinstance(encoder, CachedTextEncoder):
+ return
+
+ prompt_set = set()
+ prompt_set.add(DEFAULT_PROMPT)
+
+ examples_dir = MODEL_EXAMPLES_DIRS.get(model_name)
+ if examples_dir and os.path.isdir(examples_dir):
+ for entry in os.listdir(examples_dir):
+ example_dir = os.path.join(examples_dir, entry)
+ if not os.path.isdir(example_dir):
+ continue
+ meta_path = os.path.join(example_dir, "meta.json")
+ if not os.path.exists(meta_path):
+ continue
+ try:
+ meta = load_json(meta_path)
+ except Exception:
+ continue
+ for prompt in meta.get("prompts_text", []):
+ if isinstance(prompt, str):
+ prompt_set.add(prompt)
+
+ if prompt_set:
+ encoder.prewarm(list(prompt_set))
+
+ def build_constraint_tracks(
+ self, client: viser.ClientHandle, skeleton: SkeletonBase
+ ) -> dict[str, viser_utils.ConstraintSet]:
+ return {
+ "Full-Body": FullbodyKeyframeSet(
+ name="Full-Body",
+ server=client,
+ skeleton=skeleton,
+ ),
+ "End-Effectors": EEJointsKeyframeSet(
+ name="End-Effectors",
+ server=client,
+ skeleton=skeleton,
+ ),
+ "2D Root": RootKeyframe2DSet(
+ name="2D Root",
+ server=client,
+ skeleton=skeleton,
+ ),
+ }
+
+ def set_timeline_defaults(self, timeline, model_fps: float) -> None:
+ timeline.set_defaults(
+ default_text=DEFAULT_PROMPT,
+ default_duration=int(DEFAULT_CUR_DURATION * model_fps - 1),
+ min_duration=int(MIN_DURATION * model_fps - 1), # 2 seconds minimum,
+ max_duration=int(
+ MAX_DURATION * model_fps - 1 # - NB_TRANSITION_FRAMES
+ ), # 10 seconds maximum, minus the transition frames, if needed
+ default_num_frames_zoom=int(1.10 * 10 * model_fps), # a bit more than the max
+ max_frames_zoom=1000,
+ fps=model_fps,
+ )
+
+ def _apply_constraint_overlay_visibility(self, session: ClientSession) -> None:
+ """Apply show-all vs show-only-current-frame to constraint overlays."""
+ only_frame = session.frame_idx if session.show_only_current_constraint else None
+ for constraint in session.constraints.values():
+ constraint.set_overlay_visibility(only_frame)
+
+ def set_constraint_tracks_visible(self, session: ClientSession, visible: bool) -> None:
+ timeline = session.client.timeline
+ timeline_data = session.timeline_data
+ if timeline_data.get("constraint_tracks_visible", True) == visible:
+ return
+
+ with timeline_data["keyframe_update_lock"]:
+ if visible:
+ for track_id, track_info in timeline_data["tracks"].items():
+ timeline.add_track(
+ track_info["name"],
+ track_type=track_info.get("track_type", "keyframe"),
+ color=track_info.get("color"),
+ height_scale=track_info.get("height_scale", 1.0),
+ uuid=track_id,
+ )
+
+ for keyframe_id, keyframe_data in timeline_data["keyframes"].items():
+ timeline.add_keyframe(
+ track_id=keyframe_data["track_id"],
+ frame=keyframe_data["frame"],
+ value=keyframe_data.get("value"),
+ opacity=keyframe_data.get("opacity", 1.0),
+ locked=keyframe_data.get("locked", False),
+ uuid=keyframe_id,
+ )
+
+ for interval_id, interval_data in timeline_data["intervals"].items():
+ timeline.add_interval(
+ track_id=interval_data["track_id"],
+ start_frame=interval_data["start_frame_idx"],
+ end_frame=interval_data["end_frame_idx"],
+ value=interval_data.get("value"),
+ opacity=interval_data.get("opacity", 1.0),
+ locked=interval_data.get("locked", False),
+ uuid=interval_id,
+ )
+ else:
+ for track_id in list(timeline_data["tracks"].keys()):
+ timeline.remove_track(track_id)
+
+ timeline_data["constraint_tracks_visible"] = visible
+
+ def _cleanup_session_for_client(self, client_id: int) -> None:
+ """Remove session and scene state for a client (e.g. on session expiry)."""
+ if client_id in self.client_sessions:
+ del self.client_sessions[client_id]
+ self.start_direction_markers.pop(client_id, None)
+ self.grid_handles.pop(client_id, None)
+
+ def _setup_demo_for_client(self, client: viser.ClientHandle) -> None:
+ """Initialize scene, GUI, and session state for a client (no modals)."""
+ self.setup_scene(client)
+
+ model_bundle = self.load_model(self.default_model_name)
+
+ # Initialize each empty constraint track
+ constraint_tracks = self.build_constraint_tracks(client, model_bundle.skeleton)
+
+ # Create GUI elements for this client
+ (
+ gui_elements,
+ timeline_tracks,
+ example_dict,
+ gui_examples_dropdown,
+ gui_save_example_path_text,
+ gui_model_selector,
+ ) = ui.create_gui(
+ demo=self,
+ client=client,
+ model_name=self.default_model_name,
+ model_fps=model_bundle.model_fps,
+ )
+ timeline_data = {
+ "tracks": timeline_tracks,
+ "tracks_ids": {val["name"]: key for key, val in timeline_tracks.items()},
+ "keyframes": {},
+ "intervals": {},
+ "keyframe_update_lock": threading.Lock(),
+ "keyframe_move_timers": {},
+ "pending_keyframe_moves": {}, # keyframe_id -> new_frame
+ "constraint_tracks_visible": True,
+ "dense_path_after_release_timer": None,
+ }
+
+ # Initialize session state
+ cur_duration = DEFAULT_CUR_DURATION
+ max_frame_idx = int(cur_duration * model_bundle.model_fps - 1)
+
+ session = ClientSession(
+ client=client,
+ gui_elements=gui_elements,
+ motions={},
+ constraints=constraint_tracks,
+ timeline_data=timeline_data,
+ frame_idx=0,
+ playing=False,
+ playback_speed=DEFAULT_PLAYBACK_SPEED,
+ cur_duration=cur_duration,
+ max_frame_idx=max_frame_idx,
+ updating_motions=False,
+ edit_mode=False,
+ model_name=self.default_model_name,
+ model_fps=model_bundle.model_fps,
+ skeleton=model_bundle.skeleton,
+ motion_rep=model_bundle.motion_rep,
+ examples_base_dir=self.get_examples_base_dir(self.default_model_name, absolute=True),
+ example_dict=example_dict,
+ gui_examples_dropdown=gui_examples_dropdown,
+ gui_save_example_path_text=gui_save_example_path_text,
+ gui_model_selector=gui_model_selector,
+ )
+
+ self.client_sessions[client.client_id] = session
+
+ # Initialize default character for this client
+ self.add_character_motion(client, session.skeleton)
+
+ def on_client_connect(self, client: viser.ClientHandle) -> None:
+ """Initialize GUI and state for each new client."""
+ print(f"Client {client.client_id} connected")
+
+ if HF_MODE and self.queue_manager is not None:
+ self.queue_manager.on_client_connect(client)
+ else:
+ # Show quick start popup when a browser client connects (non-HF mode).
+ with client.gui.add_modal(
+ "Welcome — Quick Start",
+ size="xl",
+ show_close_button=True,
+ save_choice="kimodo.demo.quick_start_ack",
+ ) as modal:
+ client.gui.add_markdown(DEMO_UI_QUICK_START_MODAL_MD)
+ client.gui.add_button("Got it (don't remind me again)").on_click(lambda _event: modal.close())
+ self._setup_demo_for_client(client)
+
+ def setup_scene(self, client: viser.ClientHandle) -> None:
+ self.configure_theme(client)
+ client.camera.position = np.array(
+ [2.7417358737841426, 1.8790455698853281, 7.675741569777456],
+ dtype=np.float64,
+ )
+ client.camera.look_at = np.array([0.0, 0.0, 0.0], dtype=np.float64)
+ client.camera.up_direction = np.array(
+ [-1.1102230246251568e-16, 1.0, 1.3596310734468913e-32],
+ dtype=np.float64,
+ )
+ client.camera.fov = np.deg2rad(45.0)
+ grid_handle = client.scene.add_grid(
+ "/grid",
+ width=self.floor_len,
+ height=self.floor_len,
+ wxyz=viser.transforms.SO3.from_x_radians(-np.pi / 2.0).wxyz,
+ position=(0.0, 0.0001, 0.0),
+ fade_distance=3 * self.floor_len,
+ section_color=LIGHT_THEME["grid"],
+ infinite_grid=True,
+ )
+ self.grid_handles[client.client_id] = grid_handle
+ # marker for origin
+ origin_waypoint = viser_utils.WaypointMesh(
+ "/origin_waypoint",
+ client,
+ position=np.array([0.0, 0.0, 0.0]),
+ heading=np.array([0.0, 1.0]),
+ color=(0, 0, 255),
+ )
+ self.start_direction_markers[client.client_id] = origin_waypoint
+
+ def on_client_disconnect(self, client: viser.ClientHandle) -> None:
+ """Clean up when client disconnects."""
+ print(f"Client {client.client_id} disconnected")
+ client_id = client.client_id
+
+ if HF_MODE and self.queue_manager is not None:
+ self.queue_manager.on_client_disconnect(client_id)
+
+ self._cleanup_session_for_client(client_id)
+
+ def set_start_direction_visible(self, client_id: int, visible: bool) -> None:
+ marker = self.start_direction_markers.get(client_id)
+ if marker is None:
+ return
+ marker.set_visible(visible)
+
+ def client_active(self, client_id: int) -> bool:
+ return client_id in self.client_sessions
+
+ def add_character_motion(
+ self,
+ client: viser.ClientHandle,
+ skeleton: SkeletonBase,
+ joints_pos: Optional[torch.Tensor] = None,
+ joints_rot: Optional[torch.Tensor] = None,
+ foot_contacts: Optional[torch.Tensor] = None,
+ ) -> None:
+ client_id = client.client_id
+ if not self.client_active(client_id):
+ return
+ session = self.client_sessions[client_id]
+
+ ci = len(session.motions)
+ character_name = f"character{ci}"
+ # build character skeleton and skinning mesh
+ if "g1" in session.model_name:
+ mesh_mode = "g1_stl"
+ elif "smplx" in session.model_name:
+ mesh_mode = "smplx_skin"
+ elif "soma" in session.model_name:
+ if session.gui_elements.gui_use_soma_layer_checkbox.value:
+ mesh_mode = "soma_layer_skin"
+ else:
+ mesh_mode = "soma_skin"
+ else:
+ raise ValueError("The model name is not recognized for skinning.")
+
+ new_character = Character(
+ character_name,
+ client,
+ skeleton,
+ create_skeleton_mesh=True,
+ create_skinned_mesh=True,
+ visible_skeleton=False, # don't show immediately
+ visible_skinned_mesh=False, # don't show immediately
+ skinned_mesh_opacity=session.gui_elements.gui_viz_skinned_mesh_opacity_slider.value,
+ show_foot_contacts=session.gui_elements.gui_viz_foot_contacts_checkbox.value,
+ dark_mode=session.gui_elements.gui_dark_mode_checkbox.value,
+ mesh_mode=mesh_mode,
+ gui_use_soma_layer_checkbox=session.gui_elements.gui_use_soma_layer_checkbox,
+ )
+
+ # if no motion given, initialize to character default (rest) pose for one frame
+ init_joints_pos, init_joints_rot = new_character.get_pose()
+ if joints_pos is None:
+ joints_pos = init_joints_pos[None].repeat(session.max_frame_idx + 1, 1, 1)
+ if joints_rot is None:
+ joints_rot = init_joints_rot[None].repeat(session.max_frame_idx + 1, 1, 1, 1)
+
+ new_motion = CharacterMotion(new_character, joints_pos, joints_rot, foot_contacts)
+ # save the motion in our dict
+ session.motions[character_name] = new_motion
+
+ # put the character at the right frame
+ new_motion.set_frame(session.frame_idx)
+
+ # put them visible with a small delay
+ # so that the set_frame function has time to finish
+ def _set_visibility():
+ new_motion.character.set_skinned_mesh_visibility(session.gui_elements.gui_viz_skinned_mesh_checkbox.value)
+ new_motion.character.set_skeleton_visibility(session.gui_elements.gui_viz_skeleton_checkbox.value)
+
+ timer = threading.Timer(
+ 0.2, # 0.2s delay
+ _set_visibility,
+ )
+ timer.start()
+
+ def clear_motions(self, client_id: int) -> None:
+ if not self.client_active(client_id):
+ return
+ session = self.client_sessions[client_id]
+ for motion in list(session.motions.values()):
+ motion.clear()
+ session.motions.clear()
+
+ def compute_model_constraints_lst(
+ self,
+ session: ClientSession,
+ model_bundle: ModelBundle,
+ num_frames: int,
+ ):
+ return generation.compute_model_constraints_lst(session, model_bundle, num_frames, self.device)
+
+ def check_cuda_health(self) -> bool:
+ """Check if CUDA is still functional.
+
+ Trigger auto-restart if corrupted.
+ """
+ if self.device == "cpu":
+ return True
+ try:
+ torch.tensor([1.0], device=self.device) + torch.tensor([1.0], device=self.device)
+ return True
+ except RuntimeError as e:
+ if "device-side assert" in str(e) or "CUDA error" in str(e):
+ if self._cuda_healthy:
+ self._cuda_healthy = False
+ print("FATAL: CUDA context is corrupted (device-side assert). " "The process must be restarted.")
+ self._trigger_restart()
+ return False
+ raise
+
+ def _trigger_restart(self) -> None:
+ """Exit the process so the HF Space (or systemd/Docker) can restart it."""
+ import sys
+
+ print("Initiating automatic restart due to unrecoverable CUDA error...")
+ sys.stdout.flush()
+ sys.stderr.flush()
+ os._exit(1)
+
+ def generate(
+ self,
+ client: viser.ClientHandle,
+ prompts: list[str],
+ num_frames: list[int],
+ num_samples: int,
+ seed: int,
+ diffusion_steps: int,
+ cfg_weight: Optional[list[float]] = None,
+ cfg_type: Optional[str] = None,
+ postprocess_parameters: Optional[dict] = None,
+ transitions_parameters: Optional[dict] = None,
+ real_robot_rotations: bool = False,
+ ) -> None:
+ if not self._cuda_healthy:
+ raise RuntimeError("CUDA is in a corrupted state. The space is restarting...")
+
+ locked = self._generation_lock.acquire(blocking=False)
+ if not locked:
+ waiting_notif = client.add_notification(
+ title="Waiting for GPU...",
+ body="Another generation is in progress. Yours will start automatically.",
+ loading=True,
+ with_close_button=False,
+ )
+ self._generation_lock.acquire()
+ waiting_notif.remove()
+
+ try:
+ session = self.client_sessions[client.client_id]
+ model_bundle = self.load_model(session.model_name)
+ generation.generate(
+ client=client,
+ session=session,
+ model_bundle=model_bundle,
+ prompts=prompts,
+ num_frames=num_frames,
+ num_samples=num_samples,
+ seed=seed,
+ diffusion_steps=diffusion_steps,
+ cfg_weight=cfg_weight,
+ cfg_type=cfg_type,
+ postprocess_parameters=postprocess_parameters,
+ transitions_parameters=transitions_parameters,
+ real_robot_rotations=real_robot_rotations,
+ device=self.device,
+ clear_motions=self.clear_motions,
+ add_character_motion=self.add_character_motion,
+ )
+ finally:
+ self._generation_lock.release()
+
+ def set_frame(self, client_id: int, frame_idx: int, update_timeline: bool = True):
+ if not self.client_active(client_id):
+ return
+
+ session = self.client_sessions[client_id]
+
+ session.frame_idx = frame_idx
+ if update_timeline:
+ session.client.timeline.set_current_frame(frame_idx)
+ for motion in list(session.motions.values()):
+ motion.set_frame(frame_idx)
+ self._apply_constraint_overlay_visibility(session)
+
+ def run(self) -> None:
+ update_counter = 0
+ cuda_check_interval = 300
+ while True:
+ last_update_time = time.time()
+ if self.models:
+ # the max playback speed is 2x the model fps (from gui_playback_speed_buttons)
+ playback_fps = max(bundle.model_fps for bundle in self.models.values()) * 2.0
+ else:
+ playback_fps = 60.0
+
+ # update each client session independently
+ # copy to a list first to avoid changing size if client disconnects
+ for client_id, session in list(self.client_sessions.items()):
+ update_interval = int(playback_fps / (session.playback_speed * session.model_fps))
+ new_frame_idx = session.frame_idx
+ if session.playing and update_counter % update_interval == 0:
+ if session.frame_idx >= session.max_frame_idx:
+ new_frame_idx = 0
+ else:
+ new_frame_idx = session.frame_idx + 1
+
+ # make sure the client is still active before updating the frame
+ if self.client_active(client_id):
+ self.set_frame(client_id, new_frame_idx)
+
+ if update_counter % cuda_check_interval == 0:
+ self.check_cuda_health()
+
+ time_remaining = max(0, 1.0 / playback_fps - (time.time() - last_update_time))
+ time.sleep(time_remaining)
+ update_counter += 1
+ update_counter %= playback_fps # wrap around to 0 every second
+
+ def configure_theme(
+ self,
+ client: viser.ClientHandle,
+ dark_mode: bool = False,
+ titlebar_dark_mode_checkbox_uuid: str | None = None,
+ ):
+ # Sync grid color with theme (light vs dark)
+ theme = DARK_THEME if dark_mode else LIGHT_THEME
+ grid_handle = self.grid_handles.get(client.client_id)
+ if grid_handle is not None:
+ grid_handle.section_color = theme["grid"]
+
+ #
+ # setup theme
+ #
+ buttons = (
+ TitlebarButton(
+ text="Documentation",
+ icon="Description",
+ href="https://research.nvidia.com/labs/sil/projects/kimodo/docs/interactive_demo/index.html",
+ ),
+ TitlebarButton(
+ text="Project Page",
+ icon=None,
+ href="https://research.nvidia.com/labs/sil/projects/kimodo/",
+ ),
+ TitlebarButton(
+ text="Github",
+ icon="GitHub",
+ href="https://github.com/nv-tlabs/kimodo",
+ ),
+ )
+ assets_dir = DEMO_ASSETS_ROOT
+ logo_light_path = assets_dir / "nvidia_logo.png"
+ logo_dark_path = assets_dir / "nvidia_logo_dark.png"
+ if logo_light_path.exists():
+ light_b64 = base64.standard_b64encode(logo_light_path.read_bytes()).decode("ascii")
+ dark_b64 = (
+ base64.standard_b64encode(logo_dark_path.read_bytes()).decode("ascii")
+ if logo_dark_path.exists()
+ else None
+ )
+ image = TitlebarImage(
+ image_url_light=f"data:image/png;base64,{light_b64}",
+ image_url_dark=(f"data:image/png;base64,{dark_b64}" if dark_b64 else None),
+ image_alt="NVIDIA",
+ href="https://www.nvidia.com/",
+ )
+ else:
+ image = None
+ titlebar_theme = TitlebarConfig(buttons=buttons, image=image, title_text="Kimodo")
+ client.gui.set_panel_label("Kimodo")
+ client.gui.configure_theme(
+ titlebar_content=titlebar_theme,
+ control_layout="floating", # "floating", # ['floating', 'collapsible', 'fixed']
+ control_width="large", # ['small', 'medium', 'large']
+ dark_mode=dark_mode,
+ show_logo=False, # hide viser logo on bottom left corner
+ show_share_button=False,
+ titlebar_dark_mode_checkbox_uuid=titlebar_dark_mode_checkbox_uuid,
+ brand_color=(152, 189, 255), # (60, 131, 0), # (R, G, B) tuple
+ )
diff --git a/kimodo/demo/config.py b/kimodo/demo/config.py
new file mode 100644
index 0000000000000000000000000000000000000000..45fd9f453cb796d957f8c0bb6ce1570823796688
--- /dev/null
+++ b/kimodo/demo/config.py
@@ -0,0 +1,163 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+import os
+
+from kimodo.assets import DEMO_EXAMPLES_ROOT
+from kimodo.model.registry import (
+ AVAILABLE_MODELS,
+ DEFAULT_MODEL,
+ FRIENDLY_NAMES,
+ get_datasets,
+ get_model_info,
+ get_models_for_dataset_skeleton,
+ get_short_key_from_display_name,
+ get_skeleton_display_name,
+ get_skeleton_display_names_for_dataset,
+ get_skeleton_key_from_display_name,
+ get_skeletons_for_dataset,
+ get_versions_for_dataset_skeleton,
+ resolve_to_short_key,
+)
+
+SERVER_NAME = os.environ.get("SERVER_NAME", "0.0.0.0")
+SERVER_PORT = int(os.environ.get("SERVER_PORT", "7860"))
+HF_MODE = os.environ.get("HF_MODE", False)
+
+# HF mode: user queue and session limit (override via env in Spaces)
+MAX_ACTIVE_USERS = int(os.environ.get("MAX_ACTIVE_USERS", "5"))
+MAX_SESSION_MINUTES = float(os.environ.get("MAX_SESSION_MINUTES", "5.0"))
+
+DEFAULT_PLAYBACK_SPEED = 1.0
+# default start duration is 6.0 sec, but model can handle up to 10 sec
+DEFAULT_CUR_DURATION = 6.0
+DEFAULT_PROMPT = "A person walks forward."
+MIN_DURATION = 2.0
+MAX_DURATION = 10.0
+
+SHOW_TRANSITION_PARAMS = True
+INIT_POSTPROCESSING = True
+NB_TRANSITION_FRAMES = 5
+
+LIGHT_THEME = dict(
+ floor=(220, 220, 220),
+ grid=(180, 180, 180),
+)
+
+# Dark theme: slightly lighter grid and floor for better visibility and less flat black
+DARK_THEME = dict(
+ floor=(48, 48, 52),
+ grid=(105, 105, 110),
+)
+
+EXAMPLES_ROOT_DIR = str(DEMO_EXAMPLES_ROOT)
+
+# Model list and paths from kimodo registry (all models: Kimodo + TMR)
+MODEL_NAMES = tuple(AVAILABLE_MODELS)
+MODEL_EXAMPLES_DIRS = {name: os.path.join(EXAMPLES_ROOT_DIR, name) for name in MODEL_NAMES}
+# Display labels for backward compatibility (short_key -> display name)
+MODEL_LABELS = {name: FRIENDLY_NAMES.get(name, f"Model ({name})") for name in MODEL_NAMES}
+MODEL_LABEL_TO_NAME = {label: name for name, label in MODEL_LABELS.items()}
+
+# -----------------------------------------------------------------------------
+# Demo UI copy
+# -----------------------------------------------------------------------------
+
+DEMO_UI_QUICK_START_CORE_MD = """
+### Camera
+- **Left-drag**: rotate
+- **Right-drag**: pan
+- **Scroll**: zoom
+
+### Playback
+- **Space** to play/pause
+- **←/→** to step frames, or click the frame number.
+- **Scroll up/down** in the timeline: move left/right
+- **Shift + scroll** in the timeline: zoom in/out
+
+### Prompts
+- **Double-click** a text prompt to edit it.
+- **Click and drag** the right edge of a prompt box to extend/shorten it.
+- **Click empty space** to add a prompt.
+- **Right-click** a prompt to delete it.
+
+### Generate
+- Go to the **Generate** tab to modify options
+- It is also possible to **load** examples
+- Click **Generate** to generate a motion
+
+### Constraints
+- This is **optional**: should be use after a first generation
+- **Click** in the timeline tracks (Full-Body / 2D root etc) to add a constraint.
+- **Right-click** on a constraint to delete it.
+- To **edit** a constraint:
+ - Move playback to the target frame
+ - Click **Enter Editing Mode** in the Constraints tab.
+"""
+
+DEMO_UI_QUICK_START_MODAL_MD = (
+ DEMO_UI_QUICK_START_CORE_MD
+ + """
+
+See the **Instructions** tab for the full user manual.
+"""
+)
+
+DEMO_UI_INSTRUCTIONS_TAB_MD = (
+ """
+## How to Use This Demo
+
+"""
+ + DEMO_UI_QUICK_START_CORE_MD
+ + """
+
+---
+
+### Generating Motion (step-by-step)
+
+1. **Edit the text prompts** in the timeline (e.g., "A person walks forward.")
+2. **Modify the duration** by moving the right edge of each prompts (2–10 seconds)
+3. **Add constraints** (optional) to control the motion:
+ - Click **Enter Editing Mode** to adjust the character pose
+ - Use the timeline to place keyframes or intervals in constraint tracks (see below)
+4. **Click Generate** to create the motion
+5. If generating multiple samples, **click on a mesh** to select which one to keep
+
+### Timeline Editing
+
+**Adding Constraints:**
+1. Click anywhere on the timeline to add a keyframe at that frame. The keyframe is created based on the current character motion.
+2. Ctrl/Cmd+click+drag to add an interval constraint, or expand a keyframe into an interval
+3. Enter editing mode with the **Enter Editing Mode** button to adjust character pose before/after adding constraints.
+
+**Constraint Types:**
+- **Full-Body**: constrains the entire character pose
+- **2D Root**: constrains the character's path on the ground plane
+ - Enable **Densify** to create a continuous path
+- **End-Effectors**: constrains hands and feet positions
+ - Use separate tracks for Left/Right Hand/Foot
+
+
+**Moving & Deleting:**
+- **Drag keyframes/intervals** to move them to different frames
+- **Right-click** a keyframe or interval to delete it
+- Use **Clear All Constraints** to remove everything
+
+**Tips:**
+- The posing skeleton becomes visible in editing mode for precise positioning
+- Use **Snap to constraint** to align the current frame to a constraint
+
+### Saving & Loading
+
+You can save the current constraints or current motion to load in later from the Load/Save menu.
+Saving an **Example** will save the full constraints, motion, and generation metadata.
+
+### Visualization Options
+
+Switch to the **Visualize** tab to:
+- Toggle mesh and skeleton visibility
+- Adjust mesh opacity
+- Show/hide foot contact indicators
+- Switch between light and dark modes
+"""
+)
diff --git a/kimodo/demo/embedding_cache.py b/kimodo/demo/embedding_cache.py
new file mode 100644
index 0000000000000000000000000000000000000000..28771675baddb323821e7d941a248634bbc9bf4f
--- /dev/null
+++ b/kimodo/demo/embedding_cache.py
@@ -0,0 +1,253 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+import contextlib
+import contextvars
+import hashlib
+import json
+import os
+import threading
+import time
+from collections import OrderedDict
+from dataclasses import dataclass
+from typing import Iterable, Optional
+
+import numpy as np
+import torch
+
+from kimodo.sanitize import sanitize_texts
+
+_ACTIVE_SESSION = contextvars.ContextVar("kimodo_demo_active_session", default=None)
+
+
+@dataclass
+class CacheStats:
+ hits: int = 0
+ misses: int = 0
+ disk_hits: int = 0
+
+
+class EmbeddingCache:
+ """Disk-backed text embedding cache with a small in-memory LRU."""
+
+ def __init__(
+ self,
+ *,
+ model_name: str,
+ encoder_id: str,
+ base_dir: Optional[str] = None,
+ max_mem_entries: int = 128,
+ ) -> None:
+ cache_root = base_dir or os.environ.get(
+ "kimodo_EMBED_CACHE_DIR",
+ os.path.join("~", ".cache", "kimodo_demo", "embeddings"),
+ )
+ self.base_dir = os.path.expanduser(cache_root)
+ self.model_name = model_name
+ self.encoder_id = encoder_id
+ self.max_mem_entries = max_mem_entries
+ self.stats = CacheStats()
+
+ self._lock = threading.Lock()
+ self._mem_cache: OrderedDict[str, np.ndarray] = OrderedDict()
+ self._index = {}
+ self._index_loaded = False
+
+ def _model_dir(self) -> str:
+ return os.path.join(self.base_dir, self.model_name)
+
+ def _index_path(self) -> str:
+ return os.path.join(self._model_dir(), "index.json")
+
+ def _prewarm_marker_path(self, key: str) -> str:
+ return os.path.join(self._model_dir(), f"prewarm_{key}.json")
+
+ def has_prewarm_marker(self, key: str) -> bool:
+ return os.path.exists(self._prewarm_marker_path(key))
+
+ def write_prewarm_marker(self, key: str, *, prompt_count: int) -> None:
+ os.makedirs(self._model_dir(), exist_ok=True)
+ payload = {"prompt_count": prompt_count, "updated_at": time.time()}
+ tmp_path = f"{self._prewarm_marker_path(key)}.tmp"
+ with open(tmp_path, "w", encoding="utf-8") as f:
+ json.dump(payload, f)
+ os.replace(tmp_path, self._prewarm_marker_path(key))
+
+ def _load_index(self) -> None:
+ if self._index_loaded:
+ return
+ index_path = self._index_path()
+ if os.path.exists(index_path):
+ try:
+ with open(index_path, "r", encoding="utf-8") as f:
+ self._index = json.load(f)
+ except json.JSONDecodeError:
+ self._index = {}
+ self._index_loaded = True
+
+ def _save_index(self) -> None:
+ os.makedirs(self._model_dir(), exist_ok=True)
+ tmp_path = f"{self._index_path()}.tmp"
+ with open(tmp_path, "w", encoding="utf-8") as f:
+ json.dump(self._index, f)
+ os.replace(tmp_path, self._index_path())
+
+ def _make_key(self, text: str) -> str:
+ key_src = f"{self.model_name}|{self.encoder_id}|{text}"
+ return hashlib.sha256(key_src.encode("utf-8")).hexdigest()
+
+ def _entry_path(self, key: str) -> str:
+ return os.path.join(self._model_dir(), f"{key}.npy")
+
+ def _mem_get(self, key: str) -> Optional[np.ndarray]:
+ if key in self._mem_cache:
+ self._mem_cache.move_to_end(key)
+ return self._mem_cache[key]
+ return None
+
+ def _mem_put(self, key: str, value: np.ndarray) -> None:
+ self._mem_cache[key] = value
+ self._mem_cache.move_to_end(key)
+ while len(self._mem_cache) > self.max_mem_entries:
+ self._mem_cache.popitem(last=False)
+
+ def _disk_load(self, key: str) -> Optional[np.ndarray]:
+ path = self._entry_path(key)
+ if not os.path.exists(path):
+ return None
+ try:
+ return np.load(path)
+ except Exception:
+ return None
+
+ def _disk_save(self, key: str, value: np.ndarray) -> None:
+ os.makedirs(self._model_dir(), exist_ok=True)
+ np.save(self._entry_path(key), value)
+ self._index[key] = {
+ "length": int(value.shape[0]),
+ "dtype": str(value.dtype),
+ "updated_at": time.time(),
+ }
+
+ def _maybe_use_session_cache(self, texts: list[str]):
+ session = _ACTIVE_SESSION.get()
+ if session is None:
+ return None
+ if session.last_prompt_texts == texts and session.last_prompt_embeddings is not None:
+ return session.last_prompt_embeddings, session.last_prompt_lengths
+ return None
+
+ def _update_session_cache(self, texts: list[str], tensor: torch.Tensor, lengths: list[int]) -> None:
+ session = _ACTIVE_SESSION.get()
+ if session is None:
+ return
+ session.last_prompt_texts = texts
+ session.last_prompt_embeddings = tensor
+ session.last_prompt_lengths = lengths
+
+ def get_or_encode(self, texts: Iterable[str], encoder):
+ if isinstance(texts, str):
+ texts = [texts]
+ texts = sanitize_texts(list(texts))
+ if len(texts) == 0:
+ empty = torch.empty()
+ return empty, []
+
+ session_cache = self._maybe_use_session_cache(texts)
+ if session_cache is not None:
+ return session_cache
+
+ arrays: list[Optional[np.ndarray]] = [None] * len(texts)
+ lengths: list[int] = [0] * len(texts)
+ misses: list[tuple[int, str, str]] = []
+
+ with self._lock:
+ self._load_index()
+ for idx, text in enumerate(texts):
+ key = self._make_key(text)
+ cached = self._mem_get(key)
+ if cached is not None:
+ arrays[idx] = cached
+ lengths[idx] = cached.shape[0]
+ self.stats.hits += 1
+ continue
+
+ cached = self._disk_load(key)
+ if cached is not None:
+ arrays[idx] = cached
+ lengths[idx] = cached.shape[0]
+ self._mem_put(key, cached)
+ self.stats.disk_hits += 1
+ continue
+
+ misses.append((idx, text, key))
+ self.stats.misses += 1
+
+ if misses:
+ miss_texts = [text for _, text, _ in misses]
+ miss_tensor, miss_lengths = encoder(miss_texts)
+ miss_tensor = miss_tensor.detach().cpu()
+ miss_tensor_np = miss_tensor.numpy()
+
+ with self._lock:
+ self._load_index()
+ for miss_idx, length in enumerate(miss_lengths):
+ idx, _text, key = misses[miss_idx]
+ arr = miss_tensor_np[miss_idx, :length].copy()
+ arrays[idx] = arr
+ lengths[idx] = int(length)
+ self._mem_put(key, arr)
+ self._disk_save(key, arr)
+ self._save_index()
+
+ max_len = max(lengths) if lengths else 0
+ feat_dim = arrays[0].shape[-1] if arrays[0] is not None else 0
+ dtype = arrays[0].dtype if arrays[0] is not None else np.float32
+ padded = np.zeros((len(texts), max_len, feat_dim), dtype=dtype)
+ for idx, arr in enumerate(arrays):
+ if arr is None:
+ continue
+ padded[idx, : arr.shape[0]] = arr
+
+ result = torch.from_numpy(padded)
+ self._update_session_cache(texts, result, lengths)
+ return result, lengths
+
+
+class CachedTextEncoder:
+ """Wrapper around a text encoder to add disk-backed caching."""
+
+ def __init__(self, encoder, *, model_name: str, base_dir: Optional[str] = None):
+ self.encoder = encoder
+ self.model_name = model_name
+ encoder_id = f"{type(encoder).__name__}"
+ self.cache = EmbeddingCache(model_name=model_name, encoder_id=encoder_id, base_dir=base_dir)
+
+ def __call__(self, texts):
+ return self.cache.get_or_encode(texts, self.encoder)
+
+ def prewarm(self, texts) -> None:
+ if isinstance(texts, str):
+ texts = [texts]
+ texts = sanitize_texts(list(texts))
+ prewarm_key = hashlib.sha256("|".join(texts).encode("utf-8")).hexdigest()
+ if self.cache.has_prewarm_marker(prewarm_key):
+ return
+ self.cache.get_or_encode(texts, self.encoder)
+ self.cache.write_prewarm_marker(prewarm_key, prompt_count=len(texts))
+
+ def to(self, device=None, dtype=None):
+ if hasattr(self.encoder, "to"):
+ self.encoder.to(device=device, dtype=dtype)
+ return self
+
+ @contextlib.contextmanager
+ def session_context(self, session):
+ token = _ACTIVE_SESSION.set(session)
+ try:
+ yield
+ finally:
+ _ACTIVE_SESSION.reset(token)
+
+ def __getattr__(self, name):
+ return getattr(self.encoder, name)
diff --git a/kimodo/demo/generation.py b/kimodo/demo/generation.py
new file mode 100644
index 0000000000000000000000000000000000000000..00282e615c296ea11eba746d3e092899924d4df9
--- /dev/null
+++ b/kimodo/demo/generation.py
@@ -0,0 +1,218 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from collections import defaultdict
+from typing import Optional
+
+import numpy as np
+import torch
+
+import viser
+from kimodo.constraints import (
+ TYPE_TO_CLASS,
+ FullBodyConstraintSet,
+ Root2DConstraintSet,
+)
+from kimodo.exports.mujoco import apply_g1_real_robot_projection
+from kimodo.skeleton import G1Skeleton34, SOMASkeleton30
+from kimodo.tools import seed_everything
+
+from .embedding_cache import CachedTextEncoder
+from .state import ClientSession, ModelBundle
+
+
+def compute_model_constraints_lst(
+ session: ClientSession,
+ model_bundle: ModelBundle,
+ num_frames: int,
+ device: str,
+):
+ """Compute the lst of constraints for the model based on the constraints in viser."""
+ assert len(session.motions) == 1, "Only one motion allowed for constrained generation"
+ if not session.constraints:
+ return []
+
+ model_skeleton = model_bundle.model.skeleton
+ # For SOMA, UI uses somaskel77; extract 30-joint subset for the model
+ use_skel_slice = isinstance(model_skeleton, SOMASkeleton30) and session.skeleton.nbjoints != model_skeleton.nbjoints
+ skel_slice = model_skeleton.get_skel_slice(session.skeleton) if use_skel_slice else None
+
+ dense_smooth_root_pos_2d = None
+ if session.constraints["2D Root"].dense_path:
+ # get the full 2d root
+ dense_smooth_root_pos_2d = session.constraints["2D Root"].get_constraint_info(device=device)["root_pos"][
+ :, [0, 2]
+ ]
+
+ model_constraints = []
+ for track_name, constraint in session.constraints.items():
+ constraint_info = constraint.get_constraint_info(device=device)
+ frame_idx = constraint_info["frame_idx"]
+ # drop any constraints outside the generation range
+ valid_info = [(i, fi) for i, fi in enumerate(frame_idx) if fi < num_frames]
+ valid_idx = [i for i, _ in valid_info]
+ valid_frame_idx = [fi for _, fi in valid_info]
+
+ if len(valid_frame_idx) == 0:
+ continue
+
+ frame_indices = torch.tensor(valid_frame_idx)
+ if track_name == "2D Root":
+ smooth_root_pos_2d = constraint_info["root_pos"][valid_idx][:, [0, 2]].to(device)
+ # same as "smooth_root_2d"
+ model_constraints.append(
+ Root2DConstraintSet(
+ model_skeleton,
+ frame_indices,
+ smooth_root_pos_2d,
+ )
+ )
+ elif track_name == "Full-Body":
+ constraint_joints_pos = constraint_info["joints_pos"][valid_idx].to(device)
+ constraint_joints_rot = constraint_info["joints_rot"][valid_idx].to(device)
+ if skel_slice is not None:
+ constraint_joints_pos = constraint_joints_pos[:, skel_slice]
+ constraint_joints_rot = constraint_joints_rot[:, skel_slice]
+
+ smooth_root_pos_2d = None
+ if dense_smooth_root_pos_2d is not None:
+ smooth_root_pos_2d = dense_smooth_root_pos_2d[frame_indices]
+
+ model_constraints.append(
+ FullBodyConstraintSet(
+ model_skeleton,
+ frame_indices,
+ constraint_joints_pos,
+ constraint_joints_rot,
+ smooth_root_2d=smooth_root_pos_2d,
+ )
+ )
+ elif track_name == "End-Effectors":
+ constraint_joints_pos = constraint_info["joints_pos"][valid_idx].to(device)
+ constraint_joints_rot = constraint_info["joints_rot"][valid_idx].to(device)
+ if skel_slice is not None:
+ constraint_joints_pos = constraint_joints_pos[:, skel_slice]
+ constraint_joints_rot = constraint_joints_rot[:, skel_slice]
+
+ end_effector_type_set_lst = [
+ end_effector_type_set
+ for i, end_effector_type_set in enumerate(constraint_info["end_effector_type"])
+ if i in valid_idx
+ ]
+
+ # regroup the end effector data by type
+ cls_idx = defaultdict(list)
+ for idx, end_effector_type_set in enumerate(end_effector_type_set_lst):
+ for end_effector_type in end_effector_type_set:
+ cls_idx[TYPE_TO_CLASS[end_effector_type]].append(idx)
+
+ for cls, lst_idx in cls_idx.items():
+ frame_indices_cls = frame_indices[lst_idx]
+ smooth_root_pos_2d = None
+ if dense_smooth_root_pos_2d is not None:
+ smooth_root_pos_2d = dense_smooth_root_pos_2d[frame_indices_cls]
+
+ constraint_joints_pos_el = constraint_joints_pos[lst_idx]
+ constraint_joints_rot_el = constraint_joints_rot[lst_idx]
+
+ model_constraints.append(
+ cls(
+ model_skeleton,
+ frame_indices_cls,
+ constraint_joints_pos_el,
+ constraint_joints_rot_el,
+ smooth_root_2d=smooth_root_pos_2d,
+ )
+ )
+ else:
+ raise ValueError(f"Unsupported constraint type: {constraint.display_name}")
+ return model_constraints
+
+
+def generate(
+ *,
+ client: viser.ClientHandle,
+ session: ClientSession,
+ model_bundle: ModelBundle,
+ prompts: list[str],
+ num_frames: list[int],
+ num_samples: int,
+ seed: int,
+ diffusion_steps: int,
+ cfg_weight: Optional[list[float]] = None,
+ cfg_type: Optional[str] = None,
+ postprocess_parameters: Optional[dict] = None,
+ transitions_parameters: Optional[dict] = None,
+ real_robot_rotations: bool = False,
+ device: str,
+ clear_motions,
+ add_character_motion,
+) -> None:
+ client_id = client.client_id
+ print(
+ f"Generating {num_samples} samples for a total of {sum(num_frames)} frames with those prompt: {prompts} (client {client_id})"
+ )
+
+ seed_everything(seed)
+
+ model_constraints = compute_model_constraints_lst(session, model_bundle, sum(num_frames), device)
+ cfg_weight = cfg_weight or [2.0, 2.0]
+ postprocess_parameters = postprocess_parameters or {}
+ transitions_parameters = transitions_parameters or {}
+
+ encoder = getattr(model_bundle.model, "text_encoder", None)
+ if isinstance(encoder, CachedTextEncoder):
+ with encoder.session_context(session):
+ pred_joints_output = model_bundle.model(
+ prompts,
+ num_frames,
+ diffusion_steps,
+ multi_prompt=True,
+ constraint_lst=model_constraints,
+ cfg_weight=cfg_weight,
+ num_samples=num_samples,
+ cfg_type=cfg_type,
+ **(postprocess_parameters | transitions_parameters),
+ ) # [B, T, motion_rep_dim]
+ else:
+ pred_joints_output = model_bundle.model(
+ prompts,
+ num_frames,
+ diffusion_steps,
+ multi_prompt=True,
+ constraint_lst=model_constraints,
+ cfg_weight=cfg_weight,
+ num_samples=num_samples,
+ cfg_type=cfg_type,
+ **(postprocess_parameters | transitions_parameters),
+ ) # [B, T, motion_rep_dim]
+
+ joints_pos = pred_joints_output["posed_joints"] # [B, T, J, 3]
+ joints_rot = pred_joints_output["global_rot_mats"]
+ foot_contacts = pred_joints_output.get("foot_contacts")
+
+ # Optionally project G1 to real robot DoF (1-DoF per joint, clamped) for display.
+ if real_robot_rotations and isinstance(session.skeleton, G1Skeleton34):
+ joints_pos, joints_rot = apply_g1_real_robot_projection(
+ session.skeleton,
+ pred_joints_output["posed_joints"],
+ pred_joints_output["global_rot_mats"],
+ clamp_to_limits=True,
+ )
+
+ # Display on characters (callbacks keep this module UI-agnostic).
+ clear_motions(client_id)
+ # Keep one sample centered at the origin so constraints align.
+ spread_factor = 1.0 # meters
+ center_idx = num_samples // 2
+ x_trans = (np.arange(num_samples) - center_idx) * spread_factor
+ for i in range(num_samples):
+ cur_joints_pos = joints_pos[i]
+ cur_joints_pos[..., 0] += x_trans[i]
+ add_character_motion(
+ client,
+ session.skeleton,
+ cur_joints_pos,
+ joints_rot[i],
+ foot_contacts[i],
+ )
diff --git a/kimodo/demo/queue_manager.py b/kimodo/demo/queue_manager.py
new file mode 100644
index 0000000000000000000000000000000000000000..f36d302e7eeb0e581c0333b5dd95c37080c8ea3b
--- /dev/null
+++ b/kimodo/demo/queue_manager.py
@@ -0,0 +1,336 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""HF mode user queue and session time limit."""
+
+import math
+import threading
+import time
+from collections.abc import Callable
+from typing import Any
+
+import viser
+
+from .config import DEMO_UI_QUICK_START_MODAL_MD, MAX_SESSION_MINUTES
+
+# Link for "Duplicate this Space" on Hugging Face (used in queue and expiry modals).
+DUPLICATE_SPACE_URL = "https://huggingface.co/spaces/nvidia/Kimodo?duplicate=true"
+GITHUB_REPO_URL = "https://github.com/nv-tlabs/kimodo"
+
+# How often to refresh queue modal content (position, total, estimated wait).
+QUEUE_MODAL_REFRESH_INTERVAL_SEC = 15
+
+
+class UserQueue:
+ """Thread-safe queue: active users (with activation timestamp) and waiting queue."""
+
+ def __init__(self, max_active: int, max_minutes: float) -> None:
+ self._max_active = max_active
+ self._max_minutes = max_minutes
+ self._max_seconds = max_minutes * 60.0
+ self._active: dict[int, float] = {} # client_id -> activation timestamp
+ self._queued: list[int] = []
+ self._lock = threading.Lock()
+
+ def try_activate(self, client_id: int) -> bool:
+ """If a slot is free, add client as active and return True.
+
+ Else return False.
+ """
+ with self._lock:
+ if len(self._active) < self._max_active:
+ self._active[client_id] = time.time()
+ return True
+ return False
+
+ def enqueue(self, client_id: int) -> None:
+ with self._lock:
+ if client_id not in self._queued:
+ self._queued.append(client_id)
+
+ def remove(self, client_id: int) -> bool:
+ """Remove from active or queue.
+
+ Returns True if was active.
+ """
+ with self._lock:
+ was_active = client_id in self._active
+ self._active.pop(client_id, None)
+ if client_id in self._queued:
+ self._queued.remove(client_id)
+ return was_active
+
+ def promote_next(self) -> int | None:
+ """If queue non-empty, pop first, activate them, return their client_id.
+
+ Else None.
+ """
+ with self._lock:
+ if not self._queued:
+ return None
+ client_id = self._queued.pop(0)
+ self._active[client_id] = time.time()
+ return client_id
+
+ def get_queue_position(self, client_id: int) -> tuple[int, int] | None:
+ """(1-based position, total_in_queue) or None if not queued."""
+ with self._lock:
+ if client_id not in self._queued:
+ return None
+ pos = self._queued.index(client_id)
+ return (pos + 1, len(self._queued))
+
+ def get_estimated_wait_seconds(self, client_id: int) -> float:
+ """Estimated seconds until this queued client gets a slot."""
+ with self._lock:
+ if client_id not in self._queued:
+ return 0.0
+ pos = self._queued.index(client_id) + 1 # 1-based
+ # Expiry times of active users (when they free a slot)
+ now = time.time()
+ expiries = sorted(now + self._max_seconds - (now - t) for t in self._active.values())
+ if not expiries:
+ return 0.0
+ # Nth slot to free (1-indexed) wraps over expiries
+ idx = (pos - 1) % len(expiries)
+ cycles = (pos - 1) // len(expiries)
+ slot_free_time = expiries[idx] + cycles * self._max_seconds
+ return max(0.0, slot_free_time - now)
+
+ def is_active(self, client_id: int) -> bool:
+ with self._lock:
+ return client_id in self._active
+
+ def was_active(self, client_id: int) -> bool:
+ """True if client is currently active (for use when already holding lock)."""
+ return client_id in self._active
+
+
+def _format_wait(seconds: float) -> str:
+ if seconds < 60:
+ return "less than a minute"
+ mins = int(math.ceil(seconds / 60))
+ return f"~{mins} minute{'s' if mins != 1 else ''}"
+
+
+def _queue_modal_markdown(position: int, total: int, estimated_wait_sec: float) -> str:
+ wait_str = _format_wait(estimated_wait_sec)
+ mins = int(MAX_SESSION_MINUTES) if MAX_SESSION_MINUTES == int(MAX_SESSION_MINUTES) else MAX_SESSION_MINUTES
+ return f"""## Kimodo Demo — Please Wait
+
+This demo runs with limited capacity.
+Each user gets **{mins} minute{"s" if mins != 1 else ""}** of interactive time.
+
+**Your position in queue:** {position} / {total}
+
+**Estimated wait:** {wait_str}
+
+Please keep this tab open — the demo will start automatically when it's your turn.
+
+---
+*Want unlimited access? [Duplicate this Space]({DUPLICATE_SPACE_URL}) or clone the [GitHub repo]({GITHUB_REPO_URL}) to run locally!*
+"""
+
+
+def _welcome_modal_markdown() -> str:
+ mins = int(MAX_SESSION_MINUTES) if MAX_SESSION_MINUTES == int(MAX_SESSION_MINUTES) else MAX_SESSION_MINUTES
+ return f"""## Welcome to Kimodo Demo
+
+You have been granted a **{mins}-minute** demo session.
+Your session timer has started.
+
+Click the button below to begin!
+"""
+
+
+def _expiry_modal_markdown() -> str:
+ mins = int(MAX_SESSION_MINUTES) if MAX_SESSION_MINUTES == int(MAX_SESSION_MINUTES) else MAX_SESSION_MINUTES
+ return f"""## Session Expired
+
+Your {mins}-minute demo session has ended.
+Thank you for trying Kimodo!
+
+Refresh this page to rejoin the queue, or [duplicate this Space]({DUPLICATE_SPACE_URL}) for unlimited access.
+"""
+
+
+class QueueManager:
+ """Orchestrates HF mode: queue modals, welcome modal, session timer, promotion."""
+
+ def __init__(
+ self,
+ queue: UserQueue,
+ server: viser.ViserServer,
+ setup_demo_for_client: Callable[[viser.ClientHandle], None],
+ cleanup_session: Callable[[int], None],
+ ) -> None:
+ self._queue = queue
+ self._server = server
+ self._setup_demo_for_client = setup_demo_for_client
+ self._cleanup_session = cleanup_session
+ self._max_seconds = queue._max_seconds
+
+ self._queue_modal_handles: dict[int, tuple[Any, Any]] = {}
+ self._welcome_modal_handles: dict[int, Any] = {}
+ self._expiry_timers: dict[int, threading.Timer] = {}
+ self._lock = threading.Lock()
+ self._refresh_stop = threading.Event()
+ self._refresh_thread = threading.Thread(
+ target=self._queue_modal_refresh_loop,
+ name="queue-modal-refresh",
+ daemon=True,
+ )
+ self._refresh_thread.start()
+
+ def _queue_modal_refresh_loop(self) -> None:
+ """Periodically refresh queue modals so position, total, and estimated wait stay current."""
+ while not self._refresh_stop.wait(timeout=QUEUE_MODAL_REFRESH_INTERVAL_SEC):
+ self._update_all_queue_modals()
+
+ def on_client_connect(self, client: viser.ClientHandle) -> None:
+ """Handle new connection: activate if slot free, else enqueue and show queue modal."""
+ client_id = client.client_id
+ if self._queue.try_activate(client_id):
+ try:
+ self._setup_demo_for_client(client)
+ except RuntimeError as e:
+ if "CUDA error" in str(e):
+ print(f"CUDA error while setting up client {client_id}: {e}")
+ return
+ raise
+ self._start_session_timer(client_id)
+ self._show_welcome_modal(client)
+ else:
+ self._queue.enqueue(client_id)
+ self._show_queue_modal(client)
+ self._update_all_queue_modals()
+
+ def on_client_disconnect(self, client_id: int) -> None:
+ """Remove from queue/active, cancel timer, promote next if was active.
+
+ Session/scene cleanup is done by the demo's on_client_disconnect.
+ """
+ with self._lock:
+ self._expiry_timers.pop(client_id, None)
+ self._queue_modal_handles.pop(client_id, None)
+ self._welcome_modal_handles.pop(client_id, None)
+ was_active = self._queue.remove(client_id)
+ if was_active:
+ self._promote_next_user()
+ else:
+ self._update_all_queue_modals()
+
+ def _show_queue_modal(self, client: viser.ClientHandle) -> None:
+ client_id = client.client_id
+ pos, total = self._queue.get_queue_position(client_id) or (0, 0)
+ wait_sec = self._queue.get_estimated_wait_seconds(client_id)
+ md_content = _queue_modal_markdown(pos, total, wait_sec)
+
+ modal = client.gui.add_modal(
+ "Kimodo Demo — Please Wait",
+ size="xl",
+ show_close_button=False,
+ )
+ with modal:
+ md_handle = client.gui.add_markdown(md_content)
+ with self._lock:
+ self._queue_modal_handles[client_id] = (modal, md_handle)
+
+ def _show_quick_start_modal(self, client: viser.ClientHandle) -> None:
+ """Show the quick start instructions modal (same as non-HF mode)."""
+ with client.gui.add_modal(
+ "Welcome — Quick Start",
+ size="xl",
+ show_close_button=True,
+ save_choice="kimodo.demo.quick_start_ack",
+ ) as quick_start_modal:
+ client.gui.add_markdown(DEMO_UI_QUICK_START_MODAL_MD)
+ client.gui.add_button("Got it (don't remind me again)").on_click(lambda _: quick_start_modal.close())
+
+ def _show_welcome_modal(self, client: viser.ClientHandle) -> None:
+ client_id = client.client_id
+
+ def _on_start_demo(_: Any) -> None:
+ modal.close()
+ self._show_quick_start_modal(client)
+
+ modal = client.gui.add_modal(
+ "Welcome to Kimodo Demo",
+ size="xl",
+ show_close_button=True,
+ )
+ with modal:
+ client.gui.add_markdown(_welcome_modal_markdown())
+ client.gui.add_button("Start Demo").on_click(_on_start_demo)
+ with self._lock:
+ self._welcome_modal_handles[client_id] = modal
+
+ def _update_all_queue_modals(self) -> None:
+ with self._lock:
+ handles = list(self._queue_modal_handles.items())
+ for client_id, (modal, md_handle) in handles:
+ pos_total = self._queue.get_queue_position(client_id)
+ if pos_total is None:
+ continue
+ pos, total = pos_total
+ wait_sec = self._queue.get_estimated_wait_seconds(client_id)
+ try:
+ md_handle.content = _queue_modal_markdown(pos, total, wait_sec)
+ except Exception:
+ pass
+
+ def _promote_next_user(self) -> None:
+ promoted_id = self._queue.promote_next()
+ if promoted_id is None:
+ return
+ clients = self._server.get_clients()
+ client = clients.get(promoted_id)
+ if client is None:
+ return
+ with self._lock:
+ old = self._queue_modal_handles.pop(promoted_id, None)
+ if old is not None:
+ try:
+ old[0].close()
+ except Exception:
+ pass
+ try:
+ self._setup_demo_for_client(client)
+ except RuntimeError as e:
+ if "CUDA error" in str(e):
+ print(f"CUDA error while setting up client {promoted_id}: {e}")
+ return
+ raise
+ self._start_session_timer(promoted_id)
+ self._show_welcome_modal(client)
+ self._update_all_queue_modals()
+
+ def _start_session_timer(self, client_id: int) -> None:
+ def on_expiry() -> None:
+ self._on_session_expired(client_id)
+
+ t = threading.Timer(self._max_seconds, on_expiry)
+ t.daemon = True
+ with self._lock:
+ self._expiry_timers[client_id] = t
+ t.start()
+
+ def _on_session_expired(self, client_id: int) -> None:
+ with self._lock:
+ self._expiry_timers.pop(client_id, None)
+ if not self._queue.is_active(client_id):
+ return
+ self._queue.remove(client_id)
+ clients = self._server.get_clients()
+ client = clients.get(client_id)
+ if client is not None:
+ try:
+ with client.gui.add_modal(
+ "Session Expired",
+ size="lg",
+ show_close_button=False,
+ ) as modal_ctx:
+ client.gui.add_markdown(_expiry_modal_markdown())
+ except Exception:
+ pass
+ self._cleanup_session(client_id)
+ self._promote_next_user()
diff --git a/kimodo/demo/state.py b/kimodo/demo/state.py
new file mode 100644
index 0000000000000000000000000000000000000000..e158b99c4276a7c453a7c4ca69f768513b4ed83f
--- /dev/null
+++ b/kimodo/demo/state.py
@@ -0,0 +1,59 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from dataclasses import dataclass, field
+from typing import Optional
+
+import torch
+
+import kimodo.viz.viser_utils as viser_utils
+import viser
+from kimodo.skeleton import SkeletonBase
+from kimodo.viz.viser_utils import GuiElements
+
+from .config import (
+ DEFAULT_CUR_DURATION,
+ DEFAULT_MODEL,
+ DEFAULT_PLAYBACK_SPEED,
+)
+
+
+@dataclass(frozen=True)
+class ModelBundle:
+ model: object
+ motion_rep: object
+ skeleton: SkeletonBase
+ model_fps: float
+
+
+@dataclass
+class ClientSession:
+ """Per-client session data."""
+
+ client: viser.ClientHandle
+ gui_elements: GuiElements
+ motions: dict # character_name -> CharacterMotion
+ constraints: dict[str, viser_utils.ConstraintSet] = field(default_factory=dict)
+ timeline_data: object = None
+ frame_idx: int = 0
+ playing: bool = False
+ playback_speed: float = DEFAULT_PLAYBACK_SPEED
+ cur_duration: float = DEFAULT_CUR_DURATION
+ max_frame_idx: int = 100 # will be updated based on model_fps
+ updating_motions: bool = False
+ edit_mode: bool = False
+ model_name: str = DEFAULT_MODEL
+ model_fps: float = 0.0
+ skeleton: SkeletonBase | None = None
+ motion_rep: object | None = None
+ examples_base_dir: str = ""
+ example_dict: dict[str, str] = field(default_factory=dict)
+ gui_examples_dropdown: Optional[viser.GuiInputHandle] = None
+ gui_save_example_path_text: Optional[viser.GuiInputHandle] = None
+ gui_model_selector: Optional[viser.GuiInputHandle] = None
+ last_prompt_texts: Optional[list[str]] = None
+ last_prompt_embeddings: Optional[torch.Tensor] = None
+ last_prompt_lengths: Optional[list[int]] = None
+ edit_mode_snapshot: Optional[dict[int, dict[str, object]]] = None
+ undo_drag_snapshot: Optional[dict[str, object]] = None
+ show_only_current_constraint: bool = False # False = Show All, True = Show only Current
diff --git a/kimodo/demo/ui.py b/kimodo/demo/ui.py
new file mode 100644
index 0000000000000000000000000000000000000000..7d9f5e32c48fb7b5a2daa8bea71e2281147ecb51
--- /dev/null
+++ b/kimodo/demo/ui.py
@@ -0,0 +1,3255 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+# ruff: noqa: I001
+import math
+import os
+import threading
+from typing import Optional
+
+from kimodo.constraints import load_constraints_lst, save_constraints_lst
+from kimodo.exports.bvh import motion_to_bvh_bytes, save_motion_bvh
+from kimodo.exports.motion_io import (
+ amass_npz_to_bytes,
+ g1_csv_to_bytes,
+ kimodo_npz_to_bytes,
+ load_motion_file,
+ save_kimodo_npz,
+)
+from kimodo.model.registry import kimodo_short_key_for_skeleton_dataset, registry_skeleton_for_joint_count
+from kimodo.tools import to_torch
+from kimodo.viz import viser_utils
+from kimodo.viz.viser_utils import GuiElements
+import numpy as np
+import torch
+import viser
+from viser._timeline_api import PROMPT_COLORS
+
+from . import generation
+from .config import (
+ DEFAULT_CUR_DURATION,
+ DEMO_UI_INSTRUCTIONS_TAB_MD,
+ get_datasets,
+ get_model_info,
+ get_models_for_dataset_skeleton,
+ get_skeleton_display_name,
+ get_skeleton_display_names_for_dataset,
+ get_skeleton_key_from_display_name,
+ get_short_key_from_display_name,
+ HF_MODE,
+ INIT_POSTPROCESSING,
+ MODEL_NAMES,
+ NB_TRANSITION_FRAMES,
+ SHOW_TRANSITION_PARAMS,
+)
+from .state import ClientSession
+from kimodo.skeleton import G1Skeleton34, SOMASkeleton30, SOMASkeleton77
+
+
+def extract_intervals_and_singles(t: torch.Tensor):
+ intervals = []
+ intervals_indices = []
+ single_frames = []
+ single_frames_indices = []
+
+ start_idx = 0
+
+ for i in range(1, len(t) + 1):
+ # End of run if:
+ # - end of tensor
+ # - non-consecutive value
+ if i == len(t) or t[i] != t[i - 1] + 1:
+ run_length = i - start_idx
+
+ if run_length >= 2:
+ intervals.append((int(t[start_idx]), int(t[i - 1])))
+ intervals_indices.append((start_idx, i - 1))
+ else:
+ single_frames.append(int(t[start_idx]))
+ single_frames_indices.append(start_idx)
+
+ start_idx = i
+
+ return intervals, intervals_indices, single_frames, single_frames_indices
+
+
+def create_gui(
+ demo,
+ client: viser.ClientHandle,
+ model_name: str,
+ model_fps: float,
+):
+ """Create GUI elements for a specific client."""
+ client_id = client.client_id
+
+ def get_active_session(event_client: viser.ClientHandle | None):
+ if event_client is None:
+ return None
+ if not demo.client_active(event_client.client_id):
+ return None
+ return demo.client_sessions[event_client.client_id]
+
+ def build_timeline_tracks():
+ timeline = client.timeline
+ demo.set_timeline_defaults(timeline, model_fps)
+ timeline.set_visible(True)
+ timeline.set_current_frame(0)
+
+ timeline_tracks = {}
+ fullbody_id = timeline.add_track(
+ "Full-Body",
+ track_type="keyframe",
+ color=(219, 148, 86),
+ height_scale=0.5,
+ )
+ timeline_tracks[fullbody_id] = {
+ "name": "Full-Body",
+ "track_type": "keyframe",
+ "color": (219, 148, 86),
+ "height_scale": 0.5,
+ }
+
+ root2d_id = timeline.add_track(
+ "2D Root",
+ track_type="keyframe",
+ color=(150, 100, 200),
+ height_scale=0.5,
+ )
+ timeline_tracks[root2d_id] = {
+ "name": "2D Root",
+ "track_type": "keyframe",
+ "color": (150, 100, 200),
+ "height_scale": 0.5,
+ }
+ lefthand_id = timeline.add_track(
+ "Left Hand",
+ track_type="keyframe",
+ color=(100, 200, 150),
+ height_scale=0.5,
+ )
+ timeline_tracks[lefthand_id] = {
+ "name": "Left Hand",
+ "track_type": "keyframe",
+ "color": (100, 200, 150),
+ "height_scale": 0.5,
+ }
+ righthand_id = timeline.add_track(
+ "Right Hand",
+ track_type="keyframe",
+ color=(200, 100, 150),
+ height_scale=0.5,
+ )
+ timeline_tracks[righthand_id] = {
+ "name": "Right Hand",
+ "track_type": "keyframe",
+ "color": (200, 100, 150),
+ "height_scale": 0.5,
+ }
+ leftfoot_id = timeline.add_track(
+ "Left Foot",
+ track_type="keyframe",
+ color=(219, 148, 86),
+ height_scale=0.5,
+ )
+ timeline_tracks[leftfoot_id] = {
+ "name": "Left Foot",
+ "track_type": "keyframe",
+ "color": (219, 148, 86),
+ "height_scale": 0.5,
+ }
+ rightfoot_id = timeline.add_track(
+ "Right Foot",
+ track_type="keyframe",
+ color=(150, 100, 200),
+ height_scale=0.5,
+ )
+ timeline_tracks[rightfoot_id] = {
+ "name": "Right Foot",
+ "track_type": "keyframe",
+ "color": (150, 100, 200),
+ "height_scale": 0.5,
+ }
+ return timeline, timeline_tracks
+
+ timeline, timeline_tracks = build_timeline_tracks()
+ # These handles are part of GuiElements, but the demo currently uses timeline + buttons
+ # embedded in the Viser UI instead of custom controls.
+ gui_play_pause_button = None
+ gui_next_frame_button = None
+ gui_prev_frame_button = None
+ gui_timeline = None
+ gui_duration_slider = None
+
+ # now other gui elements
+ tab_group = client.gui.add_tab_group()
+
+ #
+ # Playback and Motion generation controls
+ #
+ with tab_group.add_tab("Generate", viser.Icon.WALK):
+ with client.gui.add_folder("Model Selection", expand_by_default=True):
+ info = get_model_info(model_name)
+ if info is None:
+ info = get_model_info(next(iter(MODEL_NAMES)))
+
+ def get_allowed_skeleton_labels(dataset_ui_label: str) -> list[str]:
+ labels = get_skeleton_display_names_for_dataset(dataset_ui_label, family="Kimodo")
+ if HF_MODE:
+ labels = [label for label in labels if get_skeleton_key_from_display_name(label) != "SMPLX"]
+ return labels
+
+ dataset_ui_label = "Rigplay" if HF_MODE else info.dataset_ui_label
+ datasets = ["Rigplay"] if HF_MODE else get_datasets(family="Kimodo")
+ skeleton_labels = get_allowed_skeleton_labels(dataset_ui_label)
+ initial_skeleton_label = get_skeleton_display_name(info.skeleton)
+ if initial_skeleton_label not in skeleton_labels and skeleton_labels:
+ initial_skeleton_label = skeleton_labels[0]
+ initial_skeleton_key = (
+ get_skeleton_key_from_display_name(initial_skeleton_label) if skeleton_labels else None
+ )
+ models_for_pair = (
+ get_models_for_dataset_skeleton(dataset_ui_label, initial_skeleton_key, family="Kimodo")
+ if initial_skeleton_key is not None
+ else []
+ )
+ version_options = [m.display_name for m in models_for_pair]
+ initial_version = (
+ info.display_name
+ if info.display_name in version_options
+ else (version_options[0] if version_options else "")
+ )
+ gui_dataset_selector = client.gui.add_dropdown(
+ "Training dataset",
+ options=datasets,
+ initial_value=dataset_ui_label,
+ visible=not HF_MODE,
+ )
+ gui_skeleton_selector = client.gui.add_dropdown(
+ "Model" if HF_MODE else "Skeleton",
+ options=skeleton_labels,
+ initial_value=initial_skeleton_label,
+ )
+ gui_version_selector = client.gui.add_dropdown(
+ "Version",
+ options=version_options,
+ initial_value=initial_version,
+ )
+ gui_version_selector.visible = len(models_for_pair) > 1
+ gui_model_display = client.gui.add_markdown(
+ content=f"**Model:** {initial_version}",
+ )
+ gui_load_model_button = client.gui.add_button(
+ "Load model",
+ hint="Load the selected model (dataset, skeleton, version).",
+ )
+
+ class ModelSelectorHandle:
+ """Wrapper so session and callbacks can treat three dropdowns as one."""
+
+ def __init__(self):
+ self._dataset = gui_dataset_selector
+ self._skeleton = gui_skeleton_selector
+ self._version = gui_version_selector
+ self._display = gui_model_display
+
+ @property
+ def value(self) -> str:
+ return get_short_key_from_display_name(self._version.value) or ""
+
+ def set_from_short_key(self, short_key: str) -> None:
+ info = get_model_info(short_key)
+ if info is None:
+ return
+ dataset_ui_label = "Rigplay" if HF_MODE else info.dataset_ui_label
+ self._dataset.value = dataset_ui_label
+ self._skeleton.options = get_allowed_skeleton_labels(dataset_ui_label)
+ skeleton_label = get_skeleton_display_name(info.skeleton)
+ if skeleton_label not in self._skeleton.options and self._skeleton.options:
+ skeleton_label = self._skeleton.options[0]
+ self._skeleton.value = skeleton_label
+ skeleton_key = get_skeleton_key_from_display_name(skeleton_label)
+ if skeleton_key is None:
+ return
+ models = get_models_for_dataset_skeleton(dataset_ui_label, skeleton_key, family="Kimodo")
+ self._version.options = [m.display_name for m in models]
+ self._version.value = (
+ info.display_name if info.display_name in self._version.options else self._version.options[0]
+ )
+ self._version.visible = len(models) > 1
+ self._display.content = f"**Model:** {self._version.value}"
+
+ gui_model_selector = ModelSelectorHandle()
+
+ with client.gui.add_folder("Examples", expand_by_default=True):
+ examples_base_dir = demo.get_examples_base_dir(model_name, absolute=True)
+ example_dict = viser_utils.load_example_cases(examples_base_dir)
+ example_names = list(example_dict.keys())
+ if not example_names:
+ example_names = [""]
+ gui_examples_dropdown = client.gui.add_dropdown(
+ "Example",
+ options=example_names,
+ initial_value=example_names[0],
+ )
+ gui_load_example_button = client.gui.add_button(
+ "Load Example",
+ hint="Load the selected example.",
+ disabled=not example_dict,
+ )
+
+ def update_examples_dropdown(
+ new_example_dict: dict[str, str],
+ keep_selection: bool = True,
+ ) -> None:
+ if not new_example_dict:
+ gui_examples_dropdown.options = [""]
+ gui_examples_dropdown.value = ""
+ gui_load_example_button.disabled = True
+ return
+ gui_load_example_button.disabled = False
+ example_names_local = list(new_example_dict.keys())
+ gui_examples_dropdown.options = example_names_local
+ if keep_selection and gui_examples_dropdown.value in example_names_local:
+ return
+ gui_examples_dropdown.value = example_names_local[0]
+
+ with client.gui.add_folder("Generate", expand_by_default=True):
+ gui_duration = client.gui.add_markdown(content=f"Total duration: {DEFAULT_CUR_DURATION:.1f} (sec)")
+
+ def update_duration_gui(duration):
+ gui_duration.content = f"Total duration: {duration:.1f} (sec)"
+
+ def compute_prompt_num_frames(prompt_values):
+ """Convert timeline prompt bounds to per-prompt frame counts.
+
+ Convention in this demo:
+ - All prompts except the last are treated as [start_frame, end_frame)
+ (end is exclusive).
+ - The last prompt is treated as [start_frame, end_frame] (end is inclusive).
+ - This assumes the prompts values are sorted by start_frame.
+ """
+ if len(prompt_values) == 0:
+ return []
+ num_frames = []
+ for i, x in enumerate(prompt_values):
+ cur = x.end_frame - x.start_frame
+ if i == len(prompt_values) - 1:
+ cur += 1
+ num_frames.append(cur)
+ return num_frames
+
+ def update_duration_auto():
+ session = demo.client_sessions[client_id]
+ prompt_values = sorted(
+ [x for x in timeline._prompts.values()],
+ key=lambda x: x.start_frame,
+ )
+ num_frames = compute_prompt_num_frames(prompt_values)
+ total_nb_frames = sum(num_frames)
+ cur_duration = total_nb_frames / session.model_fps
+ set_new_duration(client_id, cur_duration)
+ update_duration_gui(cur_duration)
+
+ gui_num_samples_slider = client.gui.add_slider(
+ "Num Samples",
+ min=1,
+ max=10,
+ step=1,
+ initial_value=1,
+ visible=not HF_MODE,
+ )
+
+ gui_use_soma_layer_checkbox = client.gui.add_checkbox(
+ "SOMA layer",
+ initial_value=False,
+ visible="soma" in (model_name or ""),
+ )
+
+ with client.gui.add_folder("Model Parameters", expand_by_default=False):
+ gui_seed = client.gui.add_number("Seed", initial_value=42)
+
+ with client.gui.add_folder("Diffusion", expand_by_default=False):
+ gui_diffusion_steps_slider = client.gui.add_slider(
+ "Denoising Steps",
+ min=2,
+ max=1000,
+ step=10,
+ initial_value=100,
+ )
+ with client.gui.add_folder("Classifier-Free Guidance", expand_by_default=False):
+ gui_cfg_checkbox = client.gui.add_checkbox(
+ "Enable",
+ initial_value=True,
+ visible=True,
+ )
+
+ gui_cfg_text_weight_slider = client.gui.add_slider(
+ "Text Weight",
+ min=0.0,
+ max=5.0,
+ step=0.1,
+ initial_value=2.0,
+ visible=True,
+ )
+ gui_cfg_constraint_weight_slider = client.gui.add_slider(
+ "Constraint Weight",
+ min=0.0,
+ max=5.0,
+ step=0.1,
+ initial_value=2.0,
+ visible=True,
+ )
+ with client.gui.add_folder(
+ "Transitions",
+ expand_by_default=False,
+ visible=SHOW_TRANSITION_PARAMS,
+ ):
+ gui_num_transition_frames_slider = client.gui.add_slider(
+ "Transition frames",
+ min=1,
+ max=10,
+ step=1,
+ initial_value=NB_TRANSITION_FRAMES,
+ visible=True,
+ )
+
+ with client.gui.add_folder("Post Processing", expand_by_default=False):
+ _model_name = model_name or ""
+ _postprocess_visible = "g1" not in _model_name
+ gui_postprocess_checkbox = client.gui.add_checkbox(
+ "Enable",
+ initial_value=INIT_POSTPROCESSING,
+ hint="Apply motion post-processing (not available for G1)",
+ visible=_postprocess_visible,
+ )
+ gui_root_margin = client.gui.add_number(
+ "Root Margin",
+ min=0.0,
+ # max=0.5,
+ step=0.01,
+ initial_value=0.04,
+ hint="Margin for root position (meters). Lower values pin root closer to target.",
+ visible=INIT_POSTPROCESSING and _postprocess_visible,
+ )
+
+ @gui_postprocess_checkbox.on_update
+ def _(event: viser.GuiEvent) -> None:
+ if get_active_session(event.client) is None:
+ return
+ # disable the slider if sharing transition is False
+ gui_root_margin.visible = gui_postprocess_checkbox.value
+
+ gui_real_robot_rotations_checkbox = client.gui.add_checkbox(
+ "Real robot rotations",
+ initial_value=False,
+ hint="Project joint rotations to G1 real robot DoF (1-DoF per joint) and clamp to axis limits from the MuJoCo XML.",
+ visible="g1" in _model_name,
+ )
+
+ gui_generate_button = client.gui.add_button("Generate", color="green")
+ with client.gui.add_folder("Constraints", expand_by_default=False):
+ gui_gizmo_space_dropdown = client.gui.add_dropdown(
+ "Gizmo space",
+ ("Local", "World"),
+ initial_value="Local",
+ visible="g1" not in _model_name,
+ )
+ gui_edit_constraint_button = client.gui.add_button("Enter Editing Mode")
+ gui_snap_to_constraint_button = client.gui.add_button(
+ "Snap to Constraint",
+ disabled=True,
+ )
+ gui_reset_constraint_button = client.gui.add_button(
+ "Reset Constraint",
+ disabled=True,
+ )
+ gui_undo_drag_button = client.gui.add_button(
+ "Undo Move",
+ disabled=True,
+ )
+
+ with client.gui.add_folder("Root 2D Options", expand_by_default=True):
+ gui_dense_path_checkbox = client.gui.add_checkbox(
+ "Make Smooth Path",
+ initial_value=False,
+ visible=True,
+ )
+
+ gui_show_only_current_constraint_checkbox = client.gui.add_checkbox(
+ "Show only Current",
+ initial_value=False,
+ hint="Show only constraint overlays at the current frame; uncheck to show all.",
+ )
+
+ def apply_constraint_overlay_visibility(session: ClientSession) -> None:
+ demo._apply_constraint_overlay_visibility(session)
+
+ @gui_show_only_current_constraint_checkbox.on_update
+ def _(event: viser.GuiEvent) -> None:
+ session = get_active_session(event.client)
+ if session is None:
+ return
+ session.show_only_current_constraint = gui_show_only_current_constraint_checkbox.value
+ apply_constraint_overlay_visibility(session)
+
+ gui_clear_all_constraints_button = client.gui.add_button(
+ "Clear All Constraints",
+ color="red",
+ )
+
+ def has_constraint_at_frame(session: ClientSession, frame_idx: int) -> bool:
+ for constraint_name in ["Full-Body", "End-Effectors", "2D Root"]:
+ constraint = session.constraints.get(constraint_name)
+ if constraint is None:
+ continue
+ if frame_idx in constraint.keyframes:
+ return True
+ return False
+
+ def update_snap_to_constraint_button(session: ClientSession) -> None:
+ gui_snap_to_constraint_button.disabled = not has_constraint_at_frame(session, session.frame_idx)
+
+ def ensure_edit_snapshot(session: ClientSession, motion, frame_idx: int) -> None:
+ if session.edit_mode_snapshot is None:
+ session.edit_mode_snapshot = {}
+ if frame_idx in session.edit_mode_snapshot:
+ return
+ session.edit_mode_snapshot[frame_idx] = {
+ "joints_pos": motion.get_joints_pos(frame_idx),
+ "joints_rot": motion.get_joints_rot(frame_idx),
+ }
+
+ def _update_dense_path(motion, session):
+ constraint_info = session.constraints["2D Root"].get_constraint_info()
+
+ if len(constraint_info["frame_idx"]) > 0:
+ min_root_frame = min(constraint_info["frame_idx"])
+ max_root_frame = max(constraint_info["frame_idx"])
+ motion.set_projected_root_pos_path(
+ constraint_info["root_pos"][:, [0, 2]],
+ min_frame_idx=min_root_frame,
+ max_frame_idx=max_root_frame,
+ )
+
+ # Delay (ms) after last keyframe/interval move before updating path = "on release".
+ DENSE_PATH_AFTER_RELEASE_MS = 300
+
+ def _schedule_dense_path_after_release(session):
+ """Schedule a single path update to run after user stops dragging."""
+ if "2D Root" not in session.constraints or not session.constraints["2D Root"].dense_path:
+ return
+ tdata = session.timeline_data
+ if tdata.get("dense_path_after_release_timer"):
+ tdata["dense_path_after_release_timer"].cancel()
+ delay = DENSE_PATH_AFTER_RELEASE_MS / 1000.0
+
+ def run():
+ if not demo.client_active(client_id):
+ return
+ sess = demo.client_sessions[client_id]
+ tdata["dense_path_after_release_timer"] = None
+ if "2D Root" not in sess.constraints or not sess.constraints["2D Root"].dense_path:
+ return
+ mot = list(sess.motions.values())[0]
+ _update_dense_path(mot, sess)
+
+ t = threading.Timer(delay, run)
+ tdata["dense_path_after_release_timer"] = t
+ t.start()
+
+ @gui_dense_path_checkbox.on_update
+ def _(event: viser.GuiEvent) -> None:
+ session = get_active_session(event.client)
+ if session is None:
+ return
+
+ if gui_dense_path_checkbox.value:
+ # Make sure 0 and max_frame_idx keyframes are added to the constraint
+ # since dense path should cover full duration for best model performance
+ root_2d_track = session.timeline_data["tracks_ids"]["2D Root"]
+
+ # add a locked keyframe at 0
+ start_keyframe_id = client.timeline.add_locked_keyframe( # noqa
+ root_2d_track,
+ 0,
+ opacity=0.0,
+ )
+ session.timeline_data["keyframes"][start_keyframe_id] = {
+ "frame": 0,
+ "track_id": root_2d_track,
+ "locked": True,
+ "opacity": 0.0,
+ "value": None,
+ }
+ add_constraint_callback(
+ start_keyframe_id,
+ "2D Root",
+ (0, 0),
+ verbose=False,
+ )
+
+ # add a locked keyframe at max_frame_idx
+ end_keyframe_id = client.timeline.add_locked_keyframe(
+ root_2d_track,
+ session.max_frame_idx,
+ opacity=0.0,
+ )
+ session.timeline_data["keyframes"][end_keyframe_id] = {
+ "frame": session.max_frame_idx,
+ "track_id": root_2d_track,
+ "locked": True,
+ "opacity": 0.0,
+ "value": None,
+ }
+ add_constraint_callback(
+ end_keyframe_id,
+ "2D Root",
+ (session.max_frame_idx, session.max_frame_idx),
+ verbose=False,
+ )
+
+ # add a locked interval only for visual purposes
+ locked_interval = client.timeline.add_locked_interval( # noqa
+ root_2d_track,
+ start_frame=0,
+ end_frame=session.max_frame_idx,
+ )
+ session.timeline_data["intervals"][locked_interval] = {
+ "track_id": root_2d_track,
+ "start_frame_idx": 0,
+ "end_frame_idx": session.max_frame_idx,
+ "locked": True,
+ "opacity": 0.3,
+ "value": None,
+ }
+
+ session.constraints["2D Root"].set_dense_path(gui_dense_path_checkbox.value)
+ if session.constraints["2D Root"].dense_path:
+ # update the character motion to reflect the full path
+ # will be full length by construction, no need to specify min/max frame idx
+ motion = list(session.motions.values())[0]
+ _update_dense_path(motion, session)
+
+ # remove locked interval and locked keyframes
+ if not gui_dense_path_checkbox.value:
+ # Get all locked keyframes
+ keyframes_to_remove = []
+ for uuid, keyframe in client.timeline._keyframes.items():
+ if keyframe.locked:
+ keyframes_to_remove.append(uuid)
+ _data = session.timeline_data["keyframes"][uuid]
+ remove_constraint_callback(
+ uuid,
+ constraint_type=session.timeline_data["tracks"][_data["track_id"]]["name"],
+ frame_range=(_data["frame"], _data["frame"]),
+ verbose=False,
+ )
+
+ intervals_to_remove = []
+ # remove all locked intervals
+ for uuid, interval in client.timeline._intervals.items():
+ if interval.locked:
+ intervals_to_remove.append(uuid)
+
+ # removing keyframes and intervals
+ for uuid in keyframes_to_remove:
+ client.timeline.remove_keyframe(uuid)
+
+ for uuid in intervals_to_remove:
+ client.timeline.remove_interval(uuid)
+
+ apply_constraint_overlay_visibility(session)
+
+ with client.gui.add_folder(
+ "Load/Save",
+ expand_by_default=False,
+ visible=not HF_MODE,
+ ):
+ with client.gui.add_folder("Motion", expand_by_default=False):
+ gui_save_motion_path_text = client.gui.add_text("Save Path", initial_value="output")
+ gui_save_motion_format_dropdown = client.gui.add_dropdown(
+ "Save Format",
+ options=(
+ ["NPZ", "CSV"]
+ if "g1" in model_name.lower()
+ else ["NPZ", "AMASS NPZ"]
+ if "smplx" in model_name.lower()
+ else ["NPZ", "BVH"]
+ ),
+ initial_value="NPZ",
+ )
+ gui_save_bvh_standard_tpose_checkbox = client.gui.add_checkbox(
+ "Standard T-pose",
+ initial_value=False,
+ hint="For BVH export, use the standard T-pose rest skeleton.",
+ visible=False,
+ )
+ gui_save_motion_button = client.gui.add_button(
+ "Save Motion",
+ hint="Save the current motion (format + path above)",
+ )
+ gui_load_motion_path_text = client.gui.add_text(
+ "Load Path",
+ initial_value="output.npz",
+ hint="SOMA .bvh, Kimodo or AMASS .npz, or G1 MuJoCo .csv",
+ )
+ gui_load_motion_button = client.gui.add_button(
+ "Load Motion",
+ hint="Load the selected motion",
+ )
+ with client.gui.add_folder("Constraints", expand_by_default=False):
+ gui_save_constraints_path_text = client.gui.add_text(
+ "Save Path", initial_value="output_constraints.json"
+ )
+ gui_save_constraints_button = client.gui.add_button("Save Constraints")
+ gui_load_constraints_path_text = client.gui.add_text(
+ "Load Path", initial_value="output_constraints.json"
+ )
+ gui_load_constraints_button = client.gui.add_button("Load Constraints")
+ with client.gui.add_folder("Example", expand_by_default=False):
+ gui_save_example_path_text = client.gui.add_text(
+ "Save Dir",
+ initial_value=os.path.join(
+ demo.get_examples_base_dir(model_name, absolute=True),
+ "custom_example_1",
+ ),
+ )
+ gui_save_example_button = client.gui.add_button("Save Example")
+ gui_load_example_path_text = client.gui.add_text(
+ "Load Dir",
+ initial_value=os.path.join(
+ demo.get_examples_base_dir(model_name, absolute=True),
+ "custom_example_1",
+ ),
+ )
+ gui_load_gt_checkbox = client.gui.add_checkbox(
+ "Load GT instead",
+ initial_value=False,
+ )
+ gui_load_example_from_path_button = client.gui.add_button("Load Example")
+
+ def _get_primary_motion(session: ClientSession):
+ return list(session.motions.values())[0]
+
+ def _motion_to_numpy_dict(motion) -> dict[str, np.ndarray]:
+ joints_pos = motion.joints_pos.detach().cpu().numpy()
+ joints_rot = motion.joints_rot.detach().cpu().numpy()
+ joints_local_rot = motion.joints_local_rot.detach().cpu().numpy()
+
+ if joints_pos.ndim != 3:
+ raise ValueError(f"Expected unbatched joints_pos with shape [T, J, 3], got {joints_pos.shape}")
+ if joints_rot.ndim != 4:
+ raise ValueError(f"Expected unbatched joints_rot with shape [T, J, 3, 3], got {joints_rot.shape}")
+ if joints_local_rot.ndim != 4:
+ raise ValueError(
+ "Expected unbatched joints_local_rot with shape " f"[T, J, 3, 3], got {joints_local_rot.shape}"
+ )
+
+ motion_data = {
+ "posed_joints": joints_pos,
+ "global_rot_mats": joints_rot,
+ "local_rot_mats": joints_local_rot,
+ "root_positions": joints_pos[:, motion.skeleton.root_idx, :],
+ }
+ if motion.foot_contacts is not None:
+ foot_contacts = motion.foot_contacts.detach().cpu().numpy()
+ if foot_contacts.ndim != 2:
+ raise ValueError(
+ f"Expected unbatched foot_contacts with shape [T, C], got {foot_contacts.shape}"
+ )
+ motion_data["foot_contacts"] = foot_contacts
+ return motion_data
+
+ def _coerce_save_path(raw_path: str, *, ext: str) -> str:
+ """Ensure the save path ends with the correct extension for the chosen format."""
+ name = (raw_path or "").strip()
+ if name == "":
+ return f"output{ext}"
+ known_exts = (".npz", ".bvh", ".csv")
+ if name.lower().endswith(known_exts):
+ return os.path.splitext(name)[0] + ext
+ if os.path.splitext(name)[1] == "":
+ return name + ext
+ return name
+
+ def save_motion(client, save_path, fmt):
+ session = demo.client_sessions[client.client_id]
+ motion = _get_primary_motion(session)
+ motion_data = _motion_to_numpy_dict(motion)
+
+ if fmt == "BVH":
+ save_path = _coerce_save_path(save_path, ext=".bvh")
+ save_motion_bvh(
+ save_path,
+ motion.joints_local_rot,
+ motion.joints_pos[:, session.skeleton.root_idx, :],
+ skeleton=session.skeleton,
+ fps=float(session.model_fps),
+ standard_tpose=bool(gui_save_bvh_standard_tpose_checkbox.value),
+ )
+ elif fmt == "CSV":
+ save_path = _coerce_save_path(save_path, ext=".csv")
+ data = g1_csv_to_bytes(motion_data, session.skeleton, demo.device)
+ with open(save_path, "wb") as f:
+ f.write(data)
+ elif fmt == "AMASS NPZ":
+ save_path = _coerce_save_path(save_path, ext=".npz")
+ data = amass_npz_to_bytes(motion_data, session.skeleton, session.model_fps)
+ with open(save_path, "wb") as f:
+ f.write(data)
+ else:
+ save_path = _coerce_save_path(save_path, ext=".npz")
+ save_kimodo_npz(save_path, motion_data)
+ return save_path
+
+ @gui_save_motion_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ if get_active_session(event_client) is None:
+ return
+
+ raw_path = gui_save_motion_path_text.value
+ fmt = str(gui_save_motion_format_dropdown.value).upper()
+ try:
+ saved_path = save_motion(event_client, raw_path, fmt)
+ event_client.add_notification(
+ title="Motion saved!",
+ body=f"Saved motion to {saved_path}",
+ auto_close_seconds=5.0,
+ color="green",
+ )
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ event_client.add_notification(
+ title="Failed to save motion!",
+ body=str(e),
+ auto_close_seconds=5.0,
+ color="red",
+ )
+
+ def load_motion(client, load_path):
+ session = demo.client_sessions[client.client_id]
+
+ fps_arg = session.model_fps if session.model_fps and session.model_fps > 0 else None
+ motion_dict, num_joints_motion = load_motion_file(load_path, target_fps=fps_arg)
+
+ target_skel = registry_skeleton_for_joint_count(num_joints_motion)
+ current_info = get_model_info(session.model_name)
+ current_skel = current_info.skeleton if current_info is not None else None
+
+ if current_skel != target_skel:
+ dataset = current_info.dataset if current_info is not None else "RP"
+ new_key = kimodo_short_key_for_skeleton_dataset(target_skel, dataset)
+ if new_key is None:
+ new_key = kimodo_short_key_for_skeleton_dataset(target_skel, "RP")
+ if new_key is None:
+ raise ValueError(
+ f"No Kimodo model found for skeleton {target_skel} (motion has J={num_joints_motion})."
+ )
+ if new_key != session.model_name:
+ gui_model_selector.set_from_short_key(new_key)
+ apply_model_selection(new_key)
+ _update_visibility_for_loaded_model(new_key)
+ client.add_notification(
+ title="Model switched",
+ body=f"Switched to {new_key} to match loaded motion (J={num_joints_motion}).",
+ auto_close_seconds=5.0,
+ color="blue",
+ )
+ session = demo.client_sessions[client.client_id]
+
+ joints_pos = motion_dict["posed_joints"].to(device=demo.device, dtype=torch.float32)
+ joints_rot = motion_dict["global_rot_mats"].to(device=demo.device, dtype=torch.float32)
+ foot_contacts = motion_dict.get("foot_contacts")
+ if foot_contacts is not None:
+ foot_contacts = foot_contacts.to(device=demo.device, dtype=torch.float32)
+
+ # Support both batched [B, T, J, 3] and unbatched [T, J, 3]; take first sample if batched
+ if joints_pos.ndim == 4:
+ joints_pos = joints_pos[0]
+ if joints_rot.ndim == 5:
+ joints_rot = joints_rot[0]
+ if foot_contacts is not None and foot_contacts.ndim == 3:
+ foot_contacts = foot_contacts[0]
+
+ # Motion must match the current model's skeleton after auto-switch
+ num_joints_loaded = joints_pos.shape[1]
+ num_joints_skeleton = session.skeleton.nbjoints
+ if num_joints_loaded != num_joints_skeleton:
+ # Backward compat: expand 30-joint SOMA motion to 77
+ if (
+ num_joints_loaded == 30
+ and num_joints_skeleton == 77
+ and isinstance(session.skeleton, SOMASkeleton77)
+ ):
+ from kimodo.skeleton import global_rots_to_local_rots
+
+ skel30 = SOMASkeleton30().to(demo.device)
+ if "local_rot_mats" in motion_dict:
+ local_rot_30 = motion_dict["local_rot_mats"].to(device=demo.device, dtype=torch.float32)
+ if local_rot_30.ndim == 4:
+ local_rot_30 = local_rot_30[0]
+ else:
+ local_rot_30 = global_rots_to_local_rots(joints_rot, skel30)
+ local_rot_77 = skel30.to_SOMASkeleton77(local_rot_30)
+ root_positions = joints_pos[:, skel30.root_idx, :]
+ joints_rot, joints_pos, _ = session.skeleton.fk(local_rot_77, root_positions)
+
+ if foot_contacts is not None and foot_contacts.shape[-1] == 4:
+ foot_contacts = torch.cat(
+ [
+ foot_contacts[..., :2],
+ foot_contacts[..., 1:2],
+ foot_contacts[..., 2:4],
+ foot_contacts[..., 3:4],
+ ],
+ dim=-1,
+ )
+ else:
+ raise ValueError(
+ f"The loaded motion has {num_joints_loaded} joints but the current model "
+ f"({session.model_name}) has {num_joints_skeleton} joints. "
+ "Load a motion generated with the same skeleton, or switch the model to match the motion."
+ )
+ elif joints_rot.shape[1] != num_joints_skeleton:
+ raise ValueError(
+ f"Rotation data has {joints_rot.shape[1]} joints but the current model has "
+ f"{num_joints_skeleton} joints. The NPZ may be corrupted or from a different skeleton."
+ )
+
+ # Apply G1 real robot projection (1-DoF per joint + axis limits) if enabled.
+ if (
+ "g1" in session.model_name
+ and isinstance(session.skeleton, G1Skeleton34)
+ and gui_real_robot_rotations_checkbox.value
+ ):
+ joints_pos, joints_rot = generation.apply_g1_real_robot_projection(
+ session.skeleton, joints_pos, joints_rot
+ )
+
+ # Update duration and frame range based on loaded motion
+ num_frames = joints_pos.shape[0]
+ duration = num_frames / session.model_fps
+
+ # Update GUI elements
+ session.cur_duration = duration
+ session.max_frame_idx = num_frames - 1
+
+ # Clear existing motions and add the loaded one
+ demo.clear_motions(client.client_id)
+ demo.add_character_motion(
+ client,
+ session.skeleton,
+ joints_pos,
+ joints_rot,
+ foot_contacts,
+ )
+
+ # Reset to frame 0
+ demo.set_frame(client.client_id, 0)
+
+ @gui_load_motion_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+
+ load_path = gui_load_motion_path_text.value
+ loading_notif = event_client.add_notification(
+ title="Loading motion...",
+ body=f"Loading from {load_path}",
+ loading=True,
+ with_close_button=False,
+ auto_close_seconds=None,
+ )
+ try:
+ load_motion(event_client, load_path)
+
+ loading_notif.title = "Motion loaded!"
+ loading_notif.body = f"Loaded motion from {load_path} ({session.max_frame_idx + 1} frames, {session.cur_duration:.2f}s)"
+ loading_notif.loading = False
+ loading_notif.with_close_button = True
+ loading_notif.auto_close_seconds = 5.0
+ loading_notif.color = "green"
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ loading_notif.title = "Failed to load motion!"
+ loading_notif.body = str(e)
+ loading_notif.loading = False
+ loading_notif.with_close_button = True
+ loading_notif.auto_close_seconds = 10.0
+ loading_notif.color = "red"
+
+ def save_constraints(client, save_path):
+ session = demo.client_sessions[client.client_id]
+ # Keep save behavior aligned with demo frame convention:
+ # valid frame indices are [0, max_frame_idx], so count is +1.
+ num_frames = session.max_frame_idx + 1
+ model_bundle = demo.load_model(session.model_name)
+ constraints_lst = demo.compute_model_constraints_lst(session, model_bundle, num_frames)
+ save_constraints_lst(save_path, constraints_lst)
+
+ @gui_save_constraints_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ if get_active_session(event_client) is None:
+ return
+
+ try:
+ save_path = gui_save_constraints_path_text.value
+ save_constraints(event_client, save_path)
+ event_client.add_notification(
+ title="Constraints saved!",
+ body=f"Saved constraints to {save_path}",
+ auto_close_seconds=5.0,
+ color="green",
+ )
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ event_client.add_notification(
+ title="Failed to save constraints!",
+ body=str(e),
+ auto_close_seconds=10.0,
+ color="red",
+ )
+
+ def load_constraints(client, load_path):
+ session = demo.client_sessions[client.client_id]
+ constraints_lst = load_constraints_lst(load_path, skeleton=session.skeleton)
+
+ # Clear existing constraints first
+ with session.timeline_data["keyframe_update_lock"]:
+ for constraint in list(session.constraints.values()):
+ constraint.clear()
+ client.timeline.clear_keyframes()
+ client.timeline.clear_intervals()
+
+ # Add loaded constraints to the session
+ # We need to directly add constraint data, not read from current motion
+ device = demo.device
+ for constraint_obj in constraints_lst:
+ constraint_type = constraint_obj.name
+
+ # decompose the frame indices into intervals or single keyframes
+ frame_indices = constraint_obj.frame_indices
+ (
+ intervals,
+ intervals_indices,
+ single_frames,
+ single_frames_indices,
+ ) = extract_intervals_and_singles(frame_indices)
+
+ load_targets: list[dict] = []
+ root_pos = None
+
+ if constraint_type == "root2d":
+ # smooth_root_2d is [T, 2] (x, z), convert to [T, 3] (x, 0, z)
+ num_frames = constraint_obj.smooth_root_2d.shape[0]
+ root_pos = torch.zeros(num_frames, 3, device=device)
+ root_pos[:, 0] = constraint_obj.smooth_root_2d[:, 0]
+ root_pos[:, 2] = constraint_obj.smooth_root_2d[:, 1]
+ load_targets = [
+ {
+ "track_name": "2D Root",
+ "constraint_track": session.constraints["2D Root"],
+ }
+ ]
+ elif constraint_type == "fullbody":
+ load_targets = [
+ {
+ "track_name": "Full-Body",
+ "constraint_track": session.constraints["Full-Body"],
+ }
+ ]
+ elif constraint_type in {
+ "left-hand",
+ "right-hand",
+ "left-foot",
+ "right-foot",
+ }:
+ track_name = {
+ "left-hand": "Left Hand",
+ "right-hand": "Right Hand",
+ "left-foot": "Left Foot",
+ "right-foot": "Right Foot",
+ }[constraint_type]
+ load_targets = [
+ {
+ "track_name": track_name,
+ "constraint_track": session.constraints["End-Effectors"],
+ "joint_names": constraint_obj.joint_names,
+ "end_effector_type": constraint_type,
+ }
+ ]
+ elif constraint_type in {"end-effector", "end-effectors"}:
+ # Backward-compatible loader:
+ # split a generic end-effector constraint into per-limb timeline tracks.
+ joint_names_set = set(constraint_obj.joint_names)
+ for jname, track_name, eff_type in [
+ ("LeftHand", "Left Hand", "left-hand"),
+ ("RightHand", "Right Hand", "right-hand"),
+ ("LeftFoot", "Left Foot", "left-foot"),
+ ("RightFoot", "Right Foot", "right-foot"),
+ ]:
+ if jname not in joint_names_set:
+ continue
+ target_joint_names = [jname]
+ if "Hips" in joint_names_set:
+ target_joint_names.append("Hips")
+ load_targets.append(
+ {
+ "track_name": track_name,
+ "constraint_track": session.constraints["End-Effectors"],
+ "joint_names": target_joint_names,
+ "end_effector_type": eff_type,
+ }
+ )
+ if not load_targets:
+ raise KeyError(
+ "No recognized end-effector joint in constraint "
+ f"joint_names={constraint_obj.joint_names}"
+ )
+ else:
+ raise KeyError(f"Unsupported constraint type in loader: {constraint_type}")
+
+ for target in load_targets:
+ track_id = session.timeline_data["tracks_ids"][target["track_name"]]
+ constraint_track = target["constraint_track"]
+
+ # add intervals
+ for (start_idx, end_idx), (start_idx_t, end_idx_t) in zip(intervals, intervals_indices):
+ # Add to timeline
+ interval_id = client.timeline.add_interval(track_id, start_idx, end_idx)
+ session.timeline_data["intervals"][interval_id] = {
+ "track_id": track_id,
+ "start_frame_idx": start_idx,
+ "end_frame_idx": end_idx,
+ "locked": False,
+ "opacity": 1.0,
+ "value": None,
+ }
+ if constraint_type == "root2d":
+ constraint_track.add_interval(
+ interval_id,
+ start_idx,
+ end_idx,
+ root_pos[start_idx_t : end_idx_t + 1],
+ )
+ elif constraint_type == "fullbody":
+ constraint_track.add_interval(
+ interval_id,
+ start_idx,
+ end_idx,
+ constraint_obj.global_joints_positions[start_idx_t : end_idx_t + 1],
+ constraint_obj.global_joints_rots[start_idx_t : end_idx_t + 1],
+ )
+ else:
+ constraint_track.add_interval(
+ interval_id,
+ start_idx,
+ end_idx,
+ constraint_obj.global_joints_positions[start_idx_t : end_idx_t + 1],
+ constraint_obj.global_joints_rots[start_idx_t : end_idx_t + 1],
+ target["joint_names"],
+ target["end_effector_type"],
+ )
+
+ # add keyframes
+ for frame, frame_t in zip(single_frames, single_frames_indices):
+ # Add to timeline
+ keyframe_id = client.timeline.add_keyframe(track_id, frame)
+ session.timeline_data["keyframes"][keyframe_id] = {
+ "track_id": track_id,
+ "frame": frame,
+ "locked": False,
+ "opacity": 1.0,
+ "value": None,
+ }
+ if constraint_type == "root2d":
+ constraint_track.add_keyframe(
+ keyframe_id,
+ frame,
+ root_pos[frame_t],
+ )
+ elif constraint_type == "fullbody":
+ constraint_track.add_keyframe(
+ keyframe_id,
+ frame,
+ constraint_obj.global_joints_positions[frame_t],
+ constraint_obj.global_joints_rots[frame_t],
+ )
+ else:
+ constraint_track.add_keyframe(
+ keyframe_id,
+ frame,
+ constraint_obj.global_joints_positions[frame_t],
+ constraint_obj.global_joints_rots[frame_t],
+ target["joint_names"],
+ target["end_effector_type"],
+ )
+
+ @gui_load_constraints_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ if get_active_session(event_client) is None:
+ return
+
+ try:
+ load_path = gui_load_constraints_path_text.value
+ load_constraints(event_client, load_path)
+ session = demo.client_sessions[event_client.client_id]
+ apply_constraint_overlay_visibility(session)
+
+ event_client.add_notification(
+ title="Constraints loaded!",
+ body=f"Loaded constraints from {load_path}",
+ auto_close_seconds=5.0,
+ color="green",
+ )
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ event_client.add_notification(
+ title="Failed to load constraints!",
+ body=str(e),
+ auto_close_seconds=10.0,
+ color="red",
+ )
+
+ with client.gui.add_folder("Exports", expand_by_default=False):
+ with client.gui.add_folder("Screenshot", expand_by_default=False, visible=not HF_MODE):
+ gui_screenshot_path_text = client.gui.add_text(
+ "Save Path",
+ initial_value="render.png",
+ hint="Filename for the screenshot (PNG).",
+ )
+ gui_screenshot_button = client.gui.add_button(
+ "Download Screenshot",
+ hint="Capture the current canvas and download a PNG.",
+ )
+ with client.gui.add_folder("Video", expand_by_default=False, visible=not HF_MODE):
+ gui_video_path_text = client.gui.add_text(
+ "Save Path",
+ initial_value="render.mp4",
+ hint="Filename for the video (MP4).",
+ )
+ gui_video_button = client.gui.add_button(
+ "Download Video",
+ hint="Render every frame and download as MP4.",
+ )
+ with client.gui.add_folder("Motion", expand_by_default=True):
+ gui_download_name_text = client.gui.add_text(
+ "Name",
+ initial_value="output",
+ hint="Base filename to save as (extension will be added based on format if omitted).",
+ )
+ gui_download_format_dropdown = client.gui.add_dropdown(
+ "Format",
+ options=(
+ ["NPZ", "CSV"]
+ if "g1" in model_name.lower()
+ else ["NPZ", "AMASS NPZ"]
+ if "smplx" in model_name.lower()
+ else ["NPZ", "BVH"]
+ ),
+ initial_value="NPZ",
+ )
+ gui_download_bvh_standard_tpose_checkbox = client.gui.add_checkbox(
+ "Standard T-pose",
+ initial_value=False,
+ hint="For BVH export, use the standard T-pose rest skeleton.",
+ visible=False,
+ )
+ gui_download_button = client.gui.add_button(
+ "Download",
+ hint="Download the current motion (format + name above).",
+ )
+
+ def _download_bytes_to_browser(
+ event_client: viser.ClientHandle,
+ *,
+ data: bytes,
+ filename: str,
+ mime_type: str = "application/octet-stream",
+ ) -> None:
+ """Trigger a browser download for an in-memory byte payload.
+
+ Important: this intentionally does NOT use `showSaveFilePicker()` to avoid
+ Chrome/Edge's file-write permission prompt ("this site can see edits you make").
+ If you want "always ask where to save", configure your browser download settings.
+ """
+ import base64
+ import json
+
+ # Base64 is the most robust way to move binary over our websocket JS channel.
+ b64 = base64.b64encode(data).decode("ascii")
+ js = f"""
+(() => {{
+ const filename = {json.dumps(filename)};
+ const mimeType = {json.dumps(mime_type)};
+ const b64 = {json.dumps(b64)};
+
+ // Decode base64 -> Uint8Array.
+ const binStr = atob(b64);
+ const bytes = new Uint8Array(binStr.length);
+ for (let i = 0; i < binStr.length; i++) bytes[i] = binStr.charCodeAt(i);
+ const blob = new Blob([bytes], {{ type: mimeType }});
+
+ // Standard browser download behavior.
+ const url = URL.createObjectURL(blob);
+ const a = document.createElement("a");
+ a.href = url;
+ a.download = filename;
+ document.body.appendChild(a);
+ a.click();
+ a.remove();
+ URL.revokeObjectURL(url);
+}})();
+"""
+ # Reuse viser’s JS execution mechanism (used for Plotly setup).
+ from viser import _messages as _viser_messages
+
+ event_client.gui._websock_interface.queue_message( # type: ignore[attr-defined]
+ _viser_messages.RunJavascriptMessage(source=js)
+ )
+
+ def _motion_to_npz_bytes(motion) -> bytes:
+ motion_data = _motion_to_numpy_dict(motion)
+ return kimodo_npz_to_bytes(motion_data)
+
+ def _motion_to_csv_bytes(motion, session: ClientSession) -> bytes:
+ motion_data = _motion_to_numpy_dict(motion)
+ return g1_csv_to_bytes(motion_data, session.skeleton, demo.device)
+
+ def _motion_to_amass_npz_bytes(motion, session: ClientSession) -> bytes:
+ motion_data = _motion_to_numpy_dict(motion)
+ return amass_npz_to_bytes(motion_data, session.skeleton, session.model_fps)
+
+ def _get_motion_export_formats(loaded_model_name: str) -> list[str]:
+ model_name_lower = (loaded_model_name or "").lower()
+ if "g1" in model_name_lower:
+ return ["NPZ", "CSV"]
+ if "smplx" in model_name_lower:
+ return ["NPZ", "AMASS NPZ"]
+ return ["NPZ", "BVH"]
+
+ def _update_format_dropdown(dropdown, loaded_model_name: str) -> None:
+ new_options = _get_motion_export_formats(loaded_model_name)
+ current_value = str(dropdown.value)
+ dropdown.options = new_options
+ dropdown.value = current_value if current_value in new_options else new_options[0]
+
+ def _update_motion_export_dropdown(loaded_model_name: str) -> None:
+ _update_format_dropdown(gui_download_format_dropdown, loaded_model_name)
+ _update_format_dropdown(gui_save_motion_format_dropdown, loaded_model_name)
+ _update_bvh_standard_tpose_visibility()
+
+ def _update_bvh_standard_tpose_visibility() -> None:
+ gui_save_bvh_standard_tpose_checkbox.visible = (
+ str(gui_save_motion_format_dropdown.value).upper() == "BVH"
+ )
+ gui_download_bvh_standard_tpose_checkbox.visible = (
+ str(gui_download_format_dropdown.value).upper() == "BVH"
+ )
+
+ @gui_save_motion_format_dropdown.on_update
+ def _(_event: viser.GuiEvent) -> None:
+ _update_bvh_standard_tpose_visibility()
+
+ @gui_download_format_dropdown.on_update
+ def _(_event: viser.GuiEvent) -> None:
+ _update_bvh_standard_tpose_visibility()
+
+ def _coerce_download_filename(raw_name: str, *, ext: str) -> str:
+ """Coerce a user-entered filename to a safe basename with the desired extension.
+
+ - If empty: uses "output{ext}"
+ - If no extension: appends ext
+ - If endswith a known export extension: rewrites extension to ext (prevents mismatches)
+ - Any provided directory components are stripped
+ """
+ import os
+
+ name = (raw_name or "").strip()
+ name = os.path.basename(name.replace("\\", "/"))
+ if name == "":
+ return f"output{ext}"
+
+ known_exts = (".npz", ".bvh", ".csv", ".png", ".mp4")
+ lower = name.lower()
+ if lower.endswith(known_exts):
+ return os.path.splitext(name)[0] + ext
+
+ root, cur_ext = os.path.splitext(name)
+ if cur_ext == "":
+ return name + ext
+ return name
+
+ def _get_render_size(event_client: viser.ClientHandle) -> tuple[int, int]:
+ width = int(event_client.camera.image_width)
+ height = int(event_client.camera.image_height)
+ if width <= 0 or height <= 0:
+ # Fall back to a reasonable default if the camera hasn't synced yet.
+ return (1280, 720)
+ return (width, height)
+
+ def _round_up_to_multiple(value: int, multiple: int) -> int:
+ if multiple <= 0:
+ return value
+ return ((value + multiple - 1) // multiple) * multiple
+
+ def _download_canvas_to_browser(event_client: viser.ClientHandle, *, filename: str) -> None:
+ """Use the client-side canvas save path to avoid server-side renders."""
+ import json
+
+ js = f"""
+(() => {{
+ const filename = {json.dumps(filename)};
+ const canvases = Array.from(document.querySelectorAll("canvas"));
+ if (!canvases.length) {{
+ console.error("No canvases found to save.");
+ return;
+ }}
+ // Pick the largest canvas by area (usually the main 3D view).
+ const canvas = canvases.reduce((best, cur) => {{
+ const bestArea = (best?.width || 0) * (best?.height || 0);
+ const curArea = (cur?.width || 0) * (cur?.height || 0);
+ return curArea > bestArea ? cur : best;
+ }}, null);
+ if (!canvas) {{
+ console.error("No canvas selected to save.");
+ return;
+ }}
+ canvas.toBlob((blob) => {{
+ if (!blob) {{
+ console.error("Export failed");
+ return;
+ }}
+ const url = URL.createObjectURL(blob);
+ const a = document.createElement("a");
+ a.href = url;
+ a.download = filename;
+ document.body.appendChild(a);
+ a.click();
+ a.remove();
+ URL.revokeObjectURL(url);
+ }}, "image/png");
+}})();
+"""
+ from viser import _messages as _viser_messages
+
+ event_client.gui._websock_interface.queue_message( # type: ignore[attr-defined]
+ _viser_messages.RunJavascriptMessage(source=js)
+ )
+
+ @gui_screenshot_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ if get_active_session(event_client) is None:
+ return
+
+ try:
+ filename = _coerce_download_filename(
+ str(gui_screenshot_path_text.value),
+ ext=".png",
+ )
+ _download_canvas_to_browser(event_client, filename=filename)
+ event_client.add_notification(
+ title="Screenshot download started",
+ body=f"Saving {filename}",
+ auto_close_seconds=5.0,
+ color="green",
+ )
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ event_client.add_notification(
+ title="Failed to download screenshot!",
+ body=str(e),
+ auto_close_seconds=10.0,
+ color="red",
+ )
+
+ @gui_video_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+ recording_notification: viser.NotificationHandle | None = None
+ try:
+ recording_notification = event_client.add_notification(
+ title="Recording video...",
+ body="Saving frames, please wait.",
+ loading=True,
+ with_close_button=False,
+ auto_close_seconds=None,
+ color="blue",
+ )
+ event_client.timeline.disable_constraints()
+ width, height = _get_render_size(event_client)
+ # Avoid ffmpeg macro block resizing warnings.
+ width = _round_up_to_multiple(width, 16)
+ height = _round_up_to_multiple(height, 16)
+ original_frame = session.frame_idx
+ frames = []
+ for frame_idx in range(session.max_frame_idx + 1):
+ demo.set_frame(
+ event_client.client_id,
+ frame_idx,
+ update_timeline=True,
+ )
+ frames.append(
+ event_client.get_render(
+ height=height,
+ width=width,
+ transport_format="jpeg",
+ )
+ )
+
+ # Restore the original frame (and timeline).
+ demo.set_frame(event_client.client_id, original_frame)
+
+ import imageio.v3 as iio
+
+ filename = _coerce_download_filename(
+ str(gui_video_path_text.value),
+ ext=".mp4",
+ )
+ payload = iio.imwrite(
+ "",
+ frames,
+ extension=".mp4",
+ fps=float(session.model_fps),
+ codec="h264",
+ plugin="pyav",
+ )
+ event_client.send_file_download(filename, payload, save_immediately=True)
+ event_client.add_notification(
+ title="Video download started",
+ body=f"Saving {filename}",
+ auto_close_seconds=5.0,
+ color="green",
+ )
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ event_client.add_notification(
+ title="Failed to download video!",
+ body=str(e),
+ auto_close_seconds=10.0,
+ color="red",
+ )
+ finally:
+ event_client.timeline.enable_constraints()
+ if recording_notification is not None:
+ recording_notification.remove()
+
+ @gui_download_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+ motion = _get_primary_motion(session)
+ try:
+ fmt = str(gui_download_format_dropdown.value).upper()
+ raw_name = str(gui_download_name_text.value)
+
+ if fmt == "BVH":
+ filename = _coerce_download_filename(raw_name, ext=".bvh")
+ payload = motion_to_bvh_bytes(
+ motion.joints_local_rot,
+ motion.joints_pos[:, session.skeleton.root_idx, :], # root positions
+ skeleton=session.skeleton,
+ fps=float(session.model_fps),
+ standard_tpose=bool(gui_download_bvh_standard_tpose_checkbox.value),
+ )
+ mime = "text/plain"
+ elif fmt == "CSV":
+ filename = _coerce_download_filename(raw_name, ext=".csv")
+ payload = _motion_to_csv_bytes(motion, session)
+ mime = "text/csv"
+ elif fmt == "AMASS NPZ":
+ filename = _coerce_download_filename(raw_name, ext=".npz")
+ payload = _motion_to_amass_npz_bytes(motion, session)
+ mime = "application/octet-stream"
+ else:
+ # Default to NPZ (most common and matches existing save/load).
+ filename = _coerce_download_filename(raw_name, ext=".npz")
+ payload = _motion_to_npz_bytes(motion)
+ mime = "application/octet-stream"
+
+ _download_bytes_to_browser(
+ event_client,
+ data=payload,
+ filename=filename,
+ mime_type=mime,
+ )
+
+ event_client.add_notification(
+ title="Download started",
+ body=f"Saving {filename}",
+ auto_close_seconds=5.0,
+ color="green",
+ )
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ event_client.add_notification(
+ title="Failed to download motion!",
+ body=str(e),
+ auto_close_seconds=10.0,
+ color="red",
+ )
+
+ @gui_save_example_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ from kimodo.tools import save_json
+
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+
+ save_dir = gui_save_example_path_text.value
+ if os.path.exists(save_dir):
+ event_client.add_notification(
+ title="Failed to save example!",
+ body="Example directory already exists",
+ auto_close_seconds=10.0,
+ color="red",
+ )
+ return
+
+ try:
+ os.makedirs(save_dir)
+ # save the constraints
+ constraint_path = os.path.join(save_dir, "constraints.json")
+ save_constraints(event_client, constraint_path)
+ # save the motion
+ motion_path = os.path.join(save_dir, "motion.npz")
+ save_motion(event_client, motion_path, "NPZ")
+ # save the gui metadata
+ meta_path = os.path.join(save_dir, "meta.json")
+ prompt_texts = []
+ prompt_durations_sec = []
+ prompt_values = sorted(
+ [x for x in client.timeline._prompts.values()],
+ key=lambda x: x.start_frame,
+ )
+ for i, prompt in enumerate(prompt_values):
+ prompt_texts.append(prompt.text)
+ # Match demo/generation convention:
+ # non-last prompts: [start, end) ; last prompt: [start, end].
+ n_frames = prompt.end_frame - prompt.start_frame
+ if i == len(prompt_values) - 1:
+ n_frames += 1
+ prompt_durations_sec.append(n_frames / session.model_fps)
+ if len(prompt_texts) == 1:
+ meta_info = {
+ "text": prompt_texts[0],
+ "duration": prompt_durations_sec[0],
+ }
+ else:
+ meta_info = {
+ "texts": prompt_texts,
+ "durations": prompt_durations_sec,
+ }
+ meta_info["num_samples"] = gui_num_samples_slider.value
+ meta_info["seed"] = gui_seed.value
+ meta_info["diffusion_steps"] = gui_diffusion_steps_slider.value
+ meta_info["cfg"] = {
+ "enabled": gui_cfg_checkbox.value,
+ "text_weight": gui_cfg_text_weight_slider.value,
+ "constraint_weight": gui_cfg_constraint_weight_slider.value,
+ }
+ save_json(meta_path, meta_info)
+
+ # update the example dropdown
+ session.example_dict = viser_utils.load_example_cases(session.examples_base_dir)
+ update_examples_dropdown(session.example_dict, keep_selection=True)
+
+ event_client.add_notification(
+ title="Example saved!",
+ body=f"Saved example to {save_dir}",
+ auto_close_seconds=5.0,
+ color="green",
+ )
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ event_client.add_notification(
+ title="Failed to save example!",
+ body=str(e),
+ auto_close_seconds=10.0,
+ color="red",
+ )
+
+ def set_new_duration(client_id, new_duration):
+ session = demo.client_sessions[client_id]
+ session.cur_duration = new_duration
+ update_duration_gui(new_duration)
+ session.max_frame_idx = int(session.cur_duration * session.model_fps - 1)
+ if session.frame_idx > session.max_frame_idx:
+ demo.set_frame(client_id, session.max_frame_idx)
+
+ def apply_model_selection(new_model_name: str) -> None:
+ session = demo.client_sessions[client_id]
+ if new_model_name == session.model_name:
+ return
+
+ session.playing = False # Pause playback when switching models.
+
+ old_model_fps = session.model_fps
+ old_duration = session.cur_duration
+ old_prompts = [
+ (prompt.text, prompt.start_frame, prompt.end_frame) for prompt in client.timeline._prompts.values()
+ ]
+ old_default_zoom_frames = client.timeline._default_num_frames_zoom
+ old_max_zoom_frames = client.timeline._max_frames_zoom
+
+ model_bundle = demo.load_model(new_model_name)
+
+ # Clear motions and constraints when switching models.
+ if session.edit_mode and session.motions:
+ exit_editing_mode(session)
+ session.edit_mode = False
+ demo.clear_motions(client_id)
+ with session.timeline_data["keyframe_update_lock"]:
+ for constraint in list(session.constraints.values()):
+ constraint.clear()
+ session.constraints = demo.build_constraint_tracks(client, model_bundle.skeleton)
+ session.timeline_data["keyframes"] = {}
+ session.timeline_data["intervals"] = {}
+ client.timeline.clear_keyframes()
+ client.timeline.clear_intervals()
+
+ session.model_name = new_model_name
+ session.model_fps = model_bundle.model_fps
+ session.skeleton = model_bundle.skeleton
+ session.motion_rep = model_bundle.motion_rep
+ session.cur_duration = old_duration
+ session.max_frame_idx = int(session.cur_duration * session.model_fps - 1)
+ session.frame_idx = 0
+ session.edit_mode = False
+
+ demo.set_timeline_defaults(client.timeline, session.model_fps)
+ client.timeline.set_current_frame(0)
+ gui_model_fps.value = session.model_fps
+ update_duration_gui(session.cur_duration)
+
+ if old_model_fps > 0:
+ default_zoom_seconds = old_default_zoom_frames / old_model_fps
+ max_zoom_seconds = old_max_zoom_frames / old_model_fps
+ new_default_zoom = int(round(default_zoom_seconds * session.model_fps))
+ new_max_zoom = int(round(max_zoom_seconds * session.model_fps))
+ new_default_zoom = max(1, new_default_zoom)
+ new_max_zoom = max(new_default_zoom, new_max_zoom)
+ client.timeline.set_zoom_settings(
+ default_num_frames_zoom=new_default_zoom,
+ max_frames_zoom=new_max_zoom,
+ )
+
+ client.timeline.clear_prompts()
+ if old_prompts and old_model_fps > 0:
+ for i, (prompt_text, start_frame, end_frame) in enumerate(old_prompts):
+ start_sec = start_frame / old_model_fps
+ end_sec = end_frame / old_model_fps
+ new_start = int(round(start_sec * session.model_fps))
+ new_end = int(round(end_sec * session.model_fps))
+ new_start = max(0, min(new_start, session.max_frame_idx))
+ new_end = max(new_start, min(new_end, session.max_frame_idx))
+ color = PROMPT_COLORS[i % len(PROMPT_COLORS)]
+ client.timeline.add_prompt(prompt_text, new_start, new_end, color=color)
+
+ session.examples_base_dir = demo.get_examples_base_dir(new_model_name, absolute=True)
+ session.example_dict = viser_utils.load_example_cases(session.examples_base_dir)
+ update_examples_dropdown(session.example_dict, keep_selection=False)
+ gui_save_example_path_text.value = os.path.join(
+ demo.get_examples_base_dir(new_model_name, absolute=True),
+ "custom_example_1",
+ )
+ gui_load_example_path_text.value = os.path.join(
+ demo.get_examples_base_dir(new_model_name, absolute=True),
+ "custom_example_1",
+ )
+
+ demo.add_character_motion(client, session.skeleton)
+ apply_constraint_overlay_visibility(session)
+
+ def _update_version_and_display_from_dataset_skeleton() -> None:
+ dataset_ui = gui_dataset_selector.value
+ skeleton_display = gui_skeleton_selector.value
+ skeleton_val = get_skeleton_key_from_display_name(skeleton_display)
+ if skeleton_val is None:
+ return
+ models = get_models_for_dataset_skeleton(dataset_ui, skeleton_val, family="Kimodo")
+ if not models:
+ return
+ gui_version_selector.options = [m.display_name for m in models]
+ gui_version_selector.value = models[0].display_name
+ gui_version_selector.visible = len(models) > 1
+ gui_model_display.content = f"**Model:** {models[0].display_name}"
+
+ def _update_visibility_for_loaded_model(loaded_model_name: str) -> None:
+ """Update model-specific controls from the currently loaded model only."""
+ if not loaded_model_name:
+ return
+ _update_motion_export_dropdown(loaded_model_name)
+ gui_use_soma_layer_checkbox.visible = "soma" in loaded_model_name
+ _is_g1 = "g1" in loaded_model_name
+ gui_real_robot_rotations_checkbox.visible = _is_g1
+ gui_postprocess_checkbox.visible = not _is_g1
+ gui_root_margin.visible = not _is_g1 and gui_postprocess_checkbox.value
+ if _is_g1:
+ gui_gizmo_space_dropdown.value = "Local"
+ gui_gizmo_space_dropdown.visible = not _is_g1
+ gui_gizmo_space_dropdown.disabled = _is_g1
+
+ def _on_load_model_click(event: viser.GuiEvent) -> None:
+ """Load the currently selected model (called from Load model button)."""
+ if get_active_session(event.client) is None:
+ return
+ new_model_name = gui_model_selector.value
+ if not new_model_name:
+ return
+ info = get_model_info(new_model_name)
+ if info is None:
+ return
+ session = demo.client_sessions[event.client.client_id]
+ if new_model_name == session.model_name:
+ return
+ loading_notif = event.client.add_notification(
+ title="Loading model...",
+ body=f"Loading {info.display_name}",
+ loading=True,
+ with_close_button=False,
+ )
+ try:
+ apply_model_selection(new_model_name)
+ _update_visibility_for_loaded_model(new_model_name)
+ loading_notif.title = "Model loaded"
+ loading_notif.body = f"{info.display_name} is ready."
+ loading_notif.loading = False
+ loading_notif.with_close_button = True
+ loading_notif.auto_close_seconds = 5.0
+ loading_notif.color = "green"
+ except Exception as e:
+ loading_notif.loading = False
+ loading_notif.with_close_button = True
+ event.client.add_notification(
+ title="Model failed to load",
+ body=str(e),
+ color="red",
+ auto_close_seconds=10.0,
+ )
+ gui_model_selector.set_from_short_key(session.model_name)
+
+ @gui_load_model_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ _on_load_model_click(event)
+
+ @gui_dataset_selector.on_update
+ def _(event: viser.GuiEvent) -> None:
+ if get_active_session(event.client) is None:
+ return
+ skeleton_labels = get_allowed_skeleton_labels(gui_dataset_selector.value)
+ gui_skeleton_selector.options = skeleton_labels
+ gui_skeleton_selector.value = skeleton_labels[0] if skeleton_labels else ""
+ _update_version_and_display_from_dataset_skeleton()
+
+ @gui_skeleton_selector.on_update
+ def _(event: viser.GuiEvent) -> None:
+ if get_active_session(event.client) is None:
+ return
+ _update_version_and_display_from_dataset_skeleton()
+
+ @gui_version_selector.on_update
+ def _(event: viser.GuiEvent) -> None:
+ if get_active_session(event.client) is None:
+ return
+ info = get_model_info(gui_model_selector.value)
+ if info is not None:
+ gui_model_display.content = f"**Model:** {info.display_name}"
+
+ @gui_use_soma_layer_checkbox.on_update
+ def _(event: viser.GuiEvent) -> None:
+ session = get_active_session(event.client)
+ if session is None or "soma" not in (session.model_name or ""):
+ return
+
+ loading_notif = event.client.add_notification(
+ title="Applying SOMA layer...",
+ body="Updating mesh.",
+ loading=True,
+ with_close_button=False,
+ )
+ try:
+ current_motion = list(session.motions.values())[0] if session.motions else None
+ current_frame_idx = session.frame_idx
+
+ # Recreate the character to apply the new SOMA mesh mode selection.
+ demo.clear_motions(event.client.client_id)
+ if current_motion is None:
+ demo.add_character_motion(event.client, session.skeleton)
+ else:
+ demo.add_character_motion(
+ event.client,
+ session.skeleton,
+ current_motion.joints_pos,
+ current_motion.joints_rot,
+ current_motion.foot_contacts,
+ )
+
+ demo.set_frame(event.client.client_id, current_frame_idx)
+ except Exception as e:
+ print(e)
+ event.client.add_notification(
+ title="SOMA layer failed",
+ body=str(e),
+ color="red",
+ auto_close_seconds=10.0,
+ )
+ gui_use_soma_layer_checkbox.value = not gui_use_soma_layer_checkbox.value
+ finally:
+ loading_notif.loading = False
+ loading_notif.with_close_button = True
+ loading_notif.auto_close_seconds = 2.0
+
+ @gui_real_robot_rotations_checkbox.on_update
+ def _(event: viser.GuiEvent) -> None:
+ session = get_active_session(event.client)
+ if session is None or "g1" not in session.model_name:
+ return
+ if not isinstance(session.skeleton, G1Skeleton34) or not session.motions:
+ return
+ if not gui_real_robot_rotations_checkbox.value:
+ return
+ # Reproject all displayed G1 motions to real robot DoF (1-DoF per joint + axis limits).
+ from kimodo.skeleton import global_rots_to_local_rots
+
+ current_frame_idx = session.frame_idx
+ for motion in session.motions.values():
+ if motion.length <= 1:
+ continue
+ rest_pos = motion.joints_pos[0:1]
+ rest_rot = motion.joints_rot[0:1]
+ same_as_rest = (motion.joints_pos - rest_pos).abs().max().item() < 1e-6 and (
+ motion.joints_rot - rest_rot
+ ).abs().max().item() < 1e-6
+ if same_as_rest:
+ continue
+ new_pos, new_rot = generation.apply_g1_real_robot_projection(
+ session.skeleton,
+ motion.joints_pos,
+ motion.joints_rot,
+ )
+ motion.joints_pos = new_pos
+ motion.joints_rot = new_rot
+ motion.joints_local_rot = global_rots_to_local_rots(new_rot, session.skeleton)
+ # Refresh skeleton and skinned mesh caches so the viz uses new positions.
+ motion.precompute_mesh_info()
+ demo.set_frame(event.client.client_id, current_frame_idx)
+ event.client.add_notification(
+ title="Real robot projection applied",
+ body="The motion is projected to G1 real robot DoF (1-DoF per joint, clamped to axis limits).",
+ auto_close_seconds=4.0,
+ color="green",
+ )
+
+ def load_example_from_path(
+ event_client: viser.ClientHandle,
+ example_path: str,
+ load_gt: bool = False,
+ ) -> None:
+ from kimodo.meta import parse_prompts_from_meta
+ from kimodo.tools import load_json
+
+ session = get_active_session(event_client)
+ if session is None:
+ return
+
+ # Pause playback when loading an example.
+ session.playing = False
+
+ if not os.path.isdir(example_path):
+ event_client.add_notification(
+ title="Example path not found",
+ body=f"Directory does not exist: {example_path}",
+ auto_close_seconds=5.0,
+ color="red",
+ )
+ return
+
+ # Long motions trigger a skinning precompute that can take several
+ # seconds; show a persistent "loading" notification so the user
+ # knows the app isn't frozen. Cleared in the finally block below.
+ loading_notif = event_client.add_notification(
+ title="Loading example...",
+ body=f"Loading {os.path.basename(example_path.rstrip(os.sep))}. This may take a moment for long motions.",
+ loading=True,
+ with_close_button=False,
+ )
+
+ try:
+ # constraints
+ constraints_path = os.path.join(example_path, "constraints.json")
+ if os.path.exists(constraints_path):
+ load_constraints(event_client, constraints_path)
+ else:
+ # clear all existing constraints
+ with session.timeline_data["keyframe_update_lock"]:
+ for constraint in list(session.constraints.values()):
+ constraint.clear()
+ event_client.timeline.clear_keyframes()
+ event_client.timeline.clear_intervals()
+ # motion
+ motion_filename = "gt_motion.npz" if load_gt else "motion.npz"
+ motion_path = os.path.join(example_path, motion_filename)
+ if os.path.exists(motion_path):
+ load_motion(event_client, motion_path)
+ # metadata
+ meta_path = os.path.join(example_path, "meta.json")
+ if os.path.exists(meta_path):
+ meta_info = load_json(meta_path)
+ event_client.timeline.clear_prompts()
+
+ texts, durations_sec = parse_prompts_from_meta(meta_info)
+ fps = session.model_fps
+ # Convert durations (seconds) to consecutive frame bounds
+ num_frames = 0
+ frame_bounds = []
+ for i, d in enumerate(durations_sec):
+ n_frames = max(1, int(round(d * fps)))
+ start_frame = num_frames
+ # Inverse of compute_prompt_num_frames():
+ # non-last prompts end at next prompt start (exclusive),
+ # last prompt includes its end frame.
+ if i == len(durations_sec) - 1:
+ end_frame = num_frames + n_frames - 1
+ else:
+ end_frame = num_frames + n_frames
+ frame_bounds.append((start_frame, end_frame))
+ num_frames += n_frames
+
+ # Adapt timeline zoom to the loaded motion.
+ target_visible_frames = int(math.ceil(1.10 * num_frames))
+ event_client.timeline.set_zoom_settings(
+ default_num_frames_zoom=target_visible_frames,
+ )
+
+ for i, (prompt_text, (start_frame, end_frame)) in enumerate(zip(texts, frame_bounds)):
+ color = PROMPT_COLORS[i % len(PROMPT_COLORS)]
+ event_client.timeline.add_prompt(prompt_text, start_frame, end_frame, color=color)
+
+ update_duration_auto()
+
+ # Only load optional fields if present
+ if "num_samples" in meta_info:
+ gui_num_samples_slider.value = meta_info["num_samples"]
+ if "seed" in meta_info:
+ gui_seed.value = meta_info["seed"]
+ if "diffusion_steps" in meta_info:
+ gui_diffusion_steps_slider.value = meta_info["diffusion_steps"]
+ if "cfg" in meta_info:
+ cfg = meta_info["cfg"]
+ if "enabled" in cfg:
+ gui_cfg_checkbox.value = cfg["enabled"]
+ if "text_weight" in cfg:
+ gui_cfg_text_weight_slider.value = cfg["text_weight"]
+ if "constraint_weight" in cfg:
+ gui_cfg_constraint_weight_slider.value = cfg["constraint_weight"]
+
+ # Set frame to 0 when example is loaded.
+ session.frame_idx = 0
+ event_client.timeline.set_current_frame(0)
+ demo.set_frame(event_client.client_id, 0)
+
+ event_client.add_notification(
+ title="Example loaded!",
+ body=f"Loaded example from {example_path}",
+ auto_close_seconds=5.0,
+ color="green",
+ )
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ event_client.add_notification(
+ title="Failed to load example!",
+ body=str(e),
+ auto_close_seconds=10.0,
+ color="red",
+ )
+ finally:
+ loading_notif.remove()
+
+ @gui_load_example_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+
+ if not session.example_dict or (gui_examples_dropdown.value not in session.example_dict):
+ event_client.add_notification(
+ title="No examples available",
+ body="No examples found for the selected model.",
+ auto_close_seconds=5.0,
+ color="red",
+ )
+ return
+
+ example_path = session.example_dict[gui_examples_dropdown.value]
+ load_example_from_path(event_client, example_path, gui_load_gt_checkbox.value)
+
+ @gui_load_example_from_path_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+
+ example_path = gui_load_example_path_text.value
+ if not example_path:
+ event_client.add_notification(
+ title="No example path",
+ body="Please provide an example directory.",
+ auto_close_seconds=5.0,
+ color="red",
+ )
+ return
+ load_example_from_path(event_client, example_path, gui_load_gt_checkbox.value)
+
+ @gui_cfg_checkbox.on_update
+ def _(_) -> None:
+ if not demo.client_active(client_id):
+ return
+ val = gui_cfg_checkbox.value
+ gui_cfg_text_weight_slider.visible = val
+ gui_cfg_constraint_weight_slider.visible = val
+
+ def exit_editing_mode(session: ClientSession):
+ gui_edit_constraint_button.label = "Enter Editing Mode"
+ gui_generate_button.disabled = False
+ gui_generate_button.label = "Generate"
+ gui_reset_constraint_button.disabled = True
+ if "g1" in session.model_name:
+ gui_gizmo_space_dropdown.value = "Local"
+ gui_gizmo_space_dropdown.disabled = True
+ gui_gizmo_space_dropdown.visible = False
+ else:
+ gui_gizmo_space_dropdown.disabled = False
+ gui_gizmo_space_dropdown.visible = True
+ gui_undo_drag_button.disabled = True
+ gui_use_soma_layer_checkbox.disabled = False
+ session.edit_mode_snapshot = None
+ session.undo_drag_snapshot = None
+
+ motion = list(session.motions.values())[0]
+ motion.clear_all_gizmos()
+ motion.character.set_skinned_mesh_wireframe(False)
+ motion.character.set_skeleton_visibility(False)
+ motion.character.set_skinned_mesh_visibility(True)
+ motion.character.set_skinned_mesh_opacity(1.0)
+ session.gui_elements.gui_viz_skinned_mesh_opacity_slider.value = 1.0
+
+ # If the path is dense, put the motion back on the path
+ if "2D Root" in session.constraints and session.constraints["2D Root"].dense_path:
+ _update_dense_path(motion, session)
+
+ gui_viz_skinned_mesh_checkbox.value = True
+ gui_viz_skeleton_checkbox.value = False
+
+ # enter editing mode callback
+ @gui_edit_constraint_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+
+ session.edit_mode = not session.edit_mode
+
+ edit_alert = "Entered editing mode"
+ no_edit_alert = "Exited editing mode"
+ edit_message = "You can now modify pose or path constraints."
+ no_edit_message = "Can now generate motions."
+ event_client.add_notification(
+ title=edit_alert if session.edit_mode else no_edit_alert,
+ body=edit_message if session.edit_mode else no_edit_message,
+ auto_close_seconds=10.0,
+ color="blue",
+ )
+
+ if session.edit_mode:
+ gui_edit_constraint_button.label = "Exit Editing Mode"
+ gui_generate_button.disabled = True
+ gui_generate_button.label = "Generate Disabled In Editing Mode"
+ if "g1" in session.model_name:
+ gui_gizmo_space_dropdown.value = "Local"
+ gui_gizmo_space_dropdown.disabled = True
+ gui_use_soma_layer_checkbox.disabled = True
+
+ assert len(session.motions) == 1, "Only one motion allowed in edit mode"
+ motion = list(session.motions.values())[0]
+ snapshot_frame_idx = min(session.frame_idx, motion.length - 1)
+ session.edit_mode_snapshot = {}
+ ensure_edit_snapshot(session, motion, snapshot_frame_idx)
+ gui_reset_constraint_button.disabled = False
+
+ motion.character.set_skeleton_visibility(True)
+ # motion.character.set_skinned_mesh_wireframe(True)
+ motion.character.set_skinned_mesh_opacity(0.65)
+ session.gui_elements.gui_viz_skinned_mesh_opacity_slider.value = 0.65
+ motion.character.set_skinned_mesh_visibility(True)
+ gui_viz_skinned_mesh_checkbox.value = True
+ gui_viz_skeleton_checkbox.value = True
+
+ # need gizmos for root translation and individual joints
+ def _on_root2d_gizmo_release():
+ if "2D Root" in session.constraints and session.constraints["2D Root"].dense_path:
+ mot = list(session.motions.values())[0]
+ _update_dense_path(mot, session)
+
+ def _on_gizmo_drag_start():
+ mot = list(session.motions.values())[0]
+ frame_idx = min(session.frame_idx, mot.length - 1)
+ session.undo_drag_snapshot = {
+ "frame_idx": frame_idx,
+ "joints_pos": mot.get_joints_pos(frame_idx),
+ "joints_rot": mot.get_joints_rot(frame_idx),
+ }
+ gui_undo_drag_button.disabled = False
+
+ motion.add_root_translation_gizmo(
+ session.constraints,
+ on_2d_root_drag_end=_on_root2d_gizmo_release,
+ on_drag_start=_on_gizmo_drag_start,
+ )
+ gizmo_space = "local" if "g1" in session.model_name else gui_gizmo_space_dropdown.value.lower()
+ motion.add_joint_gizmos(
+ session.constraints,
+ space=gizmo_space,
+ on_drag_start=_on_gizmo_drag_start,
+ )
+ else:
+ exit_editing_mode(session)
+
+ @gui_reset_constraint_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None or not session.edit_mode_snapshot:
+ return
+
+ if not session.motions:
+ return
+ motion = list(session.motions.values())[0]
+ snapshot_frame_idx = min(session.frame_idx, motion.length - 1)
+ if snapshot_frame_idx not in session.edit_mode_snapshot:
+ return
+ motion.update_pose_at_frame(
+ snapshot_frame_idx,
+ joints_pos=session.edit_mode_snapshot[snapshot_frame_idx]["joints_pos"],
+ joints_rot=session.edit_mode_snapshot[snapshot_frame_idx]["joints_rot"],
+ )
+ demo.set_frame(event_client.client_id, snapshot_frame_idx, update_timeline=False)
+
+ @gui_undo_drag_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None or session.undo_drag_snapshot is None:
+ return
+
+ if not session.motions:
+ return
+ motion = list(session.motions.values())[0]
+ frame_idx = session.undo_drag_snapshot["frame_idx"]
+ motion.update_pose_at_frame(
+ frame_idx,
+ joints_pos=session.undo_drag_snapshot["joints_pos"],
+ joints_rot=session.undo_drag_snapshot["joints_rot"],
+ )
+ demo.set_frame(event_client.client_id, frame_idx, update_timeline=False)
+ session.undo_drag_snapshot = None
+ gui_undo_drag_button.disabled = True
+
+ def validate_interval(start_frame_idx: int, end_frame_idx: int, max_frame_idx: int) -> bool:
+ if start_frame_idx < 0 or start_frame_idx > max_frame_idx:
+ return False
+ if end_frame_idx < 0 or end_frame_idx > max_frame_idx:
+ return False
+ if end_frame_idx < start_frame_idx:
+ return False
+ return True
+
+ def clamp_interval_to_range(
+ start_frame_idx: int, end_frame_idx: int, max_frame_idx: int
+ ) -> Optional[tuple[int, int]]:
+ if end_frame_idx < 0 or start_frame_idx > max_frame_idx:
+ return None
+ start_clamped = max(0, start_frame_idx)
+ end_clamped = min(max_frame_idx, end_frame_idx)
+ if end_clamped < start_clamped:
+ return None
+ return start_clamped, end_clamped
+
+ # add constraint callback
+ def add_constraint_callback(
+ constraint_id: str,
+ constraint_type: str,
+ frame_range: tuple[int, int],
+ joint_names: list[str] = None,
+ verbose: bool = True,
+ ):
+ """Add a constraint to the session.
+
+ Args:
+ constraint_type: str, the type of constraint to add
+ frame_range: tuple[int, int], the frame range to add the constraint to
+ joint_names: list[str], the names of the joints to constraint if the constraint type is End-Effectors
+ """
+ # Check if session still exists
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+
+ assert len(session.motions) == 1, "Only one motion allowed for adding constraints"
+ motion = list(session.motions.values())[0]
+
+ end_effector_type = None
+ if constraint_type in [
+ "Left Hand",
+ "Right Hand",
+ "Left Foot",
+ "Right Foot",
+ ]:
+ joint_names = [constraint_type.replace(" ", ""), "Hips"]
+ # Hips are required because of smooth root representation
+ end_effector_type = constraint_type.replace(" ", "-").lower()
+ constraint_type = "End-Effectors"
+
+ # check to make sure interval is valid
+ is_interval = frame_range[1] != frame_range[0]
+ start_frame_idx = int(frame_range[0])
+ end_frame_idx = int(frame_range[1])
+
+ if is_interval:
+ clamped = clamp_interval_to_range(start_frame_idx, end_frame_idx, session.max_frame_idx)
+ if clamped is None:
+ print("Interval outside range! Couldn't add constraint.")
+ return
+ start_frame_idx, end_frame_idx = clamped
+ else:
+ if not validate_interval(start_frame_idx, end_frame_idx, session.max_frame_idx):
+ print("Invalid interval! Couldn't add constraint.")
+ return
+
+ # collect input args for the constraint based on which track it is
+ if is_interval:
+ constraint_kwargs = {
+ "interval_id": constraint_id,
+ "start_frame_idx": start_frame_idx,
+ "end_frame_idx": end_frame_idx,
+ }
+ else:
+ constraint_kwargs = {
+ "keyframe_id": constraint_id,
+ "frame_idx": start_frame_idx,
+ }
+
+ if constraint_type in ["Full-Body", "End-Effectors"]:
+ constraint_kwargs["joints_pos"] = motion.get_joints_pos(start_frame_idx, end_frame_idx)
+ constraint_kwargs["joints_rot"] = motion.get_joints_rot(start_frame_idx, end_frame_idx)
+ if constraint_type == "End-Effectors":
+ constraint_kwargs["joint_names"] = joint_names
+ constraint_kwargs["end_effector_type"] = end_effector_type
+
+ elif constraint_type == "2D Root":
+ constraint_kwargs["root_pos"] = motion.get_projected_root_pos(start_frame_idx, end_frame_idx)
+
+ # add the keyframe(s) to the constraint track
+ constraint = session.constraints[constraint_type]
+ if is_interval:
+ constraint.add_interval(**constraint_kwargs)
+ else:
+ constraint.add_keyframe(**constraint_kwargs)
+
+ apply_constraint_overlay_visibility(session)
+
+ if verbose:
+ client.add_notification(
+ title="Constraint added",
+ body="",
+ auto_close_seconds=5.0,
+ color="blue",
+ )
+
+ # timeline callbacks for keyframes and intervals
+ @client.timeline.on_keyframe_add
+ def _(keyframe_id: str, track_id: str, frame: int):
+ """Called when a keyframe is added to a track."""
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ with session.timeline_data["keyframe_update_lock"]:
+ constraint_type = session.timeline_data["tracks"][track_id]["name"]
+ add_constraint_callback(
+ keyframe_id,
+ constraint_type,
+ (frame, frame),
+ verbose=False,
+ )
+ keyframe_data = client.timeline._keyframes.get(keyframe_id)
+ session.timeline_data["keyframes"][keyframe_id] = {
+ "frame": frame,
+ "track_id": track_id,
+ "locked": bool(keyframe_data.locked) if keyframe_data is not None else False,
+ "opacity": keyframe_data.opacity if keyframe_data is not None else 1.0,
+ "value": keyframe_data.value if keyframe_data is not None else None,
+ }
+ # Update smooth path when adding a keyframe (single action, not drag).
+ if constraint_type == "2D Root" and session.constraints["2D Root"].dense_path:
+ motion = list(session.motions.values())[0]
+ _update_dense_path(motion, session)
+
+ @client.timeline.on_interval_add
+ def handle_interval_add(interval_id: str, track_id: str, start_frame: int, end_frame: int):
+ """Called when an interval is added to a track."""
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ with session.timeline_data["keyframe_update_lock"]:
+ constraint_type = session.timeline_data["tracks"][track_id]["name"]
+ add_constraint_callback(
+ interval_id,
+ constraint_type,
+ (start_frame, end_frame),
+ verbose=False,
+ )
+ interval_data = client.timeline._intervals.get(interval_id)
+ session.timeline_data["intervals"][interval_id] = {
+ "track_id": track_id,
+ "start_frame_idx": start_frame,
+ "end_frame_idx": end_frame,
+ "locked": bool(interval_data.locked) if interval_data is not None else False,
+ "opacity": interval_data.opacity if interval_data is not None else 1.0,
+ "value": interval_data.value if interval_data is not None else None,
+ }
+ if constraint_type == "2D Root" and session.constraints["2D Root"].dense_path:
+ motion = list(session.motions.values())[0]
+ _update_dense_path(motion, session)
+
+ def remove_constraint_callback(
+ constraint_id: str,
+ constraint_type: str,
+ frame_range: tuple[int, int],
+ verbose: bool = True,
+ ) -> None:
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ session.updating_motions = True
+
+ is_interval = frame_range[1] != frame_range[0]
+ start_frame_idx = int(frame_range[0])
+ end_frame_idx = int(frame_range[1])
+
+ if is_interval:
+ clamped = clamp_interval_to_range(start_frame_idx, end_frame_idx, session.max_frame_idx)
+ if clamped is None:
+ return
+ start_frame_idx, end_frame_idx = clamped
+ else:
+ if not validate_interval(start_frame_idx, end_frame_idx, session.max_frame_idx):
+ print("Invalid interval! Couldn't remove constraint.")
+ return
+
+ if constraint_type in [
+ "Left Hand",
+ "Right Hand",
+ "Left Foot",
+ "Right Foot",
+ ]:
+ constraint_type = "End-Effectors"
+
+ constraint = session.constraints[constraint_type]
+ if is_interval:
+ constraint.remove_interval(constraint_id, start_frame_idx, end_frame_idx)
+ else:
+ constraint.remove_keyframe(constraint_id, start_frame_idx)
+
+ if verbose:
+ client.add_notification(
+ title="Constraint removed",
+ body="",
+ auto_close_seconds=5.0,
+ color="blue",
+ )
+
+ @client.timeline.on_keyframe_move
+ def handle_keyframe_move(keyframe_id: str, new_frame: int):
+ """Called when a keyframe is moved to a new frame."""
+ # print(f"Keyframe moved: {keyframe_id} to frame {new_frame}")
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+
+ # Cancel any pending timer for this keyframe
+ timeline_data = session.timeline_data
+ with timeline_data["keyframe_update_lock"]:
+ if keyframe_id in timeline_data["keyframe_move_timers"]:
+ timeline_data["keyframe_move_timers"][keyframe_id].cancel()
+
+ # Store the latest target frame
+ timeline_data["pending_keyframe_moves"][keyframe_id] = new_frame
+ # Create a new timer to execute the actual move after a delay
+ # This debounces rapid movements - only execute when user stops moving
+ timer = threading.Timer(
+ 0.03, # 10ms delay - adjust as needed
+ _execute_keyframe_move,
+ args=(client_id, keyframe_id, new_frame, session),
+ )
+ timeline_data["keyframe_move_timers"][keyframe_id] = timer
+ timer.start()
+
+ def _execute_keyframe_move(
+ client_id: int,
+ keyframe_id: str,
+ new_frame: int,
+ session: ClientSession,
+ ):
+ """Actually execute the keyframe move operation (called after debounce delay)."""
+
+ timeline_data = session.timeline_data
+ with timeline_data["keyframe_update_lock"]:
+ # Check if this move is still the latest one
+ if keyframe_id not in timeline_data["pending_keyframe_moves"]:
+ return # Move was cancelled
+
+ if timeline_data["pending_keyframe_moves"][keyframe_id] != new_frame:
+ return # A newer move superseded this one
+
+ # Remove from pending
+ del timeline_data["pending_keyframe_moves"][keyframe_id]
+ if keyframe_id in timeline_data["keyframe_move_timers"]:
+ del timeline_data["keyframe_move_timers"][keyframe_id]
+
+ # Now execute the actual move (keep it in the lock so we don't delete it while moving)
+ if keyframe_id not in timeline_data["keyframes"]:
+ # double check
+ return
+ keyframe_data = timeline_data["keyframes"][keyframe_id]
+ if not keyframe_data:
+ return
+
+ # if the frame did not move, don't do anything
+ if keyframe_data["frame"] == new_frame:
+ return
+
+ track_id = keyframe_data["track_id"]
+ constraint_type = timeline_data["tracks"][track_id]["name"]
+ cur_frame = keyframe_data["frame"]
+
+ # Remove constraint at old frame
+ remove_constraint_callback(
+ keyframe_id,
+ constraint_type,
+ (cur_frame, cur_frame),
+ verbose=False,
+ )
+ # Add constraint at new frame
+ add_constraint_callback(
+ keyframe_id,
+ constraint_type,
+ (new_frame, new_frame),
+ verbose=False,
+ )
+
+ # update our data
+ keyframe_data["frame"] = new_frame
+
+ # Schedule path update only after user stops dragging (no move for 300ms).
+ if constraint_type == "2D Root":
+ _schedule_dense_path_after_release(session)
+
+ @client.timeline.on_keyframe_delete
+ def handle_keyframe_delete(keyframe_id: str):
+ """Called when a keyframe is deleted."""
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ with session.timeline_data["keyframe_update_lock"]:
+ if keyframe_id not in session.timeline_data["keyframes"]:
+ return
+ keyframe_data = session.timeline_data["keyframes"][keyframe_id]
+ track_id = keyframe_data["track_id"]
+ constraint_type = session.timeline_data["tracks"][track_id]["name"]
+ cur_frame = keyframe_data["frame"]
+ remove_constraint_callback(
+ keyframe_id,
+ constraint_type,
+ (cur_frame, cur_frame),
+ verbose=False,
+ )
+ del session.timeline_data["keyframes"][keyframe_id]
+ if constraint_type == "2D Root" and session.constraints["2D Root"].dense_path:
+ motion = list(session.motions.values())[0]
+ _update_dense_path(motion, session)
+
+ @client.timeline.on_interval_move
+ def handle_interval_move(interval_id: str, new_start: int, new_end: int):
+ """Called when an interval is moved or resized."""
+ # print(f"Interval moved: {interval_id} to {new_start}-{new_end}")
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+
+ # Cancel any pending timer for this interval
+ # We share the same lock for keyframe and interval moves assuming the user can't move both at the same time
+ timeline_data = session.timeline_data
+ with timeline_data["keyframe_update_lock"]:
+ if interval_id in timeline_data["keyframe_move_timers"]:
+ timeline_data["keyframe_move_timers"][interval_id].cancel()
+
+ # Store the latest target frame
+ new_interval = (new_start, new_end)
+ timeline_data["pending_keyframe_moves"][interval_id] = new_interval
+ # Create a new timer to execute the actual move after a delay
+ # This debounces rapid movements - only execute when user stops moving
+ timer = threading.Timer(
+ 0.5, # 100ms delay - adding interval is much slower than moving a keyframe
+ _execute_interval_move,
+ args=(client_id, interval_id, new_interval, session),
+ )
+ timeline_data["keyframe_move_timers"][interval_id] = timer
+ timer.start()
+
+ def _execute_interval_move(
+ client_id: int,
+ interval_id: str,
+ new_interval: tuple[int, int],
+ session: ClientSession,
+ ):
+ """Actually execute the interval move operation (called after debounce delay)."""
+
+ timeline_data = session.timeline_data
+ with timeline_data["keyframe_update_lock"]:
+ # Check if this move is still the latest one
+ if interval_id not in timeline_data["pending_keyframe_moves"]:
+ return # Move was cancelled
+
+ if timeline_data["pending_keyframe_moves"][interval_id] != new_interval:
+ return # A newer move superseded this one
+
+ # Remove from pending
+ del timeline_data["pending_keyframe_moves"][interval_id]
+ if interval_id in timeline_data["keyframe_move_timers"]:
+ del timeline_data["keyframe_move_timers"][interval_id]
+
+ # Now execute the actual move
+ if interval_id not in timeline_data["intervals"]:
+ return
+ interval_data = timeline_data["intervals"][interval_id]
+ if not interval_data:
+ return
+
+ # if the interval did not move, don't do anything
+ if (
+ interval_data["start_frame_idx"] == new_interval[0]
+ and interval_data["end_frame_idx"] == new_interval[1]
+ ):
+ return
+
+ track_id = interval_data["track_id"]
+ constraint_type = timeline_data["tracks"][track_id]["name"]
+ cur_range = (
+ interval_data["start_frame_idx"],
+ interval_data["end_frame_idx"],
+ )
+
+ # Remove constraint at old frame
+ remove_constraint_callback(
+ interval_id,
+ constraint_type,
+ cur_range,
+ verbose=False,
+ )
+ # Add constraint at new frame
+ add_constraint_callback(
+ interval_id,
+ constraint_type,
+ new_interval,
+ verbose=False,
+ )
+
+ # update our data
+ interval_data["start_frame_idx"] = new_interval[0]
+ interval_data["end_frame_idx"] = new_interval[1]
+
+ # Schedule path update only after user stops dragging (no move for 300ms).
+ if constraint_type == "2D Root":
+ _schedule_dense_path_after_release(session)
+
+ @client.timeline.on_interval_delete
+ def handle_interval_delete(interval_id: str):
+ """Called when an interval is deleted."""
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ with session.timeline_data["keyframe_update_lock"]:
+ if interval_id not in session.timeline_data["intervals"]:
+ return
+ interval_data = session.timeline_data["intervals"][interval_id]
+ track_id = interval_data["track_id"]
+ constraint_type = session.timeline_data["tracks"][track_id]["name"]
+ remove_constraint_callback(
+ interval_id,
+ constraint_type,
+ (
+ interval_data["start_frame_idx"],
+ interval_data["end_frame_idx"],
+ ),
+ verbose=False,
+ )
+ del session.timeline_data["intervals"][interval_id]
+ if constraint_type == "2D Root" and session.constraints["2D Root"].dense_path:
+ motion = list(session.motions.values())[0]
+ _update_dense_path(motion, session)
+
+ @gui_snap_to_constraint_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+
+ target_character_motion = list(session.motions.values())[0]
+ frame_idx = session.frame_idx
+
+ if frame_idx >= target_character_motion.length:
+ # frame idx larger than the motion, could not snap
+ return
+
+ for constraint_name in ["Full-Body", "End-Effectors"]:
+ if (
+ constraint_name in session.constraints
+ and frame_idx in session.constraints[constraint_name].keyframes
+ ):
+ pos = session.constraints[constraint_name].keyframes[frame_idx]["joints_pos"]
+ rot = session.constraints[constraint_name].keyframes[frame_idx]["joints_rot"]
+
+ # update the full joints_pos of the character to match the constraints
+ target_character_motion.update_pose_at_frame(
+ frame_idx,
+ joints_pos=pos,
+ joints_rot=rot,
+ )
+ target_character_motion.set_frame(frame_idx)
+ return # motion already fully changed
+
+ if "2D Root" in session.constraints and frame_idx in session.constraints["2D Root"].keyframes:
+ # update only the root position
+ new_root_pos = session.constraints["2D Root"].keyframes[frame_idx]
+ old_root_pos = target_character_motion.get_projected_root_pos(frame_idx)
+ root_diff = new_root_pos - old_root_pos
+ root_diff[1] = 0.0 # don't change height
+
+ new_joints_pos = (
+ target_character_motion.joints_pos[frame_idx]
+ + to_torch(
+ root_diff,
+ device=target_character_motion.joints_pos.device,
+ dtype=target_character_motion.joints_pos.dtype,
+ )[None]
+ )
+ rot = target_character_motion.joints_rot[frame_idx]
+
+ target_character_motion.update_pose_at_frame(
+ frame_idx,
+ joints_pos=new_joints_pos,
+ joints_rot=rot,
+ )
+ target_character_motion.set_frame(frame_idx)
+
+ @gui_clear_all_constraints_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+ with session.timeline_data["keyframe_update_lock"]:
+ # use the lock here to wait for any constraint updates to finish
+ for constraint in list(session.constraints.values()):
+ constraint.clear()
+ client.timeline.clear_keyframes()
+ client.timeline.clear_intervals()
+ if gui_dense_path_checkbox.value:
+ gui_dense_path_checkbox.value = False
+ if "2D Root" in session.constraints:
+ session.constraints["2D Root"].set_dense_path(False)
+
+ # generation callback
+ @gui_generate_button.on_click
+ def _(event: viser.GuiEvent) -> None:
+ event_client = event.client
+ session = get_active_session(event_client)
+ if session is None:
+ return
+
+ generating_notif = event_client.add_notification(
+ title="Generating motion...",
+ body="Generating motions for the given prompt!",
+ loading=True,
+ with_close_button=False,
+ )
+ gui_generate_button.disabled = True
+ client.timeline.disable_constraints()
+
+ num_samples = gui_num_samples_slider.value
+ timeline = session.client.timeline
+
+ # sort them to avoid issues:
+ prompt_values = sorted([x for x in timeline._prompts.values()], key=lambda x: x.start_frame)
+
+ texts = [x.text for x in prompt_values]
+ num_frames = compute_prompt_num_frames(prompt_values)
+
+ # compute the total duration
+ total_nb_frames = sum(num_frames)
+ total_duration = total_nb_frames / session.model_fps
+
+ # update just in case
+ set_new_duration(client_id, total_duration)
+
+ transitions_parameters = {
+ "num_transition_frames": gui_num_transition_frames_slider.value,
+ }
+
+ # G1: postprocessing is disabled (does not work well for this model).
+ postprocess_parameters = {
+ "post_processing": (False if "g1" in session.model_name else gui_postprocess_checkbox.value),
+ "root_margin": gui_root_margin.value,
+ }
+ try:
+ demo.generate(
+ event_client,
+ texts,
+ num_frames,
+ num_samples,
+ gui_seed.value,
+ gui_diffusion_steps_slider.value,
+ cfg_weight=[
+ gui_cfg_text_weight_slider.value,
+ gui_cfg_constraint_weight_slider.value,
+ ],
+ cfg_type="separated" if gui_cfg_checkbox.value else "nocfg",
+ postprocess_parameters=postprocess_parameters,
+ transitions_parameters=transitions_parameters,
+ real_robot_rotations=gui_real_robot_rotations_checkbox.value,
+ )
+ session.max_frame_idx = int(session.cur_duration * session.model_fps - 1)
+ session.max_frame_idx = int(session.cur_duration * session.model_fps) - 1
+ if session.frame_idx > session.max_frame_idx:
+ session.frame_idx = session.max_frame_idx
+
+ if num_samples > 1:
+ # add mesh selector to choose character to commit
+ def commit_motion(event: viser.GuiEvent) -> None:
+ target = event.target
+ commit_name = target.name.split("/")[1] # e.g. /character0/simple_skinned
+ print(f"Committing motion for character: {commit_name}")
+ # delete non-selected motions
+ new_motion_kwargs = None
+ for character_name, motion in session.motions.items():
+ if character_name == commit_name:
+ new_motion_kwargs = {
+ "skeleton": session.skeleton,
+ "joints_rot": motion.joints_rot,
+ "foot_contacts": motion.foot_contacts,
+ }
+ root_x_offset = motion.joints_pos[0, session.skeleton.root_idx, 0]
+ new_joints_pos = motion.joints_pos.clone()
+ new_joints_pos[..., 0] -= root_x_offset
+ new_motion_kwargs["joints_pos"] = new_joints_pos
+ break
+ # clear and re-add the selected motion
+ demo.clear_motions(event_client.client_id)
+ demo.add_character_motion(event_client, **new_motion_kwargs)
+ gui_edit_constraint_button.disabled = False
+ gui_generate_button.disabled = False
+ gui_snap_to_constraint_button.disabled = False
+ client.timeline.enable_constraints()
+ gui_generate_button.label = "Generate"
+ gui_save_example_button.disabled = False
+ gui_save_motion_button.disabled = False
+ gui_download_button.disabled = False
+ gui_save_constraints_button.disabled = False
+ gui_load_example_button.disabled = False
+
+ for motion in session.motions.values():
+ char = motion.character
+ character_name = char.name # e.g. "character0"
+ if char.skinned_mesh is not None:
+ char.skinned_mesh.on_click(commit_motion)
+ elif char.g1_mesh_rig is not None:
+ # Register click on every part so any part can be clicked,
+ # and use highlight_group so the whole robot highlights together.
+ for handle in char.g1_mesh_rig.mesh_handles:
+ handle.on_click(commit_motion, highlight_group=character_name)
+
+ gui_edit_constraint_button.disabled = True
+ gui_generate_button.disabled = True
+ gui_snap_to_constraint_button.disabled = True
+ gui_generate_button.label = "Choose Sample Before Generating"
+ gui_save_example_button.disabled = True
+ gui_save_motion_button.disabled = True
+ gui_download_button.disabled = True
+ gui_save_constraints_button.disabled = True
+ gui_load_example_button.disabled = True
+ else:
+ gui_edit_constraint_button.disabled = False
+ gui_generate_button.disabled = False
+ gui_snap_to_constraint_button.disabled = False
+ client.timeline.enable_constraints()
+
+ generating_notif.title = "Motion generation finished!"
+ generating_notif.body = "Motions have been generated successfully for the given prompt."
+ if num_samples > 1:
+ generating_notif.body += " Now choose which sample to commit."
+ generating_notif.loading = False
+ generating_notif.with_close_button = True
+ generating_notif.auto_close_seconds = 5.0
+ generating_notif.color = "green"
+
+ # put the motion at zero
+ demo.set_frame(client_id, 0)
+
+ except Exception as e:
+ import traceback
+
+ traceback.print_exc()
+ print(f"Error during generation for client {event_client.client_id}: {e}")
+ # Re-enable buttons and notify the user
+ if event_client.client_id in demo.client_sessions:
+ session = demo.client_sessions[event_client.client_id]
+ gui_generate_button.disabled = False
+ gui_load_example_button.disabled = False
+ gui_save_example_button.disabled = False
+ gui_save_motion_button.disabled = False
+ gui_download_button.disabled = False
+ try:
+ event_client.add_notification(
+ title="Generation failed!",
+ body=f"Error: {str(e)}",
+ auto_close_seconds=5.0,
+ color="red",
+ )
+ except Exception:
+ pass
+ demo.check_cuda_health()
+
+ #
+ # Visualization settings
+ #
+ with tab_group.add_tab("Visualize", viser.Icon.EYE):
+ with client.gui.add_folder("Playback", expand_by_default=True):
+ gui_model_fps = client.gui.add_number("Model FPS", initial_value=model_fps, disabled=True)
+ gui_playback_speed_buttons = client.gui.add_button_group(
+ "Playback Speed",
+ options=[
+ "0.5x",
+ "1x",
+ "2x",
+ ],
+ )
+ gui_playback_speed_buttons.value = "1x"
+
+ @client.timeline.on_frame_change
+ def handle_timeline_frame_change(new_frame_idx: int):
+ """Update the frame when the user clicks on the timeline."""
+ demo.set_frame(client_id, new_frame_idx, update_timeline=False)
+ session = demo.client_sessions.get(client_id)
+ if session is not None:
+ if session.edit_mode and session.motions:
+ motion = list(session.motions.values())[0]
+ snapshot_frame_idx = min(session.frame_idx, motion.length - 1)
+ ensure_edit_snapshot(session, motion, snapshot_frame_idx)
+ update_snap_to_constraint_button(session)
+
+ @client.timeline.on_prompt_add
+ async def _on_add(
+ prompt_id: str,
+ start_frame: int,
+ end_frame: int,
+ text: str,
+ color: tuple[int, int, int] | None,
+ ) -> None:
+ update_duration_auto()
+
+ @client.timeline.on_prompt_update
+ async def _on_update(prompt_id: str, new_text: str) -> None:
+ update_duration_auto()
+
+ @client.timeline.on_prompt_resize
+ async def _on_resize(prompt_id: str, new_start: int, new_end: int) -> None:
+ update_duration_auto()
+
+ @client.timeline.on_prompt_move
+ async def _on_move(prompt_id: str, new_start: int, new_end: int) -> None:
+ update_duration_auto()
+
+ @client.timeline.on_prompt_delete
+ async def _on_delete(prompt_id: str) -> None:
+ update_duration_auto()
+
+ def play_pause_button_callback(session: ClientSession):
+ session.playing = not session.playing
+
+ def next_frame_callback(session: ClientSession):
+ if session.frame_idx < session.max_frame_idx:
+ session.frame_idx += 1
+ if session.frame_idx == session.max_frame_idx:
+ pass
+ demo.set_frame(client_id, session.frame_idx)
+
+ def prev_frame_callback(session: ClientSession):
+ if session.frame_idx > 0:
+ session.frame_idx -= 1
+ if session.frame_idx == 0:
+ pass
+ demo.set_frame(client_id, session.frame_idx)
+
+ @gui_playback_speed_buttons.on_click
+ def _(_) -> None:
+ if not demo.client_active(client_id):
+ return
+ speed_map = {
+ "0.5x": 0.5,
+ "1x": 1.0,
+ "2x": 2.0,
+ }
+ session = demo.client_sessions[client_id]
+ session.playback_speed = speed_map[gui_playback_speed_buttons.value]
+
+ with client.gui.add_folder("Body options", expand_by_default=True):
+ gui_viz_skinned_mesh_checkbox = client.gui.add_checkbox("Show Mesh", initial_value=True)
+ gui_viz_skinned_mesh_opacity_slider = client.gui.add_slider(
+ "Mesh Opacity", min=0.0, max=1.0, step=0.01, initial_value=1.0
+ )
+ gui_viz_skeleton_checkbox = client.gui.add_checkbox("Show Skeleton", initial_value=False)
+ gui_viz_foot_contacts_checkbox = client.gui.add_checkbox("Show Foot Contacts", initial_value=False)
+ gui_viz_foot_contacts_checkbox.visible = gui_viz_skeleton_checkbox.value
+ with client.gui.add_folder("Camera options", expand_by_default=True):
+ gui_camera_fov_slider = client.gui.add_slider(
+ "Camera FOV (deg)",
+ min=30.0,
+ max=90.0,
+ step=1.0,
+ initial_value=45.0,
+ )
+ client.camera.fov = np.deg2rad(gui_camera_fov_slider.value)
+ with client.gui.add_folder("Interface options", expand_by_default=True):
+ gui_show_timeline_checkbox = client.gui.add_checkbox(
+ "Show Timeline",
+ initial_value=True,
+ )
+ gui_show_constraint_tracks_checkbox = client.gui.add_checkbox(
+ "Show Constraint tracks",
+ initial_value=True,
+ )
+ gui_show_constraint_labels_checkbox = client.gui.add_checkbox(
+ "Show Constraint labels",
+ initial_value=True,
+ )
+ gui_show_starting_direction_checkbox = client.gui.add_checkbox(
+ "Show Starting Direction",
+ initial_value=True,
+ )
+ gui_dark_mode_checkbox = client.gui.add_checkbox(
+ "Dark Mode",
+ initial_value=False, # Default to light mode
+ )
+ gui_show_constraint_tracks_checkbox.visible = gui_show_timeline_checkbox.value
+ demo.set_start_direction_visible(client_id, gui_show_starting_direction_checkbox.value)
+
+ @gui_dark_mode_checkbox.on_update
+ def _(_):
+ # Apply the theme using configure_theme (pass uuid so titlebar toggle stays)
+ demo.configure_theme(
+ client,
+ gui_dark_mode_checkbox.value,
+ titlebar_dark_mode_checkbox_uuid=gui_dark_mode_checkbox.uuid,
+ )
+ session = demo.client_sessions[client.client_id]
+ for motion in session.motions.values():
+ motion.character.change_theme(gui_dark_mode_checkbox.value)
+
+ # Show dark mode toggle in titlebar (right of Github), hide sidebar checkbox
+ demo.configure_theme(
+ client,
+ gui_dark_mode_checkbox.value,
+ titlebar_dark_mode_checkbox_uuid=gui_dark_mode_checkbox.uuid,
+ )
+ gui_dark_mode_checkbox.visible = False
+
+ @gui_show_constraint_labels_checkbox.on_update
+ def _(_):
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ for constraint in session.constraints.values():
+ constraint.set_label_visibility(gui_show_constraint_labels_checkbox.value)
+
+ @gui_show_timeline_checkbox.on_update
+ def _(_):
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ session.client.timeline.set_visible(gui_show_timeline_checkbox.value)
+ gui_show_constraint_tracks_checkbox.visible = gui_show_timeline_checkbox.value
+ if gui_show_timeline_checkbox.value:
+ demo.set_constraint_tracks_visible(session, gui_show_constraint_tracks_checkbox.value)
+
+ @gui_show_constraint_tracks_checkbox.on_update
+ def _(_):
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ demo.set_constraint_tracks_visible(session, gui_show_constraint_tracks_checkbox.value)
+
+ @gui_show_starting_direction_checkbox.on_update
+ def _(_):
+ if not demo.client_active(client_id):
+ return
+ demo.set_start_direction_visible(client_id, gui_show_starting_direction_checkbox.value)
+
+ @gui_viz_skeleton_checkbox.on_update
+ def _(_) -> None:
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ gui_viz_foot_contacts_checkbox.visible = gui_viz_skeleton_checkbox.value
+ if not gui_viz_skeleton_checkbox.value:
+ gui_viz_foot_contacts_checkbox.value = False
+ for motion in session.motions.values():
+ motion.character.set_skeleton_visibility(gui_viz_skeleton_checkbox.value)
+
+ @gui_viz_foot_contacts_checkbox.on_update
+ def _(_) -> None:
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ for motion in session.motions.values():
+ motion.character.set_show_foot_contacts(
+ gui_viz_foot_contacts_checkbox.value, frame_idx=motion.cur_frame_idx
+ )
+
+ @gui_viz_skinned_mesh_checkbox.on_update
+ def _(_) -> None:
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ for motion in session.motions.values():
+ motion.character.set_skinned_mesh_visibility(gui_viz_skinned_mesh_checkbox.value)
+
+ @gui_viz_skinned_mesh_opacity_slider.on_update
+ def _(_) -> None:
+ if not demo.client_active(client_id):
+ return
+ session = demo.client_sessions[client_id]
+ for motion in session.motions.values():
+ motion.character.set_skinned_mesh_opacity(gui_viz_skinned_mesh_opacity_slider.value)
+
+ @gui_camera_fov_slider.on_update
+ def _(_) -> None:
+ if not demo.client_active(client_id):
+ return
+ client.camera.fov = np.deg2rad(gui_camera_fov_slider.value)
+
+ #
+
+ # Instructions tab
+ #
+ with tab_group.add_tab("Instructions", viser.Icon.INFO_CIRCLE):
+ client.gui.add_markdown(DEMO_UI_INSTRUCTIONS_TAB_MD)
+
+ #
+ # Keyboard events
+ #
+ space_pressed = [False]
+
+ @client.scene.on_keyboard_event("keydown", debounce_ms=100)
+ def handle_key(event: viser.KeyboardEvent) -> None:
+ # Check if client session still exists
+ if client_id not in demo.client_sessions:
+ return
+
+ session = demo.client_sessions[client_id]
+
+ if event.event_type == "keyup":
+ if event.key == " ":
+ space_pressed[0] = False
+ return
+
+ # Space bar: only toggle on FIRST press
+ if event.key == " ":
+ if not space_pressed[0]:
+ space_pressed[0] = True
+ play_pause_button_callback(session)
+ return
+
+ # Handle arrow keys: frame navigation (fast OS repeat with 50ms debounce).
+ elif event.key == "ArrowLeft":
+ prev_frame_callback(session)
+ elif event.key == "ArrowRight":
+ next_frame_callback(session)
+
+ gui_elements = GuiElements(
+ gui_play_pause_button=gui_play_pause_button,
+ gui_next_frame_button=gui_next_frame_button,
+ gui_prev_frame_button=gui_prev_frame_button,
+ gui_generate_button=gui_generate_button,
+ gui_model_fps=gui_model_fps,
+ gui_timeline=gui_timeline,
+ gui_viz_skeleton_checkbox=gui_viz_skeleton_checkbox,
+ gui_viz_foot_contacts_checkbox=gui_viz_foot_contacts_checkbox,
+ gui_viz_skinned_mesh_checkbox=gui_viz_skinned_mesh_checkbox,
+ gui_viz_skinned_mesh_opacity_slider=gui_viz_skinned_mesh_opacity_slider,
+ gui_camera_fov_slider=gui_camera_fov_slider,
+ gui_duration_slider=gui_duration_slider,
+ gui_num_samples_slider=gui_num_samples_slider,
+ gui_cfg_checkbox=gui_cfg_checkbox,
+ gui_cfg_text_weight_slider=gui_cfg_text_weight_slider,
+ gui_cfg_constraint_weight_slider=gui_cfg_constraint_weight_slider,
+ gui_diffusion_steps_slider=gui_diffusion_steps_slider,
+ gui_seed=gui_seed,
+ gui_postprocess_checkbox=gui_postprocess_checkbox,
+ gui_root_margin=gui_root_margin,
+ gui_real_robot_rotations_checkbox=gui_real_robot_rotations_checkbox,
+ gui_dark_mode_checkbox=gui_dark_mode_checkbox,
+ gui_use_soma_layer_checkbox=gui_use_soma_layer_checkbox,
+ )
+ return (
+ gui_elements,
+ timeline_tracks,
+ example_dict,
+ gui_examples_dropdown,
+ gui_save_example_path_text,
+ gui_model_selector,
+ )
diff --git a/kimodo/exports/__init__.py b/kimodo/exports/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..57ce23ddceef985d59b206f2f3fd0f14ee36ca69
--- /dev/null
+++ b/kimodo/exports/__init__.py
@@ -0,0 +1,65 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Export utilities: MuJoCo, BVH, SMPLX/AMASS, and motion I/O helpers."""
+
+from .bvh import bvh_to_kimodo_motion, motion_to_bvh_bytes, read_bvh_frame_time_seconds, save_motion_bvh
+from .motion_convert_lib import convert_motion_files
+from .motion_formats import (
+ infer_npz_kind,
+ infer_source_format_from_path,
+ infer_target_format_from_path,
+ resolve_source_fps,
+)
+from .motion_io import (
+ KIMODO_CONVERT_TARGET_FPS,
+ amass_npz_to_bytes,
+ complete_motion_dict,
+ g1_csv_to_bytes,
+ kimodo_npz_to_bytes,
+ load_amass_npz,
+ load_g1_csv,
+ load_kimodo_npz,
+ load_kimodo_npz_as_torch,
+ load_motion_file,
+ motion_dict_to_numpy,
+ save_kimodo_npz,
+ save_kimodo_npz_at_target_fps,
+)
+from .mujoco import MujocoQposConverter, apply_g1_real_robot_projection
+from .smplx import (
+ AMASSConverter,
+ amass_npz_to_kimodo_motion,
+ get_amass_parameters,
+ kimodo_y_up_to_amass_coord_rotation_matrix,
+)
+
+__all__ = [
+ "AMASSConverter",
+ "KIMODO_CONVERT_TARGET_FPS",
+ "MujocoQposConverter",
+ "amass_npz_to_bytes",
+ "amass_npz_to_kimodo_motion",
+ "apply_g1_real_robot_projection",
+ "bvh_to_kimodo_motion",
+ "complete_motion_dict",
+ "convert_motion_files",
+ "g1_csv_to_bytes",
+ "get_amass_parameters",
+ "infer_npz_kind",
+ "infer_source_format_from_path",
+ "infer_target_format_from_path",
+ "kimodo_npz_to_bytes",
+ "kimodo_y_up_to_amass_coord_rotation_matrix",
+ "load_amass_npz",
+ "load_g1_csv",
+ "load_kimodo_npz",
+ "load_kimodo_npz_as_torch",
+ "load_motion_file",
+ "motion_dict_to_numpy",
+ "motion_to_bvh_bytes",
+ "read_bvh_frame_time_seconds",
+ "resolve_source_fps",
+ "save_kimodo_npz",
+ "save_kimodo_npz_at_target_fps",
+ "save_motion_bvh",
+]
diff --git a/kimodo/exports/bvh.py b/kimodo/exports/bvh.py
new file mode 100644
index 0000000000000000000000000000000000000000..d9625cca562aa0e6de2f2a8020c9e8cc1b1c6d4e
--- /dev/null
+++ b/kimodo/exports/bvh.py
@@ -0,0 +1,298 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Export utilities for converting internal motion representations into common file formats.
+
+This module is intended to hold lightweight serialization / export helpers that can be reused
+outside of interactive demos.
+"""
+
+import os
+import tempfile
+from pathlib import Path
+from typing import Tuple, Union
+
+import numpy as np
+import torch
+
+from kimodo.geometry import matrix_to_quaternion as _matrix_to_quaternion
+
+
+def _strip_end_site_blocks(bvh_text: str) -> str:
+ """Remove all 'End Site { ... }' blocks from BVH text so output matches original format.
+
+ bvhio adds an End Site for every leaf joint when writing; we do not set EndSite on joints, so we
+ post-process the string to remove these blocks for Blender/original compatibility.
+ """
+ lines = bvh_text.splitlines(keepends=True)
+ result = []
+ i = 0
+ while i < len(lines):
+ line = lines[i]
+ if "End Site" in line:
+ # Skip this line and the following block { ... }; brace-count to find closing }
+ i += 1
+ if i < len(lines) and "{" in lines[i]:
+ i += 1
+ depth = 1
+ while i < len(lines) and depth > 0:
+ if "{" in lines[i]:
+ depth += 1
+ if "}" in lines[i]:
+ depth -= 1
+ i += 1
+ continue
+ result.append(line)
+ i += 1
+ return "".join(result)
+
+
+def _coerce_batch(name: str, x: torch.Tensor, *, expected_ndim: int) -> torch.Tensor:
+ """Coerce (T, ...) or (1, T, ...) into (T, ...)."""
+ if x.ndim == expected_ndim:
+ return x
+ if x.ndim == expected_ndim + 1:
+ if int(x.shape[0]) != 1:
+ raise ValueError(
+ f"{name} has batch dimension B={int(x.shape[0])}, but BVH export " "only supports a single clip (B==1)."
+ )
+ return x[0]
+ raise ValueError(f"{name} must have shape (T, ...) or (1, T, ...); got {tuple(x.shape)}")
+
+
+def motion_to_bvh(
+ local_rot_mats: torch.Tensor,
+ root_positions: torch.Tensor,
+ *,
+ skeleton,
+ fps: float,
+ standard_tpose: bool = False,
+) -> str:
+ """Convert local rotations and root positions to BVH format; return UTF-8 string.
+
+ Args:
+ local_rot_mats: (T, J, 3, 3) or (1, T, J, 3, 3) local rotation matrices.
+ root_positions: (T, 3) or (1, T, 3) root joint positions (e.g. from posed joints).
+ skeleton: Skeleton with bone_order_names, bvh_neutral_joints, etc.
+ fps: Frames per second for the motion.
+ standard_tpose: If True, export with the rest pose being the standard T-pose rather than the rest pose consistent with the BONES-SEED dataset.
+ Notes:
+ BVH is plain-text. Root is named "Root" with ZYX rotation order; leaf joints
+ have no End Site block.
+ """
+ try:
+ import bvhio # type: ignore[import-not-found]
+ import glm # type: ignore[import-not-found]
+ from SpatialTransform import Pose # type: ignore[import-not-found]
+ except Exception as e: # pragma: no cover
+ raise ImportError(
+ "BVH export requires `bvhio` (and its deps `PyGLM` + `SpatialTransform`). "
+ "Install with: `pip install bvhio`."
+ ) from e
+
+ local_rot_mats = local_rot_mats.detach()
+ root_positions = root_positions.detach()
+ # SOMA: accept either somaskel30 (convert to 77) or somaskel77 (use as-is)
+ if skeleton.name == "somaskel30":
+ local_rot_mats = skeleton.to_SOMASkeleton77(local_rot_mats)
+ skeleton = skeleton.somaskel77
+
+ if standard_tpose:
+ neutral = skeleton.neutral_joints.detach().cpu().numpy()
+ else:
+ # transform local rots to the original rest pose consistent with the BONES-SEED dataset
+ local_rot_mats, _ = skeleton.from_standard_tpose(local_rot_mats)
+ neutral = skeleton.bvh_neutral_joints.detach().cpu().numpy()
+
+ joint_names = list(skeleton.bone_order_names)
+ parents = skeleton.joint_parents.detach().cpu().numpy().astype(int)
+ root_idx = int(skeleton.root_idx)
+
+ local_rot_mats = _coerce_batch("local_rot_mats", local_rot_mats, expected_ndim=4)
+ T, J = local_rot_mats.shape[:2]
+ q_wxyz = _matrix_to_quaternion(local_rot_mats).detach().cpu().numpy() # [T, J, 4]
+
+ root_xyz = _coerce_batch("root_positions", root_positions, expected_ndim=2)
+ root_xyz = root_xyz.cpu().numpy() # [T, 3]
+
+ # Build BVH hierarchy: Root (wrapper at origin) -> Hips (pelvis with offset in meters) -> ...
+ # Offsets are in meters to match the original format.
+ children: dict[int, list[int]] = {i: [] for i in range(J)}
+ for i, p in enumerate(parents):
+ if p >= 0:
+ children[int(p)].append(int(i))
+
+ _ROOT_CHANNELS = [
+ "Xposition",
+ "Yposition",
+ "Zposition",
+ "Zrotation",
+ "Yrotation",
+ "Xrotation",
+ ]
+ _JOINT_CHANNELS = ["Zrotation", "Yrotation", "Xrotation"]
+
+ # Scale from meters to centimeters (match original SEED data BVH scale).
+ neutral = neutral * 100
+ root_xyz = root_xyz * 100
+
+ # Hips offset from Root: use skeleton neutral; if root is at origin (zeros), use a
+ # nominal pelvis height so the hierarchy is non-degenerate in Blender.
+ hips_offset = neutral[root_idx]
+ if (hips_offset == 0).all():
+ hips_offset = np.array([0.0, 100.0, 0.0], dtype=neutral.dtype) # 1 m in cm
+
+ def _make_joint(i: int) -> "bvhio.BvhJoint":
+ name = joint_names[i]
+ j = bvhio.BvhJoint(name, offset=glm.vec3(0, 0, 0))
+ if i == root_idx:
+ # Hips: offset from Root (origin) in cm
+ off = hips_offset
+ j.Offset = glm.vec3(float(off[0]), float(off[1]), float(off[2]))
+ j.Channels = _ROOT_CHANNELS.copy()
+ else:
+ p = int(parents[i])
+ off = neutral[i] - neutral[p]
+ j.Offset = glm.vec3(float(off[0]), float(off[1]), float(off[2]))
+ j.Channels = _JOINT_CHANNELS.copy()
+
+ for c in children[i]:
+ j.Children.append(_make_joint(c))
+ return j
+
+ # Wrapper Root at origin; single child is Hips (skeleton root).
+ root_wrapper = bvhio.BvhJoint("Root", offset=glm.vec3(0.0, 0.0, 0.0))
+ root_wrapper.Channels = _ROOT_CHANNELS.copy()
+ root_wrapper.Children.append(_make_joint(root_idx))
+ root_joint = root_wrapper
+
+ # Populate keyframes: Root = identity/zero, Hips = root motion, others = local rotation.
+ bvh_layout = root_joint.layout()
+ name_to_id = {n: idx for idx, n in enumerate(joint_names)}
+ ordered_joint_ids = []
+ for bj, _, _ in bvh_layout:
+ if bj.Name == "Root":
+ ordered_joint_ids.append(None)
+ else:
+ ordered_joint_ids.append(name_to_id[bj.Name])
+
+ bvh_joints = [bj for bj, _, _ in bvh_layout]
+ for bj in bvh_joints:
+ bj.Keyframes = [None] * T # type: ignore[list-item]
+
+ identity_quat = glm.quat(1.0, 0.0, 0.0, 0.0)
+ zero_vec = glm.vec3(0.0, 0.0, 0.0)
+ for t in range(T):
+ for bj, jid in zip(bvh_joints, ordered_joint_ids):
+ if jid is None:
+ position = zero_vec
+ rotation = identity_quat
+ elif jid == root_idx:
+ pos = root_xyz[t]
+ position = glm.vec3(float(pos[0]), float(pos[1]), float(pos[2]))
+ qw, qx, qy, qz = q_wxyz[t, jid]
+ rotation = glm.quat(float(qw), float(qx), float(qy), float(qz))
+ else:
+ position = zero_vec
+ qw, qx, qy, qz = q_wxyz[t, jid]
+ rotation = glm.quat(float(qw), float(qx), float(qy), float(qz))
+ bj.Keyframes[t] = Pose(position, rotation) # type: ignore[index]
+
+ container = bvhio.BvhContainer(root_joint, frameCount=T, frameTime=1.0 / float(fps))
+ with tempfile.NamedTemporaryFile(mode="w", suffix=".bvh", delete=False, encoding="utf-8") as f:
+ tmp_path = f.name
+ try:
+ bvhio.writeBvh(tmp_path, container, percision=6)
+ bvh_text = Path(tmp_path).read_text(encoding="utf-8")
+ return _strip_end_site_blocks(bvh_text)
+ finally:
+ try:
+ os.remove(tmp_path)
+ except Exception:
+ pass
+
+
+def motion_to_bvh_bytes(
+ local_rot_mats: torch.Tensor,
+ root_positions: torch.Tensor,
+ *,
+ skeleton,
+ fps: float,
+ standard_tpose: bool = False,
+) -> bytes:
+ """Convert local rotations and root positions to BVH bytes (UTF-8).
+
+ Convenience wrapper around :func:`motion_to_bvh`.
+ """
+ return motion_to_bvh(
+ local_rot_mats,
+ root_positions,
+ skeleton=skeleton,
+ fps=fps,
+ standard_tpose=standard_tpose,
+ ).encode("utf-8")
+
+
+def save_motion_bvh(
+ path: Union[str, Path],
+ local_rot_mats: torch.Tensor,
+ root_positions: torch.Tensor,
+ *,
+ skeleton,
+ fps: float,
+ standard_tpose: bool = False,
+) -> None:
+ """Write local rotations and root positions to a BVH file at the given path."""
+ Path(path).write_text(
+ motion_to_bvh(local_rot_mats, root_positions, skeleton=skeleton, fps=fps, standard_tpose=standard_tpose),
+ encoding="utf-8",
+ )
+
+
+def read_bvh_frame_time_seconds(path: Union[str, Path]) -> float:
+ """Read ``Frame Time`` from a BVH file (seconds per frame)."""
+ with open(path, encoding="utf-8") as f:
+ for line in f:
+ if "Frame Time:" in line:
+ parts = line.split()
+ return float(parts[-1])
+ raise ValueError(f"Could not find 'Frame Time:' in {path}")
+
+
+def bvh_to_kimodo_motion(
+ path: Union[str, Path],
+ skeleton=None,
+ *,
+ standard_tpose: bool = False,
+) -> Tuple:
+ """Load a Kimodo-style SOMA BVH into a Kimodo motion dict.
+
+ Expects the same hierarchy as :func:`save_motion_bvh` (``Root`` wrapper + SOMA77 joints).
+ The frame rate is always read from the BVH ``Frame Time`` header. Callers
+ that need a different playback rate should resample the returned motion dict
+ (see :func:`~kimodo.exports.motion_io.resample_motion_dict_to_kimodo_fps`).
+
+ Returns:
+ ``(motion_dict, source_fps)`` where ``source_fps`` is the native BVH
+ frame rate read from the file header.
+ """
+ from kimodo.exports.motion_io import complete_motion_dict
+ from kimodo.skeleton.bvh import parse_bvh_motion
+ from kimodo.skeleton.registry import build_skeleton
+
+ if skeleton is None:
+ skeleton = build_skeleton(77)
+ device = skeleton.neutral_joints.device
+
+ local_rot_mats, root_trans, bvh_fps = parse_bvh_motion(str(path))
+ local_rot_mats = local_rot_mats.to(device=device)
+ root_trans = root_trans.to(device=device)
+
+ if int(local_rot_mats.shape[1]) != int(skeleton.nbjoints):
+ raise ValueError(
+ f"BVH has {local_rot_mats.shape[1]} joints but skeleton has {skeleton.nbjoints}; "
+ "use a Kimodo-exported SOMA BVH or matching skeleton."
+ )
+ if not standard_tpose:
+ local_rot_mats, _ = skeleton.to_standard_tpose(local_rot_mats)
+
+ return complete_motion_dict(local_rot_mats, root_trans, skeleton, float(bvh_fps)), bvh_fps
diff --git a/kimodo/exports/motion_convert_lib.py b/kimodo/exports/motion_convert_lib.py
new file mode 100644
index 0000000000000000000000000000000000000000..384ef20d3b11c7e4c9a65bf4b02bc53d23aad605
--- /dev/null
+++ b/kimodo/exports/motion_convert_lib.py
@@ -0,0 +1,159 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Library API for converting between Kimodo NPZ, AMASS NPZ, SOMA BVH, and G1 MuJoCo CSV."""
+
+from __future__ import annotations
+
+import warnings
+
+import numpy as np
+
+from kimodo.exports.bvh import bvh_to_kimodo_motion, save_motion_bvh
+from kimodo.exports.motion_formats import (
+ infer_source_format_from_path,
+ infer_target_format_from_path,
+ resolve_source_fps,
+)
+from kimodo.exports.motion_io import (
+ load_amass_npz,
+ load_g1_csv,
+ load_kimodo_npz_as_torch,
+ save_kimodo_npz_at_target_fps,
+)
+from kimodo.exports.mujoco import MujocoQposConverter
+from kimodo.exports.smplx import AMASSConverter
+from kimodo.skeleton.registry import build_skeleton
+
+
+def convert_motion_files(
+ input_path: str,
+ output_path: str,
+ *,
+ from_fmt: str | None = None,
+ to_fmt: str | None = None,
+ source_fps: float | None = None,
+ z_up: bool = True,
+ mujoco_rest_zero: bool = False,
+ bvh_standard_tpose: bool = False,
+) -> None:
+ """Convert a motion file between Kimodo-supported formats.
+
+ Supported pairs (hub-and-spoke through Kimodo NPZ):
+
+ - amass <-> kimodo
+ - soma-bvh <-> kimodo
+ - g1-csv <-> kimodo
+
+ Args:
+ input_path: Source file (``.npz``, ``.bvh``, or ``.csv``).
+ output_path: Destination file.
+ from_fmt: Source format; inferred from extension/contents when ``None``.
+ to_fmt: Target format; inferred from extension when ``None``.
+ source_fps: Source motion frame rate (Hz). If provided, trusted as-is.
+ If ``None``, auto-detected from BVH ``Frame Time``, AMASS
+ ``mocap_frame_rate``, or default 30.
+ z_up: For AMASS conversions, apply the Z-up <-> Kimodo Y-up transform.
+ mujoco_rest_zero: For G1 CSV, joint angles relative to MuJoCo rest pose.
+ bvh_standard_tpose: If input or output is BVH: the BVH file uses the standard T-pose
+ as its rest pose instead of the BONES-SEED rest pose.
+ """
+ from_fmt = from_fmt or infer_source_format_from_path(input_path)
+ to_fmt = to_fmt or infer_target_format_from_path(output_path, from_fmt)
+
+ _validate_output_extension(to_fmt, output_path)
+
+ pair = (from_fmt, to_fmt)
+
+ if pair == ("amass", "kimodo"):
+ sk = build_skeleton(22)
+ effective_source = source_fps
+ if effective_source is None:
+ with np.load(input_path, allow_pickle=True) as z:
+ effective_source = float(z["mocap_frame_rate"]) if "mocap_frame_rate" in z.files else 30.0
+ motion = load_amass_npz(input_path, source_fps=effective_source, z_up=z_up)
+ save_kimodo_npz_at_target_fps(motion, sk, effective_source, output_path)
+ return
+
+ if pair == ("kimodo", "amass"):
+ data, J = load_kimodo_npz_as_torch(input_path, ensure_complete=False)
+ if J != 22:
+ raise ValueError(f"Kimodo→AMASS requires 22 joints (SMPL-X); this file has J={J}.")
+ sk = build_skeleton(22)
+ effective_source = resolve_source_fps(source_fps, "kimodo", input_path, None)
+ converter = AMASSConverter(fps=effective_source, skeleton=sk)
+ converter.convert_save_npz(data, output_path, z_up=z_up)
+ return
+
+ if pair == ("soma-bvh", "kimodo"):
+ sk = build_skeleton(77)
+ motion, bvh_fps = bvh_to_kimodo_motion(input_path, skeleton=sk, standard_tpose=bvh_standard_tpose)
+ effective_source = source_fps if source_fps is not None else bvh_fps
+ save_kimodo_npz_at_target_fps(motion, sk, effective_source, output_path)
+ return
+
+ if pair == ("kimodo", "soma-bvh"):
+ data, J = load_kimodo_npz_as_torch(input_path, ensure_complete=False)
+ if J == 30:
+ warnings.warn(
+ f"Input has 30 joints (somaskel30); expanding to somaskel77 for BVH export.",
+ UserWarning,
+ stacklevel=2,
+ )
+ sk = build_skeleton(30)
+ elif J == 77:
+ sk = build_skeleton(77)
+ else:
+ raise ValueError(f"Kimodo→BVH requires a SOMA skeleton (30 or 77 joints); this file has J={J}.")
+ effective_source = resolve_source_fps(source_fps, "kimodo", input_path, None)
+ save_motion_bvh(
+ output_path,
+ data["local_rot_mats"],
+ data["root_positions"],
+ skeleton=sk,
+ fps=effective_source,
+ standard_tpose=bvh_standard_tpose,
+ )
+ return
+
+ if pair == ("g1-csv", "kimodo"):
+ sk = build_skeleton(34)
+ effective_source = resolve_source_fps(source_fps, "g1-csv", input_path, None)
+ motion = load_g1_csv(input_path, source_fps=effective_source, mujoco_rest_zero=mujoco_rest_zero)
+ save_kimodo_npz_at_target_fps(motion, sk, effective_source, output_path)
+ return
+
+ if pair == ("kimodo", "g1-csv"):
+ data, J = load_kimodo_npz_as_torch(input_path, ensure_complete=False)
+ if J != 34:
+ raise ValueError(f"Kimodo→CSV requires G1 with 34 joints; this file has J={J}.")
+ sk = build_skeleton(34)
+ effective_source = resolve_source_fps(source_fps, "kimodo", input_path, None)
+ converter = MujocoQposConverter(sk)
+ qpos = converter.dict_to_qpos(
+ {k: v for k, v in data.items() if k in ("local_rot_mats", "root_positions")},
+ device=str(sk.neutral_joints.device),
+ numpy=True,
+ mujoco_rest_zero=mujoco_rest_zero,
+ )
+ converter.save_csv(qpos, output_path)
+ return
+
+ raise ValueError(
+ f"Unsupported conversion {from_fmt!r} → {to_fmt!r}. "
+ "Supported: amass↔kimodo (SMPL-X NPZ), soma-bvh↔kimodo, g1-csv↔kimodo."
+ )
+
+
+def _validate_output_extension(to_fmt: str, output_path: str) -> None:
+ lower = output_path.lower()
+ if to_fmt == "kimodo" and lower.endswith(".npz"):
+ return
+ if to_fmt == "amass":
+ if not lower.endswith(".npz"):
+ raise ValueError("AMASS output must use a .npz path.")
+ elif to_fmt == "soma-bvh":
+ if not lower.endswith(".bvh"):
+ raise ValueError("SOMA BVH output must use a .bvh path.")
+ elif to_fmt == "g1-csv":
+ if not lower.endswith(".csv"):
+ raise ValueError("G1 CSV output must use a .csv path.")
diff --git a/kimodo/exports/motion_formats.py b/kimodo/exports/motion_formats.py
new file mode 100644
index 0000000000000000000000000000000000000000..c2ba4cb4a0aeb214d79084fe6911a8176a7f92c9
--- /dev/null
+++ b/kimodo/exports/motion_formats.py
@@ -0,0 +1,78 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Infer motion file formats from paths and NPZ contents."""
+
+from __future__ import annotations
+
+import os
+from typing import Literal
+
+import numpy as np
+
+MotionSourceFormat = Literal["amass", "kimodo", "soma-bvh", "g1-csv"]
+MotionTargetFormat = Literal["amass", "kimodo", "soma-bvh", "g1-csv"]
+NpzMotionKind = Literal["amass", "kimodo"]
+
+
+def infer_npz_kind(path: str) -> NpzMotionKind:
+ """Classify a ``.npz`` as AMASS SMPL-X or Kimodo from required array keys."""
+ with np.load(path, allow_pickle=False) as z:
+ keys = set(z.files)
+ if "trans" in keys and "pose_body" in keys and "root_orient" in keys:
+ return "amass"
+ if "local_rot_mats" in keys or "posed_joints" in keys:
+ return "kimodo"
+ raise ValueError(
+ f"Unrecognized NPZ {path!r}: expected AMASS keys (trans, pose_body, ...) "
+ "or Kimodo keys (local_rot_mats, posed_joints, ...)."
+ )
+
+
+def infer_source_format_from_path(path: str) -> MotionSourceFormat:
+ """Infer converter input format from file extension and NPZ contents when needed."""
+ ext = os.path.splitext(path)[1].lower()
+ if ext == ".bvh":
+ return "soma-bvh"
+ if ext == ".csv":
+ return "g1-csv"
+ if ext == ".npz":
+ return infer_npz_kind(path) # type: ignore[return-value]
+ raise ValueError(f"Cannot infer format from extension of {path!r}")
+
+
+def infer_target_format_from_path(path: str, from_fmt: MotionSourceFormat) -> MotionTargetFormat:
+ """Infer converter output format from destination path and source format."""
+ ext = os.path.splitext(path)[1].lower()
+ if ext == ".bvh":
+ return "soma-bvh"
+ if ext == ".csv":
+ return "g1-csv"
+ if ext == ".npz":
+ if from_fmt == "amass":
+ return "kimodo"
+ if from_fmt == "kimodo":
+ return "amass"
+ if from_fmt in ("g1-csv", "soma-bvh"):
+ return "kimodo"
+ raise ValueError(
+ "Ambiguous .npz output: set --to to 'kimodo' or 'amass' when the input format is not amass/kimodo."
+ )
+ raise ValueError(f"Cannot infer output format from extension of {path!r}")
+
+
+def resolve_source_fps(
+ fps: float | None,
+ from_kind: str,
+ input_path: str,
+ data: dict | None,
+) -> float:
+ """Resolve source frame rate (Hz) for conversion when ``fps`` is not overridden."""
+ if fps is not None:
+ return float(fps)
+ if data is not None and "mocap_frame_rate" in data:
+ return float(np.asarray(data["mocap_frame_rate"]).item())
+ if from_kind == "soma-bvh":
+ from kimodo.exports.bvh import read_bvh_frame_time_seconds
+
+ return 1.0 / read_bvh_frame_time_seconds(input_path)
+ return 30.0
diff --git a/kimodo/exports/motion_io.py b/kimodo/exports/motion_io.py
new file mode 100644
index 0000000000000000000000000000000000000000..0b4bacc98ef84f24b77be9db01479bd1a966e877
--- /dev/null
+++ b/kimodo/exports/motion_io.py
@@ -0,0 +1,443 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Assemble Kimodo NPZ-compatible motion dicts from local rotations + root trajectory."""
+
+from __future__ import annotations
+
+import os
+import warnings
+from typing import Any, Dict, Tuple
+
+import numpy as np
+import torch
+
+from kimodo.geometry import matrix_to_quaternion, quaternion_to_matrix
+from kimodo.motion_rep.feature_utils import compute_heading_angle, compute_vel_xyz
+from kimodo.motion_rep.feet import foot_detect_from_pos_and_vel
+from kimodo.motion_rep.smooth_root import get_smooth_root_pos
+from kimodo.skeleton import SkeletonBase
+from kimodo.skeleton.registry import build_skeleton
+from kimodo.tools import to_numpy
+
+# Default motion rate for Kimodo NPZ produced by format conversion (matches common model FPS).
+KIMODO_CONVERT_TARGET_FPS = 30.0
+
+
+def _quaternion_slerp(q0: torch.Tensor, q1: torch.Tensor, t: torch.Tensor) -> torch.Tensor:
+ """Spherical linear interpolation; ``q0``, ``q1`` (..., 4) wxyz; ``t`` broadcastable to (...,
+ 1)."""
+ if t.dim() < q0.dim():
+ t = t.unsqueeze(-1)
+ dot = (q0 * q1).sum(dim=-1, keepdim=True)
+ q1 = torch.where(dot < 0, -q1, q1)
+ dot = torch.abs(dot).clamp(-1.0, 1.0)
+ theta_0 = torch.acos(dot)
+ sin_theta = torch.sin(theta_0)
+ s0 = torch.sin((1.0 - t) * theta_0) / sin_theta.clamp(min=1e-8)
+ s1 = torch.sin(t * theta_0) / sin_theta.clamp(min=1e-8)
+ q = s0 * q0 + s1 * q1
+ return q / torch.linalg.norm(q, dim=-1, keepdim=True).clamp(min=1e-8)
+
+
+def resample_motion_dict_to_kimodo_fps(
+ motion_dict: Dict[str, torch.Tensor],
+ skeleton: SkeletonBase,
+ source_fps: float,
+ target_fps: float = KIMODO_CONVERT_TARGET_FPS,
+) -> Tuple[Dict[str, torch.Tensor], bool]:
+ """Resample a Kimodo motion dict to ``target_fps``.
+
+ When the fps ratio is close to an integer (e.g. 120 / 30 = 4), the faster
+ stepping method is used (take every *step*-th frame). Otherwise falls back
+ to linear interp (root) + quaternion slerp (joints).
+
+ Re-runs :func:`complete_motion_dict` at the target rate so derived channels stay consistent.
+
+ Returns:
+ The motion dict and ``True`` if time resampling was applied, else ``False`` (already at
+ ``target_fps`` with matching frame count; only re-derived via FK).
+ """
+ local_rot_mats = motion_dict["local_rot_mats"]
+ root_positions = motion_dict["root_positions"]
+ local_rot_mats, root_positions = _coerce_time_local_root(local_rot_mats, root_positions)
+ t_in = int(local_rot_mats.shape[0])
+ if t_in < 1:
+ raise ValueError("Motion must have at least one frame.")
+ if source_fps <= 0:
+ raise ValueError(f"source_fps must be positive; got {source_fps}")
+
+ t_out = max(1, int(round(t_in * target_fps / source_fps)))
+ if t_out == t_in and abs(float(source_fps) - float(target_fps)) < 1e-3:
+ return complete_motion_dict(local_rot_mats, root_positions, skeleton, float(target_fps)), False
+
+ ratio = source_fps / target_fps
+ step = round(ratio)
+ if step >= 2 and abs(ratio - step) < 0.05:
+ local_out = local_rot_mats[::step]
+ root_out = root_positions[::step]
+ else:
+ device = local_rot_mats.device
+ dtype = local_rot_mats.dtype
+ u = torch.linspace(0, t_in - 1, t_out, device=device, dtype=dtype)
+ i0 = u.floor().long().clamp(0, t_in - 1)
+ i1 = torch.minimum(i0 + 1, torch.tensor(t_in - 1, device=device))
+ tau_1d = (u - i0.float()).unsqueeze(-1)
+ rp0 = root_positions[i0]
+ rp1 = root_positions[i1]
+ root_out = (1.0 - tau_1d) * rp0 + tau_1d * rp1
+
+ quats = matrix_to_quaternion(local_rot_mats)
+ q0 = quats[i0]
+ q1 = quats[i1]
+ tau_q = (u - i0.float()).view(t_out, 1, 1)
+ quat_out = _quaternion_slerp(q0, q1, tau_q)
+ local_out = quaternion_to_matrix(quat_out)
+
+ return complete_motion_dict(local_out, root_out, skeleton, float(target_fps)), True
+
+
+def warn_kimodo_npz_framerate(source_fps: float, t_before: int, t_after: int) -> None:
+ """Emit a warning after time resampling for Kimodo NPZ (linear root, quaternion slerp per
+ joint)."""
+ warnings.warn(
+ f"Resampled motion to {KIMODO_CONVERT_TARGET_FPS:.0f} Hz for Kimodo NPZ "
+ f"(source ~{source_fps:.4g} Hz, {t_before} input frames → {t_after} output frames). "
+ "Pass --source-fps if the detected source rate is wrong.",
+ UserWarning,
+ stacklevel=3,
+ )
+
+
+def _coerce_time_local_root(
+ local_rot_mats: torch.Tensor,
+ root_positions: torch.Tensor,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Normalize to shapes (T, J, 3, 3) and (T, 3)."""
+ if local_rot_mats.dim() == 5:
+ if int(local_rot_mats.shape[0]) != 1:
+ raise ValueError(f"local_rot_mats batch size must be 1 for single clip; got {local_rot_mats.shape[0]}")
+ local_rot_mats = local_rot_mats[0]
+ if root_positions.dim() == 3:
+ if int(root_positions.shape[0]) != 1:
+ raise ValueError(f"root_positions batch size must be 1; got {root_positions.shape[0]}")
+ root_positions = root_positions[0]
+ if local_rot_mats.dim() != 4:
+ raise ValueError(f"local_rot_mats must be (T,J,3,3); got {tuple(local_rot_mats.shape)}")
+ if root_positions.dim() != 2 or int(root_positions.shape[-1]) != 3:
+ raise ValueError(f"root_positions must be (T,3); got {tuple(root_positions.shape)}")
+ if int(local_rot_mats.shape[0]) != int(root_positions.shape[0]):
+ raise ValueError("local_rot_mats and root_positions must have the same number of frames")
+ return local_rot_mats, root_positions
+
+
+def complete_motion_dict(
+ local_rot_mats: torch.Tensor,
+ root_positions: torch.Tensor,
+ skeleton: SkeletonBase,
+ fps: float,
+) -> Dict[str, torch.Tensor]:
+ """Build the Kimodo motion output dict from local rotations and root positions.
+
+ Matches keys written by CLI generation (see docs/source/user_guide/output_formats.md).
+
+ Args:
+ local_rot_mats: (T, J, 3, 3) or (1, T, J, 3, 3) local rotation matrices.
+ root_positions: (T, 3) or (1, T, 3) root / pelvis world positions (meters).
+ skeleton: Skeleton instance (SOMA77, G1, SMPL-X, etc.).
+ fps: Sampling rate (Hz).
+
+ Returns:
+ Dict with tensors ``posed_joints``, ``global_rot_mats``, ``local_rot_mats``,
+ ``foot_contacts``, ``smooth_root_pos``, ``root_positions``, ``global_root_heading``.
+ """
+ device = local_rot_mats.device
+ dtype = local_rot_mats.dtype
+ local_rot_mats, root_positions = _coerce_time_local_root(
+ local_rot_mats.to(device=device, dtype=dtype),
+ root_positions.to(device=device, dtype=dtype),
+ )
+
+ global_rot_mats, posed_joints, _ = skeleton.fk(local_rot_mats, root_positions)
+
+ smooth_root_pos = get_smooth_root_pos(root_positions.unsqueeze(0)).squeeze(0)
+
+ lengths = torch.tensor([posed_joints.shape[0]], device=device)
+ velocities = compute_vel_xyz(posed_joints.unsqueeze(0), fps, lengths=lengths).squeeze(0)
+
+ heading_angle = compute_heading_angle(posed_joints.unsqueeze(0), skeleton).squeeze(0)
+ global_root_heading = torch.stack([torch.cos(heading_angle), torch.sin(heading_angle)], dim=-1)
+
+ foot_contacts = foot_detect_from_pos_and_vel(
+ posed_joints.unsqueeze(0),
+ velocities.unsqueeze(0),
+ skeleton,
+ 0.15,
+ 0.10,
+ ).squeeze(0)
+
+ return {
+ "posed_joints": posed_joints,
+ "global_rot_mats": global_rot_mats,
+ "local_rot_mats": local_rot_mats,
+ "foot_contacts": foot_contacts,
+ "smooth_root_pos": smooth_root_pos,
+ "root_positions": root_positions,
+ "global_root_heading": global_root_heading,
+ }
+
+
+def motion_dict_to_numpy(d: Dict[str, Any]) -> Dict[str, np.ndarray]:
+ """Convert motion dict values to numpy arrays for ``np.savez``."""
+ out: Dict[str, np.ndarray] = {}
+ for k, v in d.items():
+ if hasattr(v, "detach"):
+ out[k] = to_numpy(v)
+ elif isinstance(v, np.ndarray):
+ out[k] = v
+ else:
+ out[k] = np.asarray(v)
+ return out
+
+
+def save_kimodo_npz(path: str, motion_dict: Dict[str, Any]) -> None:
+ """Save a Kimodo-compatible motion dict to ``.npz`` (numpy arrays)."""
+ np.savez(path, **motion_dict_to_numpy(motion_dict))
+
+
+def load_kimodo_npz(path: str) -> Dict[str, np.ndarray]:
+ """Load arrays from a Kimodo ``.npz`` file."""
+ with np.load(path, allow_pickle=False) as data:
+ return {k: np.asarray(data[k]) for k in data.files}
+
+
+def load_g1_csv(
+ path: str,
+ source_fps: float = KIMODO_CONVERT_TARGET_FPS,
+ *,
+ mujoco_rest_zero: bool = False,
+) -> Dict[str, torch.Tensor]:
+ """Load a G1 MuJoCo ``qpos`` CSV (``(T, 36)``) into a Kimodo motion dict.
+
+ Args:
+ path: CSV path (comma-separated, no header).
+ source_fps: Source frame rate (Hz) of the CSV data.
+ mujoco_rest_zero: Must match how the CSV was written (see :class:`MujocoQposConverter`).
+ """
+ from kimodo.exports.mujoco import MujocoQposConverter
+
+ qpos = np.loadtxt(path, delimiter=",")
+ if qpos.ndim != 2 or qpos.shape[-1] != 36:
+ raise ValueError(f"Expected G1 CSV with shape (T, 36); got {qpos.shape}")
+ sk = build_skeleton(34)
+ converter = MujocoQposConverter(sk)
+ return converter.qpos_to_motion_dict(qpos, float(source_fps), mujoco_rest_zero=mujoco_rest_zero)
+
+
+def load_amass_npz(
+ path: str,
+ source_fps: float | None = None,
+ *,
+ z_up: bool = True,
+) -> Dict[str, torch.Tensor]:
+ """Load an AMASS-style SMPL-X ``.npz`` into a Kimodo motion dict (22 joints).
+
+ Args:
+ path: NPZ with ``trans``, ``root_orient``, ``pose_body``, etc.
+ source_fps: Source frame rate (Hz); if ``None``, uses ``mocap_frame_rate``
+ from the file when present, else 30 Hz.
+ z_up: If ``True``, apply AMASS Z-up to Kimodo Y-up transform (same as CLI).
+ """
+ from kimodo.exports.smplx import amass_npz_to_kimodo_motion
+
+ sk = build_skeleton(22)
+ return amass_npz_to_kimodo_motion(path, sk, source_fps=source_fps, z_up=z_up)
+
+
+def load_kimodo_npz_as_torch(
+ path: str,
+ source_fps: float = KIMODO_CONVERT_TARGET_FPS,
+ *,
+ ensure_complete: bool = True,
+) -> tuple[Dict[str, torch.Tensor], int]:
+ """Load a Kimodo NPZ and return all arrays as torch tensors on the skeleton device.
+
+ Args:
+ path: Kimodo NPZ file path.
+ source_fps: Source frame rate (Hz) used for derived channels when
+ ``ensure_complete=True``.
+ ensure_complete: If ``True`` and the NPZ lacks derived channels
+ (``posed_joints``, ``global_rot_mats``, …), run :func:`complete_motion_dict`
+ to fill them from ``local_rot_mats`` + ``root_positions``.
+ If ``False``, load all arrays verbatim (requires ``local_rot_mats``).
+
+ Returns:
+ ``(tensor_dict, num_joints)``
+ """
+ raw = load_kimodo_npz(path)
+ if "local_rot_mats" in raw:
+ j = int(raw["local_rot_mats"].shape[1])
+ elif "posed_joints" in raw:
+ j = int(raw["posed_joints"].shape[1])
+ else:
+ raise ValueError("Kimodo NPZ must contain 'local_rot_mats' or 'posed_joints'.")
+ sk = build_skeleton(j)
+ device = sk.neutral_joints.device
+ dtype = torch.float32
+
+ if not ensure_complete:
+ if "local_rot_mats" not in raw:
+ raise ValueError("Kimodo NPZ must contain 'local_rot_mats' (and typically 'root_positions').")
+ out: Dict[str, torch.Tensor] = {}
+ for k, v in raw.items():
+ out[k] = torch.from_numpy(np.asarray(v)).to(device=device, dtype=dtype)
+ return out, j
+
+ if "posed_joints" in raw and "global_rot_mats" in raw:
+ out = {}
+ for k, v in raw.items():
+ out[k] = torch.from_numpy(np.asarray(v)).to(device=device, dtype=dtype)
+ return out, j
+
+ if "local_rot_mats" not in raw or "root_positions" not in raw:
+ raise ValueError("Kimodo NPZ must contain posed_joints+global_rot_mats, or local_rot_mats+root_positions.")
+ local = torch.from_numpy(np.asarray(raw["local_rot_mats"])).to(device=device, dtype=dtype)
+ root = torch.from_numpy(np.asarray(raw["root_positions"])).to(device=device, dtype=dtype)
+ return complete_motion_dict(local, root, sk, float(source_fps)), j
+
+
+def save_kimodo_npz_at_target_fps(
+ motion: Dict[str, torch.Tensor],
+ skeleton: SkeletonBase,
+ source_fps: float,
+ output_path: str,
+ target_fps: float = KIMODO_CONVERT_TARGET_FPS,
+) -> None:
+ """Resample a motion dict to ``target_fps`` when needed, then save Kimodo NPZ."""
+ t_before = int(motion["local_rot_mats"].shape[0])
+ motion, did_resample = resample_motion_dict_to_kimodo_fps(motion, skeleton, source_fps, target_fps)
+ t_after = int(motion["local_rot_mats"].shape[0])
+ if did_resample:
+ warn_kimodo_npz_framerate(source_fps, t_before, t_after)
+ save_kimodo_npz(output_path, motion)
+
+
+def kimodo_npz_to_bytes(motion_dict: Dict[str, Any]) -> bytes:
+ """Serialize a Kimodo motion dict to in-memory NPZ bytes."""
+ import io
+
+ buf = io.BytesIO()
+ np.savez(buf, **motion_dict_to_numpy(motion_dict))
+ return buf.getvalue()
+
+
+def g1_csv_to_bytes(motion_dict: Dict[str, Any], skeleton: SkeletonBase, device: Any) -> bytes:
+ """Convert a motion dict to G1 MuJoCo CSV bytes via :class:`MujocoQposConverter`."""
+ import io
+
+ from kimodo.exports.mujoco import MujocoQposConverter
+
+ converter = MujocoQposConverter(skeleton)
+ qpos = converter.dict_to_qpos(
+ {k: v for k, v in motion_dict.items() if k in ("local_rot_mats", "root_positions")},
+ device,
+ numpy=True,
+ )
+ buf = io.StringIO()
+ np.savetxt(buf, qpos, delimiter=",")
+ return buf.getvalue().encode("utf-8")
+
+
+def amass_npz_to_bytes(motion_dict: Dict[str, Any], skeleton: SkeletonBase, fps: float) -> bytes:
+ """Convert a motion dict to AMASS NPZ bytes via :class:`AMASSConverter`."""
+ import io
+
+ from kimodo.exports.smplx import AMASSConverter
+
+ converter = AMASSConverter(skeleton=skeleton, fps=fps)
+ buf = io.BytesIO()
+ converter.convert_save_npz(
+ {k: v for k, v in motion_dict.items() if k in ("local_rot_mats", "root_positions")},
+ buf,
+ )
+ return buf.getvalue()
+
+
+def _read_amass_source_fps(path: str) -> float:
+ """Read the source frame rate from an AMASS NPZ, defaulting to 30 Hz."""
+ with np.load(path, allow_pickle=True) as z:
+ if "mocap_frame_rate" in z.files:
+ return float(z["mocap_frame_rate"])
+ return 30.0
+
+
+def load_motion_file(
+ path: str,
+ source_fps: float | None = None,
+ target_fps: float | None = None,
+ *,
+ z_up: bool = True,
+ mujoco_rest_zero: bool = False,
+) -> tuple[Dict[str, torch.Tensor], int]:
+ """Load a motion file and return a Kimodo motion dict plus joint count.
+
+ Supports SOMA BVH (``.bvh``), G1 MuJoCo CSV (``.csv``), Kimodo NPZ, and AMASS SMPL-X NPZ
+ (``.npz``).
+
+ The motion is loaded at its native (or overridden) source rate, then
+ resampled to ``target_fps`` when they differ.
+
+ Args:
+ path: Path to ``.bvh``, ``.csv``, or ``.npz``.
+ source_fps: Source frame rate (Hz). If provided, trusted as-is.
+ If ``None``, auto-detected per format: BVH ``Frame Time`` header,
+ AMASS ``mocap_frame_rate``, or :data:`KIMODO_CONVERT_TARGET_FPS`
+ (30 Hz) for CSV / Kimodo NPZ.
+ target_fps: Desired output frame rate (Hz). Defaults to
+ :data:`KIMODO_CONVERT_TARGET_FPS` (30 Hz). The motion is
+ resampled when ``source_fps`` and ``target_fps`` differ.
+ z_up: AMASS NPZ only; passed to :func:`load_amass_npz`.
+ mujoco_rest_zero: G1 CSV only; passed to :func:`load_g1_csv`.
+
+ Returns:
+ ``(motion_dict, num_joints)`` with the same keys as :func:`complete_motion_dict`.
+ """
+ from kimodo.exports.motion_formats import infer_npz_kind
+
+ if target_fps is None:
+ target_fps = KIMODO_CONVERT_TARGET_FPS
+
+ ext = os.path.splitext(path)[1].lower()
+ if ext == ".bvh":
+ from kimodo.exports.bvh import bvh_to_kimodo_motion
+
+ motion_dict, bvh_fps = bvh_to_kimodo_motion(path)
+ effective_source = source_fps if source_fps is not None else bvh_fps
+ num_joints = int(motion_dict["local_rot_mats"].shape[1])
+ elif ext == ".csv":
+ effective_source = source_fps if source_fps is not None else KIMODO_CONVERT_TARGET_FPS
+ motion_dict = load_g1_csv(path, source_fps=effective_source, mujoco_rest_zero=mujoco_rest_zero)
+ num_joints = 34
+ elif ext == ".npz":
+ kind = infer_npz_kind(path)
+ if kind == "amass":
+ effective_source = source_fps if source_fps is not None else _read_amass_source_fps(path)
+ motion_dict = load_amass_npz(path, source_fps=effective_source, z_up=z_up)
+ num_joints = 22
+ else:
+ effective_source = source_fps if source_fps is not None else KIMODO_CONVERT_TARGET_FPS
+ motion_dict, num_joints = load_kimodo_npz_as_torch(path, source_fps=effective_source)
+ else:
+ raise ValueError(f"Unsupported motion file {path!r}; expected .bvh, .csv, or .npz")
+
+ if abs(effective_source - target_fps) > 0.5:
+ sk = build_skeleton(num_joints)
+ motion_dict, did_resample = resample_motion_dict_to_kimodo_fps(motion_dict, sk, effective_source, target_fps)
+ if did_resample:
+ t_out = int(motion_dict["local_rot_mats"].shape[0])
+ warnings.warn(
+ f"Resampled motion from {effective_source:.4g} Hz to " f"{target_fps:.0f} Hz ({t_out} frames).",
+ UserWarning,
+ stacklevel=2,
+ )
+
+ return motion_dict, num_joints
diff --git a/kimodo/exports/mujoco.py b/kimodo/exports/mujoco.py
new file mode 100644
index 0000000000000000000000000000000000000000..77015dd24015f239529c1437d393cdc5859cdd97
--- /dev/null
+++ b/kimodo/exports/mujoco.py
@@ -0,0 +1,588 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Convert kimodo motion (y-up, z-forward) to MuJoCo qpos (z-up, x-forward) for G1 skeleton."""
+
+import os
+import xml.etree.ElementTree as ET
+from typing import Optional
+
+import numpy as np
+import torch
+from scipy.spatial.transform import Rotation
+
+from kimodo.assets import skeleton_asset_path
+from kimodo.geometry import (
+ axis_angle_to_matrix,
+ matrix_to_axis_angle,
+ matrix_to_quaternion,
+ quaternion_to_matrix,
+)
+from kimodo.skeleton import G1Skeleton34, SkeletonBase, global_rots_to_local_rots
+from kimodo.tools import ensure_batched, to_numpy, to_torch
+
+# Cache so that the same (skeleton, xml_path) returns the same converter instance.
+_converter_cache: dict[tuple[int, str], "MujocoQposConverter"] = {}
+
+
+class MujocoQposConverter:
+ """Fast batch converter from our dictionary format to mujoco qpos with precomputed transforms.
+
+ In mujoco, the coordination is z up and x forward, right handed.
+
+ Features (30 joints):
+ - root (pelvis, 7 = translation + rotation) + 29 dof joints (29)
+
+ In kimodo, the coordinate system is y up and z forward, right handed.
+ Features (34 joints):
+ - root (pelvis) + (34 - 1) joints; among these joints, 4 are end-effector joints added by kimodo.
+
+ Cached by (input_skeleton id, xml_path); repeated calls with the same args return the same instance.
+ """
+
+ def __new__(
+ cls,
+ input_skeleton: SkeletonBase,
+ xml_path: str = str(skeleton_asset_path("g1skel34", "xml", "g1.xml")),
+ ):
+ key = (id(input_skeleton), xml_path)
+ if key not in _converter_cache:
+ inst = object.__new__(cls)
+ _converter_cache[key] = inst
+ return _converter_cache[key]
+
+ def __init__(
+ self,
+ input_skeleton: SkeletonBase,
+ xml_path: str = str(skeleton_asset_path("g1skel34", "xml", "g1.xml")),
+ ):
+ """Initialize converter with precomputed transforms.
+
+ Args:
+ xml_path: Path to the mujoco XML file containing joint definitions
+ """
+ if getattr(self, "_initialized", False):
+ return
+ self.xml_path = xml_path
+ self.skeleton = input_skeleton
+ self._prepare_transforms()
+ self._subtree_joints = {}
+ self._initialized = True
+
+ def _prepare_transforms(self):
+ """Precompute all necessary transforms for efficient batch processing."""
+ # Define coordinate transformations between mujoco and kimodo space
+ # 1) R_zup_to_yup: rotation around x-axis by -90 degrees
+ # 2) x_forward_to_y_forward: rotation around z-axis by -90 degrees
+ # Combined transformation matrix: mujoco_to_kimodo = R_zup_to_yup * x_forward_to_y_forward
+ self.mujoco_to_kimodo_matrix = torch.tensor(
+ [[0.0, 1.0, 0.0], [0.0, 0.0, 1.0], [1.0, 0.0, 0.0]], dtype=torch.float32
+ )
+ self.kimodo_to_mujoco_matrix = self.mujoco_to_kimodo_matrix.T # Inverse transformation: kimodo_to_mujoco
+
+ # Parse XML once and extract joint information
+ tree = ET.parse(self.xml_path)
+ root = tree.getroot()
+
+ xml_classes = [x for x in tree.findall(".//default") if "class" in x.attrib]
+ joint_axes = dict()
+ class_ranges: dict[str, tuple[float, float]] = {}
+ for xml_class in xml_classes:
+ j = xml_class.findall("joint")
+ if j:
+ joint_axes[xml_class.get("class")] = j[0].get("axis")
+ range_str = j[0].get("range")
+ if range_str:
+ range_vals = [float(x) for x in range_str.split()]
+ if len(range_vals) == 2:
+ class_ranges[xml_class.get("class")] = (
+ range_vals[0],
+ range_vals[1],
+ )
+
+ mujoco_hinge_joints = root.find("worldbody").findall(".//joint") # skip the base joint
+ self._mujoco_joint_axis_values_kimodo_space = torch.zeros(
+ (len(mujoco_hinge_joints), 3), dtype=torch.float32
+ ) # mujoco order but kimodo space
+ self._mujoco_joint_axis_values_mujoco_space = torch.zeros(
+ (len(mujoco_hinge_joints), 3), dtype=torch.float32
+ ) # mujoco order but mujoco space
+
+ # for the below indices, mujoco_indices_to_kimodo_indices does not include mujoco root (30 - 1 = 29 elements),
+ # while kimodo_indices_to_mujoco_indices inclues the kimodo root (32 elements).
+ self._mujoco_indices_to_kimodo_indices = torch.zeros((len(mujoco_hinge_joints),), dtype=torch.int32)
+ self._kimodo_indices_to_mujoco_indices = (
+ torch.ones((self.skeleton.nbjoints,), dtype=torch.int32) * -1
+ ) # -1 means not in the csv skeleton
+
+ self._nb_joints_mujoco = len(mujoco_hinge_joints) + 1
+ self._nb_joints_kimodo = self.skeleton.nbjoints
+ self._mujoco_joint_including_root_parent_list = torch.full(
+ (len(mujoco_hinge_joints) + 1,), -1, dtype=torch.int32
+ )
+ self._mujoco_joint_including_root_list = ["pelvis_skel"]
+
+ for joint_id_in_csv, joint in enumerate(mujoco_hinge_joints):
+ joint_name_in_skeleton = joint.get("name").replace("_joint", "_skel")
+ joint_parent_name_in_skeleton = self.skeleton.bone_parents[joint_name_in_skeleton]
+
+ self._mujoco_joint_including_root_list.append(joint_name_in_skeleton)
+ self._mujoco_joint_including_root_parent_list[joint_id_in_csv + 1] = (
+ self._mujoco_joint_including_root_list.index(joint_parent_name_in_skeleton)
+ )
+
+ joint_idx_in_kimodo_skeleton = self.skeleton.bone_order_names.index(joint_name_in_skeleton)
+ axis_values = [float(x) for x in (joint.get("axis") or joint_axes[joint.get("class")]).split(" ")]
+
+ # the mapped axis in kimodo skeleton space is calculated as bones_axis = mujoco_to_kimodo.apply(axis_values)
+ # [1, 0, 0] -> [0, 0, 1]; [0, 1, 0] -> [1, 0, 0]; [0, 0, 1] -> [0, 1, 0]
+ mujoco_joint_axis_mapping_kimodo_space = [
+ torch.tensor([0, 0, 1]),
+ torch.tensor([1, 0, 0]),
+ torch.tensor([0, 1, 0]),
+ ][np.argmax(axis_values)]
+
+ self._mujoco_joint_axis_values_kimodo_space[joint_id_in_csv] = mujoco_joint_axis_mapping_kimodo_space
+ self._mujoco_joint_axis_values_mujoco_space[joint_id_in_csv] = torch.tensor(axis_values)
+
+ self._mujoco_indices_to_kimodo_indices[joint_id_in_csv] = joint_idx_in_kimodo_skeleton
+ self._kimodo_indices_to_mujoco_indices[joint_idx_in_kimodo_skeleton] = (
+ joint_id_in_csv + 1
+ ) # +1 for the root
+ self._kimodo_indices_to_mujoco_indices[0] = 0 # the root joint mapping
+
+ # Joint limits (min, max) in radians for each mujoco hinge, for clamping
+ self._joint_limits_min = torch.full((len(mujoco_hinge_joints),), float("-inf"), dtype=torch.float32)
+ self._joint_limits_max = torch.full((len(mujoco_hinge_joints),), float("inf"), dtype=torch.float32)
+ for joint_id_in_csv, joint in enumerate(mujoco_hinge_joints):
+ range_vals = None
+ if joint.get("range"):
+ range_vals = [float(x) for x in joint.get("range").split()]
+ elif joint.get("class") and joint.get("class") in class_ranges:
+ lo, hi = class_ranges[joint.get("class")]
+ range_vals = [lo, hi]
+ if range_vals is not None and len(range_vals) == 2:
+ self._joint_limits_min[joint_id_in_csv] = range_vals[0]
+ self._joint_limits_max[joint_id_in_csv] = range_vals[1]
+
+ # load the offset matrices from the xml
+ R_zup_to_yup = Rotation.from_euler("x", -90, degrees=True)
+ x_forward_to_y_forward = Rotation.from_euler("z", -90, degrees=True)
+ mujoco_to_kimodo = R_zup_to_yup * x_forward_to_y_forward
+
+ self._rot_offsets_q2t = torch.zeros(len(self._kimodo_indices_to_mujoco_indices), 3, 3, dtype=torch.float32)
+ self._rot_offsets_q2t[...] = torch.eye(3)[None]
+
+ self._rot_offsets_f2q = torch.zeros(len(self._kimodo_indices_to_mujoco_indices), 3, 3, dtype=torch.float32)
+ self._rot_offsets_f2q[...] = torch.eye(3)[None]
+ parent_map = {child: parent for parent in root.iter() for child in parent}
+ for i, joint in enumerate(mujoco_hinge_joints):
+ body = parent_map[joint]
+ if "quat" in body.attrib:
+ rot = Rotation.from_quat(
+ [float(x) for x in body.get("quat").strip().split(" ")],
+ scalar_first=True,
+ )
+ idx = self._mujoco_indices_to_kimodo_indices[i]
+ self._rot_offsets_q2t[idx] = torch.from_numpy(rot.as_matrix())
+ rot = mujoco_to_kimodo * rot * mujoco_to_kimodo.inv()
+ self._rot_offsets_f2q[idx] = torch.from_numpy(rot.as_matrix().T)
+
+ # Hinge axis in f2q space so extraction uses the same frame as joint_rot_f2q.
+ # Then extract(offset) gives the angle s.t. axis_angle(angle * axis_f2q) = offset, and
+ # reconstruction R_local = offset.T @ axis_angle(angle * axis_f2q) = I when input is identity.
+ axis_kimodo = self._mujoco_joint_axis_values_kimodo_space
+ self._mujoco_joint_axis_values_f2q_space = torch.zeros_like(axis_kimodo)
+ for i in range(len(mujoco_hinge_joints)):
+ j = self._mujoco_indices_to_kimodo_indices[i].item()
+ axis_f2q = torch.mv(self._rot_offsets_f2q[j], axis_kimodo[i])
+ n = axis_f2q.norm()
+ if n > 1e-8:
+ axis_f2q = axis_f2q / n
+ self._mujoco_joint_axis_values_f2q_space[i] = axis_f2q
+
+ # Rest-pose DOFs: angle we extract when R_local = I (t-pose). MuJoCo limits are
+ # relative to joint zero (rest pose), so we must clamp in MuJoCo space: convert
+ # joint_dofs to mujoco_angle = joint_dofs - rest_dofs, clamp, then back.
+ rest_rot_f2q = self._rot_offsets_f2q[self._mujoco_indices_to_kimodo_indices]
+ rest_rot_f2q = rest_rot_f2q.unsqueeze(0).unsqueeze(0)
+ self._rest_dofs = self._local_rots_f2q_to_joint_dofs(rest_rot_f2q).squeeze(0).squeeze(0)
+ # Axis-angle rest DOFs: angle s.t. axis_angle(angle * axis_f2q) = offset. Used in
+ # project_to_real_robot_rotations so extract+reconstruct round-trip and t-pose is preserved.
+ rest_rot_f2q_flat = self._rot_offsets_f2q[self._mujoco_indices_to_kimodo_indices]
+ full_aa = matrix_to_axis_angle(rest_rot_f2q_flat)
+ self._rest_dofs_axis_angle = (full_aa * self._mujoco_joint_axis_values_f2q_space).sum(dim=-1)
+
+ def dict_to_qpos(
+ self,
+ output: dict,
+ device: Optional[str] = None,
+ root_quat_w_first: bool = True,
+ numpy: bool = True,
+ mujoco_rest_zero: bool = False,
+ ):
+ """Convert kimodo output dict to mujoco qpos format.
+
+ Args:
+ output: dict with keys "local_rot_mats" and "root_positions".
+ device: device to use for the output.
+ root_quat_w_first: If True, quaternion in qpos is (w,x,y,z).
+ numpy: If True, convert the output to numpy array.
+ mujoco_rest_zero: If True, joint angles are written so that kimodo rest (t-pose)
+ maps to q=0 in MuJoCo. If False, write raw joint_dofs.
+
+ Returns:
+ qpos: (B, T, 7+J) mujoco qpos format.
+ """
+ local_rot_mats = to_torch(output["local_rot_mats"], device)
+ root_positions = to_torch(output["root_positions"], device)
+
+ qpos = self.to_qpos(
+ local_rot_mats,
+ root_positions,
+ root_quat_w_first=root_quat_w_first,
+ mujoco_rest_zero=mujoco_rest_zero,
+ )
+ if numpy:
+ qpos = to_numpy(qpos)
+ return qpos
+
+ def qpos_to_motion_dict(
+ self,
+ qpos: torch.Tensor | np.ndarray,
+ source_fps: float,
+ *,
+ root_quat_w_first: bool = True,
+ mujoco_rest_zero: bool = False,
+ ):
+ """Inverse of :meth:`to_qpos` / :meth:`dict_to_qpos` for MuJoCo CSV ``(T, 36)`` rows.
+
+ Args:
+ qpos: Shape ``(T, 36)`` or ``(1, T, 36)`` (root xyz, root quat wxyz, 29 joint angles).
+ source_fps: Source frame rate (Hz) of the qpos data.
+ root_quat_w_first: Must match how the CSV was written (default ``True``).
+ mujoco_rest_zero: Must match :meth:`dict_to_qpos` / :meth:`to_qpos`.
+
+ Returns:
+ Kimodo motion dict (see :func:`kimodo.exports.motion_io.complete_motion_dict`).
+ """
+ from kimodo.exports.motion_io import complete_motion_dict
+
+ qpos = to_torch(qpos, None)
+ if qpos.dim() == 2:
+ qpos = qpos.unsqueeze(0)
+ device = qpos.device
+ dtype = qpos.dtype
+ batch_size, num_frames, ncols = qpos.shape
+ if ncols != 36:
+ raise ValueError(f"Expected qpos last dim 36; got {ncols}")
+
+ kimodo_to_mujoco_matrix = self.kimodo_to_mujoco_matrix.to(device=device, dtype=dtype)
+ mujoco_to_kimodo_matrix = kimodo_to_mujoco_matrix.T
+
+ root_mujoco = qpos[..., :3]
+ root_positions = torch.matmul(mujoco_to_kimodo_matrix[None, None, ...], root_mujoco[..., None]).squeeze(-1)
+
+ quat = qpos[..., 3:7]
+ if root_quat_w_first:
+ root_rot_mujoco = quaternion_to_matrix(quat)
+ else:
+ quat_wxyz = quat[..., [3, 0, 1, 2]]
+ root_rot_mujoco = quaternion_to_matrix(quat_wxyz)
+
+ O0 = self._rot_offsets_f2q[0].to(device=device, dtype=dtype)
+ # root_rot_mujoco is (..., 3, 3) after optional batch unsqueeze (e.g. (1, T, 3, 3)).
+ # Use ``...il`` so ``k`` sums with ``kl``; ``...ik`` incorrectly keeps ``k`` in the output.
+ R_f2q_root = torch.einsum(
+ "ij,...jk,kl->...il",
+ mujoco_to_kimodo_matrix,
+ root_rot_mujoco,
+ kimodo_to_mujoco_matrix,
+ )
+ R_kimodo_root = torch.einsum("ij,...jk->...ik", O0.T, R_f2q_root)
+
+ joint_dofs = qpos[..., 7:]
+ if mujoco_rest_zero:
+ rest_dofs = self._rest_dofs.to(device=device, dtype=dtype)
+ angles = joint_dofs + rest_dofs[None, None, :]
+ use_relative = True
+ else:
+ angles = joint_dofs
+ use_relative = False
+
+ nb_joints = self.skeleton.nbjoints
+ template = torch.eye(3, device=device, dtype=dtype).expand(batch_size, num_frames, nb_joints, 3, 3).contiguous()
+ template[:, :, 0] = R_kimodo_root
+
+ local_rot_mats = self._joint_dofs_to_local_rot_mats(
+ angles,
+ template,
+ device,
+ dtype,
+ use_relative=use_relative,
+ )
+
+ if batch_size != 1:
+ raise ValueError(f"Only a single clip is supported; got batch_size={batch_size}")
+
+ return complete_motion_dict(local_rot_mats[0], root_positions[0], self.skeleton, source_fps)
+
+ def save_csv(self, qpos: torch.Tensor | np.ndarray, csv_path):
+ # comment this
+ qpos = to_numpy(qpos)
+ shape = qpos.shape
+ if len(shape) == 2:
+ # only one motion: save it
+ np.savetxt(csv_path, qpos, delimiter=",")
+ if len(shape) == 3:
+ # batch of motions
+ if shape[0] == 1:
+ # if only one motion, just save it
+ np.savetxt(csv_path, qpos[0], delimiter=",")
+ else:
+ csv_path_base, ext = os.path.splitext(csv_path)
+ for i in range(shape[0]):
+ self.save_csv(qpos[i], csv_path_base + "_" + str(i).zfill(2) + ext)
+
+ def _local_rots_to_joint_dofs(
+ self,
+ local_rot_mats: torch.Tensor,
+ axis_vals: torch.Tensor,
+ ) -> torch.Tensor:
+ """Extract per-joint single-DoF angles (radians) via Euler projection (for to_qpos/f2q)."""
+ x_joint_dof = torch.atan2(local_rot_mats[..., 2, 1], local_rot_mats[..., 2, 2])
+ y_joint_dof = torch.atan2(local_rot_mats[..., 0, 2], local_rot_mats[..., 0, 0])
+ z_joint_dof = torch.atan2(local_rot_mats[..., 1, 0], local_rot_mats[..., 1, 1])
+ xyz_joint_dofs = torch.stack([x_joint_dof, y_joint_dof, z_joint_dof], dim=-1)
+ axis_vals = axis_vals.to(device=local_rot_mats.device, dtype=local_rot_mats.dtype)
+ joint_dofs = (xyz_joint_dofs * axis_vals[None, None, :, :]).sum(dim=-1)
+ return joint_dofs
+
+ def _local_rots_to_joint_dofs_axis_angle(
+ self,
+ local_rot_mats: torch.Tensor,
+ axis_vals: torch.Tensor,
+ ) -> torch.Tensor:
+ """Extract per-joint single-DoF angles (radians) via axis-angle; round-trips with
+ axis_angle_to_matrix.
+
+ Args:
+ local_rot_mats: (..., num_hinges, 3, 3) in same frame as axis_vals.
+ axis_vals: (num_hinges, 3) unit axis per hinge.
+ Returns:
+ joint_dofs: (..., num_hinges) signed angle = dot(axis_angle(R), axis).
+ """
+ axis_vals = axis_vals.to(device=local_rot_mats.device, dtype=local_rot_mats.dtype)
+ full_aa = matrix_to_axis_angle(local_rot_mats)
+ joint_dofs = (full_aa * axis_vals).sum(dim=-1)
+ return joint_dofs
+
+ def _local_rots_f2q_to_joint_dofs(self, local_rot_mats_f2q: torch.Tensor) -> torch.Tensor:
+ """Extract per-joint single-DoF angles from local rotations in f2q space (for to_qpos)."""
+ axis_vals = self._mujoco_joint_axis_values_f2q_space
+ return self._local_rots_to_joint_dofs(local_rot_mats_f2q, axis_vals)
+
+ def _clamp_to_limits(self, joint_dofs: torch.Tensor) -> torch.Tensor:
+ """Clamp joint angles to XML limits (radians).
+
+ Angles are in kimodo convention (0 = rest).
+ """
+ device = joint_dofs.device
+ lo = self._joint_limits_min.to(device=device, dtype=joint_dofs.dtype)
+ hi = self._joint_limits_max.to(device=device, dtype=joint_dofs.dtype)
+ return torch.clamp(joint_dofs, lo[None, None, :], hi[None, None, :])
+
+ def _clamp_joint_dofs(self, joint_dofs: torch.Tensor, rest_dofs: torch.Tensor) -> torch.Tensor:
+ """Clamp joint angles to MuJoCo limits (radians), with rest_dofs conversion."""
+ device = joint_dofs.device
+ rest_dofs = rest_dofs.to(device=device, dtype=joint_dofs.dtype)
+ mujoco_dofs = joint_dofs - rest_dofs[None, None, :]
+ lo = self._joint_limits_min.to(device=device, dtype=joint_dofs.dtype)
+ hi = self._joint_limits_max.to(device=device, dtype=joint_dofs.dtype)
+ mujoco_dofs = torch.clamp(mujoco_dofs, lo[None, None, :], hi[None, None, :])
+ return mujoco_dofs + rest_dofs[None, None, :]
+
+ def _joint_dofs_to_local_rot_mats(
+ self,
+ joint_dofs: torch.Tensor,
+ original_local_rot_mats: torch.Tensor,
+ device: torch.device,
+ dtype: torch.dtype,
+ use_relative: bool = False,
+ ) -> torch.Tensor:
+ """Reconstruct full local rotation matrices from 1-DoF angles."""
+ out = original_local_rot_mats.clone()
+ axis_kimodo = self._mujoco_joint_axis_values_kimodo_space.to(device=device, dtype=dtype)
+ for i in range(joint_dofs.shape[-1]):
+ j = self._mujoco_indices_to_kimodo_indices[i].item()
+ angle = joint_dofs[..., i]
+ axis = axis_kimodo[i]
+ if use_relative:
+ axis_angle = angle[..., None] * axis[None, None, :]
+ R_local = axis_angle_to_matrix(axis_angle)
+ else:
+ rot_offsets_f2q = self._rot_offsets_f2q.to(device=device, dtype=dtype)
+ axis_in_f2q = torch.mv(rot_offsets_f2q[j], axis)
+ axis_angle = angle[..., None] * axis_in_f2q[None, None, :]
+ R_f2q = axis_angle_to_matrix(axis_angle)
+ R_local = torch.einsum("ij,btjk->btik", rot_offsets_f2q[j].T, R_f2q)
+ out[:, :, j, :, :] = R_local
+ return out
+
+ @ensure_batched(local_rot_mats=5, root_positions=3, lengths=1)
+ def project_to_real_robot_rotations(
+ self,
+ local_rot_mats: torch.Tensor,
+ root_positions: torch.Tensor,
+ clamp_to_limits: bool = True,
+ mujoco_rest_zero: bool = False,
+ ) -> dict:
+ """Project full 3D local rotations to G1 real robot DoF and back to 3D for viz.
+
+ Joint angles are extracted along each hinge axis, optionally clamped to XML limits, then
+ reconstructed to 3D rotations. When mujoco_rest_zero=False (default), raw angles are used
+ (baked-with-quat). When True, angles are relative to rest (0 = T-pose in MuJoCo).
+ """
+ device = local_rot_mats.device
+ dtype = local_rot_mats.dtype
+
+ # Transform to f2q frame and extract 1-DoF angles (axis-angle projection).
+ local_rot_f2q = torch.matmul(self._rot_offsets_f2q.to(device=device, dtype=dtype), local_rot_mats)
+ hinge_rots = local_rot_f2q[:, :, self._mujoco_indices_to_kimodo_indices, :, :]
+ axis_f2q = self._mujoco_joint_axis_values_f2q_space.to(device=device, dtype=dtype)
+ joint_dofs = self._local_rots_to_joint_dofs_axis_angle(hinge_rots, axis_f2q)
+
+ # Optionally express angles relative to rest (MuJoCo q=0 at T-pose).
+ if mujoco_rest_zero:
+ rest_dofs = self._rest_dofs_axis_angle.to(device=device, dtype=dtype)
+ angles = joint_dofs - rest_dofs[None, None, :]
+ use_relative = True
+ else:
+ angles = joint_dofs
+ use_relative = False
+
+ if clamp_to_limits:
+ if mujoco_rest_zero:
+ angles = self._clamp_to_limits(angles)
+ else:
+ rest_dofs_aa = self._rest_dofs_axis_angle.to(device=device, dtype=dtype)
+ angles = self._clamp_joint_dofs(angles, rest_dofs_aa)
+
+ # Reconstruct 3D local rotations from 1-DoF angles and run FK.
+ local_rot_mats_proj = self._joint_dofs_to_local_rot_mats(
+ angles, local_rot_mats, device, dtype, use_relative=use_relative
+ )
+ global_rot_mats, posed_joints, _ = self.skeleton.fk(local_rot_mats_proj, root_positions)
+ return {
+ "local_rot_mats": local_rot_mats_proj,
+ "global_rot_mats": global_rot_mats,
+ "posed_joints": posed_joints,
+ "root_positions": root_positions,
+ }
+
+ @ensure_batched(local_rot_mats=5, root_positions=3, lengths=1)
+ def to_qpos(
+ self,
+ local_rot_mats: torch.Tensor,
+ root_positions: torch.Tensor,
+ root_quat_w_first: bool = True,
+ mujoco_rest_zero: bool = False,
+ ) -> torch.Tensor:
+ """Fast batch conversion from kimodo features to mujoco qpos format.
+
+ Args:
+ local_rot_mats: (B, T, J, 3, 3) local rotation matrices (kimodo convention).
+ root_positions: (B, T, 3) root positions.
+ root_quat_w_first: If True, quaternion in qpos is (w,x,y,z).
+ mujoco_rest_zero: If True, joint angles are written so that kimodo rest (t-pose)
+ maps to q=0 in MuJoCo. If False, write raw joint_dofs.
+
+ Returns:
+ torch.Tensor of shape [batch, numFrames, 36] containing mujoco qpos data:
+ - root_trans (3) + root_quat (4) + joint_dofs (29) = 36 columns
+ """
+
+ batch_size, num_frames, nb_joints = local_rot_mats.shape[:3]
+ device, dtype = local_rot_mats.device, local_rot_mats.dtype
+
+ local_rot_mats = torch.matmul(self._rot_offsets_f2q.to(device), local_rot_mats)
+
+ batch_size, num_frames = root_positions.shape[0], root_positions.shape[1]
+
+ # Move precomputed matrices to the same device/dtype
+ kimodo_to_mujoco_matrix = self.kimodo_to_mujoco_matrix.to(device=device, dtype=dtype)
+
+ # Initialize output tensor: [batch, numFrames, 36]
+ qpos = torch.zeros((batch_size, num_frames, 36), dtype=dtype, device=device)
+
+ # Convert root translation: apply coordinate transformation
+ root_positions_mujoco = torch.matmul(kimodo_to_mujoco_matrix[None, None, ...], root_positions[..., None])
+ qpos[:, :, :3] = root_positions_mujoco.view(batch_size, num_frames, 3)
+
+ # Convert root rotation: apply coordinate transformation to rotation matrix
+ root_rot = local_rot_mats[:, :, 0, :] # [batch, numFrames, 3, 3]
+
+ # Apply coordinate transformation: R_mujoco = kimodo_to_mujoco * R_kimodo * kimodo_to_mujoco^T
+ mujoco_to_kimodo_matrix = kimodo_to_mujoco_matrix.T
+ root_rot_mujoco = torch.matmul(
+ torch.matmul(kimodo_to_mujoco_matrix[None, None, ...], root_rot),
+ mujoco_to_kimodo_matrix[None, None, ...],
+ )
+ root_rot_quat = matrix_to_quaternion(root_rot_mujoco) # [w, x, y, z]
+ if root_quat_w_first:
+ qpos[:, :, 3:7] = root_rot_quat[:, :, [0, 1, 2, 3]] # [w, x, y, z]
+ else:
+ qpos[:, :, 3:7] = root_rot_quat[:, :, [1, 2, 3, 0]] # [w, x, y, z] -> [x, y, z, w]
+
+ # Joint DOFs: raw angles or relative to rest (rest = q=0 in MuJoCo).
+ joint_rot_f2q = local_rot_mats[:, :, self._mujoco_indices_to_kimodo_indices, :, :]
+ joint_dofs = self._local_rots_f2q_to_joint_dofs(joint_rot_f2q)
+ if mujoco_rest_zero:
+ rest_dofs = self._rest_dofs.to(device=device, dtype=dtype)
+ qpos[:, :, 7:] = joint_dofs - rest_dofs[None, None, :]
+ else:
+ qpos[:, :, 7:] = joint_dofs
+ return qpos
+
+
+def apply_g1_real_robot_projection(
+ skeleton: G1Skeleton34,
+ joints_pos: torch.Tensor,
+ joints_rot: torch.Tensor,
+ clamp_to_limits: bool = True,
+) -> tuple[torch.Tensor, torch.Tensor]:
+ """Project G1 motion to real robot DoF (1-DoF per joint) with optional axis limits.
+
+ Extracts a single angle per hinge along its axis (1-DoF), optionally clamps to
+ joint limits from the MuJoCo XML (when clamp_to_limits=True), then reconstructs
+ 3D rotations and runs FK. T-pose (identity local rotations) is preserved.
+
+ Args:
+ skeleton: G1 skeleton instance.
+ joints_pos: (T, J, 3) or (B, T, J, 3) joint positions in global space.
+ joints_rot: (T, J, 3, 3) or (B, T, J, 3, 3) global rotation matrices.
+ clamp_to_limits: If True, clamp joint angles to XML axis limits (default True).
+
+ Returns:
+ (posed_joints, global_rot_mats) as tensors, same shape as inputs (batch preserved).
+ """
+
+ local_rot_mats = global_rots_to_local_rots(joints_rot, skeleton)
+ root_positions = joints_pos[..., skeleton.root_idx, :]
+
+ # Converter expects batch dim (B, T, ...); add and remove if single sequence.
+ single_sequence = local_rot_mats.dim() == 4
+ if single_sequence:
+ local_rot_mats = local_rot_mats.unsqueeze(0)
+ root_positions = root_positions.unsqueeze(0)
+
+ converter = MujocoQposConverter(skeleton)
+ projected = converter.project_to_real_robot_rotations(
+ local_rot_mats, root_positions, clamp_to_limits=clamp_to_limits
+ )
+
+ out_pos = projected["posed_joints"]
+ out_rot = projected["global_rot_mats"]
+ if single_sequence:
+ out_pos = out_pos.squeeze(0)
+ out_rot = out_rot.squeeze(0)
+ return out_pos, out_rot
diff --git a/kimodo/exports/smplx.py b/kimodo/exports/smplx.py
new file mode 100644
index 0000000000000000000000000000000000000000..ce1d15262fccf91800b006a9b679731e95431da6
--- /dev/null
+++ b/kimodo/exports/smplx.py
@@ -0,0 +1,251 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Convert kimodo motion to AMASS/SMPL-X compatible parameters (axis-angle, Y-up or Z-up)."""
+
+import os
+from typing import Optional
+
+import einops
+import numpy as np
+import torch
+
+from kimodo.assets import skeleton_asset_path
+from kimodo.geometry import axis_angle_to_matrix, matrix_to_axis_angle
+from kimodo.tools import ensure_batched, to_numpy, to_torch
+
+
+def kimodo_y_up_to_amass_coord_rotation_matrix() -> np.ndarray:
+ """3x3 rotation mapping Kimodo Y-up (+Z forward) to AMASS Z-up (+Y forward).
+
+ Used by :func:`get_amass_parameters` and :func:`amass_arrays_to_kimodo_motion` (inverse).
+ """
+ y_up_to_z_up = np.array(
+ [
+ [1.0, 0.0, 0.0],
+ [0.0, 0.0, -1.0],
+ [0.0, 1.0, 0.0],
+ ],
+ dtype=np.float32,
+ )
+ rot_z_180 = np.array(
+ [
+ [-1.0, 0.0, 0.0],
+ [0.0, -1.0, 0.0],
+ [0.0, 0.0, 1.0],
+ ],
+ dtype=np.float32,
+ )
+ return np.matmul(rot_z_180, y_up_to_z_up).astype(np.float32)
+
+
+@ensure_batched(local_rot_mats=5, root_positions=3, lengths=1)
+def get_amass_parameters(
+ local_rot_mats,
+ root_positions,
+ skeleton,
+ z_up=True,
+):
+ """Convert local rot mats and root positions to AMASS-style trans and pose_body; optional z_up
+ coordinate transform.
+
+ Our method generates motions with Y-up and +Z forward; if z_up=True, transform to Z-up and +Y
+ forward as in AMASS.
+ """
+ # Our method generate motions with Y-up and +Z forward
+ # if z_up = True, we transform this to: Z-up with +Y forward, as in AMASS
+ # Remove the root offset; SMPL-X FK adds pelvis offset back.
+ pelvis_offset = skeleton.neutral_joints[skeleton.root_idx].cpu().numpy()
+ trans = root_positions - pelvis_offset
+
+ root_rot_mats = to_numpy(local_rot_mats[:, :, 0])
+ local_rot_axis_angle = to_numpy(matrix_to_axis_angle(to_torch(local_rot_mats)))
+ pose_body = einops.rearrange(local_rot_axis_angle[:, :, 1:], "b t j d -> b t (j d)")
+
+ # Optionally convert from Y-up to Z-up coordinates.
+ if z_up:
+ y_up_to_z_up = kimodo_y_up_to_amass_coord_rotation_matrix()
+ root_rot_mats = np.matmul(y_up_to_z_up, root_rot_mats)
+ trans = np.matmul(trans + pelvis_offset, y_up_to_z_up.T) - pelvis_offset
+
+ root_orient = to_numpy(matrix_to_axis_angle(to_torch(root_rot_mats)))
+ return trans, root_orient, pose_body
+
+
+def amass_arrays_to_kimodo_motion(
+ trans: np.ndarray,
+ root_orient: np.ndarray,
+ pose_body: np.ndarray,
+ skeleton,
+ source_fps: float,
+ *,
+ z_up: bool = True,
+):
+ """Inverse of :func:`get_amass_parameters` for a single sequence (AMASS → Kimodo motion dict).
+
+ Args:
+ trans: ``(T, 3)`` AMASS root translation (same as ``trans`` in AMASS NPZ).
+ root_orient: ``(T, 3)`` axis-angle root orientation in AMASS coordinates (z-up when ``z_up``).
+ pose_body: ``(T, 63)`` body pose axis-angle (21 joints × 3).
+ skeleton: :class:`~kimodo.skeleton.definitions.SMPLXSkeleton22` instance.
+ source_fps: Source frame rate (Hz) of the AMASS recording.
+ z_up: If ``True``, invert the same Y-up↔Z-up transform as ``get_amass_parameters(..., z_up=True)``.
+
+ Returns:
+ Motion dict compatible with :func:`kimodo.exports.motion_io.save_kimodo_npz`.
+ """
+ from kimodo.exports.motion_io import complete_motion_dict
+
+ trans = np.asarray(trans, dtype=np.float32)
+ root_orient = np.asarray(root_orient, dtype=np.float32)
+ pose_body = np.asarray(pose_body, dtype=np.float32)
+ if trans.ndim != 2 or trans.shape[-1] != 3:
+ raise ValueError(f"trans must be (T, 3); got {trans.shape}")
+ if root_orient.shape != trans.shape:
+ raise ValueError(f"root_orient shape {root_orient.shape} must match trans {trans.shape}")
+ t = trans.shape[0]
+ if pose_body.shape != (t, 63):
+ raise ValueError(f"pose_body must be (T, 63); got {pose_body.shape}")
+
+ pelvis_offset = skeleton.neutral_joints[skeleton.root_idx].detach().cpu().numpy().astype(np.float32)
+ device = skeleton.neutral_joints.device
+ dtype = torch.float32
+
+ Y_np = kimodo_y_up_to_amass_coord_rotation_matrix()
+ if z_up:
+ y_up_to_z_up = torch.from_numpy(Y_np).to(device=device, dtype=dtype)
+ # trans_amass = root_kimodo @ Y.T - pelvis_offset => root_kimodo = (trans_amass + pelvis_offset) @ Y
+ root_positions_np = (trans + pelvis_offset) @ Y_np
+ else:
+ root_positions_np = trans + pelvis_offset
+
+ root_positions = torch.from_numpy(root_positions_np).to(device=device, dtype=dtype)
+
+ R_amass_root = axis_angle_to_matrix(torch.from_numpy(root_orient).to(device=device, dtype=dtype))
+ if z_up:
+ R_kimodo_root = torch.einsum("ij,tjk->tik", y_up_to_z_up.T, R_amass_root)
+ else:
+ R_kimodo_root = R_amass_root
+
+ nb = skeleton.nbjoints
+ if nb != 22:
+ raise ValueError(f"Expected SMPL-X body skeleton with 22 joints; got {nb}")
+
+ local_rot_mats = torch.zeros((t, nb, 3, 3), device=device, dtype=dtype)
+ local_rot_mats[:, 0] = R_kimodo_root
+
+ pose_aa = torch.from_numpy(pose_body.reshape(t, 21, 3)).to(device=device, dtype=dtype)
+ local_rot_mats[:, 1:] = axis_angle_to_matrix(pose_aa.reshape(-1, 3)).reshape(t, 21, 3, 3)
+
+ return complete_motion_dict(local_rot_mats, root_positions, skeleton, source_fps)
+
+
+def amass_npz_to_kimodo_motion(npz_path: str, skeleton, source_fps: Optional[float] = None, *, z_up: bool = True):
+ """Load an AMASS-style ``.npz`` and return a Kimodo motion dict.
+
+ Args:
+ npz_path: Path to AMASS NPZ (``trans``, ``root_orient``, ``pose_body``, ...).
+ skeleton: SMPL-X skeleton instance.
+ source_fps: Source frame rate (Hz); if ``None``, uses ``mocap_frame_rate``
+ from the file when present, else ``30.0``.
+ z_up: Same meaning as :func:`amass_arrays_to_kimodo_motion`.
+ """
+ with np.load(npz_path, allow_pickle=True) as data:
+ trans = np.asarray(data["trans"], dtype=np.float32)
+ root_orient = np.asarray(data["root_orient"], dtype=np.float32)
+ pose_body = np.asarray(data["pose_body"], dtype=np.float32)
+ if source_fps is None:
+ source_fps = float(data["mocap_frame_rate"]) if "mocap_frame_rate" in data.files else 30.0
+
+ return amass_arrays_to_kimodo_motion(trans, root_orient, pose_body, skeleton, source_fps, z_up=z_up)
+
+
+class AMASSConverter:
+ def __init__(
+ self,
+ fps,
+ skeleton,
+ beta_path=str(skeleton_asset_path("smplx22", "beta.npy")),
+ mean_hands_path=str(skeleton_asset_path("smplx22", "mean_hands.npy")),
+ ):
+ self.fps = fps
+ self.skeleton = skeleton
+ # Load betas
+ if os.path.exists(beta_path):
+ # only use first 16 betas to match AMASS
+ betas = np.load(beta_path)[:16]
+ else:
+ betas = np.zeros(16)
+
+ # Load mean hands
+ if os.path.exists(mean_hands_path):
+ mean_hands = np.load(mean_hands_path)
+ else:
+ mean_hands = np.zeros(90)
+
+ self.default_frame_params = {
+ "pose_jaw": np.zeros(3),
+ "pose_eye": np.zeros(6),
+ "pose_hand": mean_hands,
+ }
+ self.output_dict_base = {
+ "gender": "neutral",
+ "surface_model_type": "smplx",
+ "betas": betas,
+ "num_betas": len(betas),
+ "mocap_frame_rate": float(fps),
+ }
+
+ def convert_save_npz(self, output: dict, npz_path, z_up=True):
+ trans, root_orient, pose_body = get_amass_parameters(
+ output["local_rot_mats"],
+ output["root_positions"],
+ self.skeleton,
+ z_up=z_up,
+ )
+ nb_frames = trans.shape[-2]
+
+ amass_output_base = self.output_dict_base.copy()
+ for key, val in self.default_frame_params.items():
+ amass_output_base[key] = einops.repeat(val, "d -> t d", t=nb_frames)
+
+ amass_output_base["mocap_time_length"] = nb_frames / self.fps
+ self.save_npz(trans, root_orient, pose_body, amass_output_base, npz_path)
+
+ def save_npz(self, trans, root_orient, pose_body, base_output, npz_path):
+ shape = trans.shape
+ if len(shape) == 3 and shape[0] == 1:
+ # if only one motion, squeeze the data
+ trans = trans[0]
+ root_orient = root_orient[0]
+ pose_body = pose_body[0]
+ shape = trans.shape
+ if len(shape) == 2:
+ amass_output = {
+ "trans": trans,
+ "root_orient": root_orient,
+ "pose_body": pose_body,
+ } | base_output
+ np.savez(npz_path, **amass_output)
+
+ elif len(shape) == 3:
+ # real batch of motions
+ npz_path_base, ext = os.path.splitext(npz_path)
+ for i in range(shape[0]):
+ npz_path_i = npz_path_base + "_" + str(i).zfill(2) + ext
+ self.save_npz(trans[i], root_orient[i], pose_body[i], base_output, npz_path_i)
+
+
+# amass_output = {
+# "gender": "neutral",
+# "surface_model_type": "smplx",
+# "mocap_frame_rate": float(fps),
+# "mocap_time_length": len(motion) / float(fps)
+# "trans": trans,
+# "betas": betas,
+# "num_betas": len(betas),
+# "root_orient": np.array([T, 3]), # axis angle
+# "pose_body": np.array([T, 63]), # 63=21*3, axis angle 21 = 22 - root
+# "pose_hand": np.array([T, 90]), # 90=30*3=15*2*3 axis angle (load from mean_hands)
+# "pose_jaw": np.array([T, 3]), # all zeros is fine
+# "pose_eye": np.array([T, 6]), # all zeros is fine`
+# }
diff --git a/kimodo/geometry.py b/kimodo/geometry.py
new file mode 100644
index 0000000000000000000000000000000000000000..e0d2397bf5f4517fc92280caa7dfbab993452940
--- /dev/null
+++ b/kimodo/geometry.py
@@ -0,0 +1,216 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Rotation and representation conversions: axis-angle, quaternion, matrix, 6D continuous."""
+
+import torch
+import torch.nn.functional as F
+
+
+def angle_to_Y_rotation_matrix(angle: torch.Tensor) -> torch.Tensor:
+ """Build a rotation matrix around the Y axis from a scalar angle (radians).
+
+ Shape: angle.shape + (3, 3).
+ """
+ cos, sin = torch.cos(angle), torch.sin(angle)
+ one, zero = torch.ones_like(angle), torch.zeros_like(angle)
+ mat = torch.stack((cos, zero, sin, zero, one, zero, -sin, zero, cos), -1)
+ mat = mat.reshape(angle.shape + (3, 3))
+ return mat
+
+
+def matrix_to_cont6d(matrix: torch.Tensor) -> torch.Tensor:
+ """Convert rotation matrix to 6D continuous representation (first two columns).
+
+ Shape: (..., 3, 3) -> (..., 6).
+ """
+ cont_6d = torch.concat([matrix[..., 0], matrix[..., 1]], dim=-1)
+ return cont_6d
+
+
+def cont6d_to_matrix(cont6d: torch.Tensor) -> torch.Tensor:
+ """Convert 6D continuous representation to rotation matrix (Gram–Schmidt on two columns).
+
+ Last dim must be 6.
+ """
+ assert cont6d.shape[-1] == 6, "The last dimension must be 6"
+ x_raw = cont6d[..., 0:3]
+ y_raw = cont6d[..., 3:6]
+
+ x = x_raw / torch.norm(x_raw, dim=-1, keepdim=True)
+ z = torch.cross(x, y_raw, dim=-1)
+ z = z / torch.norm(z, dim=-1, keepdim=True)
+
+ y = torch.cross(z, x, dim=-1)
+
+ x = x[..., None]
+ y = y[..., None]
+ z = z[..., None]
+
+ mat = torch.cat([x, y, z], dim=-1)
+ return mat
+
+
+def axis_angle_to_matrix(axis_angle: torch.Tensor) -> torch.Tensor:
+ """Convert axis-angle to rotation matrix.
+
+ Args:
+ axis_angle: (..., 3) axis-angle vectors (angle = norm, axis = normalized)
+ Returns:
+ rotmat: (..., 3, 3) rotation matrices
+ """
+ eps = 1e-6
+ angle = torch.norm(axis_angle, dim=-1, keepdim=True) # (..., 1)
+ axis = axis_angle / (angle + eps)
+
+ x, y, z = axis.unbind(-1)
+
+ zero = torch.zeros_like(x)
+ K = torch.stack([zero, -z, y, z, zero, -x, -y, x, zero], dim=-1).reshape(*axis.shape[:-1], 3, 3)
+
+ eye = torch.eye(3, device=axis.device, dtype=axis.dtype)
+ eye = eye.expand(*axis.shape[:-1], 3, 3)
+
+ sin = torch.sin(angle)[..., None]
+ cos = torch.cos(angle)[..., None]
+
+ R = eye + sin * K + (1 - cos) * (K @ K)
+ return R
+
+
+def matrix_to_axis_angle(R: torch.Tensor) -> torch.Tensor:
+ """Convert rotation matrix to axis-angle via quaternions (more numerically stable).
+
+ Args:
+ R: (..., 3, 3) rotation matrices
+ Returns:
+ axis_angle: (..., 3)
+ """
+ # Go through quaternions for numerical stability
+ quat = matrix_to_quaternion(R) # (..., 4) with (w, x, y, z)
+ return quaternion_to_axis_angle(quat)
+
+
+def quaternion_to_axis_angle(quat: torch.Tensor) -> torch.Tensor:
+ """Convert quaternion to axis-angle representation.
+
+ Args:
+ quat: (..., 4) quaternions with real part first (w, x, y, z)
+ Returns:
+ axis_angle: (..., 3)
+ """
+ eps = 1e-6
+
+ # Ensure canonical form to avoid sign ambiguity.
+ # Primary: prefer w > 0. When w ≈ 0 (angle ≈ π), prefer first nonzero xyz > 0.
+ w = quat[..., 0:1]
+ xyz = quat[..., 1:]
+
+ # Find first significant component of xyz for tie-breaking when w ≈ 0
+ first_significant = xyz[..., 0:1] # use x component as tie-breaker
+
+ # Flip if: w < 0, OR (w ≈ 0 AND first xyz component < 0)
+ should_flip = (w < -eps) | ((w.abs() <= eps) & (first_significant < 0))
+ quat = torch.where(should_flip, -quat, quat)
+
+ w = quat[..., 0]
+ xyz = quat[..., 1:]
+
+ # sin(angle/2) = ||xyz||
+ sin_half_angle = xyz.norm(dim=-1)
+
+ # angle = 2 * atan2(sin(angle/2), cos(angle/2))
+ # This is more stable than 2 * acos(w) near angle=0
+ angle = 2.0 * torch.atan2(sin_half_angle, w)
+
+ # axis = xyz / sin(angle/2), but handle small angles
+ # For small angles: axis-angle ≈ 2 * xyz (since sin(x) ≈ x for small x)
+ small_angle = sin_half_angle.abs() < eps
+
+ # Safe division
+ scale = torch.where(
+ small_angle,
+ 2.0 * torch.ones_like(angle), # small angle: axis_angle ≈ 2 * xyz
+ angle / sin_half_angle.clamp(min=eps),
+ )
+
+ return xyz * scale.unsqueeze(-1)
+
+
+def _sqrt_positive_part(x: torch.Tensor) -> torch.Tensor:
+ """Returns torch.sqrt(torch.max(0, x)) subgradient is zero where x is 0."""
+ return torch.sqrt(x * (x > 0).to(x.dtype))
+
+
+def matrix_to_quaternion(matrix: torch.Tensor) -> torch.Tensor:
+ """Convert rotations given as rotation matrices to quaternions.
+
+ Args:
+ matrix: Rotation matrices as tensor of shape (..., 3, 3).
+ Returns:
+ quaternions with real part first, as tensor of shape (..., 4).
+ """
+ if matrix.size(-1) != 3 or matrix.size(-2) != 3:
+ raise ValueError(f"Invalid rotation matrix shape {matrix.shape}.")
+
+ batch_dim = matrix.shape[:-2]
+ m00, m01, m02, m10, m11, m12, m20, m21, m22 = torch.unbind(matrix.reshape(batch_dim + (9,)), dim=-1)
+
+ q_abs = _sqrt_positive_part(
+ torch.stack(
+ [
+ 1.0 + m00 + m11 + m22,
+ 1.0 + m00 - m11 - m22,
+ 1.0 - m00 + m11 - m22,
+ 1.0 - m00 - m11 + m22,
+ ],
+ dim=-1,
+ )
+ )
+
+ quat_by_rijk = torch.stack(
+ [
+ torch.stack([q_abs[..., 0] ** 2, m21 - m12, m02 - m20, m10 - m01], dim=-1),
+ torch.stack([m21 - m12, q_abs[..., 1] ** 2, m10 + m01, m02 + m20], dim=-1),
+ torch.stack([m02 - m20, m10 + m01, q_abs[..., 2] ** 2, m12 + m21], dim=-1),
+ torch.stack([m10 - m01, m20 + m02, m21 + m12, q_abs[..., 3] ** 2], dim=-1),
+ ],
+ dim=-2,
+ )
+
+ flr = torch.tensor(0.1).to(dtype=q_abs.dtype, device=q_abs.device)
+ quat_candidates = quat_by_rijk / (2.0 * q_abs[..., None].max(flr))
+
+ return (
+ (F.one_hot(q_abs.argmax(dim=-1), num_classes=4)[..., None] * quat_candidates)
+ .sum(dim=-2)
+ .reshape(batch_dim + (4,))
+ )
+
+
+def quaternion_to_matrix(quaternions: torch.Tensor) -> torch.Tensor:
+ """Convert rotations given as quaternions to rotation matrices.
+
+ Args:
+ quaternions: quaternions with real part first,
+ as tensor of shape (..., 4).
+ Returns:
+ Rotation matrices as tensor of shape (..., 3, 3).
+ """
+ r, i, j, k = torch.unbind(quaternions, -1)
+ two_s = 2.0 / (quaternions * quaternions).sum(-1)
+
+ o = torch.stack(
+ (
+ 1 - two_s * (j * j + k * k),
+ two_s * (i * j - k * r),
+ two_s * (i * k + j * r),
+ two_s * (i * j + k * r),
+ 1 - two_s * (i * i + k * k),
+ two_s * (j * k - i * r),
+ two_s * (i * k - j * r),
+ two_s * (j * k + i * r),
+ 1 - two_s * (i * i + j * j),
+ ),
+ -1,
+ )
+ return o.reshape(quaternions.shape[:-1] + (3, 3))
diff --git a/kimodo/meta.py b/kimodo/meta.py
new file mode 100644
index 0000000000000000000000000000000000000000..dd9ff2f75e55e8d1dd94a4b55e1475a139338d32
--- /dev/null
+++ b/kimodo/meta.py
@@ -0,0 +1,80 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Parse and normalize prompt text/duration data from meta dicts."""
+
+import os
+from typing import Any, Optional
+
+from kimodo.tools import load_json
+
+from .sanitize import sanitize_text, sanitize_texts
+
+
+def load_prompts_from_meta(meta_path: str, **kwargs):
+ """Load prompts from a meta dict or file. If fps is provided, the durations are converted to
+ frames.
+
+ Args:
+ meta_path: Path to the meta file.
+ **kwargs: Additional arguments to pass to parse_prompts_from_meta.
+
+ Returns:
+ texts: List of texts.
+ durations: List of durations in seconds or frames.
+ """
+ if not os.path.exists(meta_path):
+ raise FileNotFoundError(f"meta.json not found in input folder: {meta_path}")
+
+ meta = load_json(meta_path)
+ return parse_prompts_from_meta(meta, **kwargs)
+
+
+def parse_prompts_from_meta(
+ meta: dict[str, Any],
+ fps: Optional[float] = None,
+ sanitize: bool = False,
+) -> tuple[list[str], list[float]]:
+ """Parse prompt texts and durations from a meta dict into normalized lists. If fps is provided,
+ the durations are converted to frames.
+
+ Accepts either:
+ - Single prompt: "text" (str) and "duration" (float) in seconds.
+ - Multiple prompts: "texts" (list of str) and "durations" (list of float) in seconds.
+
+ Returns:
+ (texts, durations): texts as list of str, durations as list of float (seconds or frames).
+ Lengths of both lists are equal.
+
+ Raises:
+ ValueError: If meta does not contain a recognized format.
+ """
+ # Single prompt
+ if "text" in meta and "duration" in meta:
+ text = meta["text"]
+ duration = float(meta["duration"])
+ if fps is not None:
+ duration = int(duration * fps)
+ if isinstance(text, list):
+ raise ValueError("meta has 'text' but it is a list; use 'texts' for multiple prompts")
+
+ if sanitize:
+ text = sanitize_text(text)
+ return ([text], [duration])
+
+ # Multiple prompts
+ if "texts" in meta and "durations" in meta:
+ texts = meta["texts"]
+ durations = meta["durations"]
+ if not isinstance(texts, list) or not isinstance(durations, list):
+ raise ValueError("meta 'texts' and 'durations' must be lists")
+ if len(texts) != len(durations):
+ raise ValueError(f"meta 'texts' and 'durations' length mismatch: {len(texts)} vs {len(durations)}")
+ durations = [float(d) for d in durations]
+ if fps is not None:
+ durations = [int(d * fps) for d in durations]
+
+ if sanitize:
+ texts = sanitize_texts(texts)
+ return texts, durations
+
+ raise ValueError("meta must contain either 'text' and 'duration', or 'texts' and 'durations'.")
diff --git a/kimodo/metrics/__init__.py b/kimodo/metrics/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..159e0da823a0f1c75bce3d75da44584d45f11ffc
--- /dev/null
+++ b/kimodo/metrics/__init__.py
@@ -0,0 +1,39 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Evaluation metrics for motion quality (foot skate, contact consistency, constraint following)."""
+
+from .base import (
+ Metric,
+ aggregate_metrics,
+ clear_metrics,
+ compute_metrics,
+)
+from .constraints import ContraintFollow
+from .foot_skate import (
+ FootContactConsistency,
+ FootSkateFromContacts,
+ FootSkateFromHeight,
+ FootSkateRatio,
+)
+from .tmr import (
+ TMR_EmbeddingMetric,
+ TMR_Metric,
+ compute_tmr_per_sample_retrieval,
+ compute_tmr_retrieval_metrics,
+)
+
+__all__ = [
+ "Metric",
+ "ContraintFollow",
+ "FootContactConsistency",
+ "FootSkateFromContacts",
+ "FootSkateFromHeight",
+ "FootSkateRatio",
+ "TMR_EmbeddingMetric",
+ "TMR_Metric",
+ "aggregate_metrics",
+ "clear_metrics",
+ "compute_metrics",
+ "compute_tmr_per_sample_retrieval",
+ "compute_tmr_retrieval_metrics",
+]
diff --git a/kimodo/metrics/base.py b/kimodo/metrics/base.py
new file mode 100644
index 0000000000000000000000000000000000000000..4ca1ebc248c0bd8cbd58eca12a3458f4a65d0745
--- /dev/null
+++ b/kimodo/metrics/base.py
@@ -0,0 +1,66 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Base metric class and batch/aggregate helpers."""
+
+from __future__ import annotations
+
+from collections import defaultdict
+from typing import Dict, List
+
+import torch
+
+
+class Metric:
+ """Base class for metrics that accumulate results over multiple __call__ and expose
+ aggregate()."""
+
+ def __init__(self, **kwargs):
+ self.clear()
+
+ def __call__(self, *args, **kwargs):
+ """Compute metric for current batch, append to saved_metrics, and return the batch
+ result."""
+ metrics = self._compute(*args, **kwargs)
+ for key, val in metrics.items():
+ self.saved_metrics[key].append(val.detach().cpu().float())
+ return metrics
+
+ def _compute(self, **kwargs):
+ """Subclasses implement this to compute metric dict from batch inputs."""
+ raise NotImplementedError()
+
+ def clear(self):
+ """Reset all accumulated metric values."""
+ self.saved_metrics = defaultdict(list)
+
+ def aggregate(self):
+ """Return a dict of concatenated/stacked tensors over all accumulated batches."""
+ output = {}
+ for key, lst in self.saved_metrics.items():
+ try:
+ output[key] = torch.cat(lst)
+ except RuntimeError:
+ output[key] = torch.stack(lst)
+ return output
+
+
+def compute_metrics(metrics_list: List[Metric], metrics_in: Dict) -> Dict:
+ """Run each metric on metrics_in and return the combined dict of batch results."""
+ metrics_out = {}
+ for metric in metrics_list:
+ metrics_out.update(metric(**metrics_in))
+ return metrics_out
+
+
+def aggregate_metrics(metrics_list: List[Metric]) -> Dict:
+ """Return combined aggregated results (concatenated over batches) for all metrics."""
+ metrics_out = {}
+ for metric in metrics_list:
+ metrics_out.update(metric.aggregate())
+ return metrics_out
+
+
+def clear_metrics(metrics_list: List[Metric]) -> None:
+ """Clear accumulated values for all metrics in the list."""
+ for metric in metrics_list:
+ metric.clear()
diff --git a/kimodo/metrics/constraints.py b/kimodo/metrics/constraints.py
new file mode 100644
index 0000000000000000000000000000000000000000..d31ba99cf859cc0d22d71efcb1b324c47c5d7931
--- /dev/null
+++ b/kimodo/metrics/constraints.py
@@ -0,0 +1,87 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Constraint-following metrics."""
+
+from __future__ import annotations
+
+from collections import defaultdict
+from typing import Dict, List, Optional
+
+import torch
+from torch import Tensor
+
+from kimodo.constraints import (
+ EndEffectorConstraintSet,
+ FullBodyConstraintSet,
+ Root2DConstraintSet,
+)
+from kimodo.tools import ensure_batched
+
+from .base import Metric
+
+
+class ContraintFollow(Metric):
+ """Constraint-following metric dispatcher for kimodo constraint sets."""
+
+ def __init__(
+ self,
+ skeleton,
+ root_threshold: float = 0.10,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self.skeleton = skeleton
+ self.root_threshold = root_threshold
+
+ @ensure_batched(posed_joints=4, constraints_lst=2, lengths=1)
+ def _compute(
+ self,
+ posed_joints: Tensor,
+ constraints_lst: Optional[List],
+ lengths: Optional[Tensor] = None,
+ **kwargs,
+ ) -> Dict:
+ if not constraints_lst:
+ return {}
+
+ root_idx = self.skeleton.root_idx
+ output = defaultdict(list)
+
+ for posed_joints_s, constraint_lst_s, lengths_s in zip(posed_joints, constraints_lst, lengths):
+ output_seq = defaultdict(list)
+ for constraint in constraint_lst_s:
+ frame_idx = constraint.frame_indices.to(device=posed_joints_s.device, dtype=torch.long)
+ assert frame_idx.max() < lengths_s, "The constraint is defined outsite the lenght of the motion."
+ if frame_idx.numel() == 0:
+ continue
+
+ if isinstance(constraint, Root2DConstraintSet):
+ pred_root2d = posed_joints_s[frame_idx, root_idx][:, [0, 2]]
+ target = constraint.smooth_root_2d.to(posed_joints_s.device)
+
+ dist = torch.norm(pred_root2d - target, dim=-1)
+ output_seq["constraint_root2d_err"].append(dist)
+ hit = (dist <= self.root_threshold).float()
+ output_seq["constraint_root2d_acc"].append(hit)
+
+ elif isinstance(constraint, FullBodyConstraintSet):
+ pred = posed_joints_s[frame_idx]
+ target = constraint.global_joints_positions.to(posed_joints_s.device)
+ err = torch.norm(pred - target, dim=-1)
+ output_seq["constraint_fullbody_keyframe"].append(err)
+
+ elif isinstance(constraint, EndEffectorConstraintSet):
+ pos_idx = constraint.pos_indices.to(device=posed_joints_s.device, dtype=torch.long)
+ pred = posed_joints_s[frame_idx].index_select(1, pos_idx)
+ target = constraint.global_joints_positions.to(posed_joints_s.device).index_select(1, pos_idx)
+ err = torch.norm(pred - target, dim=-1)
+ output_seq["constraint_end_effector"].append(err)
+
+ # in case we have several same constraints in the list
+ for key, val in output_seq.items():
+ output[key].append(torch.cat(val).mean())
+
+ reduced = {}
+ for key, vals in output.items():
+ reduced[key] = torch.stack(vals, dim=0)
+ return reduced
diff --git a/kimodo/metrics/foot_skate.py b/kimodo/metrics/foot_skate.py
new file mode 100644
index 0000000000000000000000000000000000000000..acf381c97757f4c405dd641af5c81bb36119cd17
--- /dev/null
+++ b/kimodo/metrics/foot_skate.py
@@ -0,0 +1,253 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Foot skate and contact consistency metrics."""
+
+from __future__ import annotations
+
+from typing import Dict, Optional
+
+import torch
+from torch import Tensor
+
+from kimodo.motion_rep.feature_utils import compute_vel_xyz
+from kimodo.motion_rep.feet import foot_detect_from_pos_and_vel
+from kimodo.skeleton import SkeletonBase
+from kimodo.tools import ensure_batched
+
+from .base import Metric
+
+
+def get_four_contacts(fidx: list):
+ if len(fidx) == 4:
+ return fidx
+ if len(fidx) == 6:
+ # For soma77
+ # remove "LeftToeEnd" and "RightToeEnd"
+ fidx = fidx[:2] + fidx[3:5]
+ return fidx
+ raise ValueError("Expects 4 or 6 foot joints (heel/toe per foot)")
+
+
+class FootSkateFromHeight(Metric):
+ """When toe joint is near the floor, measures mean velocity of the toes."""
+
+ def __init__(
+ self,
+ skeleton: SkeletonBase,
+ fps: float,
+ height_thresh: float = 0.05,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self.height_thresh = height_thresh
+ self.skeleton = skeleton
+ self.fps = fps
+
+ @ensure_batched(posed_joints=4, lengths=1)
+ def _compute(
+ self,
+ posed_joints: Tensor,
+ lengths: Optional[Tensor] = None,
+ **kwargs,
+ ) -> Dict:
+ fidx = self.skeleton.foot_joint_idx
+ fidx = get_four_contacts(fidx)
+
+ feet_pos = posed_joints[:, :, fidx]
+ toe_pos = feet_pos[:, :, [1, 3]]
+
+ toe_on_floor = (toe_pos[..., 1] < self.height_thresh)[:, :-1] # y-up [B, T, 2] where [left right]
+
+ dt = 1.0 / self.fps
+ toe_vel = torch.norm(toe_pos[:, 1:] - toe_pos[:, :-1], dim=-1) / dt # [B, nframes-1, 2]
+
+ # compute err
+ contact_toe_vel = toe_vel * toe_on_floor # vel when corresponding toe is on ground
+
+ # account for generated length
+ # since they are velocities use length-1 to avoid inaccurate vel going one frame past len
+ device = toe_on_floor.device
+ len_mask = torch.arange(toe_on_floor.shape[1], device=device)[None, :, None].expand(toe_on_floor.shape) < (
+ lengths[:, None, None] - 1
+ )
+ toe_on_floor = toe_on_floor * len_mask
+ contact_toe_vel = contact_toe_vel * len_mask
+
+ mean_vel = torch.sum(contact_toe_vel, (1, 2)) / (torch.sum(toe_on_floor, (1, 2)) + 1e-6)
+ return {"foot_skate_from_height": mean_vel}
+
+
+class FootSkateFromContacts(Metric):
+ """Measures velocity of the toes and ankles when predicted to be in contact."""
+
+ def __init__(
+ self,
+ skeleton: SkeletonBase,
+ fps: float,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self.skeleton = skeleton
+ self.fps = fps
+
+ @ensure_batched(posed_joints=4, foot_contacts=3, lengths=1)
+ def _compute(
+ self,
+ posed_joints: Tensor,
+ foot_contacts: Tensor,
+ lengths: Optional[Tensor] = None,
+ **kwargs,
+ ) -> Dict:
+ fidx = self.skeleton.foot_joint_idx
+ fidx = get_four_contacts(fidx)
+
+ feet_pos = posed_joints[:, :, fidx]
+ dt = 1.0 / self.fps
+ foot_vel = torch.norm(feet_pos[:, 1:] - feet_pos[:, :-1], dim=-1) / dt
+
+ if foot_contacts.shape[-1] == 6:
+ # For soma77
+ # remove "LeftToeEnd" and "RightToeEnd"
+ foot_contacts = foot_contacts[..., [0, 1, 3, 4]]
+
+ foot_contacts = foot_contacts[:, :-1]
+ vel_err = foot_vel * foot_contacts
+
+ # account for generated length
+ # since they are velocities use length-1 to avoid inaccurate vel going one frame past len
+ device = foot_contacts.device
+ len_mask = torch.arange(foot_contacts.shape[1], device=device)[None, :, None].expand(foot_contacts.shape) < (
+ lengths[:, None, None] - 1
+ )
+ foot_contacts = foot_contacts * len_mask
+ vel_err = vel_err * len_mask
+
+ mean_vel = torch.sum(vel_err, (1, 2)) / (torch.sum(foot_contacts, (1, 2)) + 1e-6) # mean over contacting frames
+
+ # Compute max velocity error across all feet and frames (per batch)
+ max_vel = vel_err.amax(dim=(1, 2)) # [B]
+
+ return {
+ "foot_skate_from_pred_contacts": mean_vel,
+ "foot_skate_max_vel": max_vel,
+ }
+
+
+class FootSkateRatio(Metric):
+ """Compute fraction of frames where the foot skates when it is on the ground.
+
+ Inspired by GMD: https://github.com/korrawe/guided-motion-diffusion/blob/main/data_loaders/humanml/utils/metrics.py#L204
+ """
+
+ def __init__(
+ self,
+ skeleton: SkeletonBase,
+ fps: float,
+ height_thresh=0.05,
+ vel_thresh=0.2,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self.height_thresh = height_thresh
+ self.vel_thresh = vel_thresh
+
+ self.skeleton = skeleton
+ self.fps = fps
+
+ @ensure_batched(posed_joints=4, foot_contacts=3, lengths=1)
+ def _compute(
+ self,
+ posed_joints: Tensor,
+ foot_contacts: Tensor,
+ lengths: Optional[Tensor] = None,
+ **kwargs,
+ ) -> Dict:
+ fidx = self.skeleton.foot_joint_idx
+ fidx = get_four_contacts(fidx)
+
+ feet_pos = posed_joints[:, :, fidx]
+ toe_pos = feet_pos[:, :, [1, 3]]
+
+ toe_on_floor = toe_pos[..., 1] < self.height_thresh # y-up [B, T, 2] where [left right]
+ # current and next frame on floor to consider it in contact
+ toe_on_floor = torch.logical_and(toe_on_floor[:, :-1], toe_on_floor[:, 1:]) # [B, T-1, 2]
+
+ dt = 1.0 / self.fps
+ toe_vel = torch.norm(toe_pos[:, 1:] - toe_pos[:, :-1], dim=-1) / dt # [B, nframes-1, 2]
+
+ # compute err
+ contact_toe_vel = toe_vel * toe_on_floor # vel when corresponding toe is on ground
+
+ # account for generated length
+ # since they are velocities use length-1 to avoid inaccurate vel going one frame past len
+ device = toe_on_floor.device
+ len_mask = torch.arange(toe_on_floor.shape[1], device=device)[None, :, None].expand(toe_on_floor.shape) < (
+ lengths[:, None, None] - 1
+ )
+ toe_on_floor = toe_on_floor * len_mask
+ contact_toe_vel = contact_toe_vel * len_mask
+
+ # skating if velocity during contact > thresh
+ toe_skate = contact_toe_vel > self.vel_thresh
+ skate_ratio = torch.sum(toe_skate, (1, 2)) / (torch.sum(toe_on_floor, (1, 2)) + 1e-6)
+ return {"foot_skate_ratio": skate_ratio}
+
+
+class FootContactConsistency(Metric):
+ """Measures consistency between heuristic detected foot contacts (from height and velocity) and
+ predicted foot contacts.
+
+ i.e. accuracy of how well predicted matches heuristic.
+ """
+
+ def __init__(
+ self,
+ skeleton: SkeletonBase,
+ fps: float,
+ vel_thresh: float = 0.15,
+ height_thresh: float = 0.10,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self.vel_thresh = vel_thresh
+ self.height_thresh = height_thresh
+
+ self.skeleton = skeleton
+ self.fps = fps
+
+ @ensure_batched(posed_joints=4, foot_contacts=3, lengths=1)
+ def _compute(
+ self,
+ posed_joints: Tensor,
+ foot_contacts: Tensor,
+ lengths: Optional[Tensor] = None,
+ **kwargs,
+ ) -> Dict:
+ velocity = compute_vel_xyz(posed_joints, float(self.fps), lengths=lengths)
+ heuristic_contacts = foot_detect_from_pos_and_vel(
+ posed_joints,
+ velocity,
+ self.skeleton,
+ self.vel_thresh,
+ self.height_thresh,
+ )
+
+ if foot_contacts.shape[-1] == 6:
+ # For soma77
+ # remove "LeftToeEnd" and "RightToeEnd"
+ foot_contacts = foot_contacts[..., [0, 1, 3, 4]]
+
+ num_contacts = foot_contacts.shape[-1]
+ incorrect = torch.logical_xor(heuristic_contacts, foot_contacts)
+ # account for generated length
+ # since they are velocities, use length-1 to avoid inaccurate vel going one frame past len
+ device = foot_contacts.device
+ len_mask = torch.arange(foot_contacts.shape[1], device=device)[None, :, None].expand(foot_contacts.shape) < (
+ lengths[:, None, None] - 1
+ )
+ incorrect = incorrect * len_mask
+
+ incorrect_ratio = torch.sum(incorrect, (1, 2)) / (num_contacts * (lengths - 1))
+ accuracy = 1 - incorrect_ratio
+
+ return {"foot_contact_consistency": accuracy}
diff --git a/kimodo/metrics/tmr.py b/kimodo/metrics/tmr.py
new file mode 100644
index 0000000000000000000000000000000000000000..aa0c1ac341477430253607dc5130bc02812eb1df
--- /dev/null
+++ b/kimodo/metrics/tmr.py
@@ -0,0 +1,545 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""TMR evaluation metrics: text-motion retrieval, R-Precision, and related scores."""
+
+from __future__ import annotations
+
+from collections import defaultdict
+from typing import Any, Dict, List, Optional
+
+import numpy as np
+import torch
+from scipy import linalg
+from torch import Tensor
+
+from kimodo.model.tmr import TMR
+
+from .base import Metric
+
+
+# Scores are between 0 and 1
+def get_score_matrix_unit(x, y):
+ sim_matrix = np.einsum("b i, c i -> b c", x, y)
+ scores = sim_matrix / 2 + 0.5
+ return scores
+
+
+def get_scores_unit(x, y):
+ similarity = np.einsum("... i, ... i", x, y)
+ scores = similarity / 2 + 0.5
+ return scores
+
+
+def compute_tmr_per_sample_retrieval(
+ motion_emb: np.ndarray,
+ text_emb: np.ndarray,
+ sample_ids: List[str],
+ texts: List[str],
+ top_k: int = 5,
+) -> List[Dict[str, Any]]:
+ """For each sample (text query i), compute t2m rank of motion i and top-k retrieved motions with
+ ids and texts.
+
+ Returns list of dicts: [{"rank": int, "top_k": [{"id": str, "text": str}, ...]}, ...].
+ """
+ motion_emb = np.asarray(motion_emb).squeeze()
+ text_emb = np.asarray(text_emb).squeeze()
+ if motion_emb.ndim == 1:
+ motion_emb = motion_emb[np.newaxis, :]
+ if text_emb.ndim == 1:
+ text_emb = text_emb[np.newaxis, :]
+ n = motion_emb.shape[0]
+ assert text_emb.shape[0] == n and len(sample_ids) == n and len(texts) == n
+ scores = get_score_matrix_unit(text_emb, motion_emb)
+ out: List[Dict[str, Any]] = []
+ for i in range(n):
+ row = np.asarray(scores[i])
+ order = np.argsort(row)[::-1]
+ rank = int(np.where(order == i)[0][0]) + 1
+ top_indices = order[:top_k]
+ top_k_list = [{"id": sample_ids[j], "text": texts[j]} for j in top_indices]
+ out.append({"rank": rank, "top_k": top_k_list})
+ return out
+
+
+class TMR_Metric(Metric):
+ def __init__(
+ self,
+ tmr_model: TMR,
+ ranks: List = [1, 2, 3, 5, 10],
+ ranks_rounding=2,
+ **kwargs,
+ ):
+ super().__init__(**kwargs)
+ self.tmr_model = tmr_model
+ self.ranks = ranks
+ self.ranks_rounding = ranks_rounding
+
+ def clear(self):
+ self.saved_metrics = defaultdict(list)
+ self.saved_text_latents = []
+ self.saved_motion_gen_latents = []
+ self.saved_motion_gt_latents = []
+
+ def _compute(
+ self,
+ motion_rep,
+ pred_joints_output: Dict,
+ gt_joints_output: Dict,
+ text_x_dict: Dict,
+ lengths: Tensor,
+ **kwargs,
+ ) -> Dict:
+ pred_posed_joints = pred_joints_output["posed_joints"]
+ original_skeleton = motion_rep.skeleton if motion_rep is not None else None
+ latents_motion = self.tmr_model.encode_motion(
+ pred_posed_joints,
+ lengths=lengths,
+ original_skeleton=original_skeleton,
+ unit_vector=True,
+ )
+ latents_motion = latents_motion.cpu().numpy()
+
+ if isinstance(text_x_dict, dict) and "texts" in text_x_dict:
+ latents_text = self.tmr_model.encode_raw_text(text_x_dict["texts"], unit_vector=True)
+ else:
+ latents_text = self.tmr_model.encode_text(text_x_dict, unit_vector=True)
+ if latents_text.dim() == 1:
+ latents_text = latents_text.unsqueeze(0)
+ latents_text = latents_text.cpu().numpy()
+
+ self.saved_text_latents.append(latents_text)
+ self.saved_motion_gen_latents.append(latents_motion)
+
+ scores_text = get_scores_unit(latents_motion, latents_text)
+ output = {"TMR/t2m_sim": scores_text}
+
+ if gt_joints_output is not None and "posed_joints" in gt_joints_output:
+ gt_posed_joints = gt_joints_output["posed_joints"]
+ gt_latents_motion = self.tmr_model.encode_motion(
+ gt_posed_joints,
+ lengths=lengths,
+ original_skeleton=original_skeleton,
+ unit_vector=True,
+ )
+ gt_latents_motion = gt_latents_motion.cpu().numpy()
+ self.saved_motion_gt_latents.append(gt_latents_motion)
+
+ gt_scores_text = get_scores_unit(gt_latents_motion, latents_text)
+ scores_motion = get_scores_unit(latents_motion, gt_latents_motion)
+
+ output["TMR/t2m_gt_sim"] = gt_scores_text
+ output["TMR/m2m_sim"] = scores_motion
+
+ # pytorch tensors
+ for key, val in output.items():
+ output[key] = torch.tensor(val)
+ return output
+
+ def aggregate(self):
+ output = {}
+ for key, lst in self.saved_metrics.items():
+ output[key] = np.concatenate(lst)
+
+ assert self.saved_text_latents, "Should call the metric at least once."
+
+ text_latents = np.concatenate(self.saved_text_latents)
+ motion_gen_latents = np.concatenate(self.saved_motion_gen_latents)
+
+ batch_size = len(text_latents)
+ assert text_latents.shape == motion_gen_latents.shape
+
+ scores_t2m = get_score_matrix_unit(text_latents, motion_gen_latents)
+ scores_t2t = get_score_matrix_unit(text_latents, text_latents)
+
+ t2m_metrics = contrastive_metrics(
+ scores=scores_t2m,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=2,
+ )
+
+ for key, val in t2m_metrics.items():
+ output["TMR/t2m_R/" + key] = val
+
+ mu_gen, cov_gen = calculate_activation_statistics(motion_gen_latents)
+ mu_text, cov_text = calculate_activation_statistics(text_latents)
+
+ fid_gen_text = calculate_frechet_distance(mu_gen, cov_gen, mu_text, cov_text)
+ output["TMR/FID/gen_text"] = fid_gen_text
+
+ if self.saved_motion_gt_latents:
+ motion_gt_latents = np.concatenate(self.saved_motion_gt_latents)
+ assert motion_gt_latents.shape == motion_gen_latents.shape
+
+ scores_m2gm = get_score_matrix_unit(motion_gen_latents, motion_gt_latents)
+ scores_t2gm = get_score_matrix_unit(text_latents, motion_gt_latents)
+
+ m2gm_metrics = contrastive_metrics(
+ scores=scores_m2gm,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=2,
+ )
+ for key, val in m2gm_metrics.items():
+ output["TMR/m2m_R/" + key] = val
+
+ t2gm_metrics = contrastive_metrics(
+ scores=scores_t2gm,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=2,
+ )
+ for key, val in t2gm_metrics.items():
+ output["TMR/t2m_gt_R/" + key] = val
+
+ mu_gt_motion, cov_gt_motion = calculate_activation_statistics(motion_gt_latents)
+ fid_gen_motion = calculate_frechet_distance(
+ mu_gen,
+ cov_gen,
+ mu_gt_motion,
+ cov_gt_motion,
+ )
+ output["TMR/FID/gen_gt"] = fid_gen_motion
+
+ fid_gt_text = calculate_frechet_distance(
+ mu_gt_motion,
+ cov_gt_motion,
+ mu_text,
+ cov_text,
+ )
+ output["TMR/FID/gt_text"] = fid_gt_text
+
+ for key, val in output.items():
+ if isinstance(val, (int, float, np.integer, np.floating)):
+ val = torch.tensor([val for _ in range(batch_size)])
+
+ if isinstance(val, np.ndarray):
+ val = torch.from_numpy(val)
+
+ output[key] = val.cpu().float()
+ return output
+
+
+class TMR_EmbeddingMetric(Metric):
+ """TMR metrics from precomputed motion and text embeddings (no model load).
+
+ Use in the loop: pass motion_emb and text_emb per sample; aggregate() computes retrieval metrics.
+ """
+
+ def __init__(self, ranks_rounding: int = 2, **kwargs):
+ super().__init__(**kwargs)
+ self.ranks_rounding = ranks_rounding
+
+ def clear(self):
+ self.saved_metrics = defaultdict(list)
+ self.saved_text_latents = []
+ self.saved_motion_gen_latents = []
+ self.saved_motion_gt_latents = []
+
+ def _compute(
+ self,
+ motion_emb=None,
+ text_emb=None,
+ gt_motion_emb=None,
+ **kwargs,
+ ) -> Dict:
+ if motion_emb is None or text_emb is None:
+ return {}
+ motion_emb = np.asarray(motion_emb)
+ text_emb = np.asarray(text_emb)
+ if motion_emb.ndim == 1:
+ motion_emb = motion_emb[np.newaxis, :]
+ if text_emb.ndim == 1:
+ text_emb = text_emb[np.newaxis, :]
+ self.saved_text_latents.append(text_emb)
+ self.saved_motion_gen_latents.append(motion_emb)
+ if gt_motion_emb is not None:
+ gt_motion_emb = np.asarray(gt_motion_emb)
+ if gt_motion_emb.ndim == 1:
+ gt_motion_emb = gt_motion_emb[np.newaxis, :]
+ self.saved_motion_gt_latents.append(gt_motion_emb)
+ scores = get_scores_unit(motion_emb, text_emb)
+ return {"TMR/t2m_sim": torch.tensor(scores, dtype=torch.float32)}
+
+ def aggregate(self):
+ output = {}
+ for key, lst in self.saved_metrics.items():
+ output[key] = np.concatenate(lst)
+ if not self.saved_text_latents:
+ return output
+ text_latents = np.concatenate(self.saved_text_latents)
+ motion_gen_latents = np.concatenate(self.saved_motion_gen_latents)
+ batch_size = len(text_latents)
+ assert text_latents.shape == motion_gen_latents.shape
+ scores_t2m = get_score_matrix_unit(text_latents, motion_gen_latents)
+ scores_t2t = get_score_matrix_unit(text_latents, text_latents)
+ t2m_metrics = contrastive_metrics(
+ scores=scores_t2m,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=self.ranks_rounding,
+ )
+ for key, val in t2m_metrics.items():
+ output["TMR/t2m_R/" + key] = val
+ if batch_size >= 2:
+ mu_gen, cov_gen = calculate_activation_statistics(motion_gen_latents)
+ mu_text, cov_text = calculate_activation_statistics(text_latents)
+ output["TMR/FID/gen_text"] = calculate_frechet_distance(mu_gen, cov_gen, mu_text, cov_text)
+ else:
+ output["TMR/FID/gen_text"] = float("nan")
+ if self.saved_motion_gt_latents:
+ motion_gt_latents = np.concatenate(self.saved_motion_gt_latents)
+ assert motion_gt_latents.shape == motion_gen_latents.shape
+ scores_m2gm = get_score_matrix_unit(motion_gen_latents, motion_gt_latents)
+ scores_t2gm = get_score_matrix_unit(text_latents, motion_gt_latents)
+ m2gm_metrics = contrastive_metrics(
+ scores=scores_m2gm,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=self.ranks_rounding,
+ )
+ for key, val in m2gm_metrics.items():
+ output["TMR/m2m_R/" + key] = val
+ t2gm_metrics = contrastive_metrics(
+ scores=scores_t2gm,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=self.ranks_rounding,
+ )
+ for key, val in t2gm_metrics.items():
+ output["TMR/t2m_gt_R/" + key] = val
+ if batch_size >= 2:
+ mu_gt_motion, cov_gt_motion = calculate_activation_statistics(motion_gt_latents)
+ output["TMR/FID/gen_gt"] = calculate_frechet_distance(mu_gen, cov_gen, mu_gt_motion, cov_gt_motion)
+ output["TMR/FID/gt_text"] = calculate_frechet_distance(mu_gt_motion, cov_gt_motion, mu_text, cov_text)
+ else:
+ output["TMR/FID/gen_gt"] = float("nan")
+ output["TMR/FID/gt_text"] = float("nan")
+ for key, val in output.items():
+ if isinstance(val, (int, float, np.integer, np.floating)):
+ val = torch.tensor([val for _ in range(batch_size)])
+ if isinstance(val, np.ndarray):
+ val = torch.from_numpy(val)
+ output[key] = val.cpu().float()
+ return output
+
+
+def compute_tmr_retrieval_metrics(
+ motion_emb: np.ndarray,
+ text_emb: np.ndarray,
+ gt_motion_emb: Optional[np.ndarray] = None,
+ rounding: int = 2,
+) -> Dict[str, float]:
+ """Compute TMR retrieval metrics from precomputed embeddings."""
+ if motion_emb.shape != text_emb.shape:
+ raise ValueError(f"Expected same shape for motion/text embeddings, got {motion_emb.shape} vs {text_emb.shape}")
+
+ scores_t2m = get_score_matrix_unit(text_emb, motion_emb)
+ scores_t2t = get_score_matrix_unit(text_emb, text_emb)
+
+ output: Dict[str, float] = {}
+ t2m_metrics = contrastive_metrics(
+ scores=scores_t2m,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=rounding,
+ )
+ for key, val in t2m_metrics.items():
+ output[f"TMR/t2m_R/{key}"] = float(val)
+
+ n_samples = len(motion_emb)
+ if n_samples >= 2:
+ mu_gen, cov_gen = calculate_activation_statistics(motion_emb)
+ mu_text, cov_text = calculate_activation_statistics(text_emb)
+ output["TMR/FID/gen_text"] = float(calculate_frechet_distance(mu_gen, cov_gen, mu_text, cov_text))
+ else:
+ output["TMR/FID/gen_text"] = float("nan")
+
+ if gt_motion_emb is not None:
+ if gt_motion_emb.shape != motion_emb.shape:
+ raise ValueError(f"Expected gt motion embeddings shape {motion_emb.shape}, got {gt_motion_emb.shape}")
+
+ scores_m2gm = get_score_matrix_unit(motion_emb, gt_motion_emb)
+ scores_t2gm = get_score_matrix_unit(text_emb, gt_motion_emb)
+
+ m2gm_metrics = contrastive_metrics(
+ scores=scores_m2gm,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=rounding,
+ )
+ for key, val in m2gm_metrics.items():
+ output[f"TMR/m2m_R/{key}"] = float(val)
+
+ t2gm_metrics = contrastive_metrics(
+ scores=scores_t2gm,
+ scores_t2t=scores_t2t,
+ threshold=0.99,
+ rounding=rounding,
+ )
+ for key, val in t2gm_metrics.items():
+ output[f"TMR/t2m_gt_R/{key}"] = float(val)
+
+ if n_samples >= 2:
+ mu_gt_motion, cov_gt_motion = calculate_activation_statistics(gt_motion_emb)
+ output["TMR/FID/gen_gt"] = float(calculate_frechet_distance(mu_gen, cov_gen, mu_gt_motion, cov_gt_motion))
+ output["TMR/FID/gt_text"] = float(calculate_frechet_distance(mu_gt_motion, cov_gt_motion, mu_text, cov_text))
+ else:
+ output["TMR/FID/gen_gt"] = float("nan")
+ output["TMR/FID/gt_text"] = float("nan")
+
+ return output
+
+
+def all_contrastive_metrics(sims, emb=None, threshold=None, rounding=2, return_cols=False):
+ text_selfsim = None
+ if emb is not None:
+ text_selfsim = emb @ emb.T
+
+ t2m_m, t2m_cols = contrastive_metrics(sims, text_selfsim, threshold, return_cols=True, rounding=rounding)
+ m2t_m, m2t_cols = contrastive_metrics(sims.T, text_selfsim, threshold, return_cols=True, rounding=rounding)
+
+ all_m = {}
+ for key in t2m_m:
+ all_m[f"t2m/{key}"] = t2m_m[key]
+ all_m[f"m2t/{key}"] = m2t_m[key]
+
+ all_m["t2m/len"] = float(len(sims))
+ all_m["m2t/len"] = float(len(sims[0]))
+ if return_cols:
+ return all_m, t2m_cols, m2t_cols
+ return all_m
+
+
+def contrastive_metrics(
+ scores,
+ scores_t2t=None,
+ threshold=None,
+ rounding=2,
+):
+ n, m = scores.shape
+ assert n == m
+ num_queries = n
+
+ dists = -scores
+ sorted_dists = np.sort(dists, axis=1)
+ # GT is in the diagonal
+ gt_dists = np.diag(dists)[:, None]
+
+ if scores_t2t is not None and threshold is not None:
+ real_threshold = 2 * threshold - 1
+ idx = np.argwhere(scores_t2t > real_threshold)
+ partition = np.unique(idx[:, 0], return_index=True)[1]
+ # take as GT the minimum score of similar values
+ gt_dists = np.minimum.reduceat(dists[tuple(idx.T)], partition)
+ gt_dists = gt_dists[:, None]
+
+ rows, cols = np.where((sorted_dists - gt_dists) == 0) # find column position of GT
+
+ # if there are ties
+ if rows.size > num_queries:
+ assert np.unique(rows).size == num_queries, "issue in metric evaluation"
+ avg_cols = break_ties_average(sorted_dists, gt_dists)
+ cols = avg_cols
+
+ msg = "expected ranks to match queries ({} vs {}) "
+ assert cols.size == num_queries, msg
+
+ metrics = {}
+ vals = [str(x).zfill(2) for x in [1, 2, 3, 5, 10]]
+ for val in vals:
+ metrics[f"R{val}"] = 100 * float(np.sum(cols < int(val))) / num_queries
+
+ metrics["MedR"] = float(np.median(cols) + 1)
+ metrics["len"] = num_queries
+
+ if rounding is not None:
+ for key in metrics:
+ metrics[key] = round(metrics[key], rounding)
+ return metrics
+
+
+def break_ties_average(sorted_dists, gt_dists):
+ # fast implementation, based on this code:
+ # https://stackoverflow.com/a/49239335
+ locs = np.argwhere((sorted_dists - gt_dists) == 0)
+
+ # Find the split indices
+ steps = np.diff(locs[:, 0])
+ splits = np.nonzero(steps)[0] + 1
+ splits = np.insert(splits, 0, 0)
+
+ # Compute the result columns
+ summed_cols = np.add.reduceat(locs[:, 1], splits)
+ counts = np.diff(np.append(splits, locs.shape[0]))
+ avg_cols = summed_cols / counts
+ return avg_cols
+
+
+def calculate_activation_statistics(activations):
+ """
+ Params:
+ -- activation: num_samples x dim_feat
+ Returns:
+ -- mu: dim_feat
+ -- sigma: dim_feat x dim_feat
+ """
+ mu = np.mean(activations, axis=0)
+ cov = np.cov(activations, rowvar=False)
+ return mu, cov
+
+
+def calculate_frechet_distance(mu1, sigma1, mu2, sigma2, eps=1e-6):
+ """Numpy implementation of the Frechet Distance. The Frechet distance between two multivariate
+ Gaussians X_1 ~ N(mu_1, C_1)
+
+ and X_2 ~ N(mu_2, C_2) is
+ d^2 = ||mu_1 - mu_2||^2 + Tr(C_1 + C_2 - 2*sqrt(C_1*C_2)).
+ Stable version by Dougal J. Sutherland.
+ Params:
+ -- mu1 : Numpy array containing the activations of a layer of the
+ inception net (like returned by the function 'get_predictions')
+ for generated samples.
+ -- mu2 : The sample mean over activations, precalculated on an
+ representative dataset set.
+ -- sigma1: The covariance matrix over activations for generated samples.
+ -- sigma2: The covariance matrix over activations, precalculated on an
+ representative dataset set.
+ Returns:
+ -- : The Frechet Distance.
+ """
+
+ mu1 = np.atleast_1d(mu1)
+ mu2 = np.atleast_1d(mu2)
+
+ sigma1 = np.atleast_2d(sigma1)
+ sigma2 = np.atleast_2d(sigma2)
+
+ assert mu1.shape == mu2.shape, "Training and test mean vectors have different lengths"
+ assert sigma1.shape == sigma2.shape, "Training and test covariances have different dimensions"
+
+ diff = mu1 - mu2
+
+ # Product might be almost singular
+ covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)
+ if not np.isfinite(covmean).all():
+ msg = ("fid calculation produces singular product; " "adding %s to diagonal of cov estimates") % eps
+ print(msg)
+ offset = np.eye(sigma1.shape[0]) * eps
+ covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
+
+ # Numerical error might give slight imaginary component
+ if np.iscomplexobj(covmean):
+ if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
+ # try again with diagonal %s
+ offset = np.eye(sigma1.shape[0]) * eps
+ covmean = linalg.sqrtm((sigma1 + offset).dot(sigma2 + offset))
+ if not np.allclose(np.diagonal(covmean).imag, 0, atol=1e-3):
+ m = np.max(np.abs(covmean.imag))
+ raise ValueError("Imaginary component {}".format(m))
+ covmean = covmean.real
+
+ tr_covmean = np.trace(covmean)
+
+ return diff.dot(diff) + np.trace(sigma1) + np.trace(sigma2) - 2 * tr_covmean
diff --git a/kimodo/model/__init__.py b/kimodo/model/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..a1d4cd87748f15bd0538d256763a52b61ed480d4
--- /dev/null
+++ b/kimodo/model/__init__.py
@@ -0,0 +1,31 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Kimodo model package: main model class, text encoders, and loading utilities."""
+
+from .common import resolve_target
+from .kimodo_model import Kimodo
+from .llm2vec import LLM2VecEncoder
+from .load_model import load_model
+from .loading import (
+ AVAILABLE_MODELS,
+ DEFAULT_MODEL,
+ DEFAULT_TEXT_ENCODER_URL,
+ MODEL_NAMES,
+ load_checkpoint_state_dict,
+)
+from .tmr import TMR
+from .twostage_denoiser import TwostageDenoiser
+
+__all__ = [
+ "Kimodo",
+ "LLM2VecEncoder",
+ "TMR",
+ "TwostageDenoiser",
+ "load_model",
+ "load_checkpoint_state_dict",
+ "resolve_target",
+ "AVAILABLE_MODELS",
+ "DEFAULT_MODEL",
+ "DEFAULT_TEXT_ENCODER_URL",
+ "MODEL_NAMES",
+]
diff --git a/kimodo/model/backbone.py b/kimodo/model/backbone.py
new file mode 100644
index 0000000000000000000000000000000000000000..014f6599f2b0ff7fddfabb9a47db8d0962941b11
--- /dev/null
+++ b/kimodo/model/backbone.py
@@ -0,0 +1,312 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Transformer backbone: padding, masking, and encoder stack for the denoiser."""
+
+import logging
+from typing import Optional, Union
+
+import torch
+from omegaconf import ListConfig
+from pydantic.dataclasses import dataclass
+from torch import Tensor, nn
+from torch.nn import TransformerEncoder, TransformerEncoderLayer
+
+from kimodo.tools import validate
+
+log = logging.getLogger(__name__)
+
+
+def pad_x_and_mask_to_fixed_size(x: Tensor, mask: Tensor, size: int):
+ """Pad a feature vector x and the mask to always have the same size.
+
+ Args:
+ x (torch.Tensor): [B, T, D]
+ mask (torch.Tensor): [B, T]
+ size (int)
+ Returns:
+ torch.Tensor: [B, size, D]
+ torch.Tensor: [B, size]
+ """
+
+ batch_size, cur_max_size, dim = x.shape[0], x.shape[1], x.shape[2]
+
+ if cur_max_size == size:
+ # already padded to this size, probably in the collate function
+ return x, mask
+
+ if cur_max_size > size:
+ # This issue should have been handled in the collate function
+ # usefull as a check for test time
+ log.warn("The size of the tensor is larger than the maximum size. Cropping the input..")
+ cur_max_size = size
+
+ new_x = torch.zeros(
+ (batch_size, size, dim),
+ dtype=x.dtype,
+ device=x.device,
+ )
+ new_x[:, :cur_max_size] = x
+
+ # same for the mask
+ new_mask = torch.zeros(
+ (batch_size, size),
+ dtype=mask.dtype,
+ device=mask.device,
+ )
+ new_mask[:, :cur_max_size] = mask
+ return new_x, new_mask
+
+
+@dataclass(frozen=True, config=dict(extra="forbid", arbitrary_types_allowed=True))
+class TransformerEncoderBlockConfig:
+ """Configuration for the transformer encoder backbone."""
+
+ # input features dimension
+ input_dim: int
+ # output features dimension
+ output_dim: int
+
+ # skeleton object
+ skeleton: object
+
+ # dimension of the text embeddings
+ llm_shape: Union[list[int], ListConfig]
+
+ # mask the text or not
+ use_text_mask: bool
+
+ # latent dimension of the model
+ latent_dim: int
+ # dimension of the feedforward network in transformer
+ ff_size: int
+ # num layers in transformer
+ num_layers: int
+ # num heads in transformer
+ num_heads: int
+ # activation in transformer
+ activation: str
+ # dropout rate for the transformer
+ dropout: float
+ # dropout rate for the positional embeddings
+ pe_dropout: float
+ # use norm first or not
+ norm_first: bool = False
+ # artificially extend the number of text tokens
+ num_text_tokens_override: Optional[int] = None
+
+ # Input first heading angle
+ input_first_heading_angle: bool = False
+
+
+class TransformerEncoderBlock(nn.Module):
+ @validate(TransformerEncoderBlockConfig, save_args=True, super_init=True)
+ def __init__(self, conf):
+ self.nbjoints = self.skeleton.nbjoints
+ llm_dim = self.llm_shape[-1]
+ self.embed_text = nn.Linear(llm_dim, self.latent_dim)
+
+ self.sequence_pos_encoder = PositionalEncoding(self.latent_dim, self.pe_dropout)
+
+ # maximum number of tokens
+ self.num_text_tokens = self.llm_shape[0]
+ if self.num_text_tokens_override is not None:
+ self.num_text_tokens = self.num_text_tokens_override
+
+ self.embed_timestep = TimestepEmbedder(self.latent_dim, self.sequence_pos_encoder)
+
+ self.input_linear = nn.Linear(self.input_dim, self.latent_dim)
+ self.output_linear = nn.Linear(self.latent_dim, self.output_dim)
+ self.linear_first_heading_angle = nn.Linear(2, self.latent_dim)
+
+ trans_enc_layer = TransformerEncoderLayer(
+ d_model=self.latent_dim,
+ nhead=self.num_heads,
+ dim_feedforward=self.ff_size,
+ dropout=self.dropout,
+ activation=self.activation,
+ batch_first=True,
+ norm_first=self.norm_first,
+ )
+ self.seqTransEncoder = TransformerEncoder(
+ trans_enc_layer,
+ num_layers=self.num_layers,
+ enable_nested_tensor=False,
+ )
+
+ def forward(
+ self,
+ x: Tensor,
+ x_pad_mask: torch.Tensor,
+ text_feat: torch.Tensor,
+ text_feat_pad_mask: torch.Tensor,
+ timesteps: Tensor,
+ first_heading_angle: Optional[Tensor] = None,
+ ) -> Tensor:
+ """
+ Args:
+ x (torch.Tensor): [B, T, dim_motion] current noisy motion
+ x_pad_mask (torch.Tensor): [B, T] attention mask, positions with True are allowed to attend, False are not
+ text_feat (torch.Tensor): [B, max_text_len, llm_dim] embedded text prompts
+ text_feat_pad_mask (torch.Tensor): [B, max_text_len] attention mask, positions with True are allowed to attend, False are not
+ timesteps (torch.Tensor): [B,] current denoising step
+
+ Returns:
+ torch.Tensor: [B, T, output_dim]
+ """
+ batch_size = len(x)
+ x = self.input_linear(x) # [B, T, D]
+
+ # Pad the text tokens + mask to always have the same size == self.num_text_tokens
+ # done here if it was not done in the collate function
+ if self.num_text_tokens is not None:
+ text_feat, text_feat_pad_mask = pad_x_and_mask_to_fixed_size(
+ text_feat,
+ text_feat_pad_mask,
+ self.num_text_tokens,
+ )
+
+ # Encode the text features and the time information
+ emb_text = self.embed_text(text_feat) # [B, max_text_len, D]
+ emb_time = self.embed_timestep(timesteps) # [B, 1, D]
+
+ # Create mask for the time information
+ time_mask = torch.ones((batch_size, 1), dtype=bool, device=x.device)
+
+ # Create the prefix features (text, time, etc): [B, max_text_len + 1 + etc]
+ prefix_feats = torch.cat((emb_text, emb_time), axis=1)
+
+ # Behavior from old code: not use text mask -> True for all the tokens
+ if not self.use_text_mask:
+ text_feat_pad_mask = torch.ones(
+ (batch_size, emb_text.shape[1]),
+ dtype=torch.bool,
+ device=x.device,
+ )
+
+ prefix_mask = torch.cat((text_feat_pad_mask, time_mask), axis=1)
+
+ # add the input first heading angle
+ if self.input_first_heading_angle:
+ assert first_heading_angle is not None, "The first heading angle is mandatory for this model"
+ # cos(angle) / sin(angle)
+ first_heading_angle_feats = torch.stack(
+ [
+ torch.cos(first_heading_angle),
+ torch.sin(first_heading_angle),
+ ],
+ axis=-1,
+ )
+
+ first_heading_angle_feats = self.linear_first_heading_angle(first_heading_angle_feats)
+ first_heading_angle_feats = first_heading_angle_feats[:, None] # for cat
+ first_heading_angle_mask = torch.ones(
+ (batch_size, 1),
+ dtype=bool,
+ device=x.device,
+ )
+ prefix_feats = torch.cat((prefix_feats, first_heading_angle_feats), axis=1)
+ prefix_mask = torch.cat((prefix_mask, first_heading_angle_mask), axis=1)
+
+ # compute the number of prefix features
+ pose_start_ind = prefix_feats.shape[1]
+
+ # Concatenate prefix and x: [B, len(prefix) + T, D]
+ xseq = torch.cat((prefix_feats, x), axis=1)
+
+ # Concatenate the masks and negate them: [B, len(prefix) + T]
+ src_key_padding_mask = ~torch.cat((prefix_mask, x_pad_mask), axis=1)
+
+ # Add positional encoding
+ xseq = self.sequence_pos_encoder(xseq)
+
+ # Input to the transformer and keep the motion indexes
+ if isinstance(self.seqTransEncoder, nn.TransformerEncoder):
+ assert not self.seqTransEncoder.use_nested_tensor, "Flash attention should be disabled due to bug!"
+
+ output = self.seqTransEncoder(
+ xseq,
+ src_key_padding_mask=src_key_padding_mask,
+ )
+ output = output[:, pose_start_ind:] # [B, T, D]
+ output = self.output_linear(output) # [B, T, OD]
+ return output
+
+
+class PositionalEncoding(nn.Module):
+ """Non-learned positional encoding."""
+
+ def __init__(
+ self,
+ d_model: int,
+ dropout: Optional[float] = 0.1,
+ max_len: Optional[int] = 5000,
+ ):
+ """
+ Args:
+ d_model (int): input dim
+ dropout (Optional[float] = 0.1): dropout probability on output
+ max_len (Optional[int] = 5000): maximum sequence length
+ """
+ super(PositionalEncoding, self).__init__()
+ self.dropout = nn.Dropout(p=dropout)
+
+ pe = torch.zeros(max_len, d_model)
+ position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
+
+ # Note: have to replace torch.exp() and math.log() with torch.pow()
+ # due to MKL exp() and ln() throws floating point exceptions on certain CPUs
+ # see corresponding commit and MR
+ div_term = torch.pow(10000.0, -torch.arange(0, d_model, 2).float() / d_model)
+ # div_term = torch.exp(
+ # torch.arange(0, d_model, 2).float() * (-np.log(10000.0) / d_model)
+ # )
+
+ pe[:, 0::2] = torch.sin(position * div_term)
+ pe[:, 1::2] = torch.cos(position * div_term)
+ pe = pe.unsqueeze(0) # [1, T, D]
+
+ self.register_buffer("pe", pe, persistent=False)
+
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
+ """Apply positional encoding to input sequence.
+
+ Args:
+ x (torch.Tensor): [B, T, D] input motion sequence
+
+ Returns:
+ torch.Tensor: [B, T, D] input motion with PE added to it (and optionally dropout)
+ """
+ x = x + self.pe[:, : x.shape[1], :]
+ return self.dropout(x)
+
+
+class TimestepEmbedder(nn.Module):
+ """Encoder for diffusion step."""
+
+ def __init__(self, latent_dim: int, sequence_pos_encoder: PositionalEncoding):
+ """
+ Args:
+ latent_dim (int): dim to encode to
+ sequence_pos_encoder (PositionalEncoding): the PE to use on timesteps
+ """
+ super().__init__()
+ self.latent_dim = latent_dim
+ self.sequence_pos_encoder = sequence_pos_encoder
+
+ time_embed_dim = self.latent_dim
+ self.time_embed = nn.Sequential(
+ nn.Linear(self.latent_dim, time_embed_dim),
+ nn.SiLU(),
+ nn.Linear(time_embed_dim, time_embed_dim),
+ )
+
+ def forward(self, timesteps: torch.Tensor) -> torch.Tensor:
+ """Embed timesteps by adding PE then going through linear layers.
+
+ Args:
+ timesteps (torch.Tensor): [B]
+
+ Returns:
+ torch.Tensor: [B, 1, D]
+ """
+ return self.time_embed(self.sequence_pos_encoder.pe.transpose(0, 1)[timesteps])
diff --git a/kimodo/model/cfg.py b/kimodo/model/cfg.py
new file mode 100644
index 0000000000000000000000000000000000000000..6c39defdbbf0c074a0adb746ec4d98266e5ee463
--- /dev/null
+++ b/kimodo/model/cfg.py
@@ -0,0 +1,133 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Classifier-free guidance wrapper for the denoiser at sampling time."""
+
+from typing import Dict, Optional, Tuple, Union
+
+import torch
+import torch.nn as nn
+
+CFG_TYPES = ["nocfg", "regular", "separated"]
+
+
+class ClassifierFreeGuidedModel(nn.Module):
+ """Wrapper around denoiser to use classifier-free guidance at sampling time."""
+
+ def __init__(self, model: nn.Module, cfg_type: Optional[str] = "separated"):
+ """Wrap the denoiser for classifier-free guidance; cfg_type in CFG_TYPES (e.g. 'regular',
+ 'nocfg')."""
+ super().__init__()
+ self.model = model
+ assert cfg_type in CFG_TYPES, f"Invalid cfg_type: {cfg_type}"
+ self.cfg_type_default = cfg_type
+
+ def forward(
+ self,
+ cfg_weight: Union[float, Tuple[float, float]],
+ x: torch.Tensor,
+ x_pad_mask: torch.Tensor,
+ text_feat: torch.Tensor,
+ text_feat_pad_mask: torch.Tensor,
+ timesteps: torch.Tensor,
+ first_heading_angle: Optional[torch.Tensor] = None,
+ motion_mask: Optional[torch.Tensor] = None,
+ observed_motion: Optional[torch.Tensor] = None,
+ cfg_type: Optional[str] = None,
+ ) -> torch.Tensor:
+ """
+ Args:
+ cfg_weight (float): guidance weight float or tuple of floats with (text, constraint) weights if using separated cfg
+ x (torch.Tensor): [B, T, dim_motion] current noisy motion
+ x_pad_mask (torch.Tensor): [B, T] attention mask, positions with True are allowed to attend, False are not
+ text_feat (torch.Tensor): [B, max_text_len, llm_dim] embedded text prompts
+ text_feat_pad_mask (torch.Tensor): [B, max_text_len] attention mask, positions with True are allowed to attend, False are not
+ timesteps (torch.Tensor): [B,] current denoising step
+ motion_mask
+ observed_motion
+ neutral_joints (torch.Tensor): [B, nbjoints] The neutral joints of the motions
+
+ Returns:
+ torch.Tensor: same size as input x
+ """
+
+ if cfg_type is None:
+ cfg_type = self.cfg_type_default
+
+ assert cfg_type in CFG_TYPES, f"Invalid cfg_type: {cfg_type}"
+
+ # batched conditional and uncond pass together
+ if cfg_type == "nocfg":
+ return self.model(
+ x,
+ x_pad_mask,
+ text_feat,
+ text_feat_pad_mask,
+ timesteps,
+ first_heading_angle=first_heading_angle,
+ motion_mask=motion_mask,
+ observed_motion=observed_motion,
+ )
+ elif cfg_type == "regular":
+ assert isinstance(cfg_weight, (float, int)), "cfg_weight must be a single float for regular CFG"
+ # out_uncond + w * (out_text_and_constraint - out_uncond)
+ text_feat = torch.concatenate([text_feat, 0 * text_feat], dim=0)
+ if motion_mask is not None:
+ motion_mask = torch.concatenate([motion_mask, 0 * motion_mask], dim=0)
+ if observed_motion is not None:
+ observed_motion = torch.concatenate([observed_motion, observed_motion], dim=0)
+ if first_heading_angle is not None:
+ first_heading_angle = torch.concatenate([first_heading_angle, first_heading_angle], dim=0)
+
+ out_cond_uncond = self.model(
+ torch.concatenate([x, x], dim=0),
+ torch.concatenate([x_pad_mask, x_pad_mask], dim=0),
+ text_feat,
+ torch.concatenate([text_feat_pad_mask, False * text_feat_pad_mask], dim=0),
+ torch.concatenate([timesteps, timesteps], dim=0),
+ first_heading_angle=first_heading_angle,
+ motion_mask=motion_mask,
+ observed_motion=observed_motion,
+ )
+
+ out, out_uncond = torch.chunk(out_cond_uncond, 2)
+ out_new = out_uncond + (cfg_weight * (out - out_uncond))
+ elif cfg_type == "separated":
+ assert len(cfg_weight) == 2, "cfg_weight must be a tuple of two floats for separated CFG"
+ # out_uncond + w_text * (out_text - out_uncond) + w_constraint * (out_constraint - out_uncond)
+ text_feat = torch.concatenate([text_feat, 0 * text_feat, 0 * text_feat], dim=0)
+ if motion_mask is not None:
+ motion_mask = torch.concatenate([0 * motion_mask, motion_mask, 0 * motion_mask], dim=0)
+ if observed_motion is not None:
+ observed_motion = torch.concatenate([observed_motion, observed_motion, observed_motion], dim=0)
+ if first_heading_angle is not None:
+ first_heading_angle = torch.concatenate(
+ [first_heading_angle, first_heading_angle, first_heading_angle],
+ dim=0,
+ )
+
+ out_cond_uncond = self.model(
+ torch.concatenate([x, x, x], dim=0),
+ torch.concatenate([x_pad_mask, x_pad_mask, x_pad_mask], dim=0),
+ text_feat,
+ torch.concatenate(
+ [
+ text_feat_pad_mask,
+ False * text_feat_pad_mask,
+ False * text_feat_pad_mask,
+ ],
+ dim=0,
+ ),
+ torch.concatenate([timesteps, timesteps, timesteps], dim=0),
+ first_heading_angle=first_heading_angle,
+ motion_mask=motion_mask,
+ observed_motion=observed_motion,
+ )
+
+ out_text, out_constraint, out_uncond = torch.chunk(out_cond_uncond, 3)
+ out_new = (
+ out_uncond + (cfg_weight[0] * (out_text - out_uncond)) + (cfg_weight[1] * (out_constraint - out_uncond))
+ )
+ else:
+ raise ValueError(f"Invalid cfg_type: {cfg_type}")
+
+ return out_new
diff --git a/kimodo/model/common.py b/kimodo/model/common.py
new file mode 100644
index 0000000000000000000000000000000000000000..3b6937bb98bc676d380ef4657cfd58f42bc5f294
--- /dev/null
+++ b/kimodo/model/common.py
@@ -0,0 +1,48 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Config hydration: env vars, _target_ resolution, and recursive instantiation."""
+
+import importlib
+import os
+
+
+def get_env_var(name: str, default=None):
+ """Read env var by name and by lowercased name; return default if neither set."""
+ return os.getenv(name, os.getenv(name.lower(), default))
+
+
+def resolve_target(target: str):
+ """Import module and return the attribute named by a dotted path (e.g. 'pkg.mod.Class')."""
+ module_name, attr_name = target.rsplit(".", 1)
+ module = importlib.import_module(module_name)
+ return getattr(module, attr_name)
+
+
+def materialize_value(value):
+ """Recursively turn dicts with '_target_' into instances; lists/dicts traversed; leaves
+ unchanged."""
+ if isinstance(value, dict):
+ if "_target_" in value:
+ return instantiate_from_dict(value)
+ return {k: materialize_value(v) for k, v in value.items()}
+ if isinstance(value, list):
+ return [materialize_value(v) for v in value]
+ return value
+
+
+def instantiate_from_dict(node, overrides=None):
+ """Build an instance from a config dict: '_target_' gives the class, other keys are kwargs; overrides merged in."""
+ if not isinstance(node, dict) or "_target_" not in node:
+ raise ValueError("Config node must be a dict with a '_target_' key.")
+
+ target = resolve_target(node["_target_"])
+ kwargs = {}
+ for key, value in node.items():
+ if key == "_target_":
+ continue
+ kwargs[key] = materialize_value(value)
+
+ if overrides:
+ kwargs.update({k: v for k, v in overrides.items() if v is not None})
+
+ return target(**kwargs)
diff --git a/kimodo/model/diffusion.py b/kimodo/model/diffusion.py
new file mode 100644
index 0000000000000000000000000000000000000000..7e36d9940142b01d9447ceac2f9425d87589af4d
--- /dev/null
+++ b/kimodo/model/diffusion.py
@@ -0,0 +1,133 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Diffusion process and DDIM sampling for motion generation."""
+
+import math
+from typing import Optional, Tuple
+
+import torch
+from torch import nn
+
+
+def get_beta_schedule(
+ num_diffusion_timesteps: int,
+ max_beta: Optional[float] = 0.999,
+) -> torch.Tensor:
+ """Get cosine beta schedule."""
+
+ def alpha_bar(t):
+ return math.cos((t + 0.008) / 1.008 * math.pi / 2) ** 2
+
+ betas = []
+ for i in range(num_diffusion_timesteps):
+ t1 = i / num_diffusion_timesteps
+ t2 = (i + 1) / num_diffusion_timesteps
+ betas.append(min(1 - alpha_bar(t2) / alpha_bar(t1), max_beta))
+ return torch.tensor(betas, dtype=torch.float)
+
+
+class Diffusion(torch.nn.Module):
+ """Cosine-schedule diffusion process: betas, alphas, and DDIM step mapping."""
+
+ def __init__(self, num_base_steps: int):
+ """Set up cosine beta schedule and precompute diffusion variables for num_base_steps."""
+ super().__init__()
+ self.num_base_steps = num_base_steps
+ betas_base = get_beta_schedule(self.num_base_steps)
+ self.register_buffer("betas_base", betas_base, persistent=False)
+ alphas_cumprod_base = torch.cumprod(1.0 - self.betas_base, dim=0)
+ self.register_buffer("alphas_cumprod_base", alphas_cumprod_base, persistent=False)
+ use_timesteps, _ = self.space_timesteps(self.num_base_steps)
+ self.calc_diffusion_vars(use_timesteps)
+
+ def extra_repr(self) -> str:
+ return f"num_base_steps={self.num_base_steps}"
+
+ @property
+ def device(self):
+ return self.betas_base.device
+
+ def space_timesteps(self, num_denoising_steps: int) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Return (use_timesteps, map_tensor) for a subsampled denoising schedule of
+ num_denoising_steps."""
+ nsteps_train = self.num_base_steps
+ frac_stride = (nsteps_train - 1) / max(1, num_denoising_steps - 1)
+ use_timesteps = torch.round(torch.arange(nsteps_train, device=self.device) * frac_stride).to(torch.long)
+ use_timesteps = torch.clamp(use_timesteps, max=nsteps_train - 1)
+ map_tensor = torch.arange(nsteps_train, device=self.device, dtype=torch.long)[use_timesteps]
+ return use_timesteps, map_tensor
+
+ def calc_diffusion_vars(self, use_timesteps: torch.Tensor) -> None:
+ """Update buffers (betas, alphas, alphas_cumprod, etc.) for the given subsampled
+ timesteps."""
+ alphas_cumprod = self.alphas_cumprod_base[use_timesteps]
+ last_alpha_cumprod = torch.cat([torch.tensor([1.0]).to(alphas_cumprod), alphas_cumprod[:-1]])
+ betas = 1.0 - alphas_cumprod / last_alpha_cumprod
+ self.register_buffer("betas", betas, persistent=False)
+
+ alphas = 1.0 - self.betas
+ self.register_buffer("alphas", alphas, persistent=False)
+ alphas_cumprod = torch.cumprod(self.alphas, dim=0)
+ alphas_cumprod = torch.clamp(alphas_cumprod, min=1e-9)
+ self.register_buffer("alphas_cumprod", alphas_cumprod, persistent=False)
+
+ alphas_cumprod_prev = torch.cat([torch.tensor([1.0]).to(self.alphas_cumprod), self.alphas_cumprod[:-1]])
+ self.register_buffer("alphas_cumprod_prev", alphas_cumprod_prev, persistent=False)
+
+ sqrt_recip_alphas_cumprod = torch.rsqrt(self.alphas_cumprod)
+ self.register_buffer("sqrt_recip_alphas_cumprod", sqrt_recip_alphas_cumprod, persistent=False)
+
+ sqrt_recipm1_alphas_cumprod = torch.rsqrt(self.alphas_cumprod / (1.0 - self.alphas_cumprod))
+ self.register_buffer("sqrt_recipm1_alphas_cumprod", sqrt_recipm1_alphas_cumprod, persistent=False)
+
+ posterior_variance = self.betas * (1.0 - self.alphas_cumprod_prev) / (1.0 - self.alphas_cumprod)
+ self.register_buffer("posterior_variance", posterior_variance, persistent=False)
+
+ sqrt_alphas_cumprod = torch.rsqrt(1.0 / self.alphas_cumprod)
+ self.register_buffer("sqrt_alphas_cumprod", sqrt_alphas_cumprod, persistent=False)
+
+ sqrt_one_minus_alphas_cumprod = torch.rsqrt(1.0 / (1.0 - self.alphas_cumprod))
+ self.register_buffer(
+ "sqrt_one_minus_alphas_cumprod",
+ sqrt_one_minus_alphas_cumprod,
+ persistent=False,
+ )
+
+ def q_sample(
+ self,
+ x_start: torch.Tensor,
+ t: torch.Tensor,
+ noise: torch.Tensor = None,
+ ):
+ if noise is None:
+ noise = torch.randn_like(x_start)
+ assert noise.shape == x_start.shape
+
+ xt = (
+ self.sqrt_alphas_cumprod[t, None, None] * x_start
+ + self.sqrt_one_minus_alphas_cumprod[t, None, None] * noise
+ )
+ return xt
+
+
+class DDIMSampler(nn.Module):
+ """Deterministic DDIM sampler (eta = 0)."""
+
+ def __init__(self, diffusion: Diffusion):
+ super().__init__()
+ self.diffusion = diffusion
+
+ def __call__(
+ self,
+ use_timesteps: torch.Tensor,
+ x_t: torch.Tensor,
+ pred_xstart: torch.Tensor,
+ t: torch.Tensor,
+ ) -> torch.Tensor:
+ self.diffusion.calc_diffusion_vars(use_timesteps)
+ eps = (
+ self.diffusion.sqrt_recip_alphas_cumprod[t, None, None] * x_t - pred_xstart
+ ) / self.diffusion.sqrt_recipm1_alphas_cumprod[t, None, None]
+ alpha_bar_prev = self.diffusion.alphas_cumprod_prev[t, None, None]
+ x = pred_xstart * torch.sqrt(alpha_bar_prev) + torch.sqrt(1 - alpha_bar_prev) * eps
+ return x
diff --git a/kimodo/model/kimodo_model.py b/kimodo/model/kimodo_model.py
new file mode 100644
index 0000000000000000000000000000000000000000..7f4f2c3aad3fa81ea4b44525d2c788850d07c233
--- /dev/null
+++ b/kimodo/model/kimodo_model.py
@@ -0,0 +1,634 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Kimodo model: denoiser, text encoder, diffusion sampling, and post-processing."""
+
+import logging
+from typing import Dict, List, Optional, Tuple, Union
+
+import torch
+from torch import nn
+from tqdm.auto import tqdm
+
+from kimodo.constraints import EndEffectorConstraintSet, FullBodyConstraintSet
+from kimodo.motion_rep.feature_utils import compute_heading_angle, length_to_mask
+from kimodo.postprocess import post_process_motion
+from kimodo.sanitize import sanitize_texts
+from kimodo.skeleton import SOMASkeleton30
+from kimodo.tools import to_numpy
+
+from .cfg import ClassifierFreeGuidedModel
+from .diffusion import DDIMSampler, Diffusion
+
+log = logging.getLogger(__name__)
+
+
+class Kimodo(nn.Module):
+ """Helper class for test time."""
+
+ def __init__(
+ self,
+ denoiser: nn.Module,
+ text_encoder: nn.Module,
+ num_base_steps: int,
+ device: Optional[Union[str, torch.device]] = None,
+ cfg_type: Optional[str] = "separated",
+ ):
+ super().__init__()
+
+ self.denoiser = denoiser.eval()
+
+ if cfg_type is None:
+ cfg_type = "nocfg"
+
+ # Add Classifier-free guidance to the model if needed
+ self.denoiser = ClassifierFreeGuidedModel(self.denoiser, cfg_type=cfg_type)
+
+ self.motion_rep = denoiser.motion_rep
+ self.skeleton = self.motion_rep.skeleton
+
+ self.fps = denoiser.motion_rep.fps
+
+ self.diffusion = Diffusion(num_base_steps=num_base_steps)
+ self.sampler = DDIMSampler(self.diffusion)
+ self.text_encoder = text_encoder
+
+ self.device = device
+ # for classifier-free guidance
+
+ self.to(device)
+
+ @property
+ def output_skeleton(self):
+ """Skeleton used for model output (somaskel77 for SOMA, else unchanged)."""
+ if isinstance(self.skeleton, SOMASkeleton30):
+ return self.skeleton.somaskel77
+ return self.skeleton
+
+ def train(self, mode: bool):
+ self.denoiser.train(mode)
+ return self
+
+ def eval(self):
+ self.denoiser.eval()
+ return self
+
+ def denoising_step(
+ self,
+ motion: torch.Tensor,
+ pad_mask: torch.Tensor,
+ text_feat: torch.Tensor,
+ text_pad_mask: torch.Tensor,
+ t: torch.Tensor,
+ first_heading_angle: Optional[torch.Tensor],
+ motion_mask: torch.Tensor,
+ observed_motion: torch.Tensor,
+ num_denoising_steps: torch.Tensor,
+ cfg_weight: Union[float, Tuple[float, float]],
+ guide_masks: Optional[Dict] = None,
+ cfg_type: Optional[str] = None,
+ ) -> torch.Tensor:
+ """Single denoising step.
+
+ Returns:
+ torch.Tensor: [B, T, D] noisy motion input to t-1
+ """
+ # subsample timesteps
+ # NOTE: do this at every step due to ONNX export, i.e. num_samp_stepsmay change dynamically when
+ # running onnx version so need to account for that.
+ num_denoising_steps = num_denoising_steps[0]
+ use_timesteps, map_tensor = self.diffusion.space_timesteps(num_denoising_steps)
+ self.diffusion.calc_diffusion_vars(use_timesteps)
+
+ # first compute initial clean prediction from denoiser
+ t_map = map_tensor[t]
+
+ with torch.inference_mode():
+ pred_clean = self.denoiser(
+ cfg_weight,
+ motion,
+ pad_mask,
+ text_feat,
+ text_pad_mask,
+ t_map,
+ first_heading_angle,
+ motion_mask,
+ observed_motion,
+ cfg_type=cfg_type,
+ )
+
+ # sampler computes next step noisy motion
+ x_tm1 = self.sampler(use_timesteps, motion, pred_clean, t)
+ return x_tm1
+
+ def _multiprompt(
+ self,
+ prompts: list[str],
+ num_frames: int | list[int],
+ num_denoising_steps: int,
+ constraint_lst: Optional[list] = [],
+ cfg_weight: Optional[float] = [2.0, 2.0],
+ num_samples: Optional[int] = None,
+ cfg_type: Optional[str] = None,
+ return_numpy: bool = False,
+ first_heading_angle: Optional[torch.Tensor] = None,
+ # for transitioning
+ num_transition_frames: int = 5,
+ # for postprocess
+ post_processing: bool = False,
+ root_margin: float = 0.04,
+ # progress bar
+ progress_bar=tqdm,
+ ) -> torch.Tensor:
+ device = self.device
+
+ bs = num_samples
+ texts = sanitize_texts(prompts)
+
+ if isinstance(num_frames, int):
+ # same duration for all the segments
+ num_frames = [num_frames for _ in range(num_samples)]
+
+ tosqueeze = False
+ if num_samples is None:
+ num_samples = 1
+ tosqueeze = True
+
+ if constraint_lst is None:
+ constraint_lst = []
+
+ # Generate one chunck at a time
+ current_frame = 0
+ generated_motions = []
+
+ for idx, (text, num_frame) in enumerate(zip(texts, num_frames)):
+ texts_bs = [text for _ in range(num_samples)]
+
+ lengths = torch.tensor(
+ [num_frame for _ in range(num_samples)],
+ device=device,
+ )
+
+ is_first_motion = not generated_motions
+
+ observed_motion, motion_mask = None, None
+
+ # filter the constraint_lst to only keep the relevent ones
+ constraint_lst_base = [
+ constraint.crop_move(current_frame, current_frame + num_frame) for constraint in constraint_lst
+ ] # this move temporally but not spatially
+
+ observed_motion, motion_mask = self.motion_rep.create_conditions_from_constraints_batched(
+ constraint_lst_base,
+ lengths,
+ to_normalize=False, # don't normalize yet, it needs to be moved around
+ device=device,
+ )
+
+ if not is_first_motion:
+ nb_transition_frames = num_transition_frames
+
+ if nb_transition_frames < 1:
+ raise ValueError(f"num_transition_frames must be at least 1, got {nb_transition_frames}")
+
+ latest_motions = generated_motions.pop()
+ # remove the transition part of A (will be put back afterward)
+ generated_motions.append(latest_motions[:, :-nb_transition_frames])
+ latest_frames = latest_motions[:, -nb_transition_frames:]
+
+ last_output = self.motion_rep.inverse(
+ latest_frames,
+ is_normalized=False,
+ return_numpy=False,
+ )
+ smooth_root_2d = last_output["smooth_root_pos"][..., [0, 2]]
+
+ # add constraints at the begining to allow natural transitions
+ constraint_lst_transition = []
+ for batch_id in range(bs):
+ new_constraint = FullBodyConstraintSet(
+ self.skeleton,
+ torch.arange(num_transition_frames),
+ last_output["posed_joints"][batch_id, :num_transition_frames],
+ last_output["global_rot_mats"][batch_id, :num_transition_frames],
+ smooth_root_2d[batch_id, :num_transition_frames],
+ )
+ # separate end-effector constraint to capture hand/feet rotations
+ new_ee_constraint = EndEffectorConstraintSet(
+ self.skeleton,
+ torch.arange(num_transition_frames),
+ last_output["posed_joints"][batch_id, :num_transition_frames],
+ last_output["global_rot_mats"][batch_id, :num_transition_frames],
+ smooth_root_2d[batch_id, :num_transition_frames],
+ joint_names=["LeftHand", "RightHand", "LeftFoot", "RightFoot"],
+ )
+
+ constraint_lst_transition.append([new_constraint, new_ee_constraint])
+
+ transition_lengths = torch.tensor(
+ [nb_transition_frames for _ in range(num_samples)],
+ device=device,
+ )
+
+ observed_motion_transition, motion_mask_transition = (
+ self.motion_rep.create_conditions_from_constraints_batched(
+ constraint_lst_transition,
+ transition_lengths,
+ to_normalize=False, # don't normalize yet
+ device=device,
+ )
+ )
+
+ # concatenate the obversed motion / motion mask
+ observed_motion = torch.cat([observed_motion_transition, observed_motion], axis=1)
+ motion_mask = torch.cat([motion_mask_transition, motion_mask], axis=1)
+
+ # we need to move each observed motion in the batch to the new starting points
+ last_smooth_root_2d = smooth_root_2d[:, 0]
+ observed_motion = self.motion_rep.translate_2d(
+ observed_motion, -last_smooth_root_2d
+ ) # equivalent to: self.motion_rep.translate_2d_to_zero(observed_motion)
+
+ # remove dummy values after moving
+ observed_motion = observed_motion * motion_mask
+
+ lengths = lengths + transition_lengths
+ first_heading_angle = compute_heading_angle(last_output["posed_joints"], self.skeleton)[:, 0]
+ else:
+ if first_heading_angle is None:
+ # Start at 0 angle, but this will change afterward
+ first_heading_angle = torch.tensor([0.0] * bs, device=device)
+ else:
+ first_heading_angle = torch.as_tensor(first_heading_angle, device=device)
+ if first_heading_angle.numel() == 1:
+ first_heading_angle = first_heading_angle.repeat(bs)
+
+ observed_motion = self.motion_rep.normalize(observed_motion)
+
+ max_frames = max(lengths)
+ motion_pad_mask = length_to_mask(lengths)
+
+ motion = self._generate(
+ texts_bs,
+ max_frames,
+ num_denoising_steps=num_denoising_steps,
+ pad_mask=motion_pad_mask,
+ first_heading_angle=first_heading_angle,
+ motion_mask=motion_mask,
+ observed_motion=observed_motion,
+ cfg_weight=cfg_weight,
+ cfg_type=cfg_type,
+ )
+
+ motion = self.motion_rep.unnormalize(motion)
+
+ if not is_first_motion:
+ motion_with_transition = self.motion_rep.translate_2d(
+ motion,
+ last_smooth_root_2d,
+ )
+
+ if post_processing:
+ # Per-segment postprocessing: inverse, postprocess, re-encode.
+ # The full transition+segment is postprocessed together so the
+ # transition constraints keep the junction smooth.
+ seg_output = self.motion_rep.inverse(
+ motion_with_transition, is_normalized=False, return_numpy=False,
+ )
+ seg_constraints = [list(cl) for cl in constraint_lst_transition]
+ for bi in range(bs):
+ seg_constraints[bi].extend(
+ [c.crop_move(current_frame - nb_transition_frames,
+ current_frame - nb_transition_frames + num_frame + nb_transition_frames)
+ for c in constraint_lst]
+ )
+ corrected = post_process_motion(
+ seg_output["local_rot_mats"],
+ seg_output["root_positions"],
+ seg_output["foot_contacts"],
+ self.skeleton,
+ seg_constraints,
+ root_margin=root_margin,
+ )
+ seg_output.update(corrected)
+ motion = self.motion_rep(
+ seg_output["local_rot_mats"],
+ seg_output["root_positions"],
+ to_normalize=False,
+ lengths=lengths,
+ )
+ else:
+ motion = motion_with_transition[:, num_transition_frames:]
+ transition_frames = motion_with_transition[:, :num_transition_frames]
+
+ # linearly combine the previously generated transitions with the newly generated ones
+ alpha = torch.linspace(1, 0, num_transition_frames, device=device)[:, None]
+ new_transition_frames = (
+ latest_frames[:, :num_transition_frames] * alpha + (1 - alpha) * transition_frames
+ )
+
+ # add new transitions frames for A (merging with B prediction of the history)
+ generated_motions.append(new_transition_frames)
+
+ elif post_processing:
+ # First segment: postprocess immediately
+ seg_output = self.motion_rep.inverse(
+ motion, is_normalized=False, return_numpy=False,
+ )
+ seg_constraints = constraint_lst_base if constraint_lst_base else []
+ corrected = post_process_motion(
+ seg_output["local_rot_mats"],
+ seg_output["root_positions"],
+ seg_output["foot_contacts"],
+ self.skeleton,
+ seg_constraints,
+ root_margin=root_margin,
+ )
+ seg_output.update(corrected)
+ motion = self.motion_rep(
+ seg_output["local_rot_mats"],
+ seg_output["root_positions"],
+ to_normalize=False,
+ lengths=lengths,
+ )
+
+ generated_motions.append(motion)
+ current_frame += num_frame
+
+ generated_motions = torch.cat(generated_motions, axis=1) # temporal axis (b, t, d)
+
+ if tosqueeze:
+ generated_motions = generated_motions[0]
+
+ output = self.motion_rep.inverse(
+ generated_motions,
+ is_normalized=False,
+ return_numpy=False,
+ )
+
+ # Post-processing: already applied per-segment inside the loop above,
+ # so no additional post-processing pass is needed here.
+
+ # Convert SOMA output to somaskel77 for external API
+ if isinstance(self.skeleton, SOMASkeleton30):
+ output = self.skeleton.output_to_SOMASkeleton77(output)
+
+ # Convert to numpy if requested
+ if return_numpy:
+ output = to_numpy(output)
+ return output
+
+ def __call__(
+ self,
+ prompts: str | list[str],
+ num_frames: int | list[int],
+ num_denoising_steps: int,
+ multi_prompt: bool = False,
+ constraint_lst: Optional[list] = [],
+ cfg_weight: Optional[float] = [2.0, 2.0],
+ num_samples: Optional[int] = None,
+ cfg_type: Optional[str] = None,
+ return_numpy: bool = False,
+ first_heading_angle: Optional[torch.Tensor] = None,
+ # for transitioning
+ num_transition_frames: int = 5,
+ # for postprocess
+ post_processing: bool = False,
+ root_margin: float = 0.04,
+ # progress bar
+ progress_bar=tqdm,
+ ) -> dict:
+ """Generate motion from text prompts and optional kinematic constraints.
+
+ When a single prompt/num_frames pair is given, one motion is generated.
+ Passing lists of prompts and/or num_frames produces a batch of
+ independent motions. With ``multi_prompt=True``, the prompts are
+ treated as sequential segments that are generated and stitched together
+ with smooth transitions.
+
+ Args:
+ prompts: One or more text descriptions of the desired motion.
+ A single string generates one sample; a list generates a batch
+ (or sequential segments when ``multi_prompt=True``).
+ num_frames: Duration of the generated motion in frames. Can be a
+ single int applied to every prompt or a per-prompt list.
+ num_denoising_steps: Number of DDIM denoising steps. More steps
+ generally improve quality at the cost of speed.
+ multi_prompt: If ``True``, treat ``prompts`` as an ordered sequence
+ of segments and concatenate them with transitions.
+ constraint_lst: Per-sample list of kinematic constraints (e.g.
+ keyframe poses, end-effector targets, 2-D paths). Pass an
+ empty list for unconstrained generation.
+ cfg_weight: Classifier-free guidance scale(s). A two-element list
+ ``[text_cfg, constraint_cfg]`` controls text and constraint
+ guidance independently.
+ num_samples: Number of samples to generate.
+ cfg_type: Override the default CFG strategy set at init
+ (e.g. ``"separated"``).
+ return_numpy: If ``True``, convert all output tensors to numpy
+ arrays.
+ first_heading_angle: Initial body heading in radians. Shape
+ ``(B,)`` or scalar. Defaults to ``0`` (facing +Z).
+ num_transition_frames: Number of overlapping frames used to blend
+ consecutive segments in multi-prompt mode.
+ post_processing: If ``True``, apply post-processing
+ (foot-skate cleanup and constraint enforcement).
+ root_margin: Horizontal margin (in meters) used by the post-processor
+ to determine when to correct root motion. When root deviates more than
+ margin from the constraint, the post-processor will correct it.
+ progress_bar: Callable wrapping an iterable to display progress
+ (default: ``tqdm``). Pass a no-op to silence output.
+
+ Returns:
+ dict: A dictionary of motion tensors (or numpy arrays if
+ ``return_numpy=True``) with the following keys:
+
+ - ``local_rot_mats`` – Local joint rotations as rotation matrices.
+ - ``global_rot_mats`` – Global joint rotations as rotation matrices.
+ - ``posed_joints`` – Joint positions in world space.
+ - ``root_positions`` – Root joint positions.
+ - ``smooth_root_pos`` – Smoothed root trajectory.
+ - ``foot_contacts`` – Boolean foot-contact labels [left heel, left toe, right heel, right toe].
+ - ``global_root_heading`` – Root heading angle over time.
+ """
+ device = self.device
+
+ if multi_prompt:
+ # multi prompt generation
+ return self._multiprompt(
+ prompts,
+ num_frames,
+ num_denoising_steps,
+ constraint_lst,
+ cfg_weight,
+ num_samples,
+ cfg_type,
+ return_numpy,
+ first_heading_angle,
+ num_transition_frames,
+ post_processing,
+ root_margin,
+ progress_bar,
+ )
+
+ # Input checking
+ tosqueeze = False
+ if isinstance(prompts, list) and isinstance(num_frames, list):
+ assert len(prompts) == len(num_frames), "The number of prompts should match the number of num_frames."
+ num_samples = len(prompts)
+ elif isinstance(prompts, list):
+ num_samples = len(prompts)
+ num_frames = [num_frames for _ in range(num_samples)]
+ elif isinstance(num_frames, list):
+ num_samples = len(num_frames)
+ prompts = [prompts for _ in range(num_samples)]
+ else:
+ if num_samples is None:
+ tosqueeze = True
+ num_samples = 1
+ prompts = [prompts for _ in range(num_samples)]
+ num_frames = [num_frames for _ in range(num_samples)]
+
+ bs = num_samples
+ texts = sanitize_texts(prompts)
+
+ lengths = torch.tensor(
+ num_frames,
+ device=device,
+ )
+ max_frames = max(lengths)
+ motion_pad_mask = length_to_mask(lengths)
+
+ if first_heading_angle is None:
+ # Start at 0 angle
+ first_heading_angle = torch.tensor([0.0] * bs, device=device)
+ else:
+ first_heading_angle = torch.as_tensor(first_heading_angle, device=device)
+ if first_heading_angle.numel() == 1:
+ first_heading_angle = first_heading_angle.repeat(bs)
+
+ observed_motion, motion_mask = None, None
+ if constraint_lst:
+ observed_motion, motion_mask = self.motion_rep.create_conditions_from_constraints_batched(
+ constraint_lst,
+ lengths,
+ to_normalize=True,
+ device=device,
+ )
+
+ motion = self._generate(
+ texts,
+ max_frames,
+ num_denoising_steps=num_denoising_steps,
+ pad_mask=motion_pad_mask,
+ first_heading_angle=first_heading_angle,
+ motion_mask=motion_mask,
+ observed_motion=observed_motion,
+ cfg_weight=cfg_weight,
+ cfg_type=cfg_type,
+ progress_bar=progress_bar,
+ )
+
+ if tosqueeze:
+ motion = motion[0]
+
+ output = self.motion_rep.inverse(
+ motion,
+ is_normalized=True,
+ return_numpy=False, # Keep as tensor for potential post-processing
+ )
+
+ # Apply post-processing if requested
+ if post_processing:
+ corrected = post_process_motion(
+ output["local_rot_mats"],
+ output["root_positions"],
+ output["foot_contacts"],
+ self.skeleton,
+ constraint_lst,
+ root_margin=root_margin,
+ )
+ # key frame outputs / foot contacts are not changed
+ output.update(corrected)
+
+ # Convert SOMA output to somaskel77 for external API
+ if isinstance(self.skeleton, SOMASkeleton30):
+ output = self.skeleton.output_to_SOMASkeleton77(output)
+
+ # Convert to numpy if requested
+ if return_numpy:
+ output = to_numpy(output)
+ return output
+
+ def _generate(
+ self,
+ texts: List[str],
+ max_frames: int,
+ num_denoising_steps: int,
+ pad_mask: torch.Tensor,
+ first_heading_angle: Optional[torch.Tensor],
+ motion_mask: torch.Tensor,
+ observed_motion: torch.Tensor,
+ cfg_weight: Optional[float] = 2.0,
+ text_feat: Optional[torch.Tensor] = None,
+ text_pad_mask: Optional[torch.Tensor] = None,
+ guide_masks: Optional[Dict] = None,
+ cfg_type: Optional[str] = None,
+ progress_bar=tqdm,
+ ) -> torch.Tensor:
+ """Sample full denoising loop.
+
+ Args:
+ texts (List[str]): batch of text prompts to use for sampling (if text_feat is not passed in)
+ """
+
+ device = self.device
+ if text_feat is None:
+ assert text_pad_mask is None
+ log.info("Encoding text...")
+ text_feat, text_length = self.text_encoder(texts)
+ text_feat = text_feat.to(device)
+
+ # handle empty string (set to zero)
+ empty_text_mask = [len(text.strip()) == 0 for text in texts]
+ text_feat[empty_text_mask] = 0
+
+ # Create the pad mask for the text
+ batch_size, maxlen = text_feat.shape[:2]
+ tensor_text_length = torch.tensor(text_length, device=device)
+ tensor_text_length[empty_text_mask] = 0
+ text_pad_mask = torch.arange(maxlen, device=device).expand(batch_size, maxlen) < tensor_text_length[:, None]
+
+ if motion_mask is not None:
+ if motion_mask.dtype == torch.bool:
+ motion_mask = 1 * motion_mask
+
+ batch_size = text_feat.shape[0]
+
+ # sample loop
+ indices = list(range(num_denoising_steps))[::-1]
+ shape = (batch_size, max_frames, self.motion_rep.motion_rep_dim)
+ cur_mot = torch.randn(shape, device=self.device)
+ num_denoising_steps = torch.tensor(
+ [num_denoising_steps], device=self.device
+ ) # this and t need to be tensor for onnx export
+ # init diffusion with correct num steps before looping
+ use_timesteps = self.diffusion.space_timesteps(num_denoising_steps[0])[0]
+ self.diffusion.calc_diffusion_vars(use_timesteps)
+ for i in progress_bar(indices):
+ t = torch.tensor([i] * cur_mot.size(0), device=self.device)
+ with torch.inference_mode():
+ cur_mot = self.denoising_step(
+ cur_mot,
+ pad_mask,
+ text_feat,
+ text_pad_mask,
+ t,
+ first_heading_angle,
+ motion_mask,
+ observed_motion,
+ num_denoising_steps,
+ cfg_weight,
+ guide_masks=guide_masks,
+ cfg_type=cfg_type,
+ )
+ return cur_mot
diff --git a/kimodo/model/llm2vec/README.md b/kimodo/model/llm2vec/README.md
new file mode 100644
index 0000000000000000000000000000000000000000..450e995c0aa9cdb0f1117aa33cb30a30d84a41c0
--- /dev/null
+++ b/kimodo/model/llm2vec/README.md
@@ -0,0 +1 @@
+This is a patched version of the original [LLM2Vec](https://github.com/McGill-NLP/llm2vec) codebase so that `McGill-NLP/LLM2Vec-Meta-Llama-3-8B-Instruct-mntp-supervised` works with `transformers==5.0.0rc3`.
diff --git a/kimodo/model/llm2vec/__init__.py b/kimodo/model/llm2vec/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..5890f5848bc66237649197109d3a31328d3f77a1
--- /dev/null
+++ b/kimodo/model/llm2vec/__init__.py
@@ -0,0 +1,11 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""LLM2Vec text encoder and wrapper for Kimodo."""
+
+from .llm2vec import LLM2Vec
+from .llm2vec_wrapper import LLM2VecEncoder
+
+__all__ = [
+ "LLM2Vec",
+ "LLM2VecEncoder",
+]
diff --git a/kimodo/model/llm2vec/llm2vec.py b/kimodo/model/llm2vec/llm2vec.py
new file mode 100644
index 0000000000000000000000000000000000000000..6d01f5716ed57a6b6ea9b8cca3c9f292bdf19b5c
--- /dev/null
+++ b/kimodo/model/llm2vec/llm2vec.py
@@ -0,0 +1,477 @@
+# SPDX-FileCopyrightText: Copyright (c) 2024 McGill NLP
+# SPDX-License-Identifier: MIT
+#
+# Permission is hereby granted, free of charge, to any person obtaining a
+# copy of this software and associated documentation files (the "Software"),
+# to deal in the Software without restriction, including without limitation
+# the rights to use, copy, modify, merge, publish, distribute, sublicense,
+# and/or sell copies of the Software, and to permit persons to whom the
+# Software is furnished to do so, subject to the following conditions:
+#
+# The above copyright notice and this permission notice shall be included in
+# all copies or substantial portions of the Software.
+#
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL
+# THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING
+# FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER
+# DEALINGS IN THE SOFTWARE.
+
+
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import json
+import logging
+import os
+from functools import partial
+from typing import Dict, List, Optional, Union
+
+import numpy as np
+import torch
+import torch.multiprocessing as mp
+from peft import PeftModel
+from torch import Tensor, device, nn
+from tqdm.autonotebook import tqdm, trange
+from transformers import (
+ AutoConfig,
+ AutoModel,
+ AutoTokenizer,
+ GemmaConfig,
+ LlamaConfig,
+ MistralConfig,
+ PretrainedConfig,
+ Qwen2Config,
+)
+
+logger = logging.getLogger(__name__)
+
+
+def batch_to_device(batch, target_device: device):
+ """Send a pytorch batch to a device (CPU/GPU)"""
+ for key in batch:
+ if isinstance(batch[key], Tensor):
+ batch[key] = batch[key].to(target_device)
+ return batch
+
+
+class LLM2Vec(nn.Module):
+ def __init__(
+ self,
+ model: AutoModel,
+ tokenizer: AutoTokenizer,
+ pooling_mode: str = "mean",
+ max_length: int = 512,
+ doc_max_length: int = 400,
+ skip_instruction: bool = True,
+ ):
+ super().__init__()
+ self.model = model
+ self.tokenizer = tokenizer
+ self.pooling_mode = pooling_mode
+ self.skip_instruction = skip_instruction
+ self.max_length = max_length
+ self.doc_max_length = doc_max_length
+ self.config = model.config
+
+ @classmethod
+ def _get_model_class(cls, config_class_name, enable_bidirectional):
+ if not enable_bidirectional:
+ return AutoModel
+ if config_class_name == "MistralConfig":
+ from .models.bidirectional_mistral import MistralBiModel
+
+ return MistralBiModel
+ elif config_class_name == "LlamaConfig":
+ from .models.bidirectional_llama import LlamaBiModel
+
+ return LlamaBiModel
+ elif config_class_name == "GemmaConfig":
+ from .models.bidirectional_gemma import GemmaBiModel
+
+ return GemmaBiModel
+ elif config_class_name == "Qwen2Config":
+ from .models.bidirectional_qwen2 import Qwen2BiModel
+
+ return Qwen2BiModel
+ else:
+ raise ValueError(f"{config_class_name} is not supported yet with bidirectional models.")
+
+ @classmethod
+ def from_pretrained(
+ cls,
+ base_model_name_or_path,
+ peft_model_name_or_path=None,
+ merge_peft=False,
+ enable_bidirectional=True,
+ **kwargs,
+ ):
+ # pop out encoder args
+ keys = ["pooling_mode", "max_length", "doc_max_length", "skip_instruction"]
+ encoder_args = {key: kwargs.pop(key, None) for key in keys if kwargs.get(key) is not None}
+
+ tokenizer = AutoTokenizer.from_pretrained(base_model_name_or_path)
+ tokenizer.pad_token = tokenizer.eos_token
+ tokenizer.padding_side = "left"
+
+ config = AutoConfig.from_pretrained(base_model_name_or_path)
+ config_class_name = config.__class__.__name__
+
+ model_class = cls._get_model_class(config_class_name, enable_bidirectional=enable_bidirectional)
+
+ model = model_class.from_pretrained(base_model_name_or_path, **kwargs)
+
+ if os.path.isdir(base_model_name_or_path) and os.path.exists(f"{base_model_name_or_path}/config.json"):
+ with open(f"{base_model_name_or_path}/config.json", "r") as fIn:
+ config_dict = json.load(fIn)
+ config = PretrainedConfig.from_dict(config_dict)
+ model.config._name_or_path = config._name_or_path
+
+ # For special case where config.json and adapter weights are in the same directory
+ if hasattr(model, "peft_config"):
+ model = PeftModel.from_pretrained(
+ model,
+ base_model_name_or_path,
+ )
+ model = model.merge_and_unload()
+
+ if peft_model_name_or_path is not None:
+ model = PeftModel.from_pretrained(
+ model,
+ peft_model_name_or_path,
+ )
+ if merge_peft:
+ model = model.merge_and_unload()
+
+ config = {}
+ config_addr = peft_model_name_or_path if peft_model_name_or_path is not None else base_model_name_or_path
+ if os.path.exists(f"{config_addr}/llm2vec_config.json"):
+ with open(f"{config_addr}/llm2vec_config.json", "r") as fIn:
+ llm2vec_config = json.load(fIn)
+ config.update(llm2vec_config)
+
+ for key, value in encoder_args.items():
+ config[key] = value
+
+ return cls(model=model, tokenizer=tokenizer, **config)
+
+ def prepare_for_tokenization(self, text):
+ if self.model.config._name_or_path == "meta-llama/Meta-Llama-3-8B-Instruct":
+ text = "<|start_header_id|>user<|end_header_id|>\n\n" + text.strip() + "<|eot_id|>"
+ return text
+ if self.model.config._name_or_path in [
+ "mistralai/Mistral-7B-Instruct-v0.2",
+ "meta-llama/Llama-2-7b-chat-hf",
+ ]:
+ text = "[INST] " + text.strip() + " [/INST]"
+ if self.model.config._name_or_path in [
+ "google/gemma-2-9b-it",
+ ]:
+ text = "user\n" + text.strip() + ""
+ if self.model.config._name_or_path in [
+ "Qwen/Qwen2-1.5B-Instruct",
+ "Qwen/Qwen2-7B-Instruct",
+ ]:
+ text = "<|im_start|>user\n" + text.strip() + "<|im_end|>"
+ if self.pooling_mode == "eos_token":
+ if self.model.config._name_or_path == "meta-llama/Meta-Llama-3-8B":
+ text = text.strip() + "<|end_of_text|>"
+ elif isinstance(self.model.config, LlamaConfig) or isinstance(self.model.config, MistralConfig):
+ text = text.strip() + " "
+ elif isinstance(self.model.config, GemmaConfig):
+ text = text.strip() + ""
+ elif isinstance(self.model.config, Qwen2Config):
+ text = text.strip() + "<|endoftext|>"
+ return text
+
+ def tokenize(self, texts):
+ texts_2 = []
+ original_texts = []
+ for text in texts:
+ t = text.split("!@#$%^&*()")
+ texts_2.append(t[1] if len(t) > 1 else "")
+ original_texts.append("".join(t))
+
+ original = self.tokenizer(
+ original_texts,
+ return_tensors="pt",
+ padding=True,
+ truncation=True,
+ max_length=self.max_length,
+ )
+ embed_mask = None
+ for t_i, t in enumerate(texts_2):
+ ids = self.tokenizer(
+ [t],
+ return_tensors="pt",
+ padding=True,
+ truncation=True,
+ max_length=self.max_length,
+ add_special_tokens=False,
+ )
+ if embed_mask is None:
+ e_m = torch.zeros_like(original["attention_mask"][t_i])
+ if len(ids["input_ids"][0]) > 0:
+ e_m[-len(ids["input_ids"][0]) :] = torch.ones(len(ids["input_ids"][0]))
+ embed_mask = e_m.unsqueeze(0)
+ else:
+ e_m = torch.zeros_like(original["attention_mask"][t_i])
+ if len(ids["input_ids"][0]) > 0:
+ e_m[-len(ids["input_ids"][0]) :] = torch.ones(len(ids["input_ids"][0]))
+ embed_mask = torch.cat((embed_mask, e_m.unsqueeze(0)), dim=0)
+
+ original["embed_mask"] = embed_mask
+ return original
+
+ def _skip_instruction(self, sentence_feature):
+ assert sentence_feature["attention_mask"].shape == sentence_feature["embed_mask"].shape
+ sentence_feature["attention_mask"] = sentence_feature["embed_mask"]
+
+ def forward(self, sentence_feature: Dict[str, Tensor]):
+ embed_mask = None
+ if "embed_mask" in sentence_feature:
+ embed_mask = sentence_feature.pop("embed_mask")
+ reps = self.model(**sentence_feature)
+ sentence_feature["embed_mask"] = embed_mask
+
+ return self.get_pooling(sentence_feature, reps.last_hidden_state)
+
+ def get_pooling(self, features, last_hidden_states): # All models padded from left
+ assert self.tokenizer.padding_side == "left", "Pooling modes are implemented for padding from left."
+ if self.skip_instruction:
+ self._skip_instruction(features)
+ seq_lengths = features["attention_mask"].sum(dim=-1)
+ if self.pooling_mode == "mean":
+ return torch.stack(
+ [last_hidden_states[i, -length:, :].mean(dim=0) for i, length in enumerate(seq_lengths)],
+ dim=0,
+ )
+ elif self.pooling_mode == "weighted_mean":
+ bs, l, _ = last_hidden_states.shape
+ complete_weights = torch.zeros(bs, l, device=last_hidden_states.device)
+ for i, seq_l in enumerate(seq_lengths):
+ if seq_l > 0:
+ complete_weights[i, -seq_l:] = torch.arange(seq_l) + 1
+ complete_weights[i] /= torch.clamp(complete_weights[i].sum(), min=1e-9)
+ return torch.sum(last_hidden_states * complete_weights.unsqueeze(-1), dim=1)
+ elif self.pooling_mode == "eos_token" or self.pooling_mode == "last_token":
+ return last_hidden_states[:, -1]
+ elif self.pooling_mode == "bos_token":
+ return last_hidden_states[features["input_ids"] == self.tokenizer.bos_token_id]
+ else:
+ raise ValueError(f"{self.pooling_mode} is not implemented yet.")
+
+ def _convert_to_str(self, instruction, text):
+ tokenized_q = self.tokenizer(
+ text,
+ return_tensors="pt",
+ padding=True,
+ truncation=True,
+ max_length=self.max_length,
+ add_special_tokens=False,
+ )
+ tokenized_q_length = len(tokenized_q["input_ids"][0])
+
+ while tokenized_q_length > self.doc_max_length:
+ reduction_ratio = self.doc_max_length / tokenized_q_length
+ reduced_length = int(len(text.split()) * reduction_ratio)
+ text = " ".join(text.split()[:reduced_length])
+ tokenized_q = self.tokenizer(
+ text,
+ return_tensors="pt",
+ padding=True,
+ truncation=True,
+ max_length=self.max_length,
+ add_special_tokens=False,
+ )
+ tokenized_q_length = len(tokenized_q["input_ids"][0])
+
+ return f"{instruction.strip()} !@#$%^&*(){text}" if instruction else f"!@#$%^&*(){text}"
+
+ def encode(
+ self,
+ sentences: Union[str, List[str]],
+ batch_size: int = 32,
+ show_progress_bar: bool = True,
+ convert_to_numpy: bool = False,
+ convert_to_tensor: bool = False,
+ device: Optional[str] = None,
+ ):
+ """
+ Encode a list of sentences to their respective embeddings. The sentences can be a list of strings or a string.
+ Args:
+ sentences: sentence or sentences to encode.
+ batch_size: batch size for turning sentence tokens into embeddings.
+ show_progress_bar: whether to show progress bars during encoding steps.
+ convert_to_numpy: If true, return numpy arrays instead of torch tensors.
+ convert_to_tensor: If true, return torch tensors (default).
+ device: torch backend device identifier (e.g., 'cuda', 'cpu','mps' etc.). If not specified,
+ the default is to use cuda when available, otherwise cpu. Note that only the choice of 'cuda' supports
+ multiprocessing as currently implemented.
+
+ Returns: embeddings of the sentences. Embeddings are detached and always on the CPU (see _encode implementation).
+
+ """
+ if isinstance(sentences[0], str) and isinstance(sentences[-1], int):
+ sentences = [sentences]
+ # required for MEDI version of MTEB
+ if isinstance(sentences[0], str):
+ sentences = [[""] + [sentence] for sentence in sentences]
+
+ if device is None:
+ device = "cuda" if torch.cuda.is_available() else "cpu"
+
+ concatenated_input_texts = []
+ for sentence in sentences:
+ assert isinstance(sentence[0], str)
+ assert isinstance(sentence[1], str)
+ concatenated_input_texts.append(self._convert_to_str(sentence[0], sentence[1]))
+ sentences = concatenated_input_texts
+
+ self.eval()
+
+ if convert_to_tensor:
+ convert_to_numpy = False
+
+ length_sorted_idx = np.argsort([-self._text_length(sen) for sen in sentences])
+ sentences_sorted = [sentences[idx] for idx in length_sorted_idx]
+ all_embeddings = []
+
+ if torch.cuda.device_count() <= 1:
+ # This branch also support mps devices
+ self.to(device)
+ for start_index in trange(
+ 0,
+ len(sentences),
+ batch_size,
+ desc="Batches",
+ disable=not show_progress_bar,
+ ):
+ sentences_batch = sentences_sorted[start_index : start_index + batch_size]
+ embeddings = self._encode(sentences_batch, device=device, convert_to_numpy=convert_to_numpy)
+ all_embeddings.append(embeddings)
+ else:
+ num_proc = torch.cuda.device_count()
+ cuda_compatible_multiprocess = mp.get_context("spawn")
+ with cuda_compatible_multiprocess.Pool(num_proc) as p:
+ sentences_batches = [
+ sentences_sorted[start_index : start_index + batch_size]
+ for start_index in range(0, len(sentences), batch_size)
+ ]
+
+ progress_bar = tqdm(
+ total=len(sentences_batches),
+ desc="Batches",
+ disable=not show_progress_bar,
+ )
+ results = []
+
+ def update(*args):
+ progress_bar.update()
+
+ for batch in sentences_batches:
+ results.append(
+ p.apply_async(
+ self._encode,
+ args=(batch, None, convert_to_numpy, True),
+ callback=update,
+ )
+ )
+
+ all_embeddings = [result.get() for result in results]
+ progress_bar.close()
+
+ all_embeddings = torch.cat(all_embeddings, dim=0)
+ all_embeddings = all_embeddings[np.argsort(length_sorted_idx)]
+ all_embeddings = all_embeddings.to(torch.float32)
+ if convert_to_numpy:
+ all_embeddings = np.asarray([emb.numpy() for emb in all_embeddings])
+ return all_embeddings
+
+ def save(self, output_path, merge_before_save=False, save_config=True):
+ if merge_before_save and isinstance(self.model, PeftModel):
+ self.model = self.model.merge_and_unload()
+ # Fixes the issue of saving - https://huggingface.co/McGill-NLP/LLM2Vec-Mistral-7B-Instruct-v2-mntp-unsup-simcse/discussions/1
+ if hasattr(self.model, "_hf_peft_config_loaded"):
+ self.model._hf_peft_config_loaded = False
+
+ self.model.save_pretrained(output_path)
+ self.tokenizer.save_pretrained(output_path)
+
+ llm2vec_config = {
+ "pooling_mode": self.pooling_mode,
+ "max_length": self.max_length,
+ "doc_max_length": self.doc_max_length,
+ "skip_instruction": self.skip_instruction,
+ }
+
+ if save_config:
+ os.makedirs(output_path, exist_ok=True)
+ with open(f"{output_path}/llm2vec_config.json", "w") as fOut:
+ json.dump(llm2vec_config, fOut, indent=4)
+
+ def _encode(
+ self,
+ sentences_batch,
+ device: Optional[str] = None,
+ convert_to_numpy: bool = False,
+ multiprocessing=False,
+ ):
+ if multiprocessing:
+ # multiprocessing only supports CUDA devices at this time, so we ignore the value of device
+ # and use cuda:rank for the device
+ rank = mp.current_process()._identity[0]
+ if device is None and torch.cuda.is_available():
+ device = f"cuda:{rank % torch.cuda.device_count()}"
+
+ self.to(device)
+ features = self.tokenize([self.prepare_for_tokenization(sentence) for sentence in sentences_batch])
+ features = batch_to_device(features, device)
+
+ with torch.no_grad():
+ embeddings = self.forward(features)
+ embeddings = embeddings.detach()
+ embeddings = embeddings.cpu()
+
+ return embeddings
+
+ def _text_length(self, text: Union[List[int], List[List[int]]]):
+ """Help function to get the length for the input text.
+
+ Text can be either a string (which means a single text) a list of ints (which means a single
+ tokenized text), or a tuple of list of ints (representing several text inputs to the model).
+ """
+ if (
+ isinstance(text, str) or (isinstance(text, list) and isinstance(text[0], int)) or len(text) == 0
+ ): # Single text, list of ints, or empty
+ return len(text)
+ if isinstance(text, dict): # {key: value} case
+ return len(next(iter(text.values())))
+ elif not hasattr(text, "__len__"): # Object has no len() method
+ return 1
+ else:
+ return sum([len(t) for t in text])
+
+ def resize_token_embeddings(
+ self,
+ new_num_tokens: Optional[int] = None,
+ pad_to_multiple_of: Optional[int] = None,
+ ) -> nn.Embedding:
+ return self.model.resize_token_embeddings(new_num_tokens=new_num_tokens, pad_to_multiple_of=pad_to_multiple_of)
+
+ def gradient_checkpointing_enable(self, gradient_checkpointing_kwargs=None):
+ self.model.gradient_checkpointing_enable(gradient_checkpointing_kwargs=gradient_checkpointing_kwargs)
diff --git a/kimodo/model/llm2vec/llm2vec_wrapper.py b/kimodo/model/llm2vec/llm2vec_wrapper.py
new file mode 100644
index 0000000000000000000000000000000000000000..3b08417de8e6885739f8d345fb9615355b628c5e
--- /dev/null
+++ b/kimodo/model/llm2vec/llm2vec_wrapper.py
@@ -0,0 +1,95 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""LLM2Vec encoder wrapper for Kimodo text conditioning."""
+
+import os
+
+import numpy as np
+import torch
+
+from .llm2vec import LLM2Vec
+
+
+class LLM2VecEncoder:
+ """LLM2Vec text embeddings."""
+
+ def __init__(
+ self,
+ base_model_name_or_path: str,
+ peft_model_name_or_path: str,
+ dtype: str,
+ llm_dim: int,
+ device: str = "auto",
+ ) -> None:
+ torch_dtype = getattr(torch, dtype)
+ self.llm_dim = llm_dim
+
+ cache_dir = os.environ.get("HUGGINGFACE_CACHE_DIR")
+
+ if "TEXT_ENCODERS_DIR" in os.environ:
+ base_model_name_or_path = os.path.join(os.environ["TEXT_ENCODERS_DIR"], base_model_name_or_path)
+ peft_model_name_or_path = os.path.join(os.environ["TEXT_ENCODERS_DIR"], peft_model_name_or_path)
+
+ self.model = LLM2Vec.from_pretrained(
+ base_model_name_or_path=base_model_name_or_path,
+ peft_model_name_or_path=peft_model_name_or_path,
+ torch_dtype=torch_dtype,
+ cache_dir=cache_dir,
+ )
+
+ env_device = os.environ.get("TEXT_ENCODER_DEVICE")
+ if env_device:
+ device = env_device
+ if device == "auto":
+ device = "cuda" if torch.cuda.is_available() else "cpu"
+ self._device = device
+ if device is not None:
+ self.model = self.model.to(device)
+
+ self.model.eval()
+ for p in self.model.parameters():
+ p.requires_grad = False
+
+ def to(self, device: torch.device):
+ self.model = self.model.to(device)
+ self._device = str(device) if not isinstance(device, str) else device
+ return self
+
+ def eval(self):
+ self.model.eval()
+ return self
+
+ def get_device(self):
+ return self.model.model.device
+
+ def __call__(self, text: list[str] | str):
+ is_string = False
+ if isinstance(text, str):
+ text = [text]
+ is_string = True
+
+ with torch.no_grad():
+ encoded_text = self.model.encode(
+ text,
+ # IMPORTANT: different batch sizes unexpectedly change the output embeddings, so we always set it to 1
+ # here for repeatability no matter how many texts are being encoded. This
+ # is a fundamental issue with transformers, and is especially bad at lower
+ # precisions (https://github.com/huggingface/transformers/issues/25420#issuecomment-1775317535)
+ # note: this is an internal batch size used by llm2vec - the text list can still be of arbitrary length.
+ batch_size=1,
+ show_progress_bar=False,
+ device=self._device,
+ )
+
+ assert len(encoded_text.shape)
+ assert self.llm_dim == encoded_text.shape[-1]
+
+ encoded_text = encoded_text[:, None]
+ lengths = np.ones(len(encoded_text), dtype=int).tolist()
+
+ if is_string:
+ encoded_text = encoded_text[0]
+ lengths = lengths[0]
+
+ encoded_text = torch.tensor(encoded_text).to(self._device)
+ return encoded_text, lengths
diff --git a/kimodo/model/llm2vec/models/__init__.py b/kimodo/model/llm2vec/models/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..d2504048d0ee1addb1d3c95cbc04aae2b59e3e68
--- /dev/null
+++ b/kimodo/model/llm2vec/models/__init__.py
@@ -0,0 +1,4 @@
+# from .bidirectional_gemma import GemmaBiForMNTP, GemmaBiModel
+# from .bidirectional_llama import LlamaBiForMNTP, LlamaBiModel
+# from .bidirectional_mistral import MistralBiForMNTP, MistralBiModel
+# from .bidirectional_qwen2 import Qwen2BiForMNTP, Qwen2BiModel
diff --git a/kimodo/model/llm2vec/models/attn_mask_utils.py b/kimodo/model/llm2vec/models/attn_mask_utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..787f99172c52b3e340b037025ad6ac4c6f8f1929
--- /dev/null
+++ b/kimodo/model/llm2vec/models/attn_mask_utils.py
@@ -0,0 +1,181 @@
+# SPDX-FileCopyrightText: Copyright (c) 2024 McGill NLP
+# SPDX-License-Identifier: MIT
+#
+# Permission is hereby granted, free of charge, to any person obtaining a
+# copy of this software and associated documentation files (the "Software"),
+# to deal in the Software without restriction, including without limitation
+# the rights to use, copy, modify, merge, publish, distribute, sublicense,
+# and/or sell copies of the Software, and to permit persons to whom the
+# Software is furnished to do so, subject to the following conditions:
+#
+# The above copyright notice and this permission notice shall be included in
+# all copies or substantial portions of the Software.
+#
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL
+# THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING
+# FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER
+# DEALINGS IN THE SOFTWARE.
+
+from typing import List, Optional, Tuple, Union
+
+import torch
+from transformers.modeling_attn_mask_utils import AttentionMaskConverter
+
+
+def _prepare_4d_causal_attention_mask(
+ attention_mask: Optional[torch.Tensor],
+ input_shape: Union[torch.Size, Tuple, List],
+ inputs_embeds: torch.Tensor,
+ past_key_values_length: int,
+ sliding_window: Optional[int] = None,
+):
+ """Creates a causal 4D mask of shape `(batch_size, 1, query_length, key_value_length)` from a 2D
+ mask of shape `(batch_size, key_value_length)`
+
+ Args:
+ attention_mask (`torch.Tensor` or `None`):
+ A 2D attention mask of shape `(batch_size, key_value_length)`
+ input_shape (`tuple(int)` or `list(int)` or `torch.Size`):
+ The input shape should be a tuple that defines `(batch_size, query_length)`.
+ inputs_embeds (`torch.Tensor`):
+ The embedded inputs as a torch Tensor.
+ past_key_values_length (`int`):
+ The length of the key value cache.
+ sliding_window (`int`, *optional*):
+ If the model uses windowed attention, a sliding window should be passed.
+ """
+ attn_mask_converter = AttentionMaskConverter(
+ is_causal=False, sliding_window=sliding_window
+ ) # is_causal=True in original implementation
+
+ key_value_length = input_shape[-1] + past_key_values_length
+
+ # 4d mask is passed through the layers
+ if attention_mask is not None and len(attention_mask.shape) == 2:
+ attention_mask = attn_mask_converter.to_4d(
+ attention_mask,
+ input_shape[-1],
+ key_value_length=key_value_length,
+ dtype=inputs_embeds.dtype,
+ )
+ elif attention_mask is not None and len(attention_mask.shape) == 4:
+ expected_shape = (input_shape[0], 1, input_shape[1], key_value_length)
+ if tuple(attention_mask.shape) != expected_shape:
+ raise ValueError(
+ f"Incorrect 4D attention_mask shape: {tuple(attention_mask.shape)}; expected: {expected_shape}."
+ )
+ else:
+ # if the 4D mask has correct shape - invert it and fill with negative infinity
+ inverted_mask = 1.0 - attention_mask
+ attention_mask = inverted_mask.masked_fill(
+ inverted_mask.to(torch.bool), torch.finfo(inputs_embeds.dtype).min
+ )
+ else:
+ attention_mask = attn_mask_converter.to_causal_4d(
+ input_shape[0],
+ input_shape[-1],
+ key_value_length,
+ dtype=inputs_embeds.dtype,
+ device=inputs_embeds.device,
+ )
+
+ return attention_mask
+
+
+# Adapted from _prepare_4d_causal_attention_mask
+def _prepare_4d_causal_attention_mask_for_sdpa(
+ attention_mask: Optional[torch.Tensor],
+ input_shape: Union[torch.Size, Tuple, List],
+ inputs_embeds: torch.Tensor,
+ past_key_values_length: int,
+ sliding_window: Optional[int] = None,
+):
+ """Prepares the correct `attn_mask` argument to be used by
+ `torch.nn.functional.scaled_dot_product_attention`.
+
+ In case no token is masked in the `attention_mask` argument, we simply set it to `None` for the cases `query_length == 1` and
+ `key_value_length == query_length`, and rely instead on SDPA `is_causal` argument to use causal/non-causal masks,
+ allowing to dispatch to the flash attention kernel (that can otherwise not be used if a custom `attn_mask` is passed).
+ """
+ attn_mask_converter = AttentionMaskConverter(
+ is_causal=False, sliding_window=sliding_window
+ ) # is_causal=True in original implementation
+
+ key_value_length = input_shape[-1] + past_key_values_length
+ batch_size, query_length = input_shape
+
+ # torch.jit.trace, symbolic_trace and torchdynamo with fullgraph=True are unable to capture the controlflow `is_causal=attention_mask is None and q_len > 1`
+ # used as an SDPA argument. We keep compatibility with these tracing tools by always using SDPA's `attn_mask` argument in case we are tracing.
+ # TODO: For dynamo, rather use a check on fullgraph=True once this is possible (https://github.com/pytorch/pytorch/pull/120400).
+ is_tracing = (
+ torch.jit.is_tracing()
+ or isinstance(inputs_embeds, torch.fx.Proxy)
+ or (hasattr(torch, "_dynamo") and torch._dynamo.is_compiling())
+ )
+
+ if attention_mask is not None:
+ # 4d mask is passed through
+ if len(attention_mask.shape) == 4:
+ expected_shape = (input_shape[0], 1, input_shape[1], key_value_length)
+ if tuple(attention_mask.shape) != expected_shape:
+ raise ValueError(
+ f"Incorrect 4D attention_mask shape: {tuple(attention_mask.shape)}; expected: {expected_shape}."
+ )
+ else:
+ # if the 4D mask has correct shape - invert it and fill with negative infinity
+ inverted_mask = 1.0 - attention_mask.to(inputs_embeds.dtype)
+ attention_mask = inverted_mask.masked_fill(
+ inverted_mask.to(torch.bool), torch.finfo(inputs_embeds.dtype).min
+ )
+ return attention_mask
+
+ elif not is_tracing and torch.all(attention_mask == 1):
+ if query_length == 1:
+ # For query_length == 1, causal attention and bi-directional attention are the same.
+ attention_mask = None
+ elif key_value_length == query_length:
+ attention_mask = None
+ else:
+ # Unfortunately, for query_length > 1 and key_value_length != query_length, we cannot generally ignore the attention mask, as SDPA causal mask generation
+ # may be wrong. We will set `is_causal=False` in SDPA and rely on Transformers attention_mask instead, hence not setting it to None here.
+ # Reference: https://github.com/pytorch/pytorch/issues/108108
+ pass
+ elif query_length > 1 and key_value_length != query_length:
+ # See the comment above (https://github.com/pytorch/pytorch/issues/108108).
+ # Ugly: we set it to True here to dispatch in the following controlflow to `to_causal_4d`.
+ attention_mask = True
+ elif is_tracing:
+ raise ValueError(
+ 'Attention using SDPA can not be traced with torch.jit.trace when no attention_mask is provided. To solve this issue, please either load your model with the argument `attn_implementation="eager"` or pass an attention_mask input when tracing the model.'
+ )
+
+ if attention_mask is None:
+ expanded_4d_mask = None
+ elif attention_mask is True:
+ expanded_4d_mask = attn_mask_converter.to_causal_4d(
+ input_shape[0],
+ input_shape[-1],
+ key_value_length,
+ dtype=inputs_embeds.dtype,
+ device=inputs_embeds.device,
+ )
+ else:
+ expanded_4d_mask = attn_mask_converter.to_4d(
+ attention_mask,
+ input_shape[-1],
+ dtype=inputs_embeds.dtype,
+ key_value_length=key_value_length,
+ )
+
+ # Attend to all tokens in masked rows from the causal_mask, for example the relevant first rows when
+ # using left padding. This is required by F.scaled_dot_product_attention memory-efficient attention path.
+ # Details: https://github.com/pytorch/pytorch/issues/110213
+ if not is_tracing and expanded_4d_mask.device.type == "cuda":
+ expanded_4d_mask = AttentionMaskConverter._unmask_unattended(
+ expanded_4d_mask, min_dtype=torch.finfo(inputs_embeds.dtype).min
+ )
+
+ return expanded_4d_mask
diff --git a/kimodo/model/llm2vec/models/bidirectional_llama.py b/kimodo/model/llm2vec/models/bidirectional_llama.py
new file mode 100644
index 0000000000000000000000000000000000000000..f3e624e6d342ac37bd1eaecb85d9b76823de6864
--- /dev/null
+++ b/kimodo/model/llm2vec/models/bidirectional_llama.py
@@ -0,0 +1,224 @@
+# SPDX-FileCopyrightText: Copyright (c) 2024 McGill NLP
+# SPDX-License-Identifier: MIT
+#
+# Permission is hereby granted, free of charge, to any person obtaining a
+# copy of this software and associated documentation files (the "Software"),
+# to deal in the Software without restriction, including without limitation
+# the rights to use, copy, modify, merge, publish, distribute, sublicense,
+# and/or sell copies of the Software, and to permit persons to whom the
+# Software is furnished to do so, subject to the following conditions:
+#
+# The above copyright notice and this permission notice shall be included in
+# all copies or substantial portions of the Software.
+#
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL
+# THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING
+# FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER
+# DEALINGS IN THE SOFTWARE.
+
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+#
+# Licensed under the Apache License, Version 2.0 (the "License");
+# you may not use this file except in compliance with the License.
+# You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+
+import torch
+from peft import PeftModel
+from torch import nn
+from transformers import LlamaConfig, LlamaForCausalLM, LlamaModel, LlamaPreTrainedModel
+from transformers.cache_utils import Cache, StaticCache
+from transformers.modeling_attn_mask_utils import AttentionMaskConverter
+from transformers.models.llama.modeling_llama import (
+ LlamaAttention,
+ LlamaDecoderLayer,
+ # LlamaFlashAttention2,
+ LlamaMLP,
+ LlamaRMSNorm,
+ LlamaRotaryEmbedding,
+ # LlamaSdpaAttention,
+)
+from transformers.utils import logging
+
+from .utils import is_transformers_attn_greater_or_equal_4_43_1
+
+logger = logging.get_logger(__name__)
+
+
+class ModifiedLlamaAttention(LlamaAttention):
+ def __init__(self, *args, **kwargs):
+ super().__init__(*args, **kwargs)
+ self.is_causal = False
+
+
+# class ModifiedLlamaFlashAttention2(LlamaFlashAttention2):
+# def __init__(self, *args, **kwargs):
+# super().__init__(*args, **kwargs)
+# self.is_causal = False
+
+
+# class ModifiedLlamaSdpaAttention(LlamaSdpaAttention):
+# def __init__(self, *args, **kwargs):
+# super().__init__(*args, **kwargs)
+# self.is_causal = False
+
+
+# LLAMA_ATTENTION_CLASSES = {
+# "eager": ModifiedLlamaAttention,
+# "flash_attention_2": ModifiedLlamaFlashAttention2,
+# "sdpa": ModifiedLlamaSdpaAttention,
+# }
+
+
+class ModifiedLlamaDecoderLayer(LlamaDecoderLayer):
+ def __init__(self, config: LlamaConfig, layer_idx: int):
+ nn.Module.__init__(self)
+ self.hidden_size = config.hidden_size
+
+ self.self_attn = ModifiedLlamaAttention(config=config, layer_idx=layer_idx)
+ # self.self_attn = LLAMA_ATTENTION_CLASSES[config._attn_implementation](
+ # config=config, layer_idx=layer_idx
+ # )
+
+ self.mlp = LlamaMLP(config)
+ self.input_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+ self.post_attention_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+
+
+class LlamaBiModel(LlamaModel):
+ _no_split_modules = ["ModifiedLlamaDecoderLayer"]
+
+ def __init__(self, config: LlamaConfig):
+ if not is_transformers_attn_greater_or_equal_4_43_1():
+ raise ValueError(
+ "The current implementation of LlamaEncoderModel follows modeling_llama.py of transformers version >= 4.43.1"
+ )
+ LlamaPreTrainedModel.__init__(self, config)
+ self.padding_idx = config.pad_token_id
+ self.vocab_size = config.vocab_size
+
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
+ self.layers = nn.ModuleList(
+ [ModifiedLlamaDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
+ )
+ self.norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+ self.rotary_emb = LlamaRotaryEmbedding(config=config)
+ self.gradient_checkpointing = False
+
+ # Initialize weights and apply final processing
+ self.post_init()
+
+ def _update_causal_mask(
+ self,
+ attention_mask,
+ input_tensor,
+ cache_position,
+ past_key_values: Cache,
+ output_attentions: bool,
+ ):
+ if self.config._attn_implementation == "flash_attention_2":
+ if attention_mask is not None and 0.0 in attention_mask:
+ return attention_mask
+ return None
+
+ # For SDPA, when possible, we will rely on its `is_causal` argument instead of its `attn_mask` argument, in
+ # order to dispatch on Flash Attention 2. This feature is not compatible with static cache, as SDPA will fail
+ # to infer the attention mask.
+ past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
+ using_static_cache = isinstance(past_key_values, StaticCache)
+
+ # When output attentions is True, sdpa implementation's forward method calls the eager implementation's forward
+ # if self.config._attn_implementation == "sdpa" and not using_static_cache and not output_attentions:
+ # if AttentionMaskConverter._ignore_causal_mask_sdpa(
+ # attention_mask,
+ # inputs_embeds=input_tensor,
+ # past_key_values_length=past_seen_tokens,
+ # is_training=self.training,
+ # ):
+ # return None
+
+ dtype, device = input_tensor.dtype, input_tensor.device
+ min_dtype = torch.finfo(dtype).min
+ sequence_length = input_tensor.shape[1]
+ if using_static_cache:
+ target_length = past_key_values.get_max_length()
+ else:
+ target_length = (
+ attention_mask.shape[-1]
+ if isinstance(attention_mask, torch.Tensor)
+ else past_seen_tokens + sequence_length + 1
+ )
+
+ causal_mask = torch.zeros(
+ (sequence_length, target_length), dtype=dtype, device=device
+ ) # in original implementation - torch.full((sequence_length, target_length), fill_value=min_dtype, dtype=dtype, device=device)
+ # Commenting out next 2 lines to disable causal masking
+ # if sequence_length != 1:
+ # causal_mask = torch.triu(causal_mask, diagonal=1)
+ causal_mask *= torch.arange(target_length, device=device) > cache_position.reshape(-1, 1)
+ causal_mask = causal_mask[None, None, :, :].expand(input_tensor.shape[0], 1, -1, -1)
+ if attention_mask is not None:
+ causal_mask = causal_mask.clone() # copy to contiguous memory for in-place edit
+ if attention_mask.dim() == 2:
+ mask_length = attention_mask.shape[-1]
+ padding_mask = causal_mask[..., :mask_length].eq(0.0) * attention_mask[:, None, None, :].eq(0.0)
+ causal_mask[..., :mask_length] = causal_mask[..., :mask_length].masked_fill(padding_mask, min_dtype)
+ elif attention_mask.dim() == 4:
+ # backwards compatibility: we allow passing a 4D attention mask shorter than the input length with
+ # cache. In that case, the 4D attention mask attends to the newest tokens only.
+ if attention_mask.shape[-2] < cache_position[0] + sequence_length:
+ offset = cache_position[0]
+ else:
+ offset = 0
+ mask_shape = attention_mask.shape
+ mask_slice = (attention_mask.eq(0.0)).to(dtype=dtype) * min_dtype
+ causal_mask[
+ : mask_shape[0],
+ : mask_shape[1],
+ offset : mask_shape[2] + offset,
+ : mask_shape[3],
+ ] = mask_slice
+
+ if (
+ self.config._attn_implementation == "sdpa"
+ and attention_mask is not None
+ and attention_mask.device.type == "cuda"
+ and not output_attentions
+ ):
+ causal_mask = AttentionMaskConverter._unmask_unattended(causal_mask, min_dtype)
+
+ return causal_mask
+
+
+class LlamaBiForMNTP(LlamaForCausalLM):
+ def __init__(self, config):
+ LlamaPreTrainedModel.__init__(self, config)
+ self.model = LlamaBiModel(config)
+ self.vocab_size = config.vocab_size
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
+
+ # Initialize weights and apply final processing
+ self.post_init()
+
+ # getter for PEFT model
+ def get_model_for_peft(self):
+ return self.model
+
+ # setter for PEFT model
+ def set_model_for_peft(self, model: PeftModel):
+ self.model = model
+
+ # save the PEFT model
+ def save_peft_model(self, path):
+ self.model.save_pretrained(path)
diff --git a/kimodo/model/llm2vec/models/utils.py b/kimodo/model/llm2vec/models/utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..92ed8ec7058668697e628da342621b7599fb5ca2
--- /dev/null
+++ b/kimodo/model/llm2vec/models/utils.py
@@ -0,0 +1,32 @@
+# SPDX-FileCopyrightText: Copyright (c) 2024 McGill NLP
+# SPDX-License-Identifier: MIT
+#
+# Permission is hereby granted, free of charge, to any person obtaining a
+# copy of this software and associated documentation files (the "Software"),
+# to deal in the Software without restriction, including without limitation
+# the rights to use, copy, modify, merge, publish, distribute, sublicense,
+# and/or sell copies of the Software, and to permit persons to whom the
+# Software is furnished to do so, subject to the following conditions:
+#
+# The above copyright notice and this permission notice shall be included in
+# all copies or substantial portions of the Software.
+#
+# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL
+# THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING
+# FROM, OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER
+# DEALINGS IN THE SOFTWARE.
+
+import importlib.metadata
+
+from packaging import version
+from transformers.utils.import_utils import _is_package_available
+
+
+def is_transformers_attn_greater_or_equal_4_43_1():
+ if not _is_package_available("transformers"):
+ return False
+
+ return version.parse(importlib.metadata.version("transformers")) >= version.parse("4.43.1")
diff --git a/kimodo/model/load_model.py b/kimodo/model/load_model.py
new file mode 100644
index 0000000000000000000000000000000000000000..b732d0f67ff0474567aa170c114734e5c71d17e2
--- /dev/null
+++ b/kimodo/model/load_model.py
@@ -0,0 +1,215 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Load Kimodo diffusion models from local checkpoints or Hugging Face."""
+
+from pathlib import Path
+from typing import Optional
+
+from huggingface_hub import snapshot_download
+from omegaconf import OmegaConf
+
+from .loading import (
+ AVAILABLE_MODELS,
+ DEFAULT_MODEL,
+ DEFAULT_TEXT_ENCODER_URL,
+ MODEL_NAMES,
+ TMR_MODELS,
+ get_env_var,
+ instantiate_from_dict,
+)
+from .registry import get_model_info, resolve_model_name
+
+DEFAULT_TEXT_ENCODER = "llm2vec"
+TEXT_ENCODER_PRESETS = {
+ "llm2vec": {
+ "target": "kimodo.model.LLM2VecEncoder",
+ "kwargs": {
+ "base_model_name_or_path": "McGill-NLP/LLM2Vec-Meta-Llama-3-8B-Instruct-mntp",
+ "peft_model_name_or_path": "McGill-NLP/LLM2Vec-Meta-Llama-3-8B-Instruct-mntp-supervised",
+ "dtype": "bfloat16",
+ "llm_dim": 4096,
+ "device": "auto",
+ },
+ }
+}
+
+
+def _resolve_hf_model_path(modelname: str) -> Path:
+ """Resolve model name to a local path, using Hugging Face cache or CHECKPOINT_DIR."""
+ try:
+ repo_id = MODEL_NAMES[modelname]
+ except KeyError:
+ raise ValueError(f"Model '{modelname}' not found. Available models: {MODEL_NAMES.keys()}")
+
+ local_cache = get_env_var("LOCAL_CACHE", "False").lower() == "true"
+ if not local_cache:
+ snapshot_dir = snapshot_download(repo_id=repo_id) # will check online no matter what
+ return Path(snapshot_dir)
+
+ try:
+ snapshot_dir = snapshot_download(repo_id=repo_id, local_files_only=True) # will check local cache only
+ return Path(snapshot_dir)
+ except Exception:
+ # if local cache is not found, download from online
+ try:
+ snapshot_dir = snapshot_download(repo_id=repo_id)
+ return Path(snapshot_dir)
+ except Exception:
+ raise RuntimeError(f"Could not resolve model '{modelname}' from Hugging Face (repo: {repo_id}). ") from None
+
+
+def _build_api_text_encoder_conf(text_encoder_url: str) -> dict:
+ return {
+ "_target_": "kimodo.model.text_encoder_api.TextEncoderAPI",
+ "url": text_encoder_url,
+ }
+
+
+def _build_local_text_encoder_conf(text_encoder_fp32: bool = False) -> dict:
+ text_encoder_name = get_env_var("TEXT_ENCODER", DEFAULT_TEXT_ENCODER)
+ if text_encoder_name not in TEXT_ENCODER_PRESETS:
+ available = ", ".join(sorted(TEXT_ENCODER_PRESETS))
+ raise ValueError(f"Unknown TEXT_ENCODER='{text_encoder_name}'. Available: {available}")
+
+ preset = TEXT_ENCODER_PRESETS[text_encoder_name]
+ if text_encoder_fp32:
+ preset["kwargs"]["dtype"] = "float32"
+ return {
+ "_target_": preset["target"],
+ **preset["kwargs"],
+ }
+
+
+def _select_text_encoder_conf(text_encoder_url: str, text_encoder_fp32: bool = False) -> dict:
+ # TEXT_ENCODER_MODE options:
+ # - "api": force TextEncoderAPI
+ # - "local": force local LLM2VecEncoder
+ # - "auto": try API first, fallback to local if unreachable
+ mode = get_env_var("TEXT_ENCODER_MODE", "auto").lower()
+ if mode == "local":
+ return _build_local_text_encoder_conf(text_encoder_fp32)
+ if mode == "api":
+ return _build_api_text_encoder_conf(text_encoder_url)
+
+ api_conf = _build_api_text_encoder_conf(text_encoder_url)
+ try:
+ text_encoder = instantiate_from_dict(api_conf)
+ # Probe availability early so inference doesn't fail later.
+ text_encoder(["healthcheck"])
+ return api_conf
+ except Exception as error:
+ print(
+ "Text encoder service is unreachable, falling back to local LLM2Vec "
+ f"encoder. ({type(error).__name__}: {error})"
+ )
+ return _build_local_text_encoder_conf(text_encoder_fp32)
+
+
+def load_model(
+ modelname=None,
+ device=None,
+ eval_mode: bool = True,
+ default_family: Optional[str] = "Kimodo",
+ return_resolved_name: bool = False,
+ text_encoder=None,
+ text_encoder_fp32: bool = False,
+):
+ """Load a kimodo model by name (e.g. 'g1', 'soma').
+
+ Resolution of partial/full names (e.g. Kimodo-SOMA-RP-v1, SOMA) is done
+ inside this function using default_family when the name is not a known
+ short key.
+
+ Args:
+ modelname: Model identifier; uses DEFAULT_MODEL if None. Can be a short key,
+ a full name (e.g. Kimodo-SOMA-RP-v1), or a partial name; unknown names
+ are resolved via resolve_model_name using default_family.
+ device: Target device for the model (e.g. 'cuda', 'cpu').
+ eval_mode: If True, set model to eval mode.
+ default_family: Used when modelname is not in AVAILABLE_MODELS to resolve
+ partial names ("Kimodo" for demo/generation, "TMR" for embed script).
+ Default "Kimodo".
+ return_resolved_name: If True, return (model, resolved_short_key). If False,
+ return only the model.
+ text_encoder: Pre-built text encoder to reuse. When provided, skips
+ text encoder selection/instantiation entirely.
+ text_encoder_fp32: If True, uses fp32 for the text encoder rather than default bfloat16.
+
+ Returns:
+ Loaded model in eval mode, or (model, resolved short key) if
+ return_resolved_name is True.
+
+ Raises:
+ ValueError: If modelname is not in AVAILABLE_MODELS and cannot be resolved.
+ FileNotFoundError: If config.yaml is missing in the checkpoint folder.
+ """
+ if modelname is None:
+ modelname = DEFAULT_MODEL
+ if modelname not in AVAILABLE_MODELS:
+ if default_family is not None:
+ modelname = resolve_model_name(modelname, default_family)
+ else:
+ raise ValueError(
+ f"""The model is not recognized.
+ Please choose between: {AVAILABLE_MODELS}"""
+ )
+
+ resolved_modelname = modelname
+
+ # In case, we specify a custom checkpoint directory
+ configured_checkpoint_dir = get_env_var("CHECKPOINT_DIR")
+ if configured_checkpoint_dir:
+ print(f"CHECKPOINT_DIR is set to {configured_checkpoint_dir}, checking the local cache...")
+ # Checkpoint folders are named by display name (e.g. Kimodo-SOMA-RP-v1)
+ info = get_model_info(modelname)
+ checkpoint_folder_name = info.display_name if info is not None else modelname
+ model_path = Path(configured_checkpoint_dir) / checkpoint_folder_name
+ if not model_path.exists() and modelname != checkpoint_folder_name:
+ # Fallback: try short_key for backward compatibility
+ model_path = Path(configured_checkpoint_dir) / modelname
+ if not model_path.exists():
+ print(f"Model folder not found at '{model_path}', downloading it from Hugging Face...")
+ model_path = _resolve_hf_model_path(modelname)
+ else:
+ # Otherwise, we load the model from the local cache or download it from Hugging Face.
+ model_path = _resolve_hf_model_path(modelname)
+
+ model_config_path = model_path / "config.yaml"
+ if not model_config_path.exists():
+ raise FileNotFoundError(f"The model checkpoint folder exists but config.yaml is missing: {model_config_path}")
+
+ model_conf = OmegaConf.load(model_config_path)
+
+ if modelname in TMR_MODELS:
+ # Same process at the moment for TMR and Kimodo
+ pass
+
+ if text_encoder is not None:
+ runtime_conf = OmegaConf.create({"checkpoint_dir": str(model_path)})
+ else:
+ text_encoder_url = get_env_var("TEXT_ENCODER_URL", DEFAULT_TEXT_ENCODER_URL)
+ runtime_conf = OmegaConf.create(
+ {
+ "checkpoint_dir": str(model_path),
+ "text_encoder": _select_text_encoder_conf(text_encoder_url, text_encoder_fp32),
+ }
+ )
+
+ model_cfg = OmegaConf.to_container(OmegaConf.merge(model_conf, runtime_conf), resolve=True)
+ model_cfg.pop("checkpoint_dir", None)
+
+ if text_encoder is not None:
+ # Prevent Hydra from instantiating a new text encoder; pass None so
+ # Kimodo.__init__ receives a placeholder we replace immediately after.
+ model_cfg["text_encoder"] = None
+
+ model = instantiate_from_dict(model_cfg, overrides={"device": device})
+
+ if text_encoder is not None:
+ model.text_encoder = text_encoder
+
+ if eval_mode:
+ model = model.eval()
+ if return_resolved_name:
+ return model, resolved_modelname
+ return model
diff --git a/kimodo/model/loading.py b/kimodo/model/loading.py
new file mode 100644
index 0000000000000000000000000000000000000000..b2636a210871cc08e8ab2c667807892e5f9cc539
--- /dev/null
+++ b/kimodo/model/loading.py
@@ -0,0 +1,81 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Model loading utilities: checkpoints, registry, env, and Hydra-based instantiation."""
+
+import os
+from pathlib import Path
+from typing import Any, Dict, Optional, Union
+
+import torch
+from hydra.utils import instantiate
+from omegaconf import OmegaConf
+from safetensors.torch import load_file as load_safetensors
+
+from .registry import (
+ AVAILABLE_MODELS,
+ DEFAULT_MODEL,
+ DEFAULT_TEXT_ENCODER_URL,
+ KIMODO_MODELS,
+ MODEL_NAMES,
+ TMR_MODELS,
+)
+
+
+def get_env_var(name: str, default: Optional[str] = None) -> Optional[str]:
+ """Return environment variable value, or default if unset/empty."""
+ return os.environ.get(name) or default
+
+
+def instantiate_from_dict(
+ cfg: Dict[str, Any],
+ overrides: Optional[Dict[str, Any]] = None,
+):
+ """Instantiate an object from a config dict (e.g. from OmegaConf.to_container).
+
+ The dict must contain _target_ with a fully qualified class path. Nested configs are
+ instantiated recursively.
+ """
+ if overrides:
+ cfg = {**cfg, **overrides}
+ conf = OmegaConf.create(cfg)
+ return instantiate(conf)
+
+
+def load_checkpoint_state_dict(ckpt_path: Union[str, Path]) -> dict:
+ """Load a state dict from a checkpoint file.
+
+ If the checkpoint is a dict with a 'state_dict' key (e.g. PyTorch Lightning),
+ that is returned; otherwise the whole checkpoint is treated as the state dict.
+
+ Args:
+ ckpt_path: Path to the checkpoint file.
+
+ Returns:
+ state_dict suitable for model.load_state_dict().
+ """
+ ckpt_path = str(ckpt_path)
+
+ if ckpt_path.endswith(".safetensors"):
+ state_dict = load_safetensors(ckpt_path)
+ else:
+ checkpoint = torch.load(ckpt_path, map_location="cpu", weights_only=False)
+ if isinstance(checkpoint, dict) and "state_dict" in checkpoint:
+ state_dict = checkpoint["state_dict"]
+ elif isinstance(checkpoint, dict):
+ state_dict = checkpoint
+ else:
+ raise ValueError(f"Unsupported checkpoint format: {ckpt_path}")
+ return {key: val.detach().cpu() for key, val in state_dict.items()}
+
+
+__all__ = [
+ "get_env_var",
+ "instantiate_from_dict",
+ "KIMODO_MODELS",
+ "TMR_MODELS",
+ "AVAILABLE_MODELS",
+ "MODEL_NAMES",
+ "DEFAULT_MODEL",
+ "DEFAULT_TEXT_ENCODER_URL",
+ "load_checkpoint_state_dict",
+]
diff --git a/kimodo/model/registry.py b/kimodo/model/registry.py
new file mode 100644
index 0000000000000000000000000000000000000000..34b03011a59b309bc65a964775bf957d0bb26cfe
--- /dev/null
+++ b/kimodo/model/registry.py
@@ -0,0 +1,478 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Registry of model names and Hugging Face repo IDs for Kimodo and TMR.
+
+Canonical source of truth is the list of repo IDs. Short keys (e.g. soma-rp) and metadata (dataset,
+skeleton, version, display name) are derived by parsing.
+"""
+
+import re
+from dataclasses import dataclass
+from typing import Optional
+
+# Canonical list: repo IDs in the same syntax as Hugging Face (org/Model-Name-v1).
+# Parser expects: org/Family-SKELETON-DATASET-version (e.g. Kimodo-SOMA-RP-v1).
+KIMODO_REPO_IDS = [
+ "nvidia/Kimodo-SOMA-RP-v1",
+ "nvidia/Kimodo-SOMA-RP-v1.1",
+ "nvidia/Kimodo-SMPLX-RP-v1",
+ "nvidia/Kimodo-G1-RP-v1",
+ "nvidia/Kimodo-SOMA-SEED-v1",
+ "nvidia/Kimodo-SOMA-SEED-v1.1",
+ "nvidia/Kimodo-G1-SEED-v1",
+]
+TMR_REPO_IDS = [
+ "nvidia/TMR-SOMA-RP-v1",
+]
+
+# Repo ID without org, for display (e.g. Kimodo-SOMA-RP-v1).
+_REPO_NAME_PATTERN = re.compile(r"^(Kimodo|TMR)-([A-Za-z0-9]+)-(RP|SEED)-v(\d+(?:\.\d+)*)$")
+
+
+@dataclass
+class ModelInfo:
+ """Structured metadata for one model, derived from its repo ID."""
+
+ repo_id: str
+ short_key: str
+ family: str
+ skeleton: str
+ dataset: str
+ version: str
+ display_name: str
+
+ @property
+ def dataset_ui_label(self) -> str:
+ return "Rigplay" if self.dataset == "RP" else "SEED"
+
+
+def _parse_repo_id(repo_id: str) -> Optional[ModelInfo]:
+ """Parse a repo ID into ModelInfo.
+
+ Returns None if format is unrecognized.
+ """
+ # repo_id is "org/Model-Name-v1"
+ if "/" in repo_id:
+ _, name = repo_id.split("/", 1)
+ else:
+ name = repo_id
+ m = _REPO_NAME_PATTERN.match(name)
+ if not m:
+ return None
+ family, skeleton, dataset, ver = m.groups()
+ # Normalize skeleton for display (as is for now)
+ skeleton_display = skeleton
+ # Include family so Kimodo-SOMA-RP and TMR-SOMA-RP have distinct keys.
+ short_key = f"{family.lower()}-{skeleton.lower()}-{dataset.lower()}"
+ return ModelInfo(
+ repo_id=repo_id,
+ short_key=short_key,
+ family=family,
+ skeleton=skeleton_display,
+ dataset=dataset,
+ version=f"v{ver}",
+ display_name=name,
+ )
+
+
+def _version_tuple(v: str) -> tuple[int, ...]:
+ """Parse 'vN' or 'vN.M' into a comparable tuple of ints."""
+ if v.startswith("v"):
+ parts = v[1:].split(".")
+ if all(p.isdigit() for p in parts):
+ return tuple(int(p) for p in parts)
+ return (0,)
+
+
+def _version_key(info: ModelInfo) -> tuple[int, ...]:
+ return _version_tuple(info.version)
+
+
+def _build_registry() -> tuple[list[ModelInfo], dict[str, str], list[str]]:
+ """Build model infos, short_key -> repo_id map, and list of short keys.
+
+ When multiple versions exist for the same (family, skeleton, dataset), each ModelInfo gets a
+ version-specific short_key (e.g. kimodo-soma-rp-v1, kimodo-soma-rp-v2) and a versionless alias
+ (kimodo-soma-rp) is added to MODEL_NAMES pointing to the latest version. When only one version
+ exists, the short_key stays versionless (e.g. kimodo-smplx-rp).
+ """
+ all_repos = KIMODO_REPO_IDS + TMR_REPO_IDS
+ infos: list[ModelInfo] = []
+ for repo_id in all_repos:
+ info = _parse_repo_id(repo_id)
+ if info is None:
+ raise ValueError(f"Registry repo ID does not match expected pattern: {repo_id}")
+ infos.append(info)
+
+ # Group by base short_key to detect multi-version families.
+ base_groups: dict[str, list[ModelInfo]] = {}
+ for info in infos:
+ base_groups.setdefault(info.short_key, []).append(info)
+
+ # For groups with multiple versions, make each short_key version-specific.
+ for base_key, group in base_groups.items():
+ if len(group) > 1:
+ for info in group:
+ info.short_key = f"{base_key}-{info.version}"
+
+ # Map each (now unique) short_key to its repo_id.
+ model_names: dict[str, str] = {}
+ for info in infos:
+ model_names[info.short_key] = info.repo_id
+
+ # Add versionless aliases for multi-version groups, pointing to the latest.
+ for base_key, group in base_groups.items():
+ if len(group) > 1:
+ latest = max(group, key=_version_key)
+ model_names[base_key] = latest.repo_id
+
+ return infos, model_names, list(model_names.keys())
+
+
+MODEL_INFOS, MODEL_NAMES, _SHORT_KEYS = _build_registry()
+AVAILABLE_MODELS = _SHORT_KEYS
+
+# Short-key lists for Kimodo vs TMR (load_model uses TMR_MODELS to branch).
+KIMODO_MODELS = [info.short_key for info in MODEL_INFOS if info.family == "Kimodo"]
+TMR_MODELS = [info.short_key for info in MODEL_INFOS if info.family == "TMR"]
+
+# Backward compatibility: FRIENDLY_NAMES for any code that still expects it.
+# Includes versioned short_keys and versionless aliases (latest display name).
+FRIENDLY_NAMES = {info.short_key: info.display_name for info in MODEL_INFOS}
+for _key, _repo_id in MODEL_NAMES.items():
+ if _key not in FRIENDLY_NAMES:
+ for _info in MODEL_INFOS:
+ if _info.repo_id == _repo_id:
+ FRIENDLY_NAMES[_key] = _info.display_name
+ break
+
+DEFAULT_MODEL = "kimodo-soma-rp"
+DEFAULT_TEXT_ENCODER_URL = "http://127.0.0.1:9550/"
+
+# Friendly names for skeleton dropdown (key -> label).
+SKELETON_DISPLAY_NAMES = {
+ "SOMA": "SOMA Human Body",
+ "SMPLX": "SMPLX Human Body",
+ "G1": "Unitree G1 Humanoid Robot",
+}
+
+# Order for skeleton dropdown: SOMA, SMPLX, G1.
+SKELETON_ORDER = ("SOMA", "SMPLX", "G1")
+
+
+def get_skeleton_display_name(skeleton_key: str) -> str:
+ """Return the UI label for a skeleton key (e.g. SOMA -> SOMA Human Body)."""
+ return SKELETON_DISPLAY_NAMES.get(skeleton_key, skeleton_key)
+
+
+def get_skeleton_key_from_display_name(display_name: str) -> Optional[str]:
+ """Return the skeleton key for a UI label, or None."""
+ for key, label in SKELETON_DISPLAY_NAMES.items():
+ if label == display_name:
+ return key
+ return None
+
+
+def get_skeleton_display_names_for_dataset(dataset_ui_label: str, family: Optional[str] = None) -> list[str]:
+ """Return skeleton UI labels for the given dataset.
+
+ If family is set (e.g. "Kimodo"), only skeletons with a model of that family are included.
+ """
+ keys = get_skeletons_for_dataset(dataset_ui_label, family=family)
+ return [get_skeleton_display_name(k) for k in keys]
+
+
+def get_short_key(repo_id: str) -> Optional[str]:
+ """Return the short key for a repo ID, or None if not in registry."""
+ for info in MODEL_INFOS:
+ if info.repo_id == repo_id:
+ return info.short_key
+ return None
+
+
+def get_model_info(short_key: str) -> Optional[ModelInfo]:
+ """Return ModelInfo for a short key, or None if not found.
+
+ When multiple versions share the same short_key, returns the one used for loading (the latest
+ version), so CHECKPOINT_DIR and HF use the same version.
+ """
+ repo_id = MODEL_NAMES.get(short_key)
+ if repo_id is None:
+ return None
+ for info in MODEL_INFOS:
+ if info.repo_id == repo_id:
+ return info
+ return None
+
+
+def get_short_key_from_display_name(display_name: str) -> Optional[str]:
+ """Return short_key for a display name (e.g. Kimodo-SOMA-RP-v1), or None."""
+ for info in MODEL_INFOS:
+ if info.display_name == display_name:
+ return info.short_key
+ return None
+
+
+def get_models_for_demo() -> list[ModelInfo]:
+ """Return all model infos in registry order (for demo model list)."""
+ return list(MODEL_INFOS)
+
+
+def get_datasets(family: Optional[str] = None) -> list[str]:
+ """Return unique dataset UI labels (Rigplay, SEED) present in registry.
+
+ If family is set (e.g. "Kimodo"), only datasets that have a model of that family are included.
+ """
+ infos = MODEL_INFOS
+ if family is not None:
+ infos = [i for i in infos if i.family == family]
+ labels = set()
+ for info in infos:
+ labels.add(info.dataset_ui_label)
+ return sorted(labels)
+
+
+def get_skeletons_for_dataset(dataset_ui_label: str, family: Optional[str] = None) -> list[str]:
+ """Return skeleton names that have a model for the given dataset.
+
+ Order: SOMA, SMPLX, G1 (only those present for the dataset).
+ If family is set (e.g. "Kimodo"), only skeletons with a model of that
+ family are included.
+ """
+ dataset = "RP" if dataset_ui_label == "Rigplay" else "SEED"
+ infos = MODEL_INFOS
+ if family is not None:
+ infos = [i for i in infos if i.family == family]
+ skeletons = set()
+ for info in infos:
+ if info.dataset == dataset:
+ skeletons.add(info.skeleton)
+ return [s for s in SKELETON_ORDER if s in skeletons]
+
+
+def get_versions_for_dataset_skeleton(dataset_ui_label: str, skeleton: str) -> list[str]:
+ """Return version strings (e.g. v1) for the given dataset/skeleton.
+
+ Sorted by version number so the last element is the highest (e.g. v1, v2).
+ """
+ dataset = "RP" if dataset_ui_label == "Rigplay" else "SEED"
+ versions = []
+ for info in MODEL_INFOS:
+ if info.dataset == dataset and info.skeleton == skeleton:
+ versions.append(info.version)
+
+ return sorted(set(versions), key=_version_tuple)
+
+
+def get_models_for_dataset_skeleton(
+ dataset_ui_label: str, skeleton: str, family: Optional[str] = None
+) -> list[ModelInfo]:
+ """Return model infos for the given dataset/skeleton, sorted by version (max first).
+
+ Used to build the Version dropdown (options = full display names, one per model). If family is
+ set (e.g. "Kimodo"), only models of that family are returned.
+ """
+ dataset = "RP" if dataset_ui_label == "Rigplay" else "SEED"
+ infos = [info for info in MODEL_INFOS if info.dataset == dataset and info.skeleton == skeleton]
+ if family is not None:
+ infos = [i for i in infos if i.family == family]
+
+ return sorted(infos, key=_version_key, reverse=True)
+
+
+def resolve_to_short_key(dataset_ui_label: str, skeleton: str, version: str) -> Optional[str]:
+ """Return the short key for (dataset, skeleton, version), or None."""
+ for info in MODEL_INFOS:
+ if info.dataset_ui_label == dataset_ui_label and info.skeleton == skeleton and info.version == version:
+ return info.short_key
+ return None
+
+
+# -----------------------------------------------------------------------------
+# Flexible model name resolution (partial names, case-insensitive, defaults)
+# -----------------------------------------------------------------------------
+
+_FAMILY_ALIASES = {"kimodo": "Kimodo", "tmr": "TMR"}
+_DATASET_ALIASES = {"rp": "RP", "rigplay": "RP", "seed": "SEED"}
+_SKELETON_ALIASES = {
+ "soma": "SOMA",
+ "smplx": "SMPLX",
+ "g1": "G1",
+}
+
+
+def _normalize_family(s: str) -> Optional[str]:
+ """Return canonical family (Kimodo/TMR) or None if unknown."""
+ return _FAMILY_ALIASES.get(s.strip().lower())
+
+
+def _normalize_dataset(s: str) -> Optional[str]:
+ """Return canonical dataset (RP/SEED) or None if unknown."""
+ return _DATASET_ALIASES.get(s.strip().lower())
+
+
+def _normalize_skeleton(s: str) -> Optional[str]:
+ """Return canonical skeleton (SOMA/SMPLX/G1) or None if unknown."""
+ return _SKELETON_ALIASES.get(s.strip().lower())
+
+
+def _get_latest_for_family_skeleton_dataset(family: str, skeleton: str, dataset: str) -> Optional[ModelInfo]:
+ """Return the model info with the highest version for (family, skeleton, dataset)."""
+ candidates = [
+ info for info in MODEL_INFOS if info.family == family and info.skeleton == skeleton and info.dataset == dataset
+ ]
+ if not candidates:
+ return None
+ return max(candidates, key=_version_key)
+
+
+def kimodo_short_key_for_skeleton_dataset(skeleton: str, dataset: str) -> Optional[str]:
+ """Return the latest Kimodo model short_key for ``skeleton`` and ``dataset`` (RP/SEED), or
+ None."""
+ info = _get_latest_for_family_skeleton_dataset("Kimodo", skeleton, dataset)
+ return info.short_key if info is not None else None
+
+
+def registry_skeleton_for_joint_count(nb_joints: int) -> str:
+ """Map motion joint count to registry skeleton key (SOMA / SMPLX / G1)."""
+ if nb_joints == 34:
+ return "G1"
+ if nb_joints == 22:
+ return "SMPLX"
+ if nb_joints in (77, 30):
+ return "SOMA"
+ raise ValueError(f"No Kimodo model registered for motion with J={nb_joints}")
+
+
+# Optional version: Family-Skeleton-Dataset-vN or Family-Skeleton-Dataset
+_RESOLVE_FULL_PATTERN = re.compile(
+ r"^(Kimodo|TMR|kimodo|tmr)[\-_]" r"([A-Za-z0-9]+)[\-_]" r"(RP|SEED|rp|seed)" r"(?:[\-_]v(\d+(?:\.\d+)*))?$",
+ re.IGNORECASE,
+)
+# Partial: Skeleton-Dataset or Skeleton or Dataset (no family)
+_RESOLVE_PARTIAL_PATTERN = re.compile(
+ r"^([A-Za-z0-9]+)(?:[\-_](RP|SEED|rp|seed))?(?:[\-_]v(\d+(?:\.\d+)*))?$",
+ re.IGNORECASE,
+)
+
+
+def resolve_model_name(name: Optional[str], default_family: Optional[str] = None) -> str:
+ """Resolve a user-facing model name to a short_key.
+
+ Accepts full names (e.g. Kimodo-SOMA-RP-v1), case-insensitive matching,
+ and partial names with defaults: dataset=RP, skeleton=SOMA, family from
+ default_family (Kimodo for demo/generation, TMR for embed script).
+ Omitted version resolves to the latest for that model.
+
+ Args:
+ name: User-provided name (can be None or empty).
+ default_family: "Kimodo" or "TMR" when name is empty or omits family.
+
+ Returns:
+ Short key (e.g. kimodo-soma-rp) for use with load_model / MODEL_NAMES.
+
+ Raises:
+ ValueError: If name cannot be resolved or default_family is missing when needed.
+ """
+ if name is not None:
+ name = name.strip()
+ if not name:
+ if default_family is None:
+ raise ValueError('Model name is empty; provide a name or set default_family ("Kimodo" or "TMR").')
+ fam = _normalize_family(default_family)
+ if fam is None:
+ raise ValueError(f"default_family must be 'Kimodo' or 'TMR', got {default_family!r}")
+ info = _get_latest_for_family_skeleton_dataset(fam, "SOMA", "RP")
+ if info is None:
+ raise ValueError(f"No model found for {fam}-SOMA-RP. Available: {list(MODEL_NAMES.keys())}")
+ return info.short_key
+
+ # Exact short_key
+ if name in MODEL_NAMES:
+ return name
+
+ # Case-insensitive match against short_key or display_name
+ name_lower = name.lower()
+ matches = []
+ for info in MODEL_INFOS:
+ if name_lower == info.short_key.lower():
+ matches.append(info)
+ disp = info.display_name.lower()
+ if name_lower == disp or name_lower == ("nvidia/" + disp):
+ matches.append(info)
+ if len(matches) == 1:
+ return matches[0].short_key
+ if len(matches) > 1:
+ return matches[0].short_key
+
+ # Parsed full form: Family-Skeleton-Dataset or Family-Skeleton-Dataset-vN
+ m = _RESOLVE_FULL_PATTERN.match(name)
+ if m:
+ fam_raw, skel_raw, ds_raw, ver_num = m.groups()
+ fam = _normalize_family(fam_raw)
+ skel = _normalize_skeleton(skel_raw)
+ ds = _normalize_dataset(ds_raw)
+ if fam is not None and skel is not None and ds is not None:
+ if ver_num is not None:
+ version = f"v{ver_num}"
+ for info in MODEL_INFOS:
+ if info.family == fam and info.skeleton == skel and info.dataset == ds and info.version == version:
+ return info.short_key
+ else:
+ info = _get_latest_for_family_skeleton_dataset(fam, skel, ds)
+ if info is not None:
+ return info.short_key
+
+ # Parsed partial: Skeleton-Dataset, Skeleton, or Dataset (use default_family)
+ if default_family is not None:
+ m = _RESOLVE_PARTIAL_PATTERN.match(name)
+ if m:
+ tok1, ds_raw, ver_num = m.groups()
+ fam = _normalize_family(default_family)
+ if fam is not None:
+ skel = _normalize_skeleton(tok1)
+ ds_candidate = _normalize_dataset(ds_raw) if ds_raw else None
+ if skel is not None and ds_candidate is not None:
+ ds = ds_candidate
+ elif skel is not None:
+ ds = "RP"
+ else:
+ skel = "SOMA"
+ ds = _normalize_dataset(tok1) if tok1 else "RP"
+ if ds is None:
+ ds = "RP"
+ if ver_num is not None:
+ version = f"v{ver_num}"
+ for info in MODEL_INFOS:
+ if (
+ info.family == fam
+ and info.skeleton == skel
+ and info.dataset == ds
+ and info.version == version
+ ):
+ return info.short_key
+ else:
+ info = _get_latest_for_family_skeleton_dataset(fam, skel, ds)
+ if info is not None:
+ return info.short_key
+
+ # Single token: skeleton or dataset
+ fam = _normalize_family(default_family)
+ if fam is not None:
+ skel = _normalize_skeleton(name)
+ if skel is not None:
+ info = _get_latest_for_family_skeleton_dataset(fam, skel, "RP")
+ if info is not None:
+ return info.short_key
+ ds = _normalize_dataset(name)
+ if ds is not None:
+ info = _get_latest_for_family_skeleton_dataset(fam, "SOMA", ds)
+ if info is not None:
+ return info.short_key
+
+ raise ValueError(
+ f"Model name {name!r} could not be resolved. "
+ f"Use a short key (e.g. {list(MODEL_NAMES.keys())[:3]}...), "
+ "a full name (e.g. Kimodo-SOMA-RP-v1), or a partial (e.g. SOMA-RP, SOMA) "
+ "with default_family set."
+ )
diff --git a/kimodo/model/text_encoder_api.py b/kimodo/model/text_encoder_api.py
new file mode 100644
index 0000000000000000000000000000000000000000..c684589ac200462770d3ec9cd75d6e640fa7a69c
--- /dev/null
+++ b/kimodo/model/text_encoder_api.py
@@ -0,0 +1,73 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Remote text encoder API client (Gradio) for motion generation."""
+
+import logging
+
+import numpy as np
+import torch
+from gradio_client import Client
+
+# Suppress the [httpx] logs (GET requests)
+logging.getLogger("httpx").setLevel(logging.WARNING)
+
+# Suppress internal gradio_client logs
+logging.getLogger("gradio_client").setLevel(logging.WARNING)
+
+
+class TextEncoderAPI:
+ """Text encoder API client for motion generation."""
+
+ def __init__(self, url: str):
+ self.client = Client(url, verbose=False)
+ self.device = "cpu"
+ self.dtype = torch.float
+
+ def _create_np_random_name(self):
+ import uuid
+
+ return str(uuid.uuid4()) + ".npy"
+
+ def to(self, device=None, dtype=None):
+ if device is not None:
+ self.device = device
+ if dtype is not None:
+ self.dtype = dtype
+ return self
+
+ def __call__(self, texts):
+ """Encode text prompts into tensors.
+
+ Args:
+ texts (str | list[str]): text prompts to encode
+
+ Returns:
+ tuple[torch.Tensor, list[int]]: encoded text tensors and their lengths
+ """
+ if isinstance(texts, str):
+ texts = [texts]
+
+ tensors = []
+ lengths = []
+ for text in texts:
+ filename = self._create_np_random_name()
+
+ result = self.client.predict(
+ text=text,
+ filename=filename,
+ api_name="/DemoWrapper",
+ )
+ path = result[0]["value"]
+ tensor = np.load(path)
+ length = tensor.shape[0]
+
+ tensors.append(tensor)
+ lengths.append(length)
+
+ padded_tensor = np.zeros((len(lengths), max(lengths), tensors[0].shape[-1]), dtype=tensors[0].dtype)
+ for idx, (tensor, length) in enumerate(zip(tensors, lengths)):
+ padded_tensor[idx, :length] = tensor
+
+ padded_tensor = torch.from_numpy(padded_tensor)
+ padded_tensor = padded_tensor.to(device=self.device, dtype=self.dtype)
+ return padded_tensor, lengths
diff --git a/kimodo/model/tmr.py b/kimodo/model/tmr.py
new file mode 100644
index 0000000000000000000000000000000000000000..442411ea6e7afdd8da1852dbd4edb71037bf0f50
--- /dev/null
+++ b/kimodo/model/tmr.py
@@ -0,0 +1,383 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""TMR model: encoder, and text-to-motion retrieval head."""
+
+import contextlib
+from pathlib import Path
+from typing import Dict, List, Optional, Tuple
+
+import torch
+import torch.nn as nn
+from einops import repeat
+from torch import Tensor
+
+from kimodo.model import load_checkpoint_state_dict
+from kimodo.motion_rep.feature_utils import length_to_mask
+from kimodo.sanitize import sanitize_texts
+from kimodo.skeleton import SkeletonBase, build_skeleton
+from kimodo.tools import ensure_batched
+
+
+class PositionalEncoding(nn.Module):
+ """Sinusoidal positional encoding for sequences (batch_first optional)."""
+
+ def __init__(self, d_model, dropout=0.1, max_len=5000, batch_first=False) -> None:
+ super().__init__()
+ self.batch_first = batch_first
+
+ self.dropout = nn.Dropout(p=dropout)
+
+ pe = torch.zeros(max_len, d_model)
+ position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
+ # Note: have to replace torch.exp() and math.log() with torch.pow()
+ # due to MKL exp() and ln() throws floating point exceptions on certain CPUs
+ div_term = torch.pow(10000.0, -torch.arange(0, d_model, 2).float() / d_model)
+ # div_term = torch.exp(
+ # torch.arange(0, d_model, 2).float() * (-np.log(10000.0) / d_model)
+ # )
+
+ pe[:, 0::2] = torch.sin(position * div_term)
+ pe[:, 1::2] = torch.cos(position * div_term)
+ pe = pe.unsqueeze(0).transpose(0, 1)
+ self.register_buffer("pe", pe, persistent=False)
+
+ def forward(self, x: Tensor) -> Tensor:
+ if self.batch_first:
+ x = x + self.pe.permute(1, 0, 2)[:, : x.shape[1], :]
+ else:
+ x = x + self.pe[: x.shape[0], :]
+ return self.dropout(x)
+
+
+def load_ckpt(self, ckpt_path):
+ """Load model weights from checkpoint path."""
+ state_dict = load_checkpoint_state_dict(ckpt_path)
+ self.load_state_dict(state_dict)
+
+
+class ACTORStyleEncoder(nn.Module):
+ """Motion encoder in ACTOR style: optional motion_rep projection, VAE/MLP tokens, transformer."""
+
+ def __init__(
+ self,
+ motion_rep: Optional[nn.Module],
+ llm_shape: Optional[Tuple],
+ vae: bool,
+ latent_dim: int = 256,
+ ff_size: int = 1024,
+ num_layers: int = 4,
+ num_heads: int = 4,
+ dropout: float = 0.1,
+ activation: str = "gelu",
+ ckpt_path: Optional[str] = None,
+ ) -> None:
+ super().__init__()
+
+ self.motion_rep = motion_rep
+ if motion_rep is not None and llm_shape is None:
+ nfeats = motion_rep.motion_rep_dim
+ elif motion_rep is None and llm_shape is not None:
+ nfeats = llm_shape[-1]
+ else:
+ raise ValueError
+
+ self.nfeats = nfeats
+ self.projection = nn.Linear(nfeats, latent_dim)
+
+ self.vae = vae
+ self.nbtokens = 2 if vae else 1
+ self.tokens = nn.Parameter(torch.randn(self.nbtokens, latent_dim))
+
+ self.sequence_pos_encoding = PositionalEncoding(latent_dim, dropout=dropout, batch_first=True)
+
+ seq_trans_encoder_layer = nn.TransformerEncoderLayer(
+ d_model=latent_dim,
+ nhead=num_heads,
+ dim_feedforward=ff_size,
+ dropout=dropout,
+ activation=activation,
+ batch_first=True,
+ )
+
+ self.seqTransEncoder = nn.TransformerEncoder(
+ seq_trans_encoder_layer,
+ num_layers=num_layers,
+ enable_nested_tensor=False,
+ )
+
+ if ckpt_path is not None:
+ load_ckpt(self, ckpt_path)
+
+ def forward(self, x_dict: Dict) -> Tensor:
+ x = x_dict["x"]
+ mask = x_dict["mask"]
+
+ x = self.projection(x)
+
+ device = x.device
+ bs = len(x)
+
+ tokens = repeat(self.tokens, "nbtoken dim -> bs nbtoken dim", bs=bs)
+ xseq = torch.cat((tokens, x), 1)
+
+ token_mask = torch.ones((bs, self.nbtokens), dtype=bool, device=device)
+ aug_mask = torch.cat((token_mask, mask), 1)
+
+ # add positional encoding
+ xseq = self.sequence_pos_encoding(xseq)
+ final = self.seqTransEncoder(xseq, src_key_padding_mask=~aug_mask)
+ return final[:, : self.nbtokens]
+
+
+class TMR(nn.Module):
+ r"""TMR: Text-to-Motion Retrieval inference code (no decoder)
+ Find more information about the model on the following website:
+ https://mathis.petrovich.fr/tmr
+ """
+
+ @classmethod
+ def from_args(
+ cls,
+ motion_rep: nn.Module,
+ llm_shape: tuple | list,
+ vae: bool,
+ latent_dim: int = 256,
+ ff_size: int = 1024,
+ num_layers: int = 4,
+ num_heads: int = 4,
+ dropout: float = 0.1,
+ activation: str = "gelu",
+ ckpt_folder: Optional[str] = None,
+ device: Optional[str] = None,
+ **kwargs,
+ ):
+ motion_encoder, top_text_encoder = None, None
+
+ motion_encoder = ACTORStyleEncoder(
+ motion_rep=motion_rep,
+ llm_shape=None,
+ vae=vae,
+ latent_dim=latent_dim,
+ ff_size=ff_size,
+ num_layers=num_layers,
+ num_heads=num_heads,
+ dropout=dropout,
+ activation=activation,
+ ckpt_path=Path(ckpt_folder) / "motion_encoder.pt",
+ ).to(device)
+
+ top_text_encoder = ACTORStyleEncoder(
+ motion_rep=None,
+ llm_shape=llm_shape,
+ vae=vae,
+ latent_dim=latent_dim,
+ ff_size=ff_size,
+ num_layers=num_layers,
+ num_heads=num_heads,
+ dropout=dropout,
+ activation=activation,
+ ckpt_path=Path(ckpt_folder) / "text_encoder.pt",
+ ).to(device)
+ return cls(
+ motion_encoder,
+ top_text_encoder,
+ vae,
+ device=device,
+ **kwargs,
+ )
+
+ def __init__(
+ self,
+ motion_encoder: nn.Module,
+ top_text_encoder: nn.Module,
+ vae: bool,
+ text_encoder: Optional = None,
+ fact: Optional[float] = None,
+ sample_mean: Optional[bool] = True,
+ unit_vector: Optional[bool] = False,
+ compute_grads: bool = False,
+ device: Optional[str] = None,
+ ) -> None:
+ super().__init__()
+
+ self.motion_encoder = motion_encoder
+ self.text_encoder = top_text_encoder
+ self.raw_text_encoder = text_encoder
+
+ self.motion_rep = None
+ self.skeleton = None
+ if self.motion_encoder is not None:
+ self.motion_rep = self.motion_encoder.motion_rep
+ if self.motion_rep is not None:
+ self.skeleton = self.motion_rep.skeleton
+
+ self.compute_grads = compute_grads
+
+ self.device = device
+
+ # sampling parameters
+ self.vae = vae
+ self.fact = fact if fact is not None else 1.0
+ self.sample_mean = sample_mean
+ self.unit_vector = unit_vector
+
+ def full_text_encoder(self, texts: list[str]):
+ assert isinstance(texts, list), "The input should be batched."
+ # sanitize the texts first
+ # then encode the text, and then use the top text encoder
+ texts = sanitize_texts(texts)
+ text_feat, text_length = self.raw_text_encoder(texts)
+ if isinstance(text_length, list):
+ text_length = torch.tensor(text_length, device=self.device)
+ else:
+ text_length = text_length.to(self.device)
+ inputs = {
+ "x": text_feat.to(self.device),
+ "mask": length_to_mask(text_length, device=self.device),
+ }
+ return self.text_encoder(inputs)
+
+ def _find_encoder(self, inputs, modality):
+ assert modality in ["text", "motion", "raw_text", "auto"]
+
+ if modality == "text":
+ return self.text_encoder
+ elif modality == "motion":
+ return self.motion_encoder
+ elif modality == "raw_text":
+ return self.full_text_encoder
+
+ if isinstance(inputs[0], str):
+ return self.full_text_encoder
+
+ m_nfeats = self.motion_encoder.nfeats
+ t_nfeats = self.text_encoder.nfeats
+
+ if m_nfeats == t_nfeats:
+ raise ValueError("Cannot automatically find the encoder, as they share the same input space.")
+
+ nfeats = inputs["x"].shape[-1]
+ if nfeats == m_nfeats:
+ return self.motion_encoder
+ elif nfeats == t_nfeats:
+ return self.text_encoder
+ else:
+ raise ValueError("The inputs is not recognized.")
+
+ def _encode(
+ self,
+ inputs,
+ modality: str = "auto",
+ sample_mean: Optional[bool] = None,
+ fact: Optional[float] = None,
+ return_distribution: bool = False,
+ unit_vector: Optional[bool] = None,
+ ):
+ sample_mean = self.sample_mean if sample_mean is None else sample_mean
+ fact = self.fact if fact is None else fact
+ unit_vector = self.unit_vector if unit_vector is None else unit_vector
+
+ # Encode the inputs
+ encoder = self._find_encoder(inputs, modality)
+ encoded = encoder(inputs)
+
+ # Sampling
+ if self.vae:
+ dists = encoded.unbind(1)
+ mu, logvar = dists
+ if sample_mean:
+ latent_vectors = mu
+ else:
+ # Reparameterization trick
+ std = logvar.exp().pow(0.5)
+ eps = std.data.new(std.size()).normal_()
+ latent_vectors = mu + fact * eps * std
+ else:
+ dists = None
+ (latent_vectors,) = encoded.unbind(1)
+
+ if unit_vector:
+ latent_vectors = torch.nn.functional.normalize(latent_vectors, dim=-1)
+
+ if return_distribution:
+ return latent_vectors, dists
+
+ return latent_vectors
+
+ @ensure_batched(posed_joints=4, lengths=1)
+ def encode_motion(
+ self,
+ posed_joints: torch.Tensor,
+ original_skeleton: Optional[SkeletonBase] = None,
+ lengths: Optional[torch.Tensor] = None,
+ unit_vector: Optional[bool] = None,
+ ):
+ # TODO here.
+ convert_ctx = torch.no_grad() if not self.compute_grads else contextlib.nullcontext()
+
+ if original_skeleton is None:
+ original_skeleton = build_skeleton(posed_joints.shape[-2])
+
+ if lengths is None:
+ nbatch, nbframes = posed_joints.shape[:2]
+ device = posed_joints.device
+ assert nbatch == 1, "If lenghts is not provided, the input should not be batched."
+ lengths = torch.tensor([nbframes], device=device)
+
+ # slice the posed joints if we use less joints
+ skel_slice = self.motion_rep.skeleton.get_skel_slice(original_skeleton)
+ posed_joints = posed_joints[..., skel_slice, :]
+
+ with convert_ctx:
+ features = self.motion_rep(
+ posed_joints=posed_joints,
+ to_canonicalize=True,
+ to_normalize=True,
+ lengths=lengths,
+ )
+ mask = length_to_mask(lengths, device=features.device)
+ x_dict = {"x": features, "mask": mask}
+ latent_vectors = self._encode(
+ x_dict,
+ modality="motion",
+ unit_vector=unit_vector,
+ )
+ return latent_vectors
+
+ def encode_text(
+ self,
+ x_dict: Dict,
+ unit_vector: Optional[bool] = None,
+ ):
+ # TODO: make it ensure batched
+ convert_ctx = torch.no_grad() if not self.compute_grads else contextlib.nullcontext()
+
+ with convert_ctx:
+ latent_vectors = self._encode(
+ x_dict,
+ modality="text",
+ unit_vector=unit_vector,
+ )
+ return latent_vectors
+
+ def encode_raw_text(
+ self,
+ texts: List[str],
+ unit_vector: Optional[bool] = None,
+ ):
+ is_batched = True
+ if isinstance(texts, str):
+ is_batched = False
+ texts = [texts]
+
+ convert_ctx = torch.no_grad() if not self.compute_grads else contextlib.nullcontext()
+
+ with convert_ctx:
+ latent_vectors = self._encode(
+ texts,
+ modality="raw_text",
+ unit_vector=unit_vector,
+ )
+ if not is_batched:
+ latent_vectors = latent_vectors[0]
+ return latent_vectors
diff --git a/kimodo/model/twostage_denoiser.py b/kimodo/model/twostage_denoiser.py
new file mode 100644
index 0000000000000000000000000000000000000000..d14cf76085cfc1dbba486cb6470644bade4038d0
--- /dev/null
+++ b/kimodo/model/twostage_denoiser.py
@@ -0,0 +1,153 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Two-stage transformer denoiser: root stage then body stage for motion diffusion."""
+
+import contextlib
+from typing import Optional
+
+import torch
+from torch import nn
+
+from .backbone import TransformerEncoderBlock
+from .loading import load_checkpoint_state_dict
+
+
+class TwostageDenoiser(nn.Module):
+ """Two-stage denoiser: first predicts global root features, then body features conditioned on local root."""
+
+ def __init__(
+ self,
+ motion_rep,
+ motion_mask_mode,
+ ckpt_path: Optional[str] = None,
+ **kwargs,
+ ):
+ """Build root and body transformer blocks; optionally load checkpoint from ckpt_path."""
+ super().__init__()
+ self.motion_rep = motion_rep
+ self.motion_mask_mode = motion_mask_mode
+
+ # it should be a dual motion_rep
+ # and be global by default
+ # global motion_rep as inpnut
+ input_dim = motion_rep.motion_rep_dim
+ will_concatenate = motion_mask_mode == "concat"
+
+ # stage 1: root only
+ root_input_dim = input_dim * 2 if will_concatenate else input_dim
+ root_output_dim = motion_rep.global_root_dim
+
+ self.root_model = TransformerEncoderBlock(
+ input_dim=root_input_dim,
+ output_dim=root_output_dim,
+ skeleton=self.motion_rep.skeleton,
+ **kwargs,
+ )
+
+ # replace the global root by the local root
+ local_motion_rep_dim = input_dim - motion_rep.global_root_dim + motion_rep.local_root_dim
+
+ # stage 2: local body
+ body_input_dim = local_motion_rep_dim + (
+ input_dim if will_concatenate else 0
+ ) # body stage always takes in local root info for motion (but still the global mask)
+
+ body_output_dim = input_dim - motion_rep.global_root_dim
+ self.body_model = TransformerEncoderBlock(
+ input_dim=body_input_dim,
+ output_dim=body_output_dim,
+ skeleton=self.motion_rep.skeleton,
+ **kwargs,
+ )
+
+ if ckpt_path:
+ self.load_ckpt(ckpt_path)
+
+ def load_ckpt(self, ckpt_path: str) -> None:
+ """Load checkpoint from path; state dict keys are stripped of 'denoiser.backbone.'
+ prefix."""
+ state_dict = load_checkpoint_state_dict(ckpt_path)
+ state_dict = {key.replace("denoiser.backbone.", ""): val for key, val in state_dict.items()}
+ self.load_state_dict(state_dict)
+
+ def forward(
+ self,
+ x: torch.Tensor,
+ x_pad_mask: torch.Tensor,
+ text_feat: torch.Tensor,
+ text_feat_pad_mask: torch.Tensor,
+ timesteps: torch.Tensor,
+ first_heading_angle: Optional[torch.Tensor] = None,
+ motion_mask: Optional[torch.Tensor] = None,
+ observed_motion: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ """
+ Args:
+ x (torch.Tensor): [B, T, dim_motion] current noisy motion
+ x_pad_mask (torch.Tensor): [B, T] attention mask, positions with True are allowed to attend, False are not
+ text_feat (torch.Tensor): [B, max_text_len, llm_dim] embedded text prompts
+ text_feat_pad_mask (torch.Tensor): [B, max_text_len] attention mask, positions with True are allowed to attend, False are not
+ timesteps (torch.Tensor): [B,] current denoising step
+ motion_mask
+ observed_motion
+
+ Returns:
+ torch.Tensor: same size as input x
+ """
+
+ if self.motion_mask_mode == "concat":
+ if motion_mask is None or observed_motion is None:
+ motion_mask = torch.zeros_like(x)
+ observed_motion = torch.zeros_like(x)
+ x = x * (1 - motion_mask) + observed_motion * motion_mask
+ x_extended = torch.cat([x, motion_mask], axis=-1)
+ else:
+ x_extended = x
+
+ # Stage 1: predict root motion in global
+ root_motion_pred = self.root_model(
+ x_extended,
+ x_pad_mask,
+ text_feat,
+ text_feat_pad_mask,
+ timesteps,
+ first_heading_angle,
+ ) # [B, T, 5]
+
+ # Maybe pass this as argument instead of recomputing it
+ lengths = x_pad_mask.sum(-1)
+
+ # Convert root pred to local rep
+ # At test-time want to allow gradient through for guidance
+ convert_ctx = torch.no_grad() if self.training else contextlib.nullcontext()
+ with convert_ctx:
+ root_motion_local = self.motion_rep.global_root_to_local_root(
+ root_motion_pred,
+ normalized=True,
+ lengths=lengths,
+ )
+ if self.training:
+ root_motion_local = root_motion_local.detach()
+
+ # concatenate the predicted local root with the body motion
+ body_x = x[..., self.motion_rep.body_slice]
+ x_new = torch.cat([root_motion_local, body_x], axis=-1)
+
+ if self.motion_mask_mode == "concat":
+ x_new_extended = torch.cat([x_new, motion_mask], axis=-1)
+ else:
+ x_new_extended = x_new
+
+ # Stage 2: predict local body motion based on local root
+ predicted_body = self.body_model(
+ x_new_extended,
+ x_pad_mask,
+ text_feat,
+ text_feat_pad_mask,
+ timesteps,
+ first_heading_angle,
+ )
+
+ # concatenate the predicted local body with the predicted root
+ output = torch.cat([root_motion_pred, predicted_body], axis=-1)
+ return output
diff --git a/kimodo/motion_rep/__init__.py b/kimodo/motion_rep/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..721e0392fe88404410942e80d7ad120d1a91a88b
--- /dev/null
+++ b/kimodo/motion_rep/__init__.py
@@ -0,0 +1,11 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Motion representation utilities."""
+
+from .reps import KimodoMotionRep, MotionRepBase, TMRMotionRep
+
+__all__ = [
+ "MotionRepBase",
+ "KimodoMotionRep",
+ "TMRMotionRep",
+]
diff --git a/kimodo/motion_rep/conditioning.py b/kimodo/motion_rep/conditioning.py
new file mode 100644
index 0000000000000000000000000000000000000000..c7c64da24dd143b7684b16a96429ab41ed15db3f
--- /dev/null
+++ b/kimodo/motion_rep/conditioning.py
@@ -0,0 +1,28 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Constraint conditioning: build index and data dicts from constraint sets for the denoiser."""
+
+from collections import defaultdict
+
+import torch
+
+
+def build_condition_dicts(constraints_lst: list):
+ index_dict = defaultdict(list)
+ data_dict = defaultdict(list)
+ for constraint in constraints_lst:
+ constraint.update_constraints(data_dict, index_dict)
+ return index_dict, data_dict
+
+
+def get_unique_index_and_data(indices_lst, data):
+ # unique + sort them by t
+ indices_unique, inverse = torch.unique(indices_lst, dim=0, return_inverse=True)
+ # pick first value for each unique (t, j)
+ first_idx = torch.zeros(indices_unique.size(0), dtype=torch.long, device=inverse.device)
+ first_idx.scatter_(0, inverse, torch.arange(len(inverse), device=inverse.device))
+ assert (indices_lst[first_idx] == indices_unique).all()
+ # get the data
+ indices_lst = indices_lst[first_idx]
+ data = data[first_idx]
+ return indices_lst, data
diff --git a/kimodo/motion_rep/feature_utils.py b/kimodo/motion_rep/feature_utils.py
new file mode 100644
index 0000000000000000000000000000000000000000..fd204360ddfdab1cf9738ae2a96744f137baecbf
--- /dev/null
+++ b/kimodo/motion_rep/feature_utils.py
@@ -0,0 +1,212 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Motion representation helpers: velocity, heading, masks, and rotation of features."""
+
+from typing import List, Optional, Union
+
+import einops
+import torch
+
+from kimodo.geometry import cont6d_to_matrix, matrix_to_cont6d
+from kimodo.skeleton import SkeletonBase
+from kimodo.tools import ensure_batched
+
+
+def diff_angles(angles: torch.Tensor, fps: float) -> torch.Tensor:
+ """Compute frame-to-frame angular differences in radians, scaled by fps.
+
+ Args:
+ angles: [..., T] batched sequences of rotation angles in radians.
+ fps: Sampling rate used to convert frame differences to per-second rate.
+
+ Returns:
+ [..., T-1] difference between consecutive angles (rad/s).
+ """
+
+ cos = torch.cos(angles)
+ sin = torch.sin(angles)
+
+ cos_diff = cos[..., 1:] * cos[..., :-1] + sin[..., 1:] * sin[..., :-1]
+ sin_diff = sin[..., 1:] * cos[..., :-1] - cos[..., 1:] * sin[..., :-1]
+
+ # should be close to angles.diff() but more robust
+ # multiply by fps = 1 / dt
+ angles_diff = fps * torch.arctan2(sin_diff, cos_diff)
+ return angles_diff
+
+
+@ensure_batched(positions=4, lengths=1)
+def compute_vel_xyz(
+ positions: torch.Tensor,
+ fps: float,
+ lengths: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ """Compute the velocities from positions: dx/dt. Works with batches. The last velocity is duplicated to keep the same size.
+
+ Args:
+ positions (torch.Tensor): [..., T, J, 3] xyz positions of a human skeleton
+ fps (float): frame per seconds
+ lengths (Optional[torch.Tensor]): [...] size of each input batched. If not provided, positions should not be batched
+
+ Returns:
+ velocity (torch.Tensor): [..., T, J, 3] velocities computed from the positions
+ """
+ device = positions.device
+
+ if lengths is None:
+ assert positions.shape[0] == 1, "If lengths is not provided, the input should not be batched."
+ lengths = torch.tensor([len(positions)], device=device)
+
+ # useful for indexing
+ range_len = torch.arange(len(lengths))
+
+ # compute velocities with fps
+ velocity = fps * (positions[:, 1:] - positions[:, :-1])
+ # pading the velocity vector
+ vel_pad = torch.zeros_like(velocity[:, 0])
+ velocity, _ = einops.pack([velocity, vel_pad], "batch * nbjoints dim")
+
+ # repeat the last velocities
+ # with special care for different lengths with batches
+ velocity[(range_len, lengths - 1)] = velocity[(range_len, lengths - 2)]
+ return velocity
+
+
+@ensure_batched(root_rot_angles=2, lengths=1)
+def compute_vel_angle(
+ root_rot_angles: torch.Tensor,
+ fps: float,
+ lengths: Optional[torch.Tensor] = None,
+) -> torch.Tensor:
+ """Compute the local root rotation velocity: dtheta/dt.
+
+ Args:
+ root_rot_angles (torch.Tensor): [..., T] rotation angle (in radian)
+ fps (float): frame per seconds
+ lengths (Optional[torch.Tensor]): [...] size of each input batched. If not provided, root_rot_angles should not be batched
+
+ Returns:
+ local_root_rot_vel (torch.Tensor): [..., T] local root rotation velocity (in radian/s)
+ """
+ device = root_rot_angles.device
+ if lengths is None:
+ assert root_rot_angles.shape[0] == 1, "If lengths is not provided, the input should not be batched."
+ lengths = torch.tensor([len(root_rot_angles)], device=device)
+
+ # useful for indexing
+ range_len = torch.arange(len(lengths))
+
+ local_root_rot_vel = diff_angles(root_rot_angles, fps)
+ pad_rot_vel_angles = torch.zeros_like(root_rot_angles[:, 0])
+ local_root_rot_vel, _ = einops.pack(
+ [local_root_rot_vel, pad_rot_vel_angles],
+ "batch *",
+ )
+ # repeat the last rotation angle
+ # with special care for different lengths with batches
+ local_root_rot_vel[(range_len, lengths - 1)] = local_root_rot_vel[(range_len, lengths - 2)]
+ return local_root_rot_vel
+
+
+@ensure_batched(posed_joints=4)
+def compute_heading_angle(posed_joints: torch.Tensor, skeleton: SkeletonBase) -> torch.Tensor:
+ """Compute the heading direction from joint positions using the hip vector.
+
+ Args:
+ posed_joints: [B, T, J, 3] global joint positions.
+ skeleton: Skeleton instance used to get hip joint indices.
+
+ Returns:
+ [B] heading angle in radians.
+ """
+ # compute root heading for the sequence from hip positions
+ r_hip, l_hip = skeleton.hip_joint_idx
+ diff = posed_joints[:, :, r_hip] - posed_joints[:, :, l_hip]
+ heading_angle = torch.atan2(diff[..., 2], -diff[..., 0])
+ return heading_angle
+
+
+def length_to_mask(
+ length: Union[torch.Tensor, List],
+ max_len: Optional[int] = None,
+ device=None,
+) -> torch.Tensor:
+ """Convert sequence lengths to a boolean validity mask.
+
+ Args:
+ length: Sequence lengths, either a tensor ``[B]`` or a Python list.
+ max_len: Optional mask width. If omitted, uses ``max(length)``.
+ device: Optional device. When ``length`` is a list, this controls where
+ the new tensor is created.
+
+ Returns:
+ A boolean tensor of shape ``[B, max_len]`` where ``True`` marks valid
+ timesteps.
+ """
+ if isinstance(length, list):
+ if device is None:
+ device = "cpu"
+ length = torch.tensor(length, device=device)
+
+ # Use requested device for output; move length if needed so mask and length match
+ if device is not None:
+ target = torch.device(device)
+ if length.device != target:
+ length = length.to(target)
+ device = length.device
+
+ if max_len is None:
+ max_len = max(length)
+
+ mask = torch.arange(max_len, device=device).expand(len(length), max_len) < length.unsqueeze(1)
+ return mask
+
+
+class RotateFeatures:
+ """Helper that applies a global heading rotation to motion features."""
+
+ def __init__(self, angle: torch.Tensor):
+ """Precompute 2D and 3D rotation matrices for a batch of angles.
+
+ Args:
+ angle: Rotation angle(s) in radians, shaped ``[B]``.
+ """
+ self.angle = angle
+
+ ## Create the necessary rotations matrices
+ cos, sin = torch.cos(angle), torch.sin(angle)
+ one, zero = torch.ones_like(angle), torch.zeros_like(angle)
+
+ # 2D rotation transposed (sin are -sin)
+ self.corrective_mat_2d_T = torch.stack((cos, sin, -sin, cos), -1).reshape(angle.shape + (2, 2))
+ # 3D rotation on Y axis
+ self.corrective_mat_Y = torch.stack((cos, zero, sin, zero, one, zero, -sin, zero, cos), -1).reshape(
+ angle.shape + (3, 3)
+ )
+ self.corrective_mat_Y_T = self.corrective_mat_Y.transpose(-2, -1).contiguous()
+
+ def rotate_positions(self, positions: torch.Tensor):
+ """Rotate 3D positions around the Y axis."""
+ return positions @ self.corrective_mat_Y_T
+
+ def rotate_2d_positions(self, positions_2d: torch.Tensor):
+ """Rotate 2D ``(x, z)`` vectors in the ground plane."""
+ return positions_2d @ self.corrective_mat_2d_T
+
+ def rotate_rotations(self, rotations: torch.Tensor):
+ """Left-multiply global rotation matrices by the heading correction."""
+ # "Rotate" the global rotations
+ # which means add an extra Y rotation after the transform
+ # so at the left R' = R_y R
+ # (since we use the convention x' = R x)
+ # "bik,btdkj->btdij"
+
+ B, T, J = rotations.shape[:3]
+ BTJ = B * T * J
+ return (
+ self.corrective_mat_Y[:, None, None].expand(B, T, J, 3, 3).reshape(BTJ, 3, 3) @ rotations.reshape(BTJ, 3, 3)
+ ).reshape(B, T, J, 3, 3)
+
+ def rotate_6d_rotations(self, rotations_6d: torch.Tensor):
+ """Rotate 6D rotation features via matrix conversion."""
+ return matrix_to_cont6d(self.rotate_rotations(cont6d_to_matrix(rotations_6d)))
diff --git a/kimodo/motion_rep/feet.py b/kimodo/motion_rep/feet.py
new file mode 100644
index 0000000000000000000000000000000000000000..89e6bdb155fb7e5e6be219c709a5856702515600
--- /dev/null
+++ b/kimodo/motion_rep/feet.py
@@ -0,0 +1,60 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Foot contact detection from joint positions and velocities."""
+
+import torch
+
+from ..tools import ensure_batched
+
+
+@ensure_batched(positions=4, velocity=4)
+def foot_detect_from_pos_and_vel(
+ positions: torch.Tensor,
+ velocity: torch.Tensor,
+ skeleton,
+ vel_thres: float,
+ height_thresh: float,
+) -> torch.Tensor:
+ """Compute foot contact labels using heuristics combining joint height and velocities.
+
+ Args:
+ positions (torch.Tensor): [X, T, J, 3] global joint positions
+ velocity (torch.Tensor): [X, T, J, 3] velocities (already padded correctly), already multiplied by 1 / dt
+ vel_thres (float): threshold for joint velocity
+ height_thresh (float): threshold for joint height
+
+ Returns:
+ torch.Tensor: [X, T, 4] contact labels for left and right foot joints
+ (heel/toe order follows the skeleton joint index definition), where
+ ``1`` denotes contact.
+ """
+
+ device = positions.device
+ # Use at most 2 foot joints per side (ankle + toe); SOMA77 defines a
+ # third end-effector (ToeEnd) that SOMA30 and other skeletons omit.
+ fid_l = skeleton.left_foot_joint_idx[:2]
+ fid_r = skeleton.right_foot_joint_idx[:2]
+
+ velfactor, heightfactor = (
+ torch.tensor([vel_thres, vel_thres], device=device),
+ torch.tensor([height_thresh, height_thresh], device=device),
+ )
+
+ feet_l_v = torch.linalg.norm(velocity[:, :, fid_l], axis=-1)
+ feet_l_h = positions[:, :, fid_l, 1]
+
+ feet_l = torch.logical_and(
+ feet_l_v < velfactor,
+ feet_l_h < heightfactor,
+ ).to(positions.dtype)
+
+ feet_r_v = torch.linalg.norm(velocity[:, :, fid_r], axis=-1)
+ feet_r_h = positions[:, :, fid_r, 1]
+
+ feet_r = torch.logical_and(
+ feet_r_v < velfactor,
+ feet_r_h < heightfactor,
+ ).to(positions.dtype)
+
+ foot_contacts = torch.cat((feet_l, feet_r), axis=-1)
+ return foot_contacts
diff --git a/kimodo/motion_rep/reps/__init__.py b/kimodo/motion_rep/reps/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..3bd6ac79c5f0392c501a99018223fe6095a42c19
--- /dev/null
+++ b/kimodo/motion_rep/reps/__init__.py
@@ -0,0 +1,13 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Motion representation implementations: base, Kimodo, and TMR."""
+
+from .base import MotionRepBase
+from .kimodo_motionrep import KimodoMotionRep
+from .tmr_motionrep import TMRMotionRep
+
+__all__ = [
+ "MotionRepBase",
+ "KimodoMotionRep",
+ "TMRMotionRep",
+]
diff --git a/kimodo/motion_rep/reps/base.py b/kimodo/motion_rep/reps/base.py
new file mode 100644
index 0000000000000000000000000000000000000000..6911cead058d7eb21da4675156d6b49325d30758
--- /dev/null
+++ b/kimodo/motion_rep/reps/base.py
@@ -0,0 +1,299 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Base motion representation: feature layout, normalization, and conditioning helpers."""
+
+import os
+from typing import Optional
+
+import einops
+import numpy as np
+import torch
+from einops import repeat
+
+from ...tools import ensure_batched
+from ..conditioning import build_condition_dicts
+from ..feature_utils import compute_vel_angle, compute_vel_xyz
+from ..stats import Stats
+
+
+def _require_split_stats_layout(stats_path: str) -> None:
+ """Raise if stats_path does not contain the required global_root, local_root, body subdirs."""
+ subdirs = ("global_root", "local_root", "body")
+ missing = []
+ for name in subdirs:
+ subpath = os.path.join(stats_path, name)
+ mean_path = os.path.join(subpath, "mean.npy")
+ if not os.path.isfile(mean_path):
+ missing.append(f"{subpath}/ (mean.npy)")
+ if missing:
+ raise FileNotFoundError(
+ f"Checkpoint stats must use the split layout with subfolders "
+ f"global_root/, local_root/, and body/ under '{stats_path}'. "
+ f"Missing or incomplete: {', '.join(missing)}. "
+ )
+
+
+class MotionRepBase:
+ """Base class for motion representations used in generation and conditioning.
+
+ Subclasses define:
+ - ``size_dict``: feature blocks and their shapes,
+ - ``last_root_feature``: last entry of the root block,
+ - ``local_root_size_dict``: local-root feature layout,
+ and implement transform-specific methods such as ``__call__``, ``inverse``,
+ ``rotate``, ``translate_2d`` and ``create_conditions``.
+ """
+
+ def __init__(
+ self,
+ skeleton,
+ fps,
+ stats_path: Optional[str] = None,
+ ):
+ """Initialize feature slicing metadata and optional normalization stats."""
+
+ self.skeleton = skeleton
+ self.fps = fps
+ self.nbjoints = skeleton.nbjoints
+
+ self.feature_names = list(self.size_dict.keys())
+ self.ps = list(self.size_dict.values())
+ self.nfeats_dict = {key: val.numel() for key, val in self.size_dict.items()}
+ feats_cumsum = np.cumsum([0] + list(self.nfeats_dict.values())).tolist()
+ self.slice_dict = {key: slice(feats_cumsum[i], feats_cumsum[i + 1]) for i, key in enumerate(self.feature_names)}
+
+ self.motion_rep_dim = sum(self.nfeats_dict.values())
+ self.root_slice = slice(0, self.slice_dict[self.last_root_feature].stop)
+ self.body_slice = slice(self.root_slice.stop, self.motion_rep_dim)
+ self.body_dim = self.body_slice.stop - self.body_slice.start
+ self.global_root_dim = self.root_slice.stop
+ self.local_root_dim = sum(val.numel() for val in self.local_root_size_dict.values())
+
+ if stats_path:
+ _require_split_stats_layout(stats_path)
+ self.global_root_stats = Stats(os.path.join(stats_path, "global_root"))
+ self.local_root_stats = Stats(os.path.join(stats_path, "local_root"))
+ self.body_stats = Stats(os.path.join(stats_path, "body"))
+
+ # Global stats
+ mean = torch.cat([self.global_root_stats.mean, self.body_stats.mean])
+ std = torch.cat([self.global_root_stats.std, self.body_stats.std])
+ assert len(mean) == len(std) == self.motion_rep_dim, "There is an stat issue."
+ self.stats = Stats()
+ self.stats.register_from_tensors(mean, std)
+
+ def get_root_pos(self, features: torch.Tensor, fallback_to_smooth: bool = True):
+ """Extract root positions from a feature tensor.
+
+ Supports both ``root_pos`` and ``smooth_root_pos`` representations.
+ """
+ if "root_pos" in self.slice_dict:
+ return features[..., self.slice_dict["root_pos"]]
+
+ if "smooth_root_pos" not in self.slice_dict:
+ raise TypeError("This motion rep should have either a root_pos or smooth_root_pos field")
+
+ if fallback_to_smooth:
+ return features[:, :, self.slice_dict["smooth_root_pos"]]
+
+ # else compute the root pos from the smooth root and local joints offset
+ smooth_root_pos = features[:, :, self.slice_dict["smooth_root_pos"]].clone()
+ local_joints_positions_flatten = features[..., self.slice_dict["local_joints_positions"]]
+ hips_offset = local_joints_positions_flatten[..., self.skeleton.root_idx : self.skeleton.root_idx + 3]
+ root_pos = torch.stack(
+ [
+ smooth_root_pos[..., 0] + hips_offset[..., 0],
+ smooth_root_pos[..., 1],
+ smooth_root_pos[..., 2] + hips_offset[..., 2],
+ ],
+ axis=-1,
+ )
+ return root_pos
+
+ @ensure_batched(root_features=3, lengths=1)
+ def global_root_to_local_root(
+ self,
+ root_features: torch.Tensor,
+ normalized: bool,
+ lengths: Optional[torch.Tensor],
+ ):
+ """Convert global root features to local-root motion features.
+
+ Args:
+ root_features: Root feature tensor containing root position and
+ global heading, shaped ``[B, T, D_root]``.
+ normalized: Whether ``root_features`` are normalized.
+ lengths: Optional valid lengths per sequence.
+
+ Returns:
+ Tensor ``[B, T, 4]`` with local root rotational velocity, planar
+ velocity, and global root height.
+ """
+ if normalized:
+ root_features = self.global_root_stats.unnormalize(root_features)
+
+ [root_pos, global_root_heading] = einops.unpack(root_features, self.ps[:2], "batch time *")
+ cos, sin = global_root_heading.unbind(-1)
+ heading_angle = torch.arctan2(sin, cos)
+
+ local_root_rot_vel = compute_vel_angle(heading_angle, self.fps, lengths=lengths)
+ local_root_vel = compute_vel_xyz(
+ root_pos[..., None, :],
+ self.fps,
+ lengths=lengths,
+ )[..., 0, [0, 2]]
+ global_root_y = root_pos[..., 1]
+ local_root_motion = torch.cat(
+ [
+ local_root_rot_vel[..., None],
+ local_root_vel,
+ global_root_y[..., None],
+ ],
+ axis=-1,
+ )
+
+ if normalized:
+ local_root_motion = self.local_root_stats.normalize(local_root_motion)
+ return local_root_motion
+
+ def get_root_heading_angle(self, features: torch.Tensor) -> torch.Tensor:
+ """Compute root heading angle from cosine/sine heading features."""
+ global_root_heading = features[:, :, self.slice_dict["global_root_heading"]]
+ cos, sin = global_root_heading.unbind(-1)
+ return torch.arctan2(sin, cos)
+
+ @ensure_batched(features=3)
+ def rotate_to(
+ self,
+ features: torch.Tensor,
+ target_angle: torch.Tensor,
+ return_delta_angle=False,
+ ):
+ """Rotate each sequence so frame-0 heading matches ``target_angle``."""
+ # rotate so that the first frame angle is the target
+ # it put the motion_rep to the angle
+ current_first_angle = self.get_root_heading_angle(features)[:, 0]
+ delta_angle = target_angle - current_first_angle
+ rotated_features = self.rotate(features, delta_angle)
+ if return_delta_angle:
+ return rotated_features, delta_angle
+ return rotated_features
+
+ @ensure_batched(features=3)
+ def rotate_to_zero(
+ self,
+ features: torch.Tensor,
+ return_delta_angle=False,
+ ):
+ """Rotate each sequence so frame-0 heading becomes zero."""
+ target_angle = torch.zeros(len(features), device=features.device)
+ return self.rotate_to(features, target_angle, return_delta_angle=return_delta_angle)
+
+ @ensure_batched(features=3)
+ def randomize_first_heading(
+ self,
+ features: torch.Tensor,
+ return_delta_angle=False,
+ ) -> torch.Tensor:
+ """Rotate each sequence to a random frame-0 heading."""
+ target_heading_angle = torch.rand(features.shape[0]) * 2 * np.pi
+ return self.rotate_to(
+ features,
+ target_heading_angle,
+ return_delta_angle=return_delta_angle,
+ )
+
+ @ensure_batched(features=3, target_2d_pos=2)
+ def translate_2d_to(
+ self,
+ features: torch.Tensor,
+ target_2d_pos: torch.Tensor,
+ return_delta_pos: bool = False,
+ ) -> torch.Tensor:
+ """Translate each sequence so frame-0 root ``(x, z)`` matches a target."""
+ root_pos = self.get_root_pos(features)
+ current_first_2d_pos = root_pos[:, 0, [0, 2]].clone()
+ delta_2d_pos = target_2d_pos - current_first_2d_pos
+ translated_features = self.translate_2d(features, delta_2d_pos)
+ if return_delta_pos:
+ return translated_features, delta_2d_pos
+ return translated_features
+
+ @ensure_batched(features=3)
+ def translate_2d_to_zero(
+ self,
+ features: torch.Tensor,
+ return_delta_pos: bool = False,
+ ) -> torch.Tensor:
+ """Translate each sequence so frame-0 root ``(x, z)`` is at the origin."""
+ target_2d_pos = torch.zeros(len(features), 2, device=features.device)
+ return self.translate_2d_to(features, target_2d_pos, return_delta_pos=return_delta_pos)
+
+ @ensure_batched(features=3)
+ def canonicalize(self, features: torch.Tensor, normalized: bool = False):
+ """Canonicalize heading and planar position at frame 0."""
+ if normalized:
+ features = self.unnormalize(features)
+ rotated_features = self.rotate_to_zero(features)
+ canonicalized_features = self.translate_2d_to_zero(rotated_features)
+ if normalized:
+ canonicalized_features = self.normalize(canonicalized_features)
+ return canonicalized_features
+
+ def normalize(self, features):
+ """Normalize features."""
+ return self.stats.normalize(features)
+
+ def unnormalize(self, features):
+ """Undo feature normalization."""
+ return self.stats.unnormalize(features)
+
+ def create_conditions_from_constraints(
+ self,
+ constraints_lst: list,
+ length: int,
+ to_normalize: bool,
+ device: str,
+ ):
+ """Create a conditioning tensor and mask from constraint objects."""
+ index_dict, data_dict = build_condition_dicts(constraints_lst)
+ return self.create_conditions(index_dict, data_dict, length, to_normalize, device)
+
+ def create_conditions_from_constraints_batched(
+ self,
+ constraints_lst: list | list[list],
+ lengths: torch.Tensor,
+ to_normalize: bool,
+ device: str,
+ ):
+ """Batched version of ``create_conditions_from_constraints``.
+
+ Supports either one shared constraint list for all batch elements, or a per-sample list of
+ constraint lists.
+ """
+ num_samples = len(lengths)
+ if not constraints_lst or not isinstance(constraints_lst[0], list):
+ # If no constraints, or constraints are shared across the batch,
+ # build once and repeat.
+ observed_motion, motion_mask = self.create_conditions_from_constraints(
+ constraints_lst, int(lengths.max()), to_normalize, device
+ )
+ observed_motion = repeat(observed_motion, "t d -> b t d", b=num_samples)
+ motion_mask = repeat(motion_mask, "t d -> b t d", b=num_samples)
+ return observed_motion, motion_mask
+
+ length = int(lengths.max())
+ observed_motion_lst = []
+ motion_mask_lst = []
+ for constraints_lst_el in constraints_lst:
+ observed_motion, motion_mask = self.create_conditions_from_constraints(
+ constraints_lst_el,
+ length,
+ to_normalize,
+ device,
+ )
+ observed_motion_lst.append(observed_motion)
+ motion_mask_lst.append(motion_mask)
+ observed_motion = torch.stack(observed_motion_lst, axis=0)
+ motion_mask = torch.stack(motion_mask_lst, axis=0)
+ return observed_motion, motion_mask
diff --git a/kimodo/motion_rep/reps/kimodo_motionrep.py b/kimodo/motion_rep/reps/kimodo_motionrep.py
new file mode 100644
index 0000000000000000000000000000000000000000..5cc6a05b3a662455cfda258ed8729af6f860584a
--- /dev/null
+++ b/kimodo/motion_rep/reps/kimodo_motionrep.py
@@ -0,0 +1,306 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from typing import Optional
+
+import einops
+import torch
+from torch import Tensor
+
+from kimodo.tools import to_numpy
+
+from ...geometry import cont6d_to_matrix, matrix_to_cont6d
+from ...skeleton.kinematics import fk
+from ...skeleton.transforms import global_rots_to_local_rots
+from ...tools import ensure_batched
+from ..conditioning import get_unique_index_and_data
+from ..feature_utils import RotateFeatures, compute_heading_angle, compute_vel_xyz
+from ..feet import foot_detect_from_pos_and_vel
+from ..smooth_root import get_smooth_root_pos
+from .base import MotionRepBase
+
+
+class KimodoMotionRep(MotionRepBase):
+ """Global root / global joints rotations representation, relative to a smooth root."""
+
+ def __init__(
+ self,
+ skeleton,
+ fps,
+ stats_path: Optional[str] = None,
+ ):
+ nbjoints = skeleton.nbjoints
+
+ self.size_dict = {
+ "smooth_root_pos": torch.Size([3]),
+ "global_root_heading": torch.Size([2]),
+ "local_joints_positions": torch.Size([nbjoints, 3]),
+ "global_rot_data": torch.Size([nbjoints, 6]),
+ "velocities": torch.Size([nbjoints, 3]),
+ "foot_contacts": torch.Size([4]),
+ }
+ self.last_root_feature = "global_root_heading"
+ self.local_root_size_dict = {
+ "local_root_rot_vel": torch.Size([1]),
+ "local_root_vel": torch.Size([2]),
+ "global_root_y": torch.Size([1]),
+ }
+ super().__init__(skeleton, fps, stats_path)
+
+ @ensure_batched(local_joint_rots=5, root_positions=3, lengths=1)
+ def __call__(
+ self,
+ local_joint_rots: torch.Tensor,
+ root_positions: torch.Tensor,
+ to_normalize: bool,
+ to_canonicalize: bool = False,
+ lengths: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ """Convert local rotations and root trajectory into smooth-root features.
+
+ Args:
+ local_joint_rots: Local joint rotation matrices ``[B, T, J, 3, 3]``.
+ root_positions: Root positions ``[B, T, 3]``.
+ to_normalize: Whether to normalize output features.
+ to_canonicalize: Whether to canonicalize output features (False by default).
+ lengths: Optional valid lengths for variable-length batches.
+
+ Returns:
+ Motion features with shape ``[B, T, motion_rep_dim]``.
+ """
+ device = local_joint_rots.device
+ if lengths is None:
+ assert local_joint_rots.shape[0] == 1, "If lenghts is not provided, the input should not be batched."
+ lengths = torch.tensor([local_joint_rots.shape[1]], device=device)
+
+ (
+ global_joints_rots,
+ global_joints_positions,
+ local_joints_positions_origin_is_pelvis,
+ ) = fk(local_joint_rots, root_positions, self.skeleton)
+
+ root_heading_angle = compute_heading_angle(global_joints_positions, self.skeleton)
+ global_root_heading = torch.stack([torch.cos(root_heading_angle), torch.sin(root_heading_angle)], dim=-1)
+
+ smooth_root_pos = get_smooth_root_pos(root_positions)
+ hips_offset = root_positions - smooth_root_pos
+ hips_offset[..., 1] = root_positions[..., 1]
+ local_joints_positions = local_joints_positions_origin_is_pelvis + hips_offset[:, :, None]
+
+ velocities = compute_vel_xyz(global_joints_positions, self.fps, lengths=lengths)
+ foot_contacts = foot_detect_from_pos_and_vel(global_joints_positions, velocities, self.skeleton, 0.15, 0.10)
+ global_rot_data = matrix_to_cont6d(global_joints_rots)
+
+ features, _ = einops.pack(
+ [
+ smooth_root_pos,
+ global_root_heading,
+ local_joints_positions,
+ global_rot_data,
+ velocities,
+ foot_contacts,
+ ],
+ "batch time *",
+ )
+
+ if to_canonicalize:
+ features = self.canonicalize(features, normalized=False)
+
+ if to_normalize:
+ features = self.normalize(features)
+ return features
+
+ @ensure_batched(features=3, angle=1)
+ def rotate(self, features: torch.Tensor, angle: torch.Tensor):
+ """Rotate root/joint positional and rotational features by heading."""
+ # assume it is not normalized
+ bs = features.shape[0]
+ device = features.device
+ [
+ smooth_root_pos,
+ global_root_heading,
+ local_joints_positions,
+ global_rot_data,
+ velocities,
+ foot_contacts,
+ ] = einops.unpack(features, self.ps, "batch time *")
+
+ if not isinstance(angle, torch.Tensor):
+ angle = torch.tensor(angle, device=device)
+ if len(angle.shape) == 0:
+ angle = angle.repeat(bs)
+
+ RF = RotateFeatures(angle)
+ new_features, _ = einops.pack(
+ [
+ RF.rotate_positions(smooth_root_pos),
+ RF.rotate_2d_positions(global_root_heading),
+ RF.rotate_positions(local_joints_positions),
+ RF.rotate_6d_rotations(global_rot_data),
+ RF.rotate_positions(velocities),
+ foot_contacts,
+ ],
+ "batch time *",
+ )
+ return new_features
+
+ @ensure_batched(features=3, translation_2d=2)
+ def translate_2d(
+ self,
+ features: torch.Tensor,
+ translation_2d: torch.Tensor,
+ ) -> torch.Tensor:
+ """Translate smooth root planar position by ``(dx, dz)``."""
+ # only move on the ground
+ # If we need a translate_3D function, we should not forget to move the local_joints_positions as well
+ bs = features.shape[0]
+ if len(translation_2d.shape) == 1:
+ translation_2d = translation_2d.repeat(bs, 1)
+
+ new_features = features.clone()
+ new_smooth_root_pos = new_features[:, :, self.slice_dict["smooth_root_pos"]]
+ new_smooth_root_pos[:, :, 0] += translation_2d[:, [0]]
+ new_smooth_root_pos[:, :, 2] += translation_2d[:, [1]]
+ return new_features
+
+ @ensure_batched(features=3)
+ def inverse(
+ self,
+ features: torch.Tensor,
+ is_normalized: bool,
+ posed_joints_from="rotations",
+ return_numpy: bool = False,
+ ) -> torch.Tensor:
+ """Decode smooth-root features into motion tensors."""
+ assert posed_joints_from in [
+ "rotations",
+ "positions",
+ ], "posed_joints_from should 'rotations' or 'positions'"
+
+ if is_normalized:
+ features = self.unnormalize(features)
+
+ [
+ smooth_root_pos,
+ global_root_heading,
+ local_joints_positions,
+ global_rot_data,
+ velocities,
+ foot_contacts,
+ ] = einops.unpack(features, self.ps, "batch time *")
+
+ global_rot_mats = cont6d_to_matrix(global_rot_data)
+ local_rot_mats = global_rots_to_local_rots(global_rot_mats, self.skeleton)
+
+ posed_joints_from_pos = local_joints_positions.clone()
+ posed_joints_from_pos[..., 0] += smooth_root_pos[..., None, 0]
+ posed_joints_from_pos[..., 2] += smooth_root_pos[..., None, 2]
+ root_positions = posed_joints_from_pos[..., self.skeleton.root_idx, :]
+ foot_contacts = foot_contacts > 0.5
+
+ if posed_joints_from == "rotations":
+ _, posed_joints, _ = self.skeleton.fk(
+ local_rot_mats,
+ root_positions,
+ )
+ else:
+ posed_joints = posed_joints_from_pos
+
+ output_tensor_dict = {
+ "local_rot_mats": local_rot_mats,
+ "global_rot_mats": global_rot_mats,
+ "posed_joints": posed_joints,
+ "root_positions": root_positions,
+ "smooth_root_pos": smooth_root_pos,
+ "foot_contacts": foot_contacts,
+ "global_root_heading": global_root_heading,
+ }
+ if return_numpy:
+ return to_numpy(output_tensor_dict)
+ return output_tensor_dict
+
+ def create_conditions(
+ self,
+ index_dict: dict[Tensor],
+ data_dict: dict[Tensor],
+ length: int,
+ to_normalize: bool,
+ device: str,
+ ):
+ """Build sparse conditioning tensors for smooth-root representation."""
+ # create empty features and mask to be filled in
+ observed_motion = torch.zeros(length, self.motion_rep_dim, device=device)
+ motion_mask = torch.zeros(length, self.motion_rep_dim, dtype=bool, device=device)
+
+ def _cat_indices(indices_list: list[Tensor]) -> Tensor:
+ indices = torch.cat([torch.tensor(x) if not isinstance(x, Tensor) else x for x in indices_list])
+ return indices.to(device=device, dtype=torch.long)
+
+ def _match_obs_dtype(tensor: Tensor) -> Tensor:
+ return tensor.to(device=device, dtype=observed_motion.dtype)
+
+ if (fname := "smooth_root_2d") in index_dict and index_dict[fname]:
+ indices = _cat_indices(index_dict[fname])
+ indices, smooth_root_2d = get_unique_index_and_data(indices, torch.cat(data_dict[fname]))
+ smooth_root_2d = _match_obs_dtype(smooth_root_2d)
+ f_sliced = observed_motion[:, self.slice_dict["smooth_root_pos"]]
+ f_sliced[indices, 0] = smooth_root_2d[:, 0]
+ f_sliced[indices, 2] = smooth_root_2d[:, 1]
+ m_sliced = motion_mask[:, self.slice_dict["smooth_root_pos"]]
+ m_sliced[indices, 0] = True
+ m_sliced[indices, 2] = True
+
+ if (fname := "root_y_pos") in index_dict and index_dict[fname]:
+ indices = _cat_indices(index_dict[fname])
+ indices, root_pos_Y = get_unique_index_and_data(indices, torch.cat(data_dict[fname]))
+ root_pos_Y = _match_obs_dtype(root_pos_Y)
+ f_sliced = observed_motion[:, self.slice_dict["smooth_root_pos"]]
+ f_sliced[indices, 1] = root_pos_Y
+ m_sliced = motion_mask[:, self.slice_dict["smooth_root_pos"]]
+ m_sliced[indices, 1] = True
+
+ if (fname := "global_root_heading") in index_dict and index_dict[fname]:
+ indices = _cat_indices(index_dict[fname])
+ indices, global_root_heading = get_unique_index_and_data(indices, torch.cat(data_dict[fname]))
+ global_root_heading = _match_obs_dtype(global_root_heading)
+ f_sliced = observed_motion[:, self.slice_dict[fname]]
+ f_sliced[indices] = global_root_heading
+ m_sliced = motion_mask[:, self.slice_dict[fname]]
+ m_sliced[indices] = True
+
+ if (fname := "global_joints_rots") in index_dict and index_dict[fname]:
+ indices_lst = _cat_indices(index_dict[fname])
+ indices_lst, global_joints_rots = get_unique_index_and_data(indices_lst, torch.cat(data_dict[fname]))
+ global_joints_rots = _match_obs_dtype(global_joints_rots)
+ global_rot_data = matrix_to_cont6d(global_joints_rots)
+ f_sliced = observed_motion[:, self.slice_dict["global_rot_data"]]
+ masking = torch.zeros(len(f_sliced) * self.nbjoints, 6, device=device, dtype=bool)
+ masking[indices_lst.T[0] * self.nbjoints + indices_lst.T[1]] = True
+ masking = masking.reshape(len(f_sliced), self.nbjoints * 6)
+ f_sliced[masking] = global_rot_data.flatten()
+ m_sliced = motion_mask[:, self.slice_dict["global_rot_data"]]
+ m_sliced[masking] = True
+
+ if (fname := "global_joints_positions") in index_dict and index_dict[fname]:
+ indices_lst = _cat_indices(index_dict[fname])
+ indices_lst, global_joints_positions = get_unique_index_and_data(indices_lst, torch.cat(data_dict[fname]))
+ global_joints_positions = _match_obs_dtype(global_joints_positions)
+ T_indices = indices_lst[:, 0].contiguous()
+ _test = motion_mask[T_indices, self.slice_dict["smooth_root_pos"]]
+ if not _test[:, [0, 2]].all():
+ raise ValueError("For constraining global positions, the smooth root should also be constrained.")
+ smooth_root_pos = observed_motion[T_indices, self.slice_dict["smooth_root_pos"]].clone()
+ local_reference = smooth_root_pos.clone()
+ local_reference[..., 1] = 0.0
+ local_joints_positions = global_joints_positions - local_reference
+ f_sliced = observed_motion[:, self.slice_dict["local_joints_positions"]]
+ masking = torch.zeros(len(f_sliced) * self.nbjoints, 3, device=device, dtype=bool)
+ masking[indices_lst.T[0] * self.nbjoints + indices_lst.T[1]] = True
+ masking = masking.reshape(len(f_sliced), self.nbjoints * 3)
+ f_sliced[masking] = local_joints_positions.flatten()
+ m_sliced = motion_mask[:, self.slice_dict["local_joints_positions"]]
+ m_sliced[masking] = True
+
+ if to_normalize:
+ observed_motion = self.normalize(observed_motion)
+ return observed_motion, motion_mask
diff --git a/kimodo/motion_rep/reps/tmr_motionrep.py b/kimodo/motion_rep/reps/tmr_motionrep.py
new file mode 100644
index 0000000000000000000000000000000000000000..8ff38caf96910ae452cc89303df6be9e0461810d
--- /dev/null
+++ b/kimodo/motion_rep/reps/tmr_motionrep.py
@@ -0,0 +1,234 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""TMR motion representation: global root, global joints, velocities, and foot contacts."""
+
+from typing import Optional
+
+import einops
+import torch
+
+from ...skeleton.kinematics import fk
+from ...tools import ensure_batched, to_numpy
+from ..feature_utils import RotateFeatures, compute_heading_angle, compute_vel_xyz
+from ..feet import foot_detect_from_pos_and_vel
+from .base import MotionRepBase
+
+
+class TMRMotionRep(MotionRepBase):
+ """Motion representation with global root and local joint positions.
+ The local joint positions are rotation invariant (they all face z+)
+
+ Feature layout:
+ - root position ``(x, y, z)``
+ - root heading as ``(cos(theta), sin(theta))``
+ - local joint positions (root and rotation removed)
+ - local joint velocities (rotation removed)
+ - binary foot contacts
+ """
+
+ def __init__(
+ self,
+ skeleton,
+ fps,
+ stats_path: Optional[str] = None,
+ ):
+ nbjoints = skeleton.nbjoints
+
+ self.size_dict = {
+ "root_pos": torch.Size([3]),
+ "global_root_heading": torch.Size([2]),
+ "local_joints_positions": torch.Size([nbjoints - 1, 3]),
+ "velocities": torch.Size([nbjoints, 3]),
+ "foot_contacts": torch.Size([4]),
+ }
+ self.last_root_feature = "global_root_heading"
+ self.local_root_size_dict = {
+ "local_root_rot_vel": torch.Size([1]),
+ "local_root_vel": torch.Size([2]),
+ "global_root_y": torch.Size([1]),
+ }
+ super().__init__(skeleton, fps, stats_path)
+
+ @ensure_batched(local_joint_rots=5, root_positions=3, posed_joints=4, lengths=1)
+ def __call__(
+ self,
+ local_joint_rots: Optional[torch.Tensor] = None,
+ root_positions: Optional[torch.Tensor] = None,
+ posed_joints: Optional[torch.Tensor] = None,
+ *,
+ to_normalize: bool,
+ to_canonicalize: bool = False,
+ lengths: Optional[torch.Tensor] = None,
+ ) -> torch.Tensor:
+ """Convert motion inputs to this feature representation.
+
+ Args:
+ local_joint_rots: Local joint rotation matrices ``[B, T, J, 3, 3]``.
+ Required when ``posed_joints`` is not provided.
+ root_positions: Root translations ``[B, T, 3]``. Required when
+ ``posed_joints`` is not provided.
+ posed_joints: Optional precomputed global joint positions
+ ``[B, T, J, 3]``. If passed, FK is skipped.
+ to_normalize: Whether to normalize output features.
+ to_canonicalize: Whether to canonicalize output features (False by default).
+ lengths: Optional valid lengths for variable-length batches.
+
+ Returns:
+ Motion features with shape ``[B, T, motion_rep_dim]``.
+ """
+ if posed_joints is not None:
+ device = posed_joints.device
+ nbatch, nbframes, nbjoints = posed_joints.shape[:3]
+ else:
+ device = local_joint_rots.device
+ nbatch, nbframes, nbjoints = local_joint_rots.shape[:3]
+
+ if lengths is None:
+ assert nbatch == 1, "If lenghts is not provided, the input should not be batched."
+ lengths = torch.tensor([nbframes], device=device)
+
+ if posed_joints is None:
+ _, global_positions, local_joints_positions_origin_is_pelvis = fk(
+ local_joint_rots, root_positions, self.skeleton
+ )
+ else:
+ global_positions = posed_joints
+ root_positions = posed_joints[:, :, 0]
+ local_joints_positions_origin_is_pelvis = posed_joints - root_positions[:, :, None]
+
+ root_heading_angle = compute_heading_angle(global_positions, self.skeleton)
+ global_root_heading = torch.stack([torch.cos(root_heading_angle), torch.sin(root_heading_angle)], dim=-1)
+
+ ground_offset = 0 * root_positions
+ ground_offset[..., 1] = root_positions[..., 1]
+
+ local_joints_positions = local_joints_positions_origin_is_pelvis[:, :, 1:] + ground_offset[:, :, None]
+ velocities = compute_vel_xyz(global_positions, self.fps, lengths=lengths)
+
+ # Remove the heading angle for each frame
+ RF = RotateFeatures(-root_heading_angle)
+ local_joints_positions = RF.rotate_positions(local_joints_positions)
+ velocities = RF.rotate_positions(velocities)
+
+ foot_contacts = foot_detect_from_pos_and_vel(global_positions, velocities, self.skeleton, 0.15, 0.10)
+ features, _ = einops.pack(
+ [
+ root_positions,
+ global_root_heading,
+ local_joints_positions,
+ velocities,
+ foot_contacts,
+ ],
+ "batch time *",
+ )
+
+ if to_canonicalize:
+ features = self.canonicalize(features, normalized=False)
+
+ if to_normalize:
+ features = self.normalize(features)
+ return features
+
+ @ensure_batched(features=3, angle=1)
+ def rotate(self, features: torch.Tensor, angle: torch.Tensor):
+ """Rotate all spatial features by a heading delta (radians)."""
+ # rotate by the angle
+ # it add the angle to the current features
+ # assume it is not normalized
+ bs = features.shape[0]
+ device = features.device
+ [
+ root_pos,
+ global_root_heading,
+ local_joints_positions,
+ velocities,
+ foot_contacts,
+ ] = einops.unpack(features, self.ps, "batch time *")
+
+ if not isinstance(angle, torch.Tensor):
+ angle = torch.tensor(angle, device=device)
+ if len(angle.shape) == 0:
+ angle = angle.repeat(bs)
+
+ RF = RotateFeatures(angle)
+ new_features, _ = einops.pack(
+ [
+ RF.rotate_positions(root_pos),
+ RF.rotate_2d_positions(global_root_heading),
+ local_joints_positions, # already rotation invariant
+ velocities, # already rotation invariant
+ foot_contacts,
+ ],
+ "batch time *",
+ )
+ return new_features
+
+ @ensure_batched(features=3, translation_2d=2)
+ def translate_2d(
+ self,
+ features: torch.Tensor,
+ translation_2d: torch.Tensor,
+ ) -> torch.Tensor:
+ """Translate root planar position by ``(dx, dz)``."""
+ # only move on the ground
+ # For 3D, we should not forget to move the local_joints_positions as well
+ bs = features.shape[0]
+ if len(translation_2d.shape) == 1:
+ translation_2d = translation_2d.repeat(bs, 1)
+
+ new_features = features.clone()
+ new_root_pos = new_features[:, :, self.slice_dict["root_pos"]]
+ new_root_pos[:, :, 0] += translation_2d[:, 0]
+ new_root_pos[:, :, 2] += translation_2d[:, 1]
+ return new_features
+
+ @ensure_batched(features=3)
+ def inverse(
+ self,
+ features: torch.Tensor,
+ is_normalized: bool,
+ posed_joints_from="positions",
+ return_numpy: bool = False,
+ ) -> torch.Tensor:
+ """Decode features back to a motion dictionary.
+
+ Args:
+ features: Feature tensor ``[B, T, D]``.
+ is_normalized: Whether input features are normalized.
+ posed_joints_from: Must be ``"positions"`` for this representation.
+ return_numpy: Whether to convert tensors to numpy arrays.
+
+ Returns:
+ Dictionary containing reconstructed positions and auxiliary data.
+ """
+ assert posed_joints_from == "positions"
+ if is_normalized:
+ features = self.unnormalize(features)
+
+ [
+ root_positions,
+ global_root_heading,
+ local_joints_positions,
+ velocities,
+ foot_contacts,
+ ] = einops.unpack(features, self.ps, "batch time *")
+
+ dummy_root = 0 * local_joints_positions[:, :, [0]]
+ posed_joints_from_pos = torch.stack([dummy_root, local_joints_positions], axis=2)
+ posed_joints_from_pos[..., 0] += root_positions[..., None, 0]
+ posed_joints_from_pos[..., 2] += root_positions[..., None, 2]
+ root_positions = posed_joints_from_pos[..., self.skeleton.root_idx, :]
+ foot_contacts = foot_contacts > 0.5
+ posed_joints = posed_joints_from_pos
+
+ output_tensor_dict = {
+ "local_rot_mats": None,
+ "global_rot_mats": None,
+ "posed_joints": posed_joints,
+ "root_positions": root_positions,
+ "foot_contacts": foot_contacts,
+ "global_root_heading": global_root_heading,
+ }
+ if return_numpy:
+ return to_numpy(output_tensor_dict)
+ return output_tensor_dict
diff --git a/kimodo/motion_rep/smooth_root.py b/kimodo/motion_rep/smooth_root.py
new file mode 100644
index 0000000000000000000000000000000000000000..dd23f6f35d5a87241afbbf4d6e3106d636927bc7
--- /dev/null
+++ b/kimodo/motion_rep/smooth_root.py
@@ -0,0 +1,234 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Smooth root trajectory: ADMM-based smoother with margin constraints and get_smooth_root_pos helper."""
+
+import math
+
+import numpy as np
+import torch
+from scipy import sparse
+from scipy.sparse.linalg import splu
+
+from kimodo.tools import ensure_batched
+
+
+class TrajectorySmoother:
+ """Modify trajectories to hit target values while respecting soft constraints.
+
+ This smoother keeps the trajectory close to the original positions while minimizing
+ accelerations. Targets are enforced at specified frames via soft constraints.
+ """
+
+ def __init__(
+ self,
+ margins,
+ pos_weight=0.0,
+ loop=False,
+ admm_iters=100,
+ alpha_overrelax=1.0,
+ circle_project=False,
+ ):
+ """Initialize the TrajectorySmoother.
+
+ Args:
+ margins: Array of margin values for each frame.
+ margins[i] < 0: unconstrained
+ margins[i] == 0: pinned on this frame
+ margins[i] > 0: can deviate within the margin
+ pos_weight: Weight for position preservation
+ loop: Whether the trajectory should loop
+ admm_iters: Number of ADMM iterations
+ """
+ self.pos_weight = pos_weight
+ self.admm_iters = admm_iters
+ self.alpha_overrelax = alpha_overrelax
+ self.circle_project = circle_project
+ N = len(margins)
+
+ # Store margin information as numpy arrays
+ self.margin_vals = margins
+
+ # Build acceleration matrix A
+ a_data = []
+ a_rows = []
+ a_cols = []
+
+ for i in range(1, N - 1):
+ scale = 1.0
+ a_data.extend([-scale, 2.0 * scale, -scale])
+ a_rows.extend([i, i, i])
+ a_cols.extend([i - 1, i, i + 1])
+
+ if loop:
+ # Add periodic accelerations
+ scale = 1.0
+ a_data.extend([-scale, 2.0 * scale, -scale])
+ a_rows.extend([0, 0, 0])
+ a_cols.extend([N - 1, 0, 1])
+
+ scale = 1.0
+ a_data.extend([-scale, 2.0 * scale, -scale])
+ a_rows.extend([N - 1, N - 1, N - 1])
+ a_cols.extend([N - 2, N - 1, 0])
+
+ A = sparse.csr_matrix((a_data, (a_rows, a_cols)), shape=(N, N))
+
+ # Build identity matrix
+ identity_matrix = sparse.eye(N)
+
+ # Build system matrix M
+ M = pos_weight * identity_matrix + A.T @ A
+
+ # Calculate ADMM step size
+ diag_max = max(abs(M.diagonal()))
+ self.admm_stepsize = 0.25 * np.sqrt(diag_max)
+
+ M = M + self.admm_stepsize * identity_matrix
+ self.system_lu = splu(M.tocsc())
+
+ def smooth(self, targets, x0):
+ """Interpolate between reference positions while satisfying constraints.
+
+ Args:
+ observations: Target positions for constrained frames (numpy array)
+ ref_positions: Reference positions defining original shape
+ (numpy array)
+
+ Returns:
+ Interpolated positions (numpy array)
+ """
+ x_target = targets.copy()
+ x = x0.copy()
+ z = np.zeros_like(x)
+ u = np.zeros_like(x)
+
+ for _ in range(self.admm_iters):
+ self.z_update(z, x, x_target, u)
+ self.u_update(u, x, z)
+ self.x_update(x, z, u, x_target)
+
+ return x
+
+ def x_update(self, x, z, u, x_t):
+ """Update x in the ADMM iteration."""
+
+ # x = (wp * I + A^T A + p I)^-1 (wp * x_orig + p (z - u))
+ r = self.pos_weight * x_t + self.admm_stepsize * (z - u)
+ x[:] = self.system_lu.solve(r)
+
+ def z_update(self, z, x, z_t, u):
+ """Update z in the ADMM iteration using vectorized operations."""
+ # Compute the difference from target for all margin locations at once
+ z[:] = x + u - z_t
+
+ # Check if we need to project back to margin
+ z_diff_norms = np.linalg.norm(z, axis=1)
+ mask = z_diff_norms > self.margin_vals
+ if np.any(mask):
+ scale_factors = self.margin_vals[mask] / z_diff_norms[mask]
+ z[mask] *= scale_factors[:, np.newaxis]
+
+ # Add back the target
+ z[:] += z_t
+
+ if self.circle_project:
+ z[:] = z / (np.linalg.norm(z, axis=1, keepdims=True) + 1.0e-6)
+
+ def u_update(self, u, x, z):
+ """Update u in the ADMM iteration using vectorized operations."""
+ u[:] += self.alpha_overrelax * (x - z)
+
+
+def smooth_signal(x, margins, pos_weight=0, alpha_overrelax=1.8, admm_iters=500, circle_project=False):
+ """Multigrid trajectory smoothing with margin constraints.
+
+ Args:
+ x: Input trajectory ``[T, D]`` as a NumPy array.
+ margins: Allowed radius around each target frame ``[T]``.
+ pos_weight: Weight for staying close to the original signal.
+ alpha_overrelax: ADMM over-relaxation coefficient.
+ admm_iters: ADMM iterations per multigrid level.
+ circle_project: If ``True``, project each vector to the unit sphere.
+
+ Returns:
+ Smoothed trajectory of shape ``[T, D]``.
+ """
+ x_smoothed = x.copy()
+ x_smoothed[:] = x.mean(axis=0, keepdims=True)
+
+ # smooth the signal, multigrid style by starting out coarse,
+ # doubling the resolution and repeating until we're at the full
+ # resolution, using the previous result as the initial guess.
+ levels = int(math.floor(math.log2(len(x))))
+ levels = max(levels - 4, 1)
+
+ stepsize = 2**levels
+ while True:
+ # smooth signals at this level:
+ num_steps = len(x_smoothed[::stepsize])
+ smoother = TrajectorySmoother(
+ margins=margins[::stepsize],
+ pos_weight=pos_weight,
+ alpha_overrelax=alpha_overrelax,
+ admm_iters=admm_iters,
+ circle_project=circle_project,
+ )
+ x_smoothed[::stepsize] = smoother.smooth(x[::stepsize], x_smoothed[::stepsize])
+
+ # interpolate to next level:
+ next_stepsize = stepsize // 2
+ num_interleaved = len(x_smoothed[next_stepsize::stepsize])
+ if num_interleaved == num_steps:
+ # linearly extrapolate the last value if we have to:
+ x_smoothed[next_stepsize::stepsize][-1] = (
+ x_smoothed[::stepsize][-1] + (x_smoothed[::stepsize][-1] - x_smoothed[::stepsize][-2]) / 2
+ )
+ num_interleaved = num_interleaved - 1
+
+ # linearly interpolate the remaining values:
+ x_smoothed[next_stepsize::stepsize][:num_interleaved] = (
+ x_smoothed[::stepsize][:-1] + x_smoothed[::stepsize][1:]
+ ) / 2
+
+ if stepsize == 1:
+ break
+
+ stepsize //= 2
+
+ return x_smoothed
+
+
+@ensure_batched(hip_translations=3)
+def get_smooth_root_pos(hip_translations):
+ """Smooth root trajectory in the ground plane while preserving height.
+
+ Args:
+ hip_translations: Root translations ``[B, T, 3]``.
+
+ Returns:
+ Smoothed root translations ``[B, T, 3]`` where ``x/z`` are smoothed and
+ ``y`` remains unchanged.
+ """
+ root_translations_xz = hip_translations[..., [0, 2]]
+ root_translations_y = hip_translations[..., [1]]
+
+ batch_size, nframes = root_translations_xz.shape[:2]
+ margins = np.full(root_translations_xz.shape[1], 0.06)
+
+ root_translations_smoothed_xz = []
+ for batch in range(batch_size):
+ root_translations_smoothed_xz.append(
+ smooth_signal(root_translations_xz[batch].detach().cpu().numpy(), margins)[None]
+ )
+
+ root_translations_smoothed_xz = torch.tensor(np.concatenate(root_translations_smoothed_xz))
+
+ root_translations = torch.cat(
+ [
+ root_translations_smoothed_xz.to(root_translations_y.device),
+ root_translations_y,
+ ],
+ dim=-1,
+ )[..., [0, 2, 1]]
+
+ return root_translations
diff --git a/kimodo/motion_rep/stats.py b/kimodo/motion_rep/stats.py
new file mode 100644
index 0000000000000000000000000000000000000000..eecd9d417c1a4e54546e75afe3fb82f9e5e57e31
--- /dev/null
+++ b/kimodo/motion_rep/stats.py
@@ -0,0 +1,123 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Feature normalization statistics (mean/std) for motion representations."""
+
+import logging
+import os
+from typing import Optional
+
+import numpy as np
+import torch
+
+log = logging.getLogger(__name__)
+
+
+class Stats(torch.nn.Module):
+ """Utility module for feature normalization statistics.
+
+ Normalization follows:
+ ``(data - mean) / sqrt(std**2 + eps)``
+ """
+
+ def __init__(
+ self,
+ folder: Optional[str] = None,
+ load: bool = True,
+ eps=1e-05,
+ ):
+ super().__init__()
+ self.folder = folder
+ self.eps = eps
+ if folder is not None and load:
+ self.load()
+
+ def sliced(self, indices):
+ """Return a new ``Stats`` object containing selected feature indices."""
+ new_stats = Stats(folder=self.folder, load=False, eps=self.eps)
+ new_stats.register_from_tensors(
+ self.mean[..., indices].clone(),
+ self.std[..., indices].clone(),
+ )
+ return new_stats
+
+ def load(self):
+ """Load ``mean.npy`` and ``std.npy`` from ``self.folder``."""
+ mean_path = os.path.join(self.folder, "mean.npy")
+ std_path = os.path.join(self.folder, "std.npy")
+ if not os.path.exists(mean_path) or not os.path.exists(std_path):
+ raise FileNotFoundError(
+ f"Missing stats files in '{self.folder}'. Expected:\n"
+ f" - {mean_path}\n"
+ f" - {std_path}\n\n"
+ "Make sure the checkpoint/stats have been downloaded and are mounted into the container.\n"
+ "If you're using Docker Compose, run it from the repo root so `./:/workspace` mounts the correct directory."
+ )
+
+ mean = torch.from_numpy(np.load(mean_path))
+ std = torch.from_numpy(np.load(std_path))
+ self.register_from_tensors(mean, std)
+
+ def register_from_tensors(self, mean: torch.Tensor, std: torch.Tensor):
+ """Register mean/std tensors as non-persistent buffers."""
+ self.register_buffer("mean", mean, persistent=False)
+ self.register_buffer("std", std, persistent=False)
+
+ def normalize(self, data: torch.Tensor) -> torch.Tensor:
+ """Normalize data using the stored statistics."""
+ mean = self.mean.to(device=data.device, dtype=data.dtype)
+ std = self.std.to(device=data.device, dtype=data.dtype)
+ # adjust std with eps
+ return (data - mean) / torch.sqrt(std**2 + self.eps)
+
+ def unnormalize(self, data: torch.Tensor) -> torch.Tensor:
+ """Undo normalization using the stored statistics."""
+ mean = self.mean.to(device=data.device, dtype=data.dtype)
+ std = self.std.to(device=data.device, dtype=data.dtype)
+ # adjust std with eps
+ return data * torch.sqrt(std**2 + self.eps) + mean
+
+ def is_loaded(self):
+ """Return whether statistics are currently available."""
+ return hasattr(self, "mean")
+
+ def get_dim(self):
+ """Return feature dimensionality."""
+ return self.mean.shape[0]
+
+ def save(
+ self,
+ folder: Optional[str] = None,
+ mean: Optional[torch.Tensor] = None,
+ std: Optional[torch.Tensor] = None,
+ ):
+ """Save statistics to ``folder`` as ``mean.npy`` and ``std.npy``."""
+ if folder is None:
+ folder = self.folder
+ if folder is None:
+ raise ValueError("No folder to save stats")
+
+ if mean is None and std is None:
+ try:
+ mean = self.mean.cpu().numpy()
+ std = self.std.cpu().numpy()
+ except AttributeError:
+ raise ValueError("Stats were not loaded")
+
+ # don't override stats folder
+ os.makedirs(folder, exist_ok=False)
+
+ np.save(os.path.join(folder, "mean.npy"), mean)
+ np.save(os.path.join(folder, "std.npy"), std)
+
+ def __eq__(self, other):
+ return (self.mean.cpu() == other.mean.cpu()).all() and (self.std.cpu() == other.std.cpu()).all()
+
+ # should define a hash value for pytorch, as we defined __eq__
+ def __hash__(self):
+ # Convert mean and std to bytes for a consistent hash value
+ mean_hash = hash(self.mean.detach().cpu().numpy().tobytes())
+ std_hash = hash(self.std.detach().cpu().numpy().tobytes())
+ return hash((mean_hash, std_hash))
+
+ def __repr__(self):
+ return f'Stats(folder="{self.folder}")'
diff --git a/kimodo/postprocess.py b/kimodo/postprocess.py
new file mode 100644
index 0000000000000000000000000000000000000000..5d0b6dcacbe655f6b95e3e839e02ee05a77f354b
--- /dev/null
+++ b/kimodo/postprocess.py
@@ -0,0 +1,351 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Post-processing utilities for motion generation output."""
+
+from types import SimpleNamespace
+from typing import Dict, List, Optional, Tuple
+
+import numpy as np
+import torch
+
+from .constraints import (
+ EndEffectorConstraintSet,
+ FullBodyConstraintSet,
+ Root2DConstraintSet,
+)
+from .geometry import matrix_to_quaternion, quaternion_to_matrix
+from .skeleton import (
+ G1Skeleton34,
+ SkeletonBase,
+ SMPLXSkeleton22,
+ SOMASkeleton30,
+ SOMASkeleton77,
+ fk,
+)
+
+
+def extract_input_motion_from_constraints(
+ constraint_lst: List,
+ skeleton: SkeletonBase,
+ num_frames: int,
+ num_joints: int,
+) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Extract hip translations and local rotations from constraints for postprocessing.
+
+ Args:
+ constraint_lst: List of constraints (FullBodyConstraintSet, EndEffectorConstraintSet, etc.)
+ skeleton: Skeleton instance
+ num_frames: Total number of frames in the motion
+ num_joints: Number of joints
+
+ Returns:
+ Tuple of (hip_translations_input, rotations_input):
+ - hip_translations_input: Hip translations, shape (T, 3)
+ - rotations_input: Local joint rotations as quaternions, shape (T, J, 4)
+ """
+ # Initialize with zeros for all frames
+ hip_translations_input = torch.zeros(num_frames, 3)
+ rotations_input = torch.zeros(num_frames, num_joints, 4)
+ rotations_input[..., 0] = 1.0 # Initialize as identity quaternions (w=1, x=y=z=0)
+
+ def _match_hip_dtype(tensor: torch.Tensor) -> torch.Tensor:
+ return tensor.to(device=hip_translations_input.device, dtype=hip_translations_input.dtype)
+
+ def _match_rot_dtype(tensor: torch.Tensor) -> torch.Tensor:
+ return tensor.to(device=rotations_input.device, dtype=rotations_input.dtype)
+
+ if not constraint_lst:
+ return hip_translations_input, rotations_input
+
+ # Sort constraints to ensure FullBodyConstraintSet is processed last
+ # This ensures it will get the last say on whether hip translations need to be exact root or smoothed root
+ sorted_constraints = sorted(constraint_lst, key=lambda c: isinstance(c, FullBodyConstraintSet))
+ for constraint in sorted_constraints:
+ frame_indices = constraint.frame_indices
+ if isinstance(frame_indices, torch.Tensor):
+ valid_mask = frame_indices < num_frames
+ if valid_mask.sum() == 0:
+ continue
+ frame_indices = frame_indices[valid_mask]
+ else:
+ valid_positions = [i for i, idx in enumerate(frame_indices) if idx < num_frames]
+ if not valid_positions:
+ continue
+ frame_indices = [frame_indices[i] for i in valid_positions]
+
+ # Handle Root2DConstraintSet separately - only assign smooth_root_2d at xz dimensions
+ if isinstance(constraint, Root2DConstraintSet):
+ smooth_root_2d = constraint.smooth_root_2d # (K, 2) where K = len(frame_indices)
+ if isinstance(frame_indices, torch.Tensor):
+ smooth_root_2d = smooth_root_2d[valid_mask]
+ else:
+ smooth_root_2d = smooth_root_2d[valid_positions]
+ smooth_root_2d = _match_hip_dtype(smooth_root_2d)
+ hip_translations_input[frame_indices, 0] = smooth_root_2d[:, 0] # x
+ hip_translations_input[frame_indices, 2] = smooth_root_2d[:, 1] # z
+ continue
+ elif isinstance(constraint, FullBodyConstraintSet) or isinstance(constraint, EndEffectorConstraintSet):
+ global_rots = constraint.global_joints_rots # (K, J, 3, 3) where K = len(frame_indices)
+ global_positions = constraint.global_joints_positions # (K, J, 3)
+ if isinstance(frame_indices, torch.Tensor):
+ global_rots = global_rots[valid_mask]
+ global_positions = global_positions[valid_mask]
+ smooth_root_2d = constraint.smooth_root_2d[valid_mask]
+ else:
+ global_rots = global_rots[valid_positions]
+ global_positions = global_positions[valid_positions]
+ smooth_root_2d = constraint.smooth_root_2d[valid_positions]
+
+ root_positions = global_positions[:, skeleton.root_idx] # (K, 3)
+ # replace xz with smooth_root_2d values for EE constraints that do not include Hips
+ # since the hips themselves are not actually constrained in the model conditioning
+ if isinstance(constraint, EndEffectorConstraintSet) and "Hips" not in constraint.joint_names:
+ root_positions[:, 0] = smooth_root_2d[:, 0] # x
+ root_positions[:, 2] = smooth_root_2d[:, 1] # z
+
+ local_rot_mats = skeleton.global_rots_to_local_rots(global_rots) # (K, J, 3, 3)
+ local_rot_quats = matrix_to_quaternion(local_rot_mats) # (K, J, 4)
+
+ hip_translations_input[frame_indices] = _match_hip_dtype(root_positions)
+ rotations_input[frame_indices] = _match_rot_dtype(local_rot_quats)
+ else:
+ NotImplementedError(f"Constraint {constraint.name} is not supported")
+
+ return hip_translations_input, rotations_input
+
+
+def create_working_rig_from_skeleton(
+ skeleton: SkeletonBase, above_ground_offset: float = 0.007
+) -> List[SimpleNamespace]:
+ """Create the working rig as a list of SimpleNamespace objects from skeleton.
+
+ Args:
+ skeleton: SkeletonBase instance with bone_order_names, neutral_joints, joint_parents
+ above_ground_offset: Additional offset to position the rig slightly above ground
+ Returns:
+ List of SimpleNamespace objects representing the working rig
+ """
+ working_rig_joints = []
+
+ joint_names = skeleton.bone_order_names
+ neutral_positions = skeleton.neutral_joints.cpu().numpy()
+ parent_indices = skeleton.joint_parents.cpu().numpy()
+
+ if isinstance(skeleton, (G1Skeleton34, SMPLXSkeleton22)):
+ retarget_map = {
+ skeleton.bone_order_names[skeleton.root_idx]: "Hips",
+ skeleton.left_hand_joint_names[0]: "LeftHand",
+ skeleton.right_hand_joint_names[0]: "RightHand",
+ skeleton.left_foot_joint_names[0]: "LeftFoot",
+ skeleton.right_foot_joint_names[0]: "RightFoot",
+ }
+ else:
+ # works for SOMA
+ retarget_map = {
+ "Hips": "Hips",
+ "Head": "Head",
+ "LeftHand": "LeftHand",
+ "RightHand": "RightHand",
+ "LeftFoot": "LeftFoot",
+ "RightFoot": "RightFoot",
+ }
+
+ for i, joint_name in enumerate(joint_names):
+ parent_name = None if parent_indices[i] == -1 else joint_names[parent_indices[i]]
+
+ # Calculate local translation relative to parent
+ if parent_indices[i] == -1:
+ # Move the rig so that the lowest point (toe) is at ground level (y=0),
+ # plus a small offset to position the rig slightly above ground
+ toe_height = neutral_positions[:, 1].min() # lowest y-coordinate (toe)
+ local_translation = (
+ neutral_positions[i] + np.array([0.0, -toe_height + above_ground_offset, 0.0])
+ ).tolist()
+ else:
+ parent_idx = parent_indices[i]
+ parent_position = neutral_positions[parent_idx]
+ joint_position = neutral_positions[i]
+ local_translation = (joint_position - parent_position).tolist()
+
+ # Default rotation (identity quaternion: x=0, y=0, z=0, w=1)
+ default_rotation = [0.0, 0.0, 0.0, 1.0]
+
+ joint_info = SimpleNamespace(
+ name=joint_name,
+ parent=parent_name,
+ t_pose_rotation=default_rotation,
+ t_pose_translation=local_translation,
+ retarget_tag=retarget_map.get(joint_name),
+ )
+
+ working_rig_joints.append(joint_info)
+
+ return working_rig_joints
+
+
+def post_process_motion(
+ local_rot_mats: torch.Tensor,
+ root_positions: torch.Tensor,
+ contacts: torch.Tensor,
+ skeleton: SkeletonBase,
+ constraint_lst: Optional[List] = None,
+ contact_threshold: float = 0.5,
+ root_margin: float = 0.04,
+) -> Dict[str, torch.Tensor]:
+ """Post-process generated motion to reduce foot skating and improve quality.
+
+ Args:
+ local_rot_mats: Local joint rotation matrices, shape (B, T, J, 3, 3)
+ root_positions: Root joint positions, shape (B, T, 3)
+ contacts: Foot contact labels, shape (B, T, num_contacts)
+ skeleton: Skeleton instance
+ constraint_lst: Optional list of constraints (or list of lists of constraints for batched inference)(FullBodyConstraintSet, etc.)
+ contact_threshold: Threshold for foot contact detection
+ root_margin: Margin for root position correction
+
+ Returns:
+ Dictionary with corrected motion data:
+ - local_rot_mats: Corrected local rotation matrices (B, T, J, 3, 3)
+ - root_positions: Corrected root positions (B, T, 3)
+ - posed_joints: Corrected global joint positions (B, T, J, 3)
+ - global_rot_mats: Corrected global rotation matrices (B, T, J, 3, 3)
+ """
+ # Ensure batch dimension
+ assert local_rot_mats.dim() == 5, "local_rot_mats should be 5D, make sure to include the batch dimension"
+
+ batch_size, num_frames, num_joints = local_rot_mats.shape[:3]
+
+ def _build_constraint_masks_dict(constraints: List) -> Dict[str, torch.Tensor]:
+ out = {
+ key: torch.zeros(num_frames, dtype=torch.float32)
+ for key in [
+ "FullBody",
+ "LeftFoot",
+ "RightFoot",
+ "LeftHand",
+ "RightHand",
+ "Root",
+ ]
+ }
+ for constraint in constraints:
+ frame_indices = constraint.frame_indices
+ if isinstance(frame_indices, torch.Tensor):
+ frame_indices = frame_indices[frame_indices < num_frames]
+ if frame_indices.numel() == 0:
+ continue
+ else:
+ frame_indices = [idx for idx in frame_indices if idx < num_frames]
+ if not frame_indices:
+ continue
+ if constraint.name == "fullbody":
+ out["FullBody"][frame_indices] = 1.0
+ elif constraint.name == "left-foot":
+ out["LeftFoot"][frame_indices] = 1.0
+ elif constraint.name == "right-foot":
+ out["RightFoot"][frame_indices] = 1.0
+ elif constraint.name == "left-hand":
+ out["LeftHand"][frame_indices] = 1.0
+ elif constraint.name == "right-hand":
+ out["RightHand"][frame_indices] = 1.0
+ elif constraint.name == "root2d":
+ out["Root"][frame_indices] = 1.0
+ return out
+
+ # Create constraint masks from constraint_lst (one dict per batch item when batched)
+ batched_constraints = bool(constraint_lst) and isinstance(constraint_lst[0], list)
+ if batched_constraints:
+ constraint_masks_dict_lst = [_build_constraint_masks_dict(constraint_lst[b]) for b in range(batch_size)]
+ else:
+ constraint_masks_dict = (
+ _build_constraint_masks_dict(constraint_lst)
+ if constraint_lst
+ else {
+ key: torch.zeros(num_frames, dtype=torch.float32)
+ for key in [
+ "FullBody",
+ "LeftFoot",
+ "RightFoot",
+ "LeftHand",
+ "RightHand",
+ "Root",
+ ]
+ }
+ )
+
+ # Create working rig
+ above_ground_offset = 0.02 if isinstance(skeleton, (SOMASkeleton30, SOMASkeleton77)) else 0.007
+ # larger offset for SOMA since model tends to generate lower to the ground
+ working_rig = create_working_rig_from_skeleton(skeleton, above_ground_offset=above_ground_offset)
+ has_double_ankle_joints = isinstance(skeleton, G1Skeleton34)
+
+ # Prepare input tensors. The generated motion will be modified in place. Clone first.
+ neutral_joints_pelvis_offset = skeleton.neutral_joints[0].cpu().clone()
+ hip_translations_corrected = root_positions.cpu().clone()
+ rotations_corrected = matrix_to_quaternion(local_rot_mats).cpu().clone() # (B, T, J, 4)
+ contacts = contacts.cpu()
+
+ # Extract input motion (target keyframes) from constraints for each batch
+ # For constrained keyframes, use the original motion from constraints
+ # For non-constrained frames, zeros are used
+ hip_translations_input = torch.zeros(batch_size, num_frames, 3)
+ rotations_input = torch.zeros(batch_size, num_frames, num_joints, 4)
+ rotations_input[..., 0] = 1.0 # Initialize as identity quaternions (w=1, x=y=z=0)
+
+ if constraint_lst:
+ for b in range(batch_size):
+ # Get constraints for this batch item (if batched) or use the same list
+ constraints_lst_el = (
+ constraint_lst[b]
+ if isinstance(
+ constraint_lst[0], list
+ ) # when the constraint_list is in batch format, each item in a list is a constraintlist for one sample
+ else constraint_lst # single constraint list shared for all samples in the batch
+ )
+ hip_translations_input[b], rotations_input[b] = extract_input_motion_from_constraints(
+ constraints_lst_el,
+ skeleton,
+ num_frames,
+ num_joints,
+ )
+
+ # Call the motion correction for each batch (optional package)
+ try:
+ from motion_correction import motion_postprocess
+ except ImportError as e:
+ raise RuntimeError(
+ "Motion correction is required for this postprocessing path but the "
+ "motion_correction package is not installed. Install with: pip install -e ."
+ ) from e
+ for b in range(batch_size):
+ masks_b = constraint_masks_dict_lst[b] if batched_constraints else constraint_masks_dict
+ motion_postprocess.correct_motion(
+ hip_translations_corrected[b : b + 1],
+ rotations_corrected[b : b + 1],
+ contacts[b : b + 1],
+ hip_translations_input[b : b + 1],
+ rotations_input[b : b + 1],
+ masks_b,
+ contact_threshold,
+ root_margin,
+ working_rig,
+ has_double_ankle_joints,
+ )
+
+ local_rot_mats_corrected = quaternion_to_matrix(rotations_corrected)
+
+ # Compute posed joints using FK
+ device = local_rot_mats.device
+ global_rot_mats, posed_joints, _ = fk(
+ local_rot_mats_corrected.to(device),
+ hip_translations_corrected.to(device),
+ skeleton,
+ )
+
+ result = {
+ "local_rot_mats": local_rot_mats_corrected.to(device),
+ "root_positions": hip_translations_corrected.to(device),
+ "posed_joints": posed_joints,
+ "global_rot_mats": global_rot_mats,
+ }
+
+ return result
diff --git a/kimodo/sanitize.py b/kimodo/sanitize.py
new file mode 100644
index 0000000000000000000000000000000000000000..f3728e2fd1b3c3c949aa9d69a94db1570c58e230
--- /dev/null
+++ b/kimodo/sanitize.py
@@ -0,0 +1,95 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+"""Text prompt sanitization for motion generation (whitespace, punctuation, capitalization)."""
+
+
+def sanitize_text(text: str, paragraph: bool = True) -> str:
+ """Sanitize a text prompt: strip, collapse spaces, capitalize, trim non-alphanumeric, add/fix final punctuation.
+
+ Args:
+ text: Input text prompt.
+ paragraph: If True, capitalize after each sentence break and normalize spacing between sentences.
+
+ Returns:
+ Sanitized text.
+ """
+ # remove any trailing or leading whitespace
+ text = text.strip()
+
+ # https://stackoverflow.com/a/1546251
+ # replace duplicate spaces by one space
+ text = " ".join(text.split())
+
+ if text == "":
+ return text
+
+ # removing leading non alpha numeric characters
+ for i, c in enumerate(text):
+ if not str.isalnum(c):
+ continue
+ break
+ text = text[i:]
+
+ # Capitalize
+ text = text.capitalize()
+
+ final_punctuations = ".!?\"])'"
+ # removing trailing non alpha numeric characters
+ # expect final punctuations
+ for i, c in reversed(list(enumerate(text))):
+ if not str.isalnum(c) and c not in final_punctuations:
+ continue
+ break
+ text = text[: i + 1]
+
+ # Adding period at the end if needed
+ if text[-1] not in ".!?":
+ text = text + "."
+
+ if paragraph:
+ # fix end of sentences if several sentences
+ for sentence_break in ".!?":
+ subtexts = text.split(sentence_break)
+ text = f"{sentence_break} ".join( # put back a space after the break
+ [
+ y[0].capitalize() + y[1:] # only capitalize the first character
+ if y
+ else y # y is empty at the end
+ for x in subtexts
+ for y in [x.strip()] # remove extra spaces
+ ]
+ ).strip() # remove extra space at the end
+ return text
+
+
+def sanitize_texts(texts: list[str]) -> list[str]:
+ """Sanitize each text prompt in the list (see sanitize_text).
+
+ Args:
+ texts: List of input text prompts.
+
+ Returns:
+ List of sanitized texts.
+ """
+ return [sanitize_text(text) for text in texts]
+
+
+if __name__ == "__main__":
+ texts = [
+ " A person is walking.",
+ "someone go forward",
+ "jump",
+ "jumping!",
+ "jumping)",
+ "-go",
+ "blocasdji -----",
+ "",
+ ]
+
+ print("Old texts")
+ print("\n".join(texts))
+ print()
+
+ new_texts = sanitize_texts(texts)
+ print("Sanitized texts")
+ print("\n".join(new_texts))
diff --git a/kimodo/scripts/__init__.py b/kimodo/scripts/__init__.py
new file mode 100644
index 0000000000000000000000000000000000000000..e69de29bb2d1d6434b8b29ae775ad8c2e48c5391
diff --git a/kimodo/scripts/docker-entrypoint.sh b/kimodo/scripts/docker-entrypoint.sh
new file mode 100644
index 0000000000000000000000000000000000000000..db2246e558d0aad8a6b1c0c6241e429e85e7898d
--- /dev/null
+++ b/kimodo/scripts/docker-entrypoint.sh
@@ -0,0 +1,26 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+HOST_UID="${HOST_UID:-}"
+HOST_GID="${HOST_GID:-}"
+HOST_USER="${HOST_USER:-user}"
+
+if [[ -z "${HOST_UID}" || -z "${HOST_GID}" ]]; then
+ if [[ -d /workspace ]]; then
+ HOST_UID="$(stat -c %u /workspace)"
+ HOST_GID="$(stat -c %g /workspace)"
+ else
+ HOST_UID="${HOST_UID:-1000}"
+ HOST_GID="${HOST_GID:-1000}"
+ fi
+fi
+
+if ! getent group "${HOST_GID}" >/dev/null 2>&1; then
+ groupadd -g "${HOST_GID}" "${HOST_USER}"
+fi
+
+if ! getent passwd "${HOST_UID}" >/dev/null 2>&1; then
+ useradd -m -u "${HOST_UID}" -g "${HOST_GID}" -s /bin/bash "${HOST_USER}"
+fi
+
+exec gosu "${HOST_UID}:${HOST_GID}" "$@"
diff --git a/kimodo/scripts/generate.py b/kimodo/scripts/generate.py
new file mode 100644
index 0000000000000000000000000000000000000000..e74f951932de9af08f7d78fb574b1f24d4dd0891
--- /dev/null
+++ b/kimodo/scripts/generate.py
@@ -0,0 +1,502 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+import argparse
+import os
+import shutil
+from typing import Any, Dict, Optional
+
+import torch
+
+from kimodo import DEFAULT_MODEL, load_model
+from kimodo.constraints import load_constraints_lst
+from kimodo.exports.motion_io import save_kimodo_npz
+from kimodo.meta import load_prompts_from_meta
+from kimodo.model.cfg import CFG_TYPES
+from kimodo.model.registry import get_model_info
+from kimodo.tools import load_json, save_json, seed_everything
+
+
+def parse_args():
+ parser = argparse.ArgumentParser(description="Cmd line API for generation motions with kimodo")
+ parser.add_argument(
+ "prompt",
+ nargs="?",
+ type=str,
+ default=None,
+ help="Text prompt describing the motion to generate, or several prompts separated by periods.",
+ )
+ parser.add_argument(
+ "--model",
+ type=str,
+ default=DEFAULT_MODEL,
+ help="Name of the model (e.g. Kimodo-SOMA-RP-v1, etc).",
+ )
+ parser.add_argument(
+ "--duration",
+ type=str,
+ default="5.0",
+ help="Duration in seconds (default: 5.0). Separate by spaces in a string for different durations per prompts",
+ )
+ parser.add_argument(
+ "--num_samples",
+ type=int,
+ default=1,
+ help="Number of samples to generate (default: 1)",
+ )
+ parser.add_argument(
+ "--diffusion_steps",
+ type=int,
+ default=100,
+ help="Number of diffusion steps (default: 100)",
+ )
+ parser.add_argument(
+ "--num_transition_frames",
+ type=int,
+ default=5,
+ help="Number of frames to help transitioning (default: 5)",
+ )
+ parser.add_argument(
+ "--constraints",
+ type=str,
+ default=None,
+ help="Saved constraint list",
+ )
+ parser.add_argument(
+ "--output",
+ type=str,
+ default="output",
+ help="Output stem name: with one sample writes a single file per format (e.g. test.npz, test.csv); with multiple samples creates a folder and writes test_00.npz, test_01.npz, ... inside it. Used for NPZ, AMASS NPZ, CSV, and BVH.",
+ )
+ parser.add_argument(
+ "--save_example_dir",
+ action="store_true",
+ help=(
+ "Save demo-compatible example directories (each contains motion.npz, constraints.json, meta.json). "
+ "With one sample, writes