Diffusers
Safetensors
AR / test.py
xfcghj's picture
Upload folder using huggingface_hub
394918c verified
Raw
History Blame Contribute Delete
11.5 kB
import os
import argparse
import io
import pickle
import tarfile
import numpy as np
import torch
import torch.nn as nn
import trimesh
import zstandard
from scipy.spatial import KDTree
from tqdm import tqdm
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from models.auto import StackedTrajectorySpatialTextRefiner
from models.vae import TripoSGVaeWrapper
from models.clip import FrozenCLIPTextEncoder
from triposg.inference_utils import hierarchical_extract_geometry
class TextProjectionLayer(nn.Module):
def __init__(self, in_dim=512, out_dim=64):
super().__init__()
self.net = nn.Sequential(
nn.Linear(in_dim, in_dim),
nn.ReLU(),
nn.Linear(in_dim, out_dim),
)
self.residual = nn.Linear(in_dim, out_dim) if in_dim != out_dim else nn.Identity()
def forward(self, x):
return self.net(x) + self.residual(x)
def get_full_trajectories(archive_file, limit=20, min_frames=16):
data_list = []
with open(archive_file, "rb") as fh:
dctx = zstandard.ZstdDecompressor()
with dctx.stream_reader(fh) as reader:
with tarfile.open(fileobj=reader, mode="r|") as tar:
for member in tar:
if len(data_list) >= limit:
break
if not member.isfile():
continue
f = tar.extractfile(member)
if f is None:
continue
data = pickle.load(io.BytesIO(f.read()))
if "vertices" in data and "faces" in data and data["vertices"].shape[0] >= min_frames:
data_list.append(data)
return data_list
def save_mesh_as_png(mesh, save_path):
fig = plt.figure(figsize=(6, 6))
ax = fig.add_subplot(111, projection="3d")
vertices = np.asarray(mesh.vertices)
faces = np.asarray(mesh.faces)
ax.plot_trisurf(
vertices[:, 0],
vertices[:, 1],
vertices[:, 2],
triangles=faces,
cmap="viridis",
edgecolor="none",
alpha=0.85,
)
max_range = np.array(
[
vertices[:, 0].max() - vertices[:, 0].min(),
vertices[:, 1].max() - vertices[:, 1].min(),
vertices[:, 2].max() - vertices[:, 2].min(),
]
).max() / 2.0
center = vertices.mean(axis=0)
ax.set_xlim(center[0] - max_range, center[0] + max_range)
ax.set_ylim(center[1] - max_range, center[1] + max_range)
ax.set_zlim(center[2] - max_range, center[2] + max_range)
ax.view_init(elev=20, azim=45)
plt.axis("off")
plt.savefig(save_path, bbox_inches="tight", pad_inches=0, dpi=150)
plt.close(fig)
def to_tensor(x, device, dtype=None):
if isinstance(x, torch.Tensor):
x = x.detach().clone()
return x.to(device=device, dtype=dtype) if dtype is not None else x.to(device=device)
return torch.as_tensor(x, device=device, dtype=dtype)
@torch.no_grad()
def encode_sequence_latents(vae, vertices_seq, faces, device, num_frames, num_samples):
latents = []
f_t = to_tensor(faces, device=device, dtype=torch.float32)
for frame_idx in range(num_frames):
v_t = to_tensor(vertices_seq[frame_idx], device=device, dtype=torch.float32)
latent = vae.encode_mesh(v_t, f_t, num_samples=num_samples)
if isinstance(latent, np.ndarray):
latent = torch.from_numpy(latent).to(device)
latents.append(latent.squeeze(0))
return torch.stack(latents, dim=0)
@torch.no_grad()
def autoregressive_predict(auto_model, initial_latent, text_embed, num_future_frames, num_tokens, device):
current_context = initial_latent.unsqueeze(0).unsqueeze(0).to(device)
shortest_path_matrix = torch.zeros((1, num_tokens, num_tokens), device=device, dtype=torch.long)
pred_latents = [current_context[:, 0:1]]
for _ in range(num_future_frames):
pred_seq = auto_model(
current_context,
text_embed=text_embed,
shortest_path_matrix=shortest_path_matrix,
)
next_latent = pred_seq[:, -1:, :, :]
pred_latents.append(next_latent)
current_context = torch.cat([current_context, next_latent], dim=1)
return pred_latents
def compute_mesh_metrics(mesh_recon, mesh_gt, num_points=2048, tau=0.05):
pts_gt, _ = trimesh.sample.sample_surface(mesh_gt, num_points)
pts_recon, _ = trimesh.sample.sample_surface(mesh_recon, num_points)
tree_gt = KDTree(pts_gt)
tree_recon = KDTree(pts_recon)
d_r2r = np.mean(tree_gt.query(pts_recon, k=1)[0])
d_raw2r = np.mean(tree_recon.query(pts_gt, k=1)[0])
cd = d_r2r + d_raw2r
precision = np.mean(tree_gt.query(pts_recon, k=1)[0] < tau)
recall = np.mean(tree_recon.query(pts_gt, k=1)[0] < tau)
f_score = (2 * precision * recall) / (precision + recall) if precision + recall > 0 else 0.0
return d_r2r, d_raw2r, cd, f_score
def build_model(args, device):
return StackedTrajectorySpatialTextRefiner(
num_groups=args.num_refine_groups,
temporal_kwargs={
"num_tokens": args.num_tokens,
"token_dim": args.token_dim,
"max_seq_len": args.max_seq_len,
"num_blocks": args.temporal_blocks,
"compress_layers": args.compress_layers,
},
spatial_kwargs={
"num_tokens": args.num_tokens,
"token_dim": args.token_dim,
"num_blocks": args.spatial_blocks,
},
text_kwargs={
"token_dim": args.token_dim,
"text_dim": args.token_dim,
},
).to(device)
def load_checkpoint(auto_model, text_proj, checkpoint_path):
ckpt = torch.load(checkpoint_path, map_location="cpu")
if "auto" not in ckpt or "proj" not in ckpt:
raise KeyError("Checkpoint must contain keys 'auto' and 'proj'.")
auto_model.load_state_dict(ckpt["auto"])
text_proj.load_state_dict(ckpt["proj"])
return ckpt
def evaluate(args):
model_device = torch.device(args.model_device)
vae_device = torch.device(args.vae_device)
clip_device = torch.device(args.clip_device)
os.makedirs(args.output_dir, exist_ok=True)
vae = TripoSGVaeWrapper(device=str(vae_device))
vae.eval()
text_encoder = FrozenCLIPTextEncoder(args.clip_model_path, device=str(clip_device))
text_proj = TextProjectionLayer(in_dim=512, out_dim=args.token_dim).to(clip_device)
auto_model = build_model(args, model_device)
load_checkpoint(auto_model, text_proj, args.checkpoint_path)
auto_model.eval()
text_proj.eval()
dataset = get_full_trajectories(args.archive_path, limit=args.limit, min_frames=args.num_frames)
metrics = {"mae_r2r": [], "mae_raw2r": [], "cd": [], "fscore": []}
print(f"Loaded {len(dataset)} trajectories.")
print(f"Model device: {model_device}, VAE device: {vae_device}, CLIP device: {clip_device}")
for traj_idx, traj_data in enumerate(tqdm(dataset)):
if args.visualize_every > 0 and traj_idx % args.visualize_every != 0:
continue
traj_vis_dir = os.path.join(args.output_dir, f"{traj_idx:04d}")
os.makedirs(traj_vis_dir, exist_ok=True)
vertices_seq = traj_data["vertices"][: args.num_frames]
faces = traj_data["faces"]
caption = traj_data.get("caption", "")
gt_latents = encode_sequence_latents(
vae=vae,
vertices_seq=vertices_seq,
faces=faces,
device=vae_device,
num_frames=args.num_frames,
num_samples=args.num_surface_samples,
)
text_tokens = text_encoder([caption])
text_embed = text_proj(text_tokens.to(clip_device)).to(model_device)
pred_latents = autoregressive_predict(
auto_model=auto_model,
initial_latent=gt_latents[0].to(model_device),
text_embed=text_embed,
num_future_frames=args.num_frames - 1,
num_tokens=args.num_tokens,
device=model_device,
)
mesh_gt_0 = trimesh.Trimesh(vertices=vertices_seq[0], faces=faces, process=False)
save_mesh_as_png(mesh_gt_0, os.path.join(traj_vis_dir, "000_gt.png"))
for frame_idx in range(1, args.num_frames):
latent_t = pred_latents[frame_idx].to(vae_device).squeeze(1)
geometric_func = lambda x, latent=latent_t: vae.decode_latent(latent, sampled_points=x)
output_geometry = hierarchical_extract_geometry(
geometric_func,
device=str(vae_device),
bounds=(-1.005, -1.005, -1.005, 1.005, 1.005, 1.005),
)
recon_v, recon_f = output_geometry[0]
mesh_recon = trimesh.Trimesh(vertices=recon_v, faces=recon_f, process=False)
mesh_gt = trimesh.Trimesh(vertices=vertices_seq[frame_idx], faces=faces, process=False)
save_mesh_as_png(mesh_recon, os.path.join(traj_vis_dir, f"{frame_idx:03d}_pred.png"))
if args.save_gt_frames:
save_mesh_as_png(mesh_gt, os.path.join(traj_vis_dir, f"{frame_idx:03d}_gt.png"))
d_r2r, d_raw2r, cd, f_score = compute_mesh_metrics(
mesh_recon,
mesh_gt,
num_points=args.metric_points,
tau=args.fscore_tau,
)
metrics["mae_r2r"].append(d_r2r)
metrics["mae_raw2r"].append(d_raw2r)
metrics["cd"].append(cd)
metrics["fscore"].append(f_score)
running = {k: float(np.mean(v)) for k, v in metrics.items() if len(v) > 0}
print(f"Trajectory {traj_idx:04d} | {running}")
print("\n--- Evaluation Results ---")
for key, values in metrics.items():
print(f"{key}: {np.mean(values):.6f}" if values else f"{key}: no samples")
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--checkpoint_path", type=str, required=True)
parser.add_argument("--archive_path", type=str, required=True)
parser.add_argument("--clip_model_path", type=str, default="/home/dataset-assist-0/usr/lh/ysh/dw/RL/AR/pretrain/clip-vit-base-patch32")
parser.add_argument("--output_dir", type=str, default="/home/dataset-assist-0/usr/lh/ysh/dw/RL/AR/mesh_vis_results")
parser.add_argument("--num_tokens", type=int, default=512)
parser.add_argument("--token_dim", type=int, default=64)
parser.add_argument("--num_refine_groups", type=int, default=3)
parser.add_argument("--temporal_blocks", type=int, default=6)
parser.add_argument("--compress_layers", type=int, default=3)
parser.add_argument("--spatial_blocks", type=int, default=3)
parser.add_argument("--max_seq_len", type=int, default=16)
parser.add_argument("--num_frames", type=int, default=16)
parser.add_argument("--limit", type=int, default=400)
parser.add_argument("--visualize_every", type=int, default=74)
parser.add_argument("--num_surface_samples", type=int, default=51200)
parser.add_argument("--metric_points", type=int, default=2048)
parser.add_argument("--fscore_tau", type=float, default=0.05)
parser.add_argument("--save_gt_frames", action="store_true")
parser.add_argument("--model_device", type=str, default="cuda:0")
parser.add_argument("--vae_device", type=str, default="cuda:0")
parser.add_argument("--clip_device", type=str, default="cuda:0")
return parser.parse_args()
if __name__ == "__main__":
evaluate(parse_args())