始める無料で始める

GANモデルの学習

PyBooks のチームは、Generative Adversarial Network (GAN) を使ったテキスト生成に向けて順調に進んでいます。ジェネレーターとディスクリミネーターのネットワーク定義は完了しました。次は、それらを学習させる段階です。最後のステップとして、偽データを生成し、実データと比較して GAN の学習状況を確認します。入力にはテンソルを用い、出力は入力テンソルに似た形になるようにします。これにより、PyBooks のチームは、テキストデータと同じ関係性をもつ特徴量を備えた合成データをテキスト分析に活用できます。

ジェネレーターとディスクリミネーターはそれぞれ generatordiscriminator に初期化・保存されています。

この演習では次の変数が初期化されています。

  • seq_length = 5: 各合成データ系列の長さ
  • num_sequences = 100: 生成する系列の総数
  • num_epochs = 50: データセットを何周するか
  • print_every = 10: 出力表示の頻度(10エポックごとに結果を表示)

この演習はコースの一部です

PyTorch で学ぶテキストの Deep Learning

コースを見る

実践的なインタラクティブ演習

このサンプルコードを完成させて、この演習に挑戦してみましょう。

# Define the loss function and optimizer
criterion = nn.____()
optimizer_gen = ____(generator.parameters(), lr=0.001)
optimizer_disc = ____(discriminator.parameters(), lr=0.001)

for epoch in range(num_epochs):
    for real_data in data:
      	# Unsqueezing real_data and prevent gradient recalculations
        real_data = real_data.____(0)
        noise = torch.rand((1, seq_length))
        fake_data = generator(noise)
        disc_real = discriminator(real_data)
        disc_fake = discriminator(fake_data.____())
        loss_disc = criterion(disc_real, torch.ones_like(disc_real)) + criterion(disc_fake, torch.zeros_like(disc_fake))
        optimizer_disc.zero_grad()
        loss_disc.backward()
        optimizer_disc.step()
コードを編集して実行