{"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":"# Kmers as features for Single cell sequence data","metadata":{}},{"cell_type":"code","source":"import time\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom joblib import dump\nfrom itertools import product\nfrom sklearn import preprocessing as pr\nfrom sklearn.model_selection import train_test_split\n\nfrom typing import Sequence\n\nimport jax\nimport optax\nimport jax.numpy as jnp\nfrom jax import random\nfrom flax import linen as nn","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2022-11-16T06:54:05.072797Z","iopub.execute_input":"2022-11-16T06:54:05.073311Z","iopub.status.idle":"2022-11-16T06:54:07.359101Z","shell.execute_reply.started":"2022-11-16T06:54:05.073205Z","shell.execute_reply":"2022-11-16T06:54:07.357805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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=12)\n    Axes.yaxis.set_tick_params(labelsize=12)\n\nclass Coder(nn.Module):\n    \n    Units: Sequence[int]\n    name: str \n    train: bool = True \n    \n    def setup(self):\n        self.layers = [nn.Dense(feat,use_bias=False,name = self.name+' layer_'+str(k)) for k,feat in enumerate(self.Units)]\n        self.norms = [nn.BatchNorm(use_running_average=not self.train,name = self.name+' norm_'+str(k)) for k,feat in enumerate(self.Units)]\n        \n    @nn.compact\n    def __call__(self,inputs):\n        x = inputs\n        for k,block in enumerate(zip(self.layers,self.norms)):\n            lay,norm = block\n            x = lay(x)\n            x = norm(x)\n            x = nn.relu(x)\n        return x\n\nclass Encoder(nn.Module):\n    \n    Units: Sequence[int]\n    name: str \n    train: bool = True\n    \n    def setup(self):\n        self.encoder = Coder(self.Units[1::],self.name,self.train)\n        self.mean = nn.Dense(self.Units[-1], name='mean')\n        self.logvar = nn.Dense(self.Units[-1], name='logvar')\n    \n    @nn.compact\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 Decoder(nn.Module):\n    \n    Units: Sequence[int]\n    name: str \n    train: bool = True\n    \n    def setup(self):\n        self.decoder = Coder(self.Units[0:-1],self.name,self.train)\n        self.out = nn.Dense(self.Units[-1],use_bias=False, name='out')\n    \n    @nn.compact\n    def __call__(self, inputs):\n        x = inputs\n        decoded_1 = self.decoder(x)\n        \n        out =self.out(decoded_1)\n        out = nn.BatchNorm(use_running_average=not self.train,name = 'outnorm')(out)\n        out = nn.sigmoid(out)\n        \n        return out\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 VAE(nn.Module):\n    \n    Units: Sequence[int]\n    name: str \n    train: bool = True\n    \n    def setup(self):\n        self.encoder = Encoder(self.Units,self.name+'encoder',self.train)\n        self.decoder = Decoder(self.Units[::-1],self.name+'decoder',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","metadata":{"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2022-11-16T06:54:07.361572Z","iopub.execute_input":"2022-11-16T06:54:07.362726Z","iopub.status.idle":"2022-11-16T06:54:07.388892Z","shell.execute_reply.started":"2022-11-16T06:54:07.362676Z","shell.execute_reply":"2022-11-16T06:54:07.387732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@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(axis=-1)\n    loss_value = optax.l2_loss(recon_x, batch).mean(axis=-1)\n    total_loss = loss_value.mean() + kld_loss.mean()\n    \n    return total_loss,newbatchst['batch_stats']\n\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(5)]\n    \n    Scheduler = optax.sgdr_schedule(esp)\n    localOptimizer = optax.adam(learning_rate=Scheduler)\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,TrainData.shape[0],batch_size):\n    \n            stb = time.time()\n            batch = TrainData[k:k+batch_size]\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,TestData.shape[0],batch_size):\n            rng, key = random.split(rng)\n            val_batch = TestData[i:i+batch_size]\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":{"iopub.status.busy":"2022-11-16T06:54:07.390304Z","iopub.execute_input":"2022-11-16T06:54:07.390774Z","iopub.status.idle":"2022-11-16T06:54:07.412697Z","shell.execute_reply.started":"2022-11-16T06:54:07.390727Z","shell.execute_reply":"2022-11-16T06:54:07.411559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mainUnits  = [340,170,85,21,5,2]\nsh = 0.00005\n\nlr = 0.005\nminlr = 0.0005\nbatchSize = 256\nepochs = 20\nInputShape = 340\n\nKmerData = pd.read_csv('../input/single-cells-kmers-cite/CITETrainKmersPerCell.csv')\nKmerData = KmerData.set_index('ids')\n\ntrainSamps, valSamps, _, _ = train_test_split(KmerData.index,np.arange(len(KmerData.index)) , test_size=0.1, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T06:54:07.415023Z","iopub.execute_input":"2022-11-16T06:54:07.415355Z","iopub.status.idle":"2022-11-16T06:54:17.098948Z","shell.execute_reply.started":"2022-11-16T06:54:07.415326Z","shell.execute_reply":"2022-11-16T06:54:17.097610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainData = np.array(KmerData.loc[trainSamps])\ntestData = np.array(KmerData.loc[valSamps])\n\nscaler = pr.MinMaxScaler()\nscaler.fit(trainData)\n\ntrainData = scaler.transform(trainData)\ntestData = scaler.transform(testData)\n\ndef VAEModel():\n    return VAE(mainUnits,'test')\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\ninit_data = jnp.ones((batchSize, InputShape), 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.01,epochs=50,\n                                    batch_size=256)\n\nfinalParams = {'params':params0,'batch_stats':batchStats}","metadata":{"execution":{"iopub.status.busy":"2022-11-16T06:54:17.100387Z","iopub.execute_input":"2022-11-16T06:54:17.100734Z","iopub.status.idle":"2022-11-16T06:56:57.146950Z","shell.execute_reply.started":"2022-11-16T06:54:17.100703Z","shell.execute_reply":"2022-11-16T06:56:57.145684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def EcoderModel(trainparams,batch):\n    return Encoder(mainUnits,'testencoder',train=False).apply(trainparams, batch)\n\nlocalparams = {'params':finalParams['params']['testencoder'],'batch_stats':finalParams['batch_stats']['testencoder']}\n\nfulldata = scaler.transform(np.array(KmerData))\nmu,logvar = EcoderModel(localparams,fulldata)\nVariationalRepresentation = reparameterize(rng,mu,logvar)\n \nplt.figure(figsize=(15,5))\naxs = plt.gca()\naxs.scatter(VariationalRepresentation[:,0],VariationalRepresentation[:,1],alpha=0.15)\naxs.title.set_text('Latent Space')\nPlotStyle(axs)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T06:56:57.148547Z","iopub.execute_input":"2022-11-16T06:56:57.149019Z","iopub.status.idle":"2022-11-16T06:57:00.104014Z","shell.execute_reply.started":"2022-11-16T06:56:57.148972Z","shell.execute_reply":"2022-11-16T06:57:00.102842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Meta Data ","metadata":{}},{"cell_type":"code","source":"MetaData = pd.read_csv('../input/open-problems-multimodal/metadata.csv')\nMetaData = MetaData.set_index('cell_id')\n\ncellTypes = MetaData['cell_type'].unique()\ntypeToNumber = {}\n\nfor k,val in enumerate(cellTypes):\n    typeToNumber[val] = k\n    \nMetaData['typenumber'] = [typeToNumber[val] for val in MetaData['cell_type']]","metadata":{"execution":{"iopub.status.busy":"2022-11-16T06:57:00.105692Z","iopub.execute_input":"2022-11-16T06:57:00.106328Z","iopub.status.idle":"2022-11-16T06:57:00.644023Z","shell.execute_reply.started":"2022-11-16T06:57:00.106286Z","shell.execute_reply":"2022-11-16T06:57:00.642825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Day","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(15,5))\naxs = plt.gca()\naxs.scatter(VariationalRepresentation[:,0],VariationalRepresentation[:,1],alpha=0.25,c=MetaData.loc[KmerData.index]['day'])\naxs.title.set_text('Latent Space')\nPlotStyle(axs)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T06:57:00.645292Z","iopub.execute_input":"2022-11-16T06:57:00.645963Z","iopub.status.idle":"2022-11-16T06:57:03.651442Z","shell.execute_reply.started":"2022-11-16T06:57:00.645924Z","shell.execute_reply":"2022-11-16T06:57:03.650317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Cell type","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(15,5))\naxs = plt.gca()\naxs.scatter(VariationalRepresentation[:,0],VariationalRepresentation[:,1],alpha=0.25,c=MetaData.loc[KmerData.index]['typenumber'])\naxs.title.set_text('Latent Space')\nPlotStyle(axs)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T06:57:03.653123Z","iopub.execute_input":"2022-11-16T06:57:03.654579Z","iopub.status.idle":"2022-11-16T06:57:06.781292Z","shell.execute_reply.started":"2022-11-16T06:57:03.654531Z","shell.execute_reply":"2022-11-16T06:57:06.777038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Donor","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(15,5))\naxs = plt.gca()\naxs.scatter(VariationalRepresentation[:,0],VariationalRepresentation[:,1],alpha=0.25,c=MetaData.loc[KmerData.index]['donor'])\naxs.title.set_text('Latent Space')\nPlotStyle(axs)","metadata":{"execution":{"iopub.status.busy":"2022-11-16T06:57:06.784882Z","iopub.execute_input":"2022-11-16T06:57:06.785662Z","iopub.status.idle":"2022-11-16T06:57:09.735125Z","shell.execute_reply.started":"2022-11-16T06:57:06.785613Z","shell.execute_reply":"2022-11-16T06:57:09.733684Z"},"trusted":true},"execution_count":null,"outputs":[]}]}