INNER CODE UNIT · Python
intersection
castorini/daam · daam/evaluate.py:20
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
intersection = (a * b).sum()
area = a.sum()
return (intersection / (area + 1e-8)).item()