virtualmodelcontrol 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- virtualmodelcontrol/__init__.py +75 -0
- virtualmodelcontrol/_version.py +24 -0
- virtualmodelcontrol/compiler.py +158 -0
- virtualmodelcontrol/control/__init__.py +5 -0
- virtualmodelcontrol/control/controller.py +76 -0
- virtualmodelcontrol/core/__init__.py +24 -0
- virtualmodelcontrol/core/params.py +244 -0
- virtualmodelcontrol/core/registry.py +54 -0
- virtualmodelcontrol/core/signals.py +46 -0
- virtualmodelcontrol/core/space.py +162 -0
- virtualmodelcontrol/core/symbolic.py +67 -0
- virtualmodelcontrol/core/units.py +34 -0
- virtualmodelcontrol/dynamics.py +149 -0
- virtualmodelcontrol/mechanisms/__init__.py +67 -0
- virtualmodelcontrol/mechanisms/components/__init__.py +33 -0
- virtualmodelcontrol/mechanisms/components/base.py +71 -0
- virtualmodelcontrol/mechanisms/components/dissipation.py +50 -0
- virtualmodelcontrol/mechanisms/components/inertance.py +56 -0
- virtualmodelcontrol/mechanisms/components/sources.py +63 -0
- virtualmodelcontrol/mechanisms/components/storage.py +272 -0
- virtualmodelcontrol/mechanisms/coordinates/__init__.py +24 -0
- virtualmodelcontrol/mechanisms/coordinates/base.py +117 -0
- virtualmodelcontrol/mechanisms/coordinates/frames.py +62 -0
- virtualmodelcontrol/mechanisms/coordinates/joints.py +51 -0
- virtualmodelcontrol/mechanisms/coordinates/ops.py +146 -0
- virtualmodelcontrol/mechanisms/coordinates/references.py +43 -0
- virtualmodelcontrol/mechanisms/mechanism.py +88 -0
- virtualmodelcontrol/models/__init__.py +22 -0
- virtualmodelcontrol/models/actuation.py +196 -0
- virtualmodelcontrol/models/assembly.py +154 -0
- virtualmodelcontrol/models/continuum/__init__.py +5 -0
- virtualmodelcontrol/models/continuum/pcc.py +120 -0
- virtualmodelcontrol/models/kinematic.py +41 -0
- virtualmodelcontrol/models/rigid/__init__.py +6 -0
- virtualmodelcontrol/models/rigid/couplings.py +64 -0
- virtualmodelcontrol/models/rigid/poe.py +95 -0
- virtualmodelcontrol/py.typed +0 -0
- virtualmodelcontrol/robots/__init__.py +5 -0
- virtualmodelcontrol/robots/adapt.py +98 -0
- virtualmodelcontrol/robots/helyx.py +79 -0
- virtualmodelcontrol/sim/__init__.py +7 -0
- virtualmodelcontrol/sim/model_plant.py +79 -0
- virtualmodelcontrol/sim/plant.py +40 -0
- virtualmodelcontrol/sim/run.py +73 -0
- virtualmodelcontrol/system.py +55 -0
- virtualmodelcontrol-0.1.0.dist-info/METADATA +82 -0
- virtualmodelcontrol-0.1.0.dist-info/RECORD +49 -0
- virtualmodelcontrol-0.1.0.dist-info/WHEEL +4 -0
- virtualmodelcontrol-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
"""Virtual Model Control: robots and controllers as mechanisms of springs, dampers, inertances."""
|
|
2
|
+
|
|
3
|
+
from . import sim
|
|
4
|
+
from .compiler import Compiled, compile
|
|
5
|
+
from .control import VMCController
|
|
6
|
+
from .core import SO2, Euclidean, Param, ParamSet, Product, Signals, register
|
|
7
|
+
from .dynamics import Dynamics, compile_dynamics
|
|
8
|
+
from .mechanisms import (
|
|
9
|
+
Custom,
|
|
10
|
+
ForceSource,
|
|
11
|
+
FramePoint,
|
|
12
|
+
GaussianSpring,
|
|
13
|
+
Gravity,
|
|
14
|
+
GravityCompensation,
|
|
15
|
+
Inertance,
|
|
16
|
+
Joint,
|
|
17
|
+
LimitSpring,
|
|
18
|
+
LinearDamper,
|
|
19
|
+
LinearSpring,
|
|
20
|
+
Mechanism,
|
|
21
|
+
Norm,
|
|
22
|
+
PointMass,
|
|
23
|
+
PolynomialSpring,
|
|
24
|
+
Projection,
|
|
25
|
+
Ref,
|
|
26
|
+
SigmoidSpring,
|
|
27
|
+
Stack,
|
|
28
|
+
TanhDamper,
|
|
29
|
+
TanhSpring,
|
|
30
|
+
)
|
|
31
|
+
from .system import VirtualMechanismSystem
|
|
32
|
+
|
|
33
|
+
try:
|
|
34
|
+
from ._version import __version__
|
|
35
|
+
except ImportError: # a source tree that was never installed
|
|
36
|
+
__version__ = "0.0.0+unknown"
|
|
37
|
+
|
|
38
|
+
__all__ = [
|
|
39
|
+
"SO2",
|
|
40
|
+
"Compiled",
|
|
41
|
+
"Custom",
|
|
42
|
+
"Dynamics",
|
|
43
|
+
"Euclidean",
|
|
44
|
+
"ForceSource",
|
|
45
|
+
"FramePoint",
|
|
46
|
+
"GaussianSpring",
|
|
47
|
+
"Gravity",
|
|
48
|
+
"GravityCompensation",
|
|
49
|
+
"Inertance",
|
|
50
|
+
"Joint",
|
|
51
|
+
"LimitSpring",
|
|
52
|
+
"LinearDamper",
|
|
53
|
+
"LinearSpring",
|
|
54
|
+
"Mechanism",
|
|
55
|
+
"Norm",
|
|
56
|
+
"Param",
|
|
57
|
+
"ParamSet",
|
|
58
|
+
"PointMass",
|
|
59
|
+
"PolynomialSpring",
|
|
60
|
+
"Product",
|
|
61
|
+
"Projection",
|
|
62
|
+
"Ref",
|
|
63
|
+
"SigmoidSpring",
|
|
64
|
+
"Signals",
|
|
65
|
+
"Stack",
|
|
66
|
+
"TanhDamper",
|
|
67
|
+
"TanhSpring",
|
|
68
|
+
"VMCController",
|
|
69
|
+
"VirtualMechanismSystem",
|
|
70
|
+
"__version__",
|
|
71
|
+
"compile",
|
|
72
|
+
"compile_dynamics",
|
|
73
|
+
"register",
|
|
74
|
+
"sim",
|
|
75
|
+
]
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
# file generated by vcs-versioning
|
|
2
|
+
# don't change, don't track in version control
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
__all__ = [
|
|
6
|
+
"__version__",
|
|
7
|
+
"__version_tuple__",
|
|
8
|
+
"version",
|
|
9
|
+
"version_tuple",
|
|
10
|
+
"__commit_id__",
|
|
11
|
+
"commit_id",
|
|
12
|
+
]
|
|
13
|
+
|
|
14
|
+
version: str
|
|
15
|
+
__version__: str
|
|
16
|
+
__version_tuple__: tuple[int | str, ...]
|
|
17
|
+
version_tuple: tuple[int | str, ...]
|
|
18
|
+
commit_id: str | None
|
|
19
|
+
__commit_id__: str | None
|
|
20
|
+
|
|
21
|
+
__version__ = version = '0.1.0'
|
|
22
|
+
__version_tuple__ = version_tuple = (0, 1, 0)
|
|
23
|
+
|
|
24
|
+
__commit_id__ = commit_id = None
|
|
@@ -0,0 +1,158 @@
|
|
|
1
|
+
"""Compile a virtual mechanism system into CasADi functions: the control law and its energetics."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Iterable
|
|
6
|
+
from dataclasses import dataclass
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
import casadi as ca
|
|
10
|
+
import numpy as np
|
|
11
|
+
|
|
12
|
+
from .core.params import Binding, ParamSet
|
|
13
|
+
from .mechanisms.coordinates.base import Context
|
|
14
|
+
from .system import VirtualMechanismSystem
|
|
15
|
+
|
|
16
|
+
ARGS = ["q", "v", "z", "p", "t"]
|
|
17
|
+
OPTS = {"cse": True} # merges the repeated derivatives of shared coordinates
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass
|
|
21
|
+
class Compiled:
|
|
22
|
+
"""CasADi functions of a compiled system.
|
|
23
|
+
|
|
24
|
+
``law`` (→ u, ż), ``tau``, ``energy`` (→ V, T), ``power`` (→ port, dissipation, source) and
|
|
25
|
+
``forces`` take (q, v, z, p, t): configuration, velocity, virtual state z = [positions,
|
|
26
|
+
velocities], live Params p and time t [s]. ``fast`` maps one vector [θ, θ̇, z, p, t] of motor
|
|
27
|
+
angles and rates to [u, ż]; ``fast_energy`` maps it to V + T.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
system: VirtualMechanismSystem
|
|
31
|
+
params: ParamSet
|
|
32
|
+
live: list[str]
|
|
33
|
+
law: ca.Function
|
|
34
|
+
tau: ca.Function
|
|
35
|
+
energy: ca.Function
|
|
36
|
+
power: ca.Function
|
|
37
|
+
forces: ca.Function
|
|
38
|
+
fast: ca.Function
|
|
39
|
+
fast_energy: ca.Function
|
|
40
|
+
component_names: list[str]
|
|
41
|
+
z0: np.ndarray
|
|
42
|
+
n_motors: tuple[int, int]
|
|
43
|
+
n_u: int
|
|
44
|
+
|
|
45
|
+
def live_values(self) -> np.ndarray:
|
|
46
|
+
"""Current values of the live Params, packed like p."""
|
|
47
|
+
return self.params.vector(self.live)
|
|
48
|
+
|
|
49
|
+
def live_slices(self) -> dict[str, slice]:
|
|
50
|
+
"""Where each live Param sits in p."""
|
|
51
|
+
out, offset = {}, 0
|
|
52
|
+
for name in self.live:
|
|
53
|
+
size = self.params[name].size
|
|
54
|
+
out[name] = slice(offset, offset + size)
|
|
55
|
+
offset += size
|
|
56
|
+
return out
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
def compile(system: VirtualMechanismSystem, runtime: Iterable[str] = ()) -> Compiled:
|
|
60
|
+
"""Compile ``system`` into CasADi functions.
|
|
61
|
+
|
|
62
|
+
``stage`` Params, and those matching the glob patterns in ``runtime``, stay live inputs; every
|
|
63
|
+
other Param is folded in at its current value (compile again after changing one).
|
|
64
|
+
"""
|
|
65
|
+
params = system.params
|
|
66
|
+
live = params.select(patterns=runtime, scopes=["stage"])
|
|
67
|
+
binding = Binding(params, live)
|
|
68
|
+
space, actuation = system.robot.model.space, system.actuation
|
|
69
|
+
states = system.states
|
|
70
|
+
nz = sum(s.dim for s in states)
|
|
71
|
+
q, v, t = ca.SX.sym("q", space.nq), ca.SX.sym("v", space.nv), ca.SX.sym("t")
|
|
72
|
+
z = ca.SX.sym("z", 2 * nz)
|
|
73
|
+
zpos, zvel = z[:nz], z[nz:]
|
|
74
|
+
slices, offset = {}, 0
|
|
75
|
+
for state in states:
|
|
76
|
+
slices[id(state)] = slice(offset, offset + state.dim)
|
|
77
|
+
offset += state.dim
|
|
78
|
+
ctx = Context(q, binding, z=zpos, t=t, states=slices)
|
|
79
|
+
G = space.velocity_map(q)
|
|
80
|
+
|
|
81
|
+
tau = ca.SX.zeros(space.nv, 1)
|
|
82
|
+
fz = ca.SX.zeros(nz, 1)
|
|
83
|
+
Mz = ca.SX.zeros(nz, nz)
|
|
84
|
+
V, P_diss, P_src = ca.SX(0), ca.SX(0), ca.SX(0)
|
|
85
|
+
per_component: list[Any] = []
|
|
86
|
+
names: list[str] = []
|
|
87
|
+
for name, comp in system.components:
|
|
88
|
+
y = ctx.value(comp.coord)
|
|
89
|
+
Jq = ca.mtimes(ca.jacobian(y, q), G)
|
|
90
|
+
Jz = ca.jacobian(y, zpos)
|
|
91
|
+
if comp.kind == "inertance":
|
|
92
|
+
if ca.depends_on(y, ca.vertcat(q, t)):
|
|
93
|
+
raise ValueError(f"{name}: a controller's inertances may act on its states only")
|
|
94
|
+
Mz += ca.mtimes([Jz.T, comp.inertance(ctx, y), Jz])
|
|
95
|
+
continue
|
|
96
|
+
yd = ca.mtimes(Jq, v) + ca.mtimes(Jz, zvel) + ca.jacobian(y, t)
|
|
97
|
+
f = comp.force(ctx, y, yd)
|
|
98
|
+
tau_k = ca.mtimes(Jq.T, f)
|
|
99
|
+
tau += tau_k
|
|
100
|
+
fz += ca.mtimes(Jz.T, f)
|
|
101
|
+
if comp.kind == "storage":
|
|
102
|
+
V += comp.energy(ctx, y)
|
|
103
|
+
elif comp.kind == "dissipation":
|
|
104
|
+
P_diss += ca.dot(f, yd)
|
|
105
|
+
else:
|
|
106
|
+
P_src += ca.dot(f, yd)
|
|
107
|
+
per_component += [y, yd, f, tau_k]
|
|
108
|
+
names.append(name)
|
|
109
|
+
|
|
110
|
+
z0 = np.concatenate([s.initial for s in states] + [np.zeros(nz)])
|
|
111
|
+
T = 0.5 * ca.dot(zvel, ca.mtimes(Mz, zvel))
|
|
112
|
+
if nz:
|
|
113
|
+
M0 = np.array(ca.Function("M", [z, binding.p], [Mz])(z0, binding.values()))
|
|
114
|
+
if np.linalg.matrix_rank(M0) < nz:
|
|
115
|
+
raise ValueError("every virtual state needs an inertance (the state mass is singular)")
|
|
116
|
+
coriolis = ca.jtimes(ca.mtimes(Mz, zvel), zpos, zvel) - ca.gradient(T, zpos)
|
|
117
|
+
zdot = ca.vertcat(zvel, ca.solve(Mz, fz - coriolis))
|
|
118
|
+
else:
|
|
119
|
+
zdot = ca.SX.zeros(0, 1)
|
|
120
|
+
u = actuation.allocate(tau, q, binding.view(actuation.params))
|
|
121
|
+
|
|
122
|
+
args = [q, v, z, binding.p, t]
|
|
123
|
+
law = ca.Function("law", args, [u, zdot], ARGS, ["u", "zdot"], OPTS)
|
|
124
|
+
energy = ca.Function("energy", args, [V, T], ARGS, ["V", "T"], OPTS)
|
|
125
|
+
|
|
126
|
+
# Fast path: motor angles and rates in, through the exact inverse of the transmission.
|
|
127
|
+
n_angles, n_rates = actuation.motor_sizes(space)
|
|
128
|
+
theta, theta_dot = ca.SX.sym("theta", n_angles), ca.SX.sym("theta_dot", n_rates)
|
|
129
|
+
pa = binding.view(actuation.params)
|
|
130
|
+
qa = actuation.config_from_motors(theta, pa)
|
|
131
|
+
va = actuation.velocity_from_motors(qa, theta_dot, pa)
|
|
132
|
+
x = ca.vertcat(theta, theta_dot, z, binding.p, t)
|
|
133
|
+
ua, zdot_a = law(qa, va, z, binding.p, t)
|
|
134
|
+
Va, Ta = energy(qa, va, z, binding.p, t)
|
|
135
|
+
|
|
136
|
+
return Compiled(
|
|
137
|
+
system=system,
|
|
138
|
+
params=params,
|
|
139
|
+
live=binding.live,
|
|
140
|
+
law=law,
|
|
141
|
+
tau=ca.Function("tau", args, [tau], ARGS, ["tau"], OPTS),
|
|
142
|
+
energy=energy,
|
|
143
|
+
power=ca.Function(
|
|
144
|
+
"power",
|
|
145
|
+
args,
|
|
146
|
+
[ca.dot(tau, v), P_diss, P_src],
|
|
147
|
+
ARGS,
|
|
148
|
+
["port", "dissipation", "source"],
|
|
149
|
+
OPTS,
|
|
150
|
+
),
|
|
151
|
+
forces=ca.Function("forces", args, per_component, OPTS),
|
|
152
|
+
fast=ca.Function("fast", [x], [ca.vertcat(ua, zdot_a)], ["x"], ["out"], OPTS),
|
|
153
|
+
fast_energy=ca.Function("fast_energy", [x], [Va + Ta], ["x"], ["E"], OPTS),
|
|
154
|
+
component_names=names,
|
|
155
|
+
z0=z0,
|
|
156
|
+
n_motors=(n_angles, n_rates),
|
|
157
|
+
n_u=int(u.numel()),
|
|
158
|
+
)
|
|
@@ -0,0 +1,76 @@
|
|
|
1
|
+
"""VMCController: runs a compiled virtual mechanism step by step."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Mapping
|
|
6
|
+
from typing import Any
|
|
7
|
+
|
|
8
|
+
import numpy as np
|
|
9
|
+
from numpy.typing import ArrayLike
|
|
10
|
+
|
|
11
|
+
from ..core.signals import Signals
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
class VMCController:
|
|
15
|
+
"""Runs a compiled system: motor angles and rates in, motor torques out.
|
|
16
|
+
|
|
17
|
+
``step`` reads ``motor_position`` [rad] and ``motor_velocity`` [rad/s] and returns
|
|
18
|
+
``motor_torque`` [N·m]. Virtual states are integrated by semi-implicit Euler over the measured
|
|
19
|
+
dt. Live Params start at their values when compiled; ``set`` changes them here only.
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
def __init__(self, compiled: Any) -> None:
|
|
23
|
+
self.compiled = compiled
|
|
24
|
+
self.params = compiled.live_values()
|
|
25
|
+
self._slices = compiled.live_slices()
|
|
26
|
+
self.z = compiled.z0.copy()
|
|
27
|
+
self.t: float | None = None
|
|
28
|
+
self._x: np.ndarray | None = None
|
|
29
|
+
|
|
30
|
+
def reset(self, t: float, meas: Signals | None = None, z0: ArrayLike | None = None) -> None:
|
|
31
|
+
"""Restart at time ``t`` [s] from the initial virtual state (or ``z0``)."""
|
|
32
|
+
self.z = self.compiled.z0.copy() if z0 is None else np.array(z0, dtype=float)
|
|
33
|
+
self._x = None if meas is None else self._pack(meas, t)
|
|
34
|
+
self.t = t
|
|
35
|
+
|
|
36
|
+
def step(self, t: float, meas: Signals) -> Signals:
|
|
37
|
+
"""One control step at time ``t`` [s]; returns the motor torques."""
|
|
38
|
+
dt = 0.0 if self.t is None else t - self.t
|
|
39
|
+
x = self._pack(meas, t)
|
|
40
|
+
out = np.asarray(self.compiled.fast(x)).ravel()
|
|
41
|
+
nu, nz = self.compiled.n_u, self.z.size // 2
|
|
42
|
+
if nz:
|
|
43
|
+
self.z[nz:] += dt * out[nu + nz :]
|
|
44
|
+
self.z[:nz] += dt * self.z[nz:]
|
|
45
|
+
self.t, self._x = t, x
|
|
46
|
+
return Signals(t, motor_torque=out[:nu])
|
|
47
|
+
|
|
48
|
+
def set(self, values: Mapping[str, ArrayLike] | None = None, **kwargs: ArrayLike) -> float:
|
|
49
|
+
"""Change live Params at once; returns the exact jump of the controller's energy [J]."""
|
|
50
|
+
new = self.params.copy()
|
|
51
|
+
for name, value in {**(values or {}), **kwargs}.items():
|
|
52
|
+
if name not in self._slices:
|
|
53
|
+
raise KeyError(
|
|
54
|
+
f"{name!r} is not a live Param here; compile with runtime=[{name!r}] to change "
|
|
55
|
+
"it while running"
|
|
56
|
+
)
|
|
57
|
+
new[self._slices[name]] = np.ravel(value, order="F")
|
|
58
|
+
jump = 0.0 if self._x is None else self._energy(new) - self._energy(self.params)
|
|
59
|
+
self.params = new
|
|
60
|
+
return jump
|
|
61
|
+
|
|
62
|
+
def energy(self) -> float:
|
|
63
|
+
"""Energy of the controller at the last step [J] (stored plus virtual kinetic)."""
|
|
64
|
+
return 0.0 if self._x is None else self._energy(self.params)
|
|
65
|
+
|
|
66
|
+
def _energy(self, params: np.ndarray) -> float:
|
|
67
|
+
assert self._x is not None
|
|
68
|
+
x = self._x.copy()
|
|
69
|
+
n = self._x.size - 1 - params.size
|
|
70
|
+
x[n:-1] = params
|
|
71
|
+
return float(self.compiled.fast_energy(x))
|
|
72
|
+
|
|
73
|
+
def _pack(self, meas: Signals, t: float) -> np.ndarray:
|
|
74
|
+
return np.concatenate(
|
|
75
|
+
[meas["motor_position"], meas["motor_velocity"], self.z, self.params, [t]]
|
|
76
|
+
)
|
|
@@ -0,0 +1,24 @@
|
|
|
1
|
+
"""Core: parameters, spaces, symbolic helpers, signals, registry and units."""
|
|
2
|
+
|
|
3
|
+
from .params import SCOPES, Binding, Param, ParamSet, as_param, constants
|
|
4
|
+
from .registry import get, load_plugins, names, register
|
|
5
|
+
from .signals import Signals
|
|
6
|
+
from .space import SO2, Euclidean, Product, Space
|
|
7
|
+
|
|
8
|
+
__all__ = [
|
|
9
|
+
"SCOPES",
|
|
10
|
+
"SO2",
|
|
11
|
+
"Binding",
|
|
12
|
+
"Euclidean",
|
|
13
|
+
"Param",
|
|
14
|
+
"ParamSet",
|
|
15
|
+
"Product",
|
|
16
|
+
"Signals",
|
|
17
|
+
"Space",
|
|
18
|
+
"as_param",
|
|
19
|
+
"constants",
|
|
20
|
+
"get",
|
|
21
|
+
"load_plugins",
|
|
22
|
+
"names",
|
|
23
|
+
"register",
|
|
24
|
+
]
|
|
@@ -0,0 +1,244 @@
|
|
|
1
|
+
"""Parameters: named numbers with units, bounds and scopes, packed into one flat vector."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import fnmatch
|
|
6
|
+
from collections.abc import Iterable, Iterator, Mapping
|
|
7
|
+
from typing import Any, Literal
|
|
8
|
+
|
|
9
|
+
import casadi as ca
|
|
10
|
+
import numpy as np
|
|
11
|
+
from numpy.typing import ArrayLike
|
|
12
|
+
|
|
13
|
+
Scope = Literal["fixed", "design", "episode", "stage"]
|
|
14
|
+
SCOPES: tuple[str, ...] = ("fixed", "design", "episode", "stage")
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class Param:
|
|
18
|
+
"""A named number (scalar, vector or matrix) with a unit, bounds and a scope.
|
|
19
|
+
|
|
20
|
+
The scope is the time scale on which the value may change: ``fixed`` (never), ``design``
|
|
21
|
+
(with the hardware), ``episode`` (between runs) or ``stage`` (at every step).
|
|
22
|
+
"""
|
|
23
|
+
|
|
24
|
+
__slots__ = ("_value", "bounds", "name", "scale", "scope", "unit")
|
|
25
|
+
|
|
26
|
+
def __init__(
|
|
27
|
+
self,
|
|
28
|
+
name: str,
|
|
29
|
+
value: ArrayLike,
|
|
30
|
+
*,
|
|
31
|
+
unit: str = "",
|
|
32
|
+
bounds: tuple[Any, Any] = (-np.inf, np.inf),
|
|
33
|
+
scale: float = 1.0,
|
|
34
|
+
scope: Scope = "fixed",
|
|
35
|
+
) -> None:
|
|
36
|
+
if scope not in SCOPES:
|
|
37
|
+
raise ValueError(f"scope must be one of {SCOPES}, got {scope!r}")
|
|
38
|
+
value = np.array(value, dtype=float)
|
|
39
|
+
if value.ndim > 2:
|
|
40
|
+
raise ValueError(f"Param {name!r} must be a scalar, vector or matrix")
|
|
41
|
+
self.name = name
|
|
42
|
+
self._value = value
|
|
43
|
+
self.unit = unit
|
|
44
|
+
self.bounds = (bounds[0], bounds[1])
|
|
45
|
+
self.scale = float(scale)
|
|
46
|
+
self.scope = scope
|
|
47
|
+
|
|
48
|
+
@property
|
|
49
|
+
def value(self) -> np.ndarray:
|
|
50
|
+
"""Current value, as a float array of fixed shape."""
|
|
51
|
+
return self._value
|
|
52
|
+
|
|
53
|
+
@value.setter
|
|
54
|
+
def value(self, value: ArrayLike) -> None:
|
|
55
|
+
new = np.array(value, dtype=float)
|
|
56
|
+
if new.shape != self._value.shape:
|
|
57
|
+
if new.size != self._value.size:
|
|
58
|
+
raise ValueError(
|
|
59
|
+
f"Param {self.name!r} has shape {self._value.shape}, got a value of shape "
|
|
60
|
+
f"{new.shape}"
|
|
61
|
+
)
|
|
62
|
+
new = new.reshape(self._value.shape)
|
|
63
|
+
self._value = new
|
|
64
|
+
|
|
65
|
+
@property
|
|
66
|
+
def shape(self) -> tuple[int, ...]:
|
|
67
|
+
"""Shape of the value: () for a scalar, (n,) for a vector, (r, c) for a matrix."""
|
|
68
|
+
return self._value.shape
|
|
69
|
+
|
|
70
|
+
@property
|
|
71
|
+
def size(self) -> int:
|
|
72
|
+
"""Number of entries in the value."""
|
|
73
|
+
return self._value.size
|
|
74
|
+
|
|
75
|
+
def __repr__(self) -> str:
|
|
76
|
+
unit = f", unit={self.unit!r}" if self.unit else ""
|
|
77
|
+
return f"Param({self.name!r}, {self._value.tolist()!r}{unit}, scope={self.scope!r})"
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def as_param(
|
|
81
|
+
value: Any,
|
|
82
|
+
name: str,
|
|
83
|
+
*,
|
|
84
|
+
unit: str = "",
|
|
85
|
+
bounds: tuple[Any, Any] = (-np.inf, np.inf),
|
|
86
|
+
scope: Scope = "fixed",
|
|
87
|
+
) -> Param:
|
|
88
|
+
"""Return ``value`` if it already is a Param, else wrap it in a new Param."""
|
|
89
|
+
if isinstance(value, Param):
|
|
90
|
+
return value
|
|
91
|
+
return Param(name, value, unit=unit, bounds=bounds, scope=scope)
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
class ParamSet(Mapping[str, Param]):
|
|
95
|
+
"""Ordered, named collection of Params; packs their values into one flat vector.
|
|
96
|
+
|
|
97
|
+
A Param object appears once: adding it again under another name keeps the first name.
|
|
98
|
+
Matrices are packed column by column, as CasADi stores them.
|
|
99
|
+
"""
|
|
100
|
+
|
|
101
|
+
def __init__(self, params: Iterable[Param] = ()) -> None:
|
|
102
|
+
self._params: dict[str, Param] = {}
|
|
103
|
+
self._names: dict[int, str] = {}
|
|
104
|
+
for param in params:
|
|
105
|
+
self.add(param)
|
|
106
|
+
|
|
107
|
+
def add(self, param: Param, name: str | None = None, *, rename: bool = False) -> str:
|
|
108
|
+
"""Add ``param`` under ``name`` (default: its own name) and return the name used.
|
|
109
|
+
|
|
110
|
+
With ``rename``, a name already taken gets a numeric suffix instead of raising.
|
|
111
|
+
"""
|
|
112
|
+
if id(param) in self._names:
|
|
113
|
+
return self._names[id(param)]
|
|
114
|
+
name = param.name if name is None else name
|
|
115
|
+
if name in self._params:
|
|
116
|
+
if not rename:
|
|
117
|
+
raise ValueError(f"duplicate parameter name {name!r}")
|
|
118
|
+
k = 2
|
|
119
|
+
while f"{name}{k}" in self._params:
|
|
120
|
+
k += 1
|
|
121
|
+
name = f"{name}{k}"
|
|
122
|
+
self._params[name] = param
|
|
123
|
+
self._names[id(param)] = name
|
|
124
|
+
return name
|
|
125
|
+
|
|
126
|
+
def merge(self, other: ParamSet, prefix: str = "") -> None:
|
|
127
|
+
"""Add every Param of ``other``, with names prefixed by ``prefix.``."""
|
|
128
|
+
for name, param in other.items():
|
|
129
|
+
self.add(param, f"{prefix}.{name}" if prefix else name)
|
|
130
|
+
|
|
131
|
+
def __getitem__(self, name: str) -> Param:
|
|
132
|
+
try:
|
|
133
|
+
return self._params[name]
|
|
134
|
+
except KeyError:
|
|
135
|
+
raise KeyError(f"no parameter {name!r}; known: {list(self._params)}") from None
|
|
136
|
+
|
|
137
|
+
def __iter__(self) -> Iterator[str]:
|
|
138
|
+
return iter(self._params)
|
|
139
|
+
|
|
140
|
+
def __len__(self) -> int:
|
|
141
|
+
return len(self._params)
|
|
142
|
+
|
|
143
|
+
def __repr__(self) -> str:
|
|
144
|
+
return f"ParamSet({list(self._params)})"
|
|
145
|
+
|
|
146
|
+
def name_of(self, param: Param) -> str:
|
|
147
|
+
"""Name under which ``param`` is stored."""
|
|
148
|
+
return self._names[id(param)]
|
|
149
|
+
|
|
150
|
+
def has(self, param: Param) -> bool:
|
|
151
|
+
"""True if this exact Param object is in the set."""
|
|
152
|
+
return id(param) in self._names
|
|
153
|
+
|
|
154
|
+
def select(self, patterns: Iterable[str] = (), scopes: Iterable[str] = ()) -> list[str]:
|
|
155
|
+
"""Names matching any glob pattern (``ctrl.*.stiffness``) or having any of the scopes."""
|
|
156
|
+
patterns, scopes = list(patterns), set(scopes)
|
|
157
|
+
return [
|
|
158
|
+
name
|
|
159
|
+
for name, param in self._params.items()
|
|
160
|
+
if param.scope in scopes or any(fnmatch.fnmatchcase(name, pat) for pat in patterns)
|
|
161
|
+
]
|
|
162
|
+
|
|
163
|
+
def size(self, names: Iterable[str] | None = None) -> int:
|
|
164
|
+
"""Total number of entries of the named Params (default: all)."""
|
|
165
|
+
return sum(self[n].size for n in self._names_or_all(names))
|
|
166
|
+
|
|
167
|
+
def vector(self, names: Iterable[str] | None = None) -> np.ndarray:
|
|
168
|
+
"""Values of the named Params (default: all) packed into one flat vector."""
|
|
169
|
+
parts = [self[n].value.ravel(order="F") for n in self._names_or_all(names)]
|
|
170
|
+
return np.concatenate(parts) if parts else np.zeros(0)
|
|
171
|
+
|
|
172
|
+
def set_vector(self, x: ArrayLike, names: Iterable[str] | None = None) -> None:
|
|
173
|
+
"""Unpack a flat vector into the named Params (default: all)."""
|
|
174
|
+
x = np.asarray(x, dtype=float).ravel()
|
|
175
|
+
names = self._names_or_all(names)
|
|
176
|
+
if x.size != self.size(names):
|
|
177
|
+
raise ValueError(f"expected {self.size(names)} values, got {x.size}")
|
|
178
|
+
offset = 0
|
|
179
|
+
for n in names:
|
|
180
|
+
param = self[n]
|
|
181
|
+
param.value = x[offset : offset + param.size].reshape(param.shape, order="F")
|
|
182
|
+
offset += param.size
|
|
183
|
+
|
|
184
|
+
def to_dict(self) -> dict[str, Any]:
|
|
185
|
+
"""Values by name, as plain floats and lists."""
|
|
186
|
+
return {name: param.value.tolist() for name, param in self._params.items()}
|
|
187
|
+
|
|
188
|
+
def update(self, values: Mapping[str, ArrayLike]) -> None:
|
|
189
|
+
"""Set values by name."""
|
|
190
|
+
for name, value in values.items():
|
|
191
|
+
self[name].value = value
|
|
192
|
+
|
|
193
|
+
def _names_or_all(self, names: Iterable[str] | None) -> list[str]:
|
|
194
|
+
return list(self._params) if names is None else list(names)
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
def _shaped(block: Any, shape: tuple[int, ...]) -> Any:
|
|
198
|
+
return ca.reshape(block, shape[0], shape[1]) if len(shape) == 2 else block
|
|
199
|
+
|
|
200
|
+
|
|
201
|
+
class Binding:
|
|
202
|
+
"""CasADi expressions for the Params of a set.
|
|
203
|
+
|
|
204
|
+
Live Params are slices of one symbol ``p``; the others are folded in as constants at their
|
|
205
|
+
current values.
|
|
206
|
+
"""
|
|
207
|
+
|
|
208
|
+
def __init__(self, params: ParamSet, live: Iterable[str] = (), symbol: Any = ca.SX) -> None:
|
|
209
|
+
live = set(live)
|
|
210
|
+
unknown = live - set(params)
|
|
211
|
+
if unknown:
|
|
212
|
+
raise KeyError(f"unknown live parameters {sorted(unknown)}")
|
|
213
|
+
self.params = params
|
|
214
|
+
self.live = [name for name in params if name in live]
|
|
215
|
+
self.p = symbol.sym("p", params.size(self.live))
|
|
216
|
+
self._expr: dict[int, Any] = {}
|
|
217
|
+
offset = 0
|
|
218
|
+
for name, param in params.items():
|
|
219
|
+
if name in live:
|
|
220
|
+
block = self.p[offset : offset + param.size]
|
|
221
|
+
offset += param.size
|
|
222
|
+
else:
|
|
223
|
+
block = ca.DM(param.value.ravel(order="F"))
|
|
224
|
+
self._expr[id(param)] = _shaped(block, param.shape)
|
|
225
|
+
|
|
226
|
+
def __call__(self, param: Param) -> Any:
|
|
227
|
+
"""Expression of ``param``."""
|
|
228
|
+
try:
|
|
229
|
+
return self._expr[id(param)]
|
|
230
|
+
except KeyError:
|
|
231
|
+
raise KeyError(f"{param!r} is not in this binding") from None
|
|
232
|
+
|
|
233
|
+
def view(self, params: Mapping[str, Param]) -> dict[str, Any]:
|
|
234
|
+
"""Expressions of a model's Params, keyed by the model's own names."""
|
|
235
|
+
return {name: self(param) for name, param in params.items()}
|
|
236
|
+
|
|
237
|
+
def values(self) -> np.ndarray:
|
|
238
|
+
"""Current values of the live Params, packed like ``p``."""
|
|
239
|
+
return self.params.vector(self.live)
|
|
240
|
+
|
|
241
|
+
|
|
242
|
+
def constants(params: Mapping[str, Param]) -> dict[str, Any]:
|
|
243
|
+
"""Current values of Params as CasADi constants, keyed by name (a numeric view)."""
|
|
244
|
+
return {name: _shaped(ca.DM(p.value.ravel(order="F")), p.shape) for name, p in params.items()}
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
"""Registry: names usable from configuration files, extended by plugins through entry points."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
from collections.abc import Callable
|
|
6
|
+
from importlib.metadata import entry_points
|
|
7
|
+
from typing import Any, TypeVar
|
|
8
|
+
|
|
9
|
+
T = TypeVar("T")
|
|
10
|
+
|
|
11
|
+
GROUP = "virtualmodelcontrol.plugins"
|
|
12
|
+
"""Entry-point group scanned for plugins; loading an entry point runs its ``register`` calls."""
|
|
13
|
+
|
|
14
|
+
_REGISTRY: dict[str, dict[str, Any]] = {}
|
|
15
|
+
_plugins_loaded = False
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def register(kind: str, name: str) -> Callable[[T], T]:
|
|
19
|
+
"""Class or function decorator: make ``obj`` available as ``get(kind, name)``."""
|
|
20
|
+
|
|
21
|
+
def decorate(obj: T) -> T:
|
|
22
|
+
table = _REGISTRY.setdefault(kind, {})
|
|
23
|
+
if name in table and table[name] is not obj:
|
|
24
|
+
raise ValueError(f"a {kind} named {name!r} is already registered")
|
|
25
|
+
table[name] = obj
|
|
26
|
+
return obj
|
|
27
|
+
|
|
28
|
+
return decorate
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def get(kind: str, name: str) -> Any:
|
|
32
|
+
"""Registered object; plugins are loaded on the first miss."""
|
|
33
|
+
if name not in _REGISTRY.get(kind, {}):
|
|
34
|
+
load_plugins()
|
|
35
|
+
table = _REGISTRY.get(kind, {})
|
|
36
|
+
if name not in table:
|
|
37
|
+
raise KeyError(f"no {kind} named {name!r}; known: {sorted(table)}")
|
|
38
|
+
return table[name]
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def names(kind: str) -> list[str]:
|
|
42
|
+
"""Registered names of one kind, plugins included."""
|
|
43
|
+
load_plugins()
|
|
44
|
+
return sorted(_REGISTRY.get(kind, {}))
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def load_plugins() -> None:
|
|
48
|
+
"""Import every entry point in the group ``virtualmodelcontrol.plugins`` (once)."""
|
|
49
|
+
global _plugins_loaded
|
|
50
|
+
if _plugins_loaded:
|
|
51
|
+
return
|
|
52
|
+
_plugins_loaded = True
|
|
53
|
+
for ep in entry_points(group=GROUP):
|
|
54
|
+
ep.load()
|