{"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.11.11"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":6818,"databundleVersionId":1960702,"sourceType":"competition"},{"sourceId":8242683,"sourceType":"kernelVersion"}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"ef9c9674","cell_type":"code","source":"\nfrom fastai.vision.all import *\nimport pandas as pd\nimport numpy as np\nimport os\nfrom sklearn.model_selection import train_test_split\nimport matplotlib.pyplot as plt\n","metadata":{"execution":{"iopub.status.busy":"2025-04-18T20:07:34.382161Z","iopub.execute_input":"2025-04-18T20:07:34.382389Z","iopub.status.idle":"2025-04-18T20:07:46.175861Z","shell.execute_reply.started":"2025-04-18T20:07:34.382366Z","shell.execute_reply":"2025-04-18T20:07:46.175290Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"bde55b41-a3bd-468f-b728-2e275015aa31","cell_type":"code","source":"import fastai\nprint(fastai.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T20:07:56.504798Z","iopub.execute_input":"2025-04-18T20:07:56.505108Z","iopub.status.idle":"2025-04-18T20:07:56.509830Z","shell.execute_reply.started":"2025-04-18T20:07:56.505085Z","shell.execute_reply":"2025-04-18T20:07:56.509007Z"}},"outputs":[],"execution_count":null},{"id":"b988a5ba-eda3-48bf-9405-3f241124aa8d","cell_type":"code","source":"import torch\nprint(torch.cuda.get_device_name(0))","metadata":{"execution":{"iopub.status.busy":"2025-04-18T20:07:58.541379Z","iopub.execute_input":"2025-04-18T20:07:58.542055Z","iopub.status.idle":"2025-04-18T20:07:58.618028Z","shell.execute_reply.started":"2025-04-18T20:07:58.542033Z","shell.execute_reply":"2025-04-18T20:07:58.617409Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"ebf307ab-446d-4a9d-9635-cf7eb348df61","cell_type":"code","source":"MODEL_NAME = 'Resnet50'\nTRAIN = Path('../input/humpback-whale-identification/train/')\nTEST = Path('../input/humpback-whale-identification/test/')\nLABELS = Path('../input/humpback-whale-identification/train.csv')\nSAMPLE_SUB = Path('../input/humpback-whale-identification/sample_submission.csv')\nBBOX = Path('../input/generating-whale-bounding-boxes/bounding_boxes.csv')\n# Backbone architecture\narch = resnet50\n# Number of workers for data preprocessing\nnum_workers = 4","metadata":{"execution":{"iopub.status.busy":"2025-04-18T20:08:00.573060Z","iopub.execute_input":"2025-04-18T20:08:00.573720Z","iopub.status.idle":"2025-04-18T20:08:00.578809Z","shell.execute_reply.started":"2025-04-18T20:08:00.573685Z","shell.execute_reply":"2025-04-18T20:08:00.578033Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"b0bdd1a2-15a8-464f-a8c6-d886ced711da","cell_type":"code","source":"df = pd.read_csv(LABELS).set_index('Image')\nnew_whale_df = df[df.Id == \"new_whale\"] # only new_whale dataset\ntrain_df = df[df.Id != \"new_whale\"].copy()\n\nunique_labels = np.unique(train_df.Id.values)\nlabels_list = unique_labels.tolist()\nlabels_dict = {label: i for i, label in enumerate(unique_labels)}\n# labels_dict = dict()\n# labels_list = []\n# for i in range(len(unique_labels)):\n#     labels_dict[unique_labels[i]] = i\n#     labels_list.append(unique_labels[i])\n# print(\"Number of classes: {}\".format(len(unique_labels)))\n# train_df.Id = train_df.Id.apply(lambda x: labels_dict[x])\n# train_labels = np.asarray(train_df.Id.values)\n# test_names = [f for f in os.listdir(TEST)]\n# train_df['image_name'] = train_df.index","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T20:08:03.674631Z","iopub.execute_input":"2025-04-18T20:08:03.674898Z","iopub.status.idle":"2025-04-18T20:08:03.736805Z","shell.execute_reply.started":"2025-04-18T20:08:03.674878Z","shell.execute_reply":"2025-04-18T20:08:03.736071Z"}},"outputs":[],"execution_count":null},{"id":"da3cf81c-9d06-4563-8d33-d3c7b6d8bbc7","cell_type":"code","source":"# labels_count = train_df.Id.value_counts()\nlabels_count = df[df.Id != \"new_whale\"].Id.value_counts()\n\nplt.figure(figsize=(18, 4))\nplt.subplot(121)\n_, _,_ = plt.hist(labels_count.values)\nplt.ylabel(\"frequency\")\nplt.xlabel(\"class size\")\n\nplt.title('class distribution; log scale')\nlabels_count.head()\n\nplt.subplot(122)\n_ = plt.plot(labels_count[1:].values)\nplt.title('w/o class new_whale; log scale')\nplt.xlabel(\"class\")\nplt.ylabel(\"log(size)\")\nplt.gca().set_yscale('log')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T20:08:08.027115Z","iopub.execute_input":"2025-04-18T20:08:08.027416Z","iopub.status.idle":"2025-04-18T20:08:08.565655Z","shell.execute_reply.started":"2025-04-18T20:08:08.027395Z","shell.execute_reply":"2025-04-18T20:08:08.564840Z"}},"outputs":[],"execution_count":null},{"id":"5600c99b-6e78-4e25-a08b-cd24b36aed4d","cell_type":"code","source":"train_df['image_name'] = train_df.index\nbbox_df = pd.read_csv(BBOX).set_index('Image')\n\nrs = np.random.RandomState(42) # set random seed to be equal to the sense of life\nperm = rs.permutation(len(train_df))\n\ntr_n = train_df['image_name'].values\n\nval_n = train_df['image_name'].values[perm][:1000]\n\ntrain_labels = set(train_df.loc[tr_n, 'Id'])\nvalid_df = train_df.loc[val_n]\nvalid_df = valid_df[valid_df['Id'].isin(train_labels)]\nval_n = valid_df['image_name'].values\n# train_ids_for_vocab = train_df.loc[train_n, 'Id']\n# label_vocab = train_ids_for_vocab.unique().tolist()\n\n# valid_df = train_df.loc[val_n]\n# valid_df = valid_df[valid_df['Id'].isin(set(label_vocab))]\n# val_n = valid_df['image_name'].values\n\n\nprint('Train/val:', len(tr_n), len(val_n))\nprint('Train classes', len(train_df.loc[tr_n].Id.unique()))\nprint('Val classes', len(train_df.loc[val_n].Id.unique()))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T20:08:13.683848Z","iopub.execute_input":"2025-04-18T20:08:13.684700Z","iopub.status.idle":"2025-04-18T20:08:13.775316Z","shell.execute_reply.started":"2025-04-18T20:08:13.684667Z","shell.execute_reply":"2025-04-18T20:08:13.774546Z"}},"outputs":[],"execution_count":null},{"id":"1ff3b8db-77ed-4f92-a5ad-baa4c2b7a6c2","cell_type":"code","source":"from fastai.vision.augment import *\nimport random\nimport cv2\n\n# 自定义模糊变换类\nclass RandomBlur(Transform):\n    def __init__(self, blur_strengths=3, p=0.5):\n        super().__init__()\n        self.blur_strengths = blur_strengths\n        self.p = p\n        \n    def encodes(self, img):\n        if random.random() < self.p:\n            # 将PIL图像转换为opencv格式\n            img_array = np.array(img)\n            \n            # 随机高斯模糊\n            blur_amount = random.randint(1, self.blur_strengths) * 2 + 1  # 必须是奇数\n            img_array = cv2.GaussianBlur(img_array, (blur_amount, blur_amount), 0)\n            \n            # 转回PIL图像\n            return PILImage.create(img_array)\n        return img\n\n# 基本数据增强\nbase_aug = aug_transforms(\n    max_rotate=20,    # 20度旋转\n    max_zoom=2,       # 2倍缩放\n    max_warp=0,       # 不使用warp变换\n    max_lighting=0.2, # 亮度变化\n    do_flip=True,     # 启用翻转\n    p_affine=0.75,    # 仿射变换概率\n    p_lighting=0.75   # 亮度变换概率\n)\n\nblur_transform = RandomBlur(blur_strengths=3, p=0.5)\n\n# 组合所有变换\nfinal_transforms = base_aug + [blur_transform]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T20:42:19.923750Z","iopub.execute_input":"2025-04-18T20:42:19.924567Z","iopub.status.idle":"2025-04-18T20:42:20.234378Z","shell.execute_reply.started":"2025-04-18T20:42:19.924535Z","shell.execute_reply":"2025-04-18T20:42:20.233612Z"}},"outputs":[],"execution_count":null},{"id":"10af929b-52e9-4779-becf-896954d895d8","cell_type":"code","source":"# def open_cropped_image(fname):\n#     # 拼接完整路径并读取图像\n#     img_path = TRAIN / fname\n#     img = PILImage.create(img_path)\n\n#     # 查找裁剪框\n#     bbox = bbox_df.loc[fname]  # fname 是纯图片名，如 '0000e88ab.jpg'\n#     x0, y0, x1, y1 = bbox['x0'], bbox['y0'], bbox['x1'], bbox['y1']\n\n#     # 如果 bbox 合法，就裁剪\n#     if x0 < x1 and y0 < y1:\n#         img = img.crop((x0, y0, x1, y1))  # PIL 图像支持 crop\n#     return img\n\ndef open_cropped_image(fname):\n    img_path = TRAIN / fname\n    # 使用PIL加载图像\n    img = PILImage.create(img_path)\n    \n    # 获取边界框\n    bbox = bbox_df.loc[fname]\n    x0, y0, x1, y1 = bbox['x0'], bbox['y0'], bbox['x1'], bbox['y1']\n    \n    # 裁剪(使用与参考代码相同的条件检查)\n    if not (x0 >= x1 or y0 >= y1):\n        img = img.crop((x0, y0, x1, y1))\n    \n    # 转换为numpy并调整大小(模拟参考代码行为)\n    img_array = np.array(img)\n    img_array = cv2.resize(img_array, (384, 384))\n    \n    # 转回PIL图像格式\n    return PILImage.create(img_array)\n\nwhale_df = train_df.copy()\n\nvocab = CategoryMap(labels_list, sort=False)\n\nwhale_block = DataBlock(\n    blocks=(ImageBlock, CategoryBlock(vocab=vocab)),\n    get_items=lambda df: df.index.tolist(),  # 从 df 拿图片名\n    get_x=open_cropped_image,\n    # 关键修改：确保返回的是字符串标签，而不是数字ID\n    # get_y=lambda o: df.loc[o, 'Id'] if o in df.index else None,\n    get_y=lambda o: train_df.loc[o, 'Id'] if o in train_df.index else None,\n    splitter=FuncSplitter(lambda o: o in val_n),  # 用你已有的验证集划分\n    item_tfms=Resize(384),\n    batch_tfms=aug_transforms(do_flip=True, max_rotate=20, max_zoom=2,\n                              max_lighting=0.2, max_warp=0.2, p_affine=0.75, p_lighting=0.75)\n    # batch_tfms=final_transforms\n)\n\n# dls = whale_block.dataloaders(df.loc[df.Id != \"new_whale\"], bs=32, num_workers=4)\ndls = whale_block.dataloaders(train_df, bs=32, num_workers=4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T20:58:15.644295Z","iopub.execute_input":"2025-04-18T20:58:15.644585Z","iopub.status.idle":"2025-04-18T20:58:16.206568Z","shell.execute_reply.started":"2025-04-18T20:58:15.644562Z","shell.execute_reply":"2025-04-18T20:58:16.205773Z"}},"outputs":[],"execution_count":null},{"id":"dc028d03-f32f-4a58-ba64-4bb814c4cc87","cell_type":"code","source":"learn = vision_learner(\n    dls,\n    resnet50,\n    metrics=accuracy,\n    pretrained=True,\n    opt_func=Adam,\n    lin_ftrs=[]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T20:58:19.258551Z","iopub.execute_input":"2025-04-18T20:58:19.258826Z","iopub.status.idle":"2025-04-18T20:58:20.096327Z","shell.execute_reply.started":"2025-04-18T20:58:19.258806Z","shell.execute_reply":"2025-04-18T20:58:20.095524Z"}},"outputs":[],"execution_count":null},{"id":"3c2abd99","cell_type":"code","source":"lrs = slice(1e-4,1e-3)\nlearn.freeze()\nlearn.fit_one_cycle(2, lr_max=1e-3)\n\nlearn.unfreeze()\nlearn.fit_one_cycle(16, lr_max=lrs)\n","metadata":{"execution":{"iopub.status.busy":"2025-04-18T20:58:22.128668Z","iopub.execute_input":"2025-04-18T20:58:22.129335Z","iopub.status.idle":"2025-04-18T22:03:15.882645Z","shell.execute_reply.started":"2025-04-18T20:58:22.129303Z","shell.execute_reply":"2025-04-18T22:03:15.881909Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"8a34d285","cell_type":"code","source":"test_files = get_image_files(TEST)\ntest_dl = dls.test_dl(test_files, with_labels=False)\npreds_t, _ = learn.tta(dl=test_dl, n=8)\nprobs = preds_t.softmax(dim=1).numpy()\n\n# 插入 new_whale\nbest_th = 0.38\nprobs = np.concatenate([np.full((probs.shape[0], 1), best_th), probs], axis=1)\n\nlabels_list_full = [\"new_whale\"] + labels_list\n\n# 拼 top5\ntop5_preds = [[labels_list_full[i] for i in p.argsort()[-5:][::-1]] for p in probs]\ntest_fnames = [f.name for f in test_files]\npred_dic = dict(zip(test_fnames, top5_preds))\n\nsample_df = pd.read_csv(SAMPLE_SUB)\nsample_list = list(sample_df.Image)\npred_list_cor = [' '.join(pred_dic[img]) for img in sample_list]\n\ndf_sub = pd.DataFrame({'Image': sample_list, 'Id': pred_list_cor})\ndf_sub.to_csv(f'submission_{MODEL_NAME}.csv', index=False)\ndf_sub.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-18T22:11:21.424986Z","iopub.execute_input":"2025-04-18T22:11:21.425305Z","iopub.status.idle":"2025-04-18T22:18:25.363375Z","shell.execute_reply.started":"2025-04-18T22:11:21.425279Z","shell.execute_reply":"2025-04-18T22:18:25.362605Z"}},"outputs":[],"execution_count":null}]}