{"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":"# Sequence clustering with JAX","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\ndata = pd.read_csv('/kaggle/input/stanford-ribonanza-rna-folding/train_data.csv',usecols=['sequence_id','sequence'])\n\nuniquedata = data.groupby('sequence_id')['sequence'].value_counts().index\n\nuniquenames = [val[0] for val in uniquedata]\nuniqueseqs = [val[1] for val in uniquedata]\n\ndataframe = pd.DataFrame()\ndataframe['id'] = uniquenames\ndataframe['seq'] = uniqueseqs","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport time\nimport numpy as np\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\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::],'encoderlmlp',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='decoder_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\n###############################################################################\n# Loading packages \n###############################################################################\n\ndef LoadBatch(paths):\n    \n    container = []    \n    for pth in paths:\n        container.append(np.load(pth))\n    container = np.stack(container)\n    \n    return container.reshape((-1,16,16,1))\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        \n###############################################################################\n# Loading packages \n###############################################################################\n\nsh = 10**-7\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    #\n    localOptimizer = optax.chain(optax.clip_by_global_norm(1),\n                                 optax.scale_by_adam(),  \n                                 optax.scale_by_schedule(Scheduler), \n                                 optax.scale(-1.0)\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        valloss = []\n        for i in range(0,len(TestData),batch_size):            \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":"code","source":"import matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\n\nbatchSize = 32\n\noutdir = '/kaggle/input/4merribonanza/4mer/train/'\n\nInxtrain,Inxtest, _, _ = train_test_split(uniquenames, uniquenames, test_size=0.10, random_state=42)\n\ntrainSamps = np.array([outdir+val+'.npy' for val in Inxtrain[0:12000*batchSize]])\ntestSamps = np.array([outdir+val+'.npy' for val in Inxtest[0:2000*batchSize]])\n\ndatadirs = np.array([outdir+val+'.npy' for val in uniquenames])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nInputShape = (16,16,1)\n\ndepth = 2\nmainUnits = [8,12,12,8]\n\nnp.random.seed(4128)\n\ntrainData = trainSamps[0:batchSize*(trainSamps.shape[0]//batchSize)]\ntestData = testSamps[0:batchSize*(testSamps.shape[0]//batchSize)]\n\ndef VAEModel():\n    return ConvVAE(mainUnits,(3,3),(2,2),InputShape,1,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\nfinalParams = {'params':params0,'batch_stats':batchStats}\nlocalparams = {'params':finalParams['params']['encoder'],'batch_stats':finalParams['batch_stats']['encoder']}\n\ndef EncoderModel(batch):\n    return CONVEncoder(mainUnits,(3,3),(2,2),InputShape,depth,batchSize,train=False).apply(localparams, batch)","metadata":{},"execution_count":null,"outputs":[]},{"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\nbottleneck = TransformData(EncoderModel,testSamps,batchSize)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\nax = plt.gca()\nax.scatter(bottleneck[:,0],bottleneck[:,1],alpha=0.1)","metadata":{},"execution_count":null,"outputs":[]}]}