{"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":"gpu","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"},{"sourceId":11510285,"sourceType":"datasetVersion","datasetId":7217458},{"sourceId":11511096,"sourceType":"datasetVersion","datasetId":7218098},{"sourceId":12017632,"sourceType":"datasetVersion","datasetId":7560801}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport sys\nimport json\nimport glob\n\nimport torch\nimport numpy as np\nimport pandas as pd\nfrom torchvision.transforms import Compose\nsys.path.append('/kaggle/input/openfwi/OpenFWI-main')\nimport network\nimport transforms as T\nfrom dataset import FWIDataset","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-31T14:40:19.277841Z","iopub.execute_input":"2025-05-31T14:40:19.278091Z","iopub.status.idle":"2025-05-31T14:40:27.985977Z","shell.execute_reply.started":"2025-05-31T14:40:19.278072Z","shell.execute_reply":"2025-05-31T14:40:27.985203Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEBUG = False\nWEIGHTS_PATH = '/kaggle/input/fwi-pretrained-weight/PretrainedModel/'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T14:40:31.867751Z","iopub.execute_input":"2025-05-31T14:40:31.868069Z","iopub.status.idle":"2025-05-31T14:40:31.871699Z","shell.execute_reply.started":"2025-05-31T14:40:31.868047Z","shell.execute_reply":"2025-05-31T14:40:31.871053Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_NAME = 'InversionNet'\n\nTEST_DATA_DIR = '/kaggle/input/waveform-inversion/test'\nOUTPUT_CSV = '/kaggle/working/submission.csv'\nDATASET_CONFIG = '/kaggle/input/openfwi/OpenFWI-main/dataset_config.json'\n\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nK_TRANSFORM = 1 # k value for LogTransform, adjust if needed (from test.py --k)\n\n# Sample spatial/temporal might be needed depending on the model architecture used for fva_l1.pth\nSAMPLE_SPATIAL = 1.0 # Adjust if needed (from test.py --sample-spatial)\nSAMPLE_TEMPORAL = 1  # Adjust if needed (from test.py --sample-temporal)\nNORM_LAYER = 'bn' # Adjust if needed (from test.py --norm)\nUP_MODE = None # Adjust if needed (from test.py --up-mode)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T14:40:34.973520Z","iopub.execute_input":"2025-05-31T14:40:34.973795Z","iopub.status.idle":"2025-05-31T14:40:35.042005Z","shell.execute_reply.started":"2025-05-31T14:40:34.973776Z","shell.execute_reply":"2025-05-31T14:40:35.041008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_dataset_config(config_path, dataset_name):\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        return ctx\n\n    except FileNotFoundError:\n        print(f\"Error: {config_path} not found.\")\n        sys.exit(1)\n    except KeyError:\n        print(f\"Error: {dataset_name} not found in {config_path}.\")\n        sys.exit(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T14:40:38.026942Z","iopub.execute_input":"2025-05-31T14:40:38.027334Z","iopub.status.idle":"2025-05-31T14:40:38.032068Z","shell.execute_reply.started":"2025-05-31T14:40:38.027306Z","shell.execute_reply":"2025-05-31T14:40:38.031489Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"LabelsMap = {\n    0: \"curvefault-a\",\n    1: \"curvefault-b\",\n    2: \"curvevel-a\",\n    3: \"curvevel-b\",\n    4: \"flatfault-a\",\n    5: \"flatfault-b\",\n    6: \"flatvel-a\",\n    7: \"flatvel-b\",\n    8: \"style-a\",\n    9: \"style-b\",\n}\n\nLabelToNum = {v: k for k, v in LabelsMap.items()}\n\nWeightsMap = {\n    \"curvefault-a\":\"cfa_l1.pth\",\n    \"curvefault-b\":\"cfb_l1.pth\",\n    \"curvevel-a\":\"cva_l1.pth\",\n    \"curvevel-b\":\"cvb_l1.pth\",\n    \"flatfault-a\":\"ffa_l1.pth\",\n    \"flatfault-b\":\"ffb_l1.pth\",\n    \"flatvel-a\":\"fva_l1.pth\",\n    \"flatvel-b\":\"fvb_l1.pth\",\n    \"style-a\":\"sta_l1_new.pth\",\n    \"style-b\":\"stb_l1.pth\",\n}\n\n\ndef main():\n    print(f\"Using device: {DEVICE}\")\n    print(f\"Test data directory: {TEST_DATA_DIR}\")\n    print(f\"Output CSV {OUTPUT_CSV}\")\n    print(f\"Loading model: {MODEL_NAME}\")\n\n    # Initialize Model\n    if MODEL_NAME not in network.model_dict:\n        print(f\"Error: Unsupported mode '{MODEL_NAME}'. Check network.py\")\n        sys.exit(1)\n\n    model = network.model_dict[MODEL_NAME](\n        upsample_mode=UP_MODE,\n        sample_temporal=SAMPLE_TEMPORAL,\n        norm=NORM_LAYER\n    ).to(DEVICE)\n\n    # Find test files\n    test_files = glob.glob(os.path.join(TEST_DATA_DIR, '*.npy'))\n    if not test_files:\n        print(f\"Error: No .npy files found in {TEST_DATA_DIR}\")\n        sys.exit(1)\n\n    result = [None]*len(test_files)\n\n    # Read the list of styles of test data\n    test_types = [int(line.strip()) for line in open(\"/kaggle/input/test-type-fwi/test_type.txt\", 'r', encoding='utf-8')]\n    \n    print(\"Loop all styles\")\n    for style in LabelsMap.values():\n        W_PATH = WEIGHTS_PATH + WeightsMap[style]\n        print(f\"Style {style}\")\n        print(f\"Loading weights from {W_PATH}\")\n\n        ctx = load_dataset_config(DATASET_CONFIG, style)\n\n        #Load Weights\n        try:\n            checkpoint = torch.load(W_PATH, map_location=DEVICE, weights_only=False)\n            if 'model' in checkpoint:\n                state_dict = checkpoint['model']\n            elif 'state_dict' in checkpint:\n                state_dict = checkpoint['state_dict']\n            else:\n                state_dict = checkpoint\n\n            model.load_state_dict(state_dict)\n            print(\"Model weights loaded successfully.\")\n\n            if 'epoch' in checkpoint and 'step' in checkpoint:\n                print(f\"Weights from Epoch {checkpoint['epoch']}/Step {checkpoint['step']}.\")\n\n        except FileNotFoundError:\n            print(f\"Error: Weights file not found at {W_PATH}\")\n            sys.exit(1)\n\n        except Exception as e:\n            print(f\"Error loading weights: {e}\")\n            sys.exit(1)\n\n        model.eval()\n\n        # Inference\n        with torch.no_grad():\n            for i, (file_path, t_style) in enumerate(zip(test_files, test_types)):\n                # If style isn't match next case.\n                if LabelsMap[t_style] != style:\n                    continue\n\n                else:\n                    oid = os.path.splitext(os.path.basename(file_path))[0]\n                    \n                    try:\n                        seismic_data = np.load(file_path)\n                        if seismic_data.ndim == 3:\n                            seismic_data = seismic_data[np.newaxis, :, :, :]\n                        elif seismic_data.ndim != 4:\n                            print(f\"Warning: Unexpected data dimension {seismic_data.ndim} for {oid}. Skipping\")\n                            continue\n\n                        # Convert to tensor and move to device\n                        data_tensor = torch.from_numpy(seismic_data).type(torch.FloatTensor)\n                        data_tensor = T.log_transform(data_tensor, k=K_TRANSFORM)\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                        data_tensor = T.minmax_normalize(data_tensor, log_data_min, log_data_max)\n\n                        data_tensor = data_tensor.to(DEVICE)\n\n                        # Prediction\n                        pred_tensor = model(data_tensor)\n\n                        # Denormalize prediction\n                        pred_np = T.tonumpy_denormalize(pred_tensor, ctx['label_min'], ctx['label_max'], exp=False)\n                        velo_map = pred_np[0,0]\n\n                        # Format for submission\n                        height, width = velo_map.shape\n                        if width < 70:\n                            print(f\"Warning: Predicted width {width} for {oid} is less than 70. Padding or check model output.\")\n\n                        tmp_result = []\n                        for y_pos in range(height):\n                            # Processing each row\n                            row_data = {'oid_ypos': f\"{oid}_y_{y_pos}\"}\n                            odd_indices = range(1, min(width, 70), 2)\n\n                            for x_idx in odd_indices:\n                                col_name = f\"x_{x_idx}\"\n                                row_data[col_name]=velo_map[y_pos, x_idx]\n\n                            # Handle missing velue\n                            last_valid_x = odd_indices[-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                                    fill_value = velo_map[y_pos, last_valid_x] if last_valid_x >= 0 else 3000\n                                    row_data[col_name_req] = fill_value\n\n\n                            tmp_result.append(row_data)\n                        result[i] = tmp_result\n                    except Exception as e:\n                        print(f\"Error processing file {file_path}: {e}\")\n                        continue\n\n    if not result:\n        print(\"No results generated. Exiting.\")\n        sys.exit(1)\n\n\n    # Make submission\n    submission = []\n    for res in result:\n        for row in res:\n            submission.append(row)\n\n    submission_df = pd.DataFrame(submission)\n            \n    # Ensure correct column order\n    cols = ['oid_ypos'] + [f'x_{i}' for i in range(1, 70, 2)]\n    submission_df = submission_df[cols]\n    submission_df.to_csv(OUTPUT_CSV, index=False, float_format='%.4f')\n\n    print(f\"Submission file save to {OUTPUT_CSV}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T07:52:37.029131Z","iopub.execute_input":"2025-05-31T07:52:37.029403Z","iopub.status.idle":"2025-05-31T07:52:37.044894Z","shell.execute_reply.started":"2025-05-31T07:52:37.029383Z","shell.execute_reply":"2025-05-31T07:52:37.044168Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == '__main__':\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T07:52:37.046040Z","iopub.execute_input":"2025-05-31T07:52:37.046263Z","iopub.status.idle":"2025-05-31T08:32:13.780175Z","shell.execute_reply.started":"2025-05-31T07:52:37.046248Z","shell.execute_reply":"2025-05-31T08:32:13.777533Z"}},"outputs":[],"execution_count":null}]}