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_prompt_encoder.py
| class TripletPromptEncoder(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
# TODO add to config
self.code_embedder = torch.nn.Embedding(cfg.vocab_size + 2, 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):
code = batch["code"]
numeric_value = batch["numeric_value"]
numeric_value_mask = batch["numeric_value_mask"]
# 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 = code_emb + val_emb
assert embedding.isfinite().all(), "Embedding is not finite"
return embedding
def forward(self, batch):
embedding = self.get_embedding(batch)
batch[INPUT_ENCODER_MASK_KEY] = batch["mask"]
batch[INPUT_ENCODER_TOKENS_KEY] = embedding
return batch
|