MCPcopy Create free account
hub / github.com/Rex-sys-hk/PlanScope / AgentEncoder

Class AgentEncoder

src/models/pluto/modules/agent_encoder.py:8–94  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6
7
8class AgentEncoder(nn.Module):
9 def __init__(
10 self,
11 state_channel=6,
12 history_channel=9,
13 dim=128,
14 hist_steps=21,
15 use_ego_history=False,
16 drop_path=0.2,
17 state_attn_encoder=True,
18 state_dropout=0.75,
19 ) -> None:
20 super().__init__()
21 self.dim = dim
22 self.state_channel = state_channel
23 self.use_ego_history = use_ego_history
24 self.hist_steps = hist_steps
25 self.state_attn_encoder = state_attn_encoder
26
27 self.history_encoder = NATSequenceEncoder(
28 in_chans=history_channel, embed_dim=dim // 4, drop_path_rate=drop_path
29 )
30
31 if not use_ego_history:
32 if not self.state_attn_encoder:
33 self.ego_state_emb = build_mlp(state_channel, [dim] * 2, norm="bn")
34 else:
35 self.ego_state_emb = StateAttentionEncoder(
36 state_channel, dim, state_dropout
37 )
38
39 self.type_emb = nn.Embedding(4, dim)
40
41 @staticmethod
42 def to_vector(feat, valid_mask):
43 vec_mask = valid_mask[..., :-1] & valid_mask[..., 1:]
44
45 while len(vec_mask.shape) < len(feat.shape):
46 vec_mask = vec_mask.unsqueeze(-1)
47
48 return torch.where(
49 vec_mask,
50 feat[:, :, 1:, ...] - feat[:, :, :-1, ...],
51 torch.zeros_like(feat[:, :, 1:, ...]),
52 )
53
54 def forward(self, data):
55 T = self.hist_steps
56
57 position = data["agent"]["position"][:, :, :T]
58 heading = data["agent"]["heading"][:, :, :T]
59 velocity = data["agent"]["velocity"][:, :, :T]
60 shape = data["agent"]["shape"][:, :, :T]
61 category = data["agent"]["category"].long()
62 valid_mask = data["agent"]["valid_mask"][:, :, :T]
63
64 heading_vec = self.to_vector(heading, valid_mask)
65 valid_mask_vec = valid_mask[..., 1:] & valid_mask[..., :-1]

Callers 2

__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected