MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / train_one_epoch

Function train_one_epoch

tensorflow/contrib/eager/python/examples/gan/mnist.py:197–262  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

195
196
197def 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))

Callers 1

mainFunction · 0.70

Calls 9

discriminator_lossFunction · 0.85
generator_lossFunction · 0.85
random_uniformMethod · 0.80
reshapeMethod · 0.80
gradientMethod · 0.80
deviceMethod · 0.45
assign_addMethod · 0.45
GradientTapeMethod · 0.45
apply_gradientsMethod · 0.45

Tested by

no test coverage detected