Anonymous-123/ImageNet-Editing
1
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 