class ViT(d2l.Classifier):
"""Vision transformer."""
def __init__(self, img_size, patch_size, num_hiddens, mlp_num_hiddens,
num_heads, num_blks, emb_dropout, blk_dropout, lr=0.1,
use_bias=False, num_classes=10, num_channels=1, rngs=None):
super().__init__()
self.save_hyperparameters(ignore=['rngs'])
rngs = nnx.Rngs(params=0, dropout=1) if rngs is None else rngs
self.patch_embedding = PatchEmbedding(
img_size, patch_size, num_hiddens, num_channels, rngs=rngs)
self.cls_token = nnx.Param(jnp.zeros((1, 1, num_hiddens)))
num_steps = self.patch_embedding.num_patches + 1 # Add the cls token
# Positional embeddings are learnable, initialized to small noise
self.pos_embedding = nnx.Param(
rngs.params.normal((1, num_steps, num_hiddens)) * 0.02)
self.embedding_dropout = nnx.Dropout(emb_dropout, rngs=rngs)
self.blks = nnx.List([
ViTBlock(num_hiddens, mlp_num_hiddens, num_heads, blk_dropout,
use_bias, rngs=rngs) for _ in range(num_blks)])
self.head = nnx.Sequential(
nnx.LayerNorm(num_hiddens, rngs=rngs),
nnx.Linear(num_hiddens, num_classes, rngs=rngs))
def forward(self, X):
X = self.patch_embedding(X)
X = d2l.concat((jnp.tile(self.cls_token, (X.shape[0], 1, 1)), X), 1)
X = self.embedding_dropout(X + self.pos_embedding)
for blk in self.blks:
X = blk(X)
return self.head(X[:, 0])