CoolFace
Apppublic

ZJW666/ProjectedGANCLC

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
legacy.py332 linesDownload Raw Back to root
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