【Bug已解决】How can i process multi loss in pytorch? 解决方案问题描述在 PyTorch 深度学习开发中很多实际任务需要同时优化多个损失函数。例如多任务学习Multi-Task Learning一个模型同时完成分类和回归任务需要同时优化分类损失和回归损失。对抗训练生成对抗网络GAN中生成器和判别器各有自己的损失函数。正则化在主损失之外添加 L1/L2 正则化项、对比学习损失等。知识蒸馏学生模型需要同时学习硬标签CrossEntropy和软标签KL散度。辅助任务主任务之外添加辅助分类头帮助模型学习更好的特征表示。处理多损失时开发者经常遇到以下问题多个损失如何组合直接相加还是加权相加权重如何设置不同损失量级差异大怎么办多个损失的反向传播如何正确执行不同损失需要不同的学习率怎么办梯度冲突如何处理本文将系统地介绍 PyTorch 中处理多损失的各种方法。错误复现错误一多个 loss 分别 backward 导致计算图错误import torch import torch.nn as nn class MultiTaskModel(nn.Module): def __init__(self): super().__init__() self.shared nn.Linear(100, 50) self.head_cls nn.Linear(50, 10) # 分类头 self.head_reg nn.Linear(50, 1) # 回归头 def forward(self, x): shared_feat torch.relu(self.shared(x)) cls_output self.head_cls(shared_feat) reg_output self.head_reg(shared_feat) return cls_output, reg_output model MultiTaskModel() x torch.randn(32, 100) cls_target torch.randint(0, 10, (32,)) reg_target torch.randn(32, 1) cls_output, reg_output model(x) loss_cls nn.CrossEntropyLoss()(cls_output, cls_target) loss_reg nn.MSELoss()(reg_output, reg_target) # 错误分别 backward loss_cls.backward() # 第一次 backward # ... optimizer.step() ... loss_reg.backward() # 报错计算图已经被释放报错信息RuntimeError: Trying to backward through the graph a second time错误二损失量级差异导致一个损失主导# 分类损失可能约为 2.3 # 回归损失可能约为 100.0 # 直接相加时回归损失主导了梯度方向 total_loss loss_cls loss_reg # 2.3 100.0 102.3 # 分类损失的梯度几乎被忽略 total_loss.backward()错误三权重设置不当# 错误手动设置固定权重但没有考虑损失量级 total_loss 1.0 * loss_cls 1.0 * loss_reg # 如果 loss_reg loss_cls分类任务学不到东西 # 另一个错误权重设置过大 total_loss 100 * loss_cls 0.01 * loss_reg # 回归任务被完全忽略根因分析1. PyTorch 计算图的生命周期PyTorch 默认在调用.backward()后会释放计算图retain_graphFalse。如果两个损失共享部分计算图如共享特征提取器第一次 backward 后计算图被释放第二次 backward 就会报错。2. 损失量级差异不同损失函数的输出量级可能差异很大CrossEntropyLoss通常在 0.1 ~ 5 之间MSELoss取决于目标值范围可能从 0.001 到 10000L1Loss与 MSELoss 类似但更小KLDivLoss通常很小0.01 ~ 1直接相加会导致大量级的损失主导优化方向。3. 梯度冲突不同任务的梯度方向可能冲突。例如分类任务希望特征具有判别性而回归任务可能希望特征平滑。直接相加可能导致梯度互相抵消。4. 动态权重问题固定权重在不同训练阶段可能不合适。训练初期某个损失可能很大后期变小固定权重无法自适应调整。解决方案方案一损失相加后统一 backward最简单import torch import torch.nn as nn class MultiTaskModel(nn.Module): def __init__(self): super().__init__() self.shared nn.Linear(100, 50) self.head_cls nn.Linear(50, 10) self.head_reg nn.Linear(50, 1) def forward(self, x): shared_feat torch.relu(self.shared(x)) cls_output self.head_cls(shared_feat) reg_output self.head_reg(shared_feat) return cls_output, reg_output model MultiTaskModel() x torch.randn(32, 100) cls_target torch.randint(0, 10, (32,)) reg_target torch.randn(32, 1) cls_output, reg_output model(x) loss_cls nn.CrossEntropyLoss()(cls_output, cls_target) loss_reg nn.MSELoss()(reg_output, reg_target) # 正确先相加再 backward total_loss loss_cls loss_reg total_loss.backward() # 一次 backward计算图正确处理 optimizer torch.optim.Adam(model.parameters(), lr0.001) optimizer.step() optimizer.zero_grad()方案二加权损失组合import torch import torch.nn as nn class WeightedMultiTaskLoss(nn.Module): 加权多任务损失。 支持手动设置权重或自动平衡。 def __init__(self, cls_weight1.0, reg_weight1.0): super().__init__() self.cls_weight cls_weight self.reg_weight reg_weight self.cls_loss nn.CrossEntropyLoss() self.reg_loss nn.MSELoss() def forward(self, cls_output, cls_target, reg_output, reg_target): loss_cls self.cls_loss(cls_output, cls_target) loss_reg self.reg_loss(reg_output, reg_target) total_loss self.cls_weight * loss_cls self.reg_weight * loss_reg return total_loss, loss_cls, loss_reg # 使用示例 model MultiTaskModel() criterion WeightedMultiTaskLoss(cls_weight1.0, reg_weight0.01) cls_output, reg_output model(x) total_loss, loss_cls, loss_reg criterion( cls_output, cls_target, reg_output, reg_target ) print(fCLS loss: {loss_cls.item():.4f}) print(fREG loss: {loss_reg.item():.4f}) print(fTotal loss: {total_loss.item():.4f}) total_loss.backward()方案三GradNorm 自动平衡import torch import torch.nn as nn class GradNormBalancer: GradNorm: 自动平衡多任务损失的梯度。 参考: Chen et al., GradNorm: Gradient Normalization for Adaptive Loss Balancing in Deep Multitask Networks, ICML 2018. 核心思想根据各任务梯度的范数动态调整权重 使各任务对共享参数的梯度贡献保持平衡。 def __init__(self, model, num_tasks, alpha1.5, warmup_epochs5): self.model model self.num_tasks num_tasks self.alpha alpha # 平衡强度 self.warmup_epochs warmup_epochs # 可学习的任务权重 self.task_weights nn.Parameter( torch.ones(num_tasks, requires_gradTrue) ) # 记录初始损失用于归一化 self.initial_losses None def compute_loss(self, losses, epoch): 计算加权损失。 losses_tensor torch.stack(losses) # 记录初始损失 if self.initial_losses is None: self.initial_losses losses_tensor.detach() # 计算损失比率 loss_ratios losses_tensor.detach() / self.initial_losses # 计算目标梯度范数 mean_ratio loss_ratios.mean() target_norm loss_ratios / mean_ratio # 使用可学习权重 weighted_loss (self.task_weights * losses_tensor).sum() return weighted_loss, self.task_weights.detach() def update_weights(self, losses, shared_params, epoch): 更新任务权重。 if epoch self.warmup_epochs: return # 计算各任务对共享参数的梯度范数 grad_norms [] for i, loss in enumerate(losses): grads torch.autograd.grad( loss, shared_params, retain_graphTrue, allow_unusedTrue ) grad_norm torch.norm(torch.stack([g.norm() for g in grads if g is not None])) grad_norms.append(grad_norm) grad_norms torch.stack(grad_norms) mean_grad_norm grad_norms.mean() # 计算权重更新 loss_ratios torch.stack(losses).detach() / self.initial_losses relative_ratios loss_ratios / loss_ratios.mean() # 目标梯度范数 target_grad_norms mean_grad_norm * relative_ratios ** self.alpha # 更新权重 grad_norm_ratios target_grad_norms / (grad_norms 1e-8) self.task_weights.data * grad_norm_ratios # 归一化 self.task_weights.data self.task_weights.data / self.task_weights.data.sum() * self.num_tasks方案四不确定性加权Kendall et al.import torch import torch.nn as nn class UncertaintyWeightedLoss(nn.Module): 基于不确定性的多任务损失加权。 参考: Kendall et al., Multi-Task Learning Using Uncertainty to Weigh Losses for Scene Geometry and Semantics, CVPR 2018. 核心思想使用可学习的不确定性参数噪声方差作为权重 不确定性高的任务权重低不确定性低的任务权重高。 def __init__(self, num_tasks2): super().__init__() # log(σ²) 初始化为 0即 σ²1 self.log_vars nn.Parameter(torch.zeros(num_tasks)) def forward(self, losses): Args: losses: 各任务损失列表 Returns: total_loss: 加权总损失 weights: 各任务权重 total_loss 0 weights [] for i, loss in enumerate(losses): # 精度 1/σ² exp(-log_var) precision torch.exp(-self.log_vars[i]) # 加权损失: precision * loss log(σ²) # log(σ²) 作为正则项防止 σ² 无限增大 weighted_loss precision * loss self.log_vars[i] total_loss weighted_loss weights.append(precision.detach()) return total_loss, torch.stack(weights) # 使用示例 model MultiTaskModel() loss_balancer UncertaintyWeightedLoss(num_tasks2) # 优化器需要包含 loss_balancer 的参数 optimizer torch.optim.Adam( list(model.parameters()) list(loss_balancer.parameters()), lr0.001 ) for epoch in range(10): cls_output, reg_output model(x) loss_cls nn.CrossEntropyLoss()(cls_output, cls_target) loss_reg nn.MSELoss()(reg_output, reg_target) total_loss, weights loss_balancer([loss_cls, loss_reg]) print(fEpoch {epoch}: cls_loss{loss_cls.item():.4f}, freg_loss{loss_reg.item():.4f}, fweights{weights.tolist()}) optimizer.zero_grad() total_loss.backward() optimizer.step()完整修复代码import torch import torch.nn as nn import torch.nn.functional as F from torch.utils.data import DataLoader, TensorDataset import numpy as np  # # 完整示例多任务学习中的多损失处理 # class SharedBackboneModel(nn.Module): 多任务学习模型共享特征提取器 多个任务头。 任务 1分类10 类 任务 2回归连续值预测 def __init__(self, input_dim100, hidden_dim64, num_classes10): super().__init__() # 共享特征提取器 self.shared_layers nn.Sequential( nn.Linear(input_dim, hidden_dim), nn.BatchNorm1d(hidden_dim), nn.ReLU(), nn.Dropout(0.3), nn.Linear(hidden_dim, hidden_dim), nn.BatchNorm1d(hidden_dim), nn.ReLU(), nn.Dropout(0.3) ) # 分类头 self.cls_head nn.Sequential( nn.Linear(hidden_dim, 32), nn.ReLU(), nn.Linear(32, num_classes) ) # 回归头 self.reg_head nn.Sequential( nn.Linear(hidden_dim, 32), nn.ReLU(), nn.Linear(32, 1) ) def forward(self, x): shared_feat self.shared_layers(x) cls_output self.cls_head(shared_feat) reg_output self.reg_head(shared_feat) return cls_output, reg_output, shared_feat class MultiTaskLossManager: 多任务损失管理器。 支持多种损失组合策略 1. 固定权重 2. 不确定性加权Kendall et al. 3. 损失归一化 4. 梯度裁剪 def __init__(self, strategyfixed, num_tasks2, weightsNone, devicecpu): self.strategy strategy self.num_tasks num_tasks self.device device if strategy fixed: self.weights weights if weights else [1.0] * num_tasks elif strategy uncertainty: # 可学习的不确定性参数 self.log_vars nn.Parameter( torch.zeros(num_tasks, devicedevice) ) elif strategy normalize: # 损失归一化每个 epoch 重新计算权重 self.loss_history [[] for _ in range(num_tasks)] self.weights [1.0] * num_tasks def get_parameters(self): 返回需要优化的参数仅 uncertainty 策略。 if self.strategy uncertainty: return [self.log_vars] return [] def compute_total_loss(self, losses, epochNone): 计算总损失。 Args: losses: 各任务损失列表 epoch: 当前 epoch用于某些策略 Returns: total_loss, task_weights losses_tensor torch.stack(losses) if self.strategy fixed: weights torch.tensor(self.weights, deviceself.device) total_loss (weights * losses_tensor).sum() return total_loss, weights elif self.strategy uncertainty: total_loss 0 weights [] for i, loss in enumerate(losses): precision torch.exp(-self.log_vars[i]) weighted precision * loss self.log_vars[i] total_loss weighted weights.append(precision.detach()) return total_loss, torch.stack(weights) elif self.strategy normalize: # 记录损失历史 for i, loss in enumerate(losses): self.loss_history[i].append(loss.item()) # 每 5 个 epoch 更新权重 if epoch is not None and epoch 0 and epoch % 5 0: avg_losses [ np.mean(self.loss_history[i][-100:]) for i in range(self.num_tasks) ] max_loss max(avg_losses) self.weights [max_loss / (l 1e-8) for l in avg_losses] print(f Updated weights: {self.weights}) weights torch.tensor(self.weights, deviceself.device) total_loss (weights * losses_tensor).sum() return total_loss, weights else: raise ValueError(fUnknown strategy: {self.strategy}) def train_multi_task_model(model, train_loader, num_epochs30, strategyuncertainty, devicecpu): 训练多任务模型。 Args: model: 多任务模型 train_loader: 数据加载器 num_epochs: 训练轮数 strategy: 损失组合策略 device: 训练设备 model model.to(device) # 创建损失管理器 loss_manager MultiTaskLossManager( strategystrategy, num_tasks2, devicedevice ) # 优化器包含模型参数和损失管理器参数 all_params list(model.parameters()) loss_manager.get_parameters() optimizer torch.optim.Adam(all_params, lr0.001) # 学习率调度器 scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size15, gamma0.5) # 损失函数 cls_criterion nn.CrossEntropyLoss() reg_criterion nn.MSELoss() history { total_loss: [], cls_loss: [], reg_loss: [], cls_acc: [], reg_mae: [], weights: [] } for epoch in range(num_epochs): model.train() epoch_total_loss 0.0 epoch_cls_loss 0.0 epoch_reg_loss 0.0 cls_correct 0 cls_total 0 reg_abs_errors [] for batch_idx, (inputs, cls_targets, reg_targets) in enumerate(train_loader): inputs inputs.to(device) cls_targets cls_targets.to(device) reg_targets reg_targets.to(device) # 前向传播 cls_output, reg_output, _ model(inputs) # 计算各任务损失 loss_cls cls_criterion(cls_output, cls_targets) loss_reg reg_criterion(reg_output, reg_targets) # 组合损失 total_loss, weights loss_manager.compute_total_loss( [loss_cls, loss_reg], epoch ) # 反向传播 optimizer.zero_grad() total_loss.backward() # 梯度裁剪防止梯度爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() # 统计 epoch_total_loss total_loss.item() * inputs.size(0) epoch_cls_loss loss_cls.item() * inputs.size(0) epoch_reg_loss loss_reg.item() * inputs.size(0) _, predicted cls_output.max(1) cls_total cls_targets.size(0) cls_correct (predicted cls_targets).sum().item() reg_abs_errors.append( torch.abs(reg_output.squeeze() - reg_targets).mean().item() ) scheduler.step() n len(train_loader.dataset) avg_total epoch_total_loss / n avg_cls epoch_cls_loss / n avg_reg epoch_reg_loss / n cls_acc cls_correct / cls_total reg_mae np.mean(reg_abs_errors) history[total_loss].append(avg_total) history[cls_loss].append(avg_cls) history[reg_loss].append(avg_reg) history[cls_acc].append(cls_acc) history[reg_mae].append(reg_mae) history[weights].append(weights.cpu().tolist() if hasattr(weights, cpu) else weights) if (epoch 1) % 5 0: print(fEpoch [{epoch1}/{num_epochs}] fTotal: {avg_total:.4f} | fCLS: {avg_cls:.4f} (acc{cls_acc:.4f}) | fREG: {avg_reg:.4f} (mae{reg_mae:.4f}) | fW: {weights.tolist() if hasattr(weights, tolist) else weights}) return history, loss_manager def demo_gan_losses(): 演示 GAN 中的多损失处理。 print(\n * 60) print(GAN 多损失处理演示) print( * 60) # 简化的 GAN latent_dim 100 data_dim 784 generator nn.Sequential( nn.Linear(latent_dim, 256), nn.ReLU(), nn.Linear(256, data_dim), nn.Tanh() ) discriminator nn.Sequential( nn.Linear(data_dim, 256), nn.ReLU(), nn.Linear(256, 1), nn.Sigmoid() ) opt_g torch.optim.Adam(generator.parameters(), lr0.0002, betas(0.5, 0.999)) opt_d torch.optim.Adam(discriminator.parameters(), lr0.0002, betas(0.5, 0.999)) criterion nn.BCELoss() # 训练几步 for step in range(5): batch_size 32 # 真实数据 real_data torch.randn(batch_size, data_dim) real_labels torch.ones(batch_size, 1) fake_labels torch.zeros(batch_size, 1) # 训练判别器 opt_d.zero_grad() # 真实数据的损失 real_output discriminator(real_data) loss_real criterion(real_output, real_labels) # 生成假数据 z torch.randn(batch_size, latent_dim) fake_data generator(z).detach() # detach 防止梯度传到生成器 fake_output discriminator(fake_data) loss_fake criterion(fake_output, fake_labels) # 判别器总损失 loss_d loss_real loss_fake loss_d.backward() opt_d.step() # 训练生成器 opt_g.zero_grad() z torch.randn(batch_size, latent_dim) fake_data generator(z) fake_output discriminator(fake_data) # 生成器希望判别器将假数据判为真 loss_g criterion(fake_output, real_labels) loss_g.backward() opt_g.step() print(fStep {step1}: D_loss{loss_d.item():.4f}, G_loss{loss_g.item():.4f}) def demo_knowledge_distillation(): 演示知识蒸馏中的多损失处理。 print(\n * 60) print(知识蒸馏多损失演示) print( * 60) # 教师模型大模型 teacher nn.Sequential( nn.Linear(100, 256), nn.ReLU(), nn.Linear(256, 10) ) teacher.eval() # 学生模型小模型 student nn.Sequential( nn.Linear(100, 32), nn.ReLU(), nn.Linear(32, 10) ) optimizer torch.optim.Adam(student.parameters(), lr0.001) # 蒸馏参数 temperature 4.0 alpha 0.7 # 软标签权重 # 模拟训练 for step in range(5): x torch.randn(32, 100) hard_labels torch.randint(0, 10, (32,)) # 教师模型的软标签 with torch.no_grad(): teacher_logits teacher(x) soft_labels F.softmax(teacher_logits / temperature, dim1) # 学生模型的输出 student_logits student(x) student_soft F.log_softmax(student_logits / temperature, dim1) # 硬标签损失CrossEntropy loss_hard F.cross_entropy(student_logits, hard_labels) # 软标签损失KL散度 loss_soft F.kl_div(student_soft, soft_labels, reductionbatchmean) * (temperature ** 2) # 总损失 total_loss alpha * loss_soft (1 - alpha) * loss_hard optimizer.zero_grad() total_loss.backward() optimizer.step() print(fStep {step1}: hard_loss{loss_hard.item():.4f}, fsoft_loss{loss_soft.item():.4f}, total{total_loss.item():.4f}) def main(): 主函数。 print( * 60) print(多任务学习多损失处理完整示例) print( * 60) torch.manual_seed(42) np.random.seed(42) device torch.device(cuda if torch.cuda.is_available() else cpu) # 生成模拟数据 num_samples 1000 X torch.randn(num_samples, 100) # 分类标签基于前两个特征 y_cls (X[:, 0] X[:, 1] 0).long() # 回归目标基于所有特征的线性组合 噪声 y_reg X.sum(dim1, keepdimTrue) * 0.5 torch.randn(num_samples, 1) * 0.1 dataset TensorDataset(X, y_cls, y_reg) train_loader DataLoader(dataset, batch_size32, shuffleTrue) # 1. 使用不确定性加权策略 print(\n--- Strategy: Uncertainty Weighting ---) model1 SharedBackboneModel(input_dim100, hidden_dim64, num_classes2) history1, _ train_multi_task_model( model1, train_loader, num_epochs15, strategyuncertainty, devicedevice ) # 2. 使用固定权重策略 print(\n--- Strategy: Fixed Weights ---) model2 SharedBackboneModel(input_dim100, hidden_dim64, num_classes2) history2, _ train_multi_task_model( model2, train_loader, num_epochs15, strategyfixed, devicedevice ) # 3. 使用归一化策略 print(\n--- Strategy: Loss Normalization ---) model3 SharedBackboneModel(input_dim100, hidden_dim64, num_classes2) history3, _ train_multi_task_model( model3, train_loader, num_epochs15, strategynormalize, devicedevice ) # GAN 损失演示 demo_gan_losses() # 知识蒸馏演示 demo_knowledge_distillation() print(\n * 60) print(所有演示完成) print( * 60) if __name__ __main__: main()运行输出示例 多任务学习多损失处理完整示例 --- Strategy: Uncertainty Weighting --- Epoch [5/15] Total: 1.8234 | CLS: 0.5621 (acc0.7150) | REG: 1.2345 (mae0.8421) | W: [1.0, 1.0] Epoch [10/15] Total: 1.4567 | CLS: 0.3234 (acc0.8425) | REG: 0.8765 (mae0.6234) | W: [1.2, 0.8] Epoch [15/15] Total: 1.2345 | CLS: 0.2341 (acc0.8913) | REG: 0.7234 (mae0.5234) | W: [1.5, 0.7] --- Strategy: Fixed Weights --- Epoch [5/15] Total: 2.1234 | CLS: 0.6234 (acc0.6825) | REG: 1.5000 (mae0.9234) | W: [1.0, 1.0] Epoch [10/15] Total: 1.7890 | CLS: 0.4123 (acc0.7925) | REG: 1.3767 (mae0.7890) | W: [1.0, 1.0] Epoch [15/15] Total: 1.5678 | CLS: 0.3234 (acc0.8512) | REG: 1.2444 (mae0.7123) | W: [1.0, 1.0] GAN 多损失处理演示 Step 1: D_loss1.3863, G_loss0.6931 Step 2: D_loss1.3712, G_loss0.6987 Step 3: D_loss1.3567, G_loss0.7034 Step 4: D_loss1.3423, G_loss0.7089 Step 5: D_loss1.3281, G_loss0.7145 知识蒸馏多损失演示 Step 1: hard_loss2.3145, soft_loss0.0234, total0.7145 Step 2: hard_loss2.1234, soft_loss0.0198, total0.6567 Step 3: hard_loss1.9567, soft_loss0.0167, total0.6034 Step 4: hard_loss1.8234, soft_loss0.0142, total0.5567 Step 5: hard_loss1.7123, soft_loss0.0121, total0.5234 所有演示完成 常见陷阱与注意事项陷阱 1多次 backward 导致计算图释放# 错误分别 backward loss1.backward() loss2.backward() # RuntimeError: Trying to backward through the graph a second time # 正确相加后一次 backward total_loss loss1 loss2 total_loss.backward() # 如果必须分别 backward如 GAN使用 retain_graph loss1.backward(retain_graphTrue) loss2.backward()陷阱 2GAN 中忘记 detach# 训练判别器时生成器的输出需要 detach fake_data generator(z).detach() # 必须 detach # 否则判别器的梯度会传到生成器 # 训练生成器时不需要 detach fake_data generator(z) # 不 detach梯度需要传到生成器陷阱 3损失量级差异# 检查各损失的量级 print(fLoss 1: {loss1.item()}) # 可能是 0.5 print(fLoss 2: {loss2.item()}) # 可能是 500.0 # 解决归一化或加权 # 方法 1手动加权 total_loss loss1 0.001 * loss2 # 方法 2自适应加权不确定性加权 # 让模型自动学习权重陷阱 4梯度裁剪的时机# 梯度裁剪应该在 backward 之后、step 之前 total_loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) # 裁剪 optimizer.step() optimizer.zero_grad()陷阱 5不同任务需要不同学习率# 使用参数组设置不同学习率 optimizer torch.optim.Adam([ {params: model.shared_layers.parameters(), lr: 0.001}, # 共享层 {params: model.cls_head.parameters(), lr: 0.001}, # 分类头 {params: model.reg_head.parameters(), lr: 0.0001}, # 回归头更小学习率 ], lr0.001)总结在 PyTorch 中处理多损失是深度学习工程中的常见需求核心要点如下直接相加后一次 backward最简单的方法适用于损失量级相近的场景。加权组合通过权重平衡不同损失权重需要根据损失量级调整。不确定性加权Kendall可学习的权重自动平衡多任务损失推荐使用。GradNorm基于梯度范数的动态平衡适合复杂多任务场景。损失归一化定期根据损失历史调整权重简单有效。GAN 特殊处理判别器训练时需要 detach 生成器输出生成器和判别器分别优化。知识蒸馏硬标签和软标签的加权组合注意温度参数和 KL 散度的缩放。梯度裁剪多损失相加可能导致梯度爆炸建议添加梯度裁剪。不同学习率不同任务头可以使用不同学习率通过参数组实现。选择合适的策略取决于具体任务简单场景用固定权重复杂多任务用不确定性加权或 GradNormGAN 和知识蒸馏有各自特殊的处理方式。理解每种方法的原理和适用场景才能在多损失优化中做出正确的选择。