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

1 line
205 KiB
Plaintext
Raw Normal View History

2026-07-09 21:42:42 +08:00
{"metadata":{"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"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isGpuEnabled":false,"isInternetEnabled":false,"language":"python","sourceType":"notebook"},"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_minor":5,"nbformat":4,"cells":[{"id":"befdc3e8","cell_type":"code","source":"import tensorflow as tf\nimport torch\nimport torchvision\nfrom torch.utils.data import IterableDataset, DataLoader\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport torch.nn as nn\nfrom PIL import Image\ndef 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\ndef 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\nclass 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","metadata":{"execution":{"iopub.status.busy":"2026-07-09T04:07:12.551342Z","iopub.execute_input":"2026-07-09T04:07:12.551925Z","iopub.status.idle":"2026-07-09T04:07:48.180247Z","shell.execute_reply.started":"2026-07-09T04:07:12.551894Z","shell.execute_reply":"2026-07-09T04:07:48.179571Z"},"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":[],"trusted":true},"outputs":[],"execution_count":1},{"id":"92b8d6c3-f964-4e11-97ed-2482ad60772b","cell_type":"code","source":"from transformers import ViTForImageClassification, ViTImageProcessor\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n \n\nmodel_name = \"google/vit-base-patch16-224-in21k\" # 在ImageNet21k上预训练\n\nmodel = ViTForImageClassification.from_pret