Skip to content
Closed

Adico #176

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
69 changes: 69 additions & 0 deletions vmas/scenarios/balance.py
Original file line number Diff line number Diff line change
Expand Up @@ -257,6 +257,75 @@ def observation(self, agent: Agent):
dim=-1,
)

def observation_from_pos(self, pos: torch.Tensor, env_index: int = None):
"""
Get observation from a given position for balance scenario.

Args:
pos: Position tensor of shape [batch_size, 2] or [2]
env_index: Index of the environment (optional)

Returns:
Observation tensor as if an agent were at the given position
"""
# Ensure pos has correct shape
if pos.dim() == 1:
pos = pos.unsqueeze(0) # [2] -> [1, 2]

batch_size = pos.shape[0]

# Get states for the specified environment
if env_index is None:
package_pos = self.package.state.pos[0].unsqueeze(0) # [1, 2]
package_vel = self.package.state.vel[0].unsqueeze(0) # [1, 2]
line_pos = self.line.state.pos[0].unsqueeze(0) # [1, 2]
line_vel = self.line.state.vel[0].unsqueeze(0) # [1, 2]
line_ang_vel = self.line.state.ang_vel[0].unsqueeze(0) # [1, 1]
line_rot = self.line.state.rot[0].unsqueeze(0) # [1, 1]
goal_pos = self.package.goal.state.pos[0].unsqueeze(0) # [1, 2]
else:
package_pos = self.package.state.pos[env_index].unsqueeze(0) # [1, 2]
package_vel = self.package.state.vel[env_index].unsqueeze(0) # [1, 2]
line_pos = self.line.state.pos[env_index].unsqueeze(0) # [1, 2]
line_vel = self.line.state.vel[env_index].unsqueeze(0) # [1, 2]
line_ang_vel = self.line.state.ang_vel[env_index].unsqueeze(0) # [1, 1]
line_rot = self.line.state.rot[env_index].unsqueeze(0) # [1, 1]
goal_pos = self.package.goal.state.pos[env_index].unsqueeze(0) # [1, 2]

# Expand to match batch size
package_pos = package_pos.expand(batch_size, -1)
package_vel = package_vel.expand(batch_size, -1)
line_pos = line_pos.expand(batch_size, -1)
line_vel = line_vel.expand(batch_size, -1)
line_ang_vel = line_ang_vel.expand(batch_size, -1)
line_rot = line_rot.expand(batch_size, -1)
goal_pos = goal_pos.expand(batch_size, -1)

# Create zero velocity for the hypothetical agent at pos
agent_vel = torch.zeros_like(pos)

# Match the structure of the observation method:
# [agent.state.pos, agent.state.vel,
# agent.state.pos - package.state.pos,
# agent.state.pos - line.state.pos,
# package.state.pos - goal.state.pos,
# package.state.vel, line.state.vel, line.state.ang_vel,
# line.state.rot % pi]
return torch.cat(
[
pos, # agent position
agent_vel, # agent velocity (zero for static query point)
pos - package_pos, # relative position to package
pos - line_pos, # relative position to line
package_pos - goal_pos, # package to goal
package_vel, # package velocity
line_vel, # line velocity
line_ang_vel, # line angular velocity
line_rot % torch.pi, # line rotation
],
dim=-1,
)

def done(self):
return self.on_the_ground + self.world.is_overlapping(
self.package, self.package.goal
Expand Down
55 changes: 55 additions & 0 deletions vmas/scenarios/ball_passage.py
Original file line number Diff line number Diff line change
Expand Up @@ -287,6 +287,61 @@ def observation(self, agent: Agent):
dim=-1,
)

def observation_from_pos(self, pos: torch.Tensor, env_index: int = None):
"""
Get observation from a given position for ball passage scenario.

Args:
pos: Position tensor of shape [batch_size, 2] or [2]
env_index: Index of the environment (optional)

Returns:
Observation tensor as if an agent were at the given position
"""
# Ensure pos has correct shape [batch_size, 2]
if pos.dim() == 1:
pos = pos.unsqueeze(0)

batch_size = pos.shape[0]
device = self.world.device

# Helper to get the correct state slice
def get_state(entity_state_attr):
if env_index is None:
# If no index, we assume we want the first env or the whole batch
# To match pos batch_size, we take the first and expand
return entity_state_attr[0].unsqueeze(0).expand(batch_size, -1)
return entity_state_attr[env_index].unsqueeze(0).expand(batch_size, -1)

# Extract states
ball_pos = get_state(self.ball.state.pos)
goal_pos = get_state(self.goal.state.pos)

# Hypothetical agent velocity is zero for a static position query
agent_vel = torch.zeros_like(pos)

# Get positions of all "open" passages (where collide is False)
# Note: In VMAS, the list of passages is the same across batch dims,
# but their positions vary per env_index
passage_obs = []
for passage in self.passages:
if not passage.collide:
p_pos = get_state(passage.state.pos)
passage_obs.append(pos - p_pos)

# Match the structure of Scenario.observation():
# [pos, vel, pos - goal, pos - ball, *(pos - open_passages)]
return torch.cat(
[
pos,
agent_vel,
pos - goal_pos,
pos - ball_pos,
*passage_obs,
],
dim=-1,
)

def done(self):
return (
(
Expand Down
44 changes: 44 additions & 0 deletions vmas/scenarios/ball_trajectory.py
Original file line number Diff line number Diff line change
Expand Up @@ -210,6 +210,50 @@ def observation(self, agent: Agent):
dim=-1,
)

def observation_from_pos(self, pos: torch.Tensor, env_index: int = None):
"""
Get observation from a given position for ball trajectory scenario.

Args:
pos: Position tensor of shape [batch_size, 2] or [2]
env_index: Index of the environment (optional)

Returns:
Observation tensor as if an agent were at the given position
"""
# Ensure pos has correct shape [batch_size, 2]
if pos.dim() == 1:
pos = pos.unsqueeze(0)

batch_size = pos.shape[0]

# Helper to get the correct state slice (single env vs whole batch)
def get_state(entity_state_attr):
if env_index is None:
# Use the first environment's state and expand to match batch_size
return entity_state_attr[0].unsqueeze(0).expand(batch_size, -1)
# Use the specific environment's state
return entity_state_attr[env_index].unsqueeze(0).expand(batch_size, -1)

# Extract ball state
ball_pos = get_state(self.ball.state.pos)

# Hypothetical agent velocity is zero for a static position query
agent_vel = torch.zeros_like(pos)

# Match the structure of Scenario.observation():
# [agent.pos, agent.vel, agent.pos - ball.pos, agent.pos]
# Note: Your original observation duplicates agent.pos at the end.
return torch.cat(
[
pos, # agent.state.pos
agent_vel, # agent.state.vel
pos - ball_pos, # agent.state.pos - ball.state.pos
pos, # agent.state.pos (the duplicate from your code)
],
dim=-1,
)

def info(self, agent: Agent) -> Dict[str, Tensor]:
return {
"pos_rew": self.pos_rew,
Expand Down
42 changes: 42 additions & 0 deletions vmas/scenarios/buzz_wire.py
Original file line number Diff line number Diff line change
Expand Up @@ -245,6 +245,48 @@ def observation(self, agent: Agent):
dim=-1,
)

def observation_from_pos(self, pos: torch.Tensor, env_index: int = None):
"""
Get observation from a given position for the buzz wire scenario.

Args:
pos: Position tensor of shape [batch_size, 2] or [2]
env_index: Index of the environment (optional)

Returns:
Observation tensor as if an agent were at the given position
"""
# Ensure pos has correct shape [batch_size, 2]
if pos.dim() == 1:
pos = pos.unsqueeze(0)

batch_size = pos.shape[0]

# Helper to get the correct state slice (single env vs whole batch)
def get_state(entity_state_attr):
if env_index is None:
# Use the first environment's state and expand to match batch_size
return entity_state_attr[0].unsqueeze(0).expand(batch_size, -1)
# Use the specific environment's state
return entity_state_attr[env_index].unsqueeze(0).expand(batch_size, -1)

# Extract goal state
goal_pos = get_state(self.goal.state.pos)

# Hypothetical agent velocity is zero for a static position query
agent_vel = torch.zeros_like(pos)

# Match the structure of Scenario.observation():
# [agent.pos, agent.vel, agent.pos - goal.pos]
return torch.cat(
[
pos, # agent.state.pos
agent_vel, # agent.state.vel
pos - goal_pos, # agent.state.pos - goal.state.pos
],
dim=-1,
)

def done(self):
return (
torch.linalg.vector_norm(self.ball.state.pos - self.goal.state.pos, dim=1)
Expand Down
65 changes: 65 additions & 0 deletions vmas/scenarios/debug/empty.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@

import torch

from vmas import render_interactively

from vmas.simulator.core import Agent, Sphere, World
from vmas.simulator.scenario import BaseScenario
from vmas.simulator.utils import Color, ScenarioUtils


class Scenario(BaseScenario):
def make_world(self, batch_dim: int, device: torch.device, **kwargs):

self.n_agents = 4
self.agent_radius = 0.16

# Make world
world = World(batch_dim, device, x_semidim=0, y_semidim=0)

self.colors = [Color.GREEN, Color.BLUE, Color.RED, Color.GRAY]

# Add agents
for i in range(self.n_agents):
agent = Agent(
name=f"agent_{i}",
rotatable=False,
shape=Sphere(radius=self.agent_radius),
render_action=True,
color=self.colors[i],
collide=False,
)
world.add_agent(agent)

return world

def reset_world_at(self, env_index: int = None):

ScenarioUtils.spawn_entities_randomly(
self.world.agents,
self.world,
env_index,
min_dist_between_entities=0,
x_bounds=(
0,
0,
),
y_bounds=(
0,
0,
),
)

def reward(self, agent: Agent):
return torch.zeros(
self.world.batch_dim, device=self.world.device, dtype=torch.float32
)

def observation(self, agent: Agent):
return torch.zeros(
self.world.batch_dim, 1, device=self.world.device, dtype=torch.float32
)


if __name__ == "__main__":
render_interactively(__file__, control_two_agents=True)
50 changes: 50 additions & 0 deletions vmas/scenarios/discovery.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,6 +250,56 @@ def observation(self, agent: Agent):
dim=-1,
)

def observation_from_pos( self, pos: Tensor, env_index: int = None, agent_index: int = 0 ):
env_index = 0 if env_index is None else env_index
agent = self.world.agents[agent_index]
pos = pos.to(
device=self.world.device, dtype=agent.state.pos.dtype
).reshape(-1, 2)
n_points = pos.shape[0]

observations = [pos, torch.zeros_like(pos)]

for sensor in agent.sensors:
entities = [
entity
for entity in self.world.entities
if entity is not agent and sensor.entity_filter(entity)
]

if not entities:
measurements = pos.new_full(
(n_points, sensor._angles.shape[-1]),
sensor._max_range,
)
else:
sphere_pos = torch.stack(
[entity.state.pos[env_index] for entity in entities]
).unsqueeze(0).expand(n_points, -1, -1)

sphere_radius = pos.new_tensor(
[entity.shape.radius for entity in entities]
).unsqueeze(0).expand(n_points, -1)

angles = (
sensor._angles[env_index] + agent.state.rot[env_index]
).unsqueeze(0).expand(n_points, -1)

distances = self.world._cast_rays_to_sphere(
sphere_pos,
sphere_radius,
pos,
angles,
sensor._max_range,
)
measurements = distances.min(dim=-2).values.clamp(
max=sensor._max_range
)

observations.append(measurements)

return torch.cat(observations, dim=-1)

def info(self, agent: Agent) -> Dict[str, Tensor]:
info = {
"covering_reward": (
Expand Down
Loading