machinelearning:: temporal Related ideas: Transformers Temporal convolutional networks allow early prediction of events in critical care - PMC (nih.gov) Temporal context matters
Summary¶
- Temporal model with TCN used to extract a feature set describing disease trajectory
- Snapshot model with self-supervised ViT to describe image features at a single timepoint
- Recalibration model to align trajectory features and snapshot features
-
ViTs useful for capturing long range dependencies in images, but comparatively data-hungry compared to CNNs¶
Recalibration Model¶
Uses maximum mean discrepancy ( #MMD ) to measure the difference in feature space between two feature embeddings. In this paper, they use MMD as a regularization #loss term to align the snapshot embedding with the trajectory embedding, such that \(L = L_{prediction} + \lambda_1 L_{MMD}\) The #MMD term seems pretty complicated with some terms I don't really understand but by assuming "the data are all represented in the same latent space with Euclidean metric", the MMD loss term becomes $$ L_{MMD} = || \frac{1}{N}\sum_{s=1}^{N}x_{s} - \frac{1}{M} \sum_{t=1}^{M} y_{t}||^2 $$ where \(x\) is the snapshot feature embedding and \(y\) is the temporal embedding, with both \(x_{s}, y_t \in {R}^{512}\). N is the total number of snapshot embeddings across all patients and M is the number of temporal embeddings (# patients probably)
i'm not sure yet why N and M are not equal. Deconstructing this equation, recall that \(x_s\) and \(y_t\) are both 512 feature vectors. So we take the average embdeding across all snapshot embeddings \(\hat x\) and temporal embeddings \(\hat y\) , then examine the difference between \(\hat x\) and \(\hat y\) by taking the L2 norm of their difference vector.
Temporal Model¶
Temporal model uses a hierarchicical self-attention module, using attention at different levels to aggregrate embeddings into a single optimal representation.
Questions at the outset:¶
- what are levels
- what is being fed into the temporal model
- how does the model deal with time between inputs?
TCN¶
Goal of TCN is to gather relationships across time through causal dilated convolutions (I have notes on this in onenote) The causal nature means the computation of an output at timepoint t depends only on timepoints 1...t
The module takes a sequence of images from one patient \(\{x_{1}, x_{2}, ..., x_{t}\}\), feeds each \(x_i\) into a pretrained ResNet to extract 512 features where \(p_i = ResNet(x_i)\) to serve as input to the TCN.
Their module uses 3 dilated causal convolution layers wl dilation factors 1, 2, 4 and kernel size=3.
Hierarchical Attention¶
They insert a multi-head #self_attention block between every two convolution layers of the TCN. This constitutes a single dilated causal convolution layer, which they repeat 3 times.
The goal of the multihead #self_attention block is to weight each input by their relevance to each other
Input features are transformed into query f and key g via 1x1 convolutions
self_attention relies on queries/keys to compute the attention score for each entry. (compare query's similarity to keys after some feature transformation)¶
The query/key for #self_attention is up to the engineer to decide (often they're the same).
Typically we use a linear layer to transform the input into a query and key. In their case (perhaps to reduce the computational complexity since they already have an informative embedding, or maybe to preserve spatial relationships), they just apply a 1x1 Convolution
They also don't use the value vector, nor the typical scaled dot product. $$ A = softmax (f^Tg)$$ in contrast to traditional attention $$ A = softmax(\dfrac{W_q q (W_k k)^T}{\sqrt d})V $$ Rather than learning a matrix \(W^v\) to obtain \(V\), they simply apply another \(1 \times 1\) convolution to \(p\) to obtain \(h\).
Then using a residual connection they finally obtain the output \(o \in \mathbb{R}^{512 \times T}\) of the attention head \(\(o = p + A^Th\)\) where each row in \(o\) is an embedding for input image \(t\).
Snapshot Model¶
Adopt a self-supervised image transformer to extract representations from snapshot images.
Ideas¶
- My current approach seeks to generate a single output from the temporal model. But this approach separately constructs a temporal embedding and snapshot embedding, then simply uses the temporal embedding to
- causal self-attention? analogous to causal convolution where we use masking so that for the \(i'th\) input sequence, we only compute self-attention according to elements \(1...i\)
- this is nice because we inject an inductive bias whereby the representation at timepoint \(i\) could only have arisen as a consequence of inputs \(1...i\).
- self-attention is typically performed to identify the relative importance of one timepoint to another.
- given this understanding, does it make sense to perform any causal censoring?
- This causal censoring isnt done in NLP because order doesn't matter as much - there is no directional relationship (since the word at the end of the sentence can still be used to modify a word at the start of one).
- So the result of MSA is an NxN matrix, where \(A_{i,j}\) = 0 for all \(j > i\)
- would this cause us to over/underweight the early elements in the sequence?
- i need to look into the decoder's implementation of self-attention for this
- maybe i can just yoink the decoder's implementation of masking and use it in the encoder region
- is softmax applied over the entire matrix, or a single row (over each query)
- i think it's likely over each query so that no one word dominates.
- would this cause us to over/underweight the early elements in the sequence?
- Hierarchical MIL attention
- we use MIL attention to learn a WSI-level weighted embedding, and then apply it again on the sequence of inputs to learn sequence-level importance.
- note however that this sequence-level importance is akin to a 'bag of words'.
- we could try imposing an inductive bias, like the attention must monotonically increase with time.
- note however that this sequence-level importance is akin to a 'bag of words'.
- we use MIL attention to learn a WSI-level weighted embedding, and then apply it again on the sequence of inputs to learn sequence-level importance.