19 Commits

Author SHA1 Message Date
Logic acbd7c605a feat(sim): save air-insert and rollout validation updates 2026-05-05 20:52:53 +08:00
Logic 73f5b6e3d9 merge: socket-peg sim and imf attnres into main 2026-05-03 10:14:48 +08:00
Logic bc4caf355b fix(sim): split transfer and air-insert camera sets 2026-05-03 10:14:04 +08:00
Logic b1116e489f merge: imf attnres policy
# Conflicts:
#	roboimi/demos/vla_scripts/eval_vla.py
#	roboimi/envs/double_base.py
2026-05-02 22:23:29 +08:00
Logic a2c018acce merge: socket peg air insert simulation 2026-05-02 22:13:36 +08:00
Logic 4890f54b13 fix(sim): align socket peg collection settings 2026-05-02 21:38:16 +08:00
Logic 5c5cb299e9 feat(sim): switch air insert task to socket peg 2026-05-02 17:34:43 +08:00
Logic 4c3646a3d5 fix(policy): perform stable horizontal air insertion 2026-04-24 09:41:37 +08:00
Logic 4936cf2635 fix(policy): stabilize air insert scripted success 2026-04-24 09:20:50 +08:00
Logic d245d64def fix(policy): avoid cross-arm collision in air insert rollout 2026-04-23 18:04:54 +08:00
Logic 8145c9eb62 feat(policy): add scripted air insertion policy 2026-04-23 17:44:53 +08:00
Logic a837a982f7 feat(env): add strict air insertion reward and success logic 2026-04-23 17:40:46 +08:00
Logic f1ede7690f feat(scene): add ring and bar insertion scene assets 2026-04-23 17:32:43 +08:00
Logic 06ac6c6d18 docs(plan): cover rollout entrypoint and eval regressions 2026-04-23 17:14:49 +08:00
Logic fce6839daa feat(env): register sim air insert ring bar task 2026-04-23 17:05:33 +08:00
Logic 3eb1a83940 docs(plan): tighten task ordering and smoke checks 2026-04-23 16:55:24 +08:00
Logic 636290d36a docs: clarify ring bar task state contracts 2026-04-23 16:47:05 +08:00
Logic 4ea75966ee docs(plan): add sim air insert ring bar implementation plan 2026-04-23 16:43:02 +08:00
Logic 27f4a07632 docs(spec): add sim air insert ring bar design 2026-04-23 16:33:58 +08:00
26 changed files with 4308 additions and 2039 deletions
@@ -0,0 +1,311 @@
# sim_air_insert_ring_bar Implementation Plan
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
**Goal:** Add an independent dual-Diana MuJoCo task `sim_air_insert_ring_bar` with a square ring block, a square bar block, staged rewards, strict finite-geometry in-air insertion success detection, and a task-specific scripted policy.
**Architecture:** Reuse the current dual-Diana EE-control stack and environment factory, but add a task-specific scene XML, robot asset entrypoint, sampling helpers, and a new task-specific environment module. Keep `sim_transfer` untouched while introducing pure-Python geometry helpers and focused tests so reward/success behavior can be regression tested without requiring a full MuJoCo rollout in every test.
**Tech Stack:** Python, unittest, MuJoCo XML assets, existing dual-Diana environment classes, Hydra-compatible task naming/config patterns.
---
## File Structure / Responsibilities
- **Create:** `roboimi/assets/models/manipulators/DianaMed/ring_bar_objects.xml`
- Defines the rigid ring body and bar body, each with a free joint and stable box-based geoms.
- **Create:** `roboimi/assets/models/manipulators/DianaMed/bi_diana_ring_bar_ee.xml`
- Scene entrypoint that includes the shared world/table/robot assets plus the new object XML.
- **Modify:** `roboimi/assets/robots/diana_med.py`
- Add a task-specific robot asset class for the new scene XML without changing existing `BiDianaMed` behavior.
- **Modify:** `roboimi/utils/act_ex_utils.py`
- Add deterministic helpers to sample left/right planar placement regions for ring and bar objects.
- **Modify:** `roboimi/utils/constants.py`
- Register the new task name and default metadata.
- **Create:** `roboimi/envs/double_air_insert_env.py`
- New task-specific environment, finite-geometry success helpers, reset logic, reward logic, and task factory branch.
- **Modify:** `roboimi/envs/double_pos_ctrl_env.py`
- Route `make_sim_env()` to the new task-specific environment while keeping current `sim_transfer` logic unchanged.
- **Create:** `roboimi/demos/diana_air_insert_policy.py`
- Task-specific waypoint/open-loop scripted policy for grasp-lift-align-insert.
- **Modify:** `roboimi/demos/vla_scripts/eval_vla.py`
- Reset the new task with the correct sampled task state instead of assuming a single transfer box pose.
- **Create:** `tests/test_air_insert_env.py`
- Focused unit tests for sampling, reset helpers, reward progression, and strict success detection.
- **Modify:** `tests/test_eval_vla_headless.py`
- Add coverage that headless evaluation dispatches the correct reset sampler for the new task.
- **Modify:** `tests/test_robot_asset_paths.py`
- Verify the new robot asset class resolves its XML path correctly independent of cwd.
---
### Task 1: Add failing tests for task registration, samplers, and asset wiring
**Files:**
- Create: `tests/test_air_insert_env.py`
- Modify: `tests/test_eval_vla_headless.py`
- Modify: `tests/test_robot_asset_paths.py`
- Modify: `roboimi/utils/act_ex_utils.py` (later in implementation)
- Modify: `roboimi/utils/constants.py` (later in implementation)
- Modify: `roboimi/assets/robots/diana_med.py` (later in implementation)
- Modify: `roboimi/envs/double_pos_ctrl_env.py` (later in implementation)
- Create: `roboimi/envs/double_air_insert_env.py` (minimal stub in this task)
- [ ] **Step 1: Write failing tests for task config and sampling helpers**
Add tests in `tests/test_air_insert_env.py` covering:
- `SIM_TASK_CONFIGS['sim_air_insert_ring_bar']` exists
- `sample_air_insert_ring_bar_pose()` (or equivalent helper) returns ring/bar positions with fixed z and correct left/right planar ranges
- output structure is explicit and easy for reset/eval code to consume
- [ ] **Step 2: Write failing tests for environment factory dispatch and robot asset resolution**
Add tests covering:
- `make_sim_env('sim_air_insert_ring_bar', headless=True)` dispatches to the new environment with rendering disabled
- a new robot asset class resolves the new XML path independent of cwd, similar to the existing `BiDianaMed` test pattern
- [ ] **Step 3: Write failing tests for eval reset helper dispatch**
Extend `tests/test_eval_vla_headless.py` so headless eval can reset the new task using the new sampler instead of hard-coding `sample_transfer_pose()`.
- [ ] **Step 4: Run the targeted tests to verify they fail for the expected missing-feature reasons**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env tests.test_eval_vla_headless tests.test_robot_asset_paths -v`
Expected:
- FAIL because the new task config/helper/class/dispatch branch does not exist yet
- [ ] **Step 5: Implement the minimal production code to satisfy the new task registration and helper tests**
Implement only enough to make the new tests pass:
- add new task config entry
- add the new placement sampler
- add the new robot asset class
- create a minimal importable `double_air_insert_env.py` stub and class/function surface needed for factory dispatch tests
- add the factory dispatch branch / headless wiring
- update eval reset dispatch for the new task
- [ ] **Step 6: Re-run the targeted tests to verify they pass**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env tests.test_eval_vla_headless tests.test_robot_asset_paths -v`
Expected:
- PASS for the new registration/sampler/dispatch/asset tests
- [ ] **Step 7: Commit Task 1**
Run:
`git add tests/test_air_insert_env.py tests/test_eval_vla_headless.py tests/test_robot_asset_paths.py roboimi/utils/act_ex_utils.py roboimi/utils/constants.py roboimi/assets/robots/diana_med.py roboimi/envs/double_pos_ctrl_env.py roboimi/envs/double_air_insert_env.py roboimi/demos/vla_scripts/eval_vla.py && git commit -m "feat(env): register sim air insert ring bar task"`
---
### Task 2: Add the MuJoCo ring+bar scene assets and reset helpers
**Files:**
- Create: `roboimi/assets/models/manipulators/DianaMed/ring_bar_objects.xml`
- Create: `roboimi/assets/models/manipulators/DianaMed/bi_diana_ring_bar_ee.xml`
- Create or Modify: `roboimi/envs/double_air_insert_env.py`
- Modify: `tests/test_air_insert_env.py`
- [ ] **Step 1: Write failing tests for object reset helpers and scene-specific joint naming assumptions**
In `tests/test_air_insert_env.py`, add unit tests for helper functions that:
- write ring pose to `ring_block_joint` from the named task-state mapping
- write bar pose to `bar_block_joint` from the named task-state mapping
- read back `env_state` as a stable 14D vector `[ring_pos, ring_quat, bar_pos, bar_quat]`
Use fake `mj_data` objects so tests stay fast and deterministic.
- [ ] **Step 2: Run the focused test slice and verify it fails**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env -v`
Expected:
- FAIL because reset/state helper functions and joint conventions are not implemented yet
- [ ] **Step 3: Implement the scene XML files and reset/state helper code**
Implement:
- the object XML with one rigid ring body and one rigid bar body
- the task scene XML entrypoint using the shared world/table/robot includes
- reset helper(s) in `double_air_insert_env.py` that set qpos for both free joints with fixed quaternions
- task-state accessor(s) returning both object poses in a stable structure
- [ ] **Step 4: Re-run the focused test slice and verify it passes**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env -v`
Expected:
- PASS for reset/state helper tests
- [ ] **Step 5: Commit Task 2**
Run:
`git add roboimi/assets/models/manipulators/DianaMed/ring_bar_objects.xml roboimi/assets/models/manipulators/DianaMed/bi_diana_ring_bar_ee.xml roboimi/envs/double_air_insert_env.py tests/test_air_insert_env.py && git commit -m "feat(scene): add ring and bar insertion scene assets"`
---
### Task 3: Implement strict reward and finite-geometry success detection
**Files:**
- Modify: `roboimi/envs/double_air_insert_env.py`
- Modify: `tests/test_air_insert_env.py`
- [ ] **Step 1: Write failing tests for reward stages and strict success detection**
Add tests in `tests/test_air_insert_env.py` for:
- left contact stage reward
- right contact stage reward
- ring lifted off table stage
- bar lifted off table stage
- positive success case where a finite bar truly passes through the aperture
- negative case where the centerline would pass but the finite square body would clip
- negative case where the bar has not crossed the ring thickness direction enough
- negative case where one/both objects are still on the table
Structure the tests around pure helper functions and light fake contact/state objects so the geometry logic is directly regression tested.
- [ ] **Step 2: Run the focused tests and verify they fail for missing reward/success logic**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env -v`
Expected:
- FAIL because the staged reward and finite-geometry insertion logic are not implemented yet
- [ ] **Step 3: Implement minimal strict success helpers and reward logic**
Implement in `roboimi/envs/double_air_insert_env.py`:
- pure helper(s) for transforming bar geometry into ring-local coordinates
- finite-geometry insertion predicate (not centerline-only)
- table-contact / airborne checks
- staged reward function returning the highest achieved stage with `max_reward = 5`
- [ ] **Step 4: Re-run the focused tests to verify the logic passes**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env -v`
Expected:
- PASS for reward and success-detection regression tests
- [ ] **Step 5: Commit Task 3**
Run:
`git add roboimi/envs/double_air_insert_env.py tests/test_air_insert_env.py && git commit -m "feat(env): add strict air insertion reward and success logic"`
---
### Task 4: Add the scripted policy and integration smoke coverage
**Files:**
- Create: `roboimi/demos/diana_air_insert_policy.py`
- Modify: `roboimi/demos/diana_record_sim_episodes.py`
- Modify: `tests/test_air_insert_env.py`
- Optionally Modify: `roboimi/demos/vla_scripts/eval_vla.py` (only if integration gaps remain after Task 1)
- [ ] **Step 1: Write failing tests for scripted-policy action shape and basic generation**
Add tests covering:
- the new policy produces a 16D action
- trajectory generation accepts sampled named task state without error
- the first action is a valid open-gripper safe pose command
- a deterministic nominal smoke path (with canonical sampled state or fake env shim) reaches the intended terminal interface contract without shape/reward mismatches
Keep the tests unit-level; do not require a full MuJoCo rollout for every assertion.
- [ ] **Step 2: Write failing tests for the scripted rollout entrypoint and a real headless smoke path**
Add coverage for both:
- the standard scripted rollout entrypoint (`roboimi/demos/diana_record_sim_episodes.py`) can select the new task sampler/policy instead of remaining sim_transfer-only
- a deterministic integration/smoke test that instantiates `make_sim_env('sim_air_insert_ring_bar', headless=True)`, resets with sampled named task state, and steps a few actions or scripted-policy outputs using the real task XML and task-specific wiring
- [ ] **Step 3: Run the scripted-policy tests and verify they fail**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env -v`
Expected:
- FAIL because the new scripted policy does not exist yet
- [ ] **Step 4: Implement the waypoint-based scripted policy**
Implement a conservative open-loop policy with phases:
- safe wait pose
- above-target approach
- descend + grasp
- dual lift
- airborne meeting alignment
- bar push-through insertion
Use fixed orientations for version 1 and follow the existing repository style from `diana_policy.py`.
- [ ] **Step 5: Re-run the scripted-policy tests to verify they pass**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env -v`
Expected:
- PASS for scripted-policy tests
- [ ] **Step 6: Run the combined verification suite for this feature**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env tests.test_eval_vla_headless tests.test_eval_vla_rollout_artifacts tests.test_train_vla_rollout_validation tests.test_robot_asset_paths -v`
Expected:
- PASS with 0 failures
- [ ] **Step 6b: Run the mandatory real headless smoke check**
Run a focused smoke command that instantiates the real task, resets with sampled state, and steps a few actions using the new scripted policy or a deterministic action sequence.
Example command (adjust module/test helper if needed):
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env.AirInsertEnvSmokeTest -v`
Expected:
- PASS, proving the real XML/assets/env wiring instantiate and step correctly in headless mode
- [ ] **Step 7: Commit Task 4**
Run:
`git add roboimi/demos/diana_air_insert_policy.py tests/test_air_insert_env.py tests/test_eval_vla_headless.py tests/test_robot_asset_paths.py roboimi/demos/vla_scripts/eval_vla.py && git commit -m "feat(policy): add scripted air insertion policy"`
---
### Task 5: Final verification and implementation review
**Files:**
- Review all files touched above
- [ ] **Step 1: Run fresh end-to-end verification before claiming completion**
Run:
`/home/droid/.conda/envs/roboimi/bin/python -m unittest tests.test_air_insert_env tests.test_eval_vla_headless tests.test_robot_asset_paths -v`
Expected:
- PASS with 0 failures
- [ ] **Step 2: Inspect git status and recent commits**
Run:
`git status --short && git log --oneline --decorate -n 8`
Expected:
- only intended feature files modified / committed
- [ ] **Step 3: Request final code review for the completed feature**
Use the requesting-code-review skill against the full diff from the feature branch starting point to current HEAD.
- [ ] **Step 4: Address any review findings and re-run verification if code changes**
If fixes are made, repeat the unittest command from Step 1.
- [ ] **Step 5: Hand off using finishing-a-development-branch**
After verification and review, use the finishing-a-development-branch skill to decide merge / PR / cleanup.
@@ -0,0 +1,316 @@
# sim_air_insert_ring_bar Design
## Summary
Add a new independent MuJoCo simulation task named `sim_air_insert_ring_bar` that keeps the existing dual-Diana tabletop setup but replaces the single transfer box with two randomized objects:
- a square ring block grasped by the left arm
- a square bar block grasped by the right arm
The task is to pick both objects off the table and complete an in-air insertion where the bar truly passes through the ring aperture. The existing `sim_transfer` task must remain unchanged.
## Goals
- Reuse the current dual-Diana EE-control simulation stack
- Keep the same table/base robot arrangement as the existing transfer task
- Add an independent task entrypoint and scene definition
- Randomize planar placement of both objects within left/right task-specific regions
- Implement reward staging for contact, lift, and successful in-air insertion
- Add a scripted policy that performs pick, lift, align, and in-air insertion
- Preserve compatibility with existing environment creation, evaluation, and rollout patterns
## Non-Goals
- No random yaw in the first version
- No visual servoing or closed-loop insertion controller
- No general multi-task environment framework refactor
- No guarantee that the VLA training stack is immediately tuned for this new task
- No replacement or behavior change for `sim_transfer`
## Task Name
Use a new task name:
- `sim_air_insert_ring_bar`
This task should be exposed alongside `sim_transfer`, not as a replacement.
## Scene Geometry
### Shared Base Scene
Keep the dual Diana robot, the table, and the existing camera layout conceptually unchanged.
### Ring Block
Represent the square ring as a rigid free body composed from simple MuJoCo box geoms rather than an external mesh.
Dimensions:
- outer side length: 68 mm
- inner aperture side length: 32 mm
- thickness: 18 mm
- ring wall width: 18 mm
The ring should behave as a single object body with a single free joint.
### Bar Block
Represent the bar as a rigid free body with a single box geom.
Dimensions:
- length: 90 mm
- cross-section: 18 mm x 18 mm
The bar should also be a single free-joint body.
## Initial Placement / Reset
The first version uses position-only randomization with fixed orientation. Reset sampling stays **caller-driven**, matching the existing `sim_transfer` usage pattern in rollout/eval code: a helper samples task state, then callers pass that state into `env.reset(...)`.
Use an explicit sampled task-state structure with named fields:
- `ring_pos`: 3D position
- `ring_quat`: fixed 4D quaternion for version 1
- `bar_pos`: 3D position
- `bar_quat`: fixed 4D quaternion for version 1
Behavior:
- ring block: randomized only in a left-side planar sampling region
- bar block: randomized only in a right-side planar sampling region
- both objects start flat on the table
- both objects use fixed orientation at reset
- no random yaw, tilt, or flip in this version
The sampling regions should be chosen conservatively so that:
- the left arm can comfortably reach and grasp the ring
- the right arm can comfortably reach and grasp the bar
- scripted open-loop pick trajectories remain feasible
## Control / Action Interface
Reuse the current 16D EE-space action convention already used by the dual-Diana position-control environment:
- left arm EE pose: 7D (`xyz + quat`)
- right arm EE pose: 7D (`xyz + quat`)
- left gripper command: 1D
- right gripper command: 1D
The new task should continue using EE targets transformed through the existing IK-based control path.
## Environment Structure
Implement this as a new task-specific environment path while reusing the existing dual-Diana simulation base where possible.
Expected responsibilities:
- scene instantiation for the ring+bar setup
- task reset for randomized object placement
- environment-state accessors for both objects
- reward computation
- in-air insertion success detection
The environment factory must dispatch by task name and leave the `sim_transfer` branch unchanged.
## Observation / Environment State
The task should retain the current observation structure style used by the dual-Diana environment:
- `qpos`
- multi-camera images
For task state access, the environment should expose a stable `env_state` vector with this exact order:
- `ring_pos[0:3]`
- `ring_quat[3:7]`
- `bar_pos[7:10]`
- `bar_quat[10:14]`
This 14D state should be sufficient for scripted-policy debugging and future rollout analysis, while reset itself remains caller-driven via the named task-state helper structure above.
## Reward Design
Use staged rewards in the same spirit as the current task, returning the highest achieved stage rather than accumulating one-time sparse bonuses per event.
Maximum reward:
- `max_reward = 5`
Reward stages:
1. left gripper touches the ring block
2. right gripper touches the bar block
3. ring block is lifted off the table
4. bar block is lifted off the table
5. while both objects are off the table, the bar truly passes through the ring aperture
Notes:
- contact rewards are intended as grasp-progress stages
- lift rewards require the object to be off the table, not merely touched
- final success reward only applies when both objects are airborne
## Success Detection
Success must **not** be based on a centerline-only check.
A centerline-only test is insufficient because:
- the bar has thickness, so a centerline can pass through while the body cannot
- a square bar with imperfect orientation can have its centerline inside the aperture while its corners still collide with the ring
### Required Success Semantics
A successful insertion requires all of the following:
1. the ring is off the table
2. the bar is off the table
3. the bar has actually crossed through the ring thickness direction
4. the bars finite square cross-section fits through the square aperture during that crossing
### Recommended Detection Approach
Use a task-level geometric check in Python rather than relying on contact alone.
Implementation intent:
- transform the bar geometry into the rings local frame
- reason about the bar as a finite oriented box (not a line)
- verify that the bar has crossed the ring thickness direction
- verify that the portion of the bar passing the aperture fits within the inner square opening, accounting for the bars cross-section and orientation
This geometric check is the primary success test.
### Role of Contacts
Contacts may still be used for:
- grasp-stage rewards
- debugging / diagnostics
But contact alone should **not** be the sole criterion for insertion success, since:
- a true clean insertion may have limited aperture-wall contact
- persistent contact can also happen while the bar is jammed and not actually inserted
## Scripted Policy
Add a new task-specific scripted policy for `sim_air_insert_ring_bar`.
### Policy Intent
The first version prioritizes a conservative, reliable open-loop demonstration rather than an optimized trajectory.
### Action Phases
The scripted policy should follow these phases:
1. move both arms to safe initial / waiting poses with grippers open
2. move left arm above the ring and right arm above the bar
3. descend and grasp the assigned objects
4. lift both objects clear of the table
5. move both objects to an airborne meeting region above the table
6. hold the ring stably while aligning the bar with the aperture
7. push the bar along the intended insertion direction until the geometric success condition is met
### Grasp Assignment
- left arm: ring only
- right arm: bar only
### Motion Style
Keep the current repository style:
- waypoint-based trajectory definition
- open-loop interpolation between waypoints
- fixed grasp orientation in the first version
No adaptive replanning is required for the first version.
## Files / Integration Scope
The implementation is expected to add task-specific files rather than broadly refactoring the codebase.
Likely additions / changes:
- a new MuJoCo scene XML for the ring+bar task
- one or more XML fragments defining the two new objects
- a new task-specific dual-Diana environment file
- robot asset wiring for the new scene XML
- reset sampling helpers for the new task
- task registration in constants / environment factory paths
- a new scripted policy file
- focused tests for task creation, reset, rewards, success detection, and scripted policy shape/smoke behavior
## Testing Requirements
At minimum, add regression coverage for:
### Environment Creation
- the new task can be created via the task factory
- the existing `sim_transfer` task remains unchanged
### Reset / Sampling
- ring reset positions are inside the left sampling region
- bar reset positions are inside the right sampling region
- reset orientation is fixed as intended
### Environment State
- environment-state access returns both object poses in the expected structure
### Success Detection
Must include both positive and negative cases.
Positive case:
- a configuration where the finite bar truly passes through the ring aperture is detected as success
Negative cases:
- centerline-inside but finite body would clip the aperture
- not enough depth / not actually crossing the ring thickness direction
- one or both objects still on the table
### Reward Logic
- left contact stage
- right contact stage
- ring lift stage
- bar lift stage
- final success stage with `max_reward = 5`
### Scripted Policy
At minimum:
- policy emits valid 16D actions
- trajectory generation does not error
- rollout smoke path can step through the new environment
## Risks / Constraints
- MuJoCo contact naming must remain stable enough for stage rewards
- geometric insertion checks must be strict enough to avoid false positives but not so brittle that numerically valid insertions are missed
- scripted open-loop insertion may require conservative alignment and lift heights to keep the first version reliable
## Acceptance Criteria
The feature is complete when all of the following are true:
- `sim_air_insert_ring_bar` is creatable as an independent task
- the scene contains the dual Diana, table, ring block, and bar block
- reset randomizes ring and bar positions in left/right planar regions with fixed orientation
- the environment exposes task state for both objects
- staged rewards progress to `max_reward = 5`
- final success is based on finite-geometry insertion semantics, not a centerline-only shortcut
- a new scripted policy can execute the intended pick-lift-align-insert behavior in the new environment
- a canonical nominal smoke path (unit-level or deterministic integration-level) exists for the new scripted-policy interface so success is not judged purely by interpretation
- existing `sim_transfer` behavior is preserved
+200
View File
@@ -0,0 +1,200 @@
import mujoco
from mujoco import viewer
import sys
import numpy as np
import time
import threading
class MjBasicRenderer:
def __new__(cls, *args, **kwargs):
return super().__new__(cls)
def __init__(self, mj_model=None, mj_data=None):
# keyboard flag
self.render_paused = True
self.exit_flag = False
# init param
self.mj_model = mj_model
self.mj_data = mj_data
self.renderer = "viewer" # default
self.viewer = None
self._image = None
# Set up mujoco viewer
self.image_renderer = mujoco.Renderer(self.mj_model)
def __del__(self):
pass
def _init_renderer(self):
"""Initialize renderer, choose official renderer with "viewer"(joined from version 2.3.3),
another renderer with "mujoco_viewer"
"""
def key_callback(keycode):
if keycode == 32: # space
self.render_paused = not self.render_paused
elif keycode == 256: # escape
self.exit_flag = not self.exit_flag
if self.renderer == "viewer":
# This function does not block, allowing user code to continue execution.
self.viewer = viewer.launch_passive(
self.mj_model,
self.mj_data,
key_callback=key_callback,
show_left_ui=False,
show_right_ui=False,
)
self.set_renderer_config()
else:
raise ValueError("Invalid renderer for some reason.")
def render(self):
"""mujoco render"""
if self.viewer is not None and self.render_paused is True:
if self.viewer.is_running() and self.exit_flag is False:
self.viewer: viewer.Handle
self.viewer.sync()
else:
self.viewer.close()
def set_renderer_config(self):
"""Setup mujoco global config while using viewer as renderer.
It should be noted that the render thread need locked.
"""
self.viewer.cam.lookat = np.array([0.4, 0, 0.5])
self.viewer.cam.azimuth -= 0.005
with self.viewer.lock():
self.viewer.opt.flags[mujoco.mjtVisFlag.mjVIS_CONTACTPOINT] = int(
self.mj_data.time % 2
)
try:
import cv2
except ImportError:
print("Could not import cv2, please install it to enable camera viewer.")
class MjMultiRenderer(MjBasicRenderer):
# __slots__=('mj_model','mj_data','renderer','enable_camera_viewer')
def __new__(cls, *args, **kwargs):
return super().__new__(cls)
def __init__(
self,
mj_model=None,
mj_data=None,
renderer=None,
enable_camera_viewer=False,
enable_depth=False,
):
super().__init__(mj_model, mj_data)
self._depth = None
self.renderer = renderer
self._init_renderer()
self.enable_camera_viewer = enable_camera_viewer
if self.enable_camera_viewer:
self.enable_depth = enable_depth
self._init_window()
else:
self.enable_depth = False
print("No Camera View")
def __del__(self):
self.close()
def _init_renderer(self):
"""
Initialize renderer, choose official renderer with "viewer"(joined from version 2.3.3)
"""
if self.renderer == "unity":
# TODO: Support unity renderer.
raise ValueError("Unity renderer init failed for no supporting reason")
elif self.renderer == "viewer":
super()._init_renderer()
print("mujoco viewer init !")
else:
raise ValueError("renderer init failed for some reason.")
def _init_window(self, name="Camera view"):
if not self.enable_depth:
cv2.namedWindow(name, cv2.WINDOW_NORMAL)
else:
cv2.namedWindow(name, cv2.WINDOW_NORMAL)
cv2.namedWindow("Camera depth view", cv2.WINDOW_NORMAL)
def render(self):
"""render mujoco"""
if self.renderer == "viewer":
super().render()
elif self.renderer == "unity":
# TODO: Support unity renderer.
raise ValueError("Unity renderer not supported now.")
else:
raise ValueError("Invalid renderer for some reason.")
def render(self):
"""mujoco render"""
if self.viewer is not None and self.render_paused is True:
if self.viewer.is_running() and self.exit_flag is False:
self.viewer: viewer.Handle
self.viewer.sync()
else:
self.viewer.close()
def camera_render(self, cam=None):
if self.enable_camera_viewer:
if not self.enable_depth:
rgb, depth = self.render_from_camera(cam)
rgb = cv2.resize(rgb, (1920, 1600))
cv2.imshow("Camera view", rgb)
cv2.waitKey(1)
else:
rgb, depth = self.render_from_camera(cam)
cv2.imshow("Camera view", rgb)
cv2.imshow("Camera depth view", depth)
cv2.waitKey(1)
else:
print("camera info disable")
return
def render_from_camera(self, cam=None):
self.image_renderer.update_scene(self.mj_data, camera=cam)
if self.enable_depth is True:
self.image_renderer.enable_depth_rendering()
org = self.image_renderer.render()
depth = org[:, :]
self.image_renderer.disable_depth_rendering()
org = self.image_renderer.render()
image = org[:, :, ::-1]
else:
org = self.image_renderer.render()
image = org[:, :, ::-1]
depth = np.zeros([240, 320])
return image, depth
def close(self):
"""close the environment."""
if self.enable_camera_viewer and self.viewer.is_running() == False:
cv2.destroyAllWindows()
self.viewer.close()
# sys.exit(0)
# def get_cam_intrinsic(self, fovy=45.0, width=320, height=240):
# aspect = width * 1.0 / height
# fovx = np.degrees(2 * np.arctan(aspect * np.tan(np.radians(fovy / 2))))
# cx = 0.5 * width
# cy = 0.5 * height
# fx = cx / np.tan(fovx * np.pi / 180 * 0.5)
# fy = cy / np.tan(fovy * np.pi / 180 * 0.5)
# K = np.array([[fx, 0, cx],
# [0, fy, cy],
# [0, 0, 1]], dtype=np.float32)
@@ -76,7 +76,7 @@
<body name="ee_cam_left" pos="0.00 0.046 -0.075" euler="0.0 0.0 -0.0"> <body name="ee_cam_left" pos="0.00 0.046 -0.075" euler="0.0 0.0 -0.0">
<inertial pos="0 0 0" quat="1 0 0 0" mass="0" diaginertia="0 0 0" /> <inertial pos="0 0 0" quat="1 0 0 0" mass="0" diaginertia="0 0 0" />
<geom type="mesh" contype="1" conaffinity="1" group="1" rgba="0.69804 0.69804 0.69804 1" mesh="realsense_cam" /> <geom type="mesh" contype="1" conaffinity="1" group="1" rgba="0.69804 0.69804 0.69804 1" mesh="realsense_cam" />
<camera name="rs_cam_left" mode="fixed" pos="0.0 0.0 -0.25" euler="0.0 9.4 0.0 " fovy="15" resolution="1920 1200"/> <camera name="rs_cam_left" mode="fixed" pos="0.0 0.0 0.01" euler="0.0 9.4 0.0 " fovy="50" resolution="1920 1200"/>
</body> </body>
</body> </body>
<body name="l_finger_left" pos="0 0.01 0.0444"> <body name="l_finger_left" pos="0 0.01 0.0444">
@@ -0,0 +1,6 @@
<mujoco model="bi_diana_socket_peg">
<include file="./empty_world.xml" />
<include file="./table_square.xml" />
<include file="./socket_peg_objects.xml" />
<include file="./BiDianaMed_rethink.xml" />
</mujoco>
@@ -0,0 +1,19 @@
<mujoco model="socket_peg_objects">
<worldbody>
<body name="peg" pos="0.12 0.90 0.46">
<joint name="red_peg_joint" type="free" frictionloss="0.01" />
<inertial pos="0 0 0" mass="0.05" diaginertia="0.002 0.002 0.002" />
<geom condim="4" solimp="2 1 0.01" solref="0.01 1" friction="1 0.005 0.0001" pos="0 0 0" size="0.06 0.01 0.01" type="box" name="red_peg" rgba="1 0 0 1" />
</body>
<body name="socket" pos="-0.12 0.90 0.472">
<joint name="blue_socket_joint" type="free" frictionloss="0.01" />
<inertial pos="0 0 0" mass="0.05" diaginertia="0.002 0.002 0.002" />
<geom condim="4" solimp="2 1 0.01" solref="0.01 1" friction="1 0.05 0.001" pos="0 0 -0.02" size="0.06 0.018 0.002" type="box" name="socket-1" rgba="0 0 1 1" />
<geom condim="4" solimp="2 1 0.01" solref="0.01 1" friction="1 0.05 0.001" pos="0 0 0.02" size="0.06 0.018 0.002" type="box" name="socket-2" rgba="0 0 1 1" />
<geom condim="4" solimp="2 1 0.01" solref="0.01 1" friction="1 0.05 0.001" pos="0 0.02 0" size="0.06 0.002 0.018" type="box" name="socket-3" rgba="0 0 1 1" />
<geom condim="4" solimp="2 1 0.01" solref="0.01 1" friction="1 0.05 0.001" pos="0 -0.02 0" size="0.06 0.002 0.018" type="box" name="socket-4" rgba="0 0 1 1" />
<geom condim="4" solimp="2 1 0.01" solref="0.01 1" friction="1 0.005 0.0001" pos="0 0 0" size="0.04 0.01 0.01" type="box" name="pin" rgba="1 0 0 1" />
</body>
</worldbody>
</mujoco>
@@ -7,7 +7,6 @@
<geom name="table" condim="4" contype="1" conaffinity="1" type="box" rgba="0.4 0.4 0.4 1" size="0.62 0.62 0.01" density="1500" friction="0.9 0.9 0.9"/> <geom name="table" condim="4" contype="1" conaffinity="1" type="box" rgba="0.4 0.4 0.4 1" size="0.62 0.62 0.01" density="1500" friction="0.9 0.9 0.9"/>
</body> </body>
<camera name="top" pos="0.0 1.0 2.0" fovy="44" mode="targetbody" target="table"/> <camera name="top" pos="0.0 1.0 2.0" fovy="44" mode="targetbody" target="table"/>
<camera name="angle" pos="0.0 0.0 2.0" fovy="37" mode="targetbody" target="table"/>
<camera name="front" pos="0 0 0.8" fovy="65" mode="fixed" quat="0.7071 0.7071 0 0"/> <camera name="front" pos="0 0 0.8" fovy="65" mode="fixed" quat="0.7071 0.7071 0 0"/>
</worldbody> </worldbody>
</mujoco> </mujoco>
+36
View File
@@ -91,3 +91,39 @@ class BiDianaMed(ArmBase):
""" Robot's init joint position. """ """ Robot's init joint position. """
return np.array([0.0, 0.0, 0.0, 1.57, 0.0, 0.0, 0.0]) return np.array([0.0, 0.0, 0.0, 1.57, 0.0, 0.0, 0.0])
class BiDianaMedSocketPeg(ArmBase):
def __init__(self):
super().__init__(
name="Bidiana_socket_peg",
urdf_path="roboimi/assets/models/manipulators/DianaMed/DualDianaMed.urdf",
xml_path="roboimi/assets/models/manipulators/DianaMed/bi_diana_socket_peg_ee.xml",
gripper=None
)
self.left_arm = self.Arm(self, 'single', self.urdf_path)
self.left_arm.set_Arm_base_link('left_base_link')
self.left_arm.set_Arm_ee_link('left_link7')
self.left_arm.InitKDL
self.left_arm.joint_index = ['l_j1','l_j2','l_j3','l_j4','l_j5','l_j6','l_j7']
self.left_arm.gripper_index = ['l_finger_joint_left','r_finger_joint_left']
self.left_arm.actuator_index = ['a1_l','a2_l','a3_l','a4_l','a5_l','a6_l','a7_l','gripper_left']
self.left_arm.setArmInitPose(self.init_qpos)
self.arms.append(self.left_arm)
self.right_arm = self.Arm(self,'single', self.urdf_path)
self.right_arm.set_Arm_base_link('right_base_link')
self.right_arm.set_Arm_ee_link('right_link7')
self.right_arm.InitKDL
self.right_arm.joint_index = ['r_j1','r_j2','r_j3','r_j4','r_j5','r_j6','r_j7']
self.right_arm.gripper_index = ['l_finger_joint_right','r_finger_joint_right']
self.right_arm.actuator_index = ['a1_r','a2_r','a3_r','a4_r','a5_r','a6_r','a7_r','gripper_right']
self.right_arm.setArmInitPose(self.init_qpos)
self.arms.append(self.right_arm)
self.jnt_num = self.left_arm.jnt_num + self.right_arm.jnt_num
self.kp = 500 * np.ones(self.jnt_num)
self.kd = 44.57 * np.ones(self.jnt_num)
@property
def init_qpos(self):
""" Robot's init joint position. """
return np.array([0.0, 0.0, 0.0, 1.57, 0.0, 0.0, 0.0])
+184
View File
@@ -0,0 +1,184 @@
import numpy as np
from pyquaternion import Quaternion
from roboimi.demos.diana_policy import PolicyBase
class TestAirInsertPolicy(PolicyBase):
ACTION_OBJECT_Z_OFFSET = 0.078
SOCKET_GRASP_OFFSET = np.array([0.0, 0.0, 0.0], dtype=np.float64)
PEG_GRASP_OFFSET = np.array([0.0, 0.0, 0.0], dtype=np.float64)
SOCKET_OUTER_GRASP_STRATEGY = "socket_outer"
LEGACY_GRASP_STRATEGY = "legacy"
SOCKET_HOLD_Z = 0.85
PEG_INSERT_START_OFFSET = np.array([0.105, 0.0, 0.0], dtype=np.float64)
INSERT_END_T = 580
LEFT_SOCKET_GRIPPER_CLOSED = -100
RIGHT_PEG_GRIPPER_CLOSED = -100
SOCKET_APPROACH_Z = 1.05
EPISODE_END_T = 600
def __init__(self, inject_noise=False, grasp_strategy=SOCKET_OUTER_GRASP_STRATEGY):
super().__init__(inject_noise=inject_noise)
valid_strategies = {
self.SOCKET_OUTER_GRASP_STRATEGY,
self.LEGACY_GRASP_STRATEGY,
}
if grasp_strategy not in valid_strategies:
raise ValueError(
f"Unsupported air insert grasp_strategy={grasp_strategy!r}; "
f"expected one of {sorted(valid_strategies)}"
)
self.grasp_strategy = grasp_strategy
def generate_trajectory(self, task_state):
return self._generate_socket_peg_trajectory(task_state)
def _generate_socket_peg_trajectory(self, task_state):
socket_xyz = np.asarray(task_state["socket_pos"], dtype=np.float64)
peg_xyz = np.asarray(task_state["peg_pos"], dtype=np.float64)
init_mocap_pose_left = np.array(
[
-0.17297014,
1.00485877,
1.32773627,
7.06825181e-01,
8.20281078e-06,
-7.07388269e-01,
-5.20399313e-06,
],
dtype=np.float64,
)
init_mocap_pose_right = np.array(
[
0.17297014,
0.9951369,
1.32773623,
2.59463975e-06,
7.07388269e-01,
5.59551158e-06,
7.06825181e-01,
],
dtype=np.float64,
)
left_init_quat = Quaternion(init_mocap_pose_left[3:])
right_init_quat = Quaternion(init_mocap_pose_right[3:])
left_pick_quat = (
left_init_quat * Quaternion(axis=[0.0, 1.0, 0.0], degrees=45)
).elements
right_pick_quat = (
right_init_quat * Quaternion(axis=[0.0, 1.0, 0.0], degrees=45)
).elements
socket_hold_action = np.array(
[socket_xyz[0] - 0.078, socket_xyz[1], self.SOCKET_HOLD_Z], dtype=np.float64
)
peg_init_xyz = peg_xyz + np.array(
[0.078, 0.0, self.ACTION_OBJECT_Z_OFFSET + 0.01]
)
peg_lift_center = np.array(
[peg_xyz[0] + 0.078, socket_hold_action[1], self.SOCKET_HOLD_Z - 0.01],
dtype=np.float64,
)
# The front camera looks along +Y, so visual right-to-left insertion is
# world +X -> -X. With the socket XML in identity orientation, its
# tunnel axis is local/world X, so the peg approaches from +X and stops
# when its leading face reaches the socket's internal pin.
peg_insert_end_center = np.array(
[
socket_hold_action[0] + 0.078 * 2 + 0.04 + 0.06 - 0.01,
socket_hold_action[1],
self.SOCKET_HOLD_Z - 0.01,
],
dtype=np.float64,
)
self.left_trajectory = [
{
"t": 1,
"xyz": init_mocap_pose_left[:3],
"quat": init_mocap_pose_left[3:],
"gripper": 100,
},
{
"t": 130,
"xyz": socket_xyz
+ np.array([-0.078, 0.0, self.ACTION_OBJECT_Z_OFFSET]),
"quat": left_pick_quat,
"gripper": 100,
},
{
"t": 180,
"xyz": socket_xyz
+ np.array([-0.078, 0.0, self.ACTION_OBJECT_Z_OFFSET]),
"quat": left_pick_quat,
"gripper": self.LEFT_SOCKET_GRIPPER_CLOSED,
},
{
"t": 350,
"xyz": socket_hold_action,
"quat": left_pick_quat,
"gripper": self.LEFT_SOCKET_GRIPPER_CLOSED,
},
{
"t": self.EPISODE_END_T,
"xyz": socket_hold_action,
"quat": left_pick_quat,
"gripper": self.LEFT_SOCKET_GRIPPER_CLOSED,
},
]
self.right_trajectory = [
{
"t": 1,
"xyz": init_mocap_pose_right[:3],
"quat": init_mocap_pose_right[3:],
"gripper": 100,
},
{
"t": 80,
"xyz": peg_init_xyz,
"quat": right_pick_quat,
"gripper": 100,
},
{
"t": 150,
"xyz": peg_init_xyz,
"quat": right_pick_quat,
"gripper": 100,
},
{
"t": 180,
"xyz": peg_init_xyz,
"quat": right_pick_quat,
"gripper": self.RIGHT_PEG_GRIPPER_CLOSED,
},
{
"t": 350,
"xyz": peg_init_xyz,
"quat": right_pick_quat,
"gripper": self.RIGHT_PEG_GRIPPER_CLOSED,
},
{
"t": 450,
"xyz": peg_lift_center,
"quat": right_pick_quat,
"gripper": self.RIGHT_PEG_GRIPPER_CLOSED,
},
{
"t": self.INSERT_END_T,
"xyz": peg_insert_end_center,
"quat": right_pick_quat,
"gripper": self.RIGHT_PEG_GRIPPER_CLOSED,
},
{
"t": self.EPISODE_END_T,
"xyz": peg_insert_end_center,
"quat": right_pick_quat,
"gripper": self.RIGHT_PEG_GRIPPER_CLOSED,
},
]
+33 -13
View File
@@ -2,9 +2,11 @@ import time
import os import os
import numpy as np import numpy as np
from roboimi.envs.double_pos_ctrl_env import make_sim_env from roboimi.envs.double_pos_ctrl_env import make_sim_env
from diana_policy import TestPickAndTransferPolicy from roboimi.demos.diana_air_insert_policy import TestAirInsertPolicy
from roboimi.demos.diana_policy import TestPickAndTransferPolicy
import cv2 import cv2
from roboimi.utils.act_ex_utils import sample_transfer_pose from roboimi.utils.act_ex_utils import sample_air_insert_socket_peg_state, sample_transfer_pose
from roboimi.utils.constants import SIM_TASK_CONFIGS
from roboimi.utils.streaming_episode_writer import StreamingEpisodeWriter from roboimi.utils.streaming_episode_writer import StreamingEpisodeWriter
import pathlib import pathlib
@@ -12,16 +14,34 @@ HOME_PATH = str(pathlib.Path(__file__).parent.resolve())
DATASET_DIR = HOME_PATH + '/dataset' DATASET_DIR = HOME_PATH + '/dataset'
def main(): def sample_task_state(task_name):
task_name = 'sim_transfer' if task_name == 'sim_transfer':
dataset_dir = DATASET_DIR + '/sim_transfer' #SIM_TASK_CONFIGS[task_name]['dataset_dir'] return sample_transfer_pose()
num_episodes = 100 #SIM_TASK_CONFIGS[task_name]['num_episodes'] if task_name == 'sim_air_insert_socket_peg':
return sample_air_insert_socket_peg_state()
raise NotImplementedError(f'Unsupported scripted rollout task: {task_name}')
def make_policy(task_name, inject_noise=False, grasp_strategy=None):
if task_name == 'sim_transfer':
return TestPickAndTransferPolicy(inject_noise)
if task_name == 'sim_air_insert_socket_peg':
if grasp_strategy is None:
return TestAirInsertPolicy(inject_noise)
return TestAirInsertPolicy(inject_noise, grasp_strategy=grasp_strategy)
raise NotImplementedError(f'Unsupported scripted rollout task: {task_name}')
def main(task_name='sim_transfer'):
task_cfg = SIM_TASK_CONFIGS[task_name]
dataset_dir = task_cfg['dataset_dir']
num_episodes = 100
inject_noise = False inject_noise = False
episode_len = 700 #SIM_TASK_CONFIGS[task_name]['episode_len'] episode_len = task_cfg['episode_len']
camera_names = ['angle','r_vis', 'top', 'front'] #SIM_TASK_CONFIGS[task_name]['camera_names'] camera_names = task_cfg['camera_names']
image_size = (256, 256) image_size = (256, 256)
if task_name == 'sim_transfer': if task_name in {'sim_transfer', 'sim_air_insert_socket_peg'}:
print(task_name) print(task_name)
else: else:
raise NotImplementedError raise NotImplementedError
@@ -29,7 +49,7 @@ def main():
success = [] success = []
env = make_sim_env(task_name) env = make_sim_env(task_name)
policy = TestPickAndTransferPolicy(inject_noise) policy = make_policy(task_name, inject_noise=inject_noise)
# 等待osmesa完全启动后再开始收集数据 # 等待osmesa完全启动后再开始收集数据
print("等待osmesa线程启动...") print("等待osmesa线程启动...")
@@ -41,8 +61,8 @@ def main():
max_reward = float('-inf') max_reward = float('-inf')
print(f'\n{episode_idx=}') print(f'\n{episode_idx=}')
print('Rollout out EE space scripted policy') print('Rollout out EE space scripted policy')
box_pos = sample_transfer_pose() task_state = sample_task_state(task_name)
env.reset(box_pos) env.reset(task_state)
episode_writer = StreamingEpisodeWriter( episode_writer = StreamingEpisodeWriter(
dataset_path=os.path.join(dataset_dir, f'episode_{episode_idx}.hdf5'), dataset_path=os.path.join(dataset_dir, f'episode_{episode_idx}.hdf5'),
max_timesteps=episode_len, max_timesteps=episode_len,
@@ -50,7 +70,7 @@ def main():
image_size=image_size, image_size=image_size,
) )
for step in range(episode_len): for step in range(episode_len):
raw_action = policy.predict(box_pos,step) raw_action = policy.predict(task_state, step)
env.step(raw_action) env.step(raw_action)
env.render() env.render()
sum_reward += env.rew sum_reward += env.rew
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+23 -3
View File
@@ -702,10 +702,28 @@ def _run_training(cfg: DictConfig):
from roboimi.demos.vla_scripts import eval_vla from roboimi.demos.vla_scripts import eval_vla
rollout_cfg = OmegaConf.create(OmegaConf.to_container(cfg, resolve=False)) rollout_cfg = OmegaConf.create(OmegaConf.to_container(cfg, resolve=False))
rollout_num_episodes = int(cfg.train.get('rollout_num_episodes', 1))
rollout_device = str(cfg.train.get('rollout_device', cfg.train.device))
configured_rollout_workers = cfg.train.get('rollout_num_workers', None)
if configured_rollout_workers is None:
if rollout_device.startswith('cuda'):
rollout_num_workers = min(max(rollout_num_episodes, 1), 8)
else:
rollout_num_workers = 1
else:
rollout_num_workers = int(configured_rollout_workers)
rollout_cfg.eval.ckpt_path = str(checkpoint_path) rollout_cfg.eval.ckpt_path = str(checkpoint_path)
rollout_cfg.eval.num_episodes = int(cfg.train.get('rollout_num_episodes', 1)) rollout_cfg.eval.num_episodes = rollout_num_episodes
rollout_cfg.eval.num_workers = rollout_num_workers
rollout_cfg.eval.headless = True rollout_cfg.eval.headless = True
rollout_cfg.eval.device = 'cpu' rollout_cfg.eval.device = rollout_device
rollout_cfg.eval.cuda_devices = cfg.train.get('rollout_cuda_devices', None)
rollout_cfg.eval.response_timeout_s = float(
cfg.train.get('rollout_response_timeout_s', 300.0)
)
rollout_cfg.eval.server_startup_timeout_s = float(
cfg.train.get('rollout_server_startup_timeout_s', 300.0)
)
rollout_cfg.eval.verbose_action = False rollout_cfg.eval.verbose_action = False
rollout_cfg.eval.record_video = False rollout_cfg.eval.record_video = False
rollout_cfg.eval.save_trajectory_image = True rollout_cfg.eval.save_trajectory_image = True
@@ -716,9 +734,11 @@ def _run_training(cfg: DictConfig):
) )
log.info( log.info(
"🎯 开始 checkpoint rollout 验证: %s (episodes=%s, headless=True)", "🎯 开始 checkpoint rollout 验证: %s (episodes=%s, device=%s, workers=%s, headless=True)",
checkpoint_path, checkpoint_path,
rollout_cfg.eval.num_episodes, rollout_cfg.eval.num_episodes,
rollout_cfg.eval.device,
rollout_cfg.eval.num_workers,
) )
return eval_vla._run_eval(rollout_cfg) return eval_vla._run_eval(rollout_cfg)
+157
View File
@@ -0,0 +1,157 @@
import copy as cp
import time
import numpy as np
from roboimi.envs.double_base import DualDianaMed
from roboimi.envs.double_pos_ctrl_env import DualDianaMed_Pos_Ctrl
SOCKET_JOINT_NAME = "blue_socket_joint"
PEG_JOINT_NAME = "red_peg_joint"
REQUIRED_TASK_STATE_KEYS = ("socket_pos", "socket_quat", "peg_pos", "peg_quat")
SOCKET_GEOM_NAMES = ("socket-1", "socket-2", "socket-3", "socket-4")
SOCKET_SUCCESS_GEOM_NAMES = ("pin",)
SOCKET_BODY_GEOM_NAMES = SOCKET_GEOM_NAMES + SOCKET_SUCCESS_GEOM_NAMES
PEG_GEOM_NAMES = ("red_peg",)
LEFT_GRIPPER_GEOM_NAMES = (
"l_finger_left",
"r_finger_left",
"l_fingertip_g0_left",
"r_fingertip_g0_left",
"l_fingerpad_g0_left",
"r_fingerpad_g0_left",
"l_fingertip_g0_vis_left",
"r_fingertip_g0_vis_left",
)
RIGHT_GRIPPER_GEOM_NAMES = (
"l_finger_right",
"r_finger_right",
"l_fingertip_g0_right",
"r_fingertip_g0_right",
"l_fingerpad_g0_right",
"r_fingerpad_g0_right",
"l_fingertip_g0_vis_right",
"r_fingertip_g0_vis_right",
)
TABLE_GEOM_NAME = "table"
def _set_free_joint_pose(joint, position, quat):
joint.qpos[:3] = np.asarray(position, dtype=np.float64)
joint.qpos[3:7] = np.asarray(quat, dtype=np.float64)
def set_socket_peg_task_state(mj_data, task_state):
if not isinstance(task_state, dict) or tuple(task_state.keys()) != REQUIRED_TASK_STATE_KEYS:
raise ValueError(
"task_state must be an ordered dict-like mapping with keys "
"socket_pos, socket_quat, peg_pos, peg_quat"
)
_set_free_joint_pose(
mj_data.joint(SOCKET_JOINT_NAME),
task_state["socket_pos"],
task_state["socket_quat"],
)
_set_free_joint_pose(
mj_data.joint(PEG_JOINT_NAME),
task_state["peg_pos"],
task_state["peg_quat"],
)
def get_socket_peg_env_state(mj_data):
socket_qpos = cp.deepcopy(np.asarray(mj_data.joint(SOCKET_JOINT_NAME).qpos[:7], dtype=np.float64))
peg_qpos = cp.deepcopy(np.asarray(mj_data.joint(PEG_JOINT_NAME).qpos[:7], dtype=np.float64))
return np.concatenate([socket_qpos, peg_qpos], dtype=np.float64)
def _normalize_contact_pairs(contact_pairs):
return {frozenset(pair) for pair in contact_pairs}
def _has_any_object_contact(contact_set, object_geom_names, other_geom_names):
return any(
frozenset((object_geom_name, other_geom_name)) in contact_set
for object_geom_name in object_geom_names
for other_geom_name in other_geom_names
)
def _object_is_airborne(contact_set, object_geom_names):
return not _has_any_object_contact(contact_set, object_geom_names, (TABLE_GEOM_NAME,))
def peg_inserted_into_socket(contact_pairs):
contact_set = _normalize_contact_pairs(contact_pairs)
return frozenset((PEG_GEOM_NAMES[0], SOCKET_SUCCESS_GEOM_NAMES[0])) in contact_set
def compute_air_insert_reward(contact_pairs, env_state=None):
del env_state # kept for API compatibility with rollout/eval code paths
contact_set = _normalize_contact_pairs(contact_pairs)
reward = 0
if _has_any_object_contact(contact_set, SOCKET_GEOM_NAMES, LEFT_GRIPPER_GEOM_NAMES):
reward += 1
if _has_any_object_contact(contact_set, PEG_GEOM_NAMES, RIGHT_GRIPPER_GEOM_NAMES):
reward += 1
socket_airborne = _object_is_airborne(contact_set, SOCKET_BODY_GEOM_NAMES)
peg_airborne = _object_is_airborne(contact_set, PEG_GEOM_NAMES)
if socket_airborne:
reward += 1
if peg_airborne:
reward += 1
if socket_airborne and peg_airborne and peg_inserted_into_socket(contact_pairs):
reward += 1
return reward
class DualDianaMed_Air_Insert(DualDianaMed_Pos_Ctrl):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.max_reward = 5
def reset(self, task_state):
set_socket_peg_task_state(self.mj_data, task_state)
DualDianaMed.reset(self)
self.top = None
self.r_vis = None
self.l_vis = None
self.front = None
if not self.is_render:
self._update_camera_images_sync()
return
self.cam_flage = True
while self.cam_flage:
if (
type(self.top) == type(None)
or type(self.r_vis) == type(None)
or type(self.l_vis) == type(None)
or type(self.front) == type(None)
):
time.sleep(0.001)
else:
self.cam_flage = False
def step(self, action=np.zeros(16)):
super().step(action)
self.rew = self._get_reward()
self.obs = self._get_obs()
def get_env_state(self):
return get_socket_peg_env_state(self.mj_data)
def _get_reward(self):
contact_pairs = []
for collision_num in range(self.mj_data.ncon):
geom1 = self.mj_data.contact[collision_num].geom1
geom2 = self.mj_data.contact[collision_num].geom2
contact_pairs.append(
(self.getID2Name("geom", geom1), self.getID2Name("geom", geom2))
)
return compute_air_insert_reward(contact_pairs, self.get_env_state())
+6 -13
View File
@@ -52,7 +52,6 @@ class DualDianaMed(MujocoEnv):
self.r_vis = None self.r_vis = None
self.l_vis = None self.l_vis = None
self.top = None self.top = None
self.angle = None
self.front = None self.front = None
self.obs = None self.obs = None
@@ -92,7 +91,6 @@ class DualDianaMed(MujocoEnv):
def step(self,action): def step(self,action):
self.compute_qpos = action #for observation ! self.compute_qpos = action #for observation !
self.obs = self._get_obs()
if self.interpolator_left is not None and self.interpolator_right is not None: if self.interpolator_left is not None and self.interpolator_right is not None:
self.interpolator_left.updateInput(action[:7], control_cycle=self.base_time) self.interpolator_left.updateInput(action[:7], control_cycle=self.base_time)
self.interpolator_right.updateInput(action[7:-2], control_cycle=self.base_time) self.interpolator_right.updateInput(action[7:-2], control_cycle=self.base_time)
@@ -105,6 +103,7 @@ class DualDianaMed(MujocoEnv):
super().step(action) super().step(action)
self.base_time = time.time() - ctrl_cur_time self.base_time = time.time() - ctrl_cur_time
self.obs = self._get_obs()
def preStep(self, action): def preStep(self, action):
@@ -168,10 +167,9 @@ class DualDianaMed(MujocoEnv):
obs['qpos'] = self.get_obs_qpos obs['qpos'] = self.get_obs_qpos
obs['action'] = self.compute_qpos obs['action'] = self.compute_qpos
obs['images'] = dict() obs['images'] = dict()
obs['images']['top'] = self.top
obs['images']['angle'] = self.angle
obs['images']['r_vis'] = self.r_vis obs['images']['r_vis'] = self.r_vis
obs['images']['l_vis'] = self.l_vis obs['images']['l_vis'] = self.l_vis
obs['images']['top'] = self.top
obs['images']['front'] = self.front obs['images']['front'] = self.front
return obs return obs
@@ -180,10 +178,9 @@ class DualDianaMed(MujocoEnv):
self._update_camera_images_sync() self._update_camera_images_sync()
obs = collections.OrderedDict() obs = collections.OrderedDict()
obs['images'] = dict() obs['images'] = dict()
obs['images']['top'] = self.top
obs['images']['angle'] = self.angle
obs['images']['r_vis'] = self.r_vis obs['images']['r_vis'] = self.r_vis
obs['images']['l_vis'] = self.l_vis obs['images']['l_vis'] = self.l_vis
obs['images']['top'] = self.top
obs['images']['front'] = self.front obs['images']['front'] = self.front
return obs return obs
@@ -202,14 +199,12 @@ class DualDianaMed(MujocoEnv):
@property @property
def cam_view(self): def cam_view(self):
if self.cam == 'top': if self.cam == 'r_vis':
return self.top
elif self.cam == 'angle':
return self.angle
elif self.cam == 'r_vis':
return self.r_vis return self.r_vis
elif self.cam == 'l_vis': elif self.cam == 'l_vis':
return self.l_vis return self.l_vis
elif self.cam == 'top':
return self.top
elif self.cam == 'front': elif self.cam == 'front':
return self.front return self.front
else: else:
@@ -230,8 +225,6 @@ class DualDianaMed(MujocoEnv):
self.l_vis = img_renderer.render()[:, :, ::-1] self.l_vis = img_renderer.render()[:, :, ::-1]
img_renderer.update_scene(self.mj_data, camera="top") img_renderer.update_scene(self.mj_data, camera="top")
self.top = img_renderer.render()[:, :, ::-1] self.top = img_renderer.render()[:, :, ::-1]
img_renderer.update_scene(self.mj_data, camera="angle")
self.angle = img_renderer.render()[:, :, ::-1]
img_renderer.update_scene(self.mj_data, camera="front") img_renderer.update_scene(self.mj_data, camera="front")
self.front = img_renderer.render()[:, :, ::-1] self.front = img_renderer.render()[:, :, ::-1]
+15 -4
View File
@@ -73,8 +73,8 @@ class DualDianaMed_Pos_Ctrl(DualDianaMed):
self.mj_data.joint('red_box_joint').qpos[6] = 0.0 self.mj_data.joint('red_box_joint').qpos[6] = 0.0
super().reset() super().reset()
self.top = None self.top = None
self.angle = None
self.r_vis = None self.r_vis = None
self.l_vis = None
self.front = None self.front = None
if not self.is_render: if not self.is_render:
self._update_camera_images_sync() self._update_camera_images_sync()
@@ -83,8 +83,8 @@ class DualDianaMed_Pos_Ctrl(DualDianaMed):
t=0 t=0
while self.cam_flage: while self.cam_flage:
if(type(self.top)==type(None) if(type(self.top)==type(None)
or type(self.angle)==type(None)
or type(self.r_vis)==type(None) or type(self.r_vis)==type(None)
or type(self.l_vis)==type(None)
or type(self.front)==type(None)): or type(self.front)==type(None)):
time.sleep(0.001) time.sleep(0.001)
t+=1 t+=1
@@ -137,6 +137,18 @@ class DualDianaMed_Pos_Ctrl(DualDianaMed):
def make_sim_env(task_name, headless=False): def make_sim_env(task_name, headless=False):
if task_name == 'sim_air_insert_socket_peg':
from roboimi.assets.robots.diana_med import BiDianaMedSocketPeg
from roboimi.envs.double_air_insert_env import DualDianaMed_Air_Insert
env = DualDianaMed_Air_Insert(
robot=BiDianaMedSocketPeg(),
is_render=not headless,
control_freq=30,
is_interpolate=True,
cam_view='front'
)
return env
if 'sim_transfer' in task_name: if 'sim_transfer' in task_name:
from roboimi.assets.robots.diana_med import BiDianaMed from roboimi.assets.robots.diana_med import BiDianaMed
env = DualDianaMed_Pos_Ctrl( env = DualDianaMed_Pos_Ctrl(
@@ -144,7 +156,7 @@ def make_sim_env(task_name, headless=False):
is_render=not headless, is_render=not headless,
control_freq=30, control_freq=30,
is_interpolate=True, is_interpolate=True,
cam_view='angle' cam_view='top'
) )
return env return env
else: else:
@@ -170,4 +182,3 @@ if __name__ == "__main__":
env.step(action) env.step(action)
if env.is_render: if env.is_render:
env.render() env.render()
+20
View File
@@ -1,5 +1,6 @@
import numpy as np import numpy as np
def sample_insertion_pose(): def sample_insertion_pose():
# Peg # Peg
x_range = [0.1, 0.2] x_range = [0.1, 0.2]
@@ -36,3 +37,22 @@ def sample_transfer_pose():
return box_position return box_position
def sample_air_insert_socket_peg_state():
socket_position = np.random.uniform(
low=np.array([-0.20, 0.80, 0.472], dtype=np.float32),
high=np.array([-0.10, 1.00, 0.472], dtype=np.float32),
)
peg_position = np.random.uniform(
low=np.array([0.10, 0.80, 0.46], dtype=np.float32),
high=np.array([0.20, 1.00, 0.46], dtype=np.float32),
)
socket_quat = np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32)
peg_quat = np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32)
return {
"socket_pos": socket_position.astype(np.float32, copy=False),
"socket_quat": socket_quat,
"peg_pos": peg_position.astype(np.float32, copy=False),
"peg_quat": peg_quat,
}
+8 -11
View File
@@ -20,7 +20,14 @@ SIM_TASK_CONFIGS = {
'dataset_dir': DATASET_DIR + '/sim_transfer', 'dataset_dir': DATASET_DIR + '/sim_transfer',
'num_episodes': 20, 'num_episodes': 20,
'episode_len': 700, 'episode_len': 700,
'camera_names': ['top','r_vis','front'], 'camera_names': ['r_vis', 'top', 'front'],
'xml_dir': HOME_PATH + '/assets'
},
'sim_air_insert_socket_peg': {
'dataset_dir': DATASET_DIR + '/sim_air_insert_socket_peg',
'num_episodes': 20,
'episode_len': 750,
'camera_names': ['l_vis', 'r_vis', 'front'],
'xml_dir': HOME_PATH + '/assets' 'xml_dir': HOME_PATH + '/assets'
}, },
@@ -52,13 +59,3 @@ PUPPET_GRIPPER_JOINT_NORMALIZE_FN = lambda x: (x - PUPPET_GRIPPER_JOINT_CLOSE) /
MASTER_GRIPPER_JOINT_UNNORMALIZE_FN = lambda x: x * (MASTER_GRIPPER_JOINT_OPEN - MASTER_GRIPPER_JOINT_CLOSE) + MASTER_GRIPPER_JOINT_CLOSE MASTER_GRIPPER_JOINT_UNNORMALIZE_FN = lambda x: x * (MASTER_GRIPPER_JOINT_OPEN - MASTER_GRIPPER_JOINT_CLOSE) + MASTER_GRIPPER_JOINT_CLOSE
PUPPET_GRIPPER_JOINT_UNNORMALIZE_FN = lambda x: x * (PUPPET_GRIPPER_JOINT_OPEN - PUPPET_GRIPPER_JOINT_CLOSE) + PUPPET_GRIPPER_JOINT_CLOSE PUPPET_GRIPPER_JOINT_UNNORMALIZE_FN = lambda x: x * (PUPPET_GRIPPER_JOINT_OPEN - PUPPET_GRIPPER_JOINT_CLOSE) + PUPPET_GRIPPER_JOINT_CLOSE
MASTER2PUPPET_JOINT_FN = lambda x: PUPPET_GRIPPER_JOINT_UNNORMALIZE_FN(MASTER_GRIPPER_JOINT_NORMALIZE_FN(x)) MASTER2PUPPET_JOINT_FN = lambda x: PUPPET_GRIPPER_JOINT_UNNORMALIZE_FN(MASTER_GRIPPER_JOINT_NORMALIZE_FN(x))
MASTER_GRIPPER_VELOCITY_NORMALIZE_FN = lambda x: x / (MASTER_GRIPPER_POSITION_OPEN - MASTER_GRIPPER_POSITION_CLOSE)
PUPPET_GRIPPER_VELOCITY_NORMALIZE_FN = lambda x: x / (PUPPET_GRIPPER_POSITION_OPEN - PUPPET_GRIPPER_POSITION_CLOSE)
MASTER_POS2JOINT = lambda x: MASTER_GRIPPER_POSITION_NORMALIZE_FN(x) * (MASTER_GRIPPER_JOINT_OPEN - MASTER_GRIPPER_JOINT_CLOSE) + MASTER_GRIPPER_JOINT_CLOSE
MASTER_JOINT2POS = lambda x: MASTER_GRIPPER_POSITION_UNNORMALIZE_FN((x - MASTER_GRIPPER_JOINT_CLOSE) / (MASTER_GRIPPER_JOINT_OPEN - MASTER_GRIPPER_JOINT_CLOSE))
PUPPET_POS2JOINT = lambda x: PUPPET_GRIPPER_POSITION_NORMALIZE_FN(x) * (PUPPET_GRIPPER_JOINT_OPEN - PUPPET_GRIPPER_JOINT_CLOSE) + PUPPET_GRIPPER_JOINT_CLOSE
PUPPET_JOINT2POS = lambda x: PUPPET_GRIPPER_POSITION_UNNORMALIZE_FN((x - PUPPET_GRIPPER_JOINT_CLOSE) / (PUPPET_GRIPPER_JOINT_OPEN - PUPPET_GRIPPER_JOINT_CLOSE))
MASTER_GRIPPER_JOINT_MID = (MASTER_GRIPPER_JOINT_OPEN + MASTER_GRIPPER_JOINT_CLOSE)/2
+5
View File
@@ -29,6 +29,11 @@ train:
rollout_val_freq_epochs: 50 # 每隔多少个 epoch 执行一次 rollout 验证 rollout_val_freq_epochs: 50 # 每隔多少个 epoch 执行一次 rollout 验证
rollout_validate_on_checkpoint: false # 是否在保存 checkpoint 后立即运行 rollout 验证 rollout_validate_on_checkpoint: false # 是否在保存 checkpoint 后立即运行 rollout 验证
rollout_num_episodes: 3 # rollout 验证的回合数 rollout_num_episodes: 3 # rollout 验证的回合数
rollout_device: ${train.device} # rollout 使用的设备;默认跟随训练设备
rollout_num_workers: null # rollout 并行 worker 数;null 时 CUDA 自动推断,CPU 保持 1
rollout_cuda_devices: null # rollout CUDA 并行使用的逻辑 device 列表;null 时默认 [0]
rollout_response_timeout_s: 300.0 # rollout worker 等待 inference server 响应的超时时间
rollout_server_startup_timeout_s: 300.0 # rollout 等待 inference server 就绪的超时时间
# 学习率调度器(带预热) # 学习率调度器(带预热)
warmup_steps: 2000 # 预热步数(Transformer建议更长) warmup_steps: 2000 # 预热步数(Transformer建议更长)
+4
View File
@@ -2,6 +2,10 @@
# 评估配置 # 评估配置
ckpt_path: "checkpoints/vla_model_best.pt" # 模型检查点路径 ckpt_path: "checkpoints/vla_model_best.pt" # 模型检查点路径
num_episodes: 3 # 评估回合数 num_episodes: 3 # 评估回合数
num_workers: 1 # 并行 worker 数;1 表示保持单进程评估
cuda_devices: null # CUDA 并行评估时使用的逻辑设备列表;null 表示默认 [0]
response_timeout_s: 300.0 # worker 等待 inference server 响应的超时时间(秒)
server_startup_timeout_s: 300.0 # parent 等待 inference server 就绪的超时时间(秒)
max_timesteps: 700 # 每回合最大时间步 max_timesteps: 700 # 每回合最大时间步
device: ${train.device} # 与训练保持一致 device: ${train.device} # 与训练保持一致
task_name: "sim_transfer" # 环境任务名称 task_name: "sim_transfer" # 环境任务名称
+510
View File
@@ -0,0 +1,510 @@
import importlib
import inspect
import pathlib
import unittest
from unittest import mock
import xml.etree.ElementTree as ET
import numpy as np
from roboimi.envs.double_pos_ctrl_env import make_sim_env
from roboimi.utils import act_ex_utils
from roboimi.utils.constants import SIM_TASK_CONFIGS
TASK_NAME = "sim_air_insert_socket_peg"
class AirInsertTaskRegistrationTest(unittest.TestCase):
def test_sim_task_configs_registers_air_insert_socket_peg(self):
self.assertIn(TASK_NAME, SIM_TASK_CONFIGS)
self.assertNotIn("sim_air_insert_ring_bar", SIM_TASK_CONFIGS)
self.assertEqual(SIM_TASK_CONFIGS[TASK_NAME]["episode_len"], 750)
self.assertEqual(SIM_TASK_CONFIGS[TASK_NAME]["camera_names"], ["l_vis", "r_vis", "front"])
self.assertTrue(SIM_TASK_CONFIGS[TASK_NAME]["dataset_dir"].endswith("/sim_air_insert_socket_peg"))
def test_sample_air_insert_socket_peg_state_returns_explicit_named_mapping(self):
sampler = getattr(act_ex_utils, "sample_air_insert_socket_peg_state", None)
self.assertIsNotNone(
sampler,
"Expected roboimi.utils.act_ex_utils.sample_air_insert_socket_peg_state()",
)
self.assertFalse(
hasattr(act_ex_utils, "sample_air_insert_ring_bar_state"),
"air insert sampler should use socket/peg naming after the task rename",
)
task_state = sampler()
self.assertEqual(
list(task_state.keys()),
["socket_pos", "socket_quat", "peg_pos", "peg_quat"],
)
self.assertEqual(task_state["socket_pos"].shape, (3,))
self.assertEqual(task_state["socket_quat"].shape, (4,))
self.assertEqual(task_state["peg_pos"].shape, (3,))
self.assertEqual(task_state["peg_quat"].shape, (4,))
def test_sample_air_insert_socket_peg_state_uses_fixed_quats_and_left_right_planar_ranges(self):
sampler = getattr(act_ex_utils, "sample_air_insert_socket_peg_state", None)
self.assertIsNotNone(sampler)
task_state = sampler()
np.testing.assert_array_equal(task_state["socket_quat"], np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32))
np.testing.assert_array_equal(task_state["peg_quat"], np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32))
self.assertGreaterEqual(task_state["socket_pos"][0], -0.20)
self.assertLessEqual(task_state["socket_pos"][0], -0.05)
self.assertGreaterEqual(task_state["socket_pos"][1], 0.70)
self.assertLessEqual(task_state["socket_pos"][1], 1.00)
self.assertAlmostEqual(float(task_state["socket_pos"][2]), 0.472)
self.assertGreaterEqual(task_state["peg_pos"][0], 0.05)
self.assertLessEqual(task_state["peg_pos"][0], 0.20)
self.assertGreaterEqual(task_state["peg_pos"][1], 0.70)
self.assertLessEqual(task_state["peg_pos"][1], 1.00)
self.assertAlmostEqual(float(task_state["peg_pos"][2]), 0.46)
def test_make_sim_env_dispatches_air_insert_socket_peg_headless(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
air_insert_cls = getattr(air_insert_env, "DualDianaMed_Air_Insert", None)
self.assertIsNotNone(air_insert_cls)
diana_med = importlib.import_module("roboimi.assets.robots.diana_med")
socket_peg_robot_cls = getattr(diana_med, "BiDianaMedSocketPeg", None)
self.assertIsNotNone(
socket_peg_robot_cls,
"Expected roboimi.assets.robots.diana_med.BiDianaMedSocketPeg",
)
fake_env = object()
with mock.patch.object(
diana_med,
"BiDianaMedSocketPeg",
return_value="robot",
), mock.patch.object(
air_insert_env,
"DualDianaMed_Air_Insert",
return_value=fake_env,
) as env_cls:
env = make_sim_env(TASK_NAME, headless=True)
self.assertIs(env, fake_env)
env_cls.assert_called_once_with(
robot="robot",
is_render=False,
control_freq=30,
is_interpolate=True,
cam_view="front",
)
def test_diana_table_scene_exposes_only_top_and_front_scene_cameras(self):
xml_path = (
pathlib.Path(__file__).resolve().parents[1]
/ "roboimi/assets/models/manipulators/DianaMed/table_square.xml"
)
root = ET.parse(xml_path).getroot()
cameras = {camera.attrib["name"]: camera.attrib for camera in root.findall(".//camera")}
self.assertNotIn("angle", cameras, "DianaMed scene should stop exposing the old angle camera")
self.assertNotIn("left_side", cameras, "DianaMed scene should no longer expose left_side")
self.assertIn("top", cameras)
self.assertIn("front", cameras)
self.assertEqual(cameras["top"].get("mode"), "targetbody")
self.assertEqual(cameras["top"].get("target"), "table")
class AirInsertResetAndStateHelpersTest(unittest.TestCase):
def test_set_socket_peg_task_state_writes_free_joint_qpos(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
setter = getattr(air_insert_env, "set_socket_peg_task_state", None)
self.assertIsNotNone(
setter,
"Expected roboimi.envs.double_air_insert_env.set_socket_peg_task_state",
)
socket_qpos = np.zeros(7, dtype=np.float64)
peg_qpos = np.zeros(7, dtype=np.float64)
class _FakeJoint:
def __init__(self, qpos):
self.qpos = qpos
class _FakeData:
def joint(self, name):
if name == "blue_socket_joint":
return _FakeJoint(socket_qpos)
if name == "red_peg_joint":
return _FakeJoint(peg_qpos)
raise AssertionError(f"Unexpected joint name: {name}")
task_state = {
"socket_pos": np.array([-0.12, 0.90, 0.472], dtype=np.float64),
"socket_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float64),
"peg_pos": np.array([0.12, 0.91, 0.46], dtype=np.float64),
"peg_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float64),
}
setter(_FakeData(), task_state)
np.testing.assert_array_equal(
socket_qpos,
np.array([-0.12, 0.90, 0.472, 1.0, 0.0, 0.0, 0.0], dtype=np.float64),
)
np.testing.assert_array_equal(
peg_qpos,
np.array([0.12, 0.91, 0.46, 1.0, 0.0, 0.0, 0.0], dtype=np.float64),
)
def test_get_socket_peg_env_state_returns_stable_14d_vector(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
getter = getattr(air_insert_env, "get_socket_peg_env_state", None)
self.assertIsNotNone(
getter,
"Expected roboimi.envs.double_air_insert_env.get_socket_peg_env_state",
)
socket_qpos = np.array([-0.12, 0.90, 0.472, 1.0, 0.0, 0.0, 0.0], dtype=np.float64)
peg_qpos = np.array([0.12, 0.91, 0.46, 1.0, 0.0, 0.0, 0.0], dtype=np.float64)
class _FakeJoint:
def __init__(self, qpos):
self.qpos = qpos
class _FakeData:
def joint(self, name):
if name == "blue_socket_joint":
return _FakeJoint(socket_qpos)
if name == "red_peg_joint":
return _FakeJoint(peg_qpos)
raise AssertionError(f"Unexpected joint name: {name}")
env_state = getter(_FakeData())
self.assertEqual(env_state.shape, (14,))
np.testing.assert_array_equal(
env_state,
np.array(
[-0.12, 0.90, 0.472, 1.0, 0.0, 0.0, 0.0, 0.12, 0.91, 0.46, 1.0, 0.0, 0.0, 0.0],
dtype=np.float64,
),
)
def test_air_insert_env_does_not_script_attach_or_assist_objects_after_reset(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
env_cls = getattr(air_insert_env, "DualDianaMed_Air_Insert", None)
self.assertIsNotNone(env_cls)
source = inspect.getsource(env_cls)
self.assertNotIn("_update_scripted_grasped_objects", source)
self.assertNotIn("_scripted_", source)
self.assertNotIn("_stabilize_ring_grasp", source)
self.assertNotIn("_ring_grasp_locked", source)
get_reward_source = inspect.getsource(env_cls._get_reward)
self.assertNotIn("ring_block", get_reward_source)
self.assertNotIn("bar_block", get_reward_source)
def test_socket_peg_xml_defines_active_socket_and_peg_objects(self):
asset_dir = pathlib.Path(__file__).resolve().parents[1] / "roboimi/assets/models/manipulators/DianaMed"
xml_path = asset_dir / "socket_peg_objects.xml"
self.assertTrue(xml_path.exists(), "socket/peg objects should live in socket_peg_objects.xml")
self.assertFalse((asset_dir / "ring_bar_objects.xml").exists(), "old ring_bar_objects.xml should be renamed")
root = ET.parse(xml_path).getroot()
body_names = {body.attrib.get("name") for body in root.findall(".//body")}
geom_names = {geom.attrib.get("name") for geom in root.findall(".//geom")}
joint_names = {joint.attrib.get("name") for joint in root.findall(".//joint")}
self.assertIn("socket", body_names)
self.assertIn("peg", body_names)
self.assertNotIn("ring_block", body_names)
self.assertNotIn("bar_block", body_names)
self.assertIn("blue_socket_joint", joint_names)
self.assertIn("red_peg_joint", joint_names)
for geom_name in ("socket-1", "socket-2", "socket-3", "socket-4", "pin", "red_peg"):
self.assertIn(geom_name, geom_names)
def test_socket_peg_wrapper_includes_socket_peg_objects(self):
xml_path = (
pathlib.Path(__file__).resolve().parents[1]
/ "roboimi/assets/models/manipulators/DianaMed/bi_diana_socket_peg_ee.xml"
)
self.assertTrue(xml_path.exists(), "socket/peg wrapper XML should use the new task name")
root = ET.parse(xml_path).getroot()
includes = [include.attrib.get("file") for include in root.findall(".//include")]
self.assertIn("./socket_peg_objects.xml", includes)
self.assertNotIn("./ring_bar_objects.xml", includes)
class AirInsertRewardAndSuccessTest(unittest.TestCase):
@staticmethod
def _make_env_state(
socket_pos=(0.0, 0.0, 0.472),
socket_quat=(1.0, 0.0, 0.0, 0.0),
peg_pos=(0.0, 0.0, 0.46),
peg_quat=(1.0, 0.0, 0.0, 0.0),
):
return np.array([*socket_pos, *socket_quat, *peg_pos, *peg_quat], dtype=np.float64)
def test_compute_air_insert_reward_counts_left_contact_stage(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
reward_fn = getattr(air_insert_env, "compute_air_insert_reward", None)
self.assertIsNotNone(reward_fn)
reward = reward_fn(
contact_pairs=[
("socket-1", "l_finger_left"),
("socket-1", "table"),
("red_peg", "table"),
],
env_state=self._make_env_state(),
)
self.assertEqual(reward, 1)
def test_compute_air_insert_reward_counts_right_contact_stage(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
reward_fn = getattr(air_insert_env, "compute_air_insert_reward", None)
reward = reward_fn(
contact_pairs=[
("socket-1", "l_finger_left"),
("red_peg", "l_finger_right"),
("socket-1", "table"),
("red_peg", "table"),
],
env_state=self._make_env_state(),
)
self.assertEqual(reward, 2)
def test_compute_air_insert_reward_counts_lift_stages(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
reward_fn = getattr(air_insert_env, "compute_air_insert_reward", None)
reward = reward_fn(
contact_pairs=[
("socket-1", "l_finger_left"),
("red_peg", "l_finger_right"),
],
env_state=self._make_env_state(),
)
self.assertEqual(reward, 4)
def test_compute_air_insert_reward_counts_visual_fingertip_contacts_as_gripper_contacts(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
reward_fn = getattr(air_insert_env, "compute_air_insert_reward", None)
reward = reward_fn(
contact_pairs=[
("socket-3", "r_fingertip_g0_vis_left"),
("red_peg", "l_fingertip_g0_vis_right"),
],
env_state=self._make_env_state(),
)
self.assertEqual(
reward,
4,
"visual fingertip geoms are collidable in the Diana XML and should count as gripper-object contacts",
)
def test_peg_inserted_into_socket_uses_pin_contact(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
success_fn = getattr(air_insert_env, "peg_inserted_into_socket", None)
self.assertIsNotNone(
success_fn,
"Expected roboimi.envs.double_air_insert_env.peg_inserted_into_socket",
)
self.assertTrue(success_fn([("red_peg", "pin")]))
self.assertTrue(success_fn([("pin", "red_peg")]))
self.assertFalse(success_fn([("red_peg", "socket-1")]))
def test_compute_air_insert_reward_requires_airborne_success_for_final_point(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
reward_fn = getattr(air_insert_env, "compute_air_insert_reward", None)
reward = reward_fn(
contact_pairs=[
("socket-1", "l_finger_left"),
("red_peg", "l_finger_right"),
("socket-1", "table"),
("red_peg", "pin"),
],
env_state=self._make_env_state(),
)
self.assertEqual(reward, 3)
def test_compute_air_insert_reward_returns_full_score_on_true_airborne_insert(self):
air_insert_env = importlib.import_module("roboimi.envs.double_air_insert_env")
reward_fn = getattr(air_insert_env, "compute_air_insert_reward", None)
reward = reward_fn(
contact_pairs=[
("socket-1", "l_finger_left"),
("red_peg", "l_finger_right"),
("red_peg", "pin"),
],
env_state=self._make_env_state(),
)
self.assertEqual(reward, 5)
class AirInsertPolicyAndSmokeTest(unittest.TestCase):
@staticmethod
def _canonical_task_state():
return {
"socket_pos": np.array([-0.12, 0.90, 0.472], dtype=np.float32),
"socket_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32),
"peg_pos": np.array([0.12, 0.90, 0.46], dtype=np.float32),
"peg_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32),
}
def test_air_insert_policy_emits_valid_16d_action(self):
policy_module = importlib.import_module("roboimi.demos.diana_air_insert_policy")
policy_cls = getattr(policy_module, "TestAirInsertPolicy", None)
self.assertIsNotNone(policy_cls)
task_state = act_ex_utils.sample_air_insert_socket_peg_state()
policy = policy_cls(inject_noise=False)
action = policy.predict(task_state, 0)
self.assertEqual(action.shape, (16,))
np.testing.assert_array_equal(action[-2:], np.array([100, 100]))
def test_air_insert_policy_inserts_peg_front_view_right_to_left_along_world_x(self):
policy_module = importlib.import_module("roboimi.demos.diana_air_insert_policy")
policy_cls = getattr(policy_module, "TestAirInsertPolicy", None)
self.assertIsNotNone(policy_cls)
task_state = self._canonical_task_state()
policy = policy_cls(inject_noise=False)
policy.generate_trajectory(task_state)
start_waypoint = next(wp for wp in policy.right_trajectory if wp["t"] == policy.INSERT_START_T)
end_waypoint = next(wp for wp in policy.right_trajectory if wp["t"] == policy.INSERT_END_T)
self.assertLess(
end_waypoint["xyz"][0],
start_waypoint["xyz"][0] - 0.10,
"front-view right-to-left peg insertion should decrease world x substantially",
)
self.assertAlmostEqual(float(end_waypoint["xyz"][1]), float(start_waypoint["xyz"][1]), delta=0.02)
expected_insert_end_x = float(task_state["socket_pos"][0] + 0.168)
self.assertAlmostEqual(float(end_waypoint["xyz"][0]), expected_insert_end_x, delta=0.02)
self.assertGreater(float(start_waypoint["xyz"][2]), 0.70)
def test_air_insert_policy_default_left_grasps_socket_and_right_grasps_peg(self):
policy_module = importlib.import_module("roboimi.demos.diana_air_insert_policy")
policy_cls = getattr(policy_module, "TestAirInsertPolicy", None)
self.assertIsNotNone(policy_cls)
task_state = {
"socket_pos": np.array([-0.18, 0.78, 0.472], dtype=np.float32),
"socket_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32),
"peg_pos": np.array([0.16, 0.98, 0.46], dtype=np.float32),
"peg_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32),
}
policy = policy_cls(inject_noise=False)
policy.generate_trajectory(task_state)
left_close = next(wp for wp in policy.left_trajectory if wp["t"] == 180)
right_close = next(wp for wp in policy.right_trajectory if wp["t"] == 180)
action_z_offset = getattr(policy_cls, "ACTION_OBJECT_Z_OFFSET", 0.11)
expected_socket_pick = task_state["socket_pos"] + np.array([-0.078, 0.0, action_z_offset])
expected_peg_pick = task_state["peg_pos"] + np.array([0.078, 0.0, action_z_offset + 0.01])
np.testing.assert_allclose(left_close["xyz"], expected_socket_pick, atol=1e-6)
np.testing.assert_allclose(right_close["xyz"], expected_peg_pick, atol=1e-6)
self.assertLess(left_close["gripper"], 0, "default policy should close the left gripper on the socket")
self.assertLess(right_close["gripper"], 0, "default policy should close the right gripper on the peg")
def test_air_insert_policy_socket_hold_tracks_socket_xy_without_sweeping_laterally(self):
policy_module = importlib.import_module("roboimi.demos.diana_air_insert_policy")
policy_cls = getattr(policy_module, "TestAirInsertPolicy", None)
self.assertIsNotNone(policy_cls)
base_state = {
"socket_pos": np.array([-0.20, 0.72, 0.472], dtype=np.float32),
"socket_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32),
"peg_pos": np.array([0.14, 0.76, 0.46], dtype=np.float32),
"peg_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32),
}
shifted_state = dict(base_state)
shifted_state["socket_pos"] = np.array([-0.06, 0.99, 0.472], dtype=np.float32)
base_policy = policy_cls(inject_noise=False)
base_policy.generate_trajectory(base_state)
shifted_policy = policy_cls(inject_noise=False)
shifted_policy.generate_trajectory(shifted_state)
base_hold = next(wp for wp in base_policy.left_trajectory if wp["t"] == 450)
shifted_hold = next(wp for wp in shifted_policy.left_trajectory if wp["t"] == 450)
np.testing.assert_allclose(
base_hold["xyz"][:2],
base_state["socket_pos"][:2] + np.array([-0.078, 0.0]),
atol=1e-6,
)
np.testing.assert_allclose(
shifted_hold["xyz"][:2],
shifted_state["socket_pos"][:2] + np.array([-0.078, 0.0]),
atol=1e-6,
)
def test_air_insert_policy_predicts_through_full_episode_without_exhausting_waypoints(self):
policy_module = importlib.import_module("roboimi.demos.diana_air_insert_policy")
policy_cls = getattr(policy_module, "TestAirInsertPolicy", None)
self.assertIsNotNone(policy_cls)
task_state = self._canonical_task_state()
policy = policy_cls(inject_noise=False)
for step in range(SIM_TASK_CONFIGS[TASK_NAME]["episode_len"]):
action = policy.predict(task_state, step)
self.assertEqual(action.shape, (16,))
def test_scripted_rollout_entrypoint_selects_socket_peg_sampler_and_policy(self):
rollout_module = importlib.import_module("roboimi.demos.diana_record_sim_episodes")
sampler_fn = getattr(rollout_module, "sample_task_state", None)
policy_factory = getattr(rollout_module, "make_policy", None)
self.assertIsNotNone(sampler_fn)
self.assertIsNotNone(policy_factory)
task_state = sampler_fn(TASK_NAME)
self.assertEqual(list(task_state.keys()), ["socket_pos", "socket_quat", "peg_pos", "peg_quat"])
policy = policy_factory(TASK_NAME, inject_noise=False)
self.assertEqual(policy.__class__.__name__, "TestAirInsertPolicy")
def test_real_headless_smoke_instantiates_resets_and_steps_new_task_once(self):
policy_module = importlib.import_module("roboimi.demos.diana_air_insert_policy")
policy_cls = getattr(policy_module, "TestAirInsertPolicy", None)
self.assertIsNotNone(policy_cls)
task_state = act_ex_utils.sample_air_insert_socket_peg_state()
env = make_sim_env(TASK_NAME, headless=True)
policy = policy_cls(inject_noise=False)
try:
env.reset(task_state)
action = policy.predict(task_state, 0)
env.step(action)
self.assertIsNotNone(env.obs)
self.assertIn("qpos", env.obs)
self.assertIn("images", env.obs)
finally:
env.exit_flag = True
cam_thread = getattr(env, "cam_thread", None)
if cam_thread is not None:
cam_thread.join(timeout=1.0)
viewer = getattr(env, "viewer", None)
if viewer is not None:
viewer.close()
if __name__ == "__main__":
unittest.main()
+704
View File
@@ -1,5 +1,11 @@
import unittest import unittest
from unittest import mock
import numpy as np
import torch
from omegaconf import OmegaConf
from roboimi.demos.vla_scripts import eval_vla
from roboimi.vla.eval_utils import execute_policy_action from roboimi.vla.eval_utils import execute_policy_action
@@ -14,6 +20,48 @@ class _FakeEnv:
self.calls.append(("step_jnt", action)) self.calls.append(("step_jnt", action))
class _FakeQueue:
def __init__(self, initial_items=None):
self.items = list(initial_items or [])
self.put_calls = []
def put(self, item):
self.put_calls.append(item)
self.items.append(item)
def get(self, timeout=None):
del timeout
if not self.items:
raise AssertionError("queue unexpectedly empty")
return self.items.pop(0)
def _make_parallel_cfg(**eval_overrides):
eval_cfg = {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 5,
"num_workers": 2,
"max_timesteps": 1,
"device": "cpu",
"task_name": "sim_transfer",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": True,
"artifact_dir": None,
"save_artifacts": False,
"save_summary_json": False,
"save_timing": False,
"save_trajectory": False,
"save_trajectory_npz": False,
"record_video": False,
"save_trajectory_image": False,
}
eval_cfg.update(eval_overrides)
return OmegaConf.create({"agent": {}, "eval": eval_cfg})
class EvalVLAExecutionTest(unittest.TestCase): class EvalVLAExecutionTest(unittest.TestCase):
def test_execute_policy_action_uses_ee_step(self): def test_execute_policy_action_uses_ee_step(self):
env = _FakeEnv() env = _FakeEnv()
@@ -23,6 +71,662 @@ class EvalVLAExecutionTest(unittest.TestCase):
self.assertEqual(env.calls, [("step", action)]) self.assertEqual(env.calls, [("step", action)])
def test_split_episode_indices_balances_workers(self):
self.assertEqual(
eval_vla._split_episode_indices(num_episodes=10, num_workers=3),
[[0, 1, 2, 3], [4, 5, 6], [7, 8, 9]],
)
def test_normalize_num_workers_caps_worker_count_to_episode_count(self):
self.assertEqual(eval_vla._normalize_num_workers(num_workers=5, num_episodes=2), 2)
def test_plan_episode_box_poses_uses_global_episode_order(self):
planned_poses = [
np.array([0.1, 0.2, 0.3], dtype=np.float32),
np.array([1.1, 1.2, 1.3], dtype=np.float32),
np.array([2.1, 2.2, 2.3], dtype=np.float32),
]
sampler = mock.Mock(side_effect=planned_poses)
result = eval_vla._plan_episode_box_poses(num_episodes=3, sampler=sampler)
self.assertEqual(sampler.call_count, 3)
self.assertEqual(len(result), 3)
for expected, actual in zip(planned_poses, result):
np.testing.assert_array_equal(actual, expected)
def test_resolve_policy_camera_names_matches_vlaagent_fallback_sorting(self):
cfg = OmegaConf.create(
{
"agent": {
"_target_": "roboimi.vla.agent.VLAAgent",
},
"eval": {
"camera_names": ["r_vis", "top", "front"],
},
}
)
self.assertEqual(
eval_vla._resolve_policy_camera_names(cfg),
["front", "r_vis", "top"],
)
def test_resolve_policy_camera_names_matches_gr00t_fallback_input_order(self):
cfg = OmegaConf.create(
{
"agent": {
"_target_": "roboimi.vla.agent_gr00t_dit.VLAAgentGr00tDiT",
},
"eval": {
"camera_names": ["r_vis", "top", "front"],
},
}
)
self.assertEqual(
eval_vla._resolve_policy_camera_names(cfg),
["r_vis", "top", "front"],
)
def test_build_episode_plans_without_box_poses_keeps_serial_sampling_lazy(self):
plans = eval_vla._build_episode_plans(num_episodes=3)
self.assertEqual(
plans,
[
{"episode_index": 0},
{"episode_index": 1},
{"episode_index": 2},
],
)
def test_prepare_local_policy_batch_pads_latest_observation_to_obs_horizon(self):
queues = eval_vla._new_local_policy_queues(obs_horizon=3)
observation = {
"qpos": torch.tensor([1.0, 2.0], dtype=torch.float32),
"images": {
"front": torch.tensor([[[1.0]]], dtype=torch.float32),
},
}
eval_vla._populate_local_policy_queues(queues, observation)
batch = eval_vla._prepare_local_policy_batch(
queues,
obs_horizon=3,
camera_names=["front"],
)
self.assertEqual(tuple(batch["qpos"].shape), (1, 3, 2))
self.assertEqual(tuple(batch["images"]["front"].shape), (1, 3, 1, 1, 1))
np.testing.assert_array_equal(
batch["qpos"][0].cpu().numpy(),
np.array([[1.0, 2.0], [1.0, 2.0], [1.0, 2.0]], dtype=np.float32),
)
np.testing.assert_array_equal(
batch["images"]["front"][0].cpu().numpy(),
np.array([[[[1.0]]], [[[1.0]]], [[[1.0]]]], dtype=np.float32),
)
def test_enqueue_predicted_actions_uses_executable_slice(self):
queues = eval_vla._new_local_policy_queues(obs_horizon=2)
predicted_actions = torch.tensor(
[[[10.0], [20.0], [30.0], [40.0]]],
dtype=torch.float32,
)
eval_vla._enqueue_predicted_actions(
queues,
predicted_actions=predicted_actions,
obs_horizon=2,
num_action_steps=2,
)
self.assertEqual(len(queues["action"]), 2)
np.testing.assert_array_equal(queues["action"].popleft().numpy(), np.array([20.0], dtype=np.float32))
np.testing.assert_array_equal(queues["action"].popleft().numpy(), np.array([30.0], dtype=np.float32))
def test_remote_policy_runner_only_requests_server_inference_when_local_action_queue_is_empty(self):
request_queue = _FakeQueue()
response_queue = _FakeQueue(
[
{
"type": "predict_chunk_result",
"actions": np.asarray([[[10.0], [20.0], [30.0]]], dtype=np.float32),
}
]
)
runner = eval_vla._RemotePolicyRunner(
worker_index=3,
server_index=1,
request_queue=request_queue,
response_queue=response_queue,
camera_names=["front"],
obs_horizon=2,
num_action_steps=2,
)
first_observation = {
"qpos": torch.tensor([1.0, 2.0], dtype=torch.float32),
"images": {"front": torch.tensor([[[1.0]]], dtype=torch.float32)},
}
second_observation = {
"qpos": torch.tensor([3.0, 4.0], dtype=torch.float32),
"images": {"front": torch.tensor([[[2.0]]], dtype=torch.float32)},
}
first_action, first_forward = runner.select_action(
first_observation,
episode_index=7,
timestep=0,
)
second_action, second_forward = runner.select_action(
second_observation,
episode_index=7,
timestep=1,
)
self.assertTrue(first_forward)
self.assertFalse(second_forward)
self.assertEqual(len(request_queue.put_calls), 1)
self.assertEqual(request_queue.put_calls[0]["type"], "predict_chunk")
self.assertEqual(request_queue.put_calls[0]["worker_index"], 3)
self.assertEqual(request_queue.put_calls[0]["server_index"], 1)
np.testing.assert_array_equal(first_action.numpy(), np.array([20.0], dtype=np.float32))
np.testing.assert_array_equal(second_action.numpy(), np.array([30.0], dtype=np.float32))
def test_merge_worker_summaries_sorts_episodes_and_recomputes_aggregates(self):
worker_summaries = [
{
"avg_inference_fps": 999.0,
"avg_control_fps": 999.0,
"avg_obs_read_time_ms": 999.0,
"avg_total_time_ms": 999.0,
"timing_summary": {"count": 999, "model_forward_count": 999},
"episodes": [
{
"episode_index": 2,
"episode_reward": 9.0,
"episode_max_reward": 4.0,
"inference_fps": 30.0,
"control_fps": 15.0,
}
],
"_merge_state": {
"obs_read_time_ms": [9.0],
"preprocess_time_ms": [1.0],
"inference_time_ms": [3.0],
"env_step_time_ms": [4.0],
"total_time_ms": [10.0],
"model_forward_flags": [False],
},
},
{
"avg_inference_fps": 888.0,
"avg_control_fps": 888.0,
"avg_obs_read_time_ms": 888.0,
"avg_total_time_ms": 888.0,
"timing_summary": {"count": 888, "model_forward_count": 888},
"episodes": [
{
"episode_index": 1,
"episode_reward": 6.0,
"episode_max_reward": 3.0,
"inference_fps": 20.0,
"control_fps": 10.0,
},
{
"episode_index": 0,
"episode_reward": 5.0,
"episode_max_reward": 2.0,
"inference_fps": 10.0,
"control_fps": 5.0,
},
],
"_merge_state": {
"obs_read_time_ms": [1.0, 2.0, 12.0],
"preprocess_time_ms": [2.0, 3.0, 4.0],
"inference_time_ms": [4.0, 5.0, 6.0],
"env_step_time_ms": [6.0, 7.0, 8.0],
"total_time_ms": [8.0, 9.0, 20.0],
"model_forward_flags": [True, False, True],
},
},
]
artifact_paths = {
"output_dir": "/tmp/merged",
"summary_json": "/tmp/merged/rollout_summary.json",
"timing_json": "/tmp/merged/timing.json",
"trajectory_npz": None,
"video_mp4": None,
"video_camera_name": None,
}
merged = eval_vla._merge_worker_summaries(worker_summaries, artifact_paths)
self.assertEqual([episode["episode_index"] for episode in merged["episodes"]], [0, 1, 2])
self.assertEqual(merged["episode_rewards"], [5.0, 6.0, 9.0])
self.assertEqual(merged["episode_max_rewards"], [2.0, 3.0, 4.0])
self.assertAlmostEqual(merged["avg_reward"], 20.0 / 3.0)
self.assertAlmostEqual(merged["avg_max_reward"], 3.0)
self.assertAlmostEqual(merged["avg_inference_fps"], 20.0)
self.assertAlmostEqual(merged["avg_control_fps"], 10.0)
self.assertAlmostEqual(merged["avg_obs_read_time_ms"], 6.0)
self.assertAlmostEqual(merged["avg_total_time_ms"], 47.0 / 4.0)
self.assertEqual(merged["timing_summary"]["count"], 4)
self.assertEqual(merged["timing_summary"]["model_forward_count"], 2)
self.assertEqual(merged["artifact_dir"], "/tmp/merged")
self.assertEqual(merged["artifacts"], artifact_paths)
def test_build_cuda_server_payloads_uses_round_robin_worker_assignment(self):
cfg = _make_parallel_cfg(num_episodes=4, num_workers=4, device="cuda", cuda_devices=[0, 1])
artifact_paths = {"output_dir": None}
with mock.patch.object(
eval_vla,
"sample_transfer_pose",
side_effect=[
np.array([0.1, 0.2, 0.3], dtype=np.float32),
np.array([0.4, 0.5, 0.6], dtype=np.float32),
np.array([0.7, 0.8, 0.9], dtype=np.float32),
np.array([1.0, 1.1, 1.2], dtype=np.float32),
],
):
worker_payloads, _ = eval_vla._build_parallel_worker_payloads(cfg, artifact_paths)
server_payloads, assigned_workers = eval_vla._build_cuda_server_payloads(
cfg,
worker_payloads=worker_payloads,
cuda_devices=[0, 1],
)
self.assertEqual([payload["device_index"] for payload in server_payloads], [0, 1])
self.assertEqual([payload["worker_index"] for payload in assigned_workers], [0, 1, 2, 3])
self.assertEqual([payload["server_index"] for payload in assigned_workers], [0, 1, 0, 1])
self.assertEqual(server_payloads[0]["worker_indices"], [0, 2])
self.assertEqual(server_payloads[1]["worker_indices"], [1, 3])
def test_run_eval_parallel_dispatches_episode_splits_and_box_poses(self):
cfg = _make_parallel_cfg(num_episodes=5, num_workers=2, artifact_dir="/tmp/parallel-root")
planned_poses = [
np.array([float(index), float(index) + 0.1, float(index) + 0.2], dtype=np.float32)
for index in range(5)
]
observed_payloads = []
def fake_run_spawn_jobs(payloads, max_workers, worker_fn):
del worker_fn
self.assertEqual(max_workers, 2)
observed_payloads.extend(payloads)
return [
{
"episodes": [
{
"episode_index": 4,
"episode_reward": 5.0,
"episode_max_reward": 5.0,
"inference_fps": 50.0,
"control_fps": 25.0,
},
{
"episode_index": 3,
"episode_reward": 4.0,
"episode_max_reward": 4.0,
"inference_fps": 40.0,
"control_fps": 20.0,
},
],
"_merge_state": {
"obs_read_time_ms": [4.0, 5.0],
"preprocess_time_ms": [1.0, 1.0],
"inference_time_ms": [2.0, 2.0],
"env_step_time_ms": [3.0, 3.0],
"total_time_ms": [4.0, 5.0],
"model_forward_flags": [True, True],
},
},
{
"episodes": [
{
"episode_index": 2,
"episode_reward": 3.0,
"episode_max_reward": 3.0,
"inference_fps": 30.0,
"control_fps": 15.0,
},
{
"episode_index": 1,
"episode_reward": 2.0,
"episode_max_reward": 2.0,
"inference_fps": 20.0,
"control_fps": 10.0,
},
{
"episode_index": 0,
"episode_reward": 1.0,
"episode_max_reward": 1.0,
"inference_fps": 10.0,
"control_fps": 5.0,
},
],
"_merge_state": {
"obs_read_time_ms": [1.0, 2.0, 3.0],
"preprocess_time_ms": [1.0, 1.0, 1.0],
"inference_time_ms": [2.0, 2.0, 2.0],
"env_step_time_ms": [3.0, 3.0, 3.0],
"total_time_ms": [1.0, 2.0, 3.0],
"model_forward_flags": [False, True, False],
},
},
]
with mock.patch.object(
eval_vla,
"sample_transfer_pose",
side_effect=planned_poses,
), mock.patch.object(
eval_vla,
"_run_spawn_jobs",
side_effect=fake_run_spawn_jobs,
):
summary = eval_vla._run_eval_parallel(cfg)
self.assertEqual(len(observed_payloads), 2)
self.assertEqual(
[[plan["episode_index"] for plan in payload["episode_plans"]] for payload in observed_payloads],
[[0, 1, 2], [3, 4]],
)
for payload in observed_payloads:
for plan in payload["episode_plans"]:
np.testing.assert_array_equal(
np.asarray(plan["box_pos"], dtype=np.float32),
planned_poses[plan["episode_index"]],
)
self.assertEqual([episode["episode_index"] for episode in summary["episodes"]], [0, 1, 2, 3, 4])
self.assertEqual(summary["episode_rewards"], [1.0, 2.0, 3.0, 4.0, 5.0])
self.assertEqual(summary["num_episodes"], 5)
def test_run_eval_parallel_allows_trajectory_images_and_keeps_worker_artifact_paths(self):
cfg = _make_parallel_cfg(
num_episodes=2,
num_workers=2,
artifact_dir="/tmp/parallel-images",
save_summary_json=True,
save_trajectory_image=True,
)
observed_payloads = []
def fake_run_spawn_jobs(payloads, max_workers, worker_fn):
del worker_fn
self.assertEqual(max_workers, 2)
observed_payloads.extend(payloads)
return [
{
"episodes": [
{
"episode_index": 0,
"episode_reward": 1.0,
"episode_max_reward": 1.0,
"inference_fps": 10.0,
"control_fps": 5.0,
"artifact_paths": {
"trajectory_image": f"{payloads[0]['artifact_dir']}/rollout_front_ep01_trajectory.png",
},
},
],
"_merge_state": {
"obs_read_time_ms": [1.0],
"preprocess_time_ms": [1.0],
"inference_time_ms": [1.0],
"env_step_time_ms": [1.0],
"total_time_ms": [1.0],
"model_forward_flags": [True],
},
},
{
"episodes": [
{
"episode_index": 1,
"episode_reward": 2.0,
"episode_max_reward": 2.0,
"inference_fps": 20.0,
"control_fps": 10.0,
"artifact_paths": {
"trajectory_image": f"{payloads[1]['artifact_dir']}/rollout_front_ep02_trajectory.png",
},
},
],
"_merge_state": {
"obs_read_time_ms": [2.0],
"preprocess_time_ms": [2.0],
"inference_time_ms": [2.0],
"env_step_time_ms": [2.0],
"total_time_ms": [2.0],
"model_forward_flags": [False],
},
},
]
with mock.patch.object(
eval_vla,
"sample_transfer_pose",
side_effect=[
np.array([0.1, 0.2, 0.3], dtype=np.float32),
np.array([0.4, 0.5, 0.6], dtype=np.float32),
],
), mock.patch.object(
eval_vla,
"_run_spawn_jobs",
side_effect=fake_run_spawn_jobs,
):
summary = eval_vla._run_eval_parallel(cfg)
self.assertEqual(len(observed_payloads), 2)
self.assertTrue(observed_payloads[0]["artifact_dir"].endswith("workers/worker_00"))
self.assertTrue(observed_payloads[1]["artifact_dir"].endswith("workers/worker_01"))
self.assertTrue(
summary["episodes"][0]["artifact_paths"]["trajectory_image"].endswith(
"workers/worker_00/rollout_front_ep01_trajectory.png"
)
)
self.assertTrue(
summary["episodes"][1]["artifact_paths"]["trajectory_image"].endswith(
"workers/worker_01/rollout_front_ep02_trajectory.png"
)
)
def test_run_eval_parallel_surfaces_worker_failures(self):
cfg = _make_parallel_cfg(num_episodes=2, num_workers=2)
with mock.patch.object(
eval_vla,
"sample_transfer_pose",
side_effect=[
np.array([0.1, 0.2, 0.3], dtype=np.float32),
np.array([0.4, 0.5, 0.6], dtype=np.float32),
],
), mock.patch.object(
eval_vla,
"_run_spawn_jobs",
side_effect=RuntimeError("boom"),
):
with self.assertRaisesRegex(RuntimeError, "Parallel rollout worker failed"):
eval_vla._run_eval_parallel(cfg)
def test_run_eval_parallel_cuda_builds_server_payloads_and_merges_worker_results(self):
cfg = _make_parallel_cfg(
num_episodes=4,
num_workers=4,
device="cuda",
cuda_devices=[0],
artifact_dir="/tmp/cuda-root",
)
observed_server_payloads = []
observed_worker_payloads = []
def fake_run_cuda_parallel_processes(server_payloads, worker_payloads):
observed_server_payloads.extend(server_payloads)
observed_worker_payloads.extend(worker_payloads)
return [
{
"episodes": [
{
"episode_index": 2,
"episode_reward": 3.0,
"episode_max_reward": 3.0,
"inference_fps": 30.0,
"control_fps": 15.0,
},
{
"episode_index": 0,
"episode_reward": 1.0,
"episode_max_reward": 1.0,
"inference_fps": 10.0,
"control_fps": 5.0,
},
],
"_merge_state": {
"obs_read_time_ms": [1.0, 2.0],
"preprocess_time_ms": [1.0, 1.0],
"inference_time_ms": [2.0, 2.0],
"env_step_time_ms": [3.0, 3.0],
"total_time_ms": [4.0, 4.0],
"model_forward_flags": [True, False],
},
},
{
"episodes": [
{
"episode_index": 3,
"episode_reward": 4.0,
"episode_max_reward": 4.0,
"inference_fps": 40.0,
"control_fps": 20.0,
},
{
"episode_index": 1,
"episode_reward": 2.0,
"episode_max_reward": 2.0,
"inference_fps": 20.0,
"control_fps": 10.0,
},
],
"_merge_state": {
"obs_read_time_ms": [3.0, 4.0],
"preprocess_time_ms": [1.0, 1.0],
"inference_time_ms": [2.0, 2.0],
"env_step_time_ms": [3.0, 3.0],
"total_time_ms": [4.0, 4.0],
"model_forward_flags": [True, True],
},
},
]
with mock.patch.object(
eval_vla,
"sample_transfer_pose",
side_effect=[
np.array([0.1, 0.2, 0.3], dtype=np.float32),
np.array([0.4, 0.5, 0.6], dtype=np.float32),
np.array([0.7, 0.8, 0.9], dtype=np.float32),
np.array([1.0, 1.1, 1.2], dtype=np.float32),
],
), mock.patch.object(
eval_vla,
"_run_cuda_parallel_processes",
side_effect=fake_run_cuda_parallel_processes,
create=True,
):
summary = eval_vla._run_eval_parallel_cuda(cfg)
self.assertEqual(len(observed_server_payloads), 1)
self.assertEqual(observed_server_payloads[0]["device_index"], 0)
self.assertEqual(len(observed_worker_payloads), 4)
self.assertTrue(all(payload["server_index"] == 0 for payload in observed_worker_payloads))
self.assertEqual([episode["episode_index"] for episode in summary["episodes"]], [0, 1, 2, 3])
self.assertEqual(summary["episode_rewards"], [1.0, 2.0, 3.0, 4.0])
self.assertEqual(summary["num_episodes"], 4)
def test_run_eval_parallel_cuda_surfaces_server_failures(self):
cfg = _make_parallel_cfg(num_episodes=2, num_workers=2, device="cuda", cuda_devices=[0])
with mock.patch.object(
eval_vla,
"sample_transfer_pose",
side_effect=[
np.array([0.1, 0.2, 0.3], dtype=np.float32),
np.array([0.4, 0.5, 0.6], dtype=np.float32),
],
), mock.patch.object(
eval_vla,
"_run_cuda_parallel_processes",
side_effect=RuntimeError("server boom"),
create=True,
):
with self.assertRaisesRegex(RuntimeError, "Parallel CUDA rollout failed"):
eval_vla._run_eval_parallel_cuda(cfg)
def test_run_spawn_jobs_supports_real_spawn_with_actual_eval_worker_entry(self):
payloads = [
{"_spawn_probe": True, "probe_value": 1, "worker_index": 0},
{"_spawn_probe": True, "probe_value": 2, "worker_index": 1},
]
results = eval_vla._run_spawn_jobs(
payloads=payloads,
max_workers=2,
worker_fn=eval_vla._run_eval_worker_entry,
)
self.assertEqual(sorted(result["probe_value"] for result in results), [1, 2])
self.assertEqual(sorted(result["worker_index"] for result in results), [0, 1])
def test_cuda_server_and_env_worker_entrypoints_support_real_spawn_probe(self):
ctx = eval_vla.multiprocessing.get_context("spawn")
request_queue = ctx.Queue()
response_queue = ctx.Queue()
result_queue = ctx.Queue()
server = ctx.Process(
target=eval_vla._inference_server_main,
args=(
{
"_spawn_probe": True,
"server_index": 0,
"request_queue": request_queue,
"response_queues": [response_queue],
},
),
)
worker = ctx.Process(
target=eval_vla._env_worker_main,
args=(
{
"_spawn_probe": True,
"worker_index": 0,
"server_index": 0,
"request_queue": request_queue,
"response_queue": response_queue,
"result_queue": result_queue,
},
),
)
server.start()
worker.start()
result = result_queue.get(timeout=10.0)
worker.join(timeout=10.0)
request_queue.put({"type": "shutdown_server"})
server.join(timeout=10.0)
self.assertEqual(result["kind"], "worker_result")
self.assertEqual(result["summary"]["probe_worker_index"], 0)
self.assertEqual(result["summary"]["probe_server_index"], 0)
self.assertEqual(result["summary"]["probe_actions"], [[[11.0], [22.0], [33.0]]])
self.assertEqual(worker.exitcode, 0)
self.assertEqual(server.exitcode, 0)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+281 -10
View File
@@ -36,8 +36,8 @@ class _FakeEnv:
self.render_calls = 0 self.render_calls = 0
self.reset_calls = [] self.reset_calls = []
def reset(self, box_pos): def reset(self, task_state):
self.reset_calls.append(np.array(box_pos)) self.reset_calls.append(task_state)
def _get_image_obs(self): def _get_image_obs(self):
self.image_obs_calls += 1 self.image_obs_calls += 1
@@ -74,7 +74,7 @@ class _FakeRenderer:
self._env = env self._env = env
self._frames = [ self._frames = [
np.full((4, 4, 3), fill_value=index, dtype=np.uint8) np.full((4, 4, 3), fill_value=index, dtype=np.uint8)
for index in range(5) for index in range(8)
] ]
self._index = 0 self._index = 0
@@ -144,7 +144,7 @@ class EvalVLAHeadlessTest(unittest.TestCase):
is_render=False, is_render=False,
control_freq=30, control_freq=30,
is_interpolate=True, is_interpolate=True,
cam_view="angle", cam_view="top",
) )
def test_headless_sync_camera_capture_populates_images_without_gui_calls(self): def test_headless_sync_camera_capture_populates_images_without_gui_calls(self):
@@ -153,11 +153,10 @@ class EvalVLAHeadlessTest(unittest.TestCase):
env.mj_data = object() env.mj_data = object()
env.exit_flag = False env.exit_flag = False
env.is_render = False env.is_render = False
env.cam = 'angle' env.cam = 'top'
env.r_vis = None env.r_vis = None
env.l_vis = None env.l_vis = None
env.top = None env.top = None
env.angle = None
env.front = None env.front = None
env._offscreen_renderer = None env._offscreen_renderer = None
@@ -176,7 +175,6 @@ class EvalVLAHeadlessTest(unittest.TestCase):
self.assertIsNotNone(env.r_vis) self.assertIsNotNone(env.r_vis)
self.assertIsNotNone(env.l_vis) self.assertIsNotNone(env.l_vis)
self.assertIsNotNone(env.top) self.assertIsNotNone(env.top)
self.assertIsNotNone(env.angle)
self.assertIsNotNone(env.front) self.assertIsNotNone(env.front)
def test_cam_start_skips_background_thread_when_headless(self): def test_cam_start_skips_background_thread_when_headless(self):
@@ -196,11 +194,10 @@ class EvalVLAHeadlessTest(unittest.TestCase):
env.mj_data = object() env.mj_data = object()
env.exit_flag = False env.exit_flag = False
env.is_render = False env.is_render = False
env.cam = "angle" env.cam = "top"
env.r_vis = None env.r_vis = None
env.l_vis = None env.l_vis = None
env.top = None env.top = None
env.angle = None
env.front = None env.front = None
with mock.patch( with mock.patch(
@@ -217,9 +214,33 @@ class EvalVLAHeadlessTest(unittest.TestCase):
self.assertIsNotNone(env.r_vis) self.assertIsNotNone(env.r_vis)
self.assertIsNotNone(env.l_vis) self.assertIsNotNone(env.l_vis)
self.assertIsNotNone(env.top) self.assertIsNotNone(env.top)
self.assertIsNotNone(env.angle)
self.assertIsNotNone(env.front) self.assertIsNotNone(env.front)
def test_dual_diana_step_refreshes_obs_after_physics_step(self):
env = DualDianaMed.__new__(DualDianaMed)
env.compute_qpos = np.zeros(16)
env.interpolator_left = None
env.interpolator_right = None
env.control_timestep = 0.001
env.model_timestep = 0.001
env.base_time = 0.0
events = []
def fake_get_obs():
events.append("obs")
return {"images": {}, "qpos": np.zeros(16, dtype=np.float32)}
env._get_obs = fake_get_obs
with mock.patch(
"roboimi.envs.double_base.MujocoEnv.step",
autospec=True,
side_effect=lambda _self, _action: events.append("physics"),
):
env.step(np.zeros(16))
self.assertEqual(events, ["physics", "obs"])
def test_eval_main_headless_skips_render_and_still_executes_policy(self): def test_eval_main_headless_skips_render_and_still_executes_policy(self):
fake_env = _FakeEnv() fake_env = _FakeEnv()
fake_agent = _FakeAgent() fake_agent = _FakeAgent()
@@ -328,5 +349,255 @@ class EvalVLAHeadlessTest(unittest.TestCase):
self.assertEqual(summary["num_episodes"], 2) self.assertEqual(summary["num_episodes"], 2)
def test_eval_config_exposes_num_workers_default(self):
eval_cfg = OmegaConf.load(Path("roboimi/vla/conf/eval/eval.yaml"))
self.assertIn("num_workers", eval_cfg)
self.assertEqual(eval_cfg.num_workers, 1)
def test_eval_config_exposes_cuda_devices_default(self):
eval_cfg = OmegaConf.load(Path("roboimi/vla/conf/eval/eval.yaml"))
self.assertIn("cuda_devices", eval_cfg)
self.assertIsNone(eval_cfg.cuda_devices)
def test_eval_config_exposes_parallel_timeout_defaults(self):
eval_cfg = OmegaConf.load(Path("roboimi/vla/conf/eval/eval.yaml"))
self.assertIn("response_timeout_s", eval_cfg)
self.assertIn("server_startup_timeout_s", eval_cfg)
self.assertEqual(eval_cfg.response_timeout_s, 300.0)
self.assertEqual(eval_cfg.server_startup_timeout_s, 300.0)
def test_run_eval_uses_serial_path_when_num_workers_is_one(self):
cfg = OmegaConf.create(
{
"eval": {
"num_workers": 1,
"num_episodes": 3,
}
}
)
with mock.patch.object(
eval_vla,
"_run_eval_serial",
return_value={"mode": "serial"},
) as run_eval_serial, mock.patch.object(
eval_vla,
"_run_eval_parallel",
) as run_eval_parallel:
result = eval_vla._run_eval(cfg)
self.assertEqual(result, {"mode": "serial"})
run_eval_serial.assert_called_once_with(cfg)
run_eval_parallel.assert_not_called()
def test_run_eval_uses_serial_path_when_requested_workers_collapse_to_one(self):
cfg = OmegaConf.create(
{
"eval": {
"num_workers": 8,
"num_episodes": 1,
}
}
)
with mock.patch.object(
eval_vla,
"_run_eval_serial",
return_value={"mode": "serial"},
) as run_eval_serial, mock.patch.object(
eval_vla,
"_run_eval_parallel",
) as run_eval_parallel:
result = eval_vla._run_eval(cfg)
self.assertEqual(result, {"mode": "serial"})
run_eval_serial.assert_called_once_with(cfg)
run_eval_parallel.assert_not_called()
def test_run_eval_parallel_requires_headless_true(self):
cfg = OmegaConf.create(
{
"agent": {},
"eval": {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 2,
"num_workers": 2,
"max_timesteps": 1,
"device": "cpu",
"task_name": "sim_transfer",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": False,
},
}
)
with self.assertRaisesRegex(ValueError, "headless=true"):
eval_vla._run_eval_parallel(cfg)
def test_run_eval_parallel_dispatches_to_cpu_workers_when_device_is_cpu(self):
cfg = OmegaConf.create(
{
"agent": {},
"eval": {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 2,
"num_workers": 2,
"max_timesteps": 1,
"device": "cpu",
"task_name": "sim_transfer",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": True,
"cuda_devices": None,
},
}
)
with mock.patch.object(
eval_vla,
"_run_eval_parallel_cpu",
return_value={"mode": "cpu"},
create=True,
) as run_cpu_parallel, mock.patch.object(
eval_vla,
"_run_eval_parallel_cuda",
create=True,
) as run_cuda_parallel:
result = eval_vla._run_eval_parallel(cfg)
self.assertEqual(result, {"mode": "cpu"})
run_cpu_parallel.assert_called_once_with(cfg)
run_cuda_parallel.assert_not_called()
def test_run_eval_parallel_dispatches_to_cuda_servers_when_device_is_cuda(self):
cfg = OmegaConf.create(
{
"agent": {},
"eval": {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 2,
"num_workers": 2,
"max_timesteps": 1,
"device": "cuda",
"task_name": "sim_transfer",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": True,
"cuda_devices": [0],
},
}
)
with mock.patch.object(
eval_vla,
"_run_eval_parallel_cpu",
create=True,
) as run_cpu_parallel, mock.patch.object(
eval_vla,
"_run_eval_parallel_cuda",
return_value={"mode": "cuda"},
create=True,
) as run_cuda_parallel:
result = eval_vla._run_eval_parallel(cfg)
self.assertEqual(result, {"mode": "cuda"})
run_cpu_parallel.assert_not_called()
run_cuda_parallel.assert_called_once_with(cfg)
def test_resolve_cuda_devices_defaults_to_single_logical_gpu(self):
cfg = OmegaConf.create(
{
"device": "cuda",
"cuda_devices": None,
}
)
self.assertEqual(eval_vla._resolve_cuda_devices(cfg), [0])
def test_resolve_cuda_devices_rejects_empty_selection(self):
cfg = OmegaConf.create(
{
"device": "cuda",
"cuda_devices": [],
}
)
with self.assertRaisesRegex(ValueError, "cuda_devices"):
eval_vla._resolve_cuda_devices(cfg)
def test_run_eval_uses_air_insert_sampler_for_socket_peg_task(self):
self.assertTrue(
hasattr(eval_vla, "sample_air_insert_socket_peg_state"),
"Expected eval_vla to expose the new socket/peg reset sampler",
)
fake_env = _FakeEnv()
fake_agent = _FakeAgent()
sampled_task_state = {
"socket_pos": np.array([-0.10, 0.80, 0.47], dtype=np.float32),
"socket_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32),
"peg_pos": np.array([0.10, 0.82, 0.47], dtype=np.float32),
"peg_quat": np.array([1.0, 0.0, 0.0, 0.0], dtype=np.float32),
}
cfg = OmegaConf.create(
{
"agent": {},
"eval": {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 1,
"max_timesteps": 1,
"device": "cpu",
"task_name": "sim_air_insert_socket_peg",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": True,
},
}
)
with mock.patch.object(
eval_vla,
"load_checkpoint",
return_value=(fake_agent, None),
), mock.patch.object(
eval_vla,
"make_sim_env",
return_value=fake_env,
) as make_env, mock.patch.object(
eval_vla,
"sample_air_insert_socket_peg_state",
return_value=sampled_task_state,
) as socket_peg_sampler, mock.patch.object(
eval_vla,
"sample_transfer_pose",
side_effect=AssertionError("sample_transfer_pose should not be used for sim_air_insert_socket_peg"),
), mock.patch.object(
eval_vla,
"execute_policy_action",
) as execute_policy_action, mock.patch.object(
eval_vla,
"tqdm",
side_effect=lambda iterable, **kwargs: iterable,
):
eval_vla._run_eval(cfg)
make_env.assert_called_once_with("sim_air_insert_socket_peg", headless=True)
socket_peg_sampler.assert_called_once_with()
execute_policy_action.assert_called_once()
self.assertEqual(fake_env.reset_calls, [sampled_task_state])
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
+237 -96
View File
@@ -102,10 +102,8 @@ class EvalVLARolloutArtifactsTest(unittest.TestCase):
self.assertIn('artifact_dir', eval_cfg) self.assertIn('artifact_dir', eval_cfg)
self.assertFalse(eval_cfg.save_summary_json) self.assertFalse(eval_cfg.save_summary_json)
self.assertFalse(eval_cfg.save_trajectory_npz) self.assertFalse(eval_cfg.save_trajectory_npz)
self.assertFalse(eval_cfg.save_trajectory_image)
self.assertFalse(eval_cfg.record_video) self.assertFalse(eval_cfg.record_video)
self.assertIsNone(eval_cfg.artifact_dir) self.assertIsNone(eval_cfg.artifact_dir)
self.assertIsNone(eval_cfg.trajectory_image_camera_name)
self.assertIsNone(eval_cfg.video_camera_name) self.assertIsNone(eval_cfg.video_camera_name)
self.assertEqual(eval_cfg.video_fps, 30) self.assertEqual(eval_cfg.video_fps, 30)
@@ -135,8 +133,6 @@ class EvalVLARolloutArtifactsTest(unittest.TestCase):
'artifact_dir': tmpdir, 'artifact_dir': tmpdir,
'save_summary_json': True, 'save_summary_json': True,
'save_trajectory_npz': True, 'save_trajectory_npz': True,
'save_trajectory_image': True,
'trajectory_image_camera_name': 'front',
'record_video': True, 'record_video': True,
'video_camera_name': 'front', 'video_camera_name': 'front',
'video_fps': 12, 'video_fps': 12,
@@ -180,14 +176,12 @@ class EvalVLARolloutArtifactsTest(unittest.TestCase):
trajectory_path = Path(artifacts['trajectory_npz']) trajectory_path = Path(artifacts['trajectory_npz'])
summary_path = Path(artifacts['summary_json']) summary_path = Path(artifacts['summary_json'])
video_path = Path(artifacts['video_mp4']) video_path = Path(artifacts['video_mp4'])
trajectory_image_path = Path(summary['episodes'][0]['artifact_paths']['trajectory_image'])
self.assertEqual(Path(artifacts['output_dir']), Path(tmpdir)) self.assertEqual(Path(artifacts['output_dir']), Path(tmpdir))
self.assertEqual(artifacts['video_camera_name'], 'front') self.assertEqual(artifacts['video_camera_name'], 'front')
self.assertTrue(trajectory_path.exists()) self.assertTrue(trajectory_path.exists())
self.assertTrue(summary_path.exists()) self.assertTrue(summary_path.exists())
self.assertTrue(video_path.exists()) self.assertTrue(video_path.exists())
self.assertTrue(trajectory_image_path.exists())
rollout_npz = np.load(trajectory_path) rollout_npz = np.load(trajectory_path)
np.testing.assert_array_equal(rollout_npz['episode_index'], np.array([0, 0])) np.testing.assert_array_equal(rollout_npz['episode_index'], np.array([0, 0]))
@@ -224,120 +218,267 @@ class EvalVLARolloutArtifactsTest(unittest.TestCase):
saved_summary = json.load(fh) saved_summary = json.load(fh)
self.assertEqual(saved_summary['artifacts']['trajectory_npz'], str(trajectory_path)) self.assertEqual(saved_summary['artifacts']['trajectory_npz'], str(trajectory_path))
self.assertEqual(saved_summary['artifacts']['video_mp4'], str(video_path)) self.assertEqual(saved_summary['artifacts']['video_mp4'], str(video_path))
self.assertEqual(
saved_summary['episodes'][0]['artifact_paths']['trajectory_image'],
str(trajectory_image_path),
)
self.assertEqual(saved_summary['episode_rewards'], [3.0]) self.assertEqual(saved_summary['episode_rewards'], [3.0])
self.assertAlmostEqual(summary['avg_reward'], 3.0) self.assertAlmostEqual(summary['avg_reward'], 3.0)
self.assertIn('avg_obs_read_time_ms', summary) self.assertIn('avg_obs_read_time_ms', summary)
self.assertIn('avg_env_step_time_ms', summary) self.assertIn('avg_env_step_time_ms', summary)
def test_run_eval_exports_front_trajectory_images_without_video_dependency(self): def test_run_eval_parallel_rejects_trajectory_and_video_exports(self):
actions = [ unsupported_flags = [
np.arange(16, dtype=np.float32), "record_video",
np.arange(16, dtype=np.float32) + 10.0, "save_trajectory",
np.arange(16, dtype=np.float32) + 100.0, "save_trajectory_npz",
np.arange(16, dtype=np.float32) + 110.0,
] ]
fake_agent = _FakeAgent(actions)
fake_env = _FakeEnv()
for flag_name in unsupported_flags:
with self.subTest(flag_name=flag_name):
cfg = OmegaConf.create(
{
"agent": {},
"eval": {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 2,
"num_workers": 2,
"max_timesteps": 1,
"device": "cpu",
"task_name": "sim_transfer",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": True,
"save_artifacts": True,
flag_name: True,
},
}
)
with self.assertRaisesRegex(ValueError, flag_name):
eval_vla._run_eval_parallel(cfg)
def test_run_eval_parallel_writes_merged_summary_timing_and_worker_dirs(self):
with tempfile.TemporaryDirectory() as tmpdir: with tempfile.TemporaryDirectory() as tmpdir:
cfg = OmegaConf.create( cfg = OmegaConf.create(
{ {
'agent': {}, "agent": {},
'eval': { "eval": {
'ckpt_path': 'checkpoints/vla_model_best.pt', "ckpt_path": "checkpoints/vla_model_best.pt",
'num_episodes': 2, "num_episodes": 3,
'max_timesteps': 2, "num_workers": 2,
'device': 'cpu', "max_timesteps": 1,
'task_name': 'sim_transfer', "device": "cpu",
'camera_names': ['top', 'front'], "task_name": "sim_transfer",
'use_smoothing': True, "camera_names": ["front"],
'smooth_alpha': 0.5, "use_smoothing": False,
'verbose_action': False, "smooth_alpha": 0.3,
'headless': True, "verbose_action": False,
'artifact_dir': tmpdir, "headless": True,
'save_trajectory_image': True, "artifact_dir": tmpdir,
'record_video': False, "save_summary_json": True,
"save_timing": True,
}, },
} }
) )
trajectory_image_calls = [] def fake_run_spawn_jobs(payloads, max_workers, worker_fn):
del max_workers, worker_fn
def fake_save_rollout_trajectory_image( return [
env,
output_path,
raw_actions,
camera_name,
*,
line_radius=0.004,
max_markers=1500,
):
del env, line_radius, max_markers
trajectory_image_calls.append(
{ {
'output_path': output_path, "episodes": [
'camera_name': camera_name, {
'raw_actions': [np.array(action, copy=True) for action in raw_actions], "episode_index": 2,
} "episode_reward": 3.0,
) "episode_max_reward": 3.0,
if output_path is None: "inference_fps": 30.0,
return None "control_fps": 15.0,
output_path = Path(output_path) }
output_path.parent.mkdir(parents=True, exist_ok=True) ],
output_path.write_bytes(b'fake-png') "_merge_state": {
return str(output_path) "obs_read_time_ms": [3.0],
"preprocess_time_ms": [1.0],
"inference_time_ms": [2.0],
"env_step_time_ms": [4.0],
"total_time_ms": [5.0],
"model_forward_flags": [True],
},
},
{
"episodes": [
{
"episode_index": 1,
"episode_reward": 2.0,
"episode_max_reward": 2.0,
"inference_fps": 20.0,
"control_fps": 10.0,
},
{
"episode_index": 0,
"episode_reward": 1.0,
"episode_max_reward": 1.0,
"inference_fps": 10.0,
"control_fps": 5.0,
},
],
"_merge_state": {
"obs_read_time_ms": [1.0, 2.0],
"preprocess_time_ms": [1.0, 1.0],
"inference_time_ms": [2.0, 2.0],
"env_step_time_ms": [4.0, 4.0],
"total_time_ms": [5.0, 5.0],
"model_forward_flags": [False, True],
},
},
]
with mock.patch.object( with mock.patch.object(
eval_vla, eval_vla,
'load_checkpoint', "sample_transfer_pose",
return_value=(fake_agent, None), side_effect=[
np.array([0.1, 0.2, 0.3], dtype=np.float32),
np.array([0.4, 0.5, 0.6], dtype=np.float32),
np.array([0.7, 0.8, 0.9], dtype=np.float32),
],
), mock.patch.object( ), mock.patch.object(
eval_vla, eval_vla,
'make_sim_env', "_run_spawn_jobs",
return_value=fake_env, side_effect=fake_run_spawn_jobs,
), mock.patch.object( ):
eval_vla, summary = eval_vla._run_eval_parallel(cfg)
'sample_transfer_pose',
return_value=np.array([0.1, 0.2, 0.3], dtype=np.float32),
), mock.patch.object(
eval_vla,
'tqdm',
side_effect=lambda iterable, **kwargs: iterable,
), mock.patch.object(
eval_vla,
'_save_rollout_trajectory_image',
side_effect=fake_save_rollout_trajectory_image,
) as save_trajectory_image_mock, mock.patch.object(
eval_vla,
'_open_video_writer',
) as open_video_writer_mock:
summary = eval_vla._run_eval(cfg)
self.assertEqual(save_trajectory_image_mock.call_count, 2) summary_path = Path(tmpdir) / "rollout_summary.json"
open_video_writer_mock.assert_not_called() timing_path = Path(tmpdir) / "timing.json"
self.assertIsNone(summary['artifacts']['video_mp4']) worker_00_dir = Path(tmpdir) / "workers" / "worker_00"
self.assertEqual(summary['artifacts']['trajectory_image_camera_name'], 'front') worker_01_dir = Path(tmpdir) / "workers" / "worker_01"
self.assertEqual(
[call['camera_name'] for call in trajectory_image_calls], self.assertTrue(summary_path.exists())
['front', 'front'], self.assertTrue(timing_path.exists())
self.assertTrue(worker_00_dir.is_dir())
self.assertTrue(worker_01_dir.is_dir())
self.assertEqual(summary["episode_rewards"], [1.0, 2.0, 3.0])
with summary_path.open("r", encoding="utf-8") as fh:
saved_summary = json.load(fh)
with timing_path.open("r", encoding="utf-8") as fh:
saved_timing = json.load(fh)
self.assertEqual(saved_summary["episode_rewards"], [1.0, 2.0, 3.0])
self.assertEqual(saved_summary["artifact_dir"], tmpdir)
self.assertEqual(saved_timing["count"], 3)
self.assertEqual(saved_timing["model_forward_count"], 2)
def test_run_eval_parallel_cuda_writes_merged_summary_timing_and_worker_dirs(self):
with tempfile.TemporaryDirectory() as tmpdir:
cfg = OmegaConf.create(
{
"agent": {},
"eval": {
"ckpt_path": "checkpoints/vla_model_best.pt",
"num_episodes": 3,
"num_workers": 2,
"cuda_devices": [0],
"max_timesteps": 1,
"device": "cuda",
"task_name": "sim_transfer",
"camera_names": ["front"],
"use_smoothing": False,
"smooth_alpha": 0.3,
"verbose_action": False,
"headless": True,
"artifact_dir": tmpdir,
"save_summary_json": True,
"save_timing": True,
},
}
) )
first_episode_path = Path(summary['episodes'][0]['artifact_paths']['trajectory_image']) def fake_run_cuda_parallel_processes(server_payloads, worker_payloads):
second_episode_path = Path(summary['episodes'][1]['artifact_paths']['trajectory_image']) self.assertEqual(len(server_payloads), 1)
self.assertTrue(first_episode_path.exists()) self.assertEqual(server_payloads[0]["device_index"], 0)
self.assertTrue(second_episode_path.exists()) self.assertEqual([payload["server_index"] for payload in worker_payloads], [0, 0])
self.assertNotEqual(first_episode_path, second_episode_path) return [
self.assertEqual(first_episode_path.parent, Path(tmpdir)) {
self.assertEqual(second_episode_path.parent, Path(tmpdir)) "episodes": [
{
"episode_index": 2,
"episode_reward": 3.0,
"episode_max_reward": 3.0,
"inference_fps": 30.0,
"control_fps": 15.0,
}
],
"_merge_state": {
"obs_read_time_ms": [3.0],
"preprocess_time_ms": [1.0],
"inference_time_ms": [2.0],
"env_step_time_ms": [4.0],
"total_time_ms": [5.0],
"model_forward_flags": [True],
},
},
{
"episodes": [
{
"episode_index": 1,
"episode_reward": 2.0,
"episode_max_reward": 2.0,
"inference_fps": 20.0,
"control_fps": 10.0,
},
{
"episode_index": 0,
"episode_reward": 1.0,
"episode_max_reward": 1.0,
"inference_fps": 10.0,
"control_fps": 5.0,
},
],
"_merge_state": {
"obs_read_time_ms": [1.0, 2.0],
"preprocess_time_ms": [1.0, 1.0],
"inference_time_ms": [2.0, 2.0],
"env_step_time_ms": [4.0, 4.0],
"total_time_ms": [5.0, 5.0],
"model_forward_flags": [False, True],
},
},
]
np.testing.assert_array_equal(trajectory_image_calls[0]['raw_actions'][0], actions[0]) with mock.patch.object(
np.testing.assert_array_equal(trajectory_image_calls[0]['raw_actions'][1], actions[1]) eval_vla,
np.testing.assert_array_equal(trajectory_image_calls[1]['raw_actions'][0], actions[2]) "sample_transfer_pose",
np.testing.assert_array_equal(trajectory_image_calls[1]['raw_actions'][1], actions[3]) side_effect=[
np.array([0.1, 0.2, 0.3], dtype=np.float32),
np.array([0.4, 0.5, 0.6], dtype=np.float32),
np.array([0.7, 0.8, 0.9], dtype=np.float32),
],
), mock.patch.object(
eval_vla,
"_run_cuda_parallel_processes",
side_effect=fake_run_cuda_parallel_processes,
create=True,
):
summary = eval_vla._run_eval_parallel_cuda(cfg)
summary_path = Path(tmpdir) / "rollout_summary.json"
timing_path = Path(tmpdir) / "timing.json"
worker_00_dir = Path(tmpdir) / "workers" / "worker_00"
worker_01_dir = Path(tmpdir) / "workers" / "worker_01"
self.assertTrue(summary_path.exists())
self.assertTrue(timing_path.exists())
self.assertTrue(worker_00_dir.is_dir())
self.assertTrue(worker_01_dir.is_dir())
self.assertEqual(summary["episode_rewards"], [1.0, 2.0, 3.0])
with summary_path.open("r", encoding="utf-8") as fh:
saved_summary = json.load(fh)
with timing_path.open("r", encoding="utf-8") as fh:
saved_timing = json.load(fh)
self.assertEqual(saved_summary["episode_rewards"], [1.0, 2.0, 3.0])
self.assertEqual(saved_summary["artifact_dir"], tmpdir)
self.assertEqual(saved_timing["count"], 3)
self.assertEqual(saved_timing["model_forward_count"], 2)
if __name__ == '__main__': if __name__ == '__main__':
+43 -1
View File
@@ -4,7 +4,7 @@ import unittest
from pathlib import Path from pathlib import Path
from unittest import mock from unittest import mock
from roboimi.assets.robots.diana_med import BiDianaMed from roboimi.assets.robots import diana_med
class _FakeKDL: class _FakeKDL:
@@ -24,6 +24,7 @@ class RobotAssetPathResolutionTest(unittest.TestCase):
_FakeKDL.reset_calls = [] _FakeKDL.reset_calls = []
def test_bidianamed_resolves_robot_asset_paths_independent_of_cwd(self): def test_bidianamed_resolves_robot_asset_paths_independent_of_cwd(self):
BiDianaMed = diana_med.BiDianaMed
repo_root = Path(__file__).resolve().parents[1] repo_root = Path(__file__).resolve().parents[1]
expected_xml = repo_root / 'roboimi/assets/models/manipulators/DianaMed/bi_diana_transfer_ee.xml' expected_xml = repo_root / 'roboimi/assets/models/manipulators/DianaMed/bi_diana_transfer_ee.xml'
expected_urdf = repo_root / 'roboimi/assets/models/manipulators/DianaMed/DualDianaMed.urdf' expected_urdf = repo_root / 'roboimi/assets/models/manipulators/DianaMed/DualDianaMed.urdf'
@@ -58,6 +59,47 @@ class RobotAssetPathResolutionTest(unittest.TestCase):
self.assertEqual({Path(path) for path in _FakeKDL.init_calls}, {expected_urdf}) self.assertEqual({Path(path) for path in _FakeKDL.init_calls}, {expected_urdf})
self.assertTrue(all(Path(path).is_absolute() for path in _FakeKDL.init_calls)) self.assertTrue(all(Path(path).is_absolute() for path in _FakeKDL.init_calls))
def test_bidianamed_socket_peg_resolves_robot_asset_paths_independent_of_cwd(self):
BiDianaMedSocketPeg = getattr(diana_med, 'BiDianaMedSocketPeg', None)
self.assertIsNotNone(
BiDianaMedSocketPeg,
'Expected roboimi.assets.robots.diana_med.BiDianaMedSocketPeg',
)
repo_root = Path(__file__).resolve().parents[1]
expected_xml = repo_root / 'roboimi/assets/models/manipulators/DianaMed/bi_diana_socket_peg_ee.xml'
expected_urdf = repo_root / 'roboimi/assets/models/manipulators/DianaMed/DualDianaMed.urdf'
xml_calls = []
def fake_from_xml_path(*, filename, assets=None):
xml_calls.append((filename, assets))
return object()
with tempfile.TemporaryDirectory() as tempdir:
previous_cwd = os.getcwd()
try:
os.chdir(tempdir)
with mock.patch(
'roboimi.assets.robots.arm_base.mujoco.MjModel.from_xml_path',
side_effect=fake_from_xml_path,
), mock.patch(
'roboimi.assets.robots.arm_base.mujoco.MjData',
return_value=object(),
), mock.patch(
'roboimi.assets.robots.arm_base.KDL_utils',
_FakeKDL,
):
BiDianaMedSocketPeg()
finally:
os.chdir(previous_cwd)
self.assertEqual(len(xml_calls), 1)
self.assertEqual(Path(xml_calls[0][0]), expected_xml)
self.assertTrue(Path(xml_calls[0][0]).is_absolute())
self.assertGreaterEqual(len(_FakeKDL.init_calls), 2)
self.assertEqual({Path(path) for path in _FakeKDL.init_calls}, {expected_urdf})
self.assertTrue(all(Path(path).is_absolute() for path in _FakeKDL.init_calls))
if __name__ == '__main__': if __name__ == '__main__':
unittest.main() unittest.main()
@@ -158,6 +158,101 @@ class TrainVLARolloutValidationTest(unittest.TestCase):
self.assertGreater(float(cfg.train.lr), 5e-5) self.assertGreater(float(cfg.train.lr), 5e-5)
self.assertGreater(cfg.train.num_workers, 8) self.assertGreater(cfg.train.num_workers, 8)
self.assertEqual(cfg.train.rollout_val_freq_epochs, 50) self.assertEqual(cfg.train.rollout_val_freq_epochs, 50)
self.assertEqual(cfg.train.rollout_device, cfg.train.device)
self.assertIsNone(cfg.train.rollout_num_workers)
self.assertIsNone(cfg.train.rollout_cuda_devices)
def test_run_training_rollout_validation_propagates_gpu_parallel_settings(self):
cfg = OmegaConf.create(
{
'train': {
'device': 'cpu',
'batch_size': 1,
'num_workers': 0,
'val_split': 0.0,
'seed': 0,
'lr': 1e-3,
'max_steps': 2,
'log_freq': 1,
'save_freq': 1000,
'warmup_steps': 1,
'scheduler_type': 'constant',
'min_lr': 0.0,
'grad_clip': 1.0,
'weight_decay': 0.0,
'pretrained_ckpt': None,
'resume_ckpt': None,
'use_swanlab': False,
'rollout_val_freq_epochs': 2,
'rollout_num_episodes': 5,
'rollout_device': 'cuda',
'rollout_num_workers': 4,
'rollout_cuda_devices': [0, 1],
'rollout_response_timeout_s': 123.0,
'rollout_server_startup_timeout_s': 456.0,
},
'data': {
'camera_names': ['front'],
},
'agent': {
'_target_': 'fake.agent',
},
'eval': {
'ckpt_path': 'unused.pt',
'num_episodes': 99,
'max_timesteps': 1,
'device': 'cpu',
'task_name': 'sim_transfer',
'camera_names': ['front'],
'use_smoothing': False,
'smooth_alpha': 0.3,
'verbose_action': False,
'headless': False,
},
}
)
rollout_mock = mock.Mock(return_value={'avg_reward': 1.0})
def fake_instantiate(config_node, **_kwargs):
if config_node is cfg.data:
return _FakeDataset()
if config_node is cfg.agent:
return _FakeAgent()
raise AssertionError(f'unexpected instantiate config: {config_node!r}')
def fake_dataloader(_dataset, *, shuffle, **_kwargs):
del shuffle, _kwargs
return _FakeLoader(
{
'observation.front': torch.zeros(1, 3, 2, 2),
'observation.state': torch.zeros(1, 4),
'action': torch.zeros(1, 2),
'action_is_pad': torch.zeros(1, 1, dtype=torch.bool),
},
length=1,
)
with tempfile.TemporaryDirectory() as tempdir:
previous_cwd = os.getcwd()
try:
os.chdir(tempdir)
with mock.patch.object(train_vla, 'instantiate', side_effect=fake_instantiate), mock.patch.object(train_vla, 'DataLoader', side_effect=fake_dataloader), mock.patch.object(train_vla, 'build_training_optimizer', return_value=_FakeOptimizer(cfg.train.lr)), mock.patch.object(train_vla, 'get_lr_schedule_with_warmup', return_value=_FakeScheduler()), mock.patch.object(train_vla, 'tqdm', side_effect=lambda iterable, **kwargs: _FakeProgressBar(iterable)), mock.patch.object(train_vla.torch, 'save', return_value=None), mock.patch.object(eval_vla, '_run_eval', rollout_mock, create=True):
train_vla._run_training(cfg)
finally:
os.chdir(previous_cwd)
rollout_cfg = rollout_mock.call_args.args[0]
self.assertEqual(rollout_cfg.eval.device, 'cuda')
self.assertEqual(rollout_cfg.eval.num_workers, 4)
self.assertEqual(list(rollout_cfg.eval.cuda_devices), [0, 1])
self.assertEqual(float(rollout_cfg.eval.response_timeout_s), 123.0)
self.assertEqual(float(rollout_cfg.eval.server_startup_timeout_s), 456.0)
self.assertTrue(rollout_cfg.eval.headless)
self.assertEqual(rollout_cfg.eval.num_episodes, 5)
self.assertFalse(rollout_cfg.eval.record_video)
self.assertTrue(rollout_cfg.eval.save_summary_json)
self.assertTrue(rollout_cfg.eval.save_trajectory_image)
def test_training_passes_backbone_image_resize_override_to_dataset_instantiation(self): def test_training_passes_backbone_image_resize_override_to_dataset_instantiation(self):
cfg = OmegaConf.create( cfg = OmegaConf.create(