{"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"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":51294,"databundleVersionId":6923401,"sourceType":"competition"}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 3D Convolution for reactivity prediction","metadata":{}},{"cell_type":"code","source":"import time\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\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 StratifiedKFold\nfrom sklearn.model_selection import train_test_split\n\nclass Regressor(nn.Module):\n    \n    train: bool = True \n    \n    @nn.compact\n    def __call__(self,inputs):\n        \n        x = inputs\n        x = nn.Conv(16,(3,3,3),padding='SAME',use_bias=False,name='conv3D0')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D0')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(16,(3,3,3),strides=(1,1,2),padding='SAME',use_bias=False,name='conv3D2')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D2')(x)\n        x = nn.leaky_relu(x)\n        \n        x = x.reshape((x.shape[0],6,6,48))\n        x = nn.Conv(32,(3,3),padding='SAME',use_bias=False,name='conv2D0')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv2D0')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(24,(3,3),padding='SAME',use_bias=False,name='conv2D2')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv2D2')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(24,(3,3),strides=(2,2),padding='SAME',use_bias=False,name='conv2D4')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv2D4')(x)\n        x = nn.leaky_relu(x)\n        \n        x = x.reshape((x.shape[0],-1))\n        x = nn.Dense(206,use_bias=True,name='out')(x)\n\n        return x\n\ndef MakeSequenceEncoding(bSequence,tokenDict=None):\n    \n    stringFrags = [val for val in bSequence]\n    nToAdd = (6*6*6) - len(stringFrags)\n    \n    encoded = [tokenDict[val] for val in stringFrags] \n\n    toaddvec = [0 for _ in range(len(encoded[0]))]\n    toAdd = [toaddvec for k in range(nToAdd)]\n    \n    encoded = np.array(encoded + toAdd).reshape(6,6,6,len(toaddvec))\n    \n    return encoded\n\ndef LoadBatch(data,tokenDict):\n    \n    xdata,ydata = data\n    container = []\n    for pth in xdata:\n        encoded = MakeSequenceEncoding(pth,tokenDict=tokenDict)\n        container.append(encoded)\n    container = np.stack(container)\n    \n    return container.astype(np.float32),ydata\n\ndef DataLoader(data,batchsize,tokenDict):\n    xdata,ydata = data\n    for k in range(0,len(xdata),batchsize):\n        localdata = [xdata[k:k+batchsize],ydata[k:k+batchsize]]\n        yield LoadBatch(localdata,tokenDict)\n\ndef MainLoss(Model,params,batchStats,batch):\n  \n    xbatch,ybatch = batch\n    \n    block, newbatchst = Model().apply({'params': params, 'batch_stats': batchStats}, xbatch,mutable=['batch_stats'])\n    recon_x = block\n    \n    loss_recon = optax.huber_loss(recon_x, ybatch).mean()\n    \n    total_loss = loss_recon\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,tokens=None):\n    \n    totalSteps = epochs*(TrainData[0].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    \n    optState = localOptimizer.init(params)\n    \n    @jax.jit\n    def step(params,batchStats ,optState, batch):\n        \n        (loss_value,batchStats), grads = jax.value_and_grad(Loss,has_aux=True)(params,batchStats, 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, batch):\n        (loss_value,_), _ = jax.value_and_grad(Loss,has_aux=True)(params,batchStats, 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 batch in DataLoader(TrainData,batch_size,tokens):\n\n            stb = time.time()\n        \n            rng, key = random.split(rng)\n            params,batchStats ,optState, lossval = step(params,batchStats,optState,batch)\n            losses.append(lossval)\n            batchtime.append(time.time()-stb)\n\n        \n        valloss = []\n        for val_batch in DataLoader(TestData,batch_size,tokens):\n            \n            rng, key = random.split(rng)\n            valloss.append(getloss(params,batchStats,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        index = np.arange(len(TrainData[0]))\n        np.random.shuffle(index)\n        TrainData[0] = TrainData[0][index]\n        TrainData[1] = TrainData[1][index]\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":"## Model architecture","metadata":{}},{"cell_type":"code","source":"batchSize = 16\nInputShape = (6,6,6,4)\n\nx0 = jnp.ones((16,6,6,6,4))\nprint(Regressor().tabulate(jax.random.key(0), x0,console_kwargs={'width':150}))\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data loading","metadata":{}},{"cell_type":"code","source":"data = pd.read_csv('/kaggle/input/stanford-ribonanza-rna-folding/train_data_QUICK_START.csv')\ndata = data.fillna(0)\n\ndataseqs = data.groupby('sequence_id')['sequence'].unique().to_frame()\ndataseqs['sequence'] = [val[0] for val in dataseqs['sequence']]\ndataseqs['size'] = dataseqs['sequence'].apply(len)\n\ncats = ['2A3_MaP', 'DMS_MaP']\n\nexp0 = data[data['experiment_type']==cats[0]]\nexp0 = exp0.set_index('sequence_id')\n\nexp1 = data[data['experiment_type']==cats[1]]\nexp1 = exp1.set_index('sequence_id')\n\ndatacols = exp0.columns[3:209]","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_train, X_test, _, _ = train_test_split(dataseqs.index,dataseqs.index,\n                                         test_size=0.1, random_state=42,\n                                         stratify=dataseqs['size'].values)\n\nX_train = X_train[0:50000]\nX_test = X_test\n\nydata0 = exp0[datacols].loc[X_train].values\nydata1= exp1[datacols].loc[X_train].values\n\nyval0 = exp0[datacols].loc[X_test].values\nyval1 = exp1[datacols].loc[X_test].values\n\nagroup_y = [ydata0,yval0,ydata1,yval1]\n\nagroup_x = [dataseqs['sequence'].loc[X_train].values, dataseqs['sequence'].loc[X_test].values,dataseqs['size'].loc[X_train].values]\n","metadata":{"_kg_hide-input":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model training","metadata":{}},{"cell_type":"code","source":"def InitAndTrainModel(train,test,tokens):\n    \n    xtr,ytr = train\n    xtst,ytst = test\n    \n    def Model():\n        return Regressor()\n \n    def loss(params,batchStats,batch):\n        return MainLoss(Model,params,batchStats,batch)\n \n    rng = random.PRNGKey(0)\n    rng, key = random.split(rng)\n \n    finalShape = tuple([batchSize]+list(InputShape))\n    init_data = jnp.ones(finalShape, jnp.float32)\n    initModel = Model().init(key, init_data)\n\n    params0 = initModel['params']\n    batchStats = initModel['batch_stats']\n\n    trloss,tstloss,params0,batchStats = TrainModel([xtr,ytr],[xtst,ytst],loss,params0,\n                                                   batchStats,rng,lr=0.005,\n                                                   epochs=10,batch_size=batchSize,\n                                                   tokens=tokens)\n    \n    return trloss,tstloss,params0,batchStats\n\ndef ModelPrediction(Model,data,batch_size,tokens):\n    container = []\n    for batch in DataLoader(data,batch_size,tokens):\n        xd,yd = batch\n        preds = Model(xd)\n        container.append(preds)\n    return np.vstack(container)\n\ndef MAEClipped(y,ypred):\n    \n    yy = np.array([max(min(val,1),0) for val in y])\n    yypred = np.array([max(min(val,1),0) for val in ypred])\n    \n    return np.mean(np.abs(yy-yypred))\n    \ndef TestSequenceEncodings(group_x,group_y,tokens):\n    \n    aydata0,ayval0,aydata1,ayval1 = group_y\n    \n    nsplits = 5\n    kf = StratifiedKFold(n_splits=nsplits,random_state=1,shuffle=True)\n    \n    fig,axs = plt.subplots(nsplits+1,2,figsize=(16,16))\n    \n    for i, (train_index, test_index) in enumerate(kf.split(group_x[0],group_x[2])):\n        \n        xtrain,xtest = group_x[0][train_index],group_x[0][test_index]\n        aytrain0,aytest0 = aydata0[train_index],aydata0[test_index]\n        aytrain1,aytest1 = aydata1[train_index],aydata1[test_index]\n        \n        trdata0 = InitAndTrainModel([xtrain,aytrain0],[xtest,aytest0],tokens)\n        trdata1 = InitAndTrainModel([xtrain,aytrain1],[xtest,aytest1],tokens)\n        \n        axs[0,0].plot(trdata0[0],color='navy')\n        axs[0,0].plot(trdata0[1],color='black',alpha=0.25)\n        axs[0,1].plot(trdata1[0],color='navy')\n        axs[0,1].plot(trdata1[1],color='black',alpha=0.25)\n        \n        params0 = {'params':trdata0[2],'batch_stats':trdata0[3]}\n        params1 = {'params':trdata1[2],'batch_stats':trdata1[3]}\n    \n        def Model0(batch):\n            return Regressor(train=False).apply(params0, batch)\n        \n        def Model1(batch):\n            return Regressor(train=False).apply(params1, batch)\n        \n        dt0 = ModelPrediction(Model0,[group_x[1],ayval0],batchSize,tokens)\n        serror0 = np.mean(np.abs(ayval0.ravel()-dt0.ravel()))\n        error0 = MAEClipped(ayval0.ravel(),dt0.ravel())\n        \n        axs[i+1,0].scatter(dt0.ravel(),ayval0.ravel())\n        axs[i+1,0].text(0.05,0.85,'val MAE = '+str(round(serror0,3)),transform = axs[i+1,0].transAxes)\n        axs[i+1,0].text(0.05,0.65,'val clip MAE = '+str(round(error0,3)),transform = axs[i+1,0].transAxes)\n        axs[i+1,0].set_title('Fold = '+str(i))\n        \n        dt1 = ModelPrediction(Model1,[group_x[1],ayval1],batchSize,tokens)\n        serror1 = np.mean(np.abs(ayval1.ravel()-dt1.ravel()))\n        error1 = MAEClipped(ayval1.ravel(),dt1.ravel())\n        \n        axs[i+1,1].scatter(dt1.ravel(),ayval1.ravel())\n        axs[i+1,1].text(0.05,0.85,'val MAE = '+str(round(serror1,3)),transform = axs[i+1,1].transAxes)\n        axs[i+1,1].text(0.05,0.65,'val clip MAE = '+str(round(error1,3)),transform = axs[i+1,1].transAxes)\n        axs[i+1,1].set_title('Fold = '+str(i))\n","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## One hot encoded sequences¶","metadata":{}},{"cell_type":"code","source":"Alphabet = ['A','C','U','G']\n\nTokenDictionary = {}\n\nfor k,val in enumerate(Alphabet):\n    currentVec = [0 for j in range(len(Alphabet))]\n    currentVec[k] = 1\n    TokenDictionary[val]=currentVec\n\nTestSequenceEncodings(agroup_x,agroup_y,TokenDictionary)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Basic chemical properties\nFirst two element lists encodes for chemical structure(purines or pyrimidines), second one number bonds","metadata":{}},{"cell_type":"code","source":"TokenDictionary2 = {}\n\nTokenDictionary2['A'] = [1,0] + [0,1]\nTokenDictionary2['C'] = [0,1] + [1,0]\nTokenDictionary2['U'] = [0,1] + [0,1]\nTokenDictionary2['G'] = [1,0] + [1,0]\n\nTestSequenceEncodings(agroup_x,agroup_y,TokenDictionary2)","metadata":{},"execution_count":null,"outputs":[]}]}