INNER CODE UNIT · Python

compute_iou

castorini/daam · daam/evaluate.py:14

def compute_iou(a: torch.Tensor, b: torch.Tensor) -> float:
    if a.shape[0] != b.shape[0]:
        a = F.interpolate(a.unsqueeze(0).unsqueeze(0).float(), size=b.shape, mode='bicubic').squeeze()
        a[a < 1] = 0
        a[a >= 1] = 1

    intersection = (a * b).sum()
    union = a.sum() + b.sum() - intersection

    return (intersection / (union + 1e-8)).item()


def compute_ioa(a: torch.Tensor, b: torch.Tensor) -> float:
    if a.shape[0] != b.shape[0]:
        a = F.interpolate(a.unsqueeze(0).unsqueeze(0).float(), size=b.shape, mode='bicubic').squeeze()
        a[a < 1] = 0
        a[a >= 1] = 1

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…