{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"},{"sourceId":11569667,"sourceType":"datasetVersion","datasetId":7253605},{"sourceId":11569755,"sourceType":"datasetVersion","datasetId":7253661},{"sourceId":12038896,"sourceType":"datasetVersion","datasetId":7377931},{"sourceId":11568812,"sourceType":"datasetVersion","datasetId":7253205}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Classifier\n\nThis is just being used to classify the waveforms into their folders for further processing ([or searching](https://www.kaggle.com/code/johnnyhyland/faiss-search-baseline)). Useful if you're training separate models for each source type. There could be a dataleak here, 99% accuracy.","metadata":{}},{"cell_type":"code","source":"from types import SimpleNamespace\nimport torch\n\ncfg = SimpleNamespace()\ncfg.device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ncfg.local_rank = 0\ncfg.seed = 123\ncfg.subsample = None \n\nimport os\nimport glob\nimport torch\nfrom torch.utils.data import Dataset\nimport numpy as np\nimport torch.nn as nn\nfrom torchvision.models import efficientnet_b0, EfficientNet_B0_Weights\nfrom torch.utils.data import DataLoader\nimport torch.optim as optim\nimport pandas as pd\nfrom tqdm import tqdm\n\nclass_map = {\n    \"CurveFault_A\": 0,\n    \"CurveFault_B\": 1,\n    \"CurveVel_A\": 2,\n    \"CurveVel_B\": 3,\n    \"FlatFault_A\": 4,\n    \"FlatFault_B\": 5,\n    \"FlatVel_A\": 6,\n    \"FlatVel_B\": 7,\n    \"Style_A\": 8,\n    \"Style_B\": 9,\n}\n\nclass AlignedSeismicClassificationDataset(Dataset):\n    def __init__(self, cfg, mode=\"train\"):\n        self.cfg = cfg\n        self.mode = mode\n        self.samples = self.load_aligned_samples()\n\n    def load_aligned_samples(self):\n        df = pd.read_csv(\"/kaggle/input/openfwi-preprocessed-72x72/folds.csv\")\n        \n        if self.cfg.subsample is not None:\n            df = df.groupby([\"dataset\", \"fold\"]).head(self.cfg.subsample)\n\n        if self.mode == \"train\":\n            df = df[df[\"fold\"] != 0]\n        else:\n            df = df[df[\"fold\"] == 0]\n\n        samples = []\n        \n        for idx, row in tqdm(df.iterrows(), total=len(df), desc=f\"Loading {self.mode} samples\"):\n            dataset = row[\"dataset\"]\n            label = class_map.get(dataset)\n            if label is None:\n                continue\n\n            p1 = os.path.join(\"/kaggle/input/open-wfi-1/openfwi_float16_1/\", row[\"data_fpath\"])\n            p2 = os.path.join(\"/kaggle/input/open-wfi-1/openfwi_float16_1/\", row[\"data_fpath\"].split(\"/\")[0], \"*\", row[\"data_fpath\"].split(\"/\")[-1])\n            p3 = os.path.join(\"/kaggle/input/open-wfi-2/openfwi_float16_2/\", row[\"data_fpath\"])\n            p4 = os.path.join(\"/kaggle/input/open-wfi-2/openfwi_float16_2/\", row[\"data_fpath\"].split(\"/\")[0], \"*\", row[\"data_fpath\"].split(\"/\")[-1])\n            farr = glob.glob(p1) + glob.glob(p2) + glob.glob(p3) + glob.glob(p4)\n            \n            if farr:\n                file_path = farr[0]\n                for sample_idx in range(500):\n                    samples.append((file_path, sample_idx, label))\n\n        return samples\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        file_path, sample_idx, label = self.samples[idx]\n        data = np.load(file_path, mmap_mode='r')[sample_idx]  # shape: (sources, time, receivers)\n        data = torch.from_numpy(data).float()\n        # Take mean across sources dimension to get single channel\n        data = data.mean(dim=0, keepdim=True)  # shape: (1, time, receivers)\n        return data, label\n\nclass SeismicEfficientNetClassifier(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        self.backbone = efficientnet_b0(weights=EfficientNet_B0_Weights.IMAGENET1K_V1)\n        \n        if self.backbone.features[0][0].in_channels != 1:\n            self.backbone.features[0][0] = nn.Conv2d(\n                1, 32, kernel_size=3, stride=2, padding=1, bias=False\n            )\n        \n        in_features = self.backbone.classifier[1].in_features\n        self.backbone.classifier[1] = nn.Linear(in_features, num_classes)\n\n    def forward(self, x):\n        return self.backbone(x)\n\ndef train_aligned_classifier():\n    print(\"Creating datasets...\")\n    train_dataset = AlignedSeismicClassificationDataset(cfg, mode=\"train\")\n    val_dataset = AlignedSeismicClassificationDataset(cfg, mode=\"valid\")\n    \n    print(f\"Train samples: {len(train_dataset)}\")\n    print(f\"Val samples: {len(val_dataset)}\")\n    \n    train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=4)\n    val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=4)\n    \n    model = SeismicEfficientNetClassifier(num_classes=10).to(cfg.device)\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(model.parameters(), lr=1e-3)\n    num_epochs = 10\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=num_epochs)\n    \n    best_val_acc = 0.0\n    \n    for epoch in range(num_epochs):\n        print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n        \n        # Training phase\n        model.train()\n        running_loss = 0.0\n        correct = 0\n        total = 0\n        \n        for batch_idx, (data, labels) in enumerate(tqdm(train_loader, desc=\"Training\")):\n            data, labels = data.to(cfg.device), labels.to(cfg.device)\n            \n            optimizer.zero_grad()\n            outputs = model(data)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            \n            running_loss += loss.item() * data.size(0)\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n        \n        train_loss = running_loss / total\n        train_acc = correct / total\n        \n        # Validation phase\n        model.eval()\n        val_loss = 0.0\n        val_correct = 0\n        val_total = 0\n        \n        with torch.no_grad():\n            for data, labels in tqdm(val_loader, desc=\"Validation\"):\n                data, labels = data.to(cfg.device), labels.to(cfg.device)\n                outputs = model(data)\n                loss = criterion(outputs, labels)\n                \n                val_loss += loss.item() * data.size(0)\n                _, predicted = outputs.max(1)\n                val_total += labels.size(0)\n                val_correct += predicted.eq(labels).sum().item()\n        \n        val_loss /= val_total\n        val_acc = val_correct / val_total\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | \"\n              f\"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}\")\n        \n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            torch.save(model.state_dict(), \"aligned_classifier.pth\")\n            print(f\"New best model saved! Val Acc: {val_acc:.4f}\")\n    \n    print(f\"\\nTraining completed. Best validation accuracy: {best_val_acc:.4f}\")\n    return model\n\n# Train the model\nprint(\"Starting training...\")\nmodel = train_aligned_classifier()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-04T14:14:58.696852Z","iopub.execute_input":"2025-06-04T14:14:58.697199Z","execution_failed":"2025-06-04T14:15:35.708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}