{ "cells": [ { "metadata": {}, "cell_type": "markdown", "source": "### 前面的13.1 13.2 因为不可抗力事件消失了(保存的时候乱码了)", "id": "b23af7657da5adfb" }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:55.425519349Z", "start_time": "2026-07-27T09:47:53.873571522Z" } }, "cell_type": "code", "source": [ "\n", "import torch\n", "from d2l import torch as d2l" ], "id": "b39ef644f6c80c7c", "outputs": [], "execution_count": 1 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:55.671442114Z", "start_time": "2026-07-27T09:47:55.428215388Z" } }, "cell_type": "code", "source": [ "d2l.set_figsize()\n", "img = d2l.plt.imread('../data/catdog.jpg')\n", "d2l.plt.imshow(img);" ], "id": "4284e525ebcc7e87", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 2 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:55.747558006Z", "start_time": "2026-07-27T09:47:55.673122195Z" } }, "cell_type": "code", "source": [ "def box_corner_to_center(boxes):\n", " \"\"\"从(左上,右下)转换到(中间,宽度,高度)\"\"\"\n", " x1, y1, x2, y2 = boxes[:, 0], boxes[:, 1], boxes[:, 2], boxes[:, 3]\n", " cx = (x1 + x2) / 2\n", " cy = (y1 + y2) / 2\n", " w = x2 - x1\n", " h = y2 - y1\n", " boxes = torch.stack((cx, cy, w, h), axis=-1)\n", " return boxes\n", "def box_center_to_corner(boxes):\n", " \"\"\"从(中间,宽度,高度)转换到(左上,右下)\"\"\"\n", " cx, cy, w, h = boxes[:, 0], boxes[:, 1], boxes[:, 2], boxes[:, 3]\n", " x1 = cx - 0.5 * w\n", " y1 = cy - 0.5 * h\n", " x2 = cx + 0.5 * w\n", " y2 = cy + 0.5 * h\n", " boxes = torch.stack((x1, y1, x2, y2), axis=-1)\n", " return boxes" ], "id": "ee53061c33ad70d7", "outputs": [], "execution_count": 3 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:55.853579240Z", "start_time": "2026-07-27T09:47:55.749428064Z" } }, "cell_type": "code", "source": "dog_bbox, cat_bbox = [60.0, 45.0, 378.0, 516.0], [400.0, 112.0, 655.0, 493.0]", "id": "2e45aa7592a4d7b1", "outputs": [], "execution_count": 4 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:55.959910400Z", "start_time": "2026-07-27T09:47:55.859086964Z" } }, "cell_type": "code", "source": [ "boxes = torch.tensor((dog_bbox, cat_bbox))\n", "box_center_to_corner(box_corner_to_center(boxes)) == boxes" ], "id": "22f2ce93eb46eca2", "outputs": [ { "data": { "text/plain": [ "tensor([[True, True, True, True],\n", " [True, True, True, True]])" ] }, "execution_count": 5, "metadata": {}, "output_type": "execute_result" } ], "execution_count": 5 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:57.527570250Z", "start_time": "2026-07-27T09:47:57.492961332Z" } }, "cell_type": "code", "source": [ "def bbox_to_rect(bbox, color):\n", " # 将边界框(左上x,左上y,右下x,右下y)格式转换成matplotlib格式:\n", " # ((左上x,左上y),宽,高)\n", " return d2l.plt.Rectangle(\n", " xy=(bbox[0], bbox[1]), width=bbox[2]-bbox[0], height=bbox[3]-bbox[1],\n", " fill=False, edgecolor=color, linewidth=2)" ], "id": "cc21e1e13103291", "outputs": [], "execution_count": 6 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:57.635071582Z", "start_time": "2026-07-27T09:47:57.528293612Z" } }, "cell_type": "code", "source": [ "fig = d2l.plt.imshow(img)\n", "fig.axes.add_patch(bbox_to_rect(dog_bbox, 'blue'))\n", "fig.axes.add_patch(bbox_to_rect(cat_bbox, 'red'))" ], "id": "a4f3af406f8b5b1c", "outputs": [ { "data": { "text/plain": [ "" ] }, "execution_count": 7, "metadata": {}, "output_type": "execute_result" }, { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 7 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:57.702782485Z", "start_time": "2026-07-27T09:47:57.650116830Z" } }, "cell_type": "code", "source": [ "#@save\n", "def multibox_prior(data, sizes, ratios):\n", " \"\"\"生成以每个像素为中心具有不同形状的锚框\"\"\"\n", " in_height, in_width = data.shape[-2:]\n", " device, num_sizes, num_ratios = data.device, len(sizes), len(ratios)\n", " boxes_per_pixel = (num_sizes + num_ratios - 1)\n", " size_tensor = torch.tensor(sizes, device=device)\n", " ratio_tensor = torch.tensor(ratios, device=device)\n", "\n", " # 为了将锚点移动到像素的中心,需要设置偏移量。\n", " # 因为一个像素的高为1且宽为1,我们选择偏移我们的中心0.5\n", " offset_h, offset_w = 0.5, 0.5\n", " steps_h = 1.0 / in_height # 在y轴上缩放步长\n", " steps_w = 1.0 / in_width # 在x轴上缩放步长\n", "\n", " # 生成锚框的所有中心点\n", " center_h = (torch.arange(in_height, device=device) + offset_h) * steps_h\n", " center_w = (torch.arange(in_width, device=device) + offset_w) * steps_w\n", " shift_y, shift_x = torch.meshgrid(center_h, center_w, indexing='ij')\n", " shift_y, shift_x = shift_y.reshape(-1), shift_x.reshape(-1)\n", "\n", " # 生成 “boxes_per_pixel” 个高和宽,\n", " # 之后用于创建锚框的四角坐标(xmin,xmax,ymin,ymax)\n", " w = torch.cat((size_tensor * torch.sqrt(ratio_tensor[0]),\n", " sizes[0] * torch.sqrt(ratio_tensor[1:])))\\\n", " * in_height / in_width # 处理矩形输入\n", "\n", " h = torch.cat((size_tensor / torch.sqrt(ratio_tensor[0]),\n", " sizes[0] / torch.sqrt(ratio_tensor[1:])))\n", "\n", " # 除以2来获得半高和半宽\n", " anchor_manipulations = torch.stack((-w, -h, w, h)).T.repeat(\n", " in_height * in_width, 1) / 2\n", "\n", " # 每个中心点都将有 “boxes_per_pixel” 个锚框,\n", " # 所以生成含所有锚框中心的网格,重复了 “boxes_per_pixel” 次\n", " out_grid = torch.stack([shift_x, shift_y, shift_x, shift_y],\n", " dim=1).repeat_interleave(boxes_per_pixel, dim=0)\n", " output = out_grid + anchor_manipulations\n", " return output.unsqueeze(0)" ], "id": "874f386629d410fb", "outputs": [], "execution_count": 8 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:57.754193331Z", "start_time": "2026-07-27T09:47:57.704574735Z" } }, "cell_type": "code", "source": [ "img = d2l.plt.imread('../data/catdog.jpg')\n", "h, w = img.shape[:2]" ], "id": "df3d1ba4aeb6c2f7", "outputs": [], "execution_count": 9 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:57.820044630Z", "start_time": "2026-07-27T09:47:57.755572144Z" } }, "cell_type": "code", "source": [ "print(h, w)\n", "X = torch.rand(size=(1, 3, h, w))\n", "Y = multibox_prior(X, sizes=[0.75, 0.5, 0.25], ratios=[1, 2, 0.5])\n", "Y.shape" ], "id": "e8737b3210446831", "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "561 728\n" ] }, { "data": { "text/plain": [ "torch.Size([1, 2042040, 4])" ] }, "execution_count": 10, "metadata": {}, "output_type": "execute_result" } ], "execution_count": 10 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:57.877041771Z", "start_time": "2026-07-27T09:47:57.821962159Z" } }, "cell_type": "code", "source": [ "def show_bboxes(axes, bboxes, labels=None, colors=None):\n", " \"\"\"显示所有边界框\"\"\"\n", " def _make_list(obj, default_values=None):\n", " if obj is None:\n", " obj = default_values\n", " elif not isinstance(obj, (list, tuple)):\n", " obj = [obj]\n", " return obj\n", "\n", " labels = _make_list(labels)\n", " colors = _make_list(colors, ['b', 'g', 'r', 'm', 'c'])\n", " for i, bbox in enumerate(bboxes):\n", " color = colors[i % len(colors)]\n", " rect = d2l.bbox_to_rect(bbox.detach().numpy(), color)\n", " axes.add_patch(rect)\n", " if labels and len(labels) > i:\n", " text_color = 'k' if color == 'w' else 'w'\n", " axes.text(rect.xy[0], rect.xy[1], labels[i],\n", " va='center', ha='center', fontsize=9, color=text_color,\n", " bbox=dict(facecolor=color, lw=0))" ], "id": "e7070cbdb1d68616", "outputs": [], "execution_count": 11 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:57.924927279Z", "start_time": "2026-07-27T09:47:57.877778001Z" } }, "cell_type": "code", "source": "boxes = Y.reshape(h, w, 5, 4)", "id": "4cc20a550931aef7", "outputs": [], "execution_count": 12 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.080082507Z", "start_time": "2026-07-27T09:47:57.926102351Z" } }, "cell_type": "code", "source": [ "bbox_scale = torch.tensor((w, h, w, h))\n", "fig = d2l.plt.imshow(img)\n", "show_bboxes(fig.axes, boxes[250, 250, :, :] * bbox_scale,\n", " ['s=0.75, r=1', 's=0.5, r=1', 's=0.25, r=1', 's=0.75, r=2',\n", " 's=0.75, r=0.5'])" ], "id": "2daf2903aa5b1516", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 13 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.134645045Z", "start_time": "2026-07-27T09:47:58.116726506Z" } }, "cell_type": "code", "source": [ "def box_iou(boxes1, boxes2):\n", " \"\"\"计算两个锚框或边界框列表中成对的交并比\"\"\"\n", " box_area = lambda boxes: ((boxes[:, 2] - boxes[:, 0]) *\n", " (boxes[:, 3] - boxes[:, 1]))\n", " # boxes1,boxes2,areas1,areas2的形状:\n", " # boxes1:(boxes1的数量,4),\n", " # boxes2:(boxes2的数量,4),\n", " # areas1:(boxes1的数量,),\n", " # areas2:(boxes2的数量,)\n", " areas1 = box_area(boxes1)\n", " areas2 = box_area(boxes2)\n", " # inter_upperlefts,inter_lowerrights,inters的形状:\n", " # (boxes1的数量,boxes2的数量,2)\n", " inter_upperlefts = torch.max(boxes1[:, None, :2], boxes2[:, :2])\n", " inter_lowerrights = torch.min(boxes1[:, None, 2:], boxes2[:, 2:])\n", " inters = (inter_lowerrights - inter_upperlefts).clamp(min=0)\n", " # inter_areasandunion_areas的形状:(boxes1的数量,boxes2的数量)\n", " inter_areas = inters[:, :, 0] * inters[:, :, 1]\n", " union_areas = areas1[:, None] + areas2 - inter_areas\n", " return inter_areas / union_areas" ], "id": "ffe3178858e5a9b1", "outputs": [], "execution_count": 14 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.193234091Z", "start_time": "2026-07-27T09:47:58.136182636Z" } }, "cell_type": "code", "source": [ "#@save\n", "def assign_anchor_to_bbox(ground_truth, anchors, device, iou_threshold=0.5):\n", " \"\"\"将最接近的真实边界框分配给锚框\"\"\"\n", " num_anchors, num_gt_boxes = anchors.shape[0], ground_truth.shape[0]\n", " # 位于第i行和第j列的元素x_ij是锚框i和真实边界框j的IoU\n", " jaccard = box_iou(anchors, ground_truth)\n", " # 对于每个锚框,分配的真实边界框的张量\n", " anchors_bbox_map = torch.full((num_anchors,), -1, dtype=torch.long,\n", " device=device)\n", " # 根据阈值,决定是否分配真实边界框\n", " max_ious, indices = torch.max(jaccard, dim=1)\n", " anc_i = torch.nonzero(max_ious >= iou_threshold).reshape(-1)\n", " box_j = indices[max_ious >= iou_threshold]\n", " anchors_bbox_map[anc_i] = box_j\n", " col_discard = torch.full((num_anchors,), -1)\n", " row_discard = torch.full((num_gt_boxes,), -1)\n", " for _ in range(num_gt_boxes):\n", " max_idx = torch.argmax(jaccard)\n", " box_idx = (max_idx % num_gt_boxes).long()\n", " anc_idx = (max_idx / num_gt_boxes).long()\n", " anchors_bbox_map[anc_idx] = box_idx\n", " jaccard[:, box_idx] = col_discard\n", " jaccard[anc_idx, :] = row_discard\n", " return anchors_bbox_map\n" ], "id": "2ebf82c1768e24c6", "outputs": [], "execution_count": 15 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.260217572Z", "start_time": "2026-07-27T09:47:58.212670597Z" } }, "cell_type": "code", "source": [ "def offset_boxes(anchors, assigned_bb, eps=1e-6):\n", " \"\"\"对锚框偏移量的转换\"\"\"\n", " c_anc = d2l.box_corner_to_center(anchors)\n", " c_assigned_bb = d2l.box_corner_to_center(assigned_bb)\n", " offset_xy = 10 * (c_assigned_bb[:, :2] - c_anc[:, :2]) / c_anc[:, 2:]\n", " offset_wh = 5 * torch.log(eps + c_assigned_bb[:, 2:] / c_anc[:, 2:])\n", " offset = torch.cat([offset_xy, offset_wh], axis=1)\n", " return offset" ], "id": "b83220e939515b39", "outputs": [], "execution_count": 16 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.309313551Z", "start_time": "2026-07-27T09:47:58.261240696Z" } }, "cell_type": "code", "source": [ "def multibox_target(anchors, labels):\n", " \"\"\"使用真实边界框标记锚框\"\"\"\n", " batch_size, anchors = labels.shape[0], anchors.squeeze(0) # 问题:这里赋值给 anchors 会覆盖,但原代码就是这样\n", " batch_offset, batch_mask, batch_class_labels = [], [], []\n", " device, num_anchors = anchors.device, anchors.shape[0]\n", " for i in range(batch_size):\n", " label = labels[i, :, :]\n", " anchors_bbox_map = assign_anchor_to_bbox(\n", " label[:, 1:], anchors, device)\n", " bbox_mask = ((anchors_bbox_map >= 0).float().unsqueeze(-1)).repeat(\n", " 1, 4)\n", " # 将类标签和分配的边界框坐标初始化为零\n", " class_labels = torch.zeros(num_anchors, dtype=torch.long,\n", " device=device)\n", " assigned_bb = torch.zeros((num_anchors, 4), dtype=torch.float32,\n", " device=device)\n", " # 使用真实边界框来标记锚框的类别。\n", " # 如果一个锚框没有被分配,标记其为背景(值为零)\n", " indices_true = torch.nonzero(anchors_bbox_map >= 0)\n", " bb_idx = anchors_bbox_map[indices_true]\n", " class_labels[indices_true] = label[bb_idx, 0].long() + 1\n", " assigned_bb[indices_true] = label[bb_idx, 1:]\n", " # 偏移量转换\n", " offset = offset_boxes(anchors, assigned_bb) * bbox_mask\n", " batch_offset.append(offset.reshape(-1))\n", " batch_mask.append(bbox_mask.reshape(-1))\n", " batch_class_labels.append(class_labels)\n", " bbox_offset = torch.stack(batch_offset)\n", " bbox_mask = torch.stack(batch_mask)\n", " class_labels = torch.stack(batch_class_labels)\n", " return (bbox_offset, bbox_mask, class_labels)" ], "id": "57e69b1186f23243", "outputs": [], "execution_count": 17 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.451959046Z", "start_time": "2026-07-27T09:47:58.310247636Z" } }, "cell_type": "code", "source": [ "ground_truth = torch.tensor([[0, 0.1, 0.08, 0.52, 0.92],\n", "[1, 0.55, 0.2, 0.9, 0.88]])\n", "anchors = torch.tensor([[0, 0.1, 0.2, 0.3], [0.15, 0.2, 0.4, 0.4],\n", "[0.63, 0.05, 0.88, 0.98], [0.66, 0.45, 0.8, 0.8],\n", "[0.57, 0.3, 0.92, 0.9]])\n", "fig = d2l.plt.imshow(img)\n", "show_bboxes(fig.axes, ground_truth[:, 1:] * bbox_scale, ['dog', 'cat'], 'k')\n", "show_bboxes(fig.axes, anchors * bbox_scale, ['0', '1', '2', '3', '4']);" ], "id": "260673166da2aa71", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 18 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.499427639Z", "start_time": "2026-07-27T09:47:58.453550913Z" } }, "cell_type": "code", "source": [ "labels = multibox_target(anchors.unsqueeze(dim=0),\n", " ground_truth.unsqueeze(dim=0))" ], "id": "aafe21dd7b10dde3", "outputs": [], "execution_count": 19 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.560984562Z", "start_time": "2026-07-27T09:47:58.500638893Z" } }, "cell_type": "code", "source": [ "def offset_inverse(anchors, offset_preds):\n", " \"\"\"根据带有预测偏移量的锚框来预测边界框\"\"\"\n", " anc = d2l.box_corner_to_center(anchors)\n", " pred_bbox_xy = (offset_preds[:, :2] * anc[:, 2:] / 10) + anc[:, :2]\n", " pred_bbox_wh = torch.exp(offset_preds[:, 2:] / 5) * anc[:, 2:]\n", " pred_bbox = torch.cat((pred_bbox_xy, pred_bbox_wh), axis=1)\n", " predicted_bbox = d2l.box_center_to_corner(pred_bbox)\n", " return predicted_bbox" ], "id": "f15edcb302b1f95e", "outputs": [], "execution_count": 20 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.627497913Z", "start_time": "2026-07-27T09:47:58.578335925Z" } }, "cell_type": "code", "source": [ "def nms(boxes, scores, iou_threshold):\n", " \"\"\"对预测边界框的置信度进行排序\"\"\"\n", " B = torch.argsort(scores, dim=-1, descending=True)\n", " keep = [] # 保留预测边界框的指标\n", " while B.numel() > 0:\n", " i = B[0]\n", " keep.append(i)\n", " if B.numel() == 1: break\n", " iou = box_iou(boxes[i, :].reshape(-1, 4),\n", " boxes[B[1:], :].reshape(-1, 4)).reshape(-1)\n", " inds = torch.nonzero(iou <= iou_threshold).reshape(-1)\n", " B = B[inds + 1]\n", " return torch.tensor(keep, device=boxes.device)" ], "id": "7dfcdd54114d8a06", "outputs": [], "execution_count": 21 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.687688493Z", "start_time": "2026-07-27T09:47:58.628509966Z" } }, "cell_type": "code", "source": [ "#@save\n", "def multibox_detection(cls_probs, offset_preds, anchors, nms_threshold=0.5,\n", " pos_threshold=0.009999999):\n", " \"\"\"使用非极大值抑制来预测边界框\"\"\"\n", " device, batch_size = cls_probs.device, cls_probs.shape[0]\n", " anchors = anchors.squeeze(0)\n", " num_classes, num_anchors = cls_probs.shape[1], cls_probs.shape[2]\n", " out = []\n", " for i in range(batch_size):\n", " cls_prob, offset_pred = cls_probs[i], offset_preds[i].reshape(-1, 4)\n", " conf, class_id = torch.max(cls_prob[1:], 0)\n", " predicted_bb = offset_inverse(anchors, offset_pred)\n", " keep = nms(predicted_bb, conf, nms_threshold)\n", " # 找到所有的non_keep索引,并将类设置为背景\n", " all_idx = torch.arange(num_anchors, dtype=torch.long, device=device)\n", " combined = torch.cat((keep, all_idx))\n", " uniques, counts = combined.unique(return_counts=True)\n", " non_keep = uniques[counts == 1]\n", " all_id_sorted = torch.cat((keep, non_keep))\n", " class_id[non_keep] = -1\n", " class_id = class_id[all_id_sorted]\n", " conf, predicted_bb = conf[all_id_sorted], predicted_bb[all_id_sorted]\n", " # pos_threshold是一个用于非背景预测的阈值\n", " below_min_idx = (conf < pos_threshold)\n", " class_id[below_min_idx] = -1\n", " conf[below_min_idx] = 1 - conf[below_min_idx]\n", " pred_info = torch.cat((class_id.unsqueeze(1),\n", " conf.unsqueeze(1),\n", " predicted_bb), dim=1)\n", "\n", " out.append(pred_info)\n", " return torch.stack(out)" ], "id": "8313f60149168f3d", "outputs": [], "execution_count": 22 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.737335139Z", "start_time": "2026-07-27T09:47:58.689268726Z" } }, "cell_type": "code", "source": [ "anchors = torch.tensor([[0.1, 0.08, 0.52, 0.92], [0.08, 0.2, 0.56, 0.95],\n", "[0.15, 0.3, 0.62, 0.91], [0.55, 0.2, 0.9, 0.88]])\n", "offset_preds = torch.tensor([0] * anchors.numel())\n", "cls_probs = torch.tensor([[0] * 4, # 背景的预测概率\n", "[0.9, 0.8, 0.7, 0.1], # 狗的预测概率\n", "[0.1, 0.2, 0.3, 0.9]]) # 猫的预测概率" ], "id": "46de8ece74b27875", "outputs": [], "execution_count": 23 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.871856579Z", "start_time": "2026-07-27T09:47:58.738750970Z" } }, "cell_type": "code", "source": [ "fig = d2l.plt.imshow(img)\n", "show_bboxes(fig.axes, anchors * bbox_scale,\n", "['dog=0.9', 'dog=0.8', 'dog=0.7', 'cat=0.9'])" ], "id": "ed43634b61e3ec27", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 24 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:58.961912712Z", "start_time": "2026-07-27T09:47:58.872770917Z" } }, "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]]])" ] }, "execution_count": 25, "metadata": {}, "output_type": "execute_result" } ], "execution_count": 25 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:59.200171491Z", "start_time": "2026-07-27T09:47:58.962631809Z" } }, "cell_type": "code", "source": [ "fig = d2l.plt.imshow(img)\n", "for i in output[0].detach().numpy():\n", " if i[0] == -1:\n", " continue\n", " label = ('dog=', 'cat=')[int(i[0])] + str(i[1])\n", " show_bboxes(fig.axes, [torch.tensor(i[2:]) * bbox_scale], label)" ], "id": "fec1c1f8349a44fe", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 26 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:59.227792428Z", "start_time": "2026-07-27T09:47:59.205845108Z" } }, "cell_type": "code", "source": [ "img = d2l.plt.imread('../Pictures/1.jpg')\n", "h, w = img.shape[:2]\n", "h, w" ], "id": "ea166e0233574bf3", "outputs": [ { "data": { "text/plain": [ "(640, 640)" ] }, "execution_count": 27, "metadata": {}, "output_type": "execute_result" } ], "execution_count": 27 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:59.462849505Z", "start_time": "2026-07-27T09:47:59.228706763Z" } }, "cell_type": "code", "source": [ "def display_anchors(fmap_w, fmap_h, s):\n", " d2l.set_figsize()\n", " # 前两个维度上的值不影响输出\n", " fmap = torch.zeros((1, 10, fmap_h, fmap_w))\n", " anchors = d2l.multibox_prior(fmap, sizes=s, ratios=[1, 2, 0.5,0.3])\n", " bbox_scale = torch.tensor((w, h, w, h))\n", " d2l.show_bboxes(d2l.plt.imshow(img).axes,\n", " anchors[0] * bbox_scale)\n", "display_anchors(fmap_w=3, fmap_h=3, s=[0.15,0.2,0.3])" ], "id": "41cab7a08738ed58", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 28 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:47:59.984065826Z", "start_time": "2026-07-27T09:47:59.567538110Z" } }, "cell_type": "code", "source": "display_anchors(fmap_w=2, fmap_h=2, s=[0.4])", "id": "6620db5aa275bf52", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 29 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:48:00.334968551Z", "start_time": "2026-07-27T09:48:00.005932344Z" } }, "cell_type": "code", "source": "display_anchors(fmap_w=1, fmap_h=1, s=[0.8])", "id": "8a0c4555b2752c99", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "execution_count": 30 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:48:00.402774849Z", "start_time": "2026-07-27T09:48:00.390428177Z" } }, "cell_type": "code", "source": [ "import torchvision,os\n", "import pandas as pd\n", "def read_data_bananas(is_train=True):\n", " \"\"\"读取香蕉检测数据集中的图像和标签\"\"\"\n", " data_dir = d2l.download_extract('banana-detection')\n", " csv_fname = os.path.join(data_dir, 'bananas_train' if is_train\n", " else 'bananas_val', 'label.csv')\n", " csv_data = pd.read_csv(csv_fname)\n", " csv_data = csv_data.set_index('img_name')\n", " images, targets = [], []\n", " for img_name, target in csv_data.iterrows():\n", " images.append(torchvision.io.read_image(\n", " os.path.join(data_dir, 'bananas_train' if is_train else\n", " 'bananas_val', 'images', f'{img_name}')))\n", " # 这里的target包含(类别,左上角x,左上角y,右下角x,右下角y),\n", " # 其中所有图像都具有相同的香蕉类(索引为0)\n", " targets.append(list(target))\n", " return images, torch.tensor(targets).unsqueeze(1) / 256" ], "id": "70e7bc1342f90cea", "outputs": [], "execution_count": 31 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:48:00.463250908Z", "start_time": "2026-07-27T09:48:00.403783355Z" } }, "cell_type": "code", "source": [ "class BananasDataset(torch.utils.data.Dataset):\n", " \"\"\"一个用于加载香蕉检测数据集的自定义数据集\"\"\"\n", " def __init__(self, is_train):\n", " self.features, self.labels = read_data_bananas(is_train)\n", " print('read ' + str(len(self.features)) + (f' training examples' if\n", " is_train else f' validation examples'))\n", " def __getitem__(self, idx):\n", " return (self.features[idx].float(), self.labels[idx])\n", " def __len__(self):\n", " return len(self.features)\n", "def load_data_bananas(batch_size):\n", " \"\"\"加载香蕉检测数据集\"\"\"\n", " train_iter = torch.utils.data.DataLoader(BananasDataset(is_train=True),\n", " batch_size, shuffle=True)\n", " val_iter = torch.utils.data.DataLoader(BananasDataset(is_train=False),\n", " batch_size)\n", " return train_iter, val_iter" ], "id": "1c80340fca8af59a", "outputs": [], "execution_count": 32 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:48:03.208374192Z", "start_time": "2026-07-27T09:48:00.491960801Z" } }, "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]))" ] }, "execution_count": 33, "metadata": {}, "output_type": "execute_result" } ], "execution_count": 33 }, { "metadata": { "ExecuteTime": { "end_time": "2026-07-27T09:48:03.533990096Z", "start_time": "2026-07-27T09:48:03.269951948Z" } }, "cell_type": "code", "source": [ "imgs = (batch[0][0:10].permute(0, 2, 3, 1)) / 255\n", "axes = d2l.show_images(imgs, 2, 5, scale=2)\n", "for ax, label in zip(axes, batch[1][0:10]):\n", " d2l.show_bboxes(ax, [label[0][1:5] * edge_size], colors=['b'])" ], "id": "528cb4f32e152b5", "outputs": [ { "data": { "text/plain": [ "
" ], "image/svg+xml": "\n\n\n \n \n \n \n 2026-07-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", "jetTransient": { "display_id": null } } ], "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": 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": "4ec2f61624100d07", "outputs": [], "execution_count": 54 } ], "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 }