Files
nn/chapter13.ipynb
T

1803 lines
1.9 MiB
Plaintext
Raw Normal View History

2026-07-27 14:18:15 +08:00
{
"cells": [
{
"metadata": {},
"cell_type": "markdown",
"source": "### 前面的13.1 13.2 因为不可抗力事件消失了(保存的时候乱码了)",
"id": "b23af7657da5adfb"
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:55.425519349Z",
"start_time": "2026-07-27T09:47:53.873571522Z"
2026-07-27 14:18:15 +08:00
}
},
"cell_type": "code",
"source": [
"\n",
"import torch\n",
"from d2l import torch as d2l"
],
"id": "b39ef644f6c80c7c",
2026-07-27 17:57:55 +08:00
"outputs": [],
2026-07-27 14:18:15 +08:00
"execution_count": 1
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:55.671442114Z",
"start_time": "2026-07-27T09:47:55.428215388Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"220.346324pt\" height=\"173.353814pt\" viewBox=\"0 0 220.346324 173.353814\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:47:55.619678</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 173.353814 \nL 220.346324 173.353814 \nL 220.346324 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 149.475689 \nL 213.146324 149.475689 \nL 213.146324 10.875689 \nL 33.2875 10.875689 \nz\n\"/>\n </g>\n <g clip-path=\"url(#p5074bdf23f)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAPoAAADBCAYAAADmUa1TAADeOElEQVR4nOz9d7NtWXbdif2W3faY655LUw4FoAoADZpks7vlOhStkKhu8YNI+l5SKKRQM8RQSGJINE0QLBAEYaqAqspKn/nsdcdst5z+WPucd19WFgpoZGYBqLcybt7zzj127z3XnHPMMccUKaXE6/V6vV5/q5f8RX+A1+v1er2+/PXa0F+v1+uXYL029Nfr9folWK8N/fV6vX4J1mtDf71er1+C9drQX6/X65dgvTb01+v1+iVYrw399Xq9fgnWa0N/vV6vX4Klf9Ef4Ktef1kaoPirvvBf+AVer9fry1u/FIaeUuLA9E2z4d21SwEIxE8ZpUifb7+fxxoWn33R+c7PIxgL8fKNfh4DWQjxyuNfr9frf8z6pTD0w0pkw0vip++HV61akA39r+TRD/d99v1etxe8Xl/x+htr6H9ZYzk8XohswHkdLDB97usdNobPrs/zsIeN4bPP/1le/eXn+fyt5C/i9V97+tfrL7r+2hv6n2fQnzWWuyH6Zx93fCyz4c02IviMM07H/70a3s/vIYQghICUr+KYn2vMQPwLfKc/z2BfG/nr9UWsv/aGDn++AUO+6F/Jw1Mixni8L4RACAHvPdM4Mg0jKSWUUkgpMcZgjEEphVIKACllDvGFQEr5imEJIYgxHt/3eP/nGbv4/Kj+taG+Xl/l+mtr6HcNKKV0x2XOPnL2vClBRBJTIkVPCNmY97uO3XbL7eaW280tm82Wy8trnj15xovnz+j3O4xRXJyfo42mqhqapqGqKmxZUVY1RVWxWi05Wa9ZLFrKoqAsC6RSSJl/coEyJ/6zz0fcAfzizwz9D3+QJF5GGxxv5bDjp/C9zwnpX28ar9fPW3+tDP0Vrx1fGrdM2RREisAEcf63UMSU8NEQhCD6gbG75erZFT/605/wzg9/yAcfvcuzy+c8v7ql7wL73uHHnkJ5LtYV5tvf5ONPPmaz82glkAKCLLjuA1VZUxnN22894tu/8iucn665f/8e69M13/7VX6NqWpKEJBNy/k8ImXEAIZFSINPLNOD4/fIDjvcf1tFgX/m7BCGPf78btbzynJ+zXm8Gv9zrr5WhA8SYPbZI6RgKH+7LOa8ipZi94Byix+gJMTD2O549fcrvf+/3+b3f/X0+ePd9TCHZDTtubvdI1aAFCC341W//KotKcv/8BFzHU3HDum1oKsumd+jLLaPb0e0nrgrB9/uO//gHf8C9e/dYtBX/h//j/57v/OZvYMoCpQ1CaiTiWA4Tc8h/WK8Y2h1DfuW738EBjoZ+B8z7vFJbSunnGvFn8YTX65dv/UIN/bN59yGnTikh00/flwAvFCkJSAGRPCkEgu8Yh54ffP/7/LN/9s/43r//HtvbDffWJxS2xUpYNRUhKCbvQEu22y0Xp49Yr5b4/RX7mxsuFobTZYOLiTcv1nRB8cnT57y4eoGUifWq5enTx3SbmssXT9nvHlCmFq1rlLazF5dHIz/cvmugB6/8ecZ5MMi7f49CkOZ05bOv8/PWX6Ze//Oe/2WuP++z/UU/wxdVsvzbGvn8Qgz97km5C2odPPdLr84roWpMiZACgogME8kNhKmn63Z87/d+n//HP///8P4HH7Hf71mtW5ZFomKgqRXF6YoQLbvtjuvtlsl5Jh9YtQ2uspTWYI1iWVuMMYxBsNcL/vS9DxEC/tFv/ybTOPCv/uW/Zhwl+80NMfQQLIiSKBKIkDepOwZ719jhpbH+RYgycDD0fFtKeXzez/PSdy/Yw/H87P0/7/l/EeP5vMf9vOferZLcXT/vdX7eZ/+rGvvfViOHL83Qf0ZB+c5xfMVTf+bngJiLxCvoeUoRUiC4AT/uGbaXvHjyCS+uX/C93/kdbi4vqasGIS6RMtLKiV996xEP3nhAszglpoLrmy2/8wd/zO2UGIaJ89MT0nbBe7ZEmRyGL9qWIsLVDvokOTs94Y37Z9xbVdx8/A7vfjrR7W7p9zc0TZuJNZ9DxIGXG9lnDf3PW6964kiac/TjBsjLsp3IT/jMCwDIzzWcP+/fnwVA/7zPeXzfO5HX4TUPr/K5DMKf8Z4/q6rycyOY9DNu/xyS0t9mo/689eUYegIId/6ZSaak2TvF+TExQoqEA8qcIkSPSJEUBSFCDJ4UA9E7vJ9wqcdtdnSXH3D5+Ec8++QZ6/M3+O3v/CrjbuDZdUdlClSS3DuzvHHR8I1HJ5i6YfCKKDqqKjEl0CSMMRR1TXSBlAJROqqFZmkqNgzcW9WcrpbcOzvhXhX5e99+RL3oWMkd/vYK7n2TJCRKJBICkSIizhfWZ8C1Q+6dDr/vXHx3vfUrF6WIvCz6y3wYAZFe3j4S8D5TAvw8HCDfJ0hSHSOFu/9PIb78t8jIv5DAvNGK+R1TAqQEqeYNSObXky/PeN6Ys7GKJOf9/g44eaw3/Axi0t3vRsZtfuoxdw9VOuAan1+Ovbt+1t//MhvC36TN4ksx9FeLRXMIGrN9x5QgptkgAikGPCJfRMnn+4IjRvDO48aB4Eb82DOOHfvxkrgZGK/e4+bJn1JS8cbFEm02PDgvePLsY87WJWO/52tvv8XDewtq60H0BK24OCv4R7/9q/zwR59wcXrOuim4jAERAt5NJEq0FbSN5bcfvMF2
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
"execution_count": 2
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:55.747558006Z",
"start_time": "2026-07-27T09:47:55.673122195Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:55.853579240Z",
"start_time": "2026-07-27T09:47:55.749428064Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:55.959910400Z",
"start_time": "2026-07-27T09:47:55.859086964Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:57.527570250Z",
"start_time": "2026-07-27T09:47:57.492961332Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:57.635071582Z",
"start_time": "2026-07-27T09:47:57.528293612Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
2026-07-27 17:57:55 +08:00
"<matplotlib.patches.Rectangle at 0x7f64a187bf90>"
2026-07-27 14:18:15 +08:00
]
},
"execution_count": 7,
"metadata": {},
"output_type": "execute_result"
},
{
"data": {
"text/plain": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"220.346324pt\" height=\"173.353814pt\" viewBox=\"0 0 220.346324 173.353814\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:47:57.594042</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 173.353814 \nL 220.346324 173.353814 \nL 220.346324 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 149.475689 \nL 213.146324 149.475689 \nL 213.146324 10.875689 \nL 33.2875 10.875689 \nz\n\"/>\n </g>\n <g clip-path=\"url(#p30492a3230)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAPoAAADBCAYAAADmUa1TAADeOElEQVR4nOz9d7NtWXbdif2W3faY655LUw4FoAoADZpks7vlOhStkKhu8YNI+l5SKKRQM8RQSGJINE0QLBAEYaqAqspKn/nsdcdst5z+WPucd19WFgpoZGYBqLcybt7zzj127z3XnHPMMccUKaXE6/V6vV5/q5f8RX+A1+v1er2+/PXa0F+v1+uXYL029Nfr9folWK8N/fV6vX4J1mtDf71er1+C9drQX6/X65dgvTb01+v1+iVYrw399Xq9fgnWa0N/vV6vX4Klf9Ef4Ktef1kaoPirvvBf+AVer9fry1u/FIaeUuLA9E2z4d21SwEIxE8ZpUifb7+fxxoWn33R+c7PIxgL8fKNfh4DWQjxyuNfr9frf8z6pTD0w0pkw0vip++HV61akA39r+TRD/d99v1etxe8Xl/x+htr6H9ZYzk8XohswHkdLDB97usdNobPrs/zsIeN4bPP/1le/eXn+fyt5C/i9V97+tfrL7r+2hv6n2fQnzWWuyH6Zx93fCyz4c02IviMM07H/70a3s/vIYQghICUr+KYn2vMQPwLfKc/z2BfG/nr9UWsv/aGDn++AUO+6F/Jw1Mixni8L4RACAHvPdM4Mg0jKSWUUkgpMcZgjEEphVIKACllDvGFQEr5imEJIYgxHt/3eP/nGbv4/Kj+taG+Xl/l+mtr6HcNKKV0x2XOPnL2vClBRBJTIkVPCNmY97uO3XbL7eaW280tm82Wy8trnj15xovnz+j3O4xRXJyfo42mqhqapqGqKmxZUVY1RVWxWi05Wa9ZLFrKoqAsC6RSSJl/coEyJ/6zz0fcAfzizwz9D3+QJF5GGxxv5bDjp/C9zwnpX28ar9fPW3+tDP0Vrx1fGrdM2RREisAEcf63UMSU8NEQhCD6gbG75erZFT/605/wzg9/yAcfvcuzy+c8v7ql7wL73uHHnkJ5LtYV5tvf5ONPPmaz82glkAKCLLjuA1VZUxnN22894tu/8iucn665f/8e69M13/7VX6NqWpKEJBNy/k8ImXEAIZFSINPLNOD4/fIDjvcf1tFgX/m7BCGPf78btbzynJ+zXm8Gv9zrr5WhA8SYPbZI6RgKH+7LOa8ipZi94Byix+gJMTD2O549fcrvf+/3+b3f/X0+ePd9TCHZDTtubvdI1aAFCC341W//KotKcv/8BFzHU3HDum1oKsumd+jLLaPb0e0nrgrB9/uO//gHf8C9e/dYtBX/h//j/57v/OZvYMoCpQ1CaiTiWA4Tc8h/WK8Y2h1DfuW738EBjoZ+B8z7vFJbSunnGvFn8YTX65dv/UIN/bN59yGnTikh00/flwAvFCkJSAGRPCkEgu8Yh54ffP/7/LN/9s/43r//HtvbDffWJxS2xUpYNRUhKCbvQEu22y0Xp49Yr5b4/RX7mxsuFobTZYOLiTcv1nRB8cnT57y4eoGUifWq5enTx3SbmssXT9nvHlCmFq1rlLazF5dHIz/cvmugB6/8ecZ5MMi7f49CkOZ05bOv8/PWX6Ze//Oe/2WuP++z/UU/wxdVsvzbGvn8Qgz97km5C2odPPdLr84roWpMiZACgogME8kNhKmn63Z87/d+n//HP///8P4HH7Hf71mtW5ZFomKgqRXF6YoQLbvtjuvtlsl5Jh9YtQ2uspTWYI1iWVuMMYxBsNcL/vS9DxEC/tFv/ybTOPCv/uW/Zhwl+80NMfQQLIiSKBKIkDepOwZ719jhpbH+RYgycDD0fFtKeXzez/PSdy/Yw/H87P0/7/l/EeP5vMf9vOferZLcXT/vdX7eZ/+rGvvfViOHL83Qf0ZB+c5xfMVTf+bngJiLxCvoeUoRUiC4AT/uGbaXvHjyCS+uX/C93/kdbi4vqasGIS6RMtLKiV996xEP3nhAszglpoLrmy2/8wd/zO2UGIaJ89MT0nbBe7ZEmRyGL9qWIsLVDvokOTs94Y37Z9xbVdx8/A7vfjrR7W7p9zc0TZuJNZ9DxIGXG9lnDf3PW6964kiac/TjBsjLsp3IT/jMCwDIzzWcP+/fnwVA/7zPeXzfO5HX4TUPr/K5DMKf8Z4/q6rycyOY9DNu/xyS0t9mo/689eUYegIId/6ZSaak2TvF+TExQoqEA8qcIkSPSJEUBSFCDJ4UA9E7vJ9wqcdtdnSXH3D5+Ec8++QZ6/M3+O3v/CrjbuDZdUdlClSS3DuzvHHR8I1HJ5i6YfCKKDqqKjEl0CSMMRR1TXSBlAJROqqFZmkqNgzcW9WcrpbcOzvhXhX5e99+RL3oWMkd/vYK7n2TJCRKJBICkSIizhfWZ8C1Q+6dDr/vXHx3vfUrF6WIvCz6y3wYAZFe3j4S8D5TAvw8HCDfJ0hSHSOFu/9PIb78t8jIv5DAvNGK+R1TAqQEqeYNSObXky/PeN6Ys7GKJOf9/g44eaw3/Axi0t3vRsZtfuoxdw9VOuAan1+Ovbt+1t//MhvC36TN4ksx9FeLRXMIGrN9x5QgptkgAikGPCJfRMnn+4IjRvDO48aB4Eb82DOOHfvxkrgZGK/e4+bJn1JS8cbFEm02PDgvePLsY87WJWO/52tvv8XDewtq60H0BK24OCv4R7/9q/zwR59wcXrOuim4jAERAt5NJEq0FbSN5bcfvMF2
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
"execution_count": 7
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:57.702782485Z",
"start_time": "2026-07-27T09:47:57.650116830Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:57.754193331Z",
"start_time": "2026-07-27T09:47:57.704574735Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:57.820044630Z",
"start_time": "2026-07-27T09:47:57.755572144Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:57.877041771Z",
"start_time": "2026-07-27T09:47:57.821962159Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:57.924927279Z",
"start_time": "2026-07-27T09:47:57.877778001Z"
2026-07-27 14:18:15 +08:00
}
},
"cell_type": "code",
"source": "boxes = Y.reshape(h, w, 5, 4)",
"id": "4cc20a550931aef7",
"outputs": [],
"execution_count": 12
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.080082507Z",
"start_time": "2026-07-27T09:47:57.926102351Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"233.229634pt\" height=\"185.525265pt\" viewBox=\"0 0 233.229634 185.525265\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:47:58.021417</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M -0 185.525265 \nL 233.229634 185.525265 \nL 233.229634 0 \nL -0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 46.170811 161.64714 \nL 226.029634 161.64714 \nL 226.029634 23.04714 \nL 46.170811 23.04714 \nz\n\"/>\n </g>\n <g clip-path=\"url(#p81be786314)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAPoAAADBCAYAAADmUa1TAADeOElEQVR4nOz9d7NtWXbdif2W3faY655LUw4FoAoADZpks7vlOhStkKhu8YNI+l5SKKRQM8RQSGJINE0QLBAEYaqAqspKn/nsdcdst5z+WPucd19WFgpoZGYBqLcybt7zzj127z3XnHPMMccUKaXE6/V6vV5/q5f8RX+A1+v1er2+/PXa0F+v1+uXYL029Nfr9folWK8N/fV6vX4J1mtDf71er1+C9drQX6/X65dgvTb01+v1+iVYrw399Xq9fgnWa0N/vV6vX4Klf9Ef4Ktef1kaoPirvvBf+AVer9fry1u/FIaeUuLA9E2z4d21SwEIxE8ZpUifb7+fxxoWn33R+c7PIxgL8fKNfh4DWQjxyuNfr9frf8z6pTD0w0pkw0vip++HV61akA39r+TRD/d99v1etxe8Xl/x+htr6H9ZYzk8XohswHkdLDB97usdNobPrs/zsIeN4bPP/1le/eXn+fyt5C/i9V97+tfrL7r+2hv6n2fQnzWWuyH6Zx93fCyz4c02IviMM07H/70a3s/vIYQghICUr+KYn2vMQPwLfKc/z2BfG/nr9UWsv/aGDn++AUO+6F/Jw1Mixni8L4RACAHvPdM4Mg0jKSWUUkgpMcZgjEEphVIKACllDvGFQEr5imEJIYgxHt/3eP/nGbv4/Kj+taG+Xl/l+mtr6HcNKKV0x2XOPnL2vClBRBJTIkVPCNmY97uO3XbL7eaW280tm82Wy8trnj15xovnz+j3O4xRXJyfo42mqhqapqGqKmxZUVY1RVWxWi05Wa9ZLFrKoqAsC6RSSJl/coEyJ/6zz0fcAfzizwz9D3+QJF5GGxxv5bDjp/C9zwnpX28ar9fPW3+tDP0Vrx1fGrdM2RREisAEcf63UMSU8NEQhCD6gbG75erZFT/605/wzg9/yAcfvcuzy+c8v7ql7wL73uHHnkJ5LtYV5tvf5ONPPmaz82glkAKCLLjuA1VZUxnN22894tu/8iucn665f/8e69M13/7VX6NqWpKEJBNy/k8ImXEAIZFSINPLNOD4/fIDjvcf1tFgX/m7BCGPf78btbzynJ+zXm8Gv9zrr5WhA8SYPbZI6RgKH+7LOa8ipZi94Byix+gJMTD2O549fcrvf+/3+b3f/X0+ePd9TCHZDTtubvdI1aAFCC341W//KotKcv/8BFzHU3HDum1oKsumd+jLLaPb0e0nrgrB9/uO//gHf8C9e/dYtBX/h//j/57v/OZvYMoCpQ1CaiTiWA4Tc8h/WK8Y2h1DfuW738EBjoZ+B8z7vFJbSunnGvFn8YTX65dv/UIN/bN59yGnTikh00/flwAvFCkJSAGRPCkEgu8Yh54ffP/7/LN/9s/43r//HtvbDffWJxS2xUpYNRUhKCbvQEu22y0Xp49Yr5b4/RX7mxsuFobTZYOLiTcv1nRB8cnT57y4eoGUifWq5enTx3SbmssXT9nvHlCmFq1rlLazF5dHIz/cvmugB6/8ecZ5MMi7f49CkOZ05bOv8/PWX6Ze//Oe/2WuP++z/UU/wxdVsvzbGvn8Qgz97km5C2odPPdLr84roWpMiZACgogME8kNhKmn63Z87/d+n//HP///8P4HH7Hf71mtW5ZFomKgqRXF6YoQLbvtjuvtlsl5Jh9YtQ2uspTWYI1iWVuMMYxBsNcL/vS9DxEC/tFv/ybTOPCv/uW/Zhwl+80NMfQQLIiSKBKIkDepOwZ719jhpbH+RYgycDD0fFtKeXzez/PSdy/Yw/H87P0/7/l/EeP5vMf9vOferZLcXT/vdX7eZ/+rGvvfViOHL83Qf0ZB+c5xfMVTf+bngJiLxCvoeUoRUiC4AT/uGbaXvHjyCS+uX/C93/kdbi4vqasGIS6RMtLKiV996xEP3nhAszglpoLrmy2/8wd/zO2UGIaJ89MT0nbBe7ZEmRyGL9qWIsLVDvokOTs94Y37Z9xbVdx8/A7vfjrR7W7p9zc0TZuJNZ9DxIGXG9lnDf3PW6964kiac/TjBsjLsp3IT/jMCwDIzzWcP+/fnwVA/7zPeXzfO5HX4TUPr/K5DMKf8Z4/q6rycyOY9DNu/xyS0t9mo/689eUYegIId/6ZSaak2TvF+TExQoqEA8qcIkSPSJEUBSFCDJ4UA9E7vJ9wqcdtdnSXH3D5+Ec8++QZ6/M3+O3v/CrjbuDZdUdlClSS3DuzvHHR8I1HJ5i6YfCKKDqqKjEl0CSMMRR1TXSBlAJROqqFZmkqNgzcW9WcrpbcOzvhXhX5e99+RL3oWMkd/vYK7n2TJCRKJBICkSIizhfWZ8C1Q+6dDr/vXHx3vfUrF6WIvCz6y3wYAZFe3j4S8D5TAvw8HCDfJ0hSHSOFu/9PIb78t8jIv5DAvNGK+R1TAqQEqeYNSObXky/PeN6Ys7GKJOf9/g44eaw3/Axi0t3vRsZtfuoxdw9VOuAan1+Ovbt+1t//MhvC36TN4ksx9FeLRXMIGrN9x5QgptkgAikGPCJfRMnn+4IjRvDO48aB4Eb82DOOHfvxkrgZGK/e4+bJn1JS8cbFEm02PDgvePLsY87WJWO/52tvv8XDewtq60H0BK24OCv4R7/9q/zwR59wcXrOuim4jAERAt5NJEq0FbSN5bcfvM
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
"execution_count": 13
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.134645045Z",
"start_time": "2026-07-27T09:47:58.116726506Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.193234091Z",
"start_time": "2026-07-27T09:47:58.136182636Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.260217572Z",
"start_time": "2026-07-27T09:47:58.212670597Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.309313551Z",
"start_time": "2026-07-27T09:47:58.261240696Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.451959046Z",
"start_time": "2026-07-27T09:47:58.310247636Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"220.346324pt\" height=\"173.353814pt\" viewBox=\"0 0 220.346324 173.353814\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:47:58.388372</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 173.353814 \nL 220.346324 173.353814 \nL 220.346324 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 149.475689 \nL 213.146324 149.475689 \nL 213.146324 10.875689 \nL 33.2875 10.875689 \nz\n\"/>\n </g>\n <g clip-path=\"url(#p9228ade140)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAPoAAADBCAYAAADmUa1TAADeOElEQVR4nOz9d7NtWXbdif2W3faY655LUw4FoAoADZpks7vlOhStkKhu8YNI+l5SKKRQM8RQSGJINE0QLBAEYaqAqspKn/nsdcdst5z+WPucd19WFgpoZGYBqLcybt7zzj127z3XnHPMMccUKaXE6/V6vV5/q5f8RX+A1+v1er2+/PXa0F+v1+uXYL029Nfr9folWK8N/fV6vX4J1mtDf71er1+C9drQX6/X65dgvTb01+v1+iVYrw399Xq9fgnWa0N/vV6vX4Klf9Ef4Ktef1kaoPirvvBf+AVer9fry1u/FIaeUuLA9E2z4d21SwEIxE8ZpUifb7+fxxoWn33R+c7PIxgL8fKNfh4DWQjxyuNfr9frf8z6pTD0w0pkw0vip++HV61akA39r+TRD/d99v1etxe8Xl/x+htr6H9ZYzk8XohswHkdLDB97usdNobPrs/zsIeN4bPP/1le/eXn+fyt5C/i9V97+tfrL7r+2hv6n2fQnzWWuyH6Zx93fCyz4c02IviMM07H/70a3s/vIYQghICUr+KYn2vMQPwLfKc/z2BfG/nr9UWsv/aGDn++AUO+6F/Jw1Mixni8L4RACAHvPdM4Mg0jKSWUUkgpMcZgjEEphVIKACllDvGFQEr5imEJIYgxHt/3eP/nGbv4/Kj+taG+Xl/l+mtr6HcNKKV0x2XOPnL2vClBRBJTIkVPCNmY97uO3XbL7eaW280tm82Wy8trnj15xovnz+j3O4xRXJyfo42mqhqapqGqKmxZUVY1RVWxWi05Wa9ZLFrKoqAsC6RSSJl/coEyJ/6zz0fcAfzizwz9D3+QJF5GGxxv5bDjp/C9zwnpX28ar9fPW3+tDP0Vrx1fGrdM2RREisAEcf63UMSU8NEQhCD6gbG75erZFT/605/wzg9/yAcfvcuzy+c8v7ql7wL73uHHnkJ5LtYV5tvf5ONPPmaz82glkAKCLLjuA1VZUxnN22894tu/8iucn665f/8e69M13/7VX6NqWpKEJBNy/k8ImXEAIZFSINPLNOD4/fIDjvcf1tFgX/m7BCGPf78btbzynJ+zXm8Gv9zrr5WhA8SYPbZI6RgKH+7LOa8ipZi94Byix+gJMTD2O549fcrvf+/3+b3f/X0+ePd9TCHZDTtubvdI1aAFCC341W//KotKcv/8BFzHU3HDum1oKsumd+jLLaPb0e0nrgrB9/uO//gHf8C9e/dYtBX/h//j/57v/OZvYMoCpQ1CaiTiWA4Tc8h/WK8Y2h1DfuW738EBjoZ+B8z7vFJbSunnGvFn8YTX65dv/UIN/bN59yGnTikh00/flwAvFCkJSAGRPCkEgu8Yh54ffP/7/LN/9s/43r//HtvbDffWJxS2xUpYNRUhKCbvQEu22y0Xp49Yr5b4/RX7mxsuFobTZYOLiTcv1nRB8cnT57y4eoGUifWq5enTx3SbmssXT9nvHlCmFq1rlLazF5dHIz/cvmugB6/8ecZ5MMi7f49CkOZ05bOv8/PWX6Ze//Oe/2WuP++z/UU/wxdVsvzbGvn8Qgz97km5C2odPPdLr84roWpMiZACgogME8kNhKmn63Z87/d+n//HP///8P4HH7Hf71mtW5ZFomKgqRXF6YoQLbvtjuvtlsl5Jh9YtQ2uspTWYI1iWVuMMYxBsNcL/vS9DxEC/tFv/ybTOPCv/uW/Zhwl+80NMfQQLIiSKBKIkDepOwZ719jhpbH+RYgycDD0fFtKeXzez/PSdy/Yw/H87P0/7/l/EeP5vMf9vOferZLcXT/vdX7eZ/+rGvvfViOHL83Qf0ZB+c5xfMVTf+bngJiLxCvoeUoRUiC4AT/uGbaXvHjyCS+uX/C93/kdbi4vqasGIS6RMtLKiV996xEP3nhAszglpoLrmy2/8wd/zO2UGIaJ89MT0nbBe7ZEmRyGL9qWIsLVDvokOTs94Y37Z9xbVdx8/A7vfjrR7W7p9zc0TZuJNZ9DxIGXG9lnDf3PW6964kiac/TjBsjLsp3IT/jMCwDIzzWcP+/fnwVA/7zPeXzfO5HX4TUPr/K5DMKf8Z4/q6rycyOY9DNu/xyS0t9mo/689eUYegIId/6ZSaak2TvF+TExQoqEA8qcIkSPSJEUBSFCDJ4UA9E7vJ9wqcdtdnSXH3D5+Ec8++QZ6/M3+O3v/CrjbuDZdUdlClSS3DuzvHHR8I1HJ5i6YfCKKDqqKjEl0CSMMRR1TXSBlAJROqqFZmkqNgzcW9WcrpbcOzvhXhX5e99+RL3oWMkd/vYK7n2TJCRKJBICkSIizhfWZ8C1Q+6dDr/vXHx3vfUrF6WIvCz6y3wYAZFe3j4S8D5TAvw8HCDfJ0hSHSOFu/9PIb78t8jIv5DAvNGK+R1TAqQEqeYNSObXky/PeN6Ys7GKJOf9/g44eaw3/Axi0t3vRsZtfuoxdw9VOuAan1+Ovbt+1t//MhvC36TN4ksx9FeLRXMIGrN9x5QgptkgAikGPCJfRMnn+4IjRvDO48aB4Eb82DOOHfvxkrgZGK/e4+bJn1JS8cbFEm02PDgvePLsY87WJWO/52tvv8XDewtq60H0BK24OCv4R7/9q/zwR59wcXrOuim4jAERAt5NJEq0FbSN5bcfvMF2
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
"execution_count": 18
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.499427639Z",
"start_time": "2026-07-27T09:47:58.453550913Z"
2026-07-27 14:18:15 +08:00
}
},
"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": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.560984562Z",
"start_time": "2026-07-27T09:47:58.500638893Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [],
2026-07-27 17:57:55 +08:00
"execution_count": 20
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.627497913Z",
"start_time": "2026-07-27T09:47:58.578335925Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [],
2026-07-27 17:57:55 +08:00
"execution_count": 21
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.687688493Z",
"start_time": "2026-07-27T09:47:58.628509966Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [],
2026-07-27 17:57:55 +08:00
"execution_count": 22
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.737335139Z",
"start_time": "2026-07-27T09:47:58.689268726Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [],
2026-07-27 17:57:55 +08:00
"execution_count": 23
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.871856579Z",
"start_time": "2026-07-27T09:47:58.738750970Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"220.346324pt\" height=\"173.353814pt\" viewBox=\"0 0 220.346324 173.353814\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:47:58.828265</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 173.353814 \nL 220.346324 173.353814 \nL 220.346324 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 149.475689 \nL 213.146324 149.475689 \nL 213.146324 10.875689 \nL 33.2875 10.875689 \nz\n\"/>\n </g>\n <g clip-path=\"url(#pf5ff0b9d7c)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAPoAAADBCAYAAADmUa1TAADeOElEQVR4nOz9d7NtWXbdif2W3faY655LUw4FoAoADZpks7vlOhStkKhu8YNI+l5SKKRQM8RQSGJINE0QLBAEYaqAqspKn/nsdcdst5z+WPucd19WFgpoZGYBqLcybt7zzj127z3XnHPMMccUKaXE6/V6vV5/q5f8RX+A1+v1er2+/PXa0F+v1+uXYL029Nfr9folWK8N/fV6vX4J1mtDf71er1+C9drQX6/X65dgvTb01+v1+iVYrw399Xq9fgnWa0N/vV6vX4Klf9Ef4Ktef1kaoPirvvBf+AVer9fry1u/FIaeUuLA9E2z4d21SwEIxE8ZpUifb7+fxxoWn33R+c7PIxgL8fKNfh4DWQjxyuNfr9frf8z6pTD0w0pkw0vip++HV61akA39r+TRD/d99v1etxe8Xl/x+htr6H9ZYzk8XohswHkdLDB97usdNobPrs/zsIeN4bPP/1le/eXn+fyt5C/i9V97+tfrL7r+2hv6n2fQnzWWuyH6Zx93fCyz4c02IviMM07H/70a3s/vIYQghICUr+KYn2vMQPwLfKc/z2BfG/nr9UWsv/aGDn++AUO+6F/Jw1Mixni8L4RACAHvPdM4Mg0jKSWUUkgpMcZgjEEphVIKACllDvGFQEr5imEJIYgxHt/3eP/nGbv4/Kj+taG+Xl/l+mtr6HcNKKV0x2XOPnL2vClBRBJTIkVPCNmY97uO3XbL7eaW280tm82Wy8trnj15xovnz+j3O4xRXJyfo42mqhqapqGqKmxZUVY1RVWxWi05Wa9ZLFrKoqAsC6RSSJl/coEyJ/6zz0fcAfzizwz9D3+QJF5GGxxv5bDjp/C9zwnpX28ar9fPW3+tDP0Vrx1fGrdM2RREisAEcf63UMSU8NEQhCD6gbG75erZFT/605/wzg9/yAcfvcuzy+c8v7ql7wL73uHHnkJ5LtYV5tvf5ONPPmaz82glkAKCLLjuA1VZUxnN22894tu/8iucn665f/8e69M13/7VX6NqWpKEJBNy/k8ImXEAIZFSINPLNOD4/fIDjvcf1tFgX/m7BCGPf78btbzynJ+zXm8Gv9zrr5WhA8SYPbZI6RgKH+7LOa8ipZi94Byix+gJMTD2O549fcrvf+/3+b3f/X0+ePd9TCHZDTtubvdI1aAFCC341W//KotKcv/8BFzHU3HDum1oKsumd+jLLaPb0e0nrgrB9/uO//gHf8C9e/dYtBX/h//j/57v/OZvYMoCpQ1CaiTiWA4Tc8h/WK8Y2h1DfuW738EBjoZ+B8z7vFJbSunnGvFn8YTX65dv/UIN/bN59yGnTikh00/flwAvFCkJSAGRPCkEgu8Yh54ffP/7/LN/9s/43r//HtvbDffWJxS2xUpYNRUhKCbvQEu22y0Xp49Yr5b4/RX7mxsuFobTZYOLiTcv1nRB8cnT57y4eoGUifWq5enTx3SbmssXT9nvHlCmFq1rlLazF5dHIz/cvmugB6/8ecZ5MMi7f49CkOZ05bOv8/PWX6Ze//Oe/2WuP++z/UU/wxdVsvzbGvn8Qgz97km5C2odPPdLr84roWpMiZACgogME8kNhKmn63Z87/d+n//HP///8P4HH7Hf71mtW5ZFomKgqRXF6YoQLbvtjuvtlsl5Jh9YtQ2uspTWYI1iWVuMMYxBsNcL/vS9DxEC/tFv/ybTOPCv/uW/Zhwl+80NMfQQLIiSKBKIkDepOwZ719jhpbH+RYgycDD0fFtKeXzez/PSdy/Yw/H87P0/7/l/EeP5vMf9vOferZLcXT/vdX7eZ/+rGvvfViOHL83Qf0ZB+c5xfMVTf+bngJiLxCvoeUoRUiC4AT/uGbaXvHjyCS+uX/C93/kdbi4vqasGIS6RMtLKiV996xEP3nhAszglpoLrmy2/8wd/zO2UGIaJ89MT0nbBe7ZEmRyGL9qWIsLVDvokOTs94Y37Z9xbVdx8/A7vfjrR7W7p9zc0TZuJNZ9DxIGXG9lnDf3PW6964kiac/TjBsjLsp3IT/jMCwDIzzWcP+/fnwVA/7zPeXzfO5HX4TUPr/K5DMKf8Z4/q6rycyOY9DNu/xyS0t9mo/689eUYegIId/6ZSaak2TvF+TExQoqEA8qcIkSPSJEUBSFCDJ4UA9E7vJ9wqcdtdnSXH3D5+Ec8++QZ6/M3+O3v/CrjbuDZdUdlClSS3DuzvHHR8I1HJ5i6YfCKKDqqKjEl0CSMMRR1TXSBlAJROqqFZmkqNgzcW9WcrpbcOzvhXhX5e99+RL3oWMkd/vYK7n2TJCRKJBICkSIizhfWZ8C1Q+6dDr/vXHx3vfUrF6WIvCz6y3wYAZFe3j4S8D5TAvw8HCDfJ0hSHSOFu/9PIb78t8jIv5DAvNGK+R1TAqQEqeYNSObXky/PeN6Ys7GKJOf9/g44eaw3/Axi0t3vRsZtfuoxdw9VOuAan1+Ovbt+1t//MhvC36TN4ksx9FeLRXMIGrN9x5QgptkgAikGPCJfRMnn+4IjRvDO48aB4Eb82DOOHfvxkrgZGK/e4+bJn1JS8cbFEm02PDgvePLsY87WJWO/52tvv8XDewtq60H0BK24OCv4R7/9q/zwR59wcXrOuim4jAERAt5NJEq0FbSN5bcfvMF2
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 24
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:58.961912712Z",
"start_time": "2026-07-27T09:47:58.872770917Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
{
"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]]])"
]
},
2026-07-27 17:57:55 +08:00
"execution_count": 25,
2026-07-27 14:18:15 +08:00
"metadata": {},
"output_type": "execute_result"
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 25
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:59.200171491Z",
"start_time": "2026-07-27T09:47:58.962631809Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"220.346324pt\" height=\"173.353814pt\" viewBox=\"0 0 220.346324 173.353814\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:47:59.051198</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 173.353814 \nL 220.346324 173.353814 \nL 220.346324 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 149.475689 \nL 213.146324 149.475689 \nL 213.146324 10.875689 \nL 33.2875 10.875689 \nz\n\"/>\n </g>\n <g clip-path=\"url(#p9bf1eeea83)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAPoAAADBCAYAAADmUa1TAADeOElEQVR4nOz9d7NtWXbdif2W3faY655LUw4FoAoADZpks7vlOhStkKhu8YNI+l5SKKRQM8RQSGJINE0QLBAEYaqAqspKn/nsdcdst5z+WPucd19WFgpoZGYBqLcybt7zzj127z3XnHPMMccUKaXE6/V6vV5/q5f8RX+A1+v1er2+/PXa0F+v1+uXYL029Nfr9folWK8N/fV6vX4J1mtDf71er1+C9drQX6/X65dgvTb01+v1+iVYrw399Xq9fgnWa0N/vV6vX4Klf9Ef4Ktef1kaoPirvvBf+AVer9fry1u/FIaeUuLA9E2z4d21SwEIxE8ZpUifb7+fxxoWn33R+c7PIxgL8fKNfh4DWQjxyuNfr9frf8z6pTD0w0pkw0vip++HV61akA39r+TRD/d99v1etxe8Xl/x+htr6H9ZYzk8XohswHkdLDB97usdNobPrs/zsIeN4bPP/1le/eXn+fyt5C/i9V97+tfrL7r+2hv6n2fQnzWWuyH6Zx93fCyz4c02IviMM07H/70a3s/vIYQghICUr+KYn2vMQPwLfKc/z2BfG/nr9UWsv/aGDn++AUO+6F/Jw1Mixni8L4RACAHvPdM4Mg0jKSWUUkgpMcZgjEEphVIKACllDvGFQEr5imEJIYgxHt/3eP/nGbv4/Kj+taG+Xl/l+mtr6HcNKKV0x2XOPnL2vClBRBJTIkVPCNmY97uO3XbL7eaW280tm82Wy8trnj15xovnz+j3O4xRXJyfo42mqhqapqGqKmxZUVY1RVWxWi05Wa9ZLFrKoqAsC6RSSJl/coEyJ/6zz0fcAfzizwz9D3+QJF5GGxxv5bDjp/C9zwnpX28ar9fPW3+tDP0Vrx1fGrdM2RREisAEcf63UMSU8NEQhCD6gbG75erZFT/605/wzg9/yAcfvcuzy+c8v7ql7wL73uHHnkJ5LtYV5tvf5ONPPmaz82glkAKCLLjuA1VZUxnN22894tu/8iucn665f/8e69M13/7VX6NqWpKEJBNy/k8ImXEAIZFSINPLNOD4/fIDjvcf1tFgX/m7BCGPf78btbzynJ+zXm8Gv9zrr5WhA8SYPbZI6RgKH+7LOa8ipZi94Byix+gJMTD2O549fcrvf+/3+b3f/X0+ePd9TCHZDTtubvdI1aAFCC341W//KotKcv/8BFzHU3HDum1oKsumd+jLLaPb0e0nrgrB9/uO//gHf8C9e/dYtBX/h//j/57v/OZvYMoCpQ1CaiTiWA4Tc8h/WK8Y2h1DfuW738EBjoZ+B8z7vFJbSunnGvFn8YTX65dv/UIN/bN59yGnTikh00/flwAvFCkJSAGRPCkEgu8Yh54ffP/7/LN/9s/43r//HtvbDffWJxS2xUpYNRUhKCbvQEu22y0Xp49Yr5b4/RX7mxsuFobTZYOLiTcv1nRB8cnT57y4eoGUifWq5enTx3SbmssXT9nvHlCmFq1rlLazF5dHIz/cvmugB6/8ecZ5MMi7f49CkOZ05bOv8/PWX6Ze//Oe/2WuP++z/UU/wxdVsvzbGvn8Qgz97km5C2odPPdLr84roWpMiZACgogME8kNhKmn63Z87/d+n//HP///8P4HH7Hf71mtW5ZFomKgqRXF6YoQLbvtjuvtlsl5Jh9YtQ2uspTWYI1iWVuMMYxBsNcL/vS9DxEC/tFv/ybTOPCv/uW/Zhwl+80NMfQQLIiSKBKIkDepOwZ719jhpbH+RYgycDD0fFtKeXzez/PSdy/Yw/H87P0/7/l/EeP5vMf9vOferZLcXT/vdX7eZ/+rGvvfViOHL83Qf0ZB+c5xfMVTf+bngJiLxCvoeUoRUiC4AT/uGbaXvHjyCS+uX/C93/kdbi4vqasGIS6RMtLKiV996xEP3nhAszglpoLrmy2/8wd/zO2UGIaJ89MT0nbBe7ZEmRyGL9qWIsLVDvokOTs94Y37Z9xbVdx8/A7vfjrR7W7p9zc0TZuJNZ9DxIGXG9lnDf3PW6964kiac/TjBsjLsp3IT/jMCwDIzzWcP+/fnwVA/7zPeXzfO5HX4TUPr/K5DMKf8Z4/q6rycyOY9DNu/xyS0t9mo/689eUYegIId/6ZSaak2TvF+TExQoqEA8qcIkSPSJEUBSFCDJ4UA9E7vJ9wqcdtdnSXH3D5+Ec8++QZ6/M3+O3v/CrjbuDZdUdlClSS3DuzvHHR8I1HJ5i6YfCKKDqqKjEl0CSMMRR1TXSBlAJROqqFZmkqNgzcW9WcrpbcOzvhXhX5e99+RL3oWMkd/vYK7n2TJCRKJBICkSIizhfWZ8C1Q+6dDr/vXHx3vfUrF6WIvCz6y3wYAZFe3j4S8D5TAvw8HCDfJ0hSHSOFu/9PIb78t8jIv5DAvNGK+R1TAqQEqeYNSObXky/PeN6Ys7GKJOf9/g44eaw3/Axi0t3vRsZtfuoxdw9VOuAan1+Ovbt+1t//MhvC36TN4ksx9FeLRXMIGrN9x5QgptkgAikGPCJfRMnn+4IjRvDO48aB4Eb82DOOHfvxkrgZGK/e4+bJn1JS8cbFEm02PDgvePLsY87WJWO/52tvv8XDewtq60H0BK24OCv4R7/9q/zwR59wcXrOuim4jAERAt5NJEq0FbSN5bcfvMF2
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 26
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:59.227792428Z",
"start_time": "2026-07-27T09:47:59.205845108Z"
2026-07-27 14:18:15 +08:00
}
},
"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)"
]
},
2026-07-27 17:57:55 +08:00
"execution_count": 27,
2026-07-27 14:18:15 +08:00
"metadata": {},
"output_type": "execute_result"
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 27
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:59.462849505Z",
"start_time": "2026-07-27T09:47:59.228706763Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"180.077031pt\" height=\"173.369063pt\" viewBox=\"0 0 180.077031 173.369063\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:47:59.341993</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 173.369063 \nL 180.077031 173.369063 \nL 180.077031 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 149.490938 \nL 171.8875 149.490938 \nL 171.8875 10.890938 \nL 33.2875 10.890938 \nz\n\"/>\n </g>\n <g clip-path=\"url(#pcb3e935286)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAMEAAADBCAYAAAB2QtScAAEAAElEQVR4nLT9x7MsV57niX2OcPeQN65+WkKLBJCiMquyqjunpqeH3TZmM0PjLGjknhsajfwPuOCSS/4FpBlpRiM3XAynh6zqri7BElkCmUgACeBpefUN7eqc8+PiuMeN94DMqrKxcZgj4sYL4X7OT35/Sv1f/vf/B6E5lFIoFIhCKYVI/CeDwmoDCgRAK7xWiDIolSAYBIsojWjDuK54MpkwR1BW8we/9xN+50efkCYAniCewjtmlWdWwaJW5Chq5ynmBWFZkfb7lJ2UUAWWZ+dcHva5uz1kaD1IjQEUBgiICiiEIB5jNJlOccHjQ8B5z4uDl3TSFJVY7GDE6czx8MkJx+MaZy39XsreQLOZafqZpZ8YrFEoY1lUjtoJIoaqDpS1Y1YXzPKC+bKgKD1F6XEyRGyCWIsDvAQIAUdAjGBNzeZGl629bbr9hNQISXCgwIkQhLj2WqOUoAjE1W7PeMftEUKgrmsArLVorUGBF01VwfHhlNPTGeItW1sjLl/rYNPA8dGY06M5vkpQ0iUEB6bEoUnRXNtVdLOCP/0Pv+DkLHDl+nVUYvBK4ZXC1SWHZ6e4uub2pauMrOKHn9zgs18/xdHn1nuXCdLn+OgM7wLS0JLWlvVDrd1LQ2ZorRBVNfepUaIJCpwKWIR8PIGy4nvvvcHdGxskdYlyHqcMtROKoiR3DkvgUr/HnWt7KKNwSrFwNcfjMZWrSJIOic2oKs/J8Rl2/UpCCKsLVFrFZwpAXzBA+9i88+JRoxsG2uz08Srh2XTMpMz587/8S2xm+f4nH2C0xkvACQSlEARa5lOKrNelDoqyrpDEYNOEwc6Io8k5Z4/PeOfyJa71UnQoEV0haJDme1CIQOEdQYQAFM4x3N5FbI/xbMr84JzpvGR/u8e71xK8MiglXNkdYpUgImg8QsApjcwCroYgmm6WgOqxo/oEhCCaIIq8qJjOCs6nC86XcxaVoxKFUxkGg4hGnOL8tGIyO6I3zNjZ6rM56JAasCrglRAQUOG1tZVXiH9dYKVpSgiBsiyx1mKMQWlIrOLatS02NlKODs85OzugKDOuXt1nb3uffmfAwYtT8mUByoCkiHLUQTg6Lrh5rcfv/M7H/PG//3tOjg/ZuXwFiSSAAoL3IHGtUJogAa1BCQQvFGVJCO11R9r4rYdq3qtAKb1224JCoZUiBE8VPIHAw4NDAhU3dzbpKo14j0YwRqF9pMGgFF4Uofaczuec5XOCNWRJilLgvcMYw87uFlaLrDhRRYpHiDeoFIhEmeRDABUJNb4nEq6IBmWiFpHmtj1spV30yPJ4csyszPmPf/4XmNTy3jtvI0pTeYUP8XsiV8WLcARUJyUUJaquCSKIgXR7i/PzCX/z4Cnvb2/y5uUtjHGoAErahRNCgKAgCORFgagoXctiSaYVm3tb9G50SHTASs14POPw+Jgzhlzau4yVeF9BGZwH5QUdAkoFRAkiNZrQ3HuUYv2OZr+b4vYyajFM85rnx+ecLGG+LKmdQpTFq4S6hLmvyWenTLop2xs9RqMeOtEoI0DLBC0taNSa2Im3KatHrTWdToeqqijLkjRNSdIEKBltWrrdLU6OFpwe5zx+cMju3ga7+wNu3t7n5YszZhOHhLTRQopaMl4cLLlza5Mf/fB9/vyvP+V8cs5oawcd4tUE7xFFFARACD4yQRBCEMrCgawL0ZYZ1u5Dtc+FKAcbRhC9xviBqO+jlVB7Bxpqbfj68THFwvH29X0yDVCjtcJoBV7htWZWOSaTMdMqR1IT10VaCg9opUgSjfXORwJupfHqmi8uUqRR71qhG60Qr1M3N6oQFEEJGkGJYD2MbIc3t6/weHzEWbHkj//kz3Bo7r5xBxeg8hDa35Qoy4NWBCOk3Q6urEgUVCEgJiHtDym94peHZxxM5rx/9wZ7XR0lt+hmwQVEUZYlZVFikwSbGIaZkOgMaxSJKkEJTmA4GKJMxl/98pd8lG6wPxyiRVAYxAkSdCR+BUJA1IUq1zSbJx4IJCgSBd1ewvb1S4xLj9KKk/MZL8+mHE/nYDsEpxAPy6rCzxyz0yX9zS69UYe0b9C6XeDWLHpNcDaCaP0xyzKsTVguCvK8pNfrYq0hTQ2Xr4zo9/q8eH7K8eEZRbFk7/IW12/scZiMOTtZkoQE0RVBaQrp8PDFGW/duc7xeMYvf/2YJO0w6vej2SbgEZz3BAzSagKgKj2uVo1Ej4KiXS/VSPsV8bf3qC5eUa3Kaf9dCYihriq0NEzlhWVZ8/D5AUYF3rq5h27ea7TGBaFwnienZzhXk3Q62NREsa0Etfq10OyjaggwBCQEkACERtrEMxKBrEyhdntE1swhiYwgDUNopbAoNm2Pu7vX2Bts40rPn/7ZX/D4yTNcAC8QRC7WpZUVRoNRJInF+5pEG3SABI1Yi93b58my5i8/e8TReY4TTXspohTBeaplgUHTTTNSbTFaI4BzQl46fAAtgURBL+2wu73Hn//13zD3BbWuCarCh4qAIMo0Z7y/oDQeEwlAWVCWoBIEixKNFaGrPEO1YCstefvakB+/d5Xvv7nPhq0hnyNFjQ4KVwuLWc3hwTnPnx1xfDhhuSjxThC50DZIlPzSSM61hwvfzViGgw1S2+Hs
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 28
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:47:59.984065826Z",
"start_time": "2026-07-27T09:47:59.567538110Z"
2026-07-27 14:18:15 +08:00
}
},
"cell_type": "code",
"source": "display_anchors(fmap_w=2, fmap_h=2, s=[0.4])",
"id": "6620db5aa275bf52",
"outputs": [
{
"data": {
"text/plain": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"180.077031pt\" height=\"173.369063pt\" viewBox=\"0 0 180.077031 173.369063\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:47:59.675822</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 173.369063 \nL 180.077031 173.369063 \nL 180.077031 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 149.490938 \nL 171.8875 149.490938 \nL 171.8875 10.890938 \nL 33.2875 10.890938 \nz\n\"/>\n </g>\n <g clip-path=\"url(#p15c37a89fa)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAMEAAADBCAYAAAB2QtScAAEAAElEQVR4nLT9x7MsV57niX2OcPeQN65+WkKLBJCiMquyqjunpqeH3TZmM0PjLGjknhsajfwPuOCSS/4FpBlpRiM3XAynh6zqri7BElkCmUgACeBpefUN7eqc8+PiuMeN94DMqrKxcZgj4sYL4X7OT35/Sv1f/vf/B6E5lFIoFIhCKYVI/CeDwmoDCgRAK7xWiDIolSAYBIsojWjDuK54MpkwR1BW8we/9xN+50efkCYAniCewjtmlWdWwaJW5Chq5ynmBWFZkfb7lJ2UUAWWZ+dcHva5uz1kaD1IjQEUBgiICiiEIB5jNJlOccHjQ8B5z4uDl3TSFJVY7GDE6czx8MkJx+MaZy39XsreQLOZafqZpZ8YrFEoY1lUjtoJIoaqDpS1Y1YXzPKC+bKgKD1F6XEyRGyCWIsDvAQIAUdAjGBNzeZGl629bbr9hNQISXCgwIkQhLj2WqOUoAjE1W7PeMftEUKgrmsArLVorUGBF01VwfHhlNPTGeItW1sjLl/rYNPA8dGY06M5vkpQ0iUEB6bEoUnRXNtVdLOCP/0Pv+DkLHDl+nVUYvBK4ZXC1SWHZ6e4uub2pauMrOKHn9zgs18/xdHn1nuXCdLn+OgM7wLS0JLWlvVDrd1LQ2ZorRBVNfepUaIJCpwKWIR8PIGy4nvvvcHdGxskdYlyHqcMtROKoiR3DkvgUr/HnWt7KKNwSrFwNcfjMZWrSJIOic2oKs/J8Rl2/UpCCKsLVFrFZwpAXzBA+9i88+JRoxsG2uz08Srh2XTMpMz587/8S2xm+f4nH2C0xkvACQSlEARa5lOKrNelDoqyrpDEYNOEwc6Io8k5Z4/PeOfyJa71UnQoEV0haJDme1CIQOEdQYQAFM4x3N5FbI/xbMr84JzpvGR/u8e71xK8MiglXNkdYpUgImg8QsApjcwCroYgmm6WgOqxo/oEhCCaIIq8qJjOCs6nC86XcxaVoxKFUxkGg4hGnOL8tGIyO6I3zNjZ6rM56JAasCrglRAQUOG1tZVXiH9dYKVpSgiBsiyx1mKMQWlIrOLatS02NlKODs85OzugKDOuXt1nb3uffmfAwYtT8mUByoCkiHLUQTg6Lrh5rcfv/M7H/PG//3tOjg/ZuXwFiSSAAoL3IHGtUJogAa1BCQQvFGVJCO11R9r4rYdq3qtAKb1224JCoZUiBE8VPIHAw4NDAhU3dzbpKo14j0YwRqF9pMGgFF4Uofaczuec5XOCNWRJilLgvcMYw87uFlaLrDhRRYpHiDeoFIhEmeRDABUJNb4nEq6IBmWiFpHmtj1spV30yPJ4csyszPmPf/4XmNTy3jtvI0pTeYUP8XsiV8WLcARUJyUUJaquCSKIgXR7i/PzCX/z4Cnvb2/y5uUtjHGoAErahRNCgKAgCORFgagoXctiSaYVm3tb9G50SHTASs14POPw+Jgzhlzau4yVeF9BGZwH5QUdAkoFRAkiNZrQ3HuUYv2OZr+b4vYyajFM85rnx+ecLGG+LKmdQpTFq4S6hLmvyWenTLop2xs9RqMeOtEoI0DLBC0taNSa2Im3KatHrTWdToeqqijLkjRNSdIEKBltWrrdLU6OFpwe5zx+cMju3ga7+wNu3t7n5YszZhOHhLTRQopaMl4cLLlza5Mf/fB9/vyvP+V8cs5oawcd4tUE7xFFFARACD4yQRBCEMrCgawL0ZYZ1u5Dtc+FKAcbRhC9xviBqO+jlVB7Bxpqbfj68THFwvH29X0yDVCjtcJoBV7htWZWOSaTMdMqR1IT10VaCg9opUgSjfXORwJupfHqmi8uUqRR71qhG60Qr1M3N6oQFEEJGkGJYD2MbIc3t6/weHzEWbHkj//kz3Bo7r5xBxeg8hDa35Qoy4NWBCOk3Q6urEgUVCEgJiHtDym94peHZxxM5rx/9wZ7XR0lt+hmwQVEUZYlZVFikwSbGIaZkOgMaxSJKkEJTmA4GKJMxl/98pd8lG6wPxyiRVAYxAkSdCR+BUJA1IUq1zSbJx4IJCgSBd1ewvb1S4xLj9KKk/MZL8+mHE/nYDsEpxAPy6rCzxyz0yX9zS69UYe0b9C6XeDWLHpNcDaCaP0xyzKsTVguCvK8pNfrYq0hTQ2Xr4zo9/q8eH7K8eEZRbFk7/IW12/scZiMOTtZkoQE0RVBaQrp8PDFGW/duc7xeMYvf/2YJO0w6vej2SbgEZz3BAzSagKgKj2uVo1Ej4KiXS/VSPsV8bf3qC5eUa3Kaf9dCYihriq0NEzlhWVZ8/D5AUYF3rq5h27ea7TGBaFwnienZzhXk3Q62NREsa0Etfq10OyjaggwBCQEkACERtrEMxKBrEyhdntE1swhiYwgDUNopbAoNm2Pu7vX2Bts40rPn/7ZX/D4yTNcAC8QRC7WpZUVRoNRJInF+5pEG3SABI1Yi93b58my5i8/e8TReY4TTXspohTBeaplgUHTTTNSbTFaI4BzQl46fAAtgURBL+2wu73Hn//13zD3BbWuCarCh4qAIMo0Z7y/oDQeEwlAWVCWoBIEixKNFaGrPEO1YCstefvakB+/d5Xvv7nPhq0hnyNFjQ4KVwuLWc3hwTnPnx1xfDhhuSjxThC50DZIlPzSSM61hwvfzViGgw1S2+Hs
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 29
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:48:00.334968551Z",
"start_time": "2026-07-27T09:48:00.005932344Z"
2026-07-27 14:18:15 +08:00
}
},
"cell_type": "code",
"source": "display_anchors(fmap_w=1, fmap_h=1, s=[0.8])",
"id": "8a0c4555b2752c99",
"outputs": [
{
"data": {
"text/plain": [
"<Figure size 350x250 with 1 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"180.077031pt\" height=\"173.369063pt\" viewBox=\"0 0 180.077031 173.369063\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:48:00.079197</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 173.369063 \nL 180.077031 173.369063 \nL 180.077031 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 149.490938 \nL 171.8875 149.490938 \nL 171.8875 10.890938 \nL 33.2875 10.890938 \nz\n\"/>\n </g>\n <g clip-path=\"url(#pb525e67899)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAMEAAADBCAYAAAB2QtScAAEAAElEQVR4nLT9x7MsV57niX2OcPeQN65+WkKLBJCiMquyqjunpqeH3TZmM0PjLGjknhsajfwPuOCSS/4FpBlpRiM3XAynh6zqri7BElkCmUgACeBpefUN7eqc8+PiuMeN94DMqrKxcZgj4sYL4X7OT35/Sv1f/vf/B6E5lFIoFIhCKYVI/CeDwmoDCgRAK7xWiDIolSAYBIsojWjDuK54MpkwR1BW8we/9xN+50efkCYAniCewjtmlWdWwaJW5Chq5ynmBWFZkfb7lJ2UUAWWZ+dcHva5uz1kaD1IjQEUBgiICiiEIB5jNJlOccHjQ8B5z4uDl3TSFJVY7GDE6czx8MkJx+MaZy39XsreQLOZafqZpZ8YrFEoY1lUjtoJIoaqDpS1Y1YXzPKC+bKgKD1F6XEyRGyCWIsDvAQIAUdAjGBNzeZGl629bbr9hNQISXCgwIkQhLj2WqOUoAjE1W7PeMftEUKgrmsArLVorUGBF01VwfHhlNPTGeItW1sjLl/rYNPA8dGY06M5vkpQ0iUEB6bEoUnRXNtVdLOCP/0Pv+DkLHDl+nVUYvBK4ZXC1SWHZ6e4uub2pauMrOKHn9zgs18/xdHn1nuXCdLn+OgM7wLS0JLWlvVDrd1LQ2ZorRBVNfepUaIJCpwKWIR8PIGy4nvvvcHdGxskdYlyHqcMtROKoiR3DkvgUr/HnWt7KKNwSrFwNcfjMZWrSJIOic2oKs/J8Rl2/UpCCKsLVFrFZwpAXzBA+9i88+JRoxsG2uz08Srh2XTMpMz587/8S2xm+f4nH2C0xkvACQSlEARa5lOKrNelDoqyrpDEYNOEwc6Io8k5Z4/PeOfyJa71UnQoEV0haJDme1CIQOEdQYQAFM4x3N5FbI/xbMr84JzpvGR/u8e71xK8MiglXNkdYpUgImg8QsApjcwCroYgmm6WgOqxo/oEhCCaIIq8qJjOCs6nC86XcxaVoxKFUxkGg4hGnOL8tGIyO6I3zNjZ6rM56JAasCrglRAQUOG1tZVXiH9dYKVpSgiBsiyx1mKMQWlIrOLatS02NlKODs85OzugKDOuXt1nb3uffmfAwYtT8mUByoCkiHLUQTg6Lrh5rcfv/M7H/PG//3tOjg/ZuXwFiSSAAoL3IHGtUJogAa1BCQQvFGVJCO11R9r4rYdq3qtAKb1224JCoZUiBE8VPIHAw4NDAhU3dzbpKo14j0YwRqF9pMGgFF4Uofaczuec5XOCNWRJilLgvcMYw87uFlaLrDhRRYpHiDeoFIhEmeRDABUJNb4nEq6IBmWiFpHmtj1spV30yPJ4csyszPmPf/4XmNTy3jtvI0pTeYUP8XsiV8WLcARUJyUUJaquCSKIgXR7i/PzCX/z4Cnvb2/y5uUtjHGoAErahRNCgKAgCORFgagoXctiSaYVm3tb9G50SHTASs14POPw+Jgzhlzau4yVeF9BGZwH5QUdAkoFRAkiNZrQ3HuUYv2OZr+b4vYyajFM85rnx+ecLGG+LKmdQpTFq4S6hLmvyWenTLop2xs9RqMeOtEoI0DLBC0taNSa2Im3KatHrTWdToeqqijLkjRNSdIEKBltWrrdLU6OFpwe5zx+cMju3ga7+wNu3t7n5YszZhOHhLTRQopaMl4cLLlza5Mf/fB9/vyvP+V8cs5oawcd4tUE7xFFFARACD4yQRBCEMrCgawL0ZYZ1u5Dtc+FKAcbRhC9xviBqO+jlVB7Bxpqbfj68THFwvH29X0yDVCjtcJoBV7htWZWOSaTMdMqR1IT10VaCg9opUgSjfXORwJupfHqmi8uUqRR71qhG60Qr1M3N6oQFEEJGkGJYD2MbIc3t6/weHzEWbHkj//kz3Bo7r5xBxeg8hDa35Qoy4NWBCOk3Q6urEgUVCEgJiHtDym94peHZxxM5rx/9wZ7XR0lt+hmwQVEUZYlZVFikwSbGIaZkOgMaxSJKkEJTmA4GKJMxl/98pd8lG6wPxyiRVAYxAkSdCR+BUJA1IUq1zSbJx4IJCgSBd1ewvb1S4xLj9KKk/MZL8+mHE/nYDsEpxAPy6rCzxyz0yX9zS69UYe0b9C6XeDWLHpNcDaCaP0xyzKsTVguCvK8pNfrYq0hTQ2Xr4zo9/q8eH7K8eEZRbFk7/IW12/scZiMOTtZkoQE0RVBaQrp8PDFGW/duc7xeMYvf/2YJO0w6vej2SbgEZz3BAzSagKgKj2uVo1Ej4KiXS/VSPsV8bf3qC5eUa3Kaf9dCYihriq0NEzlhWVZ8/D5AUYF3rq5h27ea7TGBaFwnienZzhXk3Q62NREsa0Etfq10OyjaggwBCQEkACERtrEMxKBrEyhdntE1swhiYwgDUNopbAoNm2Pu7vX2Bts40rPn/7ZX/D4yTNcAC8QRC7WpZUVRoNRJInF+5pEG3SABI1Yi93b58my5i8/e8TReY4TTXspohTBeaplgUHTTTNSbTFaI4BzQl46fAAtgURBL+2wu73Hn//13zD3BbWuCarCh4qAIMo0Z7y/oDQeEwlAWVCWoBIEixKNFaGrPEO1YCstefvakB+/d5Xvv7nPhq0hnyNFjQ4KVwuLWc3hwTnPnx1xfDhhuSjxThC50DZIlPzSSM61hwvfzViGgw1S2+Hs
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 30
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:48:00.402774849Z",
"start_time": "2026-07-27T09:48:00.390428177Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [],
2026-07-27 17:57:55 +08:00
"execution_count": 31
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:48:00.463250908Z",
"start_time": "2026-07-27T09:48:00.403783355Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [],
2026-07-27 17:57:55 +08:00
"execution_count": 32
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:48:03.208374192Z",
"start_time": "2026-07-27T09:48:00.491960801Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
"read 1000 training examples\n",
"read 100 validation examples\n"
]
},
{
"data": {
"text/plain": [
"(torch.Size([32, 3, 256, 256]), torch.Size([32, 1, 5]))"
]
},
2026-07-27 17:57:55 +08:00
"execution_count": 33,
2026-07-27 14:18:15 +08:00
"metadata": {},
"output_type": "execute_result"
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 33
2026-07-27 14:18:15 +08:00
},
{
"metadata": {
"ExecuteTime": {
2026-07-27 17:57:55 +08:00
"end_time": "2026-07-27T09:48:03.533990096Z",
"start_time": "2026-07-27T09:48:03.269951948Z"
2026-07-27 14:18:15 +08:00
}
},
"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": [
"<Figure size 1000x400 with 10 Axes>"
],
2026-07-27 17:57:55 +08:00
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"572.4pt\" height=\"231.566897pt\" viewBox=\"0 0 572.4 231.566897\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:48:03.414413</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 231.566897 \nL 572.4 231.566897 \nL 572.4 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 7.2 103.406897 \nL 103.406897 103.406897 \nL 103.406897 7.2 \nL 7.2 7.2 \nz\n\"/>\n </g>\n <g clip-path=\"url(#pfeae1e10d8)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAIYAAACGCAYAAAAYefKRAACkk0lEQVR4nGT9169u6Zbeh/3GG2b44so7V+3Kdc6pEzvwNEOTapFNybQA2aIBwxYIGJIAAwZ853/B8J3/BPnCAYZt2SJNuZkkUp27D7tPrpx27bT2Sl+e4Q3DF/PbdZr2BjaqateqWt+a831HeMbzPEP+J/8r0XWjZANV6dCYKS28erfg/lGiGgkbtaxCyc1mCwRUHG2fKArFZJiODLNRwXYTeXKR2EVLDo52naCH0ATuH5UcH5V89TTy5FGLlZKm6xEVuk2kcAbvlDunIx68UtOawBfPAs+fJyZe+eF3Djk7HHF6+m1mR+9SjUqkyHg3psgF4nsKZ0mSEFVqVyIkQm4pCg8506UEZMiRcTmii4ZSwE/GiCS6Xhk5i4kNdUgc5Zov/vAXvDDCn3/1BaHr+OmHn/DzL59R1IYsltVOSQrWKa601EWJ9wVFWTOZ1iiJF1cv0AxlOeb07BZvPXyDd1454s7Zjj/4/X/J548u6a3FzTzGlRRaM60tB+PEwazi+eMlDx++zne+9R5tf8N6+ynb5QWbHAlR8eMjmt2O3CfW64YXzxo0JMqRJ4SELhzHP008/Eo5CZYjX2MdSDlmMzvkD1Y3/MIEzsvMVzc3ZLW4yheknGhzpC5hfuA5Pih47eEBRjY0IXFz03KzXmKMcnTgidkwnTqUBEk4OvAYJ6w7oSiFNip9EEwSHJEHDyacHU756tENz560GGNIOYARiOAQDg4sd18ZMS4nPH2+4unVhk2naBKcF2bjCuMM08MxD1+9R+kLCi+ELIgYjFP6GMg5Uzgh54DBUkqB855ClN12Q7TgrGKALu3Q0iMCZGEiDrdNVKGEpze0z27Qv3jBn330S/7x5WNanzgYG27PCnYxc90muqAYCwaDZgExWFfg/fA7EQElxoS1gXbXEPqe69WK68UTnl2vwbv91wRCn0iVJTUBPyrpNzuumg59/oIXyz9lXGXeeb1mNjkg7lpsCkhKSOzRlEghDe+lENQKaafIDraryDIIPgk29WQyKgVt7nGV53Q6ZR1abFrgRXFWDPNx5rj0HBx46jJxOFOqEVyuIueLHetWiaocTx11KSgZvKePihWLsZbVLrHYBroIpIz2AULmzq2SWyc1z58vubiM5GxBM1kyIWTGznJ62/Hg4QHXi5ZPP3nGZpOxhUM0MZ9ajieejz94gjrPeH7DO2/eUGjgtLyDsQf0VmhjQ1nWgEHJpBQwWRBj6OOOTdgBmRwyIpbYbbESgEy/WFNUU3KM/Oxf/hlvHbzK7CrSf7Vg+lT5RnvIP118yeU4cjCtmRVCFwJ9gJAUL2CNxYjFGEdRVBRliRpDDJGUlZAiJkRyirRty5PnS4zboH6MaMDmiGomhsiWHVrDujeEoOzU0C3XjNrEyaxiueo4O5pwbCf0dsM29OQu0PeBHPNwUJ0lpUwMiqxh3QrPrWULLFRJGDpJ9KnhJkWePltyHnvImW9/95s4UwQevjLFmZaoEWyiVzi/WnK+2NBkxRUFzkeqkSPmSFkK4jMhRUIMLFaWm0Vmu1VCD6HNlFk5OKwRNXz5xTWLm4xKiSszoVcMQlUIZ0cFB7OKx0+WXN30bBvIahAyh0cVpUuE1FMdzqkOPY/Pf8nPfrLh3tEIf3Sbwh5iJoeMT2dE8fSNISRPoQ5rPVEjDsG6AkyiSCBtT5EcqSzod5EP/ugX2NZw++SQf/tP/4Q/e/Gv+Z13fo23JweUUXm3OuTtas7n4YpV2zMpSro+EdMQeQQBNQgWxZJUcEVFSoGm6cgKMSlKjwrsug1d3OF9RFxF3/QkBbIl9RkrSgxwcbUmJ4tQ0PWBFDu8cTw570mpw4rhdDammkx4cXFN6JTQRiyCJghtJEVDv0sYX7OoKp5vN/gkqLNsbWDXLLh9eIs+BnxZ8OrpEf/L/81/jjs5g9lBIkVht850UelyJoYdIYNzgrWKtxbEsG4TUhYUxoCAWFhuAikacqsQDF4Nd29PyMHw6ScbchYQQSxY70g54zVzejzGi+WLL5dsOlARbAmltUwnJYVPGKv8td/6Hr/2G9+lmlskrHj60fvcNEumrWDiJTYVnB2fsesDJt2lDHe4uVhzdP+QZtPSXHUUI0dtLfOi5tnHH/DFX7zPqvbcee1NZF2x/NkjPrz8MV+ebzDLxCePn3I4b3klCbOc+LX5KX90fsW6TWwn0GZQVUprsd4NkSoP0QgszhU0XU/MQxrpY8Zqou07bpaXjEfQtT05N+SUidnSRyXlApuHfxeioqpgAqKKFc9m23JulV63OAPWW26d3aK0BRebLX0UskKKSkqKMUIfoG97boyyDoGUE8Y42pCJ3rPb7Ghz4hvffYPf/Y9+m9tvHOIm00wSZdFGVl2mKB0qhqABrMVakBwp6pK2z7RZMH1iJMJulykKt78tDmLL2BoObtUUheOTD5dst0NRVtZCDtB3LULm3t05tnd88vENvYVOwYrDFzCeeAqj9E3Hb/zwTb7/m9+lmMzwhcFXwts/eIduuwRTI1jcWOjtFcluaNcOeV7wz//J7/F3/ud/jXYDf/SPP+bi+VO+e/8WP3j4BqNtw4v/+sd8uN3w/3F/QKzGvDG/zQfPL/gqdPhtoHryhHt95qyeUAd4u5hwppanKXOzTXQpYw0UpcN6j4pgrFDsi8/h
2026-07-27 14:18:15 +08:00
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
2026-07-27 17:57:55 +08:00
"execution_count": 34
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:03.590014764Z",
"start_time": "2026-07-27T09:48:03.541625963Z"
}
},
"cell_type": "code",
"source": [
"from torch import nn\n",
"from torch.nn import functional as F\n",
"def cls_predictor(num_inputs,num_anchors,num_classes):\n",
" return nn.Conv2d(num_inputs,num_anchors*(num_classes+1), kernel_size=3,padding=1)\n",
"# 不改变原来图像的面积,提升通道数,每一个通道对应的是每一种锚框的一种类别的可能性\n",
"def bbox_predictor(num_inputs, num_anchors):\n",
" return nn.Conv2d(num_inputs, num_anchors * 4, kernel_size=3, padding=1)\n",
"# 不改变原来图像的面积,提升通道数,每一个通道对应的是每一种锚框距离正确标记框的偏移量\n"
],
"id": "eb1131ed49817939",
"outputs": [],
"execution_count": 35
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:03.659914611Z",
"start_time": "2026-07-27T09:48:03.591656234Z"
}
},
"cell_type": "code",
"source": [
"def forward(x, block):\n",
" return block(x)\n",
"Y1 = forward(torch.zeros((2, 8, 20, 20)), cls_predictor(8, 5, 10))\n",
"Y2 = forward(torch.zeros((2, 16, 10, 10)), cls_predictor(16, 3, 10))\n",
"Y1.shape, Y2.shape\n",
"#类别预测输出中的通道数分别为5 × (10 + 1) = 55和3 × (10 + 1) = 33,其中任一输出的形\n",
"#状是(批量大小,通道数,高度,宽度)"
],
"id": "f2a7c69cdd4aafa6",
"outputs": [
{
"data": {
"text/plain": [
"(torch.Size([2, 55, 20, 20]), torch.Size([2, 33, 10, 10]))"
]
},
"execution_count": 36,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 36
2026-07-27 14:18:15 +08:00
},
{
"metadata": {},
2026-07-27 17:57:55 +08:00
"cell_type": "markdown",
"source": "为了方便不同类别的图像进行训练,我们要把张量拍平再cat",
"id": "47e46c98c2e7d5e2"
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:03.761906859Z",
"start_time": "2026-07-27T09:48:03.661658265Z"
}
},
2026-07-27 14:18:15 +08:00
"cell_type": "code",
2026-07-27 17:57:55 +08:00
"source": [
"def flatten_pred(pred):\n",
" #要先调换维数才能拍平 0 2 3 1 (batch w h channel)\n",
" return torch.flatten(pred.permute(0,2,3,1), start_dim=1)\n",
"def concat_preds(preds):\n",
" return torch.cat([flatten_pred(p) for p in preds], dim=1)\n",
"concat_preds([Y1, Y2]).shape\n"
],
"id": "60a8cd1d9e8c1ad5",
"outputs": [
{
"data": {
"text/plain": [
"torch.Size([2, 25300])"
]
},
"execution_count": 37,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 37
},
{
"metadata": {},
"cell_type": "markdown",
"source": "下采样 提高感受野提高大物品的识别能力",
"id": "cb572fcd8e8d90a7"
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:03.827127144Z",
"start_time": "2026-07-27T09:48:03.784059223Z"
}
},
"cell_type": "code",
"source": [
"def down_sample_blk(in_channels,out_channels):\n",
" blk = []\n",
" for _ in range(2):\n",
" blk.append(nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)) #不改变大小 提升通道数哈 传统VGG做法\n",
" blk.append(nn.BatchNorm2d(out_channels))\n",
" blk.append(nn.ReLU())\n",
" in_channels = out_channels\n",
" blk.append(nn.MaxPool2d(2)) #降维\n",
" return nn.Sequential(*blk)"
],
"id": "a591f38441189d0c",
2026-07-27 14:18:15 +08:00
"outputs": [],
2026-07-27 17:57:55 +08:00
"execution_count": 38
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:03.921203148Z",
"start_time": "2026-07-27T09:48:03.827866637Z"
}
},
"cell_type": "code",
"source": [
"forward(torch.zeros((2, 3, 20, 20)), down_sample_blk(3, 10)).shape\n",
"#降维升channel"
],
"id": "8a771cd261a46563",
"outputs": [
{
"data": {
"text/plain": [
"torch.Size([2, 10, 10, 10])"
]
},
"execution_count": 39,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 39
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:04.128732201Z",
"start_time": "2026-07-27T09:48:03.922820914Z"
}
},
"cell_type": "code",
"source": [
"def base_net():\n",
" blk = []\n",
" num_filters = [3,16,32,64]\n",
" for i in range(len(num_filters)-1):\n",
" blk.append(down_sample_blk(num_filters[i], num_filters[i+1]))\n",
" return nn.Sequential(*blk)\n",
"\n",
"forward(torch.zeros((2, 3, 256, 256)), base_net()).shape"
],
"id": "d10ba91021ef0d0c",
"outputs": [
{
"data": {
"text/plain": [
"torch.Size([2, 64, 32, 32])"
]
},
"execution_count": 40,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 40
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:04.177851376Z",
"start_time": "2026-07-27T09:48:04.129824175Z"
}
},
"cell_type": "code",
"source": [
"def get_blk(i):\n",
" if i == 0:\n",
" blk = base_net()\n",
" elif i == 1:\n",
" blk = down_sample_blk(64, 128)\n",
" elif i == 4:\n",
" blk = nn.AdaptiveMaxPool2d((1,1))\n",
" else:\n",
" blk = down_sample_blk(128, 128)\n",
" return blk"
],
"id": "c17a584cd71a450a",
"outputs": [],
"execution_count": 41
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:04.244078688Z",
"start_time": "2026-07-27T09:48:04.179531133Z"
}
},
"cell_type": "code",
"source": [
"def blk_forward(X,blk,size,ratio,cls_predictor,bbox_predictor):\n",
" Y =blk(X)\n",
" anchors = d2l.multibox_prior(Y,sizes=size,ratios=ratio)\n",
" cls_preds = cls_predictor(Y)\n",
" bbox_preds = bbox_predictor(Y)\n",
" return (Y,anchors,cls_preds,bbox_preds)\n",
"\"\"\"\n",
"Y : CNN特征图\n",
"anchors : 在当前尺度下生成的锚框\n",
"cls_preds :每个锚框的类别\n",
"bbox_preds :每个锚框的偏移量\n",
"\"\"\""
],
"id": "d686dd6c196b17a1",
"outputs": [
{
"data": {
"text/plain": [
"'\\nY : CNN特征图\\nanchors : 在当前尺度下生成的锚框\\ncls_preds :每个锚框的类别\\nbbox_preds :每个锚框的偏移量\\n'"
]
},
"execution_count": 42,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 42
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:04.311459490Z",
"start_time": "2026-07-27T09:48:04.245645919Z"
}
},
"cell_type": "code",
"source": [
"sizes = [[0.2, 0.272], [0.37, 0.447], [0.54, 0.619], [0.71, 0.79],\n",
" [0.88, 0.961]]\n",
"ratios = [[1, 2, 0.5]] * 5\n",
"num_anchors = len(sizes[0]) + len(ratios[0]) - 1"
],
"id": "24fbc8fb3e72cbaf",
"outputs": [],
"execution_count": 43
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:04.357710443Z",
"start_time": "2026-07-27T09:48:04.348460151Z"
}
},
"cell_type": "code",
"source": [
"class TinySSD(nn.Module):\n",
" def __init__(self,num_classes,**kwargs):\n",
" super(TinySSD,self).__init__(**kwargs)\n",
" self.num_classes = num_classes\n",
" idx_to_in_channels = [64,128,128,128,128]\n",
" for i in range(5):\n",
" setattr(self,f\"blk_{i}\",get_blk(i))\n",
" setattr(self,f\"cls_{i}\",cls_predictor(idx_to_in_channels[i],num_anchors,num_classes))\n",
" setattr(self, f'bbox_{i}', bbox_predictor(idx_to_in_channels[i],num_anchors))\n",
" def forward(self,X):\n",
" anchors,cls_preds,bbox_preds=[None]*5,[None]*5,[None]*5\n",
" for i in range(5):\n",
" X,anchors[i],cls_preds[i],bbox_preds[i]=blk_forward(X,getattr(self,f'blk_{i}'),sizes[i],ratios[i],getattr(self, f'cls_{i}'), getattr(self, f'bbox_{i}'))\n",
" anchors = torch.cat(anchors,dim=1)\n",
" cls_preds = concat_preds(cls_preds)\n",
" cls_preds = cls_preds.reshape(cls_preds.shape[0],-1,self.num_classes+1)\n",
" bbox_preds = concat_preds(bbox_preds)\n",
" return anchors,cls_preds,bbox_preds"
],
"id": "642c9acb03ef05af",
"outputs": [],
"execution_count": 44
},
{
"metadata": {},
"cell_type": "markdown",
"source": "应该是生成(32^2 + 16^2 + 8^2 + 4^2 + 1) × 4 = 5444个锚框\n",
"id": "bb9119aecefba2b9"
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:05.245395003Z",
"start_time": "2026-07-27T09:48:04.358575367Z"
}
},
"cell_type": "code",
"source": [
"net = TinySSD(num_classes=1)\n",
"X = torch.zeros((32, 3, 256, 256))\n",
"anchors, cls_preds, bbox_preds = net(X)\n",
"print('output anchors:', anchors.shape)\n",
"print('output class preds:', cls_preds.shape)\n",
"print('output bbox preds:', bbox_preds.shape)"
],
"id": "c0f4abb77f5df63f",
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"output anchors: torch.Size([1, 5444, 4])\n",
"output class preds: torch.Size([32, 5444, 2])\n",
"output bbox preds: torch.Size([32, 21776])\n"
]
}
],
"execution_count": 45
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:08.014082271Z",
"start_time": "2026-07-27T09:48:05.279839380Z"
}
},
"cell_type": "code",
"source": [
"batch_size = 32\n",
"train_iter, _ = d2l.load_data_bananas(batch_size)"
],
"id": "8cdaf93f4dce3208",
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"read 1000 training examples\n",
"read 100 validation examples\n"
]
}
],
"execution_count": 46
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:08.178564673Z",
"start_time": "2026-07-27T09:48:08.120907893Z"
}
},
"cell_type": "code",
"source": [
"device, net = d2l.try_gpu(), TinySSD(num_classes=1)\n",
"trainer = torch.optim.SGD(net.parameters(), lr=0.2, weight_decay=5e-4)"
],
"id": "4d40936938c9b028",
"outputs": [],
"execution_count": 47
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:08.336189213Z",
"start_time": "2026-07-27T09:48:08.180393347Z"
}
},
"cell_type": "code",
"source": [
"\n",
"cls_loss = nn.CrossEntropyLoss(reduction='none')\n",
"bbox_loss = nn.L1Loss(reduction='none')\n",
"def calc_loss(cls_preds,cls_labels,bbox_preds,bbox_labels,bbox_masks):\n",
" batch_size,num_classes = cls_preds.shape[0] , cls_preds.shape[2]\n",
" cls = cls_loss(cls_preds.reshape(-1,num_classes),cls_labels.reshape(-1)).reshape(batch_size,-1).mean(dim=1)\n",
" bbox = bbox_loss(bbox_preds*bbox_masks,bbox_labels*bbox_masks).mean(dim=1)\n",
" return cls + bbox"
],
"id": "8e07cd2f6db67147",
"outputs": [],
"execution_count": 48
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:48:08.399558542Z",
"start_time": "2026-07-27T09:48:08.343338340Z"
}
},
"cell_type": "code",
"source": [
"# 验证 类别预测用acc 锚框偏移量预测用L1范数\n",
"def cls_eval(cls_preds,cls_labels)->float:\n",
" return float((cls_preds.argmax(dim=-1).type(cls_labels.dtype)==cls_labels).sum())\n",
"def bbox_eval(bbox_preds, bbox_labels, bbox_masks):\n",
" return float((torch.abs((bbox_labels - bbox_preds) * bbox_masks)).sum())"
],
"id": "e2dd1576866afef9",
"outputs": [],
"execution_count": 49
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:50:07.366475114Z",
"start_time": "2026-07-27T09:48:08.401037720Z"
}
},
"cell_type": "code",
"source": [
"num_epochs, timer = 20, d2l.Timer()\n",
"animator = d2l.Animator(xlabel='epoch', xlim=[1, num_epochs],\n",
"legend=['class error', 'bbox mae'])\n",
"net = net.to(device)\n",
"for epoch in range(num_epochs):\n",
" metric = d2l.Accumulator(4)\n",
" net.train()\n",
" for feature , target in train_iter:\n",
" # feature -> (batch,3,h,w) target (batch,num,5) 其中5维数据分别为 label h1 w1 h2 w2\n",
" timer.start()\n",
" trainer.zero_grad()\n",
" X,Y=feature.to(device),target.to(device)\n",
" anchors,cls_preds,bbox_preds = net(X)\n",
" bbox_labels, bbox_masks, cls_labels = d2l.multibox_target(anchors, Y)\n",
" l = calc_loss(cls_preds, cls_labels, bbox_preds, bbox_labels,\n",
" bbox_masks)\n",
" l.mean().backward()\n",
" trainer.step()\n",
" metric.add(cls_eval(cls_preds, cls_labels), cls_labels.numel(),\n",
" bbox_eval(bbox_preds, bbox_labels, bbox_masks),\n",
" bbox_labels.numel())\n",
" #计数器依次是 cls正确率 分类任务总个数 box偏移量l1范数 框选任务总个数\n",
" cls_err, bbox_mae = 1 - metric[0] / metric[1], metric[2] / metric[3]\n",
" animator.add(epoch + 1, (cls_err, bbox_mae))\n",
"print(f'class err {cls_err:.2e}, bbox mae {bbox_mae:.2e}')"
],
"id": "720fb1fc6b3bfaed",
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"class err 3.31e-03, bbox mae 3.15e-03\n"
]
},
{
"data": {
"text/plain": [
"<Figure size 350x250 with 1 Axes>"
],
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"251.690625pt\" height=\"183.35625pt\" viewBox=\"0 0 251.690625 183.35625\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:50:07.189022</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 183.35625 \nL 251.690625 183.35625 \nL 251.690625 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 42.828125 145.8 \nL 238.128125 145.8 \nL 238.128125 7.2 \nL 42.828125 7.2 \nz\n\"/>\n </g>\n <g id=\"matplotlib.axis_1\">\n <g id=\"xtick_1\">\n <g id=\"line2d_1\">\n <path d=\"M 83.943914 145.8 \nL 83.943914 7.2 \n\" clip-path=\"url(#pebfd2d9840)\" style=\"fill: none; stroke: #ffffff; stroke-width: 0.8; stroke-linecap: square\"/>\n </g>\n <g id=\"line2d_2\">\n <defs>\n <path id=\"md80e273421\" d=\"M 0 0 \nL 0 3.5 \n\" style=\"stroke: #ffffff; stroke-width: 0.8\"/>\n </defs>\n <g>\n <use xlink:href=\"#md80e273421\" x=\"83.943914\" y=\"145.8\" style=\"fill: #ffffff; stroke: #ffffff; stroke-width: 0.8\"/>\n </g>\n </g>\n <g id=\"text_1\">\n <!-- 5 -->\n <g style=\"fill: #ffffff\" transform=\"translate(80.762664 160.398438) scale(0.1 -0.1)\">\n <defs>\n <path id=\"DejaVuSans-35\" d=\"M 691 4666 \nL 3169 4666 \nL 3169 4134 \nL 1269 4134 \nL 1269 2991 \nQ 1406 3038 1543 3061 \nQ 1681 3084 1819 3084 \nQ 2600 3084 3056 2656 \nQ 3513 2228 3513 1497 \nQ 3513 744 3044 326 \nQ 2575 -91 1722 -91 \nQ 1428 -91 1123 -41 \nQ 819 9 494 109 \nL 494 744 \nQ 775 591 1075 516 \nQ 1375 441 1709 441 \nQ 2250 441 2565 725 \nQ 2881 1009 2881 1497 \nQ 2881 1984 2565 2268 \nQ 2250 2553 1709 2553 \nQ 1456 2553 1204 2497 \nQ 953 2441 691 2322 \nL 691 4666 \nz\n\" transform=\"scale(0.015625)\"/>\n </defs>\n <use xlink:href=\"#DejaVuSans-35\"/>\n </g>\n </g>\n </g>\n <g id=\"xtick_2\">\n <g id=\"line2d_3\">\n <path d=\"M 135.338651 145.8 \nL 135.338651 7.2 \n\" clip-path=\"url(#pebfd2d9840)\" style=\"fill: none; stroke: #ffffff; stroke-width: 0.8; stroke-linecap: square\"/>\n </g>\n <g id=\"line2d_4\">\n <g>\n <use xlink:href=\"#md80e273421\" x=\"135.338651\" y=\"145.8\" style=\"fill: #ffffff; stroke: #ffffff; stroke-width: 0.8\"/>\n </g>\n </g>\n <g id=\"text_2\">\n <!-- 10 -->\n <g style=\"fill: #ffffff\" transform=\"translate(128.976151 160.398438) scale(0.1 -0.1)\">\n <defs>\n <path id=\"DejaVuSans-31\" d=\"M 794 531 \nL 1825 531 \nL 1825 4091 \nL 703 3866 \nL 703 4441 \nL 1819 4666 \nL 2450 4666 \nL 2450 531 \nL 3481 531 \nL 3481 0 \nL 794 0 \nL 794 531 \nz\n\" transform=\"scale(0.015625)\"/>\n <path id=\"DejaVuSans-30\" d=\"M 2034 4250 \nQ 1547 4250 1301 3770 \nQ 1056 3291 1056 2328 \nQ 1056 1369 1301 889 \nQ 1547 409 2034 409 \nQ 2525 409 2770 889 \nQ 3016 1369 3016 2328 \nQ 3016 3291 2770 3770 \nQ 2525 4250 2034 4250 \nz\nM 2034 4750 \nQ 2819 4750 3233 4129 \nQ 3647 3509 3647 2328 \nQ 3647 1150 3233 529 \nQ 2819 -91 2034 -91 \nQ 1250 -91 836 529 \nQ 422 1150 422 2328 \nQ 422 3509 836 4129 \nQ 1250 4750 2034 4750 \nz\n\" transform=\"scale(0.015625)\"/>\n </defs>\n <use xlink:href=\"#DejaVuSans-31\"/>\n <use xlink:href=\"#DejaVuSans-30\" x=\"63.6
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
"execution_count": 50
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:50:07.767970800Z",
"start_time": "2026-07-27T09:50:07.458706602Z"
}
},
"cell_type": "code",
"source": [
"X = torchvision.io.read_image('../data/banana.jpg').unsqueeze(0).float()\n",
"img = X.squeeze(0).permute(1,2,0).long()\n"
],
"id": "b6a5ae7b16b60e50",
"outputs": [],
"execution_count": 51
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:50:08.153935112Z",
"start_time": "2026-07-27T09:50:07.880617944Z"
}
},
"cell_type": "code",
"source": [
"def predict(X):\n",
" net.eval()\n",
" anchors, cls_preds, bbox_preds = net(X.to(device))\n",
" cls_probs = F.softmax(cls_preds, dim=2).permute(0, 2, 1)\n",
" output = d2l.multibox_detection(cls_probs, bbox_preds, anchors)\n",
" idx = [i for i, row in enumerate(output[0]) if row[0] != -1]\n",
" return output[0, idx]\n",
"output = predict(X)"
],
"id": "64a9460651b1b585",
"outputs": [],
"execution_count": 52
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:50:08.240708202Z",
"start_time": "2026-07-27T09:50:08.155681155Z"
}
},
"cell_type": "code",
"source": "output",
"id": "e8dc16888118c7bc",
"outputs": [
{
"data": {
"text/plain": [
"tensor([[ 0.0000, 0.9985, 0.0642, 0.7531, 0.2869, 0.9532],\n",
" [ 0.0000, 0.9973, 0.4608, 0.5666, 0.6662, 0.7837],\n",
" [ 0.0000, 0.9970, 0.5375, 0.0591, 0.7508, 0.2796],\n",
" [ 0.0000, 0.9902, 0.7037, 0.3652, 0.9121, 0.5747],\n",
" [ 0.0000, 0.4860, 0.5226, 0.0069, 0.7186, 0.2129],\n",
" [ 0.0000, 0.4443, 0.4504, 0.6233, 0.6205, 0.8134],\n",
" [ 0.0000, 0.2552, 0.5831, 0.0923, 0.8302, 0.3042],\n",
" [ 0.0000, 0.2356, 0.5915, -0.0019, 0.7679, 0.2016],\n",
" [ 0.0000, 0.2331, 0.4954, 0.6350, 0.7170, 0.8248],\n",
" [ 0.0000, 0.1502, 0.7096, 0.2895, 0.8824, 0.5102],\n",
" [ 0.0000, 0.1475, 0.1230, 0.7664, 0.3295, 0.9989],\n",
" [ 0.0000, 0.1150, 0.5567, -0.0957, 0.7636, 0.1090],\n",
" [ 0.0000, 0.1137, 0.1259, 0.6939, 0.3409, 0.9125],\n",
" [ 0.0000, 0.1115, 0.4028, 0.6006, 0.5825, 0.7848],\n",
" [ 0.0000, 0.1093, 0.4543, 0.5068, 0.7193, 0.7340],\n",
" [ 0.0000, 0.1016, 0.5291, 0.1110, 0.7004, 0.3227],\n",
" [ 0.0000, 0.0900, 0.4455, 0.4652, 0.6391, 0.6875],\n",
" [ 0.0000, 0.0780, 0.4589, 0.6589, 0.6681, 0.8913],\n",
" [ 0.0000, 0.0652, 0.5598, 0.1803, 0.7430, 0.3703],\n",
" [ 0.0000, 0.0605, 0.6292, 0.3817, 0.8251, 0.5817],\n",
" [ 0.0000, 0.0514, 0.4952, 0.7167, 0.6964, 0.9586],\n",
" [ 0.0000, 0.0469, 0.6734, 0.4222, 0.8798, 0.6158],\n",
" [ 0.0000, 0.0416, 0.4532, 0.0716, 0.6357, 0.2691],\n",
" [ 0.0000, 0.0391, 0.4798, -0.0191, 0.6770, 0.1747],\n",
" [ 0.0000, 0.0389, 0.0044, 0.7313, 0.1998, 0.9244],\n",
" [ 0.0000, 0.0389, 0.5262, 0.4644, 0.6896, 0.6895],\n",
" [ 0.0000, 0.0374, 0.6148, 0.1649, 0.7902, 0.3841],\n",
" [ 0.0000, 0.0359, 0.0650, 0.8109, 0.2880, 1.0482],\n",
" [ 0.0000, 0.0359, 0.0698, 0.6747, 0.2262, 0.9140],\n",
" [ 0.0000, 0.0319, 0.6276, -0.1184, 0.7822, 0.1588],\n",
" [ 0.0000, 0.0319, -0.1519, 0.2342, 1.2012, 0.8651],\n",
" [ 0.0000, 0.0285, 0.1491, 0.6391, 0.3408, 0.8407],\n",
" [ 0.0000, 0.0282, 0.6711, 0.2857, 0.8190, 0.5322],\n",
" [ 0.0000, 0.0281, 0.1795, -0.1376, 0.7880, 1.1077],\n",
" [ 0.0000, 0.0271, 0.4184, 0.6820, 0.6104, 0.8625],\n",
" [ 0.0000, 0.0266, 0.6531, 0.2491, 0.8567, 0.4469],\n",
" [ 0.0000, 0.0215, 0.3868, 0.5555, 0.6020, 0.7265],\n",
" [ 0.0000, 0.0164, 0.5367, -0.1162, 0.6877, 0.1575],\n",
" [ 0.0000, 0.0160, 0.3315, 0.6025, 0.5351, 0.7996],\n",
" [ 0.0000, 0.0156, 0.2120, 0.6672, 0.3861, 0.8899],\n",
" [ 0.0000, 0.0154, -0.3605, -0.1560, 0.6862, 0.3646],\n",
" [ 0.0000, 0.0149, 0.6149, 0.0149, 0.8310, 0.2521],\n",
" [ 0.0000, 0.0141, 0.1019, 0.6182, 0.2584, 0.8786],\n",
" [ 0.0000, 0.0131, -0.1225, -0.3201, 0.2735, 0.4565],\n",
" [ 0.0000, 0.0130, 0.4430, 0.5082, 1.2768, 1.2750],\n",
" [ 0.0000, 0.0125, 0.1837, 0.7573, 0.3872, 0.9679],\n",
" [ 0.0000, 0.0123, 0.4169, 0.4787, 0.5627, 0.7412],\n",
" [ 0.0000, 0.0122, 0.7223, 0.2193, 0.9493, 0.4536],\n",
" [ 0.0000, 0.0121, 0.4850, 0.4219, 0.7411, 0.6612],\n",
" [ 0.0000, 0.0112, 0.6911, 0.4813, 0.9275, 0.7128],\n",
" [ 0.0000, 0.0105, 0.5652, 0.7101, 0.7331, 0.9001],\n",
" [ 0.0000, 0.0102, 0.0054, 0.8066, 0.2270, 0.9863],\n",
" [ 0.0000, 0.0101, 0.6626, 0.0780, 0.8909, 0.3222]],\n",
" device='cuda:0', grad_fn=<IndexBackward0>)"
]
},
"execution_count": 53,
"metadata": {},
"output_type": "execute_result"
}
],
"execution_count": 53
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:50:08.489689738Z",
"start_time": "2026-07-27T09:50:08.274956342Z"
}
},
"cell_type": "code",
"source": [
"def display(img, output, threshold):\n",
" d2l.set_figsize((5, 5))\n",
" fig = d2l.plt.imshow(img)\n",
" for row in output:\n",
" score = float(row[1])\n",
" if score < threshold:\n",
" continue\n",
" h, w = img.shape[0:2]\n",
" bbox = [row[2:6] * torch.tensor((w, h, w, h), device=row.device)]\n",
" d2l.show_bboxes(fig.axes, bbox, '%.2f' % score, 'w')\n",
"display(img, output.cpu(), threshold=0.9)"
],
"id": "add91b659a0efa1c",
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor([0.0000, 0.9985, 0.0642, 0.7531, 0.2869, 0.9532],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.9973, 0.4608, 0.5666, 0.6662, 0.7837],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.9970, 0.5375, 0.0591, 0.7508, 0.2796],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.9902, 0.7037, 0.3652, 0.9121, 0.5747],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.4860, 0.5226, 0.0069, 0.7186, 0.2129],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.4443, 0.4504, 0.6233, 0.6205, 0.8134],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.2552, 0.5831, 0.0923, 0.8302, 0.3042],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.2356, 0.5915, -0.0019, 0.7679, 0.2016],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.2331, 0.4954, 0.6350, 0.7170, 0.8248],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.1502, 0.7096, 0.2895, 0.8824, 0.5102],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.1475, 0.1230, 0.7664, 0.3295, 0.9989],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.1150, 0.5567, -0.0957, 0.7636, 0.1090],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.1137, 0.1259, 0.6939, 0.3409, 0.9125],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.1115, 0.4028, 0.6006, 0.5825, 0.7848],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.1093, 0.4543, 0.5068, 0.7193, 0.7340],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.1016, 0.5291, 0.1110, 0.7004, 0.3227],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0900, 0.4455, 0.4652, 0.6391, 0.6875],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0780, 0.4589, 0.6589, 0.6681, 0.8913],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0652, 0.5598, 0.1803, 0.7430, 0.3703],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0605, 0.6292, 0.3817, 0.8251, 0.5817],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0514, 0.4952, 0.7167, 0.6964, 0.9586],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0469, 0.6734, 0.4222, 0.8798, 0.6158],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0416, 0.4532, 0.0716, 0.6357, 0.2691],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.0391, 0.4798, -0.0191, 0.6770, 0.1747],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0389, 0.0044, 0.7313, 0.1998, 0.9244],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0389, 0.5262, 0.4644, 0.6896, 0.6895],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0374, 0.6148, 0.1649, 0.7902, 0.3841],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0359, 0.0650, 0.8109, 0.2880, 1.0482],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0359, 0.0698, 0.6747, 0.2262, 0.9140],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.0319, 0.6276, -0.1184, 0.7822, 0.1588],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.0319, -0.1519, 0.2342, 1.2012, 0.8651],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0285, 0.1491, 0.6391, 0.3408, 0.8407],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0282, 0.6711, 0.2857, 0.8190, 0.5322],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.0281, 0.1795, -0.1376, 0.7880, 1.1077],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0271, 0.4184, 0.6820, 0.6104, 0.8625],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0266, 0.6531, 0.2491, 0.8567, 0.4469],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0215, 0.3868, 0.5555, 0.6020, 0.7265],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.0164, 0.5367, -0.1162, 0.6877, 0.1575],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0160, 0.3315, 0.6025, 0.5351, 0.7996],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0156, 0.2120, 0.6672, 0.3861, 0.8899],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.0154, -0.3605, -0.1560, 0.6862, 0.3646],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0149, 0.6149, 0.0149, 0.8310, 0.2521],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0141, 0.1019, 0.6182, 0.2584, 0.8786],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([ 0.0000, 0.0131, -0.1225, -0.3201, 0.2735, 0.4565],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0130, 0.4430, 0.5082, 1.2768, 1.2750],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0125, 0.1837, 0.7573, 0.3872, 0.9679],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0123, 0.4169, 0.4787, 0.5627, 0.7412],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0122, 0.7223, 0.2193, 0.9493, 0.4536],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0121, 0.4850, 0.4219, 0.7411, 0.6612],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0112, 0.6911, 0.4813, 0.9275, 0.7128],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0105, 0.5652, 0.7101, 0.7331, 0.9001],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0102, 0.0054, 0.8066, 0.2270, 0.9863],\n",
" grad_fn=<UnbindBackward0>)\n",
"tensor([0.0000, 0.0101, 0.6626, 0.0780, 0.8909, 0.3222],\n",
" grad_fn=<UnbindBackward0>)\n"
]
},
{
"data": {
"text/plain": [
"<Figure size 500x500 with 1 Axes>"
],
"image/svg+xml": "<?xml version=\"1.0\" encoding=\"utf-8\" standalone=\"no\"?>\n<!DOCTYPE svg PUBLIC \"-//W3C//DTD SVG 1.1//EN\"\n \"http://www.w3.org/Graphics/SVG/1.1/DTD/svg11.dtd\">\n<svg xmlns:xlink=\"http://www.w3.org/1999/xlink\" width=\"321.275781pt\" height=\"311.535938pt\" viewBox=\"0 0 321.275781 311.535938\" xmlns=\"http://www.w3.org/2000/svg\" version=\"1.1\">\n <metadata>\n <rdf:RDF xmlns:dc=\"http://purl.org/dc/elements/1.1/\" xmlns:cc=\"http://creativecommons.org/ns#\" xmlns:rdf=\"http://www.w3.org/1999/02/22-rdf-syntax-ns#\">\n <cc:Work>\n <dc:type rdf:resource=\"http://purl.org/dc/dcmitype/StillImage\"/>\n <dc:date>2026-07-27T17:50:08.344236</dc:date>\n <dc:format>image/svg+xml</dc:format>\n <dc:creator>\n <cc:Agent>\n <dc:title>Matplotlib v3.7.2, https://matplotlib.org/</dc:title>\n </cc:Agent>\n </dc:creator>\n </cc:Work>\n </rdf:RDF>\n </metadata>\n <defs>\n <style type=\"text/css\">*{stroke-linejoin: round; stroke-linecap: butt}</style>\n </defs>\n <g id=\"figure_1\">\n <g id=\"patch_1\">\n <path d=\"M 0 311.535938 \nL 321.275781 311.535938 \nL 321.275781 0 \nL 0 0 \nz\n\"/>\n </g>\n <g id=\"axes_1\">\n <g id=\"patch_2\">\n <path d=\"M 33.2875 287.657813 \nL 310.4875 287.657813 \nL 310.4875 10.457813 \nL 33.2875 10.457813 \nz\n\"/>\n </g>\n <g clip-path=\"url(#pf88e60ce7b)\">\n <image xlink:href=\"data:image/png;base64,\niVBORw0KGgoAAAANSUhEUgAAAYEAAAGBCAYAAACAWQ0kAAEAAElEQVR4nFT9WbNm23Wmhz2zW+3X7j670+IUAAIkUUVSVClUtkuybvwnbP8Y+dbXtsNXtm4U4VDIkiIsqylLVjUsklVFAiCAg9Nnu7uvX+3sfDFXZlmByAjEOScz997fWnOO8Y73fYZ48QerOAw9fd8TIqzXc1arBdvNASklSima9kA9r5gtam5vb8mKjKqu2G2PrNYrLi/O2dw/EkMgxshpP9B1lr4fKcocISMQKcsSgBACmZ6htUYbxbHZobUkLzJcgBgjAEopvPdYa8mEoe8HTqcTs2rGarXi6uaS7emBtm9o2hPr1TnOWo7HA5989DF2HLh/uOXpzRPsaHl82FCaBSHAMAxoDUKClBHnHcYo6rqi7Xu0ViyrGVgY+5Hm2CBV+tqcczTHFiUUy9mMthkJDkTQDIMgBIEQgq4/IiXM5gXBDWRGs16vCaOHCEIIlsuaIs8oyozdZss4DHgfOexPxBipqgoberJc8/z5M+4fN7RNS9v2zKol8/mCjz/+mLt3r+n7Bm0El+cXlGWJUILjcc84WoTWCCGxo+Xt23vqYoZWBmsd+8ORYRyBiNQKKSVSgdYKKQXOWZ4+uWG5XGCM4cvffsn9/T1FbcjKDJWl35OpGkFG3zmi0CA1WVlxPJ7o+wGAw/5I33XkeYaNliA99bpEVxkyV3gs/TAyWsc4Os7PL1ivz1gu5zSnPV3bkBnJ0HXYfqDKCoIPeJeek7zIyfKMzlqQmoDg1etbvPUoofj5z/6Aw2bL9vEB8CgFKn2pKCPRRlPM5uz2J7a7E3Vd8/z5c57cPOF//B//KauLiicvzpEGQBICvLu9xVpLCIGzszOkEIgQiR4W9Zz1csXjw45M5xSm5G/+5m8QUlPP51w/uURnAkzg2YsXhCDoegeAFAIlBd9+/S1GGpb1is27PbuHPXevH5ASFmczbj665PMff0aR5yihuH/c0o0DzdjzL//53zKfVfzsZ58S/IAxUJYKGztUrigWJSCwvaM9dFRmTpkVzIoZf/u3f0uWaT770Se8fXePdRGlKjabHfvDkdv7Bz7/4pL5oqCsM7RRxBgZxpGHh0e6bqDtLc+eX6NUZHd84JNPPsbawK9/+T2zWUWM0Bxbrq7OyHNDiI5+OJHnmucfP+Hu7h3ee87Oznj79gHv4fryhhACSkoWiwXfffcdTXOkrCRXV1cYnbF52PPmTUPXO8pKc35xjVKKu7dvMZnG5IZyNsNFizKS65tzXr58xX53xI/w5Po5RV7y7t07zs7OmM1mzJdL/vW//JLHux1/8NMX7A8HooCf/OIPmc0rtJHstvc0hxNu9JwtL3Cjp206/u7Xv+Lm6TWL5YyHh1uavmGwIxHB0+fPmM3n/PJXv+X6ZsnZWY33HiEFQoJzjqqqyYuKEGA+W1DXMzKjOBx2HE97un5LlhUURc2zmy/45S9/w5dffsUnn1xRFBnGGBjh8XHD4bjnx59/QggO50e0lBIQACwWM4iw3x8JIeK9w1pHJB1qQkiKqkQpRYwCY3Lc6NluD3gXkEIgpSLGAa0VVVViMk0kEGPAWovWhjwv0UoTQmQYerRSSCWIEWKIINLXY63Fe4dzDiHSPy7LktmsJssNzluGYWAYRuzo8N4hhKDIc/quwzmHFIqu6wg+oJSiHwbKouLZs2dsdw94bwGfXlwkMUJmDEIIhnEkDoEYIkWRM4w9AFVV0Tc9xIj3nuA9wQMBpMxACLwP5EWJUiClAuGQQiIRDNaipGK1WlKVBUJC0zQ45wBBUeT03YBzfvpcFphMMwwDSkmUNvT9ibpKF6XWGjUd3iF4nPd479FKgxBIJSmqis3jlqZpAYhAiJFAJMRIJGKMRsj0g5YSlJIIAd5ZRjtiraXIcy4uzimKHKkiQQakFswXS0TMiV7SZiPWgwuRYRwYhp5hGD58dlVZEqMnyJwgAzpTSJUuHK0LYkzPmhQKrRQxBE7HI23TMAwtmZmhpCJoTdt3ENITXJYFgUg/DFjn8DiikDx5ck3X9Iz9yND3eO+QSuKsJcs0eW6IwoOI+JCeUwBjNN57jscDRimUCoTg6PqOi8U5XdfTNA3L5ZLmdOJwPHI8njBKoaVEC4Nzjr4faE4tcqYpFiWffv4FzntijBhj6MeW0+HA9ZMnmKygqjS73QEpBcZIpBQE72iaE8PQo6RkvV7hvCXLMmKM2NEikUgCVVkxOsdhf2C5qijLHOcsu+2GotBk2QIfAlHAOFiGfoQoyLP0Xnof2O32zOoZUgvatkvP/ODY7xuUktR1wXmYkxmN1pqiqJBKYN2Ic5bROaz3xAin5oQ2AiElznlGG/ABrHPp5BFwOHaYfqQoBNZ6pJLEAEbnCOHouxEhJFoLEND3HTGmZ1RrRV2XZHk6N3zwIAR5oZBKUtUV49ATIxRlnt4VpXDOcWx6hBKcnTvKsiIGOGw6nHU45VgulywWC4qixI6W5bIg+pr7+y1SRXSmub19i1Q3zBc1RmvOz89RKGTU1OczQOLsSIgOZx2z2QwbHIHI9ZMnhBjYbjdoHfHe0vcDzlnqWUWRF1CA94G2bcjzis1mw+PjhrP1EqUEdV3TtFtAIIXkeDyRZzlXl5cMg0UrjVGw3e4YhwElBIfDEalA
},
"metadata": {},
"output_type": "display_data",
"jetTransient": {
"display_id": null
}
}
],
"execution_count": 54
},
{
"metadata": {
"ExecuteTime": {
"end_time": "2026-07-27T09:50:08.699133568Z",
"start_time": "2026-07-27T09:50:08.648954560Z"
}
},
"cell_type": "code",
2026-07-27 14:18:15 +08:00
"source": "",
2026-07-27 17:57:55 +08:00
"id": "4ec2f61624100d07",
"outputs": [],
"execution_count": 54
2026-07-27 14:18:15 +08:00
}
],
"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
}