Current section
Files
Jump to
Current section
Files
lib/trainers/trainer_sarsa.ex
defmodule Gyx.Trainers.TrainerSarsa do
@moduledoc """
This module describes an entire training process,
tune accordingly to your particular environment and agent
"""
use GenServer
alias Gyx.Core.Exp
require Logger
@enforce_keys [:environment, :agent]
defstruct environment: nil, agent: nil, trajectory: nil, rewards: nil
@type t :: %__MODULE__{
environment: any(),
agent: any(),
trajectory: list(Exp),
rewards: list(number())
}
@env_module Gyx.Environments.Gym
@agent Gyx.Agents.SARSA.Agent
def init(_) do
{:ok,
%__MODULE__{
environment: @env_module,
agent: @agent,
trajectory: [],
rewards: []
}}
end
def start_link(_, opts) do
GenServer.start_link(__MODULE__, [], opts)
end
def train(episodes) do
GenServer.call(__MODULE__, {:train, episodes})
end
def handle_call({:train, episodes}, _from, t = %__MODULE__{}) do
{:reply, trainer(t, episodes), t}
end
@spec trainer(__MODULE__.t(), integer) :: __MODULE__.t()
defp trainer(t, 0), do: t
defp trainer(t, num_episodes) do
t.environment.reset()
t
|> initialize_trajectory()
#|> IO.inspect(label: "Trajectory initialized")
|> run_episode(false)
#|> IO.inspect(label: "Episode finished")
|> log_stats()
|> trainer(num_episodes - 1)
end
defp run_episode(t = %__MODULE__{}, true), do: t
defp run_episode(t = %__MODULE__{}, false) do
exp =
%Exp{done: done, state: s, action: a, reward: r, next_state: ss} =
%{
observation: t.environment.observe(),
action_space: t.environment.get_state().action_space
}
|> t.agent.act_epsilon_greedy()
|> t.environment.step
aa =
t.agent.act_epsilon_greedy(%{
observation: ss,
action_space: t.environment.get_state().action_space
})
t.agent.td_learn({s, a, r, ss, aa})
t = %{t | trajectory: [exp | t.trajectory]}
run_episode(t, done)
end
defp initialize_trajectory(t), do: %{t | trajectory: []}
defp log_stats(t) do
reward_sum = t.trajectory |> Enum.map(& &1.reward) |> Enum.sum()
t = %{t | rewards: [reward_sum | t.rewards]}
k = 100
Logger.info("Reward: " <> to_string((t.rewards |> Enum.take(k) |> Enum.sum()) / k))
#Gyx.Qstorage.QGenServer.print_q_matrix()
t
end
end