Bases: Module, Module
Container module with an encoder, a recurrent or transformer module, and a decoder.
Copied from: https://github.com/pytorch/examples/blob/main/word_language_model/model.py
Source code in meds_torch/input_encoder/triplet_encoder.py
| class TripletEncoder(nn.Module, Module):
"""Container module with an encoder, a recurrent or transformer module, and a decoder.
Copied from: https://github.com/pytorch/examples/blob/main/word_language_model/model.py
"""
def __init__(self, cfg: DictConfig):
super().__init__()
self.cfg = cfg
# Define Triplet Embedders
self.date_embedder = CVE(cfg)
self.code_embedder = torch.nn.Embedding(cfg.vocab_size, embedding_dim=cfg.token_dim)
self.numeric_value_embedder = CVE(cfg)
def embed_func(self, embedder, x):
out = embedder.forward(x[None, :].transpose(2, 0)).permute(1, 2, 0)
return out
def get_embedding(self, batch):
static_mask = batch["static_mask"]
code = batch["code"]
numeric_value = batch["numeric_value"]
time_delta_days = batch["time_delta_days"]
numeric_value_mask = batch["numeric_value_mask"]
# Embed times and mask static value times
time_emb = self.embed_func(self.date_embedder, time_delta_days) * ~static_mask.unsqueeze(dim=1)
# Embed codes
code_emb = self.code_embedder.forward(code).permute(0, 2, 1)
# Embed numerical values and mask nan values
val_emb = self.embed_func(self.numeric_value_embedder, numeric_value) * numeric_value_mask.unsqueeze(
dim=1
)
# Sum the (time, code, value) triplets and
embedding = time_emb + code_emb + val_emb
assert embedding.isfinite().all(), "Embedding is not finite"
if embedding.shape[-1] > self.cfg.max_seq_len:
raise ValueError(
f"Triplet embedding length {embedding.shape[-1]} "
"is greater than max_seq_len {self.cfg.max_seq_len}"
)
return embedding.transpose(1, 2)
def forward(self, batch):
embedding = self.get_embedding(batch)
batch[INPUT_ENCODER_MASK_KEY] = batch["mask"]
batch[INPUT_ENCODER_TOKENS_KEY] = embedding
return batch
|