File size: 7,789 Bytes
00801a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
"""
Masked Slice Modeling (MSM) — Self-Supervised Pretraining for ACL-LKNet.

This is the PRIMARY RESEARCH CONTRIBUTION.

Core idea: MRI volumes have ordered slices with anatomical continuity.
We exploit this structure by masking some slices and training the model
to reconstruct their features from the remaining (unmasked) context.

Research question:
    Can a self-supervised objective that explicitly models inter-slice
    anatomical context produce better transferable representations for
    ACL injury detection than established pretraining?

Masking strategies (research axis):
    - random:      Mask 50% of slices uniformly at random
    - contiguous:  Mask contiguous blocks of 3-5 adjacent slices
    - structured:  Preferentially mask central slices (clinically relevant)
    - mixed:       Alternate random and contiguous per batch
"""

import math
import random as py_random
from typing import Tuple, Optional

import torch
import torch.nn as nn
import torch.nn.functional as F


class MaskedSliceModeling(nn.Module):
    """
    Self-supervised pretraining via Masked Slice Modeling.
    
    Architecture:
        1. Encode all slices through the shared backbone → slice features
        2. Replace masked slice features with learnable [MASK] tokens
        3. Add positional encoding (sinusoidal — respects slice ordering)
        4. Pass through lightweight Transformer decoder
        5. Predict the original features of masked slices
    
    Loss: MSE between predicted and actual features of masked slices
    """

    def __init__(
        self,
        feature_dim: int,
        decoder_dim: int = 256,
        decoder_layers: int = 2,
        decoder_heads: int = 4,
        max_slices: int = 48,
        mask_ratio: float = 0.5,
        mask_strategy: str = "random",
    ):
        super().__init__()
        self.feature_dim = feature_dim
        self.decoder_dim = decoder_dim
        self.mask_ratio = mask_ratio
        self.mask_strategy = mask_strategy
        self.max_slices = max_slices

        # Learnable [MASK] token
        self.mask_token = nn.Parameter(torch.zeros(1, 1, feature_dim))
        nn.init.trunc_normal_(self.mask_token, std=0.02)

        # Project encoder features → decoder dimension
        self.encoder_to_decoder = nn.Linear(feature_dim, decoder_dim)

        # Sinusoidal positional encoding (respects spatial ordering of slices)
        self.register_buffer(
            "pos_encoding", self._sinusoidal_encoding(max_slices, decoder_dim)
        )

        # Lightweight Transformer decoder
        decoder_layer = nn.TransformerEncoderLayer(
            d_model=decoder_dim,
            nhead=decoder_heads,
            dim_feedforward=decoder_dim * 4,
            dropout=0.1,
            activation="gelu",
            batch_first=True,
            norm_first=True,
        )
        self.decoder = nn.TransformerEncoder(
            decoder_layer,
            num_layers=decoder_layers,
        )

        # Predict original features from decoded representations
        self.predictor = nn.Sequential(
            nn.LayerNorm(decoder_dim),
            nn.Linear(decoder_dim, feature_dim),
        )

    @staticmethod
    def _sinusoidal_encoding(max_len: int, dim: int) -> torch.Tensor:
        """Generate sinusoidal positional encoding."""
        pe = torch.zeros(max_len, dim)
        position = torch.arange(0, max_len).unsqueeze(1).float()
        div_term = torch.exp(
            torch.arange(0, dim, 2).float() * (-math.log(10000.0) / dim)
        )
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        return pe.unsqueeze(0)  # (1, max_len, dim)

    def generate_mask(
        self, num_slices: int, strategy: Optional[str] = None
    ) -> torch.Tensor:
        """
        Generate a binary mask indicating which slices to mask.
        
        Args:
            num_slices: Number of slices in the volume
            strategy: Override the default masking strategy
            
        Returns:
            mask: (num_slices,) boolean tensor. True = masked (to predict)
        """
        strategy = strategy or self.mask_strategy
        num_mask = max(1, int(num_slices * self.mask_ratio))
        # Always keep at least 2 slices unmasked for context
        num_mask = min(num_mask, num_slices - 2)

        mask = torch.zeros(num_slices, dtype=torch.bool)

        if strategy == "random":
            indices = torch.randperm(num_slices)[:num_mask]
            mask[indices] = True

        elif strategy == "contiguous":
            # Mask contiguous blocks of 3-5 slices
            remaining = num_mask
            while remaining > 0:
                block_size = min(py_random.randint(3, 5), remaining)
                max_start = num_slices - block_size
                if max_start <= 0:
                    start = 0
                else:
                    start = py_random.randint(0, max_start)
                mask[start : start + block_size] = True
                remaining = num_mask - mask.sum().item()

        elif strategy == "structured":
            # Preferentially mask central slices (where ACL is typically visible)
            center = num_slices // 2
            # Create probability distribution peaked at center
            positions = torch.arange(num_slices).float()
            probs = torch.exp(-0.5 * ((positions - center) / (num_slices / 4)) ** 2)
            probs = probs / probs.sum()
            indices = torch.multinomial(probs, num_mask, replacement=False)
            mask[indices] = True

        elif strategy == "mixed":
            # Randomly choose between random and contiguous per call
            sub_strategy = py_random.choice(["random", "contiguous"])
            mask = self.generate_mask(num_slices, strategy=sub_strategy)

        else:
            raise ValueError(f"Unknown mask strategy: {strategy}")

        return mask

    def forward(
        self,
        slice_features: torch.Tensor,
        slice_mask: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        """
        Forward pass for MSM pretraining.
        
        Args:
            slice_features: (B, S, D) — encoded slice features from backbone
            slice_mask: (B, S) — optional padding mask (True = valid)
            
        Returns:
            loss: scalar MSE loss on masked slices
            predictions: (B, S, D) — predicted features for all slices
            mask: (B, S) — boolean mask of which slices were masked
        """
        B, S, D = slice_features.shape

        # Generate masks for each sample in the batch
        masks = torch.stack([self.generate_mask(S) for _ in range(B)])  # (B, S)
        masks = masks.to(slice_features.device)

        # Replace masked positions with [MASK] token
        mask_tokens = self.mask_token.expand(B, S, -1)  # (B, S, D)
        masked_features = slice_features.clone()
        masked_features[masks] = mask_tokens[masks]

        # Project to decoder dimension
        x = self.encoder_to_decoder(masked_features)  # (B, S, decoder_dim)

        # Add positional encoding
        x = x + self.pos_encoding[:, :S, :]

        # Transformer decoder
        x = self.decoder(x)  # (B, S, decoder_dim)

        # Predict original features
        predictions = self.predictor(x)  # (B, S, D)

        # Compute loss only on masked positions
        if masks.any():
            pred_masked = predictions[masks]     # (num_masked, D)
            target_masked = slice_features[masks]  # (num_masked, D)
            loss = F.mse_loss(pred_masked, target_masked)
        else:
            loss = torch.tensor(0.0, device=slice_features.device)

        return loss, predictions, masks