diff --git a/chapter10.ipynb b/chapter10.ipynb index 4c660b1..70b1d19 100644 --- a/chapter10.ipynb +++ b/chapter10.ipynb @@ -1156,13 +1156,16 @@ "execution_count": 39 }, { - "metadata": {}, + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T08:27:12.972888743Z", + "start_time": "2026-06-30T08:27:12.847760059Z" + } + }, "cell_type": "code", - "outputs": [], - "execution_count": null, "source": [ - "class EncoderBlock(nn.Moudle):\n", - " def __init(self,key_size,query_size,value_size,num_hiddens,norm_shape,ffn_num_input,ffn_num_hiddens,num_heads,dropout,use_bias=False,**kwargs):\n", + "class EncoderBlock(nn.Module):\n", + " def __init__(self,key_size,query_size,value_size,num_hiddens,norm_shape,ffn_num_input,ffn_num_hiddens,num_heads,dropout,use_bias=False,**kwargs):\n", " super(EncoderBlock,self).__init__(**kwargs)\n", " self.attention = MultiHeadAttention(key_size,query_size,value_size,num_hiddens,num_heads,dropout,bias=use_bias)\n", " self.addnorm1 = AddNorm(norm_shape,dropout)\n", @@ -1171,9 +1174,386 @@ " self.addnorm2 = AddNorm(norm_shape, dropout)\n", " def forward(self, X, valid_lens):\n", " Y = self.addnorm1(X, self.attention(X, X, X, valid_lens))\n", - " return self.addnorm2(Y, self.ffn(Y))" + " return self.addnorm2(Y, self.ffn(Y))\n" ], - "id": "c54dc4e20e9ffb90" + "id": "c54dc4e20e9ffb90", + "outputs": [], + "execution_count": 63 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T08:27:14.277524478Z", + "start_time": "2026-06-30T08:27:14.055208277Z" + } + }, + "cell_type": "code", + "source": [ + "X = torch.ones((2, 100, 24))\n", + "valid_lens = torch.tensor([3, 2])\n", + "encoder_blk = EncoderBlock(24, 24, 24, 24, [100, 24], 24, 48, 8, 0.5)\n", + "encoder_blk.eval()\n", + "encoder_blk(X, valid_lens).shape" + ], + "id": "bcac7f8a8b78a96", + "outputs": [ + { + "data": { + "text/plain": [ + "torch.Size([2, 100, 24])" + ] + }, + "execution_count": 64, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 64 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T08:27:14.839568803Z", + "start_time": "2026-06-30T08:27:14.780931324Z" + } + }, + "cell_type": "code", + "source": [ + "class TransformerEncoder(d2l.Encoder):\n", + " def __init__(self,vocab_size,key_size,query_size,value_size,num_hiddens,norm_shape, ffn_num_input,ffn_num_hiddens,num_head,num_layers,dropout,use_bias=False,**kwargs):\n", + " super().__init__(**kwargs)\n", + " self.num_hiddens = num_hiddens\n", + " self.embedding = nn.Embedding(vocab_size, num_hiddens)\n", + " self.pos_encoding = PositionalEncoding(num_hiddens,dropout)\n", + " self.blks = nn.Sequential()\n", + " for i in range(num_layers):\n", + " self.blks.add_module(\"block\"+str(i),EncoderBlock(key_size,query_size,value_size,num_hiddens,norm_shape,ffn_num_input,ffn_num_hiddens,num_head,dropout,use_bias))\n", + " def forward(self,X,valid_lens,*args):\n", + " X = self.pos_encoding(self.embedding(X)*math.sqrt(self.num_hiddens))\n", + " self.attention_weights = [None] * len(self.blks)\n", + " for i, blk in enumerate(self.blks):\n", + " X = blk(X, valid_lens)\n", + " self.attention_weights[i] = blk.attention.attention.attention_weights\n", + " return X" + ], + "id": "2cae7ead361bd28f", + "outputs": [], + "execution_count": 65 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T08:27:15.891564402Z", + "start_time": "2026-06-30T08:27:15.593169746Z" + } + }, + "cell_type": "code", + "source": [ + "encoder = TransformerEncoder(\n", + "200, 24, 24, 24, 24, [100, 24], 24, 48, 8, 2, 0.5)\n", + "encoder.eval()\n", + "encoder(torch.ones((2, 100), dtype=torch.long), valid_lens).shape" + ], + "id": "c45e0dc170c4da58", + "outputs": [ + { + "data": { + "text/plain": [ + "torch.Size([2, 100, 24])" + ] + }, + "execution_count": 66, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 66 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T08:27:16.188376703Z", + "start_time": "2026-06-30T08:27:16.145494924Z" + } + }, + "cell_type": "code", + "source": [ + "class DecoderBlock(nn.Module):\n", + " \"\"\"解码器中第i个块\"\"\"\n", + " def __init__(self,key_size,query_size,value_size,num_hiddens,norm_shape,ffn_num_input,ffn_num_hiddens,num_heads,dropout,i,**kwargs):\n", + " super().__init__()\n", + " self.i = i\n", + " self.attention1 = MultiHeadAttention(key_size,query_size,value_size,num_hiddens,num_heads,dropout)\n", + " self.addnorm1 = AddNorm(norm_shape, dropout)\n", + " self.attention2 = MultiHeadAttention(\n", + " key_size, query_size, value_size, num_hiddens, num_heads, dropout)\n", + " self.addnorm2 = AddNorm(norm_shape, dropout)\n", + " self.ffn = PositionWiseFFN(ffn_num_input, ffn_num_hiddens,\n", + " num_hiddens)\n", + " self.addnorm3 = AddNorm(norm_shape, dropout)\n", + " def forward(self,X,state):\n", + " enc_outputs,enc_valid_lens=state[0],state[1]\n", + " if state[2][self.i] is None:\n", + " key_values = X\n", + " else:\n", + " key_values = torch.cat((state[2][self.i],X),axis=1)\n", + " state[2][self.i] = key_values\n", + " if self.training:\n", + " batch_size, num_steps, _ = X.shape\n", + " # dec_valid_lens的开头:(batch_size,num_steps),\n", + " # 其中每一行是[1,2,...,num_steps]\n", + " dec_valid_lens = torch.arange(\n", + " 1, num_steps + 1, device=X.device).repeat(batch_size, 1)\n", + " else:\n", + " dec_valid_lens = None\n", + " # 自注意力\n", + " X2 = self.attention1(X, key_values, key_values, dec_valid_lens)\n", + " Y = self.addnorm1(X, X2)\n", + " # 编码器-解码器注意力。\n", + " # enc_outputs的开头:(batch_size,num_steps,num_hiddens)\n", + " Y2 = self.attention2(Y, enc_outputs, enc_outputs, enc_valid_lens)\n", + " Z = self.addnorm2(Y, Y2)\n", + " return self.addnorm3(Z, self.ffn(Z)), state" + ], + "id": "d04208208c75522", + "outputs": [], + "execution_count": 67 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T08:27:17.425149476Z", + "start_time": "2026-06-30T08:27:16.913474140Z" + } + }, + "cell_type": "code", + "source": [ + "decoder_blk = DecoderBlock(24, 24, 24, 24, [100, 24], 24, 48, 8, 0.5, 0)\n", + "decoder_blk.eval()\n", + "X = torch.ones((2, 100, 24))\n", + "state = [encoder_blk(X, valid_lens), valid_lens, [None]]\n", + "decoder_blk(X, state)[0].shape" + ], + "id": "9da08a05f9a4ba8f", + "outputs": [ + { + "data": { + "text/plain": [ + "torch.Size([2, 100, 24])" + ] + }, + "execution_count": 68, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 68 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T08:44:43.960218851Z", + "start_time": "2026-06-30T08:44:43.906108634Z" + } + }, + "cell_type": "code", + "source": [ + "class TransformerDecoder(d2l.AttentionDecoder):\n", + " def __init__(self, vocab_size, key_size, query_size, value_size,\n", + " num_hiddens, norm_shape, ffn_num_input, ffn_num_hiddens,\n", + " num_heads, num_layers, dropout, **kwargs):\n", + " super(TransformerDecoder, self).__init__(**kwargs)\n", + " self.num_hiddens = num_hiddens\n", + " self.num_layers = num_layers\n", + " self.embedding = nn.Embedding(vocab_size, num_hiddens)\n", + " self.pos_encoding = PositionalEncoding(num_hiddens, dropout)\n", + " self.blks = nn.Sequential()\n", + " for i in range(num_layers):\n", + " self.blks.add_module(\"block\"+str(i),\n", + " DecoderBlock(key_size, query_size, value_size, num_hiddens,\n", + " norm_shape, ffn_num_input, ffn_num_hiddens,\n", + " num_heads, dropout, i))\n", + " self.dense = nn.Linear(num_hiddens, vocab_size)\n", + " def init_state(self, enc_outputs, enc_valid_lens, *args):\n", + " return [enc_outputs, enc_valid_lens, [None] * self.num_layers]\n", + " def forward(self, X, state):\n", + " X = self.pos_encoding(self.embedding(X) * math.sqrt(self.num_hiddens))\n", + " self._attention_weights = [[None] * len(self.blks) for _ in range (2)]\n", + " for i, blk in enumerate(self.blks):\n", + " X, state = blk(X, state)\n", + " # 解码器自注意力权重\n", + " self._attention_weights[0][\n", + " i] = blk.attention1.attention.attention_weights\n", + " # “编码器-解码器”自注意力权重\n", + " self._attention_weights[1][\n", + " i] = blk.attention2.attention.attention_weights\n", + " return self.dense(X), state\n", + " @property\n", + " def attention_weights(self):\n", + " return self._attention_weights" + ], + "id": "852e3f51929e836e", + "outputs": [], + "execution_count": 95 + }, + { + "metadata": {}, + "cell_type": "code", + "source": "", + "id": "ed66867772d2ee57", + "outputs": [], + "execution_count": null + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T09:39:48.129188146Z", + "start_time": "2026-06-30T09:23:47.749767443Z" + } + }, + "cell_type": "code", + "source": [ + "num_hiddens, num_layers, dropout, batch_size, num_steps = 32, 2, 0.1, 64, 10\n", + "lr, num_epochs, device = 0.005, 200, d2l.try_gpu()\n", + "ffn_num_input, ffn_num_hiddens, num_heads = 32, 64, 4\n", + "key_size, query_size, value_size = 32, 32, 32\n", + "norm_shape = [32]\n", + "train_iter, src_vocab, tgt_vocab = d2l.load_data_nmt(batch_size, num_steps)\n", + "encoder = TransformerEncoder(\n", + " len(src_vocab), key_size, query_size, value_size, num_hiddens,\n", + " norm_shape, ffn_num_input, ffn_num_hiddens, num_heads,\n", + " num_layers, dropout)\n", + "decoder = TransformerDecoder(\n", + " len(tgt_vocab), key_size, query_size, value_size, num_hiddens,\n", + " norm_shape, ffn_num_input, ffn_num_hiddens, num_heads,\n", + " num_layers, dropout)\n", + "net = d2l.EncoderDecoder(encoder, decoder)\n", + "d2l.train_seq2seq(net, train_iter, lr, num_epochs, tgt_vocab, device)" + ], + "id": "b3255616345e942d", + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "loss 0.079, 1913.5 tokens/sec on cpu\n" + ] + }, + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-06-30T17:39:48.068546\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 104 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T09:49:14.605254502Z", + "start_time": "2026-06-30T09:49:13.756893757Z" + } + }, + "cell_type": "code", + "source": [ + "engs = ['go .', \"i lost .\", 'he\\'s calm .', 'i\\'m home .']\n", + "fras = ['va !', 'j\\'ai perdu .', 'il est calme .', 'je suis chez moi .']\n", + "for eng, fra in zip(engs, fras):\n", + " translation, dec_attention_weight_seq = d2l.predict_seq2seq(\n", + " net, eng, src_vocab, tgt_vocab, num_steps, device, True)\n", + " print(f'{eng} => {translation}, ',\n", + " f'bleu {d2l.bleu(translation, fra, k=2):.3f}')" + ], + "id": "d39611b6fe62a021", + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "go . => va doucement !, bleu 0.000\n", + "i lost . => je ., bleu 0.000\n", + "he's calm . => il est mouillé ., bleu 0.658\n", + "i'm home . => je suis malade ., bleu 0.512\n" + ] + } + ], + "execution_count": 106 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T09:49:47.809858481Z", + "start_time": "2026-06-30T09:49:47.715028120Z" + } + }, + "cell_type": "code", + "source": [ + "enc_attention_weights = torch.cat(net.encoder.attention_weights, 0).reshape((num_layers, num_heads,\n", + "-1, num_steps))\n", + "enc_attention_weights.shape" + ], + "id": "13b9109f31fba1e4", + "outputs": [ + { + "data": { + "text/plain": [ + "torch.Size([2, 4, 10, 10])" + ] + }, + "execution_count": 107, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 107 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-06-30T09:49:57.124929605Z", + "start_time": "2026-06-30T09:49:56.662979114Z" + } + }, + "cell_type": "code", + "source": [ + "d2l.show_heatmaps(\n", + "enc_attention_weights.cpu(), xlabel='Key positions',\n", + "ylabel='Query positions', titles=['Head %d' % i for i in range(1, 5)],\n", + "figsize=(7, 3.5))" + ], + "id": "2ec3af6002bd8b14", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-06-30T17:49:56.997734\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 108 + }, + { + "metadata": {}, + "cell_type": "code", + "outputs": [], + "execution_count": null, + "source": "", + "id": "69d3473471fa28ea" } ], "metadata": { diff --git a/chapter13.ipynb b/chapter13.ipynb new file mode 100644 index 0000000..d8aff04 --- /dev/null +++ b/chapter13.ipynb @@ -0,0 +1,1080 @@ +{ + "cells": [ + { + "metadata": {}, + "cell_type": "markdown", + "source": "### 前面的13.1 13.2 因为不可抗力事件消失了(保存的时候乱码了)", + "id": "b23af7657da5adfb" + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:15.242189006Z", + "start_time": "2026-07-26T11:57:12.026651705Z" + } + }, + "cell_type": "code", + "source": [ + "\n", + "import torch\n", + "from d2l import torch as d2l" + ], + "id": "b39ef644f6c80c7c", + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "/home/yukun/.conda/envs/nn/lib/python3.11/site-packages/torch/cuda/__init__.py:1007: UserWarning: Can't initialize NVML\n", + " raw_cnt = _raw_device_count_nvml()\n" + ] + } + ], + "execution_count": 1 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:15.585413376Z", + "start_time": "2026-07-26T11:57:15.270700840Z" + } + }, + "cell_type": "code", + "source": [ + "d2l.set_figsize()\n", + "img = d2l.plt.imread('../data/catdog.jpg')\n", + "d2l.plt.imshow(img);" + ], + "id": "4284e525ebcc7e87", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T19:57:15.520588\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 2 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:15.692690288Z", + "start_time": "2026-07-26T11:57:15.588725031Z" + } + }, + "cell_type": "code", + "source": [ + "def box_corner_to_center(boxes):\n", + " \"\"\"从(左上,右下)转换到(中间,宽度,高度)\"\"\"\n", + " x1, y1, x2, y2 = boxes[:, 0], boxes[:, 1], boxes[:, 2], boxes[:, 3]\n", + " cx = (x1 + x2) / 2\n", + " cy = (y1 + y2) / 2\n", + " w = x2 - x1\n", + " h = y2 - y1\n", + " boxes = torch.stack((cx, cy, w, h), axis=-1)\n", + " return boxes\n", + "def box_center_to_corner(boxes):\n", + " \"\"\"从(中间,宽度,高度)转换到(左上,右下)\"\"\"\n", + " cx, cy, w, h = boxes[:, 0], boxes[:, 1], boxes[:, 2], boxes[:, 3]\n", + " x1 = cx - 0.5 * w\n", + " y1 = cy - 0.5 * h\n", + " x2 = cx + 0.5 * w\n", + " y2 = cy + 0.5 * h\n", + " boxes = torch.stack((x1, y1, x2, y2), axis=-1)\n", + " return boxes" + ], + "id": "ee53061c33ad70d7", + "outputs": [], + "execution_count": 3 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:15.705559306Z", + "start_time": "2026-07-26T11:57:15.694772245Z" + } + }, + "cell_type": "code", + "source": "dog_bbox, cat_bbox = [60.0, 45.0, 378.0, 516.0], [400.0, 112.0, 655.0, 493.0]", + "id": "2e45aa7592a4d7b1", + "outputs": [], + "execution_count": 4 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:15.794704197Z", + "start_time": "2026-07-26T11:57:15.707000004Z" + } + }, + "cell_type": "code", + "source": [ + "boxes = torch.tensor((dog_bbox, cat_bbox))\n", + "box_center_to_corner(box_corner_to_center(boxes)) == boxes" + ], + "id": "22f2ce93eb46eca2", + "outputs": [ + { + "data": { + "text/plain": [ + "tensor([[True, True, True, True],\n", + " [True, True, True, True]])" + ] + }, + "execution_count": 5, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 5 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:15.815124864Z", + "start_time": "2026-07-26T11:57:15.799009123Z" + } + }, + "cell_type": "code", + "source": [ + "def bbox_to_rect(bbox, color):\n", + " # 将边界框(左上x,左上y,右下x,右下y)格式转换成matplotlib格式:\n", + " # ((左上x,左上y),宽,高)\n", + " return d2l.plt.Rectangle(\n", + " xy=(bbox[0], bbox[1]), width=bbox[2]-bbox[0], height=bbox[3]-bbox[1],\n", + " fill=False, edgecolor=color, linewidth=2)" + ], + "id": "cc21e1e13103291", + "outputs": [], + "execution_count": 6 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.101228908Z", + "start_time": "2026-07-26T11:57:18.498541915Z" + } + }, + "cell_type": "code", + "source": [ + "fig = d2l.plt.imshow(img)\n", + "fig.axes.add_patch(bbox_to_rect(dog_bbox, 'blue'))\n", + "fig.axes.add_patch(bbox_to_rect(cat_bbox, 'red'))" + ], + "id": "a4f3af406f8b5b1c", + "outputs": [ + { + "data": { + "text/plain": [ + "" + ] + }, + "execution_count": 7, + "metadata": {}, + "output_type": "execute_result" + }, + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T19:57:18.591473\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 7 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.221032646Z", + "start_time": "2026-07-26T11:57:19.137271191Z" + } + }, + "cell_type": "code", + "source": [ + "#@save\n", + "def multibox_prior(data, sizes, ratios):\n", + " \"\"\"生成以每个像素为中心具有不同形状的锚框\"\"\"\n", + " in_height, in_width = data.shape[-2:]\n", + " device, num_sizes, num_ratios = data.device, len(sizes), len(ratios)\n", + " boxes_per_pixel = (num_sizes + num_ratios - 1)\n", + " size_tensor = torch.tensor(sizes, device=device)\n", + " ratio_tensor = torch.tensor(ratios, device=device)\n", + "\n", + " # 为了将锚点移动到像素的中心,需要设置偏移量。\n", + " # 因为一个像素的高为1且宽为1,我们选择偏移我们的中心0.5\n", + " offset_h, offset_w = 0.5, 0.5\n", + " steps_h = 1.0 / in_height # 在y轴上缩放步长\n", + " steps_w = 1.0 / in_width # 在x轴上缩放步长\n", + "\n", + " # 生成锚框的所有中心点\n", + " center_h = (torch.arange(in_height, device=device) + offset_h) * steps_h\n", + " center_w = (torch.arange(in_width, device=device) + offset_w) * steps_w\n", + " shift_y, shift_x = torch.meshgrid(center_h, center_w, indexing='ij')\n", + " shift_y, shift_x = shift_y.reshape(-1), shift_x.reshape(-1)\n", + "\n", + " # 生成 “boxes_per_pixel” 个高和宽,\n", + " # 之后用于创建锚框的四角坐标(xmin,xmax,ymin,ymax)\n", + " w = torch.cat((size_tensor * torch.sqrt(ratio_tensor[0]),\n", + " sizes[0] * torch.sqrt(ratio_tensor[1:])))\\\n", + " * in_height / in_width # 处理矩形输入\n", + "\n", + " h = torch.cat((size_tensor / torch.sqrt(ratio_tensor[0]),\n", + " sizes[0] / torch.sqrt(ratio_tensor[1:])))\n", + "\n", + " # 除以2来获得半高和半宽\n", + " anchor_manipulations = torch.stack((-w, -h, w, h)).T.repeat(\n", + " in_height * in_width, 1) / 2\n", + "\n", + " # 每个中心点都将有 “boxes_per_pixel” 个锚框,\n", + " # 所以生成含所有锚框中心的网格,重复了 “boxes_per_pixel” 次\n", + " out_grid = torch.stack([shift_x, shift_y, shift_x, shift_y],\n", + " dim=1).repeat_interleave(boxes_per_pixel, dim=0)\n", + " output = out_grid + anchor_manipulations\n", + " return output.unsqueeze(0)" + ], + "id": "874f386629d410fb", + "outputs": [], + "execution_count": 8 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.258339667Z", + "start_time": "2026-07-26T11:57:19.246797120Z" + } + }, + "cell_type": "code", + "source": [ + "img = d2l.plt.imread('../data/catdog.jpg')\n", + "h, w = img.shape[:2]" + ], + "id": "df3d1ba4aeb6c2f7", + "outputs": [], + "execution_count": 9 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.382932850Z", + "start_time": "2026-07-26T11:57:19.281193532Z" + } + }, + "cell_type": "code", + "source": [ + "print(h, w)\n", + "X = torch.rand(size=(1, 3, h, w))\n", + "Y = multibox_prior(X, sizes=[0.75, 0.5, 0.25], ratios=[1, 2, 0.5])\n", + "Y.shape" + ], + "id": "e8737b3210446831", + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "561 728\n" + ] + }, + { + "data": { + "text/plain": [ + "torch.Size([1, 2042040, 4])" + ] + }, + "execution_count": 10, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 10 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.433949316Z", + "start_time": "2026-07-26T11:57:19.384210708Z" + } + }, + "cell_type": "code", + "source": [ + "def show_bboxes(axes, bboxes, labels=None, colors=None):\n", + " \"\"\"显示所有边界框\"\"\"\n", + " def _make_list(obj, default_values=None):\n", + " if obj is None:\n", + " obj = default_values\n", + " elif not isinstance(obj, (list, tuple)):\n", + " obj = [obj]\n", + " return obj\n", + "\n", + " labels = _make_list(labels)\n", + " colors = _make_list(colors, ['b', 'g', 'r', 'm', 'c'])\n", + " for i, bbox in enumerate(bboxes):\n", + " color = colors[i % len(colors)]\n", + " rect = d2l.bbox_to_rect(bbox.detach().numpy(), color)\n", + " axes.add_patch(rect)\n", + " if labels and len(labels) > i:\n", + " text_color = 'k' if color == 'w' else 'w'\n", + " axes.text(rect.xy[0], rect.xy[1], labels[i],\n", + " va='center', ha='center', fontsize=9, color=text_color,\n", + " bbox=dict(facecolor=color, lw=0))" + ], + "id": "e7070cbdb1d68616", + "outputs": [], + "execution_count": 11 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.483140764Z", + "start_time": "2026-07-26T11:57:19.435383512Z" + } + }, + "cell_type": "code", + "source": "boxes = Y.reshape(h, w, 5, 4)", + "id": "4cc20a550931aef7", + "outputs": [], + "execution_count": 12 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.621504492Z", + "start_time": "2026-07-26T11:57:19.484527807Z" + } + }, + "cell_type": "code", + "source": [ + "bbox_scale = torch.tensor((w, h, w, h))\n", + "fig = d2l.plt.imshow(img)\n", + "show_bboxes(fig.axes, boxes[250, 250, :, :] * bbox_scale,\n", + " ['s=0.75, r=1', 's=0.5, r=1', 's=0.25, r=1', 's=0.75, r=2',\n", + " 's=0.75, r=0.5'])" + ], + "id": "2daf2903aa5b1516", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T19:57:19.568775\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 13 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.688402612Z", + "start_time": "2026-07-26T11:57:19.637442252Z" + } + }, + "cell_type": "code", + "source": [ + "def box_iou(boxes1, boxes2):\n", + " \"\"\"计算两个锚框或边界框列表中成对的交并比\"\"\"\n", + " box_area = lambda boxes: ((boxes[:, 2] - boxes[:, 0]) *\n", + " (boxes[:, 3] - boxes[:, 1]))\n", + " # boxes1,boxes2,areas1,areas2的形状:\n", + " # boxes1:(boxes1的数量,4),\n", + " # boxes2:(boxes2的数量,4),\n", + " # areas1:(boxes1的数量,),\n", + " # areas2:(boxes2的数量,)\n", + " areas1 = box_area(boxes1)\n", + " areas2 = box_area(boxes2)\n", + " # inter_upperlefts,inter_lowerrights,inters的形状:\n", + " # (boxes1的数量,boxes2的数量,2)\n", + " inter_upperlefts = torch.max(boxes1[:, None, :2], boxes2[:, :2])\n", + " inter_lowerrights = torch.min(boxes1[:, None, 2:], boxes2[:, 2:])\n", + " inters = (inter_lowerrights - inter_upperlefts).clamp(min=0)\n", + " # inter_areasandunion_areas的形状:(boxes1的数量,boxes2的数量)\n", + " inter_areas = inters[:, :, 0] * inters[:, :, 1]\n", + " union_areas = areas1[:, None] + areas2 - inter_areas\n", + " return inter_areas / union_areas" + ], + "id": "ffe3178858e5a9b1", + "outputs": [], + "execution_count": 14 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.741523163Z", + "start_time": "2026-07-26T11:57:19.690408939Z" + } + }, + "cell_type": "code", + "source": [ + "#@save\n", + "def assign_anchor_to_bbox(ground_truth, anchors, device, iou_threshold=0.5):\n", + " \"\"\"将最接近的真实边界框分配给锚框\"\"\"\n", + " num_anchors, num_gt_boxes = anchors.shape[0], ground_truth.shape[0]\n", + " # 位于第i行和第j列的元素x_ij是锚框i和真实边界框j的IoU\n", + " jaccard = box_iou(anchors, ground_truth)\n", + " # 对于每个锚框,分配的真实边界框的张量\n", + " anchors_bbox_map = torch.full((num_anchors,), -1, dtype=torch.long,\n", + " device=device)\n", + " # 根据阈值,决定是否分配真实边界框\n", + " max_ious, indices = torch.max(jaccard, dim=1)\n", + " anc_i = torch.nonzero(max_ious >= iou_threshold).reshape(-1)\n", + " box_j = indices[max_ious >= iou_threshold]\n", + " anchors_bbox_map[anc_i] = box_j\n", + " col_discard = torch.full((num_anchors,), -1)\n", + " row_discard = torch.full((num_gt_boxes,), -1)\n", + " for _ in range(num_gt_boxes):\n", + " max_idx = torch.argmax(jaccard)\n", + " box_idx = (max_idx % num_gt_boxes).long()\n", + " anc_idx = (max_idx / num_gt_boxes).long()\n", + " anchors_bbox_map[anc_idx] = box_idx\n", + " jaccard[:, box_idx] = col_discard\n", + " jaccard[anc_idx, :] = row_discard\n", + " return anchors_bbox_map\n" + ], + "id": "2ebf82c1768e24c6", + "outputs": [], + "execution_count": 15 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.791841369Z", + "start_time": "2026-07-26T11:57:19.743006544Z" + } + }, + "cell_type": "code", + "source": [ + "def offset_boxes(anchors, assigned_bb, eps=1e-6):\n", + " \"\"\"对锚框偏移量的转换\"\"\"\n", + " c_anc = d2l.box_corner_to_center(anchors)\n", + " c_assigned_bb = d2l.box_corner_to_center(assigned_bb)\n", + " offset_xy = 10 * (c_assigned_bb[:, :2] - c_anc[:, :2]) / c_anc[:, 2:]\n", + " offset_wh = 5 * torch.log(eps + c_assigned_bb[:, 2:] / c_anc[:, 2:])\n", + " offset = torch.cat([offset_xy, offset_wh], axis=1)\n", + " return offset" + ], + "id": "b83220e939515b39", + "outputs": [], + "execution_count": 16 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.844930037Z", + "start_time": "2026-07-26T11:57:19.792995780Z" + } + }, + "cell_type": "code", + "source": [ + "def multibox_target(anchors, labels):\n", + " \"\"\"使用真实边界框标记锚框\"\"\"\n", + " batch_size, anchors = labels.shape[0], anchors.squeeze(0) # 问题:这里赋值给 anchors 会覆盖,但原代码就是这样\n", + " batch_offset, batch_mask, batch_class_labels = [], [], []\n", + " device, num_anchors = anchors.device, anchors.shape[0]\n", + " for i in range(batch_size):\n", + " label = labels[i, :, :]\n", + " anchors_bbox_map = assign_anchor_to_bbox(\n", + " label[:, 1:], anchors, device)\n", + " bbox_mask = ((anchors_bbox_map >= 0).float().unsqueeze(-1)).repeat(\n", + " 1, 4)\n", + " # 将类标签和分配的边界框坐标初始化为零\n", + " class_labels = torch.zeros(num_anchors, dtype=torch.long,\n", + " device=device)\n", + " assigned_bb = torch.zeros((num_anchors, 4), dtype=torch.float32,\n", + " device=device)\n", + " # 使用真实边界框来标记锚框的类别。\n", + " # 如果一个锚框没有被分配,标记其为背景(值为零)\n", + " indices_true = torch.nonzero(anchors_bbox_map >= 0)\n", + " bb_idx = anchors_bbox_map[indices_true]\n", + " class_labels[indices_true] = label[bb_idx, 0].long() + 1\n", + " assigned_bb[indices_true] = label[bb_idx, 1:]\n", + " # 偏移量转换\n", + " offset = offset_boxes(anchors, assigned_bb) * bbox_mask\n", + " batch_offset.append(offset.reshape(-1))\n", + " batch_mask.append(bbox_mask.reshape(-1))\n", + " batch_class_labels.append(class_labels)\n", + " bbox_offset = torch.stack(batch_offset)\n", + " bbox_mask = torch.stack(batch_mask)\n", + " class_labels = torch.stack(batch_class_labels)\n", + " return (bbox_offset, bbox_mask, class_labels)" + ], + "id": "57e69b1186f23243", + "outputs": [], + "execution_count": 17 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T11:57:19.991066598Z", + "start_time": "2026-07-26T11:57:19.846926644Z" + } + }, + "cell_type": "code", + "source": [ + "ground_truth = torch.tensor([[0, 0.1, 0.08, 0.52, 0.92],\n", + "[1, 0.55, 0.2, 0.9, 0.88]])\n", + "anchors = torch.tensor([[0, 0.1, 0.2, 0.3], [0.15, 0.2, 0.4, 0.4],\n", + "[0.63, 0.05, 0.88, 0.98], [0.66, 0.45, 0.8, 0.8],\n", + "[0.57, 0.3, 0.92, 0.9]])\n", + "fig = d2l.plt.imshow(img)\n", + "show_bboxes(fig.axes, ground_truth[:, 1:] * bbox_scale, ['dog', 'cat'], 'k')\n", + "show_bboxes(fig.axes, anchors * bbox_scale, ['0', '1', '2', '3', '4']);" + ], + "id": "260673166da2aa71", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T19:57:19.937795\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 18 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T12:01:16.948464226Z", + "start_time": "2026-07-26T12:01:16.911551343Z" + } + }, + "cell_type": "code", + "source": [ + "labels = multibox_target(anchors.unsqueeze(dim=0),\n", + " ground_truth.unsqueeze(dim=0))" + ], + "id": "aafe21dd7b10dde3", + "outputs": [], + "execution_count": 19 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T12:06:54.735214956Z", + "start_time": "2026-07-26T12:06:54.681936998Z" + } + }, + "cell_type": "code", + "source": [ + "def offset_inverse(anchors, offset_preds):\n", + " \"\"\"根据带有预测偏移量的锚框来预测边界框\"\"\"\n", + " anc = d2l.box_corner_to_center(anchors)\n", + " pred_bbox_xy = (offset_preds[:, :2] * anc[:, 2:] / 10) + anc[:, :2]\n", + " pred_bbox_wh = torch.exp(offset_preds[:, 2:] / 5) * anc[:, 2:]\n", + " pred_bbox = torch.cat((pred_bbox_xy, pred_bbox_wh), axis=1)\n", + " predicted_bbox = d2l.box_center_to_corner(pred_bbox)\n", + " return predicted_bbox" + ], + "id": "f15edcb302b1f95e", + "outputs": [], + "execution_count": 21 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:27:09.970459950Z", + "start_time": "2026-07-26T13:27:09.946476638Z" + } + }, + "cell_type": "code", + "source": [ + "def nms(boxes, scores, iou_threshold):\n", + " \"\"\"对预测边界框的置信度进行排序\"\"\"\n", + " B = torch.argsort(scores, dim=-1, descending=True)\n", + " keep = [] # 保留预测边界框的指标\n", + " while B.numel() > 0:\n", + " i = B[0]\n", + " keep.append(i)\n", + " if B.numel() == 1: break\n", + " iou = box_iou(boxes[i, :].reshape(-1, 4),\n", + " boxes[B[1:], :].reshape(-1, 4)).reshape(-1)\n", + " inds = torch.nonzero(iou <= iou_threshold).reshape(-1)\n", + " B = B[inds + 1]\n", + " return torch.tensor(keep, device=boxes.device)" + ], + "id": "7dfcdd54114d8a06", + "outputs": [], + "execution_count": 30 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:27:10.364513632Z", + "start_time": "2026-07-26T13:27:10.338997724Z" + } + }, + "cell_type": "code", + "source": [ + "#@save\n", + "def multibox_detection(cls_probs, offset_preds, anchors, nms_threshold=0.5,\n", + " pos_threshold=0.009999999):\n", + " \"\"\"使用非极大值抑制来预测边界框\"\"\"\n", + " device, batch_size = cls_probs.device, cls_probs.shape[0]\n", + " anchors = anchors.squeeze(0)\n", + " num_classes, num_anchors = cls_probs.shape[1], cls_probs.shape[2]\n", + " out = []\n", + " for i in range(batch_size):\n", + " cls_prob, offset_pred = cls_probs[i], offset_preds[i].reshape(-1, 4)\n", + " conf, class_id = torch.max(cls_prob[1:], 0)\n", + " predicted_bb = offset_inverse(anchors, offset_pred)\n", + " keep = nms(predicted_bb, conf, nms_threshold)\n", + " # 找到所有的non_keep索引,并将类设置为背景\n", + " all_idx = torch.arange(num_anchors, dtype=torch.long, device=device)\n", + " combined = torch.cat((keep, all_idx))\n", + " uniques, counts = combined.unique(return_counts=True)\n", + " non_keep = uniques[counts == 1]\n", + " all_id_sorted = torch.cat((keep, non_keep))\n", + " class_id[non_keep] = -1\n", + " class_id = class_id[all_id_sorted]\n", + " conf, predicted_bb = conf[all_id_sorted], predicted_bb[all_id_sorted]\n", + " # pos_threshold是一个用于非背景预测的阈值\n", + " below_min_idx = (conf < pos_threshold)\n", + " class_id[below_min_idx] = -1\n", + " conf[below_min_idx] = 1 - conf[below_min_idx]\n", + " pred_info = torch.cat((class_id.unsqueeze(1),\n", + " conf.unsqueeze(1),\n", + " predicted_bb), dim=1)\n", + "\n", + " out.append(pred_info)\n", + " return torch.stack(out)" + ], + "id": "8313f60149168f3d", + "outputs": [], + "execution_count": 31 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:27:10.685183804Z", + "start_time": "2026-07-26T13:27:10.651321744Z" + } + }, + "cell_type": "code", + "source": [ + "anchors = torch.tensor([[0.1, 0.08, 0.52, 0.92], [0.08, 0.2, 0.56, 0.95],\n", + "[0.15, 0.3, 0.62, 0.91], [0.55, 0.2, 0.9, 0.88]])\n", + "offset_preds = torch.tensor([0] * anchors.numel())\n", + "cls_probs = torch.tensor([[0] * 4, # 背景的预测概率\n", + "[0.9, 0.8, 0.7, 0.1], # 狗的预测概率\n", + "[0.1, 0.2, 0.3, 0.9]]) # 猫的预测概率" + ], + "id": "46de8ece74b27875", + "outputs": [], + "execution_count": 32 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:27:11.044810335Z", + "start_time": "2026-07-26T13:27:10.930620323Z" + } + }, + "cell_type": "code", + "source": [ + "fig = d2l.plt.imshow(img)\n", + "show_bboxes(fig.axes, anchors * bbox_scale,\n", + "['dog=0.9', 'dog=0.8', 'dog=0.7', 'cat=0.9'])" + ], + "id": "ed43634b61e3ec27", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:27:10.997118\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 33 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:27:12.775142785Z", + "start_time": "2026-07-26T13:27:12.685045907Z" + } + }, + "cell_type": "code", + "source": [ + "output = multibox_detection(cls_probs.unsqueeze(dim=0),\n", + "offset_preds.unsqueeze(dim=0),\n", + "anchors.unsqueeze(dim=0),\n", + "nms_threshold=0.5)\n", + "output" + ], + "id": "7eb0e18519f4b55f", + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "tensor([0, 3])\n", + "tensor([0, 1, 2, 3])\n" + ] + }, + { + "data": { + "text/plain": [ + "tensor([[[ 0.0000, 0.9000, 0.1000, 0.0800, 0.5200, 0.9200],\n", + " [ 1.0000, 0.9000, 0.5500, 0.2000, 0.9000, 0.8800],\n", + " [-1.0000, 0.8000, 0.0800, 0.2000, 0.5600, 0.9500],\n", + " [-1.0000, 0.7000, 0.1500, 0.3000, 0.6200, 0.9100]]])" + ] + }, + "execution_count": 34, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 34 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:31:29.890705615Z", + "start_time": "2026-07-26T13:31:29.770882766Z" + } + }, + "cell_type": "code", + "source": [ + "fig = d2l.plt.imshow(img)\n", + "for i in output[0].detach().numpy():\n", + " if i[0] == -1:\n", + " continue\n", + " label = ('dog=', 'cat=')[int(i[0])] + str(i[1])\n", + " show_bboxes(fig.axes, [torch.tensor(i[2:]) * bbox_scale], label)" + ], + "id": "fec1c1f8349a44fe", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:31:29.841516\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 35 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:45:58.790473045Z", + "start_time": "2026-07-26T13:45:58.732856164Z" + } + }, + "cell_type": "code", + "source": [ + "img = d2l.plt.imread('../Pictures/1.jpg')\n", + "h, w = img.shape[:2]\n", + "h, w" + ], + "id": "ea166e0233574bf3", + "outputs": [ + { + "data": { + "text/plain": [ + "(640, 640)" + ] + }, + "execution_count": 52, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 52 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:46:00.668782428Z", + "start_time": "2026-07-26T13:46:00.505299012Z" + } + }, + "cell_type": "code", + "source": [ + "def display_anchors(fmap_w, fmap_h, s):\n", + " d2l.set_figsize()\n", + " # 前两个维度上的值不影响输出\n", + " fmap = torch.zeros((1, 10, fmap_h, fmap_w))\n", + " anchors = d2l.multibox_prior(fmap, sizes=s, ratios=[1, 2, 0.5,0.3])\n", + " bbox_scale = torch.tensor((w, h, w, h))\n", + " d2l.show_bboxes(d2l.plt.imshow(img).axes,\n", + " anchors[0] * bbox_scale)\n", + "display_anchors(fmap_w=3, fmap_h=3, s=[0.15,0.2,0.3])" + ], + "id": "41cab7a08738ed58", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:46:00.601996\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 53 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:46:02.600951835Z", + "start_time": "2026-07-26T13:46:02.468950170Z" + } + }, + "cell_type": "code", + "source": "display_anchors(fmap_w=2, fmap_h=2, s=[0.4])", + "id": "6620db5aa275bf52", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:46:02.539171\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 54 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:46:03.940576114Z", + "start_time": "2026-07-26T13:46:03.820486087Z" + } + }, + "cell_type": "code", + "source": "display_anchors(fmap_w=1, fmap_h=1, s=[0.8])", + "id": "8a0c4555b2752c99", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:46:03.881177\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 55 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:48:39.530086370Z", + "start_time": "2026-07-26T13:48:39.490640006Z" + } + }, + "cell_type": "code", + "source": [ + "import torchvision,os\n", + "import pandas as pd\n", + "def read_data_bananas(is_train=True):\n", + " \"\"\"读取香蕉检测数据集中的图像和标签\"\"\"\n", + " data_dir = d2l.download_extract('banana-detection')\n", + " csv_fname = os.path.join(data_dir, 'bananas_train' if is_train\n", + " else 'bananas_val', 'label.csv')\n", + " csv_data = pd.read_csv(csv_fname)\n", + " csv_data = csv_data.set_index('img_name')\n", + " images, targets = [], []\n", + " for img_name, target in csv_data.iterrows():\n", + " images.append(torchvision.io.read_image(\n", + " os.path.join(data_dir, 'bananas_train' if is_train else\n", + " 'bananas_val', 'images', f'{img_name}')))\n", + " # 这里的target包含(类别,左上角x,左上角y,右下角x,右下角y),\n", + " # 其中所有图像都具有相同的香蕉类(索引为0)\n", + " targets.append(list(target))\n", + " return images, torch.tensor(targets).unsqueeze(1) / 256" + ], + "id": "70e7bc1342f90cea", + "outputs": [], + "execution_count": 56 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T13:49:42.340692779Z", + "start_time": "2026-07-26T13:49:42.292033898Z" + } + }, + "cell_type": "code", + "source": [ + "class BananasDataset(torch.utils.data.Dataset):\n", + " \"\"\"一个用于加载香蕉检测数据集的自定义数据集\"\"\"\n", + " def __init__(self, is_train):\n", + " self.features, self.labels = read_data_bananas(is_train)\n", + " print('read ' + str(len(self.features)) + (f' training examples' if\n", + " is_train else f' validation examples'))\n", + " def __getitem__(self, idx):\n", + " return (self.features[idx].float(), self.labels[idx])\n", + " def __len__(self):\n", + " return len(self.features)\n", + "def load_data_bananas(batch_size):\n", + " \"\"\"加载香蕉检测数据集\"\"\"\n", + " train_iter = torch.utils.data.DataLoader(BananasDataset(is_train=True),\n", + " batch_size, shuffle=True)\n", + " val_iter = torch.utils.data.DataLoader(BananasDataset(is_train=False),\n", + " batch_size)\n", + " return train_iter, val_iter" + ], + "id": "1c80340fca8af59a", + "outputs": [], + "execution_count": 57 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T15:34:46.746160081Z", + "start_time": "2026-07-26T15:34:01.758563722Z" + } + }, + "cell_type": "code", + "source": [ + "batch_size, edge_size = 32, 256\n", + "train_iter, _ = load_data_bananas(batch_size)\n", + "batch = next(iter(train_iter))\n", + "batch[0].shape, batch[1].shape" + ], + "id": "e6aacbbedaa732f6", + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Downloading ../data/banana-detection.zip from http://d2l-data.s3-accelerate.amazonaws.com/banana-detection.zip...\n", + "read 1000 training examples\n", + "read 100 validation examples\n" + ] + }, + { + "data": { + "text/plain": [ + "(torch.Size([32, 3, 256, 256]), torch.Size([32, 1, 5]))" + ] + }, + "execution_count": 58, + "metadata": {}, + "output_type": "execute_result" + } + ], + "execution_count": 58 + }, + { + "metadata": { + "ExecuteTime": { + "end_time": "2026-07-26T15:36:52.903923701Z", + "start_time": "2026-07-26T15:36:52.509530881Z" + } + }, + "cell_type": "code", + "source": [ + "imgs = (batch[0][0:10].permute(0, 2, 3, 1)) / 255\n", + "axes = d2l.show_images(imgs, 2, 5, scale=2)\n", + "for ax, label in zip(axes, batch[1][0:10]):\n", + " d2l.show_bboxes(ax, [label[0][1:5] * edge_size], colors=['b'])" + ], + "id": "528cb4f32e152b5", + "outputs": [ + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T23:36:52.688576\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "metadata": {}, + "output_type": "display_data", + "jetTransient": { + "display_id": null + } + } + ], + "execution_count": 60 + }, + { + "metadata": {}, + "cell_type": "code", + "outputs": [], + "execution_count": null, + "source": "", + "id": "eb1131ed49817939" + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 2 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython2", + "version": "2.7.6" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +}