File size: 1,501 Bytes
0813b4d
46dc1ab
0813b4d
 
46dc1ab
0813b4d
 
 
46dc1ab
 
0813b4d
 
 
 
 
 
 
 
46dc1ab
0813b4d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
46dc1ab
0813b4d
b3e3848
0813b4d
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
from stable_baselines3 import A2C
from stable_baselines3 import PPO
from tetris_gym.wrappers.observation import ExtendedObservationWrapper
from agent.MyWrapper import MyWrapper
from stable_baselines3 import DQN

class Agent:
    """
    Az Agent tanítása során PPO-t használtam 2M steppel. Reward function ként az alábbi cikkhez hasonlót használtam.
    http://cs231n.stanford.edu/reports/2016/pdfs/121_Report.pdf
    """

    def __init__(self, env) -> None:
        """
        A konsztruktorban van lehetőség például a modell betöltésére
        vagy a környezet wrapper-ekkel való kiterjesztésére.
        """
        
        self.model = PPO.load("agent/A2CPaper")
        
        # A környezetet kiterjeszthetjük wrapper-ek segítségével.
        # Ha tanításkor modosítottuk a megfigyeléseket,
        # akkor azt a módosítást kiértékeléskor is meg kell adnunk.
        # self.observation_wrapper = ExtendedObservationWrapper(env)
        self.observation_wrapper = MyWrapper(env)


    def act(self, observation):
        """
        A megfigyelés alapján visszaadja a következő lépést.
        Ez a függvény fogja megadni az ágens működését.
        """

        # Ha tanításkor modosítottuk a megfigyeléseket,
        # akkor azt a módosítást kiértékeléskor is meg kell adnunk.
        # extended_obsetvation = self.observation_wrapper.observation(observation)

        return self.model.predict(observation, deterministic=True)