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

909 lines
218 KiB
Text
Raw Normal View History

2026-07-08 13:58:01 +00: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
}