{"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":11568812,"sourceType":"datasetVersion","datasetId":7253205},{"sourceId":11569667,"sourceType":"datasetVersion","datasetId":7253605},{"sourceId":12038896,"sourceType":"datasetVersion","datasetId":7377931},{"sourceId":78406525,"sourceType":"kernelVersion"},{"sourceId":6119,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":4602,"modelId":2797}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============== Configuration ==============\nTRAIN = True\nPREDICT = True\nUSE_FULL_DATA = True  # Set False for quick testing\nENSEMBLE_FOLDS = 5  # Reduce for faster training\n\n# ============== Install Dependencies ==============\nimport subprocess\nimport sys\n\ndef install_packages():\n    \"\"\"Install required packages\"\"\"\n    packages = [\n        'einops',\n        'timm==0.9.16',\n        'segmentation-models-pytorch',\n        'pywavelets',\n        'torchmetrics',\n        'albumentations'\n    ]\n    \n    for package in packages:\n        subprocess.check_call([sys.executable, '-m', 'pip', 'install', package, '-q'])\n    \n    print(\"All packages installed successfully!\")\n\n# Run installation\ninstall_packages()\n\n# ============== Import Everything ==============\nimport os\nimport gc\nimport glob\nimport time\nimport json\nimport random\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, OneCycleLR\n\n# Check GPU\nprint(f\"PyTorch version: {torch.__version__}\")\nprint(f\"CUDA available: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    print(f\"CUDA version: {torch.version.cuda}\")\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"GPU memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\n\n# ============== Quick Test Function ==============\ndef quick_test():\n    \"\"\"Quick test to ensure everything works\"\"\"\n    print(\"\\nRunning quick test...\")\n    \n    # Test data loading\n    test_files = glob.glob(\"/kaggle/input/waveform-inversion/test/*.npy\")\n    print(f\"Found {len(test_files)} test files\")\n    \n    if len(test_files) > 0:\n        # Load sample\n        sample = np.load(test_files[0])\n        print(f\"Test sample shape: {sample.shape}\")\n        print(f\"Test sample dtype: {sample.dtype}\")\n        print(f\"Test sample range: [{sample.min():.2f}, {sample.max():.2f}]\")\n    \n    # Test model creation\n    print(\"\\nTesting model creation...\")\n    dummy_model = nn.Conv2d(5, 1, 3, padding=1).cuda()\n    dummy_input = torch.randn(1, 5, 70, 70).cuda()\n    with torch.no_grad():\n        output = dummy_model(dummy_input)\n    print(f\"Model output shape: {output.shape}\")\n    \n    print(\"\\nQuick test passed! ✓\")\n    \n    # Clean up\n    del dummy_model, dummy_input\n    gc.collect()\n    torch.cuda.empty_cache()\n\n# Run quick test\nquick_test()\n\n# ============== Data Statistics ==============\ndef analyze_competition_data():\n    \"\"\"Analyze competition data structure\"\"\"\n    print(\"\\n\" + \"=\"*50)\n    print(\"Competition Data Analysis\")\n    print(\"=\"*50)\n    \n    # Training data\n    train_path = \"/kaggle/input/waveform-inversion/train_samples/\"\n    if os.path.exists(train_path):\n        datasets = ['CurveFault_A', 'CurveFault_B', 'CurveVel_A', 'CurveVel_B',\n                   'FlatFault_A', 'FlatFault_B', 'FlatVel_A', 'FlatVel_B',\n                   'Style_A', 'Style_B']\n        \n        print(\"\\nTraining datasets:\")\n        for dataset in datasets:\n            dataset_path = os.path.join(train_path, dataset)\n            if os.path.exists(dataset_path):\n                n_files = len(glob.glob(os.path.join(dataset_path, \"**/*.npy\"), recursive=True))\n                print(f\"  {dataset}: {n_files} files\")\n    \n    # Test data\n    test_files = glob.glob(\"/kaggle/input/waveform-inversion/test/*.npy\")\n    print(f\"\\nTest files: {len(test_files)}\")\n    \n    # Submission format\n    sub_df = pd.read_csv(\"/kaggle/input/waveform-inversion/sample_submission.csv\")\n    print(f\"\\nSubmission rows: {len(sub_df)}\")\n    print(f\"Submission columns: {len(sub_df.columns)}\")\n    print(f\"Expected predictions per file: {len(sub_df) // len(test_files) if len(test_files) > 0 else 'N/A'}\")\n\n# Analyze data\nanalyze_competition_data()\n\n# ============== Memory Management ==============\ndef get_memory_usage():\n    \"\"\"Get current GPU memory usage\"\"\"\n    if torch.cuda.is_available():\n        allocated = torch.cuda.memory_allocated() / 1e9\n        reserved = torch.cuda.memory_reserved() / 1e9\n        return f\"GPU Memory: {allocated:.2f}GB allocated, {reserved:.2f}GB reserved\"\n    return \"No GPU available\"\n\nprint(f\"\\n{get_memory_usage()}\")\n\n# ============== Ready to Train Message ==============\nprint(\"\\n\" + \"=\"*50)\nprint(\"✓ Environment setup complete!\")\nprint(\"✓ All dependencies installed!\")\nprint(\"✓ Data paths verified!\")\nprint(\"\\n→ You can now run the main training script\")\nprint(\"→ Recommended: Start with ENSEMBLE_FOLDS=1 for testing\")\nprint(\"=\"*50)\n\n# ============== Helper Functions ==============\ndef create_folds_csv():\n    \"\"\"Create folds.csv if it doesn't exist\"\"\"\n    print(\"\\nCreating folds.csv...\")\n    \n    # This is a simplified version - adjust based on your data\n    data_info = []\n    \n    # Add your data loading logic here\n    # For each file, assign to a fold\n    \n    # Save as CSV\n    df = pd.DataFrame(data_info)\n    df.to_csv('folds.csv', index=False)\n    print(\"Folds created!\")\n\ndef visualize_predictions(pred, target, save_path=None):\n    \"\"\"Visualize prediction vs target\"\"\"\n    fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n    \n    # Prediction\n    im1 = axes[0].imshow(pred, cmap='jet', vmin=1500, vmax=6000)\n    axes[0].set_title('Prediction')\n    axes[0].axis('off')\n    plt.colorbar(im1, ax=axes[0])\n    \n    # Target\n    im2 = axes[1].imshow(target, cmap='jet', vmin=1500, vmax=6000)\n    axes[1].set_title('Target')\n    axes[1].axis('off')\n    plt.colorbar(im2, ax=axes[1])\n    \n    # Difference\n    diff = np.abs(pred - target)\n    im3 = axes[2].imshow(diff, cmap='hot')\n    axes[2].set_title(f'Absolute Error (MAE: {diff.mean():.1f})')\n    axes[2].axis('off')\n    plt.colorbar(im3, ax=axes[2])\n    \n    plt.tight_layout()\n    \n    if save_path:\n        plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.show()\n\n# ============== Training Configuration Template ==============\ntraining_config = {\n    \"model\": {\n        \"backbone\": \"tf_efficientnetv2_l\",\n        \"pretrained\": True,\n        \"in_channels\": 5,\n        \"out_channels\": 1,\n    },\n    \"training\": {\n        \"epochs\": 150,\n        \"batch_size\": 16,\n        \"learning_rate\": 1e-4,\n        \"weight_decay\": 1e-5,\n        \"scheduler\": \"CosineAnnealingWarmRestarts\",\n        \"warmup_epochs\": 5,\n    },\n    \"augmentation\": {\n        \"mixup\": True,\n        \"mixup_alpha\": 0.2,\n        \"physics_augment\": True,\n        \"tta_transforms\": 4,\n    },\n    \"physics\": {\n        \"use_physics_loss\": True,\n        \"physics_weight\": 0.1,\n        \"velocity_bounds\": [1500, 6000],\n        \"smoothness_weight\": 0.01,\n    },\n    \"ensemble\": {\n        \"n_folds\": 5,\n        \"voting\": \"weighted\",\n        \"post_process\": True,\n    }\n}\n\n# Save config\nwith open('training_config.json', 'w') as f:\n    json.dump(training_config, f, indent=4)\n\nprint(\"\\n✓ Training configuration saved to 'training_config.json'\")\nprint(\"\\nYou're all set! Happy training! 🚀\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-06-22T16:58:38.303428Z","iopub.execute_input":"2025-06-22T16:58:38.303681Z","iopub.status.idle":"2025-06-22T17:00:38.815324Z","shell.execute_reply.started":"2025-06-22T16:58:38.303660Z","shell.execute_reply":"2025-06-22T17:00:38.814598Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nFWI Competition - Data Explorer and Smart Loader\nThis script will help us understand the data structure and load it correctly\n\"\"\"\n\nimport os\nimport glob\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\n\n# ============== Step 1: Explore Data Structure ==============\ndef explore_competition_data():\n    \"\"\"Thoroughly explore the competition data structure\"\"\"\n    \n    base_path = \"/kaggle/input/waveform-inversion/\"\n    train_path = os.path.join(base_path, \"train_samples\")\n    \n    print(\"=\"*70)\n    print(\"FWI COMPETITION DATA EXPLORER\")\n    print(\"=\"*70)\n    \n    # Check if paths exist\n    print(f\"\\n1. Checking paths:\")\n    print(f\"   Base path exists: {os.path.exists(base_path)}\")\n    print(f\"   Train path exists: {os.path.exists(train_path)}\")\n    \n    if not os.path.exists(train_path):\n        print(\"\\nERROR: Training path doesn't exist!\")\n        return None\n    \n    # List all subdirectories\n    print(f\"\\n2. Subdirectories in {train_path}:\")\n    subdirs = [d for d in os.listdir(train_path) if os.path.isdir(os.path.join(train_path, d))]\n    for subdir in sorted(subdirs):\n        print(f\"   - {subdir}\")\n    \n    # Explore each subdirectory\n    all_files = []\n    print(f\"\\n3. Files in each subdirectory:\")\n    \n    for subdir in sorted(subdirs):\n        subdir_path = os.path.join(train_path, subdir)\n        files = sorted(glob.glob(os.path.join(subdir_path, \"*\")))\n        \n        print(f\"\\n   {subdir}:\")\n        print(f\"   Total files: {len(files)}\")\n        \n        # Show first few files\n        for i, f in enumerate(files[:3]):\n            filename = os.path.basename(f)\n            print(f\"     [{i+1}] {filename}\")\n            \n            # Check if it's an .npy file and analyze it\n            if f.endswith('.npy'):\n                try:\n                    # Load with memory mapping to check shape without loading full data\n                    data = np.load(f, mmap_mode='r')\n                    print(f\"         Shape: {data.shape}\")\n                    print(f\"         Dtype: {data.dtype}\")\n                    \n                    # Sample a small portion to check value range\n                    if len(data) > 0:\n                        sample = data[0]\n                        if hasattr(sample, 'shape'):\n                            print(f\"         Sample shape: {sample.shape}\")\n                            print(f\"         Value range: [{np.min(sample):.2f}, {np.max(sample):.2f}]\")\n                    \n                    all_files.append({\n                        'subdir': subdir,\n                        'filepath': f,\n                        'filename': filename,\n                        'shape': data.shape,\n                        'dtype': str(data.dtype)\n                    })\n                    \n                except Exception as e:\n                    print(f\"         Error loading: {str(e)}\")\n            \n        if len(files) > 3:\n            print(f\"     ... and {len(files) - 3} more files\")\n    \n    return all_files\n\n# ============== Step 2: Smart Pattern Detection ==============\ndef detect_file_patterns(all_files):\n    \"\"\"Detect patterns in the file naming and structure\"\"\"\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\"FILE PATTERN ANALYSIS\")\n    print(\"=\"*70)\n    \n    if not all_files:\n        print(\"No files to analyze!\")\n        return None\n    \n    # Convert to DataFrame for easier analysis\n    df = pd.DataFrame(all_files)\n    \n    # Analyze shapes\n    print(\"\\n1. Unique shapes found:\")\n    shape_counts = df['shape'].value_counts()\n    for shape, count in shape_counts.items():\n        print(f\"   {shape}: {count} files\")\n    \n    # Detect seismic vs velocity based on shape\n    seismic_files = []\n    velocity_files = []\n    \n    for _, row in df.iterrows():\n        shape = eval(str(row['shape']))  # Convert string back to tuple\n        \n        # Common patterns:\n        # Seismic: (N, 5, T, X) where 5 is channels, T is time, X is spatial\n        # Velocity: (N, X, Y) where X, Y are spatial dimensions\n        \n        if len(shape) == 4 and shape[1] == 5:\n            seismic_files.append(row)\n        elif len(shape) == 3 and shape[1] == shape[2]:  # Square spatial dimensions\n            velocity_files.append(row)\n        else:\n            # Try to infer from filename\n            if any(x in row['filename'].lower() for x in ['seismic', 'data', 'input', 'trace']):\n                seismic_files.append(row)\n            elif any(x in row['filename'].lower() for x in ['velocity', 'vel', 'model', 'label', 'target']):\n                velocity_files.append(row)\n    \n    print(f\"\\n2. Detected file types:\")\n    print(f\"   Seismic files: {len(seismic_files)}\")\n    print(f\"   Velocity files: {len(velocity_files)}\")\n    \n    # Show examples\n    if seismic_files:\n        print(f\"\\n   Example seismic file:\")\n        example = seismic_files[0]\n        print(f\"     File: {example['filename']}\")\n        print(f\"     Shape: {example['shape']}\")\n    \n    if velocity_files:\n        print(f\"\\n   Example velocity file:\")\n        example = velocity_files[0]\n        print(f\"     File: {example['filename']}\")\n        print(f\"     Shape: {example['shape']}\")\n    \n    return seismic_files, velocity_files\n\n# ============== Step 3: Create File Pairs ==============\ndef create_file_pairs(seismic_files, velocity_files):\n    \"\"\"Match seismic and velocity files into pairs\"\"\"\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\"CREATING FILE PAIRS\")\n    print(\"=\"*70)\n    \n    pairs = []\n    \n    # Group by subdirectory\n    seismic_by_dir = {}\n    velocity_by_dir = {}\n    \n    for s in seismic_files:\n        if s['subdir'] not in seismic_by_dir:\n            seismic_by_dir[s['subdir']] = []\n        seismic_by_dir[s['subdir']].append(s)\n    \n    for v in velocity_files:\n        if v['subdir'] not in velocity_by_dir:\n            velocity_by_dir[v['subdir']] = []\n        velocity_by_dir[v['subdir']].append(v)\n    \n    # Match files within each directory\n    for subdir in sorted(set(list(seismic_by_dir.keys()) + list(velocity_by_dir.keys()))):\n        seismic = seismic_by_dir.get(subdir, [])\n        velocity = velocity_by_dir.get(subdir, [])\n        \n        print(f\"\\n{subdir}: {len(seismic)} seismic, {len(velocity)} velocity files\")\n        \n        # Simple matching - assume they're in the same order\n        n_pairs = min(len(seismic), len(velocity))\n        for i in range(n_pairs):\n            pairs.append({\n                'dataset': subdir,\n                'seismic_file': seismic[i]['filepath'],\n                'velocity_file': velocity[i]['filepath'],\n                'seismic_shape': seismic[i]['shape'],\n                'velocity_shape': velocity[i]['shape']\n            })\n    \n    print(f\"\\nTotal pairs created: {len(pairs)}\")\n    return pairs\n\n# ============== Step 4: Alternative Loading Strategies ==============\ndef load_data_alternative():\n    \"\"\"Try alternative loading strategies if pattern matching fails\"\"\"\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\"TRYING ALTERNATIVE LOADING STRATEGIES\")\n    print(\"=\"*70)\n    \n    train_path = \"/kaggle/input/waveform-inversion/train_samples\"\n    all_npy_files = glob.glob(os.path.join(train_path, \"**/*.npy\"), recursive=True)\n    \n    print(f\"Found {len(all_npy_files)} total .npy files\")\n    \n    # Strategy 1: Check if files come in pairs (alternating)\n    print(\"\\nStrategy 1: Checking for alternating file pairs...\")\n    \n    pairs = []\n    datasets = ['CurveFault_A', 'CurveFault_B', 'CurveVel_A', 'CurveVel_B',\n                'FlatFault_A', 'FlatFault_B', 'FlatVel_A', 'FlatVel_B',\n                'Style_A', 'Style_B']\n    \n    for dataset in datasets:\n        dataset_files = [f for f in all_npy_files if dataset in f]\n        dataset_files.sort()\n        \n        if len(dataset_files) >= 2:\n            # Check first two files\n            f1_shape = np.load(dataset_files[0], mmap_mode='r').shape\n            f2_shape = np.load(dataset_files[1], mmap_mode='r').shape\n            \n            print(f\"\\n{dataset}:\")\n            print(f\"  File 1 shape: {f1_shape}\")\n            print(f\"  File 2 shape: {f2_shape}\")\n            \n            # Assume alternating pattern if shapes are different\n            if f1_shape != f2_shape:\n                for i in range(0, len(dataset_files)-1, 2):\n                    pairs.append({\n                        'dataset': dataset,\n                        'file1': dataset_files[i],\n                        'file2': dataset_files[i+1],\n                        'shape1': np.load(dataset_files[i], mmap_mode='r').shape,\n                        'shape2': np.load(dataset_files[i+1], mmap_mode='r').shape\n                    })\n    \n    return pairs\n\n# ============== Main Execution ==============\ndef main():\n    \"\"\"Main exploration function\"\"\"\n    \n    # Step 1: Explore data\n    all_files = explore_competition_data()\n    \n    if all_files:\n        # Step 2: Detect patterns\n        seismic_files, velocity_files = detect_file_patterns(all_files)\n        \n        # Step 3: Create pairs\n        if seismic_files and velocity_files:\n            pairs = create_file_pairs(seismic_files, velocity_files)\n            \n            # Save pairs information\n            if pairs:\n                pairs_df = pd.DataFrame(pairs)\n                pairs_df.to_csv('fwi_file_pairs.csv', index=False)\n                print(f\"\\nFile pairs saved to 'fwi_file_pairs.csv'\")\n                \n                # Show sample pair\n                print(\"\\nSample pair:\")\n                sample = pairs[0]\n                print(f\"  Dataset: {sample['dataset']}\")\n                print(f\"  Seismic: {os.path.basename(sample['seismic_file'])} {sample['seismic_shape']}\")\n                print(f\"  Velocity: {os.path.basename(sample['velocity_file'])} {sample['velocity_shape']}\")\n        else:\n            print(\"\\nCouldn't detect seismic/velocity files automatically.\")\n            print(\"Trying alternative strategies...\")\n            \n            # Try alternative loading\n            alt_pairs = load_data_alternative()\n            if alt_pairs:\n                print(f\"\\nFound {len(alt_pairs)} pairs using alternative strategy\")\n    else:\n        print(\"\\nNo files found! Please check the data path.\")\n    \n    # Additional diagnostics\n    print(\"\\n\" + \"=\"*70)\n    print(\"ADDITIONAL DIAGNOSTICS\")\n    print(\"=\"*70)\n    \n    # Check test data\n    test_path = \"/kaggle/input/waveform-inversion/test\"\n    if os.path.exists(test_path):\n        test_files = glob.glob(os.path.join(test_path, \"*.npy\"))\n        print(f\"\\nTest files: {len(test_files)}\")\n        if test_files:\n            test_sample = np.load(test_files[0], mmap_mode='r')\n            print(f\"Test file shape: {test_sample.shape}\")\n            print(f\"Test file dtype: {test_sample.dtype}\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-22T17:00:38.816610Z","iopub.execute_input":"2025-06-22T17:00:38.816848Z","iopub.status.idle":"2025-06-22T17:00:39.141958Z","shell.execute_reply.started":"2025-06-22T17:00:38.816830Z","shell.execute_reply":"2025-06-22T17:00:39.141174Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nFWI Competition - Simplified Effective Solution\nOptimized for P100 GPU with better preprocessing\n\"\"\"\n\nimport os\nimport gc\nimport glob\nimport random\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torch.optim.lr_scheduler import OneCycleLR\nfrom sklearn.model_selection import KFold\nfrom scipy.ndimage import gaussian_filter, median_filter\n\n# Configuration\nTRAIN = True\nPREDICT = True\nUSE_FULL_DATA = True\nENSEMBLE_FOLDS = 3\n\nclass Config:\n    train_path = \"/kaggle/input/waveform-inversion/train_samples/\"\n    test_path = \"/kaggle/input/waveform-inversion/test/\"\n    \n    # Model params\n    in_channels = 5\n    out_channels = 1\n    \n    # Training params\n    batch_size = 16\n    val_batch_size = 32\n    epochs = 30 if USE_FULL_DATA else 5\n    lr = 1e-4\n    weight_decay = 1e-4\n    \n    # Advanced params\n    use_mixup = True\n    mixup_alpha = 0.4\n    use_tta = True\n    n_folds = ENSEMBLE_FOLDS\n    \n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    seed = 42\n    num_workers = 0\n\ncfg = Config()\n\ndef set_seed(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nset_seed(cfg.seed)\n\n# ============== Data Loading ==============\ndef load_all_training_data():\n    \"\"\"Load all training data files\"\"\"\n    print(\"Loading training data...\")\n    file_pairs = []\n    \n    # Load .npy files\n    dirs_with_npy = ['CurveFault_A', 'CurveFault_B', 'FlatFault_A', 'FlatFault_B']\n    \n    for dataset in dirs_with_npy:\n        dataset_path = os.path.join(cfg.train_path, dataset)\n        seismic_files = sorted(glob.glob(os.path.join(dataset_path, \"seis*.npy\")))\n        velocity_files = sorted(glob.glob(os.path.join(dataset_path, \"vel*.npy\")))\n        \n        for vel_file in velocity_files:\n            vel_num = os.path.basename(vel_file).split('_')[0].replace('vel', '')\n            for seis_file in seismic_files:\n                seis_num = os.path.basename(seis_file).split('_')[0].replace('seis', '')\n                if seis_num == vel_num:\n                    file_pairs.append({\n                        'dataset': dataset,\n                        'seismic': seis_file,\n                        'velocity': vel_file\n                    })\n                    break\n    \n    # Load data/model folders\n    dirs_with_folders = ['CurveVel_A', 'CurveVel_B', 'FlatVel_A', 'FlatVel_B', 'Style_A', 'Style_B']\n    \n    for dataset in dirs_with_folders:\n        dataset_path = os.path.join(cfg.train_path, dataset)\n        data_path = os.path.join(dataset_path, \"data\")\n        model_path = os.path.join(dataset_path, \"model\")\n        \n        if os.path.isdir(data_path) and os.path.isdir(model_path):\n            data_files = sorted(glob.glob(os.path.join(data_path, \"*.npy\")))\n            model_files = sorted(glob.glob(os.path.join(model_path, \"*.npy\")))\n            for df, mf in zip(data_files, model_files):\n                file_pairs.append({\n                    'dataset': dataset,\n                    'seismic': df,\n                    'velocity': mf\n                })\n        elif os.path.isfile(data_path) and os.path.isfile(model_path):\n            file_pairs.append({\n                'dataset': dataset,\n                'seismic': data_path,\n                'velocity': model_path\n            })\n    \n    print(f\"Total file pairs: {len(file_pairs)}\")\n    return file_pairs\n\n# ============== Dataset with Better Processing ==============\nclass FWIDataset(torch.utils.data.Dataset):\n    def __init__(self, file_pairs, indices=None, mode='train'):\n        self.file_pairs = file_pairs\n        self.indices = indices if indices is not None else range(len(file_pairs))\n        self.mode = mode\n        self.files = [file_pairs[i] for i in self.indices]\n        self.samples_per_file = 500\n        \n    def __len__(self):\n        return len(self.files) * self.samples_per_file\n    \n    def __getitem__(self, idx):\n        file_idx = idx // self.samples_per_file\n        sample_idx = idx % self.samples_per_file\n        \n        file_info = self.files[file_idx % len(self.files)]\n        \n        # Load data\n        seismic_data = np.load(file_info['seismic'], mmap_mode='r')\n        velocity_data = np.load(file_info['velocity'], mmap_mode='r')\n        \n        # Get sample\n        if sample_idx < len(seismic_data):\n            seismic = seismic_data[sample_idx]\n            velocity = velocity_data[sample_idx]\n        else:\n            seismic = seismic_data[-1]\n            velocity = velocity_data[-1]\n        \n        # Make copies\n        seismic = np.array(seismic, copy=True, dtype=np.float32)\n        velocity = np.array(velocity, copy=True, dtype=np.float32)\n        \n        # Process\n        seismic = self._process_seismic(seismic)\n        \n        if len(velocity.shape) == 3 and velocity.shape[0] == 1:\n            velocity = velocity[0]\n        \n        # Scale velocity to [0, 1] range for training stability\n        velocity = (velocity - 1500) / 4500\n        \n        if self.mode == 'train':\n            seismic, velocity = self._augment(seismic, velocity)\n        \n        seismic = np.ascontiguousarray(seismic)\n        velocity = np.ascontiguousarray(velocity)\n        \n        return torch.from_numpy(seismic).float(), torch.from_numpy(velocity).float().unsqueeze(0)\n    \n    def _process_seismic(self, data):\n        \"\"\"Advanced seismic processing\"\"\"\n        if data.shape[0] == 5 and data.shape[1] == 1000:\n            # Multi-scale temporal sampling\n            # Early reflections (0-300ms)\n            early_idx = np.linspace(0, 300, 35, dtype=int)\n            # Middle reflections (300-700ms)\n            middle_idx = np.linspace(300, 700, 35, dtype=int)\n            \n            # Combine for 70 time samples\n            indices = np.concatenate([early_idx, middle_idx])\n            processed = data[:, indices, :]\n            \n            # Alternative: use frequency-based sampling\n            # This captures both high and low frequency content\n            # freq_indices = np.array([i for i in range(0, 1000, 14)])[:70]\n            # processed = data[:, freq_indices, :]\n        else:\n            processed = data\n        \n        # Robust normalization per channel\n        for i in range(processed.shape[0]):\n            channel = processed[i]\n            # Use median and MAD for robustness\n            median = np.median(channel)\n            mad = np.median(np.abs(channel - median))\n            processed[i] = (channel - median) / (mad * 1.4826 + 1e-8)\n            processed[i] = np.clip(processed[i], -5, 5)\n        \n        return processed\n    \n    def _augment(self, seismic, velocity):\n        \"\"\"Data augmentation\"\"\"\n        # Horizontal flip\n        if np.random.random() < 0.5:\n            seismic = seismic[:, :, ::-1]\n            velocity = velocity[:, ::-1]\n        \n        # Add noise\n        if np.random.random() < 0.3:\n            noise = np.random.normal(0, 0.05, seismic.shape)\n            seismic = seismic + noise\n        \n        # Random gain\n        if np.random.random() < 0.3:\n            gain = np.random.uniform(0.9, 1.1)\n            seismic = seismic * gain\n        \n        return seismic, velocity\n\n# ============== Effective Model ==============\nclass EfficientUNet(nn.Module):\n    def __init__(self, in_channels=5, out_channels=1):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = self._double_conv(in_channels, 64)\n        self.pool1 = nn.MaxPool2d(2)\n        \n        self.enc2 = self._double_conv(64, 128)\n        self.pool2 = nn.MaxPool2d(2)\n        \n        self.enc3 = self._double_conv(128, 256)\n        self.pool3 = nn.MaxPool2d(2)\n        \n        self.enc4 = self._double_conv(256, 512)\n        self.pool4 = nn.MaxPool2d(2)\n        \n        # Bottleneck\n        self.bottleneck = self._double_conv(512, 1024)\n        \n        # Decoder\n        self.up4 = nn.Conv2d(1024, 512, 1)\n        self.dec4 = self._double_conv(1024, 512)\n        \n        self.up3 = nn.Conv2d(512, 256, 1)\n        self.dec3 = self._double_conv(512, 256)\n        \n        self.up2 = nn.Conv2d(256, 128, 1)\n        self.dec2 = self._double_conv(256, 128)\n        \n        self.up1 = nn.Conv2d(128, 64, 1)\n        self.dec1 = self._double_conv(128, 64)\n        \n        # Output\n        self.out = nn.Conv2d(64, out_channels, 1)\n        \n    def _double_conv(self, in_ch, out_ch):\n        return nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Encoder\n        e1 = self.enc1(x)\n        e2 = self.enc2(self.pool1(e1))\n        e3 = self.enc3(self.pool2(e2))\n        e4 = self.enc4(self.pool3(e3))\n        \n        # Bottleneck\n        b = self.bottleneck(self.pool4(e4))\n        \n        # Decoder\n        d4 = self.up4(F.interpolate(b, size=e4.shape[2:], mode='bilinear', align_corners=False))\n        d4 = self.dec4(torch.cat([d4, e4], dim=1))\n        \n        d3 = self.up3(F.interpolate(d4, size=e3.shape[2:], mode='bilinear', align_corners=False))\n        d3 = self.dec3(torch.cat([d3, e3], dim=1))\n        \n        d2 = self.up2(F.interpolate(d3, size=e2.shape[2:], mode='bilinear', align_corners=False))\n        d2 = self.dec2(torch.cat([d2, e2], dim=1))\n        \n        d1 = self.up1(F.interpolate(d2, size=e1.shape[2:], mode='bilinear', align_corners=False))\n        d1 = self.dec1(torch.cat([d1, e1], dim=1))\n        \n        # Output - sigmoid to ensure [0, 1] range\n        out = torch.sigmoid(self.out(d1))\n        \n        return out\n\n# ============== Loss Function ==============\nclass CombinedLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.mae = nn.L1Loss()\n        self.mse = nn.MSELoss()\n        \n    def forward(self, pred, target):\n        # Scale back to velocity range for loss calculation\n        pred_vel = pred * 4500 + 1500\n        target_vel = target * 4500 + 1500\n        \n        # MAE loss\n        loss_mae = self.mae(pred_vel, target_vel)\n        \n        # RMSE loss\n        loss_rmse = torch.sqrt(self.mse(pred_vel, target_vel) + 1e-8)\n        \n        # Gradient loss for smoothness\n        pred_dx = pred[:, :, :, 1:] - pred[:, :, :, :-1]\n        target_dx = target[:, :, :, 1:] - target[:, :, :, :-1]\n        pred_dy = pred[:, :, 1:, :] - pred[:, :, :-1, :]\n        target_dy = target[:, :, 1:, :] - target[:, :, :-1, :]\n        \n        grad_loss = F.l1_loss(pred_dx, target_dx) + F.l1_loss(pred_dy, target_dy)\n        \n        # Combined loss\n        total_loss = loss_mae + 0.1 * loss_rmse + 0.1 * grad_loss\n        \n        return total_loss\n\n# ============== Training ==============\ndef train_epoch(model, loader, criterion, optimizer, scaler, cfg):\n    model.train()\n    losses = []\n    \n    for data, target in tqdm(loader, desc='Training'):\n        data, target = data.to(cfg.device), target.to(cfg.device)\n        \n        # Mixup\n        if cfg.use_mixup and np.random.random() < 0.5:\n            lam = np.random.beta(cfg.mixup_alpha, cfg.mixup_alpha)\n            idx = torch.randperm(data.size(0)).to(cfg.device)\n            data = lam * data + (1 - lam) * data[idx]\n            target = lam * target + (1 - lam) * target[idx]\n        \n        optimizer.zero_grad()\n        \n        with autocast():\n            output = model(data)\n            loss = criterion(output, target)\n        \n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        \n        losses.append(loss.item())\n    \n    return np.mean(losses)\n\ndef validate(model, loader, cfg):\n    model.eval()\n    losses = []\n    \n    with torch.no_grad():\n        for data, target in tqdm(loader, desc='Validation'):\n            data, target = data.to(cfg.device), target.to(cfg.device)\n            output = model(data)\n            \n            # Calculate MAE in velocity space\n            pred_vel = output * 4500 + 1500\n            target_vel = target * 4500 + 1500\n            loss = F.l1_loss(pred_vel, target_vel)\n            \n            losses.append(loss.item())\n    \n    return np.mean(losses)\n\ndef train_fold(cfg, file_pairs, fold):\n    print(f\"\\nTraining Fold {fold+1}/{cfg.n_folds}\")\n    \n    # Split data\n    kf = KFold(n_splits=cfg.n_folds, shuffle=True, random_state=cfg.seed)\n    train_idx, val_idx = list(kf.split(file_pairs))[fold]\n    \n    # Datasets\n    train_dataset = FWIDataset(file_pairs, train_idx, mode='train')\n    val_dataset = FWIDataset(file_pairs, val_idx, mode='val')\n    \n    # Loaders\n    train_loader = torch.utils.data.DataLoader(\n        train_dataset, batch_size=cfg.batch_size, shuffle=True,\n        num_workers=cfg.num_workers, pin_memory=True\n    )\n    val_loader = torch.utils.data.DataLoader(\n        val_dataset, batch_size=cfg.val_batch_size, shuffle=False,\n        num_workers=cfg.num_workers, pin_memory=True\n    )\n    \n    # Model\n    model = EfficientUNet().to(cfg.device)\n    \n    # Optimizer and scheduler\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n    scheduler = OneCycleLR(\n        optimizer, max_lr=cfg.lr * 10, epochs=cfg.epochs,\n        steps_per_epoch=len(train_loader), pct_start=0.1\n    )\n    \n    # Loss and scaler\n    criterion = CombinedLoss()\n    scaler = GradScaler()\n    \n    # Training\n    best_loss = float('inf')\n    for epoch in range(cfg.epochs):\n        print(f\"\\nEpoch {epoch+1}/{cfg.epochs}\")\n        \n        train_loss = train_epoch(model, train_loader, criterion, optimizer, scaler, cfg)\n        val_loss = validate(model, val_loader, cfg)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.1f}, Val Loss: {val_loss:.1f}\")\n        \n        if val_loss < best_loss:\n            best_loss = val_loss\n            torch.save(model.state_dict(), f'model_fold{fold}.pth')\n            print(f\"Saved best model (loss: {val_loss:.1f})\")\n    \n    # Load best model\n    model.load_state_dict(torch.load(f'model_fold{fold}.pth', weights_only=False))\n    return model\n\n# ============== Inference ==============\ndef predict_test(models, cfg):\n    test_files = sorted(glob.glob(f\"{cfg.test_path}/*.npy\"))\n    print(f\"\\nPredicting {len(test_files)} test files...\")\n    \n    submission_rows = []\n    \n    for test_file in tqdm(test_files):\n        test_data = np.load(test_file)\n        oid = os.path.basename(test_file).replace('.npy', '')\n        \n        # Process\n        indices = np.concatenate([\n            np.linspace(0, 300, 35, dtype=int),\n            np.linspace(300, 700, 35, dtype=int)\n        ])\n        processed = test_data[:, indices, :].astype(np.float32)\n        \n        # Normalize\n        for i in range(5):\n            channel = processed[i]\n            median = np.median(channel)\n            mad = np.median(np.abs(channel - median))\n            processed[i] = (channel - median) / (mad * 1.4826 + 1e-8)\n            processed[i] = np.clip(processed[i], -5, 5)\n        \n        data = torch.from_numpy(processed).float().unsqueeze(0).to(cfg.device)\n        \n        # Predict with ensemble\n        predictions = []\n        for model in models:\n            model.eval()\n            with torch.no_grad():\n                if cfg.use_tta:\n                    # TTA\n                    preds = []\n                    preds.append(model(data))\n                    preds.append(torch.flip(model(torch.flip(data, dims=[3])), dims=[3]))\n                    pred = torch.stack(preds).mean(0)\n                else:\n                    pred = model(data)\n            predictions.append(pred)\n        \n        # Average and convert to velocity\n        prediction = torch.stack(predictions).mean(0).squeeze().cpu().numpy()\n        prediction = prediction * 4500 + 1500\n        \n        # Post-process\n        prediction = np.clip(prediction, 1500, 6000)\n        prediction = median_filter(prediction, size=3)\n        \n        # Create submission\n        for y in range(70):\n            row = {'oid_ypos': f'{oid}_y_{y}'}\n            for x in range(1, 70, 2):\n                row[f'x_{x}'] = float(prediction[y, x])\n            submission_rows.append(row)\n    \n    return pd.DataFrame(submission_rows)\n\n# ============== Main ==============\nif __name__ == \"__main__\":\n    print(\"FWI Simple Effective Solution\")\n    print(f\"Device: {cfg.device}\")\n    \n    # Load data\n    file_pairs = load_all_training_data()\n    \n    # Train\n    models = []\n    if TRAIN:\n        for fold in range(cfg.n_folds):\n            model = train_fold(cfg, file_pairs, fold)\n            models.append(model)\n            gc.collect()\n            torch.cuda.empty_cache()\n    else:\n        for fold in range(cfg.n_folds):\n            model = EfficientUNet().to(cfg.device)\n            model.load_state_dict(torch.load(f'model_fold{fold}.pth', weights_only=False))\n            models.append(model)\n    \n    # Predict\n    if PREDICT:\n        submission = predict_test(models, cfg)\n        submission.to_csv('submission.csv', index=False)\n        print(f\"\\nSubmission saved! Shape: {submission.shape}\")\n        print(submission.head())\n    \n    print(\"\\n✅ Done!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-22T18:08:44.636515Z","iopub.execute_input":"2025-06-22T18:08:44.637191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nAdvanced utilities for physics-informed FWI solution\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nfrom scipy import signal\nfrom scipy.ndimage import gaussian_filter, median_filter\nimport pywt  # for wavelet transforms\n\n# ============== Physics-Informed Components ==============\n\nclass WaveEquationLoss(nn.Module):\n    \"\"\"Full wave equation loss for physics-informed training\"\"\"\n    def __init__(self, dx=1.0, dt=0.001):\n        super().__init__()\n        self.dx = dx\n        self.dt = dt\n        \n        # Define finite difference operators\n        self.register_buffer('laplacian_x', self._create_fd_kernel('xx'))\n        self.register_buffer('laplacian_y', self._create_fd_kernel('yy'))\n        self.register_buffer('time_diff', self._create_fd_kernel('tt'))\n        \n    def _create_fd_kernel(self, derivative_type):\n        \"\"\"Create finite difference kernels\"\"\"\n        if derivative_type == 'xx' or derivative_type == 'yy':\n            # Second derivative in space\n            kernel = torch.tensor([\n                [0, 0, 0],\n                [1, -2, 1],\n                [0, 0, 0]\n            ], dtype=torch.float32)\n            if derivative_type == 'yy':\n                kernel = kernel.T\n        elif derivative_type == 'tt':\n            # Second derivative in time (simplified)\n            kernel = torch.tensor([1, -2, 1], dtype=torch.float32)\n        \n        return kernel.unsqueeze(0).unsqueeze(0)\n    \n    def forward(self, velocity, wavefield=None):\n        \"\"\"\n        Compute wave equation residual:\n        ∂²u/∂t² = v²(∇²u) + f\n        \"\"\"\n        b, c, h, w = velocity.shape\n        \n        # Compute spatial derivatives\n        d2u_dx2 = F.conv2d(velocity, self.laplacian_x, padding=1) / (self.dx ** 2)\n        d2u_dy2 = F.conv2d(velocity, self.laplacian_y, padding=1) / (self.dx ** 2)\n        laplacian_u = d2u_dx2 + d2u_dy2\n        \n        # Wave equation residual (simplified without time component)\n        # In practice, you'd need the full wavefield evolution\n        wave_residual = torch.abs(laplacian_u)\n        \n        # Additional physics constraints\n        # 1. Smoothness constraint (velocities should be locally smooth)\n        grad_x = torch.abs(velocity[:, :, :, 1:] - velocity[:, :, :, :-1])\n        grad_y = torch.abs(velocity[:, :, 1:, :] - velocity[:, :, :-1, :])\n        smoothness_loss = (grad_x.mean() + grad_y.mean()) * 0.1\n        \n        # 2. Boundary conditions (e.g., absorbing boundaries)\n        boundary_loss = (\n            velocity[:, :, :5, :].abs().mean() + \n            velocity[:, :, -5:, :].abs().mean() +\n            velocity[:, :, :, :5].abs().mean() + \n            velocity[:, :, :, -5:].abs().mean()\n        ) * 0.01\n        \n        return wave_residual.mean() + smoothness_loss + boundary_loss\n\n\nclass FourierFeatures(nn.Module):\n    \"\"\"Fourier feature extraction for better frequency representation\"\"\"\n    def __init__(self, in_channels, out_channels, modes=16):\n        super().__init__()\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.modes = modes\n        \n        # Fourier coefficients\n        self.scale = 1 / (in_channels * out_channels)\n        self.weights = nn.Parameter(\n            self.scale * torch.rand(in_channels, out_channels, modes, modes, 2)\n        )\n        \n    def forward(self, x):\n        # FFT\n        x_ft = torch.fft.rfft2(x)\n        \n        # Multiply relevant Fourier modes\n        out_ft = torch.zeros(\n            x.shape[0], self.out_channels, x.shape[-2], x.shape[-1]//2 + 1,\n            dtype=torch.cfloat, device=x.device\n        )\n        \n        out_ft[:, :, :self.modes, :self.modes] = self._complex_mul2d(\n            x_ft[:, :, :self.modes, :self.modes], \n            torch.view_as_complex(self.weights)\n        )\n        \n        # IFFT\n        x = torch.fft.irfft2(out_ft, s=(x.shape[-2], x.shape[-1]))\n        return x\n    \n    def _complex_mul2d(self, input, weights):\n        \"\"\"Complex multiplication\"\"\"\n        return torch.einsum(\"bixy,ioxy->boxy\", input, weights)\n\n\nclass MultiScaleWaveletTransform(nn.Module):\n    \"\"\"Multi-scale wavelet decomposition for seismic data\"\"\"\n    def __init__(self, wavelet='db4', levels=3):\n        super().__init__()\n        self.wavelet = wavelet\n        self.levels = levels\n        \n    def forward(self, x):\n        \"\"\"Apply wavelet transform and return multi-scale features\"\"\"\n        b, c, h, w = x.shape\n        features = []\n        \n        for i in range(b):\n            for j in range(c):\n                # 2D wavelet transform\n                coeffs = pywt.wavedec2(x[i, j].cpu().numpy(), self.wavelet, level=self.levels)\n                \n                # Extract features at each scale\n                for level_coeffs in coeffs:\n                    if isinstance(level_coeffs, tuple):\n                        for coeff in level_coeffs:\n                            features.append(torch.from_numpy(coeff).to(x.device))\n                    else:\n                        features.append(torch.from_numpy(level_coeffs).to(x.device))\n        \n        return features\n\n\n# ============== Advanced Data Augmentation ==============\n\nclass SeismicAugmentation:\n    \"\"\"Physics-consistent augmentations for seismic data\"\"\"\n    \n    @staticmethod\n    def add_coherent_noise(data, snr_db=20):\n        \"\"\"Add coherent noise (ground roll, multiples)\"\"\"\n        signal_power = np.mean(data ** 2)\n        noise_power = signal_power / (10 ** (snr_db / 10))\n        \n        # Generate coherent noise (e.g., linear events)\n        noise = np.zeros_like(data)\n        for _ in range(np.random.randint(1, 4)):\n            # Random linear event\n            slope = np.random.uniform(-0.5, 0.5)\n            for i in range(data.shape[0]):\n                for j in range(data.shape[2]):\n                    t_idx = int(i + slope * j)\n                    if 0 <= t_idx < data.shape[0]:\n                        noise[t_idx, :, j] += np.random.normal(0, np.sqrt(noise_power))\n        \n        return data + noise\n    \n    @staticmethod\n    def apply_frequency_filter(data, low_freq=5, high_freq=50, fs=1000):\n        \"\"\"Apply bandpass filter\"\"\"\n        nyquist = fs / 2\n        low = low_freq / nyquist\n        high = high_freq / nyquist\n        \n        b, a = signal.butter(4, [low, high], btype='band')\n        filtered = signal.filtfilt(b, a, data, axis=0)\n        \n        return filtered\n    \n    @staticmethod\n    def random_static_shift(data, max_shift=5):\n        \"\"\"Apply random static shifts to traces\"\"\"\n        shifts = np.random.randint(-max_shift, max_shift, size=data.shape[2])\n        shifted = np.zeros_like(data)\n        \n        for i, shift in enumerate(shifts):\n            if shift > 0:\n                shifted[shift:, :, i] = data[:-shift, :, i]\n            elif shift < 0:\n                shifted[:shift, :, i] = data[-shift:, :, i]\n            else:\n                shifted[:, :, i] = data[:, :, i]\n                \n        return shifted\n    \n    @staticmethod\n    def apply_agc(data, window_size=50):\n        \"\"\"Apply Automatic Gain Control\"\"\"\n        eps = 1e-10\n        agc_data = np.zeros_like(data)\n        \n        for i in range(data.shape[2]):\n            trace = data[:, 0, i]\n            \n            # Compute envelope\n            envelope = np.abs(signal.hilbert(trace))\n            \n            # Smooth envelope\n            envelope = gaussian_filter(envelope, window_size/4)\n            \n            # Apply AGC\n            agc_data[:, 0, i] = trace / (envelope + eps)\n            \n        return agc_data\n\n\n# ============== Advanced Post-Processing ==============\n\nclass VelocityPostProcessor:\n    \"\"\"Post-processing for velocity models\"\"\"\n    \n    @staticmethod\n    def apply_physical_constraints(velocity, vmin=1500, vmax=6000):\n        \"\"\"Apply physical constraints to velocity model\"\"\"\n        # Clip to physical bounds\n        velocity = np.clip(velocity, vmin, vmax)\n        \n        # Apply median filter to remove spikes\n        velocity = median_filter(velocity, size=3)\n        \n        # Apply slight Gaussian smoothing\n        velocity = gaussian_filter(velocity, sigma=0.5)\n        \n        return velocity\n    \n    @staticmethod\n    def enforce_layer_continuity(velocity, threshold=200):\n        \"\"\"Enforce geological layer continuity\"\"\"\n        # Detect large velocity contrasts\n        grad_y = np.abs(np.diff(velocity, axis=0))\n        \n        # Find layer boundaries\n        boundaries = grad_y > threshold\n        \n        # Smooth within layers\n        smoothed = velocity.copy()\n        for i in range(1, velocity.shape[0]-1):\n            if not boundaries[i-1].any() and not boundaries[i].any():\n                # Average with neighbors if not at boundary\n                smoothed[i] = 0.6 * velocity[i] + 0.2 * velocity[i-1] + 0.2 * velocity[i+1]\n                \n        return smoothed\n    \n    @staticmethod\n    def apply_geological_priors(velocity, depth_axis=0):\n        \"\"\"Apply geological priors (velocity generally increases with depth)\"\"\"\n        # Sort velocity columns to ensure general increase with depth\n        for j in range(velocity.shape[1]):\n            col = velocity[:, j]\n            \n            # Apply slight increasing trend\n            trend = np.linspace(0, 100, len(col))\n            col_sorted = np.sort(col) + trend\n            \n            # Blend original with sorted (preserve features while enforcing trend)\n            velocity[:, j] = 0.7 * col + 0.3 * col_sorted\n            \n        return velocity\n\n\n# ============== Model Architecture Improvements ==============\n\nclass FNOBlock(nn.Module):\n    \"\"\"Fourier Neural Operator block for efficient convolution\"\"\"\n    def __init__(self, in_channels, out_channels, modes=16):\n        super().__init__()\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.modes = modes\n        \n        # Fourier layer\n        self.fourier = FourierFeatures(in_channels, out_channels, modes)\n        \n        # Regular convolution path\n        self.conv = nn.Conv2d(in_channels, out_channels, 1)\n        \n        # Activation and normalization\n        self.bn = nn.BatchNorm2d(out_channels)\n        self.activation = nn.GELU()\n        \n    def forward(self, x):\n        # Fourier path\n        x_fourier = self.fourier(x)\n        \n        # Conv path\n        x_conv = self.conv(x)\n        \n        # Combine\n        x = x_fourier + x_conv\n        x = self.bn(x)\n        x = self.activation(x)\n        \n        return x\n\n\nclass SelfAttention2D(nn.Module):\n    \"\"\"Self-attention module for global context\"\"\"\n    def __init__(self, in_channels, reduction=8):\n        super().__init__()\n        self.in_channels = in_channels\n        \n        self.query = nn.Conv2d(in_channels, in_channels // reduction, 1)\n        self.key = nn.Conv2d(in_channels, in_channels // reduction, 1)\n        self.value = nn.Conv2d(in_channels, in_channels, 1)\n        \n        self.gamma = nn.Parameter(torch.zeros(1))\n        \n    def forward(self, x):\n        b, c, h, w = x.shape\n        \n        # Compute attention\n        proj_query = self.query(x).view(b, -1, h*w).permute(0, 2, 1)\n        proj_key = self.key(x).view(b, -1, h*w)\n        energy = torch.bmm(proj_query, proj_key)\n        attention = F.softmax(energy, dim=-1)\n        \n        proj_value = self.value(x).view(b, -1, h*w)\n        out = torch.bmm(proj_value, attention.permute(0, 2, 1))\n        out = out.view(b, c, h, w)\n        \n        # Apply attention with learnable weight\n        out = self.gamma * out + x\n        \n        return out\n\n\n# ============== Training Utilities ==============\n\nclass EarlyStopping:\n    \"\"\"Early stopping with patience\"\"\"\n    def __init__(self, patience=10, min_delta=0, mode='min'):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.mode = mode\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n        \n    def __call__(self, score):\n        if self.best_score is None:\n            self.best_score = score\n        elif self.mode == 'min' and score > self.best_score - self.min_delta:\n            self.counter += 1\n            if self.counter >= self.patience:\n                self.early_stop = True\n        elif self.mode == 'max' and score < self.best_score + self.min_delta:\n            self.counter += 1\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_score = score\n            self.counter = 0\n            \n        return self.early_stop\n\n\nclass ModelCheckpoint:\n    \"\"\"Save best models during training\"\"\"\n    def __init__(self, filepath, monitor='val_loss', mode='min', save_best_only=True):\n        self.filepath = filepath\n        self.monitor = monitor\n        self.mode = mode\n        self.save_best_only = save_best_only\n        self.best = float('inf') if mode == 'min' else float('-inf')\n        \n    def __call__(self, score, model, epoch):\n        if self.mode == 'min' and score < self.best:\n            self.best = score\n            self._save_model(model, epoch, score)\n        elif self.mode == 'max' and score > self.best:\n            self.best = score\n            self._save_model(model, epoch, score)\n        elif not self.save_best_only:\n            self._save_model(model, epoch, score)\n            \n    def _save_model(self, model, epoch, score):\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'score': score,\n        }, self.filepath.format(epoch=epoch, score=score))\n\n\n# ============== Evaluation Metrics ==============\n\ndef compute_metrics(pred, target):\n    \"\"\"Compute various metrics for evaluation\"\"\"\n    mae = F.l1_loss(pred, target)\n    mse = F.mse_loss(pred, target)\n    \n    # Relative error\n    relative_error = torch.abs(pred - target) / (torch.abs(target) + 1e-8)\n    mre = relative_error.mean()\n    \n    # Structural similarity (simplified)\n    ssim = 1 - F.l1_loss(pred, target) / (pred.abs().mean() + target.abs().mean())\n    \n    return {\n        'mae': mae.item(),\n        'mse': mse.item(),\n        'rmse': torch.sqrt(mse).item(),\n        'mre': mre.item(),\n        'ssim': ssim.item()\n    }\n\n\nif __name__ == \"__main__\":\n    print(\"Advanced FWI utilities loaded successfully!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-22T17:00:46.423715Z","iopub.status.idle":"2025-06-22T17:00:46.423983Z","shell.execute_reply.started":"2025-06-22T17:00:46.423877Z","shell.execute_reply":"2025-06-22T17:00:46.423887Z"}},"outputs":[],"execution_count":null}]}