diff --git a/docs/design/input_action_spaces.md b/docs/design/input_action_spaces.md index 143dce5..50f142e 100644 --- a/docs/design/input_action_spaces.md +++ b/docs/design/input_action_spaces.md @@ -16,8 +16,7 @@ Global inputs, always broadcasted to all nodes: $$ tilt = sqrt(roll^2 + pitch^2) $$ -- Goal vector: Instead of just a scalar distance, the goal is represented asa a vector (distance and ange/direction) to - the target. +- Goal vector: Instead of just a scalar distance, the goal is represented as an angle/direction to the target. Local inputs, routed directly to specific nodes: @@ -70,6 +69,10 @@ When designing the state space, we must ask: *Could a human operator perform thi as a normalized unit vector bounds the values to the $[-1, 1]$ range, which stabilizes neural network training. Providing only a scalar "distance to the goal" would force the agent to learning localized searching behaviors (e.g. random walks or spiraling) to deduce the direction, drastically increasing the difficulty of the learning task. + + **NOTE:** We later dropped the "distance to vector", switching to only a direction as the input. Our reasoning is + the agent should always move towards the goal (it should not learn to stop at the goal), which allows for this + simplification that decreases the model input size. - Contact sensors: Segment contact detects external ground interaction and is biologically vital for timing gait transitions. - **Zero-Centered Rescaling ($[-1, 1]$):** Using a zero-centered range is standard best practice for continuous control diff --git a/src/brittle_star_project/environment/env_config.py b/src/brittle_star_project/environment/env_config.py index 0af387e..b9a621d 100644 --- a/src/brittle_star_project/environment/env_config.py +++ b/src/brittle_star_project/environment/env_config.py @@ -72,7 +72,6 @@ class ObservationBoundsConfig: joint_actuator_force: list[float] = field(default_factory=lambda: [-5.0, 5.0]) segment_contact: list[float] = field(default_factory=lambda: [0.0, 1.0]) unit_xy_direction_to_target: list[float] = field(default_factory=lambda: [-1.0, 1.0]) - xy_distance_to_target: list[float] = field(default_factory=lambda: [0.0, 20.0]) disk_z_tilt: list[float] = field(default_factory=lambda: [0.0, 3.141592653589793]) def to_bounds_dict(self) -> dict[str, tuple[float, float]]: @@ -82,6 +81,5 @@ class ObservationBoundsConfig: "joint_actuator_force": tuple(self.joint_actuator_force), "segment_contact": tuple(self.segment_contact), "unit_xy_direction_to_target": tuple(self.unit_xy_direction_to_target), - "xy_distance_to_target": tuple(self.xy_distance_to_target), "disk_z_tilt": tuple(self.disk_z_tilt), } diff --git a/src/brittle_star_project/environment/obs_processing.py b/src/brittle_star_project/environment/obs_processing.py index 0ef6e3d..37a775f 100644 --- a/src/brittle_star_project/environment/obs_processing.py +++ b/src/brittle_star_project/environment/obs_processing.py @@ -63,7 +63,6 @@ def create_obs_processor( "joint_velocity", "segment_contact", "unit_xy_direction_to_target", - "xy_distance_to_target", ] values = [] for key in ordered_keys: