शुरू करेंमुफ़्त में शुरू करें

ट्रेनिंग लूप

आखिरकार, आपने मॉडल आर्किटेक्चर और loss फंक्शनों को परिभाषित करने में जो मेहनत की थी, अब उसका फल मिलेगा: यह ट्रेनिंग का समय है! आपका कार्य GAN का ट्रेनिंग लूप इम्प्लीमेंट करना और चलाना है। नोट: लंबे रनटाइम से बचने के लिए पहले बैच के बाद एक break स्टेटमेंट रखा गया है.

दो ऑप्टिमाइज़र disc_opt और gen_opt को Adam() ऑप्टिमाइज़र के रूप में इनिशियलाइज़ किया गया है। आपने पहले जो loss निकालने वाले फंक्शन परिभाषित किए थे, gen_loss() और disc_loss(), वे उपलब्ध हैं। एक dataloader भी आपके लिए तैयार है.

ध्यान रखें कि:

  • disc_loss() के आर्ग्युमेंट्स हैं: gen, disc, real, cur_batch_size, z_dim.
  • gen_loss() के आर्ग्युमेंट्स हैं: gen, disc, cur_batch_size, z_dim.

यह अभ्यास पाठ्यक्रम का हिस्सा है

PyTorch के साथ इमेज के लिए डीप लर्निंग

पाठ्यक्रम देखें

अभ्यास निर्देश

  • जनरेटर, डिस्क्रिमिनेटर, रियल इमेज के सैंपल, करंट बैच साइज, और 16 के noise साइज को इसी क्रम में disc_loss() में पास करके डिस्क्रिमिनेटर loss निकालें और परिणाम d_loss में असाइन करें.
  • d_loss का उपयोग करके gradients निकालें.
  • जनरेटर, डिस्क्रिमिनेटर, करंट बैच साइज, और 16 के noise साइज को इसी क्रम में gen_loss() में पास करके जनरेटर loss निकालें और परिणाम g_loss में असाइन करें.
  • g_loss का उपयोग करके gradients निकालें.

इंटरैक्टिव व्यावहारिक अभ्यास

इस अभ्यास को इस नमूना कोड को पूरा करके आज़माएँ।

for epoch in range(1):
    for real in dataloader:
        cur_batch_size = len(real)
        
        disc_opt.zero_grad()
        # Calculate discriminator loss
        d_loss = ____
        # Compute gradients
        ____
        disc_opt.step()

        gen_opt.zero_grad()
        # Calculate generator loss
        g_loss = ____
        # Compute generator gradients
        ____
        gen_opt.step()

        print(f"Generator loss: {g_loss}")
        print(f"Discriminator loss: {d_loss}")
        break
कोड संपादित करें और चलाएँ