| from __future__ import annotations |
|
|
| import torch |
| from torch import nn |
|
|
|
|
| class DuelingQNetwork(nn.Module): |
| def __init__(self, observation_size: int = 17, actions: int = 4): |
| super().__init__() |
| self.encoder = nn.Sequential( |
| nn.Linear(observation_size, 128), |
| nn.LayerNorm(128), |
| nn.SiLU(), |
| nn.Linear(128, 128), |
| nn.SiLU(), |
| ) |
| self.value = nn.Sequential( |
| nn.Linear(128, 64), |
| nn.SiLU(), |
| nn.Linear(64, 1), |
| ) |
| self.advantage = nn.Sequential( |
| nn.Linear(128, 64), |
| nn.SiLU(), |
| nn.Linear(64, actions), |
| ) |
|
|
| def forward(self, observations: torch.Tensor) -> torch.Tensor: |
| features = self.encoder(observations) |
| value = self.value(features) |
| advantage = self.advantage(features) |
| return value + advantage - advantage.mean(dim=-1, keepdim=True) |
|
|
|
|
| class MicroPolicy(nn.Module): |
| def __init__(self, observation_size: int = 17, actions: int = 4): |
| super().__init__() |
| self.network = nn.Sequential( |
| nn.Linear(observation_size, 32), |
| nn.SiLU(), |
| nn.Linear(32, 32), |
| nn.SiLU(), |
| nn.Linear(32, actions), |
| ) |
|
|
| def forward(self, observations: torch.Tensor) -> torch.Tensor: |
| return self.network(observations) |
|
|