Skip to content

Commit 86f1ef9

Browse files
committed
Simplify parametrize structure
1 parent d721745 commit 86f1ef9

10 files changed

Lines changed: 42 additions & 228 deletions

File tree

drone_models/core.py

Lines changed: 19 additions & 59 deletions
Original file line numberDiff line numberDiff line change
@@ -2,11 +2,12 @@
22

33
from __future__ import annotations
44

5+
import inspect
56
import tomllib
67
import warnings
78
from functools import partial, wraps
89
from pathlib import Path
9-
from typing import TYPE_CHECKING, Any, Callable, ParamSpec, Protocol, TypeVar, runtime_checkable
10+
from typing import TYPE_CHECKING, Any, Callable, ParamSpec, TypeVar
1011

1112
import numpy as np
1213

@@ -49,16 +50,6 @@ def wrapper(
4950
return decorator
5051

5152

52-
model_parameter_registry: dict[str, type[ModelParams]] = {}
53-
54-
55-
def named_tuple2xp(params: ModelParams, xp: ModuleType, device: str | None = None) -> ModelParams:
56-
"""Convert a named tuple to an array API framework."""
57-
return params.__class__(
58-
**{k: xp.asarray(v, device=device) for k, v in params._asdict().items()}
59-
)
60-
61-
6253
def parametrize(
6354
fn: Callable[P, R], drone_model: str, xp: ModuleType | None = None, device: str | None = None
6455
) -> Callable[P, R]:
@@ -87,59 +78,28 @@ def parametrize(
8778
Returns:
8879
The parametrized controller function with all keyword argument only parameters filled in.
8980
"""
90-
model_id = fn.__module__ + "." + fn.__name__
9181
try:
92-
params = model_parameter_registry[model_id].load(drone_model)
93-
if xp is not None: # Convert to any array API framework
94-
params = named_tuple2xp(params, xp=xp, device=device)
82+
xp = np if xp is None else xp
83+
# physics = Path(sys.modules[fn.__module__].__file__).parent.name
84+
physics = fn.__module__.split(".")[-2]
85+
sig = inspect.signature(fn)
86+
kwonly_params = [
87+
name
88+
for name, param in sig.parameters.items()
89+
if param.kind == inspect.Parameter.KEYWORD_ONLY
90+
]
91+
params = load_params(physics, drone_model)
92+
93+
params = {k: xp.asarray(v, device=device) for k, v in params.items() if k in kwonly_params}
94+
# if xp is not None: # Convert to any array API framework
95+
# params = named_tuple2xp(params, xp=xp, device=device)
9596
except KeyError as e:
9697
raise KeyError(
97-
f"Model `{model_id}` does not exist in the parameter registry for drone `{drone_model}`"
98+
f"Model `{physics}` does not exist in the parameter registry for drone `{drone_model}`"
9899
) from e
99100
except ValueError as e:
100-
raise ValueError(f"Drone model `{drone_model}` not supported for `{model_id}`") from e
101-
return partial(fn, **params._asdict())
102-
103-
104-
@runtime_checkable
105-
class ModelParams(Protocol):
106-
"""Protocol for model parameters."""
107-
108-
@staticmethod
109-
def load(drone_model: str) -> ModelParams:
110-
"""Load the parameters from the config file."""
111-
112-
def _asdict(self) -> dict[str, Any]:
113-
"""Convert the parameters to a dictionary."""
114-
115-
116-
def register_model_parameters(
117-
params: ModelParams | type[ModelParams],
118-
) -> Callable[[Callable[P, R]], Callable[P, R]]:
119-
"""Register the default model parameters for this model.
120-
121-
Warning:
122-
The model parameters **must** be a named tuple with a function `load` that takes in the
123-
drone model name and returns an instance of itself, or a class that implements the
124-
ModelParams protocol.
125-
126-
Args:
127-
params: The model parameter type.
128-
129-
Returns:
130-
A decorator function that registers the parameters and returns the function unchanged.
131-
"""
132-
if not isinstance(params, ModelParams):
133-
raise ValueError(f"{params} does not implement the ModelParams protocol")
134-
135-
def decorator(fn: Callable[P, R]) -> Callable[P, R]:
136-
controller_id = fn.__module__ + "." + fn.__name__
137-
if controller_id in model_parameter_registry:
138-
raise ValueError(f"Model `{controller_id}` already registered")
139-
model_parameter_registry[controller_id] = params
140-
return fn
141-
142-
return decorator
101+
raise ValueError(f"Drone model `{drone_model}` not supported for `{physics}`") from e
102+
return partial(fn, **params)
143103

144104

145105
def load_params(physics: str, drone_model: str, xp: ModuleType | None = None) -> dict:

drone_models/first_principles/model.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -9,15 +9,13 @@
99
from scipy.spatial.transform import Rotation as R
1010

1111
import drone_models.symbols as symbols
12-
from drone_models.core import register_model_parameters, supports
13-
from drone_models.first_principles.params import FirstPrinciplesParams
12+
from drone_models.core import supports
1413
from drone_models.utils import rotation, to_xp
1514

1615
if TYPE_CHECKING:
1716
from array_api_typing import Array
1817

1918

20-
@register_model_parameters(FirstPrinciplesParams)
2119
@supports(rotor_dynamics=True)
2220
def dynamics(
2321
pos: Array,
@@ -110,7 +108,6 @@ def dynamics(
110108
return pos_dot, quat_dot, vel_dot, ang_vel_dot, rotor_vel_dot
111109

112110

113-
@register_model_parameters(FirstPrinciplesParams)
114111
def symbolic_dynamics(
115112
model_rotor_vel: bool = True,
116113
model_dist_f: bool = False,

drone_models/first_principles/params.py

Lines changed: 0 additions & 37 deletions
This file was deleted.

drone_models/so_rpy/model.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -9,15 +9,13 @@
99
from scipy.spatial.transform import Rotation as R
1010

1111
import drone_models.symbols as symbols
12-
from drone_models.core import register_model_parameters, supports
13-
from drone_models.so_rpy.params import SoRpyParams
12+
from drone_models.core import supports
1413
from drone_models.utils import rotation, to_xp
1514

1615
if TYPE_CHECKING:
1716
from array_api_typing import Array
1817

1918

20-
@register_model_parameters(SoRpyParams)
2119
@supports(rotor_dynamics=False)
2220
def dynamics(
2321
pos: Array,
@@ -108,7 +106,6 @@ def dynamics(
108106
return pos_dot, quat_dot, vel_dot, ang_vel_dot, rotor_vel_dot
109107

110108

111-
@register_model_parameters(SoRpyParams)
112109
def symbolic_dynamics(
113110
model_rotor_vel: bool = False,
114111
model_dist_f: bool = False,

drone_models/so_rpy/params.py

Lines changed: 0 additions & 34 deletions
This file was deleted.

drone_models/so_rpy_rotor/model.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,16 +10,14 @@
1010
from scipy.spatial.transform import Rotation as R
1111

1212
import drone_models.symbols as symbols
13-
from drone_models.core import register_model_parameters, supports
14-
from drone_models.so_rpy_rotor.params import SoRpyRotorParams
13+
from drone_models.core import supports
1514
from drone_models.transform import motor_force2rotor_vel
1615
from drone_models.utils import rotation, to_xp
1716

1817
if TYPE_CHECKING:
1918
from array_api_typing import Array
2019

2120

22-
@register_model_parameters(SoRpyRotorParams)
2321
@supports(rotor_dynamics=True)
2422
def dynamics(
2523
pos: Array,
@@ -128,7 +126,6 @@ def dynamics(
128126
return pos_dot, quat_dot, vel_dot, ang_vel_dot, rotor_vel_dot
129127

130128

131-
@register_model_parameters(SoRpyRotorParams)
132129
def symbolic_dynamics(
133130
model_rotor_vel: bool = False,
134131
model_dist_f: bool = False,

drone_models/so_rpy_rotor/params.py

Lines changed: 0 additions & 40 deletions
This file was deleted.

drone_models/so_rpy_rotor_drag/model.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,15 +10,13 @@
1010
from scipy.spatial.transform import Rotation as R
1111

1212
import drone_models.symbols as symbols
13-
from drone_models.core import register_model_parameters, supports
14-
from drone_models.so_rpy_rotor_drag.params import SoRpyRotorDragParams
13+
from drone_models.core import supports
1514
from drone_models.utils import rotation, to_xp
1615

1716
if TYPE_CHECKING:
1817
from array_api_typing import Array
1918

2019

21-
@register_model_parameters(SoRpyRotorDragParams)
2220
@supports(rotor_dynamics=True)
2321
def dynamics(
2422
pos: Array,
@@ -142,7 +140,6 @@ def dynamics(
142140
return pos_dot, quat_dot, vel_dot, ang_vel_dot, rotor_vel_dot
143141

144142

145-
@register_model_parameters(SoRpyRotorDragParams)
146143
def symbolic_dynamics(
147144
model_rotor_vel: bool = True,
148145
model_dist_f: bool = False,

drone_models/so_rpy_rotor_drag/params.py

Lines changed: 0 additions & 42 deletions
This file was deleted.

tests/unit/test_parametrization.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,19 @@
1+
"""Tests of the parametrization of the models."""
2+
3+
from __future__ import annotations
4+
5+
from typing import Callable
6+
7+
import pytest
8+
9+
from drone_models import available_models
10+
from drone_models.core import parametrize
11+
from drone_models.drones import available_drones
12+
13+
14+
@pytest.mark.unit
15+
@pytest.mark.parametrize("model_name, model", available_models.items())
16+
@pytest.mark.parametrize("drone_type", available_drones)
17+
def test_model_parametrization(model_name: str, model: Callable, drone_type: str):
18+
"""TODO."""
19+
parametrize(model, drone_type)

0 commit comments

Comments
 (0)