| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| """This script demonstrates how to train a Diffusion Policy on the PushT environment, |
| using a dataset processed in streaming mode.""" |
|
|
| from pathlib import Path |
|
|
| import torch |
|
|
| from lerobot.configs import FeatureType |
| from lerobot.datasets import LeRobotDatasetMetadata, StreamingLeRobotDataset |
| from lerobot.policies import make_pre_post_processors |
| from lerobot.policies.act import ACTConfig, ACTPolicy |
| from lerobot.utils.constants import ACTION |
| from lerobot.utils.feature_utils import dataset_to_policy_features |
|
|
|
|
| def main(): |
| |
| output_directory = Path("outputs/train/example_streaming_dataset") |
| output_directory.mkdir(parents=True, exist_ok=True) |
|
|
| |
| device = ( |
| torch.device("cuda") |
| if torch.cuda.is_available() |
| else torch.device("mps") |
| if torch.backends.mps.is_available() |
| else torch.device("cpu") |
| ) |
| print(f"Using device: {device}") |
|
|
| training_steps = 10 |
| log_freq = 1 |
|
|
| dataset_id = "lerobot/droid_1.0.1" |
| dataset_metadata = LeRobotDatasetMetadata(dataset_id) |
| features = dataset_to_policy_features(dataset_metadata.features) |
| output_features = {key: ft for key, ft in features.items() if ft.type is FeatureType.ACTION} |
| input_features = {key: ft for key, ft in features.items() if key not in output_features} |
|
|
| |
| cfg = ACTConfig(input_features=input_features, output_features=output_features) |
| policy = ACTPolicy(cfg) |
| policy.train() |
| policy.to(device) |
| preprocessor, postprocessor = make_pre_post_processors(cfg, dataset_stats=dataset_metadata.stats) |
|
|
| |
| |
| delta_timestamps = { |
| ACTION: [t / dataset_metadata.fps for t in range(cfg.n_action_steps)], |
| } |
|
|
| |
| |
| dataset = StreamingLeRobotDataset(dataset_id, delta_timestamps=delta_timestamps, tolerance_s=1e-3) |
|
|
| optimizer = torch.optim.Adam(policy.parameters(), lr=1e-4) |
| dataloader = torch.utils.data.DataLoader( |
| dataset, |
| num_workers=4, |
| batch_size=16, |
| pin_memory=device.type != "cpu", |
| drop_last=True, |
| prefetch_factor=2, |
| ) |
|
|
| |
| step = 0 |
| done = False |
| while not done: |
| for batch in dataloader: |
| batch = preprocessor(batch) |
| loss, _ = policy.forward(batch) |
| loss.backward() |
| optimizer.step() |
| optimizer.zero_grad() |
|
|
| if step % log_freq == 0: |
| print(f"step: {step} loss: {loss.item():.3f}") |
| step += 1 |
| if step >= training_steps: |
| done = True |
| break |
|
|
| |
| policy.save_pretrained(output_directory) |
| preprocessor.save_pretrained(output_directory) |
| postprocessor.save_pretrained(output_directory) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|