ARotting's picture
Publish 1.7K parameter distilled navigation policy
a26d6ed verified
Raw
History Blame Contribute Delete
1.39 kB
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)