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)