"""Loss terms: Hungarian-matched classification (cross-entropy) + curve regression (Smooth L1) + endpoint regression. Reasoning-module auxiliary loss is Phase 2. """ from __future__ import annotations import torch import torch.nn as nn import torch.nn.functional as F from models.head import eval_curve from models.matcher import hungarian_match def compute_loss( outputs: dict, targets: dict, sample_ys: torch.Tensor, cls_weight: float = 2.0, reg_weight: float = 5.0, endpoint_weight: float = 1.0, bg_class_weight: float = 0.2, ) -> dict: cls_logits = outputs["cls_logits"] # (B, N, 2) curve_coeffs = outputs["curve_coeffs"] # (B, N, 4) endpoints = outputs["endpoints"] # (B, N, 2) target_xs = targets["target_xs"] target_valid_mask = targets["target_valid_mask"] target_lane_valid = targets["target_lane_valid"] target_endpoints = targets["target_endpoints"] B, N, _ = cls_logits.shape device = cls_logits.device matches = hungarian_match( cls_logits, curve_coeffs, target_xs, target_valid_mask, target_lane_valid, sample_ys, cls_weight=cls_weight, reg_weight=reg_weight, ) cls_targets = torch.zeros(B, N, dtype=torch.long, device=device) reg_losses = [] endpoint_losses = [] for b, (pred_idx, gt_idx) in enumerate(matches): if pred_idx.numel() == 0: continue pred_idx = pred_idx.to(device) gt_idx = gt_idx.to(device) cls_targets[b, pred_idx] = 1 matched_coeffs = curve_coeffs[b, pred_idx] # (M, 4) pred_xs = eval_curve(matched_coeffs, sample_ys.unsqueeze(0).expand(pred_idx.numel(), -1)) # (M, S) gt_xs = target_xs[b, gt_idx] # (M, S) gt_mask = target_valid_mask[b, gt_idx] # (M, S) if gt_mask.any(): reg_loss = F.smooth_l1_loss(pred_xs[gt_mask], gt_xs[gt_mask], reduction="mean") reg_losses.append(reg_loss) pred_endpoints = endpoints[b, pred_idx] # (M, 2) gt_endpoints = target_endpoints[b, gt_idx] # (M, 2) endpoint_losses.append(F.smooth_l1_loss(pred_endpoints, gt_endpoints, reduction="mean")) class_weights = torch.tensor([bg_class_weight, 1.0], device=device) cls_loss = F.cross_entropy(cls_logits.reshape(-1, 2), cls_targets.reshape(-1), weight=class_weights) reg_loss = torch.stack(reg_losses).mean() if reg_losses else torch.zeros((), device=device) endpoint_loss = torch.stack(endpoint_losses).mean() if endpoint_losses else torch.zeros((), device=device) total = cls_weight * cls_loss + reg_weight * reg_loss + endpoint_weight * endpoint_loss return { "total": total, "cls_loss": cls_loss.detach(), "reg_loss": reg_loss.detach(), "endpoint_loss": endpoint_loss.detach(), "num_matched": sum(p.numel() for p, _ in matches), }