diff --git a/libs/architectures/architectures/supervised.py b/libs/architectures/architectures/supervised.py index 9f745fddd..3c5105a3b 100644 --- a/libs/architectures/architectures/supervised.py +++ b/libs/architectures/architectures/supervised.py @@ -302,13 +302,17 @@ def __init__( zero_init_residual: bool = False, groups: int = 1, width_per_group: int = 64, + top_k: int | None = None, stride_type: list[Literal["stride", "dilation"]] | None = None, norm_layer: NormLayer | None = None, **kwargs, ) -> None: super().__init__() + + num_channels = top_k if top_k is not None else num_chirp_masses + self.time_domain_resnet = ResNet1D( - in_channels=num_ifos * num_chirp_masses, + in_channels=num_ifos * num_channels, layers=layers, classes=1, kernel_size=kernel_size, diff --git a/libs/utils/tests/test_preprocessing.py b/libs/utils/tests/test_preprocessing.py index a7544fb1f..9cc3b6c68 100644 --- a/libs/utils/tests/test_preprocessing.py +++ b/libs/utils/tests/test_preprocessing.py @@ -1,5 +1,6 @@ import pytest import torch +import torch.nn.functional as F from ml4gw.transforms import Heterodyne from utils.augmentation import HeterodyneAugmentor from utils.preprocessing import ( @@ -302,6 +303,84 @@ def test_with_heterodyne_augmentor(self): kernels_aug = whitener_aug(x) assert torch.allclose(kernels_aug, kernels) + def test_with_heterodyne_augmentor_top_k(self): + + heterodyne_augmentor = HeterodyneAugmentor( + sample_rate=self.sample_rate, + kernel_length=self.kernel_length, + chirp_mass_low=1.0, + chirp_mass_high=2.5, + num_chirp_masses=10, + chirp_mass_spacing="log", + keep_last_n_seconds=1.5, + top_k=5, + ) + + whitener = BatchWhitener( + kernel_length=self.kernel_length, + sample_rate=self.sample_rate, + inference_sampling_rate=self.inference_sampling_rate, + batch_size=self.batch_size, + fduration=self.fduration, + fftlength=self.fftlength, + ) + + whitener_aug = BatchWhitener( + kernel_length=self.kernel_length, + sample_rate=self.sample_rate, + inference_sampling_rate=self.inference_sampling_rate, + batch_size=self.batch_size, + fduration=self.fduration, + fftlength=self.fftlength, + augmentor=heterodyne_augmentor, + ) + + channels = 2 + total_samples = int( + ( + (self.batch_size - 1) * whitener_aug.stride_size + + whitener_aug.kernel_size + ) + * 2 + ) + + x = torch.randn(channels, total_samples) + + heterodyne = Heterodyne( + sample_rate=self.sample_rate, + kernel_length=self.kernel_length, + chirp_mass=heterodyne_augmentor.chirp_mass_grid, + return_type="time", + ) + + kernels = whitener_aug(x) + + kernels_heterodyned = heterodyne(whitener(x)) + _B, _C, _M, _T = kernels_heterodyned.shape + avgpool = F.avg_pool1d( + torch.abs(kernels_heterodyned.reshape(_B, _C * _M, _T)), + 31, + stride=1, + padding=15, + ).reshape(_B, _C, _M, _T) + avgpool_snr = torch.sqrt( + (avgpool[..., -int(1.5 * self.sample_rate) :] ** 2).sum(dim=1) + ) + vals = torch.max(avgpool_snr, dim=-1).values + idx = torch.topk(vals, k=5, dim=-1).indices + idx_expand = idx[:, None, :, None].expand(-1, _C, -1, _T) + kernels_heterodyned = torch.gather( + kernels_heterodyned, dim=2, index=idx_expand + ) + kernels_heterodyned = kernels_heterodyned.reshape(_B, _C * 5, _T) + kernels_heterodyned = kernels_heterodyned[ + ..., -int(1.5 * self.sample_rate) : + ] + kernels = kernels_heterodyned + + kernels_aug = whitener_aug(x) + assert torch.allclose(kernels_aug, kernels) + class TestMultiModalPreprocessor: """Test suite for MultiModalPreprocessor module.""" diff --git a/libs/utils/utils/augmentation.py b/libs/utils/utils/augmentation.py index bd13a1e34..4507fe6d9 100644 --- a/libs/utils/utils/augmentation.py +++ b/libs/utils/utils/augmentation.py @@ -2,10 +2,76 @@ from typing import Literal import torch +import torch.nn.functional as F from ml4gw.transforms import Heterodyne from torch import Tensor +def select_top_k( + X: torch.Tensor, + top_k: int, + kernel_size: int = 31, + stride: int = 1, + padding: int = 15, + keep_last_n_samples: int | None = None, +) -> torch.Tensor: + """ + Select top `k` chirp mass channels from `M` heterodyned timeseries. + + Args: + X (Tensor): + Input tensor of shape (B, C, M, T) where B is the batch + size, C is the number of channels, M is the number of chirp + mass channels, and T is the number of time samples. + top_k (int): + Number of chirp mass channels to retain. + kernel_size (int): + Size of the running average (average pooling) window used to + smooth the absolute value of the input before computing the + selection statistic. + stride (int): + Stride of the average pooling operation. + padding (int): + Zero-padding applied to both sides of the timeseries before + the average pooling operation. + keep_last_n_samples (int, optional): + If provided, only the final `n` samples are used when + computing the statistic used to select the top k chirp mass + channels. If `None`, all samples are used. + + Returns: + Tensor: + Output tensor of shape (B, C * top_k, T) containing the + selected top k chirp mass channels for each batch and channel. + """ + + B, C, M, T = X.shape + + avgpool = torch.stack( + [ + F.avg_pool1d( + torch.abs(x.reshape(C * M, T)), + kernel_size=kernel_size, + stride=stride, + padding=padding, + ) + for x in X + ] + ).reshape(B, C, M, T) + + if keep_last_n_samples is not None: + avgpool_snr = torch.sqrt( + (avgpool[..., -keep_last_n_samples:] ** 2).sum(dim=1) + ) + else: + avgpool_snr = torch.sqrt((avgpool**2).sum(dim=1)) + + vals = avgpool_snr.max(dim=-1).values + idx = torch.topk(vals, k=top_k, dim=-1).indices + idx = idx[:, None, :, None].expand(B, C, top_k, T) + return torch.gather(X, dim=2, index=idx).reshape(B, C * top_k, T) + + class HeterodyneAugmentor(torch.nn.Module): """ Apply a heterodyne transform over a grid of chirp masses to a batch @@ -26,6 +92,10 @@ class HeterodyneAugmentor(torch.nn.Module): keep_last_n_seconds (float): If provided, only the last `n` seconds of the kernel_length are returned. Otherwise, the full kernel_length is returned. + top_k (int): + If provided, the top `k` chirp mass channels are selected based on + their absolute amplitude, corresponding to the `k` matching chirp + masses with. If `None`, all chirp mass channels are returned. Shape: Input: (batch_size, channels, time) Output: (batch_size, channels * num_chirp_masses, time_out) @@ -58,12 +128,13 @@ def __init__( chirp_mass_high: float = 2.5, num_chirp_masses: int = 100, chirp_mass_spacing: Literal["linear", "log"] = "log", - keep_last_n_seconds: float = None, + keep_last_n_seconds: float | None = None, + top_k: int | None = None, ): super().__init__() self.sample_rate = sample_rate self.kernel_length = kernel_length - self.keep_last_n_seconds = keep_last_n_seconds + self.top_k = top_k self.num_chirp_masses = num_chirp_masses self.keep_last_n_seconds = keep_last_n_seconds @@ -78,6 +149,8 @@ def __init__( self.keep_last_n_samples = int( self.keep_last_n_seconds * sample_rate ) + else: + self.keep_last_n_samples = None self.heterodyne_transform = Heterodyne( sample_rate=sample_rate, @@ -120,12 +193,22 @@ def forward(self, x: Tensor) -> Tensor: or determined by `keep_last_n_seconds`. """ _B, _C, _T = x.shape - x_heterodyned = torch.empty((_B, _C * self.num_chirp_masses, _T)) + if self.top_k is not None: + x_heterodyned = torch.empty((_B, _C * self.top_k, _T)) + else: + x_heterodyned = torch.empty((_B, _C * self.num_chirp_masses, _T)) # Heterodyne the whitened timeseries x = self.heterodyne_transform(x) - # Reshaping x from (batch_size, channels, num_chirp_mass, kernel_size) - # to (batch_size, channels x num_chirp_mass, kernel_size) - x = x.reshape(_B, _C * self.num_chirp_masses, _T) + if self.top_k is not None: + # Select the top k chirp mass channels + x = select_top_k( + x, self.top_k, keep_last_n_samples=self.keep_last_n_samples + ) + else: + # Reshaping x from + # (batch_size, channels, num_chirp_mass, kernel_size) + # to (batch_size, channels x num_chirp_mass, kernel_size) + x = x.reshape(_B, _C * self.num_chirp_masses, _T) x_heterodyned[:, :, :] = x # Returning the desired length of heterodyned strain in the # time dimension diff --git a/projects/export/export_heterodyne.yaml b/projects/export/export_heterodyne.yaml index 1e3414634..fd4fcfcf5 100644 --- a/projects/export/export_heterodyne.yaml +++ b/projects/export/export_heterodyne.yaml @@ -9,7 +9,7 @@ weights: batch_file: repository_directory: num_ifos: 2 -kernel_length: 4.0 +kernel_length: 8.0 inference_sampling_rate: 4 sample_rate: 2048 batch_size: 128 @@ -18,7 +18,7 @@ psd_length: 64 preprocessor: class_path: utils.preprocessing.BatchWhitener init_args: - kernel_length: 4.0 + kernel_length: 8.0 sample_rate: 2048 inference_sampling_rate: 4 batch_size: 128 @@ -28,12 +28,13 @@ preprocessor: class_path: utils.augmentation.HeterodyneAugmentor init_args: sample_rate: 2048 - kernel_length: 4.0 + kernel_length: 8.0 chirp_mass_low: 1.0 chirp_mass_high: 2.5 num_chirp_masses: 100 chirp_mass_spacing: "log" keep_last_n_seconds: null + top_k: 10 highpass: 32.0 streams_per_gpu: 12 # num_outputs: diff --git a/projects/online/config.yaml b/projects/online/config.yaml index acecb9b29..d10ac9696 100644 --- a/projects/online/config.yaml +++ b/projects/online/config.yaml @@ -78,3 +78,13 @@ output_buffer_length: 8 device: "cuda" verbose: true matmul_precision: "high" +# optional augmentor arguments +# augmentor: +# class_path: utils.augmentation.HeterodyneAugmentor +# init_args: +# chirp_mass_low: 1.0 +# chirp_mass_high: 2.5 +# num_chirp_masses: 100 +# chirp_mass_spacing: "log" +# keep_last_n_seconds: 1.5 +# top_k: 10 diff --git a/projects/online/online/cli.py b/projects/online/online/cli.py index f0cb60f8f..ce22e9fe8 100644 --- a/projects/online/online/cli.py +++ b/projects/online/online/cli.py @@ -23,6 +23,27 @@ def build_parser(): apply_on="parse", ) + # TODO: This is a workaround for linking sample_rate and + # kernel_length argument between the augmentor and + # the global parameters in the config. + try: + parser.link_arguments( + "sample_rate", + "augmentor.init_args.sample_rate", + apply_on="parse", + ) + except Exception: + pass + + try: + parser.link_arguments( + "kernel_length", + "augmentor.init_args.kernel_length", + apply_on="parse", + ) + except Exception: + pass + return parser diff --git a/projects/online/online/main.py b/projects/online/online/main.py index 3a71ea166..93ca4d95c 100644 --- a/projects/online/online/main.py +++ b/projects/online/online/main.py @@ -386,6 +386,7 @@ def main( integration_window_length: float, astro_event_rate: float, data_source: Literal["frames", "arrakis"] = "frames", + augmentor: torch.nn.Module | None = None, state_channels: list[str] | None = None, fftlength: float | None = None, highpass: float | None = None, @@ -450,6 +451,10 @@ def main( offline_inference_rate: Rate at which inference was performed offline when establishing the background and foreground distributions + augmentor: + Optional augmentor module that performs augmentation + of the input data that need to be analyzed by Aframe. + If not provided, no augmentation will be performed. psd_length: Length of PSD estimation window in seconds for PSD used to whiten aframe data @@ -842,6 +847,7 @@ def main( batch_size=int(update_size * online_inference_rate), fduration=fduration, fftlength=fftlength, + augmentor=augmentor, highpass=highpass, lowpass=lowpass, ).to(device) diff --git a/projects/train/configs/time_heterodyne.yaml b/projects/train/configs/time_heterodyne.yaml index 928ad26c5..2692dcfc8 100644 --- a/projects/train/configs/time_heterodyne.yaml +++ b/projects/train/configs/time_heterodyne.yaml @@ -13,7 +13,7 @@ model: class_path: architectures.supervised.SupervisedHeterodyneTimeDomainResNet init_args: layers: [3, 4, 6, 3] - kernel_size: 25 + kernel_size: 3 norm_layer: class_path: ml4gw.nn.norm.GroupNorm1D init_args: @@ -37,12 +37,12 @@ data: waveforms_dir: ifos: [H1,L1] sample_rate: 2048 - batches_per_epoch: 100 + batches_per_epoch: 200 num_files_per_batch: 10 - chunk_size: 1000 + chunk_size: 5000 chunks_per_epoch: 20 # preprocessing args - batch_size: 100 + batch_size: 2000 kernel_length: 8 psd_length: 20 fduration: 2 @@ -57,25 +57,26 @@ data: chirp_mass_high: 2.5 num_chirp_masses: 100 chirp_mass_spacing: "log" - keep_last_n_seconds: null + keep_last_n_seconds: 1.5 + top_k: 10 # highpass: # lowpass: fftlength: 2 - snr_sampler: - class_path: ml4gw.distributions.PowerLaw - init_args: - minimum: 8 - maximum: 100 - index: -3 - # curriculum learning for snr sampler # snr_sampler: - # class_path: train.augmentations.SnrSampler - # init_args: - # max_min_snr: 30 - # min_min_snr: 8 - # max_snr: 100 - # alpha: -3 - # decay_steps: 600 + # class_path: ml4gw.distributions.PowerLaw + # init_args: + # minimum: 8 + # maximum: 100 + # index: -3 + # curriculum learning for snr sampler + snr_sampler: + class_path: train.augmentations.SnrSampler + init_args: + max_min_snr: 30 + min_min_snr: 8 + max_snr: 100 + alpha: -3 + decay_steps: 600 waveform_sampler: class_path: train.data.waveforms.WaveformLoader init_args: diff --git a/projects/train/train/cli.py b/projects/train/train/cli.py index e0dfd1b2e..5fdf4080c 100644 --- a/projects/train/train/cli.py +++ b/projects/train/train/cli.py @@ -57,8 +57,9 @@ def add_arguments_to_parser(self, parser): except Exception: pass - # TODO: This is a workaround for linking num_chirp_masses between - # the model and the architecture for in_channels. + # TODO: This is a workaround for linking num_chirp_masses and + # top_k channel argument between the model and the architecture + # for in_channels. try: parser.link_arguments( "data.init_args.num_chirp_masses", @@ -67,6 +68,14 @@ def add_arguments_to_parser(self, parser): except Exception: pass + try: + parser.link_arguments( + "data.init_args.top_k", + "model.init_args.arch.init_args.top_k", + ) + except Exception: + pass + parser.link_arguments( "data.init_args.sample_rate", "data.init_args.waveform_sampler.init_args.sample_rate", diff --git a/projects/train/train/data/supervised/time_domain.py b/projects/train/train/data/supervised/time_domain.py index 50076ab24..a2b823f6b 100644 --- a/projects/train/train/data/supervised/time_domain.py +++ b/projects/train/train/data/supervised/time_domain.py @@ -3,6 +3,7 @@ import torch from ml4gw.transforms import Heterodyne +from utils.augmentation import select_top_k from train.data.supervised.supervised import SupervisedAframeDataset @@ -46,6 +47,10 @@ class HeterodyneTimeDomainSupervisedAframeDataset(SupervisedAframeDataset): keep_last_n_seconds (float): If provided, only the last `n` seconds of the kernel_length are returned. Otherwise, the full kernel_length is returned. + top_k (int): + If provided, the top `k` chirp mass channels are selected based on + their absolute amplitude, corresponding to the `k` matching chirp + masses with. If `None`, all chirp mass channels are returned. """ def __init__( @@ -55,6 +60,7 @@ def __init__( num_chirp_masses: int = 100, chirp_mass_spacing: Literal["linear", "log"] = "log", keep_last_n_seconds: float = None, + top_k: int = None, *args, **kwargs, ): @@ -68,11 +74,13 @@ def __init__( ) self.keep_last_n_seconds = keep_last_n_seconds - + self.top_k = top_k if self.keep_last_n_seconds is not None: self.keep_last_n_samples = int( self.keep_last_n_seconds * self.hparams.sample_rate ) + else: + self.keep_last_n_samples = None def build_transforms(self, *args, **kwargs): super().build_transforms(*args, **kwargs) @@ -109,17 +117,23 @@ def build_val_batches(self, background, signals): X_bg, X_inj, psds = super().build_val_batches(background, signals) X_bg = self.whitener(X_bg, psds) X_bg = self.heterodyne_transform(X_bg) - _B_bg, _C_bg, _M_bg, _T_bg = X_bg.shape - X_bg = X_bg.view(_B_bg, _C_bg * _M_bg, _T_bg) + if self.top_k is not None: + X_bg = select_top_k( + X_bg, self.top_k, keep_last_n_samples=self.keep_last_n_samples + ) # whiten each view of injections X_fg = [] for inj in X_inj: inj = self.whitener(inj, psds) inj = self.heterodyne_transform(inj) + if self.top_k is not None: + inj = select_top_k( + inj, + self.top_k, + keep_last_n_samples=self.keep_last_n_samples, + ) X_fg.append(inj) X_fg = torch.stack(X_fg) - _V_fg, _B_fg, _C_fg, _M_fg, _T_fg = X_fg.shape - X_fg = X_fg.view(_V_fg, _B_fg, _C_fg * _M_fg, _T_fg) if self.keep_last_n_seconds is not None: return X_bg[..., -self.keep_last_n_samples :], X_fg[ @@ -132,8 +146,10 @@ def inject(self, X, waveforms=None): X, y, psds = super().inject(X, waveforms) X = self.whitener(X, psds) X = self.heterodyne_transform(X) - _B, _C, _M, _T = X.shape - X = X.view(_B, _C * _M, _T) + if self.top_k is not None: + X = select_top_k( + X, self.top_k, keep_last_n_samples=self.keep_last_n_samples + ) if self.keep_last_n_seconds is not None: return X[..., -self.keep_last_n_samples :], y