labhamlet/wavjepa-base
27.6k
1import copy2import numpy as np 3 4from typing import Any, Optional5 6import torch7from torch import nn8 9 10from .pos_embed import get_1d_sincos_pos_embed_from_grid, get_2d_sincos_pos_embed, get_binaural_pos_embed11from .audio_extractor import Extractor12from .types import TransformerLayerCFG, TransformerEncoderCFG13from .utils import normalize, calculate_padding_mask, get_timestamps14 15class WavJEPA(nn.Module):16 """17 Joint-Embedding Predictive Architecture (JEPA).18 19 This implementation is inspired by:20 * I-JEPA http://arxiv.org/abs/2301.0824321 * Data2vec 2.0 http://arxiv.org/abs/2212.0752522 """23 24 teacher_encoder: nn.Module25 sample_rate : int = 1600026 process_audio_seconds : float = 2.0127 in_channels : int = 128 29 30 def __init__(31 self,32 feature_extractor: Extractor,33 transformer_encoder_layers_cfg : TransformerLayerCFG,34 transformer_encoder_cfg : TransformerEncoderCFG,35 transformer_decoder_layers_cfg : TransformerLayerCFG,36 transformer_decoder_cfg : TransformerEncoderCFG,37 size : str = "base",38 **kwargs : dict[str, Any],39 ):40 super().__init__(**kwargs)41 42 self.is_spectrogram = False43 self.target_length = int(self.sample_rate * self.process_audio_seconds)44 self.extract_audio = feature_extractor45 self.total_patches = 20046 self.feature_norms : nn.Module = nn.LayerNorm(self.extract_audio.embedding_dim)47 48 self.n_encoder_heads = transformer_encoder_layers_cfg["nhead"]49 self.encoder_embedding_dim = transformer_encoder_layers_cfg["d_model"]50 self.n_decoder_heads = transformer_decoder_layers_cfg["nhead"]51 self.decoder_embedding_dim = transformer_decoder_layers_cfg["d_model"]52 53 encoder_layer = nn.TransformerEncoderLayer(**transformer_encoder_layers_cfg, activation=nn.GELU())54 self.encoder = nn.TransformerEncoder(encoder_layer, norm = nn.LayerNorm(self.encoder_embedding_dim), **transformer_encoder_cfg)55 self.post_extraction_mapper : Optional[nn.Module] = nn.Linear(feature_extractor.embedding_dim, self.encoder_embedding_dim) if feature_extractor.embedding_dim != self.encoder_embedding_dim else None56 decoder_layer = nn.TransformerEncoderLayer(**transformer_decoder_layers_cfg, activation=nn.GELU())57 self.decoder = nn.TransformerEncoder(decoder_layer, norm = nn.LayerNorm(self.decoder_embedding_dim), **transformer_decoder_cfg)58 self.decoder_to_encoder_mapper = nn.Linear(self.decoder_embedding_dim, self.encoder_embedding_dim, bias=True)59 self.encoder_to_decoder_mapper = nn.Linear(self.encoder_embedding_dim, self.decoder_embedding_dim)60 61 # For the autocast add batch dimensions.62 self.mask_token = nn.Parameter(63 torch.zeros(1, 1, self.decoder_embedding_dim, requires_grad=True)64 )65 self.pos_encoding_encoder = self._get_pos_embed_params(self.encoder_embedding_dim)66 self.pos_encoding_decoder = self._get_pos_embed_params(self.decoder_embedding_dim)67 self.output_steps = self.extract_audio.total_patches(self.target_length) // self.in_channels68 69 self._init_teacher()70 71 72 def _get_pos_embed_params(self, embedding_dim):73 """Calculates the pos embedding embedding parameters and returns them."""74 # Update positional embedding75 pos_embed = nn.Parameter(76 torch.zeros(77 1,78 self.total_patches,79 embedding_dim,80 ),81 requires_grad=False,82 )83 positions = np.arange(self.total_patches, dtype=np.float64)84 if self.is_spectrogram:85 # If it is a spectrogram, we use 2d sincos embeddings.86 pos_embed_data = get_2d_sincos_pos_embed(87 embedding_dim, self.extract_audio.grid_size, cls_token_num=088 )89 #TODO! Remove this total patches later.90 elif not self.is_spectrogram and self.in_channels == 2 and (self.total_patches == 400):91 # We use 1D sincos embeddings with channel number indicated on the last 384 dimensions.92 pos_embed_data = get_binaural_pos_embed(embedding_dim, time_steps=self.total_patches // self.in_channels93 )94 elif not self.is_spectrogram and self.in_channels == 2 and (self.total_patches == 200):95 #Use 1D pos_embeddings if channel-mixing feature extractor96 pos_embed_data = get_1d_sincos_pos_embed_from_grid(97 embedding_dim,98 positions,99 ) 100 elif not self.is_spectrogram and self.in_channels == 1 and (self.total_patches == 200):101 # IF it is plain audio, we used 1d sincos embeddings102 pos_embed_data = get_1d_sincos_pos_embed_from_grid(103 embedding_dim,104 positions,105 )106 else:107 raise Exception(f"Not implemented for more in_channels, {self.in_channels}, {self.total_patches}")108 pos_embed.data.copy_(torch.from_numpy(pos_embed_data).float().unsqueeze(0))109 return pos_embed110 111 def _init_teacher(self):112 self.teacher_encoder = copy.deepcopy(self.encoder)113 self.teacher_encoder.requires_grad_(False)114 115 116 117 @torch.inference_mode()118 def _get_segment_representation(self, audio : torch.Tensor, padding_mask : torch.tensor):119 # Get the audio representatin of waveform x.120 local_features = self.extract_audio(audio)121 local_features = self.feature_norms(local_features)122 if self.post_extraction_mapper:123 local_features = self.post_extraction_mapper(local_features)124 local_features = local_features + self.pos_encoding_encoder125 # Encoder and decoder forward126 contextual_features = self.encoder(local_features, src_key_padding_mask = padding_mask)127 return contextual_features128 129 @torch.inference_mode()130 def get_audio_representation(self, audio : torch.Tensor):131 B = audio.shape[0]132 input_audio_len = audio.shape[-1]133 # Assert audio is of correct shape134 if audio.ndim != 3:135 raise ValueError(136 "audio input tensor must be 2D with shape (n_sounds, n_channels, num_samples)"137 )138 cur_frames = audio.shape[-1]139 pad_frames = self.target_length - (cur_frames % self.target_length)140 if pad_frames > 0:141 # Padding with constant 0s142 pad_arg = (143 0,144 pad_frames,145 ) # (channel, channel, height, height, width, width)146 audio = torch.nn.functional.pad(audio, pad_arg, mode="constant")147 embeddings = []148 padding_mask, cut_off = calculate_padding_mask(pad_frames = pad_frames, 149 total_frames = audio.shape[-1], 150 sr = self.sample_rate,151 output_steps = self.total_patches,152 process_seconds = self.target_length // self.sample_rate, 153 device = audio.device, 154 B = B)155 mask_idx = 0156 masked_mean = torch.zeros(audio.shape, dtype = torch.bool)157 masked_mean[..., cur_frames:] = True158 mt = torch.masked.masked_tensor(audio, masked_mean)159 # Now get the embeddings o the model.160 for i in range(audio.shape[-1] // self.target_length):161 mt = audio[..., i * self.target_length : (i + 1) * self.target_length]162 mask = padding_mask[...,mask_idx : mask_idx + self.output_steps]163 with torch.no_grad():164 # We do not include padding tokens in the mean and std calculation.165 embedding = self._get_segment_representation(166 normalize(mt),167 mask168 )169 mask_idx = mask_idx + self.output_steps170 embeddings.append(embedding)171 172 x = torch.hstack(embeddings)173 x = x[:, :cut_off, :]174 ts = get_timestamps(self.sample_rate, B, input_audio_len, x)175 return x, ts 176 177 178 179 