-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathevaluate.py
More file actions
103 lines (84 loc) · 3.75 KB
/
Copy pathevaluate.py
File metadata and controls
103 lines (84 loc) · 3.75 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
"""
Evaluation script for Weco optimization.
Runs the AI controller as a local CPU simulation across multiple seeds and reports the average score.
This is what Weco uses to measure the quality of each candidate solution.
The metric is average_score (higher is better).
"""
import sys
import os
# Ensure we're in the right directory
script_dir = os.path.dirname(os.path.abspath(__file__))
os.chdir(script_dir)
sys.path.insert(0, script_dir)
from game import FlappyBirdGame # noqa: E402
# Evaluation parameters
DEFAULT_EVAL_WORLDS = ("normal", "moving_gap", "wind", "variable")
EVAL_WORLDS = tuple(
world.strip()
for world in os.environ.get("FLAPPY_EVAL_WORLDS", ",".join(DEFAULT_EVAL_WORLDS)).split(",")
if world.strip()
)
GAMES_PER_WORLD = int(os.environ.get("FLAPPY_GAMES_PER_WORLD", "5"))
EVAL_SEED_START = int(os.environ.get("FLAPPY_SEED_START", "100"))
MAX_FRAMES = int(os.environ.get("FLAPPY_MAX_FRAMES", "5000"))
MAX_SCORE = int(os.environ.get("FLAPPY_MAX_SCORE", "50"))
def _eval_cases() -> list[tuple[str, int]]:
seeds = range(EVAL_SEED_START, EVAL_SEED_START + GAMES_PER_WORLD)
return [(world, seed) for world in EVAL_WORLDS for seed in seeds]
def evaluate():
"""Run evaluation and print the metric."""
try:
# Reload the AI controller to pick up Weco's changes
if "ai_controller" in sys.modules:
del sys.modules["ai_controller"]
from ai_controller import should_flap
except Exception as e:
print("average_score: 0", flush=True)
print(f"Error importing ai_controller: {e}", file=sys.stderr)
return
scores = []
frames_list = []
scores_by_world = {world: [] for world in EVAL_WORLDS}
for world, seed in _eval_cases():
try:
game = FlappyBirdGame(seed=seed, world=world)
obs = game.get_observation()
frame_count = 0
while (
not game.game_over
and frame_count < MAX_FRAMES
and game.score < MAX_SCORE
):
action = should_flap(obs)
obs, _reward, _done = game.step(action)
frame_count += 1
scores.append(game.score)
scores_by_world[world].append(game.score)
frames_list.append(frame_count)
except Exception as e:
# If the AI crashes, it gets 0 for this game
scores.append(0)
scores_by_world[world].append(0)
frames_list.append(0)
print(f"Error in game world={world} seed={seed}: {e}", file=sys.stderr)
avg_score = sum(scores) / len(scores) if scores else 0
max_score = max(scores) if scores else 0
min_score = min(scores) if scores else 0
avg_frames = sum(frames_list) / len(frames_list) if frames_list else 0
# Print the metric Weco expects (exactly one line: metric_name: value)
print(f"average_score: {avg_score}", flush=True)
# Additional info on stderr (won't affect Weco's metric parsing)
print(f" max_score: {max_score}", file=sys.stderr)
print(f" min_score: {min_score}", file=sys.stderr)
print(f" avg_frames: {avg_frames:.0f}", file=sys.stderr)
for world, world_scores in scores_by_world.items():
world_avg = sum(world_scores) / len(world_scores) if world_scores else 0
print(f" {world}_avg: {world_avg}", file=sys.stderr)
print(f" eval_worlds: {','.join(EVAL_WORLDS)}", file=sys.stderr)
print(f" games_per_world: {GAMES_PER_WORLD}", file=sys.stderr)
print(f" eval_seeds: {EVAL_SEED_START}-{EVAL_SEED_START + GAMES_PER_WORLD - 1}", file=sys.stderr)
print(f" max_frames: {MAX_FRAMES}", file=sys.stderr)
print(f" score_cap: {MAX_SCORE}", file=sys.stderr)
print(f" all_scores: {scores}", file=sys.stderr)
if __name__ == "__main__":
evaluate()