{"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":"# From sequence to 2D structure with 3D convolution","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 train_test_split\n\ndef HammingVec(seq1,seq2):\n    return [0 if val==sal else 1 for val,sal in zip(seq1,seq2)]\n\ndef Hamming(vec):\n    return sum(vec)/len(vec)\n\ndef NumberOfUniqueCharacters(seq):\n    try:\n        return len(set([val for val in seq]))\n    except TypeError:\n        return 0","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Exploratory analysis ","metadata":{}},{"cell_type":"code","source":"dta = pd.read_csv('/kaggle/input/stanford-ribonanza-rna-folding/supplementary_silico_predictions/R1_silico_predictions.csv')\n\nstructcols = ['vienna2_mfe','contrafold2_mfe','eternafold_mfe','hotknots','ipknots','knotty','spotrna','nupack_pk', 'vienna_2[threshknot]','vienna_2[hungarian]','eternafold[threshknot]','eternafold[hungarian]',\n              'contrafold_2[threshknot]','contrafold_2[hungarian]','nupack[threshknot]','nupack[hungarian]','nupack-pk[threshknot]','nupack-pk[hungarian]','shapify-hfold']\n\ncont =[]\nfor val in structcols:\n    cont.append(dta[val].unique().shape[0])\n\nplt.figure()\nax = plt.gca()\nax.bar(np.arange(len(cont)),cont)\nax.set_xticks(np.arange(len(cont)),structcols,rotation=80)\nax.set_ylabel('Unique Structures')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure()\nax = plt.gca()\ndta[structcols].isna().sum().plot.bar(ax=ax)\nax.set_ylabel('Missing values')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig,axs = plt.subplots(5,4,figsize=(15,15),sharex=True)\naxs = axs.ravel()\n\nfor val,ax in zip(structcols,axs):\n    dta[val].apply(NumberOfUniqueCharacters).hist(ax=ax)\n    ax.set_xlabel('Number oif Unique Characters')\n    ax.set_ylabel('Frequency')\n    ax.set_title(val)\nplt.tight_layout()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"selectedcols = ['contrafold2_mfe','eternafold_mfe']\n\nhvecs = []\nhdist = []\nfor val,sal in zip(dta['contrafold2_mfe'],dta['eternafold_mfe']):\n    vec = HammingVec(val, sal)\n    hvecs.append(vec)\n    hdist.append(Hamming(vec))\n\nfig,axs = plt.subplots(1,2,figsize=(15,7))\n\naxs[0].plot(hdist)\naxs[0].set_title(selectedcols[0] + ' vs ' + selectedcols[1])\naxs[0].set_ylabel('Hamming Distance')\n\naxs[1].hist(hdist,bins=50)\naxs[1].set_title(selectedcols[0] + ' vs ' + selectedcols[1])\naxs[1].set_xlabel('Hamming Distance')\naxs[1].set_ylabel('Frequency')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Discrepancy between methods","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(15,15))\nax = plt.gca()\nax.imshow(hvecs,aspect='auto')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3D convolution model ","metadata":{}},{"cell_type":"code","source":"class Regressor(nn.Module):\n    \n    train: bool = True \n    \n    @nn.compact\n    def __call__(self,inputs):\n        \n        x = inputs\n        \n        x = nn.Conv(48,(3,3,3),padding='SAME',use_bias=False,name='conv3D1')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D1')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(24,(3,3,3),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 = nn.Conv(12,(3,3,3),padding='SAME',use_bias=False,name='conv3D5')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D5')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(6,(3,3,3),padding='SAME',use_bias=False,name='conv3D6')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D6')(x)\n        x = nn.leaky_relu(x)\n                \n        x = nn.Conv(3,(3,3,3),padding='SAME',use_bias=False,name='conv3D8')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D8')(x)\n        x = nn.softmax(x)\n        \n        return x\n\nAlphabet = ['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\nAlphabet2 = ['(',')','.']\n\nTokenDictionary2 = {}\n\nfor k,val in enumerate(Alphabet2):\n    currentVec = [0 for j in range(len(Alphabet2))]\n    currentVec[k] = 1\n    TokenDictionary2[val]=currentVec\n\ndef MakeSequenceEncoding(bSequence,tokenDict=None):\n    \n    stringFrags = [val for val in bSequence]\n    nToAdd = (6*6*6) - len(stringFrags)\n    encoded = [tokenDict[val] for val in stringFrags] \n    toaddvec = [0 for _ in range(len(encoded[0]))]\n    toAdd = [toaddvec for k in range(nToAdd)]    \n    encoded = np.array(encoded + toAdd).reshape(6,6,6,len(toaddvec))\n    \n    return encoded\n\ndef LoadBatch(data):\n    \n    xdata,ydata = data\n    encx = []\n    ency = []\n    for aseq,bseq in zip(xdata,ydata):\n        encx.append(MakeSequenceEncoding(aseq,tokenDict=TokenDictionary))\n        ency.append(MakeSequenceEncoding(bseq,tokenDict=TokenDictionary2))        \n    encx = np.stack(encx)\n    ency = np.stack(ency)\n    \n    return encx,ency\n\ndef DataLoader(data,batchsize):\n    xdata,ydata = data\n    for k in range(0,len(xdata),batchsize):\n        xfrag,yfrag = xdata[k:k+batchsize],ydata[k:k+batchsize]\n        yield xfrag,yfrag\n\ndef MainLoss(Model,params,batchStats,batch):\n  \n    xbatch,ybatch = batch\n    block, newbatchst = Model().apply({'params': params, 'batch_stats': batchStats}, xbatch,mutable=['batch_stats'])\n    recon_x = block    \n    loss_recon = optax.sigmoid_binary_cross_entropy(recon_x, ybatch).mean()\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):\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    index = np.arange(len(TrainData[0]))\n    \n    for epoch in range(epochs):\n        \n        st = time.time()\n        batchtime = []\n        losses = []    \n        \n        xtrd,ytrd = TrainData\n        \n        for batch in DataLoader([xtrd[index],ytrd[index]],batch_size):\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):\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        np.random.shuffle(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\ndef ModelPrediction(Model,data,batch_size):\n    container = []\n    for batch in DataLoader(data,batch_size):\n        xd,yd = batch\n        preds = Model(xd)\n        container.append(preds)\n    return np.vstack(container)\n","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dta = dta.set_index('rowID')\n\nX_train, X_test, _, _ = train_test_split(dta.index,dta.index,test_size=0.1, random_state=42)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Contrafold2 data","metadata":{}},{"cell_type":"code","source":"trainData = LoadBatch([dta['sequence'].loc[X_train].values,dta[selectedcols[0]].loc[X_train].values])\ntestData = LoadBatch([dta['sequence'].loc[X_test].values,dta[selectedcols[0]].loc[X_test].values])\n\nbatchSize = 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\ndef Model():\n    return Regressor()\n\ndef loss(params,batchStats,batch):\n    return MainLoss(Model,params,batchStats,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 = Model().init(key, init_data)\n\nparams0 = initModel['params']\nbatchStats = initModel['batch_stats']\n\ntrloss,tstloss,params0,batchStats = TrainModel(trainData,testData,loss,params0,batchStats,rng,lr=0.005,epochs=20,batch_size=batchSize)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fprms = {'params':params0,'batch_stats':batchStats}\n\ndef Model0(batch):\n    return Regressor(train=False).apply(fprms, batch)\n\ndt0 = ModelPrediction(Model0,testData,batchSize)\ndt0 = dt0.reshape(dt0.shape[0],-1,3).argmax(axis=-1)\npreds0 = [''.join(Alphabet2[val] for val in sal) for sal in dt0]\n\nacont1 = []\ncont2 = []\nfor val,sal in zip(dta[selectedcols[0]].loc[X_test].values,preds0):\n    vec = HammingVec(val, sal[0:177])\n    acont1.append(vec)\n    cont2.append(Hamming(vec))\n    \nfig,axs = plt.subplots(1,2,figsize=(15,7))\n\naxs[0].plot(cont2)\naxs[0].set_title(selectedcols[1])\naxs[0].set_ylabel('Hamming Distance')\naxs[0].set_xlabel('Sequences')\n\naxs[1].hist(cont2,bins=50)\naxs[1].set_title(selectedcols[1])\naxs[1].set_xlabel('Hamming Distance')\naxs[1].set_ylabel('Frequency')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,15))\nax = plt.gca()\nax.imshow(acont1,aspect='auto')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Eternafold data","metadata":{}},{"cell_type":"code","source":"trainData = LoadBatch([dta['sequence'].loc[X_train].values,dta[selectedcols[1]].loc[X_train].values])\ntestData = LoadBatch([dta['sequence'].loc[X_test].values,dta[selectedcols[1]].loc[X_test].values])\n\nbatchSize = 16\nInputShape = (6,6,6,4)\n\ndef Model():\n    return Regressor()\n\ndef loss(params,batchStats,batch):\n    return MainLoss(Model,params,batchStats,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 = Model().init(key, init_data)\n\nparams0 = initModel['params']\nbatchStats = initModel['batch_stats']\n\ntrloss,tstloss,params0,batchStats = TrainModel(trainData,testData,loss,params0,batchStats,rng,lr=0.005,epochs=20,batch_size=batchSize)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fprms = {'params':params0,'batch_stats':batchStats}\n\ndef Model0(batch):\n    return Regressor(train=False).apply(fprms, batch)\n\ndt0 = ModelPrediction(Model0,testData,batchSize)\ndt0 = dt0.reshape(dt0.shape[0],-1,3).argmax(axis=-1)\npreds1 = [''.join(Alphabet2[val] for val in sal) for sal in dt0]\n\nbcont1 = []\ncont2 = []\nfor val,sal in zip(dta[selectedcols[1]].loc[X_test].values,preds1):\n    vec = HammingVec(val, sal[0:177])\n    bcont1.append(vec)\n    cont2.append(Hamming(vec))\n    \nfig,axs = plt.subplots(1,2,figsize=(15,7))\n\naxs[0].plot(cont2)\naxs[0].set_title(selectedcols[1])\naxs[0].set_ylabel('Hamming Distance')\naxs[0].set_xlabel('Sequences')\n\naxs[1].hist(cont2,bins=50)\naxs[1].set_title(selectedcols[1])\naxs[1].set_xlabel('Hamming Distance')\naxs[1].set_ylabel('Frequency')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(15,15))\nax = plt.gca()\nax.imshow(bcont1,aspect='auto')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Discrepancy locations between different models ","metadata":{}},{"cell_type":"code","source":"fig,axs = plt.subplots(1,3,figsize=(15,7))\n\naxs[0].imshow(hvecs,aspect='auto')\naxs[0].set_title(selectedcols[0] + ' vs ' + selectedcols[1])\naxs[1].imshow(acont1,aspect='auto')\naxs[1].set_title(selectedcols[0] + ' vs 3D conv model')\naxs[2].imshow(bcont1,aspect='auto')\naxs[2].set_title(selectedcols[1] + ' vs 3D conv model')","metadata":{},"execution_count":null,"outputs":[]}]}