{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":37333,"databundleVersionId":3949526,"sourceType":"competition"}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:32:58.316835Z","iopub.execute_input":"2024-12-03T07:32:58.317769Z","iopub.status.idle":"2024-12-03T07:33:01.315611Z","shell.execute_reply.started":"2024-12-03T07:32:58.317727Z","shell.execute_reply":"2024-12-03T07:33:01.314666Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip uninstall -y lightning-flash pytorch-lightning torchvision torch lightning-bolts","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:33:01.317204Z","iopub.execute_input":"2024-12-03T07:33:01.318074Z","iopub.status.idle":"2024-12-03T07:33:21.451193Z","shell.execute_reply.started":"2024-12-03T07:33:01.318009Z","shell.execute_reply":"2024-12-03T07:33:21.450226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install torch==1.13.1 torchvision==0.14.1\n!pip install lightning-flash","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:33:21.452633Z","iopub.execute_input":"2024-12-03T07:33:21.452997Z","iopub.status.idle":"2024-12-03T07:34:50.972325Z","shell.execute_reply.started":"2024-12-03T07:33:21.452953Z","shell.execute_reply":"2024-12-03T07:34:50.97124Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torchvision\nimport torchvision.transforms as transforms\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n# 设置数据集路径\nDATASET_FOLDER = \"/kaggle/input/mayo-clinic-strip-ai/\"\nDATASET_SMALL_FOLDER = \"/kaggle/input/newtrain\"\n\n# 验证路径是否存在\nif not os.path.exists(DATASET_FOLDER):\n    print(f\"警告: {DATASET_FOLDER} 不存在\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:34:50.974478Z","iopub.execute_input":"2024-12-03T07:34:50.97479Z","iopub.status.idle":"2024-12-03T07:34:52.251608Z","shell.execute_reply.started":"2024-12-03T07:34:50.97476Z","shell.execute_reply":"2024-12-03T07:34:52.250891Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nimport keras.backend as K #to define custom loss function\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n%matplotlib inline\nimport warnings\nwarnings.filterwarnings('ignore')\n\nfrom pprint import pprint\nfrom collections import defaultdict\nimport openslide\nfrom openslide import OpenSlide\n\nfrom glob import glob\n\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv2D, MaxPooling2D, Dropout, Flatten, Dense\nfrom tensorflow.keras.layers import GlobalMaxPooling2D\nfrom keras.models import load_model\n\nprint(keras.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:34:52.252543Z","iopub.execute_input":"2024-12-03T07:34:52.252771Z","iopub.status.idle":"2024-12-03T07:34:52.915989Z","shell.execute_reply.started":"2024-12-03T07:34:52.252746Z","shell.execute_reply":"2024-12-03T07:34:52.914819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\n\ntrain_df = pd.read_csv('../input/mayo-clinic-strip-ai/train.csv')\ntest_df  = pd.read_csv('../input/mayo-clinic-strip-ai/test.csv')\n\n# Specify patient_ids to remove\npatient_ids_to_remove = ['006388', '008e5c', '00c058', '01adc5']\n\n# Filter out the rows with specified patient_ids\ntrain_df = train_df[~train_df['patient_id'].isin(patient_ids_to_remove)].reset_index(drop=True)\ntrain_df = train_df.drop_duplicates(subset=['patient_id'])\n\ndf1 = train_df[train_df['label'] == 'CE']\ndf2 = train_df[train_df['label'] == 'LAA']\n#adjust n to change number of CE data\nsampled= df1.sample(n=200, random_state=42)\ntrain_df = pd.concat([sampled, df2],ignore_index=True)\n# Print the cleaned DataFrame\nprint(\"Cleaned DataFrame:\")\nprint(train_df.head())\ntrain_df['label'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:34:52.917434Z","iopub.execute_input":"2024-12-03T07:34:52.917937Z","iopub.status.idle":"2024-12-03T07:34:52.964801Z","shell.execute_reply.started":"2024-12-03T07:34:52.917908Z","shell.execute_reply":"2024-12-03T07:34:52.963961Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df['patient_id'].nunique","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:34:52.965915Z","iopub.execute_input":"2024-12-03T07:34:52.966183Z","iopub.status.idle":"2024-12-03T07:34:52.972157Z","shell.execute_reply.started":"2024-12-03T07:34:52.966158Z","shell.execute_reply":"2024-12-03T07:34:52.971364Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ntrain_images = glob(\"/kaggle/input/mayo-clinic-strip-ai/train/*\")\ntest_images = glob(\"/kaggle/input/mayo-clinic-strip-ai/test/*\")\nother_images = glob(\"/kaggle/input/mayo-clinic-strip-ai/other/*\")\nprint(f\"Number of images in a training set: {len(train_images)}\")\nprint(f\"Number of images in a training set: {len(test_images)}\")\nprint(f\"Number of other: {len(other_images)}\")\n\n# Filtering out images based on the cleaned patient_ids in train_df\nimages_to_remove = [\n    '/kaggle/input/mayo-clinic-strip-ai/train/006388_0.tif',\n    '/kaggle/input/mayo-clinic-strip-ai/train/008e5c_0.tif',\n    '/kaggle/input/mayo-clinic-strip-ai/train/00c058_0.tif',\n    '/kaggle/input/mayo-clinic-strip-ai/train/01adc5_0.tif',\n]\n\n# Remove images associated with the patient_ids\ntrain_images = [img for img in train_images if img not in images_to_remove]\n\n# Check the total number of images after deletion\ntotal_images_after_deletion = len(train_images)\nprint(\"Total number of images after deletion:\", total_images_after_deletion)\n\n# Print the paths of the cleaned list of images\nprint(\"First 5 image paths after deletion:\")\nprint(train_images[:5])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:34:52.973331Z","iopub.execute_input":"2024-12-03T07:34:52.973659Z","iopub.status.idle":"2024-12-03T07:34:52.987785Z","shell.execute_reply.started":"2024-12-03T07:34:52.973622Z","shell.execute_reply":"2024-12-03T07:34:52.987027Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df[\"file_path\"] = train_df[\"image_id\"].apply(lambda x: \"../input/mayo-clinic-strip-ai/train/\" + x + \".tif\")\ntest_df[\"file_path\"]  = test_df[\"image_id\"].apply(lambda x: \"../input/mayo-clinic-strip-ai/test/\" + x + \".tif\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:34:52.988825Z","iopub.execute_input":"2024-12-03T07:34:52.989161Z","iopub.status.idle":"2024-12-03T07:34:52.996997Z","shell.execute_reply.started":"2024-12-03T07:34:52.989126Z","shell.execute_reply":"2024-12-03T07:34:52.99637Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# labelling CE class as 1 and LAA as 0\ntrain_df[\"target\"] = train_df[\"label\"].apply(lambda x : 1 if x==\"CE\" else 0)\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:34:52.999786Z","iopub.execute_input":"2024-12-03T07:34:53.000097Z","iopub.status.idle":"2024-12-03T07:34:53.015041Z","shell.execute_reply.started":"2024-12-03T07:34:53.000071Z","shell.execute_reply":"2024-12-03T07:34:53.013821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%time\ntrain_df[\"file_path\"] = train_df[\"image_id\"].apply(lambda x: \"../input/mayo-clinic-strip-ai/train/\" + x + \".tif\")\n\ndef preprocess(image_path):\n    slide=OpenSlide(image_path)\n    region= (2500,2500)    \n    size  = (5000, 5000)\n    image = slide.read_region(region, 0, size)\n    image = image.resize((128, 128))\n    image = np.array(image)    \n    return image\n\nX_train=[]\nfor i in tqdm(train_df['file_path']):\n    x1=preprocess(i)\n    X_train.append(x1)\n\nY_train=[]    \nY_train=train_df['target']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:34:53.016205Z","iopub.execute_input":"2024-12-03T07:34:53.016527Z","iopub.status.idle":"2024-12-03T07:57:50.224326Z","shell.execute_reply.started":"2024-12-03T07:34:53.016481Z","shell.execute_reply":"2024-12-03T07:57:50.223404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"X_train=np.array(X_train)\nX_train=X_train/255.0\nY_train = np.array(Y_train)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T14:25:54.180627Z","iopub.execute_input":"2024-12-03T14:25:54.180959Z","iopub.status.idle":"2024-12-03T14:25:54.200481Z","shell.execute_reply.started":"2024-12-03T14:25:54.180933Z","shell.execute_reply":"2024-12-03T14:25:54.199275Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nX_train,X_test,Y_train,Y_test=train_test_split(X_train,Y_train, test_size=0.3, random_state=42)\nX_train,X_val,Y_train,Y_val=train_test_split(X_train,Y_train, test_size=0.5, random_state=42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T14:25:10.928938Z","iopub.execute_input":"2024-12-03T14:25:10.929309Z","iopub.status.idle":"2024-12-03T14:25:12.517938Z","shell.execute_reply.started":"2024-12-03T14:25:10.929267Z","shell.execute_reply":"2024-12-03T14:25:12.516747Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install tensorflow numpy pandas tqdm\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:57:50.374215Z","iopub.execute_input":"2024-12-03T07:57:50.37449Z","iopub.status.idle":"2024-12-03T07:57:58.951654Z","shell.execute_reply.started":"2024-12-03T07:57:50.374464Z","shell.execute_reply":"2024-12-03T07:57:58.950476Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras.applications import EfficientNetB6\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.layers import Dense, Dropout, GlobalAveragePooling2D\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.utils import to_categorical\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Y_train shape:\", Y_train.shape)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Unique values in Y_train:\", np.unique(Y_train))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tensorflow.keras.utils import to_categorical\n\n# 将整数标签转换为 one-hot 编码\nY_train = to_categorical(Y_train, num_classes=2)\n\nprint(\"Y_train shape after one-hot encoding:\", Y_train.shape)\n# 输出: (373, 2)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nfrom tensorflow.keras.utils import to_categorical\n\n\n\n# 数据增强\ndatagen = ImageDataGenerator(\n    rotation_range=20,\n    width_shift_range=0.2,\n    height_shift_range=0.2,\n    shear_range=0.2,\n    zoom_range=0.2,\n    horizontal_flip=True,\n    fill_mode='nearest'\n)\n\n# 生成增强数据\ntrain_generator = datagen.flow(\n    X_train, \n    Y_train, \n    batch_size=32  # 去掉 class_mode\n)\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tensorflow.keras.applications import EfficientNetB6\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.layers import Dense, Dropout, GlobalAveragePooling2D\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.utils import to_categorical\n\n\n\n# 加载 EfficientNet-B6 的预训练模型\nbase_model = EfficientNetB6(weights=None, include_top=False, input_shape=(128, 128, 4))\n\n# 冻结预训练模型的卷积层\nbase_model.trainable = False\n\n# 添加自定义分类头\nx = base_model.output\nx = GlobalAveragePooling2D()(x)\nx = Dropout(0.5)(x)\nx = Dense(256, activation='relu')(x)\nx = Dropout(0.5)(x)\npredictions = Dense(Y_train.shape[1], activation='softmax')(x)  # 输出层\n\n# 定义完整模型\nmodel = Model(inputs=base_model.input, outputs=predictions)\n\n# 编译模型\nmodel.compile(optimizer=Adam(learning_rate=0.001), loss='categorical_crossentropy', metrics=['accuracy'])\n\n# 打印模型结构\nmodel.summary()\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:57:59.080076Z","iopub.execute_input":"2024-12-03T07:57:59.080332Z","iopub.status.idle":"2024-12-03T07:58:03.184419Z","shell.execute_reply.started":"2024-12-03T07:57:59.080305Z","shell.execute_reply":"2024-12-03T07:58:03.183537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport subprocess\nimport tensorflow as tf\n\ndef check_cuda_version():\n    try:\n        # Run nvcc command to get CUDA version\n        result = subprocess.run(['nvcc', '--version'], stdout=subprocess.PIPE, text=True)\n        print(\"Current CUDA version:\\n\", result.stdout)\n    except FileNotFoundError:\n        print(\"CUDA not found. Please ensure it is installed.\")\n\ndef check_cudnn_version():\n    try:\n        with open('/usr/include/cudnn_version.h', 'r') as f:\n            for line in f:\n                if 'CUDNN_MAJOR' in line or 'CUDNN_MINOR' in line or 'CUDNN_PATCHLEVEL' in line:\n                    print(line.strip())\n    except FileNotFoundError:\n        print(\"CuDNN not found. Please ensure it is installed.\")\n\ndef check_tf_gpu():\n    print(\"Verifying TensorFlow GPU access...\")\n    gpus = tf.config.list_physical_devices('GPU')\n    print(\"Num GPUs Available: \", len(gpus))\n\nif __name__ == \"__main__\":\n    check_cuda_version()\n    check_cudnn_version()\n    check_tf_gpu()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:58:03.185556Z","iopub.execute_input":"2024-12-03T07:58:03.185847Z","iopub.status.idle":"2024-12-03T07:58:03.216957Z","shell.execute_reply.started":"2024-12-03T07:58:03.185819Z","shell.execute_reply":"2024-12-03T07:58:03.216165Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"conda install -c conda-forge cudnn=8.6.x  # Specify the version compatible with your TensorFlow\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:58:03.217934Z","iopub.execute_input":"2024-12-03T07:58:03.218254Z","iopub.status.idle":"2024-12-03T07:59:32.977849Z","shell.execute_reply.started":"2024-12-03T07:58:03.218226Z","shell.execute_reply":"2024-12-03T07:59:32.97673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint\n\n# 定义回调函数\nearly_stopping = EarlyStopping(monitor='loss', patience=5, restore_best_weights=True)  # 监控训练集的 loss\nmodel_checkpoint = ModelCheckpoint('efficientnet_b6_best_model.keras', save_best_only=True, monitor='loss')  # 保存基于训练集 loss 的最优模型\n\n# 开始训练\nhistory = model.fit(\n    train_generator,  # 训练数据生成器\n    epochs=20,        # 训练轮数\n    callbacks=[early_stopping, model_checkpoint]\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T07:59:32.979478Z","iopub.execute_input":"2024-12-03T07:59:32.979768Z","iopub.status.idle":"2024-12-03T08:03:20.377993Z","shell.execute_reply.started":"2024-12-03T07:59:32.979741Z","shell.execute_reply":"2024-12-03T08:03:20.377154Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ntest1=[]\nfor i in test_df['file_path']:\n    x1=preprocess(i)\n    test1.append(x1)\n    print(i)\n    \ntest1=np.array(test1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T08:03:20.383981Z","iopub.execute_input":"2024-12-03T08:03:20.384303Z","iopub.status.idle":"2024-12-03T08:03:37.09249Z","shell.execute_reply.started":"2024-12-03T08:03:20.384275Z","shell.execute_reply":"2024-12-03T08:03:37.091598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 假设类别索引与名称的映射\nclass_names = ['CE', 'LAA']  # 0 对应 'CE'，1 对应 'LAA'\n\n# 加载保存的模型\nfrom tensorflow.keras.models import load_model\nmodel = load_model('efficientnet_b6_best_model.keras')\nprint(\"Model loaded successfully!\")\n\n# 确保 test1 数据已经准备好并归一化（如果需要）\ntest1 = test1 / 255.0  # 如果像素值范围是 0-255，进行归一化\n\n\n\n# 获取每张图片的预测类别索引\npredicted_classes = predictions.argmax(axis=1)  # 获取每张图片的预测类别索引\n\n# 打印每张图片的预测结果\nfor i, predicted_class in enumerate(predicted_classes):\n    print(f\"Image {i+1}: Predicted Class = {class_names[predicted_class]}\")\n\n# 可视化这 4 张图片及其预测结果\nimport matplotlib.pyplot as plt\n\nfor i in range(len(selected_images)):\n    plt.imshow(selected_images[i])  # 显示图片\n    plt.title(f\"Predicted: {class_names[predicted_classes[i]]}\")  # 显示预测类别\n    plt.axis('off')\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T08:03:37.093594Z","iopub.execute_input":"2024-12-03T08:03:37.093872Z","iopub.status.idle":"2024-12-03T08:03:59.518055Z","shell.execute_reply.started":"2024-12-03T08:03:37.093845Z","shell.execute_reply":"2024-12-03T08:03:59.516853Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 假设类别索引与名称的映射\nclass_names = ['CE', 'LAA']  # 0 对应 'CE'，1 对应 'LAA'\n\n# 加载保存的模型\nfrom tensorflow.keras.models import load_model\nmodel = load_model('efficientnet_b6_best_model.keras')\nprint(\"Model loaded successfully!\")\n\n# 确保 test1 数据已经准备好并归一化（如果需要）\ntest1 = test1 / 255.0  # 如果像素值范围是 0-255，进行归一化\n\n# 获取每张图片的预测结果\npredictions = model.predict(test1)  # Perform prediction\npredicted_classes = predictions.argmax(axis=1)  # 获取每张图片的预测类别索引\n\n# 打印每张图片的预测结果\nfor i, predicted_class in enumerate(predicted_classes):\n    print(f\"Image {i + 1}: Predicted Class = {class_names[predicted_class]}\")\n\n# 可视化这 4 张图片及其预测结果\nimport matplotlib.pyplot as plt\n\n# 选择要可视化的图片，假设 selected_images 是你之前定义的变量\nfor i in range(len(selected_images)):\n    plt.imshow(selected_images[i])  # 显示图片\n    plt.title(f\"Predicted: {class_names[predicted_classes[i]]}\")  # 显示预测类别\n    plt.axis('off')  # 不显示坐标轴\n    plt.show()  # 显示图片\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T08:27:41.098602Z","iopub.execute_input":"2024-12-03T08:27:41.098966Z","iopub.status.idle":"2024-12-03T08:28:13.828834Z","shell.execute_reply.started":"2024-12-03T08:27:41.098933Z","shell.execute_reply":"2024-12-03T08:28:13.827855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(type(Y_train))  # 检查数据类型\nprint(Y_train.shape)  # 检查形状\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-03T08:26:43.452424Z","iopub.execute_input":"2024-12-03T08:26:43.452891Z","iopub.status.idle":"2024-12-03T08:26:43.458318Z","shell.execute_reply.started":"2024-12-03T08:26:43.452853Z","shell.execute_reply":"2024-12-03T08:26:43.457379Z"}},"outputs":[],"execution_count":null}]}