Trains `generator` and `discriminator` models on `dataset`. Args: generator: Generator model. discriminator: Discriminator model. generator_optimizer: Optimizer to use for generator. discriminator_optimizer: Optimizer to use for discriminator. dataset: Dataset of images to tra
(generator, discriminator, generator_optimizer,
discriminator_optimizer, dataset, step_counter,
log_interval, noise_dim)
| 195 | |
| 196 | |
| 197 | def train_one_epoch(generator, discriminator, generator_optimizer, |
| 198 | discriminator_optimizer, dataset, step_counter, |
| 199 | log_interval, noise_dim): |
| 200 | """Trains `generator` and `discriminator` models on `dataset`. |
| 201 | |
| 202 | Args: |
| 203 | generator: Generator model. |
| 204 | discriminator: Discriminator model. |
| 205 | generator_optimizer: Optimizer to use for generator. |
| 206 | discriminator_optimizer: Optimizer to use for discriminator. |
| 207 | dataset: Dataset of images to train on. |
| 208 | step_counter: An integer variable, used to write summaries regularly. |
| 209 | log_interval: How many steps to wait between logging and collecting |
| 210 | summaries. |
| 211 | noise_dim: Dimension of noise vector to use. |
| 212 | """ |
| 213 | |
| 214 | total_generator_loss = 0.0 |
| 215 | total_discriminator_loss = 0.0 |
| 216 | for (batch_index, images) in enumerate(dataset): |
| 217 | with tf.device('/cpu:0'): |
| 218 | tf.assign_add(step_counter, 1) |
| 219 | |
| 220 | with tf.contrib.summary.record_summaries_every_n_global_steps( |
| 221 | log_interval, global_step=step_counter): |
| 222 | current_batch_size = images.shape[0] |
| 223 | noise = tf.random_uniform( |
| 224 | shape=[current_batch_size, noise_dim], |
| 225 | minval=-1., |
| 226 | maxval=1., |
| 227 | seed=batch_index) |
| 228 | |
| 229 | # we can use 2 tapes or a single persistent tape. |
| 230 | # Using two tapes is memory efficient since intermediate tensors can be |
| 231 | # released between the two .gradient() calls below |
| 232 | with tf.GradientTape() as gen_tape, tf.GradientTape() as disc_tape: |
| 233 | generated_images = generator(noise) |
| 234 | tf.contrib.summary.image( |
| 235 | 'generated_images', |
| 236 | tf.reshape(generated_images, [-1, 28, 28, 1]), |
| 237 | max_images=10) |
| 238 | |
| 239 | discriminator_gen_outputs = discriminator(generated_images) |
| 240 | discriminator_real_outputs = discriminator(images) |
| 241 | discriminator_loss_val = discriminator_loss(discriminator_real_outputs, |
| 242 | discriminator_gen_outputs) |
| 243 | total_discriminator_loss += discriminator_loss_val |
| 244 | |
| 245 | generator_loss_val = generator_loss(discriminator_gen_outputs) |
| 246 | total_generator_loss += generator_loss_val |
| 247 | |
| 248 | generator_grad = gen_tape.gradient(generator_loss_val, |
| 249 | generator.variables) |
| 250 | discriminator_grad = disc_tape.gradient(discriminator_loss_val, |
| 251 | discriminator.variables) |
| 252 | |
| 253 | generator_optimizer.apply_gradients( |
| 254 | zip(generator_grad, generator.variables)) |
no test coverage detected