(
image_size,
classifier_use_fp16,
classifier_width,
classifier_depth,
classifier_attention_resolutions,
classifier_use_scale_shift_norm,
classifier_resblock_updown,
classifier_pool,
dataset,
)
| 242 | |
| 243 | |
| 244 | def create_classifier( |
| 245 | image_size, |
| 246 | classifier_use_fp16, |
| 247 | classifier_width, |
| 248 | classifier_depth, |
| 249 | classifier_attention_resolutions, |
| 250 | classifier_use_scale_shift_norm, |
| 251 | classifier_resblock_updown, |
| 252 | classifier_pool, |
| 253 | dataset, |
| 254 | ): |
| 255 | if image_size == 256: |
| 256 | channel_mult = (1, 1, 2, 2, 4, 4) |
| 257 | elif image_size == 128: |
| 258 | channel_mult = (1, 1, 2, 3, 4) |
| 259 | elif image_size == 64: |
| 260 | channel_mult = (1, 2, 3, 4) |
| 261 | else: |
| 262 | raise ValueError(f"unsupported image size: {image_size}") |
| 263 | |
| 264 | attention_ds = [] |
| 265 | for res in classifier_attention_resolutions.split(","): |
| 266 | attention_ds.append(image_size // int(res)) |
| 267 | if dataset=='brats': |
| 268 | number_in_channels=4 |
| 269 | else: |
| 270 | number_in_channels=1 |
| 271 | print('number_in_channels classifier', number_in_channels) |
| 272 | |
| 273 | |
| 274 | return EncoderUNetModel( |
| 275 | image_size=image_size, |
| 276 | in_channels=number_in_channels, |
| 277 | model_channels=classifier_width, |
| 278 | out_channels=2, |
| 279 | num_res_blocks=classifier_depth, |
| 280 | attention_resolutions=tuple(attention_ds), |
| 281 | channel_mult=channel_mult, |
| 282 | use_fp16=classifier_use_fp16, |
| 283 | num_head_channels=64, |
| 284 | use_scale_shift_norm=classifier_use_scale_shift_norm, |
| 285 | resblock_updown=classifier_resblock_updown, |
| 286 | pool=classifier_pool, |
| 287 | ) |
| 288 | |
| 289 | |
| 290 | def sr_model_and_diffusion_defaults(): |
no test coverage detected