"""PoseShield: Neural Collision Fields for Human Self-Collision Resolution. A Gradio demo that takes a colliding SMPL-H pose and resolves self-collisions using the PoseShield neural collision field with SLSQP optimization. """ import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import spaces # MUST come before torch import torch import torch.nn as nn import torch.nn.functional as F import numpy as np import pickle import tempfile import time import yaml import struct import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D from scipy.optimize import minimize from huggingface_hub import hf_hub_download import smplx import gradio as gr # --------------------------------------------------------------------------- # Rotation utilities (from poseshield.common.utils) # --------------------------------------------------------------------------- def normalize(x, axis=-1, eps=1e-8): norm = np.linalg.norm(x, axis=axis, keepdims=True) + eps return x / norm def axis_angle_to_matrix(axis_angle): """Convert (N, 3) axis-angle to (N, 3, 3) rotation matrices (Rodrigues).""" aa = np.asarray(axis_angle, dtype=np.float64) N = aa.shape[0] theta = np.linalg.norm(aa, axis=1, keepdims=True) eps = 1e-8 k = aa / (theta + eps) kx, ky, kz = k[:, 0], k[:, 1], k[:, 2] K = np.zeros((N, 3, 3), dtype=np.float64) K[:, 0, 1] = -kz K[:, 0, 2] = ky K[:, 1, 0] = kz K[:, 1, 2] = -kx K[:, 2, 0] = -ky K[:, 2, 1] = kx I = np.eye(3, dtype=np.float64)[None, :, :] sin_t = np.sin(theta)[:, None].reshape(N, 1, 1) cos_t = np.cos(theta)[:, None].reshape(N, 1, 1) K2 = K @ K small = (theta.reshape(N) < 1e-4) A = np.empty((N, 1, 1), dtype=np.float64) B = np.empty((N, 1, 1), dtype=np.float64) A[~small] = sin_t[~small] B[~small] = (1.0 - cos_t[~small]) th = theta.reshape(N, 1, 1) A[small] = th[small] - (th[small]**3) / 6.0 B[small] = (th[small]**2) / 2.0 - (th[small]**4) / 24.0 R = I + A * K + B * K2 return R def matrix_to_axis_angle(R): """Convert (N, 3, 3) rotation matrices to (N, 3) axis-angle vectors.""" trace = np.trace(R, axis1=1, axis2=2) trace = np.clip(trace, -1.0, 3.0) angles = np.arccos((trace - 1.0) / 2.0) rx = R[:, 2, 1] - R[:, 1, 2] ry = R[:, 0, 2] - R[:, 2, 0] rz = R[:, 1, 0] - R[:, 0, 1] axes = np.stack([rx, ry, rz], axis=1) sin_angles = np.linalg.norm(axes, axis=1, keepdims=True) / 2.0 axes = axes / (2.0 * (sin_angles + 1e-8)) axis_angle = axes * angles[:, None] return axis_angle def rotation_6d_to_matrix(d6): """Convert (..., 6) 6D rotation to (..., 3, 3) rotation matrix (Gram-Schmidt).""" a1 = d6[..., :3] a2 = d6[..., 3:] b1 = normalize(a1, axis=-1) dot = np.sum(b1 * a2, axis=-1, keepdims=True) b2 = a2 - dot * b1 b2 = normalize(b2, axis=-1) b3 = np.cross(b1, b2, axis=-1) rotation_mats = np.stack((b1, b2, b3), axis=-2) return rotation_mats def matrix_to_rotation_6d(R): """Convert (..., 3, 3) rotation matrix to (..., 6) 6D rotation.""" b1 = R[..., 0, :] b2 = R[..., 1, :] d6 = np.concatenate([b1, b2], axis=-1) return d6 def rotation_6d_to_matrix_torch(d6): """Torch differentiable version of rotation_6d_to_matrix.""" a1, a2 = d6[..., :3], d6[..., 3:] b1 = F.normalize(a1, dim=-1) b2 = a2 - (b1 * a2).sum(dim=-1, keepdim=True) * b1 b2 = F.normalize(b2, dim=-1) b3 = torch.cross(b1, b2, dim=-1) return torch.stack((b1, b2, b3), dim=-2) def matrix_to_axis_angle_torch(R): """Torch differentiable version of matrix_to_axis_angle.""" trace = R[:, 0, 0] + R[:, 1, 1] + R[:, 2, 2] cos_angle = ((trace - 1.0) / 2.0).clamp(-1.0 + 1e-7, 1.0 - 1e-7) angle = torch.acos(cos_angle) rx = R[:, 2, 1] - R[:, 1, 2] ry = R[:, 0, 2] - R[:, 2, 0] rz = R[:, 1, 0] - R[:, 0, 1] axis_raw = torch.stack([rx, ry, rz], dim=1) safe_sin = torch.sin(angle).abs().clamp(min=1e-7).unsqueeze(1) unit_axis = axis_raw / (2.0 * safe_sin) return unit_axis * angle.unsqueeze(1) # --------------------------------------------------------------------------- # PoseShield model (from poseshield.common.network) # --------------------------------------------------------------------------- class ResidualMLP(nn.Module): """Residual MLP: 21x6 joint rotations -> scalar collision field value.""" def __init__(self, in_dim=126, hidden_dim=512, num_layers=12, activation="relu"): super().__init__() self.input_layer = nn.Linear(in_dim, hidden_dim) self.hidden_layers = nn.ModuleList([ nn.Linear(hidden_dim, hidden_dim) for _ in range(num_layers) ]) if activation == "relu": self.act = nn.ReLU() elif activation == "leaky_relu": self.act = nn.LeakyReLU() elif activation == "elu": self.act = nn.ELU() else: raise ValueError(f"Unsupported activation: {activation}") self.output_layer = nn.Linear(hidden_dim, 1) def forward(self, x): bs = x.shape[0] x_reshaped = x.reshape(-1, 6) x_raw = x_reshaped[:, :3] x_norm = x_raw / x_raw.norm(dim=1, keepdim=True) y_raw = x_reshaped[:, 3:] dot = (x_norm * y_raw).sum(dim=1, keepdim=True) y_perp = y_raw - dot * x_norm y_norm = y_perp / y_perp.norm(dim=1, keepdim=True) x_valid = torch.cat([x_norm, y_norm], dim=1).reshape(bs, -1) x = self.act(self.input_layer(x_valid)) for layer in self.hidden_layers: x = self.act(layer(x)) + x return self.output_layer(x) # --------------------------------------------------------------------------- # Cost & constraint functions (from poseshield.pose.utils) # --------------------------------------------------------------------------- _SUBTREE_SIZES = [4, 4, 13, 3, 3, 12, 2, 2, 11, 1, 1, 2, 4, 4, 1, 3, 3, 2, 2, 1, 1] SMPLH_POSE_WEIGHTS = torch.tensor(_SUBTREE_SIZES, dtype=torch.float32) SMPLH_POSE_WEIGHTS /= SMPLH_POSE_WEIGHTS.sum() def constraint_function(model, x): """Collision field value for a pose. Positive = collision-free.""" output = model(x.unsqueeze(0)) return output.squeeze(0).squeeze(0) def cost_function_weighted(x, x_ref, weights=None): """Weighted L2 pose distance preserving kinematic-chain importance.""" if weights is None: weights = SMPLH_POSE_WEIGHTS.to(x.device) diff = (x - x_ref).reshape(21, 6) per_joint_norm = torch.linalg.norm(diff, dim=-1) return (per_joint_norm * weights).sum() def cost_function(x, x_ref): return cost_function_weighted(x, x_ref) # --------------------------------------------------------------------------- # SLSQP optimizer (from poseshield.pose.resolve_slsqp) # --------------------------------------------------------------------------- def optimize_slsqp(sample, model, device, max_itr=300, threshold=0.1, cost_type="normal", tol=0.03): """SLSQP optimization to resolve collisions while preserving pose.""" x0 = sample.reshape(-1).astype(np.float64) x_ref_np = x0.copy() def to_torch(x_np, requires_grad=False): return torch.tensor(x_np, dtype=torch.float32, device=device, requires_grad=requires_grad) x_ref_t = to_torch(x_ref_np, requires_grad=False) def cost_fn_np(x_np): x_t = to_torch(x_np) val = cost_function(x_t, x_ref_t) if cost_type != "weighted" else cost_function_weighted(x_t, x_ref_t) return float(val.detach().cpu().item()) def cost_fn_jac(x_np): x_t = to_torch(x_np, requires_grad=True) val = cost_function(x_t, x_ref_t) if cost_type != "weighted" else cost_function_weighted(x_t, x_ref_t) grad = torch.autograd.grad(val, x_t)[0] return grad.detach().cpu().numpy().astype(np.float64) def cons_ineq_fun(x_np): x_t = to_torch(x_np) val = constraint_function(model, x_t) - threshold return float(val.detach().cpu().item()) def cons_ineq_jac(x_np): x_t = to_torch(x_np, requires_grad=True) cons_val = constraint_function(model, x_t) grad = torch.autograd.grad(cons_val, x_t)[0] return grad.detach().cpu().numpy().astype(np.float64) def view6(x_t): return x_t.view(-1, 6) def ineq_r1_upper_fun(x_np): x_t = to_torch(x_np) r1 = view6(x_t)[:, :3] return float((1.0 + tol - r1.norm(dim=1).mean()).detach().cpu().item()) def ineq_r1_upper_jac(x_np): x_t = to_torch(x_np, requires_grad=True) r1 = view6(x_t)[:, :3] val = 1.0 + tol - r1.norm(dim=1).mean() grad = torch.autograd.grad(val, x_t)[0] return grad.detach().cpu().numpy().astype(np.float64) def ineq_r1_lower_fun(x_np): x_t = to_torch(x_np) r1 = view6(x_t)[:, :3] return float((r1.norm(dim=1).mean() - (1.0 - tol)).detach().cpu().item()) def ineq_r1_lower_jac(x_np): x_t = to_torch(x_np, requires_grad=True) r1 = view6(x_t)[:, :3] val = r1.norm(dim=1).mean() - (1.0 - tol) grad = torch.autograd.grad(val, x_t)[0] return grad.detach().cpu().numpy().astype(np.float64) def ineq_r2_upper_fun(x_np): x_t = to_torch(x_np) r2 = view6(x_t)[:, 3:] return float((1.0 + tol - r2.norm(dim=1).mean()).detach().cpu().item()) def ineq_r2_upper_jac(x_np): x_t = to_torch(x_np, requires_grad=True) r2 = view6(x_t)[:, 3:] val = 1.0 + tol - r2.norm(dim=1).mean() grad = torch.autograd.grad(val, x_t)[0] return grad.detach().cpu().numpy().astype(np.float64) def ineq_r2_lower_fun(x_np): x_t = to_torch(x_np) r2 = view6(x_t)[:, 3:] return float((r2.norm(dim=1).mean() - (1.0 - tol)).detach().cpu().item()) def ineq_r2_lower_jac(x_np): x_t = to_torch(x_np, requires_grad=True) r2 = view6(x_t)[:, 3:] val = r2.norm(dim=1).mean() - (1.0 - tol) grad = torch.autograd.grad(val, x_t)[0] return grad.detach().cpu().numpy().astype(np.float64) def ineq_orth_upper_fun(x_np): x_t = to_torch(x_np) r1, r2 = view6(x_t).split(3, dim=1) dot_mean = (r1 * r2).sum(dim=1).mean() return float((tol - dot_mean).detach().cpu().item()) def ineq_orth_upper_jac(x_np): x_t = to_torch(x_np, requires_grad=True) r1, r2 = view6(x_t).split(3, dim=1) val = tol - (r1 * r2).sum(dim=1).mean() grad = torch.autograd.grad(val, x_t)[0] return grad.detach().cpu().numpy().astype(np.float64) def ineq_orth_lower_fun(x_np): x_t = to_torch(x_np) r1, r2 = view6(x_t).split(3, dim=1) dot_mean = (r1 * r2).sum(dim=1).mean() return float((tol + dot_mean).detach().cpu().item()) def ineq_orth_lower_jac(x_np): x_t = to_torch(x_np, requires_grad=True) r1, r2 = view6(x_t).split(3, dim=1) val = tol + (r1 * r2).sum(dim=1).mean() grad = torch.autograd.grad(val, x_t)[0] return grad.detach().cpu().numpy().astype(np.float64) constraints = [ {"type": "ineq", "fun": cons_ineq_fun, "jac": cons_ineq_jac}, {"type": "ineq", "fun": ineq_r1_upper_fun, "jac": ineq_r1_upper_jac}, {"type": "ineq", "fun": ineq_r1_lower_fun, "jac": ineq_r1_lower_jac}, {"type": "ineq", "fun": ineq_r2_upper_fun, "jac": ineq_r2_upper_jac}, {"type": "ineq", "fun": ineq_r2_lower_fun, "jac": ineq_r2_lower_jac}, {"type": "ineq", "fun": ineq_orth_upper_fun, "jac": ineq_orth_upper_jac}, {"type": "ineq", "fun": ineq_orth_lower_fun, "jac": ineq_orth_lower_jac}, ] loss_history, cons_history = [], [] def callback(xk): try: loss_history.append(cost_fn_np(xk)) cons_history.append(constraint_function(model, to_torch(xk)).detach().cpu().item()) except Exception: pass res = minimize( cost_fn_np, x0, method="SLSQP", jac=cost_fn_jac, constraints=constraints, bounds=None, options={"maxiter": max_itr, "ftol": 1e-6, "disp": False}, callback=callback, ) x_opt = res.x.astype(np.float32) if len(loss_history) == 0: loss_history.append(cost_fn_np(x_opt)) if len(cons_history) == 0: cons_history.append(constraint_function(model, to_torch(x_opt)).detach().cpu().item()) return x_opt, loss_history, cons_history, bool(res.success), str(res.message) # --------------------------------------------------------------------------- # Mesh visualization # --------------------------------------------------------------------------- def pose_to_mesh(smpl_model, r_6d, device): """Convert 21x6 6D rotation to SMPL mesh vertices and faces.""" rot_mats = rotation_6d_to_matrix(r_6d) axis_angles = matrix_to_axis_angle(rot_mats) body_pose = torch.from_numpy(axis_angles.reshape(1, -1)).float().to(device) output = smpl_model( global_orient=None, body_pose=body_pose, betas=None, transl=None, return_verts=True, ) vertices = output.vertices[0].detach().cpu().numpy() faces = smpl_model.faces.astype(np.int32) return vertices, faces def visualize_smpl(vertices, faces, save_path, color="#ff7675"): """Render SMPL mesh to PNG using matplotlib.""" fig = plt.figure(figsize=(8, 6), facecolor="white") ax = fig.add_subplot(111, projection="3d") ax.set_facecolor("white") x_plt = vertices[:, 0] y_plt = vertices[:, 2] z_plt = vertices[:, 1] ax.plot_trisurf(x_plt, y_plt, z_plt, triangles=faces, shade=True, color=color, edgecolor="none", alpha=0.9) all_coords = np.stack([x_plt, y_plt, z_plt], axis=-1) min_vals = np.min(all_coords, axis=0) max_vals = np.max(all_coords, axis=0) ranges = max_vals - min_vals max_range = max(ranges) mid = (max_vals + min_vals) / 2 ax.set_xlim(mid[0] - max_range / 2, mid[0] + max_range / 2) ax.set_ylim(mid[1] - max_range / 2, mid[1] + max_range / 2) ax.set_zlim(mid[2] - max_range / 2, mid[2] + max_range / 2) ax.view_init(elev=15, azim=90) ax.axis("off") ax.grid(False) plt.tight_layout() plt.savefig(save_path, dpi=150, bbox_inches="tight", facecolor="white") plt.close(fig) # --------------------------------------------------------------------------- # Model loading at module scope # --------------------------------------------------------------------------- # Download PoseShield model from HF _model_dir = hf_hub_download("ZYYY99/PoseShield", "model.pth", repo_type="model") _config_path = hf_hub_download("ZYYY99/PoseShield", "config.yaml", repo_type="model") with open(_config_path, "r") as f: _config = yaml.safe_load(f) MODEL_HIDDEN_DIM = _config["MODEL"]["HIDDEN_DIM"] MODEL_NUM_LAYERS = _config["MODEL"]["NUM_LAYERS"] MODEL_ACTIVATION = _config["MODEL"].get("ACTIVATION", "relu") # Load collision field model collision_model = ResidualMLP( in_dim=126, hidden_dim=MODEL_HIDDEN_DIM, num_layers=MODEL_NUM_LAYERS, activation=MODEL_ACTIVATION, ).to("cuda") _ckpt = torch.load(_model_dir, map_location="cuda", weights_only=True) collision_model.load_state_dict(_ckpt) collision_model.eval() print(f"PoseShield collision field loaded: hidden_dim={MODEL_HIDDEN_DIM}, layers={MODEL_NUM_LAYERS}") # Download SMPL-H neutral body model and set up directory structure for smplx # The community SMPLH model lacks hand PCA components, so we add dummy ones. _smplh_file = hf_hub_download("Tevior/smplh", "neutral/model.npz", repo_type="model") _smplh_root = tempfile.mkdtemp() _smplh_target_dir = os.path.join(_smplh_root, "smplh") os.makedirs(_smplh_target_dir, exist_ok=True) # Load the npz, add missing hand component keys, and re-save for smplx _orig_smplh = dict(np.load(_smplh_file, allow_pickle=True)) _orig_smplh["hands_componentsl"] = np.zeros((1, 45), dtype=np.float32) _orig_smplh["hands_componentsr"] = np.zeros((1, 45), dtype=np.float32) _orig_smplh["hands_meanl"] = np.zeros((45,), dtype=np.float32) _orig_smplh["hands_meanr"] = np.zeros((45,), dtype=np.float32) np.savez(os.path.join(_smplh_target_dir, "SMPLH_NEUTRAL.npz"), **_orig_smplh) _smpl_model = smplx.create( _smplh_root, model_type="smplh", gender="neutral", ext="npz", use_pca=False, ).to("cuda") print("SMPL-H neutral body model loaded") # --------------------------------------------------------------------------- # Pre-bundled example poses (from the PoseShield demo_asset directory) # --------------------------------------------------------------------------- EXAMPLE_POSES = { "Colliding Pose #210": "x_ori_210.pkl", "Colliding Pose #408": "x_ori_408.pkl", "Colliding Pose #436": "x_ori_436.pkl", } # The example .pkl files are bundled in the Space repo root (same dir as app.py) _APP_DIR = os.path.dirname(os.path.abspath(__file__)) _example_files = {} for name, fname in EXAMPLE_POSES.items(): local_path = os.path.join(_APP_DIR, fname) if os.path.exists(local_path): _example_files[name] = local_path print(f"Loaded example pose: {name} -> {local_path}") else: # Fallback: download from the Space repo _path = hf_hub_download("hugging-apps/poseshield-collision-fix", fname, repo_type="space") _example_files[name] = _path print(f"Downloaded example pose: {name} -> {_path}") def load_pose_pickle(path): """Load a pose from a pickle file (format: dict with 'pose' key, shape (63,)).""" with open(path, "rb") as f: data = pickle.load(f, encoding="latin1") pose_aa = data["pose"].reshape(21, 3) rot_mat = axis_angle_to_matrix(pose_aa) rot_6d = matrix_to_rotation_6d(rot_mat) sample_flat = rot_6d.reshape(-1) return sample_flat, data # --------------------------------------------------------------------------- # Inference # --------------------------------------------------------------------------- @spaces.GPU(duration=60) def resolve_pose(pose_file, threshold=0.1, max_itr=150, progress=gr.Progress()): """Resolve self-collisions in a SMPL-H pose using the PoseShield collision field. Args: pose_file: A .pkl file containing a SMPL-H pose with a 'pose' key (21 joints x 3 axis-angle). threshold: Collision field threshold (lower = stricter constraint). max_itr: Maximum SLSQP optimization iterations. """ if pose_file is None: return None, None, "Please upload a pose file or select an example.", "" progress(0.1, desc="Loading pose...") sample_flat, raw_data = load_pose_pickle(pose_file) device = torch.device("cuda") progress(0.2, desc="Computing initial collision score...") with torch.no_grad(): init_val = constraint_function( collision_model, torch.from_numpy(sample_flat).float().to(device), ).item() progress(0.3, desc="Running SLSQP optimization...") start_time = time.time() optimized_x, loss_hist, cons_hist, success, message = optimize_slsqp( sample_flat, collision_model, device, max_itr=max_itr, threshold=threshold, ) elapsed = time.time() - start_time progress(0.7, desc="Computing final collision score...") final_val = constraint_function( collision_model, torch.from_numpy(optimized_x).float().to(device), ).item() final_error = cost_function( torch.from_numpy(optimized_x).to(device), torch.from_numpy(sample_flat).float().to(device), ).item() constraint_satisfied = final_val >= threshold progress(0.85, desc="Rendering meshes...") # Render before and after meshes with tempfile.NamedTemporaryFile(suffix="_before.png", delete=False) as f_before: before_path = f_before.name with tempfile.NamedTemporaryFile(suffix="_after.png", delete=False) as f_after: after_path = f_after.name r_6d_orig = sample_flat.reshape(21, 6) r_6d_opt = optimized_x.reshape(21, 6) verts_orig, faces = pose_to_mesh(_smpl_model, r_6d_orig, device) visualize_smpl(verts_orig, faces, before_path, color="#ff7675") verts_opt, _ = pose_to_mesh(_smpl_model, r_6d_opt, device) visualize_smpl(verts_opt, faces, after_path, color="#2ed573") progress(1.0, desc="Done!") status_text = ( f"**Optimization Results**\n" f"- Time: {elapsed:.2f}s\n" f"- Solver success: {success}\n" f"- Solver message: {message}\n" f"- Initial collision score: {init_val:.6f}\n" f"- Final collision score: {final_val:.6f} (threshold: {threshold})\n" f"- Constraint satisfied: {constraint_satisfied}\n" f"- Mean Vertex Deviation: {final_error:.6f}\n" f"- Iterations: {len(loss_hist)}" ) return before_path, after_path, status_text, f"{elapsed:.2f}s" # --------------------------------------------------------------------------- # Gradio UI # --------------------------------------------------------------------------- CSS = """ #col-container { max-width: 1100px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ with gr.Blocks(theme=gr.themes.Citrus(), css=CSS) as demo: gr.Markdown(""" # PoseShield: Neural Collision Fields for Human Self-Collision Resolution Upload a colliding SMPL-H pose (`.pkl` with a `pose` key of 21×3 axis-angle rotations) or try one of the example poses below. PoseShield resolves self-collisions using a learned neural collision field as a differentiable constraint. [Paper](https://arxiv.org/abs/2606.29686) | [GitHub](https://github.com/lzhyu/PoseShield) | [Model](https://huggingface.co/ZYYY99/PoseShield) """) with gr.Row(): with gr.Column(): pose_input = gr.File(label="Upload SMPL-H Pose (.pkl)", file_types=[".pkl"]) with gr.Accordion("Advanced Settings", open=False): threshold = gr.Slider(0.01, 1.0, value=0.1, step=0.01, label="Constraint Threshold (lower = stricter)") max_itr = gr.Slider(50, 300, value=150, step=10, label="Max Optimization Iterations") run_btn = gr.Button("Resolve Collisions", variant="primary") with gr.Column(): before_img = gr.Image(label="Before (Colliding)", type="filepath") after_img = gr.Image(label="After (Resolved)", type="filepath") status_output = gr.Markdown(label="Status") time_output = gr.Textbox(label="Inference Time", visible=True) gr.Examples( examples=[ [_example_files["Colliding Pose #210"]], [_example_files["Colliding Pose #408"]], [_example_files["Colliding Pose #436"]], ], inputs=[pose_input], outputs=[before_img, after_img, status_output, time_output], fn=resolve_pose, cache_examples=True, cache_mode="lazy", ) run_btn.click( fn=resolve_pose, inputs=[pose_input, threshold, max_itr], outputs=[before_img, after_img, status_output, time_output], api_name="resolve", ) demo.launch(mcp_server=True)