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

Method forward

src/models/pluto/pluto_model.py:127–233  ·  view source on GitHub ↗
(self, data)

Source from the content-addressed store, hash-verified

125 nn.init.normal_(m.weight, mean=0.0, std=0.02)
126
127 def forward(self, data):
128 agent_pos = data["agent"]["position"][:, :, self.history_steps - 1]
129 agent_heading = data["agent"]["heading"][:, :, self.history_steps - 1]
130 agent_mask = data["agent"]["valid_mask"][:, :, : self.history_steps]
131 polygon_center = data["map"]["polygon_center"]
132 polygon_mask = data["map"]["valid_mask"]
133
134 bs, A = agent_pos.shape[0:2]
135
136 position = torch.cat([agent_pos, polygon_center[..., :2]], dim=1)
137 angle = torch.cat([agent_heading, polygon_center[..., 2]], dim=1)
138 angle = (angle + math.pi) % (2 * math.pi) - math.pi
139 pos = torch.cat([position, angle.unsqueeze(-1)], dim=-1)
140
141 agent_key_padding = ~(agent_mask.any(-1))
142 polygon_key_padding = ~(polygon_mask.any(-1))
143 key_padding_mask = torch.cat([agent_key_padding, polygon_key_padding], dim=-1)
144
145 x_agent = self.agent_encoder(data)
146 x_polygon = self.map_encoder(data)
147 x_static, static_pos, static_key_padding = self.static_objects_encoder(data)
148
149 x = torch.cat([x_agent, x_polygon, x_static], dim=1)
150
151 pos = torch.cat([pos, static_pos], dim=1)
152 pos_embed = self.pos_emb(pos)
153
154 key_padding_mask = torch.cat([key_padding_mask, static_key_padding], dim=-1)
155 x = x + pos_embed
156
157 for blk in self.encoder_blocks:
158 x = blk(x, key_padding_mask=key_padding_mask, return_attn_weights=False)
159 x = self.norm(x)
160
161 prediction = self.agent_predictor(x[:, 1:A])
162
163 ref_line_available = data["reference_line"]["position"].shape[1] > 0
164
165 if ref_line_available:
166 trajectory, probability = self.planning_decoder(
167 data, {"enc_emb": x, "enc_key_padding_mask": key_padding_mask}
168 )
169 else:
170 trajectory, probability = None, None
171
172 out = {
173 "trajectory": trajectory,
174 "probability": probability, # (bs, R, M)
175 "prediction": prediction, # (bs, A-1, T, 2)
176 }
177
178 if self.use_hidden_proj:
179 out["hidden"] = self.hidden_proj(x[:, 0])
180
181 if self.ref_free_traj:
182 ref_free_traj = self.ref_free_decoder(x[:, 0]).reshape(
183 bs, self.future_steps, 4
184 )

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected