czyoung's picture
Update sonogram.py
0fa9c00 verified
Raw
History Blame
12.2 kB
import sonogram_utility as su
from pyannote.audio import Pipeline
import pickle
import torch
import soundfile as sf
import numpy as np
from pyannote.core import Segment
from pyannote.audio.models.segmentation import PyanNet
from pyannote.audio import Inference
from pyannote.pipeline.parameter import ParamDict
import torch
class Sonogram():
def __init__(self,version='1.0'):
self.earlyCleanup = True
self.isTPU = False
self.isGPU = False
try:
raise(RuntimeError("Not an error"))
#device = xm.xla_device()
print("TPU is available.")
self.isTPU = True
except RuntimeError as e:
print(f"TPU is not available: {e}")
# Fallback to CPU or other devices if needed
self.isGPU = torch.cuda.is_available()
if not self.isGPU:
print(f"GPU is not available")
self.device = torch.device("cuda" if self.isGPU else "cpu")
print(f"Using {self.device} instead.")
self.version = version
if version == 'speaker-diarization-3.1':
self.pipeline = Pipeline.from_pretrained("pyannote/speaker-diarization-3.1")
elif version == '1.0':
baselinePipeline = Pipeline.from_pretrained("pyannote/speaker-diarization-3.1")
newSpecs = baselinePipeline._segmentation.model.specifications
segModel = PyanNet.from_pretrained('20251208_Sonogram_Segmentation.ckpt')
segModel.specifications = newSpecs
segmentation_duration = segModel.specifications.duration
baselinePipeline._segmentation = Inference(
segModel,
duration=segmentation_duration,
step=baselinePipeline.segmentation_step * segmentation_duration,
skip_aggregation=True,
batch_size=1,
)
baselinePipeline.segmentation = ParamDict(min_duration_off=0.0,)
self.pipeline = baselinePipeline
# Should manage this outside class now
#self.pipeline.to(self.device)
# Load SVM classifier
with open('05062026_groupClassifier.pkl', 'rb') as f:
self.groupClassifier = pickle.load(f)
def classifyEmbedding(self,embedding):
return int(self.groupClassifier.predict(embedding.reshape(1, -1)).item())
def processFile(self,filePath):
# Loading audio file
print(f"Loading file: {filePath}")
data, sample_rate = sf.read(filePath, dtype="float32", always_2d=True)
waveform = torch.from_numpy(data.T) # shape: [channels, samples]
# Wrapping as AudioFile
audioFile = {"waveform": waveform, "sample_rate": sample_rate}
print("Detecting Voices")
segmentations = self.pipeline.get_segmentations(audioFile)
print("Generating vocal embeddings")
embeddings = self.pipeline.get_embeddings(audioFile,segmentations,exclude_overlap=False)
print("Clustering Speakers")
hardC, softC, centroids = self.pipeline.clustering(embeddings = embeddings,segmentations = segmentations)
count = self.pipeline.speaker_count(
segmentations,
self.pipeline._segmentation.model.receptive_field,
warm_up=(0.0, 0.0),)
print("Classifying Embeddings")
embeddingClasses = np.zeros((embeddings.shape[0],embeddings.shape[1]))
# Dumb loop
# Timestep
tempCount = [0,0]
for i,e in enumerate(embeddings):
# Speaker, skip empty
for j,eS in enumerate(e):
if np.any(segmentations.data[i,:,j] > 0):
groupClass = self.classifyEmbedding(eS)
embeddingClasses[i][j] = groupClass
# Correct silence detections, remove group for later addition
if groupClass == 2:
segmentations.data[i,:,j] = 0
tempCount[1] += 1
elif groupClass == 0 and np.mean(segmentations.data[i,:,j]) < 0.5:
segmentations.data[i,:,j] = 0
tempCount[0] += 1
print("Generating Annotation")
# shape: (num_chunks, num_speakers)
# keep track of inactive speakers
inactive_speakers = np.sum(segmentations.data, axis=1) == 0
hardC[inactive_speakers] = -2
discrete_diarization = self.pipeline.reconstruct(
segmentations,
hardC,
count,)
diarization = self.pipeline.to_annotation(
discrete_diarization,
min_duration_on=0.0,
min_duration_off=self.pipeline.segmentation.min_duration_off,
)
# keep track of group speakers
group_speakers = np.any(embeddingClasses >= 2,axis=1)
start = None
for timeStep in range(group_speakers.shape[0]):
if group_speakers[timeStep] > 0:
if start is None:
start = timeStep
elif start is not None:
segment = Segment(start, timeStep)
diarization[segment] = 'group'
start = None
# Catch end case
if start is not None:
segment = Segment(start, embeddingClasses.shape[0]-1)
diarization[segment] = 'group'
totalTimeInSeconds = int(waveform.shape[-1]/sample_rate)
# Rename labels
currId = 0
mapping = {}
for label in diarization.labels():
if label == 'group':
continue
else:
currId += 1
newLabel = f'SPEAKER_{currId:03d}'
mapping[label] = newLabel
diarization = diarization.rename_labels(mapping)
print("Time in seconds calculated")
return diarization, totalTimeInSeconds, waveform, sample_rate
def activeSpeaker(self,inAnnotation,step=1):
speakerAtStep = [None]
stepTime = [0]
speakerHierarchy = [label for label,_ in inAnnotation.chart()]
for label in speakerHierarchy:
# Move group labels to beginning of hierarchy
if label == 'group' or label == 99:
speakerHierarchy.remove(label)
speakerHierarchy.insert(0,label)
for segment,_,label in inAnnotation.itertracks(yield_label=True):
startI = int(segment.start / step)
# Lazy end assumption, always assumes one more step
endI = int(segment.end / step) + 1
while len(stepTime) < endI+1:
stepTime.append(stepTime[-1]+step)
speakerAtStep.append(None)
for i in range(startI,endI+1):
if speakerAtStep[i] == None:
speakerAtStep[i] = label
else:
currHier = speakerHierarchy.index(speakerAtStep[i])
newHier = speakerHierarchy.index(label)
if newHier < currHier:
speakerAtStep[i] = label
return speakerAtStep, stepTime, speakerHierarchy
def annotationToNoiseList(self,inAnnotation,maxTime,stepSize=2,windowSize=90):
sas, st, sh = self.activeSpeaker(inAnnotation,step=stepSize)
timeStepAggregate = []
timeStepClass = []
timeStepMembers = []
categories = ['group','individual','silence']
for i in st:
timeStepAggregate.append({'individual':0,'group':0,'silence':0})
for i,_ in enumerate(sas):
decision = None
groupCount = 0
individuals = set()
silenceCount = 0
end = min(i+windowSize,len(sas))
memberSet = set()
for j in range(i,end):
if sas[j] is not None:
memberSet.add(sas[i])
if sas[j] == None:
silenceCount += 1
elif sas[j] == 'group' or sas[j] == 99:
groupCount += 1
else:
individuals.add(sas[j])
if silenceCount > windowSize / 2:
decision = 'silence'
elif groupCount > windowSize / 2 or len(individuals) > 2:
decision = 'group'
else:
decision = 'individual'
for j in range(i,end):
timeStepAggregate[j][decision] += 1
# Convert to list and sort for convenience
memberSet = list(memberSet)
memberSet.sort()
timeStepMembers.append(memberSet)
for i,item in enumerate(timeStepAggregate):
cat = max(item, key=item.get)
timeStepClass.append(cat)
# For group decisions, members include all in window
if cat == 'group':
timeStepMembers[i] = '+'.join(timeStepMembers[i])
elif cat == 'individual':
timeStepMembers[i] = timeStepMembers[i][0]
else:
timeStepMembers[i] = None
# For debug
endTime = 0
singleDimList = []
categorySegmentList = []
for c in categories:
currList = []
start = None
currMembers = None
duration = 0
tracking = False
for stepClass,timeIncrement,members in zip(timeStepClass,st,timeStepMembers):
# Check for case of exact end of audio
if st == maxTime:
continue
if stepClass == c:
if currMembers == members:
duration += min(stepSize,maxTime-timeIncrement)
else:
if tracking:
singleDimList.append((currMembers,Segment(start,start+duration)))
currList.append((currMembers,Segment(start,start+duration)))
if start+duration > endTime:
endTime = start+duration
start = None
currMembers == None
duration = 0
tracking = False
start = timeIncrement
currMembers = members
duration += min(stepSize,maxTime-timeIncrement)
tracking = True
else:
if tracking:
singleDimList.append((currMembers,Segment(start,start+duration)))
currList.append((currMembers,Segment(start,start+duration)))
if start+duration > endTime:
endTime = start+duration
start = None
currMembers == None
duration = 0
tracking = False
if tracking:
singleDimList.append((currMembers,Segment(start,start+duration)))
currList.append((currMembers,Segment(start,start+duration)))
if start+duration > endTime:
endTime = start+duration
categorySegmentList.append(currList)
# Check where we left off
if endTime != maxTime:
singleDimList.append((None,Segment(endTime,maxTime)))
categorySegmentList[2].append((None,Segment(endTime,maxTime)))
# For debug
singleDimList = sorted(singleDimList,key=lambda index : index[1].start)
print(singleDimList)
return categorySegmentList, st
def __call__(self,audioPath):
annotation, totalTimeInSeconds, waveform, sampleRate = self.processFile(audioPath)
return annotation, totalTimeInSeconds, waveform, sampleRate
def toDevice(self):
self.pipeline.to(self.device)
print(f"Sonogram moved to {self.device}")
def toCPU(self):
self.pipeline.to(torch.device('cpu'))
print(f"Sonogram moved to CPU")