import torch import torch.nn.functional as F from speechbrain.inference.interfaces import Pretrained class SENSEEncoder(Pretrained): """SENSE encoder for extracting semantic speech embeddings. SENSE maps speech utterances into a multilingual semantic embedding space aligned with BGE-M3 text representations. Arguments --------- See ``Pretrained``. Example ------- >>> from speechbrain.inference.interfaces import foreign_class >>> sense = foreign_class( ... source="LIA-AvignonUniversity/SENSE", ... pymodule_file="custom.py", ... classname="SENSEEncoder", ... ) # doctest: +SKIP >>> embedding = sense.encode_file("example.wav") # doctest: +SKIP """ MODULES_NEEDED = ["wav2vec2", "attn_pooling"] def encode_file(self, path, **kwargs): """Encode an audio file into a SENSE embedding. Arguments --------- path : str Path to the audio file. **kwargs : dict Additional arguments passed to ``load_audio``. Returns ------- torch.Tensor L2-normalized semantic embedding with shape ``[1, embedding_dim]``. """ waveform = self.load_audio(path, **kwargs) # Add a batch dimension. batch = waveform.unsqueeze(0) rel_length = torch.tensor([1.0], device=self.device) return self.encode_batch(batch, rel_length) def encode_batch(self, wavs, wav_lens=None): """Encode input waveforms into SENSE embeddings. Arguments --------- wavs : torch.Tensor Batch of waveforms with shape ``[batch, time]``, or a single waveform with shape ``[time]``. wav_lens : torch.Tensor, optional Relative waveform lengths in the range ``[0, 1]``. If ``None``, the full length is assumed for every waveform. Returns ------- torch.Tensor L2-normalized semantic embeddings with shape ``[batch, embedding_dim]``. """ # Add a batch dimension for a single waveform. if len(wavs.shape) == 1: wavs = wavs.unsqueeze(0) # Assume full-length waveforms when lengths are not provided. if wav_lens is None: wav_lens = torch.ones( wavs.shape[0], device=self.device, ) # Move inputs to the inference device. wavs = wavs.to(self.device).float() wav_lens = wav_lens.to(self.device) # Extract frame-level representations with w2v-BERT 2.0. feats = self.mods.wav2vec2( wavs, wav_lens, ) # Aggregate frame-level representations into utterance embeddings. embeddings = self.mods.attn_pooling(feats) # L2-normalize the utterance embeddings. embeddings = F.normalize( embeddings, p=2, dim=-1, ) return embeddings def forward(self, wavs, wav_lens=None): """Encode input waveforms into SENSE embeddings.""" return self.encode_batch(wavs, wav_lens)