import torch from transformers import PreTrainedModel from .configuration import MusicDetectionConfig import transformers import os import librosa import numpy as np from huggingface_hub import hf_hub_download class MusicDetectionModel(PreTrainedModel): config_class = MusicDetectionConfig def __init__(self, pretrained_model, feature_extractor, classifier): super().__init__(MusicDetectionConfig()) self.pretrained_model = pretrained_model self.feature_extractor = feature_extractor self.classifier = classifier @classmethod def post_process(cls, output, audio, sampling_rate): input_duration = audio.shape[0] / sampling_rate output_sample_rate = len(output) / input_duration segments = [] active_label = None active_start = None assert len(output.shape) == 1 for index, label in enumerate(output.tolist()): if label != active_label: if active_label is not None: segments.append({ 'start': active_start / output_sample_rate, 'stop': index / output_sample_rate, # 'duration': (index / output_sample_rate) - (active_start / output_sample_rate), 'label': active_label, }) active_label = label active_start = index if active_label is not None and (index - active_start) > 0: segments.append({ 'start': active_start / output_sample_rate, 'stop': index / output_sample_rate, # 'duration': (index / output_sample_rate) - (active_start / output_sample_rate), 'label': active_label }) return segments @classmethod def from_pretrained(cls, repo_id, **kwargs): pretrained_model = transformers.AutoModel.from_pretrained(f'{repo_id}', subfolder='pretrained', trust_remote_code=kwargs.get("trust_remote_code"), token=kwargs.get("token")).eval() feature_extractor = transformers.AutoFeatureExtractor.from_pretrained(f'{repo_id}', subfolder='pretrained', trust_remote_code=kwargs.get("trust_remote_code"), token=kwargs.get("token")) if os.path.exists(repo_id): classifier_path = f"{repo_id}/classifier/model.pt2" else: classifier_path = hf_hub_download(repo_id, filename='model.pt2', subfolder='classifier', token=kwargs.get("token")) classifier = torch.export.load(classifier_path).module() return cls(pretrained_model, feature_extractor, classifier) def forward(self, audio, sampling_rate): if isinstance(audio, torch.Tensor): audio = audio.cpu().detach().numpy() with torch.no_grad(): inputs = self.feature_extractor( raw_speech=audio, padding='longest', truncation=False, return_attention_mask=True, return_tensors='pt', sampling_rate=sampling_rate, ).to(self.device) bsz = inputs['input_values'].shape[0] if bsz != 1: raise NotImplementedError(f"Batch processing is not implemented. I was given {bsz} audio sequences") with torch.no_grad(): features = self.pretrained_model( **inputs, output_hidden_states=True ) assert len(features.hidden_states) == 1 + len(self.pretrained_model.encoder.layers), f"{self.model_class} does not contains the CNN output." cnn_proj_features = features.hidden_states[0] features.hidden_states = features.hidden_states[1:] embeddings = torch.stack([cnn_proj_features] + list(features.hidden_states), dim=1) with torch.no_grad(): logits, prediction = self.classifier(embeddings) smoothed = torch.tensor(librosa.sequence.viterbi_binary( prob=prediction.cpu().detach().numpy(), transition=np.array([[0.95, 0.05], [0.05, 0.95]]), p_init=None, p_state=None, )).to(self.device).bool().squeeze(0) return self.post_process(smoothed, audio, sampling_rate)