mirror of
https://github.com/priyanshujain/notebooks.git
synced 2026-10-02 11:07:08 +00:00
Created using Colab
This commit is contained in:
1 parent
0e7a408260
commit
95e872955f
1 file changed
+38
-5
+38
-5
@@ -2592,16 +2592,46 @@
|
|||||||
"> You should spend up to 10-15 minutes on this exercise.\n",
|
"> You should spend up to 10-15 minutes on this exercise.\n",
|
||||||
"> ```\n",
|
"> ```\n",
|
||||||
"\n",
|
"\n",
|
||||||
"Now, we can put together the attention, MLP and layernorms into a single transformer block. Remember to implement the residual connections correctly!"
|
"Now, we can put together the attention, MLP and layernorms into a single transformer block. Remember to implement the residual connections correctly!\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"torch.Size([2, 4, 768])\n",
|
||||||
|
"torch.Size([12, 768, 64])\n",
|
||||||
|
"torch.Size([12, 768, 64])\n",
|
||||||
|
"torch.Size([12, 768, 64])\n",
|
||||||
|
"Output shape: torch.Size([2, 4, 768]) \n",
|
||||||
|
"\n",
|
||||||
|
"Input shape: torch.Size([1, 35, 768])\n",
|
||||||
|
"torch.Size([1, 35, 768])\n",
|
||||||
|
"torch.Size([12, 768, 64])\n",
|
||||||
|
"torch.Size([12, 768, 64])\n",
|
||||||
|
"torch.Size([12, 768, 64])\n",
|
||||||
|
"Output shape: torch.Size([1, 35, 768])\n",
|
||||||
|
"Reference output shape: torch.Size([1, 35, 768]) \n",
|
||||||
|
"\n",
|
||||||
|
"100.00% of the values are correct\n",
|
||||||
|
"\n"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
"source": [
|
"source": [
|
||||||
"class TransformerBlock(nn.Module):\n",
|
"class TransformerBlock(nn.Module):\n",
|
||||||
" def __init__(self, cfg: Config):\n",
|
" def __init__(self, cfg: Config):\n",
|
||||||
@@ -2613,7 +2643,10 @@
|
|||||||
" self.mlp = MLP(cfg)\n",
|
" self.mlp = MLP(cfg)\n",
|
||||||
"\n",
|
"\n",
|
||||||
" def forward(self, resid_pre: Float[Tensor, \"batch position d_model\"]) -> Float[Tensor, \"batch position d_model\"]:\n",
|
" def forward(self, resid_pre: Float[Tensor, \"batch position d_model\"]) -> Float[Tensor, \"batch position d_model\"]:\n",
|
||||||
" raise NotImplementedError()\n",
|
" l1 = self.ln1.forward(resid_pre)\n",
|
||||||
|
" l2 = self.attn.forward(l1) + resid_pre\n",
|
||||||
|
" l3 = self.ln2.forward(l2)\n",
|
||||||
|
" return self.mlp.forward(l3)+l2\n",
|
||||||
"\n",
|
"\n",
|
||||||
"\n",
|
"\n",
|
||||||
"tests.rand_float_test(TransformerBlock, [2, 4, 768])\n",
|
"tests.rand_float_test(TransformerBlock, [2, 4, 768])\n",
|
||||||
|
|||||||
Reference in new issue
Block a user