Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1import math2import torch3 4 5def diou_loss(6 boxes1: torch.Tensor,7 boxes2: torch.Tensor,8 reduction: str = "none",9 eps: float = 1e-7,10) -> torch.Tensor:11 """12 Distance Intersection over Union Loss (Zhaohui Zheng et. al)13 https://arxiv.org/abs/1911.0828714 Args:15 boxes1, boxes2 (Tensor): box locations in XYXY format, shape (N, 4) or (4,).16 reduction: 'none' | 'mean' | 'sum'17 'none': No reduction will be applied to the output.18 'mean': The output will be averaged.19 'sum': The output will be summed.20 eps (float): small number to prevent division by zero21 """22 23 x1, y1, x2, y2 = boxes1.unbind(dim=-1)24 x1g, y1g, x2g, y2g = boxes2.unbind(dim=-1)25 26 # TODO: use torch._assert_async() when pytorch 1.8 support is dropped27 assert (x2 >= x1).all(), "bad box: x1 larger than x2"28 assert (y2 >= y1).all(), "bad box: y1 larger than y2"29 30 # Intersection keypoints31 xkis1 = torch.max(x1, x1g)32 ykis1 = torch.max(y1, y1g)33 xkis2 = torch.min(x2, x2g)34 ykis2 = torch.min(y2, y2g)35 36 intsct = torch.zeros_like(x1)37 mask = (ykis2 > ykis1) & (xkis2 > xkis1)38 intsct[mask] = (xkis2[mask] - xkis1[mask]) * (ykis2[mask] - ykis1[mask])39 union = (x2 - x1) * (y2 - y1) + (x2g - x1g) * (y2g - y1g) - intsct + eps40 iou = intsct / union41 42 # smallest enclosing box43 xc1 = torch.min(x1, x1g)44 yc1 = torch.min(y1, y1g)45 xc2 = torch.max(x2, x2g)46 yc2 = torch.max(y2, y2g)47 diag_len = ((xc2 - xc1) ** 2) + ((yc2 - yc1) ** 2) + eps48 49 # centers of boxes50 x_p = (x2 + x1) / 251 y_p = (y2 + y1) / 252 x_g = (x1g + x2g) / 253 y_g = (y1g + y2g) / 254 distance = ((x_p - x_g) ** 2) + ((y_p - y_g) ** 2)55 56 # Eqn. (7)57 loss = 1 - iou + (distance / diag_len)58 if reduction == "mean":59 loss = loss.mean() if loss.numel() > 0 else 0.0 * loss.sum()60 elif reduction == "sum":61 loss = loss.sum()62 63 return loss64 65 66def ciou_loss(67 boxes1: torch.Tensor,68 boxes2: torch.Tensor,69 reduction: str = "none",70 eps: float = 1e-7,71) -> torch.Tensor:72 """73 Complete Intersection over Union Loss (Zhaohui Zheng et. al)74 https://arxiv.org/abs/1911.0828775 Args:76 boxes1, boxes2 (Tensor): box locations in XYXY format, shape (N, 4) or (4,).77 reduction: 'none' | 'mean' | 'sum'78 'none': No reduction will be applied to the output.79 'mean': The output will be averaged.80 'sum': The output will be summed.81 eps (float): small number to prevent division by zero82 """83 84 x1, y1, x2, y2 = boxes1.unbind(dim=-1)85 x1g, y1g, x2g, y2g = boxes2.unbind(dim=-1)86 87 # TODO: use torch._assert_async() when pytorch 1.8 support is dropped88 assert (x2 >= x1).all(), "bad box: x1 larger than x2"89 assert (y2 >= y1).all(), "bad box: y1 larger than y2"90 91 # Intersection keypoints92 xkis1 = torch.max(x1, x1g)93 ykis1 = torch.max(y1, y1g)94 xkis2 = torch.min(x2, x2g)95 ykis2 = torch.min(y2, y2g)96 97 intsct = torch.zeros_like(x1)98 mask = (ykis2 > ykis1) & (xkis2 > xkis1)99 intsct[mask] = (xkis2[mask] - xkis1[mask]) * (ykis2[mask] - ykis1[mask])100 union = (x2 - x1) * (y2 - y1) + (x2g - x1g) * (y2g - y1g) - intsct + eps101 iou = intsct / union102 103 # smallest enclosing box104 xc1 = torch.min(x1, x1g)105 yc1 = torch.min(y1, y1g)106 xc2 = torch.max(x2, x2g)107 yc2 = torch.max(y2, y2g)108 diag_len = ((xc2 - xc1) ** 2) + ((yc2 - yc1) ** 2) + eps109 110 # centers of boxes111 x_p = (x2 + x1) / 2112 y_p = (y2 + y1) / 2113 x_g = (x1g + x2g) / 2114 y_g = (y1g + y2g) / 2115 distance = ((x_p - x_g) ** 2) + ((y_p - y_g) ** 2)116 117 # width and height of boxes118 w_pred = x2 - x1119 h_pred = y2 - y1120 w_gt = x2g - x1g121 h_gt = y2g - y1g122 v = (4 / (math.pi**2)) * torch.pow((torch.atan(w_gt / h_gt) - torch.atan(w_pred / h_pred)), 2)123 with torch.no_grad():124 alpha = v / (1 - iou + v + eps)125 126 # Eqn. (10)127 loss = 1 - iou + (distance / diag_len) + alpha * v128 if reduction == "mean":129 loss = loss.mean() if loss.numel() > 0 else 0.0 * loss.sum()130 elif reduction == "sum":131 loss = loss.sum()132 133 return loss134 