Если скрытое состояние представляет из себя суперпозицию, то возведение её в квадрат, например в ReLU², даст суперпозицию попарных взаимодействий элементов.
S² = (a + b + c + d) * (a + b + c + d) = aa + 2ab + 2ac + 2ad + bb + 2bc + 2bd + cc + 2cd + dd
import torch
import torch.nn.functional as F
D = 1024
# Векторы концепций
a, b, c, d = F.normalize(torch.randn(4, D))
# Исходные суперпозиции
S = a + b + c + d
# Квадрат суперпозиции
S **= 2
S = F.normalize(S.unsqueeze(0)).squeeze(0)
# Пары концепций
aa = a * a
bb = b * b
cc = c * c
dd = d * d
ab = a * b
ac = a * c
ad = a * d
bc = b * c
bd = b * d
cd = c * d
z = torch.vstack([aa, bb, cc, dd, ab, ac, ad, bc, bd, cd])
z = F.normalize(z)
# Находим пары концепций в результирующей суперпозиции
r = S @ z.T
print(r)
# [0.50, 0.52, 0.50, 0.53, 0.27, 0.25, 0.30, 0.28, 0.26, 0.29]
S² = (a + b + c + d) * (a + b + c + d) = aa + 2ab + 2ac + 2ad + bb + 2bc + 2bd + cc + 2cd + dd
import torch
import torch.nn.functional as F
D = 1024
# Векторы концепций
a, b, c, d = F.normalize(torch.randn(4, D))
# Исходные суперпозиции
S = a + b + c + d
# Квадрат суперпозиции
S **= 2
S = F.normalize(S.unsqueeze(0)).squeeze(0)
# Пары концепций
aa = a * a
bb = b * b
cc = c * c
dd = d * d
ab = a * b
ac = a * c
ad = a * d
bc = b * c
bd = b * d
cd = c * d
z = torch.vstack([aa, bb, cc, dd, ab, ac, ad, bc, bd, cd])
z = F.normalize(z)
# Находим пары концепций в результирующей суперпозиции
r = S @ z.T
print(r)
# [0.50, 0.52, 0.50, 0.53, 0.27, 0.25, 0.30, 0.28, 0.26, 0.29]