{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":11372645,"sourceType":"datasetVersion","datasetId":7112154},{"sourceId":11388938,"sourceType":"datasetVersion","datasetId":7126650},{"sourceId":233221031,"sourceType":"kernelVersion"},{"sourceId":332421,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":278650,"modelId":299553}],"dockerImageVersionId":30919,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nimport csv \nimport argparse\nimport torch.serialization\nimport json \nimport sys ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:42.280518Z","iopub.execute_input":"2025-04-15T22:31:42.280771Z","iopub.status.idle":"2025-04-15T22:31:46.532926Z","shell.execute_reply.started":"2025-04-15T22:31:42.280750Z","shell.execute_reply":"2025-04-15T22:31:46.532230Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEBUG = False \nOID_DEBUG = '95de33eda8' # OID to process in DEBUG mode","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:46.533660Z","iopub.execute_input":"2025-04-15T22:31:46.534089Z","iopub.status.idle":"2025-04-15T22:31:46.537525Z","shell.execute_reply.started":"2025-04-15T22:31:46.534065Z","shell.execute_reply":"2025-04-15T22:31:46.536626Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sys.path.append('/kaggle/input/openfwi-pretrainedmodel/OpenFWI')\n\n\ntry:\n    from network import InversionNet, replace_legacy\n    import transforms as T\n    OPENFWI_AVAILABLE = True\n    print(\"Successfully imported InversionNet, replace_legacy, and transforms from OpenFWI.\")\nexcept ImportError:\n    print(\"Warning: OpenFWI could not be imported. InversionNet, transforms, and related functions will not be available.\")\n    InversionNet = None\n    replace_legacy = None\n    T = None\n    OPENFWI_AVAILABLE = False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:46.538422Z","iopub.execute_input":"2025-04-15T22:31:46.538726Z","iopub.status.idle":"2025-04-15T22:31:47.248374Z","shell.execute_reply.started":"2025-04-15T22:31:46.538697Z","shell.execute_reply":"2025-04-15T22:31:47.247678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATASET_CONFIG = '/kaggle/input/openfwi-pretrainedmodel/OpenFWI/dataset_config.json' # Path to dataset_config.json (adjust as needed)\nDATASET_NAME = 'flatfault-b' # Dataset name used for training the weights (adjust as needed)\nK_TRANSFORM = 1 # k value for LogTransform (from Submission.py)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.249941Z","iopub.execute_input":"2025-04-15T22:31:47.250279Z","iopub.status.idle":"2025-04-15T22:31:47.253665Z","shell.execute_reply.started":"2025-04-15T22:31:47.250259Z","shell.execute_reply":"2025-04-15T22:31:47.252817Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nMODEL_CLASSES = [\n    \"InversionNet\",\n    \"DumbNet\",\n]\n\nPATHS = [\n    \"/kaggle/input/gwi-pretrainedmodel/PretrainedModel/ffb_l2.pth\",\n    \"/kaggle/input/dumbernet/pytorch/default/1/model_84.pth\"\n    ]\n\nPATH_WEIGHTS = [\n    0.5,\n    0.5\n    ]\n\nassert len(PATHS) == len(MODEL_CLASSES), \"Number of model paths and class names must match.\"\nassert len(PATHS) == len(PATH_WEIGHTS), \"Number of model paths and weights must match.\"\n\navailable_model_classes = {\"DumbNet\"}\nif OPENFWI_AVAILABLE:\n    available_model_classes.add(\"InversionNet\")\n\nfor model_cls_name in MODEL_CLASSES:\n    if model_cls_name not in available_model_classes:\n        raise ValueError(f\"Specified model class '{model_cls_name}' is not available. Available classes: {available_model_classes}\")\n\n\ntotal_weight = sum(PATH_WEIGHTS)\nif total_weight > 0:\n    PATH_WEIGHTS = [w / total_weight for w in PATH_WEIGHTS]\nelse:\n    # Assign equal weights if sum is zero (error prevention)\n    print(\"Warning: Sum of PATH_WEIGHTS is zero. Assigning equal weights.\")\n    PATH_WEIGHTS = [1.0 / len(PATHS)] * len(PATHS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.255139Z","iopub.execute_input":"2025-04-15T22:31:47.255463Z","iopub.status.idle":"2025-04-15T22:31:47.270229Z","shell.execute_reply.started":"2025-04-15T22:31:47.255413Z","shell.execute_reply":"2025-04-15T22:31:47.269493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_dataset_config(config_path, dataset_name):\n    \"\"\"Loads normalization parameters from dataset_config.json.\"\"\"\n    if not Path(config_path).exists():\n         print(f\"Error: Dataset config file not found at {config_path}. Required for InversionNet preprocessing.\")\n         # Exit if InversionNet is selected but config is missing\n         if \"InversionNet\" in MODEL_CLASSES:\n             sys.exit(f\"Exiting: InversionNet requires {config_path}, but it was not found.\")\n         else:\n             print(\"Continuing without dataset config as InversionNet is not selected.\")\n             return None # OK to continue if only using models that don't need config\n\n    try:\n        with open(config_path) as f:\n            ctx = json.load(f)[dataset_name]\n        print(f\"Loaded config for dataset: {dataset_name}\")\n        # Check for required keys\n        required_keys = ['data_min', 'data_max', 'label_min', 'label_max']\n        if not all(key in ctx for key in required_keys):\n             raise KeyError(f\"Dataset config for '{dataset_name}' is missing one or more required keys: {required_keys}\")\n        return ctx\n    except FileNotFoundError: # Should be caught by Path().exists(), but just in case\n        print(f\"Error: {config_path} not found.\")\n        if \"InversionNet\" in MODEL_CLASSES:\n             sys.exit(1)\n        return None\n    except KeyError as e:\n        print(f\"Error: Dataset '{dataset_name}' not found in {config_path} or missing keys: {e}.\")\n        if \"InversionNet\" in MODEL_CLASSES:\n             sys.exit(1)\n        return None\n    except json.JSONDecodeError:\n        print(f\"Error: Failed to decode JSON from {config_path}.\")\n        if \"InversionNet\" in MODEL_CLASSES:\n             sys.exit(1)\n        return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.270993Z","iopub.execute_input":"2025-04-15T22:31:47.271247Z","iopub.status.idle":"2025-04-15T22:31:47.283910Z","shell.execute_reply.started":"2025-04-15T22:31:47.271226Z","shell.execute_reply":"2025-04-15T22:31:47.283199Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Find files to load and create Dataset","metadata":{}},{"cell_type":"code","source":"all_inputs = [\n    f\n    for f in\n    Path('/kaggle/input/waveform-inversion/train_samples').rglob('*.npy')\n    if ('seis' in f.stem) or ('data' in f.stem)\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.284636Z","iopub.execute_input":"2025-04-15T22:31:47.284937Z","iopub.status.idle":"2025-04-15T22:31:47.384394Z","shell.execute_reply.started":"2025-04-15T22:31:47.284908Z","shell.execute_reply":"2025-04-15T22:31:47.383852Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def inputs_files_to_output_files(input_files):\n    return [\n        Path(str(f).replace('seis', 'vel').replace('data', 'model'))\n        for f in input_files\n    ]\n\nall_outputs = inputs_files_to_output_files(all_inputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.385113Z","iopub.execute_input":"2025-04-15T22:31:47.385393Z","iopub.status.idle":"2025-04-15T22:31:47.389385Z","shell.execute_reply.started":"2025-04-15T22:31:47.385366Z","shell.execute_reply":"2025-04-15T22:31:47.388619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"assert all(f.exists() for f in all_outputs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.390110Z","iopub.execute_input":"2025-04-15T22:31:47.390304Z","iopub.status.idle":"2025-04-15T22:31:47.406273Z","shell.execute_reply.started":"2025-04-15T22:31:47.390287Z","shell.execute_reply":"2025-04-15T22:31:47.405646Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"existing_inputs = [f for f, o in zip(all_inputs, all_outputs) if o.exists()]\nexisting_outputs = [o for f, o in zip(all_inputs, all_outputs) if o.exists()]\n\nif not existing_inputs:\n     raise FileNotFoundError(\"No valid training data pairs found.\")\n\ntrain_indices = range(0, len(existing_inputs), 2)\ntrain_inputs = [existing_inputs[i] for i in train_indices]\ntrain_outputs = [existing_outputs[i] for i in train_indices]\n\nvalid_inputs = [f for i, f in enumerate(existing_inputs) if i not in train_indices]\nvalid_outputs = [o for i, o in enumerate(existing_outputs) if i not in train_indices]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.407025Z","iopub.execute_input":"2025-04-15T22:31:47.407272Z","iopub.status.idle":"2025-04-15T22:31:47.420914Z","shell.execute_reply.started":"2025-04-15T22:31:47.407251Z","shell.execute_reply":"2025-04-15T22:31:47.420128Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SeismicDataset(Dataset):\n    def __init__(self, inputs_files, output_files, n_examples_per_file=500):\n        assert len(inputs_files) == len(output_files)\n        self.inputs_files = inputs_files\n        self.output_files = output_files\n        self.n_examples_per_file = n_examples_per_file\n\n    def __len__(self):\n        return len(self.inputs_files) * self.n_examples_per_file\n\n    def __getitem__(self, idx):\n        file_idx = idx // self.n_examples_per_file\n        sample_idx = idx % self.n_examples_per_file\n\n        input_path = self.inputs_files[file_idx]\n        output_path = self.output_files[file_idx]\n\n        # Assumes valid pairs are passed to constructor\n        if not input_path.exists() or not output_path.exists():\n             raise FileNotFoundError(f\"File not found in __getitem__: {input_path} or {output_path}\")\n\n        X = np.load(input_path, mmap_mode='r')\n        y = np.load(output_path, mmap_mode='r')\n\n        try:\n            # Ensure sample_idx is within bounds\n            if sample_idx >= X.shape[0] or sample_idx >= y.shape[0]:\n                sample_idx = sample_idx % X.shape[0] # Assume X and y have same length\n\n            # Convert to float32 for the model\n            return X[sample_idx].copy().astype(np.float32), y[sample_idx].copy().astype(np.float32)\n        finally:\n            # Ensure memory maps are closed\n            if 'X' in locals() and hasattr(X, '_mmap') and X._mmap is not None:\n                X._mmap.close()\n            if 'y' in locals() and hasattr(y, '_mmap') and y._mmap is not None:\n                y._mmap.close()\n            del X, y # Explicitly delete references","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.421683Z","iopub.execute_input":"2025-04-15T22:31:47.421890Z","iopub.status.idle":"2025-04-15T22:31:47.433630Z","shell.execute_reply.started":"2025-04-15T22:31:47.421872Z","shell.execute_reply":"2025-04-15T22:31:47.432916Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_workers = 4\n\n# Filter out None values if __getitem__ can return None (currently it raises error)\ndef collate_fn(batch):\n    batch = list(filter(lambda x: x is not None and x[0] is not None, batch))\n    if not batch:\n        return None # Or handle empty batch case\n    return torch.utils.data.dataloader.default_collate(batch)\n\n# Ensure datasets are not empty before creating DataLoaders\ndltrain = None\nif train_inputs:\n    dstrain = SeismicDataset(train_inputs, train_outputs)\n    dltrain = DataLoader(dstrain, batch_size=128, shuffle=True, pin_memory=True, drop_last=True, num_workers=num_workers, persistent_workers=num_workers > 0, collate_fn=collate_fn)\n\ndlvalid = None\nif valid_inputs:\n    dsvalid = SeismicDataset(valid_inputs, valid_outputs)\n    dlvalid = DataLoader(dsvalid, batch_size=128, shuffle=False, pin_memory=True, drop_last=False, num_workers=num_workers, persistent_workers=num_workers > 0, collate_fn=collate_fn)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.434279Z","iopub.execute_input":"2025-04-15T22:31:47.434572Z","iopub.status.idle":"2025-04-15T22:31:47.453246Z","shell.execute_reply.started":"2025-04-15T22:31:47.434539Z","shell.execute_reply":"2025-04-15T22:31:47.452511Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## DumbNet","metadata":{}},{"cell_type":"markdown","source":"from [DumberNet-SUB, PUN](https://www.kaggle.com/code/pshikk/dumbernet-sub)","metadata":{}},{"cell_type":"code","source":"class ResidualBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, downsample=False):\n        super().__init__()\n        stride = 2 if downsample else 1\n\n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1, stride=stride)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n\n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n\n        self.downsample = None\n        if downsample or in_channels != out_channels:\n            self.downsample = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride),\n                nn.BatchNorm2d(out_channels)\n            )\n\n    def forward(self, x):\n        identity = x\n\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n\n        if self.downsample is not None:\n            identity = self.downsample(identity)\n\n        out += identity\n        return self.relu(out)\n\nclass DumbNet(nn.Module):\n    '''Deep CNN with residual blocks and dense classifier'''\n    def __init__(self, input_channels=5, output_size=70 * 70):\n        super().__init__()\n\n        self.stem = nn.Sequential(\n            nn.Conv2d(input_channels, 64, kernel_size=7, stride=2, padding=3),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(kernel_size=3, stride=2, padding=1)  # 1000x70 -> ~250x18\n        )\n\n        self.layer1 = ResidualBlock(64, 128, downsample=True)  # ~125x9\n        self.layer2 = ResidualBlock(128, 256, downsample=True)  # ~63x5\n        self.layer3 = ResidualBlock(256, 512, downsample=True)  # ~32x3\n        self.layer4 = ResidualBlock(512, 1024, downsample=True)\n        self.layer5 = ResidualBlock(1024, 1024, downsample = False)# same spatial\n\n        self.global_pool = nn.AdaptiveAvgPool2d((4, 4))  # fixed output\n\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(1024 * 4 * 4, 2048),\n            nn.GELU(),\n            nn.Dropout(0.5),\n\n            nn.Linear(2048, 1024),\n            nn.GELU(),\n            nn.Dropout(0.25),\n\n            nn.Linear(1024, output_size)\n        )\n\n    def forward(self, x):\n        bs = x.size(0)\n\n        x = self.stem(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        x = self.layer5(x)\n        x = self.global_pool(x)\n\n        x = self.classifier(x)\n        # Ensure the output is reshaped correctly and scaled\n        # Adding a small epsilon for stability if needed, but typically not required here\n        return x.view(bs, 1, 70, 70) * 1000.0 + 1500.0\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.455419Z","iopub.execute_input":"2025-04-15T22:31:47.455694Z","iopub.status.idle":"2025-04-15T22:31:47.474605Z","shell.execute_reply.started":"2025-04-15T22:31:47.455675Z","shell.execute_reply":"2025-04-15T22:31:47.473848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.475757Z","iopub.execute_input":"2025-04-15T22:31:47.476042Z","iopub.status.idle":"2025-04-15T22:31:47.562193Z","shell.execute_reply.started":"2025-04-15T22:31:47.476013Z","shell.execute_reply":"2025-04-15T22:31:47.561331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_zoo = {\"DumbNet\": DumbNet}\nif OPENFWI_AVAILABLE:\n    model_zoo[\"InversionNet\"] = InversionNet","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.563089Z","iopub.execute_input":"2025-04-15T22:31:47.563386Z","iopub.status.idle":"2025-04-15T22:31:47.576838Z","shell.execute_reply.started":"2025-04-15T22:31:47.563357Z","shell.execute_reply":"2025-04-15T22:31:47.576018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SAMPLE_SPATIAL = 1.0\nSAMPLE_TEMPORAL = 1\nNORM_LAYER = 'bn'\nUP_MODE = None # Or the correct value if known","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.577738Z","iopub.execute_input":"2025-04-15T22:31:47.578022Z","iopub.status.idle":"2025-04-15T22:31:47.590864Z","shell.execute_reply.started":"2025-04-15T22:31:47.577992Z","shell.execute_reply":"2025-04-15T22:31:47.590058Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = []\nmodel_class_names_loaded = [] \n\ntry:\n    if argparse.Namespace not in torch.serialization.get_safe_globals():\n         torch.serialization.add_safe_globals([argparse.Namespace])\n         print(\"argparse.Namespace added to safe globals for torch.load.\")\n    else:\n         print(\"argparse.Namespace is already in safe globals.\")\nexcept AttributeError:\n    print(\"Warning: torch.serialization.add_safe_globals not found or failed.\")\n# --- End safe type addition ---\n\nfor i, (path, model_cls_name) in enumerate(zip(PATHS, MODEL_CLASSES)):\n    print(f\"\\n--- Loading model {i+1}/{len(PATHS)} ---\")\n    print(f\"  Class: '{model_cls_name}'\")\n    print(f\"  Path:  '{path}'\")\n\n    if model_cls_name not in model_zoo:\n         print(f\"Warning: Model class '{model_cls_name}' not found in model_zoo. Skipping.\")\n         continue\n\n    ModelClass = model_zoo[model_cls_name]\n    try:\n        # --- Modification Start ---\n        if model_cls_name == \"InversionNet\":\n             # Initialize InversionNet with specific parameters if available\n             # Check if OPENFWI_AVAILABLE before attempting to instantiate\n             if OPENFWI_AVAILABLE:\n                 model = ModelClass(\n                     upsample_mode=UP_MODE,\n                     sample_spatial=SAMPLE_SPATIAL,\n                     sample_temporal=SAMPLE_TEMPORAL,\n                     norm=NORM_LAYER\n                 ).to(device)\n                 print(f\"  Instantiated InversionNet with specific parameters.\")\n             else:\n                 print(\"Error: Cannot instantiate InversionNet as OpenFWI is not available.\")\n                 continue # Skip this model if OpenFWI unavailable\n        else:\n             # Default instantiation for other models (like DumbNet)\n             model = ModelClass().to(device)\n             print(f\"  Instantiated model {model_cls_name}.\")\n        # --- Modification End ---\n    except Exception as e:\n        print(f\"Error: Failed to instantiate model class '{model_cls_name}': {e}. Skipping.\")\n        continue\n\n    state_dict = None\n    load_success = False\n    e_true = None # Keep error from weights_only=True attempt\n\n    # Attempt 1: weights_only=True (safer)\n    try:\n        print(f\"  Attempting to load state_dict with weights_only=True...\")\n        loaded_object_true = torch.load(path, map_location=device, weights_only=True)\n\n        # Check if loaded object is a dict and potentially a checkpoint\n        if isinstance(loaded_object_true, dict):\n            print(\"  Object loaded with weights_only=True is a dict.\")\n            # Heuristic check for checkpoint structure\n            is_checkpoint_dict = ('model' in loaded_object_true or\n                                  'optimizer' in loaded_object_true or\n                                  'state_dict' in loaded_object_true or # Key used in Submission.py\n                                  'state_dict_G' in loaded_object_true) # Possible GAN key\n\n            if is_checkpoint_dict:\n                print(\"  Detected checkpoint structure. Extracting state_dict...\")\n                if 'model' in loaded_object_true and isinstance(loaded_object_true['model'], dict):\n                    state_dict = loaded_object_true['model']\n                    print(\"  Extracted state_dict from 'model' key.\")\n                elif 'state_dict_G' in loaded_object_true and isinstance(loaded_object_true['state_dict_G'], dict) and model_cls_name == \"InversionNet\":\n                     state_dict = loaded_object_true['state_dict_G']\n                     print(\"  Extracted state_dict from 'state_dict_G' key.\")\n                elif 'state_dict' in loaded_object_true and isinstance(loaded_object_true['state_dict'], dict):\n                    state_dict = loaded_object_true['state_dict']\n                    print(\"  Extracted state_dict from 'state_dict' key.\")\n                else:\n                    print(\"Error: Checkpoint structure detected, but could not find standard state_dict keys ('model', 'state_dict', 'state_dict_G').\")\n                    state_dict = None # Extraction failed\n            else:\n                # Assume loaded dict is the state_dict itself\n                print(\"  Assuming loaded dict is the state_dict directly.\")\n                state_dict = loaded_object_true\n\n            # Validate and apply state_dict\n            if isinstance(state_dict, dict):\n                # Replace legacy keys if necessary (for InversionNet)\n                if model_cls_name == \"InversionNet\" and replace_legacy:\n                    try:\n                        print(\"  Applying replace_legacy...\")\n                        state_dict = replace_legacy(state_dict)\n                        print(\"  replace_legacy applied.\")\n                    except Exception as legacy_err:\n                        print(f\"Warning: Error applying replace_legacy: {legacy_err}\")\n\n                # Load state_dict into model\n                model.load_state_dict(state_dict) # Potential point of failure\n                print(\"  Successfully applied state_dict to model (weights_only=True path).\")\n                load_success = True\n            else:\n                 print(\"Error: Failed to extract/identify a valid state_dict (dictionary).\")\n                 load_success = False\n\n        else:\n             print(f\"Error: Loaded object with weights_only=True is not a dictionary (type: {type(loaded_object_true)}). Cannot load.\")\n\n    except RuntimeError as err_true:\n        e_true = err_true\n        print(f\"Warning: Failed during weights_only=True path. Error: {e_true}\")\n        # Check if error suggests fallback might help\n        if \"Unsupported global\" not in str(e_true) and \"WeightsUnpickler\" not in str(e_true):\n            print(\"  Error did not seem related to unsupported types. Fallback might not help.\")\n    except FileNotFoundError:\n        print(f\"Error: Model file not found at {path}. Skipping.\")\n        continue\n    except Exception as e_other_true:\n        print(f\"Error: An unexpected error occurred during weights_only=True load/processing: {e_other_true}. Skipping.\")\n        continue\n\n    # Attempt 2: weights_only=False (Fallback - less safe)\n    should_fallback = (\n        not load_success and\n        e_true is not None and\n        (\"Unsupported global\" in str(e_true) or \"WeightsUnpickler\" in str(e_true))\n    )\n    if should_fallback:\n        print(\"\\n  Attempting fallback: Loading with weights_only=False...\")\n        print(\"  >>> Warning: This can execute arbitrary code if the file is from an untrusted source! <<<\")\n        try:\n            loaded_object_false = torch.load(path, map_location=device, weights_only=False)\n            print(f\"  Loaded object type with weights_only=False: {type(loaded_object_false)}\")\n            if isinstance(loaded_object_false, dict):\n                if 'model' in loaded_object_false:\n                    print(\"  Found 'model' key. Extracting state_dict from checkpoint['model']...\")\n                    state_dict = loaded_object_false['model']\n                elif 'state_dict_G' in loaded_object_false and model_cls_name == \"InversionNet\":\n                    print(\"  Found 'state_dict_G' key. Extracting...\")\n                    state_dict = loaded_object_false['state_dict_G']\n                elif 'state_dict' in loaded_object_false:\n                    print(\"  Found 'state_dict' key. Extracting...\")\n                    state_dict = loaded_object_false['state_dict']\n                else:\n                    print(\"  No standard keys found. Assuming loaded dict is the state_dict.\")\n                    state_dict = loaded_object_false\n            elif isinstance(loaded_object_false, nn.Module):\n                 print(\"Warning: Loaded object is an nn.Module. Extracting state_dict from it.\")\n                 state_dict = loaded_object_false.state_dict()\n            else:\n                print(f\"Error: Loaded object with weights_only=False is not a dictionary or nn.Module (type: {type(loaded_object_false)}).\")\n                state_dict = None\n\n            if isinstance(state_dict, dict):\n                 print(\"  Successfully extracted/identified state_dict.\")\n                 if model_cls_name == \"InversionNet\" and replace_legacy:\n                     try:\n                         print(\"  Applying replace_legacy...\")\n                         state_dict = replace_legacy(state_dict)\n                         print(\"  replace_legacy applied.\")\n                     except Exception as legacy_err:\n                         print(f\"Warning: Error applying replace_legacy: {legacy_err}\")\n                 model.load_state_dict(state_dict)\n                 print(\"  Successfully applied state_dict to model (using weights_only=False fallback).\")\n                 load_success = True\n            elif state_dict is not None:\n                 print(f\"Error: Extracted state_dict is not a dictionary (type: {type(state_dict)}). Load failed.\")\n\n        except Exception as e_false:\n            print(f\"Error: Failed during weights_only=False fallback processing. Error: {e_false}. Skipping this model.\")\n            load_success = False\n\n    # --- Loading finished for this model ---\n    if load_success:\n        model.eval()\n        models.append(model)\n        model_class_names_loaded.append(model_cls_name)\n        print(f\"--- Successfully loaded and prepared model {i+1} ({model_cls_name}) ---\")\n    else:\n        print(f\"--- Failed to load model {i+1} ({model_cls_name}) ---\")\n# --- Model Loading Loop End ---\n\n# --- Post-loading checks and weight filtering ---\nif not models:\n    raise RuntimeError(\"No models could be loaded successfully. Please check logs for errors.\")\noriginal_indices = [i for i, cls_name in enumerate(MODEL_CLASSES) if cls_name in model_class_names_loaded]\nfiltered_weights = [PATH_WEIGHTS[i] for i in original_indices]\ntotal_weight = sum(filtered_weights)\nif total_weight > 0:\n     PATH_WEIGHTS_USED = [w / total_weight for w in filtered_weights]\nelse:\n     if models:\n         PATH_WEIGHTS_USED = [1.0 / len(models)] * len(models) # Equal weights if original sum was zero\n     else:\n          PATH_WEIGHTS_USED = [] # Should not happen if models list is checked above\nprint(f\"\\nSuccessfully loaded {len(models)} models in total.\")\nprint(f\"Models loaded: {model_class_names_loaded}\")\nprint(f\"Weights used (normalized): {PATH_WEIGHTS_USED}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:47.591654Z","iopub.execute_input":"2025-04-15T22:31:47.591870Z","iopub.status.idle":"2025-04-15T22:31:53.592138Z","shell.execute_reply.started":"2025-04-15T22:31:47.591843Z","shell.execute_reply":"2025-04-15T22:31:53.591454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ctx = None\nif \"InversionNet\" in model_class_names_loaded:\n    if OPENFWI_AVAILABLE and T is not None:\n        ctx = load_dataset_config(DATASET_CONFIG, DATASET_NAME)\n        if ctx is None:\n             print(\"Error: Failed to load dataset config, which is required for InversionNet. Exiting.\")\n             sys.exit(1)\n        # Precompute normalization parameters\n        log_data_min = T.log_transform(ctx['data_min'], k=K_TRANSFORM)\n        log_data_max = T.log_transform(ctx['data_max'], k=K_TRANSFORM)\n        label_min = ctx['label_min']\n        label_max = ctx['label_max']\n    else:\n         print(\"Error: InversionNet is selected, but OpenFWI transforms (T) are not available. Cannot perform preprocessing. Exiting.\")\n         sys.exit(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:53.592942Z","iopub.execute_input":"2025-04-15T22:31:53.593217Z","iopub.status.idle":"2025-04-15T22:31:53.600668Z","shell.execute_reply.started":"2025-04-15T22:31:53.593193Z","shell.execute_reply":"2025-04-15T22:31:53.599848Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference and Submission","metadata":{}},{"cell_type":"code","source":"def write_submission_csv(rows_data, filename, fieldnames):\n    \"\"\"Writes a list of row dictionaries to a CSV file.\"\"\"\n    if not rows_data:\n        print(f\"Warning: No data provided to write to {filename}. Skipping file creation.\")\n        return\n\n    print(f\"Writing {len(rows_data)} rows to {filename}...\")\n    try:\n        with open(filename, 'wt', newline='') as csvfile:\n            writer = csv.DictWriter(csvfile, fieldnames=fieldnames)\n            writer.writeheader()\n            writer.writerows(rows_data)\n        print(f\"Successfully wrote {filename}\")\n    except Exception as e:\n        print(f\"Error writing file {filename}: {e}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:53.601528Z","iopub.execute_input":"2025-04-15T22:31:53.601793Z","iopub.status.idle":"2025-04-15T22:31:53.615623Z","shell.execute_reply.started":"2025-04-15T22:31:53.601765Z","shell.execute_reply":"2025-04-15T22:31:53.614863Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dir = Path('/kaggle/input/waveform-inversion/test')\nif test_dir.exists():\n    test_files = list(test_dir.glob('*.npy'))\nelse:\n    print(f\"Warning: Test directory not found: {test_dir}\")\n    test_files = []\n\n# --- Debugging Filter ---\nif DEBUG and test_files:\n    print(f\"*** DEBUG MODE: Processing only OID: {OID_DEBUG} ***\")\n    original_count = len(test_files)\n    test_files = [f for f in test_files if f.stem == OID_DEBUG]\n    if not test_files:\n        print(f\"Warning: Debug OID '{OID_DEBUG}' not found in test files. No files to process.\")\n    else:\n        print(f\"Filtered {original_count} files down to {len(test_files)} for debug.\")\n# --- Debugging Filter End ---\n\nif not test_files:\n    print(\"No test files found. Skipping prediction and submission.\")\nelse:\n    print(f\"Found {len(test_files)} test files.\")\n\n    x_cols = [f'x_{i}' for i in range(1, 70, 2)]\n    fieldnames = ['oid_ypos'] + x_cols\n\n    class TestDataset(Dataset):\n        def __init__(self, test_files):\n            self.test_files = test_files\n\n        def __len__(self):\n            return len(self.test_files)\n\n        def __getitem__(self, i):\n            test_file = self.test_files[i]\n            if not test_file.exists():\n                raise FileNotFoundError(f\"Test file not found: {test_file}\")\n\n            try:\n                data = np.load(test_file)\n                return data.astype(np.float32), test_file.stem # Convert to float32\n            except Exception as e:\n                print(f\"Error loading test file {test_file}: {e}\")\n                raise # Re-raise error for DataLoader to handle (or skip)\n\n    dl_test = None\n    if test_files: # Only create dataset and dataloader if there are files\n        ds_test = TestDataset(test_files)\n        dl_test = DataLoader(ds_test, batch_size=128, num_workers=num_workers, pin_memory=True, collate_fn=collate_fn)\n\n        # --- Ensemble Inference with Pre/Post-processing ---\n        output_filename = 'submission.csv'\n        print(f\"Starting prediction loop. Output will be saved to {output_filename}\")\n        plot_counter = 0\n        MAX_PLOTS = 5\n        # VEL_MIN/MAX removed, using label_min/label_max from ctx for InversionNet\n\n        # --- Initialize storage for debug results if DEBUG is True ---\n        debug_model_results = {}\n        if DEBUG:\n            print(\"DEBUG MODE: Initializing storage for individual model results.\")\n            for i in range(len(models)):\n                debug_model_results[i] = [] # Create an empty list for each model index\n        # --- Debug Storage Initialization End ---\n\n        all_final_rows = [] # Collect all rows for the final submission.csv\n\n        if dl_test is not None:\n            for batch_data in tqdm(dl_test, desc='test'):\n                if batch_data is None: continue\n                inputs, oids_test = batch_data\n                if inputs is None or not len(inputs): continue\n\n                inputs = inputs.to(device)\n\n                # --- Ensemble Inference Logic (Revised) ---\n                batch_model_outputs_np = [] # Store numpy results (batch, H, W) for each model\n\n                with torch.inference_mode():\n                    processing_successful = True # Flag to track if all models processed ok\n                    for i, (model, weight, model_cls_name) in enumerate(zip(models, PATH_WEIGHTS_USED, model_class_names_loaded)):\n                        try:\n                            model_output_np = None # Reset for current model\n\n                            # --- Model-specific Preprocessing, Inference, Postprocessing ---\n                            if model_cls_name == \"DumbNet\":\n                                # DumbNet uses raw inputs directly? Confirm this assumption.\n                                inputs_processed = inputs\n                                outputs_raw = model(inputs_processed) # Forward pass includes scaling\n\n                                # Postprocessing: Just move to CPU and NumPy, remove channel dim\n                                if outputs_raw.shape[1] != 1:\n                                     print(f\"Warning: DumbNet output channel is not 1 ({outputs_raw.shape}). Check model def.\")\n                                # Ensure tensor before detaching/moving\n                                if isinstance(outputs_raw, torch.Tensor):\n                                    model_output_np = outputs_raw[:, 0].detach().cpu().numpy()\n                                else:\n                                    raise TypeError(\"DumbNet output was not a tensor.\")\n\n\n                            elif model_cls_name == \"InversionNet\":\n                                if ctx is None or T is None:\n                                    raise RuntimeError(\"InversionNet requires dataset config and transforms.\")\n\n                                # Preprocessing (CPU based, as required by transforms)\n                                inputs_cpu = inputs.cpu()\n                                inputs_log_cpu = T.log_transform(inputs_cpu, k=K_TRANSFORM)\n                                inputs_processed_cpu = T.minmax_normalize(inputs_log_cpu, log_data_min, log_data_max)\n                                inputs_processed = inputs_processed_cpu.to(device)\n\n                                # Inference\n                                outputs_raw = model(inputs_processed)\n\n                                # Postprocessing using T.tonumpy_denormalize\n                                model_output_np = T.tonumpy_denormalize(outputs_raw, ctx['label_min'], ctx['label_max'], exp=False)\n\n                                # Ensure shape is (batch, height, width) after denormalize\n                                if isinstance(model_output_np, np.ndarray):\n                                    if model_output_np.ndim == 4 and model_output_np.shape[1] == 1:\n                                        model_output_np = model_output_np[:, 0, :, :] # Remove channel dim\n                                    elif model_output_np.ndim != 3:\n                                         raise ValueError(f\"InversionNet denormalized shape unexpected: {model_output_np.shape}\")\n                                else:\n                                     raise TypeError(\"T.tonumpy_denormalize did not return numpy array.\")\n\n                            else:\n                                print(f\"Warning: Unknown model class '{model_cls_name}'. Skipping this model.\")\n                                processing_successful = False # Mark failure for this model\n                                continue # Skip to next model in ensemble\n\n                            # --- Validation and Storage ---\n                            if model_output_np is not None:\n                                # Basic shape check - should be (batch_size, 70, 70)\n                                expected_shape_hw = (70, 70)\n                                if model_output_np.shape[1:] != expected_shape_hw:\n                                    print(f\"Warning: Output shape mismatch for {model_cls_name}. Got {model_output_np.shape}, expected (*, 70, 70). Skipping model.\")\n                                    processing_successful = False # Mark failure\n                                else:\n                                    batch_model_outputs_np.append(model_output_np) # Store the numpy result\n                                    # --- Store individual model results if DEBUG is True ---\n                                    if DEBUG:\n                                         # Generate rows from model_output_np for this batch\n                                         for k in range(len(oids_test)): # Iterate through samples in the batch\n                                             oid_test_debug = oids_test[k]\n                                             y_pred_debug = model_output_np[k] # Prediction for one sample\n                                             for y_pos_debug in range(y_pred_debug.shape[0]):\n                                                  try:\n                                                       width_debug = y_pred_debug.shape[1]\n                                                       x_indices_debug = range(1, min(width_debug, 70), 2)\n                                                       row_data_debug = {}\n                                                       for x_idx_debug in x_indices_debug:\n                                                            col_name_debug = f\"x_{x_idx_debug}\"\n                                                            row_data_debug[col_name_debug] = np.nan_to_num(y_pred_debug[y_pos_debug, x_idx_debug])\n\n                                                       last_valid_x_debug = x_indices_debug[-1] if x_indices_debug else -1\n                                                       for x_idx_req_debug in range(1, 70, 2):\n                                                            col_name_req_debug = f\"x_{x_idx_req_debug}\"\n                                                            if x_idx_req_debug > last_valid_x_debug:\n                                                                fill_value_debug = np.nan_to_num(y_pred_debug[y_pos_debug, last_valid_x_debug]) if last_valid_x_debug >= 0 else 0.0\n                                                                row_data_debug[col_name_req_debug] = fill_value_debug\n\n                                                       row_debug = {'oid_ypos': f\"{oid_test_debug}_y_{y_pos_debug}\"}\n                                                       row_debug.update(row_data_debug)\n                                                       debug_model_results[i].append(row_debug)\n                                                  except Exception as e_debug_write:\n                                                      print(f\"Error generating debug row for model {i}, oid {oid_test_debug}, y_pos {y_pos_debug}: {e_debug_write}\")\n                                    # --- Debug Storage End ---\n                            else:\n                                 print(f\"Error: No output generated for model {model_cls_name}. Skipping model.\")\n                                 processing_successful = False # Mark failure\n\n                        except Exception as e:\n                            print(f\"Error processing model {i} ({model_cls_name}): {e}\")\n                            processing_successful = False # Mark failure for this batch\n                            break # Stop processing models for this batch if one fails severely\n\n                    # --- Ensemble Averaging ---\n                    y_preds = None # Final ensembled numpy prediction for the batch\n                    if processing_successful and len(batch_model_outputs_np) == len(models):\n                        if len(models) > 1:\n                            try:\n                                # Ensure weights list matches number of results\n                                if len(PATH_WEIGHTS_USED) == len(batch_model_outputs_np):\n                                     print(f\"Ensembling {len(models)} model outputs with weights: {PATH_WEIGHTS_USED}\")\n                                     y_preds = np.average(batch_model_outputs_np, axis=0, weights=PATH_WEIGHTS_USED)\n                                else:\n                                     print(f\"Error: Weight count ({len(PATH_WEIGHTS_USED)}) mismatch with model output count ({len(batch_model_outputs_np)}). Cannot average.\")\n                                     y_preds = None # Averaging failed\n                            except Exception as avg_err:\n                                print(f\"Error during numpy averaging: {avg_err}\")\n                                y_preds = None # Averaging failed\n                        elif len(models) == 1:\n                            print(\"Only one model processed successfully. Using its output.\")\n                            y_preds = batch_model_outputs_np[0]\n                        else:\n                            print(\"No models to ensemble.\") # Should not happen if models list existed\n                    elif not processing_successful:\n                         print(\"Skipping ensemble averaging for this batch due to errors during model processing.\")\n                    elif len(batch_model_outputs_np) != len(models):\n                         print(f\"Skipping ensemble averaging: Expected {len(models)} results, but got {len(batch_model_outputs_np)}.\")\n\n\n                # Check if ensemble calculation was successful (y_preds is a numpy array)\n                if y_preds is None:\n                    print(\"Skipping batch due to errors in processing or ensembling.\")\n                    continue\n\n                # --- Plotting (uses the final ensembled numpy prediction y_preds) ---\n                if plot_counter < MAX_PLOTS:\n                    num_to_plot = min(MAX_PLOTS - plot_counter, len(oids_test))\n                    for i in range(num_to_plot):\n                        if plot_counter >= MAX_PLOTS: break\n                        plt.figure(figsize=(10, 4))\n                        # Plot input\n                        plt.subplot(1, 2, 1)\n                        try:\n                            input_to_plot = inputs[i, 0].cpu().numpy() # Input sample\n                            plt.imshow(input_to_plot, cmap='seismic', aspect='auto')\n                            plt.title(f\"Input Sample {plot_counter + 1} (Ch 0)\")\n                            plt.colorbar(label='Amplitude')\n                        except IndexError:\n                            print(f\"Warning: Could not plot input for index {i}.\")\n                            plt.title(f\"Input Sample {plot_counter + 1} (Error)\")\n                        # Plot final ensembled prediction\n                        plt.subplot(1, 2, 2)\n                        try:\n                            pred_to_plot = y_preds[i] # Ensembled prediction sample\n                            # Determine vmin/vmax based on involved models or use defaults\n                            vmin_plot, vmax_plot = (label_min, label_max) if \"InversionNet\" in model_class_names_loaded and ctx else (1500, 4500) # Example fallback\n                            im = plt.imshow(pred_to_plot, cmap='viridis', aspect='auto', vmin=vmin_plot, vmax=vmax_plot)\n                            plt.title(f\"Ensembled Prediction {oids_test[i]}\")\n                            plt.colorbar(im, label='Velocity (m/s)')\n                        except IndexError:\n                            print(f\"Warning: Could not plot prediction for index {i}.\")\n                            plt.title(f\"Prediction {oids_test[i]} (Error)\")\n                        plt.tight_layout()\n                        plt.show()\n                        plot_counter += 1\n                # --- Plotting End ---\n\n                # --- CSV Writing (uses the final ensembled numpy prediction y_preds) ---\n                if y_preds is not None: # Double check y_preds exists\n                    for y_pred, oid_test in zip(y_preds, oids_test): # y_pred is a 2D numpy array (H, W)\n                        for y_pos in range(y_pred.shape[0]): # Iterate through height\n                            try:\n                                # Ensure width allows indexing up to x_69\n                                width = y_pred.shape[1]\n                                x_indices = range(1, min(width, 70), 2)\n\n                                row_data = {}\n                                for x_idx in x_indices:\n                                    col_name = f\"x_{x_idx}\"\n                                    # Replace NaN/inf with 0 before saving\n                                    row_data[col_name] = np.nan_to_num(y_pred[y_pos, x_idx])\n\n                                # Handle missing columns if width < 70 (similar to first_starter.py)\n                                last_valid_x = x_indices[-1] if x_indices else -1\n                                for x_idx_req in range(1, 70, 2):\n                                    col_name_req = f\"x_{x_idx_req}\"\n                                    if x_idx_req > last_valid_x:\n                                        # Use nan_to_num on the fill value source as well\n                                        fill_value = np.nan_to_num(y_pred[y_pos, last_valid_x]) if last_valid_x >= 0 else 0.0 # Use last valid or 0\n                                        row_data[col_name_req] = fill_value\n\n                                # Add oid_ypos and write row\n                                row = {'oid_ypos': f\"{oid_test}_y_{y_pos}\"}\n                                row.update(row_data) # Add the x_ values\n                                all_final_rows.append(row)\n\n                            except IndexError:\n                                print(f\"Warning: Index error during CSV writing for oid {oid_test}, y_pos {y_pos}. Skipping row.\")\n                                continue\n                            except Exception as e_write:\n                                print(f\"Error writing row for oid {oid_test}, y_pos {y_pos}: {e_write}\")\n                                continue\n                # --- Writing End ---\n\n        # --- Write all rows to CSV at the end of the loop ---\n        write_submission_csv(all_final_rows, output_filename, fieldnames)\n\n        # --- Write individual model CSVs if DEBUG is True ---\n        if DEBUG:\n            print(\"\\n--- Writing individual model CSVs (DEBUG MODE) ---\")\n            for i, rows_data in debug_model_results.items():\n                model_cls_name = model_class_names_loaded[i] # Get class name for filename\n                debug_filename = f\"submission_debug_model_{i}_{model_cls_name}.csv\"\n                write_submission_csv(rows_data, debug_filename, fieldnames)\n            print(\"--- Finished writing individual model CSVs ---\")\n        # --- Individual CSV Writing End ---\n\n    else:\n        print(\"No test files found or DataLoader could not be created, skipping prediction loop.\")\n# --- Prediction End ---\n\n\n\n# Optional: Clean up\n# del models, ds_test, dl_test, dstrain, dltrain, dsvalid, dlvalid\n# torch.cuda.empty_cache()\n\nprint(\"Script finished.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T22:31:53.616412Z","iopub.execute_input":"2025-04-15T22:31:53.616647Z","iopub.status.idle":"2025-04-15T22:31:56.251695Z","shell.execute_reply.started":"2025-04-15T22:31:53.616628Z","shell.execute_reply":"2025-04-15T22:31:56.250724Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}