ZJW666/ProjectedGANCLC
0
1# Copyright (c) 2021, NVIDIA CORPORATION & AFFILIATES. All rights reserved.2#3# NVIDIA CORPORATION and its licensors retain all intellectual property4# and proprietary rights in and to this software, related documentation5# and any modifications thereto. Any use, reproduction, disclosure or6# distribution of this software and related documentation without an express7# license agreement from NVIDIA CORPORATION is strictly prohibited.8 9"""Converting legacy network pickle into the new format."""10 11import click12import pickle13import re14import copy15import numpy as np16import torch17import io18import dnnlib19import misc20 21#----------------------------------------------------------------------------22 23def load_network_pkl(f, force_fp16=False):24 data = _LegacyUnpickler(f).load()25 26 # Legacy TensorFlow pickle => convert.27 if isinstance(data, tuple) and len(data) == 3 and all(isinstance(net, _TFNetworkStub) for net in data):28 tf_G, tf_D, tf_Gs = data29 G = convert_tf_generator(tf_G)30 D = convert_tf_discriminator(tf_D)31 G_ema = convert_tf_generator(tf_Gs)32 data = dict(G=G, D=D, G_ema=G_ema)33 34 # Add missing fields.35 if 'training_set_kwargs' not in data:36 data['training_set_kwargs'] = None37 if 'augment_pipe' not in data:38 data['augment_pipe'] = None39 40 # Validate contents.41 assert isinstance(data['G'], torch.nn.Module)42 assert isinstance(data['D'], torch.nn.Module)43 assert isinstance(data['G_ema'], torch.nn.Module)44 assert isinstance(data['training_set_kwargs'], (dict, type(None)))45 assert isinstance(data['augment_pipe'], (torch.nn.Module, type(None)))46 47 # Force FP16.48 if force_fp16:49 for key in ['G', 'D', 'G_ema']:50 old = data[key]51 kwargs = copy.deepcopy(old.init_kwargs)52 fp16_kwargs = kwargs.get('synthesis_kwargs', kwargs)53 fp16_kwargs.num_fp16_res = 454 fp16_kwargs.conv_clamp = 25655 if kwargs != old.init_kwargs:56 new = type(old)(**kwargs).eval().requires_grad_(False)57 misc.copy_params_and_buffers(old, new, require_all=True)58 data[key] = new59 return data60 61#----------------------------------------------------------------------------62 63class _TFNetworkStub(dnnlib.EasyDict):64 pass65 66class _LegacyUnpickler(pickle.Unpickler):67 def find_class(self, module, name):68 # print(module,name)69 if module == '__builtin__':70 return 71 if module == 'dnnlib.tflib.network' and name == 'Network':72 return _TFNetworkStub73 if module == 'torch.storage' and name == '_load_from_bytes':74 return lambda b: torch.load(io.BytesIO(b), map_location='cpu')75 return super().find_class(module, name)76 77#----------------------------------------------------------------------------78 79def _collect_tf_params(tf_net):80 # pylint: disable=protected-access81 tf_params = dict()82 def recurse(prefix, tf_net):83 for name, value in tf_net.variables:84 tf_params[prefix + name] = value85 for name, comp in tf_net.components.items():86 recurse(prefix + name + '/', comp)87 recurse('', tf_net)88 return tf_params89 90#----------------------------------------------------------------------------91 92def _populate_module_params(module, *patterns):93 for name, tensor in misc.named_params_and_buffers(module):94 found = False95 value = None96 for pattern, value_fn in zip(patterns[0::2], patterns[1::2]):97 match = re.fullmatch(pattern, name)98 if match:99 found = True100 if value_fn is not None:101 value = value_fn(*match.groups())102 break103 try:104 assert found105 if value is not None:106 tensor.copy_(torch.from_numpy(np.array(value)))107 except:108 print(name, list(tensor.shape))109 raise110 111#----------------------------------------------------------------------------112 113def convert_tf_generator(tf_G):114 if tf_G.version < 4:115 raise ValueError('TensorFlow pickle version too low')116 117 # Collect kwargs.118 tf_kwargs = tf_G.static_kwargs119 known_kwargs = set()120 def kwarg(tf_name, default=None, none=None):121 known_kwargs.add(tf_name)122 val = tf_kwargs.get(tf_name, default)123 return val if val is not None else none124 125 # Convert kwargs.126 from pg_modules import networks_stylegan2127 network_class = networks_stylegan2.Generator128 kwargs = dnnlib.EasyDict(129 z_dim = kwarg('latent_size', 512),130 c_dim = kwarg('label_size', 0),131 w_dim = kwarg('dlatent_size', 512),132 img_resolution = kwarg('resolution', 1024),133 img_channels = kwarg('num_channels', 3),134 channel_base = kwarg('fmap_base', 16384) * 2,135 channel_max = kwarg('fmap_max', 512),136 num_fp16_res = kwarg('num_fp16_res', 0),137 conv_clamp = kwarg('conv_clamp', None),138 architecture = kwarg('architecture', 'skip'),139 resample_filter = kwarg('resample_kernel', [1,3,3,1]),140 use_noise = kwarg('use_noise', True),141 activation = kwarg('nonlinearity', 'lrelu'),142 mapping_kwargs = dnnlib.EasyDict(143 num_layers = kwarg('mapping_layers', 8),144 embed_features = kwarg('label_fmaps', None),145 layer_features = kwarg('mapping_fmaps', None),146 activation = kwarg('mapping_nonlinearity', 'lrelu'),147 lr_multiplier = kwarg('mapping_lrmul', 0.01),148 w_avg_beta = kwarg('w_avg_beta', 0.995, none=1),149 ),150 )151 152 # Check for unknown kwargs.153 kwarg('truncation_psi')154 kwarg('truncation_cutoff')155 kwarg('style_mixing_prob')156 kwarg('structure')157 kwarg('conditioning')158 kwarg('fused_modconv')159 unknown_kwargs = list(set(tf_kwargs.keys()) - known_kwargs)160 if len(unknown_kwargs) > 0:161 raise ValueError('Unknown TensorFlow kwarg', unknown_kwargs[0])162 163 # Collect params.164 tf_params = _collect_tf_params(tf_G)165 for name, value in list(tf_params.items()):166 match = re.fullmatch(r'ToRGB_lod(\d+)/(.*)', name)167 if match:168 r = kwargs.img_resolution // (2 ** int(match.group(1)))169 tf_params[f'{r}x{r}/ToRGB/{match.group(2)}'] = value170 kwargs.synthesis.kwargs.architecture = 'orig'171 #for name, value in tf_params.items(): print(f'{name:<50s}{list(value.shape)}')172 173 # Convert params.174 G = network_class(**kwargs).eval().requires_grad_(False)175 # pylint: disable=unnecessary-lambda176 # pylint: disable=f-string-without-interpolation177 _populate_module_params(G,178 r'mapping\.w_avg', lambda: tf_params[f'dlatent_avg'],179 r'mapping\.embed\.weight', lambda: tf_params[f'mapping/LabelEmbed/weight'].transpose(),180 r'mapping\.embed\.bias', lambda: tf_params[f'mapping/LabelEmbed/bias'],181 r'mapping\.fc(\d+)\.weight', lambda i: tf_params[f'mapping/Dense{i}/weight'].transpose(),182 r'mapping\.fc(\d+)\.bias', lambda i: tf_params[f'mapping/Dense{i}/bias'],183 r'synthesis\.b4\.const', lambda: tf_params[f'synthesis/4x4/Const/const'][0],184 r'synthesis\.b4\.conv1\.weight', lambda: tf_params[f'synthesis/4x4/Conv/weight'].transpose(3, 2, 0, 1),185 r'synthesis\.b4\.conv1\.bias', lambda: tf_params[f'synthesis/4x4/Conv/bias'],186 r'synthesis\.b4\.conv1\.noise_const', lambda: tf_params[f'synthesis/noise0'][0, 0],187 r'synthesis\.b4\.conv1\.noise_strength', lambda: tf_params[f'synthesis/4x4/Conv/noise_strength'],188 r'synthesis\.b4\.conv1\.affine\.weight', lambda: tf_params[f'synthesis/4x4/Conv/mod_weight'].transpose(),189 r'synthesis\.b4\.conv1\.affine\.bias', lambda: tf_params[f'synthesis/4x4/Conv/mod_bias'] + 1,190 r'synthesis\.b(\d+)\.conv0\.weight', lambda r: tf_params[f'synthesis/{r}x{r}/Conv0_up/weight'][::-1, ::-1].transpose(3, 2, 0, 1),191 r'synthesis\.b(\d+)\.conv0\.bias', lambda r: tf_params[f'synthesis/{r}x{r}/Conv0_up/bias'],192 r'synthesis\.b(\d+)\.conv0\.noise_const', lambda r: tf_params[f'synthesis/noise{int(np.log2(int(r)))*2-5}'][0, 0],193 r'synthesis\.b(\d+)\.conv0\.noise_strength', lambda r: tf_params[f'synthesis/{r}x{r}/Conv0_up/noise_strength'],194 r'synthesis\.b(\d+)\.conv0\.affine\.weight', lambda r: tf_params[f'synthesis/{r}x{r}/Conv0_up/mod_weight'].transpose(),195 r'synthesis\.b(\d+)\.conv0\.affine\.bias', lambda r: tf_params[f'synthesis/{r}x{r}/Conv0_up/mod_bias'] + 1,196 r'synthesis\.b(\d+)\.conv1\.weight', lambda r: tf_params[f'synthesis/{r}x{r}/Conv1/weight'].transpose(3, 2, 0, 1),197 r'synthesis\.b(\d+)\.conv1\.bias', lambda r: tf_params[f'synthesis/{r}x{r}/Conv1/bias'],198 r'synthesis\.b(\d+)\.conv1\.noise_const', lambda r: tf_params[f'synthesis/noise{int(np.log2(int(r)))*2-4}'][0, 0],199 r'synthesis\.b(\d+)\.conv1\.noise_strength', lambda r: tf_params[f'synthesis/{r}x{r}/Conv1/noise_strength'],200 r'synthesis\.b(\d+)\.conv1\.affine\.weight', lambda r: tf_params[f'synthesis/{r}x{r}/Conv1/mod_weight'].transpose(),201 r'synthesis\.b(\d+)\.conv1\.affine\.bias', lambda r: tf_params[f'synthesis/{r}x{r}/Conv1/mod_bias'] + 1,202 r'synthesis\.b(\d+)\.torgb\.weight', lambda r: tf_params[f'synthesis/{r}x{r}/ToRGB/weight'].transpose(3, 2, 0, 1),203 r'synthesis\.b(\d+)\.torgb\.bias', lambda r: tf_params[f'synthesis/{r}x{r}/ToRGB/bias'],204 r'synthesis\.b(\d+)\.torgb\.affine\.weight', lambda r: tf_params[f'synthesis/{r}x{r}/ToRGB/mod_weight'].transpose(),205 r'synthesis\.b(\d+)\.torgb\.affine\.bias', lambda r: tf_params[f'synthesis/{r}x{r}/ToRGB/mod_bias'] + 1,206 r'synthesis\.b(\d+)\.skip\.weight', lambda r: tf_params[f'synthesis/{r}x{r}/Skip/weight'][::-1, ::-1].transpose(3, 2, 0, 1),207 r'.*\.resample_filter', None,208 r'.*\.act_filter', None,209 )210 return G211 212#----------------------------------------------------------------------------213 214def convert_tf_discriminator(tf_D):215 if tf_D.version < 4:216 raise ValueError('TensorFlow pickle version too low')217 218 # Collect kwargs.219 tf_kwargs = tf_D.static_kwargs220 known_kwargs = set()221 def kwarg(tf_name, default=None):222 known_kwargs.add(tf_name)223 return tf_kwargs.get(tf_name, default)224 225 # Convert kwargs.226 kwargs = dnnlib.EasyDict(227 c_dim = kwarg('label_size', 0),228 img_resolution = kwarg('resolution', 1024),229 img_channels = kwarg('num_channels', 3),230 architecture = kwarg('architecture', 'resnet'),231 channel_base = kwarg('fmap_base', 16384) * 2,232 channel_max = kwarg('fmap_max', 512),233 num_fp16_res = kwarg('num_fp16_res', 0),234 conv_clamp = kwarg('conv_clamp', None),235 cmap_dim = kwarg('mapping_fmaps', None),236 block_kwargs = dnnlib.EasyDict(237 activation = kwarg('nonlinearity', 'lrelu'),238 resample_filter = kwarg('resample_kernel', [1,3,3,1]),239 freeze_layers = kwarg('freeze_layers', 0),240 ),241 mapping_kwargs = dnnlib.EasyDict(242 num_layers = kwarg('mapping_layers', 0),243 embed_features = kwarg('mapping_fmaps', None),244 layer_features = kwarg('mapping_fmaps', None),245 activation = kwarg('nonlinearity', 'lrelu'),246 lr_multiplier = kwarg('mapping_lrmul', 0.1),247 ),248 epilogue_kwargs = dnnlib.EasyDict(249 mbstd_group_size = kwarg('mbstd_group_size', None),250 mbstd_num_channels = kwarg('mbstd_num_features', 1),251 activation = kwarg('nonlinearity', 'lrelu'),252 ),253 )254 255 # Check for unknown kwargs.256 kwarg('structure')257 kwarg('conditioning')258 unknown_kwargs = list(set(tf_kwargs.keys()) - known_kwargs)259 if len(unknown_kwargs) > 0:260 raise ValueError('Unknown TensorFlow kwarg', unknown_kwargs[0])261 262 # Collect params.263 tf_params = _collect_tf_params(tf_D)264 for name, value in list(tf_params.items()):265 match = re.fullmatch(r'FromRGB_lod(\d+)/(.*)', name)266 if match:267 r = kwargs.img_resolution // (2 ** int(match.group(1)))268 tf_params[f'{r}x{r}/FromRGB/{match.group(2)}'] = value269 kwargs.architecture = 'orig'270 #for name, value in tf_params.items(): print(f'{name:<50s}{list(value.shape)}')271 272 # Convert params.273 #from pg_modules import networks_stylegan2274 from pg_modules.discriminator import ProjectedDiscriminator275 276 D = ProjectedDiscriminator(**kwargs).eval().requires_grad_(False)277 # pylint: disable=unnecessary-lambda278 # pylint: disable=f-string-without-interpolation279 _populate_module_params(D,280 r'b(\d+)\.fromrgb\.weight', lambda r: tf_params[f'{r}x{r}/FromRGB/weight'].transpose(3, 2, 0, 1),281 r'b(\d+)\.fromrgb\.bias', lambda r: tf_params[f'{r}x{r}/FromRGB/bias'],282 r'b(\d+)\.conv(\d+)\.weight', lambda r, i: tf_params[f'{r}x{r}/Conv{i}{["","_down"][int(i)]}/weight'].transpose(3, 2, 0, 1),283 r'b(\d+)\.conv(\d+)\.bias', lambda r, i: tf_params[f'{r}x{r}/Conv{i}{["","_down"][int(i)]}/bias'],284 r'b(\d+)\.skip\.weight', lambda r: tf_params[f'{r}x{r}/Skip/weight'].transpose(3, 2, 0, 1),285 r'mapping\.embed\.weight', lambda: tf_params[f'LabelEmbed/weight'].transpose(),286 r'mapping\.embed\.bias', lambda: tf_params[f'LabelEmbed/bias'],287 r'mapping\.fc(\d+)\.weight', lambda i: tf_params[f'Mapping{i}/weight'].transpose(),288 r'mapping\.fc(\d+)\.bias', lambda i: tf_params[f'Mapping{i}/bias'],289 r'b4\.conv\.weight', lambda: tf_params[f'4x4/Conv/weight'].transpose(3, 2, 0, 1),290 r'b4\.conv\.bias', lambda: tf_params[f'4x4/Conv/bias'],291 r'b4\.fc\.weight', lambda: tf_params[f'4x4/Dense0/weight'].transpose(),292 r'b4\.fc\.bias', lambda: tf_params[f'4x4/Dense0/bias'],293 r'b4\.out\.weight', lambda: tf_params[f'Output/weight'].transpose(),294 r'b4\.out\.bias', lambda: tf_params[f'Output/bias'],295 r'.*\.resample_filter', None,296 )297 return D298 299#----------------------------------------------------------------------------300 301@click.command()302@click.option('--source', help='Input pickle', required=True, metavar='PATH')303@click.option('--dest', help='Output pickle', required=True, metavar='PATH')304@click.option('--force-fp16', help='Force the networks to use FP16', type=bool, default=False, metavar='BOOL', show_default=True)305def convert_network_pickle(source, dest, force_fp16):306 """Convert legacy network pickle into the native PyTorch format.307 308 The tool is able to load the main network configurations exported using the TensorFlow version of StyleGAN2 or StyleGAN2-ADA.309 It does not support e.g. StyleGAN2-ADA comparison methods, StyleGAN2 configs A-D, or StyleGAN1 networks.310 311 Example:312 313 \b314 python legacy.py \\315 --source=https://nvlabs-fi-cdn.nvidia.com/stylegan2/networks/stylegan2-cat-config-f.pkl \\316 --dest=stylegan2-cat-config-f.pkl317 """318 print(f'Loading "{source}"...')319 with dnnlib.util.open_url(source) as f:320 data = load_network_pkl(f, force_fp16=force_fp16)321 print(f'Saving "{dest}"...')322 with open(dest, 'wb') as f:323 pickle.dump(data, f)324 print('Done.')325 326#----------------------------------------------------------------------------327 328if __name__ == "__main__":329 convert_network_pickle() # pylint: disable=no-value-for-parameter330 331#----------------------------------------------------------------------------332 