{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":51294,"databundleVersionId":6923401,"sourceType":"competition"},{"sourceId":6837645,"sourceType":"datasetVersion","datasetId":3931165},{"sourceId":7080526,"sourceType":"datasetVersion","datasetId":4078782}],"dockerImageVersionId":30587,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 3D structure and convolutions for RNA data","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nids = pd.read_csv('/kaggle/input/ribonanza-validids/validids.csv')\nlocations = ['RhoFold_PDBs/' + val + '/unrelaxed_model.pdb' for val in ids['sequence_id']]","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:47:03.206758Z","iopub.execute_input":"2023-11-29T08:47:03.207604Z","iopub.status.idle":"2023-11-29T08:47:03.639875Z","shell.execute_reply.started":"2023-11-29T08:47:03.207558Z","shell.execute_reply":"2023-11-29T08:47:03.638746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import zipfile\n\nwith zipfile.ZipFile('/kaggle/input/ribonanza-3d-coords/RhoFold_PDBs_zip') as z:\n    for file in locations:\n        z.extract(file,'/kaggle/working/')","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:47:03.641846Z","iopub.execute_input":"2023-11-29T08:47:03.642213Z","iopub.status.idle":"2023-11-29T08:49:46.416022Z","shell.execute_reply.started":"2023-11-29T08:47:03.642184Z","shell.execute_reply":"2023-11-29T08:49:46.415160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport time\nimport numpy as np\nimport pandas as pd\nimport Bio.PDB as pdb\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 Bio.SVDSuperimposer import SVDSuperimposer\nfrom sklearn.model_selection import train_test_split\n\nclass EncoderNetwork(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(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),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),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 = nn.Conv(24,(3,3,3),padding='SAME',use_bias=False,name='conv3D3')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D3')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(32,(3,3,3),strides=(1,1,2),padding='SAME',use_bias=False,name='conv3D4')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D4')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(32,(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(40,(3,3,3),strides=(1,1,2),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(40,(3,3,3),padding='SAME',use_bias=False,name='conv3D7')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D7')(x)\n        x = nn.leaky_relu(x)\n        \n        x = x.reshape((x.shape[0],16,16,80))\n        \n        x = nn.Conv(64,(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.leaky_relu(x)\n        \n        x = nn.Conv(64,(3,3),padding='SAME',use_bias=False,name='conv3D9')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D9')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(32,(3,3),padding='SAME',use_bias=False,name='conv3D10')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D10')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(32,(3,3),padding='SAME',use_bias=False,name='conv3D11')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D11')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(16,(3,3),padding='SAME',use_bias=False,name='conv3D12')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D12')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(16,(3,3),padding='SAME',use_bias=False,name='conv3D13')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D13')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(16,(3,3),padding='SAME',strides=(2,2),use_bias=False,name='conv3D14')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D14')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(16,(3,3),padding='SAME',strides=(2,2),use_bias=False,name='conv3D15')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D15')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(16,(3,3),padding='SAME',strides=(2,2),use_bias=False,name='conv3D16')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3D16')(x)\n        x = nn.leaky_relu(x)\n    \n        x = x.reshape((x.shape[0],-1))\n        \n        x = nn.Dense(2,use_bias=False,name='dn01')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bndn01')(x)\n        x = nn.leaky_relu(x)\n        \n        mean_x = nn.Dense(2, name='mean')(x)\n        logvar_x = nn.Dense(2, name='logvar')(x)\n\n        return mean_x,logvar_x\n\nclass DecoderNetwork(nn.Module):\n    \n    train: bool = True \n    \n    @nn.compact\n    def __call__(self,inputs):\n        \n        x = inputs\n        \n        x = nn.Dense(64,use_bias=False,name='ddn01')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbndn01')(x)\n        x = nn.leaky_relu(x)\n        x = x.reshape((x.shape[0],2,2,16))\n        \n        x = nn.ConvTranspose(16,(3,3),padding='SAME',strides=(2,2),use_bias=False,name='dconv3D01')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D01')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.ConvTranspose(16,(3,3),padding='SAME',strides=(2,2),use_bias=False,name='dconv3D02')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D02')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.ConvTranspose(16,(3,3),padding='SAME',strides=(2,2),use_bias=False,name='dconv3D03')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D03')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(32,(3,3),padding='SAME',use_bias=False,name='dconv3D04')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dnconv3D04')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(32,(3,3),padding='SAME',use_bias=False,name='dconv3D05')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D05')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(64,(3,3),padding='SAME',use_bias=False,name='dconv3D06')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D06')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(64,(3,3),padding='SAME',use_bias=False,name='dconv3D07')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D07')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(80,(3,3),padding='SAME',use_bias=False,name='dconv3D08')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D08')(x)\n        x = nn.leaky_relu(x)\n        \n        x = x.reshape((x.shape[0],16,16,2,40))\n        \n        x = nn.Conv(40,(3,3,3),padding='SAME',use_bias=False,name='dconv3D09')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D09')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.ConvTranspose(40,(3,3,3),strides=(1,1,2),padding='SAME',use_bias=False,name='dconv3D10')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D10')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(32,(3,3,3),padding='SAME',use_bias=False,name='dconv3D11')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D11')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.ConvTranspose(32,(3,3,3),strides=(1,1,2),padding='SAME',use_bias=False,name='dconv3D12')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D12')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(24,(3,3,3),padding='SAME',use_bias=False,name='dconv3D13')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D13')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.ConvTranspose(24,(3,3,3),strides=(1,1,2),padding='SAME',use_bias=False,name='dconv3D14')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D14')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(16,(3,3,3),padding='SAME',use_bias=False,name='dconv3D15')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D15')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(16,(3,3,3),padding='SAME',use_bias=False,name='dconv3D16')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='dbnconv3D16')(x)\n        x = nn.leaky_relu(x)\n        \n        x = nn.Conv(3,(3,3,3),padding='SAME',use_bias=False,name='conv3Dout')(x)\n        x = nn.BatchNorm(use_running_average=not self.train,name='bnconv3Dout')(x)\n        x = nn.tanh(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    train: bool = True \n    \n    def setup(self):\n        self.encoder = EncoderNetwork(self.train)\n        self.decoder = DecoderNetwork(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 AlignStructures(structurexyz,fixed=None):\n    \n    sup = SVDSuperimposer()\n    maxfrag = 2300\n    sup.set(fixed[0:maxfrag], structurexyz[0:maxfrag])\n    \n    sup.run()\n    \n    rot, tran = sup.get_rotran()\n    atomdata = np.dot(structurexyz, rot) + tran\n    \n    return atomdata\n\ndef GetCoords(structure):\n    \n    container = []\n    for model in structure:\n        for chain in model:\n            for residue in chain:                             \n                for atom in residue:\n                    container.append(atom.get_vector().get_array())\n    container = np.vstack(container)\n    \n    return container\n\ndef MakeStructureEncoding(StructureCoords,toadd=4096):\n    \n    nToAdd = (toadd) - len(StructureCoords)\n    toaddvec = [0 for _ in range(len(StructureCoords[0]))]\n    toAdd = np.array([toaddvec for k in range(nToAdd)])\n    encoded = np.vstack([StructureCoords,toAdd])\n    \n    return encoded.reshape(16,16,16,3)\n\ndef GetStructureencoding(path,fixed=None):\n    \n    parser = pdb.PDBParser()\n    structure = parser.get_structure('test',path)\n    currentcoords = GetCoords(structure)\n    currentcoords = AlignStructures(currentcoords,fixed=fixed)\n    \n    return MakeStructureEncoding(currentcoords)\n\ndef LoadBatch(paths,fixed):\n    \n    container = []\n    for pth in paths:\n        encoded = GetStructureencoding(pth,fixed=fixed)\n        container.append(encoded)\n    container = np.stack(container)\n    \n    return container\n\ndef DataLoader(paths,batchsize):\n    for k in range(0,len(paths),batchsize):\n        localdata = paths[k:k+batchsize]\n        yield localdata\n\nsh = 10**-4\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    \n    kld_loss = kl_divergence(mean, logvar).mean()\n    loss_recon = optax.l2_loss(recon_x, batch).mean()\n    total_loss = loss_recon + 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    \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    order = np.arange(len(TrainData))\n    for epoch in range(epochs):\n        \n        st = time.time()\n        batchtime = []\n        losses = []\n        \n        for batch in DataLoader(TrainData[order],batch_size):\n\n            stb = time.time()\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 val_batch in DataLoader(TestData,batch_size):\n            \n            rng, key = random.split(rng)\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        \n        np.random.shuffle(order)\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":{"iopub.status.busy":"2023-11-29T08:49:46.417607Z","iopub.execute_input":"2023-11-29T08:49:46.417925Z","iopub.status.idle":"2023-11-29T08:49:49.101660Z","shell.execute_reply.started":"2023-11-29T08:49:46.417899Z","shell.execute_reply":"2023-11-29T08:49:49.100745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Arranging the 3D structure data\n\nRNA 3D structures are aligned using a random structure as a fixed comparison point. Then data is scaled to a -1,1 range ","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\ndtapath = '/kaggle/working/RhoFold_PDBs'\nfilenames = os.listdir(dtapath)\n\nfindex = dataseqs.index.intersection(filenames)\ndataseqs = dataseqs.loc[findex]\n\ndataseqs = dataseqs[dataseqs['size']==177]\n\nproblemids = ['027719d4faa7', '04feca622c05', '1b1e24d5c6a5', '1c15c37d3854',\n              '3259e28604f0', '34b6473f45d4', '3693ed3e692e', '7654f95d205b',\n              '7c6212078869', 'b1938d9d5758', 'caff66dbb64d']\n\nnewindex = dataseqs.index.difference(problemids)\ndataseqs = dataseqs.loc[newindex]","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2023-11-29T08:49:49.104177Z","iopub.execute_input":"2023-11-29T08:49:49.104828Z","iopub.status.idle":"2023-11-29T08:50:20.886481Z","shell.execute_reply.started":"2023-11-29T08:49:49.104789Z","shell.execute_reply":"2023-11-29T08:50:20.885608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filepaths = np.array([os.path.join(dtapath, val,'unrelaxed_model.pdb') for val in dataseqs.index])\n\nX_train, X_test, _, _ = train_test_split(filepaths,filepaths,test_size=0.1, random_state=42)\n\n#X_train = X_train[0:100]\n#X_test = X_test[0:100]\n\ntestids =[val[29:41] for val in X_test]\n\nparser = pdb.PDBParser()\nstructure = parser.get_structure('test',X_train[0])\natomdatafixed = GetCoords(structure)\n\ntrainData = LoadBatch(X_train,fixed=atomdatafixed)\ntestData = LoadBatch(X_test,fixed=atomdatafixed)\n\nmaxnorm = np.sqrt(np.sum(trainData.reshape(trainData.shape[0],-1,3)**2,axis=-1)).max()\n\ntrainData = trainData/maxnorm\ntestData = testData/maxnorm","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-11-29T08:50:20.887656Z","iopub.execute_input":"2023-11-29T08:50:20.887952Z","iopub.status.idle":"2023-11-29T08:50:44.038192Z","shell.execute_reply.started":"2023-11-29T08:50:20.887927Z","shell.execute_reply":"2023-11-29T08:50:44.037030Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"parser = pdb.PDBParser()\nstructure0 = parser.get_structure('test',X_train[0])\natomdatafixed0 = GetCoords(structure0)\n\nparser = pdb.PDBParser()\nstructure1 = parser.get_structure('test',X_train[10])\natomdatafixed1 = GetCoords(structure1)\n\n\nfig,axs = plt.subplots(1,3,figsize=(15,7))\n\naxs[0].scatter(atomdatafixed0[:,0],atomdatafixed0[:,1],alpha=0.1,s=15,color='black')\naxs[0].scatter(atomdatafixed1[:,0],atomdatafixed1[:,1],alpha=0.1,s=15,color='red')\n\naligned = AlignStructures(atomdatafixed1,fixed=atomdatafixed0)\n\naxs[1].scatter(atomdatafixed0[:,0],atomdatafixed0[:,1],alpha=0.1,s=15,color='black')\naxs[1].scatter(aligned[:,0],aligned[:,1],alpha=0.1,s=15,color='red')\n\nms = np.sqrt(np.sum(aligned**2,axis=-1)).max()\n\naxs[2].scatter(atomdatafixed0[:,0]/ms,atomdatafixed0[:,1]/ms,alpha=0.1,s=15,color='black')\naxs[2].scatter(aligned[:,0]/ms,aligned[:,1]/ms,alpha=0.1,s=15,color='red')\n\n\naxs[0].set_title('raw')\naxs[1].set_title('aligned')\naxs[2].set_title('scaled')","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:50:44.039640Z","iopub.execute_input":"2023-11-29T08:50:44.040036Z","iopub.status.idle":"2023-11-29T08:50:45.301838Z","shell.execute_reply.started":"2023-11-29T08:50:44.040000Z","shell.execute_reply":"2023-11-29T08:50:45.300804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Variational autoencoder","metadata":{}},{"cell_type":"code","source":"batchSize = 16\nInputShape = (16,16,16,3)\n\ndef Model():\n    return ConvVAE()\n\ndef loss(params,batchStats,z_rng ,batch):\n    return MainLoss(Model,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 = Model().init(key, init_data, rng)\n\nparams0 = initModel['params']\nbatchStats = initModel['batch_stats']\n\ntrloss,tstloss,params0,batchStats = TrainModel(trainData,testData,loss,params0,batchStats,rng,lr=0.0025,epochs=40,batch_size=batchSize)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:50:45.303137Z","iopub.execute_input":"2023-11-29T08:50:45.303483Z","iopub.status.idle":"2023-11-29T08:52:24.028571Z","shell.execute_reply.started":"2023-11-29T08:50:45.303453Z","shell.execute_reply":"2023-11-29T08:52:24.027706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"localparams = {'params':params0['encoder'],'batch_stats':batchStats['encoder']}\n\ndef EncoderModel(batch):\n    return EncoderNetwork(train=False).apply(localparams, batch)\n\ndef 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","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:24.032511Z","iopub.execute_input":"2023-11-29T08:52:24.032827Z","iopub.status.idle":"2023-11-29T08:52:24.039950Z","shell.execute_reply.started":"2023-11-29T08:52:24.032801Z","shell.execute_reply":"2023-11-29T08:52:24.039039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"varrep = TransformData(EncoderModel,testData,Bsize=batchSize)\n\nplt.figure(figsize=(15,15))\nplt.scatter(varrep[:,0],varrep[:,1])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:24.041261Z","iopub.execute_input":"2023-11-29T08:52:24.041582Z","iopub.status.idle":"2023-11-29T08:52:29.291272Z","shell.execute_reply.started":"2023-11-29T08:52:24.041558Z","shell.execute_reply":"2023-11-29T08:52:29.290302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cats = ['2A3_MaP', 'DMS_MaP']\n\nexp0 = data[data['experiment_type']==cats[0]]\nexp0 = exp0.set_index('sequence_id')\nexp0 = exp0.loc[newindex]\n\nexp1 = data[data['experiment_type']==cats[1]]\nexp1 = exp1.set_index('sequence_id')\nexp1 = exp1.loc[newindex]\n\ndatacols = exp0.columns[3:209]","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:29.295754Z","iopub.execute_input":"2023-11-29T08:52:29.296067Z","iopub.status.idle":"2023-11-29T08:52:30.533248Z","shell.execute_reply.started":"2023-11-29T08:52:29.296041Z","shell.execute_reply":"2023-11-29T08:52:30.532463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def MakeReactivityPanel(columns,repdata,colordata):\n    \n    fig,axs = plt.subplots(5,10,figsize=(20,15))\n    axs = axs.ravel()\n    \n    for i,ax in enumerate(axs):\n        z = colordata[columns[i]]\n        ax.scatter(repdata[:,0],repdata[:,1],\n                   c=z,\n                   alpha=0.1)\n        ax.set_title(columns[i])\n        ax.set_axis_off()","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2023-11-29T08:52:30.534434Z","iopub.execute_input":"2023-11-29T08:52:30.535077Z","iopub.status.idle":"2023-11-29T08:52:30.541801Z","shell.execute_reply.started":"2023-11-29T08:52:30.535042Z","shell.execute_reply":"2023-11-29T08:52:30.540721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2A3 MaP reactivity per location ","metadata":{}},{"cell_type":"code","source":"MakeReactivityPanel(datacols[0:50],varrep,exp0.loc[testids])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:30.543073Z","iopub.execute_input":"2023-11-29T08:52:30.543702Z","iopub.status.idle":"2023-11-29T08:52:35.067092Z","shell.execute_reply.started":"2023-11-29T08:52:30.543668Z","shell.execute_reply":"2023-11-29T08:52:35.066092Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[50:100],varrep,exp0.loc[testids])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:35.068359Z","iopub.execute_input":"2023-11-29T08:52:35.068669Z","iopub.status.idle":"2023-11-29T08:52:39.497395Z","shell.execute_reply.started":"2023-11-29T08:52:35.068642Z","shell.execute_reply":"2023-11-29T08:52:39.496399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[100:150],varrep,exp0.loc[testids])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:39.498846Z","iopub.execute_input":"2023-11-29T08:52:39.499515Z","iopub.status.idle":"2023-11-29T08:52:43.505511Z","shell.execute_reply.started":"2023-11-29T08:52:39.499483Z","shell.execute_reply":"2023-11-29T08:52:43.504579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[150:200],varrep,exp0.loc[testids])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:43.506697Z","iopub.execute_input":"2023-11-29T08:52:43.506997Z","iopub.status.idle":"2023-11-29T08:52:47.540778Z","shell.execute_reply.started":"2023-11-29T08:52:43.506972Z","shell.execute_reply":"2023-11-29T08:52:47.539824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DMS MaP reactivity per location ","metadata":{}},{"cell_type":"code","source":"MakeReactivityPanel(datacols[0:50],varrep,exp1.loc[testids])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:47.542119Z","iopub.execute_input":"2023-11-29T08:52:47.542467Z","iopub.status.idle":"2023-11-29T08:52:52.369142Z","shell.execute_reply.started":"2023-11-29T08:52:47.542431Z","shell.execute_reply":"2023-11-29T08:52:52.368217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[50:100],varrep,exp1.loc[testids])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:52.370223Z","iopub.execute_input":"2023-11-29T08:52:52.370527Z","iopub.status.idle":"2023-11-29T08:52:56.922527Z","shell.execute_reply.started":"2023-11-29T08:52:52.370502Z","shell.execute_reply":"2023-11-29T08:52:56.921600Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[100:150],varrep,exp1.loc[testids])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:52:56.923756Z","iopub.execute_input":"2023-11-29T08:52:56.924041Z","iopub.status.idle":"2023-11-29T08:53:00.987774Z","shell.execute_reply.started":"2023-11-29T08:52:56.924016Z","shell.execute_reply":"2023-11-29T08:53:00.986794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[150:200],varrep,exp1.loc[testids])","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:53:00.989279Z","iopub.execute_input":"2023-11-29T08:53:00.989707Z","iopub.status.idle":"2023-11-29T08:53:05.100877Z","shell.execute_reply.started":"2023-11-29T08:53:00.989671Z","shell.execute_reply":"2023-11-29T08:53:05.099949Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Reactivity ratio per location","metadata":{}},{"cell_type":"code","source":"norm0 = (exp0[datacols].loc[testids] - exp0[datacols].loc[testids].min())/(exp0[datacols].loc[testids].max() - exp0[datacols].loc[testids].min())\nnorm1 = (exp1[datacols].loc[testids] - exp1[datacols].loc[testids].min())/(exp1[datacols].loc[testids].max() - exp1[datacols].loc[testids].min())\n\nnormr = norm0/(norm1+1)\n\nMakeReactivityPanel(datacols[0:50],varrep,normr)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:53:05.102192Z","iopub.execute_input":"2023-11-29T08:53:05.102569Z","iopub.status.idle":"2023-11-29T08:53:09.797484Z","shell.execute_reply.started":"2023-11-29T08:53:05.102537Z","shell.execute_reply":"2023-11-29T08:53:09.796448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[50:100],varrep,normr)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:53:09.798854Z","iopub.execute_input":"2023-11-29T08:53:09.799166Z","iopub.status.idle":"2023-11-29T08:53:14.276020Z","shell.execute_reply.started":"2023-11-29T08:53:09.799139Z","shell.execute_reply":"2023-11-29T08:53:14.275119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[100:150],varrep,normr)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:53:14.277573Z","iopub.execute_input":"2023-11-29T08:53:14.277901Z","iopub.status.idle":"2023-11-29T08:53:17.995878Z","shell.execute_reply.started":"2023-11-29T08:53:14.277872Z","shell.execute_reply":"2023-11-29T08:53:17.994870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MakeReactivityPanel(datacols[150:200],varrep,normr)","metadata":{"execution":{"iopub.status.busy":"2023-11-29T08:53:17.997401Z","iopub.execute_input":"2023-11-29T08:53:17.998151Z","iopub.status.idle":"2023-11-29T08:53:21.516948Z","shell.execute_reply.started":"2023-11-29T08:53:17.998112Z","shell.execute_reply":"2023-11-29T08:53:21.515922Z"},"trusted":true},"execution_count":null,"outputs":[]}]}