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
()
| 149 | |
| 150 | |
| 151 | def 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 |