Skip to content

triplet_prompt_encoder

CVE

Bases: Module

Continuous Value Encoder (CVE) module.

Assumes input is a single continuous value, and encodes it as an output_dim size embedding vector.

Source code in meds_torch/input_encoder/triplet_prompt_encoder.py
class CVE(nn.Module):
    """Continuous Value Encoder (CVE) module.

    Assumes input is a single continuous value, and encodes it as an `output_dim` size embedding vector.
    """

    def __init__(self, cfg):
        super().__init__()
        self.layer = nn.Linear(1, cfg.token_dim)

    def forward(self, x):
        return self.layer(x)

TripletPromptEncoder

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