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

Class ConvStack

eval_code/recons/models/moge/model/modules.py:186–250  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

184
185
186class ConvStack(nn.Module):
187 def __init__(self,
188 dim_in: List[Optional[int]],
189 dim_res_blocks: List[int],
190 dim_out: List[Optional[int]],
191 resamplers: Union[Literal['pixel_shuffle', 'nearest', 'bilinear', 'conv_transpose', 'pixel_unshuffle', 'avg_pool', 'max_pool'], List],
192 dim_times_res_block_hidden: int = 1,
193 num_res_blocks: int = 1,
194 res_block_in_norm: Literal['layer_norm', 'group_norm' , 'instance_norm', 'none'] = 'layer_norm',
195 res_block_hidden_norm: Literal['layer_norm', 'group_norm' , 'instance_norm', 'none'] = 'group_norm',
196 activation: Literal['relu', 'leaky_relu', 'silu', 'elu'] = 'relu',
197 ):
198 super().__init__()
199 self.input_blocks = nn.ModuleList([
200 nn.Conv2d(dim_in_, dim_res_block_, kernel_size=1, stride=1, padding=0) if dim_in_ is not None else nn.Identity()
201 for dim_in_, dim_res_block_ in zip(dim_in if isinstance(dim_in, Sequence) else itertools.repeat(dim_in), dim_res_blocks)
202 ])
203 self.resamplers = nn.ModuleList([
204 Resampler(dim_prev, dim_succ, scale_factor=2, type_=resampler)
205 for i, (dim_prev, dim_succ, resampler) in enumerate(zip(
206 dim_res_blocks[:-1],
207 dim_res_blocks[1:],
208 resamplers if isinstance(resamplers, Sequence) else itertools.repeat(resamplers)
209 ))
210 ])
211 self.res_blocks = nn.ModuleList([
212 nn.Sequential(
213 *(
214 ResidualConvBlock(
215 dim_res_block_, dim_res_block_, dim_times_res_block_hidden * dim_res_block_,
216 activation=activation, in_norm=res_block_in_norm, hidden_norm=res_block_hidden_norm
217 ) for _ in range(num_res_blocks[i] if isinstance(num_res_blocks, list) else num_res_blocks)
218 )
219 ) for i, dim_res_block_ in enumerate(dim_res_blocks)
220 ])
221 self.output_blocks = nn.ModuleList([
222 nn.Conv2d(dim_res_block_, dim_out_, kernel_size=1, stride=1, padding=0) if dim_out_ is not None else nn.Identity()
223 for dim_out_, dim_res_block_ in zip(dim_out if isinstance(dim_out, Sequence) else itertools.repeat(dim_out), dim_res_blocks)
224 ])
225
226 def enable_gradient_checkpointing(self):
227 for i in range(len(self.resamplers)):
228 self.resamplers[i] = wrap_module_with_gradient_checkpointing(self.resamplers[i])
229 for i in range(len(self.res_blocks)):
230 for j in range(len(self.res_blocks[i])):
231 self.res_blocks[i][j] = wrap_module_with_gradient_checkpointing(self.res_blocks[i][j])
232
233 def forward(self, in_features: List[torch.Tensor]):
234 batch_shape = in_features[0].shape[:-3]
235 in_features = [x.reshape(-1, *x.shape[-3:]) for x in in_features]
236
237 out_features = []
238 for i in range(len(self.res_blocks)):
239 feature = self.input_blocks[i](in_features[i])
240 if i == 0:
241 x = feature
242 elif feature is not None:
243 x = x + feature

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected