MCPcopy Create free account
hub / github.com/OpenGVLab/DragGAN / MappingNetwork

Class MappingNetwork

draggan/stylegan2/training/networks.py:192–270  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

190
191@persistence.persistent_class
192class MappingNetwork(torch.nn.Module):
193 def __init__(self,
194 z_dim, # Input latent (Z) dimensionality, 0 = no latent.
195 c_dim, # Conditioning label (C) dimensionality, 0 = no label.
196 w_dim, # Intermediate latent (W) dimensionality.
197 num_ws, # Number of intermediate latents to output, None = do not broadcast.
198 num_layers = 8, # Number of mapping layers.
199 embed_features = None, # Label embedding dimensionality, None = same as w_dim.
200 layer_features = None, # Number of intermediate features in the mapping layers, None = same as w_dim.
201 activation = 'lrelu', # Activation function: 'relu', 'lrelu', etc.
202 lr_multiplier = 0.01, # Learning rate multiplier for the mapping layers.
203 w_avg_beta = 0.998, # Decay for tracking the moving average of W during training, None = do not track.
204 ):
205 super().__init__()
206 self.z_dim = z_dim
207 self.c_dim = c_dim
208 self.w_dim = w_dim
209 self.num_ws = num_ws
210 self.num_layers = num_layers
211 self.w_avg_beta = w_avg_beta
212
213 if embed_features is None:
214 embed_features = w_dim
215 if c_dim == 0:
216 embed_features = 0
217 if layer_features is None:
218 layer_features = w_dim
219 features_list = [z_dim + embed_features] + [layer_features] * (num_layers - 1) + [w_dim]
220
221 if c_dim > 0:
222 self.embed = FullyConnectedLayer(c_dim, embed_features)
223 for idx in range(num_layers):
224 in_features = features_list[idx]
225 out_features = features_list[idx + 1]
226 layer = FullyConnectedLayer(in_features, out_features, activation=activation, lr_multiplier=lr_multiplier)
227 setattr(self, f'fc{idx}', layer)
228
229 if num_ws is not None and w_avg_beta is not None:
230 self.register_buffer('w_avg', torch.zeros([w_dim]))
231
232 def forward(self, z, c, truncation_psi=1, truncation_cutoff=None, update_emas=False):
233 # Embed, normalize, and concat inputs.
234 x = None
235 with torch.autograd.profiler.record_function('input'):
236 if self.z_dim > 0:
237 misc.assert_shape(z, [None, self.z_dim])
238 x = normalize_2nd_moment(z.to(torch.float32))
239 if self.c_dim > 0:
240 misc.assert_shape(c, [None, self.c_dim])
241 y = normalize_2nd_moment(self.embed(c.to(torch.float32)))
242 x = torch.cat([x, y], dim=1) if x is not None else y
243
244 # Main layers.
245 for idx in range(self.num_layers):
246 layer = getattr(self, f'fc{idx}')
247 x = layer(x)
248
249 # Update moving average of W.

Callers 2

__init__Method · 0.70
__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected