{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11947718,"sourceType":"datasetVersion","datasetId":7511213},{"sourceId":11955387,"sourceType":"datasetVersion","datasetId":7516624}],"dockerImageVersionId":31040,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport torchaudio\nimport os\nimport pandas as pd\nfrom tqdm import tqdm\nimport torch\nfrom pathlib import Path\nimport gc\nimport psutil\n\nclass CFG:\n    NUM_CLASSES = 206\n    SAMPLE_RATE = 32000\n    WINDOW_SIZE = 5  # 秒\n    WINDOW_SAMPLES = SAMPLE_RATE * WINDOW_SIZE\n    DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\n    WEIGHTS_PATH = \"/kaggle/input/pthfewe/best_model_fold0_auc0.9103_epoch14.pth\"  # TODO: 替换为你的权重路径\n    test_soundscapes = \"/kaggle/input/birdclef-2025/test_soundscapes\"  # 路径变量与主流notebook一致\n    sample_submission_csv = '/kaggle/input/birdclef-2025/sample_submission.csv'  # 官方样例\n    submission_csv = '/kaggle/working/submission.csv'  # 你要生成的提交文件\n    taxonomy_csv = '/kaggle/input/birdclef-2025/taxonomy.csv'\n    model_path = '/kaggle/input/efficnetpth/efficientnet_v2_s-dd5fe13b.pth'  # 主模型权重目录\n    effnet_path = '/kaggle/input/efficnetpth/efficientnet_v2_s-dd5fe13b.pth'  # EfficientNet权重文件路径\n    FS = 32000\n    N_FFT = 1024\n    HOP_LENGTH = 320\n    N_MELS = 24\n    FMIN = 50\n    FMAX = 10000\n    TARGET_SHAPE = (256, 256)\n    model_name = 'tf_efficientnetv2_s.in21k'\n    in_channels = 3\n    batch_size = 1\n    use_tta = False\n    tta_count = 1\n    threshold = 0.5\n    use_specific_folds = False\n    folds = [0]\n    debug = False\n    debug_count = 3\n\ncfg = CFG()\n\ndef print_mem():\n    process = psutil.Process(os.getpid())\n    mem = process.memory_info().rss / 1024 / 1024\n    print(f\"[内存占用] {mem:.2f} MB\")\n\n# ===== 1. Cnn14实现 =====\nclass ConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv = nn.Conv2d(\n            in_channels=in_channels,\n            out_channels=out_channels,\n            kernel_size=(3, 3),\n            stride=(1, 1),\n            padding=(1, 1),\n            bias=False\n        )\n        self.bn = nn.BatchNorm2d(out_channels)\n        self.activation = nn.ReLU()\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.bn(x)\n        x = self.activation(x)\n        return x\n\nclass Cnn14(nn.Module):\n    def __init__(self, num_classes=527, in_channels=1):\n        super().__init__()\n        self.conv1 = ConvBlock(in_channels, 64)\n        self.conv2 = ConvBlock(64, 128)\n        self.conv3 = ConvBlock(128, 256)\n        self.conv4 = ConvBlock(256, 512)\n        self.pool = nn.AvgPool2d(kernel_size=(2, 2), stride=(2, 2))\n        self.fc = nn.Linear(512, 2048)\n        self.classifier = nn.Linear(2048, num_classes)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.pool(x)\n        x = self.conv2(x)\n        x = self.pool(x)\n        x = self.conv3(x)\n        x = self.pool(x)\n        x = self.conv4(x)\n        x = self.pool(x)\n        x = torch.mean(x, dim=(2, 3))\n        embedding = self.fc(x)\n        return {'embedding': embedding}\n\n# ===== 2. 音频预处理模块 =====\nclass ComplexAbs(nn.Module):\n    def forward(self, x):\n        return torch.abs(x)\n\nclass AudioProcessor(nn.Module):\n    def __init__(self, sample_rate=32000):\n        super().__init__()\n        self.sample_rate = sample_rate\n        self.stft = torchaudio.transforms.Spectrogram(\n            n_fft=1024,\n            hop_length=320,\n            power=None\n        )\n        self.to_complex_abs = ComplexAbs()\n        self.mel_scale = torchaudio.transforms.MelScale(\n            n_mels=24,\n            sample_rate=sample_rate,\n            f_min=50,\n            f_max=10000,\n            n_stft=513\n        )\n        self.log_compress = torchaudio.transforms.AmplitudeToDB()\n\n    def forward(self, waveform):\n        # waveform: (batch, 160000)  # 5秒音频，采样率32000\n        x = self.stft(waveform)\n        x = self.to_complex_abs(x)\n        x = self.mel_scale(x)\n        x = self.log_compress(x)\n        x = x.unsqueeze(1)  # (batch, 1, mel_bins, time)\n        return x  # (batch, 1, 64, time)\n\n# ===== 3. 音频特征提取模块 =====\nclass AudioFeatureExtractor(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.processor = AudioProcessor()\n        self.cnn14 = Cnn14()\n\n    def forward(self, waveform):\n        x = self.processor(waveform)  # (batch, 1, 64, time)\n        output = self.cnn14(x)\n        return output['embedding']  # (batch, 2048)\n\n# ===== 4. 深度可分离投影 =====\nclass DepthwiseProjection(nn.Module):\n    def __init__(self, in_dim=2048, out_dim=1280):\n        super().__init__()\n        self.proj = nn.Sequential(\n            nn.Conv1d(in_dim, in_dim, kernel_size=3, padding=1, groups=in_dim),\n            nn.GELU(),\n            nn.Conv1d(in_dim, out_dim, kernel_size=1)\n        )\n\n    def forward(self, x):\n        return self.proj(x.unsqueeze(-1)).squeeze(-1)\n\n# ===== 5. FPN视觉特征提取 =====\nclass FeaturePyramid(nn.Module):\n    def __init__(self, backbone):\n        super().__init__()\n        self.backbone = backbone\n        self.stage_indices = [2, 3, 4]\n        self.channels = sum([backbone.feature_info[i]['num_chs'] for i in self.stage_indices])\n        self.fuse_conv = nn.Sequential(\n            nn.Conv2d(self.channels, 1280, 1),\n            nn.BatchNorm2d(1280),\n            nn.GELU()\n        )\n\n    def forward(self, x):\n        feats = self.backbone(x)\n        selected = [feats[i] for i in self.stage_indices]\n        target_size = selected[-1].shape[2:]\n        resized = [\n            F.adaptive_avg_pool2d(f, target_size)\n            for f in selected[:-1]\n        ] + [selected[-1]]\n        fused = self.fuse_conv(torch.cat(resized, dim=1))\n        return fused.mean(dim=[2, 3])\n\n# ===== 6. 双向融合模块 =====\nclass BidirectionalFusion(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.vis_gate = nn.Sequential(nn.Linear(dim * 2, dim), nn.Sigmoid())\n        self.aud_gate = nn.Sequential(nn.Linear(dim * 2, dim), nn.Sigmoid())\n\n    def forward(self, vis, aud):\n        vis_gate = self.vis_gate(torch.cat([vis, aud], dim=1))\n        vis_out = vis * vis_gate + aud * (1 - vis_gate)\n        aud_gate = self.aud_gate(torch.cat([aud, vis], dim=1))\n        aud_out = aud * aud_gate + vis * (1 - aud_gate)\n        return torch.cat([vis_out, aud_out], dim=1)\n\n# ===== 7. 随机模态丢弃 =====\nclass ModalityDropout(nn.Module):\n    def __init__(self, p=0.2):\n        super().__init__()\n        self.p = p\n\n    def forward(self, vis, aud):\n        if self.training:\n            vis = torch.where(torch.rand_like(vis) < self.p, 0., vis)\n            aud = torch.where(torch.rand_like(aud) < self.p, 0., aud)\n        return vis, aud\n\n# ===== 8. 完整模型结构 =====\nclass EnhancedBirdNet(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        # pretrained=False，加载本地权重\n        self.vis_backbone = timm.create_model('tf_efficientnetv2_s.in21k', features_only=True, pretrained=False)\n        \n        # 加载本地 EfficientNet 权重\n        effnet_path = cfg.effnet_path\n        try:\n            # 方法1: 使用weights_only=False\n            print(\"尝试使用 weights_only=False 加载EfficientNet权重...\")\n            effnet_weights = torch.load(effnet_path, map_location='cpu', weights_only=False)\n        except Exception as e1:\n            print(f\"EfficientNet加载方法1失败: {str(e1)}\")\n            try:\n                # 方法2: 使用safe_globals上下文管理器\n                print(\"尝试使用 safe_globals 上下文管理器加载EfficientNet权重...\")\n                import torch.serialization\n                with torch.serialization.safe_globals(['numpy', 'numpy._core.multiarray.scalar']):\n                    effnet_weights = torch.load(effnet_path, map_location='cpu')\n            except Exception as e2:\n                print(f\"EfficientNet加载方法2失败: {str(e2)}\")\n                # 方法3: 使用pickle直接加载\n                print(\"尝试使用 pickle 直接加载EfficientNet权重...\")\n                import pickle\n                with open(effnet_path, 'rb') as f:\n                    effnet_weights = pickle.load(f)\n        \n        if 'model' in effnet_weights:\n            effnet_state = effnet_weights['model']\n        elif 'state_dict' in effnet_weights:\n            effnet_state = effnet_weights['state_dict']\n        else:\n            effnet_state = effnet_weights\n        load_result = self.vis_backbone.load_state_dict(effnet_state, strict=False)\n        print('EfficientNet 权重加载结果:')\n        print('  missing_keys:', load_result.missing_keys)\n        print('  unexpected_keys:', load_result.unexpected_keys)\n        self.pyramid = FeaturePyramid(self.vis_backbone)\n        self.audio_extractor = AudioFeatureExtractor()\n        self.audio_proj = DepthwiseProjection()\n        self.modality_drop = ModalityDropout(p=0.2)\n        self.bi_fusion = BidirectionalFusion(1280)\n        self.classifier = nn.Sequential(\n            nn.Linear(2560, 1024),\n            nn.BatchNorm1d(1024),\n            nn.LeakyReLU(0.1),\n            nn.Dropout(0.3),\n            nn.Linear(1024, num_classes)\n        )\n\n    def forward(self, img, audio_waveform, use_modality_dropout=True):\n        # img: (batch, 3, 64, 256) 或 (batch, 3, 256, 256)  # 与训练时一致\n        # audio_waveform: (batch, 160000)\n        vis = self.pyramid(img)\n        aud = self.audio_proj(self.audio_extractor(audio_waveform))\n        if use_modality_dropout:\n            vis, aud = self.modality_drop(vis, aud)\n        fused = self.bi_fusion(vis, aud)\n        return self.classifier(fused)\n\ndef main():\n    import torch\n    import os\n    print(f\"使用设备: {cfg.DEVICE}\")\n    print(f\"模型权重目录: {cfg.model_path}\")\n    print(f\"测试音频目录: {cfg.test_soundscapes}\")\n    print_mem()\n    # 读取 sample_submission.csv 获取列名，只读取一次\n    sample = pd.read_csv(cfg.sample_submission_csv)\n    class_names = list(sample.columns)[1:]  # 除去 row_id\n    print(f\"类别数量: {len(class_names)}\")\n    print(f\"实际查找的测试集路径: {cfg.test_soundscapes}\")\n    if os.path.exists(cfg.test_soundscapes):\n        files_in_dir = os.listdir(cfg.test_soundscapes)\n        print(f\"该目录下文件: {files_in_dir}\")\n    else:\n        print(\"测试集路径不存在！\")\n    test_files = list(Path(cfg.test_soundscapes).glob('*.ogg'))\n    print(f\"找到 {len(test_files)} 个测试音频文件: {[str(f) for f in test_files]}\")\n    if len(test_files) == 0:\n        print(f\"[错误] 没有找到测试音频文件! 请检查测试集路径: {cfg.test_soundscapes}\")\n        print(f\"当前目录下文件: {os.listdir(cfg.test_soundscapes) if os.path.exists(cfg.test_soundscapes) else '目录不存在'}\")\n        # raise FileNotFoundError(f\"未找到任何.ogg测试音频文件, 请检查路径和Kaggle数据集挂载!\")\n    # 加载模型\n    try:\n        model = EnhancedBirdNet(num_classes=cfg.NUM_CLASSES)\n        try:\n            print(\"尝试使用 weights_only=False 加载模型...\")\n            state = torch.load(cfg.WEIGHTS_PATH, map_location=cfg.DEVICE, weights_only=False)\n        except Exception as e1:\n            print(f\"方法1失败: {str(e1)}\")\n            try:\n                print(\"尝试使用 safe_globals 上下文管理器加载模型...\")\n                import torch.serialization\n                with torch.serialization.safe_globals(['numpy', 'numpy._core.multiarray.scalar']):\n                    state = torch.load(cfg.WEIGHTS_PATH, map_location=cfg.DEVICE)\n            except Exception as e2:\n                print(f\"方法2失败: {str(e2)}\")\n                print(\"尝试使用 pickle 直接加载模型...\")\n                import pickle\n                with open(cfg.WEIGHTS_PATH, 'rb') as f:\n                    state = pickle.load(f)\n        if 'model_state_dict' in state:\n            load_result = model.load_state_dict(state['model_state_dict'], strict=False)\n            print('主模型权重加载结果:')\n            print('  missing_keys:', load_result.missing_keys)\n            print('  unexpected_keys:', load_result.unexpected_keys)\n        elif 'state_dict' in state:\n            load_result = model.load_state_dict(state['state_dict'], strict=False)\n            print('主模型权重加载结果:')\n            print('  missing_keys:', load_result.missing_keys)\n            print('  unexpected_keys:', load_result.unexpected_keys)\n        else:\n            load_result = model.load_state_dict(state, strict=False)\n            print('主模型权重加载结果:')\n            print('  missing_keys:', load_result.missing_keys)\n            print('  unexpected_keys:', load_result.unexpected_keys)\n        model.to(cfg.DEVICE)\n        model.eval()\n        print(\"模型加载成功\")\n    except Exception as e:\n        print(f\"模型加载失败: {str(e)}\")\n        sample.to_csv(cfg.submission_csv, index=False)\n        print(f\"已创建空的提交文件: {cfg.submission_csv}\")\n        return\n    results = []\n    for file_path in tqdm(test_files, desc='Processing files'):\n        try:\n            fname = os.path.basename(file_path)\n            waveform, sr = torchaudio.load(file_path)\n            waveform = waveform.mean(0)\n            if sr != cfg.SAMPLE_RATE:\n                waveform = torchaudio.functional.resample(waveform, sr, cfg.SAMPLE_RATE)\n            total_samples = waveform.shape[0]\n            num_windows = (total_samples - 1) // cfg.WINDOW_SAMPLES + 1\n            for i in range(num_windows):\n                start = i * cfg.WINDOW_SAMPLES\n                end = start + cfg.WINDOW_SAMPLES\n                window = waveform[start:end]\n                if window.shape[0] < cfg.WINDOW_SAMPLES:\n                    window = torch.nn.functional.pad(window, (0, cfg.WINDOW_SAMPLES - window.shape[0]))\n                window = window.unsqueeze(0).to(cfg.DEVICE)\n                mel_transform = torchaudio.transforms.MelSpectrogram(\n                    sample_rate=cfg.SAMPLE_RATE,\n                    n_fft=1024,\n                    hop_length=320,\n                    n_mels=24,\n                    f_min=50,\n                    f_max=10000\n                ).to(cfg.DEVICE)\n                db_transform = torchaudio.transforms.AmplitudeToDB().to(cfg.DEVICE)\n                mel = mel_transform(window)\n                mel_db = db_transform(mel)\n                mel_db = (mel_db - mel_db.min()) / (mel_db.max() - mel_db.min() + 1e-8)\n                spec = torch.stack([mel_db] * 3, dim=1)\n                spec = torch.nn.functional.interpolate(\n                    spec, size=(256, 256), mode='bilinear', align_corners=False\n                )\n                with torch.no_grad():\n                    logits = model(spec, window, use_modality_dropout=False)\n                    probs = torch.sigmoid(logits).cpu().numpy()[0]\n                end_time = (i + 1) * 5\n                row_id = f'{fname}_{end_time}'\n                row = {'row_id': row_id}\n                for idx, class_name in enumerate(class_names):\n                    row[class_name] = probs[idx]\n                results.append(row)\n                del window, mel, mel_db, spec, logits, probs\n                torch.cuda.empty_cache()\n            del waveform\n            gc.collect()\n            print_mem()\n        except Exception as e:\n            print(f\"处理文件 {file_path} 时出错: {str(e)}\")\n            gc.collect()\n            print_mem()\n    submission = pd.DataFrame(results)\n    submission = submission.reindex(columns=['row_id'] + class_names, fill_value=0.0)\n    submission.to_csv('/kaggle/working/submission.csv', index=False)\n    print(f\"已保存提交文件: /kaggle/working/submission.csv\")\n    print_mem()\n    import os\n    if not os.path.exists('/kaggle/working/submission.csv'):\n        print(\"未检测到 submission.csv，兜底生成一个空的提交文件。\")\n        sample = pd.read_csv(cfg.sample_submission_csv)\n        sample.to_csv('/kaggle/working/submission.csv', index=False)\n        print(f\"兜底生成空的提交文件: /kaggle/working/submission.csv\")\n    else:\n        print(\"submission.csv 已正常生成。\")\n\n# 直接调用 main()\nmain()\n\n        ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-02T07:55:36.262827Z","iopub.execute_input":"2025-06-02T07:55:36.263128Z","iopub.status.idle":"2025-06-02T07:55:57.232210Z","shell.execute_reply.started":"2025-06-02T07:55:36.263104Z","shell.execute_reply":"2025-06-02T07:55:57.231225Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}