tgt: (bs, R, M, dim) tgt_key_padding_mask: (bs, R)
(
self,
tgt,
memory,
tgt_key_padding_mask: Optional[Tensor] = None,
memory_key_padding_mask: Optional[Tensor] = None,
m_pos: Optional[Tensor] = None,
)
| 40 | self.dropout3 = nn.Dropout(dropout) |
| 41 | |
| 42 | def forward( |
| 43 | self, |
| 44 | tgt, |
| 45 | memory, |
| 46 | tgt_key_padding_mask: Optional[Tensor] = None, |
| 47 | memory_key_padding_mask: Optional[Tensor] = None, |
| 48 | m_pos: Optional[Tensor] = None, |
| 49 | ): |
| 50 | """ |
| 51 | tgt: (bs, R, M, dim) |
| 52 | tgt_key_padding_mask: (bs, R) |
| 53 | """ |
| 54 | bs, R, M, D = tgt.shape |
| 55 | |
| 56 | tgt = tgt.transpose(1, 2).reshape(bs * M, R, D) |
| 57 | tgt2 = self.norm1(tgt) |
| 58 | tgt2 = self.r2r_attn( |
| 59 | tgt2, tgt2, tgt2, key_padding_mask=tgt_key_padding_mask.repeat(M, 1) |
| 60 | )[0] |
| 61 | tgt = tgt + self.dropout1(tgt2) |
| 62 | |
| 63 | tgt_tmp = tgt.reshape(bs, M, R, D).transpose(1, 2).reshape(bs * R, M, D) |
| 64 | tgt_valid_mask = ~tgt_key_padding_mask.reshape(-1) |
| 65 | tgt_valid = tgt_tmp[tgt_valid_mask] |
| 66 | tgt2_valid = self.norm2(tgt_valid) |
| 67 | tgt2_valid, _ = self.m2m_attn( |
| 68 | tgt2_valid + m_pos, tgt2_valid + m_pos, tgt2_valid |
| 69 | ) |
| 70 | tgt_valid = tgt_valid + self.dropout2(tgt2_valid) |
| 71 | tgt = torch.zeros_like(tgt_tmp) |
| 72 | tgt[tgt_valid_mask] = tgt_valid |
| 73 | |
| 74 | tgt = tgt.reshape(bs, R, M, D).view(bs, R * M, D) |
| 75 | tgt2 = self.norm3(tgt) |
| 76 | tgt2 = self.cross_attn( |
| 77 | tgt2, memory, memory, key_padding_mask=memory_key_padding_mask |
| 78 | )[0] |
| 79 | |
| 80 | tgt = tgt + self.dropout2(tgt2) |
| 81 | tgt2 = self.norm4(tgt) |
| 82 | tgt2 = self.ffn(tgt2) |
| 83 | tgt = tgt + self.dropout3(tgt2) |
| 84 | tgt = tgt.reshape(bs, R, M, D) |
| 85 | |
| 86 | return tgt |
| 87 | |
| 88 | |
| 89 | class PlanningDecoder(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected