-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathinfer_libero.py
More file actions
44 lines (34 loc) · 1.55 KB
/
Copy pathinfer_libero.py
File metadata and controls
44 lines (34 loc) · 1.55 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
import torch
from openpi.training import config as _config
from openpi.policies import policy_config
from openpi.training import data_loader as _data
def main():
# === 1. 加载 config 和 checkpoint ===
config = _config.get_config("pi05_libero")
checkpoint_dir = "/data/wuyuzhou-20250917/openpi/checkpoints/pi05_libero/test_single/30000"
# 自动加载 PyTorch 训练好的权重
policy = policy_config.create_trained_policy(config, checkpoint_dir)
# === 2. 用数据加载器取一条数据 ===
loader = _data.create_data_loader(config, framework="pytorch", shuffle=False)
observation, actions = next(iter(loader))
print("✅ Loaded one sample from dataset")
print("🔑 Observation keys:", list(observation.to_dict().keys()))
# Observation 是 dataclass,要转 dict
obs_dict = observation.to_dict()
# 打印每个 key 的 shape
for k, v in obs_dict.items():
if torch.is_tensor(v):
print(f" {k}: shape={tuple(v.shape)}, dtype={v.dtype}")
else:
print(f" {k}: type={type(v)}")
# 搬到 GPU
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
obs_dict = {k: v.to(device) if torch.is_tensor(v) else v for k, v in obs_dict.items()}
# === 3. 推理 ===
with torch.no_grad():
output = policy.infer(obs_dict)
predicted_actions = output["actions"]
print("🎯 Predicted actions shape:", predicted_actions.shape)
print("🎯 Predicted actions (first sample):", predicted_actions[0])
if __name__ == "__main__":
main()