Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 2 additions & 7 deletions crazyflow/dynamics/first_principles/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,11 +81,6 @@
moment of inertia, \(M\) is the \(3\times 4\) mixing matrix, and \(\mathbf{m}_z\) is its last row.
"""

from crazyflow.dynamics.first_principles.dynamics import (
Params,
dynamics,
sim_dynamics,
symbolic_dynamics,
)
from crazyflow.dynamics.first_principles.dynamics import Params, dynamics, symbolic_dynamics

__all__ = ["dynamics", "symbolic_dynamics", "sim_dynamics", "Params"]
__all__ = ["dynamics", "symbolic_dynamics", "Params"]
21 changes: 0 additions & 21 deletions crazyflow/dynamics/first_principles/dynamics.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,6 @@
from jax import Device

from crazyflow._typing import Array # To be changed to array_api_typing later
from crazyflow.sim.data import SimData


@supports(rotor_dynamics=True)
Expand Down Expand Up @@ -357,23 +356,3 @@ def create(drone: str, device: Device) -> Params:
drag_matrix=jnp.asarray(p["drag_matrix"], device=device),
rotor_dyn_coef=jnp.asarray([p["rotor_dyn_coef"]], device=device),
)


def sim_dynamics(data: SimData) -> SimData:
"""Compute the forces and torques from the first principle dynamics."""
params: Params = data.params
vel, _, acc, ang_acc, rotor_acc = dynamics(
pos=data.states.pos,
quat=data.states.quat,
vel=data.states.vel,
ang_vel=data.states.ang_vel,
cmd=data.controls.rotor_vel,
rotor_vel=data.states.rotor_vel,
dist_f=data.states.force,
dist_t=data.states.torque,
**params.__dict__,
)
states_deriv = data.states_deriv.replace(
vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc
)
return data.replace(states_deriv=states_deriv)
3 changes: 1 addition & 2 deletions crazyflow/dynamics/so_rpy/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,9 +35,8 @@
from crazyflow.dynamics.so_rpy.dynamics import (
Params,
dynamics,
sim_dynamics,
symbolic_dynamics,
symbolic_dynamics_euler,
)

__all__ = ["Params", "dynamics", "sim_dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"]
__all__ = ["Params", "dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"]
20 changes: 0 additions & 20 deletions crazyflow/dynamics/so_rpy/dynamics.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,6 @@
from jax import Device

from crazyflow._typing import Array # To be changed to array_api_typing later
from crazyflow.sim.data import SimData


@supports(rotor_dynamics=False)
Expand Down Expand Up @@ -358,22 +357,3 @@ def create(drone: str, device: Device) -> Params:
rpy_rates_coef=jnp.asarray(p["rpy_rates_coef"], device=device),
cmd_rpy_coef=jnp.asarray(p["cmd_rpy_coef"], device=device),
)


def sim_dynamics(data: SimData) -> SimData:
"""Compute the forces and torques from the so_rpy dynamics."""
params: Params = data.params
vel, _, acc, ang_acc = dynamics(
pos=data.states.pos,
quat=data.states.quat,
vel=data.states.vel,
ang_vel=data.states.ang_vel,
cmd=data.controls.attitude.cmd,
dist_f=data.states.force,
dist_t=data.states.torque,
**params.__dict__,
)
states_deriv = data.states_deriv.replace(
vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc
)
return data.replace(states_deriv=states_deriv)
3 changes: 1 addition & 2 deletions crazyflow/dynamics/so_rpy_rotor/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,9 +30,8 @@
from crazyflow.dynamics.so_rpy_rotor.dynamics import (
Params,
dynamics,
sim_dynamics,
symbolic_dynamics,
symbolic_dynamics_euler,
)

__all__ = ["Params", "dynamics", "sim_dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"]
__all__ = ["Params", "dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"]
21 changes: 0 additions & 21 deletions crazyflow/dynamics/so_rpy_rotor/dynamics.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,6 @@
from jax import Device

from crazyflow._typing import Array # To be changed to array_api_typing later
from crazyflow.sim.data import SimData


@supports(rotor_dynamics=True)
Expand Down Expand Up @@ -412,23 +411,3 @@ def create(drone: str, device: Device) -> Params:
rpy_rates_coef=jnp.asarray(p["rpy_rates_coef"], device=device),
cmd_rpy_coef=jnp.asarray(p["cmd_rpy_coef"], device=device),
)


def sim_dynamics(data: SimData) -> SimData:
"""Compute the forces and torques from the so_rpy_rotor dynamics."""
params: Params = data.params
vel, _, acc, ang_acc, rotor_acc = dynamics(
pos=data.states.pos,
quat=data.states.quat,
vel=data.states.vel,
ang_vel=data.states.ang_vel,
rotor_vel=data.states.rotor_vel,
cmd=data.controls.attitude.cmd,
dist_f=data.states.force,
dist_t=data.states.torque,
**params.__dict__,
)
states_deriv = data.states_deriv.replace(
vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc
)
return data.replace(states_deriv=states_deriv)
3 changes: 1 addition & 2 deletions crazyflow/dynamics/so_rpy_rotor_drag/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,9 +32,8 @@
from crazyflow.dynamics.so_rpy_rotor_drag.dynamics import (
Params,
dynamics,
sim_dynamics,
symbolic_dynamics,
symbolic_dynamics_euler,
)

__all__ = ["Params", "dynamics", "sim_dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"]
__all__ = ["Params", "dynamics", "symbolic_dynamics", "symbolic_dynamics_euler"]
21 changes: 0 additions & 21 deletions crazyflow/dynamics/so_rpy_rotor_drag/dynamics.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,6 @@
from jax import Device

from crazyflow._typing import Array # To be changed to array_api_typing later
from crazyflow.sim.data import SimData


# Additional symbols specific to these dynamics
Expand Down Expand Up @@ -449,23 +448,3 @@ def create(drone: str, device: Device) -> Params:
cmd_rpy_coef=jnp.asarray(p["cmd_rpy_coef"], device=device),
drag_matrix=jnp.asarray(p["drag_matrix"], device=device),
)


def sim_dynamics(data: SimData) -> SimData:
"""Compute the forces and torques from the so_rpy_rotor_drag dynamics."""
params: Params = data.params
vel, _, acc, ang_acc, rotor_acc = dynamics(
pos=data.states.pos,
quat=data.states.quat,
vel=data.states.vel,
ang_vel=data.states.ang_vel,
cmd=data.controls.attitude.cmd,
rotor_vel=data.states.rotor_vel,
dist_f=data.states.force,
dist_t=data.states.torque,
**params.__dict__,
)
states_deriv = data.states_deriv.replace(
vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc
)
return data.replace(states_deriv=states_deriv)
2 changes: 0 additions & 2 deletions crazyflow/sim/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -277,8 +277,6 @@ def create(
class SimData:
states: SimState
"""State of the simulation."""
states_deriv: SimStateDeriv
"""Derivative of the state of the simulation."""
controls: SimControls
"""Drone controller data."""
params: SimParams
Expand Down
89 changes: 89 additions & 0 deletions crazyflow/sim/dynamics.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,89 @@
"""Wrappers around the drone dynamics for `SimData` compatibility."""

from __future__ import annotations

from typing import TYPE_CHECKING

import jax.numpy as jnp

from crazyflow.dynamics import first_principles, so_rpy, so_rpy_rotor, so_rpy_rotor_drag
from crazyflow.sim.data import SimStateDeriv

if TYPE_CHECKING:
from crazyflow.sim.data import SimData


def first_principles_dynamics(data: SimData) -> SimStateDeriv:
"""Wrap the first principles dynamics."""
params: first_principles.Params = data.params
vel, _, acc, ang_acc, rotor_acc = first_principles.dynamics(
pos=data.states.pos,
quat=data.states.quat,
vel=data.states.vel,
ang_vel=data.states.ang_vel,
cmd=data.controls.rotor_vel,
rotor_vel=data.states.rotor_vel,
dist_f=data.states.force,
dist_t=data.states.torque,
**params.__dict__,
)
return SimStateDeriv(
vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc
)


def so_rpy_dynamics(data: SimData) -> SimStateDeriv:
"""Wrap the so_rpy dynamics."""
params: so_rpy.Params = data.params
vel, _, acc, ang_acc = so_rpy.dynamics(
pos=data.states.pos,
quat=data.states.quat,
vel=data.states.vel,
ang_vel=data.states.ang_vel,
cmd=data.controls.attitude.cmd,
dist_f=data.states.force,
dist_t=data.states.torque,
**params.__dict__,
)
rotor_acc = jnp.zeros_like(data.states.rotor_vel)
return SimStateDeriv(
vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc
)


def so_rpy_rotor_dynamics(data: SimData) -> SimStateDeriv:
"""Wrap the so_rpy_rotor dynamics."""
params: so_rpy_rotor.Params = data.params
vel, _, acc, ang_acc, rotor_acc = so_rpy_rotor.dynamics(
pos=data.states.pos,
quat=data.states.quat,
vel=data.states.vel,
ang_vel=data.states.ang_vel,
rotor_vel=data.states.rotor_vel,
cmd=data.controls.attitude.cmd,
dist_f=data.states.force,
dist_t=data.states.torque,
**params.__dict__,
)
return SimStateDeriv(
vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc
)


def so_rpy_rotor_drag_dynamics(data: SimData) -> SimStateDeriv:
"""Wrap the so_rpy_rotor_drag dynamics."""
params: so_rpy_rotor_drag.Params = data.params
vel, _, acc, ang_acc, rotor_acc = so_rpy_rotor_drag.dynamics(
pos=data.states.pos,
quat=data.states.quat,
vel=data.states.vel,
ang_vel=data.states.ang_vel,
cmd=data.controls.attitude.cmd,
rotor_vel=data.states.rotor_vel,
dist_f=data.states.force,
dist_t=data.states.torque,
**params.__dict__,
)
return SimStateDeriv(
vel=vel, ang_vel=data.states.ang_vel, acc=acc, ang_acc=ang_acc, rotor_acc=rotor_acc
)
43 changes: 20 additions & 23 deletions crazyflow/sim/integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@
from jax.numpy import vectorize
from jax.scipy.spatial.transform import Rotation as R

from crazyflow.sim.data import SimData
from crazyflow.sim.data import SimData, SimStateDeriv


class Integrator(StrEnum):
Expand All @@ -20,7 +20,7 @@ class Integrator(StrEnum):
default = euler


def euler(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData:
def euler(data: SimData, deriv_fn: Callable[[SimData], SimStateDeriv]) -> SimData:
"""Explicit Euler integration.

Args:
Expand All @@ -33,7 +33,7 @@ def euler(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData:
return integrate(data, deriv_fn(data), dt=1 / data.core.freq)


def rk4(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData:
def rk4(data: SimData, deriv_fn: Callable[[SimData], SimStateDeriv]) -> SimData:
"""Runge-Kutta 4 integration.

Args:
Expand All @@ -44,14 +44,14 @@ def rk4(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData:
The integrated simulation data structure.
"""
dt = 1 / data.core.freq
data_d1 = deriv_fn(data)
data_d2 = deriv_fn(integrate(data, data_d1, dt=dt / 2))
data_d3 = deriv_fn(integrate(data, data_d2, dt=dt / 2))
data_d4 = deriv_fn(integrate(data, data_d3, dt=dt))
return integrate(data, rk4_average(data_d1, data_d2, data_d3, data_d4), dt=dt)
k1 = deriv_fn(data)
k2 = deriv_fn(integrate(data, k1, dt=dt / 2))
k3 = deriv_fn(integrate(data, k2, dt=dt / 2))
k4 = deriv_fn(integrate(data, k3, dt=dt))
return integrate(data, rk4_average(k1, k2, k3, k4), dt=dt)


def symplectic_euler(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> SimData:
def symplectic_euler(data: SimData, deriv_fn: Callable[[SimData], SimStateDeriv]) -> SimData:
"""Symplectic Euler integration.

Args:
Expand All @@ -61,24 +61,21 @@ def symplectic_euler(data: SimData, deriv_fn: Callable[[SimData], SimData]) -> S
return integrate_symplectic(data, deriv_fn(data), dt=1 / data.core.freq)


def rk4_average(k1: SimData, k2: SimData, k3: SimData, k4: SimData) -> SimData:
def rk4_average(
k1: SimStateDeriv, k2: SimStateDeriv, k3: SimStateDeriv, k4: SimStateDeriv
) -> SimStateDeriv:
"""Average four derivatives according to the RK4 rules."""
data = k1
k1, k2, k3, k4 = k1.states_deriv, k2.states_deriv, k3.states_deriv, k4.states_deriv
states_deriv = jax.tree.map(
lambda x1, x2, x3, x4: (x1 + 2 * x2 + 2 * x3 + x4) / 6, k1, k2, k3, k4
)
return data.replace(states_deriv=states_deriv)
return jax.tree.map(lambda x1, x2, x3, x4: (x1 + 2 * x2 + 2 * x3 + x4) / 6, k1, k2, k3, k4)


def integrate(data: SimData, deriv: SimData, dt: float) -> SimData:
def integrate(data: SimData, deriv: SimStateDeriv, dt: float) -> SimData:
"""Integrate the dynamics forward in time."""
states, states_deriv = data.states, deriv.states_deriv
states = data.states

pos, quat, vel, ang_vel = states.pos, states.quat, states.vel, states.ang_vel
rotor_vel = states.rotor_vel
dpos, drot = states_deriv.vel, states_deriv.ang_vel
dvel, dang_vel, drotor_vel = states_deriv.acc, states_deriv.ang_acc, states_deriv.rotor_acc
dpos, drot = deriv.vel, deriv.ang_vel
dvel, dang_vel, drotor_vel = deriv.acc, deriv.ang_acc, deriv.rotor_acc

next_pos, next_quat, next_vel, next_ang_vel, next_rotor_vel = _integrate(
pos, quat, vel, ang_vel, rotor_vel, dpos, drot, dvel, dang_vel, drotor_vel, dt
Expand All @@ -89,13 +86,13 @@ def integrate(data: SimData, deriv: SimData, dt: float) -> SimData:
return data.replace(states=states)


def integrate_symplectic(data: SimData, deriv: SimData, dt: float) -> SimData:
def integrate_symplectic(data: SimData, deriv: SimStateDeriv, dt: float) -> SimData:
"""Integrate the dynamics forward in time."""
states, states_deriv = data.states, deriv.states_deriv
states = data.states

pos, quat, vel, ang_vel = states.pos, states.quat, states.vel, states.ang_vel
rotor_vel = states.rotor_vel
dvel, dang_vel, drotor_vel = states_deriv.vel, states_deriv.ang_vel, states_deriv.rotor_acc
dvel, dang_vel, drotor_vel = deriv.acc, deriv.ang_acc, deriv.rotor_acc

next_pos, next_quat, next_vel, next_ang_vel, next_rotor_vel = _integrate_symplectic(
pos, quat, vel, ang_vel, rotor_vel, dvel, dang_vel, drotor_vel, dt
Expand Down
Loading
Loading