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)))]
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ủ đề
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))
passInput
- 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ểutorch.Tensortheo đúng đặc tả kỹ thuật và kích thước quy định. - Hàm
compute_generator_loss: Trả về kết quả kiểutorch.Tensortheo đú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.6931Giả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.0067Giả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ế.
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.
Góp ý & báo lỗi bài tập
Đề bài chưa rõ, test có vấn đề hay bạn có ý tưởng giúp bài tốt hơn? Gửi cho đội ngũ AI Empire nhé — mỗi góp ý đều được đọc.
