{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":24800,"databundleVersionId":1831594,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport warnings\nimport cv2\nfrom tqdm import tqdm\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import MultiLabelBinarizer\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models, optimizers, callbacks, applications, regularizers\nfrom sklearn.metrics import precision_recall_fscore_support, roc_curve, auc\nimport matplotlib.patches as mpatches\n\nwarnings.filterwarnings('ignore')\n\n# 绘图基础样式\nplt.rcParams['font.sans-serif'] = ['DejaVu Sans', 'Arial', 'Helvetica', 'sans-serif']  \nplt.rcParams['axes.unicode_minus'] = False\nplt.rcParams['figure.dpi'] = 120\ntry:\n    plt.style.use('seaborn-v0_8-whitegrid') \nexcept:\n    pass\n\n# ==================== 1. 数据路径与读取 ====================\nDATA_DIR = \"/kaggle/input/vinbigdata-chest-xray-abnormalities-detection\"                                             \nTRAIN_DIR = os.path.join(DATA_DIR, \"train\")\nTRAIN_CSV = os.path.join(DATA_DIR, \"train.csv\")\ntrain_df = pd.read_csv(TRAIN_CSV)\n\nprint(\"\\n 原始数据全量分布\")\nplt.figure(figsize=(15, 6))\ndist1 = train_df['class_name'].value_counts()\nsns.barplot(x=dist1.index, y=dist1.values, palette='viridis')\nplt.title('Original Dataset: Class Distribution', fontsize=14, fontweight='bold')\nplt.xticks(rotation=45, ha='right')\nfor i, v in enumerate(dist1.values):\n    plt.text(i, v + (max(dist1.values)*0.01), str(v), ha='center', fontsize=9)\nplt.tight_layout(); plt.show()\n\n# ==================== 2. 数据筛选 (推荐 1000+1000) ====================\nNUM_BALANCED = 2000\nNUM_RANDOM = 2000  \n\nprint(f\"\\n筛选策略：均衡 {NUM_BALANCED} 张 + 随机 {NUM_RANDOM} 张...\")\nnp.random.seed(42)\nselected_ids = set()\n\ndiseased_only = train_df[train_df['class_name'] != 'No finding'].copy()\nclass_counts = diseased_only['class_name'].value_counts().sort_values()\nper_class_quota = NUM_BALANCED // len(class_counts)\n\nfor cls in class_counts.index:\n    cls_ids = diseased_only[diseased_only['class_name'] == cls]['image_id'].unique()\n    pick = np.random.choice(cls_ids, min(len(cls_ids), per_class_quota), replace=False)\n    for idx in pick:\n        if len(selected_ids) < NUM_BALANCED: selected_ids.add(idx)\n\nremaining_pool = list(set(train_df['image_id'].unique()) - selected_ids)\nif NUM_RANDOM > 0 and len(remaining_pool) > 0:\n    random_picks = np.random.choice(remaining_pool, min(NUM_RANDOM, len(remaining_pool)), replace=False)\n    selected_ids.update(random_picks)\n\nselected_image_ids = list(selected_ids)\nselected_df = train_df[train_df['image_id'].isin(selected_image_ids)].copy()\n\nprint(\"\\n筛选后数据分布)\")\ndist2_df = selected_df.drop_duplicates(['image_id', 'class_name'])\ndist2 = dist2_df['class_name'].value_counts()\nplt.figure(figsize=(15, 6))\nsns.barplot(x=dist2.index, y=dist2.values, palette='magma')\nplt.title(f'Selected Dataset (Images: {len(selected_image_ids)})', fontsize=14, fontweight='bold')\nplt.xticks(rotation=45, ha='right')\nfor i, v in enumerate(dist2.values):\n    plt.text(i, v + (max(dist2.values)*0.01), str(v), ha='center', fontsize=9)\nplt.tight_layout(); plt.show()\n\n# ==================== 3. 固定颜色映射与CLAHE增强 ====================\ndisease_classes = sorted([c for c in train_df['class_name'].unique() if c != 'No finding'])\ncolor_palette = [\n    (255, 0, 0), (0, 255, 0), (0, 0, 255), (255, 255, 0), (255, 0, 255), (0, 255, 255),\n    (255, 128, 0), (128, 0, 255), (0, 128, 128), (128, 128, 0), (128, 0, 0), (0, 0, 128),\n    (0, 128, 0), (255, 192, 203)\n]\nGLOBAL_COLOR_MAP = {cls: color_palette[i % len(color_palette)] for i, cls in enumerate(disease_classes)}\n\nprint(\"\\n疾病颜色图例\")\nplt.figure(figsize=(12, 3))\nfor i, (cls, color) in enumerate(GLOBAL_COLOR_MAP.items()):\n    plt.bar(i, 1, color=[c/255.0 for c in color], label=cls)\nplt.xticks(range(len(GLOBAL_COLOR_MAP)), GLOBAL_COLOR_MAP.keys(), rotation=45, ha='right')\nplt.yticks([]); plt.title(\"Legend: Color Mapping\"); plt.tight_layout(); plt.show()\n\ndef dicom_to_array(path):\n    try:\n        dicom = pydicom.dcmread(path); data = apply_voi_lut(dicom.pixel_array, dicom)\n        if dicom.PhotometricInterpretation == \"MONOCHROME1\": data = np.amax(data) - data\n        data = data - np.min(data)\n        if np.max(data) > 0: data = (data / np.max(data) * 255).astype(np.uint8)\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n        return clahe.apply(data)\n    except: return None\n\ndef draw_bboxes_on_large_image(img_array, bboxes):\n    img_display = cv2.cvtColor(img_array, cv2.COLOR_GRAY2RGB) if len(img_array.shape) == 2 else img_array.copy()\n    for bbox in bboxes:\n        cls = bbox['class_name']; color = GLOBAL_COLOR_MAP.get(cls, (255, 255, 255))\n        x1, y1, x2, y2 = int(bbox['x_min']), int(bbox['y_min']), int(bbox['x_max']), int(bbox['y_max'])\n        cv2.rectangle(img_display, (x1, y1), (x2, y2), color, 8)\n        (w, h), _ = cv2.getTextSize(cls, cv2.FONT_HERSHEY_SIMPLEX, 1.5, 3)\n        cv2.rectangle(img_display, (x1, y1 - h - 20), (x1 + w, y1), color, -1)\n        cv2.putText(img_display, cls, (x1, y1 - 10), cv2.FONT_HERSHEY_SIMPLEX, 1.5, (255, 255, 255), 3)\n    return img_display\n\n# ==================== 4. 加载数据 ====================\nimage_to_labels, image_to_bboxes = {}, {}\nfor img_id in selected_image_ids:\n    img_data = selected_df[selected_df['image_id'] == img_id]\n    lbls = img_data[img_data['class_name'] != 'No finding']['class_name'].unique()\n    image_to_labels[img_id] = list(lbls) if len(lbls) > 0 else ['No finding']\n    boxes = []\n    for _, row in img_data.iterrows():\n        if row['class_name'] != 'No finding' and not pd.isna(row['x_min']):\n            boxes.append({'class_name': row['class_name'], 'x_min': row['x_min'], 'y_min': row['y_min'], 'x_max': row['x_max'], 'y_max': row['y_max']})\n    image_to_bboxes[img_id] = boxes\n\nmlb = MultiLabelBinarizer(); y_encoded = mlb.fit_transform(list(image_to_labels.values())); label_names = mlb.classes_\n\nX_list, y_list, orig_list, final_ids = [], [], [], []\nfor img_id in tqdm(selected_image_ids, desc=\"Processing Images\"):\n    path = os.path.join(TRAIN_DIR, f\"{img_id}.dicom\"); arr = dicom_to_array(path)\n    if arr is not None:\n        X_list.append(cv2.resize(np.stack([arr]*3, axis=-1), (224, 224)).astype('float32')/255.0)\n        y_list.append(y_encoded[len(final_ids)]); orig_list.append(arr); final_ids.append(img_id)\n\nX_all = np.array(X_list); y_all = np.array(y_list)\n\n# 训练样本展示\nplt.figure(figsize=(20, 10))\nfor i in range(min(8, len(final_ids))):\n    plt.subplot(2, 4, i+1); img_id = final_ids[i]; bboxes = image_to_bboxes[img_id]\n    status = f\"Findings: {', '.join(list(set([b['class_name'] for b in bboxes])))[:15]}\" if bboxes else \"Normal\"\n    res = draw_bboxes_on_large_image(orig_list[i], bboxes); plt.imshow(res); plt.title(f\"ID: {img_id[:6]}\\n{status}\", fontsize=9); plt.axis('off')\nplt.tight_layout(); plt.show()\n\n# ==================== 5. 模型训练 (增强泛化版) ====================\ndef build_model(num_classes):\n    base = applications.EfficientNetB0(weights=None, include_top=False, input_shape=(224, 224, 3))\n    x = layers.GlobalAveragePooling2D()(base.output)\n    # 【增加 Dropout 到 0.6 + L2 正则化】应对过拟合\n    x = layers.Dropout(0.6)(x)\n    x = layers.Dense(256, activation='relu', kernel_regularizer=regularizers.l2(1e-4))(x)\n    x = layers.BatchNormalization()(x)\n    outputs = layers.Dense(num_classes, activation='sigmoid')(x)\n    return models.Model(inputs=base.input, outputs=outputs)\n\nmodel = build_model(len(label_names))\n# 【优化】：使用 Label Smoothing 损失\nmodel.compile(optimizer=optimizers.Adam(0.001), \n              loss=tf.keras.losses.BinaryCrossentropy(label_smoothing=0.1), \n              metrics=['accuracy', tf.keras.metrics.AUC(name='auc', multi_label=True)])\n\ntrain_callbacks = [\n    callbacks.EarlyStopping(monitor='val_auc', patience=10, mode='max', restore_best_weights=True),\n    callbacks.ReduceLROnPlateau(monitor='val_auc', factor=0.2, patience=5, min_lr=1e-6, mode='max', verbose=1)\n]\n\nX_train, X_val, y_train, y_val, _, orig_val, _, ids_val = train_test_split(X_all, y_all, orig_list, final_ids, test_size=0.2, random_state=42)\nhistory = model.fit(X_train, y_train, validation_data=(X_val, y_val), epochs=50, batch_size=16, callbacks=train_callbacks)\n\n# 性能曲线\nfig, axes = plt.subplots(1, 3, figsize=(20, 5))\naxes[0].plot(history.history['loss'], label='Train'); axes[0].plot(history.history['val_loss'], label='Val'); axes[0].set_title('Loss'); axes[0].legend()\naxes[1].plot(history.history['accuracy'], label='Train'); axes[1].plot(history.history['val_accuracy'], label='Val'); axes[1].set_title('Accuracy'); axes[1].legend()\naxes[2].plot(history.history['auc'], label='Train'); axes[2].plot(history.history['val_auc'], label='Val'); axes[2].set_title('AUC'); axes[2].legend()\nplt.tight_layout(); plt.show()\n\n# ==================== 6. 定向诊断分析 (1正2反 ) ====================\ny_pred_proba = model.predict(X_val)\nnf_idx = list(label_names).index('No finding')\nnorm_idx = [i for i, y in enumerate(y_val) if y[nf_idx] == 1 and np.sum(y) == 1]\ndiseased_idx = [i for i, y in enumerate(y_val) if np.sum(np.delete(y, nf_idx)) > 0]\n\ndef deep_analyze_sample(idx, case_label):\n    img_id = ids_val[idx]; pred_p = y_pred_proba[idx]; true_y = y_val[idx]\n    true_lbls = [label_names[j] for j, v in enumerate(true_y) if v == 1 and label_names[j] != 'No finding']\n    pred_lbls = [label_names[j] for j, p in enumerate(pred_p) if p > 0.5 and label_names[j] != 'No finding']\n    \n    fig, axes = plt.subplots(1, 3, figsize=(26, 10), gridspec_kw={'width_ratios': [1.2, 1, 0.8]})\n    res = draw_bboxes_on_large_image(orig_val[idx], image_to_bboxes.get(img_id, []))\n    axes[0].imshow(res); axes[0].axis('off'); axes[0].set_title(f\"[{case_label}] ID: {img_id[:10]}\", fontsize=16, fontweight='bold')\n    \n    probs_show = sorted([(label_names[i], pred_p[i]) for i in range(len(label_names)) if label_names[i] != 'No finding'], key=lambda x: x[1], reverse=True)[:8]\n    axes[1].barh([x[0] for x in probs_show][::-1], [x[1] for x in probs_show][::-1], color=['#ff4d4d' if v[1]>0.5 else '#7fb3d5' for v in probs_show][::-1])\n    axes[1].axvline(x=0.5, color='red', linestyle='--'); axes[1].set_xlim(0, 1.1)\n    \n    axes[2].axis('off')\n    report = f\"DIAGNOSTIC REPORT\\n\" + \"=\"*25 + \"\\n[Ground Truth]:\\n\"\n    report += (\"\\n\".join([f\" - {l}\" for l in true_lbls]) if true_lbls else \" - Normal Case\") + \"\\n\\n[Model Prediction]:\\n\"\n    report += (\"\\n\".join([f\" - {l}\" for l in pred_lbls]) if pred_lbls else \" - Normal Case\") + \"\\n\\n[Verdict]: \"\n    report += \"SUCCESS\" if (set(true_lbls) == set(pred_lbls)) else \"REVIEW REQUIRED\"\n    axes[2].text(0, 0.95, report, fontsize=18, verticalalignment='top', family='monospace', bbox=dict(boxstyle=\"round,pad=1\", facecolor=\"white\"))\n    plt.tight_layout(); plt.show()\n\nif len(norm_idx) >= 1: deep_analyze_sample(norm_idx[0], \"NORMAL CASE\")\nif len(diseased_idx) >= 2: deep_analyze_sample(diseased_idx[0], \"ABNORMAL 1\"); deep_analyze_sample(diseased_idx[1], \"ABNORMAL 2\")\n\n# ==================== 图 ====================\ntry:\n    plt.figure(figsize=(15, 6)); p, r, _, _ = precision_recall_fscore_support(y_val, (y_pred_proba > 0.5).astype(int), average=None, zero_division=0)\n    valid = np.sum(y_val, axis=0) > 0; n_sub = [label_names[i] for i in range(len(label_names)) if valid[i] and label_names[i] != 'No finding']\n    p_s = [p[i] for i in range(len(label_names)) if valid[i] and label_names[i] != 'No finding']; r_s = [r[i] for i in range(len(label_names)) if valid[i] and label_names[i] != 'No finding']\n    plt.bar(np.arange(len(n_sub))-0.2, p_s, 0.4, label='Precision', color='#2c3e50'); plt.bar(np.arange(len(n_sub))+0.2, r_s, 0.4, label='Recall', color='#c0392b')\n    plt.xticks(np.arange(len(n_sub)), n_sub, rotation=35, ha='right'); plt.legend(); plt.savefig('thesis_performance_v2.png', dpi=300); plt.show()\nexcept: pass\n\nmodel.save('chest_xray_disease_model.h5')\nprint(\"\\n结束\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null}]}