ageron / ageron/handson-ml2

[QUESTION] Plotting loss with the Deep Convolutional GAN

Open
#563 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
Jupyter Notebook
Stars
30k
Forks
13.1k
PR merge metrics
No merged PRs in 30d

Description

**Describe what is unclear to you**
When creating autoencoders, we were using fit() to produce the history, which could then be used to plot the loss across the training and validation periods, such as on p. 590:
```python
history = variational_ae.fit(X_train, X_train, epochs=25, batch_size=128,
validation_data=(X_valid, X_valid))
```
However, when creating the GANs and deep convolutional GANs, we do not use fit(), we use the custom train_gan function:
```python
def train_gan(gan, dataset, batch_size, codings_size, n_epochs=20):
generator, discriminator = gan.layers
for epoch in range(n_epochs):
print("Epoch {}/{}".format(epoch + 1, n_epochs))
for X_batch in dataset:
# phase 1 - training the discriminator
X_batch = tf.cast(X_batch, tf.float32)
noise = tf.random.normal(shape=[batch_size, codings_size])
generated_images = generator(noise)
X_fake_and_real = tf.concat([generated_images, X_batch], axis=0)
y1 = tf.constant([[0.]] * batch_size + [[1.]] * batch_size)
discriminator.trainable = True
discriminator.train_on_batch(X_fake_and_real, y1)
# phase 2 - training the generator
noise = tf.random.normal(shape=[batch_size, codings_size])
y2 = tf.constant([[1.]] * batch_size)
discriminator.trainable = False
gan.train_on_batch(noise, y2)
plot_multiple_images(generated_images, 8)
plt.show()
```
If we wanted to plot the loss of both the discriminator and generator across all epochs in the example on page 599, how would we go about this?

Thanks!

Contributor guide

No contributing guide indexed for this repository

Research direction

Read the train_gan function shown in the page 599 example and inspect the calls to discriminator.train_on_batch and gan.train_on_batch. Determine how per-batch results could be retained across epochs, then verify that both discriminator and generator loss histories can be plotted.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, tensorflow
Domain
machine-learning
Issue type
Documentation
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Mostly clear
Newbie friendliness
30/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.