-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpong_env.py
More file actions
276 lines (223 loc) · 10.6 KB
/
Copy pathpong_env.py
File metadata and controls
276 lines (223 loc) · 10.6 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
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
#!/usr/bin/env python3
"""
Pong-Umgebung für Reinforcement Learning
"""
import numpy as np
import gymnasium as gym
from gymnasium import spaces
from stable_baselines3 import PPO, A2C, DQN
from game_logic import PongGame
import sys
def load_model_with_auto_detection(model_path):
"""Lädt ein Modell und erkennt automatisch den Algorithmus"""
try:
# Versuche zuerst PPO
return PPO.load(model_path)
except Exception as e1:
try:
# Versuche A2C
return A2C.load(model_path)
except Exception as e2:
try:
# Versuche DQN
return DQN.load(model_path)
except Exception as e3:
# Wenn alle fehlschlagen, gib den ursprünglichen Fehler zurück
raise e1
class PongEnv(gym.Env):
"""Pong-Umgebung für Reinforcement Learning"""
def __init__(self, width: int = 600, height: int = 400, max_score: int = 1,
max_steps: int = sys.maxsize, opponent_model_path: str = None):
super().__init__()
self.width = width
self.height = height
self.max_score = max_score
self.max_steps = max_steps
self.opponent_model_path = opponent_model_path
# Action Space: 0 = nichts, 1 = hoch, 2 = runter
self.action_space = spaces.Discrete(3)
# Observation Space: [ball_x, ball_y, paddle_y, ball_vel_x, ball_vel_y]
# Normalisiert zwischen 0 und 1
self.observation_space = spaces.Box(
low=0, high=1, shape=(5,), dtype=np.float32
)
# Spielinstanz
self.game = None
self.previous_score = 0
self.previous_paddle_hit = False
self.step_count = 0
# Gegner-Modell laden falls angegeben
self.opponent_model = None
if opponent_model_path:
try:
self.opponent_model = load_model_with_auto_detection(opponent_model_path)
print(f"Gegner-Modell erfolgreich aus {opponent_model_path} geladen")
except Exception as e:
print(f"Fehler beim Laden des Gegner-Modells: {e}")
print("Verwende scripted Agent als Gegner...")
self.opponent_model = None
def reset(self, seed=None, options=None):
"""Setzt die Umgebung zurück"""
super().reset(seed=seed)
self.game = PongGame(self.width, self.height)
self.previous_score = 0
self.previous_paddle_hit = False
self.step_count = 0
# Gegner-AI aktualisieren
self._update_opponent_paddle()
return self._get_observation(), {}
def step(self, action):
"""Führt eine Aktion aus und gibt den neuen Zustand zurück"""
# Gegner-AI aktualisieren
self._update_opponent_paddle()
# RL-Agent steuert Paddle 2 (rechts)
if action == 1: # Hoch
self.game.set_paddle2_velocity(-8)
elif action == 2: # Runter
self.game.set_paddle2_velocity(8)
else: # Nichts
self.game.set_paddle2_velocity(0)
# Spiel aktualisieren
self.game.update()
# Schritt-Zähler erhöhen
self.step_count += 1
# Beobachtung und Reward berechnen
observation = self._get_observation()
reward = self._calculate_reward()
terminated = self._is_terminated()
truncated = self._is_truncated()
return observation, reward, terminated, truncated, {}
def _get_observation(self):
"""Gibt die normalisierte Beobachtung zurück"""
# Ball-Position normalisieren
ball_x = self.game.ball_pos[0] / self.width
ball_y = self.game.ball_pos[1] / self.height
# Paddle-Position normalisieren (nur Y-Koordinate)
paddle_y = self.game.paddle2_pos[1] / self.height
# Ball-Geschwindigkeit normalisieren
ball_vel_x = (self.game.ball_vel[0] + 10) / 20 # Normalisiert von [-10, 10] zu [0, 1]
ball_vel_y = (self.game.ball_vel[1] + 10) / 20 # Normalisiert von [-10, 10] zu [0, 1]
return np.array([ball_x, ball_y, paddle_y, ball_vel_x, ball_vel_y], dtype=np.float32)
def _get_opponent_observation(self):
"""Gibt die Beobachtung für den Gegner zurück (gespiegelt)"""
# Ball-Position normalisieren (gespiegelt für Gegner)
ball_x = (self.width - self.game.ball_pos[0]) / self.width
ball_y = self.game.ball_pos[1] / self.height
# Paddle-Position normalisieren (nur Y-Koordinate)
paddle_y = self.game.paddle1_pos[1] / self.height
# Ball-Geschwindigkeit normalisieren (X-Geschwindigkeit gespiegelt)
ball_vel_x = (-self.game.ball_vel[0] + 10) / 20 # Normalisiert von [-10, 10] zu [0, 1]
ball_vel_y = (self.game.ball_vel[1] + 10) / 20 # Normalisiert von [-10, 10] zu [0, 1]
return np.array([ball_x, ball_y, paddle_y, ball_vel_x, ball_vel_y], dtype=np.float32)
def _calculate_reward(self):
"""Berechnet den Reward basierend auf dem Spielzustand"""
reward = 0
# 1. Hauptrewards für wichtige Ereignisse
current_paddle_hit = self._check_paddle_hit()
if current_paddle_hit and not self.previous_paddle_hit:
# Basis-Reward für erfolgreiche Abwehr
reward += 1.0
# Zusätzliche Qualitätsbewertung
hit_quality = self._calculate_hit_quality()
reward += hit_quality
self.previous_paddle_hit = current_paddle_hit
# 2. Punktgewinn/Verlust (stärkste Rewards)
current_max_score = max(self.game.l_score, self.game.r_score)
if current_max_score > self.previous_score:
if self.game.r_score > self.game.l_score:
reward += 10.0 # Starker positiver Reward für Sieg
else:
reward -= 5.0 # Negativer Reward für Niederlage
self.previous_score = current_max_score
# 3. Vereinfachte Ball-Nähe-Bewertung
ball_x = self.game.ball_pos[0]
paddle_y = self.game.paddle2_pos[1]
ball_y = self.game.ball_pos[1]
ball_vel_x = self.game.ball_vel[0]
# Nur belohnen wenn Ball sich auf das Paddle zubewegt
if ball_vel_x > 0 and ball_x > self.width * 0.6: # Ball bewegt sich nach rechts
distance_to_ball = abs(ball_y - paddle_y)
# Einfache lineare Belohnung für Nähe
if distance_to_ball < 50:
proximity_reward = 0.1 * (1 - distance_to_ball / 50)
reward += proximity_reward
# 4. Kleiner negativer Reward für jeden Zeitschritt (ermutigt schnelles Spiel)
reward -= 0.01
return reward
def _check_paddle_hit(self):
"""Überprüft, ob der Ball das Paddle getroffen hat"""
# Vereinfachte Kollisionserkennung für Paddle 2 (rechts)
ball_x = self.game.ball_pos[0]
ball_y = self.game.ball_pos[1]
paddle_y = self.game.paddle2_pos[1]
# Prüfe ob Ball in der Nähe des Paddles ist
if (ball_x >= self.width - 30 and # Ball ist nah am rechten Paddle
abs(ball_y - paddle_y) <= 40): # Ball ist auf Höhe des Paddles
return True
return False
def _calculate_hit_quality(self):
"""Berechnet die Qualität des Paddle-Treffers basierend auf der Trefferposition"""
ball_y = self.game.ball_pos[1]
paddle_y = self.game.paddle2_pos[1]
# Paddle-Höhe (angenommen 80 Pixel)
paddle_height = 80
# Berechne relative Position des Treffers (0 = oberer Rand, 1 = unterer Rand)
relative_hit_position = (ball_y - (paddle_y - paddle_height/2)) / paddle_height
# Begrenze auf gültigen Bereich [0, 1]
relative_hit_position = np.clip(relative_hit_position, 0, 1)
# Berechne Abstand vom Zentrum (0.5 = perfektes Zentrum)
distance_from_center = abs(relative_hit_position - 0.5)
# Qualitätsbewertung: Höchster Reward im Zentrum, exponentieller Abfall zu den Rändern
# Maximale Qualität: 0.3, minimale Qualität: 0.0
hit_quality = 0.3 * np.exp(-distance_from_center * 4)
return hit_quality
def _is_terminated(self):
"""Prüft ob das Spiel beendet ist"""
return (self.game.l_score >= self.max_score or
self.game.r_score >= self.max_score)
def _is_truncated(self):
"""Prüft ob die Episode abgebrochen werden soll"""
return self.step_count >= self.max_steps
def _update_opponent_paddle(self):
"""Aktualisiert das Gegner-Paddle basierend auf Modell oder scripted AI"""
if self.opponent_model is not None:
# Verwende geladenes Modell als Gegner
opponent_obs = self._get_opponent_observation()
try:
# Verwende deterministic=True für bessere Performance
action, _ = self.opponent_model.predict(opponent_obs, deterministic=True)
if action == 1: # Hoch
self.game.set_paddle1_velocity(-8)
elif action == 2: # Runter
self.game.set_paddle1_velocity(8)
else: # Nichts
self.game.set_paddle1_velocity(0)
except Exception as e:
# Nur bei ersten Fehlern ausgeben, um Spam zu vermeiden
if not hasattr(self, '_opponent_error_reported'):
print(f"Fehler bei Modell-Vorhersage: {e}")
self._opponent_error_reported = True
# Fallback zu scripted AI
self._update_scripted_ai_paddle()
else:
# Verwende scripted AI
self._update_scripted_ai_paddle()
def _update_scripted_ai_paddle(self):
"""Einfache AI für Paddle 1 (links)"""
# AI folgt dem Ball mit einer gewissen Verzögerung
ball_y = self.game.ball_pos[1]
paddle_y = self.game.paddle1_pos[1]
# Einfache Regel: Bewege Paddle in Richtung Ball
# Sehr kleine Toleranz für bessere Reaktion
if ball_y < paddle_y - 2:
self.game.set_paddle1_velocity(-6)
elif ball_y > paddle_y + 2:
self.game.set_paddle1_velocity(6)
else:
self.game.set_paddle1_velocity(0)
def render(self):
"""Rendering-Methode (wird von der UI übernommen)"""
pass
def close(self):
"""Schließt die Umgebung"""
pass