| 190 | |
| 191 | @persistence.persistent_class |
| 192 | class 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. |