MCPcopy Create free account
hub / github.com/InternRobotics/G2VLM / decode

Method decode

eval_code/recons/models/pi3/models/pi3.py:132–171  ·  view source on GitHub ↗
(self, hidden, N, H, W)

Source from the content-addressed store, hash-verified

130
131
132 def decode(self, hidden, N, H, W):
133 BN, hw, _ = hidden.shape
134 B = BN // N
135
136 final_output = []
137
138 hidden = hidden.reshape(B*N, hw, -1)
139
140 register_token = self.register_token.repeat(B, N, 1, 1).reshape(B*N, *self.register_token.shape[-2:])
141
142 # Concatenate special tokens with patch tokens
143 hidden = torch.cat([register_token, hidden], dim=1)
144 hw = hidden.shape[1]
145
146 if self.pos_type.startswith('rope'):
147 pos = self.position_getter(B * N, H//self.patch_size, W//self.patch_size, hidden.device)
148
149 if self.patch_start_idx > 0:
150 # do not use position embedding for special tokens (camera and register tokens)
151 # so set pos to 0 for the special tokens
152 pos = pos + 1
153 pos_special = torch.zeros(B * N, self.patch_start_idx, 2).to(hidden.device).to(pos.dtype)
154 pos = torch.cat([pos_special, pos], dim=1)
155
156 for i in range(len(self.decoder)):
157 blk = self.decoder[i]
158
159 if i % 2 == 0:
160 pos = pos.reshape(B*N, hw, -1)
161 hidden = hidden.reshape(B*N, hw, -1)
162 else:
163 pos = pos.reshape(B, N*hw, -1)
164 hidden = hidden.reshape(B, N*hw, -1)
165
166 hidden = blk(hidden, xpos=pos)
167
168 if i+1 in [len(self.decoder)-1, len(self.decoder)]:
169 final_output.append(hidden.reshape(B*N, hw, -1))
170
171 return torch.cat([final_output[0], final_output[1]], dim=-1), pos.reshape(B*N, hw, -1)
172
173 def forward(self, imgs):
174 imgs = (imgs - self.image_mean) / self.image_std

Callers 4

forwardMethod · 0.95
openMethod · 0.45
_runFunction · 0.45
_runFunction · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected