AstroCo / astro_model_arch /AstromerPosEmbedding.py
AntonyT1207's picture
rename
7d3e533
Raw
History Blame Contribute Delete
1.16 kB
import torch
import torch.nn as nn
import math
class AstromerPosEmbedding(nn.Module):
def __init__(self, d_model, trainable=False, base=1000, **kwargs):
super(AstromerPosEmbedding, self).__init__()
self.d_model = d_model
initial_div_term = torch.exp(
torch.arange(0.0, d_model).float() * -(math.log(base) / d_model)
)
if trainable:
self.w = nn.Parameter(initial_div_term)
else:
self.register_buffer("w", initial_div_term)
def combine_input_embeddings(self, m_emb, t_emb):
x = m_emb + t_emb
return x
def forward(self, emb_x, t):
"""
t: Tensor de series de tiempo con dimensiones [batch_size, seq_len, 1]
"""
w = self.w.unsqueeze(0).unsqueeze(1)
pe = t * w
# Evitar operaciones inplace
pe_sin = torch.sin(pe[:, :, 0::2])
pe_cos = torch.cos(pe[:, :, 1::2])
pe = torch.cat([pe_sin.unsqueeze(-1), pe_cos.unsqueeze(-1)], dim=-1).reshape(
pe.shape
)
if emb_x == None:
return pe
return self.combine_input_embeddings(emb_x, pe)