File size: 4,145 Bytes
1e09dca
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c2c0566
 
 
 
 
 
 
 
1e09dca
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
deeed03
 
1e09dca
deeed03
 
1e09dca
 
 
 
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
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
"""Smoke-test: load make_env from a Hub snapshot (or local checkout) and step.

Usage (isaac conda, GPU)::

    ASSEMBLY_BENCH_HUB=lukasskellijs/env_assembly_bench \\
      ACCEPT_EULA=Y PRIVACY_CONSENT=Y OMNI_KIT_ACCEPT_EULA=YES \\
      python scripts/smoke_hub_load.py --from-hub --steps 5

    # Local checkout (no Hub round-trip):
    python scripts/smoke_hub_load.py --local --steps 5
"""

from __future__ import annotations

import argparse
import importlib.util
import logging
import os
import sys
from pathlib import Path
from types import SimpleNamespace

ROOT = Path(__file__).resolve().parents[1]


def _load_env_module(root: Path):
    env_py = root / "env.py"
    root_str = str(root)
    # Hub package must win over any editable/local assembly_bench install.
    if root_str in sys.path:
        sys.path.remove(root_str)
    sys.path.insert(0, root_str)
    for name in list(sys.modules):
        if name == "assembly_bench" or name.startswith("assembly_bench."):
            del sys.modules[name]
    spec = importlib.util.spec_from_file_location("env_assembly_bench_env", env_py)
    module = importlib.util.module_from_spec(spec)
    assert spec.loader is not None
    spec.loader.exec_module(module)
    return module


def main() -> None:
    logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
    p = argparse.ArgumentParser()
    p.add_argument("--from-hub", action="store_true", help="snapshot_download from Hub")
    p.add_argument("--local", action="store_true", help="use this checkout")
    p.add_argument("--steps", type=int, default=5)
    p.add_argument("--variant", default="peg_round_M1_loose")
    p.add_argument("--num-envs", type=int, default=1)
    args = p.parse_args()
    if args.from_hub == args.local:
        p.error("pass exactly one of --from-hub or --local")

    if args.from_hub:
        from huggingface_hub import snapshot_download

        hub = os.environ.get("ASSEMBLY_BENCH_HUB", "lukasskellijs/env_assembly_bench")
        # Force a re-fetch when HF_HUB_FORCE_DOWNLOAD is set; otherwise use cache.
        root = Path(
            snapshot_download(
                repo_id=hub,
                force_download=os.environ.get("HF_HUB_FORCE_DOWNLOAD", "") == "1",
            )
        )
        logging.info("Hub snapshot at %s", root)
    else:
        root = ROOT
        logging.info("Local checkout at %s", root)

    module = _load_env_module(root)
    cfg = SimpleNamespace(
        environment="assembly_bench",
        embodiment="droid_abs_joint_pos_softmimic",
        object=None,
        mimic=False,
        teleop_device=None,
        seed=0,
        device="cuda:0",
        disable_fabric=False,
        enable_cameras=True,
        headless=True,
        enable_pinocchio=False,
        episode_length=50,
        state_dim=15,
        action_dim=8,
        camera_height=720,
        camera_width=1280,
        video=False,
        video_length=10,
        video_interval=15,
        state_keys="joint_pos,gripper_pos,eef_pos,eef_quat",
        camera_keys="front_cam_rgb,wrist_camera_rgb",
        task=None,
        variant=args.variant,
        reward="none",
        hdr="asm_machine_shop",
        light_intensity=1500.0,
    )
    suites = module.make_env(n_envs=args.num_envs, use_async_envs=False, cfg=cfg)
    env = next(iter(next(iter(suites.values())).values()))
    obs, info = env.reset(seed=0)
    logging.info("reset ok; task=%r keys=%s", env.task, list(obs.keys()) if isinstance(obs, dict) else type(obs))
    for i in range(args.steps):
        action = env.action_space.sample()
        obs, reward, terminated, truncated, info = env.step(action)
        logging.info(
            "step %d reward=%s term=%s trunc=%s",
            i,
            getattr(reward, "mean", lambda: reward)(),
            terminated,
            truncated,
        )
        if terminated.any() or truncated.any():
            obs, info = env.reset()
    # Kit's SimulationApp.close() terminates the process — log success first.
    print("smoke ok", flush=True)
    logging.info("smoke ok")
    env.close()



if __name__ == "__main__":
    main()