po03087's picture
EgoLM baseline code (Ego3DLM snapshot, unmodified) + upload notes
3de4238 verified
Raw History Blame Contribute Delete
10.1 kB
import numpy as np
import matplotlib.pyplot as plt
import matplotlib
from matplotlib.animation import FuncAnimation, FFMpegWriter
from mpl_toolkits.mplot3d import Axes3D # Required for 3D plot
import os
matplotlib.use('Agg')
xsens_t2m_kinematic_chain = [
[0, 1, 15, 16, 17, 18], # left leg
[0, 1, 19, 20, 21, 22], # right leg
[0, 2, 3, 4, 5, 6], # spine -> head
[4, 7, 8, 9, 10], # left arm
[4, 11, 12, 13, 14] # right arm
]
humanml3d_t2m_kinematic_chain = [[0, 2, 5, 8, 11], [0, 1, 4, 7, 10], [0, 3, 6, 9, 12, 15], [9, 14, 17, 19, 21], [9, 13, 16, 18, 20]]
R = np.array([
[0, 0, 1],
[1, 0, 0],
[0, 1, 0]
], dtype=np.float32)
## VISUALIZE ##
def visualize(out_path, joints, prefix_idx=0, length_segment=100):
t2m_kinematic_chain = xsens_t2m_kinematic_chain if joints.shape[1] == 23 else humanml3d_t2m_kinematic_chain
def update(frame_idx):
ax.cla() # clear current frame
frame = joints[frame_idx].copy() # shape: (23, 3)
xs, ys, zs = frame[:, 0], frame[:, 1], frame[:, 2]
ax.scatter(xs, ys, zs, c='b', s=2)
color = 'lightgray' if frame_idx < prefix_idx else 'black'
for chain in t2m_kinematic_chain:
for i in range(len(chain) - 1):
j1, j2 = chain[i], chain[i + 1]
ax.plot(
[frame[j1, 0], frame[j2, 0]],
[frame[j1, 1], frame[j2, 1]],
[frame[j1, 2], frame[j2, 2]],
color=color, linewidth=1
)
# for i, (x, y, z) in enumerate(frame):
# ax.text(x, y, z, str(i), color='red', fontsize=5)
ax.set_title(f'Frame {frame_idx}', fontsize=12)
ax.set_xlim(xmin, xmax)
ax.set_ylim(ymin, ymax)
ax.set_zlim(zmin, zmax)
ax.set_box_aspect((dx, dy, dz))
ax.set_xlabel('X')
ax.set_ylabel('Y')
ax.set_zlabel('Z')
joints = joints @ R.T
num_seg = joints.shape[0] // length_segment
if num_seg == 0 :
num_seg += 1
for i in range(num_seg) :
if i >= 1: continue
start_idx = i * length_segment
end_idx = min((i+1)*length_segment, joints.shape[0])
joints_segment = joints[start_idx:end_idx]
fig = plt.figure(figsize=(6, 6), dpi=150)
ax = fig.add_subplot(111, projection='3d')
all_joints = joints_segment.reshape(-1, 3)
xmin, xmax = all_joints[:, 0].min()-0.3, all_joints[:, 0].max()+0.3
ymin, ymax = all_joints[:, 1].min()-0.3, all_joints[:, 1].max()+0.3
zmin, zmax = all_joints[:, 2].min()-0.3, all_joints[:, 2].max()+0.3
dx = max(xmax - xmin, 1e-6)
dy = max(ymax - ymin, 1e-6)
dz = max(zmax - zmin, 1e-6)
# anim = FuncAnimation(fig, update_only_joints, frames=len(joints), interval=500)
anim = FuncAnimation(fig, update, frames=[i for i in range(start_idx, end_idx)], interval=500)
# anim.save(os.path.join(out_dir, f'{base_name}_ori_{cnt}_{i}.gif'), writer="pillow", fps=10)
anim.save(out_path, writer="pillow", fps=10)
plt.close(fig)
def visualize_3pts(out_path, joints, triads_T34=None, prefix_idx=0, length_segment=100):
t2m_kinematic_chain = xsens_t2m_kinematic_chain if joints.shape[1] == 23 else humanml3d_t2m_kinematic_chain
def undo_rotate_pose_matrix(T34_rotated, R):
out = T34_rotated.copy()
out[..., :3, :3] = (R[None, ...] @ out[..., :3, :3]) # inverse of (R.T @ ·) is (R @ ·)
out[..., :3, 3] = out[..., :3, 3] @ R.T # inverse of (· @ R) is (· @ R.T)
return out
def draw_triad(ax, R, t, axis_len=0.05, lw=1.5, alpha=0.95):
"""R: (3,3), t: (3,)"""
o = t.reshape(3)
x, y, z = R[:, 0], R[:, 1], R[:, 2]
# X(red), Y(green), Z(blue)
ax.quiver(*o, *x, length=axis_len, linewidth=lw, color='r', alpha=alpha)
ax.quiver(*o, *y, length=axis_len, linewidth=lw, color='g', alpha=alpha)
ax.quiver(*o, *z, length=axis_len, linewidth=lw, color='b', alpha=alpha)
def update(frame_idx):
ax.cla() # clear current frame
frame = joints[frame_idx].copy() # shape: (23, 3)
xs, ys, zs = frame[:, 0], frame[:, 1], frame[:, 2]
ax.scatter(xs, ys, zs, c='b', s=2)
color = 'lightgray' if frame_idx < prefix_idx else 'black'
for chain in t2m_kinematic_chain:
for i in range(len(chain) - 1):
j1, j2 = chain[i], chain[i + 1]
ax.plot(
[frame[j1, 0], frame[j2, 0]],
[frame[j1, 1], frame[j2, 1]],
[frame[j1, 2], frame[j2, 2]],
color=color, linewidth=1
)
if triads_T34 is not None and frame_idx < triads_T34.shape[0]:
for j in range(triads_T34.shape[1]):
M = triads_T34[frame_idx, j]
Rloc = M[:3, :3] #orthonormalize(M[:3, :3])
tloc = M[:3, 3]
axis_len = max(dx, dy, dz) * 0.04
draw_triad(ax, Rloc, tloc, axis_len=axis_len, lw=1.4, alpha=0.95)
ax.set_title(f'Frame {frame_idx}', fontsize=12)
ax.set_xlim(xmin, xmax)
ax.set_ylim(ymin, ymax)
ax.set_zlim(zmin, zmax)
ax.set_box_aspect((dx, dy, dz))
ax.set_xlabel('X')
ax.set_ylabel('Y')
ax.set_zlabel('Z')
joints = joints @ R.T
if triads_T34 is not None :
triads_T34=undo_rotate_pose_matrix(triads_T34, R)
num_seg = joints.shape[0] // length_segment
if num_seg == 0 :
num_seg += 1
for i in range(num_seg) :
if i >= 1: continue
start_idx = i * length_segment
end_idx = min((i+1)*length_segment, joints.shape[0])
joints_segment = joints[start_idx:end_idx]
fig = plt.figure(figsize=(6, 6), dpi=150)
ax = fig.add_subplot(111, projection='3d')
all_joints = joints_segment.reshape(-1, 3)
xmin, xmax = all_joints[:, 0].min()-0.3, all_joints[:, 0].max()+0.3
ymin, ymax = all_joints[:, 1].min()-0.3, all_joints[:, 1].max()+0.3
zmin, zmax = all_joints[:, 2].min()-0.3, all_joints[:, 2].max()+0.3
dx = max(xmax - xmin, 1e-6)
dy = max(ymax - ymin, 1e-6)
dz = max(zmax - zmin, 1e-6)
# anim = FuncAnimation(fig, update_only_joints, frames=len(joints), interval=500)
anim = FuncAnimation(fig, update, frames=[i for i in range(start_idx, end_idx)], interval=500)
# anim.save(os.path.join(out_dir, f'{base_name}_ori_{cnt}_{i}.gif'), writer="pillow", fps=10)
anim.save(out_path, writer="pillow", fps=10)
plt.close(fig)
def visualize_two(out_path, joints_ref, joints_rst, prefix_idx=0, length_segment=50):
"""
joints1, joints2: (T, J, 3) numpy arrays (e.g. J=23 or 22/24, etc.)
joints1 in black, joints2 in light gray; if the two sequences differ in length, animate to the shorter one.
"""
joints1 = joints_ref
joints2 = joints_rst
assert joints1.ndim == 3 and joints1.shape[-1] == 3, "joints1 shape must be (T, J, 3)"
assert joints2.ndim == 3 and joints2.shape[-1] == 3, "joints2 shape must be (T, J, 3)"
T = min(joints1.shape[0], joints2.shape[0])
J1, J2 = joints1.shape[1], joints2.shape[1]
if J1 == 23 and J2 == 23:
t2m_kinematic_chain = xsens_t2m_kinematic_chain
else:
NotImplementedError()
joints1 = joints1 @ R.T
joints2 = joints2 @ R.T
num_seg = T // length_segment
if num_seg == 0:
num_seg = 1
for seg_idx in range(num_seg):
if seg_idx >= 1: continue
start_idx = seg_idx * length_segment
end_idx = min((seg_idx + 1) * length_segment, T)
seg1 = joints1[start_idx:end_idx] # (t, J1, 3)
seg2 = joints2[start_idx:end_idx] # (t, J2, 3)
fig = plt.figure(figsize=(6, 6), dpi=150)
ax = fig.add_subplot(111, projection='3d')
all_joints = np.concatenate([seg1.reshape(-1, 3), seg2.reshape(-1, 3)], axis=0)
xmin, xmax = all_joints[:, 0].min() - 0.3, all_joints[:, 0].max() + 0.3
ymin, ymax = all_joints[:, 1].min() - 0.3, all_joints[:, 1].max() + 0.3
zmin, zmax = all_joints[:, 2].min() - 0.3, all_joints[:, 2].max() + 0.3
dx = max(xmax - xmin, 1e-6)
dy = max(ymax - ymin, 1e-6)
dz = max(zmax - zmin, 1e-6)
def plot_skeleton(frame, color_lines, size_points=2, lw=1):
xs, ys, zs = frame[:, 0], frame[:, 1], frame[:, 2]
ax.scatter(xs, ys, zs, c='b', s=size_points)
for chain in t2m_kinematic_chain:
for i in range(len(chain) - 1):
j1, j2 = chain[i], chain[i + 1]
if j1 < frame.shape[0] and j2 < frame.shape[0]:
ax.plot(
[frame[j1, 0], frame[j2, 0]],
[frame[j1, 1], frame[j2, 1]],
[frame[j1, 2], frame[j2, 2]],
color=color_lines, linewidth=lw
)
def update(frame_idx):
ax.cla()
f1 = joints1[frame_idx]
f2 = joints2[frame_idx]
plot_skeleton(f1, color_lines='black', lw=1)
plot_skeleton(f2, color_lines='lightgray', lw=1)
ax.set_title(f'Frame {frame_idx}', fontsize=12)
ax.set_xlim(xmin, xmax)
ax.set_ylim(ymin, ymax)
ax.set_zlim(zmin, zmax)
ax.set_box_aspect((dx, dy, dz))
ax.set_xlabel('X')
ax.set_ylabel('Y')
ax.set_zlabel('Z')
frames = list(range(start_idx, end_idx))
anim = FuncAnimation(fig, update, frames=frames, interval=500)
save_path = out_path if num_seg == 1 else out_path.replace(".gif", f"_seg{seg_idx}.gif")
anim.save(save_path, writer="pillow", fps=10)
plt.close(fig)