ai-098Đọc toàn bộ đề miễn phí

Hàm Mất Mát GAN Minimax Và Non-Saturating (Goodfellow et al. 2014)

LD = - Ex [log D(x)] - Ez [log (1 - D(G(z)))]

AITrung bình35 phút

Tiến độ của tôi ở bài này

Điểm được lưu vào tài khoản sau khi chấm bài.

Đang tải điểm của bạn…

Kiến thức và chủ đề

generative-modelsganadversarial-trainingminimaxpytorch

Kiến thức tiên quyết: binary-cross-entropy-from-scratch.

Nội dung đề bài

Mục tiêu kiến thức

  • Trong GAN nguyên bản (Goodfellow 2014):
  • Mạng Phân biệt D: Nhận diện ảnh thật x là nhãn 1 và ảnh giả G(z) là nhãn 0.
  • Mất mát Discriminator:

LD = - Ex [log D(x)] - Ez [log (1 - D(G(z)))]

  • Vấn đề bão hòa gradient (Vanishing Gradients) của Minimax Generator:
  • Mục tiêu gốc của Generator là cực tiểu hóa Ez [log(1 - D(G(z)))]. Ở giai đoạn đầu huấn luyện khi D dễ dàng phân biệt ảnh giả (D(G(z)) ≈ 0), log(1 - D) có đạo hàm gần bằng 0, khiến Generator không thể học được gì!
  • Giải pháp Non-Saturating Loss:
  • Thay vì cực tiểu hóa log(1 - D), Goodfellow đề xuất cực đại hóa log D(G(z)), tức cực tiểu hóa:

LGNS = - Ez [log D(G(z))]

  • Hàm này cung cấp gradient cực lớn ở đầu quá trình huấn luyện khi D(G(z)) gần 0.

Yêu cầu

Viết 2 hàm tính mất mát từ logits thô (chưa qua sigmoid) để đảm bảo ổn định số học:

def compute_discriminator_loss(real_logits: torch.Tensor, fake_logits: torch.Tensor) -> torch.Tensor:
    # Binary cross entropy: real target = 1.0, fake target = 0.0
    pass

def compute_generator_loss(fake_logits: torch.Tensor, mode: str = "non_saturating") -> torch.Tensor:
    # mode = "non_saturating": target = 1.0
    # mode = "minimax": - log(1 - sigmoid(fake_logits))
    pass

Input

  • Hàm compute_discriminator_loss(real_logits, fake_logits): Các tham số đầu vào chứa dữ liệu Tensor/mảng NumPy hoặc giá trị siêu tham số tương ứng.
  • Hàm compute_generator_loss(fake_logits, mode): Các tham số đầu vào chứa dữ liệu Tensor/mảng NumPy hoặc giá trị siêu tham số tương ứng.

Output

  • Hàm compute_discriminator_loss: Trả về kết quả kiểu torch.Tensor theo đúng đặc tả kỹ thuật và kích thước quy định.
  • Hàm compute_generator_loss: Trả về kết quả kiểu torch.Tensor theo đúng đặc tả kỹ thuật và kích thước quy định.

Ràng buộc

  • Thời gian chạy tối đa: 2000ms.
  • Giới hạn bộ nhớ: 256MB.
  • Dữ liệu đầu vào hợp lệ theo đúng kiểu dữ liệu và miền giá trị được mô tả.

Ví dụ 1

Input

real = torch.zeros(4)
fake = torch.zeros(4)
d_loss = compute_discriminator_loss(real, fake)
g_ns = compute_generator_loss(fake, mode='non_saturating')

Output

0.6931

Giải thích

Hàm/lớp được gọi với các tham số mẫu trên và trả về kết quả số học / kích thước tensor tương ứng theo đúng thiết kế.

Ví dụ 2

Input

'Detect wrong_01 (label flip) using non-zero logits.\n\n    When D correctly classifies (real_logits >> 0, fake_logits << 0):\n    - Correct: assigns real=1, fake=0 -> low D loss\n    - Wrong_01: assigns real=0, fake=1 -> high D loss\n    '
real_logits = torch.tensor([5.0, 5.0])
fake_logits = torch.tensor([-5.0, -5.0])
d_loss = compute_discriminator_loss(real_logits, fake_logits)

Output

0.0067

Giải thích

Hàm/lớp được gọi với các tham số mẫu trên và trả về kết quả số học / kích thước tensor tương ứng theo đúng thiết kế.

3 cấp độ gợi ýMở dần khi bạn thật sự cần hỗ trợ.
Phân tích lời giảiGiải thích hướng tư duy và thuật toán.
Code tham khảoDùng để đối chiếu sau khi tự làm.

Gợi ý và lời giải chỉ mở sau khi bạn bấm Nộp bài. Giáo viên và quản trị viên mở được ngay.