Commit 0d6cbf4f authored by Ibrahim's avatar Ibrahim
Browse files

PID controller has option to persist (or not) state.;

made pyscurve optional
parent 57473f84
Loading
Loading
Loading
Loading
+5 −0
Changes for multirotor/controller/__init__.py: 5 added lines, 0 removed lines.
Original line number Diff line number Diff line
@@ -8,3 +8,8 @@ from .pid import (
    AltRateController,
    Controller
)
try:
    import pyscurve
    from .scurves import SCurveController
except ImportError:
    pass
+91 −59
Changes for multirotor/controller/pid.py: 91 added lines, 59 removed lines.
Original line number Diff line number Diff line
@@ -72,8 +72,9 @@ class PIDController:
        self.err_i = np.atleast_1d(np.zeros_like(self.k_i))
        self.err_d = np.atleast_1d(np.zeros_like(self.k_d))
        self.err = np.atleast_1d(np.zeros_like(self.k_p))
        self.dtype = self.err_p.dtype
        if self.max_err_i is None:
            self.max_err_i = np.inf
            self.max_err_i = np.atleast_1d(np.inf, self.dtype)
        else:
            self.max_err_i = np.asarray(self.max_err_i, dtype=self.err.dtype)
        self.action = np.zeros_like(self.err)
@@ -81,7 +82,7 @@ class PIDController:


    def reset(self):
        self.action = np.zeros_like(self.err)
        self.action = np.zeros_like(self.err, self.dtype)
        self.err *= 0
        self.err_p *= 0
        self.err_i *= 0
@@ -122,7 +123,7 @@ class PIDController:

    def step(
        self, reference: np.ndarray, measurement: np.ndarray, dt: float=1.,
        ref_is_error: bool=False
        ref_is_error: bool=False, persist: bool=True
    ) -> np.ndarray:
        """
        Calculate the output, based on the current measurement and the reference
@@ -136,6 +137,8 @@ class PIDController:
            The actual measurement(s).
        ref_is_error: bool
            Whether to interpret the reference input as the error.
        persist: bool
            Whether to store the current state for the next step.

        Returns
        -------
@@ -145,17 +148,20 @@ class PIDController:
        if ref_is_error:
            err = reference
        else:
            self.reference = reference
            err = reference - measurement
        self.err_p = self.k_p * err
        self.err_i = self.k_i * np.clip(
        err_p = self.k_p * err
        err_i = self.k_i * np.clip(
            self.err_i + trapezoid((self.err, err), dx=dt, axis=0),
            a_min=-self.max_err_i, a_max=self.max_err_i
        )
        self.err_d = self.k_d * (err - self.err) / dt
        err_d = self.k_d * (err - self.err) / dt
        action = self.err_p + self.err_i + self.err_d
        if persist:
            self.err_p, self.err_i, self.err_d, self.action = \
                err_p, err_i, err_d, action
            self.reference = reference
            self.err = err
        self.action = self.err_p + self.err_i + self.err_d
        return self.action
        return action



@@ -179,9 +185,10 @@ class PosController(PIDController):


    def __post_init__(self):
        self.k_p = np.ones(2) * np.asarray(self.k_p)
        self.k_i = np.ones(2) * np.asarray(self.k_i)
        self.k_d = np.ones(2) * np.asarray(self.k_d)
        self.dtype = self.vehicle.dtype
        self.k_p = np.ones(2) * np.asarray(self.k_p, self.dtype)
        self.k_i = np.ones(2) * np.asarray(self.k_i, self.dtype)
        self.k_d = np.ones(2) * np.asarray(self.k_d, self.dtype)
        if self.leashing or self.square_root_scaling:
            self.k_p[:] = 0.5 * self.max_jerk / self.max_acceleration
        super().__post_init__()
@@ -206,8 +213,9 @@ class PosController(PIDController):
        return leash


    def step(self, reference, measurement, dt):
    def step(self, reference, measurement, dt, persist: bool=True):
        # inertial frame velocity
        if persist:
            self.reference = reference
        err = reference - measurement
        err_len = np.linalg.norm(err)
@@ -216,28 +224,31 @@ class PosController(PIDController):
            err_unit = err / (err_len + 1e-6)
            err_len = min(err_len, self.leash)
            err = err_unit * err_len
            self.err = err
            if persist: self.err = err
            velocity = self.err_p = self.k_p * err
        if err_len > 0. and self.square_root_scaling:
            velocity = np.zeros_like(self.k_p)
            velocity[0] = sqrt_control(err[0], self.k_p[0], self.max_acceleration, dt)
            velocity[1] = sqrt_control(err[1], self.k_p[1], self.max_acceleration, dt)
            if persist:
                self.err = err
                self.err_p = velocity
        else:
            velocity = super().step(reference, measurement, dt=dt)
            velocity = super().step(reference, measurement, dt=dt, persist=persist)
        # convert to body-frame velocity
        roll, pitch, yaw = self.vehicle.orientation
        cos, sin = np.cos(yaw), np.sin(yaw)
        rot = np.asarray([
            [cos,   sin],
            [-sin, cos],
        ])
        ], self.dtype)
        ref_velocity = rot @ velocity
        ref_velocity_mag = np.linalg.norm(ref_velocity)
        ref_velocity_unit = ref_velocity / (ref_velocity_mag + 1e-6)
        self.action = ref_velocity_unit * min(ref_velocity_mag, self.max_velocity)
        return self.action
        action = ref_velocity_unit * min(ref_velocity_mag, self.max_velocity)
        if persist:
            self.action = action
        return action


@dataclass
@@ -259,22 +270,24 @@ class VelController(PIDController):


    def __post_init__(self):
        self.k_p = np.ones(2) * np.asarray(self.k_p)
        self.k_i = np.ones(2) * np.asarray(self.k_i)
        self.k_d = np.ones(2) * np.asarray(self.k_d)
        self.dtype = self.vehicle.dtype
        self.k_p = np.ones(2) * np.asarray(self.k_p, self.dtype)
        self.k_i = np.ones(2) * np.asarray(self.k_i, self.dtype)
        self.k_d = np.ones(2) * np.asarray(self.k_d, self.dtype)
        super().__post_init__()
        self._params = tuple(list(self._params) + ['max_tilt'])


    def step(self, reference, measurement, dt):
    def step(self, reference, measurement, dt, persist: bool=True):
        # desired pitch, roll
        pitch_roll = super().step(reference, measurement, dt=dt)
        pitch_roll = super().step(reference, measurement, dt=dt, persist=persist)
        # ctrl[0] -> x dir -> pitch -> forward
        # ctrl[1] -> y dir -> roll -> lateral
        pitch_roll[0:2] = np.clip(pitch_roll[0:2], a_min=-self.max_tilt, a_max=self.max_tilt)
        pitch_roll[1] *= -1 # +y motion requires negative roll
        self.action = pitch_roll
        return self.action # desired pitch, roll
        action = pitch_roll
        if persist: self.action = action
        return action # desired pitch, roll



@@ -295,9 +308,10 @@ class AttController(PIDController):


    def __post_init__(self):
        self.k_p = np.ones(3) * np.asarray(self.k_p)
        self.k_i = np.ones(3) * np.asarray(self.k_i)
        self.k_d = np.ones(3) * np.asarray(self.k_d)
        self.dtype = self.vehicle.dtype
        self.k_p = np.ones(3) * np.asarray(self.k_p, self.dtype)
        self.k_i = np.ones(3) * np.asarray(self.k_i, self.dtype)
        self.k_d = np.ones(3) * np.asarray(self.k_d, self.dtype)
        if self.square_root_scaling:
            self.k_p[:] = self.max_jerk / self.max_acceleration
        super().__post_init__()
@@ -305,8 +319,9 @@ class AttController(PIDController):
            ['max_acceleration', 'max_jerk', 'square_root_scaling'])


    def step(self, reference, measurement, dt):
    def step(self, reference, measurement, dt, persist: bool=True):
        err = reference - measurement
        if persist:
            self.reference = reference
        err_len = np.linalg.norm(err)
        if self.square_root_scaling and err_len > 0:
@@ -315,12 +330,14 @@ class AttController(PIDController):
            velocity[1] = sqrt_control(err[1], self.k_p[1], self.max_acceleration, dt)
            velocity[2] = sqrt_control(err[2], self.k_p[2], self.max_acceleration, dt)
            # velocity = (np.abs(err) / err_len) * velocity
            if persist:
                self.err = err
                self.err_p = velocity
        else:
            velocity = super().step(reference=reference, measurement=measurement, dt=dt)
        self.action = velocity # Euler rate
        return self.action
            velocity = super().step(reference=reference, measurement=measurement, dt=dt, persist=persist)
        action = velocity # Euler rate
        if persist: self.action = action
        return action



@@ -342,14 +359,15 @@ class RateController(PIDController):


    def __post_init__(self):
        self.k_p = np.ones(3) * np.asarray(self.k_p)
        self.k_i = np.ones(3) * np.asarray(self.k_i)
        self.k_d = np.ones(3) * np.asarray(self.k_d)
        self.dtype = self.vehicle.dtype
        self.k_p = np.ones(3) * np.asarray(self.k_p, self.dtype)
        self.k_i = np.ones(3) * np.asarray(self.k_i, self.dtype)
        self.k_d = np.ones(3) * np.asarray(self.k_d, self.dtype)
        super().__post_init__()
        self._params = tuple(list(self._params) + ['max_acceleration'])


    def step(self, reference, measurement, dt):
    def step(self, reference, measurement, dt, persist: bool=True):
        # desired angular velocity
        ref = euler_to_angular_rate(reference, self.vehicle.orientation)
        # ref = reference
@@ -358,12 +376,14 @@ class RateController(PIDController):
        # mea = measurement
        # prescribed change in velocity i.e. angular acc
        self.action = np.clip(
            super().step(reference=ref, measurement=mea, dt=dt),
            super().step(reference=ref, measurement=mea, dt=dt, persist=persist),
            -self.max_acceleration, self.max_acceleration
        )
        # torque = moment of inertia . angular_acceleration
        self.action = self.vehicle.params.inertia_matrix.dot(self.action)
        return self.action
        action = self.vehicle.params.inertia_matrix.dot(self.action)
        if persist:
            self.action = action
        return action



@@ -382,20 +402,23 @@ class AltController(PIDController):


    def __post_init__(self):
        self.k_p = np.ones(1) * np.asarray(self.k_p)
        self.dtype = self.vehicle.dtype
        self.k_p = np.asarray(self.k_p, self.dtype)
        # Alt controller is strictly a P controller
        self.k_i = np.zeros(1) * np.asarray(self.k_i)
        self.k_d = np.zeros(1) * np.asarray(self.k_d)
        self.k_i = np.zeros(1, self.dtype) * np.asarray(self.k_i, self.dtype)
        self.k_d = np.zeros(1, self.dtype) * np.asarray(self.k_d, self.dtype)
        super().__post_init__()
        self._params = tuple(list(self._params) + ['max_velocity'])


    def step(
        self, reference: np.ndarray, measurement: np.ndarray, dt: float = 1
        self, reference: np.ndarray, measurement: np.ndarray, dt: float = 1, persist: bool=True
    ) -> np.ndarray:
        self.action = super().step(reference, measurement, dt)
        self.action = np.clip(self.action, a_min=-self.max_velocity, a_max=self.max_velocity)
        return self.action
        self.action = super().step(reference, measurement, dt, persist=persist)
        action = np.clip(self.action, a_min=-self.max_velocity, a_max=self.max_velocity)
        if persist:
            self.action = action
        return action


            
@@ -408,17 +431,23 @@ class AltRateController(PIDController):
    vehicle: Multirotor


    def step(self, reference, measurement, dt):
    def __post_init__(self):
        self.dtype = self.vehicle.dtype
        super().__post_init__() # TODO: set dtypes of errs


    def step(self, reference, measurement, dt, persist: bool=True):
            roll, pitch, yaw = self.vehicle.orientation
            # change in velocity i.e. acceleration
            ctrl = super().step(reference=reference, measurement=measurement, dt=dt)
            ctrl = super().step(reference=reference, measurement=measurement, dt=dt, persist=persist)
            # convert acceleration to required z-force, given orientation
            ctrl = self.vehicle.params.mass * (
            action = self.vehicle.params.mass * (
                    ctrl / (np.cos(roll) * np.cos(pitch))
                ) + \
                self.vehicle.weight
            self.action = ctrl
            return ctrl # thrust force
            if persist:
                self.action = action
            return action # thrust force



@@ -469,6 +498,7 @@ class Controller:
        self.ctrl_z = ctrl_z
        self.ctrl_vz = ctrl_vz
        self.vehicle = self.ctrl_a.vehicle
        self.dtype = self.vehicle.dtype
        self.period_p = period_p
        self.period_a = period_a
        self.period_z = period_z
@@ -543,7 +573,7 @@ class Controller:

    def step(
        self, reference: np.ndarray, measurement=None, ref_is_error: bool=False,
        feed_forward_velocity: np.ndarray=None
        feed_forward_velocity: np.ndarray=None, persist: bool=True
    ):
        if ref_is_error:
            error = reference
@@ -558,27 +588,29 @@ class Controller:

        if self.n % self.steps_z == 0:
            dt = self.steps_z * self.vehicle.simulation.dt
            ref_vel_z = self.ctrl_z.step(ref_z, self.vehicle.position[2:], dt=dt)
            self.thrust = self.ctrl_vz.step(ref_vel_z, self.vehicle.inertial_velocity[2:], dt=dt)
            ref_vel_z = self.ctrl_z.step(ref_z, self.vehicle.position[2:], dt=dt, persist=persist)
            self.thrust = self.ctrl_vz.step(ref_vel_z, self.vehicle.inertial_velocity[2:], dt=dt, persist=persist)

        if self.n % self.steps_p == 0:
            dt = self.steps_p * self.vehicle.simulation.dt
            self._pid_vel = self._ref_vel = self.ctrl_p.step(ref_xy, self.vehicle.position[:2], dt=dt)
            self._pid_vel = self._ref_vel = self.ctrl_p.step(ref_xy, self.vehicle.position[:2], dt=dt, persist=persist)
            if feed_forward_velocity is not None:
                self._ref_vel = (self.feedforward_weight * feed_forward_velocity[:2]) + (1 - self.feedforward_weight) * self._pid_vel
            self._ref_vel = np.clip(self._ref_vel, -self.ctrl_p.max_velocity, self.ctrl_p.max_velocity)
        
        if self.n % self.steps_a == 0:
            dt = self.steps_a * self.vehicle.simulation.dt
            self._pitch_roll = self.ctrl_v.step(self._ref_vel, self.vehicle.velocity[:2], dt=dt)
            self._pitch_roll = self.ctrl_v.step(self._ref_vel, self.vehicle.velocity[:2], dt=dt, persist=persist)
            ref_orientation = np.asarray([self._pitch_roll[1], self._pitch_roll[0], ref_yaw])
            ref_rate = self.ctrl_a.step(ref_orientation, self.vehicle.orientation, dt=dt)
            self.torques = self.ctrl_r.step(ref_rate, self.vehicle.euler_rate, dt=dt)
            self.torques = self.ctrl_r.step(ref_rate, self.vehicle.euler_rate, dt=dt, persist=persist)

        self.action = np.asarray([*self.thrust, *self.torques])
        action = np.asarray([*self.thrust, *self.torques], self.dtype)
        if persist:
            self.n += 1
            self.t = self.vehicle.t
        return self.action
            self.action = action
        return action


    def predict(self, ref, deterministic=True):
+276 −0
Changes for multirotor/controller/scurves.py: 276 added lines, 0 removed lines.
Original line number Diff line number Diff line
import numpy as np
from numba import njit
from typing import Dict, Union

from .pid import Controller
from ..helpers import get_vehicle_ability
from pyscurve import ScurvePlanner
from pyscurve.scurve import PlanningError



class SCurveController:

    def __init__(self, ctrl: Controller):
        self.ctrl = ctrl
        self.vehicle = self.ctrl.vehicle
        self.ctrl_p = self.ctrl.ctrl_p
        self.ctrl_v = self.ctrl.ctrl_v
        self.ctrl_a = self.ctrl.ctrl_a
        self.ctrl_r = self.ctrl.ctrl_r
        self.ctrl_z = self.ctrl.ctrl_z
        self.ctrl_vz = self.ctrl.ctrl_vz
        self.reset()


    def get_params(self):
        p = dict(
            steps=self.steps,
            max_velocity=self.max_velocity, max_acceleration=self.max_acceleration,
            max_jerk=self.max_jerk,
        )
        p.update(ctrl=self.ctrl.get_params())
        return p


    def set_params(self,**params: Dict[str, Union[np.ndarray, bool, float, int]]):
        ctrl_params = params.get('ctrl', None)
        # if controller params are not nested under 'ctrl, make a dict
        # containing those params...
        if ctrl_params is None:
            # ...using keys from self.ctrl
            ctrls = self.ctrl.get_params().keys()
            ctrl_params = {name: params.get(name, {}) for name in ctrls}
            # delete those param names from the params dict. The remaning
            # params are for this (self) controller
            for name in ctrls:
                if name in params:
                    del params[name]
        else:
            del params['ctrl']
        self.ctrl.set_params(**ctrl_params)

        for name, value in params.items():
            if hasattr(self, name):
                setattr(self, name, value)
        # these parameters are dictated by the controller
        self.max_velocity = self.ctrl.ctrl_p.max_velocity
        self.steps = self.ctrl.steps_p


    def reset(self):
        self.ctrl.reset()
        # these parameters are dictated by the controller
        self.steps = self.ctrl.steps_p
        self.max_velocity = self.ctrl.ctrl_p.max_velocity

        self.max_acceleration = get_vehicle_ability(
            self.vehicle.params, self.vehicle.simulation,
            self.ctrl_v.max_tilt, self.ctrl_r.max_acceleration,
            max_rads=700.
        )['max_acc_xy']
        self.max_jerk = self.max_acceleration # TODO: arbitrary
        # max accelerateion is determined from the physical properties of the vehicle
        self.ctrl_p.max_acceleration = self.max_acceleration
        self.ctrl_p.max_jerk = self.max_jerk

        self.planner = ScurvePlanner()
        self.n = 0
        self.n_since_replan = 0
        self.ref_xy = np.empty(2, self.ctrl.vehicle.dtype)


    def step(self, reference: np.ndarray, ref_is_error=False):
        ref_xy = reference[:2]
        if self.n==0 or not np.array_equal(ref_xy, self.ref_xy):
            self.ref_xy = ref_xy
            try:
                self.traj = self.planner.plan_trajectory(
                    q0=self.ctrl.vehicle.position[:2],
                    q1=ref_xy,
                    v0=min(
                        np.linalg.norm(self.max_velocity),
                        np.linalg.norm(self.ctrl.vehicle.velocity[:2])
                    ) * (self.ctrl.vehicle.velocity[:2] / np.linalg.norm(self.ctrl.vehicle.velocity[:2])),
                    v1=(ref_xy-self.ctrl.vehicle.position[:2]) * self.max_velocity \
                        / np.linalg.norm(ref_xy-self.ctrl.vehicle.position[:2]),
                    v_max=self.max_velocity,
                    a_max=self.max_acceleration,
                    j_max=self.max_jerk
                )
                self.n_since_replan = 0
            except PlanningError:
                pass
        if self.n % self.steps == 0:
            target = self.traj((self.n_since_replan + self.steps) * self.ctrl.vehicle.simulation.dt)
            # point = target[:, 2]
            self._ref_vel = target[:, 1]
            # self._ref = np.concatenate((point, state[2:4])) # position, yaw
        self.ctrl.step(reference, ref_is_error=False,
                       feed_forward_velocity=self._ref_vel)
        self.n += 1
        self.n_since_replan += 1
        return self.ctrl.action


    @property
    def action(self):
        return self.ctrl.action
    @property
    def reference(self):
        return self.ctrl.reference
    @property
    def feedforward_weight(self):
        return self.ctrl.feedforward_weight




def acc_1d(amax: float, vmax: float, v0: float, v1: float, disp: float, dt: float=1e-2):
    """Solve 1d kinematics problem"""
    sign = np.sign(disp)
    dist = np.abs(disp)
    relvel = v1 - v0 # desired change in velocity
    # s = ut + 0.5 a(t) t^2 => 2 (s - ut) / t^2 = a(t)
    # 2 a(t) s = v^2 - u^2
    # distances to accelerate and decelerate from max vel to current
    # and target velocities
    dist_v0_vmax = np.abs((vmax**2 - v0**2 ) / (2 * amax))
    dist_vmax_v1 = np.abs((vmax**2 - v1**2 ) / (2 * amax))
    dist_0_v1 = (v1**2) / (2 * amax)
    dist_vmax_0 = (vmax**2) / (2 * amax)

    if dist > dist_vmax_v1:
        pass
    else:
        vmax = np.sqrt(((2 * amax * dist) + v0**2 + v1**2) / 2)

    if v0 < vmax:
        should_decelerate = False
    elif v0 > vmax:
        should_decelerate = True
    else:
        return 0
    amax = np.clip(np.abs((v0 - vmax)) / dt, 0, amax)

    if should_decelerate:
        return -sign * amax
    return sign * amax



# @njit
def _acc_1d(amax: float, vmax: float, v0: float, v1: float, disp: float):
    """Solve 1d kinematics problem"""
    sign = np.sign(disp)
    dist = np.abs(disp)
    relvel = v1 - v0 # desired change in velocity
    # s = ut + 0.5 a(t) t^2 => 2 (s - ut) / t^2 = a(t)
    # 2 a(t) s = v^2 - u^2
    # distances to accelerate and decelerate from max vel to current
    # and target velocities
    dist_v0_vmax = np.abs((vmax**2 - v0**2 ) / (2 * amax))
    dist_vmax_v1 = np.abs((vmax**2 - v1**2 ) / (2 * amax))
    dist_0_v1 = (v1**2) / (2 * amax)
    dist_vmax_0 = (vmax**2) / (2 * amax)
    # if at relative rest
    if relvel==0:
        # if self.log: print('Relative rest')
        if dist > 0:
            # if self.log: print('Accelerate. Distance > 0')
            should_decelerate = False
            amax = np.sqrt(v0**2 / (2 * dist))
        else:
            # if self.log: print('No change. Distance == 0')
            should_decelerate = False # doesn't matter, since amax=0
            amax = 0
    # if approaching. For e.g. v0=1ms, disp=10
    elif (np.sign(v0) == sign):
        # if self.log: print('Approaching')
        # moving in opposite directions to desired velocity
        # need to decelerate to 0 and reverse to match velocity at some point, except:
        # if any one velocity is 0, and they're approaching, it's like catch-up
        if (np.sign(v0) != np.sign(v1)) and (v1!=0 and v0!=0):
            # distance needed to overshoot enough such that can reverse to the 
            # currect velocity:
            dist_net = dist + dist_0_v1
            # distance to halt from max vel and accelerate to v1, changing direction
            # disp_net = dist_vmax_0 - dist_0_v1
            # If close enough that need to stop and reverse
            if dist_net > dist_v0_vmax + dist_vmax_0:
                should_decelerate = False
            else:
                should_decelerate = True
                # 2 a s = - v0^2
                amax = np.clip(np.sqrt(v0**2 / (2 * dist)), 0, amax)
        # moving in same direction (catching up)
        else:
            # if self.log: print('catch-up')
            # If far enough that can decelerate from vmax to match velocity, then
            # keep accelerating to catch up
            if dist >= dist_vmax_v1:
                # if self.log: print('Accelerate. Distance > vmax->v1 > v0->v1')
                should_decelerate = False
            else:
                # check if actually at vmax, then decelerate
                if abs(v0)==vmax:
                    # if self.log: print('Decelerate. Distance < vmax->v1')
                    should_decelerate = True
                # else if v0 < vmax and v0 > v1, but close enough not to decelerate from vmax,
                # recalculate acceleration
                else:
                    # calculate peak velocity to accelerate to such that can acceleate
                    # and decelerate at amax and cover the distance
                    v_m = np.sqrt((2 * amax * dist + v0**2 + v1**2) / 2)
                    should_decelerate = abs(v0) > v_m
                    # if self.log:
                    #     v1 = v_m
                    #     dec = 'Decelerating' if should_decelerate else 'Accelerating'
                    #     print(dec, 'from %.2f to %.2f at acc value %.2f' % (v0, v1, amax))
    # if receding
    else:
        # if self.log: print('Receding. Accelerate')
        should_decelerate = False

    if should_decelerate:
        return -sign * amax
    return sign * amax



# @njit
def acc_to_target_3d(amax: np.ndarray, vmax: np.ndarray, v0_vec: np.ndarray, v1_vec: np.ndarray, disp: float, dt: float=1e-2):
    """Solve 3d kinematics problem"""
    acc_x = acc_1d(amax[0], vmax[0], v0_vec[0], v1_vec[0], disp[0], dt)
    acc_y = acc_1d(amax[1], vmax[1], v0_vec[1], v1_vec[1], disp[1], dt)
    acc_z = acc_1d(amax[2], vmax[2], v0_vec[2], v1_vec[2], disp[2], dt)
    acc_xyz = np.asarray([acc_x, acc_y, acc_z])
    acc = acc_xyz * amax / (np.linalg.norm(acc_xyz) + 1e-6)
    return acc
# @njit
def acc_to_target_2d(amax: float, vmax: float, v0_vec: np.ndarray, v1_vec: np.ndarray, disp: float, dt: float=1e-2):
    """Solve 2d kinematics problem"""
    acc_x = acc_1d(amax, vmax, v0_vec[0], v1_vec[0], disp[0], dt)
    acc_y = acc_1d(amax, vmax, v0_vec[1], v1_vec[1], disp[1], dt)
    acc_xyz = np.asarray([acc_x, acc_y])
    acc = acc_xyz * amax / (np.linalg.norm(acc_xyz) + 1e-6)
    return acc



# @njit
def compute_accel_2d(
    amax: float,
    vmax: float,
    from_position : np.ndarray,
    from_velocity : np.ndarray,
    to_position : np.ndarray,
    to_velocity : np.ndarray,
    dt: float=1e-2
    ) -> np.ndarray:
    target_disp = to_position - from_position
    acc = acc_to_target_2d(amax, vmax, from_velocity, to_velocity, target_disp, dt)
    # if self.log:
    #     print('Target disp', target_disp)
    #     print('Target acc', acc)
    return acc
 No newline at end of file