MCPcopy Create free account
hub / github.com/PABannier/sam3.cpp / install_encoder_patch

Function install_encoder_patch

tests/dump_fenc_from_package.py:151–238  ·  view source on GitHub ↗

Monkey-patch TransformerEncoder.forward (the base class used by TransformerEncoderFusion) to dump inputs and per-layer outputs. TransformerEncoderFusion.forward() does: 1. Reshape src from seq-first to NCHW 2. (Optional pooled text — disabled with add_pooled_text_to_img_fea

()

Source from the content-addressed store, hash-verified

149
150
151def install_encoder_patch():
152 """
153 Monkey-patch TransformerEncoder.forward (the base class used by
154 TransformerEncoderFusion) to dump inputs and per-layer outputs.
155
156 TransformerEncoderFusion.forward() does:
157 1. Reshape src from seq-first to NCHW
158 2. (Optional pooled text — disabled with add_pooled_text_to_img_feat=False)
159 3. Call super().forward(src_NCHW, ..., prompt=prompt.transpose(0,1), ...)
160
161 super().forward() (TransformerEncoder.forward) does:
162 1. _prepare_multilevel_features → src_flatten [B, HW, D], lvl_pos_embed_flatten [B, HW, D]
163 2. Loop through 6 encoder layers
164
165 We patch TransformerEncoder.forward to capture the flattened tensors.
166 """
167 from sam3.model.encoder import TransformerEncoder
168 from sam3.model.act_ckpt_utils import activation_ckpt_wrapper
169
170 original_forward = TransformerEncoder.forward
171
172 def patched_forward(self, src, src_key_padding_masks=None, pos=None,
173 prompt=None, prompt_key_padding_mask=None,
174 encoder_extra_kwargs=None):
175 # --- _prepare_multilevel_features (same as original) ---
176 (
177 src_flatten,
178 key_padding_masks_flatten,
179 lvl_pos_embed_flatten,
180 level_start_index,
181 valid_ratios,
182 spatial_shapes,
183 ) = self._prepare_multilevel_features(src, src_key_padding_masks, pos)
184
185 reference_points = self.get_reference_points(
186 spatial_shapes, valid_ratios, device=src_flatten.device
187 )
188
189 # ═══ DUMP INPUTS ═══
190 # src_flatten: [B, HW, D] = [1, 5184, 256] batch-first image features
191 captured["fenc_input_tgt"] = src_flatten.detach().clone()
192 # lvl_pos_embed_flatten: [B, HW, D] = [1, 5184, 256] positional encoding
193 captured["fenc_input_pos"] = lvl_pos_embed_flatten.detach().clone()
194 # prompt: [B, T, D] = [1, 32, 256] batch-first text tokens
195 if prompt is not None:
196 captured["fenc_input_prompt"] = prompt.detach().clone()
197 # prompt_key_padding_mask: [B, T] = [1, 32] boolean (True=padding)
198 if prompt_key_padding_mask is not None:
199 captured["fenc_input_prompt_mask"] = prompt_key_padding_mask.detach().clone()
200
201 # --- Layer loop (same as original) ---
202 output = src_flatten
203 for layer_idx, layer in enumerate(self.layers):
204 layer_kwargs = {}
205
206 assert hasattr(layer, 'forward_pre') or hasattr(layer, 'forward_post')
207 layer_kwargs["memory"] = prompt
208 layer_kwargs["memory_key_padding_mask"] = prompt_key_padding_mask

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected