Repository navigation
Expand file tree
/
Copy pathtime_embedding.py
More file actions
100 lines (74 loc) · 3.65 KB
/
Copy pathtime_embedding.py
File metadata and controls
100 lines (74 loc) · 3.65 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
94
95
96
97
98
99
100
# adopted from TAMME, MICCAI 2025.
import math
import torch
from torch import nn
from einops import rearrange
from dataclasses import dataclass
@dataclass
class TimePosConfig:
emb_dim: int # token embedding dimension
seq_length: int # max sequence length for standard positional embedding
pos_emb: str # 'temporal', 'positional', 'learnable', or 'none'
n_time_features: int # how many time features per token
temperature: float = 10000 # for temporal sinusoid
def make_pos_emb(seq_len: int, emb_dim: int) -> nn.Parameter:
position = torch.arange(seq_len, dtype=torch.float32).unsqueeze(1) # [S, 1]
div_term = torch.exp(torch.arange(0, emb_dim, 2, dtype=torch.float32) * (-math.log(10000.0) / emb_dim))
pe = torch.zeros(seq_len, emb_dim, dtype=torch.float32)
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
pe = pe.unsqueeze(0) # [1, S, D]
pe = nn.Parameter(pe, requires_grad=False)
return pe
class PosFeatureEmbedding(nn.Module):
def __init__(self, n_features: int, emb_dim: int, temperature: float = 10000.0):
super().__init__()
# self.norm = nn.LayerNorm(n_features, eps=1e-6)
# each feature gets 2 * pos_dim dimensions (sin + cos),
# so emb_dim must be divisible by 2 * n_features
assert emb_dim % (n_features * 2) == 0, \
f"Embedding dimension {emb_dim} must be divisible by 2 * n_features ({2 * n_features})"
pos_dim = emb_dim//(n_features*2)
self.pos_dim = pos_dim
self.n_features = n_features
# frequencies
omega = 1.0/(temperature**(torch.arange(pos_dim, dtype=torch.float32) / pos_dim))
self.register_buffer('omega', omega) # [pos_dim]
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x = self.norm(x) # [B, S, n_features]
# outer product between features and frequencies
# -> [B, S, n_features, pos_dim]
x = torch.einsum('bsf,p->bsfp', x, self.omega)
x_sin = torch.sin(x) # [B, S, n_features, pos_dim]
x_cos = torch.cos(x) # [B, S, n_features, pos_dim]
pos_emb = torch.cat((x_sin, x_cos), dim=-1) # [B, S, n_features, 2 * pos_dim]
# flatten to emb_dim
pos_emb = rearrange(pos_emb, 'b s f p -> b s (f p)') # [B, S, emb_dim]
return pos_emb
class PosEmbedding(nn.Module):
def __init__(self, cfg: TimePosConfig, emb_dim: int):
super().__init__()
self.cfg = cfg
# if cfg.pos_emb == 'learnable':
# self.pos_emb = nn.Parameter(torch.randn(1, cfg.seq_length, emb_dim))
# elif cfg.pos_emb == 'positional':
# self.pos_emb = make_pos_emb(cfg.seq_length, emb_dim)
if cfg.pos_emb == 'temporal':
self.pos_emb = PosFeatureEmbedding(cfg.n_time_features, emb_dim, temperature=cfg.temperature)
else:
raise NotImplementedError
# elif cfg.pos_emb in ('none', None):
# self.pos_emb = nn.Parameter(
# torch.zeros((1, cfg.seq_length, emb_dim), dtype=torch.float32),
# requires_grad=False
# )
# else:
# raise NotImplementedError(f"Unknown pos embedding technique: {cfg.pos_emb}")
def forward(self, time_feats: torch.Tensor) -> torch.Tensor:
# time_feats: [B, S, n_features]
return self.pos_emb(time_feats) # [B, S, emb_dim]
def build_time_features(delta_t: torch.Tensor) -> torch.Tensor: # delta_t in years
if delta_t.dim() == 1:
delta_t = delta_t.unsqueeze(-1) # [B,1]
time_feats = delta_t.unsqueeze(1) # [B,1,1] (batch, seq_len=1, n_features=1)
return time_feats