Add basic setup.py
This commit is contained in:
42
interimit/policies/policy.py
Normal file
42
interimit/policies/policy.py
Normal file
@@ -0,0 +1,42 @@
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from interimit.nets.deepsets import DeepSetsModule, Phi
|
||||
|
||||
class Policy:
|
||||
pass
|
||||
|
||||
class DeepSetsPolicy(Policy, nn.Module):
|
||||
def __init__(self, ego_config, dynamic_config, path_config, head_config):
|
||||
"""
|
||||
Args:
|
||||
ego_config (dict): dictionary for configuring the ego network
|
||||
dynamic_config (dict): dictionary for configuring the dynamic input (deepsets) network
|
||||
path_config (dict): dictionary for configuring the path network
|
||||
head_config (dict): dictionary for configuring the common head network
|
||||
"""
|
||||
super(DeepSetsPolicy, self).__init__()
|
||||
self.ego_net = Phi.from_config(ego_config)
|
||||
self.deepsets = DeepSetsModule.from_config(dynamic_config)
|
||||
self.path_net = Phi.from_config(path_config)
|
||||
cat_dim = self.ego_net.output_dim + self.deepsets.output_dim + self.path_net.output_dim
|
||||
# head has number of concatenated features as input
|
||||
head_config["input_dim"] = cat_dim
|
||||
self.head = Phi.from_config(head_config)
|
||||
|
||||
def forward(self, ego_state, relative_states, path):
|
||||
"""
|
||||
Args:
|
||||
ego_state (torch.tensor): (ns,) state of ego vehicle
|
||||
relative_states (torch.tensor): (nv, ns) relative states of other vehicles (dynamic size)
|
||||
path (torch.tensor): (path_length, 2) coordinates (x,y) of path
|
||||
Returns:
|
||||
x (torch.tensor): (head_output_dim,) output of common head network
|
||||
"""
|
||||
x_ego = self.ego_net(ego_state)
|
||||
x_relative = self.deepsets(relative_states)
|
||||
x_path = self.path_net(path.flatten())
|
||||
x = torch.cat([x_ego, x_relative, x_path])
|
||||
x = self.head(x)
|
||||
return x
|
||||
Reference in New Issue
Block a user