C.W.K.
Stream
Lesson 04 of 06 · published

GAN 학습 패턴

~8 min · custom-train

Level 0Keras 도제
0 XP0/97 lessons0/20 achievements
0/120 XP to next level120 XP to go0% complete

GAN 이 단일 loss 틀을 깨는 이유

GAN 은 minimax 게임에 묶인 두 네트워크야. discriminator 는 real 과 generated 를 구분하는 법을 배우고, generator 는 discriminator 가 못 잡는 fake 를 만드는 법을 배워. 목적이 정반대고, weight 도 optimizer 도 따로야 — 그래서 fit() 이 가정하는 '단일 loss, 단일 optimizer' 모양으로는 표현이 안 돼. batch 마다 서로 다른 update step 두 개가 순서대로 필요해.

그래도 train_step() 이 맞는 이유

여기서 바로 manual loop 로 가고 싶지만, 이 번갈아-학습 패턴은 train_step() override 안에 완벽하게 들어가. GAN object 에 두 sub-model 을 들고, 한 step 안에서 (1) real+fake batch 로 discriminator 학습 → (2) discriminator 를 고정한 채 그 feedback 으로 generator 학습 → 두 loss 반환. fit() 을 안 떠났으니 GAN 전체가 callback / checkpoint / distribution 을 공짜로 상속해. 아래 뼈대는 두 update 를 named helper (_train_discriminator / _train_generator) 로 빼서 step 이 모델링하는 게임처럼 읽히게 했어 — 각 helper 가 자기 gradient tape 와 optimizer 를 들어.

Code

train_step() override 로 짠 GAN·python
class GAN(keras.Model):
    def __init__(self, generator, discriminator, latent_dim):
        super().__init__()
        self.generator = generator
        self.discriminator = discriminator
        self.latent_dim = latent_dim

    def train_step(self, real_images):
        batch_size = keras.ops.shape(real_images)[0]
        noise = keras.random.normal(
            shape=(batch_size, self.latent_dim)
        )

        # Train discriminator on real + fake
        fake_images = self.generator(noise)
        combined = keras.ops.concatenate([real_images, fake_images])
        labels = keras.ops.concatenate([
            keras.ops.ones((batch_size, 1)),
            keras.ops.zeros((batch_size, 1)),
        ])
        d_loss = self._train_discriminator(combined, labels)

        # Train generator to fool the discriminator
        noise = keras.random.normal(shape=(batch_size, self.latent_dim))
        misleading_labels = keras.ops.ones((batch_size, 1))
        g_loss = self._train_generator(noise, misleading_labels)

        return {"d_loss": d_loss, "g_loss": g_loss}

External links

Exercise

MNIST 용 작은 DCGAN 구현 — small generator (Dense → reshape → ConvTranspose), small discriminator. train_step override 사용. 5 epoch 학습, g_loss 감소 + sample 질적으로 개선 확인.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.