Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 18 additions & 18 deletions box_embeddings/modules/intersection/gumbel_intersection.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,17 +11,17 @@
def _compute_logaddexp_with_clipping_and_separate_forward(
t1: TBoxTensor, t2: TBoxTensor, intersection_temperature: float
) -> Tuple[torch.Tensor, torch.Tensor]:
t1_data = torch.stack((t1.z, -t1.Z), -2)
t2_data = torch.stack((t2.z, -t2.Z), -2)
t1_data = torch.stack((t1.z.clone(), -t1.Z.clone()), -2)
t2_data = torch.stack((t2.z.clone(), -t2.Z.clone()), -2)
lse = torch.logaddexp(
t1_data / intersection_temperature, t2_data / intersection_temperature
)

z = intersection_temperature * lse[..., 0, :]
Z = -intersection_temperature * lse[..., 1, :]
z = intersection_temperature * lse[..., 0, :].clone()
Z = -intersection_temperature * lse[..., 1, :].clone()

z_value = torch.max(z, torch.max(t1.z, t2.z)) # type: ignore
Z_value = torch.min(Z, torch.min(t1.Z, t2.Z))
z_value = torch.max(z, torch.max(t1.z.clone(), t2.z.clone())) # type: ignore
Z_value = torch.min(Z, torch.min(t1.Z.clone(), t2.Z.clone()))

z_final = (z - z.detach()) + z_value.detach()
Z_final = (Z - Z.detach()) + Z_value.detach()
Expand All @@ -32,32 +32,32 @@ def _compute_logaddexp_with_clipping_and_separate_forward(
def _compute_logaddexp_with_clipping(
t1: TBoxTensor, t2: TBoxTensor, intersection_temperature: float
) -> Tuple[torch.Tensor, torch.Tensor]:
t1_data = torch.stack((t1.z, -t1.Z), -2)
t2_data = torch.stack((t2.z, -t2.Z), -2)
t1_data = torch.stack((t1.z.clone(), -t1.Z.clone()), -2)
t2_data = torch.stack((t2.z.clone(), -t2.Z.clone()), -2)
lse = torch.logaddexp(
t1_data / intersection_temperature, t2_data / intersection_temperature
)

z = intersection_temperature * lse[..., 0, :]
Z = -intersection_temperature * lse[..., 1, :]
z = intersection_temperature * lse[..., 0, :].clone()
Z = -intersection_temperature * lse[..., 1, :].clone()

z_value = torch.max(z, torch.max(t1.z, t2.z)) # type: ignore
Z_value = torch.min(Z, torch.min(t1.Z, t2.Z))
z_value = torch.max(z, torch.max(t1.z.clone(), t2.z.clone())) # type: ignore
Z_value = torch.min(Z, torch.min(t1.Z.clone(), t2.Z.clone()))

return z_value, Z_value


def _compute_logaddexp(
t1: TBoxTensor, t2: TBoxTensor, intersection_temperature: float
) -> Tuple[torch.Tensor, torch.Tensor]:
t1_data = torch.stack((t1.z, -t1.Z), -2)
t2_data = torch.stack((t2.z, -t2.Z), -2)
t1_data = torch.stack((t1.z.clone(), -t1.Z.clone()), -2)
t2_data = torch.stack((t2.z.clone(), -t2.Z.clone()), -2)
lse = torch.logaddexp(
t1_data / intersection_temperature, t2_data / intersection_temperature
)

z = intersection_temperature * lse[..., 0, :]
Z = -intersection_temperature * lse[..., 1, :]
z = intersection_temperature * lse[..., 0, :].clone()
Z = -intersection_temperature * lse[..., 1, :].clone()

return z, Z

Expand Down Expand Up @@ -116,10 +116,10 @@ def gumbel_intersection(
if box_debug_level > 0:
with torch.no_grad(): # type:ignore
assert (
torch.max(t1.z, t2.z) < z
torch.max(t1.z.clone(), t2.z.clone()) < z
), "max(a,b) < beta*log(exp(a/beta) + exp(b/beta)) not holding"
assert (
torch.min(t1.z, t2.z) > Z
torch.min(t1.z.clone(), t2.z.clone()) > Z
), "min(a,b) > -beta*log(exp(-a/beta) + exp(-b/beta)) not holding"

return left.from_zZ(z, Z)
Expand Down