mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 20:49:42 +00:00
Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 34bd3198a1 |
@@ -72,19 +72,19 @@ jobs:
|
|||||||
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
@@ -95,7 +95,7 @@ jobs:
|
|||||||
# from source-copy, so code-only changes skip the slow uv-sync layer
|
# from source-copy, so code-only changes skip the slow uv-sync layer
|
||||||
# when the runner has a warm Docker daemon cache.
|
# when the runner has a warm Docker daemon cache.
|
||||||
- name: Build Libero benchmark image
|
- name: Build Libero benchmark image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: docker/Dockerfile.benchmark.libero
|
file: docker/Dockerfile.benchmark.libero
|
||||||
@@ -151,7 +151,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload Libero rollout video
|
- name: Upload Libero rollout video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: libero-rollout-video
|
name: libero-rollout-video
|
||||||
path: /tmp/libero-artifacts/videos/
|
path: /tmp/libero-artifacts/videos/
|
||||||
@@ -159,7 +159,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload Libero eval metrics
|
- name: Upload Libero eval metrics
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: libero-metrics
|
name: libero-metrics
|
||||||
path: /tmp/libero-artifacts/metrics.json
|
path: /tmp/libero-artifacts/metrics.json
|
||||||
@@ -214,7 +214,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload Libero train-smoke eval video
|
- name: Upload Libero train-smoke eval video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: libero-train-smoke-video
|
name: libero-train-smoke-video
|
||||||
path: /tmp/libero-train-smoke-artifacts/eval/
|
path: /tmp/libero-train-smoke-artifacts/eval/
|
||||||
@@ -230,19 +230,19 @@ jobs:
|
|||||||
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
@@ -250,7 +250,7 @@ jobs:
|
|||||||
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
|
|
||||||
- name: Build MetaWorld benchmark image
|
- name: Build MetaWorld benchmark image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: docker/Dockerfile.benchmark.metaworld
|
file: docker/Dockerfile.benchmark.metaworld
|
||||||
@@ -303,7 +303,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload MetaWorld rollout video
|
- name: Upload MetaWorld rollout video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: metaworld-rollout-video
|
name: metaworld-rollout-video
|
||||||
path: /tmp/metaworld-artifacts/videos/
|
path: /tmp/metaworld-artifacts/videos/
|
||||||
@@ -311,7 +311,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload MetaWorld eval metrics
|
- name: Upload MetaWorld eval metrics
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: metaworld-metrics
|
name: metaworld-metrics
|
||||||
path: /tmp/metaworld-artifacts/metrics.json
|
path: /tmp/metaworld-artifacts/metrics.json
|
||||||
@@ -332,19 +332,19 @@ jobs:
|
|||||||
ROBOTWIN_TASKS: beat_block_hammer,click_bell,handover_block,stack_blocks_two,click_alarmclock,open_microwave,adjust_bottle,lift_pot,stamp_seal,turn_switch
|
ROBOTWIN_TASKS: beat_block_hammer,click_bell,handover_block,stack_blocks_two,click_alarmclock,open_microwave,adjust_bottle,lift_pot,stamp_seal,turn_switch
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
@@ -355,7 +355,7 @@ jobs:
|
|||||||
# simulation assets (~4 GB). Layer cache lives in the runner's local
|
# simulation assets (~4 GB). Layer cache lives in the runner's local
|
||||||
# Docker daemon — reused across re-runs on the same machine.
|
# Docker daemon — reused across re-runs on the same machine.
|
||||||
- name: Build RoboTwin 2.0 benchmark image
|
- name: Build RoboTwin 2.0 benchmark image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: docker/Dockerfile.benchmark.robotwin
|
file: docker/Dockerfile.benchmark.robotwin
|
||||||
@@ -413,7 +413,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload RoboTwin rollout video
|
- name: Upload RoboTwin rollout video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4
|
uses: actions/upload-artifact@v7
|
||||||
with:
|
with:
|
||||||
name: robotwin-rollout-video
|
name: robotwin-rollout-video
|
||||||
path: /tmp/robotwin-artifacts/videos/
|
path: /tmp/robotwin-artifacts/videos/
|
||||||
@@ -421,7 +421,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload RoboTwin eval metrics
|
- name: Upload RoboTwin eval metrics
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4
|
uses: actions/upload-artifact@v7
|
||||||
with:
|
with:
|
||||||
name: robotwin-metrics
|
name: robotwin-metrics
|
||||||
path: /tmp/robotwin-artifacts/metrics.json
|
path: /tmp/robotwin-artifacts/metrics.json
|
||||||
@@ -439,19 +439,19 @@ jobs:
|
|||||||
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
@@ -459,7 +459,7 @@ jobs:
|
|||||||
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
|
|
||||||
- name: Build RoboCasa365 benchmark image
|
- name: Build RoboCasa365 benchmark image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: docker/Dockerfile.benchmark.robocasa
|
file: docker/Dockerfile.benchmark.robocasa
|
||||||
@@ -514,7 +514,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload RoboCasa365 rollout video
|
- name: Upload RoboCasa365 rollout video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: robocasa-rollout-video
|
name: robocasa-rollout-video
|
||||||
path: /tmp/robocasa-artifacts/videos/
|
path: /tmp/robocasa-artifacts/videos/
|
||||||
@@ -522,7 +522,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload RoboCasa365 eval metrics
|
- name: Upload RoboCasa365 eval metrics
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: robocasa-metrics
|
name: robocasa-metrics
|
||||||
path: /tmp/robocasa-artifacts/metrics.json
|
path: /tmp/robocasa-artifacts/metrics.json
|
||||||
@@ -540,19 +540,19 @@ jobs:
|
|||||||
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
@@ -560,7 +560,7 @@ jobs:
|
|||||||
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
|
|
||||||
- name: Build RoboCerebra benchmark image
|
- name: Build RoboCerebra benchmark image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: docker/Dockerfile.benchmark.robocerebra
|
file: docker/Dockerfile.benchmark.robocerebra
|
||||||
@@ -621,7 +621,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload RoboCerebra rollout video
|
- name: Upload RoboCerebra rollout video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: robocerebra-rollout-video
|
name: robocerebra-rollout-video
|
||||||
path: /tmp/robocerebra-artifacts/videos/
|
path: /tmp/robocerebra-artifacts/videos/
|
||||||
@@ -629,7 +629,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload RoboCerebra eval metrics
|
- name: Upload RoboCerebra eval metrics
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: robocerebra-metrics
|
name: robocerebra-metrics
|
||||||
path: /tmp/robocerebra-artifacts/metrics.json
|
path: /tmp/robocerebra-artifacts/metrics.json
|
||||||
@@ -648,19 +648,19 @@ jobs:
|
|||||||
ROBOMME_TASKS: PickXtimes,BinFill,StopCube,MoveCube,InsertPeg,SwingXtimes,VideoUnmask,ButtonUnmask,PickHighlight,PatternLock
|
ROBOMME_TASKS: PickXtimes,BinFill,StopCube,MoveCube,InsertPeg,SwingXtimes,VideoUnmask,ButtonUnmask,PickHighlight,PatternLock
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
@@ -668,7 +668,7 @@ jobs:
|
|||||||
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
|
|
||||||
- name: Build RoboMME benchmark image
|
- name: Build RoboMME benchmark image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: docker/Dockerfile.benchmark.robomme
|
file: docker/Dockerfile.benchmark.robomme
|
||||||
@@ -726,7 +726,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload RoboMME rollout video
|
- name: Upload RoboMME rollout video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: robomme-rollout-video
|
name: robomme-rollout-video
|
||||||
path: /tmp/robomme-artifacts/videos/
|
path: /tmp/robomme-artifacts/videos/
|
||||||
@@ -734,7 +734,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload RoboMME eval metrics
|
- name: Upload RoboMME eval metrics
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: robomme-metrics
|
name: robomme-metrics
|
||||||
path: /tmp/robomme-artifacts/metrics.json
|
path: /tmp/robomme-artifacts/metrics.json
|
||||||
@@ -754,19 +754,19 @@ jobs:
|
|||||||
LIBERO_PLUS_TASK_IDS: "[0,100,260,500,1000,1500,2000,2400]"
|
LIBERO_PLUS_TASK_IDS: "[0,100,260,500,1000,1500,2000,2400]"
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
@@ -774,7 +774,7 @@ jobs:
|
|||||||
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
|
|
||||||
- name: Build LIBERO-plus benchmark image
|
- name: Build LIBERO-plus benchmark image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: docker/Dockerfile.benchmark.libero_plus
|
file: docker/Dockerfile.benchmark.libero_plus
|
||||||
@@ -834,7 +834,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload LIBERO-plus rollout video
|
- name: Upload LIBERO-plus rollout video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: libero-plus-rollout-video
|
name: libero-plus-rollout-video
|
||||||
path: /tmp/libero-plus-artifacts/videos/
|
path: /tmp/libero-plus-artifacts/videos/
|
||||||
@@ -842,7 +842,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload LIBERO-plus eval metrics
|
- name: Upload LIBERO-plus eval metrics
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: libero-plus-metrics
|
name: libero-plus-metrics
|
||||||
path: /tmp/libero-plus-artifacts/metrics.json
|
path: /tmp/libero-plus-artifacts/metrics.json
|
||||||
@@ -858,19 +858,19 @@ jobs:
|
|||||||
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
|
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
if: ${{ env.DOCKERHUB_USERNAME != '' }}
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
@@ -878,7 +878,7 @@ jobs:
|
|||||||
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
DOCKERHUB_USERNAME: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
|
|
||||||
- name: Build VLABench benchmark image
|
- name: Build VLABench benchmark image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: docker/Dockerfile.benchmark.vlabench
|
file: docker/Dockerfile.benchmark.vlabench
|
||||||
@@ -936,7 +936,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload VLABench rollout video
|
- name: Upload VLABench rollout video
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: vlabench-rollout-video
|
name: vlabench-rollout-video
|
||||||
path: /tmp/vlabench-artifacts/videos/
|
path: /tmp/vlabench-artifacts/videos/
|
||||||
@@ -944,7 +944,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload VLABench eval metrics
|
- name: Upload VLABench eval metrics
|
||||||
if: always()
|
if: always()
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: vlabench-metrics
|
name: vlabench-metrics
|
||||||
path: /tmp/vlabench-artifacts/metrics.json
|
path: /tmp/vlabench-artifacts/metrics.json
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ jobs:
|
|||||||
timeout-minutes: 30
|
timeout-minutes: 30
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
|
|||||||
@@ -52,21 +52,21 @@ jobs:
|
|||||||
sudo apt-get update
|
sudo apt-get update
|
||||||
sudo apt-get install git-lfs
|
sudo apt-get install git-lfs
|
||||||
git lfs install
|
git lfs install
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
with:
|
with:
|
||||||
lfs: true
|
lfs: true
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
- name: Build and push Docker image CPU
|
- name: Build and push Docker image CPU
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./docker/Dockerfile.user
|
file: ./docker/Dockerfile.user
|
||||||
@@ -87,21 +87,21 @@ jobs:
|
|||||||
sudo apt-get update
|
sudo apt-get update
|
||||||
sudo apt-get install git-lfs
|
sudo apt-get install git-lfs
|
||||||
git lfs install
|
git lfs install
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
with:
|
with:
|
||||||
lfs: true
|
lfs: true
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
- name: Build and push Docker image GPU
|
- name: Build and push Docker image GPU
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./docker/Dockerfile.internal
|
file: ./docker/Dockerfile.internal
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ jobs:
|
|||||||
github.event.workflow_run.event == 'pull_request' &&
|
github.event.workflow_run.event == 'pull_request' &&
|
||||||
github.event.workflow_run.conclusion == 'success' &&
|
github.event.workflow_run.conclusion == 'success' &&
|
||||||
github.repository == 'huggingface/lerobot'
|
github.repository == 'huggingface/lerobot'
|
||||||
uses: huggingface/doc-builder/.github/workflows/upload_pr_documentation.yml@2430c1ec91d04667414e2fa31ecfc36c153ea391 # main
|
uses: huggingface/doc-builder/.github/workflows/upload_pr_documentation.yml@6108e850ae1cf2f71bb0815a600bcd50c39abfa7 # main
|
||||||
with:
|
with:
|
||||||
package_name: lerobot
|
package_name: lerobot
|
||||||
secrets:
|
secrets:
|
||||||
|
|||||||
@@ -55,7 +55,7 @@ jobs:
|
|||||||
github.repository == 'huggingface/lerobot'
|
github.repository == 'huggingface/lerobot'
|
||||||
permissions:
|
permissions:
|
||||||
contents: read
|
contents: read
|
||||||
uses: huggingface/doc-builder/.github/workflows/build_main_documentation.yml@e60a538eea9817ab312196d0d233604b01697265 # main
|
uses: huggingface/doc-builder/.github/workflows/build_main_documentation.yml@6108e850ae1cf2f71bb0815a600bcd50c39abfa7 # main
|
||||||
with:
|
with:
|
||||||
commit_sha: ${{ github.sha }}
|
commit_sha: ${{ github.sha }}
|
||||||
package: lerobot
|
package: lerobot
|
||||||
@@ -78,7 +78,7 @@ jobs:
|
|||||||
permissions:
|
permissions:
|
||||||
contents: read
|
contents: read
|
||||||
pull-requests: write
|
pull-requests: write
|
||||||
uses: huggingface/doc-builder/.github/workflows/build_pr_documentation.yml@e60a538eea9817ab312196d0d233604b01697265 # main
|
uses: huggingface/doc-builder/.github/workflows/build_pr_documentation.yml@6108e850ae1cf2f71bb0815a600bcd50c39abfa7 # main
|
||||||
with:
|
with:
|
||||||
commit_sha: ${{ github.event.pull_request.head.sha }}
|
commit_sha: ${{ github.event.pull_request.head.sha }}
|
||||||
pr_number: ${{ github.event.number }}
|
pr_number: ${{ github.event.number }}
|
||||||
|
|||||||
@@ -69,7 +69,7 @@ jobs:
|
|||||||
HF_LEROBOT_HOME: /mnt/cache/.cache/huggingface/lerobot
|
HF_LEROBOT_HOME: /mnt/cache/.cache/huggingface/lerobot
|
||||||
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
lfs: true
|
lfs: true
|
||||||
@@ -87,7 +87,7 @@ jobs:
|
|||||||
libusb-1.0-0-dev speech-dispatcher libgeos-dev portaudio19-dev
|
libusb-1.0-0-dev speech-dispatcher libgeos-dev portaudio19-dev
|
||||||
|
|
||||||
- name: Setup uv and Python
|
- name: Setup uv and Python
|
||||||
uses: astral-sh/setup-uv@d0cc045d04ccac9d8b7881df0226f9e82c39688e # v6
|
uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2
|
||||||
with:
|
with:
|
||||||
enable-cache: true
|
enable-cache: true
|
||||||
version: ${{ env.UV_VERSION }}
|
version: ${{ env.UV_VERSION }}
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ jobs:
|
|||||||
HF_LEROBOT_HOME: /mnt/cache/.cache/huggingface/lerobot
|
HF_LEROBOT_HOME: /mnt/cache/.cache/huggingface/lerobot
|
||||||
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
lfs: true
|
lfs: true
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
@@ -80,7 +80,7 @@ jobs:
|
|||||||
speech-dispatcher libgeos-dev portaudio19-dev
|
speech-dispatcher libgeos-dev portaudio19-dev
|
||||||
|
|
||||||
- name: Setup uv and Python
|
- name: Setup uv and Python
|
||||||
uses: astral-sh/setup-uv@d0cc045d04ccac9d8b7881df0226f9e82c39688e # v6
|
uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2
|
||||||
with:
|
with:
|
||||||
enable-cache: true
|
enable-cache: true
|
||||||
version: ${{ env.UV_VERSION }}
|
version: ${{ env.UV_VERSION }}
|
||||||
@@ -137,21 +137,21 @@ jobs:
|
|||||||
sudo apt-get update
|
sudo apt-get update
|
||||||
sudo apt-get install git-lfs
|
sudo apt-get install git-lfs
|
||||||
git lfs install
|
git lfs install
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
lfs: true
|
lfs: true
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3
|
uses: docker/setup-buildx-action@bb05f3f5519dd87d3ba754cc423b652a5edd6d2c # v4.2.0
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
uses: docker/login-action@c94ce9fb468520275223c153574b00df6fe4bcc9 # v3
|
uses: docker/login-action@af1e73f918a031802d376d3c8bbc3fe56130a9b0 # v4.4.0
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
- name: Build and push Docker image
|
- name: Build and push Docker image
|
||||||
uses: docker/build-push-action@10e90e3645eae34f1e60eeb005ba3a3d33f178e8 # v6
|
uses: docker/build-push-action@53b7df96c91f9c12dcc8a07bcb9ccacbed38856a # v7.3.0
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./docker/Dockerfile.internal
|
file: ./docker/Dockerfile.internal
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: github.repository == 'huggingface/lerobot'
|
if: github.repository == 'huggingface/lerobot'
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/github-script@v8
|
- uses: actions/github-script@v9
|
||||||
with:
|
with:
|
||||||
script: |
|
script: |
|
||||||
// Setup Input Text
|
// Setup Input Text
|
||||||
|
|||||||
@@ -48,12 +48,12 @@ jobs:
|
|||||||
outputs:
|
outputs:
|
||||||
changed: ${{ steps.diff.outputs.changed }}
|
changed: ${{ steps.diff.outputs.changed }}
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Setup uv and Python
|
- name: Setup uv and Python
|
||||||
uses: astral-sh/setup-uv@v6 # zizmor: ignore[unpinned-uses]
|
uses: astral-sh/setup-uv@v8.3.2 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
version: ${{ env.UV_VERSION }}
|
version: ${{ env.UV_VERSION }}
|
||||||
python-version: ${{ env.PYTHON_VERSION }}
|
python-version: ${{ env.PYTHON_VERSION }}
|
||||||
@@ -74,7 +74,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Upload updated lockfile
|
- name: Upload updated lockfile
|
||||||
if: steps.diff.outputs.changed == 'true'
|
if: steps.diff.outputs.changed == 'true'
|
||||||
uses: actions/upload-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/upload-artifact@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: uv-lock
|
name: uv-lock
|
||||||
path: uv.lock
|
path: uv.lock
|
||||||
@@ -93,13 +93,13 @@ jobs:
|
|||||||
HF_LEROBOT_HOME: /mnt/cache/.cache/huggingface/lerobot
|
HF_LEROBOT_HOME: /mnt/cache/.cache/huggingface/lerobot
|
||||||
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
HF_USER_TOKEN: ${{ secrets.LEROBOT_HF_USER }}
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
with:
|
with:
|
||||||
lfs: true
|
lfs: true
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Download updated lockfile
|
- name: Download updated lockfile
|
||||||
uses: actions/download-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/download-artifact@v8 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: uv-lock
|
name: uv-lock
|
||||||
|
|
||||||
@@ -115,7 +115,7 @@ jobs:
|
|||||||
speech-dispatcher libgeos-dev portaudio19-dev
|
speech-dispatcher libgeos-dev portaudio19-dev
|
||||||
|
|
||||||
- name: Setup uv and Python
|
- name: Setup uv and Python
|
||||||
uses: astral-sh/setup-uv@v6 # zizmor: ignore[unpinned-uses]
|
uses: astral-sh/setup-uv@v8.3.2 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
enable-cache: true
|
enable-cache: true
|
||||||
version: ${{ env.UV_VERSION }}
|
version: ${{ env.UV_VERSION }}
|
||||||
@@ -153,27 +153,27 @@ jobs:
|
|||||||
sudo apt-get update
|
sudo apt-get update
|
||||||
sudo apt-get install git-lfs
|
sudo apt-get install git-lfs
|
||||||
git lfs install
|
git lfs install
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
with:
|
with:
|
||||||
lfs: true
|
lfs: true
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Download updated lockfile
|
- name: Download updated lockfile
|
||||||
uses: actions/download-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/download-artifact@v8 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: uv-lock
|
name: uv-lock
|
||||||
|
|
||||||
- name: Set up Docker Buildx
|
- name: Set up Docker Buildx
|
||||||
uses: docker/setup-buildx-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/setup-buildx-action@v4 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
cache-binary: false
|
cache-binary: false
|
||||||
- name: Login to Docker Hub
|
- name: Login to Docker Hub
|
||||||
uses: docker/login-action@v3 # zizmor: ignore[unpinned-uses]
|
uses: docker/login-action@v4.4.0 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
username: ${{ secrets.DOCKERHUB_LEROBOT_USERNAME }}
|
||||||
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
password: ${{ secrets.DOCKERHUB_LEROBOT_PASSWORD }}
|
||||||
- name: Build and push Docker image
|
- name: Build and push Docker image
|
||||||
uses: docker/build-push-action@v6 # zizmor: ignore[unpinned-uses]
|
uses: docker/build-push-action@v7 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
context: .
|
context: .
|
||||||
file: ./docker/Dockerfile.internal
|
file: ./docker/Dockerfile.internal
|
||||||
@@ -247,12 +247,12 @@ jobs:
|
|||||||
env:
|
env:
|
||||||
GH_TOKEN: ${{ secrets.UPDATE_LOCK_TOKEN }}
|
GH_TOKEN: ${{ secrets.UPDATE_LOCK_TOKEN }}
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@v6
|
- uses: actions/checkout@v7
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Download updated lockfile
|
- name: Download updated lockfile
|
||||||
uses: actions/download-artifact@v4 # zizmor: ignore[unpinned-uses]
|
uses: actions/download-artifact@v8 # zizmor: ignore[unpinned-uses]
|
||||||
with:
|
with:
|
||||||
name: uv-lock
|
name: uv-lock
|
||||||
|
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
if: github.repository == 'huggingface/lerobot' && !github.event.pull_request.draft
|
if: github.repository == 'huggingface/lerobot' && !github.event.pull_request.draft
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/labeler@v6
|
- uses: actions/labeler@v7
|
||||||
with:
|
with:
|
||||||
repo-token: ${{ secrets.GITHUB_TOKEN }}
|
repo-token: ${{ secrets.GITHUB_TOKEN }}
|
||||||
sync-labels: true # Removes labels if files are removed from the PR
|
sync-labels: true # Removes labels if files are removed from the PR
|
||||||
|
|||||||
@@ -43,12 +43,12 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Set up Python
|
- name: Set up Python
|
||||||
uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
|
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v6
|
||||||
with:
|
with:
|
||||||
python-version: '3.12'
|
python-version: '3.12'
|
||||||
|
|
||||||
|
|||||||
@@ -38,12 +38,12 @@ jobs:
|
|||||||
|
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Set up Python
|
- name: Set up Python
|
||||||
uses: actions/setup-python@a309ff8b426b58ec0e2a45f0f869d46889d02405 # v6
|
uses: actions/setup-python@5fda3b95a4ea91299a34e894583c3862153e4b97 # v6
|
||||||
with:
|
with:
|
||||||
python-version: '3.12'
|
python-version: '3.12'
|
||||||
|
|
||||||
@@ -104,7 +104,7 @@ jobs:
|
|||||||
- name: Publish to TestPyPI for pre-releases
|
- name: Publish to TestPyPI for pre-releases
|
||||||
# True for tags like 'v0.2.0-rc1'
|
# True for tags like 'v0.2.0-rc1'
|
||||||
if: startsWith(github.ref, 'refs/tags/v') && contains(github.ref, '-')
|
if: startsWith(github.ref, 'refs/tags/v') && contains(github.ref, '-')
|
||||||
uses: pypa/gh-action-pypi-publish@ed0c53931b1dc9bd32cbe73a98c7f6766f8a527e # v1.13.0
|
uses: pypa/gh-action-pypi-publish@ba38be9e461d3875417946c167d0b5f3d385a247 # v1.14.1
|
||||||
with:
|
with:
|
||||||
repository-url: https://test.pypi.org/legacy/
|
repository-url: https://test.pypi.org/legacy/
|
||||||
verbose: true
|
verbose: true
|
||||||
@@ -112,7 +112,7 @@ jobs:
|
|||||||
|
|
||||||
- name: Publish to PyPI
|
- name: Publish to PyPI
|
||||||
if: startsWith(github.ref, 'refs/tags/v') && !contains(github.ref, '-')
|
if: startsWith(github.ref, 'refs/tags/v') && !contains(github.ref, '-')
|
||||||
uses: pypa/gh-action-pypi-publish@ed0c53931b1dc9bd32cbe73a98c7f6766f8a527e # v1.13.0
|
uses: pypa/gh-action-pypi-publish@ba38be9e461d3875417946c167d0b5f3d385a247 # v1.14.1
|
||||||
with:
|
with:
|
||||||
verbose: true
|
verbose: true
|
||||||
print-hash: true
|
print-hash: true
|
||||||
@@ -127,7 +127,7 @@ jobs:
|
|||||||
env:
|
env:
|
||||||
MUJOCO_GL: egl
|
MUJOCO_GL: egl
|
||||||
steps:
|
steps:
|
||||||
- uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
- uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
lfs: true
|
lfs: true
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
@@ -137,7 +137,7 @@ jobs:
|
|||||||
git curl libglib2.0-0 libegl1-mesa-dev ffmpeg libusb-1.0-0-dev \
|
git curl libglib2.0-0 libegl1-mesa-dev ffmpeg libusb-1.0-0-dev \
|
||||||
speech-dispatcher libgeos-dev portaudio19-dev
|
speech-dispatcher libgeos-dev portaudio19-dev
|
||||||
- name: Setup uv and Python
|
- name: Setup uv and Python
|
||||||
uses: astral-sh/setup-uv@d0cc045d04ccac9d8b7881df0226f9e82c39688e # v6
|
uses: astral-sh/setup-uv@11f9893b081a58869d3b5fccaea48c9e9e46f990 # v8.3.2
|
||||||
with:
|
with:
|
||||||
enable-cache: true # zizmor: ignore[cache-poisoning]
|
enable-cache: true # zizmor: ignore[cache-poisoning]
|
||||||
version: ${{ env.UV_VERSION }}
|
version: ${{ env.UV_VERSION }}
|
||||||
|
|||||||
@@ -43,12 +43,12 @@ jobs:
|
|||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
steps:
|
steps:
|
||||||
- name: Checkout code
|
- name: Checkout code
|
||||||
uses: actions/checkout@de0fac2e4500dabe0009e67214ff5f5447ce83dd # v6.0.2
|
uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1
|
||||||
with:
|
with:
|
||||||
fetch-depth: 0
|
fetch-depth: 0
|
||||||
persist-credentials: false
|
persist-credentials: false
|
||||||
|
|
||||||
- name: Secret Scanning
|
- name: Secret Scanning
|
||||||
uses: trufflesecurity/trufflehog@eafb8c5f6a06175141c27f17bcc17941853d0047 # v3.90.0
|
uses: trufflesecurity/trufflehog@27b0417c16317ca9a472a9a8092acce143b49c55 # v3.95.9
|
||||||
with:
|
with:
|
||||||
extra_args: --only-verified
|
extra_args: --only-verified
|
||||||
|
|||||||
@@ -68,16 +68,17 @@ ENV HOME=/home/user_lerobot \
|
|||||||
# issues with MuJoCo and OpenGL drivers.
|
# issues with MuJoCo and OpenGL drivers.
|
||||||
RUN uv venv --python python${PYTHON_VERSION}
|
RUN uv venv --python python${PYTHON_VERSION}
|
||||||
|
|
||||||
# Install third-party dependencies separately for layer caching
|
# Install Python dependencies for caching
|
||||||
COPY --chown=user_lerobot:user_lerobot setup.py pyproject.toml uv.lock README.md MANIFEST.in ./
|
COPY --chown=user_lerobot:user_lerobot setup.py pyproject.toml uv.lock README.md MANIFEST.in ./
|
||||||
RUN uv sync --locked --extra all --no-install-project --no-cache
|
COPY --chown=user_lerobot:user_lerobot src/ src/
|
||||||
|
|
||||||
|
RUN uv sync --locked --extra all --no-cache
|
||||||
|
|
||||||
RUN chmod +x /lerobot/.venv/lib/python${PYTHON_VERSION}/site-packages/triton/backends/nvidia/bin/ptxas
|
RUN chmod +x /lerobot/.venv/lib/python${PYTHON_VERSION}/site-packages/triton/backends/nvidia/bin/ptxas
|
||||||
|
|
||||||
# Copy the application source code and install the local project
|
# Copy the rest of the application source code
|
||||||
# Make sure to have the git-LFS files for testing
|
# Make sure to have the git-LFS files for testing
|
||||||
COPY --chown=user_lerobot:user_lerobot . .
|
COPY --chown=user_lerobot:user_lerobot . .
|
||||||
RUN uv sync --locked --extra all --no-cache
|
|
||||||
|
|
||||||
# Set the default command
|
# Set the default command
|
||||||
CMD ["/bin/bash"]
|
CMD ["/bin/bash"]
|
||||||
|
|||||||
@@ -60,14 +60,15 @@ ENV HOME=/home/user_lerobot \
|
|||||||
# run other Python projects in the same container without dependency conflicts.
|
# run other Python projects in the same container without dependency conflicts.
|
||||||
RUN uv venv
|
RUN uv venv
|
||||||
|
|
||||||
# Install third-party dependencies separately for layer caching
|
# Install Python dependencies for caching
|
||||||
COPY --chown=user_lerobot:user_lerobot setup.py pyproject.toml uv.lock README.md MANIFEST.in ./
|
COPY --chown=user_lerobot:user_lerobot setup.py pyproject.toml uv.lock README.md MANIFEST.in ./
|
||||||
RUN uv sync --locked --extra all --no-install-project --no-cache
|
COPY --chown=user_lerobot:user_lerobot src/ src/
|
||||||
|
|
||||||
# Copy the application code and install the local project
|
RUN uv sync --locked --extra all --no-cache
|
||||||
|
|
||||||
|
# Copy the rest of the application code
|
||||||
# Make sure to have the git-LFS files for testing
|
# Make sure to have the git-LFS files for testing
|
||||||
COPY --chown=user_lerobot:user_lerobot . .
|
COPY --chown=user_lerobot:user_lerobot . .
|
||||||
RUN uv sync --locked --extra all --no-cache
|
|
||||||
|
|
||||||
# Set the default command
|
# Set the default command
|
||||||
CMD ["/bin/bash"]
|
CMD ["/bin/bash"]
|
||||||
|
|||||||
@@ -136,10 +136,6 @@ config = RealSenseCameraConfig(
|
|||||||
height=480,
|
height=480,
|
||||||
color_mode=ColorMode.RGB,
|
color_mode=ColorMode.RGB,
|
||||||
use_depth=True,
|
use_depth=True,
|
||||||
# Optional fixed color controls. Omit them to leave the current sensor settings unchanged.
|
|
||||||
exposure=120,
|
|
||||||
gain=64,
|
|
||||||
white_balance=4600,
|
|
||||||
rotation=Cv2Rotation.NO_ROTATION
|
rotation=Cv2Rotation.NO_ROTATION
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -158,15 +154,6 @@ finally:
|
|||||||
```
|
```
|
||||||
<!-- prettier-ignore-end -->
|
<!-- prettier-ignore-end -->
|
||||||
|
|
||||||
Manual color controls disable the corresponding automatic exposure or white-balance mode. Their
|
|
||||||
supported ranges vary by camera model; an invalid value raises an error at connection time that
|
|
||||||
includes the range reported by the sensor. Requesting an unsupported control also raises an error.
|
|
||||||
Omitted controls leave the sensor's existing automatic or manual setting unchanged. These options
|
|
||||||
require `use_rgb=True`.
|
|
||||||
|
|
||||||
On the RealSense D405, the color stream is provided by the Stereo Module, so changing manual
|
|
||||||
exposure or gain also affects the depth stream.
|
|
||||||
|
|
||||||
</hfoption>
|
</hfoption>
|
||||||
</hfoptions>
|
</hfoptions>
|
||||||
|
|
||||||
|
|||||||
Binary file not shown.
|
Before Width: | Height: | Size: 682 KiB |
@@ -494,19 +494,6 @@ ignore_errors = true
|
|||||||
module = "lerobot.envs.*"
|
module = "lerobot.envs.*"
|
||||||
ignore_errors = false
|
ignore_errors = false
|
||||||
|
|
||||||
[[tool.mypy.overrides]]
|
|
||||||
module = "lerobot.annotations.*"
|
|
||||||
ignore_errors = false
|
|
||||||
disallow_untyped_defs = true
|
|
||||||
disallow_incomplete_defs = true
|
|
||||||
check_untyped_defs = true
|
|
||||||
|
|
||||||
[[tool.mypy.overrides]]
|
|
||||||
module = "lerobot.transforms.*"
|
|
||||||
ignore_errors = false
|
|
||||||
disallow_untyped_defs = true
|
|
||||||
disallow_incomplete_defs = true
|
|
||||||
check_untyped_defs = true
|
|
||||||
|
|
||||||
# [[tool.mypy.overrides]]
|
# [[tool.mypy.overrides]]
|
||||||
# module = "lerobot.utils.*"
|
# module = "lerobot.utils.*"
|
||||||
|
|||||||
@@ -120,22 +120,14 @@ class OpenCVCamera(Camera):
|
|||||||
self.rotation: int | None = get_cv2_rotation(config.rotation)
|
self.rotation: int | None = get_cv2_rotation(config.rotation)
|
||||||
self.backend: int = config.backend
|
self.backend: int = config.backend
|
||||||
|
|
||||||
self.capture_width: int | None = None
|
if self.height and self.width:
|
||||||
self.capture_height: int | None = None
|
|
||||||
self._reset_connection_settings()
|
|
||||||
|
|
||||||
def __str__(self) -> str:
|
|
||||||
return f"{self.__class__.__name__}({self.index_or_path})"
|
|
||||||
|
|
||||||
def _reset_connection_settings(self) -> None:
|
|
||||||
"""Restore settings that may have been auto-detected during a failed connection."""
|
|
||||||
self.fps = self.config.fps
|
|
||||||
self.width = self.config.width
|
|
||||||
self.height = self.config.height
|
|
||||||
self.capture_width, self.capture_height = self.width, self.height
|
self.capture_width, self.capture_height = self.width, self.height
|
||||||
if self.rotation in [cv2.ROTATE_90_CLOCKWISE, cv2.ROTATE_90_COUNTERCLOCKWISE]:
|
if self.rotation in [cv2.ROTATE_90_CLOCKWISE, cv2.ROTATE_90_COUNTERCLOCKWISE]:
|
||||||
self.capture_width, self.capture_height = self.height, self.width
|
self.capture_width, self.capture_height = self.height, self.width
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
return f"{self.__class__.__name__}({self.index_or_path})"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_connected(self) -> bool:
|
def is_connected(self) -> bool:
|
||||||
"""Checks if the camera is currently connected and opened."""
|
"""Checks if the camera is currently connected and opened."""
|
||||||
@@ -172,7 +164,6 @@ class OpenCVCamera(Camera):
|
|||||||
f"Failed to open {self}.Run `lerobot-find-cameras opencv` to find available cameras."
|
f"Failed to open {self}.Run `lerobot-find-cameras opencv` to find available cameras."
|
||||||
)
|
)
|
||||||
|
|
||||||
try:
|
|
||||||
self._configure_capture_settings()
|
self._configure_capture_settings()
|
||||||
self._start_read_thread()
|
self._start_read_thread()
|
||||||
|
|
||||||
@@ -184,13 +175,6 @@ class OpenCVCamera(Camera):
|
|||||||
with self.frame_lock:
|
with self.frame_lock:
|
||||||
if self.latest_frame is None:
|
if self.latest_frame is None:
|
||||||
raise ConnectionError(f"{self} failed to capture frames during warmup.")
|
raise ConnectionError(f"{self} failed to capture frames during warmup.")
|
||||||
except BaseException:
|
|
||||||
try:
|
|
||||||
self._cleanup_resources()
|
|
||||||
except Exception:
|
|
||||||
logger.exception(f"Failed to fully clean up {self} after connect() failed.")
|
|
||||||
self._reset_connection_settings()
|
|
||||||
raise
|
|
||||||
|
|
||||||
logger.info(f"{self} connected.")
|
logger.info(f"{self} connected.")
|
||||||
|
|
||||||
@@ -328,7 +312,6 @@ class OpenCVCamera(Camera):
|
|||||||
|
|
||||||
for target in targets_to_scan:
|
for target in targets_to_scan:
|
||||||
camera = cv2.VideoCapture(target)
|
camera = cv2.VideoCapture(target)
|
||||||
try:
|
|
||||||
if camera.isOpened():
|
if camera.isOpened():
|
||||||
default_width = int(camera.get(cv2.CAP_PROP_FRAME_WIDTH))
|
default_width = int(camera.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||||
default_height = int(camera.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
default_height = int(camera.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||||
@@ -338,9 +321,7 @@ class OpenCVCamera(Camera):
|
|||||||
# Get FOURCC code and convert to string
|
# Get FOURCC code and convert to string
|
||||||
default_fourcc_code = camera.get(cv2.CAP_PROP_FOURCC)
|
default_fourcc_code = camera.get(cv2.CAP_PROP_FOURCC)
|
||||||
default_fourcc_code_int = int(default_fourcc_code)
|
default_fourcc_code_int = int(default_fourcc_code)
|
||||||
default_fourcc = "".join(
|
default_fourcc = "".join([chr((default_fourcc_code_int >> 8 * i) & 0xFF) for i in range(4)])
|
||||||
[chr((default_fourcc_code_int >> 8 * i) & 0xFF) for i in range(4)]
|
|
||||||
)
|
|
||||||
|
|
||||||
camera_info = {
|
camera_info = {
|
||||||
"name": f"OpenCV Camera @ {target}",
|
"name": f"OpenCV Camera @ {target}",
|
||||||
@@ -357,7 +338,6 @@ class OpenCVCamera(Camera):
|
|||||||
}
|
}
|
||||||
|
|
||||||
found_cameras_info.append(camera_info)
|
found_cameras_info.append(camera_info)
|
||||||
finally:
|
|
||||||
camera.release()
|
camera.release()
|
||||||
|
|
||||||
return found_cameras_info
|
return found_cameras_info
|
||||||
@@ -516,26 +496,6 @@ class OpenCVCamera(Camera):
|
|||||||
self.latest_timestamp = None
|
self.latest_timestamp = None
|
||||||
self.new_frame_event.clear()
|
self.new_frame_event.clear()
|
||||||
|
|
||||||
def _cleanup_resources(self) -> None:
|
|
||||||
"""Stop background reads and release the capture, including after partial setup."""
|
|
||||||
read_thread = self.thread
|
|
||||||
videocapture = self.videocapture
|
|
||||||
|
|
||||||
try:
|
|
||||||
self._stop_read_thread()
|
|
||||||
finally:
|
|
||||||
self.videocapture = None
|
|
||||||
try:
|
|
||||||
if videocapture is not None:
|
|
||||||
videocapture.release()
|
|
||||||
finally:
|
|
||||||
# Releasing the device may unblock a hardware read that outlived
|
|
||||||
# the first bounded join in _stop_read_thread().
|
|
||||||
if read_thread is not None and read_thread.is_alive():
|
|
||||||
read_thread.join(timeout=2.0)
|
|
||||||
if read_thread.is_alive(): # pragma: no cover
|
|
||||||
logger.warning(f"{self} read thread remained alive after releasing the capture.")
|
|
||||||
|
|
||||||
@check_if_not_connected
|
@check_if_not_connected
|
||||||
def async_read(self, timeout_ms: float = 200) -> NDArray[Any]:
|
def async_read(self, timeout_ms: float = 200) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
@@ -626,6 +586,16 @@ class OpenCVCamera(Camera):
|
|||||||
if not self.is_connected and self.thread is None:
|
if not self.is_connected and self.thread is None:
|
||||||
raise DeviceNotConnectedError(f"{self} not connected.")
|
raise DeviceNotConnectedError(f"{self} not connected.")
|
||||||
|
|
||||||
self._cleanup_resources()
|
if self.thread is not None:
|
||||||
|
self._stop_read_thread()
|
||||||
|
|
||||||
|
if self.videocapture is not None:
|
||||||
|
self.videocapture.release()
|
||||||
|
self.videocapture = None
|
||||||
|
|
||||||
|
with self.frame_lock:
|
||||||
|
self.latest_frame = None
|
||||||
|
self.latest_timestamp = None
|
||||||
|
self.new_frame_event.clear()
|
||||||
|
|
||||||
logger.info(f"{self} disconnected.")
|
logger.info(f"{self} disconnected.")
|
||||||
|
|||||||
@@ -121,9 +121,6 @@ class RealSenseCamera(Camera):
|
|||||||
|
|
||||||
self.config = config
|
self.config = config
|
||||||
|
|
||||||
self.width: int | None = config.width
|
|
||||||
self.height: int | None = config.height
|
|
||||||
|
|
||||||
if config.serial_number_or_name.isdigit():
|
if config.serial_number_or_name.isdigit():
|
||||||
self.serial_number = config.serial_number_or_name
|
self.serial_number = config.serial_number_or_name
|
||||||
else:
|
else:
|
||||||
@@ -134,9 +131,6 @@ class RealSenseCamera(Camera):
|
|||||||
self.use_rgb = config.use_rgb
|
self.use_rgb = config.use_rgb
|
||||||
self.use_depth = config.use_depth
|
self.use_depth = config.use_depth
|
||||||
self.warmup_s = config.warmup_s
|
self.warmup_s = config.warmup_s
|
||||||
self.exposure: int | None = config.exposure
|
|
||||||
self.gain: int | None = config.gain
|
|
||||||
self.white_balance: int | None = config.white_balance
|
|
||||||
|
|
||||||
self.rs_pipeline: rs.pipeline | None = None
|
self.rs_pipeline: rs.pipeline | None = None
|
||||||
self.rs_profile: rs.pipeline_profile | None = None
|
self.rs_profile: rs.pipeline_profile | None = None
|
||||||
@@ -151,23 +145,14 @@ class RealSenseCamera(Camera):
|
|||||||
|
|
||||||
self.rotation: int | None = get_cv2_rotation(config.rotation)
|
self.rotation: int | None = get_cv2_rotation(config.rotation)
|
||||||
|
|
||||||
self.capture_width: int | None = None
|
if self.height and self.width:
|
||||||
self.capture_height: int | None = None
|
|
||||||
self._reset_connection_settings()
|
|
||||||
|
|
||||||
def __str__(self) -> str:
|
|
||||||
return f"{self.__class__.__name__}({self.serial_number})"
|
|
||||||
|
|
||||||
def _reset_connection_settings(self) -> None:
|
|
||||||
"""Restore settings that may have been auto-detected during a failed connection."""
|
|
||||||
self.fps = self.config.fps
|
|
||||||
self.width = self.config.width
|
|
||||||
self.height = self.config.height
|
|
||||||
self.warmup_s = self.config.warmup_s
|
|
||||||
self.capture_width, self.capture_height = self.width, self.height
|
self.capture_width, self.capture_height = self.width, self.height
|
||||||
if self.rotation in [cv2.ROTATE_90_CLOCKWISE, cv2.ROTATE_90_COUNTERCLOCKWISE]:
|
if self.rotation in [cv2.ROTATE_90_CLOCKWISE, cv2.ROTATE_90_COUNTERCLOCKWISE]:
|
||||||
self.capture_width, self.capture_height = self.height, self.width
|
self.capture_width, self.capture_height = self.height, self.width
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
return f"{self.__class__.__name__}({self.serial_number})"
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_connected(self) -> bool:
|
def is_connected(self) -> bool:
|
||||||
"""Checks if the camera pipeline is started and streams are active."""
|
"""Checks if the camera pipeline is started and streams are active."""
|
||||||
@@ -187,8 +172,7 @@ class RealSenseCamera(Camera):
|
|||||||
|
|
||||||
Raises:
|
Raises:
|
||||||
DeviceAlreadyConnectedError: If the camera is already connected.
|
DeviceAlreadyConnectedError: If the camera is already connected.
|
||||||
ValueError: If the configuration is invalid, a requested sensor option is unsupported,
|
ValueError: If the configuration is invalid (e.g., missing serial/name, name not unique).
|
||||||
or a requested sensor value is invalid.
|
|
||||||
ConnectionError: If the camera is found but fails to start the pipeline or no RealSense devices are detected at all.
|
ConnectionError: If the camera is found but fails to start the pipeline or no RealSense devices are detected at all.
|
||||||
RuntimeError: If the pipeline starts but fails to apply requested settings.
|
RuntimeError: If the pipeline starts but fails to apply requested settings.
|
||||||
"""
|
"""
|
||||||
@@ -206,9 +190,7 @@ class RealSenseCamera(Camera):
|
|||||||
f"Failed to open {self}.Run `lerobot-find-cameras realsense` to find available cameras."
|
f"Failed to open {self}.Run `lerobot-find-cameras realsense` to find available cameras."
|
||||||
) from e
|
) from e
|
||||||
|
|
||||||
try:
|
|
||||||
self._configure_capture_settings()
|
self._configure_capture_settings()
|
||||||
self._configure_sensor_options()
|
|
||||||
self._start_read_thread()
|
self._start_read_thread()
|
||||||
|
|
||||||
# NOTE(Steven/Caroline): Enforcing at least one second of warmup as RS cameras need a bit of time before the first read. If we don't wait, the first read from the warmup will raise.
|
# NOTE(Steven/Caroline): Enforcing at least one second of warmup as RS cameras need a bit of time before the first read. If we don't wait, the first read from the warmup will raise.
|
||||||
@@ -224,13 +206,6 @@ class RealSenseCamera(Camera):
|
|||||||
self.use_depth and self.latest_depth_frame is None
|
self.use_depth and self.latest_depth_frame is None
|
||||||
):
|
):
|
||||||
raise ConnectionError(f"{self} failed to capture frames during warmup.")
|
raise ConnectionError(f"{self} failed to capture frames during warmup.")
|
||||||
except BaseException:
|
|
||||||
try:
|
|
||||||
self._cleanup_resources()
|
|
||||||
except Exception:
|
|
||||||
logger.exception(f"Failed to fully clean up {self} after connect() failed.")
|
|
||||||
self._reset_connection_settings()
|
|
||||||
raise
|
|
||||||
|
|
||||||
logger.info(f"{self} connected.")
|
logger.info(f"{self} connected.")
|
||||||
|
|
||||||
@@ -364,111 +339,6 @@ class RealSenseCamera(Camera):
|
|||||||
self.new_frame_event.clear()
|
self.new_frame_event.clear()
|
||||||
return self._async_read(timeout_ms=10000, read_depth=read_depth)
|
return self._async_read(timeout_ms=10000, read_depth=read_depth)
|
||||||
|
|
||||||
def _get_color_sensor(self) -> "rs.sensor":
|
|
||||||
"""Returns the sensor that controls the color stream.
|
|
||||||
|
|
||||||
Most RealSense cameras expose "RGB Camera" for color. The D405 has no
|
|
||||||
separate RGB module — its color stream comes from "Stereo Module".
|
|
||||||
We try RGB Camera first, then fall back to Stereo Module.
|
|
||||||
"""
|
|
||||||
if self.rs_profile is None:
|
|
||||||
raise RuntimeError(f"{self}: rs_profile must be initialized before use.")
|
|
||||||
|
|
||||||
device = self.rs_profile.get_device()
|
|
||||||
sensors = {s.get_info(rs.camera_info.name): s for s in device.query_sensors()}
|
|
||||||
|
|
||||||
for name in ("RGB Camera", "Stereo Module"):
|
|
||||||
if name in sensors:
|
|
||||||
return sensors[name]
|
|
||||||
|
|
||||||
available = list(sensors.keys())
|
|
||||||
raise RuntimeError(f"{self}: no color sensor found. Available sensors: {available}")
|
|
||||||
|
|
||||||
def _set_sensor_option(self, sensor: "rs.sensor", option: "rs.option", value: float, label: str) -> None:
|
|
||||||
"""Sets a sensor option, re-raising range errors with actionable diagnostics."""
|
|
||||||
try:
|
|
||||||
sensor.set_option(option, value)
|
|
||||||
except Exception as e:
|
|
||||||
range_info = ""
|
|
||||||
try:
|
|
||||||
option_range = sensor.get_option_range(option)
|
|
||||||
range_info = (
|
|
||||||
f" (supported range: min={option_range.min}, max={option_range.max}, "
|
|
||||||
f"step={option_range.step}, default={option_range.default})"
|
|
||||||
)
|
|
||||||
except Exception:
|
|
||||||
range_info = " (option range unavailable)"
|
|
||||||
raise ValueError(
|
|
||||||
f"{self}: failed to set {label} to {value}{range_info}. Original error: {e}"
|
|
||||||
) from e
|
|
||||||
|
|
||||||
def _configure_sensor_options(self) -> None:
|
|
||||||
"""Applies manual sensor options (exposure, gain, white balance) to the color sensor.
|
|
||||||
|
|
||||||
When exposure or gain is set, auto-exposure is disabled first. When white_balance
|
|
||||||
is set, auto white balance is disabled first. An omitted option is left unchanged,
|
|
||||||
and configuration is skipped entirely if all options are omitted.
|
|
||||||
|
|
||||||
Raises:
|
|
||||||
ValueError: If the sensor does not support a requested option or a requested
|
|
||||||
value is invalid. Invalid-value errors include the option name, requested
|
|
||||||
value, and supported range when available.
|
|
||||||
"""
|
|
||||||
if self.exposure is None and self.gain is None and self.white_balance is None:
|
|
||||||
return
|
|
||||||
|
|
||||||
color_sensor = self._get_color_sensor()
|
|
||||||
|
|
||||||
requested_options = (
|
|
||||||
(rs.option.exposure, self.exposure, "exposure"),
|
|
||||||
(rs.option.gain, self.gain, "gain"),
|
|
||||||
(rs.option.white_balance, self.white_balance, "white balance"),
|
|
||||||
)
|
|
||||||
unsupported_options = [
|
|
||||||
label
|
|
||||||
for option, value, label in requested_options
|
|
||||||
if value is not None and not color_sensor.supports(option)
|
|
||||||
]
|
|
||||||
if unsupported_options:
|
|
||||||
raise ValueError(
|
|
||||||
f"{self}: color sensor does not support requested manual options: {unsupported_options}."
|
|
||||||
)
|
|
||||||
|
|
||||||
manual_exposure_requested = self.exposure is not None or self.gain is not None
|
|
||||||
if manual_exposure_requested:
|
|
||||||
if color_sensor.supports(rs.option.enable_auto_exposure):
|
|
||||||
self._set_sensor_option(color_sensor, rs.option.enable_auto_exposure, 0, "auto-exposure")
|
|
||||||
logger.info(f"{self} auto-exposure disabled.")
|
|
||||||
else:
|
|
||||||
logger.warning(
|
|
||||||
f"{self} sensor does not support disabling auto-exposure; "
|
|
||||||
"applying manual exposure/gain directly."
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.exposure is not None:
|
|
||||||
self._set_sensor_option(color_sensor, rs.option.exposure, self.exposure, "exposure")
|
|
||||||
logger.info(f"{self} exposure set to {self.exposure}.")
|
|
||||||
|
|
||||||
if self.gain is not None:
|
|
||||||
self._set_sensor_option(color_sensor, rs.option.gain, self.gain, "gain")
|
|
||||||
logger.info(f"{self} gain set to {self.gain}.")
|
|
||||||
|
|
||||||
if self.white_balance is not None:
|
|
||||||
if color_sensor.supports(rs.option.enable_auto_white_balance):
|
|
||||||
self._set_sensor_option(
|
|
||||||
color_sensor, rs.option.enable_auto_white_balance, 0, "auto white balance"
|
|
||||||
)
|
|
||||||
logger.info(f"{self} auto white balance disabled.")
|
|
||||||
else:
|
|
||||||
logger.warning(
|
|
||||||
f"{self} sensor does not support disabling auto white balance; "
|
|
||||||
"applying manual white balance directly."
|
|
||||||
)
|
|
||||||
self._set_sensor_option(
|
|
||||||
color_sensor, rs.option.white_balance, self.white_balance, "white balance"
|
|
||||||
)
|
|
||||||
logger.info(f"{self} white balance set to {self.white_balance}.")
|
|
||||||
|
|
||||||
@check_if_not_connected
|
@check_if_not_connected
|
||||||
def read_depth(self, timeout_ms: int = 200) -> NDArray[Any]:
|
def read_depth(self, timeout_ms: int = 200) -> NDArray[Any]:
|
||||||
"""
|
"""
|
||||||
@@ -671,27 +541,6 @@ class RealSenseCamera(Camera):
|
|||||||
self.latest_timestamp = None
|
self.latest_timestamp = None
|
||||||
self.new_frame_event.clear()
|
self.new_frame_event.clear()
|
||||||
|
|
||||||
def _cleanup_resources(self) -> None:
|
|
||||||
"""Stop background reads and stop the pipeline, including after partial setup."""
|
|
||||||
read_thread = self.thread
|
|
||||||
rs_pipeline = self.rs_pipeline
|
|
||||||
|
|
||||||
try:
|
|
||||||
self._stop_read_thread()
|
|
||||||
finally:
|
|
||||||
self.rs_pipeline = None
|
|
||||||
self.rs_profile = None
|
|
||||||
try:
|
|
||||||
if rs_pipeline is not None:
|
|
||||||
rs_pipeline.stop()
|
|
||||||
finally:
|
|
||||||
# Stopping the pipeline may unblock a hardware read that outlived
|
|
||||||
# the first bounded join in _stop_read_thread().
|
|
||||||
if read_thread is not None and read_thread.is_alive():
|
|
||||||
read_thread.join(timeout=2.0)
|
|
||||||
if read_thread.is_alive(): # pragma: no cover
|
|
||||||
logger.warning(f"{self} read thread remained alive after stopping the pipeline.")
|
|
||||||
|
|
||||||
def _async_read(self, timeout_ms: float, read_depth: bool = False) -> NDArray[Any]:
|
def _async_read(self, timeout_ms: float, read_depth: bool = False) -> NDArray[Any]:
|
||||||
"""Shared helper for :meth:`async_read`/:meth:`async_read_depth`: return the latest buffered frame."""
|
"""Shared helper for :meth:`async_read`/:meth:`async_read_depth`: return the latest buffered frame."""
|
||||||
if self.thread is None or not self.thread.is_alive():
|
if self.thread is None or not self.thread.is_alive():
|
||||||
@@ -835,5 +684,18 @@ class RealSenseCamera(Camera):
|
|||||||
f"Attempted to disconnect {self}, but it appears already disconnected."
|
f"Attempted to disconnect {self}, but it appears already disconnected."
|
||||||
)
|
)
|
||||||
|
|
||||||
self._cleanup_resources()
|
if self.thread is not None:
|
||||||
|
self._stop_read_thread()
|
||||||
|
|
||||||
|
if self.rs_pipeline is not None:
|
||||||
|
self.rs_pipeline.stop()
|
||||||
|
self.rs_pipeline = None
|
||||||
|
self.rs_profile = None
|
||||||
|
|
||||||
|
with self.frame_lock:
|
||||||
|
self.latest_color_frame = None
|
||||||
|
self.latest_depth_frame = None
|
||||||
|
self.latest_timestamp = None
|
||||||
|
self.new_frame_event.clear()
|
||||||
|
|
||||||
logger.info(f"{self} disconnected.")
|
logger.info(f"{self} disconnected.")
|
||||||
|
|||||||
@@ -46,17 +46,6 @@ class RealSenseCameraConfig(CameraConfig):
|
|||||||
use_depth: Whether to enable depth stream. Defaults to False.
|
use_depth: Whether to enable depth stream. Defaults to False.
|
||||||
rotation: Image rotation setting (0°, 90°, 180°, or 270°). Defaults to no rotation.
|
rotation: Image rotation setting (0°, 90°, 180°, or 270°). Defaults to no rotation.
|
||||||
warmup_s: Time reading frames before returning from connect (in seconds)
|
warmup_s: Time reading frames before returning from connect (in seconds)
|
||||||
exposure: Manual exposure value for the color sensor. When set, auto-exposure is
|
|
||||||
disabled and this fixed value is used. Valid ranges are camera-model specific
|
|
||||||
and reported if the value is rejected. Defaults to None (leave unchanged).
|
|
||||||
gain: Manual gain value for the color sensor. When set, auto-exposure is disabled
|
|
||||||
and this fixed gain is used, which also freezes exposure at its current value
|
|
||||||
when no exposure is configured. Valid ranges are camera-model specific and
|
|
||||||
reported if the value is rejected. Defaults to None (leave unchanged).
|
|
||||||
white_balance: Manual white balance value for the color sensor. When set, auto
|
|
||||||
white balance is disabled and this fixed value is used. Valid ranges are
|
|
||||||
camera-model specific and reported if the value is rejected. Defaults to None
|
|
||||||
(leave unchanged).
|
|
||||||
|
|
||||||
Note:
|
Note:
|
||||||
- Either name or serial_number must be specified.
|
- Either name or serial_number must be specified.
|
||||||
@@ -72,9 +61,6 @@ class RealSenseCameraConfig(CameraConfig):
|
|||||||
use_depth: bool = False
|
use_depth: bool = False
|
||||||
rotation: Cv2Rotation = Cv2Rotation.NO_ROTATION
|
rotation: Cv2Rotation = Cv2Rotation.NO_ROTATION
|
||||||
warmup_s: int = 1
|
warmup_s: int = 1
|
||||||
exposure: int | None = None
|
|
||||||
gain: int | None = None
|
|
||||||
white_balance: int | None = None
|
|
||||||
|
|
||||||
def __post_init__(self) -> None:
|
def __post_init__(self) -> None:
|
||||||
self.color_mode = ColorMode(self.color_mode)
|
self.color_mode = ColorMode(self.color_mode)
|
||||||
@@ -83,18 +69,6 @@ class RealSenseCameraConfig(CameraConfig):
|
|||||||
if not self.use_rgb and not self.use_depth:
|
if not self.use_rgb and not self.use_depth:
|
||||||
raise ValueError("At least one of `use_rgb` or `use_depth` must be enabled.")
|
raise ValueError("At least one of `use_rgb` or `use_depth` must be enabled.")
|
||||||
|
|
||||||
manual_color_options = {
|
|
||||||
"exposure": self.exposure,
|
|
||||||
"gain": self.gain,
|
|
||||||
"white_balance": self.white_balance,
|
|
||||||
}
|
|
||||||
configured_color_options = [name for name, value in manual_color_options.items() if value is not None]
|
|
||||||
if configured_color_options and not self.use_rgb:
|
|
||||||
raise ValueError(
|
|
||||||
"Manual color sensor options require `use_rgb=True`. "
|
|
||||||
f"Configured options: {configured_color_options}."
|
|
||||||
)
|
|
||||||
|
|
||||||
values = (self.fps, self.width, self.height)
|
values = (self.fps, self.width, self.height)
|
||||||
if any(v is not None for v in values) and any(v is None for v in values):
|
if any(v is not None for v in values) and any(v is None for v in values):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
|
|||||||
@@ -71,19 +71,13 @@ class DatasetRecordConfig:
|
|||||||
# Number of threads per encoder instance. None = auto (codec default).
|
# Number of threads per encoder instance. None = auto (codec default).
|
||||||
# Lower values reduce CPU usage, maps to 'lp' (via svtav1-params) for libsvtav1 and 'threads' for h264/hevc..
|
# Lower values reduce CPU usage, maps to 'lp' (via svtav1-params) for libsvtav1 and 'threads' for h264/hevc..
|
||||||
encoder_threads: int | None = None
|
encoder_threads: int | None = None
|
||||||
# Skip appending the date-time tag to repo_id, keeping the user-provided name as-is
|
|
||||||
# (e.g. self-managed versioned names intended for a later `lerobot-edit-dataset merge`).
|
|
||||||
no_stamp: bool = False
|
|
||||||
|
|
||||||
def stamp_repo_id(self) -> None:
|
def stamp_repo_id(self) -> None:
|
||||||
"""Append a date-time tag to ``repo_id`` so each recording session gets a unique name.
|
"""Append a date-time tag to ``repo_id`` so each recording session gets a unique name.
|
||||||
|
|
||||||
Must be called explicitly at dataset *creation* time — not on resume,
|
Must be called explicitly at dataset *creation* time — not on resume,
|
||||||
where the existing ``repo_id`` (already stamped) must be preserved.
|
where the existing ``repo_id`` (already stamped) must be preserved.
|
||||||
No-op when ``no_stamp`` is set, preserving a user-managed ``repo_id``.
|
|
||||||
"""
|
"""
|
||||||
if self.no_stamp:
|
|
||||||
return
|
|
||||||
if self.repo_id:
|
if self.repo_id:
|
||||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||||
self.repo_id = f"{self.repo_id}_{timestamp}"
|
self.repo_id = f"{self.repo_id}_{timestamp}"
|
||||||
|
|||||||
@@ -14,7 +14,6 @@
|
|||||||
import builtins
|
import builtins
|
||||||
import datetime as dt
|
import datetime as dt
|
||||||
import json
|
import json
|
||||||
import multiprocessing
|
|
||||||
import os
|
import os
|
||||||
import tempfile
|
import tempfile
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
@@ -102,12 +101,6 @@ class TrainPipelineConfig(HubMixin):
|
|||||||
batch_size: int = 8
|
batch_size: int = 8
|
||||||
prefetch_factor: int = 4
|
prefetch_factor: int = 4
|
||||||
persistent_workers: bool = True
|
persistent_workers: bool = True
|
||||||
# DataLoader worker start method. "spawn" is safer than "fork" with
|
|
||||||
# non-fork-safe libs (PyAV / torchcodec / ffmpeg), but adds some
|
|
||||||
# worker-startup time per run since workers re-import modules instead
|
|
||||||
# of inheriting parent state. Override with `--dataloader_multiprocessing_context=fork`
|
|
||||||
# when appropriate, or set it to `null` to use Python's platform default.
|
|
||||||
dataloader_multiprocessing_context: str | None = "spawn"
|
|
||||||
steps: int = 100_000
|
steps: int = 100_000
|
||||||
# Run policy in the simulation environment every N steps to measure reward/success (0 = disabled).
|
# Run policy in the simulation environment every N steps to measure reward/success (0 = disabled).
|
||||||
env_eval_freq: int = 20_000
|
env_eval_freq: int = 20_000
|
||||||
@@ -219,17 +212,6 @@ class TrainPipelineConfig(HubMixin):
|
|||||||
self.reward_model.pretrained_path = str(policy_dir)
|
self.reward_model.pretrained_path = str(policy_dir)
|
||||||
|
|
||||||
def validate(self) -> None:
|
def validate(self) -> None:
|
||||||
available_contexts = multiprocessing.get_all_start_methods()
|
|
||||||
if (
|
|
||||||
self.dataloader_multiprocessing_context is not None
|
|
||||||
and self.dataloader_multiprocessing_context not in available_contexts
|
|
||||||
):
|
|
||||||
raise ValueError(
|
|
||||||
"`dataloader_multiprocessing_context` must be None or one of "
|
|
||||||
f"{available_contexts} on this platform, got "
|
|
||||||
f"{self.dataloader_multiprocessing_context!r}."
|
|
||||||
)
|
|
||||||
|
|
||||||
self._resolve_pretrained_from_cli()
|
self._resolve_pretrained_from_cli()
|
||||||
|
|
||||||
if self.policy is None and self.reward_model is None:
|
if self.policy is None and self.reward_model is None:
|
||||||
|
|||||||
@@ -188,8 +188,8 @@ class LeRobotDatasetMetadata:
|
|||||||
def _load_metadata(self):
|
def _load_metadata(self):
|
||||||
self.info = load_info(self.root)
|
self.info = load_info(self.root)
|
||||||
check_version_compatibility(self.repo_id, self._version, CODEBASE_VERSION)
|
check_version_compatibility(self.repo_id, self._version, CODEBASE_VERSION)
|
||||||
self.tasks = load_tasks(self.root) if self.total_tasks > 0 else None
|
self.tasks = load_tasks(self.root)
|
||||||
self.episodes = load_episodes(self.root) if self.total_episodes > 0 else None
|
self.episodes = load_episodes(self.root)
|
||||||
self.stats = load_stats(self.root)
|
self.stats = load_stats(self.root)
|
||||||
|
|
||||||
def ensure_readable(self) -> None:
|
def ensure_readable(self) -> None:
|
||||||
|
|||||||
@@ -172,23 +172,6 @@ class DatasetWriter:
|
|||||||
def _get_image_file_dir(self, episode_index: int, image_key: str) -> Path:
|
def _get_image_file_dir(self, episode_index: int, image_key: str) -> Path:
|
||||||
return self._get_image_file_path(episode_index, image_key, frame_index=0).parent
|
return self._get_image_file_path(episode_index, image_key, frame_index=0).parent
|
||||||
|
|
||||||
def _get_episode_buffer_index(self) -> int:
|
|
||||||
episode_index = self.episode_buffer["episode_index"]
|
|
||||||
# episode_index is `int` when freshly created, but becomes `np.ndarray` after
|
|
||||||
# save_episode() mutates the buffer. Handle both types here.
|
|
||||||
if isinstance(episode_index, np.ndarray):
|
|
||||||
episode_index = episode_index.item() if episode_index.size == 1 else episode_index[0]
|
|
||||||
return int(episode_index)
|
|
||||||
|
|
||||||
def _delete_camera_frame_dirs(self, camera_keys: list[str]) -> None:
|
|
||||||
if self.image_writer is not None:
|
|
||||||
self._wait_image_writer()
|
|
||||||
episode_index = self._get_episode_buffer_index()
|
|
||||||
for camera_key in camera_keys:
|
|
||||||
img_dir = self._get_image_file_dir(episode_index, camera_key)
|
|
||||||
if img_dir.is_dir():
|
|
||||||
shutil.rmtree(img_dir)
|
|
||||||
|
|
||||||
def _save_image(
|
def _save_image(
|
||||||
self, image: torch.Tensor | np.ndarray | PIL.Image.Image, fpath: Path, compress_level: int = 1
|
self, image: torch.Tensor | np.ndarray | PIL.Image.Image, fpath: Path, compress_level: int = 1
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -386,9 +369,7 @@ class DatasetWriter:
|
|||||||
self._episodes_since_last_encoding = 0
|
self._episodes_since_last_encoding = 0
|
||||||
|
|
||||||
if episode_data is None:
|
if episode_data is None:
|
||||||
if len(self._meta.image_keys) > 0:
|
self.clear_episode_buffer(delete_images=len(self._meta.image_keys) > 0)
|
||||||
self._delete_camera_frame_dirs(self._meta.image_keys)
|
|
||||||
self.episode_buffer = self._create_episode_buffer()
|
|
||||||
|
|
||||||
def _batch_save_episode_video(self, start_episode: int, end_episode: int | None = None) -> None:
|
def _batch_save_episode_video(self, start_episode: int, end_episode: int | None = None) -> None:
|
||||||
"""Batch save videos for multiple episodes."""
|
"""Batch save videos for multiple episodes."""
|
||||||
@@ -580,10 +561,10 @@ class DatasetWriter:
|
|||||||
return metadata
|
return metadata
|
||||||
|
|
||||||
def clear_episode_buffer(self, delete_images: bool = True) -> None:
|
def clear_episode_buffer(self, delete_images: bool = True) -> None:
|
||||||
"""Discard the current episode buffer and optionally delete temp camera frames.
|
"""Discard the current episode buffer and optionally delete temp images.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
delete_images: If ``True``, remove temporary camera frame directories
|
delete_images: If ``True``, remove temporary image directories
|
||||||
written for the current episode.
|
written for the current episode.
|
||||||
"""
|
"""
|
||||||
# Cancel streaming encoder if active
|
# Cancel streaming encoder if active
|
||||||
@@ -591,7 +572,17 @@ class DatasetWriter:
|
|||||||
self._streaming_encoder.cancel_episode()
|
self._streaming_encoder.cancel_episode()
|
||||||
|
|
||||||
if delete_images:
|
if delete_images:
|
||||||
self._delete_camera_frame_dirs(self._meta.camera_keys)
|
if self.image_writer is not None:
|
||||||
|
self._wait_image_writer()
|
||||||
|
episode_index = self.episode_buffer["episode_index"]
|
||||||
|
# episode_index is `int` when freshly created, but becomes `np.ndarray` after
|
||||||
|
# save_episode() mutates the buffer. Handle both types here.
|
||||||
|
if isinstance(episode_index, np.ndarray):
|
||||||
|
episode_index = episode_index.item() if episode_index.size == 1 else episode_index[0]
|
||||||
|
for cam_key in self._meta.image_keys:
|
||||||
|
img_dir = self._get_image_file_dir(episode_index, cam_key)
|
||||||
|
if img_dir.is_dir():
|
||||||
|
shutil.rmtree(img_dir)
|
||||||
|
|
||||||
self.episode_buffer = self._create_episode_buffer()
|
self.episode_buffer = self._create_episode_buffer()
|
||||||
|
|
||||||
|
|||||||
@@ -384,12 +384,7 @@ class LiberoEnv(gym.Env):
|
|||||||
|
|
||||||
def close(self):
|
def close(self):
|
||||||
if self._env is not None:
|
if self._env is not None:
|
||||||
try:
|
|
||||||
self._env.close()
|
self._env.close()
|
||||||
finally:
|
|
||||||
# LIBERO deletes its inner env on close, so this wrapper must
|
|
||||||
# be recreated before the next reset.
|
|
||||||
self._env = None
|
|
||||||
|
|
||||||
|
|
||||||
def _make_env_fns(
|
def _make_env_fns(
|
||||||
|
|||||||
@@ -155,7 +155,6 @@ class MetaworldEnv(gym.Env):
|
|||||||
env.model.cam_pos[2] = [0.75, 0.075, 0.7]
|
env.model.cam_pos[2] = [0.75, 0.075, 0.7]
|
||||||
env.reset()
|
env.reset()
|
||||||
env._freeze_rand_vec = False # otherwise no randomization
|
env._freeze_rand_vec = False # otherwise no randomization
|
||||||
env.seeded_rand_vec = True # use seeded RNG so reset(seed=X) controls object positions
|
|
||||||
self._env = env
|
self._env = env
|
||||||
|
|
||||||
def render(self) -> np.ndarray:
|
def render(self) -> np.ndarray:
|
||||||
@@ -221,8 +220,6 @@ class MetaworldEnv(gym.Env):
|
|||||||
self._ensure_env()
|
self._ensure_env()
|
||||||
super().reset(seed=seed)
|
super().reset(seed=seed)
|
||||||
|
|
||||||
if seed is not None:
|
|
||||||
self._env.seed(seed)
|
|
||||||
raw_obs, info = self._env.reset(seed=seed)
|
raw_obs, info = self._env.reset(seed=seed)
|
||||||
|
|
||||||
observation = self._format_raw_obs(raw_obs)
|
observation = self._format_raw_obs(raw_obs)
|
||||||
|
|||||||
@@ -384,9 +384,7 @@ class RoboTwinEnv(gym.Env):
|
|||||||
|
|
||||||
self._env: Any | None = None # deferred — created on first reset() inside worker
|
self._env: Any | None = None # deferred — created on first reset() inside worker
|
||||||
self._step_count: int = 0
|
self._step_count: int = 0
|
||||||
self._black_frame: np.ndarray = np.zeros(
|
self._black_frame = np.zeros((self.observation_height, self.observation_width, 3), dtype=np.uint8)
|
||||||
(self.observation_height, self.observation_width, 3), dtype=np.uint8
|
|
||||||
)
|
|
||||||
|
|
||||||
image_spaces = {
|
image_spaces = {
|
||||||
cam: spaces.Box(
|
cam: spaces.Box(
|
||||||
|
|||||||
@@ -373,7 +373,7 @@ class VLABenchEnv(gym.Env):
|
|||||||
|
|
||||||
if action.shape[0] != 7:
|
if action.shape[0] != 7:
|
||||||
# Unknown layout — fall back to zero-pad so the sim doesn't crash.
|
# Unknown layout — fall back to zero-pad so the sim doesn't crash.
|
||||||
padded: np.ndarray = np.zeros(ctrl_dim, dtype=np.float64)
|
padded = np.zeros(ctrl_dim, dtype=np.float64)
|
||||||
padded[: min(action.shape[0], ctrl_dim)] = action[:ctrl_dim]
|
padded[: min(action.shape[0], ctrl_dim)] = action[:ctrl_dim]
|
||||||
return padded
|
return padded
|
||||||
|
|
||||||
|
|||||||
@@ -122,9 +122,6 @@ MODEL_ENCODING_TABLE = {
|
|||||||
"xm430-w350": X_SERIES_ENCODINGS_TABLE,
|
"xm430-w350": X_SERIES_ENCODINGS_TABLE,
|
||||||
"xm540-w270": X_SERIES_ENCODINGS_TABLE,
|
"xm540-w270": X_SERIES_ENCODINGS_TABLE,
|
||||||
"xc430-w150": X_SERIES_ENCODINGS_TABLE,
|
"xc430-w150": X_SERIES_ENCODINGS_TABLE,
|
||||||
"xh540-w150": X_SERIES_ENCODINGS_TABLE,
|
|
||||||
"xc330-t288": X_SERIES_ENCODINGS_TABLE,
|
|
||||||
"xc330-t181": X_SERIES_ENCODINGS_TABLE,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
# {model: model_resolution}
|
# {model: model_resolution}
|
||||||
@@ -137,9 +134,6 @@ MODEL_RESOLUTION = {
|
|||||||
"xm430-w350": 4096,
|
"xm430-w350": 4096,
|
||||||
"xm540-w270": 4096,
|
"xm540-w270": 4096,
|
||||||
"xc430-w150": 4096,
|
"xc430-w150": 4096,
|
||||||
"xh540-w150": 4096,
|
|
||||||
"xc330-t288": 4096,
|
|
||||||
"xc330-t181": 4096,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
# {model: model_number}
|
# {model: model_number}
|
||||||
@@ -151,9 +145,6 @@ MODEL_NUMBER_TABLE = {
|
|||||||
"xm430-w350": 1020,
|
"xm430-w350": 1020,
|
||||||
"xm540-w270": 1120,
|
"xm540-w270": 1120,
|
||||||
"xc430-w150": 1070,
|
"xc430-w150": 1070,
|
||||||
"xh540-w150": 1110,
|
|
||||||
"xc330-t288": 1220,
|
|
||||||
"xc330-t181": 1210,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
# {model: available_operating_modes}
|
# {model: available_operating_modes}
|
||||||
@@ -165,9 +156,6 @@ MODEL_OPERATING_MODES = {
|
|||||||
"xm430-w350": [0, 1, 3, 4, 5, 16],
|
"xm430-w350": [0, 1, 3, 4, 5, 16],
|
||||||
"xm540-w270": [0, 1, 3, 4, 5, 16],
|
"xm540-w270": [0, 1, 3, 4, 5, 16],
|
||||||
"xc430-w150": [1, 3, 4, 16],
|
"xc430-w150": [1, 3, 4, 16],
|
||||||
"xh540-w150": [0, 1, 3, 4, 5, 16],
|
|
||||||
"xc330-t288": [0, 1, 3, 4, 5, 16],
|
|
||||||
"xc330-t181": [0, 1, 3, 4, 5, 16],
|
|
||||||
}
|
}
|
||||||
|
|
||||||
MODEL_CONTROL_TABLE = {
|
MODEL_CONTROL_TABLE = {
|
||||||
@@ -178,9 +166,6 @@ MODEL_CONTROL_TABLE = {
|
|||||||
"xm430-w350": X_SERIES_CONTROL_TABLE,
|
"xm430-w350": X_SERIES_CONTROL_TABLE,
|
||||||
"xm540-w270": X_SERIES_CONTROL_TABLE,
|
"xm540-w270": X_SERIES_CONTROL_TABLE,
|
||||||
"xc430-w150": X_SERIES_CONTROL_TABLE,
|
"xc430-w150": X_SERIES_CONTROL_TABLE,
|
||||||
"xh540-w150": X_SERIES_CONTROL_TABLE,
|
|
||||||
"xc330-t288": X_SERIES_CONTROL_TABLE,
|
|
||||||
"xc330-t181": X_SERIES_CONTROL_TABLE,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
MODEL_BAUDRATE_TABLE = {
|
MODEL_BAUDRATE_TABLE = {
|
||||||
@@ -191,9 +176,6 @@ MODEL_BAUDRATE_TABLE = {
|
|||||||
"xm430-w350": X_SERIES_BAUDRATE_TABLE,
|
"xm430-w350": X_SERIES_BAUDRATE_TABLE,
|
||||||
"xm540-w270": X_SERIES_BAUDRATE_TABLE,
|
"xm540-w270": X_SERIES_BAUDRATE_TABLE,
|
||||||
"xc430-w150": X_SERIES_BAUDRATE_TABLE,
|
"xc430-w150": X_SERIES_BAUDRATE_TABLE,
|
||||||
"xh540-w150": X_SERIES_BAUDRATE_TABLE,
|
|
||||||
"xc330-t288": X_SERIES_BAUDRATE_TABLE,
|
|
||||||
"xc330-t181": X_SERIES_BAUDRATE_TABLE,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
AVAILABLE_BAUDRATES = [
|
AVAILABLE_BAUDRATES = [
|
||||||
|
|||||||
@@ -44,19 +44,12 @@ from lerobot.utils.constants import (
|
|||||||
POLICY_PREPROCESSOR_DEFAULT_NAME,
|
POLICY_PREPROCESSOR_DEFAULT_NAME,
|
||||||
)
|
)
|
||||||
from lerobot.utils.feature_utils import dataset_to_policy_features
|
from lerobot.utils.feature_utils import dataset_to_policy_features
|
||||||
from lerobot.utils.import_utils import _peft_available, require_package
|
|
||||||
|
|
||||||
from .evo1.configuration_evo1 import Evo1Config
|
from .evo1.configuration_evo1 import Evo1Config
|
||||||
from .groot.configuration_groot import GrootConfig
|
from .groot.configuration_groot import GrootConfig
|
||||||
from .pretrained import PreTrainedPolicy
|
from .pretrained import PreTrainedPolicy
|
||||||
from .utils import validate_visual_features_consistency
|
from .utils import validate_visual_features_consistency
|
||||||
|
|
||||||
if TYPE_CHECKING or _peft_available:
|
|
||||||
from peft import PeftConfig, PeftModel
|
|
||||||
else:
|
|
||||||
PeftConfig = None
|
|
||||||
PeftModel = None
|
|
||||||
|
|
||||||
|
|
||||||
def _reconnect_relative_absolute_steps(
|
def _reconnect_relative_absolute_steps(
|
||||||
preprocessor: PolicyProcessorPipeline, postprocessor: PolicyProcessorPipeline
|
preprocessor: PolicyProcessorPipeline, postprocessor: PolicyProcessorPipeline
|
||||||
@@ -184,7 +177,6 @@ def make_pre_post_processors(
|
|||||||
return make_groot_pre_post_processors_from_pretrained(
|
return make_groot_pre_post_processors_from_pretrained(
|
||||||
config=policy_cfg,
|
config=policy_cfg,
|
||||||
pretrained_path=pretrained_path,
|
pretrained_path=pretrained_path,
|
||||||
revision=pretrained_revision,
|
|
||||||
dataset_stats=kwargs.get("dataset_stats"),
|
dataset_stats=kwargs.get("dataset_stats"),
|
||||||
dataset_meta=kwargs.get("dataset_meta"),
|
dataset_meta=kwargs.get("dataset_meta"),
|
||||||
preprocessor_overrides=kwargs.get("preprocessor_overrides"),
|
preprocessor_overrides=kwargs.get("preprocessor_overrides"),
|
||||||
@@ -341,15 +333,12 @@ def make_policy(
|
|||||||
# Load a pretrained PEFT model on top of the policy. The pretrained path points to the folder/repo
|
# Load a pretrained PEFT model on top of the policy. The pretrained path points to the folder/repo
|
||||||
# of the adapter and the adapter's config contains the path to the base policy. So we need the
|
# of the adapter and the adapter's config contains the path to the base policy. So we need the
|
||||||
# adapter config first, then load the correct policy and then apply PEFT.
|
# adapter config first, then load the correct policy and then apply PEFT.
|
||||||
require_package("peft", extra="peft")
|
from peft import PeftConfig, PeftModel
|
||||||
|
|
||||||
logging.info("Loading policy's PEFT adapter.")
|
logging.info("Loading policy's PEFT adapter.")
|
||||||
|
|
||||||
peft_pretrained_path = str(cfg.pretrained_path)
|
peft_pretrained_path = str(cfg.pretrained_path)
|
||||||
peft_config = PeftConfig.from_pretrained(
|
peft_config = PeftConfig.from_pretrained(peft_pretrained_path)
|
||||||
peft_pretrained_path,
|
|
||||||
revision=cfg.pretrained_revision,
|
|
||||||
)
|
|
||||||
|
|
||||||
kwargs["pretrained_name_or_path"] = peft_config.base_model_name_or_path
|
kwargs["pretrained_name_or_path"] = peft_config.base_model_name_or_path
|
||||||
if not kwargs["pretrained_name_or_path"]:
|
if not kwargs["pretrained_name_or_path"]:
|
||||||
@@ -360,14 +349,9 @@ def make_policy(
|
|||||||
"the adapter was trained."
|
"the adapter was trained."
|
||||||
)
|
)
|
||||||
|
|
||||||
kwargs["revision"] = peft_config.revision
|
|
||||||
policy = policy_cls.from_pretrained(**kwargs)
|
policy = policy_cls.from_pretrained(**kwargs)
|
||||||
policy = PeftModel.from_pretrained(
|
policy = PeftModel.from_pretrained(
|
||||||
policy,
|
policy, peft_pretrained_path, config=peft_config, is_trainable=True
|
||||||
peft_pretrained_path,
|
|
||||||
config=peft_config,
|
|
||||||
revision=cfg.pretrained_revision,
|
|
||||||
is_trainable=True,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -37,19 +37,13 @@ def is_image_feature(key: str) -> bool:
|
|||||||
@dataclass
|
@dataclass
|
||||||
class ConcurrencyConfig:
|
class ConcurrencyConfig:
|
||||||
"""Configuration for the concurrency of the actor and learner.
|
"""Configuration for the concurrency of the actor and learner.
|
||||||
|
|
||||||
Possible values are:
|
Possible values are:
|
||||||
- "threads": Use threads for the actor and learner.
|
- "threads": Use threads for the actor and learner.
|
||||||
- "processes": Use processes for the actor and learner.
|
- "processes": Use processes for the actor and learner.
|
||||||
|
|
||||||
``multiprocessing_context`` selects the process-wide start method when
|
|
||||||
processes are used. Set it to ``None`` to preserve Python's default or a
|
|
||||||
method already selected by the embedding application.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
actor: str = "threads"
|
actor: str = "threads"
|
||||||
learner: str = "threads"
|
learner: str = "threads"
|
||||||
multiprocessing_context: str | None = "spawn"
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -475,7 +475,6 @@ def make_groot_pre_post_processors_from_pretrained(
|
|||||||
config: GrootConfig,
|
config: GrootConfig,
|
||||||
pretrained_path: str,
|
pretrained_path: str,
|
||||||
*,
|
*,
|
||||||
revision: str | None = None,
|
|
||||||
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
|
dataset_stats: dict[str, dict[str, torch.Tensor]] | None = None,
|
||||||
dataset_meta: Any | None = None,
|
dataset_meta: Any | None = None,
|
||||||
preprocessor_overrides: dict[str, Any] | None = None,
|
preprocessor_overrides: dict[str, Any] | None = None,
|
||||||
@@ -512,7 +511,6 @@ def make_groot_pre_post_processors_from_pretrained(
|
|||||||
|
|
||||||
preprocessor, postprocessor = _load_groot_processor_pipelines(
|
preprocessor, postprocessor = _load_groot_processor_pipelines(
|
||||||
pretrained_path,
|
pretrained_path,
|
||||||
revision=revision,
|
|
||||||
preprocessor_overrides=preprocessor_overrides,
|
preprocessor_overrides=preprocessor_overrides,
|
||||||
postprocessor_overrides=postprocessor_overrides,
|
postprocessor_overrides=postprocessor_overrides,
|
||||||
preprocessor_config_filename=preprocessor_config_filename,
|
preprocessor_config_filename=preprocessor_config_filename,
|
||||||
@@ -528,7 +526,6 @@ def make_groot_pre_post_processors_from_pretrained(
|
|||||||
def _load_groot_processor_pipelines(
|
def _load_groot_processor_pipelines(
|
||||||
pretrained_path: str,
|
pretrained_path: str,
|
||||||
*,
|
*,
|
||||||
revision: str | None,
|
|
||||||
preprocessor_overrides: dict[str, Any],
|
preprocessor_overrides: dict[str, Any],
|
||||||
postprocessor_overrides: dict[str, Any],
|
postprocessor_overrides: dict[str, Any],
|
||||||
preprocessor_config_filename: str,
|
preprocessor_config_filename: str,
|
||||||
@@ -543,7 +540,6 @@ def _load_groot_processor_pipelines(
|
|||||||
preprocessor = PolicyProcessorPipeline.from_pretrained(
|
preprocessor = PolicyProcessorPipeline.from_pretrained(
|
||||||
pretrained_model_name_or_path=pretrained_path,
|
pretrained_model_name_or_path=pretrained_path,
|
||||||
config_filename=preprocessor_config_filename,
|
config_filename=preprocessor_config_filename,
|
||||||
revision=revision,
|
|
||||||
overrides=preprocessor_overrides,
|
overrides=preprocessor_overrides,
|
||||||
to_transition=batch_to_transition,
|
to_transition=batch_to_transition,
|
||||||
to_output=transition_to_batch,
|
to_output=transition_to_batch,
|
||||||
@@ -551,7 +547,6 @@ def _load_groot_processor_pipelines(
|
|||||||
postprocessor = PolicyProcessorPipeline.from_pretrained(
|
postprocessor = PolicyProcessorPipeline.from_pretrained(
|
||||||
pretrained_model_name_or_path=pretrained_path,
|
pretrained_model_name_or_path=pretrained_path,
|
||||||
config_filename=postprocessor_config_filename,
|
config_filename=postprocessor_config_filename,
|
||||||
revision=revision,
|
|
||||||
overrides=postprocessor_overrides,
|
overrides=postprocessor_overrides,
|
||||||
to_transition=policy_action_to_transition,
|
to_transition=policy_action_to_transition,
|
||||||
to_output=transition_to_policy_action,
|
to_output=transition_to_policy_action,
|
||||||
|
|||||||
@@ -43,22 +43,11 @@ from torch.distributions import Beta
|
|||||||
|
|
||||||
from lerobot.policies.pretrained import PreTrainedPolicy
|
from lerobot.policies.pretrained import PreTrainedPolicy
|
||||||
from lerobot.utils.constants import ACTION
|
from lerobot.utils.constants import ACTION
|
||||||
from lerobot.utils.import_utils import (
|
from lerobot.utils.import_utils import _scipy_available, _transformers_available, require_package
|
||||||
_peft_available,
|
|
||||||
_scipy_available,
|
|
||||||
_transformers_available,
|
|
||||||
require_package,
|
|
||||||
)
|
|
||||||
|
|
||||||
from ..rtc.modeling_rtc import RTCProcessor
|
from ..rtc.modeling_rtc import RTCProcessor
|
||||||
from .configuration_molmoact2 import MolmoAct2Config
|
from .configuration_molmoact2 import MolmoAct2Config
|
||||||
|
|
||||||
if TYPE_CHECKING or _peft_available:
|
|
||||||
from peft import LoraConfig, get_peft_model
|
|
||||||
else:
|
|
||||||
LoraConfig = None
|
|
||||||
get_peft_model = None
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -1742,11 +1731,13 @@ class MolmoAct2Policy(PreTrainedPolicy):
|
|||||||
|
|
||||||
def _build_inner_lora_config(self):
|
def _build_inner_lora_config(self):
|
||||||
require_package("peft", extra="molmoact2")
|
require_package("peft", extra="molmoact2")
|
||||||
|
from peft import LoraConfig
|
||||||
|
|
||||||
return LoraConfig(**self._get_inner_peft_targets())
|
return LoraConfig(**self._get_inner_peft_targets())
|
||||||
|
|
||||||
def _apply_lora_adapters(self) -> None:
|
def _apply_lora_adapters(self) -> None:
|
||||||
require_package("peft", extra="molmoact2")
|
require_package("peft", extra="molmoact2")
|
||||||
|
from peft import get_peft_model
|
||||||
|
|
||||||
peft_config = self._build_inner_lora_config()
|
peft_config = self._build_inner_lora_config()
|
||||||
self._validate_peft_config(peft_config)
|
self._validate_peft_config(peft_config)
|
||||||
|
|||||||
@@ -34,22 +34,14 @@ from lerobot.configs import PreTrainedConfig
|
|||||||
from lerobot.configs.train import TrainPipelineConfig
|
from lerobot.configs.train import TrainPipelineConfig
|
||||||
from lerobot.utils.device_utils import resolve_safetensors_device
|
from lerobot.utils.device_utils import resolve_safetensors_device
|
||||||
from lerobot.utils.hub import HubMixin
|
from lerobot.utils.hub import HubMixin
|
||||||
from lerobot.utils.import_utils import _peft_available, require_package
|
|
||||||
|
|
||||||
from .utils import log_model_loading_keys
|
from .utils import log_model_loading_keys
|
||||||
|
|
||||||
if TYPE_CHECKING or _peft_available:
|
T = TypeVar("T", bound="PreTrainedPolicy")
|
||||||
from peft import PEFT_TYPE_TO_CONFIG_MAPPING, PeftType, get_peft_model
|
|
||||||
else:
|
|
||||||
PEFT_TYPE_TO_CONFIG_MAPPING = None
|
|
||||||
PeftType = None
|
|
||||||
get_peft_model = None
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
|
from lerobot.datasets.dataset_metadata import LeRobotDatasetMetadata
|
||||||
|
|
||||||
T = TypeVar("T", bound="PreTrainedPolicy")
|
|
||||||
|
|
||||||
|
|
||||||
def _build_card_context(
|
def _build_card_context(
|
||||||
cfg: TrainPipelineConfig | None,
|
cfg: TrainPipelineConfig | None,
|
||||||
@@ -392,7 +384,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
peft_cli_overrides: Optional dict of CLI overrides (method_type, target_modules, r, etc.)
|
peft_cli_overrides: Optional dict of CLI overrides (method_type, target_modules, r, etc.)
|
||||||
These are merged with policy defaults to build the final config.
|
These are merged with policy defaults to build the final config.
|
||||||
"""
|
"""
|
||||||
require_package("peft", extra="peft")
|
from peft import get_peft_model
|
||||||
|
|
||||||
# If user provided a complete config, use it directly (with overrides)
|
# If user provided a complete config, use it directly (with overrides)
|
||||||
if peft_config is not None:
|
if peft_config is not None:
|
||||||
@@ -463,7 +455,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
Returns:
|
Returns:
|
||||||
Preprocessed dict with renamed keys and init_type mapped to method-specific key.
|
Preprocessed dict with renamed keys and init_type mapped to method-specific key.
|
||||||
"""
|
"""
|
||||||
require_package("peft", extra="peft")
|
from peft import PeftType
|
||||||
|
|
||||||
cli_overrides = cli_overrides.copy()
|
cli_overrides = cli_overrides.copy()
|
||||||
|
|
||||||
@@ -488,7 +480,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
|
|
||||||
def _build_peft_config(self, cli_overrides: dict):
|
def _build_peft_config(self, cli_overrides: dict):
|
||||||
"""Build a PEFT config from policy defaults and CLI overrides."""
|
"""Build a PEFT config from policy defaults and CLI overrides."""
|
||||||
require_package("peft", extra="peft")
|
from peft import PEFT_TYPE_TO_CONFIG_MAPPING, PeftType
|
||||||
|
|
||||||
# Determine PEFT method type (default to LORA)
|
# Determine PEFT method type (default to LORA)
|
||||||
method_type_str = cli_overrides.get("method_type") or "lora"
|
method_type_str = cli_overrides.get("method_type") or "lora"
|
||||||
@@ -515,7 +507,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC):
|
|||||||
|
|
||||||
def _apply_peft_cli_overrides(self, peft_config, cli_overrides: dict):
|
def _apply_peft_cli_overrides(self, peft_config, cli_overrides: dict):
|
||||||
"""Apply CLI overrides to an existing PEFT config."""
|
"""Apply CLI overrides to an existing PEFT config."""
|
||||||
require_package("peft", extra="peft")
|
from peft import PEFT_TYPE_TO_CONFIG_MAPPING, PeftType
|
||||||
|
|
||||||
# Get method type from existing config or CLI override
|
# Get method type from existing config or CLI override
|
||||||
method_type_str = cli_overrides.get("method_type")
|
method_type_str = cli_overrides.get("method_type")
|
||||||
|
|||||||
@@ -132,20 +132,10 @@ class MapDeltaActionToRobotActionStep(RobotActionProcessorStep):
|
|||||||
def transform_features(
|
def transform_features(
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
||||||
for axis in ["x", "y", "z"]:
|
for axis in ["x", "y", "z", "gripper"]:
|
||||||
features[PipelineFeatureType.ACTION].pop(f"delta_{axis}", None)
|
features[PipelineFeatureType.ACTION].pop(f"delta_{axis}", None)
|
||||||
features[PipelineFeatureType.ACTION].pop("gripper", None)
|
|
||||||
|
|
||||||
for feat in [
|
for feat in ["enabled", "target_x", "target_y", "target_z", "target_wx", "target_wy", "target_wz"]:
|
||||||
"enabled",
|
|
||||||
"target_x",
|
|
||||||
"target_y",
|
|
||||||
"target_z",
|
|
||||||
"target_wx",
|
|
||||||
"target_wy",
|
|
||||||
"target_wz",
|
|
||||||
"gripper_vel",
|
|
||||||
]:
|
|
||||||
features[PipelineFeatureType.ACTION][f"{feat}"] = PolicyFeature(
|
features[PipelineFeatureType.ACTION][f"{feat}"] = PolicyFeature(
|
||||||
type=FeatureType.ACTION, shape=(1,)
|
type=FeatureType.ACTION, shape=(1,)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -713,8 +713,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
ProcessorMigrationError: If the model requires migration to processor format.
|
ProcessorMigrationError: If the model requires migration to processor format.
|
||||||
"""
|
"""
|
||||||
model_id = str(pretrained_model_name_or_path)
|
model_id = str(pretrained_model_name_or_path)
|
||||||
model_path = Path(model_id)
|
|
||||||
is_local_source = model_path.is_dir() or model_path.is_file()
|
|
||||||
hub_download_kwargs = {
|
hub_download_kwargs = {
|
||||||
"force_download": force_download,
|
"force_download": force_download,
|
||||||
"resume_download": resume_download,
|
"resume_download": resume_download,
|
||||||
@@ -733,7 +731,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
|
|
||||||
# 3. Build steps with overrides
|
# 3. Build steps with overrides
|
||||||
steps, validated_overrides = cls._build_steps_with_overrides(
|
steps, validated_overrides = cls._build_steps_with_overrides(
|
||||||
loaded_config, overrides or {}, model_id, base_path, hub_download_kwargs, is_local_source
|
loaded_config, overrides or {}, model_id, base_path, hub_download_kwargs
|
||||||
)
|
)
|
||||||
|
|
||||||
# 4. Validate that all overrides were used
|
# 4. Validate that all overrides were used
|
||||||
@@ -923,7 +921,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
model_id: str,
|
model_id: str,
|
||||||
base_path: Path | None,
|
base_path: Path | None,
|
||||||
hub_download_kwargs: dict[str, Any],
|
hub_download_kwargs: dict[str, Any],
|
||||||
is_local_source: bool = False,
|
|
||||||
) -> tuple[list[ProcessorStep], set[str]]:
|
) -> tuple[list[ProcessorStep], set[str]]:
|
||||||
"""Build all processor steps with overrides and state loading.
|
"""Build all processor steps with overrides and state loading.
|
||||||
|
|
||||||
@@ -947,7 +944,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
3. **State Loading** (via _load_step_state):
|
3. **State Loading** (via _load_step_state):
|
||||||
- **If step has "state_file"**: Load tensor state from .safetensors
|
- **If step has "state_file"**: Load tensor state from .safetensors
|
||||||
- **Local first**: Check base_path/state_file.safetensors
|
- **Local first**: Check base_path/state_file.safetensors
|
||||||
- **Hub fallback**: Download state file if the pipeline was loaded from the Hub
|
- **Hub fallback**: Download state file if not found locally
|
||||||
- **Optional**: Only load if step has load_state_dict method
|
- **Optional**: Only load if step has load_state_dict method
|
||||||
|
|
||||||
4. **Override Tracking**:
|
4. **Override Tracking**:
|
||||||
@@ -965,7 +962,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
model_id: The model identifier (needed for Hub state file downloads)
|
model_id: The model identifier (needed for Hub state file downloads)
|
||||||
base_path: Local directory path for finding state files
|
base_path: Local directory path for finding state files
|
||||||
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
||||||
is_local_source: Whether model_id resolved to a local directory or config file.
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Tuple of (instantiated_steps_list, unused_override_keys)
|
Tuple of (instantiated_steps_list, unused_override_keys)
|
||||||
@@ -979,9 +975,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
steps, remaining_override_keys = cls._build_steps_from_config(loaded_config, overrides)
|
steps, remaining_override_keys = cls._build_steps_from_config(loaded_config, overrides)
|
||||||
|
|
||||||
for step_instance, step_entry in zip(steps, loaded_config["steps"], strict=True):
|
for step_instance, step_entry in zip(steps, loaded_config["steps"], strict=True):
|
||||||
cls._load_step_state(
|
cls._load_step_state(step_instance, step_entry, model_id, base_path, hub_download_kwargs)
|
||||||
step_instance, step_entry, model_id, base_path, hub_download_kwargs, is_local_source
|
|
||||||
)
|
|
||||||
|
|
||||||
return steps, remaining_override_keys
|
return steps, remaining_override_keys
|
||||||
|
|
||||||
@@ -1145,7 +1139,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
model_id: str,
|
model_id: str,
|
||||||
base_path: Path | None,
|
base_path: Path | None,
|
||||||
hub_download_kwargs: dict[str, Any],
|
hub_download_kwargs: dict[str, Any],
|
||||||
is_local_source: bool = False,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Load state dictionary for a processor step if available.
|
"""Load state dictionary for a processor step if available.
|
||||||
|
|
||||||
@@ -1164,7 +1157,7 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
- **Use case**: Loading from local saved model directory
|
- **Use case**: Loading from local saved model directory
|
||||||
|
|
||||||
2. **Hub download fallback**: Download state file from repository
|
2. **Hub download fallback**: Download state file from repository
|
||||||
- **When triggered**: Local file not found and the pipeline source is a Hub repo
|
- **When triggered**: Local file not found or base_path is None
|
||||||
- **Process**: Use hf_hub_download with same parameters as config
|
- **Process**: Use hf_hub_download with same parameters as config
|
||||||
- **Example**: Download "normalize_step_0.safetensors" from "user/repo"
|
- **Example**: Download "normalize_step_0.safetensors" from "user/repo"
|
||||||
- **Result**: Downloaded to local cache, path returned
|
- **Result**: Downloaded to local cache, path returned
|
||||||
@@ -1185,7 +1178,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
model_id: The model identifier (used for Hub downloads if needed)
|
model_id: The model identifier (used for Hub downloads if needed)
|
||||||
base_path: Local directory path for finding state files (None for Hub-only)
|
base_path: Local directory path for finding state files (None for Hub-only)
|
||||||
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
hub_download_kwargs: Parameters for hf_hub_download (tokens, cache, etc.)
|
||||||
is_local_source: Whether model_id resolved to a local directory or config file.
|
|
||||||
|
|
||||||
Note:
|
Note:
|
||||||
This method modifies step_instance in-place and returns None.
|
This method modifies step_instance in-place and returns None.
|
||||||
@@ -1199,12 +1191,6 @@ class DataProcessorPipeline[TInput, TOutput](HubMixin):
|
|||||||
# Try local file first
|
# Try local file first
|
||||||
if base_path and (base_path / state_filename).exists():
|
if base_path and (base_path / state_filename).exists():
|
||||||
state_path = str(base_path / state_filename)
|
state_path = str(base_path / state_filename)
|
||||||
elif is_local_source:
|
|
||||||
state_path = base_path / state_filename if base_path else Path(state_filename)
|
|
||||||
raise FileNotFoundError(
|
|
||||||
f"State file '{state_filename}' was not found for local processor pipeline "
|
|
||||||
f"'{model_id}' at '{state_path}'."
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
# Download from Hub
|
# Download from Hub
|
||||||
state_path = hf_hub_download(
|
state_path = hf_hub_download(
|
||||||
|
|||||||
@@ -91,7 +91,7 @@ from lerobot.robots import so_follower # noqa: F401
|
|||||||
from lerobot.teleoperators import gamepad, so_leader # noqa: F401
|
from lerobot.teleoperators import gamepad, so_leader # noqa: F401
|
||||||
from lerobot.teleoperators.utils import TeleopEvents
|
from lerobot.teleoperators.utils import TeleopEvents
|
||||||
from lerobot.utils.device_utils import get_safe_torch_device
|
from lerobot.utils.device_utils import get_safe_torch_device
|
||||||
from lerobot.utils.process import ProcessSignalHandler, ensure_multiprocessing_start_method
|
from lerobot.utils.process import ProcessSignalHandler
|
||||||
from lerobot.utils.random_utils import set_seed
|
from lerobot.utils.random_utils import set_seed
|
||||||
from lerobot.utils.robot_utils import precise_sleep
|
from lerobot.utils.robot_utils import precise_sleep
|
||||||
from lerobot.utils.transition import (
|
from lerobot.utils.transition import (
|
||||||
@@ -124,7 +124,9 @@ def actor_cli(cfg: TrainRLServerPipelineConfig):
|
|||||||
cfg.validate()
|
cfg.validate()
|
||||||
display_pid = False
|
display_pid = False
|
||||||
if not use_threads(cfg):
|
if not use_threads(cfg):
|
||||||
ensure_multiprocessing_start_method(cfg.policy.concurrency.multiprocessing_context)
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
|
mp.set_start_method("spawn")
|
||||||
display_pid = True
|
display_pid = True
|
||||||
|
|
||||||
# Create logs directory to ensure it exists
|
# Create logs directory to ensure it exists
|
||||||
|
|||||||
@@ -102,7 +102,7 @@ from lerobot.utils.constants import (
|
|||||||
)
|
)
|
||||||
from lerobot.utils.device_utils import get_safe_torch_device
|
from lerobot.utils.device_utils import get_safe_torch_device
|
||||||
from lerobot.utils.io_utils import load_json, write_json
|
from lerobot.utils.io_utils import load_json, write_json
|
||||||
from lerobot.utils.process import ProcessSignalHandler, ensure_multiprocessing_start_method
|
from lerobot.utils.process import ProcessSignalHandler
|
||||||
from lerobot.utils.random_utils import set_seed
|
from lerobot.utils.random_utils import set_seed
|
||||||
from lerobot.utils.utils import (
|
from lerobot.utils.utils import (
|
||||||
format_big_number,
|
format_big_number,
|
||||||
@@ -123,7 +123,9 @@ def train_cli(cfg: TrainRLServerPipelineConfig):
|
|||||||
# Fail fast with a friendly error if the optional ``hilserl`` extra is missing.
|
# Fail fast with a friendly error if the optional ``hilserl`` extra is missing.
|
||||||
require_package("grpcio", extra="hilserl", import_name="grpc")
|
require_package("grpcio", extra="hilserl", import_name="grpc")
|
||||||
if not use_threads(cfg):
|
if not use_threads(cfg):
|
||||||
ensure_multiprocessing_start_method(cfg.policy.concurrency.multiprocessing_context)
|
import torch.multiprocessing as mp
|
||||||
|
|
||||||
|
mp.set_start_method("spawn")
|
||||||
|
|
||||||
# Use the job_name from the config
|
# Use the job_name from the config
|
||||||
train(
|
train(
|
||||||
|
|||||||
@@ -323,10 +323,6 @@ class LeKiwiClient(Robot):
|
|||||||
np.ndarray: the action sent to the motors, potentially clipped.
|
np.ndarray: the action sent to the motors, potentially clipped.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Action values may be torch tensors (e.g. replayed from a dataset) or numpy
|
|
||||||
# scalars; json.dumps only serializes Python primitives, so coerce each value to a
|
|
||||||
# plain float before sending.
|
|
||||||
action = {key: float(value) for key, value in action.items()}
|
|
||||||
self.zmq_cmd_socket.send_string(json.dumps(action)) # action is in motor space
|
self.zmq_cmd_socket.send_string(json.dumps(action)) # action is in motor space
|
||||||
|
|
||||||
# TODO(Steven): Remove the np conversion when it is possible to record a non-numpy array value
|
# TODO(Steven): Remove the np conversion when it is possible to record a non-numpy array value
|
||||||
|
|||||||
@@ -46,12 +46,6 @@ class SOFollowerConfig:
|
|||||||
position_i_coefficient: int = 0
|
position_i_coefficient: int = 0
|
||||||
position_d_coefficient: int = 32
|
position_d_coefficient: int = 32
|
||||||
|
|
||||||
# Number of extra attempts when a `sync_read` of the motors fails. Feetech buses can occasionally
|
|
||||||
# return a corrupted status packet ("Incorrect status packet!"), especially when several joints move
|
|
||||||
# at once, which otherwise aborts the control loop. Retries are immediate (no sleep) and only happen on
|
|
||||||
# failure, so the steady-state read cost is unchanged.
|
|
||||||
num_read_retries: int = 2
|
|
||||||
|
|
||||||
|
|
||||||
@RobotConfig.register_subclass("so101_follower")
|
@RobotConfig.register_subclass("so101_follower")
|
||||||
@RobotConfig.register_subclass("so100_follower")
|
@RobotConfig.register_subclass("so100_follower")
|
||||||
|
|||||||
@@ -180,7 +180,7 @@ class SOFollower(Robot):
|
|||||||
def get_observation(self) -> RobotObservation:
|
def get_observation(self) -> RobotObservation:
|
||||||
# Read arm position
|
# Read arm position
|
||||||
start = time.perf_counter()
|
start = time.perf_counter()
|
||||||
obs_dict = self.bus.sync_read("Present_Position", num_retry=self.config.num_read_retries)
|
obs_dict = self.bus.sync_read("Present_Position")
|
||||||
obs_dict = {f"{motor}.pos": val for motor, val in obs_dict.items()}
|
obs_dict = {f"{motor}.pos": val for motor, val in obs_dict.items()}
|
||||||
dt_ms = (time.perf_counter() - start) * 1e3
|
dt_ms = (time.perf_counter() - start) * 1e3
|
||||||
logger.debug(f"{self} read state: {dt_ms:.1f}ms")
|
logger.debug(f"{self} read state: {dt_ms:.1f}ms")
|
||||||
@@ -221,7 +221,7 @@ class SOFollower(Robot):
|
|||||||
# Cap goal position when too far away from present position.
|
# Cap goal position when too far away from present position.
|
||||||
# /!\ Slower fps expected due to reading from the follower.
|
# /!\ Slower fps expected due to reading from the follower.
|
||||||
if self.config.max_relative_target is not None:
|
if self.config.max_relative_target is not None:
|
||||||
present_pos = self.bus.sync_read("Present_Position", num_retry=self.config.num_read_retries)
|
present_pos = self.bus.sync_read("Present_Position")
|
||||||
goal_present_pos = {key: (g_pos, present_pos[key]) for key, g_pos in goal_pos.items()}
|
goal_present_pos = {key: (g_pos, present_pos[key]) for key, g_pos in goal_pos.items()}
|
||||||
goal_pos = ensure_safe_goal_position(goal_present_pos, self.config.max_relative_target)
|
goal_pos = ensure_safe_goal_position(goal_present_pos, self.config.max_relative_target)
|
||||||
|
|
||||||
|
|||||||
@@ -326,17 +326,8 @@ class RolloutConfig:
|
|||||||
|
|
||||||
policy_path = parser.get_path_arg("policy")
|
policy_path = parser.get_path_arg("policy")
|
||||||
if policy_path:
|
if policy_path:
|
||||||
yaml_overrides = parser.get_yaml_overrides("policy")
|
cli_overrides = parser.get_cli_overrides("policy")
|
||||||
cli_overrides = parser.get_cli_overrides("policy") or []
|
self.policy = PreTrainedConfig.from_pretrained(policy_path, cli_overrides=cli_overrides)
|
||||||
policy_overrides = yaml_overrides + cli_overrides
|
|
||||||
pretrained_revision = parser.parse_arg("pretrained_revision", cli_overrides)
|
|
||||||
if pretrained_revision is None:
|
|
||||||
pretrained_revision = parser.parse_arg("pretrained_revision", yaml_overrides)
|
|
||||||
self.policy = PreTrainedConfig.from_pretrained(
|
|
||||||
policy_path,
|
|
||||||
revision=pretrained_revision,
|
|
||||||
cli_overrides=policy_overrides,
|
|
||||||
)
|
|
||||||
self.policy.pretrained_path = policy_path
|
self.policy.pretrained_path = policy_path
|
||||||
if self.policy is None:
|
if self.policy is None:
|
||||||
raise ValueError("--policy.path is required for rollout")
|
raise ValueError("--policy.path is required for rollout")
|
||||||
|
|||||||
@@ -24,11 +24,10 @@ from __future__ import annotations
|
|||||||
import logging
|
import logging
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from threading import Event
|
from threading import Event
|
||||||
from typing import TYPE_CHECKING
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from lerobot.configs import FeatureType, PreTrainedConfig
|
from lerobot.configs import FeatureType
|
||||||
from lerobot.datasets import (
|
from lerobot.datasets import (
|
||||||
LeRobotDataset,
|
LeRobotDataset,
|
||||||
aggregate_pipeline_dataset_features,
|
aggregate_pipeline_dataset_features,
|
||||||
@@ -48,7 +47,6 @@ from lerobot.processor.relative_action_processor import RelativeActionsProcessor
|
|||||||
from lerobot.robots import make_robot_from_config
|
from lerobot.robots import make_robot_from_config
|
||||||
from lerobot.teleoperators import Teleoperator, make_teleoperator_from_config
|
from lerobot.teleoperators import Teleoperator, make_teleoperator_from_config
|
||||||
from lerobot.utils.feature_utils import combine_feature_dicts, hw_to_dataset_features
|
from lerobot.utils.feature_utils import combine_feature_dicts, hw_to_dataset_features
|
||||||
from lerobot.utils.import_utils import _peft_available, require_package
|
|
||||||
|
|
||||||
from .configs import BaseStrategyConfig, DAggerStrategyConfig, RolloutConfig
|
from .configs import BaseStrategyConfig, DAggerStrategyConfig, RolloutConfig
|
||||||
from .inference import (
|
from .inference import (
|
||||||
@@ -59,12 +57,6 @@ from .inference import (
|
|||||||
)
|
)
|
||||||
from .robot_wrapper import ThreadSafeRobot
|
from .robot_wrapper import ThreadSafeRobot
|
||||||
|
|
||||||
if TYPE_CHECKING or _peft_available:
|
|
||||||
from peft import PeftConfig, PeftModel
|
|
||||||
else:
|
|
||||||
PeftConfig = None
|
|
||||||
PeftModel = None
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
@@ -167,35 +159,6 @@ class RolloutContext:
|
|||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def _load_pretrained_policy(policy_config: PreTrainedConfig) -> PreTrainedPolicy:
|
|
||||||
"""Load policy weights, keeping adapter and base-model revisions independent."""
|
|
||||||
pretrained_revision = policy_config.pretrained_revision
|
|
||||||
policy_class = get_policy_class(policy_config.type)
|
|
||||||
|
|
||||||
if not policy_config.use_peft:
|
|
||||||
return policy_class.from_pretrained(
|
|
||||||
policy_config.pretrained_path,
|
|
||||||
config=policy_config,
|
|
||||||
revision=pretrained_revision,
|
|
||||||
)
|
|
||||||
|
|
||||||
require_package("peft", extra="peft")
|
|
||||||
|
|
||||||
peft_path = policy_config.pretrained_path
|
|
||||||
peft_config = PeftConfig.from_pretrained(peft_path, revision=pretrained_revision)
|
|
||||||
policy = policy_class.from_pretrained(
|
|
||||||
pretrained_name_or_path=peft_config.base_model_name_or_path,
|
|
||||||
config=policy_config,
|
|
||||||
revision=peft_config.revision,
|
|
||||||
)
|
|
||||||
return PeftModel.from_pretrained(
|
|
||||||
policy,
|
|
||||||
peft_path,
|
|
||||||
config=peft_config,
|
|
||||||
revision=pretrained_revision,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def build_rollout_context(
|
def build_rollout_context(
|
||||||
cfg: RolloutConfig,
|
cfg: RolloutConfig,
|
||||||
shutdown_event: Event,
|
shutdown_event: Event,
|
||||||
@@ -213,6 +176,7 @@ def build_rollout_context(
|
|||||||
# --- 1. Policy (heavy I/O, but no hardware yet) -------------------
|
# --- 1. Policy (heavy I/O, but no hardware yet) -------------------
|
||||||
logger.info("Loading policy from '%s'...", cfg.policy.pretrained_path)
|
logger.info("Loading policy from '%s'...", cfg.policy.pretrained_path)
|
||||||
policy_config = cfg.policy
|
policy_config = cfg.policy
|
||||||
|
policy_class = get_policy_class(policy_config.type)
|
||||||
|
|
||||||
if hasattr(policy_config, "compile_model"):
|
if hasattr(policy_config, "compile_model"):
|
||||||
policy_config.compile_model = cfg.use_torch_compile
|
policy_config.compile_model = cfg.use_torch_compile
|
||||||
@@ -223,7 +187,17 @@ def build_rollout_context(
|
|||||||
"Please use `cpu` or `cuda` backend."
|
"Please use `cpu` or `cuda` backend."
|
||||||
)
|
)
|
||||||
|
|
||||||
policy = _load_pretrained_policy(policy_config)
|
if policy_config.use_peft:
|
||||||
|
from peft import PeftConfig, PeftModel
|
||||||
|
|
||||||
|
peft_path = policy_config.pretrained_path
|
||||||
|
peft_config = PeftConfig.from_pretrained(peft_path)
|
||||||
|
policy = policy_class.from_pretrained(
|
||||||
|
pretrained_name_or_path=peft_config.base_model_name_or_path, config=policy_config
|
||||||
|
)
|
||||||
|
policy = PeftModel.from_pretrained(policy, peft_path, config=peft_config)
|
||||||
|
else:
|
||||||
|
policy = policy_class.from_pretrained(policy_config.pretrained_path, config=policy_config)
|
||||||
|
|
||||||
if is_rtc:
|
if is_rtc:
|
||||||
policy.config.rtc_config = cfg.inference.rtc
|
policy.config.rtc_config = cfg.inference.rtc
|
||||||
@@ -302,22 +276,12 @@ def build_rollout_context(
|
|||||||
# ``observation_features`` values are either a tuple (camera shape) or the
|
# ``observation_features`` values are either a tuple (camera shape) or the
|
||||||
# ``float`` type itself used as a sentinel for scalar motor features —
|
# ``float`` type itself used as a sentinel for scalar motor features —
|
||||||
# see ``dict[str, type | tuple]`` annotation on ``Robot.observation_features``.
|
# see ``dict[str, type | tuple]`` annotation on ``Robot.observation_features``.
|
||||||
# Keep cameras (tuple) plus both joint-position (.pos) and base-velocity (.vel)
|
|
||||||
# scalar state features. LeKiwi's observation.state is 9-dim (6 arm .pos +
|
|
||||||
# x/y/theta.vel) and the policy was trained/normalized on all 9; the old .pos-only
|
|
||||||
# filter fed a 6-dim state into a 9-dim normalizer → RuntimeError (size 6 vs 9).
|
|
||||||
# Pure-arm robots have no .vel state keys, so this is a no-op for them.
|
|
||||||
observation_features_hw = {
|
observation_features_hw = {
|
||||||
k: v
|
k: v
|
||||||
for k, v in all_obs_features.items()
|
for k, v in all_obs_features.items()
|
||||||
if isinstance(v, tuple) or (v is float and k.endswith((".pos", ".vel")))
|
if isinstance(v, tuple) or (v is float and k.endswith(".pos"))
|
||||||
}
|
}
|
||||||
# Keep both joint-position (.pos) and base-velocity (.vel) action features so
|
action_features_hw = {k: v for k, v in robot.action_features.items() if k.endswith(".pos")}
|
||||||
# mobile manipulators command the base too (e.g. LeKiwi: 6 arm .pos +
|
|
||||||
# x/y/theta.vel = 9-dim action). Pure-arm robots have no .vel keys, so this is
|
|
||||||
# a no-op for them. Without the .vel keys the base velocities are silently
|
|
||||||
# dropped from dataset_features[ACTION]/ordered_action_keys and the base never moves.
|
|
||||||
action_features_hw = {k: v for k, v in robot.action_features.items() if k.endswith((".pos", ".vel"))}
|
|
||||||
|
|
||||||
# The action side is always needed: sync inference reads action names from
|
# The action side is always needed: sync inference reads action names from
|
||||||
# ``dataset_features[ACTION]`` to map policy tensors back to robot actions.
|
# ``dataset_features[ACTION]`` to map policy tensors back to robot actions.
|
||||||
@@ -428,7 +392,6 @@ def build_rollout_context(
|
|||||||
preprocessor, postprocessor = make_pre_post_processors(
|
preprocessor, postprocessor = make_pre_post_processors(
|
||||||
policy_cfg=policy_config,
|
policy_cfg=policy_config,
|
||||||
pretrained_path=cfg.policy.pretrained_path,
|
pretrained_path=cfg.policy.pretrained_path,
|
||||||
pretrained_revision=policy_config.pretrained_revision,
|
|
||||||
dataset_stats=dataset_stats,
|
dataset_stats=dataset_stats,
|
||||||
preprocessor_overrides={
|
preprocessor_overrides={
|
||||||
"device_processor": {"device": cfg.device},
|
"device_processor": {"device": cfg.device},
|
||||||
|
|||||||
@@ -36,7 +36,6 @@ python src/lerobot/scripts/augment_dataset_quantile_stats.py \
|
|||||||
import argparse
|
import argparse
|
||||||
import concurrent.futures
|
import concurrent.futures
|
||||||
import logging
|
import logging
|
||||||
import os
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -53,7 +52,6 @@ from lerobot.datasets import (
|
|||||||
get_feature_stats,
|
get_feature_stats,
|
||||||
write_stats,
|
write_stats,
|
||||||
)
|
)
|
||||||
from lerobot.datasets.compute_stats import sample_indices
|
|
||||||
from lerobot.utils.utils import init_logging
|
from lerobot.utils.utils import init_logging
|
||||||
|
|
||||||
|
|
||||||
@@ -79,14 +77,12 @@ def has_quantile_stats(stats: dict[str, dict] | None, quantile_list_keys: list[s
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampling: bool = True) -> dict:
|
def process_single_episode(dataset: LeRobotDataset, episode_idx: int) -> dict:
|
||||||
"""Process a single episode and return its statistics.
|
"""Process a single episode and return its statistics.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
dataset: The LeRobot dataset
|
dataset: The LeRobot dataset
|
||||||
episode_idx: Index of the episode to process
|
episode_idx: Index of the episode to process
|
||||||
use_sampling: If True, sub-sample image/video frames per episode to bound
|
|
||||||
memory. If False, use every frame (exact, higher memory).
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dictionary containing episode statistics
|
Dictionary containing episode statistics
|
||||||
@@ -96,31 +92,16 @@ def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampli
|
|||||||
start_idx = dataset.meta.episodes[episode_idx]["dataset_from_index"]
|
start_idx = dataset.meta.episodes[episode_idx]["dataset_from_index"]
|
||||||
end_idx = dataset.meta.episodes[episode_idx]["dataset_to_index"]
|
end_idx = dataset.meta.episodes[episode_idx]["dataset_to_index"]
|
||||||
|
|
||||||
episode_len = end_idx - start_idx
|
|
||||||
|
|
||||||
# Images/video are the memory hog, so sub-sample those frames per episode;
|
|
||||||
# numeric columns are cheap, so read them in full (exact).
|
|
||||||
image_keys = [k for k in dataset.features if dataset.features[k]["dtype"] in ("image", "video")]
|
|
||||||
numeric_keys = [
|
|
||||||
k for k in dataset.features if dataset.features[k]["dtype"] not in ("image", "video", "string")
|
|
||||||
]
|
|
||||||
|
|
||||||
collected_data: dict[str, list] = {}
|
collected_data: dict[str, list] = {}
|
||||||
|
for idx in range(start_idx, end_idx):
|
||||||
|
item = dataset[idx]
|
||||||
|
for key, value in item.items():
|
||||||
|
if key not in dataset.features:
|
||||||
|
continue
|
||||||
|
|
||||||
# Numeric features: every frame, read directly from the underlying table.
|
if key not in collected_data:
|
||||||
if numeric_keys:
|
collected_data[key] = []
|
||||||
numeric_cols = dataset.hf_dataset.select_columns(numeric_keys)[start_idx:end_idx]
|
collected_data[key].append(value)
|
||||||
for key in numeric_keys:
|
|
||||||
collected_data[key] = [torch.as_tensor(v) for v in numeric_cols[key]]
|
|
||||||
|
|
||||||
# Image/video features: decode only a sampled subset of frames.
|
|
||||||
if image_keys:
|
|
||||||
sampled_offsets = sample_indices(episode_len) if use_sampling else list(range(episode_len))
|
|
||||||
for offset in sampled_offsets:
|
|
||||||
item = dataset[start_idx + offset]
|
|
||||||
for key in image_keys:
|
|
||||||
if key in item:
|
|
||||||
collected_data.setdefault(key, []).append(item[key])
|
|
||||||
|
|
||||||
ep_stats = {}
|
ep_stats = {}
|
||||||
for key, data_list in collected_data.items():
|
for key, data_list in collected_data.items():
|
||||||
@@ -150,13 +131,11 @@ def process_single_episode(dataset: LeRobotDataset, episode_idx: int, use_sampli
|
|||||||
return ep_stats
|
return ep_stats
|
||||||
|
|
||||||
|
|
||||||
def compute_quantile_stats_for_dataset(dataset: LeRobotDataset, use_sampling: bool = True) -> dict[str, dict]:
|
def compute_quantile_stats_for_dataset(dataset: LeRobotDataset) -> dict[str, dict]:
|
||||||
"""Compute quantile statistics for all episodes in the dataset.
|
"""Compute quantile statistics for all episodes in the dataset.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
dataset: The LeRobot dataset to compute statistics for
|
dataset: The LeRobot dataset to compute statistics for
|
||||||
use_sampling: If True, sub-sample image/video frames per episode to bound
|
|
||||||
memory. If False, use every frame (exact, higher memory).
|
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dictionary containing aggregated statistics with quantiles
|
Dictionary containing aggregated statistics with quantiles
|
||||||
@@ -174,15 +153,15 @@ def compute_quantile_stats_for_dataset(dataset: LeRobotDataset, use_sampling: bo
|
|||||||
if has_videos:
|
if has_videos:
|
||||||
logging.info("Dataset contains video keys - using sequential processing for thread safety")
|
logging.info("Dataset contains video keys - using sequential processing for thread safety")
|
||||||
for episode_idx in tqdm(range(dataset.num_episodes), desc="Processing episodes"):
|
for episode_idx in tqdm(range(dataset.num_episodes), desc="Processing episodes"):
|
||||||
ep_stats = process_single_episode(dataset, episode_idx, use_sampling)
|
ep_stats = process_single_episode(dataset, episode_idx)
|
||||||
episode_stats_list.append(ep_stats)
|
episode_stats_list.append(ep_stats)
|
||||||
else:
|
else:
|
||||||
logging.info("Dataset has no video keys - using parallel processing for better performance")
|
logging.info("Dataset has no video keys - using parallel processing for better performance")
|
||||||
max_workers = min(dataset.num_episodes, int(os.environ.get("LEROBOT_STATS_MAX_WORKERS", 16)))
|
max_workers = min(dataset.num_episodes, 16)
|
||||||
|
|
||||||
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
|
with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||||
future_to_episode = {
|
future_to_episode = {
|
||||||
executor.submit(process_single_episode, dataset, episode_idx, use_sampling): episode_idx
|
executor.submit(process_single_episode, dataset, episode_idx): episode_idx
|
||||||
for episode_idx in range(dataset.num_episodes)
|
for episode_idx in range(dataset.num_episodes)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -209,7 +188,6 @@ def augment_dataset_with_quantile_stats(
|
|||||||
repo_id: str,
|
repo_id: str,
|
||||||
root: str | Path | None = None,
|
root: str | Path | None = None,
|
||||||
overwrite: bool = False,
|
overwrite: bool = False,
|
||||||
use_sampling: bool = True,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Augment a dataset with quantile statistics if they are missing.
|
"""Augment a dataset with quantile statistics if they are missing.
|
||||||
|
|
||||||
@@ -217,8 +195,6 @@ def augment_dataset_with_quantile_stats(
|
|||||||
repo_id: Repository ID of the dataset
|
repo_id: Repository ID of the dataset
|
||||||
root: Local root directory for the dataset
|
root: Local root directory for the dataset
|
||||||
overwrite: Overwrite existing quantile statistics if they already exist
|
overwrite: Overwrite existing quantile statistics if they already exist
|
||||||
use_sampling: If True, sub-sample image/video frames per episode to bound
|
|
||||||
memory. If False, use every frame (exact, higher memory).
|
|
||||||
"""
|
"""
|
||||||
logging.info(f"Loading dataset: {repo_id}")
|
logging.info(f"Loading dataset: {repo_id}")
|
||||||
dataset = LeRobotDataset(
|
dataset = LeRobotDataset(
|
||||||
@@ -232,7 +208,7 @@ def augment_dataset_with_quantile_stats(
|
|||||||
|
|
||||||
logging.info("Dataset does not contain quantile statistics. Computing them now...")
|
logging.info("Dataset does not contain quantile statistics. Computing them now...")
|
||||||
|
|
||||||
new_stats = compute_quantile_stats_for_dataset(dataset, use_sampling=use_sampling)
|
new_stats = compute_quantile_stats_for_dataset(dataset)
|
||||||
|
|
||||||
logging.info("Updating dataset metadata with new quantile statistics")
|
logging.info("Updating dataset metadata with new quantile statistics")
|
||||||
dataset.meta.stats = new_stats
|
dataset.meta.stats = new_stats
|
||||||
@@ -272,14 +248,6 @@ def main():
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Overwrite existing quantile statistics if they already exist",
|
help="Overwrite existing quantile statistics if they already exist",
|
||||||
)
|
)
|
||||||
parser.add_argument(
|
|
||||||
"--no-sampling",
|
|
||||||
action="store_true",
|
|
||||||
help=(
|
|
||||||
"Compute stats over every frame (exact, higher memory). By default, "
|
|
||||||
"image/video frames are sub-sampled per episode to bound memory."
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
root = Path(args.root) if args.root else None
|
root = Path(args.root) if args.root else None
|
||||||
@@ -290,7 +258,6 @@ def main():
|
|||||||
repo_id=args.repo_id,
|
repo_id=args.repo_id,
|
||||||
root=root,
|
root=root,
|
||||||
overwrite=args.overwrite,
|
overwrite=args.overwrite,
|
||||||
use_sampling=not args.no_sampling,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -94,8 +94,6 @@ from lerobot.datasets.video_utils import concatenate_video_files, get_video_dura
|
|||||||
from lerobot.utils.constants import HF_LEROBOT_HOME
|
from lerobot.utils.constants import HF_LEROBOT_HOME
|
||||||
from lerobot.utils.utils import flatten_dict, init_logging
|
from lerobot.utils.utils import flatten_dict, init_logging
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
V21 = "v2.1"
|
V21 = "v2.1"
|
||||||
V30 = "v3.0"
|
V30 = "v3.0"
|
||||||
|
|
||||||
@@ -478,11 +476,11 @@ def convert_dataset(
|
|||||||
# First check if the dataset already has a v3.0 version
|
# First check if the dataset already has a v3.0 version
|
||||||
if root is None and not force_conversion:
|
if root is None and not force_conversion:
|
||||||
try:
|
try:
|
||||||
logger.info("Trying to download v3.0 version of the dataset from the hub...")
|
print("Trying to download v3.0 version of the dataset from the hub...")
|
||||||
snapshot_download(repo_id, repo_type="dataset", revision=V30, local_dir=HF_LEROBOT_HOME / repo_id)
|
snapshot_download(repo_id, repo_type="dataset", revision=V30, local_dir=HF_LEROBOT_HOME / repo_id)
|
||||||
return
|
return
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.info("Dataset does not have an uploaded v3.0 version. Continuing with conversion.")
|
print("Dataset does not have an uploaded v3.0 version. Continuing with conversion.")
|
||||||
|
|
||||||
# Set root based on whether local dataset path is provided
|
# Set root based on whether local dataset path is provided
|
||||||
use_local_dataset = False
|
use_local_dataset = False
|
||||||
@@ -490,7 +488,7 @@ def convert_dataset(
|
|||||||
if root.exists():
|
if root.exists():
|
||||||
validate_local_dataset_version(root)
|
validate_local_dataset_version(root)
|
||||||
use_local_dataset = True
|
use_local_dataset = True
|
||||||
logger.info(f"Using local dataset at {root}")
|
print(f"Using local dataset at {root}")
|
||||||
|
|
||||||
old_root = root.parent / f"{root.name}_old"
|
old_root = root.parent / f"{root.name}_old"
|
||||||
new_root = root.parent / f"{root.name}_v30"
|
new_root = root.parent / f"{root.name}_v30"
|
||||||
@@ -525,7 +523,7 @@ def convert_dataset(
|
|||||||
try:
|
try:
|
||||||
hub_api.delete_tag(repo_id, tag=CODEBASE_VERSION, repo_type="dataset")
|
hub_api.delete_tag(repo_id, tag=CODEBASE_VERSION, repo_type="dataset")
|
||||||
except (HTTPError, RevisionNotFoundError) as e:
|
except (HTTPError, RevisionNotFoundError) as e:
|
||||||
logger.warning(f"tag={CODEBASE_VERSION} probably doesn't exist. Skipping exception ({e})")
|
print(f"tag={CODEBASE_VERSION} probably doesn't exist. Skipping exception ({e})")
|
||||||
pass
|
pass
|
||||||
hub_api.delete_files(
|
hub_api.delete_files(
|
||||||
delete_patterns=["data/chunk*/episode_*", "meta/*.jsonl", "videos/chunk*"],
|
delete_patterns=["data/chunk*/episode_*", "meta/*.jsonl", "videos/chunk*"],
|
||||||
|
|||||||
@@ -154,14 +154,14 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
|||||||
repo_id = cfg.new_repo_id or cfg.repo_id
|
repo_id = cfg.new_repo_id or cfg.repo_id
|
||||||
commit_message = cfg.push_commit_message or "Add steerable annotations (lerobot-annotate)"
|
commit_message = cfg.push_commit_message or "Add steerable annotations (lerobot-annotate)"
|
||||||
api = HfApi()
|
api = HfApi()
|
||||||
logger.info(f"[lerobot-annotate] creating/locating dataset repo {repo_id}...")
|
print(f"[lerobot-annotate] creating/locating dataset repo {repo_id}...", flush=True)
|
||||||
api.create_repo(
|
api.create_repo(
|
||||||
repo_id=repo_id,
|
repo_id=repo_id,
|
||||||
repo_type="dataset",
|
repo_type="dataset",
|
||||||
private=cfg.push_private,
|
private=cfg.push_private,
|
||||||
exist_ok=True,
|
exist_ok=True,
|
||||||
)
|
)
|
||||||
logger.info(f"[lerobot-annotate] uploading {root} -> {repo_id}...")
|
print(f"[lerobot-annotate] uploading {root} -> {repo_id}...", flush=True)
|
||||||
commit_info = api.upload_folder(
|
commit_info = api.upload_folder(
|
||||||
folder_path=str(root),
|
folder_path=str(root),
|
||||||
repo_id=repo_id,
|
repo_id=repo_id,
|
||||||
@@ -172,7 +172,7 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
|||||||
# at the source dataset; a fresh card is generated below instead.
|
# at the source dataset; a fresh card is generated below instead.
|
||||||
ignore_patterns=[".annotate_staging/**", "**/.DS_Store", "README.md"],
|
ignore_patterns=[".annotate_staging/**", "**/.DS_Store", "README.md"],
|
||||||
)
|
)
|
||||||
logger.info(f"[lerobot-annotate] uploaded to https://huggingface.co/datasets/{repo_id}")
|
print(f"[lerobot-annotate] uploaded to https://huggingface.co/datasets/{repo_id}", flush=True)
|
||||||
|
|
||||||
dataset_info = load_info(root)
|
dataset_info = load_info(root)
|
||||||
card = create_lerobot_dataset_card(dataset_info=dataset_info, license="apache-2.0", repo_id=repo_id)
|
card = create_lerobot_dataset_card(dataset_info=dataset_info, license="apache-2.0", repo_id=repo_id)
|
||||||
@@ -200,13 +200,14 @@ def _push_to_hub(root: Path, cfg: AnnotationPipelineConfig) -> None:
|
|||||||
with suppress(RevisionNotFoundError):
|
with suppress(RevisionNotFoundError):
|
||||||
api.delete_tag(repo_id, tag=version_tag, repo_type="dataset")
|
api.delete_tag(repo_id, tag=version_tag, repo_type="dataset")
|
||||||
api.create_tag(**tag_kwargs)
|
api.create_tag(**tag_kwargs)
|
||||||
logger.info(f"[lerobot-annotate] tagged {repo_id} as {version_tag}")
|
print(f"[lerobot-annotate] tagged {repo_id} as {version_tag}", flush=True)
|
||||||
except Exception as exc: # noqa: BLE001
|
except Exception as exc: # noqa: BLE001
|
||||||
logger.warning(
|
print(
|
||||||
f"[lerobot-annotate] WARNING: could not create tag {version_tag!r} on {repo_id}: {exc}. "
|
f"[lerobot-annotate] WARNING: could not create tag {version_tag!r} on {repo_id}: {exc}. "
|
||||||
"Dataset is uploaded but ``LeRobotDataset`` won't be able to load it until it's tagged. "
|
"Dataset is uploaded but ``LeRobotDataset`` won't be able to load it until it's tagged. "
|
||||||
"Run: from huggingface_hub import HfApi; "
|
"Run: from huggingface_hub import HfApi; "
|
||||||
f"HfApi().create_tag({repo_id!r}, tag={version_tag!r}, repo_type='dataset', exist_ok=True)"
|
f"HfApi().create_tag({repo_id!r}, tag={version_tag!r}, repo_type='dataset', exist_ok=True)",
|
||||||
|
flush=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -89,8 +89,6 @@ from lerobot.datasets import LeRobotDataset
|
|||||||
from lerobot.utils.constants import ACTION, DONE, OBS_STATE, REWARD, SUCCESS
|
from lerobot.utils.constants import ACTION, DONE, OBS_STATE, REWARD, SUCCESS
|
||||||
from lerobot.utils.utils import init_logging
|
from lerobot.utils.utils import init_logging
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
DEFAULT_FOXGLOVE_PORT = 8765
|
DEFAULT_FOXGLOVE_PORT = 8765
|
||||||
DEFAULT_RERUN_PORT = 9090
|
DEFAULT_RERUN_PORT = 9090
|
||||||
|
|
||||||
@@ -301,7 +299,7 @@ def visualize_dataset(
|
|||||||
while True:
|
while True:
|
||||||
time.sleep(1)
|
time.sleep(1)
|
||||||
except KeyboardInterrupt:
|
except KeyboardInterrupt:
|
||||||
logger.info("Ctrl-C received. Exiting.")
|
print("Ctrl-C received. Exiting.")
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
|
|||||||
@@ -62,7 +62,7 @@ from dataclasses import asdict
|
|||||||
from functools import partial
|
from functools import partial
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from pprint import pformat
|
from pprint import pformat
|
||||||
from typing import TYPE_CHECKING, Any, TypedDict
|
from typing import Any, TypedDict
|
||||||
|
|
||||||
import einops
|
import einops
|
||||||
import gymnasium as gym
|
import gymnasium as gym
|
||||||
@@ -87,7 +87,7 @@ from lerobot.processor import PolicyProcessorPipeline
|
|||||||
from lerobot.types import PolicyAction
|
from lerobot.types import PolicyAction
|
||||||
from lerobot.utils.constants import ACTION, DONE, OBS_IMAGE, OBS_IMAGES, OBS_STR, REWARD
|
from lerobot.utils.constants import ACTION, DONE, OBS_IMAGE, OBS_IMAGES, OBS_STR, REWARD
|
||||||
from lerobot.utils.device_utils import get_safe_torch_device
|
from lerobot.utils.device_utils import get_safe_torch_device
|
||||||
from lerobot.utils.import_utils import _peft_available, register_third_party_plugins, require_package
|
from lerobot.utils.import_utils import register_third_party_plugins
|
||||||
from lerobot.utils.io_utils import write_video
|
from lerobot.utils.io_utils import write_video
|
||||||
from lerobot.utils.random_utils import set_seed
|
from lerobot.utils.random_utils import set_seed
|
||||||
from lerobot.utils.utils import (
|
from lerobot.utils.utils import (
|
||||||
@@ -95,14 +95,6 @@ from lerobot.utils.utils import (
|
|||||||
inside_slurm,
|
inside_slurm,
|
||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING or _peft_available:
|
|
||||||
from peft import PeftModel
|
|
||||||
else:
|
|
||||||
PeftModel = None
|
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
def _env_features_to_dataset_features(env_features: dict) -> dict:
|
def _env_features_to_dataset_features(env_features: dict) -> dict:
|
||||||
"""Convert EnvConfig.features to the dict format expected by LeRobotDataset.create()."""
|
"""Convert EnvConfig.features to the dict format expected by LeRobotDataset.create()."""
|
||||||
@@ -452,16 +444,15 @@ def eval_policy(
|
|||||||
exc = ValueError(
|
exc = ValueError(
|
||||||
f"Policy of type 'PreTrainedPolicy' is expected, but type '{type(policy)}' was provided."
|
f"Policy of type 'PreTrainedPolicy' is expected, but type '{type(policy)}' was provided."
|
||||||
)
|
)
|
||||||
if not _peft_available:
|
try:
|
||||||
raise exc
|
from peft import PeftModel
|
||||||
require_package("peft", extra="peft")
|
|
||||||
if not isinstance(policy, PeftModel):
|
if not isinstance(policy, PeftModel):
|
||||||
raise exc
|
raise exc
|
||||||
|
except ImportError:
|
||||||
|
raise exc from None
|
||||||
|
|
||||||
start = time.time()
|
start = time.time()
|
||||||
# Preserve the mode for direct callers. eval_policy_all scopes the mode
|
|
||||||
# around all tasks so parallel evaluations cannot race with each other.
|
|
||||||
was_training = policy.training
|
|
||||||
policy.eval()
|
policy.eval()
|
||||||
|
|
||||||
# Determine how many batched rollouts we need to get n_episodes. Note that if n_episodes is not evenly
|
# Determine how many batched rollouts we need to get n_episodes. Note that if n_episodes is not evenly
|
||||||
@@ -564,7 +555,7 @@ def eval_policy(
|
|||||||
if seeds:
|
if seeds:
|
||||||
all_seeds.extend(seeds)
|
all_seeds.extend(seeds)
|
||||||
else:
|
else:
|
||||||
all_seeds.extend([None] * env.num_envs)
|
all_seeds.append(None)
|
||||||
|
|
||||||
# FIXME: episode_data is either None or it doesn't exist
|
# FIXME: episode_data is either None or it doesn't exist
|
||||||
if return_episode_data:
|
if return_episode_data:
|
||||||
@@ -683,8 +674,6 @@ def eval_policy(
|
|||||||
if save_predicted_video:
|
if save_predicted_video:
|
||||||
info["predicted_video_paths"] = predicted_video_paths
|
info["predicted_video_paths"] = predicted_video_paths
|
||||||
|
|
||||||
policy.train(was_training)
|
|
||||||
|
|
||||||
return info
|
return info
|
||||||
|
|
||||||
|
|
||||||
@@ -802,13 +791,13 @@ def eval_main(cfg: EvalPipelineConfig):
|
|||||||
recording_repo_id=cfg.eval.recording_repo_id,
|
recording_repo_id=cfg.eval.recording_repo_id,
|
||||||
recording_private=cfg.eval.recording_private,
|
recording_private=cfg.eval.recording_private,
|
||||||
)
|
)
|
||||||
logger.info("Overall Aggregated Metrics:")
|
print("Overall Aggregated Metrics:")
|
||||||
logger.info(info["overall"])
|
print(info["overall"])
|
||||||
|
|
||||||
# Print per-suite stats
|
# Print per-suite stats
|
||||||
for task_group, task_group_info in info.items():
|
for task_group, task_group_info in info.items():
|
||||||
logger.info(f"\nAggregated Metrics for {task_group}:")
|
print(f"\nAggregated Metrics for {task_group}:")
|
||||||
logger.info(task_group_info)
|
print(task_group_info)
|
||||||
# Close all vec envs
|
# Close all vec envs
|
||||||
close_envs(envs)
|
close_envs(envs)
|
||||||
|
|
||||||
@@ -1021,12 +1010,6 @@ def eval_policy_all(
|
|||||||
recording_private=recording_private,
|
recording_private=recording_private,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Set the shared policy's mode before launching any workers. Restoring it
|
|
||||||
# inside individual tasks would let one task enable training mode while
|
|
||||||
# another task is still evaluating.
|
|
||||||
was_training = policy.training
|
|
||||||
policy.eval()
|
|
||||||
try:
|
|
||||||
if max_parallel_tasks <= 1:
|
if max_parallel_tasks <= 1:
|
||||||
prefetch_thread: threading.Thread | None = None
|
prefetch_thread: threading.Thread | None = None
|
||||||
for i, (task_group, task_id, env) in enumerate(tasks):
|
for i, (task_group, task_id, env) in enumerate(tasks):
|
||||||
@@ -1061,8 +1044,6 @@ def eval_policy_all(
|
|||||||
per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics})
|
per_task_infos.append({"task_group": tg, "task_id": tid, "metrics": metrics})
|
||||||
finally:
|
finally:
|
||||||
env.close()
|
env.close()
|
||||||
finally:
|
|
||||||
policy.train(was_training)
|
|
||||||
|
|
||||||
# compute aggregated metrics helper (robust to lists/scalars)
|
# compute aggregated metrics helper (robust to lists/scalars)
|
||||||
def _agg_from_list(xs):
|
def _agg_from_list(xs):
|
||||||
|
|||||||
@@ -40,7 +40,6 @@ from PIL import Image
|
|||||||
from lerobot.cameras import ColorMode
|
from lerobot.cameras import ColorMode
|
||||||
from lerobot.cameras.opencv import OpenCVCamera, OpenCVCameraConfig
|
from lerobot.cameras.opencv import OpenCVCamera, OpenCVCameraConfig
|
||||||
from lerobot.cameras.realsense import RealSenseCamera, RealSenseCameraConfig
|
from lerobot.cameras.realsense import RealSenseCamera, RealSenseCameraConfig
|
||||||
from lerobot.utils.utils import init_logging
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -286,8 +285,6 @@ def save_images_from_all_cameras(
|
|||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
init_logging()
|
|
||||||
|
|
||||||
parser = argparse.ArgumentParser(
|
parser = argparse.ArgumentParser(
|
||||||
description="Unified camera utility script for listing cameras and capturing images."
|
description="Unified camera utility script for listing cameras and capturing images."
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -453,11 +453,9 @@ def record(
|
|||||||
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
|
encoder_queue_maxsize=cfg.dataset.encoder_queue_maxsize,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Connect the teleoperator before the robot so the robot isn't left idle (and possibly
|
robot.connect()
|
||||||
# tripping a firmware watchdog) during teleop init. Matches lerobot_teleoperate.py.
|
|
||||||
if teleop is not None:
|
if teleop is not None:
|
||||||
teleop.connect()
|
teleop.connect()
|
||||||
robot.connect()
|
|
||||||
|
|
||||||
listener, events = init_keyboard_listener()
|
listener, events = init_keyboard_listener()
|
||||||
|
|
||||||
|
|||||||
@@ -61,7 +61,6 @@ from lerobot.robots import ( # noqa: F401
|
|||||||
earthrover_mini_plus,
|
earthrover_mini_plus,
|
||||||
hope_jr,
|
hope_jr,
|
||||||
koch_follower,
|
koch_follower,
|
||||||
lekiwi,
|
|
||||||
make_robot_from_config,
|
make_robot_from_config,
|
||||||
omx_follower,
|
omx_follower,
|
||||||
openarm_follower,
|
openarm_follower,
|
||||||
|
|||||||
@@ -165,7 +165,6 @@ from lerobot.robots import ( # noqa: F401
|
|||||||
earthrover_mini_plus,
|
earthrover_mini_plus,
|
||||||
hope_jr,
|
hope_jr,
|
||||||
koch_follower,
|
koch_follower,
|
||||||
lekiwi,
|
|
||||||
omx_follower,
|
omx_follower,
|
||||||
openarm_follower,
|
openarm_follower,
|
||||||
reachy2,
|
reachy2,
|
||||||
|
|||||||
@@ -22,8 +22,7 @@ import dataclasses
|
|||||||
import logging
|
import logging
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
from collections.abc import Iterator
|
from contextlib import nullcontext
|
||||||
from contextlib import contextmanager, nullcontext
|
|
||||||
from pprint import pformat
|
from pprint import pformat
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any
|
||||||
|
|
||||||
@@ -58,7 +57,7 @@ from lerobot.optim.factory import make_optimizer_and_scheduler
|
|||||||
from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors
|
from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors
|
||||||
from lerobot.rewards import make_reward_pre_post_processors
|
from lerobot.rewards import make_reward_pre_post_processors
|
||||||
from lerobot.utils.collate import lerobot_collate_fn
|
from lerobot.utils.collate import lerobot_collate_fn
|
||||||
from lerobot.utils.import_utils import _peft_available, register_third_party_plugins, require_package
|
from lerobot.utils.import_utils import register_third_party_plugins
|
||||||
from lerobot.utils.logging_utils import AverageMeter, MetricsTracker
|
from lerobot.utils.logging_utils import AverageMeter, MetricsTracker
|
||||||
from lerobot.utils.random_utils import set_seed
|
from lerobot.utils.random_utils import set_seed
|
||||||
from lerobot.utils.utils import (
|
from lerobot.utils.utils import (
|
||||||
@@ -69,38 +68,9 @@ from lerobot.utils.utils import (
|
|||||||
inside_slurm,
|
inside_slurm,
|
||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING or _peft_available:
|
|
||||||
from peft import PeftModel
|
|
||||||
else:
|
|
||||||
PeftModel = None
|
|
||||||
|
|
||||||
from .lerobot_eval import eval_policy_all
|
from .lerobot_eval import eval_policy_all
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
|
||||||
def _make_eval_envs(cfg: TrainPipelineConfig) -> Iterator[dict[str, dict[int, Any]]]:
|
|
||||||
"""Create evaluation environments for one run and always dispose of them."""
|
|
||||||
envs = make_env(
|
|
||||||
cfg.env,
|
|
||||||
n_envs=cfg.eval.batch_size,
|
|
||||||
use_async_envs=cfg.eval.use_async_envs,
|
|
||||||
)
|
|
||||||
try:
|
|
||||||
yield envs
|
|
||||||
finally:
|
|
||||||
close_envs(envs)
|
|
||||||
|
|
||||||
|
|
||||||
def _dataloader_worker_kwargs(cfg: TrainPipelineConfig) -> dict[str, Any]:
|
|
||||||
"""Return worker-only DataLoader options, disabling them for single-process loading."""
|
|
||||||
workers_enabled = cfg.num_workers > 0
|
|
||||||
return {
|
|
||||||
"prefetch_factor": cfg.prefetch_factor if workers_enabled else None,
|
|
||||||
"persistent_workers": cfg.persistent_workers and workers_enabled,
|
|
||||||
"multiprocessing_context": cfg.dataloader_multiprocessing_context if workers_enabled else None,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def update_policy(
|
def update_policy(
|
||||||
train_metrics: MetricsTracker,
|
train_metrics: MetricsTracker,
|
||||||
policy: PreTrainedPolicy,
|
policy: PreTrainedPolicy,
|
||||||
@@ -227,6 +197,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
if cfg.job.is_remote:
|
if cfg.job.is_remote:
|
||||||
return submit_to_hf(cfg)
|
return submit_to_hf(cfg)
|
||||||
|
|
||||||
|
from lerobot.utils.import_utils import require_package
|
||||||
|
|
||||||
require_package("accelerate", extra="training")
|
require_package("accelerate", extra="training")
|
||||||
from accelerate import Accelerator
|
from accelerate import Accelerator
|
||||||
from accelerate.utils import DistributedDataParallelKwargs, DistributedType
|
from accelerate.utils import DistributedDataParallelKwargs, DistributedType
|
||||||
@@ -295,6 +267,14 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
if not is_main_process:
|
if not is_main_process:
|
||||||
dataset, eval_dataset = make_train_eval_datasets(cfg)
|
dataset, eval_dataset = make_train_eval_datasets(cfg)
|
||||||
|
|
||||||
|
# Create environment used for evaluating checkpoints during training on simulation data.
|
||||||
|
# On real-world data, no need to create an environment as evaluations are done outside train.py,
|
||||||
|
# using the eval.py instead, with gym_dora environment and dora-rs.
|
||||||
|
eval_env = None
|
||||||
|
if cfg.env_eval_freq > 0 and cfg.env is not None and is_main_process:
|
||||||
|
logging.info("Creating env")
|
||||||
|
eval_env = make_env(cfg.env, n_envs=cfg.eval.batch_size, use_async_envs=cfg.eval.use_async_envs)
|
||||||
|
|
||||||
if cfg.is_reward_model_training:
|
if cfg.is_reward_model_training:
|
||||||
if is_main_process:
|
if is_main_process:
|
||||||
logging.info("Creating reward model")
|
logging.info("Creating reward model")
|
||||||
@@ -322,7 +302,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
if cfg.peft is not None:
|
if cfg.peft is not None:
|
||||||
if cfg.is_reward_model_training:
|
if cfg.is_reward_model_training:
|
||||||
raise ValueError("PEFT is only supported for policy training. ")
|
raise ValueError("PEFT is only supported for policy training. ")
|
||||||
require_package("peft", extra="peft")
|
from peft import PeftModel
|
||||||
|
|
||||||
if isinstance(policy, PeftModel):
|
if isinstance(policy, PeftModel):
|
||||||
logging.info("PEFT adapter already loaded from checkpoint, skipping wrap_with_peft.")
|
logging.info("PEFT adapter already loaded from checkpoint, skipping wrap_with_peft.")
|
||||||
@@ -493,7 +473,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
pin_memory=device.type == "cuda",
|
pin_memory=device.type == "cuda",
|
||||||
drop_last=False,
|
drop_last=False,
|
||||||
collate_fn=collate_fn,
|
collate_fn=collate_fn,
|
||||||
**_dataloader_worker_kwargs(cfg),
|
prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None,
|
||||||
|
persistent_workers=cfg.persistent_workers and cfg.num_workers > 0,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Build eval dataloader if a held-out split exists
|
# Build eval dataloader if a held-out split exists
|
||||||
@@ -519,7 +500,8 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
pin_memory=device.type == "cuda",
|
pin_memory=device.type == "cuda",
|
||||||
drop_last=False,
|
drop_last=False,
|
||||||
collate_fn=eval_collate_fn,
|
collate_fn=eval_collate_fn,
|
||||||
**_dataloader_worker_kwargs(cfg),
|
prefetch_factor=cfg.prefetch_factor if cfg.num_workers > 0 else None,
|
||||||
|
persistent_workers=cfg.persistent_workers and cfg.num_workers > 0,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Prepare everything with accelerator
|
# Prepare everything with accelerator
|
||||||
@@ -702,7 +684,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
if is_main_process:
|
if is_main_process:
|
||||||
step_id = get_step_identifier(step, cfg.steps)
|
step_id = get_step_identifier(step, cfg.steps)
|
||||||
logging.info(f"Eval policy at step {step}")
|
logging.info(f"Eval policy at step {step}")
|
||||||
with _make_eval_envs(cfg) as eval_env, torch.no_grad(), accelerator.autocast():
|
with torch.no_grad(), accelerator.autocast():
|
||||||
eval_info = eval_policy_all(
|
eval_info = eval_policy_all(
|
||||||
envs=eval_env, # dict[suite][task_id] -> vec_env
|
envs=eval_env, # dict[suite][task_id] -> vec_env
|
||||||
policy=accelerator.unwrap_model(policy),
|
policy=accelerator.unwrap_model(policy),
|
||||||
@@ -750,6 +732,9 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None):
|
|||||||
if is_main_process:
|
if is_main_process:
|
||||||
progbar.close()
|
progbar.close()
|
||||||
|
|
||||||
|
if eval_env:
|
||||||
|
close_envs(eval_env)
|
||||||
|
|
||||||
is_fsdp = accelerator.distributed_type == DistributedType.FSDP
|
is_fsdp = accelerator.distributed_type == DistributedType.FSDP
|
||||||
model_state_dict = accelerator.get_state_dict(policy) if is_fsdp else None
|
model_state_dict = accelerator.get_state_dict(policy) if is_fsdp else None
|
||||||
if is_main_process:
|
if is_main_process:
|
||||||
|
|||||||
@@ -45,7 +45,6 @@ lerobot-train-tokenizer \
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import json
|
import json
|
||||||
import logging
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
@@ -64,9 +63,6 @@ else:
|
|||||||
from lerobot.configs import NormalizationMode, parser
|
from lerobot.configs import NormalizationMode, parser
|
||||||
from lerobot.datasets import LeRobotDataset
|
from lerobot.datasets import LeRobotDataset
|
||||||
from lerobot.utils.constants import ACTION, OBS_STATE
|
from lerobot.utils.constants import ACTION, OBS_STATE
|
||||||
from lerobot.utils.utils import init_logging
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -278,8 +274,11 @@ def process_episode(args):
|
|||||||
|
|
||||||
return action_chunks
|
return action_chunks
|
||||||
|
|
||||||
except Exception:
|
except Exception as e:
|
||||||
logger.exception("Error processing episode %s", ep_idx)
|
print(f"Error processing episode {ep_idx}: {e}")
|
||||||
|
import traceback
|
||||||
|
|
||||||
|
traceback.print_exc()
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
@@ -301,10 +300,10 @@ def train_fast_tokenizer(
|
|||||||
Returns:
|
Returns:
|
||||||
Trained FAST tokenizer
|
Trained FAST tokenizer
|
||||||
"""
|
"""
|
||||||
logger.info(f"Training FAST tokenizer on {len(action_chunks)} action chunks...")
|
print(f"Training FAST tokenizer on {len(action_chunks)} action chunks...")
|
||||||
logger.info(f"Action chunk shape: {action_chunks.shape}")
|
print(f"Action chunk shape: {action_chunks.shape}")
|
||||||
logger.info(f"Vocab size: {vocab_size}")
|
print(f"Vocab size: {vocab_size}")
|
||||||
logger.info(f"DCT scale: {scale}")
|
print(f"DCT scale: {scale}")
|
||||||
|
|
||||||
# download the tokenizer source code (not pretrained weights)
|
# download the tokenizer source code (not pretrained weights)
|
||||||
# we'll train a new tokenizer on our own data
|
# we'll train a new tokenizer on our own data
|
||||||
@@ -315,7 +314,7 @@ def train_fast_tokenizer(
|
|||||||
|
|
||||||
# train the new tokenizer on our action data using .fit()
|
# train the new tokenizer on our action data using .fit()
|
||||||
# this trains the BPE tokenizer on DCT coefficients
|
# this trains the BPE tokenizer on DCT coefficients
|
||||||
logger.info("Training new tokenizer (this may take a few minutes)...")
|
print("Training new tokenizer (this may take a few minutes)...")
|
||||||
tokenizer = base_tokenizer.fit(
|
tokenizer = base_tokenizer.fit(
|
||||||
action_data_list,
|
action_data_list,
|
||||||
scale=scale,
|
scale=scale,
|
||||||
@@ -323,21 +322,21 @@ def train_fast_tokenizer(
|
|||||||
time_horizon=action_chunks.shape[1], # action_horizon
|
time_horizon=action_chunks.shape[1], # action_horizon
|
||||||
action_dim=action_chunks.shape[2], # encoded dimensions
|
action_dim=action_chunks.shape[2], # encoded dimensions
|
||||||
)
|
)
|
||||||
logger.info("✓ Tokenizer training complete!")
|
print("✓ Tokenizer training complete!")
|
||||||
|
|
||||||
# validate it works
|
# validate it works
|
||||||
sample_chunk = action_chunks[0]
|
sample_chunk = action_chunks[0]
|
||||||
encoded = tokenizer(sample_chunk[None])[0]
|
encoded = tokenizer(sample_chunk[None])[0]
|
||||||
if isinstance(encoded, list):
|
if isinstance(encoded, list):
|
||||||
encoded = np.array(encoded)
|
encoded = np.array(encoded)
|
||||||
logger.info(f"Sample encoding: {len(encoded)} tokens for chunk shape {sample_chunk.shape}")
|
print(f"Sample encoding: {len(encoded)} tokens for chunk shape {sample_chunk.shape}")
|
||||||
|
|
||||||
return tokenizer
|
return tokenizer
|
||||||
|
|
||||||
|
|
||||||
def compute_compression_stats(tokenizer, action_chunks: np.ndarray):
|
def compute_compression_stats(tokenizer, action_chunks: np.ndarray):
|
||||||
"""Compute compression statistics."""
|
"""Compute compression statistics."""
|
||||||
logger.info("\nComputing compression statistics...")
|
print("\nComputing compression statistics...")
|
||||||
|
|
||||||
# sample for stats (use max 1000 chunks for speed)
|
# sample for stats (use max 1000 chunks for speed)
|
||||||
sample_size = min(1000, len(action_chunks))
|
sample_size = min(1000, len(action_chunks))
|
||||||
@@ -367,12 +366,12 @@ def compute_compression_stats(tokenizer, action_chunks: np.ndarray):
|
|||||||
"max_token_length": float(np.max(token_lengths)),
|
"max_token_length": float(np.max(token_lengths)),
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.info("Compression Statistics:")
|
print("Compression Statistics:")
|
||||||
logger.info(f" Average compression ratio: {stats['compression_ratio']:.2f}x")
|
print(f" Average compression ratio: {stats['compression_ratio']:.2f}x")
|
||||||
logger.info(f" Mean token length: {stats['mean_token_length']:.1f}")
|
print(f" Mean token length: {stats['mean_token_length']:.1f}")
|
||||||
logger.info(f" P99 token length: {stats['p99_token_length']:.0f}")
|
print(f" P99 token length: {stats['p99_token_length']:.0f}")
|
||||||
logger.info(f" Min token length: {stats['min_token_length']:.0f}")
|
print(f" Min token length: {stats['min_token_length']:.0f}")
|
||||||
logger.info(f" Max token length: {stats['max_token_length']:.0f}")
|
print(f" Max token length: {stats['max_token_length']:.0f}")
|
||||||
|
|
||||||
return stats
|
return stats
|
||||||
|
|
||||||
@@ -386,9 +385,9 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
cfg: TokenizerTrainingConfig dataclass with all configuration parameters
|
cfg: TokenizerTrainingConfig dataclass with all configuration parameters
|
||||||
"""
|
"""
|
||||||
# load dataset
|
# load dataset
|
||||||
logger.info(f"Loading dataset: {cfg.repo_id}")
|
print(f"Loading dataset: {cfg.repo_id}")
|
||||||
dataset = LeRobotDataset(repo_id=cfg.repo_id, root=cfg.root)
|
dataset = LeRobotDataset(repo_id=cfg.repo_id, root=cfg.root)
|
||||||
logger.info(f"Dataset loaded: {dataset.num_episodes} episodes, {dataset.num_frames} frames")
|
print(f"Dataset loaded: {dataset.num_episodes} episodes, {dataset.num_frames} frames")
|
||||||
|
|
||||||
# parse normalization mode
|
# parse normalization mode
|
||||||
try:
|
try:
|
||||||
@@ -398,7 +397,7 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
f"Invalid normalization_mode: {cfg.normalization_mode}. "
|
f"Invalid normalization_mode: {cfg.normalization_mode}. "
|
||||||
f"Must be one of: {', '.join([m.value for m in NormalizationMode])}"
|
f"Must be one of: {', '.join([m.value for m in NormalizationMode])}"
|
||||||
) from err
|
) from err
|
||||||
logger.info(f"Normalization mode: {norm_mode.value}")
|
print(f"Normalization mode: {norm_mode.value}")
|
||||||
|
|
||||||
# parse encoded dimensions
|
# parse encoded dimensions
|
||||||
encoded_dim_ranges = []
|
encoded_dim_ranges = []
|
||||||
@@ -407,38 +406,38 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
encoded_dim_ranges.append((start, end))
|
encoded_dim_ranges.append((start, end))
|
||||||
|
|
||||||
total_encoded_dims = sum(end - start for start, end in encoded_dim_ranges)
|
total_encoded_dims = sum(end - start for start, end in encoded_dim_ranges)
|
||||||
logger.info(f"Encoding {total_encoded_dims} dimensions: {cfg.encoded_dims}")
|
print(f"Encoding {total_encoded_dims} dimensions: {cfg.encoded_dims}")
|
||||||
|
|
||||||
# parse relative dimensions
|
# parse relative dimensions
|
||||||
relative_dim_list = None
|
relative_dim_list = None
|
||||||
if cfg.relative_dims is not None and cfg.relative_dims.strip():
|
if cfg.relative_dims is not None and cfg.relative_dims.strip():
|
||||||
relative_dim_list = [int(d.strip()) for d in cfg.relative_dims.split(",")]
|
relative_dim_list = [int(d.strip()) for d in cfg.relative_dims.split(",")]
|
||||||
logger.info(f"Relative dimensions: {relative_dim_list}")
|
print(f"Relative dimensions: {relative_dim_list}")
|
||||||
else:
|
else:
|
||||||
logger.info("No relative dimensions specified")
|
print("No relative dimensions specified")
|
||||||
|
|
||||||
logger.info(f"Use relative transform: {cfg.use_relative_transform}")
|
print(f"Use relative transform: {cfg.use_relative_transform}")
|
||||||
if cfg.use_relative_transform and (relative_dim_list is None or len(relative_dim_list) == 0):
|
if cfg.use_relative_transform and (relative_dim_list is None or len(relative_dim_list) == 0):
|
||||||
logger.warning(
|
print(
|
||||||
"Warning: use_relative_transform=True but no relative_dims specified. "
|
"Warning: use_relative_transform=True but no relative_dims specified. "
|
||||||
"No relative transform will be applied."
|
"No relative transform will be applied."
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info(f"Action horizon: {cfg.action_horizon}")
|
print(f"Action horizon: {cfg.action_horizon}")
|
||||||
logger.info(f"State key: {cfg.state_key}")
|
print(f"State key: {cfg.state_key}")
|
||||||
|
|
||||||
# determine episodes to process
|
# determine episodes to process
|
||||||
num_episodes = dataset.num_episodes
|
num_episodes = dataset.num_episodes
|
||||||
if cfg.max_episodes is not None:
|
if cfg.max_episodes is not None:
|
||||||
num_episodes = min(cfg.max_episodes, num_episodes)
|
num_episodes = min(cfg.max_episodes, num_episodes)
|
||||||
|
|
||||||
logger.info(f"Processing {num_episodes} episodes...")
|
print(f"Processing {num_episodes} episodes...")
|
||||||
|
|
||||||
# process episodes sequentially (to avoid pickling issues with dataset)
|
# process episodes sequentially (to avoid pickling issues with dataset)
|
||||||
all_chunks = []
|
all_chunks = []
|
||||||
for ep_idx in range(num_episodes):
|
for ep_idx in range(num_episodes):
|
||||||
if ep_idx % 10 == 0:
|
if ep_idx % 10 == 0:
|
||||||
logger.info(f" Processing episode {ep_idx}/{num_episodes}...")
|
print(f" Processing episode {ep_idx}/{num_episodes}...")
|
||||||
|
|
||||||
chunks = process_episode(
|
chunks = process_episode(
|
||||||
(
|
(
|
||||||
@@ -456,19 +455,19 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
|
|
||||||
# concatenate all chunks
|
# concatenate all chunks
|
||||||
all_chunks = np.concatenate(all_chunks, axis=0)
|
all_chunks = np.concatenate(all_chunks, axis=0)
|
||||||
logger.info(f"Collected {len(all_chunks)} action chunks")
|
print(f"Collected {len(all_chunks)} action chunks")
|
||||||
|
|
||||||
# extract only encoded dimensions FIRST (before normalization)
|
# extract only encoded dimensions FIRST (before normalization)
|
||||||
encoded_chunks = []
|
encoded_chunks = []
|
||||||
for start, end in encoded_dim_ranges:
|
for start, end in encoded_dim_ranges:
|
||||||
encoded_chunks.append(all_chunks[:, :, start:end])
|
encoded_chunks.append(all_chunks[:, :, start:end])
|
||||||
encoded_chunks = np.concatenate(encoded_chunks, axis=-1) # [N, H, D_encoded]
|
encoded_chunks = np.concatenate(encoded_chunks, axis=-1) # [N, H, D_encoded]
|
||||||
logger.info(f"Extracted {encoded_chunks.shape[-1]} encoded dimensions")
|
print(f"Extracted {encoded_chunks.shape[-1]} encoded dimensions")
|
||||||
|
|
||||||
# apply normalization to encoded dimensions
|
# apply normalization to encoded dimensions
|
||||||
logger.info("\nBefore normalization - overall stats:")
|
print("\nBefore normalization - overall stats:")
|
||||||
logger.info(f" Min: {np.min(encoded_chunks):.4f}, Max: {np.max(encoded_chunks):.4f}")
|
print(f" Min: {np.min(encoded_chunks):.4f}, Max: {np.max(encoded_chunks):.4f}")
|
||||||
logger.info(f" Mean: {np.mean(encoded_chunks):.4f}, Std: {np.std(encoded_chunks):.4f}")
|
print(f" Mean: {np.mean(encoded_chunks):.4f}, Std: {np.std(encoded_chunks):.4f}")
|
||||||
|
|
||||||
# get normalization stats from dataset
|
# get normalization stats from dataset
|
||||||
norm_stats = dataset.meta.stats
|
norm_stats = dataset.meta.stats
|
||||||
@@ -490,9 +489,9 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
encoded_stats[stat_name] = stat_array[encoded_dim_indices]
|
encoded_stats[stat_name] = stat_array[encoded_dim_indices]
|
||||||
|
|
||||||
if encoded_stats:
|
if encoded_stats:
|
||||||
logger.info(f"\nNormalization stats for encoded dimensions (mode: {norm_mode.value}):")
|
print(f"\nNormalization stats for encoded dimensions (mode: {norm_mode.value}):")
|
||||||
for stat_name, stat_values in encoded_stats.items():
|
for stat_name, stat_values in encoded_stats.items():
|
||||||
logger.info(
|
print(
|
||||||
f" {stat_name}: shape={stat_values.shape}, "
|
f" {stat_name}: shape={stat_values.shape}, "
|
||||||
f"range=[{np.min(stat_values):.4f}, {np.max(stat_values):.4f}]"
|
f"range=[{np.min(stat_values):.4f}, {np.max(stat_values):.4f}]"
|
||||||
)
|
)
|
||||||
@@ -500,27 +499,27 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
# apply normalization based on mode
|
# apply normalization based on mode
|
||||||
try:
|
try:
|
||||||
encoded_chunks = apply_normalization(encoded_chunks, encoded_stats, norm_mode, eps=1e-8)
|
encoded_chunks = apply_normalization(encoded_chunks, encoded_stats, norm_mode, eps=1e-8)
|
||||||
logger.info(f"\nApplied {norm_mode.value} normalization")
|
print(f"\nApplied {norm_mode.value} normalization")
|
||||||
except ValueError as e:
|
except ValueError as e:
|
||||||
logger.warning(f"Warning: {e}. Using raw actions without normalization.")
|
print(f"Warning: {e}. Using raw actions without normalization.")
|
||||||
|
|
||||||
logger.info("\nAfter normalization - overall stats:")
|
print("\nAfter normalization - overall stats:")
|
||||||
logger.info(f" Min: {np.min(encoded_chunks):.4f}, Max: {np.max(encoded_chunks):.4f}")
|
print(f" Min: {np.min(encoded_chunks):.4f}, Max: {np.max(encoded_chunks):.4f}")
|
||||||
logger.info(f" Mean: {np.mean(encoded_chunks):.4f}, Std: {np.std(encoded_chunks):.4f}")
|
print(f" Mean: {np.mean(encoded_chunks):.4f}, Std: {np.std(encoded_chunks):.4f}")
|
||||||
|
|
||||||
logger.info("\nPer-dimension stats (after normalization):")
|
print("\nPer-dimension stats (after normalization):")
|
||||||
for d in range(encoded_chunks.shape[-1]):
|
for d in range(encoded_chunks.shape[-1]):
|
||||||
dim_data = encoded_chunks[:, :, d]
|
dim_data = encoded_chunks[:, :, d]
|
||||||
logger.info(
|
print(
|
||||||
f" Dim {d}: min={np.min(dim_data):7.4f}, max={np.max(dim_data):7.4f}, "
|
f" Dim {d}: min={np.min(dim_data):7.4f}, max={np.max(dim_data):7.4f}, "
|
||||||
f"mean={np.mean(dim_data):7.4f}, std={np.std(dim_data):7.4f}"
|
f"mean={np.mean(dim_data):7.4f}, std={np.std(dim_data):7.4f}"
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
logger.warning("Warning: Could not extract stats for encoded dimensions, using raw actions")
|
print("Warning: Could not extract stats for encoded dimensions, using raw actions")
|
||||||
else:
|
else:
|
||||||
logger.warning("Warning: No normalization stats found in dataset, using raw actions")
|
print("Warning: No normalization stats found in dataset, using raw actions")
|
||||||
|
|
||||||
logger.info(f"Encoded chunks shape: {encoded_chunks.shape}")
|
print(f"Encoded chunks shape: {encoded_chunks.shape}")
|
||||||
|
|
||||||
# train FAST tokenizer
|
# train FAST tokenizer
|
||||||
tokenizer = train_fast_tokenizer(
|
tokenizer = train_fast_tokenizer(
|
||||||
@@ -562,8 +561,8 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
with open(output_path / "metadata.json", "w") as f:
|
with open(output_path / "metadata.json", "w") as f:
|
||||||
json.dump(metadata, f, indent=2)
|
json.dump(metadata, f, indent=2)
|
||||||
|
|
||||||
logger.info(f"\nSaved FAST tokenizer to {output_path}")
|
print(f"\nSaved FAST tokenizer to {output_path}")
|
||||||
logger.info(f"Metadata: {json.dumps(metadata, indent=2)}")
|
print(f"Metadata: {json.dumps(metadata, indent=2)}")
|
||||||
|
|
||||||
# push to Hugging Face Hub if requested
|
# push to Hugging Face Hub if requested
|
||||||
if cfg.push_to_hub:
|
if cfg.push_to_hub:
|
||||||
@@ -571,10 +570,10 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
hub_repo_id = cfg.hub_repo_id
|
hub_repo_id = cfg.hub_repo_id
|
||||||
if hub_repo_id is None:
|
if hub_repo_id is None:
|
||||||
hub_repo_id = output_path.name
|
hub_repo_id = output_path.name
|
||||||
logger.info(f"\nNo hub_repo_id provided, using: {hub_repo_id}")
|
print(f"\nNo hub_repo_id provided, using: {hub_repo_id}")
|
||||||
|
|
||||||
logger.info(f"\nPushing tokenizer to Hugging Face Hub: {hub_repo_id}")
|
print(f"\nPushing tokenizer to Hugging Face Hub: {hub_repo_id}")
|
||||||
logger.info(f" Private: {cfg.hub_private}")
|
print(f" Private: {cfg.hub_private}")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# use the tokenizer's push_to_hub method
|
# use the tokenizer's push_to_hub method
|
||||||
@@ -594,15 +593,14 @@ def train_tokenizer(cfg: TokenizerTrainingConfig):
|
|||||||
commit_message="Upload tokenizer metadata",
|
commit_message="Upload tokenizer metadata",
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.info(f"Successfully pushed tokenizer to: https://huggingface.co/{hub_repo_id}")
|
print(f"Successfully pushed tokenizer to: https://huggingface.co/{hub_repo_id}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error pushing to hub: {e}")
|
print(f"Error pushing to hub: {e}")
|
||||||
logger.error(" Make sure you're logged in with `huggingface-cli login`")
|
print(" Make sure you're logged in with `huggingface-cli login`")
|
||||||
|
|
||||||
|
|
||||||
def main():
|
def main():
|
||||||
"""CLI entry point that parses arguments and runs the tokenizer training."""
|
"""CLI entry point that parses arguments and runs the tokenizer training."""
|
||||||
init_logging()
|
|
||||||
train_tokenizer()
|
train_tokenizer()
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -29,12 +29,6 @@ class SOLeaderConfig:
|
|||||||
# Whether to use degrees for angles
|
# Whether to use degrees for angles
|
||||||
use_degrees: bool = True
|
use_degrees: bool = True
|
||||||
|
|
||||||
# Number of extra attempts when a `sync_read` of the motors fails. Feetech buses can occasionally
|
|
||||||
# return a corrupted status packet ("Incorrect status packet!"), especially when several joints move
|
|
||||||
# at once, which otherwise aborts the teleoperation loop. Retries are immediate (no sleep) and only
|
|
||||||
# happen on failure, so the steady-state read cost is unchanged.
|
|
||||||
num_read_retries: int = 2
|
|
||||||
|
|
||||||
|
|
||||||
@TeleoperatorConfig.register_subclass("so101_leader")
|
@TeleoperatorConfig.register_subclass("so101_leader")
|
||||||
@TeleoperatorConfig.register_subclass("so100_leader")
|
@TeleoperatorConfig.register_subclass("so100_leader")
|
||||||
|
|||||||
@@ -145,7 +145,7 @@ class SOLeader(Teleoperator):
|
|||||||
@check_if_not_connected
|
@check_if_not_connected
|
||||||
def get_action(self) -> dict[str, float]:
|
def get_action(self) -> dict[str, float]:
|
||||||
start = time.perf_counter()
|
start = time.perf_counter()
|
||||||
action = self.bus.sync_read("Present_Position", num_retry=self.config.num_read_retries)
|
action = self.bus.sync_read("Present_Position")
|
||||||
action = {f"{motor}.pos": val for motor, val in action.items()}
|
action = {f"{motor}.pos": val for motor, val in action.items()}
|
||||||
dt_ms = (time.perf_counter() - start) * 1e3
|
dt_ms = (time.perf_counter() - start) * 1e3
|
||||||
logger.debug(f"{self} read action: {dt_ms:.1f}ms")
|
logger.debug(f"{self} read action: {dt_ms:.1f}ms")
|
||||||
|
|||||||
@@ -13,34 +13,18 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
from .transforms import (
|
from .transforms import (
|
||||||
CoarseDropout,
|
|
||||||
GammaCorrection,
|
|
||||||
GaussianNoise,
|
|
||||||
GaussianPatchBrightness,
|
|
||||||
ImageTransformConfig,
|
ImageTransformConfig,
|
||||||
ImageTransforms,
|
ImageTransforms,
|
||||||
ImageTransformsConfig,
|
ImageTransformsConfig,
|
||||||
JPEGCompression,
|
|
||||||
MotionBlur,
|
|
||||||
PlanckianJitter,
|
|
||||||
RandomShadow,
|
|
||||||
RandomSubsetApply,
|
RandomSubsetApply,
|
||||||
SharpnessJitter,
|
SharpnessJitter,
|
||||||
make_transform_from_config,
|
make_transform_from_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"CoarseDropout",
|
|
||||||
"GammaCorrection",
|
|
||||||
"GaussianNoise",
|
|
||||||
"GaussianPatchBrightness",
|
|
||||||
"ImageTransformConfig",
|
"ImageTransformConfig",
|
||||||
"ImageTransforms",
|
"ImageTransforms",
|
||||||
"ImageTransformsConfig",
|
"ImageTransformsConfig",
|
||||||
"JPEGCompression",
|
|
||||||
"MotionBlur",
|
|
||||||
"PlanckianJitter",
|
|
||||||
"RandomShadow",
|
|
||||||
"RandomSubsetApply",
|
"RandomSubsetApply",
|
||||||
"SharpnessJitter",
|
"SharpnessJitter",
|
||||||
"make_transform_from_config",
|
"make_transform_from_config",
|
||||||
|
|||||||
@@ -14,13 +14,11 @@
|
|||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
import collections
|
import collections
|
||||||
import math
|
|
||||||
from collections.abc import Callable, Sequence
|
from collections.abc import Callable, Sequence
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torchvision.io import decode_image, encode_jpeg
|
|
||||||
from torchvision.transforms import v2
|
from torchvision.transforms import v2
|
||||||
from torchvision.transforms.v2 import (
|
from torchvision.transforms.v2 import (
|
||||||
Transform,
|
Transform,
|
||||||
@@ -43,7 +41,7 @@ class RandomSubsetApply(Transform):
|
|||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
transforms: Sequence[Callable[..., Any]],
|
transforms: Sequence[Callable],
|
||||||
p: list[float] | None = None,
|
p: list[float] | None = None,
|
||||||
n_subset: int | None = None,
|
n_subset: int | None = None,
|
||||||
random_order: bool = False,
|
random_order: bool = False,
|
||||||
@@ -52,7 +50,7 @@ class RandomSubsetApply(Transform):
|
|||||||
if not isinstance(transforms, Sequence):
|
if not isinstance(transforms, Sequence):
|
||||||
raise TypeError("Argument transforms should be a sequence of callables")
|
raise TypeError("Argument transforms should be a sequence of callables")
|
||||||
if p is None:
|
if p is None:
|
||||||
p = [1.0] * len(transforms)
|
p = [1] * len(transforms)
|
||||||
elif len(p) != len(transforms):
|
elif len(p) != len(transforms):
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Length of p doesn't match the number of transforms: {len(p)} != {len(transforms)}"
|
f"Length of p doesn't match the number of transforms: {len(p)} != {len(transforms)}"
|
||||||
@@ -71,7 +69,7 @@ class RandomSubsetApply(Transform):
|
|||||||
self.n_subset = n_subset
|
self.n_subset = n_subset
|
||||||
self.random_order = random_order
|
self.random_order = random_order
|
||||||
|
|
||||||
self.selected_transforms: list[Callable[..., Any]] = []
|
self.selected_transforms = None
|
||||||
|
|
||||||
def forward(self, *inputs: Any) -> Any:
|
def forward(self, *inputs: Any) -> Any:
|
||||||
needs_unpacking = len(inputs) > 1
|
needs_unpacking = len(inputs) > 1
|
||||||
@@ -121,7 +119,7 @@ class SharpnessJitter(Transform):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.sharpness = self._check_input(sharpness)
|
self.sharpness = self._check_input(sharpness)
|
||||||
|
|
||||||
def _check_input(self, sharpness: float | Sequence[float]) -> tuple[float, float]:
|
def _check_input(self, sharpness):
|
||||||
if isinstance(sharpness, (int | float)):
|
if isinstance(sharpness, (int | float)):
|
||||||
if sharpness < 0:
|
if sharpness < 0:
|
||||||
raise ValueError("If sharpness is a single number, it must be non negative.")
|
raise ValueError("If sharpness is a single number, it must be non negative.")
|
||||||
@@ -146,471 +144,6 @@ class SharpnessJitter(Transform):
|
|||||||
return self._call_kernel(F.adjust_sharpness, inpt, sharpness_factor=sharpness_factor)
|
return self._call_kernel(F.adjust_sharpness, inpt, sharpness_factor=sharpness_factor)
|
||||||
|
|
||||||
|
|
||||||
class GaussianNoise(Transform):
|
|
||||||
"""Add Gaussian noise to simulate camera sensor noise.
|
|
||||||
|
|
||||||
Models readout noise from ADC quantization, which increases in low-light conditions.
|
|
||||||
Common in real-robot setups where wrist cameras operate in suboptimal lighting.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
std: Range (min, max) for noise standard deviation in pixel-value scale (0-255).
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, std: float | Sequence[float] = (5.0, 25.0)) -> None:
|
|
||||||
super().__init__()
|
|
||||||
if isinstance(std, (int, float)):
|
|
||||||
self.std = (0.0, float(std))
|
|
||||||
elif isinstance(std, Sequence) and len(std) == 2:
|
|
||||||
self.std = (float(std[0]), float(std[1]))
|
|
||||||
else:
|
|
||||||
raise TypeError("std must be a number or a sequence with length 2.")
|
|
||||||
if not 0.0 <= self.std[0] <= self.std[1]:
|
|
||||||
raise ValueError(f"std must satisfy 0 <= min <= max, but got {self.std}.")
|
|
||||||
|
|
||||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"std": torch.empty(1).uniform_(self.std[0], self.std[1]).item(),
|
|
||||||
"seed": torch.randint(0, torch.iinfo(torch.int64).max, ()).item(),
|
|
||||||
}
|
|
||||||
|
|
||||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
|
||||||
if isinstance(inpt, torch.Tensor) and inpt.is_floating_point():
|
|
||||||
generator = torch.Generator(device=inpt.device).manual_seed(params["seed"])
|
|
||||||
noise = torch.randn(inpt.shape, device=inpt.device, dtype=inpt.dtype, generator=generator)
|
|
||||||
return (inpt + noise * (params["std"] / 255.0)).clamp(0.0, 1.0)
|
|
||||||
return inpt
|
|
||||||
|
|
||||||
|
|
||||||
class MotionBlur(Transform):
|
|
||||||
"""Apply directional motion blur to simulate fast robot or object movement.
|
|
||||||
|
|
||||||
Generates a 1D averaging kernel along a random direction, applied via depthwise convolution.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
kernel_size: An odd kernel size or a range containing at least one odd kernel size.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, kernel_size: int | Sequence[int] = (3, 11)) -> None:
|
|
||||||
super().__init__()
|
|
||||||
if isinstance(kernel_size, int):
|
|
||||||
self.kernel_size = (kernel_size, kernel_size)
|
|
||||||
elif isinstance(kernel_size, Sequence) and len(kernel_size) == 2:
|
|
||||||
self.kernel_size = (int(kernel_size[0]), int(kernel_size[1]))
|
|
||||||
else:
|
|
||||||
raise TypeError("kernel_size must be an int or a sequence with length 2.")
|
|
||||||
if not 1 <= self.kernel_size[0] <= self.kernel_size[1]:
|
|
||||||
raise ValueError(f"kernel_size must satisfy 1 <= min <= max, but got {self.kernel_size}.")
|
|
||||||
self._first_odd_kernel_size = self.kernel_size[0] + (self.kernel_size[0] + 1) % 2
|
|
||||||
if self._first_odd_kernel_size > self.kernel_size[1]:
|
|
||||||
raise ValueError(f"kernel_size range must contain an odd value, but got {self.kernel_size}.")
|
|
||||||
|
|
||||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
|
||||||
num_odd_sizes = (self.kernel_size[1] - self._first_odd_kernel_size) // 2 + 1
|
|
||||||
size_index = int(torch.randint(0, num_odd_sizes, ()).item())
|
|
||||||
ks = self._first_odd_kernel_size + 2 * size_index
|
|
||||||
angle = torch.empty(1).uniform_(0, 360).item()
|
|
||||||
return {"kernel_size": ks, "angle": angle}
|
|
||||||
|
|
||||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
|
||||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
|
||||||
return inpt
|
|
||||||
if inpt.ndim < 3:
|
|
||||||
raise ValueError(f"MotionBlur expects [..., C, H, W] input, but got shape {inpt.shape}.")
|
|
||||||
|
|
||||||
kernel_size = params["kernel_size"]
|
|
||||||
radius = kernel_size // 2
|
|
||||||
angle = math.radians(params["angle"])
|
|
||||||
positions = torch.linspace(-radius, radius, kernel_size, device=inpt.device)
|
|
||||||
x_coords = (positions * math.cos(angle)).round().to(torch.long) + radius
|
|
||||||
y_coords = (positions * math.sin(angle)).round().to(torch.long) + radius
|
|
||||||
kernel = torch.zeros((kernel_size, kernel_size), device=inpt.device, dtype=inpt.dtype)
|
|
||||||
kernel[y_coords, x_coords] = 1
|
|
||||||
kernel /= kernel.sum()
|
|
||||||
|
|
||||||
channels, height, width = inpt.shape[-3:]
|
|
||||||
flat_input = inpt.reshape(-1, channels, height, width)
|
|
||||||
depthwise_kernel = kernel.expand(channels, 1, kernel_size, kernel_size)
|
|
||||||
padded = torch.nn.functional.pad(flat_input, (radius,) * 4, mode="replicate")
|
|
||||||
output = torch.nn.functional.conv2d(padded, depthwise_kernel, groups=channels)
|
|
||||||
return output.reshape(inpt.shape).clamp(0.0, 1.0)
|
|
||||||
|
|
||||||
|
|
||||||
class JPEGCompression(Transform):
|
|
||||||
"""Simulate JPEG compression artifacts (block artifacts, color banding).
|
|
||||||
|
|
||||||
Models quality degradation from video compression in network-streamed camera feeds.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
quality: Range (min, max) for JPEG quality factor (lower = more artifacts).
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, quality: int | Sequence[int] = (15, 75)) -> None:
|
|
||||||
super().__init__()
|
|
||||||
if isinstance(quality, int):
|
|
||||||
self.quality = (quality, quality)
|
|
||||||
elif isinstance(quality, Sequence) and len(quality) == 2:
|
|
||||||
self.quality = (int(quality[0]), int(quality[1]))
|
|
||||||
else:
|
|
||||||
raise TypeError("quality must be an int or a sequence with length 2.")
|
|
||||||
if not 1 <= self.quality[0] <= self.quality[1] <= 100:
|
|
||||||
raise ValueError(f"quality must satisfy 1 <= min <= max <= 100, but got {self.quality}.")
|
|
||||||
|
|
||||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
|
||||||
return {"quality": int(torch.randint(self.quality[0], self.quality[1] + 1, (1,)).item())}
|
|
||||||
|
|
||||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
|
||||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
|
||||||
return inpt
|
|
||||||
if inpt.ndim < 3:
|
|
||||||
raise ValueError(f"JPEGCompression expects [..., C, H, W] input, but got shape {inpt.shape}.")
|
|
||||||
|
|
||||||
channels, height, width = inpt.shape[-3:]
|
|
||||||
if channels not in (1, 3):
|
|
||||||
raise ValueError(f"JPEGCompression expects 1 or 3 channels, but got {channels}.")
|
|
||||||
|
|
||||||
flat_input = inpt.reshape(-1, channels, height, width)
|
|
||||||
flat_uint8 = (flat_input.clamp(0.0, 1.0) * 255).round().to(torch.uint8).cpu()
|
|
||||||
decoded_frames = [
|
|
||||||
decode_image(encode_jpeg(frame, quality=params["quality"])) for frame in flat_uint8.unbind()
|
|
||||||
]
|
|
||||||
output = torch.stack(decoded_frames).to(device=inpt.device, dtype=inpt.dtype) / 255.0
|
|
||||||
return output.reshape(inpt.shape)
|
|
||||||
|
|
||||||
|
|
||||||
class GaussianPatchBrightness(Transform):
|
|
||||||
"""Apply spatially-varying brightness with Gaussian patches.
|
|
||||||
|
|
||||||
Simulates uneven overhead lighting, spotlights, and shadow patches commonly
|
|
||||||
encountered in real robot workspaces with multiple light sources.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
num_patches: Range (min, max) for number of brightness patches.
|
|
||||||
sigma_range: Range for Gaussian sigma as fraction of image size.
|
|
||||||
factor_range: Range for brightness factor (< 1 darkens, > 1 brightens).
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
num_patches: int | Sequence[int] = (1, 4),
|
|
||||||
sigma_range: Sequence[float] = (0.05, 0.25),
|
|
||||||
factor_range: Sequence[float] = (0.4, 1.6),
|
|
||||||
) -> None:
|
|
||||||
super().__init__()
|
|
||||||
if isinstance(num_patches, int):
|
|
||||||
self.num_patches = (num_patches, num_patches)
|
|
||||||
elif isinstance(num_patches, Sequence) and len(num_patches) == 2:
|
|
||||||
self.num_patches = (int(num_patches[0]), int(num_patches[1]))
|
|
||||||
else:
|
|
||||||
raise TypeError("num_patches must be an int or a sequence with length 2.")
|
|
||||||
if not 1 <= self.num_patches[0] <= self.num_patches[1]:
|
|
||||||
raise ValueError(f"num_patches must satisfy 1 <= min <= max, but got {self.num_patches}.")
|
|
||||||
if not isinstance(sigma_range, Sequence) or len(sigma_range) != 2:
|
|
||||||
raise TypeError("sigma_range must be a sequence with length 2.")
|
|
||||||
self.sigma_range = (float(sigma_range[0]), float(sigma_range[1]))
|
|
||||||
if not 0.0 < self.sigma_range[0] <= self.sigma_range[1]:
|
|
||||||
raise ValueError(f"sigma_range must satisfy 0 < min <= max, but got {self.sigma_range}.")
|
|
||||||
if not isinstance(factor_range, Sequence) or len(factor_range) != 2:
|
|
||||||
raise TypeError("factor_range must be a sequence with length 2.")
|
|
||||||
self.factor_range = (float(factor_range[0]), float(factor_range[1]))
|
|
||||||
if not 0.0 <= self.factor_range[0] <= self.factor_range[1]:
|
|
||||||
raise ValueError(f"factor_range must satisfy 0 <= min <= max, but got {self.factor_range}.")
|
|
||||||
|
|
||||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
|
||||||
n = int(torch.randint(self.num_patches[0], self.num_patches[1] + 1, (1,)).item())
|
|
||||||
return {
|
|
||||||
"centers": torch.rand(n, 2).tolist(),
|
|
||||||
"sigmas": torch.empty(n).uniform_(self.sigma_range[0], self.sigma_range[1]).tolist(),
|
|
||||||
"factors": torch.empty(n).uniform_(self.factor_range[0], self.factor_range[1]).tolist(),
|
|
||||||
}
|
|
||||||
|
|
||||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
|
||||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
|
||||||
return inpt
|
|
||||||
h, w = inpt.shape[-2:]
|
|
||||||
mask = torch.ones(h, w, device=inpt.device, dtype=inpt.dtype)
|
|
||||||
grid_y = torch.linspace(0, 1, h, device=inpt.device, dtype=inpt.dtype)
|
|
||||||
grid_x = torch.linspace(0, 1, w, device=inpt.device, dtype=inpt.dtype)
|
|
||||||
yy, xx = torch.meshgrid(grid_y, grid_x, indexing="ij")
|
|
||||||
for (cy, cx), sigma, factor in zip(
|
|
||||||
params["centers"], params["sigmas"], params["factors"], strict=True
|
|
||||||
):
|
|
||||||
gauss = torch.exp(-((yy - cy) ** 2 + (xx - cx) ** 2) / (2 * sigma**2))
|
|
||||||
mask = mask * (1.0 + (factor - 1.0) * gauss)
|
|
||||||
broadcast_shape = (1,) * (inpt.ndim - 2) + (h, w)
|
|
||||||
return (inpt * mask.reshape(broadcast_shape)).clamp(0.0, 1.0)
|
|
||||||
|
|
||||||
|
|
||||||
class RandomShadow(Transform):
|
|
||||||
"""Add random vertical band shadow with smooth edges.
|
|
||||||
|
|
||||||
Simulates cast shadows from objects or people near the robot workspace.
|
|
||||||
Symmetric: randomly brightens or darkens to prevent BatchNorm stats shift.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
opacity: Range (min, max) for shadow/highlight opacity.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, opacity: float | Sequence[float] = (0.3, 0.6)) -> None:
|
|
||||||
super().__init__()
|
|
||||||
if isinstance(opacity, (int, float)):
|
|
||||||
self.opacity = (float(opacity), float(opacity))
|
|
||||||
elif isinstance(opacity, Sequence) and len(opacity) == 2:
|
|
||||||
self.opacity = (float(opacity[0]), float(opacity[1]))
|
|
||||||
else:
|
|
||||||
raise TypeError("opacity must be a number or a sequence with length 2.")
|
|
||||||
if not 0.0 <= self.opacity[0] <= self.opacity[1] <= 1.0:
|
|
||||||
raise ValueError(f"opacity must satisfy 0 <= min <= max <= 1, but got {self.opacity}.")
|
|
||||||
|
|
||||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"opacity": torch.empty(1).uniform_(self.opacity[0], self.opacity[1]).item(),
|
|
||||||
"start": torch.rand(1).item(),
|
|
||||||
"width": torch.empty(1).uniform_(1 / 3, 2 / 3).item(),
|
|
||||||
"direction": -1.0 if torch.rand(1).item() < 0.5 else 1.0,
|
|
||||||
}
|
|
||||||
|
|
||||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
|
||||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
|
||||||
return inpt
|
|
||||||
if inpt.ndim < 3:
|
|
||||||
raise ValueError(f"RandomShadow expects [..., C, H, W] input, but got shape {inpt.shape}.")
|
|
||||||
|
|
||||||
h, w = inpt.shape[-2:]
|
|
||||||
band_width = max(1, min(w, round(params["width"] * w)))
|
|
||||||
x_start = round(params["start"] * (w - band_width))
|
|
||||||
x_end = x_start + band_width
|
|
||||||
mask = torch.ones(h, w, device=inpt.device, dtype=inpt.dtype)
|
|
||||||
mask[:, x_start:x_end] = 1.0 + params["direction"] * params["opacity"]
|
|
||||||
|
|
||||||
smoothing_size = min(8, h, w)
|
|
||||||
if smoothing_size > 1:
|
|
||||||
batched_mask = mask[None, None]
|
|
||||||
small = torch.nn.functional.avg_pool2d(batched_mask, smoothing_size, stride=smoothing_size)
|
|
||||||
mask = torch.nn.functional.interpolate(small, size=(h, w), mode="bilinear", align_corners=False)[
|
|
||||||
0, 0
|
|
||||||
]
|
|
||||||
|
|
||||||
broadcast_shape = (1,) * (inpt.ndim - 2) + (h, w)
|
|
||||||
return (inpt * mask.reshape(broadcast_shape)).clamp(0.0, 1.0)
|
|
||||||
|
|
||||||
|
|
||||||
class CoarseDropout(Transform):
|
|
||||||
"""Drop random rectangular patches to simulate partial occlusion.
|
|
||||||
|
|
||||||
Models objects, hands, or cables passing through the camera field of view
|
|
||||||
during robot manipulation.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
max_holes: Maximum number of rectangular patches to drop.
|
|
||||||
max_height_frac: Maximum patch height as fraction of image height.
|
|
||||||
max_width_frac: Maximum patch width as fraction of image width.
|
|
||||||
fill_value: Value to fill dropped regions with.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
max_holes: int = 8,
|
|
||||||
max_height_frac: float = 0.07,
|
|
||||||
max_width_frac: float = 0.07,
|
|
||||||
fill_value: float = 0.0,
|
|
||||||
) -> None:
|
|
||||||
super().__init__()
|
|
||||||
if not isinstance(max_holes, int):
|
|
||||||
raise TypeError("max_holes must be an int.")
|
|
||||||
if max_holes < 1:
|
|
||||||
raise ValueError(f"max_holes must be at least 1, but got {max_holes}.")
|
|
||||||
if not 0.0 < max_height_frac <= 1.0:
|
|
||||||
raise ValueError(f"max_height_frac must be in (0, 1], but got {max_height_frac}.")
|
|
||||||
if not 0.0 < max_width_frac <= 1.0:
|
|
||||||
raise ValueError(f"max_width_frac must be in (0, 1], but got {max_width_frac}.")
|
|
||||||
if not 0.0 <= fill_value <= 1.0:
|
|
||||||
raise ValueError(f"fill_value must be in [0, 1], but got {fill_value}.")
|
|
||||||
self.max_holes = max_holes
|
|
||||||
self.max_height_frac = max_height_frac
|
|
||||||
self.max_width_frac = max_width_frac
|
|
||||||
self.fill_value = fill_value
|
|
||||||
|
|
||||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
|
||||||
n = int(torch.randint(1, self.max_holes + 1, (1,)).item())
|
|
||||||
sizes = torch.rand(n, 2)
|
|
||||||
sizes[:, 0] *= self.max_height_frac
|
|
||||||
sizes[:, 1] *= self.max_width_frac
|
|
||||||
return {"sizes": sizes.tolist(), "positions": torch.rand(n, 2).tolist()}
|
|
||||||
|
|
||||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
|
||||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
|
||||||
return inpt
|
|
||||||
if inpt.ndim < 3:
|
|
||||||
raise ValueError(f"CoarseDropout expects [..., C, H, W] input, but got shape {inpt.shape}.")
|
|
||||||
|
|
||||||
h, w = inpt.shape[-2:]
|
|
||||||
result = inpt.clone()
|
|
||||||
for (height_frac, width_frac), (y_frac, x_frac) in zip(
|
|
||||||
params["sizes"], params["positions"], strict=True
|
|
||||||
):
|
|
||||||
hole_h = max(1, min(h, round(height_frac * h)))
|
|
||||||
hole_w = max(1, min(w, round(width_frac * w)))
|
|
||||||
y = round(y_frac * (h - hole_h))
|
|
||||||
x = round(x_frac * (w - hole_w))
|
|
||||||
result[..., y : y + hole_h, x : x + hole_w] = self.fill_value
|
|
||||||
return result
|
|
||||||
|
|
||||||
|
|
||||||
class GammaCorrection(Transform):
|
|
||||||
"""Apply random gamma correction to simulate exposure variation.
|
|
||||||
|
|
||||||
Models different camera auto-exposure settings and sensor response curves.
|
|
||||||
Uses log-symmetric sampling so brightening and darkening are equally likely,
|
|
||||||
preventing BatchNorm statistics shift.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
gamma: Range (min, max) for gamma value. Values < 1 brighten, > 1 darken.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, gamma: float | Sequence[float] = (0.5, 2.0)) -> None:
|
|
||||||
super().__init__()
|
|
||||||
if isinstance(gamma, (int, float)):
|
|
||||||
gamma = float(gamma)
|
|
||||||
if gamma <= 0:
|
|
||||||
raise ValueError(f"gamma must be positive, but got {gamma}.")
|
|
||||||
self.gamma = (min(gamma, 1.0 / gamma), max(gamma, 1.0 / gamma))
|
|
||||||
elif isinstance(gamma, Sequence) and len(gamma) == 2:
|
|
||||||
self.gamma = (float(gamma[0]), float(gamma[1]))
|
|
||||||
else:
|
|
||||||
raise TypeError("gamma must be a number or a sequence with length 2.")
|
|
||||||
if not 0.0 < self.gamma[0] <= self.gamma[1]:
|
|
||||||
raise ValueError(f"gamma must satisfy 0 < min <= max, but got {self.gamma}.")
|
|
||||||
|
|
||||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
|
||||||
log_lo = math.log(self.gamma[0])
|
|
||||||
log_hi = math.log(self.gamma[1])
|
|
||||||
gamma = math.exp(torch.empty(1).uniform_(log_lo, log_hi).item())
|
|
||||||
return {"gamma": gamma}
|
|
||||||
|
|
||||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
|
||||||
if isinstance(inpt, torch.Tensor) and inpt.is_floating_point():
|
|
||||||
return inpt.pow(params["gamma"]).clamp(0.0, 1.0)
|
|
||||||
return inpt
|
|
||||||
|
|
||||||
|
|
||||||
# From the paper authors' MIT-licensed reference implementation:
|
|
||||||
# https://github.com/TheZino/PlanckianJitter
|
|
||||||
_PLANCKIAN_BLACKBODY_COEFFICIENTS = (
|
|
||||||
(0.6743, 0.4029, 0.0013),
|
|
||||||
(0.6281, 0.4241, 0.1665),
|
|
||||||
(0.5919, 0.4372, 0.2513),
|
|
||||||
(0.5623, 0.4457, 0.3154),
|
|
||||||
(0.5376, 0.4515, 0.3672),
|
|
||||||
(0.5163, 0.4555, 0.4103),
|
|
||||||
(0.4979, 0.4584, 0.4468),
|
|
||||||
(0.4816, 0.4604, 0.4782),
|
|
||||||
(0.4672, 0.4619, 0.5053),
|
|
||||||
(0.4542, 0.4630, 0.5289),
|
|
||||||
(0.4426, 0.4638, 0.5497),
|
|
||||||
(0.4320, 0.4644, 0.5681),
|
|
||||||
(0.4223, 0.4648, 0.5844),
|
|
||||||
(0.4135, 0.4651, 0.5990),
|
|
||||||
(0.4054, 0.4653, 0.6121),
|
|
||||||
(0.3980, 0.4654, 0.6239),
|
|
||||||
(0.3911, 0.4655, 0.6346),
|
|
||||||
(0.3847, 0.4656, 0.6444),
|
|
||||||
(0.3787, 0.4656, 0.6532),
|
|
||||||
(0.3732, 0.4656, 0.6613),
|
|
||||||
(0.3680, 0.4655, 0.6688),
|
|
||||||
(0.3632, 0.4655, 0.6756),
|
|
||||||
(0.3586, 0.4655, 0.6820),
|
|
||||||
(0.3544, 0.4654, 0.6878),
|
|
||||||
(0.3503, 0.4653, 0.6933),
|
|
||||||
)
|
|
||||||
_PLANCKIAN_MIN_TEMPERATURE = 3_000
|
|
||||||
_PLANCKIAN_MAX_TEMPERATURE = 15_000
|
|
||||||
_PLANCKIAN_TEMPERATURE_STEP = 500
|
|
||||||
|
|
||||||
|
|
||||||
class PlanckianJitter(Transform):
|
|
||||||
"""Simulate color temperature shift along the Planckian locus.
|
|
||||||
|
|
||||||
Samples one black-body temperature and applies the corresponding correlated red
|
|
||||||
and blue channel scaling while preserving the green channel. Coefficients between
|
|
||||||
the tabulated 500 K intervals are linearly interpolated.
|
|
||||||
|
|
||||||
Reference: Zini et al., "Planckian Jitter", CVPR 2022 Workshop.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
temperature: A fixed color temperature or range in Kelvin. Supported values
|
|
||||||
are between 3000 K and 15000 K.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(self, temperature: int | Sequence[int] = (3_000, 15_000)) -> None:
|
|
||||||
super().__init__()
|
|
||||||
if isinstance(temperature, int):
|
|
||||||
self.temperature = (temperature, temperature)
|
|
||||||
elif isinstance(temperature, Sequence) and len(temperature) == 2:
|
|
||||||
self.temperature = (int(temperature[0]), int(temperature[1]))
|
|
||||||
else:
|
|
||||||
raise TypeError("temperature must be an int or a sequence with length 2.")
|
|
||||||
if not (
|
|
||||||
_PLANCKIAN_MIN_TEMPERATURE
|
|
||||||
<= self.temperature[0]
|
|
||||||
<= self.temperature[1]
|
|
||||||
<= _PLANCKIAN_MAX_TEMPERATURE
|
|
||||||
):
|
|
||||||
raise ValueError(
|
|
||||||
"temperature must satisfy "
|
|
||||||
f"{_PLANCKIAN_MIN_TEMPERATURE} <= min <= max <= {_PLANCKIAN_MAX_TEMPERATURE}, "
|
|
||||||
f"but got {self.temperature}."
|
|
||||||
)
|
|
||||||
|
|
||||||
def make_params(self, flat_inputs: list[Any]) -> dict[str, Any]:
|
|
||||||
temperature = int(torch.randint(self.temperature[0], self.temperature[1] + 1, ()).item())
|
|
||||||
return {"temperature": temperature}
|
|
||||||
|
|
||||||
def transform(self, inpt: Any, params: dict[str, Any]) -> Any:
|
|
||||||
if not isinstance(inpt, torch.Tensor) or not inpt.is_floating_point():
|
|
||||||
return inpt
|
|
||||||
if inpt.ndim < 3 or inpt.shape[-3] != 3:
|
|
||||||
raise ValueError(f"PlanckianJitter expects [..., 3, H, W] input, but got shape {inpt.shape}.")
|
|
||||||
|
|
||||||
table_position = (params["temperature"] - _PLANCKIAN_MIN_TEMPERATURE) / _PLANCKIAN_TEMPERATURE_STEP
|
|
||||||
left_index = math.floor(table_position)
|
|
||||||
right_index = min(left_index + 1, len(_PLANCKIAN_BLACKBODY_COEFFICIENTS) - 1)
|
|
||||||
interpolation_weight = table_position - left_index
|
|
||||||
|
|
||||||
left = torch.tensor(
|
|
||||||
_PLANCKIAN_BLACKBODY_COEFFICIENTS[left_index],
|
|
||||||
device=inpt.device,
|
|
||||||
dtype=inpt.dtype,
|
|
||||||
)
|
|
||||||
right = torch.tensor(
|
|
||||||
_PLANCKIAN_BLACKBODY_COEFFICIENTS[right_index],
|
|
||||||
device=inpt.device,
|
|
||||||
dtype=inpt.dtype,
|
|
||||||
)
|
|
||||||
coefficients = torch.lerp(left, right, interpolation_weight)
|
|
||||||
scale = torch.stack(
|
|
||||||
(
|
|
||||||
coefficients[0] / coefficients[1],
|
|
||||||
coefficients.new_tensor(1.0),
|
|
||||||
coefficients[2] / coefficients[1],
|
|
||||||
)
|
|
||||||
)
|
|
||||||
broadcast_shape = (1,) * (inpt.ndim - 3) + (3, 1, 1)
|
|
||||||
return (inpt * scale.reshape(broadcast_shape)).clamp(0.0, 1.0)
|
|
||||||
|
|
||||||
|
|
||||||
_CUSTOM_TRANSFORMS: dict[str, type[Transform]] = {
|
|
||||||
"SharpnessJitter": SharpnessJitter,
|
|
||||||
"GaussianNoise": GaussianNoise,
|
|
||||||
"MotionBlur": MotionBlur,
|
|
||||||
"JPEGCompression": JPEGCompression,
|
|
||||||
"GaussianPatchBrightness": GaussianPatchBrightness,
|
|
||||||
"RandomShadow": RandomShadow,
|
|
||||||
"CoarseDropout": CoarseDropout,
|
|
||||||
"GammaCorrection": GammaCorrection,
|
|
||||||
"PlanckianJitter": PlanckianJitter,
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class ImageTransformConfig:
|
class ImageTransformConfig:
|
||||||
"""
|
"""
|
||||||
@@ -682,18 +215,17 @@ class ImageTransformsConfig:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def make_transform_from_config(cfg: ImageTransformConfig) -> Transform:
|
def make_transform_from_config(cfg: ImageTransformConfig):
|
||||||
if cfg.type in _CUSTOM_TRANSFORMS:
|
if cfg.type == "SharpnessJitter":
|
||||||
return _CUSTOM_TRANSFORMS[cfg.type](**cfg.kwargs)
|
return SharpnessJitter(**cfg.kwargs)
|
||||||
|
|
||||||
transform_cls = getattr(v2, cfg.type, None)
|
transform_cls = getattr(v2, cfg.type, None)
|
||||||
if isinstance(transform_cls, type) and issubclass(transform_cls, Transform):
|
if isinstance(transform_cls, type) and issubclass(transform_cls, Transform):
|
||||||
return transform_cls(**cfg.kwargs)
|
return transform_cls(**cfg.kwargs)
|
||||||
|
|
||||||
valid_custom = ", ".join(sorted(_CUSTOM_TRANSFORMS.keys()))
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Transform '{cfg.type}' is not valid. It must be a class in "
|
f"Transform '{cfg.type}' is not valid. It must be a class in "
|
||||||
f"torchvision.transforms.v2 or one of: {valid_custom}."
|
f"torchvision.transforms.v2 or 'SharpnessJitter'."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -704,8 +236,8 @@ class ImageTransforms(Transform):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self._cfg = cfg
|
self._cfg = cfg
|
||||||
|
|
||||||
self.weights: list[float] = []
|
self.weights = []
|
||||||
self.transforms: dict[str, Transform] = {}
|
self.transforms = {}
|
||||||
for tf_name, tf_cfg in cfg.tfs.items():
|
for tf_name, tf_cfg in cfg.tfs.items():
|
||||||
if tf_cfg.weight <= 0.0:
|
if tf_cfg.weight <= 0.0:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -16,39 +16,11 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
import multiprocessing
|
|
||||||
import os
|
import os
|
||||||
import signal
|
import signal
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
|
|
||||||
def ensure_multiprocessing_start_method(start_method: str | None) -> None:
|
|
||||||
"""Set a multiprocessing start method once, or verify the existing method matches.
|
|
||||||
|
|
||||||
Passing ``None`` leaves Python's process-wide default untouched. This is useful
|
|
||||||
when LeRobot is embedded in an application that owns multiprocessing setup.
|
|
||||||
"""
|
|
||||||
if start_method is None:
|
|
||||||
return
|
|
||||||
|
|
||||||
available_methods = multiprocessing.get_all_start_methods()
|
|
||||||
if start_method not in available_methods:
|
|
||||||
raise ValueError(
|
|
||||||
f"Multiprocessing start method must be one of {available_methods} on this platform, "
|
|
||||||
f"got {start_method!r}."
|
|
||||||
)
|
|
||||||
|
|
||||||
current_method = multiprocessing.get_start_method(allow_none=True)
|
|
||||||
if current_method is None:
|
|
||||||
multiprocessing.set_start_method(start_method)
|
|
||||||
elif current_method != start_method:
|
|
||||||
raise RuntimeError(
|
|
||||||
f"Multiprocessing start method is already {current_method!r}; cannot change it to "
|
|
||||||
f"{start_method!r}. Set the configured multiprocessing context to null to keep the "
|
|
||||||
"application's existing method, or launch LeRobot in a fresh process."
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class ProcessSignalHandler:
|
class ProcessSignalHandler:
|
||||||
"""Utility class to attach graceful shutdown signal handlers.
|
"""Utility class to attach graceful shutdown signal handlers.
|
||||||
|
|
||||||
|
|||||||
@@ -133,13 +133,10 @@ def say(text: str, blocking: bool = False):
|
|||||||
else:
|
else:
|
||||||
raise RuntimeError("Unsupported operating system for text-to-speech.")
|
raise RuntimeError("Unsupported operating system for text-to-speech.")
|
||||||
|
|
||||||
try:
|
|
||||||
if blocking:
|
if blocking:
|
||||||
subprocess.run(cmd, check=True, timeout=5)
|
subprocess.run(cmd, check=True)
|
||||||
else:
|
else:
|
||||||
subprocess.Popen(cmd, creationflags=subprocess.CREATE_NO_WINDOW if system == "Windows" else 0)
|
subprocess.Popen(cmd, creationflags=subprocess.CREATE_NO_WINDOW if system == "Windows" else 0)
|
||||||
except (FileNotFoundError, subprocess.TimeoutExpired) as e:
|
|
||||||
logging.warning("Text-to-speech command failed: %s | Error: %s", cmd, e)
|
|
||||||
|
|
||||||
|
|
||||||
def log_say(text: str, play_sounds: bool = True, blocking: bool = False):
|
def log_say(text: str, play_sounds: bool = True, blocking: bool = False):
|
||||||
|
|||||||
@@ -20,7 +20,7 @@
|
|||||||
# ```
|
# ```
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import cv2
|
import cv2
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@@ -123,73 +123,6 @@ def test_invalid_width_connect():
|
|||||||
camera.connect(warmup=False)
|
camera.connect(warmup=False)
|
||||||
|
|
||||||
|
|
||||||
def test_connect_cleans_up_after_settings_failure_and_allows_retry():
|
|
||||||
config = OpenCVCameraConfig(index_or_path=DEFAULT_PNG_FILE_PATH, warmup_s=0)
|
|
||||||
camera = OpenCVCamera(config)
|
|
||||||
opened_captures = []
|
|
||||||
|
|
||||||
def fail_settings():
|
|
||||||
opened_captures.append(camera.videocapture)
|
|
||||||
raise RuntimeError("settings failed")
|
|
||||||
|
|
||||||
with (
|
|
||||||
patch.object(camera, "_configure_capture_settings", side_effect=fail_settings),
|
|
||||||
pytest.raises(RuntimeError, match="settings failed"),
|
|
||||||
):
|
|
||||||
camera.connect(warmup=False)
|
|
||||||
|
|
||||||
assert camera.videocapture is None
|
|
||||||
assert camera.thread is None
|
|
||||||
assert not camera.is_connected
|
|
||||||
assert opened_captures[0] is not None
|
|
||||||
assert not opened_captures[0].isOpened()
|
|
||||||
|
|
||||||
camera.connect(warmup=False)
|
|
||||||
assert camera.is_connected
|
|
||||||
camera.disconnect()
|
|
||||||
|
|
||||||
|
|
||||||
def test_connect_cleans_up_after_warmup_failure_and_allows_retry():
|
|
||||||
config = OpenCVCameraConfig(index_or_path=DEFAULT_PNG_FILE_PATH, warmup_s=1)
|
|
||||||
camera = OpenCVCamera(config)
|
|
||||||
read_threads = []
|
|
||||||
|
|
||||||
def fail_warmup(*_args, **_kwargs):
|
|
||||||
read_threads.append(camera.thread)
|
|
||||||
raise TimeoutError("no frame")
|
|
||||||
|
|
||||||
with (
|
|
||||||
patch.object(camera, "async_read", side_effect=fail_warmup),
|
|
||||||
pytest.raises(TimeoutError, match="no frame"),
|
|
||||||
):
|
|
||||||
camera.connect()
|
|
||||||
|
|
||||||
assert camera.videocapture is None
|
|
||||||
assert camera.thread is None
|
|
||||||
assert not camera.is_connected
|
|
||||||
assert read_threads[0] is not None
|
|
||||||
assert not read_threads[0].is_alive()
|
|
||||||
|
|
||||||
camera.connect(warmup=False)
|
|
||||||
assert camera.is_connected
|
|
||||||
camera.disconnect()
|
|
||||||
|
|
||||||
|
|
||||||
def test_find_cameras_releases_unopened_handles():
|
|
||||||
module_path = OpenCVCamera.__module__
|
|
||||||
unopened_capture = MagicMock()
|
|
||||||
unopened_capture.isOpened.return_value = False
|
|
||||||
|
|
||||||
with (
|
|
||||||
patch(f"{module_path}.platform.system", return_value="Darwin"),
|
|
||||||
patch(f"{module_path}.MAX_OPENCV_INDEX", 1),
|
|
||||||
patch(f"{module_path}.cv2.VideoCapture", return_value=unopened_capture),
|
|
||||||
):
|
|
||||||
assert OpenCVCamera.find_cameras() == []
|
|
||||||
|
|
||||||
unopened_capture.release.assert_called_once_with()
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("index_or_path", TEST_IMAGE_PATHS, ids=TEST_IMAGE_SIZES)
|
@pytest.mark.parametrize("index_or_path", TEST_IMAGE_PATHS, ids=TEST_IMAGE_SIZES)
|
||||||
def test_read(index_or_path):
|
def test_read(index_or_path):
|
||||||
config = OpenCVCameraConfig(index_or_path=index_or_path, warmup_s=0)
|
config = OpenCVCameraConfig(index_or_path=index_or_path, warmup_s=0)
|
||||||
|
|||||||
@@ -20,7 +20,7 @@
|
|||||||
# ```
|
# ```
|
||||||
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
@@ -30,8 +30,6 @@ from lerobot.utils.errors import DeviceAlreadyConnectedError, DeviceNotConnected
|
|||||||
|
|
||||||
pytest.importorskip("pyrealsense2")
|
pytest.importorskip("pyrealsense2")
|
||||||
|
|
||||||
import pyrealsense2 as rs
|
|
||||||
|
|
||||||
from lerobot.cameras.realsense import RealSenseCamera, RealSenseCameraConfig
|
from lerobot.cameras.realsense import RealSenseCamera, RealSenseCameraConfig
|
||||||
|
|
||||||
TEST_ARTIFACTS_DIR = Path(__file__).parent.parent / "artifacts" / "cameras"
|
TEST_ARTIFACTS_DIR = Path(__file__).parent.parent / "artifacts" / "cameras"
|
||||||
@@ -63,17 +61,6 @@ def test_abc_implementation():
|
|||||||
_ = RealSenseCamera(config)
|
_ = RealSenseCamera(config)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("option", ["exposure", "gain", "white_balance"])
|
|
||||||
def test_manual_color_option_requires_rgb(option):
|
|
||||||
with pytest.raises(ValueError, match="use_rgb=True"):
|
|
||||||
RealSenseCameraConfig(
|
|
||||||
serial_number_or_name="042",
|
|
||||||
use_rgb=False,
|
|
||||||
use_depth=True,
|
|
||||||
**{option: 100},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_connect():
|
def test_connect():
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", warmup_s=0)
|
config = RealSenseCameraConfig(serial_number_or_name="042", warmup_s=0)
|
||||||
|
|
||||||
@@ -96,27 +83,6 @@ def test_connect_invalid_camera_path(patch_realsense):
|
|||||||
camera.connect(warmup=False)
|
camera.connect(warmup=False)
|
||||||
|
|
||||||
|
|
||||||
def test_connect_cleans_up_when_sensor_configuration_fails():
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", exposure=120)
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
pipeline = MagicMock()
|
|
||||||
pipeline.start.return_value = MagicMock()
|
|
||||||
|
|
||||||
with (
|
|
||||||
patch("lerobot.cameras.realsense.camera_realsense.rs.pipeline", return_value=pipeline),
|
|
||||||
patch.object(camera, "_configure_rs_pipeline_config"),
|
|
||||||
patch.object(camera, "_configure_capture_settings"),
|
|
||||||
patch.object(camera, "_configure_sensor_options", side_effect=ValueError("invalid exposure")),
|
|
||||||
pytest.raises(ValueError, match="invalid exposure"),
|
|
||||||
):
|
|
||||||
camera.connect(warmup=False)
|
|
||||||
|
|
||||||
pipeline.stop.assert_called_once_with()
|
|
||||||
assert camera.rs_pipeline is None
|
|
||||||
assert camera.rs_profile is None
|
|
||||||
assert not camera.is_connected
|
|
||||||
|
|
||||||
|
|
||||||
def test_invalid_width_connect():
|
def test_invalid_width_connect():
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", width=99999, height=480, fps=30)
|
config = RealSenseCameraConfig(serial_number_or_name="042", width=99999, height=480, fps=30)
|
||||||
camera = RealSenseCamera(config)
|
camera = RealSenseCamera(config)
|
||||||
@@ -125,33 +91,6 @@ def test_invalid_width_connect():
|
|||||||
camera.connect(warmup=False)
|
camera.connect(warmup=False)
|
||||||
|
|
||||||
|
|
||||||
def test_connect_cleans_up_after_warmup_failure_and_allows_retry():
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", width=640, height=480, fps=30)
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
read_threads = []
|
|
||||||
|
|
||||||
def fail_warmup(*_args, **_kwargs):
|
|
||||||
read_threads.append(camera.thread)
|
|
||||||
raise TimeoutError("no frame")
|
|
||||||
|
|
||||||
with (
|
|
||||||
patch.object(camera, "async_read", side_effect=fail_warmup),
|
|
||||||
pytest.raises(TimeoutError, match="no frame"),
|
|
||||||
):
|
|
||||||
camera.connect()
|
|
||||||
|
|
||||||
assert camera.rs_pipeline is None
|
|
||||||
assert camera.rs_profile is None
|
|
||||||
assert camera.thread is None
|
|
||||||
assert not camera.is_connected
|
|
||||||
assert read_threads[0] is not None
|
|
||||||
assert not read_threads[0].is_alive()
|
|
||||||
|
|
||||||
camera.connect(warmup=False)
|
|
||||||
assert camera.is_connected
|
|
||||||
camera.disconnect()
|
|
||||||
|
|
||||||
|
|
||||||
def test_read():
|
def test_read():
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", width=640, height=480, fps=30, warmup_s=0)
|
config = RealSenseCameraConfig(serial_number_or_name="042", width=640, height=480, fps=30, warmup_s=0)
|
||||||
with RealSenseCamera(config) as camera:
|
with RealSenseCamera(config) as camera:
|
||||||
@@ -289,203 +228,6 @@ def test_read_latest_too_old():
|
|||||||
_ = camera.read_latest(max_age_ms=0) # immediately too old
|
_ = camera.read_latest(max_age_ms=0) # immediately too old
|
||||||
|
|
||||||
|
|
||||||
def _make_mock_sensor(name: str, supported_options: set | None = None) -> MagicMock:
|
|
||||||
"""Build a fake rs.sensor that reports a name and a configurable supported-options set."""
|
|
||||||
supported = supported_options if supported_options is not None else set()
|
|
||||||
sensor = MagicMock()
|
|
||||||
sensor.get_info.return_value = name
|
|
||||||
sensor.supports.side_effect = lambda opt: opt in supported
|
|
||||||
return sensor
|
|
||||||
|
|
||||||
|
|
||||||
def _attach_mock_color_sensor(camera: RealSenseCamera, sensor: MagicMock) -> None:
|
|
||||||
"""Wire camera.rs_profile so _get_color_sensor finds the given sensor."""
|
|
||||||
profile = MagicMock()
|
|
||||||
device = MagicMock()
|
|
||||||
device.query_sensors.return_value = [sensor]
|
|
||||||
profile.get_device.return_value = device
|
|
||||||
camera.rs_profile = profile
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_color_sensor_prefers_rgb_camera():
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042")
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
rgb = _make_mock_sensor("RGB Camera")
|
|
||||||
stereo = _make_mock_sensor("Stereo Module")
|
|
||||||
profile = MagicMock()
|
|
||||||
device = MagicMock()
|
|
||||||
device.query_sensors.return_value = [stereo, rgb]
|
|
||||||
profile.get_device.return_value = device
|
|
||||||
camera.rs_profile = profile
|
|
||||||
|
|
||||||
assert camera._get_color_sensor() is rgb
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_color_sensor_falls_back_to_stereo_module():
|
|
||||||
"""D405 has no separate RGB module; color comes from Stereo Module."""
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042")
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
stereo = _make_mock_sensor("Stereo Module")
|
|
||||||
_attach_mock_color_sensor(camera, stereo)
|
|
||||||
|
|
||||||
assert camera._get_color_sensor() is stereo
|
|
||||||
|
|
||||||
|
|
||||||
def test_get_color_sensor_raises_with_available_sensors():
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042")
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
other = _make_mock_sensor("Motion Module")
|
|
||||||
_attach_mock_color_sensor(camera, other)
|
|
||||||
|
|
||||||
with pytest.raises(RuntimeError, match="Motion Module"):
|
|
||||||
camera._get_color_sensor()
|
|
||||||
|
|
||||||
|
|
||||||
def test_configure_sensor_options_skipped_when_none():
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042")
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
with patch.object(RealSenseCamera, "_get_color_sensor") as mock_get:
|
|
||||||
camera._configure_sensor_options()
|
|
||||||
mock_get.assert_not_called()
|
|
||||||
|
|
||||||
|
|
||||||
def test_configure_sensor_options_applies_all_values():
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", exposure=120, gain=64, white_balance=4600)
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
sensor = _make_mock_sensor(
|
|
||||||
"RGB Camera",
|
|
||||||
supported_options={
|
|
||||||
rs.option.enable_auto_exposure,
|
|
||||||
rs.option.exposure,
|
|
||||||
rs.option.gain,
|
|
||||||
rs.option.enable_auto_white_balance,
|
|
||||||
rs.option.white_balance,
|
|
||||||
},
|
|
||||||
)
|
|
||||||
_attach_mock_color_sensor(camera, sensor)
|
|
||||||
|
|
||||||
camera._configure_sensor_options()
|
|
||||||
|
|
||||||
sensor.set_option.assert_any_call(rs.option.enable_auto_exposure, 0)
|
|
||||||
sensor.set_option.assert_any_call(rs.option.exposure, 120)
|
|
||||||
sensor.set_option.assert_any_call(rs.option.gain, 64)
|
|
||||||
sensor.set_option.assert_any_call(rs.option.enable_auto_white_balance, 0)
|
|
||||||
sensor.set_option.assert_any_call(rs.option.white_balance, 4600)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
("config_field", "option", "label"),
|
|
||||||
[
|
|
||||||
("exposure", rs.option.exposure, "exposure"),
|
|
||||||
("gain", rs.option.gain, "gain"),
|
|
||||||
("white_balance", rs.option.white_balance, "white balance"),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_configure_sensor_options_raises_when_requested_option_is_unsupported(config_field, option, label):
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", **{config_field: 100})
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
sensor = _make_mock_sensor("RGB Camera", supported_options=set())
|
|
||||||
_attach_mock_color_sensor(camera, sensor)
|
|
||||||
|
|
||||||
with pytest.raises(ValueError, match=label):
|
|
||||||
camera._configure_sensor_options()
|
|
||||||
|
|
||||||
sensor.supports.assert_any_call(option)
|
|
||||||
sensor.set_option.assert_not_called()
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
("config_field", "option", "value"),
|
|
||||||
[
|
|
||||||
("exposure", rs.option.exposure, 120),
|
|
||||||
("gain", rs.option.gain, 64),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_configure_sensor_options_exposure_or_gain_disables_auto_exposure(config_field, option, value):
|
|
||||||
"""white_balance=None should not touch auto white balance."""
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", **{config_field: value})
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
sensor = _make_mock_sensor(
|
|
||||||
"RGB Camera",
|
|
||||||
supported_options={rs.option.enable_auto_exposure, option},
|
|
||||||
)
|
|
||||||
_attach_mock_color_sensor(camera, sensor)
|
|
||||||
|
|
||||||
camera._configure_sensor_options()
|
|
||||||
|
|
||||||
calls = [call.args for call in sensor.set_option.call_args_list]
|
|
||||||
assert (rs.option.enable_auto_exposure, 0) in calls
|
|
||||||
assert (option, value) in calls
|
|
||||||
for opt, _ in calls:
|
|
||||||
assert opt != rs.option.enable_auto_white_balance
|
|
||||||
assert opt != rs.option.white_balance
|
|
||||||
|
|
||||||
|
|
||||||
def test_configure_sensor_options_warns_when_auto_exposure_control_is_unsupported(caplog):
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", exposure=120)
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
sensor = _make_mock_sensor("RGB Camera", supported_options={rs.option.exposure})
|
|
||||||
_attach_mock_color_sensor(camera, sensor)
|
|
||||||
|
|
||||||
with caplog.at_level("WARNING"):
|
|
||||||
camera._configure_sensor_options()
|
|
||||||
|
|
||||||
sensor.set_option.assert_called_once_with(rs.option.exposure, 120)
|
|
||||||
assert "does not support disabling auto-exposure" in caplog.text
|
|
||||||
|
|
||||||
|
|
||||||
def test_configure_sensor_options_warns_when_auto_white_balance_control_is_unsupported(caplog):
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", white_balance=4600)
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
sensor = _make_mock_sensor("RGB Camera", supported_options={rs.option.white_balance})
|
|
||||||
_attach_mock_color_sensor(camera, sensor)
|
|
||||||
|
|
||||||
with caplog.at_level("WARNING"):
|
|
||||||
camera._configure_sensor_options()
|
|
||||||
|
|
||||||
sensor.set_option.assert_called_once_with(rs.option.white_balance, 4600)
|
|
||||||
assert "does not support disabling auto white balance" in caplog.text
|
|
||||||
|
|
||||||
|
|
||||||
def test_configure_sensor_options_out_of_range_raises_value_error():
|
|
||||||
"""set_option errors should be re-raised as ValueError with range diagnostics."""
|
|
||||||
config = RealSenseCameraConfig(serial_number_or_name="042", exposure=999999)
|
|
||||||
camera = RealSenseCamera(config)
|
|
||||||
|
|
||||||
sensor = _make_mock_sensor(
|
|
||||||
"RGB Camera",
|
|
||||||
supported_options={rs.option.enable_auto_exposure, rs.option.exposure},
|
|
||||||
)
|
|
||||||
|
|
||||||
def fake_set_option(option, value):
|
|
||||||
if option == rs.option.exposure:
|
|
||||||
raise RuntimeError("value out of range")
|
|
||||||
|
|
||||||
sensor.set_option.side_effect = fake_set_option
|
|
||||||
|
|
||||||
option_range = MagicMock(min=1, max=10000, step=1, default=156)
|
|
||||||
sensor.get_option_range.return_value = option_range
|
|
||||||
|
|
||||||
_attach_mock_color_sensor(camera, sensor)
|
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="exposure") as exc_info:
|
|
||||||
camera._configure_sensor_options()
|
|
||||||
|
|
||||||
msg = str(exc_info.value)
|
|
||||||
assert "999999" in msg
|
|
||||||
assert "min=1" in msg
|
|
||||||
assert "max=10000" in msg
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"rotation",
|
"rotation",
|
||||||
[
|
[
|
||||||
|
|||||||
@@ -1,104 +0,0 @@
|
|||||||
# Copyright 2025 The HuggingFace Inc. team. All rights reserved.
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])")
|
|
||||||
|
|
||||||
from lerobot.scripts.augment_dataset_quantile_stats import (
|
|
||||||
compute_quantile_stats_for_dataset,
|
|
||||||
has_quantile_stats,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _numeric_keys(dataset):
|
|
||||||
return [k for k, v in dataset.features.items() if v["dtype"] not in ("image", "video", "string")]
|
|
||||||
|
|
||||||
|
|
||||||
def _image_keys(dataset):
|
|
||||||
return [k for k, v in dataset.features.items() if v["dtype"] in ("image", "video")]
|
|
||||||
|
|
||||||
|
|
||||||
def test_numeric_stats_are_unaffected_by_sampling(tmp_path, lerobot_dataset_factory):
|
|
||||||
"""Sampling only touches image/video frames; numeric features are read in
|
|
||||||
full either way, so their stats must be identical with and without sampling."""
|
|
||||||
dataset = lerobot_dataset_factory(
|
|
||||||
root=tmp_path / "ds", total_episodes=2, total_frames=400, use_videos=False
|
|
||||||
)
|
|
||||||
|
|
||||||
exact = compute_quantile_stats_for_dataset(dataset, use_sampling=False)
|
|
||||||
sampled = compute_quantile_stats_for_dataset(dataset, use_sampling=True)
|
|
||||||
|
|
||||||
numeric_keys = _numeric_keys(dataset)
|
|
||||||
assert numeric_keys, "fixture should expose numeric features"
|
|
||||||
for key in numeric_keys:
|
|
||||||
if key not in exact:
|
|
||||||
continue
|
|
||||||
for stat in ("mean", "std", "q01", "q50", "q99"):
|
|
||||||
if stat in exact[key]:
|
|
||||||
np.testing.assert_allclose(
|
|
||||||
sampled[key][stat],
|
|
||||||
exact[key][stat],
|
|
||||||
rtol=1e-6,
|
|
||||||
atol=1e-6,
|
|
||||||
err_msg=f"numeric feature '{key}' stat '{stat}' changed under sampling",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_image_sampling_reduces_data_but_keeps_stats_close(tmp_path, lerobot_dataset_factory):
|
|
||||||
"""For images, sampling should reduce the number of samples considered while
|
|
||||||
keeping the resulting statistics close to the exact ones."""
|
|
||||||
dataset = lerobot_dataset_factory(
|
|
||||||
root=tmp_path / "ds", total_episodes=2, total_frames=400, use_videos=False
|
|
||||||
)
|
|
||||||
|
|
||||||
exact = compute_quantile_stats_for_dataset(dataset, use_sampling=False)
|
|
||||||
sampled = compute_quantile_stats_for_dataset(dataset, use_sampling=True)
|
|
||||||
|
|
||||||
image_keys = _image_keys(dataset)
|
|
||||||
assert image_keys, "fixture should expose at least one image feature"
|
|
||||||
for key in image_keys:
|
|
||||||
# sampling actually looked at fewer pixels
|
|
||||||
assert sampled[key]["count"][0] < exact[key]["count"][0]
|
|
||||||
# but per-channel mean stays close
|
|
||||||
np.testing.assert_allclose(
|
|
||||||
sampled[key]["mean"],
|
|
||||||
exact[key]["mean"],
|
|
||||||
rtol=0.15,
|
|
||||||
err_msg=f"image feature '{key}' mean drifted too far under sampling",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_short_episodes_use_all_frames(tmp_path, lerobot_dataset_factory):
|
|
||||||
"""With episodes shorter than the sampling floor, sampling is a no-op and
|
|
||||||
must produce exactly the same stats as the exact path."""
|
|
||||||
dataset = lerobot_dataset_factory(
|
|
||||||
root=tmp_path / "ds", total_episodes=2, total_frames=40, use_videos=False
|
|
||||||
)
|
|
||||||
|
|
||||||
exact = compute_quantile_stats_for_dataset(dataset, use_sampling=False)
|
|
||||||
sampled = compute_quantile_stats_for_dataset(dataset, use_sampling=True)
|
|
||||||
|
|
||||||
for key in _image_keys(dataset):
|
|
||||||
assert sampled[key]["count"][0] == exact[key]["count"][0]
|
|
||||||
|
|
||||||
|
|
||||||
def test_quantile_stats_present_after_compute(tmp_path, lerobot_dataset_factory):
|
|
||||||
"""The computed stats should contain quantile keys for the dataset."""
|
|
||||||
dataset = lerobot_dataset_factory(
|
|
||||||
root=tmp_path / "ds", total_episodes=2, total_frames=200, use_videos=False
|
|
||||||
)
|
|
||||||
stats = compute_quantile_stats_for_dataset(dataset, use_sampling=True)
|
|
||||||
assert has_quantile_stats(stats)
|
|
||||||
@@ -204,38 +204,6 @@ def test_clear_resets_buffer(tmp_path):
|
|||||||
assert dataset.writer.episode_buffer["size"] == 0
|
assert dataset.writer.episode_buffer["size"] == 0
|
||||||
|
|
||||||
|
|
||||||
def test_clear_removes_video_frame_staging_dir(tmp_path):
|
|
||||||
"""clear_episode_buffer() removes PNG staging dirs for video features."""
|
|
||||||
video_key = "observation.images.cam"
|
|
||||||
features = {
|
|
||||||
video_key: {
|
|
||||||
"dtype": "video",
|
|
||||||
"shape": (64, 96, 3),
|
|
||||||
"names": ["height", "width", "channels"],
|
|
||||||
},
|
|
||||||
"action": {"dtype": "float32", "shape": (2,), "names": None},
|
|
||||||
}
|
|
||||||
dataset = LeRobotDataset.create(
|
|
||||||
repo_id=DUMMY_REPO_ID,
|
|
||||||
fps=DEFAULT_FPS,
|
|
||||||
features=features,
|
|
||||||
root=tmp_path / "ds",
|
|
||||||
use_videos=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
dataset.add_frame(_make_frame(features))
|
|
||||||
video_staging_dir = (
|
|
||||||
dataset.root
|
|
||||||
/ Path(DEFAULT_IMAGE_PATH.format(image_key=video_key, episode_index=0, frame_index=0)).parent
|
|
||||||
)
|
|
||||||
assert video_staging_dir.is_dir()
|
|
||||||
|
|
||||||
dataset.clear_episode_buffer()
|
|
||||||
|
|
||||||
assert dataset.writer.episode_buffer["size"] == 0
|
|
||||||
assert not video_staging_dir.exists()
|
|
||||||
|
|
||||||
|
|
||||||
def test_finalize_is_idempotent(tmp_path):
|
def test_finalize_is_idempotent(tmp_path):
|
||||||
"""Calling finalize() twice does not raise."""
|
"""Calling finalize() twice does not raise."""
|
||||||
dataset = LeRobotDataset.create(
|
dataset = LeRobotDataset.create(
|
||||||
|
|||||||
@@ -28,17 +28,9 @@ from lerobot.scripts.lerobot_imgtransform_viz import (
|
|||||||
save_each_transform,
|
save_each_transform,
|
||||||
)
|
)
|
||||||
from lerobot.transforms import (
|
from lerobot.transforms import (
|
||||||
CoarseDropout,
|
|
||||||
GammaCorrection,
|
|
||||||
GaussianNoise,
|
|
||||||
GaussianPatchBrightness,
|
|
||||||
ImageTransformConfig,
|
ImageTransformConfig,
|
||||||
ImageTransforms,
|
ImageTransforms,
|
||||||
ImageTransformsConfig,
|
ImageTransformsConfig,
|
||||||
JPEGCompression,
|
|
||||||
MotionBlur,
|
|
||||||
PlanckianJitter,
|
|
||||||
RandomShadow,
|
|
||||||
RandomSubsetApply,
|
RandomSubsetApply,
|
||||||
SharpnessJitter,
|
SharpnessJitter,
|
||||||
make_transform_from_config,
|
make_transform_from_config,
|
||||||
@@ -463,153 +455,3 @@ def test_save_each_transform(img_tensor_factory, tmp_path):
|
|||||||
assert (transform_dir / file_name).exists(), (
|
assert (transform_dir / file_name).exists(), (
|
||||||
f"{file_name} was not found in {transform} directory."
|
f"{file_name} was not found in {transform} directory."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# --- Tests for robotics-relevant augmentations ---
|
|
||||||
|
|
||||||
ROBOTICS_TRANSFORMS = [
|
|
||||||
("GaussianNoise", GaussianNoise, {"std": (5.0, 25.0)}),
|
|
||||||
("MotionBlur", MotionBlur, {"kernel_size": (3, 11)}),
|
|
||||||
("JPEGCompression", JPEGCompression, {"quality": (15, 75)}),
|
|
||||||
("GaussianPatchBrightness", GaussianPatchBrightness, {}),
|
|
||||||
("RandomShadow", RandomShadow, {"opacity": (0.3, 0.6)}),
|
|
||||||
("CoarseDropout", CoarseDropout, {"max_holes": 8}),
|
|
||||||
("GammaCorrection", GammaCorrection, {"gamma": (0.5, 2.0)}),
|
|
||||||
("PlanckianJitter", PlanckianJitter, {"temperature": (3_000, 15_000)}),
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
|
||||||
def test_robotics_transform_shape_preserved(name, cls, kwargs, img_tensor_factory):
|
|
||||||
img = img_tensor_factory()
|
|
||||||
tf = cls(**kwargs)
|
|
||||||
out = tf(img)
|
|
||||||
assert out.shape == img.shape, f"{name} changed shape: {img.shape} -> {out.shape}"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
|
||||||
def test_robotics_transform_output_range(name, cls, kwargs, img_tensor_factory):
|
|
||||||
img = img_tensor_factory()
|
|
||||||
tf = cls(**kwargs)
|
|
||||||
out = tf(img)
|
|
||||||
assert out.min() >= -0.01, f"{name} min below range: {out.min():.4f}"
|
|
||||||
assert out.max() <= 1.01, f"{name} max above range: {out.max():.4f}"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
|
||||||
def test_robotics_transform_float_output(name, cls, kwargs, img_tensor_factory):
|
|
||||||
img = img_tensor_factory()
|
|
||||||
tf = cls(**kwargs)
|
|
||||||
out = tf(img)
|
|
||||||
assert out.is_floating_point(), f"{name} output dtype={out.dtype}"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
|
||||||
def test_robotics_transform_non_float_passthrough(name, cls, kwargs):
|
|
||||||
int_img = torch.randint(0, 255, (3, 32, 32), dtype=torch.uint8)
|
|
||||||
tf = cls(**kwargs)
|
|
||||||
out = tf(int_img)
|
|
||||||
assert torch.equal(out, int_img), f"{name} modified non-float input"
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
|
||||||
def test_robotics_transform_via_config(name, cls, kwargs):
|
|
||||||
cfg = ImageTransformConfig(type=name, kwargs=kwargs)
|
|
||||||
tf = make_transform_from_config(cfg)
|
|
||||||
assert isinstance(tf, cls), f"Config produced {type(tf)}, expected {cls}"
|
|
||||||
|
|
||||||
|
|
||||||
def test_make_transform_error_message_includes_custom():
|
|
||||||
"""Error message should list all registered custom transforms."""
|
|
||||||
with pytest.raises(ValueError, match="GaussianNoise"):
|
|
||||||
make_transform_from_config(ImageTransformConfig(type="NonExistent"))
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("name,cls,kwargs", ROBOTICS_TRANSFORMS, ids=[t[0] for t in ROBOTICS_TRANSFORMS])
|
|
||||||
@pytest.mark.parametrize("shape", [(4, 3, 32, 32), (2, 4, 3, 16, 16)])
|
|
||||||
def test_robotics_transform_supports_temporal_batches(name, cls, kwargs, shape):
|
|
||||||
img = torch.rand(shape)
|
|
||||||
out = cls(**kwargs)(img)
|
|
||||||
assert out.shape == img.shape, f"{name} changed shape: {img.shape} -> {out.shape}"
|
|
||||||
assert out.min() >= 0
|
|
||||||
assert out.max() <= 1
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"cls,kwargs",
|
|
||||||
[
|
|
||||||
(GaussianNoise, {"std": (25.0, 25.0)}),
|
|
||||||
(MotionBlur, {"kernel_size": 5}),
|
|
||||||
(JPEGCompression, {"quality": 10}),
|
|
||||||
(
|
|
||||||
GaussianPatchBrightness,
|
|
||||||
{"num_patches": 1, "sigma_range": (0.2, 0.2), "factor_range": (0.5, 0.5)},
|
|
||||||
),
|
|
||||||
(RandomShadow, {"opacity": 0.5}),
|
|
||||||
(CoarseDropout, {"max_holes": 1, "fill_value": 0.0}),
|
|
||||||
(GammaCorrection, {"gamma": (2.0, 2.0)}),
|
|
||||||
(PlanckianJitter, {"temperature": 3_000}),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_robotics_transform_is_not_silent_noop(cls, kwargs):
|
|
||||||
img = torch.rand(3, 32, 32)
|
|
||||||
out = cls(**kwargs)(img)
|
|
||||||
assert not torch.equal(out, img)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"transform",
|
|
||||||
[
|
|
||||||
GaussianNoise(std=25),
|
|
||||||
RandomShadow(opacity=0.5),
|
|
||||||
CoarseDropout(max_holes=4),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_robotics_transform_random_params_are_reused(transform):
|
|
||||||
img = torch.rand(3, 32, 32)
|
|
||||||
params = transform.make_params([img])
|
|
||||||
torch.testing.assert_close(transform.transform(img, params), transform.transform(img, params))
|
|
||||||
|
|
||||||
|
|
||||||
def test_motion_blur_kernel_size_stays_in_configured_range():
|
|
||||||
transform = MotionBlur(kernel_size=(4, 10))
|
|
||||||
sampled_sizes = {transform.make_params([])["kernel_size"] for _ in range(100)}
|
|
||||||
assert sampled_sizes <= {5, 7, 9}
|
|
||||||
assert sampled_sizes
|
|
||||||
|
|
||||||
|
|
||||||
def test_gamma_correction_scalar_below_one_defines_symmetric_range():
|
|
||||||
transform = GammaCorrection(gamma=0.5)
|
|
||||||
assert transform.gamma == (0.5, 2.0)
|
|
||||||
assert transform(torch.rand(3, 8, 8)).shape == (3, 8, 8)
|
|
||||||
|
|
||||||
|
|
||||||
def test_planckian_jitter_uses_correlated_temperature_coefficients():
|
|
||||||
img = torch.full((2, 3, 8, 8), 0.25)
|
|
||||||
out = PlanckianJitter(temperature=3_000)(img)
|
|
||||||
torch.testing.assert_close(out[:, 1], img[:, 1])
|
|
||||||
assert torch.all(out[:, 0] > out[:, 1])
|
|
||||||
assert torch.all(out[:, 2] < out[:, 1])
|
|
||||||
|
|
||||||
|
|
||||||
def test_random_shadow_supports_small_images():
|
|
||||||
img = torch.rand(3, 7, 7)
|
|
||||||
assert RandomShadow()(img).shape == img.shape
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"cls,kwargs",
|
|
||||||
[
|
|
||||||
(GaussianNoise, {"std": (-1.0, 1.0)}),
|
|
||||||
(MotionBlur, {"kernel_size": 4}),
|
|
||||||
(JPEGCompression, {"quality": (0, 75)}),
|
|
||||||
(GaussianPatchBrightness, {"sigma_range": (0.0, 0.25)}),
|
|
||||||
(RandomShadow, {"opacity": (0.3, 1.1)}),
|
|
||||||
(CoarseDropout, {"max_holes": 0}),
|
|
||||||
(GammaCorrection, {"gamma": 0.0}),
|
|
||||||
(PlanckianJitter, {"temperature": (2_000, 6_500)}),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_robotics_transform_rejects_invalid_config(cls, kwargs):
|
|
||||||
with pytest.raises(ValueError):
|
|
||||||
cls(**kwargs)
|
|
||||||
|
|||||||
@@ -482,20 +482,6 @@ def test_add_frame_works_in_write_mode(tmp_path):
|
|||||||
# ── Resume mode ──────────────────────────────────────────────────────
|
# ── Resume mode ──────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
def test_resume_freshly_created_empty_dataset(tmp_path):
|
|
||||||
"""resume() accepts a local dataset created before any episode was recorded."""
|
|
||||||
root = tmp_path / "resume_empty_ds"
|
|
||||||
LeRobotDataset.create(repo_id=DUMMY_REPO_ID, fps=DEFAULT_FPS, features=SIMPLE_FEATURES, root=root)
|
|
||||||
|
|
||||||
resumed = LeRobotDataset.resume(repo_id=DUMMY_REPO_ID, root=root)
|
|
||||||
|
|
||||||
assert isinstance(resumed.writer, DatasetWriter)
|
|
||||||
assert resumed.meta.total_episodes == 0
|
|
||||||
assert resumed.meta.total_frames == 0
|
|
||||||
assert resumed.meta.tasks is None
|
|
||||||
assert resumed.meta.episodes is None
|
|
||||||
|
|
||||||
|
|
||||||
def test_resume_creates_writer(tmp_path):
|
def test_resume_creates_writer(tmp_path):
|
||||||
"""After resume(), writer is a DatasetWriter."""
|
"""After resume(), writer is a DatasetWriter."""
|
||||||
root = tmp_path / "resume_ds"
|
root = tmp_path / "resume_ds"
|
||||||
|
|||||||
@@ -294,19 +294,6 @@ def test__sync_read(addr, length, ids_values, mock_motors, dummy_motors):
|
|||||||
assert read_values == ids_values
|
assert read_values == ids_values
|
||||||
|
|
||||||
|
|
||||||
def test__sync_read_retries_after_transient_failure(mock_motors, dummy_motors):
|
|
||||||
addr, length, ids_values = (10, 4, {1: 1337})
|
|
||||||
stub = mock_motors.build_sync_read_stub(addr, length, ids_values, num_invalid_try=1)
|
|
||||||
bus = FeetechMotorsBus(port=mock_motors.port, motors=dummy_motors)
|
|
||||||
bus.connect(handshake=False)
|
|
||||||
|
|
||||||
read_values, read_comm = bus._sync_read(addr, length, list(ids_values), num_retry=1)
|
|
||||||
|
|
||||||
assert read_comm == scs.COMM_SUCCESS
|
|
||||||
assert read_values == ids_values
|
|
||||||
assert mock_motors.stubs[stub].calls == 2
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("raise_on_error", (True, False))
|
@pytest.mark.parametrize("raise_on_error", (True, False))
|
||||||
def test__sync_read_comm(raise_on_error, mock_motors, dummy_motors):
|
def test__sync_read_comm(raise_on_error, mock_motors, dummy_motors):
|
||||||
addr, length, ids_values = (10, 4, {1: 1337})
|
addr, length, ids_values = (10, 4, {1: 1337})
|
||||||
|
|||||||
@@ -1,83 +0,0 @@
|
|||||||
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
|
||||||
#
|
|
||||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
||||||
# you may not use this file except in compliance with the License.
|
|
||||||
# You may obtain a copy of the License at
|
|
||||||
#
|
|
||||||
# http://www.apache.org/licenses/LICENSE-2.0
|
|
||||||
#
|
|
||||||
# Unless required by applicable law or agreed to in writing, software
|
|
||||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
||||||
# See the License for the specific language governing permissions and
|
|
||||||
# limitations under the License.
|
|
||||||
|
|
||||||
from types import SimpleNamespace
|
|
||||||
from unittest.mock import MagicMock
|
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
import lerobot.policies.factory as policy_factory
|
|
||||||
|
|
||||||
|
|
||||||
def test_make_policy_keeps_peft_adapter_and_base_revisions_separate(monkeypatch):
|
|
||||||
cfg = SimpleNamespace(
|
|
||||||
type="mock",
|
|
||||||
device="cpu",
|
|
||||||
pretrained_path="user/adapter",
|
|
||||||
pretrained_revision="adapter-sha",
|
|
||||||
use_peft=True,
|
|
||||||
input_features={},
|
|
||||||
output_features={},
|
|
||||||
)
|
|
||||||
dataset_meta = SimpleNamespace(features={}, stats={})
|
|
||||||
|
|
||||||
base_policy = torch.nn.Linear(1, 1)
|
|
||||||
policy_from_pretrained = MagicMock(return_value=base_policy)
|
|
||||||
policy_class = SimpleNamespace(from_pretrained=policy_from_pretrained)
|
|
||||||
monkeypatch.setattr(policy_factory, "get_policy_class", lambda _: policy_class)
|
|
||||||
monkeypatch.setattr(policy_factory, "dataset_to_policy_features", lambda _: {})
|
|
||||||
monkeypatch.setattr(policy_factory, "validate_visual_features_consistency", lambda *args: None)
|
|
||||||
|
|
||||||
peft_config = SimpleNamespace(
|
|
||||||
base_model_name_or_path="user/base-policy",
|
|
||||||
revision="base-sha",
|
|
||||||
)
|
|
||||||
peft_config_from_pretrained = MagicMock(return_value=peft_config)
|
|
||||||
adapted_policy = torch.nn.Linear(1, 1)
|
|
||||||
peft_model_from_pretrained = MagicMock(return_value=adapted_policy)
|
|
||||||
require_package = MagicMock()
|
|
||||||
monkeypatch.setattr(policy_factory, "require_package", require_package)
|
|
||||||
monkeypatch.setattr(
|
|
||||||
policy_factory,
|
|
||||||
"PeftConfig",
|
|
||||||
SimpleNamespace(from_pretrained=peft_config_from_pretrained),
|
|
||||||
)
|
|
||||||
monkeypatch.setattr(
|
|
||||||
policy_factory,
|
|
||||||
"PeftModel",
|
|
||||||
SimpleNamespace(from_pretrained=peft_model_from_pretrained),
|
|
||||||
)
|
|
||||||
|
|
||||||
policy = policy_factory.make_policy(cfg, ds_meta=dataset_meta)
|
|
||||||
|
|
||||||
assert policy is adapted_policy
|
|
||||||
require_package.assert_called_once_with("peft", extra="peft")
|
|
||||||
peft_config_from_pretrained.assert_called_once_with(
|
|
||||||
"user/adapter",
|
|
||||||
revision="adapter-sha",
|
|
||||||
)
|
|
||||||
policy_from_pretrained.assert_called_once_with(
|
|
||||||
config=cfg,
|
|
||||||
dataset_stats=dataset_meta.stats,
|
|
||||||
dataset_meta=dataset_meta,
|
|
||||||
pretrained_name_or_path="user/base-policy",
|
|
||||||
revision="base-sha",
|
|
||||||
)
|
|
||||||
peft_model_from_pretrained.assert_called_once_with(
|
|
||||||
base_policy,
|
|
||||||
"user/adapter",
|
|
||||||
config=peft_config,
|
|
||||||
revision="adapter-sha",
|
|
||||||
is_trainable=True,
|
|
||||||
)
|
|
||||||
@@ -113,7 +113,6 @@ def test_gaussian_actor_config_default_initialization():
|
|||||||
# Concurrency configuration
|
# Concurrency configuration
|
||||||
assert config.concurrency.actor == "threads"
|
assert config.concurrency.actor == "threads"
|
||||||
assert config.concurrency.learner == "threads"
|
assert config.concurrency.learner == "threads"
|
||||||
assert config.concurrency.multiprocessing_context == "spawn"
|
|
||||||
|
|
||||||
assert isinstance(config.actor_network_kwargs, ActorNetworkConfig)
|
assert isinstance(config.actor_network_kwargs, ActorNetworkConfig)
|
||||||
assert isinstance(config.policy_kwargs, PolicyConfig)
|
assert isinstance(config.policy_kwargs, PolicyConfig)
|
||||||
@@ -153,7 +152,6 @@ def test_concurrency_config():
|
|||||||
config = ConcurrencyConfig()
|
config = ConcurrencyConfig()
|
||||||
assert config.actor == "threads"
|
assert config.actor == "threads"
|
||||||
assert config.learner == "threads"
|
assert config.learner == "threads"
|
||||||
assert config.multiprocessing_context == "spawn"
|
|
||||||
|
|
||||||
|
|
||||||
def test_gaussian_actor_config_custom_initialization():
|
def test_gaussian_actor_config_custom_initialization():
|
||||||
|
|||||||
@@ -26,17 +26,8 @@ import tempfile
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
import torch
|
|
||||||
from safetensors.torch import save_file
|
|
||||||
|
|
||||||
from lerobot.configs import PipelineFeatureType, PolicyFeature
|
from lerobot.processor.pipeline import DataProcessorPipeline, ProcessorMigrationError
|
||||||
from lerobot.processor.pipeline import (
|
|
||||||
DataProcessorPipeline,
|
|
||||||
ProcessorMigrationError,
|
|
||||||
ProcessorStep,
|
|
||||||
ProcessorStepRegistry,
|
|
||||||
)
|
|
||||||
from lerobot.types import EnvTransition
|
|
||||||
|
|
||||||
# Simplified Config Loading Tests
|
# Simplified Config Loading Tests
|
||||||
|
|
||||||
@@ -107,140 +98,6 @@ def test_load_config_nonexistent_path_tries_hub():
|
|||||||
DataProcessorPipeline._load_config("nonexistent/path", "processor.json", {})
|
DataProcessorPipeline._load_config("nonexistent/path", "processor.json", {})
|
||||||
|
|
||||||
|
|
||||||
def test_from_pretrained_local_directory_missing_state_does_not_call_hub(monkeypatch):
|
|
||||||
"""Local processor dirs must fail locally when a state file is missing."""
|
|
||||||
|
|
||||||
@ProcessorStepRegistry.register("local_missing_state_step")
|
|
||||||
class LocalMissingStateStep(ProcessorStep):
|
|
||||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
|
||||||
return transition
|
|
||||||
|
|
||||||
def transform_features(
|
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
|
||||||
return features
|
|
||||||
|
|
||||||
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
|
||||||
pass
|
|
||||||
|
|
||||||
try:
|
|
||||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
||||||
tmp_path = Path(tmp_dir)
|
|
||||||
config = {
|
|
||||||
"name": "LocalMissingStatePipeline",
|
|
||||||
"steps": [{"registry_name": "local_missing_state_step", "state_file": "missing.safetensors"}],
|
|
||||||
}
|
|
||||||
(tmp_path / "processor.json").write_text(json.dumps(config))
|
|
||||||
|
|
||||||
def fail_hub_download(*args, **kwargs):
|
|
||||||
pytest.fail("local missing processor state should not call hf_hub_download")
|
|
||||||
|
|
||||||
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fail_hub_download)
|
|
||||||
|
|
||||||
with pytest.raises(FileNotFoundError, match="missing.safetensors.*local processor pipeline"):
|
|
||||||
DataProcessorPipeline.from_pretrained(tmp_path, config_filename="processor.json")
|
|
||||||
finally:
|
|
||||||
ProcessorStepRegistry.unregister("local_missing_state_step")
|
|
||||||
|
|
||||||
|
|
||||||
def test_from_pretrained_local_config_file_missing_state_does_not_call_hub(monkeypatch):
|
|
||||||
"""Local single-file processor configs must also keep missing state resolution local."""
|
|
||||||
|
|
||||||
@ProcessorStepRegistry.register("local_file_missing_state_step")
|
|
||||||
class LocalFileMissingStateStep(ProcessorStep):
|
|
||||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
|
||||||
return transition
|
|
||||||
|
|
||||||
def transform_features(
|
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
|
||||||
return features
|
|
||||||
|
|
||||||
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
|
||||||
pass
|
|
||||||
|
|
||||||
try:
|
|
||||||
with tempfile.TemporaryDirectory() as tmp_dir:
|
|
||||||
tmp_path = Path(tmp_dir)
|
|
||||||
config_path = tmp_path / "processor.json"
|
|
||||||
config = {
|
|
||||||
"name": "LocalFileMissingStatePipeline",
|
|
||||||
"steps": [
|
|
||||||
{"registry_name": "local_file_missing_state_step", "state_file": "missing.safetensors"}
|
|
||||||
],
|
|
||||||
}
|
|
||||||
config_path.write_text(json.dumps(config))
|
|
||||||
|
|
||||||
def fail_hub_download(*args, **kwargs):
|
|
||||||
pytest.fail("local missing processor state should not call hf_hub_download")
|
|
||||||
|
|
||||||
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fail_hub_download)
|
|
||||||
|
|
||||||
with pytest.raises(FileNotFoundError, match="missing.safetensors.*local processor pipeline"):
|
|
||||||
DataProcessorPipeline.from_pretrained(config_path, config_filename="ignored.json")
|
|
||||||
finally:
|
|
||||||
ProcessorStepRegistry.unregister("local_file_missing_state_step")
|
|
||||||
|
|
||||||
|
|
||||||
def test_from_pretrained_hub_source_missing_local_state_still_calls_hub(monkeypatch, tmp_path):
|
|
||||||
"""Hub sources still fall back to hf_hub_download for state files."""
|
|
||||||
|
|
||||||
@ProcessorStepRegistry.register("hub_state_step")
|
|
||||||
class HubStateStep(ProcessorStep):
|
|
||||||
def __init__(self):
|
|
||||||
self.value = torch.tensor(0)
|
|
||||||
|
|
||||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
|
||||||
return transition
|
|
||||||
|
|
||||||
def transform_features(
|
|
||||||
self, features: dict[PipelineFeatureType, dict[str, PolicyFeature]]
|
|
||||||
) -> dict[PipelineFeatureType, dict[str, PolicyFeature]]:
|
|
||||||
return features
|
|
||||||
|
|
||||||
def load_state_dict(self, state: dict[str, torch.Tensor]) -> None:
|
|
||||||
self.value = state["value"]
|
|
||||||
|
|
||||||
try:
|
|
||||||
state_path = tmp_path / "downloaded.safetensors"
|
|
||||||
save_file({"value": torch.tensor(7)}, state_path)
|
|
||||||
loaded_config = {
|
|
||||||
"name": "HubStatePipeline",
|
|
||||||
"steps": [{"registry_name": "hub_state_step", "state_file": "hub_state.safetensors"}],
|
|
||||||
}
|
|
||||||
calls = []
|
|
||||||
|
|
||||||
def fake_load_config(cls, model_id, config_filename, hub_download_kwargs):
|
|
||||||
return loaded_config, tmp_path / "hub_cache"
|
|
||||||
|
|
||||||
def fake_hub_download(**kwargs):
|
|
||||||
calls.append(kwargs)
|
|
||||||
return str(state_path)
|
|
||||||
|
|
||||||
monkeypatch.setattr(DataProcessorPipeline, "_load_config", classmethod(fake_load_config))
|
|
||||||
monkeypatch.setattr("lerobot.processor.pipeline.hf_hub_download", fake_hub_download)
|
|
||||||
|
|
||||||
pipeline = DataProcessorPipeline.from_pretrained("user/repo", config_filename="processor.json")
|
|
||||||
|
|
||||||
assert calls == [
|
|
||||||
{
|
|
||||||
"repo_id": "user/repo",
|
|
||||||
"filename": "hub_state.safetensors",
|
|
||||||
"repo_type": "model",
|
|
||||||
"force_download": False,
|
|
||||||
"resume_download": None,
|
|
||||||
"proxies": None,
|
|
||||||
"token": None,
|
|
||||||
"cache_dir": None,
|
|
||||||
"local_files_only": False,
|
|
||||||
"revision": None,
|
|
||||||
}
|
|
||||||
]
|
|
||||||
assert pipeline.steps[0].value.item() == 7
|
|
||||||
finally:
|
|
||||||
ProcessorStepRegistry.unregister("hub_state_step")
|
|
||||||
|
|
||||||
|
|
||||||
# Config Validation Tests
|
# Config Validation Tests
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ def _make_bus_mock() -> MagicMock:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.fixture
|
@pytest.fixture
|
||||||
def follower(tmp_path):
|
def follower():
|
||||||
bus_mock = _make_bus_mock()
|
bus_mock = _make_bus_mock()
|
||||||
|
|
||||||
def _bus_side_effect(*_args, **kwargs):
|
def _bus_side_effect(*_args, **kwargs):
|
||||||
@@ -71,7 +71,7 @@ def follower(tmp_path):
|
|||||||
),
|
),
|
||||||
patch.object(SO100Follower, "configure", lambda self: None),
|
patch.object(SO100Follower, "configure", lambda self: None),
|
||||||
):
|
):
|
||||||
cfg = SO100FollowerConfig(port="/dev/null", calibration_dir=tmp_path)
|
cfg = SO100FollowerConfig(port="/dev/null")
|
||||||
robot = SO100Follower(cfg)
|
robot = SO100Follower(cfg)
|
||||||
yield robot
|
yield robot
|
||||||
if robot.is_connected:
|
if robot.is_connected:
|
||||||
@@ -99,27 +99,6 @@ def test_get_observation(follower):
|
|||||||
assert obs[f"{motor}.pos"] == idx
|
assert obs[f"{motor}.pos"] == idx
|
||||||
|
|
||||||
|
|
||||||
def test_get_observation_uses_read_retries(follower):
|
|
||||||
# Feetech buses can intermittently fail a sync_read; the follower should forward the configured
|
|
||||||
# retry count so transient failures don't abort the control loop (see #3131).
|
|
||||||
follower.config.num_read_retries = 7
|
|
||||||
follower.connect()
|
|
||||||
follower.get_observation()
|
|
||||||
|
|
||||||
follower.bus.sync_read.assert_called_once_with("Present_Position", num_retry=7)
|
|
||||||
|
|
||||||
|
|
||||||
def test_send_action_uses_read_retries(follower):
|
|
||||||
follower.config.max_relative_target = 10.0
|
|
||||||
follower.config.num_read_retries = 7
|
|
||||||
follower.connect()
|
|
||||||
|
|
||||||
action = {f"{motor}.pos": value * 10 for value, motor in enumerate(follower.bus.motors, 1)}
|
|
||||||
follower.send_action(action)
|
|
||||||
|
|
||||||
follower.bus.sync_read.assert_called_once_with("Present_Position", num_retry=7)
|
|
||||||
|
|
||||||
|
|
||||||
def test_send_action(follower):
|
def test_send_action(follower):
|
||||||
follower.connect()
|
follower.connect()
|
||||||
|
|
||||||
|
|||||||
@@ -17,8 +17,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import sys
|
|
||||||
from types import SimpleNamespace
|
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
@@ -108,116 +106,6 @@ def test_sentry_config_defaults():
|
|||||||
assert cfg.target_video_file_size_mb is None
|
assert cfg.target_video_file_size_mb is None
|
||||||
|
|
||||||
|
|
||||||
def test_rollout_config_passes_policy_pretrained_revision(monkeypatch):
|
|
||||||
from lerobot.configs import PreTrainedConfig, parser
|
|
||||||
from lerobot.rollout import RolloutConfig
|
|
||||||
from tests.mocks.mock_robot import MockRobotConfig
|
|
||||||
|
|
||||||
captured = {}
|
|
||||||
|
|
||||||
def fake_from_pretrained(cls, pretrained_name_or_path, **kwargs):
|
|
||||||
captured["pretrained_name_or_path"] = pretrained_name_or_path
|
|
||||||
captured.update(kwargs)
|
|
||||||
return SimpleNamespace(device="cpu", pretrained_revision=kwargs["revision"])
|
|
||||||
|
|
||||||
monkeypatch.setattr(parser, "get_yaml_overrides", lambda _: ["--pretrained_revision=yaml-sha"])
|
|
||||||
monkeypatch.setattr(
|
|
||||||
sys,
|
|
||||||
"argv",
|
|
||||||
["lerobot-rollout", "--policy.path=user/policy", "--policy.pretrained_revision=cli-sha"],
|
|
||||||
)
|
|
||||||
monkeypatch.setattr(PreTrainedConfig, "from_pretrained", classmethod(fake_from_pretrained))
|
|
||||||
|
|
||||||
cfg = RolloutConfig(robot=MockRobotConfig())
|
|
||||||
|
|
||||||
assert captured["pretrained_name_or_path"] == "user/policy"
|
|
||||||
assert captured["revision"] == "cli-sha"
|
|
||||||
assert captured["cli_overrides"] == [
|
|
||||||
"--pretrained_revision=yaml-sha",
|
|
||||||
"--pretrained_revision=cli-sha",
|
|
||||||
]
|
|
||||||
assert cfg.policy.pretrained_path == "user/policy"
|
|
||||||
assert cfg.policy.pretrained_revision == "cli-sha"
|
|
||||||
|
|
||||||
|
|
||||||
def test_load_pretrained_policy_passes_revision(monkeypatch):
|
|
||||||
import lerobot.rollout.context as rollout_context
|
|
||||||
|
|
||||||
policy_config = SimpleNamespace(
|
|
||||||
type="mock",
|
|
||||||
use_peft=False,
|
|
||||||
pretrained_path="user/policy",
|
|
||||||
pretrained_revision="policy-sha",
|
|
||||||
)
|
|
||||||
policy_class = MagicMock()
|
|
||||||
loaded_policy = MagicMock()
|
|
||||||
policy_class.from_pretrained.return_value = loaded_policy
|
|
||||||
monkeypatch.setattr(rollout_context, "get_policy_class", lambda _: policy_class)
|
|
||||||
|
|
||||||
policy = rollout_context._load_pretrained_policy(policy_config)
|
|
||||||
|
|
||||||
assert policy is loaded_policy
|
|
||||||
policy_class.from_pretrained.assert_called_once_with(
|
|
||||||
"user/policy",
|
|
||||||
config=policy_config,
|
|
||||||
revision="policy-sha",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_load_pretrained_peft_policy_keeps_adapter_and_base_revisions_separate(monkeypatch):
|
|
||||||
import lerobot.rollout.context as rollout_context
|
|
||||||
|
|
||||||
policy_config = SimpleNamespace(
|
|
||||||
type="mock",
|
|
||||||
use_peft=True,
|
|
||||||
pretrained_path="user/adapter",
|
|
||||||
pretrained_revision="adapter-sha",
|
|
||||||
)
|
|
||||||
policy_class = MagicMock()
|
|
||||||
base_policy = MagicMock()
|
|
||||||
policy_class.from_pretrained.return_value = base_policy
|
|
||||||
monkeypatch.setattr(rollout_context, "get_policy_class", lambda _: policy_class)
|
|
||||||
|
|
||||||
peft_config = SimpleNamespace(
|
|
||||||
base_model_name_or_path="user/base-policy",
|
|
||||||
revision="base-sha",
|
|
||||||
)
|
|
||||||
peft_config_from_pretrained = MagicMock(return_value=peft_config)
|
|
||||||
adapted_policy = MagicMock()
|
|
||||||
peft_model_from_pretrained = MagicMock(return_value=adapted_policy)
|
|
||||||
require_package = MagicMock()
|
|
||||||
monkeypatch.setattr(rollout_context, "require_package", require_package)
|
|
||||||
monkeypatch.setattr(
|
|
||||||
rollout_context,
|
|
||||||
"PeftConfig",
|
|
||||||
SimpleNamespace(from_pretrained=peft_config_from_pretrained),
|
|
||||||
raising=False,
|
|
||||||
)
|
|
||||||
monkeypatch.setattr(
|
|
||||||
rollout_context,
|
|
||||||
"PeftModel",
|
|
||||||
SimpleNamespace(from_pretrained=peft_model_from_pretrained),
|
|
||||||
raising=False,
|
|
||||||
)
|
|
||||||
|
|
||||||
policy = rollout_context._load_pretrained_policy(policy_config)
|
|
||||||
|
|
||||||
assert policy is adapted_policy
|
|
||||||
require_package.assert_called_once_with("peft", extra="peft")
|
|
||||||
peft_config_from_pretrained.assert_called_once_with("user/adapter", revision="adapter-sha")
|
|
||||||
policy_class.from_pretrained.assert_called_once_with(
|
|
||||||
pretrained_name_or_path="user/base-policy",
|
|
||||||
config=policy_config,
|
|
||||||
revision="base-sha",
|
|
||||||
)
|
|
||||||
peft_model_from_pretrained.assert_called_once_with(
|
|
||||||
base_policy,
|
|
||||||
"user/adapter",
|
|
||||||
config=peft_config,
|
|
||||||
revision="adapter-sha",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# RolloutRingBuffer
|
# RolloutRingBuffer
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user