Model architecture · E02 · Implementation

Porting a transformer is more than translating layer names

The PyTorch version makes the training loop explicit—and exposes three architectural differences that a line-by-line port can hide.

PyTorch nnPyTorch autograd / optim
Equal head counts can hide different internal widths. Here four Keras heads each carry 32 channels, while four PyTorch heads divide a 32-channel embedding.
Figure 1. The same “four heads”?. Equal head counts can hide different internal widths. Here four Keras heads each carry 32 channels, while four PyTorch heads divide a 32-channel embedding. Source dimensions. Original vector illustration.

Follow the information

From input to outcome

VALID extraction drops the rightmost three columns. Flattening retains token positions, so its 896 input features are tied to this patch count. This is the consistent constructor path, not the mismatched training-cell configuration.

VALID extraction drops the rightmost three columns. Flattening retains token positions, so its 896 input features are tied to this patch count. This is the consistent constructor path, not the mismatched training-cell configuration.
Figure 2. Information flow. Solid arrows carry observations, tensors or artifacts; other routes are explicitly labelled. Signal shapes, matrices and network icons are schematic, not measured samples or literal neuron counts. Open full-size SVG ↗ On narrow screens, scroll the diagram horizontally.

Read this alongside Figure 1: Equal head counts can hide different internal widths. Here four Keras heads each carry 32 channels, while four PyTorch heads divide a 32-channel embedding. The module map and layer-level figures below expand the operations in this route.

Porting a transformer is more than translating layer names: architectureChart image: B × 1 × 32 × 31 → Unfold patches: 28 tokens · VALID → Token projection: 32 channels + positions → Transformer × 4: Four heads · 8 per head → Flatten + MLP: 896 → 128 → 64 → 32 → Class logits: Two outputs. A high-level module map; comparison branches and training details are explained in the article.MODEL ARCHITECTURE / E02 / MODULE MAP01 INPUTChart imageB × 1 × 32 × 3102 MODULEUnfold patches28 tokens · VALID03 MODULEToken projection32 channels + positions04 MODULETransformer × 4Four heads · 8 per head05 MODULEFlatten + MLP896 → 128 → 64 → 3206 OUTPUTClass logitsTwo outputs
Source-grounded module map. Boxes summarize operations, not individual neurons; comparison arms and training paths are detailed below. On a small screen, scroll the diagram horizontally.
Chart image — B × 1 × 32 × 31

The architecture in context

The system we are building

The useful starting point for a framework port is the tensor contract. Here PyTorch receives channels-first images and uses unfold twice. An 8 × 4 stride produces four rows and seven columns of patches: the final three image columns are discarded. That differs from the TensorFlow model’s padded 32-token sequence. This version also uses four transformer blocks and a flattened classification head instead of three blocks and token averaging.

Who does what in the stack

PyTorch nn
Module registration, attention blocks and the classification head.
PyTorch autograd / optim
Loss derivatives, gradient clipping and AdamW updates.

The notebook owns the optimizer loop: AdamW updates follow explicit loss backpropagation, gradient clipping and learning-rate handling. nn.Module owns parameter registration and mode switching; the custom code owns batching, checkpoint decisions and validation. The final training cell is implementation evidence rather than a completed saved training run.

Framework responsibility map. Each row maps a library or custom component to its job; rows are not a sequential inference graph.
Framework responsibility map. Each row maps a library or custom component to its job; rows are not a sequential inference graph. Open full-size SVG ↗

From module map to executable structure

Inside PyTorch patch transformer

The diagram and count use consistent constructor defaults: image 32 × 31, patch 8 × 4, d=32, 4 heads, 4 blocks. The unexecuted training cell is a different, inconsistent configuration.

Layer-level implementation. B denotes batch size; parameter and shape conventions are expanded in the table.
Layer-level implementation. B denotes batch size; parameter and shape conventions are expanded in the table. Open full-size SVG ↗
Layer / tensor / operation ledger
Layer or branchOutput shapeImplementation detail
Unfold without paddingB × 28 × 32BCHW input; 4 × 7 valid patches; three rightmost columns are discarded.
Patch + position embeddingB × 28 × 32Linear 32→32, learned table of 28 positions.
Pre-LN attention × 4B × 28 × 32LayerNorm → MultiheadAttention(embed_dim=32, heads=4) → residual. Each head has width 8.
MLP inside each blockB × 28 × 32LayerNorm → 32→64 GELU → dropout → 64→32 → dropout → residual.
Flattened headB × 2 logitsLayerNorm → flatten 896 → 128 GELU → 64 GELU → 32 GELU → 2; dropout between hidden layers.

PyTorch divides embed_dim across heads. Setting embed_dim=32 and num_heads=4 therefore produces width 8, not the Keras notebook’s width 32 per head. This changes both parameter count and computation. Flattening all 28 tokens yields 896 head inputs, so the first dense head has 114,816 parameters—larger than all four transformer blocks together. The output remains unnormalized because CrossEntropyLoss performs the log-softmax internally.

Pre-LN residual connections at left; one multi-head attention operation expanded at right. These classifiers are not the encoder–decoder model from the 2017 paper.
Pre-LN residual connections at left; one multi-head attention operation expanded at right. These classifiers are not the encoder–decoder model from the 2017 paper. Open full-size SVG ↗

Diagram hierarchy informed by Figures 1–2 of Attention Is All You Need. Original figures here follow the archived classifier code; in particular its pre-LN placement differs from that paper’s post-LN architecture.

The equation and the update

Ah=softmax⁡(QhKhT/8);Y=Concat⁡h(AhVh)WOA_h=\operatorname{softmax}(Q_hK_h^T/\sqrt{8});\quad Y=\operatorname{Concat}_h(A_hV_h)W_O

CrossEntropyLoss consumes logits. AdamW: 5e-5 learning rate, 1e-4 decay; batch 256; up to 200 epochs; five-epoch linear warm-up; plateau factor .8 / patience 15; early-stop patience 40; gradient norm clipping at 1.0. The final training cell has no completed execution evidence.

Learning or solution path. A parameter-update path is different from the forward inference path; see text for target-network, frozen-feature and local-loss boundaries.
Learning or solution path. A parameter-update path is different from the forward inference path; see text for target-network, frozen-feature and local-loss boundaries. Open full-size SVG ↗

Implementation card / no invented benchmarks

Capacity, budget and execution evidence

Parameters / retained state
161,410 scalars for the default configuration, not the final training-cell configuration.
Duration and hardware evidence
MPS-if-available / CPU fallback is coded. No completed timing is claimed for the final training cell.
Source coordinates
E02 cell 2, lines 165–272 and 291–365; cells 4–5
Current reproduction context
Current workstation, supplied by the author: Apple M4, 128 GB unified RAM, 40 GPU cores and 16 CPU cores. This is context for prospective reproduction, not attribution of every archived run. Python and framework versions are not fully locked for these historical sources; declarations, when available, are identified separately.

Counts above are calculated from the stated layer shapes unless identified as saved measurements. They exclude optimizer state and nontrainable buffers. No archived training was rerun for this revision.

What these design choices change

The final cell selects 25 × 31 images, 5 × 4 patches, embedding 128 and 8 heads. However, PatchEmbedding still calculates position count from hard-coded height 32: 42 positions rather than the actual 35. The first 35 positions can be indexed, but seven embedding rows are unused. Do not describe this as a clean architecture port or reuse the default count for that run. A deliberate repair would pass the image dimensions into PatchEmbedding and test the resulting model as a new version.

Reproduction and measurement protocol

Use the same padded input, embedding widths, block count and pooling strategy before attributing a performance difference to a framework. In this archive those factors are confounded. Check the logits shape and label dtype with a single batch, then test one backward step with dropout disabled. Measure loader time separately: a fast attention kernel will not compensate for serial image decoding.

For a new run, save the resolved Python/framework versions, backend, dtype, seed, input shapes, batch size and exact source revision. Start with one batch and one update. Log training steps separately from epochs or environment steps. Do not equate the configured maximum with a completed budget or convergence.

Measure initialization/compilation, data preparation, warmed forward pass, training updates and evaluation separately. Synchronize accelerator work around timed regions using the chosen framework’s supported mechanism. Report peak process memory and framework allocation separately; parameter bytes exclude activations, gradients, optimizer state and input buffers. On a shared machine, begin with a single CPU worker and a small batch rather than claiming all available resources.

A closer look at the implementation

The code that carries the idea

The two unfold operations expose patch axes without a Python loop over patches. permute places those axes before the patch payload, then reshape produces B × tokens × patch_values. The forward method keeps extraction, embedding, repeated attention and classification separate enough to test individually.

Python · cell 2 · lines 252–272
    def extract_patches(self, x):
        B, C, H, W = x.shape
        patches = x.unfold(2, self.patch_height, self.patch_height).\
                   unfold(3, self.patch_width, self.patch_width)
        patches = patches.permute(0, 2, 3, 1, 4, 5)
        patches = patches.reshape(B, -1, C * self.patch_height * self.patch_width)
        return patches

    def forward(self, x):
        # Extract patches
        patches = self.extract_patches(x)
        
        # Encode patches
        x = self.patch_embedding(patches)
        
        # Apply transformer blocks
        for block in self.transformer_blocks:
            x = block(x)
            
        # Classification head
        return self.mlp_head(x)

Verbatim archive excerpt from pt_cnn_transformer_model_lp3_3class.ipynb. Context-dependent historical code, not a standalone runnable program. Comments retain their original wording; the article distinguishes implemented behavior from stale or overbroad comments.

The boundary that matters

PyTorch MultiheadAttention divides embed_dim across its heads. With embed_dim=32 and four heads, each head has eight channels; the Keras example specifies 32 channels per head. Matching the integer “32” in both APIs does not match the attention model.

Keep building

Other posts of interest