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

Method forward

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

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,
    )

Source from the content-addressed store, hash-verified

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
89class PlanningDecoder(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected