mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-29 12:39:41 +00:00
Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 0f8aa7d03b |
Binary file not shown.
|
Before Width: | Height: | Size: 51 KiB |
@@ -1,138 +0,0 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
|
|
||||||
"""
|
|
||||||
Example demonstrating how to use the ActionTokenizerProcessorStep to tokenize actions.
|
|
||||||
|
|
||||||
This example shows how to:
|
|
||||||
1. Load a dataset with action data
|
|
||||||
2. Apply the action tokenizer processor to tokenize actions with proper padding/truncation
|
|
||||||
3. Access both the tokenized actions and the attention mask
|
|
||||||
4. Decode tokenized actions back to their original form
|
|
||||||
"""
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from transformers import AutoProcessor
|
|
||||||
|
|
||||||
from lerobot.datasets.lerobot_dataset import LeRobotDataset
|
|
||||||
from lerobot.processor.core import EnvTransition, TransitionKey
|
|
||||||
from lerobot.processor.tokenizer_processor import ActionTokenizerProcessorStep
|
|
||||||
from lerobot.utils.constants import ACTION_TOKEN_MASK
|
|
||||||
|
|
||||||
# Define delta timestamps for the dataset
|
|
||||||
delta_timestamps = {
|
|
||||||
'action': [
|
|
||||||
0.0, 0.03333333333333333, 0.06666666666666667, 0.1, 0.13333333333333333,
|
|
||||||
0.16666666666666666, 0.2, 0.23333333333333334, 0.26666666666666666, 0.3,
|
|
||||||
0.3333333333333333, 0.36666666666666664, 0.4, 0.43333333333333335,
|
|
||||||
0.4666666666666667, 0.5, 0.5333333333333333, 0.5666666666666667, 0.6,
|
|
||||||
0.6333333333333333, 0.6666666666666666, 0.7, 0.7333333333333333,
|
|
||||||
0.7666666666666667, 0.8, 0.8333333333333334, 0.8666666666666667, 0.9,
|
|
||||||
0.9333333333333333, 0.9666666666666667, 1.0, 1.0333333333333334,
|
|
||||||
1.0666666666666667, 1.1, 1.1333333333333333, 1.1666666666666667, 1.2,
|
|
||||||
1.2333333333333334, 1.2666666666666666, 1.3, 1.3333333333333333,
|
|
||||||
1.3666666666666667, 1.4, 1.4333333333333333, 1.4666666666666666, 1.5,
|
|
||||||
1.5333333333333334, 1.5666666666666667, 1.6, 1.6333333333333333
|
|
||||||
]
|
|
||||||
}
|
|
||||||
|
|
||||||
# Load the dataset
|
|
||||||
print("Loading dataset...")
|
|
||||||
dataset = LeRobotDataset(
|
|
||||||
repo_id="local",
|
|
||||||
root="/fsx/jade_choghari/outputs/pgen_annotations1",
|
|
||||||
delta_timestamps=delta_timestamps
|
|
||||||
)
|
|
||||||
|
|
||||||
# Create a dataloader
|
|
||||||
dataloader = torch.utils.data.DataLoader(
|
|
||||||
dataset,
|
|
||||||
num_workers=0,
|
|
||||||
batch_size=4,
|
|
||||||
shuffle=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Get a batch of data
|
|
||||||
batch = next(iter(dataloader))
|
|
||||||
action_data = batch["action"] # Shape: (batch_size, action_horizon, action_dim)
|
|
||||||
|
|
||||||
print(f"\nOriginal action shape: {action_data.shape}")
|
|
||||||
print(f"Original action data (first sample, first timestep):\n{action_data[0, 0]}")
|
|
||||||
|
|
||||||
# Method 1: Using the tokenizer directly (as in fast_tokenize.py)
|
|
||||||
print("\n" + "="*80)
|
|
||||||
print("Method 1: Direct tokenizer usage")
|
|
||||||
print("="*80)
|
|
||||||
|
|
||||||
tokenizer = AutoProcessor.from_pretrained("physical-intelligence/fast", trust_remote_code=True)
|
|
||||||
|
|
||||||
# Tokenize directly
|
|
||||||
tokens = tokenizer(action_data)
|
|
||||||
print(f"\nDirect tokenization result type: {type(tokens)}")
|
|
||||||
print(f"Tokens shape/length: {tokens.shape if isinstance(tokens, torch.Tensor) else len(tokens)}")
|
|
||||||
|
|
||||||
# Decode
|
|
||||||
decoded_actions = tokenizer.decode(tokens)
|
|
||||||
print(f"Decoded actions shape: {decoded_actions.shape}")
|
|
||||||
reconstruction_error = torch.abs(action_data - decoded_actions).mean()
|
|
||||||
print(f"Mean absolute reconstruction error: {reconstruction_error.item():.6f}")
|
|
||||||
|
|
||||||
# Method 2: Using the ActionTokenizerProcessorStep with proper padding/truncation
|
|
||||||
print("\n" + "="*80)
|
|
||||||
print("Method 2: Using ActionTokenizerProcessorStep (with padding & mask)")
|
|
||||||
print("="*80)
|
|
||||||
|
|
||||||
# Create the action tokenizer processor step
|
|
||||||
action_tokenizer_processor = ActionTokenizerProcessorStep(
|
|
||||||
tokenizer_name="physical-intelligence/fast",
|
|
||||||
trust_remote_code=True,
|
|
||||||
max_action_tokens=32, # Maximum number of tokens per action
|
|
||||||
)
|
|
||||||
|
|
||||||
# Create a transition with the action data
|
|
||||||
transition = {
|
|
||||||
TransitionKey.ACTION: action_data,
|
|
||||||
TransitionKey.OBSERVATION: {}, # Empty for this example
|
|
||||||
}
|
|
||||||
|
|
||||||
# Apply the processor
|
|
||||||
processed_transition = action_tokenizer_processor(transition)
|
|
||||||
|
|
||||||
# Extract tokenized actions and mask
|
|
||||||
tokenized_actions = processed_transition[TransitionKey.ACTION]
|
|
||||||
complementary_data = processed_transition[TransitionKey.COMPLEMENTARY_DATA]
|
|
||||||
action_mask = complementary_data[ACTION_TOKEN_MASK]
|
|
||||||
|
|
||||||
print(f"\nTokenized actions shape: {tokenized_actions.shape}") # (batch_size, max_action_tokens)
|
|
||||||
print(f"Action mask shape: {action_mask.shape}") # (batch_size, max_action_tokens)
|
|
||||||
print(f"Tokenized actions dtype: {tokenized_actions.dtype}")
|
|
||||||
print(f"Action mask dtype: {action_mask.dtype}")
|
|
||||||
|
|
||||||
# Show token statistics
|
|
||||||
print(f"\nFirst sample tokens: {tokenized_actions[0]}")
|
|
||||||
print(f"First sample mask: {action_mask[0]}")
|
|
||||||
num_real_tokens = action_mask[0].sum().item()
|
|
||||||
print(f"Number of real tokens (non-padding): {num_real_tokens}")
|
|
||||||
print(f"Number of padding tokens: {action_mask.shape[1] - num_real_tokens}")
|
|
||||||
|
|
||||||
# Decode using the mask
|
|
||||||
print("\nDecoding tokenized actions...")
|
|
||||||
decoded_with_processor = tokenizer.decode(tokenized_actions)
|
|
||||||
print(f"Decoded actions shape: {decoded_with_processor.shape}")
|
|
||||||
|
|
||||||
# Calculate reconstruction error
|
|
||||||
reconstruction_error_processor = torch.abs(action_data - decoded_with_processor).mean()
|
|
||||||
print(f"Mean absolute reconstruction error: {reconstruction_error_processor.item():.6f}")
|
|
||||||
|
|
||||||
# Show that masking works correctly
|
|
||||||
print("\n" + "="*80)
|
|
||||||
print("Mask demonstration")
|
|
||||||
print("="*80)
|
|
||||||
for i in range(min(4, tokenized_actions.shape[0])):
|
|
||||||
mask_i = action_mask[i]
|
|
||||||
num_real = mask_i.sum().item()
|
|
||||||
print(f"Sample {i}: {num_real} real tokens, {len(mask_i) - num_real} padding tokens")
|
|
||||||
|
|
||||||
print("\n" + "="*80)
|
|
||||||
print("Action tokenization example completed successfully!")
|
|
||||||
print("="*80)
|
|
||||||
|
|
||||||
@@ -1402,13 +1402,6 @@ def main():
|
|||||||
action="store_true",
|
action="store_true",
|
||||||
help="Push modified dataset to HuggingFace Hub",
|
help="Push modified dataset to HuggingFace Hub",
|
||||||
)
|
)
|
||||||
# add image key
|
|
||||||
parser.add_argument(
|
|
||||||
"--image-key",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help="Image observation key to use for image mode (default: None)",
|
|
||||||
)
|
|
||||||
|
|
||||||
args = parser.parse_args()
|
args = parser.parse_args()
|
||||||
console = Console()
|
console = Console()
|
||||||
@@ -1450,9 +1443,6 @@ def main():
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Get image keys (for image mode)
|
# Get image keys (for image mode)
|
||||||
if args.image_key:
|
|
||||||
image_keys = [args.image_key]
|
|
||||||
else:
|
|
||||||
image_keys = dataset.meta.camera_keys[:args.num_image_views_per_sample]
|
image_keys = dataset.meta.camera_keys[:args.num_image_views_per_sample]
|
||||||
if not args.video_mode:
|
if not args.video_mode:
|
||||||
console.print(f"[cyan]Using image keys: {image_keys}[/cyan]")
|
console.print(f"[cyan]Using image keys: {image_keys}[/cyan]")
|
||||||
|
|||||||
@@ -1,25 +0,0 @@
|
|||||||
import numpy as np
|
|
||||||
from transformers import AutoProcessor
|
|
||||||
import torch
|
|
||||||
from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
|
|
||||||
|
|
||||||
delta_timestamps = {'action': [0.0, 0.03333333333333333, 0.06666666666666667, 0.1, 0.13333333333333333, 0.16666666666666666, 0.2, 0.23333333333333334, 0.26666666666666666, 0.3, 0.3333333333333333, 0.36666666666666664, 0.4, 0.43333333333333335, 0.4666666666666667, 0.5, 0.5333333333333333, 0.5666666666666667, 0.6, 0.6333333333333333, 0.6666666666666666, 0.7, 0.7333333333333333, 0.7666666666666667, 0.8, 0.8333333333333334, 0.8666666666666667, 0.9, 0.9333333333333333, 0.9666666666666667, 1.0, 1.0333333333333334, 1.0666666666666667, 1.1, 1.1333333333333333, 1.1666666666666667, 1.2, 1.2333333333333334, 1.2666666666666666, 1.3, 1.3333333333333333, 1.3666666666666667, 1.4, 1.4333333333333333, 1.4666666666666666, 1.5, 1.5333333333333334, 1.5666666666666667, 1.6, 1.6333333333333333]}
|
|
||||||
dataset = LeRobotDataset(repo_id="local", root="/fsx/jade_choghari/outputs/pgen_annotations1", delta_timestamps=delta_timestamps)
|
|
||||||
|
|
||||||
dataloader = torch.utils.data.DataLoader(
|
|
||||||
dataset,
|
|
||||||
num_workers=0,
|
|
||||||
batch_size=4,
|
|
||||||
shuffle=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
batch = next(iter(dataloader))
|
|
||||||
|
|
||||||
# Load the tokenizer from the Hugging Face hub
|
|
||||||
tokenizer = AutoProcessor.from_pretrained("physical-intelligence/fast", trust_remote_code=True)
|
|
||||||
|
|
||||||
# Tokenize & decode action chunks (we use dummy data here)
|
|
||||||
action_data = batch["action"] # one batch of action chunks
|
|
||||||
tokens = tokenizer(action_data) # tokens = list[int]
|
|
||||||
decoded_actions = tokenizer.decode(tokens)
|
|
||||||
print("tokenized actions: ", tokens)
|
|
||||||
@@ -10,19 +10,17 @@ from lerobot.policies.factory import make_policy, make_policy_config
|
|||||||
from lerobot.configs.policies import PreTrainedConfig
|
from lerobot.configs.policies import PreTrainedConfig
|
||||||
|
|
||||||
cfg = PreTrainedConfig.from_pretrained(
|
cfg = PreTrainedConfig.from_pretrained(
|
||||||
pretrained_name_or_path="/fsx/jade_choghari/outputs/pi0_training/checkpoints/last/pretrained_model",
|
pretrained_name_or_path="/fsx/jade_choghari/outputs/pi0_training_new/checkpoints/last/pretrained_model",
|
||||||
)
|
)
|
||||||
cfg.dtype = "bfloat16"
|
cfg.dtype = "bfloat16"
|
||||||
|
|
||||||
pre_processor, post_processor = make_pre_post_processors(
|
pre_processor, post_processor = make_pre_post_processors(
|
||||||
policy_cfg=cfg,
|
policy_cfg=cfg,
|
||||||
pretrained_path="/fsx/jade_choghari/outputs/pi0_training/checkpoints/last/pretrained_model",
|
pretrained_path="/fsx/jade_choghari/outputs/pi0_training_new/checkpoints/last/pretrained_model",
|
||||||
)
|
)
|
||||||
|
|
||||||
delta_timestamps = {'action': [0.0, 0.03333333333333333, 0.06666666666666667, 0.1, 0.13333333333333333, 0.16666666666666666, 0.2, 0.23333333333333334, 0.26666666666666666, 0.3, 0.3333333333333333, 0.36666666666666664, 0.4, 0.43333333333333335, 0.4666666666666667, 0.5, 0.5333333333333333, 0.5666666666666667, 0.6, 0.6333333333333333, 0.6666666666666666, 0.7, 0.7333333333333333, 0.7666666666666667, 0.8, 0.8333333333333334, 0.8666666666666667, 0.9, 0.9333333333333333, 0.9666666666666667, 1.0, 1.0333333333333334, 1.0666666666666667, 1.1, 1.1333333333333333, 1.1666666666666667, 1.2, 1.2333333333333334, 1.2666666666666666, 1.3, 1.3333333333333333, 1.3666666666666667, 1.4, 1.4333333333333333, 1.4666666666666666, 1.5, 1.5333333333333334, 1.5666666666666667, 1.6, 1.6333333333333333]}
|
|
||||||
|
|
||||||
dataset = LeRobotDataset(repo_id="local", root="/fsx/jade_choghari/outputs/pgen_annotations1", delta_timestamps=delta_timestamps)
|
|
||||||
|
|
||||||
|
dataset = LeRobotDataset(repo_id="local", root="/fsx/jade_choghari/outputs/pgen_annotations1")
|
||||||
# rename map --rename_map='{
|
# rename map --rename_map='{
|
||||||
# "observation.images.side": "observation.images.base_0_rgb",
|
# "observation.images.side": "observation.images.base_0_rgb",
|
||||||
# "observation.images.up": "observation.images.left_wrist_0_rgb"
|
# "observation.images.up": "observation.images.left_wrist_0_rgb"
|
||||||
@@ -45,47 +43,16 @@ dataloader = torch.utils.data.DataLoader(
|
|||||||
)
|
)
|
||||||
|
|
||||||
batch = next(iter(dataloader))
|
batch = next(iter(dataloader))
|
||||||
|
|
||||||
batch = pre_processor(batch)
|
batch = pre_processor(batch)
|
||||||
|
|
||||||
|
# Test training forward pass
|
||||||
policy.train()
|
policy.train()
|
||||||
# run inference
|
|
||||||
# action = policy.select_action(batch)
|
|
||||||
loss, loss_dict = policy.forward(batch)
|
loss, loss_dict = policy.forward(batch)
|
||||||
breakpoint()
|
print(f"Training loss: {loss_dict}")
|
||||||
# import requests
|
|
||||||
# from PIL import Image
|
|
||||||
# from transformers import AutoProcessor
|
|
||||||
# model = policy.model.paligemma_with_expert.paligemma
|
|
||||||
# model = model.to(device="cuda", dtype=torch.bfloat16)
|
|
||||||
# model.eval()
|
|
||||||
# prompt = "Describe this image."
|
|
||||||
# url = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/pipeline-cat-chonk.jpeg"
|
|
||||||
# image = Image.open(requests.get(url, stream=True).raw)
|
|
||||||
# processor = AutoProcessor.from_pretrained(
|
|
||||||
# "google/paligemma-3b-pt-224",
|
|
||||||
# )
|
|
||||||
# inputs = processor(image, prompt, return_tensors="pt").to(model.device)
|
|
||||||
# print("generating...")
|
|
||||||
# output = model.generate(
|
|
||||||
# **inputs,
|
|
||||||
# max_new_tokens=50,
|
|
||||||
# use_cache=True, # default dynamic cache
|
|
||||||
# )
|
|
||||||
# print(processor.decode(output[0], skip_special_tokens=True))
|
|
||||||
|
|
||||||
|
# Test inference
|
||||||
# # other model
|
policy.eval()
|
||||||
# from transformers import PaliGemmaForConditionalGeneration
|
with torch.no_grad():
|
||||||
# model = PaliGemmaForConditionalGeneration.from_pretrained(
|
actions = policy.predict_action_chunk(batch)
|
||||||
# "google/paligemma2-3b-pt-224",
|
print(f"Predicted actions shape: {actions.shape}")
|
||||||
# torch_dtype=torch.bfloat16,
|
|
||||||
# device_map="auto",
|
|
||||||
# )
|
|
||||||
# model.eval()
|
|
||||||
# print("generating...")
|
|
||||||
# output = model.generate(
|
|
||||||
# **inputs,
|
|
||||||
# max_new_tokens=100,
|
|
||||||
# use_cache=True, # default dynamic cache
|
|
||||||
# )
|
|
||||||
# print("Model 2 output:")
|
|
||||||
# print(processor.decode(output[0], skip_special_tokens=True))
|
|
||||||
@@ -1,18 +0,0 @@
|
|||||||
import torch
|
|
||||||
from huggingface_hub import HfApi
|
|
||||||
|
|
||||||
import lerobot
|
|
||||||
from lerobot.datasets.lerobot_dataset import LeRobotDataset, LeRobotDatasetMetadata
|
|
||||||
|
|
||||||
dataset = LeRobotDataset(repo_id="lerobot/libero")
|
|
||||||
|
|
||||||
dataloader = torch.utils.data.DataLoader(
|
|
||||||
dataset,
|
|
||||||
num_workers=0,
|
|
||||||
batch_size=4,
|
|
||||||
shuffle=True,
|
|
||||||
)
|
|
||||||
batch = next(iter(dataloader))
|
|
||||||
print(batch.keys())
|
|
||||||
|
|
||||||
breakpoint()
|
|
||||||
@@ -1,159 +0,0 @@
|
|||||||
## One-sentence answer
|
|
||||||
|
|
||||||
> `make_att_2d_masks(prefix_pad_masks, prefix_att_masks)` builds the **actual 2D attention mask** `[B, L, L]` that tells the transformer **which token positions may attend to which others**, combining **padding** and **causality**.
|
|
||||||
|
|
||||||
Everything else you’ve seen so far was just metadata.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## What goes in
|
|
||||||
|
|
||||||
### Inputs
|
|
||||||
|
|
||||||
```python
|
|
||||||
prefix_pad_masks # shape [B, L]
|
|
||||||
prefix_att_masks # shape [B, L]
|
|
||||||
```
|
|
||||||
|
|
||||||
Where:
|
|
||||||
|
|
||||||
* `prefix_pad_masks[b, i] = True`
|
|
||||||
→ token `i` exists (not padding)
|
|
||||||
|
|
||||||
* `prefix_att_masks[b, i] = False`
|
|
||||||
→ token `i` is **bidirectional**
|
|
||||||
|
|
||||||
* `prefix_att_masks[b, i] = True`
|
|
||||||
→ token `i` is **causal (autoregressive)**
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## What comes out
|
|
||||||
|
|
||||||
```python
|
|
||||||
att_2d_prefix # shape [B, L, L]
|
|
||||||
```
|
|
||||||
|
|
||||||
Each entry:
|
|
||||||
|
|
||||||
```text
|
|
||||||
att_2d_prefix[b, i, j] = True
|
|
||||||
```
|
|
||||||
|
|
||||||
means:
|
|
||||||
|
|
||||||
> “In batch `b`, **token i (query)** is allowed to attend to **token j (key)**.”
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## How it is constructed (conceptually)
|
|
||||||
|
|
||||||
For **each batch b**, **each query position i**, **each key position j**:
|
|
||||||
|
|
||||||
```python
|
|
||||||
if not prefix_pad_masks[b, j]:
|
|
||||||
att[b, i, j] = False # cannot attend to padding
|
|
||||||
else if not prefix_att_masks[b, i]:
|
|
||||||
att[b, i, j] = True # bidirectional token → can see all real tokens
|
|
||||||
else:
|
|
||||||
att[b, i, j] = (j <= i) # causal token → can see only past + itself
|
|
||||||
```
|
|
||||||
|
|
||||||
That’s it.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Tiny concrete example (exactly matching your code)
|
|
||||||
|
|
||||||
Suppose:
|
|
||||||
|
|
||||||
```python
|
|
||||||
prefix_pad_masks[0] = [T, T, T, T, T, F]
|
|
||||||
prefix_att_masks[0] = [F, F, F, T, T, T]
|
|
||||||
```
|
|
||||||
|
|
||||||
Tokens:
|
|
||||||
|
|
||||||
```
|
|
||||||
0: IMG
|
|
||||||
1: IMG
|
|
||||||
2: LANG
|
|
||||||
3: SUB0
|
|
||||||
4: SUB1
|
|
||||||
5: PAD
|
|
||||||
```
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### Resulting `att_2d_prefix[0]`
|
|
||||||
|
|
||||||
`✓ = True, ✗ = False`
|
|
||||||
|
|
||||||
| Q \ K | 0 | 1 | 2 | 3 | 4 | 5 |
|
|
||||||
| ---------- | - | - | - | - | - | - |
|
|
||||||
| 0 (bi) | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ |
|
|
||||||
| 1 (bi) | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ |
|
|
||||||
| 2 (bi) | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ |
|
|
||||||
| 3 (causal) | ✓ | ✓ | ✓ | ✓ | ✗ | ✗ |
|
|
||||||
| 4 (causal) | ✓ | ✓ | ✓ | ✓ | ✓ | ✗ |
|
|
||||||
| 5 (pad) | ✗ | ✗ | ✗ | ✗ | ✗ | ✗ |
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Why this matters for your training code
|
|
||||||
|
|
||||||
This line:
|
|
||||||
|
|
||||||
```python
|
|
||||||
att_2d_prefix_4d = self._prepare_attention_masks_4d(att_2d_prefix)
|
|
||||||
```
|
|
||||||
|
|
||||||
Converts `[B, L, L] → [B, 1, L, L]` and possibly flips True/False to `0/-inf`.
|
|
||||||
|
|
||||||
This is **exactly what Paligemma uses inside self-attention**.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## Key implications (VERY important)
|
|
||||||
|
|
||||||
### 1️⃣ This mask does **not isolate token groups**
|
|
||||||
|
|
||||||
* Bidirectional tokens can attend to **everything**
|
|
||||||
* Causal tokens only restrict *their own row*
|
|
||||||
|
|
||||||
So **flow/action tokens must be blocked separately**.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### 2️⃣ This is why your AR subtask prediction works
|
|
||||||
|
|
||||||
* Subtask tokens are causal
|
|
||||||
* Output at position `i` predicts token `i+1`
|
|
||||||
* Padding is fully ignored
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
### 3️⃣ Inference behavior
|
|
||||||
|
|
||||||
When `subtask_tokens = None`:
|
|
||||||
|
|
||||||
* `prefix_att_masks` contains only `False`
|
|
||||||
* `att_2d_prefix` becomes **fully bidirectional**
|
|
||||||
* No AR behavior remains
|
|
||||||
|
|
||||||
Exactly what you want.
|
|
||||||
|
|
||||||
---
|
|
||||||
|
|
||||||
## One-sentence takeaway (commit this)
|
|
||||||
|
|
||||||
> `make_att_2d_masks` fuses **padding** and **causality** into a concrete `[B, L, L]` attention matrix that the transformer actually uses.
|
|
||||||
|
|
||||||
If you want next, I can:
|
|
||||||
|
|
||||||
* inspect `make_att_2d_masks()` source with you
|
|
||||||
* show how to block **flow → subtask** attention
|
|
||||||
* explain how this changes when suffix tokens are added
|
|
||||||
* help you refactor this into a cleaner “grouped attention” API
|
|
||||||
|
|
||||||
You’re now at the point where the model’s behavior should feel *predictable*, not magical.
|
|
||||||
@@ -1,11 +1,10 @@
|
|||||||
python examples/dataset/annotate.py \
|
|
||||||
--repo-id jadechoghari/collect-data \
|
|
||||||
--video-key observation.images.base \
|
|
||||||
--model Qwen/Qwen3-VL-30B-A3B-Instruct \
|
|
||||||
--episodes 16 22
|
|
||||||
|
|
||||||
# python examples/dataset/annotate.py \
|
# python examples/dataset/annotate.py \
|
||||||
# --repo-id lerobot/svla_so101_pickplace \
|
# --repo-id lerobot/svla_so101_pickplace \
|
||||||
# --video-key observation.images.side \
|
# --video-key observation.images.side \
|
||||||
# --model Qwen/Qwen3-VL-30B-A3B-Instruct \
|
# --model Qwen/Qwen3-VL-30B-A3B-Instruct \
|
||||||
# --episodes 5
|
|
||||||
|
python examples/dataset/annotate.py \
|
||||||
|
--repo-id lerobot/svla_so101_pickplace \
|
||||||
|
--video-key observation.images.side \
|
||||||
|
--model Qwen/Qwen3-VL-30B-A3B-Instruct \
|
||||||
|
--episodes 3 5 7 44
|
||||||
@@ -4,12 +4,12 @@
|
|||||||
# This generates user prompts and robot utterances for hierarchical policy training
|
# This generates user prompts and robot utterances for hierarchical policy training
|
||||||
|
|
||||||
# Configuration
|
# Configuration
|
||||||
REPO_ID="jadechoghari/collect-data"
|
REPO_ID="lerobot/svla_so101_pickplace"
|
||||||
MODEL="Qwen/Qwen3-VL-30B-A3B-Instruct"
|
MODEL="Qwen/Qwen3-VL-30B-A3B-Instruct"
|
||||||
# Alternative: MODEL="Qwen/Qwen2-VL-7B-Instruct"
|
# Alternative: MODEL="Qwen/Qwen2-VL-7B-Instruct"
|
||||||
|
|
||||||
|
|
||||||
OUTPUT_DIR="/fsx/jade_choghari/outputs/collect-data-pgen"
|
OUTPUT_DIR="/fsx/jade_choghari/outputs/pgen_annotations1"
|
||||||
BATCH_SIZE=32
|
BATCH_SIZE=32
|
||||||
TEMPERATURE=0.9
|
TEMPERATURE=0.9
|
||||||
SAMPLE_INTERVAL=5.0 # Generate dialogue every 1 second (all episodes processed)
|
SAMPLE_INTERVAL=5.0 # Generate dialogue every 1 second (all episodes processed)
|
||||||
@@ -22,7 +22,6 @@ python examples/dataset/annotate_pgen.py \
|
|||||||
--temperature "$TEMPERATURE" \
|
--temperature "$TEMPERATURE" \
|
||||||
--batch-size "$BATCH_SIZE" \
|
--batch-size "$BATCH_SIZE" \
|
||||||
--sample-interval "$SAMPLE_INTERVAL" \
|
--sample-interval "$SAMPLE_INTERVAL" \
|
||||||
--image-key observation.images.base \
|
|
||||||
--num-image-views-per-sample 1
|
--num-image-views-per-sample 1
|
||||||
|
|
||||||
# For faster testing, increase sample interval:
|
# For faster testing, increase sample interval:
|
||||||
|
|||||||
@@ -1 +0,0 @@
|
|||||||
srun --time 12:00:00 --qos=high --gres=gpu:1 --mem=24G --partition=hopper-prod --container-image /fsx/michel_aractingi/docker_images/huggingface+lerobot-gpu+dev.sqsh --container-mounts /fsx/jade_choghari
|
|
||||||
@@ -1,47 +0,0 @@
|
|||||||
# Voice Assistant Examples
|
|
||||||
|
|
||||||
Voice-enabled robot assistant examples using speech-to-text (STT), and text-to-speech (TTS).
|
|
||||||
|
|
||||||
## Overview
|
|
||||||
|
|
||||||
These examples demonstrate how to build a voice interface for robot control:
|
|
||||||
|
|
||||||
1. **Hold SPACE** → Push-to-talk recording starts
|
|
||||||
2. **Release SPACE** → Recording stops
|
|
||||||
3. **STT (Whisper)** → Converts speech to text (high-level task prompt)
|
|
||||||
4. **Pi0.5** → Generates robot response/utterance
|
|
||||||
5. **TTS (Kokoro)** → Speaks the response back
|
|
||||||
|
|
||||||
## Requirements
|
|
||||||
|
|
||||||
```bash
|
|
||||||
pip install torch transformers sounddevice numpy pynput kokoro>=0.9.2
|
|
||||||
```
|
|
||||||
|
|
||||||
## Usage
|
|
||||||
|
|
||||||
### With Pi0.5 Model
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python examples/voice_assistant/voice_assistant_pi05.py \
|
|
||||||
--pretrained_path path/to/pi05/checkpoint
|
|
||||||
```
|
|
||||||
|
|
||||||
## How It Works
|
|
||||||
|
|
||||||
### Pi0.5 Voice Integration
|
|
||||||
|
|
||||||
Pi0.5 can generate robot utterances as part of its subtask prediction. The flow:
|
|
||||||
|
|
||||||
1. **High-level prompt**: User voice command is transcribed and formatted as a task prompt
|
|
||||||
2. **Subtask generation**: Pi0.5 autoregressively generates a response
|
|
||||||
3. **Utterance extraction**: If the response contains `<utterance>...</utterance>` tags, the content is extracted
|
|
||||||
4. **TTS output**: The response is spoken back to the user
|
|
||||||
|
|
||||||
## Configuration Options
|
|
||||||
|
|
||||||
| Option | Default | Description |
|
|
||||||
|--------|---------|-------------|
|
|
||||||
| `--pretrained_path` | None | Path to Pi0.5 checkpoint |
|
|
||||||
| `--record_seconds` | 5.0 | Audio recording duration |
|
|
||||||
| `--max_response_tokens` | 100 | Max tokens in generated response |
|
|
||||||
@@ -1,336 +0,0 @@
|
|||||||
#!/usr/bin/env python
|
|
||||||
"""
|
|
||||||
Voice Assistant with Pi0.5: Microphone → STT → Pi0.5 → TTS → Speaker
|
|
||||||
|
|
||||||
This example demonstrates how to use Pi0.5 as a conversational robot assistant:
|
|
||||||
1. Hold SPACE to record your voice command
|
|
||||||
2. Speech-to-text (Whisper) converts speech to text
|
|
||||||
3. Text is fed as a high-level prompt to Pi0.5
|
|
||||||
4. Pi0.5 generates a response (robot utterance)
|
|
||||||
5. Text-to-speech (Kokoro) speaks the response back
|
|
||||||
|
|
||||||
Requirements:
|
|
||||||
pip install torch transformers sounddevice numpy pynput kokoro>=0.9.2
|
|
||||||
|
|
||||||
Usage:
|
|
||||||
python examples/voice_assistant/voice_assistant_pi05.py \
|
|
||||||
--pretrained_path lerobot/pi0.5-base
|
|
||||||
"""
|
|
||||||
|
|
||||||
import os
|
|
||||||
|
|
||||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
|
||||||
|
|
||||||
import argparse
|
|
||||||
import re
|
|
||||||
import subprocess
|
|
||||||
import threading
|
|
||||||
import time
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
import sounddevice as sd
|
|
||||||
import torch
|
|
||||||
from pynput import keyboard
|
|
||||||
from transformers import AutoTokenizer, WhisperForConditionalGeneration, WhisperProcessor
|
|
||||||
|
|
||||||
from lerobot.policies.pi05.configuration_pi05 import PI05Config
|
|
||||||
from lerobot.policies.pi05.modeling_pi05 import PI05Pytorch
|
|
||||||
|
|
||||||
SAMPLE_RATE = 16000
|
|
||||||
|
|
||||||
|
|
||||||
def get_device():
|
|
||||||
if torch.cuda.is_available():
|
|
||||||
return torch.device("cuda")
|
|
||||||
elif torch.backends.mps.is_available():
|
|
||||||
return torch.device("mps")
|
|
||||||
return torch.device("cpu")
|
|
||||||
|
|
||||||
|
|
||||||
class Pi05VoiceAssistant:
|
|
||||||
"""Voice assistant using Pi0.5 for generating robot utterances."""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
pretrained_path: str | None = None,
|
|
||||||
max_response_tokens: int = 100,
|
|
||||||
max_record_seconds: float = 30.0,
|
|
||||||
):
|
|
||||||
self.device = get_device()
|
|
||||||
self.dtype = torch.float32 if self.device.type == "mps" else torch.bfloat16
|
|
||||||
self.max_response_tokens = max_response_tokens
|
|
||||||
self.max_record_seconds = max_record_seconds
|
|
||||||
|
|
||||||
# Push-to-talk state
|
|
||||||
self._recording = False
|
|
||||||
self._audio_chunks: list[np.ndarray] = []
|
|
||||||
self._stream: sd.InputStream | None = None
|
|
||||||
|
|
||||||
print(f"Using device: {self.device}")
|
|
||||||
self._load_models(pretrained_path)
|
|
||||||
|
|
||||||
def _load_models(self, pretrained_path: str | None):
|
|
||||||
print("Loading STT (Whisper tiny)...")
|
|
||||||
self.stt_processor = WhisperProcessor.from_pretrained("openai/whisper-tiny.en")
|
|
||||||
self.stt_model = WhisperForConditionalGeneration.from_pretrained(
|
|
||||||
"openai/whisper-tiny.en", torch_dtype=self.dtype
|
|
||||||
).to(self.device)
|
|
||||||
|
|
||||||
print("Loading Pi0.5 model...")
|
|
||||||
self._load_pi05(pretrained_path)
|
|
||||||
|
|
||||||
print("Loading tokenizer...")
|
|
||||||
self.tokenizer = AutoTokenizer.from_pretrained("google/paligemma-3b-pt-224")
|
|
||||||
|
|
||||||
self._load_tts()
|
|
||||||
print("Ready!\n")
|
|
||||||
|
|
||||||
def _load_pi05(self, pretrained_path: str | None):
|
|
||||||
"""Load Pi0.5 model for utterance generation."""
|
|
||||||
config = PI05Config()
|
|
||||||
config.dtype = "float32" if self.device.type == "mps" else "bfloat16"
|
|
||||||
|
|
||||||
self.pi05_model = PI05Pytorch(config)
|
|
||||||
|
|
||||||
if pretrained_path:
|
|
||||||
try:
|
|
||||||
from safetensors.torch import load_file
|
|
||||||
state_dict = load_file(f"{pretrained_path}/model.safetensors")
|
|
||||||
self.pi05_model.load_state_dict(state_dict, strict=False)
|
|
||||||
print(f"✓ Loaded Pi0.5 weights from {pretrained_path}")
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Warning: Could not load pretrained weights: {e}")
|
|
||||||
print("Using randomly initialized model for demo purposes")
|
|
||||||
|
|
||||||
self.pi05_model = self.pi05_model.to(self.device)
|
|
||||||
self.pi05_model.eval()
|
|
||||||
|
|
||||||
def _load_tts(self):
|
|
||||||
try:
|
|
||||||
print("Loading TTS (Kokoro 82M)...")
|
|
||||||
from kokoro import KPipeline
|
|
||||||
|
|
||||||
self.tts_pipeline = KPipeline(lang_code="a") # American English
|
|
||||||
self.tts_voice = "af_heart"
|
|
||||||
self.tts_type = "kokoro"
|
|
||||||
print("Kokoro loaded!")
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Kokoro not available ({e})")
|
|
||||||
print("Using macOS `say` for TTS")
|
|
||||||
self.tts_pipeline = None
|
|
||||||
self.tts_type = "system"
|
|
||||||
|
|
||||||
def _audio_callback(self, indata, frames, time_info, status):
|
|
||||||
"""Callback for audio stream - collects chunks while recording."""
|
|
||||||
if self._recording:
|
|
||||||
self._audio_chunks.append(indata.copy())
|
|
||||||
|
|
||||||
def _start_recording(self):
|
|
||||||
"""Start recording audio."""
|
|
||||||
if self._recording:
|
|
||||||
return
|
|
||||||
self._recording = True
|
|
||||||
self._audio_chunks = []
|
|
||||||
print("🎤 Recording... (release SPACE to stop)")
|
|
||||||
|
|
||||||
def _stop_recording(self) -> np.ndarray | None:
|
|
||||||
"""Stop recording and return the audio."""
|
|
||||||
if not self._recording:
|
|
||||||
return None
|
|
||||||
self._recording = False
|
|
||||||
|
|
||||||
if not self._audio_chunks:
|
|
||||||
return None
|
|
||||||
|
|
||||||
audio = np.concatenate(self._audio_chunks, axis=0).flatten()
|
|
||||||
duration = len(audio) / SAMPLE_RATE
|
|
||||||
volume = np.abs(audio).max()
|
|
||||||
print(f"Recorded {duration:.1f}s, volume: {volume:.4f}")
|
|
||||||
|
|
||||||
if volume < 0.001:
|
|
||||||
print("⚠️ Very low audio - check microphone permissions!")
|
|
||||||
return None
|
|
||||||
|
|
||||||
return audio
|
|
||||||
|
|
||||||
def wait_for_spacebar(self) -> np.ndarray | None:
|
|
||||||
"""Wait for spacebar press, record while held, return audio on release."""
|
|
||||||
audio_result = None
|
|
||||||
recording_done = threading.Event()
|
|
||||||
|
|
||||||
def on_press(key):
|
|
||||||
if key == keyboard.Key.space:
|
|
||||||
self._start_recording()
|
|
||||||
|
|
||||||
def on_release(key):
|
|
||||||
nonlocal audio_result
|
|
||||||
if key == keyboard.Key.space and self._recording:
|
|
||||||
audio_result = self._stop_recording()
|
|
||||||
recording_done.set()
|
|
||||||
return False # Stop listener
|
|
||||||
|
|
||||||
# Start audio stream
|
|
||||||
self._stream = sd.InputStream(
|
|
||||||
samplerate=SAMPLE_RATE,
|
|
||||||
channels=1,
|
|
||||||
dtype="float32",
|
|
||||||
callback=self._audio_callback,
|
|
||||||
blocksize=int(SAMPLE_RATE * 0.1), # 100ms blocks
|
|
||||||
)
|
|
||||||
|
|
||||||
with self._stream:
|
|
||||||
print("\n⏳ Press and hold SPACE to speak...")
|
|
||||||
with keyboard.Listener(on_press=on_press, on_release=on_release) as listener:
|
|
||||||
# Wait for recording to complete or timeout
|
|
||||||
recording_done.wait(timeout=self.max_record_seconds)
|
|
||||||
if self._recording:
|
|
||||||
audio_result = self._stop_recording()
|
|
||||||
|
|
||||||
return audio_result
|
|
||||||
|
|
||||||
def transcribe(self, audio: np.ndarray) -> str:
|
|
||||||
start = time.perf_counter()
|
|
||||||
inputs = self.stt_processor(audio, sampling_rate=SAMPLE_RATE, return_tensors="pt")
|
|
||||||
input_features = inputs.input_features.to(self.device, dtype=self.dtype)
|
|
||||||
tokens = self.stt_model.generate(input_features)
|
|
||||||
text = self.stt_processor.batch_decode(tokens, skip_special_tokens=True)[0]
|
|
||||||
print(f"STT: {time.perf_counter() - start:.2f}s")
|
|
||||||
return text.strip()
|
|
||||||
|
|
||||||
def _create_dummy_images(self, batch_size: int = 1) -> tuple[list[torch.Tensor], list[torch.Tensor]]:
|
|
||||||
"""Create placeholder images for Pi0.5 when no camera is available."""
|
|
||||||
image_shape = (batch_size, 3, 224, 224)
|
|
||||||
dummy_image = torch.zeros(image_shape, dtype=torch.float32, device=self.device)
|
|
||||||
dummy_mask = torch.ones(batch_size, dtype=torch.bool, device=self.device)
|
|
||||||
return [dummy_image], [dummy_mask]
|
|
||||||
|
|
||||||
def _tokenize_prompt(self, text: str) -> tuple[torch.Tensor, torch.Tensor]:
|
|
||||||
"""Tokenize the user prompt for Pi0.5."""
|
|
||||||
prompt = f"User request: {text}\nRobot response:"
|
|
||||||
tokenized = self.tokenizer(
|
|
||||||
[prompt],
|
|
||||||
max_length=200,
|
|
||||||
truncation=True,
|
|
||||||
padding="max_length",
|
|
||||||
return_tensors="pt",
|
|
||||||
)
|
|
||||||
tokens = tokenized["input_ids"].to(self.device)
|
|
||||||
masks = tokenized["attention_mask"].to(self.device, dtype=torch.bool)
|
|
||||||
return tokens, masks
|
|
||||||
|
|
||||||
def generate_response(self, user_text: str) -> str:
|
|
||||||
"""Generate robot utterance using Pi0.5's language generation."""
|
|
||||||
start = time.perf_counter()
|
|
||||||
|
|
||||||
images, img_masks = self._create_dummy_images()
|
|
||||||
tokens, masks = self._tokenize_prompt(user_text)
|
|
||||||
|
|
||||||
with torch.no_grad():
|
|
||||||
generated_tokens = self.pi05_model._generate_subtask_tokens(
|
|
||||||
images=images,
|
|
||||||
img_masks=img_masks,
|
|
||||||
tokens=tokens,
|
|
||||||
masks=masks,
|
|
||||||
tokenizer=self.tokenizer,
|
|
||||||
max_length=self.max_response_tokens,
|
|
||||||
device=self.device,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Decode generated tokens
|
|
||||||
valid_tokens = generated_tokens[0][generated_tokens[0] != 0]
|
|
||||||
response = self.tokenizer.decode(valid_tokens, skip_special_tokens=True)
|
|
||||||
|
|
||||||
# Extract utterance if marked with special tokens
|
|
||||||
response = self._extract_utterance(response)
|
|
||||||
|
|
||||||
print(f"Pi0.5: {time.perf_counter() - start:.2f}s")
|
|
||||||
return response.strip()
|
|
||||||
|
|
||||||
def _extract_utterance(self, text: str) -> str:
|
|
||||||
"""Extract utterance from between <utterance> tokens if present."""
|
|
||||||
pattern = r"<utterance>(.*?)</utterance>"
|
|
||||||
match = re.search(pattern, text, re.DOTALL)
|
|
||||||
if match:
|
|
||||||
return match.group(1).strip()
|
|
||||||
return text
|
|
||||||
|
|
||||||
def speak(self, text: str):
|
|
||||||
start = time.perf_counter()
|
|
||||||
if self.tts_type == "kokoro":
|
|
||||||
generator = self.tts_pipeline(text, voice=self.tts_voice)
|
|
||||||
audio_chunks = [audio for _, _, audio in generator]
|
|
||||||
if audio_chunks:
|
|
||||||
audio = np.concatenate(audio_chunks)
|
|
||||||
sd.play(audio, 24000)
|
|
||||||
sd.wait()
|
|
||||||
else:
|
|
||||||
subprocess.run(["say", text], check=True)
|
|
||||||
print(f"TTS: {time.perf_counter() - start:.2f}s")
|
|
||||||
|
|
||||||
def run(self):
|
|
||||||
print("=" * 50)
|
|
||||||
print("Pi0.5 Voice Assistant")
|
|
||||||
print("=" * 50)
|
|
||||||
print("• Hold SPACE to record your voice command")
|
|
||||||
print("• Release SPACE when done speaking")
|
|
||||||
print("• Press Ctrl+C to exit")
|
|
||||||
print("=" * 50)
|
|
||||||
|
|
||||||
while True:
|
|
||||||
try:
|
|
||||||
audio = self.wait_for_spacebar()
|
|
||||||
|
|
||||||
if audio is None:
|
|
||||||
print("(no audio captured)\n")
|
|
||||||
continue
|
|
||||||
|
|
||||||
user_text = self.transcribe(audio)
|
|
||||||
|
|
||||||
if not user_text:
|
|
||||||
print("(no speech detected)\n")
|
|
||||||
continue
|
|
||||||
|
|
||||||
print(f"You: {user_text}")
|
|
||||||
|
|
||||||
response = self.generate_response(user_text)
|
|
||||||
print(f"Robot: {response}\n")
|
|
||||||
|
|
||||||
self.speak(response)
|
|
||||||
|
|
||||||
except KeyboardInterrupt:
|
|
||||||
print("\nGoodbye!")
|
|
||||||
break
|
|
||||||
|
|
||||||
|
|
||||||
def main():
|
|
||||||
parser = argparse.ArgumentParser(description="Pi0.5 Voice Assistant")
|
|
||||||
parser.add_argument(
|
|
||||||
"--pretrained_path",
|
|
||||||
type=str,
|
|
||||||
default=None,
|
|
||||||
help="Path to pretrained Pi0.5 model (optional)",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--max_response_tokens",
|
|
||||||
type=int,
|
|
||||||
default=100,
|
|
||||||
help="Maximum tokens in generated response",
|
|
||||||
)
|
|
||||||
parser.add_argument(
|
|
||||||
"--max_record_seconds",
|
|
||||||
type=float,
|
|
||||||
default=30.0,
|
|
||||||
help="Maximum recording duration in seconds",
|
|
||||||
)
|
|
||||||
args = parser.parse_args()
|
|
||||||
|
|
||||||
assistant = Pi05VoiceAssistant(
|
|
||||||
pretrained_path=args.pretrained_path,
|
|
||||||
max_response_tokens=args.max_response_tokens,
|
|
||||||
max_record_seconds=args.max_record_seconds,
|
|
||||||
)
|
|
||||||
assistant.run()
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
|
||||||
@@ -1,27 +0,0 @@
|
|||||||
{
|
|
||||||
"repo_id": "local",
|
|
||||||
"vocab_size": 1024,
|
|
||||||
"scale": 10.0,
|
|
||||||
"encoded_dims": "0:7",
|
|
||||||
"encoded_dim_ranges": [
|
|
||||||
[
|
|
||||||
0,
|
|
||||||
7
|
|
||||||
]
|
|
||||||
],
|
|
||||||
"total_encoded_dims": 7,
|
|
||||||
"delta_dims": null,
|
|
||||||
"delta_dim_list": null,
|
|
||||||
"use_delta_transform": false,
|
|
||||||
"state_key": "observation.state",
|
|
||||||
"normalization_mode": "QUANTILES",
|
|
||||||
"action_horizon": 10,
|
|
||||||
"num_training_chunks": 25065,
|
|
||||||
"compression_stats": {
|
|
||||||
"compression_ratio": 3.464660463274599,
|
|
||||||
"mean_token_length": 20.204,
|
|
||||||
"p99_token_length": 36.00999999999999,
|
|
||||||
"min_token_length": 5.0,
|
|
||||||
"max_token_length": 38.0
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,158 +0,0 @@
|
|||||||
import logging
|
|
||||||
from typing import ClassVar
|
|
||||||
|
|
||||||
import numpy as np
|
|
||||||
from scipy.fft import dct
|
|
||||||
from scipy.fft import idct
|
|
||||||
from tokenizers import ByteLevelBPETokenizer
|
|
||||||
from tokenizers.trainers import BpeTrainer
|
|
||||||
from transformers import PreTrainedTokenizerFast
|
|
||||||
from transformers.processing_utils import ProcessorMixin
|
|
||||||
|
|
||||||
|
|
||||||
class UniversalActionProcessor(ProcessorMixin):
|
|
||||||
attributes: ClassVar[list[str]] = ["bpe_tokenizer"]
|
|
||||||
bpe_tokenizer_class: str = "AutoTokenizer"
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
bpe_tokenizer: PreTrainedTokenizerFast,
|
|
||||||
scale: float = 10,
|
|
||||||
vocab_size: int = 1024,
|
|
||||||
min_token: int = 0,
|
|
||||||
*,
|
|
||||||
action_dim: int | None = None,
|
|
||||||
time_horizon: int | None = None,
|
|
||||||
):
|
|
||||||
self.scale = scale
|
|
||||||
self.vocab_size = vocab_size
|
|
||||||
self.min_token = min_token
|
|
||||||
|
|
||||||
# Action horizon and dimension needed during decoding. These can be specified
|
|
||||||
# in three ways (in order of priority):
|
|
||||||
# 1. passed in as kwargs to decode()
|
|
||||||
# 2. in the constructor
|
|
||||||
# 3. cached from the last time decode() was called
|
|
||||||
self.time_horizon = time_horizon
|
|
||||||
self.action_dim = action_dim
|
|
||||||
self.called_time_horizon = time_horizon
|
|
||||||
self.called_action_dim = action_dim
|
|
||||||
|
|
||||||
super().__init__(bpe_tokenizer)
|
|
||||||
|
|
||||||
def __call__(self, action_chunk: np.array) -> np.array:
|
|
||||||
assert action_chunk.ndim <= 3, "Only 3 dimensions supported: [batch, timesteps, action_dim]"
|
|
||||||
if action_chunk.ndim == 2:
|
|
||||||
action_chunk = action_chunk[None, ...]
|
|
||||||
|
|
||||||
# Cache the time horizon and action dimension for decoding
|
|
||||||
self.called_time_horizon = action_chunk.shape[-2]
|
|
||||||
self.called_action_dim = action_chunk.shape[-1]
|
|
||||||
|
|
||||||
dct_coeff = dct(action_chunk, axis=1, norm="ortho")
|
|
||||||
dct_coeff = np.around(dct_coeff * self.scale)
|
|
||||||
tokens = []
|
|
||||||
for elem in dct_coeff:
|
|
||||||
token_str = "".join(map(chr, np.maximum(elem.flatten() - self.min_token, 0).astype(int)))
|
|
||||||
tokens.append(self.bpe_tokenizer(token_str)["input_ids"])
|
|
||||||
return tokens
|
|
||||||
|
|
||||||
def decode(
|
|
||||||
self,
|
|
||||||
tokens: list[list[int]],
|
|
||||||
*,
|
|
||||||
time_horizon: int | None = None,
|
|
||||||
action_dim: int | None = None,
|
|
||||||
) -> np.array:
|
|
||||||
self.time_horizon = time_horizon or self.time_horizon or self.called_time_horizon
|
|
||||||
self.action_dim = action_dim or self.action_dim or self.called_action_dim
|
|
||||||
|
|
||||||
# Cache the time horizon and action dimension for the next call
|
|
||||||
self.called_time_horizon = self.time_horizon
|
|
||||||
self.called_action_dim = self.action_dim
|
|
||||||
|
|
||||||
assert (
|
|
||||||
self.time_horizon is not None and self.action_dim is not None
|
|
||||||
), "Tokenizer not initialized, call encode() once or pass in time_horizon and action_dim."
|
|
||||||
|
|
||||||
decoded_actions = []
|
|
||||||
for token in tokens:
|
|
||||||
try:
|
|
||||||
decoded_tokens = self.bpe_tokenizer.decode(token)
|
|
||||||
decoded_dct_coeff = np.array(list(map(ord, decoded_tokens))) + self.min_token
|
|
||||||
decoded_dct_coeff = decoded_dct_coeff.reshape(-1, self.action_dim)
|
|
||||||
assert (
|
|
||||||
decoded_dct_coeff.shape
|
|
||||||
== (
|
|
||||||
self.time_horizon,
|
|
||||||
self.action_dim,
|
|
||||||
)
|
|
||||||
), f"Decoded DCT coefficients have shape {decoded_dct_coeff.shape}, expected ({self.time_horizon}, {self.action_dim})"
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error decoding tokens: {e}")
|
|
||||||
print(f"Tokens: {token}")
|
|
||||||
decoded_dct_coeff = np.zeros((self.time_horizon, self.action_dim))
|
|
||||||
decoded_actions.append(idct(decoded_dct_coeff / self.scale, axis=0, norm="ortho"))
|
|
||||||
return np.stack(decoded_actions)
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def fit(
|
|
||||||
cls,
|
|
||||||
action_data: list[np.array],
|
|
||||||
scale: float = 10,
|
|
||||||
vocab_size: int = 1024,
|
|
||||||
*,
|
|
||||||
time_horizon: int | None = None,
|
|
||||||
action_dim: int | None = None,
|
|
||||||
) -> "UniversalActionProcessor":
|
|
||||||
# Run DCT over all inputs
|
|
||||||
dct_tokens = [dct(a, axis=0, norm="ortho").flatten() for a in action_data]
|
|
||||||
|
|
||||||
# Quantize and find min token
|
|
||||||
max_token = int(np.around(np.concatenate(dct_tokens) * scale).max())
|
|
||||||
min_token = int(np.around(np.concatenate(dct_tokens) * scale).min())
|
|
||||||
min_vocab_size = max_token - min_token
|
|
||||||
|
|
||||||
assert (
|
|
||||||
min_vocab_size <= vocab_size
|
|
||||||
), f"Vocab size {vocab_size} is too small for the range of tokens {min_vocab_size}"
|
|
||||||
if min_vocab_size + 100 > vocab_size:
|
|
||||||
logging.warning(
|
|
||||||
f"Initial alphabet size {min_vocab_size} is almost as large as the vocab"
|
|
||||||
f"size {vocab_size}, consider increasing vocab size"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Make token iterator for BPE training
|
|
||||||
def _token_iter():
|
|
||||||
for tokens in dct_tokens:
|
|
||||||
rounded_tokens = np.around(tokens * scale) - min_token
|
|
||||||
rounded_tokens = rounded_tokens.astype(int)
|
|
||||||
string = "".join(map(chr, rounded_tokens))
|
|
||||||
yield string
|
|
||||||
|
|
||||||
# Train BPE tokenizer
|
|
||||||
bpe = ByteLevelBPETokenizer()
|
|
||||||
|
|
||||||
# Set up the entire range of possible tokens as the initial alphabet
|
|
||||||
alphabet = [chr(i) for i in range(max_token - min_token + 1)]
|
|
||||||
trainer = BpeTrainer(
|
|
||||||
vocab_size=vocab_size,
|
|
||||||
min_frequency=2,
|
|
||||||
show_progress=True,
|
|
||||||
special_tokens=[],
|
|
||||||
initial_alphabet=alphabet,
|
|
||||||
max_token_length=10000,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Train the inner tokenizer (don't use ByteLevelBPETokenizer.train_from_iterator()
|
|
||||||
# because it doesn't support custom alphabets)
|
|
||||||
bpe._tokenizer.train_from_iterator(_token_iter(), trainer=trainer)
|
|
||||||
|
|
||||||
return cls(
|
|
||||||
PreTrainedTokenizerFast(tokenizer_object=bpe, clean_up_tokenization_spaces=False),
|
|
||||||
scale=scale,
|
|
||||||
vocab_size=vocab_size,
|
|
||||||
min_token=min_token,
|
|
||||||
time_horizon=time_horizon,
|
|
||||||
action_dim=action_dim,
|
|
||||||
)
|
|
||||||
@@ -1,11 +0,0 @@
|
|||||||
{
|
|
||||||
"action_dim": 7,
|
|
||||||
"auto_map": {
|
|
||||||
"AutoProcessor": "processing_action_tokenizer.UniversalActionProcessor"
|
|
||||||
},
|
|
||||||
"min_token": -32,
|
|
||||||
"processor_class": "UniversalActionProcessor",
|
|
||||||
"scale": 10.0,
|
|
||||||
"time_horizon": 10,
|
|
||||||
"vocab_size": 1024
|
|
||||||
}
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
{}
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,11 +0,0 @@
|
|||||||
{
|
|
||||||
"added_tokens_decoder": {},
|
|
||||||
"auto_map": {
|
|
||||||
"AutoProcessor": "processing_action_tokenizer.UniversalActionProcessor"
|
|
||||||
},
|
|
||||||
"clean_up_tokenization_spaces": false,
|
|
||||||
"extra_special_tokens": {},
|
|
||||||
"model_max_length": 1000000000000000019884624838656,
|
|
||||||
"processor_class": "UniversalActionProcessor",
|
|
||||||
"tokenizer_class": "PreTrainedTokenizerFast"
|
|
||||||
}
|
|
||||||
@@ -162,7 +162,7 @@ class LeRobotDatasetMetadata:
|
|||||||
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)
|
self.tasks = load_tasks(self.root)
|
||||||
# self.tasks_high_level = load_tasks_high_level(self.root)
|
self.tasks_high_level = load_tasks_high_level(self.root)
|
||||||
self.episodes = load_episodes(self.root)
|
self.episodes = load_episodes(self.root)
|
||||||
self.stats = load_stats(self.root)
|
self.stats = load_stats(self.root)
|
||||||
|
|
||||||
|
|||||||
@@ -1,196 +0,0 @@
|
|||||||
# FAST Tokenizer Training for LeRobotDataset
|
|
||||||
|
|
||||||
This directory contains tools for training a FAST (Factorized Action Sequence Tokenizer) on LeRobot datasets.
|
|
||||||
|
|
||||||
## Files
|
|
||||||
|
|
||||||
- **`train_fast_tokenizer.py`**: Main training script (refactored for LeRobotDataset)
|
|
||||||
- **`train_fast_tokenizer_example.md`**: Usage examples and parameter documentation
|
|
||||||
- **`MIGRATION_NOTES.md`**: Migration guide from B1K to LeRobotDataset
|
|
||||||
|
|
||||||
## Quick Start
|
|
||||||
|
|
||||||
```bash
|
|
||||||
# Basic usage
|
|
||||||
python train_fast_tokenizer.py \
|
|
||||||
--repo_id "lerobot/aloha_sim_insertion_human" \
|
|
||||||
--action_horizon 10 \
|
|
||||||
--encoded_dims "0:14"
|
|
||||||
|
|
||||||
# With delta transform
|
|
||||||
python train_fast_tokenizer.py \
|
|
||||||
--repo_id "lerobot/aloha_sim_insertion_human" \
|
|
||||||
--action_horizon 10 \
|
|
||||||
--encoded_dims "0:14" \
|
|
||||||
--delta_dims "0,1,2,3,4,5,6,7,8,9,10,11,12,13" \
|
|
||||||
--state_key "observation.state" \
|
|
||||||
--vocab_size 1024
|
|
||||||
```
|
|
||||||
|
|
||||||
## What is FAST?
|
|
||||||
|
|
||||||
FAST is a tokenizer for robotic action sequences that:
|
|
||||||
1. Applies DCT (Discrete Cosine Transform) to action chunks
|
|
||||||
2. Quantizes DCT coefficients
|
|
||||||
3. Uses BPE (Byte-Pair Encoding) to compress the quantized sequence
|
|
||||||
4. Achieves high compression ratios (e.g., 10-20x) while maintaining accuracy
|
|
||||||
|
|
||||||
This enables efficient storage and processing of long action sequences in vision-language-action models.
|
|
||||||
|
|
||||||
## Requirements
|
|
||||||
|
|
||||||
- Python 3.10+
|
|
||||||
- LeRobot dataset (either local or from HuggingFace Hub)
|
|
||||||
- transformers (for AutoProcessor)
|
|
||||||
- numpy
|
|
||||||
- torch
|
|
||||||
- tyro
|
|
||||||
|
|
||||||
## Workflow
|
|
||||||
|
|
||||||
```
|
|
||||||
LeRobotDataset → Extract Episodes → Apply Delta Transform
|
|
||||||
↓
|
|
||||||
Select Dimensions → Normalize (q01, q99) → Create Chunks
|
|
||||||
↓
|
|
||||||
Train FAST Tokenizer → Compute Stats → Save
|
|
||||||
```
|
|
||||||
|
|
||||||
## Parameters Guide
|
|
||||||
|
|
||||||
### Essential Parameters
|
|
||||||
|
|
||||||
- **`repo_id`**: HuggingFace dataset repository ID
|
|
||||||
- Example: `"lerobot/aloha_sim_insertion_human"`
|
|
||||||
|
|
||||||
- **`action_horizon`**: Length of action sequences to tokenize
|
|
||||||
- Typical: 10-16 steps
|
|
||||||
|
|
||||||
- **`encoded_dims`**: Which action dimensions to encode
|
|
||||||
- Format: `"start:end,start:end"`
|
|
||||||
- Example: `"0:7"` = dimensions 0-6
|
|
||||||
- Example: `"0:3,7:10"` = dimensions 0-2 and 7-9
|
|
||||||
|
|
||||||
### Optional Parameters
|
|
||||||
|
|
||||||
- **`delta_dims`**: Apply delta transform (action - state) to these dimensions
|
|
||||||
- Format: `"0,1,2,3,4,5"`
|
|
||||||
- Use for position-based actions
|
|
||||||
|
|
||||||
- **`state_key`**: Dataset key containing state observations
|
|
||||||
- Default: `"observation.state"`
|
|
||||||
|
|
||||||
- **`vocab_size`**: BPE vocabulary size
|
|
||||||
- Default: 1024
|
|
||||||
- Larger = better compression but more memory
|
|
||||||
|
|
||||||
- **`scale`**: DCT quantization scale
|
|
||||||
- Default: 10.0
|
|
||||||
- Smaller = finer quantization, larger = coarser
|
|
||||||
|
|
||||||
- **`sample_fraction`**: Fraction of action chunks to use per episode
|
|
||||||
- Default: 0.1 (10%)
|
|
||||||
- Increase for small datasets, decrease for large datasets
|
|
||||||
|
|
||||||
## Output
|
|
||||||
|
|
||||||
The script creates a directory (default: `./fast_tokenizer_{repo_id}`) containing:
|
|
||||||
|
|
||||||
1. **Tokenizer files**: Can be loaded with `AutoProcessor.from_pretrained()`
|
|
||||||
2. **`metadata.json`**: Contains:
|
|
||||||
- Training configuration
|
|
||||||
- Compression statistics
|
|
||||||
- Dataset information
|
|
||||||
|
|
||||||
## Example Output
|
|
||||||
|
|
||||||
```
|
|
||||||
Loading dataset: lerobot/aloha_sim_insertion_human
|
|
||||||
Dataset loaded: 50 episodes, 5000 frames
|
|
||||||
Encoding 14 dimensions: 0:14
|
|
||||||
Delta dimensions: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13]
|
|
||||||
Action horizon: 10
|
|
||||||
Processing 50 episodes...
|
|
||||||
Collected 4500 action chunks
|
|
||||||
Extracted 14 encoded dimensions
|
|
||||||
|
|
||||||
Before normalization - overall stats:
|
|
||||||
Min: -2.3451, Max: 3.1234, Mean: 0.0234, Std: 0.8765
|
|
||||||
|
|
||||||
Applied quantile normalization [q01, q99] → [-1, 1]
|
|
||||||
|
|
||||||
After normalization - overall stats:
|
|
||||||
Min: -1.0000, Max: 1.0000, Mean: 0.0156, Std: 0.4321
|
|
||||||
|
|
||||||
Training FAST tokenizer on 4500 action chunks...
|
|
||||||
Action chunk shape: (4500, 10, 14)
|
|
||||||
Vocab size: 1024
|
|
||||||
DCT scale: 10.0
|
|
||||||
✓ Tokenizer training complete!
|
|
||||||
|
|
||||||
Compression Statistics:
|
|
||||||
Average compression ratio: 14.23x
|
|
||||||
Mean token length: 9.8
|
|
||||||
P99 token length: 15
|
|
||||||
Min token length: 6
|
|
||||||
Max token length: 18
|
|
||||||
|
|
||||||
✅ Saved FAST tokenizer to ./fast_tokenizer_lerobot_aloha_sim_insertion_human
|
|
||||||
```
|
|
||||||
|
|
||||||
## Using the Trained Tokenizer
|
|
||||||
|
|
||||||
```python
|
|
||||||
from transformers import AutoProcessor
|
|
||||||
|
|
||||||
# Load tokenizer
|
|
||||||
tokenizer = AutoProcessor.from_pretrained(
|
|
||||||
"./fast_tokenizer_lerobot_aloha_sim_insertion_human",
|
|
||||||
trust_remote_code=True
|
|
||||||
)
|
|
||||||
|
|
||||||
# Encode action chunk [horizon, action_dim]
|
|
||||||
action_chunk = np.random.randn(10, 14) # Example
|
|
||||||
tokens = tokenizer(action_chunk[None])[0] # Returns token IDs
|
|
||||||
|
|
||||||
# Decode tokens back to actions
|
|
||||||
reconstructed = tokenizer.decode(tokens)
|
|
||||||
```
|
|
||||||
|
|
||||||
## Tips
|
|
||||||
|
|
||||||
1. **Start Small**: Use `--max_episodes 10` for initial testing
|
|
||||||
2. **Check Dimensions**: Verify encoded dimensions match your robot's action space
|
|
||||||
3. **Delta Transform**: Use for position-based actions, not velocity-based
|
|
||||||
4. **Normalization**: Ensure dataset has proper statistics computed
|
|
||||||
5. **Compression Ratio**: Aim for 10-20x for good balance of compression and accuracy
|
|
||||||
|
|
||||||
## Troubleshooting
|
|
||||||
|
|
||||||
**Issue**: "No normalization stats found"
|
|
||||||
- **Solution**: Compute dataset statistics first, or use raw actions
|
|
||||||
|
|
||||||
**Issue**: "Episode too short for action horizon"
|
|
||||||
- **Solution**: Reduce `--action_horizon` or filter short episodes
|
|
||||||
|
|
||||||
**Issue**: "State key not found"
|
|
||||||
- **Solution**: Check dataset features and use correct `--state_key`
|
|
||||||
|
|
||||||
**Issue**: Memory error with large datasets
|
|
||||||
- **Solution**: Reduce `--sample_fraction` or `--max_episodes`
|
|
||||||
|
|
||||||
## Citation
|
|
||||||
|
|
||||||
If you use FAST in your research, please cite:
|
|
||||||
|
|
||||||
```bibtex
|
|
||||||
@article{black2023fast,
|
|
||||||
title={FAST: Factorized Action Sequence Tokenizer for Vision-Language-Action Models},
|
|
||||||
author={Black, Kevin and others},
|
|
||||||
journal={arXiv preprint},
|
|
||||||
year={2023}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
@@ -37,11 +37,6 @@ class PI05Config(PreTrainedConfig):
|
|||||||
# Shorter state and action vectors will be padded to these dimensions
|
# Shorter state and action vectors will be padded to these dimensions
|
||||||
max_state_dim: int = 32
|
max_state_dim: int = 32
|
||||||
max_action_dim: int = 32
|
max_action_dim: int = 32
|
||||||
max_action_tokens: int = 32
|
|
||||||
fast_vocab_size: int = 2048
|
|
||||||
|
|
||||||
# FAST-only mode: train with only discrete action token prediction (no flow matching, no subtask)
|
|
||||||
fast_only: bool = False
|
|
||||||
|
|
||||||
# Flow matching parameters: see openpi `PI0Pytorch`
|
# Flow matching parameters: see openpi `PI0Pytorch`
|
||||||
num_inference_steps: int = 10
|
num_inference_steps: int = 10
|
||||||
|
|||||||
@@ -1,21 +0,0 @@
|
|||||||
lerobot-train \
|
|
||||||
--dataset.repo_id=lerobot \
|
|
||||||
--dataset.root=/fsx/jade_choghari/outputs/collect-data-pgen \
|
|
||||||
--output_dir=/fsx/jade_choghari/outputs/pi0test1 \
|
|
||||||
--job_name=pi0_training \
|
|
||||||
--policy.repo_id=jade_choghari/pi0-base \
|
|
||||||
--policy.path=/fsx/jade_choghari/outputs/pi0_fast_fruit1/checkpoints/last/pretrained_model \
|
|
||||||
--policy.dtype=bfloat16 \
|
|
||||||
--steps=3000 \
|
|
||||||
--save_freq=1000 \
|
|
||||||
--rename_map='{
|
|
||||||
"observation.images.base": "observation.images.base_0_rgb",
|
|
||||||
"observation.images.left_wrist": "observation.images.left_wrist_0_rgb",
|
|
||||||
"observation.images.right_wrist": "observation.images.right_wrist_0_rgb",
|
|
||||||
}' \
|
|
||||||
--batch_size=4 \
|
|
||||||
--policy.device=cuda \
|
|
||||||
# --wandb.enable=true \
|
|
||||||
# --wandb.disable_artifact=true \
|
|
||||||
# --wandb.project=pi05hi-training \
|
|
||||||
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -33,7 +33,6 @@ from lerobot.processor import (
|
|||||||
ProcessorStep,
|
ProcessorStep,
|
||||||
ProcessorStepRegistry,
|
ProcessorStepRegistry,
|
||||||
RenameObservationsProcessorStep,
|
RenameObservationsProcessorStep,
|
||||||
ActionTokenizerProcessorStep,
|
|
||||||
TokenizerProcessorStep,
|
TokenizerProcessorStep,
|
||||||
UnnormalizerProcessorStep,
|
UnnormalizerProcessorStep,
|
||||||
)
|
)
|
||||||
@@ -55,8 +54,8 @@ class Pi05PrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
|||||||
|
|
||||||
max_state_dim: int = 32
|
max_state_dim: int = 32
|
||||||
task_key: str = "task"
|
task_key: str = "task"
|
||||||
high_level_task_key: str = "user_prompt"
|
prompt_key: str = "prompt"
|
||||||
subtask_only_key: str = "subtask"
|
target_key: str = "target"
|
||||||
|
|
||||||
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
def __call__(self, transition: EnvTransition) -> EnvTransition:
|
||||||
transition = transition.copy()
|
transition = transition.copy()
|
||||||
@@ -68,7 +67,7 @@ class Pi05PrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
|||||||
if tasks is None:
|
if tasks is None:
|
||||||
raise ValueError("No task found in complementary data")
|
raise ValueError("No task found in complementary data")
|
||||||
|
|
||||||
high_level_tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get(self.high_level_task_key)
|
high_level_tasks = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}).get("user_prompt")
|
||||||
|
|
||||||
# TODO: check if this necessary
|
# TODO: check if this necessary
|
||||||
state = deepcopy(state)
|
state = deepcopy(state)
|
||||||
@@ -87,36 +86,27 @@ class Pi05PrepareStateAndLanguageTokenizerProcessorStep(ProcessorStep):
|
|||||||
for high_level_task in high_level_tasks:
|
for high_level_task in high_level_tasks:
|
||||||
cleaned_high_level_tasks.append(high_level_task.strip().replace("_", " ").replace("\n", " "))
|
cleaned_high_level_tasks.append(high_level_task.strip().replace("_", " ").replace("\n", " "))
|
||||||
|
|
||||||
# Process low level tasks with state information
|
# Process tasks to create prompts (input) and targets (what to predict)
|
||||||
low_level_prompts = []
|
prompts = [] # Input prompts ending with "Subtask:"
|
||||||
subtask_only_prompts = [] # Store only the subtask text for prediction
|
targets = [] # Target text to predict (the subtask)
|
||||||
for i, task in enumerate(tasks):
|
for i, task in enumerate(tasks):
|
||||||
cleaned_text = task.strip().replace("_", " ").replace("\n", " ")
|
cleaned_text = task.strip().replace("_", " ").replace("\n", " ")
|
||||||
state_str = " ".join(map(str, discretized_states[i]))
|
state_str = " ".join(map(str, discretized_states[i]))
|
||||||
|
|
||||||
# Store only the subtask text (used as prediction target)
|
# Store the subtask text as target for prediction
|
||||||
subtask_only_prompts.append(cleaned_text)
|
targets.append(cleaned_text)
|
||||||
|
|
||||||
if cleaned_high_level_tasks:
|
if cleaned_high_level_tasks:
|
||||||
cleaned_high_level_task = cleaned_high_level_tasks[i]
|
cleaned_high_level_task = cleaned_high_level_tasks[i]
|
||||||
full_prompt = f"High level task: {cleaned_high_level_task}; State: {state_str}; Subtask: {cleaned_text}"
|
# Prompt ends with "Subtask:" - model will predict the target
|
||||||
|
prompt = f"High level task: {cleaned_high_level_task}; State: {state_str}; Subtask:"
|
||||||
else:
|
else:
|
||||||
full_prompt = f"Task: {cleaned_text}, State: {state_str};\n" #remove Action by jade
|
raise ValueError("No high level tasks found")
|
||||||
|
|
||||||
low_level_prompts.append(full_prompt)
|
prompts.append(prompt)
|
||||||
|
|
||||||
transition[TransitionKey.COMPLEMENTARY_DATA][self.task_key] = low_level_prompts
|
transition[TransitionKey.COMPLEMENTARY_DATA][self.prompt_key] = prompts
|
||||||
transition[TransitionKey.COMPLEMENTARY_DATA][self.subtask_only_key] = subtask_only_prompts
|
transition[TransitionKey.COMPLEMENTARY_DATA][self.target_key] = targets
|
||||||
|
|
||||||
# Process high level tasks without state information (if available)
|
|
||||||
if high_level_tasks is not None:
|
|
||||||
high_level_prompts = []
|
|
||||||
for i, cleaned_high_level_task in enumerate(cleaned_high_level_tasks):
|
|
||||||
state_str = " ".join(map(str, discretized_states[i]))
|
|
||||||
full_prompt = f"High level task: {cleaned_high_level_task}; State: {state_str}; Subtask:"
|
|
||||||
high_level_prompts.append(full_prompt)
|
|
||||||
|
|
||||||
transition[TransitionKey.COMPLEMENTARY_DATA][self.high_level_task_key] = high_level_prompts
|
|
||||||
return transition
|
return transition
|
||||||
|
|
||||||
def transform_features(
|
def transform_features(
|
||||||
@@ -159,6 +149,7 @@ def make_pi05_pre_post_processors(
|
|||||||
Returns:
|
Returns:
|
||||||
A tuple containing the configured pre-processor and post-processor pipelines.
|
A tuple containing the configured pre-processor and post-processor pipelines.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Add remaining processors
|
# Add remaining processors
|
||||||
input_steps: list[ProcessorStep] = [
|
input_steps: list[ProcessorStep] = [
|
||||||
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
|
RenameObservationsProcessorStep(rename_map={}), # To mimic the same processor as pretrained one
|
||||||
@@ -177,9 +168,6 @@ def make_pi05_pre_post_processors(
|
|||||||
padding_side="right",
|
padding_side="right",
|
||||||
padding="max_length",
|
padding="max_length",
|
||||||
),
|
),
|
||||||
ActionTokenizerProcessorStep(
|
|
||||||
tokenizer_name="/fsx/jade_choghari/outputs/fast_tokenizer", # TODO: jade put the PI
|
|
||||||
),
|
|
||||||
DeviceProcessorStep(device=config.device),
|
DeviceProcessorStep(device=config.device),
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
@@ -1,23 +0,0 @@
|
|||||||
export CUDA_LAUNCH_BLOCKING=1
|
|
||||||
lerobot-train \
|
|
||||||
--dataset.repo_id=local \
|
|
||||||
--dataset.root=/fsx/jade_choghari/outputs/collect-data-pgen \
|
|
||||||
--output_dir=/fsx/jade_choghari/outputs/pi0_fast_fruit2 \
|
|
||||||
--job_name=pi0_training \
|
|
||||||
--policy.repo_id=jade_choghari/pi0-base1 \
|
|
||||||
--policy.path=lerobot/pi05_base \
|
|
||||||
--policy.dtype=bfloat16 \
|
|
||||||
--steps=200000 \
|
|
||||||
--save_freq=5000 \
|
|
||||||
--rename_map='{
|
|
||||||
"observation.images.base": "observation.images.base_0_rgb",
|
|
||||||
"observation.images.left_wrist": "observation.images.left_wrist_0_rgb",
|
|
||||||
"observation.images.right_wrist": "observation.images.right_wrist_0_rgb",
|
|
||||||
}' \
|
|
||||||
--batch_size=16 \
|
|
||||||
--policy.device=cuda \
|
|
||||||
--policy.fast_only=true \
|
|
||||||
# --wandb.enable=true \
|
|
||||||
# --wandb.disable_artifact=true \
|
|
||||||
# --wandb.project=pi05hi-training \
|
|
||||||
# /fsx/jade_choghari/.cache/huggingface/lerobot/jadechoghari/collect-data
|
|
||||||
@@ -1,13 +0,0 @@
|
|||||||
rm -rf /fsx/jade_choghari/outputs/pi0_multi_training
|
|
||||||
lerobot-train \
|
|
||||||
--dataset.repo_id=local\
|
|
||||||
--dataset.root=/fsx/jade_choghari/data/libero \
|
|
||||||
--output_dir=/fsx/jade_choghari/outputs/pi0_multi_training \
|
|
||||||
--job_name=pi0_multi_training \
|
|
||||||
--policy.repo_id=jadechoghari/pi0-base1 \
|
|
||||||
--policy.path=/fsx/jade_choghari/outputs/libero_training_fast_6/checkpoints/last/pretrained_model/ \
|
|
||||||
--policy.dtype=bfloat16 \
|
|
||||||
--steps=50000 \
|
|
||||||
--save_freq=5000 \
|
|
||||||
--batch_size=4 \
|
|
||||||
--policy.device=cuda \
|
|
||||||
@@ -1,12 +0,0 @@
|
|||||||
python src/lerobot/policies/pi05/train_fast_tokenizer.py \
|
|
||||||
--repo_id "local" \
|
|
||||||
--root /fsx/jade_choghari/data/libero \
|
|
||||||
--action_horizon 10 \
|
|
||||||
--encoded_dims "0:7" \
|
|
||||||
--vocab_size 1024 \
|
|
||||||
--push_to_hub \
|
|
||||||
--hub_repo_id jadechoghari/fast-libero-tokenizer-quantiles \
|
|
||||||
--normalization_mode QUANTILES \
|
|
||||||
|
|
||||||
|
|
||||||
# python train_fast_tokenizer.py --repo_id my_dataset
|
|
||||||
@@ -1,533 +0,0 @@
|
|||||||
"""Train FAST tokenizer for action encoding.
|
|
||||||
|
|
||||||
This script:
|
|
||||||
1. Loads action chunks from LeRobotDataset (with sampling)
|
|
||||||
2. Applies delta transforms and per-timestamp normalization
|
|
||||||
3. Trains FAST tokenizer on specified action dimensions
|
|
||||||
4. Saves tokenizer to assets directory
|
|
||||||
5. Reports compression statistics
|
|
||||||
"""
|
|
||||||
|
|
||||||
import json
|
|
||||||
import numpy as np
|
|
||||||
import tyro
|
|
||||||
from pathlib import Path
|
|
||||||
from transformers import AutoProcessor
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from huggingface_hub import HfApi
|
|
||||||
from lerobot.configs.types import NormalizationMode
|
|
||||||
from lerobot.datasets.lerobot_dataset import LeRobotDataset
|
|
||||||
|
|
||||||
|
|
||||||
def apply_delta_transform(state: np.ndarray, actions: np.ndarray, delta_dims: list[int] | None) -> np.ndarray:
|
|
||||||
"""Apply delta transform to specified dimensions.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
state: Current state [D]
|
|
||||||
actions: Future actions [D]
|
|
||||||
delta_dims: List of dimension indices to apply delta transform to
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Transformed actions [D]
|
|
||||||
"""
|
|
||||||
if delta_dims is None or len(delta_dims) == 0:
|
|
||||||
return actions
|
|
||||||
|
|
||||||
delta_actions = actions.copy()
|
|
||||||
for dim in delta_dims:
|
|
||||||
delta_actions[dim] = actions[dim] - state[dim]
|
|
||||||
|
|
||||||
return delta_actions
|
|
||||||
|
|
||||||
|
|
||||||
def apply_normalization(
|
|
||||||
data: np.ndarray,
|
|
||||||
stats: dict[str, np.ndarray],
|
|
||||||
mode: NormalizationMode,
|
|
||||||
eps: float = 1e-8,
|
|
||||||
) -> np.ndarray:
|
|
||||||
"""Apply normalization to data based on the specified mode.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
data: Data to normalize [N, H, D] or [D]
|
|
||||||
stats: Dictionary of statistics (mean, std, min, max, q01, q99, q10, q90)
|
|
||||||
mode: Normalization mode to apply
|
|
||||||
eps: Small epsilon for numerical stability
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Normalized data with the same shape as input
|
|
||||||
"""
|
|
||||||
if mode == NormalizationMode.IDENTITY:
|
|
||||||
return data
|
|
||||||
|
|
||||||
if mode == NormalizationMode.MEAN_STD:
|
|
||||||
mean = stats.get("mean")
|
|
||||||
std = stats.get("std")
|
|
||||||
if mean is None or std is None:
|
|
||||||
raise ValueError("MEAN_STD mode requires 'mean' and 'std' in stats")
|
|
||||||
return (data - mean) / np.maximum(std, eps)
|
|
||||||
|
|
||||||
if mode == NormalizationMode.MIN_MAX:
|
|
||||||
min_val = stats.get("min")
|
|
||||||
max_val = stats.get("max")
|
|
||||||
if min_val is None or max_val is None:
|
|
||||||
raise ValueError("MIN_MAX mode requires 'min' and 'max' in stats")
|
|
||||||
denom = np.maximum(max_val - min_val, eps)
|
|
||||||
return 2.0 * (data - min_val) / denom - 1.0
|
|
||||||
|
|
||||||
if mode == NormalizationMode.QUANTILES:
|
|
||||||
q01 = stats.get("q01")
|
|
||||||
q99 = stats.get("q99")
|
|
||||||
if q01 is None or q99 is None:
|
|
||||||
raise ValueError("QUANTILES mode requires 'q01' and 'q99' in stats")
|
|
||||||
denom = np.maximum(q99 - q01, eps)
|
|
||||||
# Clip to quantile range then normalize to [-1, 1]
|
|
||||||
clipped = np.clip(data, q01, q99)
|
|
||||||
return 2.0 * (clipped - q01) / denom - 1.0
|
|
||||||
|
|
||||||
if mode == NormalizationMode.QUANTILE10:
|
|
||||||
q10 = stats.get("q10")
|
|
||||||
q90 = stats.get("q90")
|
|
||||||
if q10 is None or q90 is None:
|
|
||||||
raise ValueError("QUANTILE10 mode requires 'q10' and 'q90' in stats")
|
|
||||||
denom = np.maximum(q90 - q10, eps)
|
|
||||||
# Clip to quantile range then normalize to [-1, 1]
|
|
||||||
clipped = np.clip(data, q10, q90)
|
|
||||||
return 2.0 * (clipped - q10) / denom - 1.0
|
|
||||||
|
|
||||||
raise ValueError(f"Unsupported normalization mode: {mode}")
|
|
||||||
|
|
||||||
|
|
||||||
def process_episode(args):
|
|
||||||
"""Process single episode and return action chunks."""
|
|
||||||
dataset, ep_idx, action_horizon, delta_dims, sample_fraction, state_key, use_delta_transform = args
|
|
||||||
|
|
||||||
try:
|
|
||||||
# Get episode info
|
|
||||||
ep_info = dataset.meta.episodes[ep_idx]
|
|
||||||
from_idx = ep_info["dataset_from_index"]
|
|
||||||
to_idx = ep_info["dataset_to_index"]
|
|
||||||
ep_length = to_idx - from_idx
|
|
||||||
|
|
||||||
if ep_length < action_horizon:
|
|
||||||
return None
|
|
||||||
|
|
||||||
# Load all frames in episode
|
|
||||||
# If dataset has episode filtering, we need to use the mapping
|
|
||||||
states = []
|
|
||||||
actions = []
|
|
||||||
|
|
||||||
for abs_idx in range(from_idx, to_idx):
|
|
||||||
# Map absolute index to relative index if needed
|
|
||||||
if dataset._absolute_to_relative_idx is not None:
|
|
||||||
if abs_idx not in dataset._absolute_to_relative_idx:
|
|
||||||
# This episode's frames aren't in the filtered dataset
|
|
||||||
return None
|
|
||||||
rel_idx = dataset._absolute_to_relative_idx[abs_idx]
|
|
||||||
else:
|
|
||||||
rel_idx = abs_idx
|
|
||||||
|
|
||||||
frame = dataset.hf_dataset[rel_idx]
|
|
||||||
|
|
||||||
# Get state (could be from observation.state or other state key)
|
|
||||||
if state_key in frame:
|
|
||||||
state = frame[state_key].numpy() if torch.is_tensor(frame[state_key]) else np.array(frame[state_key])
|
|
||||||
else:
|
|
||||||
# If no state key, use zeros (no delta transform)
|
|
||||||
state = np.zeros_like(frame["action"].numpy() if torch.is_tensor(frame["action"]) else np.array(frame["action"]))
|
|
||||||
|
|
||||||
action = frame["action"].numpy() if torch.is_tensor(frame["action"]) else np.array(frame["action"])
|
|
||||||
|
|
||||||
states.append(state)
|
|
||||||
actions.append(action)
|
|
||||||
|
|
||||||
states = np.array(states)
|
|
||||||
actions = np.array(actions)
|
|
||||||
|
|
||||||
# Create action chunks (sliding window)
|
|
||||||
# All actions in a chunk are relative to the FIRST state in that chunk
|
|
||||||
action_chunks = []
|
|
||||||
|
|
||||||
for i in range(len(states) - action_horizon + 1):
|
|
||||||
current_state = states[i] # First state in chunk
|
|
||||||
future_absolute_actions = actions[i:i + action_horizon]
|
|
||||||
|
|
||||||
if use_delta_transform:
|
|
||||||
# Relative actions
|
|
||||||
delta_chunk = np.zeros_like(future_absolute_actions)
|
|
||||||
for t in range(action_horizon):
|
|
||||||
delta_chunk[t] = apply_delta_transform(
|
|
||||||
current_state,
|
|
||||||
future_absolute_actions[t],
|
|
||||||
delta_dims,
|
|
||||||
)
|
|
||||||
action_chunks.append(delta_chunk)
|
|
||||||
else:
|
|
||||||
# Absolute actions (NO delta)
|
|
||||||
action_chunks.append(future_absolute_actions)
|
|
||||||
|
|
||||||
if len(action_chunks) == 0:
|
|
||||||
return None
|
|
||||||
|
|
||||||
action_chunks = np.array(action_chunks)
|
|
||||||
|
|
||||||
# Sample chunks
|
|
||||||
if sample_fraction < 1.0:
|
|
||||||
n_chunks = len(action_chunks)
|
|
||||||
n_samples = max(1, int(n_chunks * sample_fraction))
|
|
||||||
episode_seed = hash(ep_idx) % (2**31)
|
|
||||||
rng = np.random.RandomState(episode_seed)
|
|
||||||
indices = rng.choice(n_chunks, size=n_samples, replace=False)
|
|
||||||
action_chunks = action_chunks[indices]
|
|
||||||
|
|
||||||
return action_chunks
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error processing episode {ep_idx}: {e}")
|
|
||||||
import traceback
|
|
||||||
traceback.print_exc()
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def train_fast_tokenizer(
|
|
||||||
action_chunks: np.ndarray,
|
|
||||||
vocab_size: int = 1024,
|
|
||||||
scale: float = 10.0,
|
|
||||||
) -> AutoProcessor:
|
|
||||||
"""
|
|
||||||
Train FAST tokenizer (BPE on DCT coefficients) on action chunks.
|
|
||||||
|
|
||||||
Uses the .fit() method to train a new tokenizer on the provided data.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
action_chunks: Array of action chunks [N, H, D] where N=num_chunks, H=horizon, D=action_dim
|
|
||||||
vocab_size: BPE vocabulary size
|
|
||||||
scale: DCT scaling factor for quantization
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Trained FAST tokenizer
|
|
||||||
"""
|
|
||||||
print(f"Training FAST tokenizer on {len(action_chunks)} action chunks...")
|
|
||||||
print(f"Action chunk shape: {action_chunks.shape}")
|
|
||||||
print(f"Vocab size: {vocab_size}")
|
|
||||||
print(f"DCT scale: {scale}")
|
|
||||||
|
|
||||||
# Download the tokenizer source code (not pretrained weights)
|
|
||||||
# We'll train a new tokenizer on our own data
|
|
||||||
base_tokenizer = AutoProcessor.from_pretrained(
|
|
||||||
"physical-intelligence/fast",
|
|
||||||
trust_remote_code=True
|
|
||||||
)
|
|
||||||
|
|
||||||
# Convert action_chunks array to list of arrays (expected by .fit())
|
|
||||||
action_data_list = [action_chunks[i] for i in range(len(action_chunks))]
|
|
||||||
|
|
||||||
# Train the new tokenizer on our action data using .fit()
|
|
||||||
# This trains the BPE tokenizer on DCT coefficients
|
|
||||||
print("Training new tokenizer (this may take a few minutes)...")
|
|
||||||
tokenizer = base_tokenizer.fit(
|
|
||||||
action_data_list,
|
|
||||||
scale=scale,
|
|
||||||
vocab_size=vocab_size,
|
|
||||||
time_horizon=action_chunks.shape[1], # action_horizon
|
|
||||||
action_dim=action_chunks.shape[2], # encoded dimensions
|
|
||||||
)
|
|
||||||
print("✓ Tokenizer training complete!")
|
|
||||||
|
|
||||||
# Validate it works
|
|
||||||
sample_chunk = action_chunks[0]
|
|
||||||
encoded = tokenizer(sample_chunk[None])[0]
|
|
||||||
if isinstance(encoded, list):
|
|
||||||
encoded = np.array(encoded)
|
|
||||||
print(f"Sample encoding: {len(encoded)} tokens for chunk shape {sample_chunk.shape}")
|
|
||||||
|
|
||||||
return tokenizer
|
|
||||||
|
|
||||||
|
|
||||||
def compute_compression_stats(tokenizer, action_chunks: np.ndarray):
|
|
||||||
"""Compute compression statistics."""
|
|
||||||
print("\nComputing compression statistics...")
|
|
||||||
|
|
||||||
# Sample for stats (use max 1000 chunks for speed)
|
|
||||||
sample_size = min(1000, len(action_chunks))
|
|
||||||
sample_indices = np.random.RandomState(42).choice(len(action_chunks), size=sample_size, replace=False)
|
|
||||||
sample_chunks = action_chunks[sample_indices]
|
|
||||||
|
|
||||||
token_lengths = []
|
|
||||||
for chunk in sample_chunks:
|
|
||||||
encoded = tokenizer(chunk[None])[0]
|
|
||||||
if isinstance(encoded, list):
|
|
||||||
token_lengths.append(len(encoded))
|
|
||||||
else:
|
|
||||||
token_lengths.append(encoded.shape[0] if hasattr(encoded, 'shape') else len(encoded))
|
|
||||||
|
|
||||||
token_lengths = np.array(token_lengths)
|
|
||||||
|
|
||||||
# Compression ratio: (H * D) / avg_tokens
|
|
||||||
input_size = action_chunks.shape[1] * action_chunks.shape[2]
|
|
||||||
avg_tokens = np.mean(token_lengths)
|
|
||||||
compression_ratio = input_size / avg_tokens
|
|
||||||
|
|
||||||
stats = {
|
|
||||||
'compression_ratio': float(compression_ratio),
|
|
||||||
'mean_token_length': float(np.mean(token_lengths)),
|
|
||||||
'p99_token_length': float(np.percentile(token_lengths, 99)),
|
|
||||||
'min_token_length': float(np.min(token_lengths)),
|
|
||||||
'max_token_length': float(np.max(token_lengths)),
|
|
||||||
}
|
|
||||||
|
|
||||||
print(f"Compression Statistics:")
|
|
||||||
print(f" Average compression ratio: {stats['compression_ratio']:.2f}x")
|
|
||||||
print(f" Mean token length: {stats['mean_token_length']:.1f}")
|
|
||||||
print(f" P99 token length: {stats['p99_token_length']:.0f}")
|
|
||||||
print(f" Min token length: {stats['min_token_length']:.0f}")
|
|
||||||
print(f" Max token length: {stats['max_token_length']:.0f}")
|
|
||||||
|
|
||||||
return stats
|
|
||||||
|
|
||||||
|
|
||||||
def main(
|
|
||||||
repo_id: str,
|
|
||||||
root: str | None = None,
|
|
||||||
action_horizon: int = 10,
|
|
||||||
max_episodes: int | None = None,
|
|
||||||
sample_fraction: float = 0.1,
|
|
||||||
encoded_dims: str = "0:6,7:23",
|
|
||||||
delta_dims: str | None = None,
|
|
||||||
use_delta_transform: bool = False,
|
|
||||||
state_key: str = "observation.state",
|
|
||||||
normalization_mode: str = "QUANTILES",
|
|
||||||
vocab_size: int = 1024,
|
|
||||||
scale: float = 10.0,
|
|
||||||
output_dir: str | None = None,
|
|
||||||
push_to_hub: bool = False,
|
|
||||||
hub_repo_id: str | None = None,
|
|
||||||
hub_private: bool = False,
|
|
||||||
):
|
|
||||||
"""
|
|
||||||
Train FAST tokenizer for action encoding.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
repo_id: LeRobot dataset repository ID
|
|
||||||
root: Root directory for dataset (default: ~/.cache/huggingface/lerobot)
|
|
||||||
action_horizon: Number of future actions in each chunk
|
|
||||||
max_episodes: Max episodes to use (None = all episodes in dataset)
|
|
||||||
sample_fraction: Fraction of chunks to sample per episode
|
|
||||||
encoded_dims: Comma-separated dimension ranges to encode (e.g., "0:6,7:23")
|
|
||||||
delta_dims: Comma-separated dimension indices for delta transform (e.g., "0,1,2,3,4,5")
|
|
||||||
use_delta_transform: Whether to apply delta transform (relative actions vs absolute actions)
|
|
||||||
state_key: Dataset key for state observations (default: "observation.state")
|
|
||||||
normalization_mode: Normalization mode (MEAN_STD, MIN_MAX, QUANTILES, QUANTILE10, IDENTITY)
|
|
||||||
vocab_size: FAST vocabulary size (BPE vocab size)
|
|
||||||
scale: DCT scaling factor (default: 10.0)
|
|
||||||
output_dir: Directory to save tokenizer (default: ./fast_tokenizer_{repo_id})
|
|
||||||
push_to_hub: Whether to push the tokenizer to Hugging Face Hub
|
|
||||||
hub_repo_id: Hub repository ID (e.g., "username/tokenizer-name"). If None, uses output_dir name
|
|
||||||
hub_private: Whether to create a private repository on the Hub
|
|
||||||
"""
|
|
||||||
# Load dataset
|
|
||||||
print(f"Loading dataset: {repo_id}")
|
|
||||||
dataset = LeRobotDataset(repo_id=repo_id, root=root)
|
|
||||||
print(f"Dataset loaded: {dataset.num_episodes} episodes, {dataset.num_frames} frames")
|
|
||||||
|
|
||||||
# Parse normalization mode
|
|
||||||
try:
|
|
||||||
norm_mode = NormalizationMode(normalization_mode)
|
|
||||||
except ValueError:
|
|
||||||
raise ValueError(
|
|
||||||
f"Invalid normalization_mode: {normalization_mode}. "
|
|
||||||
f"Must be one of: {', '.join([m.value for m in NormalizationMode])}"
|
|
||||||
)
|
|
||||||
print(f"Normalization mode: {norm_mode.value}")
|
|
||||||
|
|
||||||
# Parse encoded dimensions
|
|
||||||
encoded_dim_ranges = []
|
|
||||||
for range_str in encoded_dims.split(','):
|
|
||||||
start, end = map(int, range_str.strip().split(':'))
|
|
||||||
encoded_dim_ranges.append((start, end))
|
|
||||||
|
|
||||||
total_encoded_dims = sum(end - start for start, end in encoded_dim_ranges)
|
|
||||||
print(f"Encoding {total_encoded_dims} dimensions: {encoded_dims}")
|
|
||||||
|
|
||||||
# Parse delta dimensions
|
|
||||||
delta_dim_list = None
|
|
||||||
if delta_dims is not None and delta_dims.strip():
|
|
||||||
delta_dim_list = [int(d.strip()) for d in delta_dims.split(',')]
|
|
||||||
print(f"Delta dimensions: {delta_dim_list}")
|
|
||||||
else:
|
|
||||||
print("No delta dimensions specified")
|
|
||||||
|
|
||||||
print(f"Use delta transform: {use_delta_transform}")
|
|
||||||
if use_delta_transform and (delta_dim_list is None or len(delta_dim_list) == 0):
|
|
||||||
print("Warning: use_delta_transform=True but no delta_dims specified. No delta will be applied.")
|
|
||||||
|
|
||||||
print(f"Action horizon: {action_horizon}")
|
|
||||||
print(f"State key: {state_key}")
|
|
||||||
|
|
||||||
# Determine episodes to process
|
|
||||||
num_episodes = dataset.num_episodes
|
|
||||||
if max_episodes is not None:
|
|
||||||
num_episodes = min(max_episodes, num_episodes)
|
|
||||||
|
|
||||||
print(f"Processing {num_episodes} episodes...")
|
|
||||||
|
|
||||||
# Process episodes sequentially (to avoid pickling issues with dataset)
|
|
||||||
all_chunks = []
|
|
||||||
for ep_idx in range(num_episodes):
|
|
||||||
if ep_idx % 10 == 0:
|
|
||||||
print(f" Processing episode {ep_idx}/{num_episodes}...")
|
|
||||||
|
|
||||||
chunks = process_episode(
|
|
||||||
(dataset, ep_idx, action_horizon, delta_dim_list, sample_fraction, state_key, use_delta_transform)
|
|
||||||
)
|
|
||||||
if chunks is not None:
|
|
||||||
all_chunks.append(chunks)
|
|
||||||
|
|
||||||
# Concatenate all chunks
|
|
||||||
all_chunks = np.concatenate(all_chunks, axis=0)
|
|
||||||
print(f"Collected {len(all_chunks)} action chunks")
|
|
||||||
|
|
||||||
# Extract only encoded dimensions FIRST (before normalization)
|
|
||||||
encoded_chunks = []
|
|
||||||
for start, end in encoded_dim_ranges:
|
|
||||||
encoded_chunks.append(all_chunks[:, :, start:end])
|
|
||||||
encoded_chunks = np.concatenate(encoded_chunks, axis=-1) # [N, H, D_encoded]
|
|
||||||
print(f"Extracted {encoded_chunks.shape[-1]} encoded dimensions")
|
|
||||||
|
|
||||||
# Apply normalization to encoded dimensions
|
|
||||||
print(f"\nBefore normalization - overall stats:")
|
|
||||||
print(f" Min: {np.min(encoded_chunks):.4f}, Max: {np.max(encoded_chunks):.4f}")
|
|
||||||
print(f" Mean: {np.mean(encoded_chunks):.4f}, Std: {np.std(encoded_chunks):.4f}")
|
|
||||||
|
|
||||||
# Get normalization stats from dataset
|
|
||||||
norm_stats = dataset.meta.stats
|
|
||||||
if norm_stats is not None and "action" in norm_stats:
|
|
||||||
action_stats = norm_stats["action"]
|
|
||||||
|
|
||||||
# Build encoded dimension indices
|
|
||||||
encoded_dim_indices = []
|
|
||||||
for start, end in encoded_dim_ranges:
|
|
||||||
encoded_dim_indices.extend(range(start, end))
|
|
||||||
encoded_dim_indices = np.array(encoded_dim_indices)
|
|
||||||
|
|
||||||
# Extract stats for encoded dimensions only
|
|
||||||
encoded_stats = {}
|
|
||||||
for stat_name, stat_values in action_stats.items():
|
|
||||||
if isinstance(stat_values, (list, np.ndarray)):
|
|
||||||
stat_array = np.array(stat_values)
|
|
||||||
if len(stat_array) > max(encoded_dim_indices):
|
|
||||||
encoded_stats[stat_name] = stat_array[encoded_dim_indices]
|
|
||||||
|
|
||||||
if encoded_stats:
|
|
||||||
print(f"\nNormalization stats for encoded dimensions (mode: {norm_mode.value}):")
|
|
||||||
for stat_name, stat_values in encoded_stats.items():
|
|
||||||
print(f" {stat_name}: shape={stat_values.shape}, "
|
|
||||||
f"range=[{np.min(stat_values):.4f}, {np.max(stat_values):.4f}]")
|
|
||||||
|
|
||||||
# Apply normalization based on mode
|
|
||||||
try:
|
|
||||||
encoded_chunks = apply_normalization(
|
|
||||||
encoded_chunks,
|
|
||||||
encoded_stats,
|
|
||||||
norm_mode,
|
|
||||||
eps=1e-8
|
|
||||||
)
|
|
||||||
print(f"\nApplied {norm_mode.value} normalization")
|
|
||||||
except ValueError as e:
|
|
||||||
print(f"Warning: {e}. Using raw actions without normalization.")
|
|
||||||
|
|
||||||
print(f"\nAfter normalization - overall stats:")
|
|
||||||
print(f" Min: {np.min(encoded_chunks):.4f}, Max: {np.max(encoded_chunks):.4f}")
|
|
||||||
print(f" Mean: {np.mean(encoded_chunks):.4f}, Std: {np.std(encoded_chunks):.4f}")
|
|
||||||
|
|
||||||
print(f"\nPer-dimension stats (after normalization):")
|
|
||||||
for d in range(encoded_chunks.shape[-1]):
|
|
||||||
dim_data = encoded_chunks[:, :, d]
|
|
||||||
print(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}")
|
|
||||||
else:
|
|
||||||
print("Warning: Could not extract stats for encoded dimensions, using raw actions")
|
|
||||||
else:
|
|
||||||
print("Warning: No normalization stats found in dataset, using raw actions")
|
|
||||||
|
|
||||||
print(f"Encoded chunks shape: {encoded_chunks.shape}")
|
|
||||||
|
|
||||||
# Train FAST tokenizer
|
|
||||||
tokenizer = train_fast_tokenizer(
|
|
||||||
encoded_chunks,
|
|
||||||
vocab_size=vocab_size,
|
|
||||||
scale=scale,
|
|
||||||
)
|
|
||||||
|
|
||||||
# Compute compression statistics
|
|
||||||
compression_stats = compute_compression_stats(tokenizer, encoded_chunks)
|
|
||||||
|
|
||||||
# Save tokenizer
|
|
||||||
if output_dir is None:
|
|
||||||
output_dir = f"fast_tokenizer_{repo_id.replace('/', '_')}"
|
|
||||||
output_path = Path(output_dir)
|
|
||||||
output_path.mkdir(parents=True, exist_ok=True)
|
|
||||||
|
|
||||||
tokenizer.save_pretrained(output_path)
|
|
||||||
|
|
||||||
# Save metadata
|
|
||||||
metadata = {
|
|
||||||
'repo_id': repo_id,
|
|
||||||
'vocab_size': vocab_size,
|
|
||||||
'scale': scale,
|
|
||||||
'encoded_dims': encoded_dims,
|
|
||||||
'encoded_dim_ranges': encoded_dim_ranges,
|
|
||||||
'total_encoded_dims': total_encoded_dims,
|
|
||||||
'delta_dims': delta_dims,
|
|
||||||
'delta_dim_list': delta_dim_list,
|
|
||||||
'use_delta_transform': use_delta_transform,
|
|
||||||
'state_key': state_key,
|
|
||||||
'normalization_mode': norm_mode.value,
|
|
||||||
'action_horizon': action_horizon,
|
|
||||||
'num_training_chunks': len(encoded_chunks),
|
|
||||||
'compression_stats': compression_stats,
|
|
||||||
}
|
|
||||||
|
|
||||||
with open(output_path / "metadata.json", 'w') as f:
|
|
||||||
json.dump(metadata, f, indent=2)
|
|
||||||
|
|
||||||
print(f"\nSaved FAST tokenizer to {output_path}")
|
|
||||||
print(f"Metadata: {json.dumps(metadata, indent=2)}")
|
|
||||||
|
|
||||||
# Push to Hugging Face Hub if requested
|
|
||||||
if push_to_hub:
|
|
||||||
# Determine the hub repository ID
|
|
||||||
if hub_repo_id is None:
|
|
||||||
hub_repo_id = output_path.name
|
|
||||||
print(f"\nNo hub_repo_id provided, using: {hub_repo_id}")
|
|
||||||
|
|
||||||
print(f"\nPushing tokenizer to Hugging Face Hub: {hub_repo_id}")
|
|
||||||
print(f" Private: {hub_private}")
|
|
||||||
|
|
||||||
try:
|
|
||||||
# Use the tokenizer's push_to_hub method
|
|
||||||
tokenizer.push_to_hub(
|
|
||||||
repo_id=hub_repo_id,
|
|
||||||
private=hub_private,
|
|
||||||
commit_message=f"Upload FAST tokenizer trained on {repo_id}"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Also upload the metadata.json file separately
|
|
||||||
api = HfApi()
|
|
||||||
api.upload_file(
|
|
||||||
path_or_fileobj=str(output_path / "metadata.json"),
|
|
||||||
path_in_repo="metadata.json",
|
|
||||||
repo_id=hub_repo_id,
|
|
||||||
repo_type="model",
|
|
||||||
commit_message="Upload tokenizer metadata"
|
|
||||||
)
|
|
||||||
|
|
||||||
print(f"Successfully pushed tokenizer to: https://huggingface.co/{hub_repo_id}")
|
|
||||||
except Exception as e:
|
|
||||||
print(f"Error pushing to hub: {e}")
|
|
||||||
print(" Make sure you're logged in with `huggingface-cli login`")
|
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
tyro.cli(main)
|
|
||||||
@@ -1,101 +0,0 @@
|
|||||||
# Train FAST Tokenizer - Usage Examples
|
|
||||||
|
|
||||||
This script trains a FAST (Factorized Action Sequence Tokenizer) on LeRobotDataset action data.
|
|
||||||
|
|
||||||
## Basic Usage
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python src/lerobot/policies/pi05/train_fast_tokenizer.py \
|
|
||||||
--repo_id "lerobot/aloha_sim_insertion_human" \
|
|
||||||
--action_horizon 10 \
|
|
||||||
--encoded_dims "0:7" \
|
|
||||||
--vocab_size 1024 \
|
|
||||||
--scale 10.0
|
|
||||||
```
|
|
||||||
|
|
||||||
## Parameters
|
|
||||||
|
|
||||||
### Required
|
|
||||||
- `--repo_id`: LeRobot dataset repository ID (e.g., "lerobot/aloha_sim_insertion_human")
|
|
||||||
|
|
||||||
### Optional
|
|
||||||
- `--root`: Root directory for dataset (default: ~/.cache/huggingface/lerobot)
|
|
||||||
- `--action_horizon`: Number of future actions in each chunk (default: 10)
|
|
||||||
- `--max_episodes`: Maximum number of episodes to use (default: None = all)
|
|
||||||
- `--sample_fraction`: Fraction of chunks to sample per episode (default: 0.1)
|
|
||||||
- `--encoded_dims`: Comma-separated dimension ranges to encode (default: "0:6,7:23")
|
|
||||||
- Example: "0:7" encodes dimensions 0-6
|
|
||||||
- Example: "0:3,6:9" encodes dimensions 0-2 and 6-8
|
|
||||||
- `--delta_dims`: Comma-separated dimension indices for delta transform (default: None)
|
|
||||||
- Example: "0,1,2,3,4,5" applies delta transform to first 6 dimensions
|
|
||||||
- Delta transform: action[i] - state[i] for specified dimensions
|
|
||||||
- `--state_key`: Dataset key for state observations (default: "observation.state")
|
|
||||||
- `--vocab_size`: FAST vocabulary size / BPE vocab size (default: 1024)
|
|
||||||
- `--scale`: DCT scaling factor (default: 10.0)
|
|
||||||
- `--output_dir`: Directory to save tokenizer (default: ./fast_tokenizer_{repo_id})
|
|
||||||
|
|
||||||
## Examples
|
|
||||||
|
|
||||||
### Example 1: Train on full action space
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python src/lerobot/policies/pi05/train_fast_tokenizer.py \
|
|
||||||
--repo_id "lerobot/pusht" \
|
|
||||||
--action_horizon 16 \
|
|
||||||
--encoded_dims "0:2" \
|
|
||||||
--vocab_size 512 \
|
|
||||||
--max_episodes 100
|
|
||||||
```
|
|
||||||
|
|
||||||
### Example 2: Train with delta transform
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python src/lerobot/policies/pi05/train_fast_tokenizer.py \
|
|
||||||
--repo_id "lerobot/aloha_sim_insertion_human" \
|
|
||||||
--action_horizon 10 \
|
|
||||||
--encoded_dims "0:14" \
|
|
||||||
--delta_dims "0,1,2,3,4,5,6,7,8,9,10,11,12,13" \
|
|
||||||
--state_key "observation.state" \
|
|
||||||
--vocab_size 1024 \
|
|
||||||
--scale 10.0 \
|
|
||||||
--sample_fraction 0.2
|
|
||||||
```
|
|
||||||
|
|
||||||
### Example 3: Train on subset of dimensions
|
|
||||||
|
|
||||||
```bash
|
|
||||||
python src/lerobot/policies/pi05/train_fast_tokenizer.py \
|
|
||||||
--repo_id "lerobot/aloha_sim_insertion_human" \
|
|
||||||
--action_horizon 10 \
|
|
||||||
--encoded_dims "0:7" \
|
|
||||||
--vocab_size 1024 \
|
|
||||||
--output_dir "./my_tokenizer"
|
|
||||||
```
|
|
||||||
|
|
||||||
## Output
|
|
||||||
|
|
||||||
The script saves:
|
|
||||||
1. **Tokenizer files**: Trained FAST tokenizer (can be loaded with `AutoProcessor.from_pretrained()`)
|
|
||||||
2. **metadata.json**: Contains:
|
|
||||||
- Configuration parameters
|
|
||||||
- Compression statistics (compression ratio, token lengths)
|
|
||||||
- Training dataset information
|
|
||||||
|
|
||||||
## Understanding the Process
|
|
||||||
|
|
||||||
1. **Load Dataset**: Loads the LeRobotDataset from HuggingFace
|
|
||||||
2. **Extract Action Chunks**: Creates sliding windows of actions with specified horizon
|
|
||||||
3. **Apply Delta Transform**: (Optional) Computes action deltas relative to current state
|
|
||||||
4. **Select Encoded Dimensions**: Extracts only the dimensions to be encoded
|
|
||||||
5. **Normalize**: Applies quantile normalization ([q01, q99] → [-1, 1])
|
|
||||||
6. **Train Tokenizer**: Trains BPE tokenizer on DCT coefficients
|
|
||||||
7. **Compute Stats**: Reports compression ratio and token length statistics
|
|
||||||
8. **Save**: Saves tokenizer and metadata
|
|
||||||
|
|
||||||
## Notes
|
|
||||||
|
|
||||||
- **Normalization**: The script uses quantile normalization (q01, q99) from the dataset's statistics
|
|
||||||
- **Sampling**: To speed up training, you can sample a fraction of chunks per episode
|
|
||||||
- **Delta Transform**: Applied per-dimension to make actions relative to current state
|
|
||||||
- **Compression**: FAST uses DCT + BPE to compress action sequences efficiently
|
|
||||||
|
|
||||||
@@ -1,28 +0,0 @@
|
|||||||
#!/bin/bash
|
|
||||||
|
|
||||||
# FSDP training script for PI05 with aggressive memory optimization
|
|
||||||
# Use this for large models that OOM with standard DDP
|
|
||||||
|
|
||||||
accelerate launch --config_file /admin/home/jade_choghari/lerobot/fsdp_config.yaml \
|
|
||||||
$(which lerobot-train) \
|
|
||||||
--dataset.repo_id=local \
|
|
||||||
--dataset.root=/fsx/jade_choghari/data/libero \
|
|
||||||
--output_dir=/fsx/jade_choghari/outputs/libero_training_fsdp \
|
|
||||||
--job_name=libero_training_fsdp \
|
|
||||||
--policy.repo_id=jade_choghari/pi05-fast-libero-fsdp \
|
|
||||||
--policy.path=/fsx/jade_choghari/models/libero-pi-fast \
|
|
||||||
--policy.dtype=bfloat16 \
|
|
||||||
--steps=100000 \
|
|
||||||
--save_freq=10 \
|
|
||||||
--batch_size=8 \
|
|
||||||
--policy.device=cuda \
|
|
||||||
--policy.fast_only=true \
|
|
||||||
--policy.scheduler_warmup_steps=2000 \
|
|
||||||
--policy.scheduler_decay_steps=60000 \
|
|
||||||
--policy.scheduler_decay_lr=1e-5 \
|
|
||||||
--policy.gradient_checkpointing=false \
|
|
||||||
--wandb.enable=true \
|
|
||||||
--wandb.disable_artifact=true \
|
|
||||||
--wandb.project=pi05-libero-training-fsdp
|
|
||||||
|
|
||||||
|
|
||||||
@@ -1,24 +0,0 @@
|
|||||||
export CUDA_LAUNCH_BLOCKING=1
|
|
||||||
lerobot-train \
|
|
||||||
--dataset.repo_id=local \
|
|
||||||
--dataset.root=/fsx/jade_choghari/data/libero \
|
|
||||||
--output_dir=/fsx/jade_choghari/outputs/libero_training_fast_4 \
|
|
||||||
--job_name=libero_training_fast \
|
|
||||||
--policy.repo_id=jade_choghari/pi05-fast-libero \
|
|
||||||
--policy.path=/fsx/jade_choghari/models/pi05-base \
|
|
||||||
--policy.dtype=bfloat16 \
|
|
||||||
--steps=100000 \
|
|
||||||
--save_freq=20000 \
|
|
||||||
--batch_size=4 \
|
|
||||||
--policy.device=cuda \
|
|
||||||
--policy.fast_only=true \
|
|
||||||
--policy.scheduler_warmup_steps=1000 \
|
|
||||||
--policy.scheduler_decay_steps=30000 \
|
|
||||||
--policy.scheduler_decay_lr=1e-5 \
|
|
||||||
--policy.gradient_checkpointing=true \
|
|
||||||
--rename_map='{
|
|
||||||
"observation.images.image1": "observation.images.base_0_rgb",
|
|
||||||
"observation.images.image2": "observation.images.left_wrist_0_rgb",
|
|
||||||
}' \
|
|
||||||
--policy.empty_cameras=1 \
|
|
||||||
# /fsx/jade_choghari/.cache/huggingface/lerobot/jadechoghari/collect-data
|
|
||||||
@@ -1,15 +0,0 @@
|
|||||||
#!/bin/bash
|
|
||||||
#SBATCH --job-name=pi05-train
|
|
||||||
#SBATCH --time=24:00:00
|
|
||||||
#SBATCH --qos=high
|
|
||||||
#SBATCH --gres=gpu:8
|
|
||||||
#SBATCH --mem=256G
|
|
||||||
#SBATCH --partition=hopper-prod
|
|
||||||
#SBATCH --output=/fsx/jade_choghari/logs/%x-%j.out
|
|
||||||
#SBATCH --error=/fsx/jade_choghari/logs/%x-%j.err
|
|
||||||
|
|
||||||
srun \
|
|
||||||
--container-image=/fsx/michel_aractingi/docker_images/huggingface+lerobot-gpu+dev.sqsh \
|
|
||||||
--container-mounts=/fsx/jade_choghari \
|
|
||||||
--container-workdir=$HOME/lerobot \
|
|
||||||
bash /admin/home/jade_choghari/lerobot/src/lerobot/policies/pi05/train_multi.sh
|
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
#!/bin/bash
|
|
||||||
set -euxo pipefail
|
|
||||||
|
|
||||||
# Source YOUR Miniforge conda (mounted from FSX)
|
|
||||||
source /fsx/jade_choghari/miniforge3/etc/profile.d/conda.sh
|
|
||||||
|
|
||||||
conda activate lerobot
|
|
||||||
accelerate launch --mixed_precision=bf16 --multi_gpu --num_processes=8 \
|
|
||||||
$(which lerobot-train) \
|
|
||||||
--dataset.repo_id=local \
|
|
||||||
--dataset.root=/fsx/jade_choghari/data/libero \
|
|
||||||
--output_dir=/fsx/jade_choghari/outputs/libero_training_fast_mean_1 \
|
|
||||||
--job_name=libero_training_fast \
|
|
||||||
--policy.repo_id=jade_choghari/pi05-fast-libero \
|
|
||||||
--policy.path=/fsx/jade_choghari/models/pi05-base \
|
|
||||||
--policy.dtype=bfloat16 \
|
|
||||||
--steps=100000 \
|
|
||||||
--save_freq=20000 \
|
|
||||||
--batch_size=4 \
|
|
||||||
--policy.device=cuda \
|
|
||||||
--policy.fast_only=true \
|
|
||||||
--policy.scheduler_warmup_steps=4000 \
|
|
||||||
--policy.scheduler_decay_steps=100000 \
|
|
||||||
--policy.scheduler_decay_lr=1e-5 \
|
|
||||||
--policy.gradient_checkpointing=true \
|
|
||||||
--policy.chunk_size=10 \
|
|
||||||
--policy.n_action_steps=10 \
|
|
||||||
--policy.max_action_tokens=256 \
|
|
||||||
--rename_map='{
|
|
||||||
"observation.images.image1": "observation.images.base_0_rgb",
|
|
||||||
"observation.images.image2": "observation.images.left_wrist_0_rgb",
|
|
||||||
}' \
|
|
||||||
--policy.empty_cameras=1 \
|
|
||||||
--wandb.enable=true \
|
|
||||||
--wandb.disable_artifact=true \
|
|
||||||
--wandb.project=pi05-libero-training \
|
|
||||||
@@ -75,7 +75,7 @@ from .policy_robot_bridge import (
|
|||||||
RobotActionToPolicyActionProcessorStep,
|
RobotActionToPolicyActionProcessorStep,
|
||||||
)
|
)
|
||||||
from .rename_processor import RenameObservationsProcessorStep
|
from .rename_processor import RenameObservationsProcessorStep
|
||||||
from .tokenizer_processor import TokenizerProcessorStep, ActionTokenizerProcessorStep
|
from .tokenizer_processor import TokenizerProcessorStep
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"ActionProcessorStep",
|
"ActionProcessorStep",
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -173,7 +173,6 @@ def rollout(
|
|||||||
observation = env_preprocessor(observation)
|
observation = env_preprocessor(observation)
|
||||||
|
|
||||||
observation = preprocessor(observation)
|
observation = preprocessor(observation)
|
||||||
|
|
||||||
with torch.inference_mode():
|
with torch.inference_mode():
|
||||||
action = policy.select_action(observation)
|
action = policy.select_action(observation)
|
||||||
action = postprocessor(action)
|
action = postprocessor(action)
|
||||||
|
|||||||
@@ -62,7 +62,6 @@ def update_policy(
|
|||||||
accelerator: Accelerator,
|
accelerator: Accelerator,
|
||||||
lr_scheduler=None,
|
lr_scheduler=None,
|
||||||
lock=None,
|
lock=None,
|
||||||
postprocessor = None,
|
|
||||||
) -> tuple[MetricsTracker, dict]:
|
) -> tuple[MetricsTracker, dict]:
|
||||||
"""
|
"""
|
||||||
Performs a single training step to update the policy's weights.
|
Performs a single training step to update the policy's weights.
|
||||||
@@ -91,10 +90,6 @@ def update_policy(
|
|||||||
# Let accelerator handle mixed precision
|
# Let accelerator handle mixed precision
|
||||||
with accelerator.autocast():
|
with accelerator.autocast():
|
||||||
loss, output_dict = policy.forward(batch)
|
loss, output_dict = policy.forward(batch)
|
||||||
# action = policy.predict_action_chunk(batch)
|
|
||||||
# if postprocessor is not None:
|
|
||||||
# action = postprocessor(action)
|
|
||||||
# breakpoint()
|
|
||||||
# TODO(rcadene): policy.unnormalize_outputs(out_dict)
|
# TODO(rcadene): policy.unnormalize_outputs(out_dict)
|
||||||
|
|
||||||
# Use accelerator's backward method
|
# Use accelerator's backward method
|
||||||
@@ -156,7 +151,7 @@ def train(cfg: TrainPipelineConfig, accelerator: Accelerator | None = None):
|
|||||||
from accelerate.utils import DistributedDataParallelKwargs
|
from accelerate.utils import DistributedDataParallelKwargs
|
||||||
|
|
||||||
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
|
ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
|
||||||
accelerator = Accelerator(step_scheduler_with_optimizer=False, gradient_accumulation_steps=4, kwargs_handlers=[ddp_kwargs])
|
accelerator = Accelerator(step_scheduler_with_optimizer=False, kwargs_handlers=[ddp_kwargs])
|
||||||
|
|
||||||
init_logging(accelerator=accelerator)
|
init_logging(accelerator=accelerator)
|
||||||
|
|
||||||
@@ -212,7 +207,6 @@ def train(cfg: TrainPipelineConfig, accelerator: Accelerator | None = None):
|
|||||||
rename_map=cfg.rename_map,
|
rename_map=cfg.rename_map,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# Wait for all processes to finish policy creation before continuing
|
# Wait for all processes to finish policy creation before continuing
|
||||||
accelerator.wait_for_everyone()
|
accelerator.wait_for_everyone()
|
||||||
|
|
||||||
@@ -250,7 +244,6 @@ def train(cfg: TrainPipelineConfig, accelerator: Accelerator | None = None):
|
|||||||
**postprocessor_kwargs,
|
**postprocessor_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
if is_main_process:
|
if is_main_process:
|
||||||
logging.info("Creating optimizer and scheduler")
|
logging.info("Creating optimizer and scheduler")
|
||||||
optimizer, lr_scheduler = make_optimizer_and_scheduler(cfg, policy)
|
optimizer, lr_scheduler = make_optimizer_and_scheduler(cfg, policy)
|
||||||
@@ -350,7 +343,6 @@ def train(cfg: TrainPipelineConfig, accelerator: Accelerator | None = None):
|
|||||||
cfg.optimizer.grad_clip_norm,
|
cfg.optimizer.grad_clip_norm,
|
||||||
accelerator=accelerator,
|
accelerator=accelerator,
|
||||||
lr_scheduler=lr_scheduler,
|
lr_scheduler=lr_scheduler,
|
||||||
postprocessor=postprocessor,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Note: eval and checkpoint happens *after* the `step`th training update has completed, so we
|
# Note: eval and checkpoint happens *after* the `step`th training update has completed, so we
|
||||||
|
|||||||
@@ -26,15 +26,13 @@ OBS_IMAGES = OBS_IMAGE + "s"
|
|||||||
OBS_LANGUAGE = OBS_STR + ".language"
|
OBS_LANGUAGE = OBS_STR + ".language"
|
||||||
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
|
OBS_LANGUAGE_TOKENS = OBS_LANGUAGE + ".tokens"
|
||||||
OBS_LANGUAGE_ATTENTION_MASK = OBS_LANGUAGE + ".attention_mask"
|
OBS_LANGUAGE_ATTENTION_MASK = OBS_LANGUAGE + ".attention_mask"
|
||||||
OBS_LANGUAGE_HIGH_LEVEL_TASK = OBS_STR + ".user_prompt"
|
OBS_LANGUAGE_PROMPT = OBS_STR + ".prompt"
|
||||||
OBS_LANGUAGE_HIGH_LEVEL_TASK_TOKENS = OBS_LANGUAGE_HIGH_LEVEL_TASK + ".tokens"
|
OBS_LANGUAGE_PROMPT_TOKENS = OBS_LANGUAGE_PROMPT + ".tokens"
|
||||||
OBS_LANGUAGE_HIGH_LEVEL_TASK_ATTENTION_MASK = OBS_LANGUAGE_HIGH_LEVEL_TASK + ".attention_mask"
|
OBS_LANGUAGE_PROMPT_ATTENTION_MASK = OBS_LANGUAGE_PROMPT + ".attention_mask"
|
||||||
OBS_LANGUAGE_SUBTASK_ONLY = OBS_STR + ".subtask"
|
OBS_LANGUAGE_TARGET = OBS_STR + ".target"
|
||||||
OBS_LANGUAGE_SUBTASK_ONLY_TOKENS = OBS_LANGUAGE_SUBTASK_ONLY + ".tokens"
|
OBS_LANGUAGE_TARGET_TOKENS = OBS_LANGUAGE_TARGET + ".tokens"
|
||||||
OBS_LANGUAGE_SUBTASK_ONLY_ATTENTION_MASK = OBS_LANGUAGE_SUBTASK_ONLY + ".attention_mask"
|
OBS_LANGUAGE_TARGET_ATTENTION_MASK = OBS_LANGUAGE_TARGET + ".attention_mask"
|
||||||
ACTION = "action"
|
ACTION = "action"
|
||||||
ACTION_TOKENS = ACTION + ".tokens"
|
|
||||||
ACTION_TOKEN_MASK = ACTION + ".token_mask"
|
|
||||||
REWARD = "next.reward"
|
REWARD = "next.reward"
|
||||||
TRUNCATED = "next.truncated"
|
TRUNCATED = "next.truncated"
|
||||||
DONE = "next.done"
|
DONE = "next.done"
|
||||||
|
|||||||
Reference in New Issue
Block a user