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

Class PlanningDecoder

src/models/pluto/modules/planning_decoder.py:89–188  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

87
88
89class PlanningDecoder(nn.Module):
90 def __init__(
91 self,
92 num_mode,
93 decoder_depth,
94 dim,
95 num_heads,
96 mlp_ratio,
97 dropout,
98 future_steps,
99 yaw_constraint=False,
100 cat_x=False,
101 ) -> None:
102 super().__init__()
103
104 self.num_mode = num_mode
105 self.future_steps = future_steps
106 self.yaw_constraint = yaw_constraint
107 self.cat_x = cat_x
108
109 self.decoder_blocks = nn.ModuleList(
110 [
111 DecoderLayer(dim, num_heads, mlp_ratio, dropout)
112 for _ in range(decoder_depth)
113 ]
114 )
115
116 self.r_pos_emb = FourierEmbedding(3, dim, 64)
117 self.r_encoder = PointsEncoder(6, dim)
118
119 self.q_proj = nn.Linear(2 * dim, dim)
120
121 self.m_emb = nn.Parameter(torch.Tensor(1, 1, num_mode, dim))
122 self.m_pos = nn.Parameter(torch.Tensor(1, num_mode, dim))
123
124 if self.cat_x:
125 self.cat_x_proj = nn.Linear(2 * dim, dim)
126
127 self.loc_head = MLPLayer(dim, 2 * dim, self.future_steps * 2)
128 self.yaw_head = MLPLayer(dim, 2 * dim, self.future_steps * 2)
129 self.vel_head = MLPLayer(dim, 2 * dim, self.future_steps * 2)
130 self.pi_head = MLPLayer(dim, dim, 1)
131
132 nn.init.normal_(self.m_emb, mean=0.0, std=0.01)
133 nn.init.normal_(self.m_pos, mean=0.0, std=0.01)
134
135 def forward(self, data, enc_data):
136 enc_emb = enc_data["enc_emb"]
137 enc_key_padding_mask = enc_data["enc_key_padding_mask"]
138
139 r_position = data["reference_line"]["position"]
140 r_vector = data["reference_line"]["vector"]
141 r_orientation = data["reference_line"]["orientation"]
142 r_valid_mask = data["reference_line"]["valid_mask"]
143 r_key_padding_mask = ~r_valid_mask.any(-1)
144
145 r_feature = torch.cat(
146 [

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected