| 272 | |
| 273 | |
| 274 | def DCGAN(generator, discriminator_model, img_dim, patch_size, image_dim_ordering): |
| 275 | |
| 276 | gen_input = Input(shape=img_dim, name="DCGAN_input") |
| 277 | |
| 278 | generated_image = generator(gen_input) |
| 279 | |
| 280 | if image_dim_ordering == "channels_first": |
| 281 | h, w = img_dim[1:] |
| 282 | else: |
| 283 | h, w = img_dim[:-1] |
| 284 | ph, pw = patch_size |
| 285 | |
| 286 | list_row_idx = [(i * ph, (i + 1) * ph) for i in range(h // ph)] |
| 287 | list_col_idx = [(i * pw, (i + 1) * pw) for i in range(w // pw)] |
| 288 | |
| 289 | list_gen_patch = [] |
| 290 | for row_idx in list_row_idx: |
| 291 | for col_idx in list_col_idx: |
| 292 | if image_dim_ordering == "channels_last": |
| 293 | x_patch = Lambda(lambda z: z[:, row_idx[0]:row_idx[1], col_idx[0]:col_idx[1], :])(generated_image) |
| 294 | else: |
| 295 | x_patch = Lambda(lambda z: z[:, :, row_idx[0]:row_idx[1], col_idx[0]:col_idx[1]])(generated_image) |
| 296 | list_gen_patch.append(x_patch) |
| 297 | |
| 298 | DCGAN_output = discriminator_model(list_gen_patch) |
| 299 | |
| 300 | DCGAN = Model(inputs=[gen_input], |
| 301 | outputs=[generated_image, DCGAN_output], |
| 302 | name="DCGAN") |
| 303 | |
| 304 | return DCGAN |
| 305 | |
| 306 | |
| 307 | def load(model_name, img_dim, nb_patch, bn_mode, use_mbd, batch_size, do_plot): |