| 23 | NUM_CLASSES = 1000 |
| 24 | |
| 25 | def create_model( |
| 26 | image_size, |
| 27 | num_channels, |
| 28 | num_res_blocks, |
| 29 | channel_mult="", |
| 30 | learn_sigma=False, |
| 31 | class_cond=False, |
| 32 | use_checkpoint=False, |
| 33 | attention_resolutions="16", |
| 34 | num_heads=1, |
| 35 | num_head_channels=-1, |
| 36 | num_heads_upsample=-1, |
| 37 | use_scale_shift_norm=False, |
| 38 | dropout=0, |
| 39 | resblock_updown=False, |
| 40 | use_fp16=False, |
| 41 | use_new_attention_order=False, |
| 42 | model_path='', |
| 43 | ): |
| 44 | if channel_mult == "": |
| 45 | if image_size == 512: |
| 46 | channel_mult = (0.5, 1, 1, 2, 2, 4, 4) |
| 47 | elif image_size == 256: |
| 48 | channel_mult = (1, 1, 2, 2, 4, 4) |
| 49 | elif image_size == 128: |
| 50 | channel_mult = (1, 1, 2, 3, 4) |
| 51 | elif image_size == 64: |
| 52 | channel_mult = (1, 2, 3, 4) |
| 53 | else: |
| 54 | raise ValueError(f"unsupported image size: {image_size}") |
| 55 | else: |
| 56 | channel_mult = tuple(int(ch_mult) for ch_mult in channel_mult.split(",")) |
| 57 | |
| 58 | attention_ds = [] |
| 59 | if isinstance(attention_resolutions, int): |
| 60 | attention_ds.append(image_size // attention_resolutions) |
| 61 | elif isinstance(attention_resolutions, str): |
| 62 | for res in attention_resolutions.split(","): |
| 63 | attention_ds.append(image_size // int(res)) |
| 64 | else: |
| 65 | raise NotImplementedError |
| 66 | |
| 67 | model= UNetModel( |
| 68 | image_size=image_size, |
| 69 | in_channels=3, |
| 70 | model_channels=num_channels, |
| 71 | out_channels=(3 if not learn_sigma else 6), |
| 72 | num_res_blocks=num_res_blocks, |
| 73 | attention_resolutions=tuple(attention_ds), |
| 74 | dropout=dropout, |
| 75 | channel_mult=channel_mult, |
| 76 | num_classes=(NUM_CLASSES if class_cond else None), |
| 77 | use_checkpoint=use_checkpoint, |
| 78 | use_fp16=use_fp16, |
| 79 | num_heads=num_heads, |
| 80 | num_head_channels=num_head_channels, |
| 81 | num_heads_upsample=num_heads_upsample, |
| 82 | use_scale_shift_norm=use_scale_shift_norm, |