Communicative_CRN / src /evaluator.py
medimaging's picture
Upload evaluator.py
403f471 verified
Raw
History Blame Contribute Delete
5.2 kB
import numpy as np
import torch
from itertools import chain
class Evaluator(object):
def __init__(self, environment, model, logger, agents, max_steps):
self.env = environment
self.model = model
self.logger = logger
self.agents = agents
self.max_steps = max_steps
def play_n_episodes(self, render=False, fixed_spawn=None, silent=False):
"""
wraps play_one_episode, playing a single episode at a time and logs
results used when playing demos.
"""
if fixed_spawn is None:
num_runs = 1
fixed_spawn = [None]
else:
# fixed_spawn should be, for example, [0.5 , 0.5 , 0.5, 0, 0, 0] for 2 runs
# In the first run agents spawn in the middle and in the second they will spawn from the corner
fixed_spawn = np.array(fixed_spawn).reshape((-1, 3)) # 3 dimensions
num_runs = fixed_spawn.shape[0]
# Set all the agents to the same spawn point
fixed_spawn = np.stack([fixed_spawn for _ in range(self.agents)], axis=-1)
num_files = self.env.files.num_files
self.model.train(False)
headers = ["number"] + list(chain.from_iterable(zip(
[f"Filename {i}" for i in range(self.agents)],
[f"Agent {i} pos x" for i in range(self.agents)],
[f"Agent {i} pos y" for i in range(self.agents)],
[f"Agent {i} pos z" for i in range(self.agents)],
[f"Landmark {i} pos x" for i in range(self.agents)],
[f"Landmark {i} pos y" for i in range(self.agents)],
[f"Landmark {i} pos z" for i in range(self.agents)],
[f"Distance {i}" for i in range(self.agents)])))
self.logger.write_locations(headers)
distances = []
for j in range(num_runs):
for k in range(num_files):
score, start_dists, q_values, info = self.play_one_episode(render, fixed_spawn=fixed_spawn[j])
row = [j * num_files + k + 1] + list(chain.from_iterable(zip(
[info[f"filename_{i}"] for i in range(self.agents)],
[info[f"agent_xpos_{i}"] for i in range(self.agents)],
[info[f"agent_ypos_{i}"] for i in range(self.agents)],
[info[f"agent_zpos_{i}"] for i in range(self.agents)],
[info.get(f"landmark_xpos_{i}", "N/A") for i in range(self.agents)],
[info.get(f"landmark_ypos_{i}", "N/A") for i in range(self.agents)],
[info.get(f"landmark_zpos_{i}", "N/A") for i in range(self.agents)],
[info.get(f"distError_{i}", "N/A") for i in range(self.agents)])))
for i in range(self.agents):
key = f"distError_{i}"
if key in info:
distances.append(info[key])
self.logger.write_locations(row)
if len(distances) == 0:
return None, None # No distance mean and std for task "play" as there is no ground truth
mean = np.mean(distances, 0)
std = np.std(distances, 0, ddof=1)
if not silent:
self.logger.log(f"mean distances {mean}")
self.logger.log(f"Std distances {std}")
return mean, std
def play_one_episode(self, render=False, frame_history=4, fixed_spawn=None):
device = next(self.model.parameters()).device
def predict(obs_stack):
"""
Run a full episode, mapping observation to action,
using greedy policy.
"""
inputs = torch.from_numpy(obs_stack).float().permute(
0, 4, 1, 2, 3).unsqueeze(0).to(device)
with torch.no_grad():
q_vals = self.model(inputs)
idx = torch.max(q_vals, -1)[1]
greedy_steps = np.array(idx, dtype=np.int32).flatten()
return greedy_steps, q_vals.detach().cpu().numpy().squeeze(0)
obs_stack = self.env.reset(fixed_spawn)
# Here obs have shape (agent, *image_size, frame_history)
sum_r = np.zeros((self.agents))
isOver = [False] * self.agents
start_dists = None
steps = 0
while steps < self.max_steps and not np.all(isOver):
acts, q_values = predict(obs_stack)
obs_stack, r, isOver, info = self.env.step(acts, q_values, isOver)
steps += 1
# Machine-parseable trajectory log line (agent 0 only). Consumed
# by app.py to render the agent's search path as an animation.
print("STEP_LOC: {} {} {} {}".format(
steps,
info.get('agent_xpos_0', 'NA'),
info.get('agent_ypos_0', 'NA'),
info.get('agent_zpos_0', 'NA'),
))
if start_dists is None:
start_dists = [
info.get('distError_' + str(i), "N/A") for i in range(self.agents)]
if render:
self.env.render()
for i in range(self.agents):
if not isOver[i]:
sum_r[i] += r[i]
return sum_r, start_dists, q_values, info