diff --git a/chapter13.ipynb b/chapter13.ipynb index d8aff04..d3bacef 100644 --- a/chapter13.ipynb +++ b/chapter13.ipynb @@ -9,8 +9,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:15.242189006Z", - "start_time": "2026-07-26T11:57:12.026651705Z" + "end_time": "2026-07-27T09:47:55.425519349Z", + "start_time": "2026-07-27T09:47:53.873571522Z" } }, "cell_type": "code", @@ -20,23 +20,14 @@ "from d2l import torch as d2l" ], "id": "b39ef644f6c80c7c", - "outputs": [ - { - "name": "stderr", - "output_type": "stream", - "text": [ - "/home/yukun/.conda/envs/nn/lib/python3.11/site-packages/torch/cuda/__init__.py:1007: UserWarning: Can't initialize NVML\n", - " raw_cnt = _raw_device_count_nvml()\n" - ] - } - ], + "outputs": [], "execution_count": 1 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:15.585413376Z", - "start_time": "2026-07-26T11:57:15.270700840Z" + "end_time": "2026-07-27T09:47:55.671442114Z", + "start_time": "2026-07-27T09:47:55.428215388Z" } }, "cell_type": "code", @@ -52,7 +43,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T19:57:15.520588\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:47:55.619678\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -66,8 +57,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:15.692690288Z", - "start_time": "2026-07-26T11:57:15.588725031Z" + "end_time": "2026-07-27T09:47:55.747558006Z", + "start_time": "2026-07-27T09:47:55.673122195Z" } }, "cell_type": "code", @@ -98,8 +89,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:15.705559306Z", - "start_time": "2026-07-26T11:57:15.694772245Z" + "end_time": "2026-07-27T09:47:55.853579240Z", + "start_time": "2026-07-27T09:47:55.749428064Z" } }, "cell_type": "code", @@ -111,8 +102,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:15.794704197Z", - "start_time": "2026-07-26T11:57:15.707000004Z" + "end_time": "2026-07-27T09:47:55.959910400Z", + "start_time": "2026-07-27T09:47:55.859086964Z" } }, "cell_type": "code", @@ -139,8 +130,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:15.815124864Z", - "start_time": "2026-07-26T11:57:15.799009123Z" + "end_time": "2026-07-27T09:47:57.527570250Z", + "start_time": "2026-07-27T09:47:57.492961332Z" } }, "cell_type": "code", @@ -159,8 +150,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.101228908Z", - "start_time": "2026-07-26T11:57:18.498541915Z" + "end_time": "2026-07-27T09:47:57.635071582Z", + "start_time": "2026-07-27T09:47:57.528293612Z" } }, "cell_type": "code", @@ -174,7 +165,7 @@ { "data": { "text/plain": [ - "" + "" ] }, "execution_count": 7, @@ -186,7 +177,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T19:57:18.591473\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:47:57.594042\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -200,8 +191,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.221032646Z", - "start_time": "2026-07-26T11:57:19.137271191Z" + "end_time": "2026-07-27T09:47:57.702782485Z", + "start_time": "2026-07-27T09:47:57.650116830Z" } }, "cell_type": "code", @@ -254,8 +245,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.258339667Z", - "start_time": "2026-07-26T11:57:19.246797120Z" + "end_time": "2026-07-27T09:47:57.754193331Z", + "start_time": "2026-07-27T09:47:57.704574735Z" } }, "cell_type": "code", @@ -270,8 +261,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.382932850Z", - "start_time": "2026-07-26T11:57:19.281193532Z" + "end_time": "2026-07-27T09:47:57.820044630Z", + "start_time": "2026-07-27T09:47:57.755572144Z" } }, "cell_type": "code", @@ -306,8 +297,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.433949316Z", - "start_time": "2026-07-26T11:57:19.384210708Z" + "end_time": "2026-07-27T09:47:57.877041771Z", + "start_time": "2026-07-27T09:47:57.821962159Z" } }, "cell_type": "code", @@ -340,8 +331,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.483140764Z", - "start_time": "2026-07-26T11:57:19.435383512Z" + "end_time": "2026-07-27T09:47:57.924927279Z", + "start_time": "2026-07-27T09:47:57.877778001Z" } }, "cell_type": "code", @@ -353,8 +344,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.621504492Z", - "start_time": "2026-07-26T11:57:19.484527807Z" + "end_time": "2026-07-27T09:47:58.080082507Z", + "start_time": "2026-07-27T09:47:57.926102351Z" } }, "cell_type": "code", @@ -372,7 +363,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T19:57:19.568775\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:47:58.021417\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -386,8 +377,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.688402612Z", - "start_time": "2026-07-26T11:57:19.637442252Z" + "end_time": "2026-07-27T09:47:58.134645045Z", + "start_time": "2026-07-27T09:47:58.116726506Z" } }, "cell_type": "code", @@ -420,8 +411,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.741523163Z", - "start_time": "2026-07-26T11:57:19.690408939Z" + "end_time": "2026-07-27T09:47:58.193234091Z", + "start_time": "2026-07-27T09:47:58.136182636Z" } }, "cell_type": "code", @@ -458,8 +449,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.791841369Z", - "start_time": "2026-07-26T11:57:19.743006544Z" + "end_time": "2026-07-27T09:47:58.260217572Z", + "start_time": "2026-07-27T09:47:58.212670597Z" } }, "cell_type": "code", @@ -480,8 +471,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.844930037Z", - "start_time": "2026-07-26T11:57:19.792995780Z" + "end_time": "2026-07-27T09:47:58.309313551Z", + "start_time": "2026-07-27T09:47:58.261240696Z" } }, "cell_type": "code", @@ -525,8 +516,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T11:57:19.991066598Z", - "start_time": "2026-07-26T11:57:19.846926644Z" + "end_time": "2026-07-27T09:47:58.451959046Z", + "start_time": "2026-07-27T09:47:58.310247636Z" } }, "cell_type": "code", @@ -547,7 +538,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T19:57:19.937795\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:47:58.388372\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -561,8 +552,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T12:01:16.948464226Z", - "start_time": "2026-07-26T12:01:16.911551343Z" + "end_time": "2026-07-27T09:47:58.499427639Z", + "start_time": "2026-07-27T09:47:58.453550913Z" } }, "cell_type": "code", @@ -577,8 +568,8 @@ { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T12:06:54.735214956Z", - "start_time": "2026-07-26T12:06:54.681936998Z" + "end_time": "2026-07-27T09:47:58.560984562Z", + "start_time": "2026-07-27T09:47:58.500638893Z" } }, "cell_type": "code", @@ -594,13 +585,13 @@ ], "id": "f15edcb302b1f95e", "outputs": [], - "execution_count": 21 + "execution_count": 20 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:27:09.970459950Z", - "start_time": "2026-07-26T13:27:09.946476638Z" + "end_time": "2026-07-27T09:47:58.627497913Z", + "start_time": "2026-07-27T09:47:58.578335925Z" } }, "cell_type": "code", @@ -621,13 +612,13 @@ ], "id": "7dfcdd54114d8a06", "outputs": [], - "execution_count": 30 + "execution_count": 21 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:27:10.364513632Z", - "start_time": "2026-07-26T13:27:10.338997724Z" + "end_time": "2026-07-27T09:47:58.687688493Z", + "start_time": "2026-07-27T09:47:58.628509966Z" } }, "cell_type": "code", @@ -667,13 +658,13 @@ ], "id": "8313f60149168f3d", "outputs": [], - "execution_count": 31 + "execution_count": 22 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:27:10.685183804Z", - "start_time": "2026-07-26T13:27:10.651321744Z" + "end_time": "2026-07-27T09:47:58.737335139Z", + "start_time": "2026-07-27T09:47:58.689268726Z" } }, "cell_type": "code", @@ -687,13 +678,13 @@ ], "id": "46de8ece74b27875", "outputs": [], - "execution_count": 32 + "execution_count": 23 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:27:11.044810335Z", - "start_time": "2026-07-26T13:27:10.930620323Z" + "end_time": "2026-07-27T09:47:58.871856579Z", + "start_time": "2026-07-27T09:47:58.738750970Z" } }, "cell_type": "code", @@ -709,7 +700,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:27:10.997118\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:47:58.828265\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -718,13 +709,13 @@ } } ], - "execution_count": 33 + "execution_count": 24 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:27:12.775142785Z", - "start_time": "2026-07-26T13:27:12.685045907Z" + "end_time": "2026-07-27T09:47:58.961912712Z", + "start_time": "2026-07-27T09:47:58.872770917Z" } }, "cell_type": "code", @@ -737,14 +728,6 @@ ], "id": "7eb0e18519f4b55f", "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "tensor([0, 3])\n", - "tensor([0, 1, 2, 3])\n" - ] - }, { "data": { "text/plain": [ @@ -754,18 +737,18 @@ " [-1.0000, 0.7000, 0.1500, 0.3000, 0.6200, 0.9100]]])" ] }, - "execution_count": 34, + "execution_count": 25, "metadata": {}, "output_type": "execute_result" } ], - "execution_count": 34 + "execution_count": 25 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:31:29.890705615Z", - "start_time": "2026-07-26T13:31:29.770882766Z" + "end_time": "2026-07-27T09:47:59.200171491Z", + "start_time": "2026-07-27T09:47:58.962631809Z" } }, "cell_type": "code", @@ -784,7 +767,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:31:29.841516\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:47:59.051198\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -793,13 +776,13 @@ } } ], - "execution_count": 35 + "execution_count": 26 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:45:58.790473045Z", - "start_time": "2026-07-26T13:45:58.732856164Z" + "end_time": "2026-07-27T09:47:59.227792428Z", + "start_time": "2026-07-27T09:47:59.205845108Z" } }, "cell_type": "code", @@ -816,18 +799,18 @@ "(640, 640)" ] }, - "execution_count": 52, + "execution_count": 27, "metadata": {}, "output_type": "execute_result" } ], - "execution_count": 52 + "execution_count": 27 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:46:00.668782428Z", - "start_time": "2026-07-26T13:46:00.505299012Z" + "end_time": "2026-07-27T09:47:59.462849505Z", + "start_time": "2026-07-27T09:47:59.228706763Z" } }, "cell_type": "code", @@ -849,7 +832,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:46:00.601996\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:47:59.341993\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -858,13 +841,13 @@ } } ], - "execution_count": 53 + "execution_count": 28 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:46:02.600951835Z", - "start_time": "2026-07-26T13:46:02.468950170Z" + "end_time": "2026-07-27T09:47:59.984065826Z", + "start_time": "2026-07-27T09:47:59.567538110Z" } }, "cell_type": "code", @@ -876,7 +859,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:46:02.539171\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:47:59.675822\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -885,13 +868,13 @@ } } ], - "execution_count": 54 + "execution_count": 29 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:46:03.940576114Z", - "start_time": "2026-07-26T13:46:03.820486087Z" + "end_time": "2026-07-27T09:48:00.334968551Z", + "start_time": "2026-07-27T09:48:00.005932344Z" } }, "cell_type": "code", @@ -903,7 +886,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T21:46:03.881177\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:48:00.079197\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -912,13 +895,13 @@ } } ], - "execution_count": 55 + "execution_count": 30 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:48:39.530086370Z", - "start_time": "2026-07-26T13:48:39.490640006Z" + "end_time": "2026-07-27T09:48:00.402774849Z", + "start_time": "2026-07-27T09:48:00.390428177Z" } }, "cell_type": "code", @@ -944,13 +927,13 @@ ], "id": "70e7bc1342f90cea", "outputs": [], - "execution_count": 56 + "execution_count": 31 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T13:49:42.340692779Z", - "start_time": "2026-07-26T13:49:42.292033898Z" + "end_time": "2026-07-27T09:48:00.463250908Z", + "start_time": "2026-07-27T09:48:00.403783355Z" } }, "cell_type": "code", @@ -975,13 +958,13 @@ ], "id": "1c80340fca8af59a", "outputs": [], - "execution_count": 57 + "execution_count": 32 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T15:34:46.746160081Z", - "start_time": "2026-07-26T15:34:01.758563722Z" + "end_time": "2026-07-27T09:48:03.208374192Z", + "start_time": "2026-07-27T09:48:00.491960801Z" } }, "cell_type": "code", @@ -997,7 +980,6 @@ "name": "stdout", "output_type": "stream", "text": [ - "Downloading ../data/banana-detection.zip from http://d2l-data.s3-accelerate.amazonaws.com/banana-detection.zip...\n", "read 1000 training examples\n", "read 100 validation examples\n" ] @@ -1008,18 +990,18 @@ "(torch.Size([32, 3, 256, 256]), torch.Size([32, 1, 5]))" ] }, - "execution_count": 58, + "execution_count": 33, "metadata": {}, "output_type": "execute_result" } ], - "execution_count": 58 + "execution_count": 33 }, { "metadata": { "ExecuteTime": { - "end_time": "2026-07-26T15:36:52.903923701Z", - "start_time": "2026-07-26T15:36:52.509530881Z" + "end_time": "2026-07-27T09:48:03.533990096Z", + "start_time": "2026-07-27T09:48:03.269951948Z" } }, "cell_type": "code", @@ -1036,7 +1018,7 @@ "text/plain": [ "
" ], - "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-26T23:36:52.688576\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:48:03.414413\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" }, "metadata": {}, "output_type": "display_data", @@ -1045,15 +1027,755 @@ } } ], - "execution_count": 60 + "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 }, { "metadata": {}, + "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" + } + }, "cell_type": "code", + "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", "outputs": [], - "execution_count": null, + "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": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:50:07.189022\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "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=)" + ] + }, + "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=)\n", + "tensor([0.0000, 0.9973, 0.4608, 0.5666, 0.6662, 0.7837],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.9970, 0.5375, 0.0591, 0.7508, 0.2796],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.9902, 0.7037, 0.3652, 0.9121, 0.5747],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.4860, 0.5226, 0.0069, 0.7186, 0.2129],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.4443, 0.4504, 0.6233, 0.6205, 0.8134],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.2552, 0.5831, 0.0923, 0.8302, 0.3042],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.2356, 0.5915, -0.0019, 0.7679, 0.2016],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.2331, 0.4954, 0.6350, 0.7170, 0.8248],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.1502, 0.7096, 0.2895, 0.8824, 0.5102],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.1475, 0.1230, 0.7664, 0.3295, 0.9989],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.1150, 0.5567, -0.0957, 0.7636, 0.1090],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.1137, 0.1259, 0.6939, 0.3409, 0.9125],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.1115, 0.4028, 0.6006, 0.5825, 0.7848],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.1093, 0.4543, 0.5068, 0.7193, 0.7340],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.1016, 0.5291, 0.1110, 0.7004, 0.3227],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0900, 0.4455, 0.4652, 0.6391, 0.6875],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0780, 0.4589, 0.6589, 0.6681, 0.8913],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0652, 0.5598, 0.1803, 0.7430, 0.3703],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0605, 0.6292, 0.3817, 0.8251, 0.5817],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0514, 0.4952, 0.7167, 0.6964, 0.9586],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0469, 0.6734, 0.4222, 0.8798, 0.6158],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0416, 0.4532, 0.0716, 0.6357, 0.2691],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.0391, 0.4798, -0.0191, 0.6770, 0.1747],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0389, 0.0044, 0.7313, 0.1998, 0.9244],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0389, 0.5262, 0.4644, 0.6896, 0.6895],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0374, 0.6148, 0.1649, 0.7902, 0.3841],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0359, 0.0650, 0.8109, 0.2880, 1.0482],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0359, 0.0698, 0.6747, 0.2262, 0.9140],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.0319, 0.6276, -0.1184, 0.7822, 0.1588],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.0319, -0.1519, 0.2342, 1.2012, 0.8651],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0285, 0.1491, 0.6391, 0.3408, 0.8407],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0282, 0.6711, 0.2857, 0.8190, 0.5322],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.0281, 0.1795, -0.1376, 0.7880, 1.1077],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0271, 0.4184, 0.6820, 0.6104, 0.8625],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0266, 0.6531, 0.2491, 0.8567, 0.4469],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0215, 0.3868, 0.5555, 0.6020, 0.7265],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.0164, 0.5367, -0.1162, 0.6877, 0.1575],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0160, 0.3315, 0.6025, 0.5351, 0.7996],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0156, 0.2120, 0.6672, 0.3861, 0.8899],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.0154, -0.3605, -0.1560, 0.6862, 0.3646],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0149, 0.6149, 0.0149, 0.8310, 0.2521],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0141, 0.1019, 0.6182, 0.2584, 0.8786],\n", + " grad_fn=)\n", + "tensor([ 0.0000, 0.0131, -0.1225, -0.3201, 0.2735, 0.4565],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0130, 0.4430, 0.5082, 1.2768, 1.2750],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0125, 0.1837, 0.7573, 0.3872, 0.9679],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0123, 0.4169, 0.4787, 0.5627, 0.7412],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0122, 0.7223, 0.2193, 0.9493, 0.4536],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0121, 0.4850, 0.4219, 0.7411, 0.6612],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0112, 0.6911, 0.4813, 0.9275, 0.7128],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0105, 0.5652, 0.7101, 0.7331, 0.9001],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0102, 0.0054, 0.8066, 0.2270, 0.9863],\n", + " grad_fn=)\n", + "tensor([0.0000, 0.0101, 0.6626, 0.0780, 0.8909, 0.3222],\n", + " grad_fn=)\n" + ] + }, + { + "data": { + "text/plain": [ + "
" + ], + "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-27T17:50:08.344236\n image/svg+xml\n \n \n Matplotlib v3.7.2, https://matplotlib.org/\n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n \n\n" + }, + "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", "source": "", - "id": "eb1131ed49817939" + "id": "4ec2f61624100d07", + "outputs": [], + "execution_count": 54 } ], "metadata": {