MCPcopy Create free account
hub / github.com/JuliaWolleb/diffusion-anomaly / create_classifier

Function create_classifier

guided_diffusion/script_util.py:244–287  ·  view source on GitHub ↗
(
    image_size,
    classifier_use_fp16,
    classifier_width,
    classifier_depth,
    classifier_attention_resolutions,
    classifier_use_scale_shift_norm,
    classifier_resblock_updown,
    classifier_pool,
    dataset,
)

Source from the content-addressed store, hash-verified

242
243
244def 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
290def sr_model_and_diffusion_defaults():

Callers 2

mainFunction · 0.90

Calls 1

EncoderUNetModelClass · 0.85

Tested by

no test coverage detected