The transformer encoder layer from CLIP.
| 149 | |
| 150 | |
| 151 | class EncoderLayer(nn.Module): |
| 152 | """The transformer encoder layer from CLIP.""" |
| 153 | |
| 154 | def __init__(self, config: CLIPTextConfig): |
| 155 | super().__init__() |
| 156 | self.embed_dim = config.hidden_size |
| 157 | # Add biases to the attention projections |
| 158 | self.self_attn = Attention( |
| 159 | config.hidden_size, config.num_attention_heads, bias=True |
| 160 | ) |
| 161 | self.layer_norm1 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps) |
| 162 | self.mlp = MLP(config) |
| 163 | self.layer_norm2 = nn.LayerNorm(self.embed_dim, eps=config.layer_norm_eps) |
| 164 | |
| 165 | def __call__(self, x: mx.array, mask: Optional[mx.array] = None) -> mx.array: |
| 166 | y = self.layer_norm1(x) |
| 167 | y = self.self_attn(y, y, y, mask) |
| 168 | x = x + y |
| 169 | y = self.layer_norm2(x) |
| 170 | y = self.mlp(y) |
| 171 | return x + y |
| 172 | |
| 173 | |
| 174 | class TextEmbeddings(nn.Module): |