{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Realistic Simulation of Test Noise\n\nAs discovered in https://www.kaggle.com/code/vslaykovsky/g2net-winning-strategy-with-external-data, test noise comes in 2 flavors:\n* ~80% of test samples use generated stationary Gaussian noise.\n* ~20% of test samples borrow samples from real detectors.\n\nIn this notebook we generate noise that closely mimics samples from the second group. For each source test image this is done in the following steps:\n1. Bucketize SFT data into 256 time buckets in order to get 360x256 images in the end.\n2. For each of 256 time buckets we generate noise with the same amplitude as in the source spectrogram. e.g. np.std(test_sft[:, i]) == np.std(generated_sft[:, i]) (See `bucketize_real_noise_asd`, `simulate_real_noise`). This produces noise spectrograms that closely track non-stationary patterns of source spectrograms.\n3. We find persistent monochromatic detector artefacts (horizontal lines) in source spectrograms and reproduce them in the target spectrogram. Some custom logic is used here as pyfstat `LineWriter` has very limited functionality.\n\n# Why is it needed for the competition?\n\nThe 600 test set spectrograms only contain generated data, so if you want to improve performance of your model on the remaining 20% of samples collected from real detectors, you need spectrogram patterns that come from a similar distribution.\n\n\n**Consider upvoting if you have a mouse and you feel like clicking ▲ buttons today!**","metadata":{}},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"IS_KAGGLE = False\ntry:\n    import kaggle_secrets\n    G2NET_ROOT = '/kaggle/input/g2net-detecting-continuous-gravitational-waves'\n    TEST_CSV = '/kaggle/input/g2net-winning-strategy-with-external-data/test.csv'\n    IS_KAGGLE = True\n    !pip install git+https://github.com/PyFstat/PyFstat@python37\nexcept Exception as ex:\n    G2NET_ROOT = '/mnt/g2net'\n    TEST_CSV = 'test.csv'\n","metadata":{"execution":{"iopub.status.busy":"2022-12-10T23:07:03.993151Z","iopub.execute_input":"2022-12-10T23:07:03.993626Z","iopub.status.idle":"2022-12-10T23:07:43.623031Z","shell.execute_reply.started":"2022-12-10T23:07:03.993537Z","shell.execute_reply":"2022-12-10T23:07:43.621971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import copy\nimport glob\nimport logging\nimport multiprocessing\nimport os\nfrom pathlib import Path\n\nimport h5py\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport pyfstat\nfrom matplotlib import colors\nfrom pyfstat.utils import get_sft_as_arrays\nfrom scipy import stats\nfrom tqdm import tqdm\n\nloggers = [logging.getLogger(name) for name in logging.root.manager.loggerDict]\nfor logger in loggers:\n    logger.setLevel(logging.WARNING)\n\nC_SQRSX = 26.5\n\n\nTEST_IMAGE = f'{G2NET_ROOT}/test/56b090eaf.hdf5'\nDEBUG = True\nBUCKETS = 256\nNOISE_PATH = 'data/realistic_noise/'\n\nplt.rcParams[\"font.family\"] = \"serif\"\nplt.rcParams[\"font.size\"] = 20\nmatplotlib.rcParams['figure.figsize'] = (16, 10)\nNPY_TO_PNG_NORM = 2e-22\n\n%matplotlib inline","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:07:43.624665Z","iopub.execute_input":"2022-12-10T23:07:43.624956Z","iopub.status.idle":"2022-12-10T23:07:46.384398Z","shell.execute_reply.started":"2022-12-10T23:07:43.624929Z","shell.execute_reply":"2022-12-10T23:07:46.383328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_amplitude_phase_spectrograms(timestamps, frequency, fourier_data):\n    fig, axs = plt.subplots(1, 2, figsize=(16, 10))\n\n    for ax in axs:\n        ax.set(xlabel=\"SFT index\", ylabel=\"Frequency [Hz]\")\n\n    time_in_days = (timestamps - timestamps[0]) / 1800\n\n    axs[0].set_title(\"SFT absolute value\")\n    c = axs[0].pcolorfast(\n        time_in_days, frequency, np.absolute(fourier_data.astype(np.complex128))[:fourier_data.shape[0]-1, :fourier_data.shape[1]-1], norm=colors.Normalize(0, 1e-21), cmap='gray'\n    )\n    fig.colorbar(c, ax=axs[0], orientation=\"horizontal\", label=\"Value\")\n\n    axs[1].set_title(\"SFT phase\")\n    c = axs[1].pcolorfast(\n        time_in_days, frequency, np.angle(fourier_data)[:fourier_data.shape[0]-1, :fourier_data.shape[1]-1], norm=colors.CenteredNorm(), cmap='gray'\n    )\n\n    fig.colorbar(c, ax=axs[1], orientation=\"horizontal\", label=\"Value\")\n\n    return fig, axs\n\n\ndef plot_amplitude_spectrogram(timestamps, frequency, fourier_data, ax=None):\n    if ax is None:\n        fig, ax = plt.subplots(1, 1, figsize=(8, 10))\n    else:\n        fig = ax.get_figure()\n\n    ax.set(xlabel=\"SFT index\", ylabel=\"Frequency [Hz]\")\n\n    time_in_days = (timestamps - timestamps[0]) / 1800\n\n    ax.set_title(\"SFT absolute value\")\n    c = ax.pcolorfast(\n        time_in_days, frequency, np.absolute(fourier_data.astype(np.complex128))[:fourier_data.shape[0]-1, :fourier_data.shape[1]-1], norm=colors.Normalize(0, 1e-21), cmap='gray'\n    )\n    fig.colorbar(c, ax=ax, orientation=\"horizontal\", label=\"Value\")\n\n","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:07:46.385897Z","iopub.execute_input":"2022-12-10T23:07:46.386079Z","iopub.status.idle":"2022-12-10T23:07:46.398318Z","shell.execute_reply.started":"2022-12-10T23:07:46.386056Z","shell.execute_reply":"2022-12-10T23:07:46.397183Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_data(file):\n    file = Path(file)\n    with h5py.File(file, \"r\") as f:\n        filename = file.stem\n        f = f[filename]\n        h1 = f[\"H1\"]\n        l1 = f[\"L1\"]\n        freq_hz = list(f[\"frequency_Hz\"])\n\n        h1_stft = h1[\"SFTs\"][()]\n        h1_timestamp = h1[\"timestamps_GPS\"][()]\n        # H2 data\n        l1_stft = l1[\"SFTs\"][()]\n        l1_timestamp = l1[\"timestamps_GPS\"][()]\n\n        return [h1_stft, h1_timestamp],            [l1_stft, l1_timestamp], np.array(freq_hz)\n\nif DEBUG:\n    (h1_sfts, h1_ts), (l1_sfts, l1_ts), freq = read_data(TEST_IMAGE)\n    print('~Test dataset period:', min(h1_ts), max(h1_ts))\n\n    plot_amplitude_phase_spectrograms(\n        h1_ts, freq, h1_sfts\n    )","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:07:46.401546Z","iopub.execute_input":"2022-12-10T23:07:46.401783Z","iopub.status.idle":"2022-12-10T23:07:47.873314Z","shell.execute_reply.started":"2022-12-10T23:07:46.401752Z","shell.execute_reply":"2022-12-10T23:07:47.872378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Simulate real noise\n\nCopying temporal dynamics of noise from source to the generated spectrogram.\nThis logic takes gaps into account. Buckets are distributed evently through time regardless of potential gaps.\n\nThis helps us generate less \"jaggy\" signals in both test and generated spectrograms.","metadata":{}},{"cell_type":"code","source":"df_test = pd.read_csv(TEST_CSV)\ndf_test","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:07:47.874614Z","iopub.execute_input":"2022-12-10T23:07:47.87481Z","iopub.status.idle":"2022-12-10T23:07:47.914685Z","shell.execute_reply.started":"2022-12-10T23:07:47.874784Z","shell.execute_reply":"2022-12-10T23:07:47.914029Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def bucketize_real_noise_asd(sfts, ts, buckets=256):\n    bucket_size = (ts.max() - ts.min()) // buckets\n    idx = np.searchsorted(ts, [ts[0] + bucket_size * i for i in range(buckets)])\n    global_noise_amp = np.mean(np.abs(sfts))\n    return np.array([np.mean(np.abs(i)) if i.shape[1] > 0 else global_noise_amp for i in np.array_split(sfts, idx[1:], axis=1)]), bucket_size\n\nif DEBUG:\n    asd, bucket_size = bucketize_real_noise_asd(h1_sfts, h1_ts)\n    plt.plot(asd), h1_ts[0], h1_ts[0] + len(asd) * bucket_size, h1_ts.min(), h1_ts.max()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:07:47.91589Z","iopub.execute_input":"2022-12-10T23:07:47.916243Z","iopub.status.idle":"2022-12-10T23:07:48.077766Z","shell.execute_reply.started":"2022-12-10T23:07:47.91621Z","shell.execute_reply":"2022-12-10T23:07:48.077096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\n\nTMP_PATH='data/tmp/PyFstat_example_data'\n\ndef generate_segment(writer_kwargs):\n    writer = pyfstat.Writer(**writer_kwargs)\n    writer.make_data()\n    return writer.sftfilepath\n\n\ndef simulate_real_noise(frequency, timestamps, fourier_data, detector, buckets=256):\n    asd, bucket_size = bucketize_real_noise_asd(fourier_data, timestamps, buckets=buckets)\n    !rm -rf $TMP_PATH\n    os.makedirs(TMP_PATH, exist_ok=True)\n    writer_kwargs = {\n        \"outdir\": TMP_PATH,\n        # \"tstart\": timestamps[0],\n        \"detectors\": detector,  # Detector to simulate, in this case LIGO Hanford\n        \"F0\": np.mean(frequency),  # Central frequency of the band to be generated [Hz]\n        'Band': 1/5.01,  # Frequency band-width around F0 [Hz]\n        # \"sqrtSX\": 1e-23,  # Single-sided Amplitude Spectral Density of the noise\n        \"Tsft\": 1800,  # Fourier transform time duration\n        \"SFTWindowType\": \"tukey\",\n        \"SFTWindowBeta\": 0.01,\n        \"duration\": bucket_size\n    }\n\n    all_args = []\n    for segment in range(buckets):\n        args = copy.deepcopy(writer_kwargs)\n        args[\"label\"] = f\"segment_{segment}\"\n        args[\"sqrtSX\"] = asd[segment] / C_SQRSX\n        args[\"tstart\"] = timestamps[0] + segment * bucket_size\n        all_args.append(args)\n\n    with multiprocessing.Pool() as p:\n        sft_path = p.map(generate_segment, all_args)\n\n    sft_path = \";\".join(sorted(sft_path))  # Concatenate different files using ;\n    frequency, timestamps, fourier_data = get_sft_as_arrays(sft_path)\n    ts, sft = timestamps[detector], fourier_data[detector]\n    if len(sft) == 361:\n        # print('Cutting 361 to 360')\n        sft = sft[1:]\n        frequency = frequency[1:]\n    return frequency, ts, sft\n\nif DEBUG:\n    sim_freq, sim_ts, sim_sfts = simulate_real_noise(freq, h1_ts, h1_sfts, 'H1', buckets=256)\n    plt.figure(figsize=(16, 10))\n    plt.suptitle('Real (left) vs simulated (right) noise')\n    plot_amplitude_spectrogram(\n        h1_ts, freq, h1_sfts, ax=plt.subplot(1, 2, 1)\n    )\n\n    plot_amplitude_spectrogram(\n        sim_ts, sim_freq, sim_sfts, ax=plt.subplot(1, 2, 2)\n    )\n\n    plt.subplot(121).plot(np.mean(np.absolute(h1_sfts), axis=0))\n    plt.subplot(121).set_ylim(1e-22, 2e-22)\n    plt.subplot(121).set_title(os.path.basename(TEST_IMAGE))\n\n    plt.subplot(122).plot(np.mean(np.absolute(sim_sfts), axis=0))\n    plt.subplot(122).set_ylim(1e-22, 2e-22)\n    plt.subplot(122).set_title('Simulated spectrogram')","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:07:48.078978Z","iopub.execute_input":"2022-12-10T23:07:48.079316Z","iopub.status.idle":"2022-12-10T23:09:19.633477Z","shell.execute_reply.started":"2022-12-10T23:07:48.079289Z","shell.execute_reply":"2022-12-10T23:09:19.632437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Simulate Narrow Instrumental Artifacts\n\nCopying narrow artifacts (horizontal lines) from source the the generated spectrogram","metadata":{}},{"cell_type":"code","source":"\ndef find_lines(sfts, n_sigmas=5):\n    sft_stds = np.std(sfts.astype(np.complex128)[:, :4096].reshape(360, 32, -1), axis=-1)\n    # print(sft_stds)\n    mean_std = np.mean(sft_stds)\n    std_std = np.std(sft_stds)\n    line_idx = np.where(sft_stds > mean_std + std_std * n_sigmas)\n    line_amp = sft_stds[line_idx] / mean_std\n    return line_idx, line_amp\n\nif DEBUG:\n    cnt = 0\n    (h1_sfts, h1_ts), (l1_sfts, l1_ts), freq = read_data(TEST_IMAGE)\n\n    line_idx, line_amp = find_lines(h1_sfts)\n\n    print(line_idx, line_amp)\n\n    plt.plot(np.std(h1_sfts.astype(np.complex128), axis=1))\n    plt.title('Average amplitude by frequency bucket. Note the peak at 281-282Hz that corresponds to the horizontal line')\n    plt.show()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:09:19.63489Z","iopub.execute_input":"2022-12-10T23:09:19.635057Z","iopub.status.idle":"2022-12-10T23:09:20.641179Z","shell.execute_reply.started":"2022-12-10T23:09:19.635033Z","shell.execute_reply":"2022-12-10T23:09:20.639951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def write_line(sfts, idx, amp):\n    sfts = np.copy(sfts)\n    for y, x, a in zip(idx[0].tolist(), idx[1].tolist(), amp.tolist()):\n        if x == 31:\n            sfts[y, x * 128:] *= a\n        else:\n            sfts[y, x * 128:(x+1) * 128] *= a\n    return sfts\n\nif DEBUG:\n    line_idx, line_amp = find_lines(h1_sfts)\n    sft_line = write_line(sim_sfts, line_idx, line_amp)\n    plt.figure(figsize=(16, 10))\n    plt.suptitle('Real (left) vs simulated (right) noise')\n    \n    plot_amplitude_spectrogram(\n        h1_ts, freq, h1_sfts, ax=plt.subplot(1, 2, 1)\n    )\n    plot_amplitude_spectrogram(\n        sim_ts, sim_freq, sft_line, ax=plt.subplot(1, 2, 2)\n    )","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:09:20.642415Z","iopub.execute_input":"2022-12-10T23:09:20.642687Z","iopub.status.idle":"2022-12-10T23:09:21.329666Z","shell.execute_reply.started":"2022-12-10T23:09:20.642656Z","shell.execute_reply":"2022-12-10T23:09:21.328379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Generate noise samples for every test spectrogram\n\n* Note that we only look at `is_generated_noise == False` test spectrograms as they come from real detectors\n* Here we generate only one random spectrogram, the rest of samples were processed locally","metadata":{}},{"cell_type":"code","source":"df_nongen = df_test.query('is_generated_noise == False')\ndf_nongen = df_nongen.sample(1, random_state=42)\nIDS = df_nongen.id.values\n\ndef save_noise(dir, detector, sfts):\n    os.makedirs(dir, exist_ok=True)    \n    np.save(f'{dir}/{detector}', sfts)\n\nfor i, id in enumerate(tqdm(IDS)):\n    (h1_sfts, h1_ts), (l1_sfts, l1_ts), freq = read_data(f'{G2NET_ROOT}/test/{id}.hdf5')\n\n    if not os.path.exists(f'{NOISE_PATH}/npy/{id}/H1.npy'):\n        line_idx, line_amp = find_lines(h1_sfts)\n        sim_freq, sim_ts, sim_sfts = simulate_real_noise(freq, h1_ts, h1_sfts, 'H1', buckets=BUCKETS)\n        sim_sfts = write_line(sim_sfts, line_idx, line_amp)\n        save_noise(f'{NOISE_PATH}/npy/{id}/', 'H1', sim_sfts)\n   \n    if not os.path.exists(f'{NOISE_PATH}/npy/{id}/L1.npy'):\n        line_idx, line_amp = find_lines(l1_sfts)\n        sim_freq, sim_ts, sim_sfts = simulate_real_noise(freq, l1_ts, l1_sfts, 'L1', buckets=BUCKETS)\n        sim_sfts = write_line(sim_sfts, line_idx, line_amp)\n        save_noise(f'{NOISE_PATH}/npy/{id}/', 'L1', sim_sfts)\n","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:09:21.332787Z","iopub.execute_input":"2022-12-10T23:09:21.33307Z","iopub.status.idle":"2022-12-10T23:11:47.786057Z","shell.execute_reply.started":"2022-12-10T23:09:21.333034Z","shell.execute_reply":"2022-12-10T23:11:47.784089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if DEBUG:\n    np.random.seed(123)\n    for fname in np.random.choice(glob.glob(f'{NOISE_PATH}/npy/*/H1.npy'), size=10):\n        id = os.path.basename(os.path.dirname(fname))\n        (h1_sfts, h1_ts), (l1_sfts, l1_ts), freq = read_data(f'{G2NET_ROOT}/test/{id}.hdf5')\n\n        plt.figure(figsize=(16, 10))\n        plt.suptitle(f'{id} (left) vs generated (right)')\n        plot_amplitude_spectrogram(h1_ts, freq, h1_sfts, ax=plt.subplot(121))\n        plot_amplitude_spectrogram(h1_ts, freq, np.load(fname)[:, :len(h1_ts)], ax=plt.subplot(122))\n        plt.show()","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2022-12-10T23:11:47.787833Z","iopub.status.idle":"2022-12-10T23:11:47.788594Z","shell.execute_reply.started":"2022-12-10T23:11:47.78834Z","shell.execute_reply":"2022-12-10T23:11:47.788367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Produce Images from Precalculated Samples\n\n`gs://vslaykovsky/test` contains precalculated .npy files with generated noise. In this section we produce .PNG files suitable for training. ","metadata":{}},{"cell_type":"code","source":"from PIL import Image\nimport re \n\nfor df in tqdm(np.array_split(df_test.query('is_generated_noise == False'), 100), desc='Processing .NPY files'):\n    uris = ' '.join([f'gs://vslaykovsky/test/{id}' for id in df.id])\n    os.makedirs('data/tmp/npy', exist_ok=True)\n    !gsutil -q -m cp -r $uris data/tmp/npy\n    for npy in glob.glob('data/tmp/npy/*/*'):\n        try:\n            id, detector = re.findall('data/tmp/npy/(.*)/(.*).npy', npy)[0]\n            img = np.absolute(np.load(npy, allow_pickle=True))[:, :5632].reshape(360, 512, -1).mean(axis=2)\n\n            mean, std = np.mean(img), np.std(img.astype(np.float64))\n            # print(mean, std)    \n\n            img = img - mean\n            img = img / std / 5 # 5 sigma\n            img *= 128 \n            img += 128\n\n            \n            img = np.clip(img, 0, 255).astype(np.uint8)\n            # print(np.min(img), np.max(img), np.mean(img), np.std(img))\n            os.makedirs(f'data/realistic_noise/images/{id}', exist_ok=True)\n            Image.fromarray(img).save(f'data/realistic_noise/images/{id}/{detector}.png')\n        except Exception as ex:\n            print('Couldnt process', id, detector, ex)\n    !rm -rf data/tmp/npy","metadata":{"execution":{"iopub.status.busy":"2022-12-10T23:14:27.276489Z","iopub.execute_input":"2022-12-10T23:14:27.276788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 50))\nfor i, fname in enumerate(np.random.choice(glob.glob('data/realistic_noise/images/*/*.png'), size=32)):\n    plt.subplot(8, 4, i + 1).imshow(np.array(Image.open(fname)), cmap='gray')","metadata":{"execution":{"iopub.status.busy":"2022-12-10T23:11:47.791995Z","iopub.status.idle":"2022-12-10T23:11:47.792647Z","shell.execute_reply.started":"2022-12-10T23:11:47.792411Z","shell.execute_reply":"2022-12-10T23:11:47.792434Z"},"trusted":true},"execution_count":null,"outputs":[]}]}