From 12b076d41b779d8ab303f1cb3725668db8392898 Mon Sep 17 00:00:00 2001 From: Davide iafrate Date: Sun, 29 Mar 2026 23:34:09 +0000 Subject: [PATCH] feat: make torch an optional dependency for the batched simulator torch, torchdiffeq, roma, and opt-einsum are only required for the batched simulator (BatchedMultirotor, simulate_batch, BatchedSE3Control, etc.). Moving them to an optional extra avoids forcing all users to install PyTorch (~2 GB) when they only need the standard single-drone simulator. - pyproject.toml: remove the four packages from core dependencies and add a new `batched` extra; also include them in `all` - Wrap top-level torch/roma/torchdiffeq imports with try/except in every file that mixes batched and non-batched classes, so those modules remain importable without the extra installed Install the batched simulator with: pip install rotorpy[batched] Co-Authored-By: Claude Sonnet 4.6 --- pyproject.toml | 14 ++++++++++---- rotorpy/controllers/quadrotor_control.py | 7 +++++-- rotorpy/sensors/imu.py | 5 ++++- rotorpy/simulate.py | 7 +++++-- rotorpy/trajectories/circular_traj.py | 5 ++++- rotorpy/trajectories/hover_traj.py | 5 ++++- rotorpy/trajectories/lissajous_traj.py | 5 ++++- rotorpy/trajectories/minsnap.py | 5 ++++- rotorpy/trajectories/traj_template.py | 5 ++++- rotorpy/vehicles/multirotor.py | 9 ++++++--- rotorpy/wind/default_winds.py | 5 ++++- rotorpy/wind/dryden_winds.py | 5 ++++- 12 files changed, 58 insertions(+), 19 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 2120c69..4e0d424 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,14 +29,16 @@ dependencies = [ 'pandas', 'tqdm', 'gymnasium', - 'roma', # For batched sim - 'torch>=1.11.0', # For batched sim - 'torchdiffeq', # For batched sim - 'opt-einsum', # For batched sim 'timed_count', # Only for ardupilot sitl example ] [project.optional-dependencies] +batched = [ + 'torch>=1.11.0', + 'torchdiffeq', + 'roma', + 'opt-einsum', +] learning = [ 'stable_baselines3', 'tensorboard', @@ -64,6 +66,10 @@ px4 = [ 'pymavlink', ] all = [ + "torch>=1.11.0", + "torchdiffeq", + "roma", + "opt-einsum", "stable_baselines3", "tensorboard", "pytest", diff --git a/rotorpy/controllers/quadrotor_control.py b/rotorpy/controllers/quadrotor_control.py index e203485..9d8e05c 100644 --- a/rotorpy/controllers/quadrotor_control.py +++ b/rotorpy/controllers/quadrotor_control.py @@ -1,6 +1,9 @@ import numpy as np -import torch -import roma +try: + import torch + import roma +except ImportError: + pass from scipy.spatial.transform import Rotation class SE3Control(object): diff --git a/rotorpy/sensors/imu.py b/rotorpy/sensors/imu.py index f07148b..b2edf40 100644 --- a/rotorpy/sensors/imu.py +++ b/rotorpy/sensors/imu.py @@ -1,6 +1,9 @@ import numpy as np from scipy.spatial.transform import Rotation -import torch +try: + import torch +except ImportError: + pass import copy class Imu: diff --git a/rotorpy/simulate.py b/rotorpy/simulate.py index 4674a87..1552168 100644 --- a/rotorpy/simulate.py +++ b/rotorpy/simulate.py @@ -2,8 +2,11 @@ from enum import Enum import copy import numpy as np -import roma -import torch +try: + import roma + import torch +except ImportError: + pass from numpy.linalg import norm from scipy.spatial.transform import Rotation from time import perf_counter diff --git a/rotorpy/trajectories/circular_traj.py b/rotorpy/trajectories/circular_traj.py index a09a950..e6e37e5 100644 --- a/rotorpy/trajectories/circular_traj.py +++ b/rotorpy/trajectories/circular_traj.py @@ -1,5 +1,8 @@ import numpy as np -import torch +try: + import torch +except ImportError: + pass import sys class ThreeDCircularTraj(object): diff --git a/rotorpy/trajectories/hover_traj.py b/rotorpy/trajectories/hover_traj.py index 0b1b67c..653962b 100644 --- a/rotorpy/trajectories/hover_traj.py +++ b/rotorpy/trajectories/hover_traj.py @@ -1,5 +1,8 @@ import numpy as np -import torch +try: + import torch +except ImportError: + pass class HoverTraj(object): """ diff --git a/rotorpy/trajectories/lissajous_traj.py b/rotorpy/trajectories/lissajous_traj.py index 2b5db6a..d45833d 100644 --- a/rotorpy/trajectories/lissajous_traj.py +++ b/rotorpy/trajectories/lissajous_traj.py @@ -1,5 +1,8 @@ import numpy as np -import torch +try: + import torch +except ImportError: + pass """ Lissajous curves are defined by trigonometric functions parameterized in time. diff --git a/rotorpy/trajectories/minsnap.py b/rotorpy/trajectories/minsnap.py index 83308a6..dda4c8d 100644 --- a/rotorpy/trajectories/minsnap.py +++ b/rotorpy/trajectories/minsnap.py @@ -5,7 +5,10 @@ import cvxopt from scipy.linalg import block_diag from typing import List -import torch +try: + import torch +except ImportError: + pass def cvxopt_solve_qp(P, q, G=None, h=None, A=None, b=None): """ diff --git a/rotorpy/trajectories/traj_template.py b/rotorpy/trajectories/traj_template.py index e11516b..45d68e0 100644 --- a/rotorpy/trajectories/traj_template.py +++ b/rotorpy/trajectories/traj_template.py @@ -2,7 +2,10 @@ Imports """ import numpy as np -import torch +try: + import torch +except ImportError: + pass class TrajTemplate(object): """ diff --git a/rotorpy/vehicles/multirotor.py b/rotorpy/vehicles/multirotor.py index 8bc2171..1e28c27 100644 --- a/rotorpy/vehicles/multirotor.py +++ b/rotorpy/vehicles/multirotor.py @@ -8,9 +8,12 @@ from scipy.spatial.transform import Rotation as R # imports for Batched Dynamics -import torch -from torchdiffeq import odeint -import roma +try: + import torch + from torchdiffeq import odeint + import roma +except ImportError: + pass import time diff --git a/rotorpy/wind/default_winds.py b/rotorpy/wind/default_winds.py index 5a66e0d..a2d0e22 100644 --- a/rotorpy/wind/default_winds.py +++ b/rotorpy/wind/default_winds.py @@ -1,6 +1,9 @@ import numpy as np import sys -import torch +try: + import torch +except ImportError: + pass import math import random diff --git a/rotorpy/wind/dryden_winds.py b/rotorpy/wind/dryden_winds.py index fe045cc..1bb0bef 100644 --- a/rotorpy/wind/dryden_winds.py +++ b/rotorpy/wind/dryden_winds.py @@ -1,5 +1,8 @@ import numpy as np -import torch +try: + import torch +except ImportError: + pass import os import sys