Files
Kaggle/Petals-to-the-Metal /main.ipynb
T

909 lines
218 KiB
Plaintext
Raw Normal View History

2026-07-08 21:58:01 +08:00
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"id": "befdc3e8",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:24:31.746883Z",
"iopub.status.busy": "2026-07-08T12:24:31.746615Z",
"iopub.status.idle": "2026-07-08T12:25:02.131639Z",
"shell.execute_reply": "2026-07-08T12:25:02.130586Z"
},
"papermill": {
"duration": 30.393653,
"end_time": "2026-07-08T12:25:02.135962+00:00",
"exception": false,
"start_time": "2026-07-08T12:24:31.742309+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"WARNING: All log messages before absl::InitializeLog() is called are written to STDERR\n",
"I0000 00:00:1783513501.576520 23 gpu_device.cc:2020] Created device /job:localhost/replica:0/task:0/device:GPU:0 with 13756 MB memory: -> device: 0, name: Tesla T4, pci bus id: 0000:00:04.0, compute capability: 7.5\n",
"I0000 00:00:1783513501.579453 23 gpu_device.cc:2020] Created device /job:localhost/replica:0/task:0/device:GPU:1 with 13756 MB memory: -> device: 1, name: Tesla T4, pci bus id: 0000:00:05.0, compute capability: 7.5\n",
"Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-2.0836544..2.64].\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAakAAAGhCAYAAADbf0s2AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAAPYQAAD2EBqD+naQABAABJREFUeJzs/XmU5dlV34l+ftOdh5iHnLMya55UVaoqSUhCQhIC2WKS2xgwD9w8Y95DekbC3SA/01j2s9XtZRs8t6cFtoEG/JZp2vCgLUQjkNCA5irVnFVZOUbGfOfpN7w/9tlxTkRFZNzIzKpSSbnXuisibtz7G87vnD1893fv42VZlnFTbspNuSk35aZ8HYr/al/ATbkpN+Wm3JSbspfcNFI35abclJtyU75u5aaRuik35abclJvydSs3jdRNuSk35abclK9buWmkbspNuSk35aZ83cpNI3VTbspNuSk35etWbhqpm3JTbspNuSlft3LTSN2Um3JTbspN+bqVm0bqptyUm3JTbsrXrdw0UjflptyUm3JTvm7lVTNS/+Jf/AtOnDhBoVDg0Ucf5XOf+9yrdSk35abclJtyU75O5VUxUr/xG7/Bhz70IX7+53+eL37xi9x///28+93vZnl5+dW4nJtyU27KTbkpX6fivRoNZh999FEefvhh/vk//+cApGnK0aNH+cAHPsDP/uzP7vv9NE25dOkS1WoVz/Ne7su9KTflptyUm3KDJcsyWq0Whw4dwvf3jpfCV/CaABgOh3zhC1/gwx/+8NZ7vu/zzne+k09/+tO7fmcwGDAYDLb+vnjxInfdddfLfq035abclJtyU15eOX/+PEeOHNnz/6+4kVpdXSVJEubn57e9Pz8/z1NPPbXrdz760Y/ykY98ZLwT5IEKcASYAI4hdxkA0+b/eSjUoFiHIIQkho1NYB3YAFpADCRA2/z9OaADXC3u9IAZYBY4AfTM978CDM1nQizImjOvnvluFZg0r7K55rz56ZufQ2ANeA64bL4zA9wJpTdB7ihsfgFYAp4Hzph7ejWkDtyC3EOEjGkXuf4hMAKayDjvlAJQgQe/AxaOwenj0G7B5gb8yedh/RIkj70ytwHIcztkfk+B48izWwJWgFXnsz4wh8y3VeT5Nvc5/gwyPpsHuKZpcw0tYICM5zgSAIeRexoi8zoB7gf6wDPIvdbMsRvAxQNc137iI3Pi+7D3rWvvEjIG55zPz7C1hqMTEExBv4vMkRlz7R4wCfki1Eqw+TiMLgKPm3tVnzZEdMMIGEBQAz+AMIbRAOKhOV4PWV99+Rxr5joBiubcATKvq0DJ/L5u/jdtPpcTXZNswugyMldScz1VoAa5KngJDJ41318DziLPRe+zBDxofm8ia6tsrsmTa8nXIVeS6x2tQP9r5ji6/iuAqt0EeJHd9VnOXHubl65NzxznNjNGeo2ZGY+W+fsAUq1Wr/r/V9xIXYt8+MMf5kMf+tDW381mk6NHj+7+YVWEXWSgI2RQy+YVAXkYBpCOIMwgTZEBLpljlJBJ3EUWzKY57n7AqIcswBay0DvIBI+dzyTIJPURBRGbc5nrIjTH8cx7NXNtnrmONjLphuazi4hCnIdBG0bnzbFLiLE7ZMZhyZz3lZQu8IK5/gBZWBGiWHzk3p/HKtiSeb+NjMUEbLQgWoHTszBoQmMdTh6D6To8vQFZgozpOtvH+UaIOgbqsKwjc6iKKIs84pAMEEWu59dVlSJjn7G/kWpx8Ofjm5fOl3ElQ55DZL6fIPfwNHJPR5H785zP30iZxc5L99pTRIm7Y+Uh4xnJzziDJEGeC4iiNP8jNH7PCJIImU+huaeyuc8IwhKkDUg7kPqQ5cDPQRY458rsMUmc34tYpytw3tfxrJrPTJjLD8UIjQIYeeYeFRQy+iLJm/NVzP/65nox55gApiA4bI49lGtO9TozIIb8DBQqMGiI0eOwudYm4ixVEQPdQHSTh322xtBtGb8qsn57iO7IIXNiyvz0zf82zTWrThvXUXJkv5TNK26kZmZmCIKAK1eubHv/ypUrLCws7PqdfD5PPp8f7wTmgTHCKo0Au+gyIIV0AMMhpAFkmfl8ggy0PoQ+Ygxa7O7t73X+PjIRurxUcbqTIsUaLfVE3InjI5ND30/M9fTM/3NY41uApAuJ61HnkYU6MMd6pY3UCBsZqDItYxe5GgGNLNWYRWwtitEABh0xUIOWGOLyhDFOZayhf7lSk+p4ZMg45rHGQR0bdS4S7HPE/K2KbD8ZYBWFzof9JDvAZ1V2nkO/m2EV2QLyDPSzN4pepQanhjhWIHNE1+wQcVAUWXDXihnzLIVshF0/fez6ANIEBu66D52Xec/zjBLvGcMUG2Ola9E53pboMy8gczN03nPvT+duBCTgZeB74KnRzMn7xPaV6n26xk6PA1sIi18Bzzg9qQeZb4bIOBlBHqICjHrgF8CfkvvLNOJSAwsv1Us5+Z8/DZQhq4ohJASumDFRfRNhjWkP+wxfJnnFjVQul+Ohhx7i4x//ON/zPd8DCBHi4x//OO9///uv8+BsDTYhdvJ75r2283cX6JhnpcZAJ7NGn10kalllfCOlhqTN7g9OlVuNLcgBkAfvikZ/dazS6yCLqIBd5OoJg13oK+ZvVeJt7JO+0dHGuJICy1jPWV9qhPPIeAfASba840NTMFWBc8+LcqkGsHQBGk0sFNvm5THAKTJ+ujD1OtUYacRqYBvWzfslbASmhnccySFecwNRAPtJH6vsxr3/qnmp45WY46iiGSH3cRzxmjWSvzDm8a8meeAU4s0fxiq4KWScN8zfHqJQB841Zub7GolopNFHIvMJZM20EKU6QsZzChsdDLHOno5XE7IQ4tA4q7qW9Kc6ku48LWAVtRpedcJKWMdlaHyboaQUtiRD5npojqVISs2ca2TuRQ3gHLAgRsjPQ1CEOBUDCIIIJQmEBSiYiLFQgFIdWhvi2GUaUZ1HIFWN3H1z7tPgz0LtNsg8MZzds2YYPHNNHfOMMkTH6HO5FtExG+73wVcJ7vvQhz7Ej/zIj/D617+eRx55hF/8xV+k0+nwV/7KX7m+A7sRiZvHcT2pFFnUTWTA1dstI5OriDyQITbc7bP/w9C8l7uodhP1fhUGVDwdZGKrsdTfVQnpxG5hJ4p6kzrRc1i4TKMpXQAFXgo9vtLieu4qLtQDW8/LC8EPIR9BPoRuC6IAohCSIYz6kB1UQV+ruDCeGq4B2+Eg2G64cs7nX67rU5hLle/VRB01nUfuPHTnqzppJp+iENmBRJWre9/TiLE4hijdScQYJohR1vk7iV1zmo8tYPMysN1zT9ke7er61nlRNvcSGr9I14zmxXTsUucVO59RR1d1ijt/Y/OZEIu66Lo115Rl5nBuxOpecyA5Mc8YrNS8tu41s/efGEMYhuCnkqoYDY1x9a09zefBy0nkVanBqAuXLpu0hjo27vMuAlXIKjA0BjwdQLaC5MfAzq+eM0bXIwqZX97/o6+Kkfr+7/9+VlZW+J/+p/+JpaUlXve61/H7v//7LyFTHFhirMejBkpDb4Uv9EFtIljrwHxOCQvq3WuyW+GH/USjODU8e4l6a23nPV3U6oG50EAfa/jWEQ/mkvm859xz37n3Tee6dZGrIhrnXl5JcT0pVbpm0QYR5CIxTqkxulEokGY8MN99JSBMNe4F7PPV5HCIVfTItVPFzkGNfscRnRvjeqeGBLSleK4mAeKddxBHp4qFMF3JITmjsvk95eBGSuEyfbYeEj0tAHeY61AY2o18ykiOtYnkMjXq0/xf1Vx/H4k6wI6Xjp2+1NDU5XxeDoKc0bUp8tzKbEsBbGl5/b5C5COswXSfj8Louv7UKCpMZ9b1tjSNnkeNWw5Cc22UxeikPewcMjCdl4fYkygn50MQgJ/BoAupudYEibBqRXHuikUolwUeXfmMgUHV4db7zpmxrAss2NXx7WBJFzqX1QEeF1W6mkwj0e/Xq5ECeP/733/98N5uoko9hywEnbxD7GTWPEaCPAzFmnXhdpCHscZ4sAvYB6mTdT9RL00Tprpg3GtTeKGBLNxnze8qGWKQIvOzb45
"text/plain": [
"<Figure size 640x480 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"import tensorflow as tf\n",
"import torch\n",
"import torchvision\n",
"from torch.utils.data import IterableDataset, DataLoader\n",
"from matplotlib import pyplot as plt\n",
"import numpy as np\n",
"import torch.nn as nn\n",
"from PIL import Image\n",
"def parse_tfrecord(example_proto):\n",
" feature_description = {\n",
" 'image': tf.io.FixedLenFeature([], tf.string),\n",
" 'class': tf.io.FixedLenFeature([], tf.int64),\n",
" 'id' : tf.io.FixedLenFeature([], tf.string),\n",
" }\n",
" parsed = tf.io.parse_single_example(example_proto, feature_description)\n",
" image = tf.image.decode_jpeg(parsed['image'], channels=3)\n",
" image = tf.image.resize(image, [224, 224])\n",
" #image = tf.image.convert_image_dtype(image, tf.float32)\n",
" label = parsed['class']\n",
" idd = parsed['id']\n",
" return image, label,idd\n",
"\n",
"def load_tfrecord_dataset(pattern):\n",
" files = tf.io.gfile.glob(pattern)\n",
" if not files:\n",
" raise ValueError(f\"No files found for pattern {pattern}\")\n",
" dataset = tf.data.TFRecordDataset(files)\n",
" dataset = dataset.map(parse_tfrecord)\n",
" # 可选:打乱、批处理等,但此处我们只返回样本级别的数据集\n",
" return dataset\n",
"\n",
"class TFRecordToPyTorch(IterableDataset):\n",
" def __init__(self, tfrecord_pattern,transform=None):\n",
" self.tfrecord_pattern = tfrecord_pattern\n",
" self.transform=transform\n",
"\n",
" def __iter__(self):\n",
" # 每次迭代创建新的数据集,保证可重复使用\n",
" dataset = load_tfrecord_dataset(self.tfrecord_pattern)\n",
" # 使用 as_numpy_iterator() 获取 NumPy 数组,便于转换为 PyTorch 张量\n",
" for image_np, label_np,idd in dataset.as_numpy_iterator():\n",
" # image_np shape: (224,224,3), dtype float32, label_np scalar int64\n",
" # 转为 PyTorch 张量,并调整为 CxHxW\n",
" image_pil = Image.fromarray((image_np).astype('uint8')) \n",
" if self.transform:\n",
" image_tensor = self.transform(image_pil)\n",
" else:\n",
" # 如果不需要 transform,至少转为 tensor\n",
" image_tensor = torch.from_numpy(image_np).permute(2,0,1)\n",
" #image_torch = torch.from_numpy(image_np).permute(2, 0, 1) # (3,224,224)\n",
" label_torch = torch.tensor(label_np, dtype=torch.long)\n",
" id_torch = idd\n",
" yield image_tensor, label_torch,id_torch\n",
"\n",
"# 使用\n",
"transform = torchvision.transforms.Compose([\n",
" torchvision.transforms.ToTensor(),\n",
" torchvision.transforms.Normalize(mean=[0.485, 0.456, 0.406],\n",
" std=[0.229, 0.224, 0.225])\n",
"])\n",
"tfrecord_path = '/kaggle/input/competitions/tpu-getting-started/tfrecords-jpeg-224x224/train/*'\n",
"dataset = TFRecordToPyTorch(tfrecord_path,transform)\n",
"tfrecord_path = '/kaggle/input/competitions/tpu-getting-started/tfrecords-jpeg-224x224/val/*'\n",
"dataset2 = TFRecordToPyTorch(tfrecord_path,transform)\n",
"# 可以配合 DataLoader 使用\n",
"train_dataloader = DataLoader(dataset, batch_size=32, num_workers=0) # num_workers 设为0,因为 TF 数据集内部已并行\n",
"val_dataloader = DataLoader(dataset2, batch_size=32, num_workers=0)\n",
"for batch in train_dataloader:\n",
" plt.imshow(batch[0][1].permute(1,2,0).numpy())\n",
" break\n",
" plt.axis('off')\n",
" plt.show()"
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "7e062233",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:25:02.145618Z",
"iopub.status.busy": "2026-07-08T12:25:02.145357Z",
"iopub.status.idle": "2026-07-08T12:25:03.234393Z",
"shell.execute_reply": "2026-07-08T12:25:03.233285Z"
},
"papermill": {
"duration": 1.09574,
"end_time": "2026-07-08T12:25:03.236312+00:00",
"exception": false,
"start_time": "2026-07-08T12:25:02.140572+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/usr/local/lib/python3.12/dist-packages/torchvision/models/_utils.py:208: UserWarning: The parameter 'pretrained' is deprecated since 0.13 and may be removed in the future, please use 'weights' instead.\n",
" warnings.warn(\n",
"/usr/local/lib/python3.12/dist-packages/torchvision/models/_utils.py:223: UserWarning: Arguments other than a weight enum or `None` for 'weights' are deprecated since 0.13 and may be removed in the future. The current behavior is equivalent to passing `weights=ResNet50_Weights.IMAGENET1K_V1`. You can also use `weights=ResNet50_Weights.DEFAULT` to get the most up-to-date weights.\n",
" warnings.warn(msg)\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Downloading: \"https://download.pytorch.org/models/resnet50-0676ba61.pth\" to /root/.cache/torch/hub/checkpoints/resnet50-0676ba61.pth\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"100%|██████████| 97.8M/97.8M [00:00<00:00, 182MB/s]\n"
]
}
],
"source": [
"pretrained_net = torchvision.models.resnet50(pretrained=True)"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "606beb4f",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:25:03.246016Z",
"iopub.status.busy": "2026-07-08T12:25:03.245773Z",
"iopub.status.idle": "2026-07-08T12:25:03.279527Z",
"shell.execute_reply": "2026-07-08T12:25:03.278755Z"
},
"papermill": {
"duration": 0.040261,
"end_time": "2026-07-08T12:25:03.280979+00:00",
"exception": false,
"start_time": "2026-07-08T12:25:03.240718+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [
{
"data": {
"text/plain": [
"Parameter containing:\n",
"tensor([[ 0.0391, -0.0359, 0.0388, ..., -0.0067, -0.0466, 0.0074],\n",
" [ 0.0360, -0.0080, 0.0314, ..., 0.0155, 0.0126, -0.0466],\n",
" [ 0.0393, 0.0281, -0.0347, ..., 0.0169, -0.0164, 0.0289],\n",
" ...,\n",
" [-0.0022, -0.0320, -0.0400, ..., -0.0068, 0.0455, -0.0202],\n",
" [-0.0064, 0.0433, 0.0035, ..., 0.0013, -0.0382, 0.0487],\n",
" [ 0.0182, 0.0425, 0.0161, ..., 0.0343, -0.0364, 0.0126]],\n",
" requires_grad=True)"
]
},
"execution_count": 3,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"pretrained_net.fc=nn.Linear(pretrained_net.fc.in_features,104)\n",
"nn.init.xavier_uniform_(pretrained_net.fc.weight)"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "3caa2b59",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:25:03.290688Z",
"iopub.status.busy": "2026-07-08T12:25:03.290007Z",
"iopub.status.idle": "2026-07-08T12:25:03.295201Z",
"shell.execute_reply": "2026-07-08T12:25:03.294433Z"
},
"papermill": {
"duration": 0.011388,
"end_time": "2026-07-08T12:25:03.296615+00:00",
"exception": false,
"start_time": "2026-07-08T12:25:03.285227+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [],
"source": [
"for parm in pretrained_net.conv1.parameters():\n",
" parm.requires_grad=False\n",
"for parm in pretrained_net.bn1.parameters():\n",
" parm.requires_grad=False\n",
"for parm in pretrained_net.layer1.parameters():\n",
" parm.requires_grad=False\n",
"for parm in pretrained_net.layer2.parameters():\n",
" parm.requires_grad=False\n",
"for parm in pretrained_net.layer3.parameters():\n",
" parm.requires_grad=False"
]
},
{
"cell_type": "code",
"execution_count": 5,
"id": "799bad35",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:25:03.310453Z",
"iopub.status.busy": "2026-07-08T12:25:03.309534Z",
"iopub.status.idle": "2026-07-08T12:25:03.314616Z",
"shell.execute_reply": "2026-07-08T12:25:03.313963Z"
},
"papermill": {
"duration": 0.014735,
"end_time": "2026-07-08T12:25:03.316029+00:00",
"exception": false,
"start_time": "2026-07-08T12:25:03.301294+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [],
"source": [
" def print_trainable_info(model):\n",
" frozen = sum(p.numel() for p in model.parameters() if not p.requires_grad)\n",
" trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n",
" total = frozen + trainable\n",
" print(f\" 冻结参数: {frozen:,} 可训练参数: {trainable:,} ({100.*trainable/total:.1f}%)\")"
]
},
{
"cell_type": "code",
"execution_count": 6,
"id": "5e78f3a1",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:25:03.325175Z",
"iopub.status.busy": "2026-07-08T12:25:03.324965Z",
"iopub.status.idle": "2026-07-08T12:25:03.330029Z",
"shell.execute_reply": "2026-07-08T12:25:03.329021Z"
},
"papermill": {
"duration": 0.011349,
"end_time": "2026-07-08T12:25:03.331384+00:00",
"exception": false,
"start_time": "2026-07-08T12:25:03.320035+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
" 冻结参数: 8,543,296 可训练参数: 15,177,832 (64.0%)\n"
]
}
],
"source": [
"print_trainable_info(pretrained_net)"
]
},
{
"cell_type": "code",
"execution_count": 7,
"id": "3eac7215",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:25:03.340381Z",
"iopub.status.busy": "2026-07-08T12:25:03.339863Z",
"iopub.status.idle": "2026-07-08T12:25:03.344586Z",
"shell.execute_reply": "2026-07-08T12:25:03.344009Z"
},
"papermill": {
"duration": 0.010776,
"end_time": "2026-07-08T12:25:03.345906+00:00",
"exception": false,
"start_time": "2026-07-08T12:25:03.335130+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [],
"source": [
"@torch.no_grad()\n",
"def validate(model,loader):\n",
" model.eval()\n",
" acc=0\n",
" total=0\n",
" for batch in loader:\n",
" X = batch[0]\n",
" labels = batch[1]\n",
" X = X.to(device)\n",
" labels = labels.to(device)\n",
" pred=torch.argmax(model(X),dim=1)\n",
" acc+=pred.eq(labels).sum()\n",
" total+=labels.size(0)\n",
" print(f\"acc:{acc/total}\")"
]
},
{
"cell_type": "code",
"execution_count": 8,
"id": "19aaeb28",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:25:03.355217Z",
"iopub.status.busy": "2026-07-08T12:25:03.355008Z",
"iopub.status.idle": "2026-07-08T12:55:31.896478Z",
"shell.execute_reply": "2026-07-08T12:55:31.895418Z"
},
"papermill": {
"duration": 1828.548066,
"end_time": "2026-07-08T12:55:31.898198+00:00",
"exception": false,
"start_time": "2026-07-08T12:25:03.350132+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 1/20: 399it [01:18, 5.06it/s, loss=1.0582, avg_loss=1.2128]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 1 train_loss: 0.0379\n",
"acc:0.8380926847457886\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 2/20: 399it [01:18, 5.07it/s, loss=0.2247, avg_loss=0.2473]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 2 train_loss: 0.0077\n",
"acc:0.860722005367279\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 3/20: 399it [01:18, 5.06it/s, loss=0.0718, avg_loss=0.0589]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 3 train_loss: 0.0018\n",
"acc:0.8809267282485962\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 4/20: 399it [01:18, 5.07it/s, loss=0.0066, avg_loss=0.0290]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 4 train_loss: 0.0009\n",
"acc:0.8741918206214905\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 5/20: 399it [01:18, 5.06it/s, loss=0.0016, avg_loss=0.0107]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 5 train_loss: 0.0003\n",
"acc:0.904633641242981\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 6/20: 399it [01:18, 5.07it/s, loss=0.0014, avg_loss=0.0027]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 6 train_loss: 0.0001\n",
"acc:0.9043642282485962\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 7/20: 399it [01:18, 5.07it/s, loss=0.0009, avg_loss=0.0011]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 7 train_loss: 0.0000\n",
"acc:0.907597005367279\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 8/20: 399it [01:18, 5.07it/s, loss=0.0008, avg_loss=0.0007]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 8 train_loss: 0.0000\n",
"acc:0.9067887663841248\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 9/20: 399it [01:18, 5.07it/s, loss=0.0007, avg_loss=0.0006]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 9 train_loss: 0.0000\n",
"acc:0.9073275923728943\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 10/20: 399it [01:18, 5.09it/s, loss=0.0007, avg_loss=0.0006]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 10 train_loss: 0.0000\n",
"acc:0.9073275923728943\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 11/20: 399it [01:18, 5.08it/s, loss=0.0007, avg_loss=0.0006]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 11 train_loss: 0.0000\n",
"acc:0.9073275923728943\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 12/20: 399it [01:18, 5.06it/s, loss=0.0007, avg_loss=0.0006]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 12 train_loss: 0.0000\n",
"acc:0.9073275923728943\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 13/20: 399it [01:18, 5.08it/s, loss=0.0006, avg_loss=0.0006]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 13 train_loss: 0.0000\n",
"acc:0.9073275923728943\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 14/20: 399it [01:18, 5.08it/s, loss=0.0005, avg_loss=0.0005]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 14 train_loss: 0.0000\n",
"acc:0.9073275923728943\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 15/20: 399it [01:18, 5.09it/s, loss=0.0003, avg_loss=0.0004]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 15 train_loss: 0.0000\n",
"acc:0.9089439511299133\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 16/20: 399it [01:18, 5.07it/s, loss=0.0002, avg_loss=0.0002]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 16 train_loss: 0.0000\n",
"acc:0.9105603694915771\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 17/20: 399it [01:19, 5.04it/s, loss=0.0002, avg_loss=0.0002]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 17 train_loss: 0.0000\n",
"acc:0.9110991358757019\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 18/20: 399it [01:18, 5.07it/s, loss=0.0001, avg_loss=0.0001]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 18 train_loss: 0.0000\n",
"acc:0.9119073152542114\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 19/20: 399it [01:19, 5.05it/s, loss=0.0001, avg_loss=0.0001]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 19 train_loss: 0.0000\n",
"acc:0.912446141242981\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"Epoch 20/20: 399it [01:18, 5.08it/s, loss=0.0000, avg_loss=0.0000]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Epoch 20 train_loss: 0.0000\n",
"acc:0.9135236740112305\n"
]
}
],
"source": [
"from tqdm import tqdm \n",
"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
"pretrained_net=pretrained_net.to(device)\n",
"loss_func = nn.CrossEntropyLoss()\n",
"optimizer = torch.optim.AdamW(pretrained_net.parameters(), lr=2e-4)\n",
"scheduler=torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=10)\n",
"epochs = 20\n",
"for epoch in range(epochs):\n",
" pretrained_net.train()\n",
" training_loss = 0\n",
" \n",
" # 使用 tqdm 包装 dataloader,并设置描述信息\n",
" progress_bar = tqdm(train_dataloader, desc=f\"Epoch {epoch+1}/{epochs}\")\n",
" lens=0\n",
" for batch in progress_bar:\n",
" optimizer.zero_grad()\n",
" X = batch[0].to(device)\n",
" labels = batch[1].to(device)\n",
" \n",
" outputs = pretrained_net(X)\n",
" loss = loss_func(outputs, labels)\n",
" loss.backward()\n",
" optimizer.step()\n",
" \n",
" training_loss += loss.item()\n",
" lens+=labels.size(0)\n",
" # 更新进度条显示当前 batch 的损失\n",
" progress_bar.set_postfix({\n",
" 'loss': f'{loss.item():.4f}',\n",
" 'avg_loss': f'{training_loss / (progress_bar.n+1):.4f}' # progress_bar.n 是已处理 batch 数\n",
" })\n",
" \n",
" scheduler.step()\n",
" \n",
" # 计算平均训练损失(注意:len(train_dataloader) 才是 batch 总数)\n",
" avg_train_loss = training_loss / lens\n",
" print(f\"Epoch {epoch+1} train_loss: {avg_train_loss:.4f}\")\n",
" \n",
" # 验证(你也可以为验证添加进度条,见下方建议)\n",
" validate(pretrained_net, val_dataloader)\n",
" "
]
},
{
"cell_type": "code",
"execution_count": 9,
"id": "981cd43a",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:55:33.032330Z",
"iopub.status.busy": "2026-07-08T12:55:33.031985Z",
"iopub.status.idle": "2026-07-08T12:56:06.103071Z",
"shell.execute_reply": "2026-07-08T12:56:06.102141Z"
},
"papermill": {
"duration": 33.636176,
"end_time": "2026-07-08T12:56:06.104939+00:00",
"exception": false,
"start_time": "2026-07-08T12:55:32.468763+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [],
"source": [
"import pandas as pd\n",
"def parse_tfrecord_test(example_proto):\n",
" feature_description = {\n",
" 'image': tf.io.FixedLenFeature([], tf.string),\n",
" 'id' : tf.io.FixedLenFeature([], tf.string)\n",
" }\n",
" parsed = tf.io.parse_single_example(example_proto, feature_description)\n",
" image = tf.image.decode_jpeg(parsed['image'], channels=3)\n",
" image = tf.image.resize(image, [224, 224])\n",
" #image = tf.image.convert_image_dtype(image, tf.float32)\n",
" idd = parsed['id']\n",
" return image,idd\n",
"def load_tfrecord_dataset_test(pattern):\n",
" files = tf.io.gfile.glob(pattern)\n",
" if not files:\n",
" raise ValueError(f\"No files found for pattern {pattern}\")\n",
" dataset = tf.data.TFRecordDataset(files)\n",
" dataset = dataset.map(parse_tfrecord_test)\n",
" # 可选:打乱、批处理等,但此处我们只返回样本级别的数据集\n",
" return dataset\n",
"class TFRecordToPyTorchTest(IterableDataset):\n",
" def __init__(self, tfrecord_pattern,transform=None):\n",
" self.tfrecord_pattern = tfrecord_pattern\n",
" self.transform=transform\n",
"\n",
" def __iter__(self):\n",
" # 每次迭代创建新的数据集,保证可重复使用\n",
" dataset = load_tfrecord_dataset_test(self.tfrecord_pattern)\n",
" # 使用 as_numpy_iterator() 获取 NumPy 数组,便于转换为 PyTorch 张量\n",
" for image_np,idd in dataset.as_numpy_iterator():\n",
" # image_np shape: (224,224,3), dtype float32, label_np scalar int64\n",
" # 转为 PyTorch 张量,并调整为 CxHxW\n",
" image_pil = Image.fromarray((image_np).astype('uint8')) \n",
" if self.transform:\n",
" image_tensor = self.transform(image_pil)\n",
" else:\n",
" # 如果不需要 transform,至少转为 tensor\n",
" image_tensor = torch.from_numpy(image_np).permute(2,0,1)\n",
" #image_torch = torch.from_numpy(image_np).permute(2, 0, 1) # (3,224,224)\n",
" #label_torch = torch.tensor(label_np, dtype=torch.long)\n",
" id_torch = idd\n",
" yield image_tensor,id_torch\n",
"tfrecord_path = '/kaggle/input/competitions/tpu-getting-started/tfrecords-jpeg-224x224/test/*'\n",
"dataset3 = TFRecordToPyTorchTest(tfrecord_path,transform)\n",
"test_dataloader = DataLoader(dataset3, batch_size=32, num_workers=0)\n",
"id_array=[]\n",
"all_preds=[]\n",
"with torch.no_grad():\n",
" for batch in test_dataloader:\n",
" input_ids = batch[0].to(device)\n",
" idd = batch[1]\n",
" outputs = pretrained_net(input_ids)\n",
" preds = torch.argmax(outputs, dim=1)\n",
" all_preds.extend(preds.cpu().numpy())\n",
" id_array.extend(idd)\n",
"submission = pd.DataFrame({\n",
" 'id':id_array,\n",
" 'label': all_preds\n",
"})\n"
]
},
{
"cell_type": "code",
"execution_count": 10,
"id": "1e950456",
"metadata": {
"execution": {
"iopub.execute_input": "2026-07-08T12:56:07.318686Z",
"iopub.status.busy": "2026-07-08T12:56:07.318282Z",
"iopub.status.idle": "2026-07-08T12:56:07.361026Z",
"shell.execute_reply": "2026-07-08T12:56:07.360031Z"
},
"papermill": {
"duration": 0.611337,
"end_time": "2026-07-08T12:56:07.362666+00:00",
"exception": false,
"start_time": "2026-07-08T12:56:06.751329+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
" id label\n",
"0 59d1b6146 46\n",
"1 48c96bd6b 15\n",
"2 7b437ba4e 9\n",
"3 1b7aef8e8 79\n",
"4 d6143b4d4 4\n",
"... ... ...\n",
"7377 2a608c0db 103\n",
"7378 d82a21bbd 93\n",
"7379 f9c931893 53\n",
"7380 18c7b92b8 41\n",
"7381 523df966b 102\n",
"\n",
"[7382 rows x 2 columns]\n",
"Submission saved!\n"
]
}
],
"source": [
"submission['id'] = submission['id'].apply(lambda x: x.decode('utf-8'))\n",
"print(submission)\n",
"submission.to_csv('submission.csv', index=False)\n",
"print(\"Submission saved!\")"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "79c03587",
"metadata": {
"papermill": {
"duration": 0.645772,
"end_time": "2026-07-08T12:56:08.568042+00:00",
"exception": false,
"start_time": "2026-07-08T12:56:07.922270+00:00",
"status": "completed"
},
"tags": []
},
"outputs": [],
"source": []
}
],
"metadata": {
"kaggle": {
"accelerator": "none",
"dataSources": [],
"dockerImageVersionId": 28755,
"isGpuEnabled": false,
"isInternetEnabled": false,
"language": "python",
"sourceType": "notebook"
},
"kernelspec": {
"display_name": "Python 3",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.13"
},
"papermill": {
"default_parameters": {},
"duration": 1903.252205,
"end_time": "2026-07-08T12:56:12.294916+00:00",
"environment_variables": {},
"exception": null,
"input_path": "__notebook__.ipynb",
"output_path": "__notebook__.ipynb",
"parameters": {},
"start_time": "2026-07-08T12:24:29.042711+00:00",
"version": "2.7.0"
}
},
"nbformat": 4,
"nbformat_minor": 5
}