{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Clustering Audio with 3d Convolution\n\nLarge audio files add a constraint to the kind of models that can be applied. The use of spectrograms and other feature engineering techniques has been proposed to resolve that issue. However, reshaping the data from a one-dimensional array to a multi-dimensional array can bring closer together highly correlated pieces of data. It also allows to use of network architectures optimized to similar data sources. The following describes a simple convolutional VAE model trained with reshaped and decimated audio files of similar length. ","metadata":{}},{"cell_type":"code","source":"import os\nimport time\nimport librosa\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom scipy import signal\n\nfrom typing import Sequence,Tuple\n\nimport jax\nimport optax\nimport jax.numpy as jnp\nfrom jax import random\nfrom flax import linen as nn\n\nfrom sklearn.model_selection import train_test_split\n\nfontsize = 16\n\ndef PlotStyle(Axes): \n    \"\"\"\n    Parameters\n    ----------\n    Axes : Matplotlib axes object\n        Applies a general style to the matplotlib object\n\n    Returns\n    -------\n    None.\n    \"\"\"    \n    Axes.spines['top'].set_visible(False)\n    Axes.spines['bottom'].set_visible(True)\n    Axes.spines['left'].set_visible(True)\n    Axes.spines['right'].set_visible(False)\n    Axes.xaxis.set_tick_params(labelsize=fontsize)\n    Axes.yaxis.set_tick_params(labelsize=fontsize)\n\nclass CoderCONV(nn.Module):\n    \n    Units: Sequence[int]\n    Ksize: Tuple[int]\n    Strides: Tuple[int]\n    \n    depth: int\n    Name: str \n    UpSampling: bool = True\n    train: bool = True\n                    \n    @nn.compact\n    def __call__(self,inputs):\n        x = inputs\n        for k,val in enumerate(self.Units):\n            finalDepth = self.depth-k+1\n            if finalDepth<=1:\n                finalDepth = 1\n            \n            for ii in range(finalDepth):\n                x = nn.Conv(val,self.Ksize,padding='SAME',use_bias=False,\n                            name=self.Name+' conv_conv_'+str(k)+str(ii))(x)\n                x = nn.BatchNorm(use_running_average=not self.train,\n                                 name = self.Name+' conv_norm_'+str(k)+str(ii))(x)\n                x = nn.leaky_relu(x)\n                \n            if self.UpSampling:\n                x = nn.ConvTranspose(val,self.Ksize,padding='SAME',\n                                     strides=self.Strides,use_bias=False,\n                                     name=self.Name+' convUp_conv_'+str(k)+str(ii))(x)\n                x = nn.BatchNorm(use_running_average=not self.train,\n                                 name = self.Name+' conv_normUp_'+str(k)+str(ii))(x)\n                x = nn.leaky_relu(x)\n            else:\n                x = nn.Conv(val,self.Ksize,padding='SAME',strides=self.Strides,\n                            use_bias=False,\n                            name=self.Name+' convDn_conv_'+str(k)+str(ii))(x)\n                x = nn.BatchNorm(use_running_average=not self.train,\n                                 name = self.Name+' conv_normDn_'+str(k)+str(ii))(x)\n                x = nn.leaky_relu(x)\n            \n        return x\n\nclass CoderMLP(nn.Module):\n    \n    Units: Sequence[int]\n    Name: str \n    train: bool = True\n    \n    @nn.compact\n    def __call__(self,inputs):\n        x = inputs\n        for k,feat in enumerate(self.Units):\n            x = nn.Dense(feat,use_bias=False,name = self.Name+' layer_'+str(k))(x)\n            x = nn.BatchNorm(use_running_average=not self.train,name = self.Name+' norm_'+str(k))(x)\n            x = nn.leaky_relu(x)\n        return x\n\nclass EncoderMLP(nn.Module):\n    \n    Units: Sequence[int]\n    train: bool = True\n    \n    def setup(self):\n        self.encoder = CoderMLP(self.Units[1::],'encodermlp',train=self.train)\n        self.mean = nn.Dense(self.Units[-1], name='mean')\n        self.logvar = nn.Dense(self.Units[-1], name='logvar')\n    \n    def __call__(self, inputs):\n        \n        x = inputs\n        mlpencoded = self.encoder(x)\n        mean_x = self.mean(mlpencoded)\n        logvar_x = self.logvar(mlpencoded)\n        \n        return mean_x, logvar_x\n\nclass DecoderMLP(nn.Module):\n    \n    Units: Sequence[int]\n    train: bool = True\n    \n    def setup(self):\n        self.decoder = CoderMLP(self.Units[0:-1],'decodermlp',train=self.train)\n        self.out = nn.Dense(self.Units[-1],use_bias=False, name='out')\n        self.outnorm = nn.BatchNorm(use_running_average=not self.train,name = 'outnorm')\n    \n    def __call__(self, inputs):\n        x = inputs\n        decoded_1 = self.decoder(x)\n        \n        out =self.out(decoded_1)\n        out = self.outnorm(out)\n        out = nn.leaky_relu(out)\n        \n        return out\n\nclass CONVEncoder(nn.Module):\n    \n    Units: Sequence[int]\n    Ksize: Tuple[int]\n    Strides: Tuple[int]\n    InputShape: Tuple[int]\n    \n    depth: int\n    BatchSize: int \n    train: bool = True \n    \n    def setup(self):\n        \n        self.localConv = CoderCONV(self.Units,self.Ksize,self.Strides,self.depth,'convencoder',UpSampling=False,train=self.train)\n        self.divFactor = 2**(len(self.Units)-1)\n        \n        self.targetShape = [val//self.divFactor for val in self.InputShape[0:-1]] + [self.Units[-1]]\n        \n        self.localShape = np.prod(np.array(self.targetShape))\n        self.EncUnits = [self.localShape,self.localShape//4,self.localShape//16,2]\n        self.localEncoder = EncoderMLP(self.EncUnits,train=self.train)\n        \n    def __call__(self,inputs):\n        \n        x = inputs\n        x = self.localConv(x)\n        x = x.reshape((x.shape[0],-1))\n        mean_x,logvar_x = self.localEncoder(x)\n        \n        return mean_x,logvar_x\n\nclass CONVDecoder(nn.Module):\n    \n    Units: Sequence[int]\n    Ksize: Tuple[int]\n    Strides: Tuple[int]\n    InputShape: Tuple[int]\n    \n    outchannels: int\n    depth: int\n    BatchSize: int\n    train: bool = True \n    \n    def setup(self):\n        \n        self.localConv = CoderCONV(self.Units[1::],self.Ksize,self.Strides,self.depth,'convdecoder',UpSampling=True,train=self.train)\n        self.divFactor = 2**(len(self.Units)-1)\n        \n        self.finalShape = [self.BatchSize]+[val//self.divFactor for val in self.InputShape[0:-1]] + [self.Units[-1]]\n        \n        self.localShape = np.prod(np.array(self.finalShape[1::]))\n        self.DecUnits = [2,self.localShape//16,self.localShape//4,self.localShape]\n        \n        self.localDecoder = DecoderMLP(self.DecUnits,train=self.train)\n        \n        self.outnorm = nn.BatchNorm(use_running_average=not self.train,name = 'outnorm')\n        self.outConv = nn.Conv(self.outchannels,self.Ksize,padding='SAME',use_bias=False,name='conv_dec_out')\n        \n    #@nn.compact\n    def __call__(self,inputs):\n        \n        x = inputs\n        x = self.localDecoder(x)\n        x = jnp.reshape(jnp.array(x),self.finalShape)\n        x = self.localConv(x)\n        x = self.outConv(x)\n        x = self.outnorm(x)\n        x = nn.sigmoid(x)\n        \n        return x\n\ndef reparameterize(rng, mean, logvar):\n    std = jnp.exp(0.5 * logvar)\n    eps = random.normal(rng, logvar.shape)\n    return mean + eps * std\n\nclass ConvVAE(nn.Module):\n    \n    Units: Sequence[int]\n    Ksize: Tuple[int]\n    Strides: Tuple[int]\n    InputShape: Tuple[int]\n    \n    outchannels: int\n    depth: int\n    BatchSize: int\n    train: bool = True \n    \n    def setup(self):\n        self.encoder = CONVEncoder(self.Units,self.Ksize,self.Strides,\n                                   self.InputShape,self.depth,self.BatchSize,\n                                   self.train)\n        self.decoder = CONVDecoder(self.Units[::-1],self.Ksize,self.Strides,\n                                   self.InputShape,self.outchannels,\n                                   self.depth,self.BatchSize,\n                                   self.train)\n\n    def __call__(self, x, z_rng):\n        mean, logvar = self.encoder(x)\n        z = reparameterize(z_rng, mean, logvar)\n        recon_x = self.decoder(z)\n        return recon_x, mean, logvar\n\ndef LoadBatch(paths):\n    \n    container = []    \n    for pth in paths:\n        y, sr = librosa.load(pth)\n        ydem = signal.decimate(y, 8)\n        ydem = (ydem - ydem.min())/(ydem.max() - ydem.min())\n        toAdd = (24*24*24)-len(ydem)\n        yout = np.append(ydem,[0 for k in range(toAdd)])\n        yout = yout.reshape(24,24,24,1)\n        container.append(yout)\n    container = np.stack(container)\n    return container\n\ndef DataLoader(datadirs,batchsize):\n    for k in range(0,len(datadirs),batchsize):\n        paths = datadirs[k:k+batchsize]\n        yield LoadBatch(paths)\n\nsh = 10**-6\n\n@jax.vmap\ndef kl_divergence(mean, logvar):\n    return -0.5 * sh * jnp.sum(1 + logvar - jnp.square(mean) - jnp.exp(logvar))\n\ndef MainLoss(Model,params,batchStats,z_rng ,batch):\n  \n    block, newbatchst = Model().apply({'params': params, 'batch_stats': batchStats}, batch, z_rng,mutable=['batch_stats'])\n    recon_x, mean, logvar = block\n    kld_loss = kl_divergence(mean, logvar).mean()\n    loss_value = optax.l2_loss(recon_x, batch).mean()\n    total_loss = loss_value + kld_loss\n    \n    return total_loss,newbatchst['batch_stats']\n\ndef TrainModel(TrainData,TestData,Loss,params,batchStats,rng,epochs=10,batch_size=64,lr=0.005):\n    \n    totalSteps = epochs*(TrainData.shape[0]//batch_size) + epochs\n    stepsPerCycle = totalSteps//4\n\n    esp = [{\"init_value\":lr/10, \n            \"peak_value\":(lr)/((k+1)), \n            \"decay_steps\":int(stepsPerCycle*0.75), \n            \"warmup_steps\":int(stepsPerCycle*0.25), \n            \"end_value\":lr/10} for k in range(4)]\n    \n    Scheduler = optax.sgdr_schedule(esp)\n    localOptimizer = optax.chain(optax.clip_by_global_norm(1.0),\n                                 optax.scale_by_adam(),  \n                                 optax.scale_by_schedule(Scheduler), \n                                 optax.scale(-1.0)\n                                 )\n\n    optState = localOptimizer.init(params)\n    \n    @jax.jit\n    def step(params,batchStats ,optState, z_rng, batch):\n        \n        (loss_value,batchStats), grads = jax.value_and_grad(Loss,has_aux=True)(params,batchStats, z_rng, batch)\n        updates, optState = localOptimizer.update(grads, optState, params)\n        params = optax.apply_updates(params, updates)\n        \n        return params,batchStats, optState, loss_value\n    \n    @jax.jit\n    def getloss(params,batchStats, z_rng, batch):\n        (loss_value,_), _ = jax.value_and_grad(Loss,has_aux=True)(params,batchStats, z_rng, batch)\n        return loss_value\n    \n    trainloss = []\n    testloss = []\n    \n    for epoch in range(epochs):\n        \n        st = time.time()\n        batchtime = []\n        losses = []\n        \n        for k in range(0,len(TrainData),batch_size):\n\n            stb = time.time()\n            batchPaths = TrainData[k:k+batch_size]\n            batch = LoadBatch(batchPaths)\n        \n            rng, key = random.split(rng)\n            params,batchStats ,optState, lossval = step(params,batchStats,optState,key,batch)\n            losses.append(lossval)\n            batchtime.append(time.time()-stb)\n\n        \n        valloss = []\n        for i in range(0,len(TestData),batch_size):\n            \n            rng, key = random.split(rng)\n            val_batchPaths = TestData[i:i+batch_size]\n            val_batch = LoadBatch(val_batchPaths)\n            valloss.append(getloss(params,batchStats,key,val_batch))\n        \n        mbatch = 1000*np.mean(batchtime)\n        meanloss = np.mean(losses)\n        meanvalloss = np.mean(valloss)\n        \n        trainloss.append(meanloss)\n        testloss.append(meanvalloss)\n        np.random.shuffle(TrainData)\n    \n        end = time.time()\n        output = 'Epoch = '+str(epoch) + ' Time per epoch = ' + str(round(end-st,3)) + 's  Time per batch = ' + str(round(mbatch,3)) + 'ms' + ' Train Loss = ' + str(meanloss) +' Test Loss = ' + str(meanvalloss)\n        print(output)\n        \n    return trainloss,testloss,params,batchStats\n","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loading the data","metadata":{}},{"cell_type":"code","source":"filenames = pd.read_csv('/kaggle/input/bengalifilenames/names.csv')\nfilenames = filenames[0:100000]\n\nXtrain,Xtest,_,_ = train_test_split(filenames['name'].values,filenames['name'].values,test_size=0.1,random_state=42) \n\nbatchSize = 16\nInputShape = (24,24,24,1)\noutChannels = 1\ndepth = 3\n\nmainUnits = [20,24,24,20]\nXtrain = Xtrain[0:batchSize*(Xtrain.shape[0]//batchSize)]\nXtest = Xtest[0:batchSize*(Xtest.shape[0]//batchSize)]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training the model","metadata":{}},{"cell_type":"code","source":"basePath = r'/kaggle/input/bengaliai-speech/train_mp3s'\nnp.random.seed(128)\n\ntrainData = np.array([os.path.join(basePath,val) for val in Xtrain])\ntestData = np.array([os.path.join(basePath,val) for val in Xtest])\n    \ndef VAEModel():\n    return ConvVAE(mainUnits,(3,3,3),(2,2,2),InputShape,outChannels,depth,batchSize)\n     \ndef loss(params,batchStats,z_rng ,batch):\n    return MainLoss(VAEModel,params,batchStats,z_rng ,batch)\n     \nrng = random.PRNGKey(0)\nrng, key = random.split(rng)\n     \nfinalShape = tuple([batchSize]+list(InputShape))\ninit_data = jnp.ones(finalShape, jnp.float32)\ninitModel = VAEModel().init(key, init_data, rng)\n    \nparams0 = initModel['params']\nbatchStats = initModel['batch_stats']\n\ntrloss,tstloss,params0,batchStats = TrainModel(trainData,testData,loss,params0,\n                                               batchStats,rng,lr=0.0025,\n                                               epochs=10,batch_size=batchSize)\n    \nstate = {'params':params0,'batch_stats':batchStats}\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Audio clustering","metadata":{}},{"cell_type":"code","source":"def TransformData(Model,Data,Bsize=10000):\n\n    VariationalRep = []\n    rng = random.PRNGKey(0)\n    \n    for batch in DataLoader(Data,Bsize):\n        \n        mu,logvar = Model(batch)\n        varfrag = reparameterize(rng,mu,logvar)\n        VariationalRep.append(varfrag)\n        rng, key = random.split(rng)\n    \n    VariationalRep = np.vstack(VariationalRep)\n    \n    return VariationalRep\n\nlocalparams = {'params':state['params']['encoder'],'batch_stats':state['batch_stats']['encoder']}\n\ndef EncoderModel(batch):\n    return CONVEncoder(mainUnits,(3,3,3),(2,2,2),InputShape,depth,batchSize,train=False).apply(localparams, batch)\n\nVariationalRep = TransformData(EncoderModel,testData,batchSize)","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"minidata = filenames[filenames['name'].isin(Xtest)].copy()\nminidata = minidata.set_index('name')\nminidata = minidata.loc[Xtest]\n\nminidata['dim0'] = VariationalRep[:,0]\nminidata['dim1'] = VariationalRep[:,1]\n\nplt.figure(figsize=(15,7))\nax = plt.gca()\nminidata.plot.scatter(x='dim0',y='dim1',alpha=0.25,color='navy',ax=ax)\nPlotStyle(ax)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Clustering","metadata":{}},{"cell_type":"code","source":"from scipy.signal import argrelextrema\nfrom sklearn.neighbors import KernelDensity\n\nminval = VariationalRep[:,0].min()\nmaxval = VariationalRep[:,0].max()\n\nkde = KernelDensity(kernel='gaussian', bandwidth=0.00001).fit(VariationalRep[:,0].reshape(-1,1))\ns = np.linspace(minval,maxval)\ne = kde.score_samples(s.reshape(-1,1))\n\nmi = argrelextrema(e, np.less)[0]\npositions = [-1+minval]+[s[val] for val in mi]+[1+maxval]\n\ndef GetCluster(sample,positions=positions):\n    \n    out = -1\n    for k in range(len(positions)-1):\n        if sample <= positions[k+1] and sample >= positions[k]:\n            out = k\n            break\n    return out\n\nminidata['cluster'] = minidata['dim0'].map(GetCluster)","metadata":{"_kg_hide-input":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,7))\nax = plt.gca()\nminidata.plot.scatter(x='dim0',y='dim1',alpha=0.25,c='cluster',cmap='Set1',ax=ax)\nPlotStyle(ax)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,7))\nax = plt.gca()\nminidata['cluster'].value_counts(sort=False).plot.bar(ax=ax)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Audio examples ","metadata":{}},{"cell_type":"code","source":"from IPython.display import Audio\nclusters = minidata['cluster'].value_counts().index[0:4]","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loopdata = minidata[minidata['cluster']==clusters[0]]\nsample = np.random.choice(loopdata.index.values)\nAudio(os.path.join(basePath,sample))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loopdata = minidata[minidata['cluster']==clusters[1]]\nsample = np.random.choice(loopdata.index.values)\nAudio(os.path.join(basePath,sample))    ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loopdata = minidata[minidata['cluster']==clusters[2]]\nsample = np.random.choice(loopdata.index.values)\nAudio(os.path.join(basePath,sample))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loopdata = minidata[minidata['cluster']==clusters[3]]\nsample = np.random.choice(loopdata.index.values)\nAudio(os.path.join(basePath,sample))","metadata":{},"execution_count":null,"outputs":[]}]}