Repository navigation
Expand file tree
/
Copy pathmemory_utils.py
More file actions
93 lines (78 loc) · 3.09 KB
/
Copy pathmemory_utils.py
File metadata and controls
93 lines (78 loc) · 3.09 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
import math
import torch
# @torch.jit.script
def get_similarity(
mk: torch.Tensor,
ms: torch.Tensor,
qk: torch.Tensor,
qe: torch.Tensor,
add_batch_dim: bool = False,
) -> torch.Tensor:
# used for training/inference and memory reading/memory potentiation
# mk: B x CK x [N] - Memory keys
# ms: B x 1 x [N] - Memory shrinkage
# qk: B x CK x [HW/P] - Query keys
# qe: B x CK x [HW/P] - Query selection
# Dimensions in [] are flattened
if add_batch_dim:
mk, ms = mk.unsqueeze(0), ms.unsqueeze(0)
qk, qe = qk.unsqueeze(0), qe.unsqueeze(0)
CK = mk.shape[1]
mk = mk.flatten(start_dim=2)
ms = ms.flatten(start_dim=1).unsqueeze(2) if ms is not None else None
qk = qk.flatten(start_dim=2)
qe = qe.flatten(start_dim=2) if qe is not None else None
if qe is not None:
# See XMem's appendix for derivation
mk = mk.transpose(1, 2)
a_sq = mk.pow(2) @ qe
two_ab = (mk @ (qk * qe)).mul_(2)
b_sq = (qe * qk.pow(2)).sum(1, keepdim=True)
similarity = two_ab.sub_(a_sq).sub_(b_sq)
else:
# similar to STCN if we don't have the selection term
a_sq = mk.pow(2).sum(1).unsqueeze(2)
two_ab = (mk.transpose(1, 2) @ qk).mul_(2)
similarity = two_ab.sub_(a_sq)
if ms is not None:
return similarity.mul_(ms).div_(math.sqrt(CK))
return similarity.div_(math.sqrt(CK)) # B*N*HW
def do_softmax(
similarity: torch.Tensor,
top_k: int | None = None,
inplace: bool = False,
return_usage: bool = False,
) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]:
# normalize similarity with top-k softmax
# similarity: B x N x [HW/P]
# use inplace with care
if top_k is not None:
values, indices = torch.topk(similarity, k=top_k, dim=1)
x_exp = values.exp_()
x_exp /= torch.sum(x_exp, dim=1, keepdim=True)
if inplace:
similarity.zero_().scatter_(1, indices, x_exp) # B*N*HW
affinity = similarity
else:
affinity = torch.zeros_like(similarity).scatter_(1, indices, x_exp) # B*N*HW
else:
maxes = torch.max(similarity, dim=1, keepdim=True)[0]
# exp_() is safe (sub() output is a fresh throwaway tensor); div_() is NOT —
# Exp's backward needs its own output value preserved, so the final
# normalization must stay out-of-place (confirmed by backward-parity test).
x_exp = similarity.sub(maxes).exp_()
x_exp_sum = torch.sum(x_exp, dim=1, keepdim=True)
affinity = x_exp / x_exp_sum
indices = None
if return_usage:
return affinity, affinity.sum(dim=2)
return affinity
def get_affinity(mk: torch.Tensor, ms: torch.Tensor, qk: torch.Tensor, qe: torch.Tensor) -> torch.Tensor:
# shorthand used in training with no top-k
similarity = get_similarity(mk, ms, qk, qe)
return do_softmax(similarity)
def readout(affinity: torch.Tensor, mv: torch.Tensor) -> torch.Tensor:
B, CV, T, H, W = mv.shape
mo = mv.view(B, CV, T * H * W)
mem = torch.bmm(mo, affinity)
return mem.view(B, CV, H, W)