CoolFace
Apppublic

Anonymous-123/ImageNet-Editing

sourceHugging Facecreativeml-openrail-mupdated 4y agoView on Hugging Face
1likes
script_util.py453 linesDownload Raw Back to guided_diffusion
1import argparse2import inspect3 4from . import gaussian_diffusion as gd5from .respace import SpacedDiffusion, space_timesteps6from .unet import SuperResModel, UNetModel, EncoderUNetModel7 8NUM_CLASSES = 10009 10 11def diffusion_defaults():12    """13    Defaults for image and classifier training.14    """15    return dict(16        learn_sigma=False,17        diffusion_steps=1000,18        noise_schedule="linear",19        timestep_respacing="",20        use_kl=False,21        predict_xstart=False,22        rescale_timesteps=False,23        rescale_learned_sigmas=False,24    )25 26 27def classifier_defaults():28    """29    Defaults for classifier models.30    """31    return dict(32        image_size=64,33        classifier_use_fp16=False,34        classifier_width=128,35        classifier_depth=2,36        classifier_attention_resolutions="32,16,8",  # 1637        classifier_use_scale_shift_norm=True,  # False38        classifier_resblock_updown=True,  # False39        classifier_pool="attention",40    )41 42 43def model_and_diffusion_defaults():44    """45    Defaults for image training.46    """47    res = dict(48        image_size=64,49        num_channels=128,50        num_res_blocks=2,51        num_heads=4,52        num_heads_upsample=-1,53        num_head_channels=-1,54        attention_resolutions="16,8",55        channel_mult="",56        dropout=0.0,57        class_cond=False,58        use_checkpoint=False,59        use_scale_shift_norm=True,60        resblock_updown=False,61        use_fp16=False,62        use_new_attention_order=False,63    )64    res.update(diffusion_defaults())65    return res66 67 68def classifier_and_diffusion_defaults():69    res = classifier_defaults()70    res.update(diffusion_defaults())71    return res72 73 74def create_model_and_diffusion(75    image_size,76    class_cond,77    learn_sigma,78    num_channels,79    num_res_blocks,80    channel_mult,81    num_heads,82    num_head_channels,83    num_heads_upsample,84    attention_resolutions,85    dropout,86    diffusion_steps,87    noise_schedule,88    timestep_respacing,89    use_kl,90    predict_xstart,91    rescale_timesteps,92    rescale_learned_sigmas,93    use_checkpoint,94    use_scale_shift_norm,95    resblock_updown,96    use_fp16,97    use_new_attention_order,98):99    model = create_model(100        image_size,101        num_channels,102        num_res_blocks,103        channel_mult=channel_mult,104        learn_sigma=learn_sigma,105        class_cond=class_cond,106        use_checkpoint=use_checkpoint,107        attention_resolutions=attention_resolutions,108        num_heads=num_heads,109        num_head_channels=num_head_channels,110        num_heads_upsample=num_heads_upsample,111        use_scale_shift_norm=use_scale_shift_norm,112        dropout=dropout,113        resblock_updown=resblock_updown,114        use_fp16=use_fp16,115        use_new_attention_order=use_new_attention_order,116    )117    diffusion = create_gaussian_diffusion(118        steps=diffusion_steps,119        learn_sigma=learn_sigma,120        noise_schedule=noise_schedule,121        use_kl=use_kl,122        predict_xstart=predict_xstart,123        rescale_timesteps=rescale_timesteps,124        rescale_learned_sigmas=rescale_learned_sigmas,125        timestep_respacing=timestep_respacing,126    )127    return model, diffusion128 129 130def create_model(131    image_size,132    num_channels,133    num_res_blocks,134    channel_mult="",135    learn_sigma=False,136    class_cond=False,137    use_checkpoint=False,138    attention_resolutions="16",139    num_heads=1,140    num_head_channels=-1,141    num_heads_upsample=-1,142    use_scale_shift_norm=False,143    dropout=0,144    resblock_updown=False,145    use_fp16=False,146    use_new_attention_order=False,147):148    if channel_mult == "":149        if image_size == 512:150            channel_mult = (0.5, 1, 1, 2, 2, 4, 4)151        elif image_size == 256:152            channel_mult = (1, 1, 2, 2, 4, 4)153        elif image_size == 128:154            channel_mult = (1, 1, 2, 3, 4)155        elif image_size == 64:156            channel_mult = (1, 2, 3, 4)157        else:158            raise ValueError(f"unsupported image size: {image_size}")159    else:160        channel_mult = tuple(int(ch_mult) for ch_mult in channel_mult.split(","))161 162    attention_ds = []163    for res in attention_resolutions.split(","):164        attention_ds.append(image_size // int(res))165 166    return UNetModel(167        image_size=image_size,168        in_channels=3,169        model_channels=num_channels,170        out_channels=(3 if not learn_sigma else 6),171        num_res_blocks=num_res_blocks,172        attention_resolutions=tuple(attention_ds),173        dropout=dropout,174        channel_mult=channel_mult,175        num_classes=(NUM_CLASSES if class_cond else None),176        use_checkpoint=use_checkpoint,177        use_fp16=use_fp16,178        num_heads=num_heads,179        num_head_channels=num_head_channels,180        num_heads_upsample=num_heads_upsample,181        use_scale_shift_norm=use_scale_shift_norm,182        resblock_updown=resblock_updown,183        use_new_attention_order=use_new_attention_order,184    )185 186 187def create_classifier_and_diffusion(188    image_size,189    classifier_use_fp16,190    classifier_width,191    classifier_depth,192    classifier_attention_resolutions,193    classifier_use_scale_shift_norm,194    classifier_resblock_updown,195    classifier_pool,196    learn_sigma,197    diffusion_steps,198    noise_schedule,199    timestep_respacing,200    use_kl,201    predict_xstart,202    rescale_timesteps,203    rescale_learned_sigmas,204):205    classifier = create_classifier(206        image_size,207        classifier_use_fp16,208        classifier_width,209        classifier_depth,210        classifier_attention_resolutions,211        classifier_use_scale_shift_norm,212        classifier_resblock_updown,213        classifier_pool,214    )215    diffusion = create_gaussian_diffusion(216        steps=diffusion_steps,217        learn_sigma=learn_sigma,218        noise_schedule=noise_schedule,219        use_kl=use_kl,220        predict_xstart=predict_xstart,221        rescale_timesteps=rescale_timesteps,222        rescale_learned_sigmas=rescale_learned_sigmas,223        timestep_respacing=timestep_respacing,224    )225    return classifier, diffusion226 227 228def create_classifier(229    image_size,230    classifier_use_fp16,231    classifier_width,232    classifier_depth,233    classifier_attention_resolutions,234    classifier_use_scale_shift_norm,235    classifier_resblock_updown,236    classifier_pool,237):238    if image_size == 512:239        channel_mult = (0.5, 1, 1, 2, 2, 4, 4)240    elif image_size == 256:241        channel_mult = (1, 1, 2, 2, 4, 4)242    elif image_size == 128:243        channel_mult = (1, 1, 2, 3, 4)244    elif image_size == 64:245        channel_mult = (1, 2, 3, 4)246    else:247        raise ValueError(f"unsupported image size: {image_size}")248 249    attention_ds = []250    for res in classifier_attention_resolutions.split(","):251        attention_ds.append(image_size // int(res))252 253    return EncoderUNetModel(254        image_size=image_size,255        in_channels=3,256        model_channels=classifier_width,257        out_channels=1000,258        num_res_blocks=classifier_depth,259        attention_resolutions=tuple(attention_ds),260        channel_mult=channel_mult,261        use_fp16=classifier_use_fp16,262        num_head_channels=64,263        use_scale_shift_norm=classifier_use_scale_shift_norm,264        resblock_updown=classifier_resblock_updown,265        pool=classifier_pool,266    )267 268 269def sr_model_and_diffusion_defaults():270    res = model_and_diffusion_defaults()271    res["large_size"] = 256272    res["small_size"] = 64273    arg_names = inspect.getfullargspec(sr_create_model_and_diffusion)[0]274    for k in res.copy().keys():275        if k not in arg_names:276            del res[k]277    return res278 279 280def sr_create_model_and_diffusion(281    large_size,282    small_size,283    class_cond,284    learn_sigma,285    num_channels,286    num_res_blocks,287    num_heads,288    num_head_channels,289    num_heads_upsample,290    attention_resolutions,291    dropout,292    diffusion_steps,293    noise_schedule,294    timestep_respacing,295    use_kl,296    predict_xstart,297    rescale_timesteps,298    rescale_learned_sigmas,299    use_checkpoint,300    use_scale_shift_norm,301    resblock_updown,302    use_fp16,303):304    model = sr_create_model(305        large_size,306        small_size,307        num_channels,308        num_res_blocks,309        learn_sigma=learn_sigma,310        class_cond=class_cond,311        use_checkpoint=use_checkpoint,312        attention_resolutions=attention_resolutions,313        num_heads=num_heads,314        num_head_channels=num_head_channels,315        num_heads_upsample=num_heads_upsample,316        use_scale_shift_norm=use_scale_shift_norm,317        dropout=dropout,318        resblock_updown=resblock_updown,319        use_fp16=use_fp16,320    )321    diffusion = create_gaussian_diffusion(322        steps=diffusion_steps,323        learn_sigma=learn_sigma,324        noise_schedule=noise_schedule,325        use_kl=use_kl,326        predict_xstart=predict_xstart,327        rescale_timesteps=rescale_timesteps,328        rescale_learned_sigmas=rescale_learned_sigmas,329        timestep_respacing=timestep_respacing,330    )331    return model, diffusion332 333 334def sr_create_model(335    large_size,336    small_size,337    num_channels,338    num_res_blocks,339    learn_sigma,340    class_cond,341    use_checkpoint,342    attention_resolutions,343    num_heads,344    num_head_channels,345    num_heads_upsample,346    use_scale_shift_norm,347    dropout,348    resblock_updown,349    use_fp16,350):351    _ = small_size  # hack to prevent unused variable352 353    if large_size == 512:354        channel_mult = (1, 1, 2, 2, 4, 4)355    elif large_size == 256:356        channel_mult = (1, 1, 2, 2, 4, 4)357    elif large_size == 64:358        channel_mult = (1, 2, 3, 4)359    else:360        raise ValueError(f"unsupported large size: {large_size}")361 362    attention_ds = []363    for res in attention_resolutions.split(","):364        attention_ds.append(large_size // int(res))365 366    return SuperResModel(367        image_size=large_size,368        in_channels=3,369        model_channels=num_channels,370        out_channels=(3 if not learn_sigma else 6),371        num_res_blocks=num_res_blocks,372        attention_resolutions=tuple(attention_ds),373        dropout=dropout,374        channel_mult=channel_mult,375        num_classes=(NUM_CLASSES if class_cond else None),376        use_checkpoint=use_checkpoint,377        num_heads=num_heads,378        num_head_channels=num_head_channels,379        num_heads_upsample=num_heads_upsample,380        use_scale_shift_norm=use_scale_shift_norm,381        resblock_updown=resblock_updown,382        use_fp16=use_fp16,383    )384 385 386def create_gaussian_diffusion(387    *,388    steps=1000,389    learn_sigma=False,390    sigma_small=False,391    noise_schedule="linear",392    use_kl=False,393    predict_xstart=False,394    rescale_timesteps=False,395    rescale_learned_sigmas=False,396    timestep_respacing="",397):398    betas = gd.get_named_beta_schedule(noise_schedule, steps)399    if use_kl:400        loss_type = gd.LossType.RESCALED_KL401    elif rescale_learned_sigmas:402        loss_type = gd.LossType.RESCALED_MSE403    else:404        loss_type = gd.LossType.MSE405    if not timestep_respacing:406        timestep_respacing = [steps]407    return SpacedDiffusion(408        use_timesteps=space_timesteps(steps, timestep_respacing),409        betas=betas,410        model_mean_type=(411            gd.ModelMeanType.EPSILON if not predict_xstart else gd.ModelMeanType.START_X412        ),413        model_var_type=(414            (415                gd.ModelVarType.FIXED_LARGE416                if not sigma_small417                else gd.ModelVarType.FIXED_SMALL418            )419            if not learn_sigma420            else gd.ModelVarType.LEARNED_RANGE421        ),422        loss_type=loss_type,423        rescale_timesteps=rescale_timesteps,424    )425 426 427def add_dict_to_argparser(parser, default_dict):428    for k, v in default_dict.items():429        v_type = type(v)430        if v is None:431            v_type = str432        elif isinstance(v, bool):433            v_type = str2bool434        parser.add_argument(f"--{k}", default=v, type=v_type)435 436 437def args_to_dict(args, keys):438    return {k: getattr(args, k) for k in keys}439 440 441def str2bool(v):442    """443    https://stackoverflow.com/questions/15008758/parsing-boolean-values-with-argparse444    """445    if isinstance(v, bool):446        return v447    if v.lower() in ("yes", "true", "t", "y", "1"):448        return True449    elif v.lower() in ("no", "false", "f", "n", "0"):450        return False451    else:452        raise argparse.ArgumentTypeError("boolean value expected")453