CoolFace
Modelpublic

labhamlet/wavjepa-base

sourceHugging Facemitupdated 11mo agoView on Hugging Face
2likes7.6kdownloads
model.py179 linesDownload Raw Back to root
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