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
25 changes: 23 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -235,6 +235,27 @@ obs, state, rewards, dones, info = env.step(key, state, actions, params)
# rewards/dones are dicts keyed by agent (+ dones["__all__"]); reward is shared.
```

### Wind & turbulence

Every plane-based env (`Plane`, `Plane3D`, `PlanePatrol`, `PlanePatrolMARL`)
inherits a wind model applied to the air-relative aerodynamics. It is an
**unobservable** disturbance by default — pass `observe_wind=True` for a
fully-observable baseline. Formations feel a single shared gust field.

```python
from target_gym import Plane, PlaneParams

params = PlaneParams(
wind_x=-15.0, # steady mean wind (m/s), world frame
wind_shear_x=0.02, # + linear altitude shear: +0.02 m/s per metre above...
shear_ref_alt=5000.0, # ...this reference altitude
turbulence_sigma=3.0, # + Ornstein-Uhlenbeck gusts (0 = off); theta = turbulence_theta
)

hidden = Plane() # wind is a hidden disturbance (POMDP)
baseline = Plane(observe_wind=True) # appends the realized wind to the observation
```

---

## Challenges Modeled
Expand All @@ -249,13 +270,13 @@ TargetGym tasks are designed to expose RL agents to **realistic control challeng
* [x] **Multi-timescale dynamics**: From millisecond neutronics to hour-long xenon transients (reactor).
* [x] **Moving / non-stationary targets**: The patrol slot tracks a maneuvering lead aircraft.
* [x] **Multi-agent coordination**: The MARL patrol task requires two learners to cooperate (formation + trackable flight).
* [ ] **Non-stationarity**: Introduce perturbations in the environments (wind, turbulence).
* [x] **Non-stationarity / perturbations**: Every plane-based env (2D, 3D, patrol, formation) inherits a full wind model as a physics-engine property — steady wind (`wind_x/y/z`), altitude-dependent **wind shear** (`wind_shear_x/y`), and **Ornstein-Uhlenbeck turbulence** (`turbulence_sigma`), all applied to the air-relative aerodynamics. Formations feel one shared gust field. Wind is an *unobservable* disturbance by default; `Plane(observe_wind=True)` exposes it for a fully-observable baseline.

---

## Roadmap

* [ ] Add perturbations (wind, turbulence) for non-stationary dynamics.
* [ ] Add microburst / spatially-varying wind fields (position-dependent, not just altitude-linear).
* [ ] Provide benchmark results for popular RL baselines.
* [ ] Mature glass furnace and reactor environments (shorter episodes, better reward shaping).
* [ ] Add random orientation variations to circle and heading tasks.
Expand Down
42 changes: 37 additions & 5 deletions src/target_gym/patrol/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@

from target_gym.base import EnvState
from target_gym.experts.pid import Plane3DPIDState, plane3d_heading_pid_step
from target_gym.plane.dynamics import advance_gust
from target_gym.plane3d.dynamics import compute_velocity_3d
from target_gym.plane3d.env import (
PlaneParams3D,
Expand Down Expand Up @@ -62,6 +63,11 @@ class PatrolState(EnvState):
slot_right: float
slot_up: float
lead_turn_rate: float # rad per step commanded onto the lead heading
# Shared formation turbulence gust (m/s): all aircraft are in the same air
# mass, so they feel one common OU gust on top of the steady params.wind.
gust_x: float = 0.0
gust_y: float = 0.0
gust_z: float = 0.0


@struct.dataclass
Expand Down Expand Up @@ -417,26 +423,52 @@ def compute_next_state_patrol(
params: PatrolParams,
lead_pid_params,
integration_method: str = "rk4_1",
key=None,
):
"""One environment step: advance the follower (from ``action``) and the lead."""
# Follower (the learner).
"""One environment step: advance the follower (from ``action``) and the lead.

A single **shared** turbulence gust (advanced with ``key``) is applied to
every aircraft through ``eff_params`` — all planes are in the same air mass —
while each still gets its own altitude-dependent wind shear inside
:func:`compute_next_state_3d`. ``key=None`` keeps a constant (steady) wind.
"""
gust = advance_gust(
jnp.array([state.gust_x, state.gust_y, state.gust_z]),
params.turbulence_theta,
params.turbulence_sigma,
params.delta_t,
key,
)
eff_params = params.replace(
wind_x=params.wind_x + gust[0],
wind_y=params.wind_y + gust[1],
wind_z=params.wind_z + gust[2],
)

# Follower (the learner) — key=None so it does not advance its own per-plane
# gust; the shared gust is already folded into eff_params.wind.
power, stick, aileron = decode_action(action)
new_follower, metrics = compute_next_state_3d(
power,
stick,
aileron,
state.follower,
params,
eff_params,
integration_method=integration_method,
)

# Lead (scripted autopilot).
new_lead, new_pid = step_lead(state, params, lead_pid_params, integration_method)
# Lead (scripted autopilot), same shared wind.
new_lead, new_pid = step_lead(
state, eff_params, lead_pid_params, integration_method
)

new_state = state.replace(
follower=new_follower,
lead=new_lead,
lead_pid=new_pid,
time=state.time + 1,
gust_x=gust[0],
gust_y=gust[1],
gust_z=gust[2],
)
return new_state, metrics
1 change: 1 addition & 0 deletions src/target_gym/patrol/env_jax.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,7 @@ def step_env(
params,
self._lead_pid_params,
integration_method=self.integration_method,
key=key,
)
reward = self.compute_reward(new_state, params)
terminated, truncated = check_is_terminal_patrol(new_state, params, xp=jnp)
Expand Down
32 changes: 29 additions & 3 deletions src/target_gym/patrol/marl.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
slot_error,
)
from target_gym.patrol.rendering import _render_scene
from target_gym.plane.dynamics import advance_gust
from target_gym.plane3d.env import (
PlaneState3D,
compute_next_state_3d,
Expand All @@ -65,6 +66,10 @@ class FormationState(EnvState):
slot_back: jnp.ndarray # (K,)
slot_right: jnp.ndarray # (K,)
slot_up: jnp.ndarray # (K,)
# Shared formation turbulence gust (m/s) applied to every aircraft.
gust_x: float = 0.0
gust_y: float = 0.0
gust_z: float = 0.0


@struct.dataclass
Expand Down Expand Up @@ -283,19 +288,40 @@ def step_env(
params = self.default_params
method = self.integration_method

# One shared turbulence gust for the whole formation (same air mass),
# advanced with the step key and folded into eff_params.wind; each
# aircraft still gets its own altitude-dependent shear in the engine.
gust = advance_gust(
jnp.array([state.gust_x, state.gust_y, state.gust_z]),
params.turbulence_theta,
params.turbulence_sigma,
params.delta_t,
key,
)
eff_params = params.replace(
wind_x=params.wind_x + gust[0],
wind_y=params.wind_y + gust[1],
wind_z=params.wind_z + gust[2],
)

lp, ls, la = decode_action(actions[LEAD])
new_lead, _ = compute_next_state_3d(
lp, ls, la, state.lead, params, integration_method=method
lp, ls, la, state.lead, eff_params, integration_method=method
)
new_wingmen = []
for i, w in enumerate(state.wingmen):
wp, ws, wa = decode_action(actions[wingman_name(i)])
nw, _ = compute_next_state_3d(
wp, ws, wa, w, params, integration_method=method
wp, ws, wa, w, eff_params, integration_method=method
)
new_wingmen.append(nw)
new_state = state.replace(
lead=new_lead, wingmen=tuple(new_wingmen), time=state.time + 1
lead=new_lead,
wingmen=tuple(new_wingmen),
time=state.time + 1,
gust_x=gust[0],
gust_y=gust[1],
gust_z=gust[2],
)

reward, terminated = self._reward_and_terminal(new_state, params)
Expand Down
55 changes: 51 additions & 4 deletions src/target_gym/plane/dynamics.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,41 @@
from target_gym.utils import compute_norm_from_coordinates


def advance_gust(gust, theta: float, sigma: float, dt: float, key):
"""Advance an Ornstein-Uhlenbeck turbulence gust one step.

``gust`` is a wind-deviation vector (m/s) that mean-reverts to zero:
``g' = g - theta*dt*g + sigma*sqrt(dt)*N``. With ``key is None`` (e.g. a
caller that does not model turbulence) or ``sigma == 0`` the gust is
unchanged, so the total wind stays the steady ``params.wind``. Shared by
the 2D and 3D engines so turbulence is a physics-engine property.
"""
if key is None:
return gust
noise = jax.random.normal(key, jnp.shape(gust))
return gust - theta * dt * gust + sigma * jnp.sqrt(dt) * noise


def total_wind_2d(z, gust_x, gust_z, params):
"""Total world-frame wind = steady mean + linear altitude shear + gust.

Single source of truth used by both the transition and the (optional)
wind observation, so they never disagree.
"""
shear = params.wind_shear_x * (z - params.shear_ref_alt)
return params.wind_x + shear + gust_x, params.wind_z + gust_z


def total_wind_3d(z, gust_x, gust_y, gust_z, params):
"""Total world-frame 3D wind = mean + linear altitude shear + gust."""
dz = z - params.shear_ref_alt
return (
params.wind_x + params.wind_shear_x * dz + gust_x,
params.wind_y + params.wind_shear_y * dz + gust_y,
params.wind_z + gust_z,
)


def compute_drag(S: float, C: float, V: float, rho: float) -> float:
"""
Compute the drag.
Expand Down Expand Up @@ -428,13 +463,25 @@ def compute_acceleration(
"""
Compute linear and angular accelerations for the aircraft.
Returns: (a_x, a_z, alpha_y, metrics)

Aerodynamics act on the **air-relative** velocity ``V_ground - wind`` (wind
read from ``params.wind_x``/``params.wind_z`` — a physics-engine property,
so any plane-based env that uses these params gets it): the angle of attack,
flight-path angle, airspeed, Mach and dynamic pressure all derive from it,
so a headwind/tailwind/crosswind changes the forces correctly. The
resulting accelerations still act on the ground velocity (Newton's law is in
the inertial frame), and position integrates the ground velocity — so with
``wind = 0`` (the default) the behaviour is unchanged.
"""
xp = jnp
thrust, stick = action
x_dot, z_dot, _ = velocities

_, z, theta = positions
alpha, gamma = compute_alpha(theta, x_dot, z_dot)
# Air-relative velocity drives all aerodynamics.
air_x = x_dot - params.wind_x
air_z = z_dot - params.wind_z
alpha, gamma = compute_alpha(theta, air_x, air_z)
# jax.debug.print(
# "{x} {y} {z}", x=jnp.rad2deg(alpha), y=jnp.rad2deg(gamma), z=jnp.rad2deg(theta)
# )
Expand All @@ -443,12 +490,12 @@ def compute_acceleration(
) # TODO : make it the actual mass when we start considering fuel consumption
rho = compute_air_density_from_altitude(z)
M = compute_Mach_from_velocity_and_speed_of_sound(
velocity=compute_velocity_from_horizontal_and_vertical_speed(x_dot, z_dot),
velocity=compute_velocity_from_horizontal_and_vertical_speed(air_x, air_z),
speed_of_sound=compute_speed_of_sound_from_altitude(z),
)
# --- Weight & velocity ---
# --- Weight & airspeed ---
P = compute_weight(m, params.gravity)
V = compute_norm_from_coordinates(xp.array([x_dot, z_dot]))
V = compute_norm_from_coordinates(xp.array([air_x, air_z]))

# ====================================================
# WINGS
Expand Down
58 changes: 54 additions & 4 deletions src/target_gym/plane/env.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
integrate_dynamics,
)
from target_gym.plane.dynamics import (
advance_gust,
compute_acceleration,
compute_air_density_from_altitude,
compute_alpha,
Expand All @@ -19,6 +20,7 @@
compute_speed_of_sound_from_altitude,
compute_thrust_output,
compute_velocity_from_horizontal_and_vertical_speed,
total_wind_2d,
)

SPEED_OF_SOUND = 343.0
Expand Down Expand Up @@ -52,6 +54,10 @@ class PlaneState(EnvState):
stick: float
fuel: float
target_altitude: float
# Ornstein-Uhlenbeck turbulence gust (m/s), mean-reverting to 0. Total wind
# acting on the aircraft is params.wind_* + gust_*. Default 0 => no gust.
gust_x: float = 0.0
gust_z: float = 0.0

@property
def rho(self):
Expand Down Expand Up @@ -111,6 +117,21 @@ class PlaneParams(EnvParams):
initial_power: float = 1.0
initial_stick: float = 0.0

# Steady mean wind (world frame, m/s). Aerodynamics use the air-relative
# velocity V_ground - (wind + gust).
wind_x: float = 0.0
wind_z: float = 0.0
# Ornstein-Uhlenbeck turbulence on top of the mean wind: sigma is the gust
# std (m/s), theta the mean-reversion rate (1/s; correlation time ~ 1/theta).
# sigma = 0 (default) => no turbulence, wind is exactly the steady mean.
turbulence_sigma: float = 0.0
turbulence_theta: float = 0.2
# Linear wind shear: the horizontal wind gains ``wind_shear_x`` m/s per metre
# of altitude above ``shear_ref_alt`` (0 => no shear). Makes the wind
# altitude-dependent, so climbing/descending is itself a disturbance.
wind_shear_x: float = 0.0
shear_ref_alt: float = 0.0

delta_t: float = 1.0


Expand Down Expand Up @@ -194,25 +215,52 @@ def compute_next_state(
state: PlaneState,
params: PlaneParams,
integration_method: str = "rk4_1",
key=None,
):
"""Compute next state and metrics using multiple sub-steps with jax.lax.scan."""
"""Compute next state and metrics using multiple sub-steps with jax.lax.scan.

Wind is a physics-engine property: the total wind acting on the aircraft is
the steady ``params.wind_*`` plus an Ornstein-Uhlenbeck turbulence gust
(advanced with ``key`` when ``params.turbulence_sigma > 0``), and the
aerodynamics/engine-Mach use the air-relative velocity while
position/observations stay in the ground frame. With ``wind = 0`` and
``turbulence_sigma = 0`` the behaviour is unchanged; ``key=None`` freezes the
gust (so callers that don't model turbulence keep a constant wind).
"""
dt = params.delta_t
power = compute_next_power(power_requested, state.power, dt)
stick = compute_next_stick(stick_requested, state.stick, dt)

# Compute thrustx
# Total wind = steady mean + altitude shear + OU turbulence gust.
gust = advance_gust(
jnp.array([state.gust_x, state.gust_z]),
params.turbulence_theta,
params.turbulence_sigma,
dt,
key,
)
total_wind_x, total_wind_z = total_wind_2d(state.z, gust[0], gust[1], params)
eff_params = params.replace(wind_x=total_wind_x, wind_z=total_wind_z)

# Engine ram/Mach effects depend on airspeed, not ground speed.
air_speed = compute_velocity_from_horizontal_and_vertical_speed(
state.x_dot - total_wind_x, state.z_dot - total_wind_z
)
M_air = compute_Mach_from_velocity_and_speed_of_sound(
air_speed, state.speed_of_sound
)
thrust = compute_thrust_output(
power=power,
thrust_output_at_sea_level=params.thrust_output_at_sea_level,
rho=state.rho,
M=state.M,
M=M_air,
)
positions = jnp.array([state.x, state.z, state.theta])
velocities = jnp.array([state.x_dot, state.z_dot, state.theta_dot])
_compute_acceleration = partial(
compute_acceleration,
action=(thrust, stick),
params=params,
params=eff_params,
clip=True,
min_clip_boundaries=(-100, -100, -1.5),
max_clip_boundaries=(100, 100, 1.5),
Expand Down Expand Up @@ -244,5 +292,7 @@ def compute_next_state(
fuel=state.fuel,
time=state.time + 1,
target_altitude=state.target_altitude,
gust_x=gust[0],
gust_z=gust[1],
)
return new_state, metrics
Loading
Loading