MCPcopy Create free account
hub / github.com/Robbyant/lingbot-depth / ConvStack

Class ConvStack

mdm/model/modules_decoder.py:126–185  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

124
125
126class ConvStack(nn.Module):
127 def __init__(self,
128 dim_in: List[Optional[int]],
129 dim_res_blocks: List[int],
130 dim_out: List[Optional[int]],
131 resamplers: Union[Literal['pixel_shuffle', 'nearest', 'bilinear', 'conv_transpose', 'pixel_unshuffle', 'avg_pool', 'max_pool'], List],
132 dim_times_res_block_hidden: int = 1,
133 num_res_blocks: int = 1,
134 res_block_in_norm: Literal['layer_norm', 'group_norm' , 'instance_norm', 'none'] = 'layer_norm',
135 res_block_hidden_norm: Literal['layer_norm', 'group_norm' , 'instance_norm', 'none'] = 'group_norm',
136 activation: Literal['relu', 'leaky_relu', 'silu', 'elu'] = 'relu',
137 ):
138 super().__init__()
139 self.input_blocks = nn.ModuleList([
140 nn.Conv2d(dim_in_, dim_res_block_, kernel_size=1, stride=1, padding=0) if dim_in_ is not None else nn.Identity()
141 for dim_in_, dim_res_block_ in zip(dim_in if isinstance(dim_in, Sequence) else itertools.repeat(dim_in), dim_res_blocks)
142 ])
143 self.resamplers = nn.ModuleList([
144 Resampler(dim_prev, dim_succ, scale_factor=2, type_=resampler)
145 for i, (dim_prev, dim_succ, resampler) in enumerate(zip(
146 dim_res_blocks[:-1],
147 dim_res_blocks[1:],
148 resamplers if isinstance(resamplers, Sequence) else itertools.repeat(resamplers)
149 ))
150 ])
151 self.res_blocks = nn.ModuleList([
152 nn.Sequential(
153 *(
154 ResidualConvBlock(
155 dim_res_block_, dim_res_block_, dim_times_res_block_hidden * dim_res_block_,
156 activation=activation, in_norm=res_block_in_norm, hidden_norm=res_block_hidden_norm
157 ) for _ in range(num_res_blocks[i] if isinstance(num_res_blocks, list) else num_res_blocks)
158 )
159 ) for i, dim_res_block_ in enumerate(dim_res_blocks)
160 ])
161 self.output_blocks = nn.ModuleList([
162 nn.Conv2d(dim_res_block_, dim_out_, kernel_size=1, stride=1, padding=0) if dim_out_ is not None else nn.Identity()
163 for dim_out_, dim_res_block_ in zip(dim_out if isinstance(dim_out, Sequence) else itertools.repeat(dim_out), dim_res_blocks)
164 ])
165
166 def enable_gradient_checkpointing(self):
167 for i in range(len(self.resamplers)):
168 self.resamplers[i] = wrap_module_with_gradient_checkpointing(self.resamplers[i])
169 for i in range(len(self.res_blocks)):
170 for j in range(len(self.res_blocks[i])):
171 self.res_blocks[i][j] = wrap_module_with_gradient_checkpointing(self.res_blocks[i][j])
172
173 def forward(self, in_features: List[torch.Tensor]):
174 out_features = []
175 for i in range(len(self.res_blocks)):
176 feature = self.input_blocks[i](in_features[i])
177 if i == 0:
178 x = feature
179 elif feature is not None:
180 x = x + feature
181 x = self.res_blocks[i](x)
182 out_features.append(self.output_blocks[i](x))
183 if i < len(self.res_blocks) - 1:

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected