เริ่มต้นใช้งานเริ่มต้นใช้งานได้ฟรี

ค่า Loss ของ Discriminator

ถึงเวลากำหนดค่า loss สำหรับ discriminator แล้ว โดย discriminator มีหน้าที่จำแนกรูปภาพว่าเป็นของจริงหรือของปลอม ดังนั้น discriminator จะเกิด loss เมื่อจำแนกผลลัพธ์ของ generator ว่าเป็นของจริง (label 1) หรือจำแนกรูปภาพจริงว่าเป็นของปลอม (label 0)

กำหนดฟังก์ชัน disc_loss() สำหรับคำนวณค่า loss ของ discriminator โดยรับอาร์กิวเมนต์ 5 ตัว ดังนี้

  • gen — โมเดล generator
  • disc — โมเดล discriminator
  • real — กลุ่มตัวอย่างรูปภาพจริงจากข้อมูลฝึก
  • num_images — จำนวนรูปภาพใน batch
  • z_dim — ขนาดของ noise แบบสุ่มที่ใช้เป็น input

แบบฝึกหัดนี้เป็นส่วนหนึ่งของหลักสูตร

Deep Learning สำหรับภาพด้วย PyTorch

ดูคอร์ส

คำแนะนำการฝึกหัด

  • ใช้ discriminator จำแนก fake images แล้วกำหนดผลการทำนายให้กับ disc_pred_fake
  • คำนวณค่า loss ส่วนของรูปภาพปลอม โดยเรียก criterion ด้วยผลการทำนายของ discriminator สำหรับรูปภาพปลอม และ tensor ของค่าศูนย์ที่มีรูปร่างเดียวกัน
  • ใช้ discriminator จำแนก real images แล้วกำหนดผลการทำนายให้กับ disc_pred_real
  • คำนวณค่า loss ส่วนของรูปภาพจริง โดยเรียก criterion ด้วยผลการทำนายของ discriminator สำหรับรูปภาพจริง และ tensor ของค่าหนึ่งที่มีรูปร่างเดียวกัน

แบบฝึกหัดเชิงโต้ตอบแบบลงมือทำ

ลองทำแบบฝึกหัดนี้โดยเติมโค้ดตัวอย่างนี้ให้สมบูรณ์

def disc_loss(gen, disc, real, num_images, z_dim):
    criterion = nn.BCEWithLogitsLoss()
    noise = torch.randn(num_images, z_dim)
    fake = gen(noise)
    # Get discriminator's predictions for fake images
    disc_pred_fake = ____
    # Calculate the fake loss component
    fake_loss = ____
    # Get discriminator's predictions for real images
    disc_pred_real = ____
    # Calculate the real loss component
    real_loss = ____
    disc_loss = (real_loss + fake_loss) / 2
    return disc_loss
แก้ไขและรันโค้ด