You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch实现创意对抗网络(CAN):自定义损失适配难题

Fixing Generator Loss for Creative Adversarial Networks (CAN) in PyTorch

Got it, let's break down your problem—you're on the right track with the discriminator using nn.CrossEntropyLoss, but the generator side is tripping you up because that loss function isn't built for what CAN's generator needs to do. Let's fix this.

1. Why nn.CrossEntropyLoss Fails for the Generator

First, let's clarify the core mismatch:

  • nn.CrossEntropyLoss expects targets as LongTensor (class indices) and inputs as classification logits. But the generator's goal in CAN isn't to predict class indices—it's to produce samples that are both realistic (fool the discriminator) and creative (deviate from known class distributions). This tensor type and objective mismatch is why you're hitting walls.

2. Why Your Discriminator's CrossEntropyLoss Makes Sense

Your discriminator is likely performing a classification task (e.g., distinguishing real samples from their known classes vs. generated samples). For this, nn.CrossEntropyLoss is perfect: it takes the discriminator's logit outputs (FloatTensor) and matches them to class index targets (LongTensor). The decreasing loss you're seeing is exactly what you want here—your discriminator is getting better at telling real/generated samples apart.

3. Implementing the Generator's Custom CAN Loss

CAN's generator loss has two key components: adversarial loss (fool the discriminator) and creative loss (push samples away from existing class distributions). Here's how to build both:

A. Adversarial Loss for Generator

Instead of CrossEntropyLoss, use nn.BCEWithLogitsLoss or nn.BCELoss for the adversarial part—these work with FloatTensor targets, which aligns with the generator's goal:

  • If your discriminator outputs raw logits (no sigmoid), use BCEWithLogitsLoss (it includes sigmoid internally for numerical stability).
  • If the discriminator outputs sigmoid-activated probabilities, use BCELoss.

Example code snippet:

# Initialize loss function
adv_loss_gen = nn.BCEWithLogitsLoss()

# During generator training:
generated_samples = generator(z)
disc_logits_gen = discriminator(generated_samples)
# Target: make discriminator think generated samples are "real" (class 1)
targets_real = torch.ones_like(disc_logits_gen, dtype=torch.float32)
adversarial_loss = adv_loss_gen(disc_logits_gen, targets_real)

B. Creative Loss (Core of CAN)

The creative loss is what sets CAN apart—you need to penalize the generator for producing samples that are too close to existing class distributions. A common implementation (aligned with the CAN paper) is to compute the distance between generated samples' features and the centers of known classes:

# Precompute class centers from real training data (run once before training)
class_centers = {}
for class_idx in range(num_classes):
    class_samples = real_data[real_labels == class_idx]
    # Extract intermediate features from discriminator (add this method to your discriminator)
    class_features = discriminator.extract_features(class_samples)
    class_centers[class_idx] = torch.mean(class_features, dim=0)

# Calculate creative loss during generator training
gen_features = discriminator.extract_features(generated_samples)
total_creative_loss = 0.0
for center in class_centers.values():
    # L2 distance between generated features and each class center
    distance = torch.norm(gen_features - center, p=2, dim=1)
    total_creative_loss += torch.mean(distance)

# Weight the creative loss (adjust alpha based on your experiment needs)
alpha = 0.5
creative_loss = alpha * total_creative_loss

C. Total Generator Loss

Combine the two losses for backpropagation:

total_gen_loss = adversarial_loss + creative_loss
total_gen_loss.backward()
optimizer_gen.step()

4. Key Notes for Paper Alignment

  • Tune the alpha hyperparameter to balance realism and creativity—if your generator produces boring samples, increase alpha; if samples are unrealistic, decrease it.
  • If your discriminator is a multi-class classifier (instead of binary real/fake), adjust the adversarial loss to target incorrect class labels (e.g., make the discriminator assign random real classes to generated samples).

内容的提问来源于stack exchange,提问作者krustybek

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.26 11:05:29